feat(android): extract structured purchase requirements
This commit is contained in:
+21
@@ -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) }
|
||||
}
|
||||
+85
@@ -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 = ""
|
||||
)
|
||||
)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user