feat(android): extract structured purchase requirements

This commit is contained in:
QiuSW
2026-07-25 21:05:51 +08:00
parent c25d63c730
commit 32d5310000
22 changed files with 2014 additions and 32 deletions
@@ -0,0 +1,21 @@
package com.roubao.autopilot.data
import org.junit.Assert.assertFalse
import org.junit.Assert.assertTrue
import org.junit.Test
class ApiProviderCapabilityTest {
@Test
fun `action model providers cannot run requirement extraction`() {
assertFalse(ApiProvider.GUI_OWL.supportsRequirementExtraction)
assertFalse(ApiProvider.MAI_UI.supportsRequirementExtraction)
}
@Test
fun `openai compatible providers can run requirement extraction`() {
assertTrue(ApiProvider.ALIYUN.supportsRequirementExtraction)
assertTrue(ApiProvider.OPENAI.supportsRequirementExtraction)
assertTrue(ApiProvider.OPENROUTER.supportsRequirementExtraction)
assertTrue(ApiProvider.CUSTOM.supportsRequirementExtraction)
}
}
@@ -0,0 +1,395 @@
package com.roubao.autopilot.vlm
import com.roubao.task.ProbeReferenceImage
import com.roubao.task.ProbeTask
import kotlinx.coroutines.ExperimentalCoroutinesApi
import kotlinx.coroutines.awaitCancellation
import kotlinx.coroutines.cancelAndJoin
import kotlinx.coroutines.launch
import kotlinx.coroutines.test.runCurrent
import kotlinx.coroutines.test.runTest
import org.json.JSONObject
import org.junit.Assert.assertArrayEquals
import org.junit.Assert.assertEquals
import org.junit.Assert.assertFalse
import org.junit.Assert.assertNull
import org.junit.Assert.assertTrue
import org.junit.Test
import java.security.MessageDigest
@OptIn(ExperimentalCoroutinesApi::class)
class RequirementExtractorTest {
@Test
fun `valid response produces versioned requirement with immutable task constraints`() =
runTest {
val input = input()
val extractor = extractorWithResponse(validResponse(confidence = 0.92))
val completed = extractor.extract(input) as RequirementExtractionResult.Completed
val requirement = completed.extraction
val encoded = JSONObject(RequirementExtractionJson.encode(requirement))
assertEquals(REQUIREMENT_SCHEMA_VERSION, requirement.schemaVersion)
assertEquals("折叠桌面手机支架", requirement.searchQuery)
assertEquals("手机支架", requirement.category)
assertEquals(input.sku, requirement.sku)
assertEquals(input.quantity, requirement.quantity)
assertNull(requirement.maxBudget)
assertFalse(requirement.manualReviewRequired)
assertTrue(
requirement.warnings.any {
it.code == RequirementWarningCode.MAX_BUDGET_NOT_PROVIDED
}
)
assertEquals(input.sku, encoded.getString("sku"))
assertEquals(input.quantity, encoded.getInt("quantity"))
assertTrue(encoded.isNull("max_budget"))
assertFalse(encoded.getJSONObject("manual_review").getBoolean("required"))
assertEquals(
REQUIREMENT_PROMPT_VERSION,
encoded.getJSONObject("provenance").getString("prompt_version")
)
}
@Test
fun `privacy mapper excludes order number and store name from provider request`() =
runTest {
val imageBytes = jpegBytes()
val task = task(
imageBytes = imageBytes,
sourceOrderNo = "ORDER_PRIVATE_SENTINEL",
sourceStoreName = "STORE_PRIVATE_SENTINEL"
)
var capturedRequest: RequirementVlmRequest? = null
var gatewayCalls = 0
val gateway = RequirementVlmGateway { request ->
gatewayCalls += 1
capturedRequest = request
Result.success(validResponse())
}
val input = RequirementExtractionInput.from(task, imageBytes)
RequirementExtractor(gateway, "test-provider", "test-model").extract(input)
val request = requireNotNull(capturedRequest)
assertFalse(request.prompt.contains(task.sourceOrderNo))
assertFalse(request.prompt.contains(task.sourceStoreName))
assertFalse(request.prompt.contains("PATH_PRIVATE_SENTINEL"))
assertTrue(request.prompt.contains(task.title))
assertTrue(request.prompt.contains(task.sku))
assertFalse(request.prompt.contains("\"quantity\""))
assertArrayEquals(imageBytes, request.imageBytes)
assertEquals(1, gatewayCalls)
}
@Test
fun `model cannot add replacement sku quantity or execution action`() = runTest {
val unsafe = JSONObject(validResponse())
.put("sku", "REPLACEMENT")
.put("quantity", 999)
.put("action", "submit_order")
.toString()
val completed = extractorWithResponse(unsafe)
.extract(input()) as RequirementExtractionResult.Completed
assertTrue(completed.extraction.manualReviewRequired)
assertEquals("SKU-ORIGINAL", completed.extraction.sku)
assertEquals(3, completed.extraction.quantity)
assertEquals(
listOf(RequirementReviewReason.INVALID_MODEL_OUTPUT),
completed.extraction.manualReviewReasons
)
}
@Test
fun `low confidence requires manual review without changing constraints`() = runTest {
val input = input()
val completed = extractorWithResponse(validResponse(confidence = 0.49))
.extract(input) as RequirementExtractionResult.Completed
assertTrue(completed.extraction.manualReviewRequired)
assertEquals(
listOf(RequirementReviewReason.LOW_CONFIDENCE),
completed.extraction.manualReviewReasons
)
assertEquals(input.sku, completed.extraction.sku)
assertEquals(input.quantity, completed.extraction.quantity)
assertTrue(
completed.extraction.warnings.any {
it.code == RequirementWarningCode.LOW_CONFIDENCE
}
)
}
@Test
fun `confidence equal to threshold does not require manual review`() = runTest {
val completed = extractorWithResponse(
validResponse(confidence = REQUIREMENT_CONFIDENCE_THRESHOLD)
).extract(input()) as RequirementExtractionResult.Completed
assertFalse(completed.extraction.manualReviewRequired)
}
@Test
fun `conflicting evidence warning requires manual review`() = runTest {
val response = JSONObject(validResponse())
response.getJSONArray("warnings")
.getJSONObject(0)
.put("code", "TITLE_IMAGE_CONFLICT")
val completed = extractorWithResponse(response.toString())
.extract(input()) as RequirementExtractionResult.Completed
assertTrue(completed.extraction.manualReviewRequired)
assertEquals(
listOf(RequirementReviewReason.CONFLICTING_EVIDENCE),
completed.extraction.manualReviewReasons
)
}
@Test
fun `invalid reference image fails before provider call`() = runTest {
var calls = 0
val gateway = RequirementVlmGateway {
calls += 1
Result.success(validResponse())
}
val invalid = input().copy(expectedImageSha256 = "0".repeat(64))
val failed = RequirementExtractor(gateway, "test-provider", "test-model")
.extract(invalid) as RequirementExtractionResult.Failed
assertEquals(RequirementExtractionFailureCode.REFERENCE_IMAGE_INVALID, failed.code)
assertFalse(failed.retryable)
assertEquals(0, calls)
}
@Test
fun `non jpeg payload fails before provider call`() = runTest {
var calls = 0
val bytes = "not-a-jpeg".toByteArray()
val invalid = input().copy(
imageBytes = bytes,
expectedImageSizeBytes = bytes.size.toLong(),
expectedImageSha256 = bytes.sha256()
)
val gateway = RequirementVlmGateway {
calls += 1
Result.success(validResponse())
}
val failed = RequirementExtractor(gateway, "test-provider", "test-model")
.extract(invalid) as RequirementExtractionResult.Failed
assertEquals(RequirementExtractionFailureCode.REFERENCE_IMAGE_INVALID, failed.code)
assertEquals(0, calls)
}
@Test
fun `oversized title fails before provider call`() = runTest {
var calls = 0
val gateway = RequirementVlmGateway {
calls += 1
Result.success(validResponse())
}
val oversized = input().copy(title = "a".repeat(2049))
val failed = RequirementExtractor(gateway, "test-provider", "test-model")
.extract(oversized) as RequirementExtractionResult.Failed
assertEquals(RequirementExtractionFailureCode.SOURCE_INPUT_INVALID, failed.code)
assertFalse(failed.retryable)
assertEquals(0, calls)
}
@Test
fun `provider failure is generic and retryable`() = runTest {
val gateway = RequirementVlmGateway {
Result.failure(IllegalStateException("private provider payload"))
}
val failed = RequirementExtractor(gateway, "test-provider", "test-model")
.extract(input()) as RequirementExtractionResult.Failed
assertEquals(RequirementExtractionFailureCode.PROVIDER_ERROR, failed.code)
assertTrue(failed.retryable)
}
@Test
fun `typed image decode failure maps to non retryable image error`() = runTest {
val gateway = RequirementVlmGateway {
Result.failure(InvalidRequirementReferenceImageException())
}
val failed = RequirementExtractor(gateway, "test-provider", "test-model")
.extract(input()) as RequirementExtractionResult.Failed
assertEquals(RequirementExtractionFailureCode.REFERENCE_IMAGE_INVALID, failed.code)
assertFalse(failed.retryable)
}
@Test
fun `non retryable structured provider failure stays non retryable`() = runTest {
val gateway = RequirementVlmGateway {
Result.failure(StructuredVlmException(retryable = false))
}
val failed = RequirementExtractor(gateway, "test-provider", "test-model")
.extract(input()) as RequirementExtractionResult.Failed
assertEquals(RequirementExtractionFailureCode.PROVIDER_ERROR, failed.code)
assertFalse(failed.retryable)
}
@Test
fun `cancellation is propagated instead of becoming provider failure`() = runTest {
val gateway = RequirementVlmGateway {
awaitCancellation()
}
val job = launch {
RequirementExtractor(gateway, "test-provider", "test-model")
.extract(input())
}
runCurrent()
job.cancelAndJoin()
assertTrue(job.isCancelled)
}
@Test
fun `fenced json and surrounding prose are rejected`() = runTest {
val fenced = "```json\n${validResponse()}\n```"
val fencedResult = extractorWithResponse(fenced)
.extract(input()) as RequirementExtractionResult.Completed
val proseResult = extractorWithResponse("Result: ${validResponse()}")
.extract(input()) as RequirementExtractionResult.Completed
assertTrue(fencedResult.extraction.manualReviewRequired)
assertTrue(proseResult.extraction.manualReviewRequired)
assertEquals(
listOf(RequirementReviewReason.INVALID_MODEL_OUTPUT),
proseResult.extraction.manualReviewReasons
)
}
@Test
fun `coordinate-like attribute is rejected as invalid model output`() = runTest {
val root = JSONObject(validResponse())
root.getJSONArray("attributes")
.getJSONObject(0)
.put("name", "click_coordinate")
val completed = extractorWithResponse(root.toString())
.extract(input()) as RequirementExtractionResult.Completed
assertTrue(completed.extraction.manualReviewRequired)
assertEquals(0.0, completed.extraction.confidence, 0.0)
}
@Test
fun `numeric strings are rejected as invalid schema types`() = runTest {
val root = JSONObject(validResponse())
.put("schema_version", "1")
.put("confidence", "0.91")
val completed = extractorWithResponse(root.toString())
.extract(input()) as RequirementExtractionResult.Completed
assertTrue(completed.extraction.manualReviewRequired)
assertEquals(
listOf(RequirementReviewReason.INVALID_MODEL_OUTPUT),
completed.extraction.manualReviewReasons
)
}
@Test
fun `execution directive in any semantic field is rejected`() = runTest {
val root = JSONObject(validResponse())
.put("search_query", "Click(100,200)")
val completed = extractorWithResponse(root.toString())
.extract(input()) as RequirementExtractionResult.Completed
assertTrue(completed.extraction.manualReviewRequired)
assertEquals(
listOf(RequirementReviewReason.INVALID_MODEL_OUTPUT),
completed.extraction.manualReviewReasons
)
}
private fun extractorWithResponse(response: String): RequirementExtractor =
RequirementExtractor(
gateway = RequirementVlmGateway { Result.success(response) },
providerId = "test-provider",
model = "test-model"
)
private fun input(): RequirementExtractionInput {
val imageBytes = jpegBytes()
return RequirementExtractionInput(
title = "可折叠桌面支架",
sku = "SKU-ORIGINAL",
quantity = 3,
imageMediaType = "image/jpeg",
imageBytes = imageBytes,
expectedImageSizeBytes = imageBytes.size.toLong(),
expectedImageSha256 = imageBytes.sha256()
)
}
private fun task(
imageBytes: ByteArray,
sourceOrderNo: String,
sourceStoreName: String
): ProbeTask =
ProbeTask(
probeId = "probe_test",
sourceOrderNo = sourceOrderNo,
sourceStoreName = sourceStoreName,
title = "可折叠桌面支架",
sku = "SKU-ORIGINAL",
quantity = 3,
referenceImage = ProbeReferenceImage(
relativePath = "assets/PATH_PRIVATE_SENTINEL.jpg",
mediaType = "image/jpeg",
sizeBytes = imageBytes.size.toLong(),
sha256 = imageBytes.sha256()
)
)
private fun validResponse(confidence: Double = 0.91): String =
"""
{
"schema_version": 1,
"search_query": "折叠桌面手机支架",
"category": "手机支架",
"attributes": [
{"name": "形态", "value": "可折叠桌面款", "source": "BOTH"},
{"name": "颜色", "value": "按参考图", "source": "IMAGE"}
],
"confidence": $confidence,
"warnings": [
{"code": "ATTRIBUTE_UNCERTAIN", "message": "颜色需要人工核对"}
]
}
""".trimIndent()
private fun jpegBytes(): ByteArray =
byteArrayOf(
0xff.toByte(),
0xd8.toByte(),
0xff.toByte(),
0xe0.toByte(),
0x01,
0x02,
0xff.toByte(),
0xd9.toByte()
)
private fun ByteArray.sha256(): String =
MessageDigest.getInstance("SHA-256")
.digest(this)
.joinToString(separator = "") { byte -> "%02x".format(byte) }
}
@@ -0,0 +1,85 @@
package com.roubao.autopilot.vlm
import org.junit.Assert.assertFalse
import org.junit.Assert.assertTrue
import org.junit.Test
class RequirementProviderEndpointPolicyTest {
@Test
fun `https endpoints are allowed`() {
assertTrue(
RequirementProviderEndpointPolicy.isAllowed(
baseUrl = "https://api.example.test/v1",
apiKey = "configured"
)
)
assertTrue(
RequirementProviderEndpointPolicy.isAllowed(
baseUrl = "api.example.test/v1",
apiKey = "configured"
)
)
}
@Test
fun `http is limited to keyless loopback`() {
assertTrue(
RequirementProviderEndpointPolicy.isAllowed(
baseUrl = "http://127.0.0.1:8765/v1",
apiKey = ""
)
)
assertTrue(
RequirementProviderEndpointPolicy.isAllowed(
baseUrl = "http://localhost:8000/v1",
apiKey = ""
)
)
assertTrue(
RequirementProviderEndpointPolicy.isAllowed(
baseUrl = "http://[::1]:8000/v1",
apiKey = ""
)
)
assertFalse(
RequirementProviderEndpointPolicy.isAllowed(
baseUrl = "http://127.0.0.1:8765/v1",
apiKey = "secret"
)
)
assertFalse(
RequirementProviderEndpointPolicy.isAllowed(
baseUrl = "http://192.168.1.5:8000/v1",
apiKey = ""
)
)
assertFalse(
RequirementProviderEndpointPolicy.isAllowed(
baseUrl = "http://api.example.test/v1",
apiKey = ""
)
)
}
@Test
fun `userinfo query fragment and unsupported schemes are rejected`() {
assertFalse(
RequirementProviderEndpointPolicy.isAllowed(
baseUrl = "https://user@example.test/v1",
apiKey = ""
)
)
assertFalse(
RequirementProviderEndpointPolicy.isAllowed(
baseUrl = "https://api.example.test/v1?token=value",
apiKey = ""
)
)
assertFalse(
RequirementProviderEndpointPolicy.isAllowed(
baseUrl = "ftp://api.example.test/v1",
apiKey = ""
)
)
}
}