feat: add multi-image vision analysis
This commit is contained in:
+328
-1
@@ -331,6 +331,12 @@ class ModelsCatalogApiTests(TestCase):
|
||||
url="https://provider-secret.example/v1/images/edits",
|
||||
model_sku="secret-sku-image-2",
|
||||
)
|
||||
vision_alias = self.create_alias(
|
||||
alias="vision-standard",
|
||||
operation_type=ModelAlias.OperationType.VISION,
|
||||
capabilities=["text", "vision"],
|
||||
model_sku="secret-sku-vision",
|
||||
)
|
||||
PricingRule.objects.create(
|
||||
operation_type=title_alias.operation_type,
|
||||
alias=title_alias.alias,
|
||||
@@ -349,13 +355,19 @@ class ModelsCatalogApiTests(TestCase):
|
||||
resolution="1k",
|
||||
points_cost=12,
|
||||
)
|
||||
PricingRule.objects.create(
|
||||
operation_type=vision_alias.operation_type,
|
||||
alias=vision_alias.alias,
|
||||
resolution="",
|
||||
points_cost=3,
|
||||
)
|
||||
|
||||
response = self.client.get(self.url, **self.auth_header())
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertNotIn(GenerateRateThrottle, ModelsView.throttle_classes)
|
||||
models = {item["alias"]: item for item in response.data["models"]}
|
||||
self.assertEqual(set(models), {"title-standard", "image-edit"})
|
||||
self.assertEqual(set(models), {"title-standard", "image-edit", "vision-standard"})
|
||||
self.assertEqual(
|
||||
set(models["title-standard"]),
|
||||
{
|
||||
@@ -384,6 +396,14 @@ class ModelsCatalogApiTests(TestCase):
|
||||
{"resolution": "1K", "points_cost": 12},
|
||||
],
|
||||
)
|
||||
self.assertEqual(models["vision-standard"]["operation_type"], "vision")
|
||||
self.assertEqual(models["vision-standard"]["capabilities"], ["text", "vision"])
|
||||
self.assertTrue(models["vision-standard"]["requires_image"])
|
||||
self.assertEqual(models["vision-standard"]["pricing_status"], "priced")
|
||||
self.assertEqual(
|
||||
models["vision-standard"]["prices"],
|
||||
[{"resolution": "default", "points_cost": 3}],
|
||||
)
|
||||
response_body = json.dumps(response.data, ensure_ascii=False)
|
||||
self.assertNotIn("secret-sku", response_body)
|
||||
self.assertNotIn("provider-secret.example", response_body)
|
||||
@@ -1056,8 +1076,10 @@ class FakeGenerationProvider:
|
||||
self._capabilities = set(capabilities or {"text", "image", "vision"})
|
||||
self.text_calls = []
|
||||
self.image_calls = []
|
||||
self.vision_calls = []
|
||||
self.text_error = None
|
||||
self.image_error = None
|
||||
self.vision_error = None
|
||||
|
||||
def capabilities(self):
|
||||
return set(self._capabilities)
|
||||
@@ -1083,6 +1105,17 @@ class FakeGenerationProvider:
|
||||
raw={"b64_json": "SECRET_RAW_SHOULD_NOT_BE_STORED"},
|
||||
)
|
||||
|
||||
def analyze_images(self, prompt, model, **kwargs):
|
||||
self.vision_calls.append({"prompt": prompt, "model": model, **kwargs})
|
||||
if self.vision_error is not None:
|
||||
raise self.vision_error
|
||||
return TextGenerationResult(
|
||||
text="第一张展示商品正面。\n第二张展示商品细节。",
|
||||
titles=(),
|
||||
model_used=model.model,
|
||||
raw={"secret": "SECRET_RAW_SHOULD_NOT_BE_STORED"},
|
||||
)
|
||||
|
||||
|
||||
class FakeImageUrlResponse:
|
||||
def __init__(self, *, status_code=200, headers=None, chunks=()):
|
||||
@@ -1143,14 +1176,26 @@ class GenerateApiTests(TestCase):
|
||||
model=f"gpt-image-{suffix}",
|
||||
capabilities=["image", "vision"],
|
||||
)
|
||||
self.vision_model = self.create_ai_model(
|
||||
name=f"vision-model-{suffix}",
|
||||
model=f"gpt-vision-{suffix}",
|
||||
capabilities=["text", "vision"],
|
||||
)
|
||||
self.title_alias = f"title-standard-{suffix}"
|
||||
self.image_alias = f"image-hd-{suffix}"
|
||||
self.vision_alias = f"vision-standard-{suffix}"
|
||||
ModelAlias.objects.create(
|
||||
operation_type=ModelAlias.OperationType.TITLE,
|
||||
alias=self.title_alias,
|
||||
ai_model=self.title_model,
|
||||
is_default=True,
|
||||
)
|
||||
ModelAlias.objects.create(
|
||||
operation_type=ModelAlias.OperationType.VISION,
|
||||
alias=self.vision_alias,
|
||||
ai_model=self.vision_model,
|
||||
is_default=True,
|
||||
)
|
||||
ModelAlias.objects.create(
|
||||
operation_type=ModelAlias.OperationType.IMAGE,
|
||||
alias=self.image_alias,
|
||||
@@ -1168,6 +1213,11 @@ class GenerateApiTests(TestCase):
|
||||
resolution="1K",
|
||||
points_cost=10,
|
||||
)
|
||||
PricingRule.objects.create(
|
||||
operation_type=CallRecord.OperationType.VISION,
|
||||
alias=self.vision_alias,
|
||||
points_cost=3,
|
||||
)
|
||||
|
||||
def create_ai_model(self, *, name, model, capabilities):
|
||||
ai_model = AiModel(
|
||||
@@ -1314,6 +1364,283 @@ class GenerateApiTests(TestCase):
|
||||
1,
|
||||
)
|
||||
|
||||
def test_analyze_images_supports_single_image_with_explicit_alias(self):
|
||||
encoded = base64.b64encode(b"single-image").decode("ascii")
|
||||
|
||||
response = self.post_with_provider(
|
||||
"/api/v1/analyze/images",
|
||||
{
|
||||
"prompt": "描述这张商品图",
|
||||
"model": self.vision_alias,
|
||||
"images": [{"image_base64": encoded}],
|
||||
},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.data["alias"], self.vision_alias)
|
||||
self.assertEqual(response.data["points_cost"], 3)
|
||||
self.assertEqual(len(self.provider.vision_calls), 1)
|
||||
self.assertEqual(
|
||||
[image.data for image in self.provider.vision_calls[0]["images"]],
|
||||
[b"single-image"],
|
||||
)
|
||||
|
||||
def test_analyze_images_supports_ordered_mixed_sources_and_charges_once(self):
|
||||
first = base64.b64encode(b"first-image").decode("ascii")
|
||||
response_from_url = FakeImageUrlResponse(
|
||||
headers={"Content-Type": "image/jpeg"},
|
||||
chunks=(b"second-", b"image"),
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"apps.api.generation.socket.getaddrinfo",
|
||||
return_value=dns_result("93.184.216.34"),
|
||||
),
|
||||
patch(
|
||||
"apps.api.generation.requests.Session.get",
|
||||
return_value=response_from_url,
|
||||
),
|
||||
):
|
||||
response = self.post_with_provider(
|
||||
"/api/v1/analyze/images",
|
||||
{
|
||||
"prompt": "比较两张商品图",
|
||||
"images": [
|
||||
{"image_base64": f"data:image/png;base64,{first}"},
|
||||
{"image_url": "https://images.example.test/detail.jpg"},
|
||||
],
|
||||
"parameters": {"temperature": 0.2},
|
||||
},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.data["text"], "第一张展示商品正面。\n第二张展示商品细节。")
|
||||
self.assertEqual(response.data["alias"], self.vision_alias)
|
||||
self.assertEqual(response.data["model_used"], self.vision_model.model)
|
||||
self.assertEqual(response.data["points_cost"], 3)
|
||||
self.assertEqual(response.data["points_balance"], 97)
|
||||
self.assertEqual(len(self.provider.vision_calls), 1)
|
||||
images = self.provider.vision_calls[0]["images"]
|
||||
self.assertEqual([image.data for image in images], [b"first-image", b"second-image"])
|
||||
self.assertEqual(
|
||||
[image.mime_type for image in images],
|
||||
["image/png", "image/jpeg"],
|
||||
)
|
||||
|
||||
self.wallet.refresh_from_db()
|
||||
self.assertEqual(self.wallet.points_balance, 97)
|
||||
call = CallRecord.objects.get(pk=response.data["call_id"])
|
||||
self.assertEqual(call.operation_type, CallRecord.OperationType.VISION)
|
||||
self.assertEqual(call.status, CallRecord.Status.SUCCESS)
|
||||
self.assertEqual(call.resolution, "")
|
||||
self.assertEqual(call.result_summary, response.data["text"])
|
||||
self.assertNotIn("first-image", call.result_summary)
|
||||
self.assertNotIn("SECRET_RAW", call.result_summary)
|
||||
self.assertEqual(
|
||||
PointsLedger.objects.filter(
|
||||
ref_call=call,
|
||||
change_type=PointsLedger.ChangeType.CONSUME,
|
||||
).count(),
|
||||
1,
|
||||
)
|
||||
|
||||
def test_analyze_images_requires_nonempty_exclusive_image_sources(self):
|
||||
empty = self.post_with_provider(
|
||||
"/api/v1/analyze/images",
|
||||
{"prompt": "分析图片", "images": []},
|
||||
)
|
||||
both = self.post_with_provider(
|
||||
"/api/v1/analyze/images",
|
||||
{
|
||||
"prompt": "分析图片",
|
||||
"images": [
|
||||
{
|
||||
"image_url": "https://images.example.test/input.jpg",
|
||||
"image_base64": "aW1hZ2U=",
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
self.assertEqual(empty.status_code, 400)
|
||||
self.assertEqual(empty.data["error"]["code"], "bad_request")
|
||||
self.assertEqual(both.status_code, 400)
|
||||
self.assertEqual(both.data["error"]["code"], "bad_request")
|
||||
self.assertEqual(self.provider.vision_calls, [])
|
||||
self.assert_generation_not_charged()
|
||||
|
||||
@override_settings(VISION_MAX_IMAGES=1)
|
||||
def test_analyze_images_rejects_too_many_images_before_charge(self):
|
||||
encoded = base64.b64encode(b"image").decode("ascii")
|
||||
|
||||
response = self.post_with_provider(
|
||||
"/api/v1/analyze/images",
|
||||
{
|
||||
"prompt": "分析图片",
|
||||
"images": [
|
||||
{"image_base64": encoded},
|
||||
{"image_base64": encoded},
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 400)
|
||||
self.assertEqual(response.data["error"]["code"], "bad_request")
|
||||
self.assert_generation_not_charged()
|
||||
|
||||
@override_settings(VISION_MAX_IMAGE_BYTES=3, VISION_MAX_TOTAL_BYTES=10)
|
||||
def test_analyze_images_rejects_oversized_single_image_before_charge(self):
|
||||
encoded = base64.b64encode(b"four").decode("ascii")
|
||||
|
||||
response = self.post_with_provider(
|
||||
"/api/v1/analyze/images",
|
||||
{"prompt": "分析图片", "images": [{"image_base64": encoded}]},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 400)
|
||||
self.assertEqual(response.data["error"]["code"], "bad_request")
|
||||
self.assert_generation_not_charged()
|
||||
|
||||
@override_settings(VISION_MAX_IMAGE_BYTES=10, VISION_MAX_TOTAL_BYTES=5)
|
||||
def test_analyze_images_rejects_oversized_total_before_charge(self):
|
||||
encoded = base64.b64encode(b"abc").decode("ascii")
|
||||
|
||||
response = self.post_with_provider(
|
||||
"/api/v1/analyze/images",
|
||||
{
|
||||
"prompt": "分析图片",
|
||||
"images": [
|
||||
{"image_base64": encoded},
|
||||
{"image_base64": encoded},
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 400)
|
||||
self.assertEqual(response.data["error"]["code"], "bad_request")
|
||||
self.assert_generation_not_charged()
|
||||
|
||||
def test_analyze_images_rejects_private_image_url_before_charge(self):
|
||||
with patch(
|
||||
"apps.api.generation.socket.getaddrinfo",
|
||||
return_value=dns_result("127.0.0.1"),
|
||||
):
|
||||
response = self.post_with_provider(
|
||||
"/api/v1/analyze/images",
|
||||
{
|
||||
"prompt": "分析图片",
|
||||
"images": [{"image_url": "http://internal.example.test/input.jpg"}],
|
||||
},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 400)
|
||||
self.assertEqual(response.data["error"]["code"], "bad_request")
|
||||
self.assertEqual(self.provider.vision_calls, [])
|
||||
self.assert_generation_not_charged()
|
||||
|
||||
def test_analyze_images_requires_api_key(self):
|
||||
encoded = base64.b64encode(b"image").decode("ascii")
|
||||
|
||||
response = self.client.post(
|
||||
"/api/v1/analyze/images",
|
||||
{"prompt": "分析图片", "images": [{"image_base64": encoded}]},
|
||||
format="json",
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 401)
|
||||
self.assertEqual(response.data["error"]["code"], "unauthorized")
|
||||
self.assert_generation_not_charged()
|
||||
|
||||
@override_settings(
|
||||
MODERATION_ENABLED=True,
|
||||
MODERATION_PROVIDER="keyword",
|
||||
MODERATION_CACHE_VERSION_KEY="test:api:vision:moderation:version",
|
||||
)
|
||||
def test_analyze_images_blocks_prompt_before_loading_images_or_charge(self):
|
||||
SensitiveWord.objects.create(word="敏感词", category="policy")
|
||||
|
||||
with (
|
||||
patch("apps.api.generation.decode_image_input") as decode_image,
|
||||
patch("apps.api.generation.download_image_input") as download_image,
|
||||
):
|
||||
response = self.post_with_provider(
|
||||
"/api/v1/analyze/images",
|
||||
{
|
||||
"prompt": "分析敏-感词图片",
|
||||
"images": [
|
||||
{"image_url": "https://images.example.test/input.jpg"}
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 400)
|
||||
self.assertEqual(response.data["error"]["code"], "content_blocked")
|
||||
decode_image.assert_not_called()
|
||||
download_image.assert_not_called()
|
||||
self.assert_generation_not_charged()
|
||||
|
||||
def test_analyze_images_rejects_model_or_provider_without_text_vision(self):
|
||||
bad_alias = f"vision-without-text-{uuid.uuid4().hex[:8]}"
|
||||
ModelAlias.objects.create(
|
||||
operation_type=ModelAlias.OperationType.VISION,
|
||||
alias=bad_alias,
|
||||
ai_model=self.image_model,
|
||||
)
|
||||
encoded = base64.b64encode(b"image").decode("ascii")
|
||||
payload = {
|
||||
"prompt": "分析图片",
|
||||
"model": bad_alias,
|
||||
"images": [{"image_base64": encoded}],
|
||||
}
|
||||
|
||||
model_rejected = self.post_with_provider("/api/v1/analyze/images", payload)
|
||||
provider_rejected = self.post_with_provider(
|
||||
"/api/v1/analyze/images",
|
||||
{**payload, "model": self.vision_alias},
|
||||
provider=FakeGenerationProvider(capabilities={"vision"}),
|
||||
)
|
||||
|
||||
self.assertEqual(model_rejected.status_code, 400)
|
||||
self.assertEqual(model_rejected.data["error"]["code"], "model_not_allowed")
|
||||
self.assertEqual(provider_rejected.status_code, 400)
|
||||
self.assertEqual(provider_rejected.data["error"]["code"], "model_not_allowed")
|
||||
self.assert_generation_not_charged()
|
||||
|
||||
def test_analyze_images_upstream_failure_refunds_once(self):
|
||||
encoded = base64.b64encode(b"image").decode("ascii")
|
||||
self.provider.vision_error = requests.Timeout("vision timeout")
|
||||
|
||||
response = self.post_with_provider(
|
||||
"/api/v1/analyze/images",
|
||||
{
|
||||
"prompt": "分析图片",
|
||||
"model": self.vision_alias,
|
||||
"images": [{"image_base64": encoded}],
|
||||
},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 502)
|
||||
self.assertEqual(response.data["error"]["code"], "upstream_timeout")
|
||||
self.wallet.refresh_from_db()
|
||||
self.assertEqual(self.wallet.points_balance, 100)
|
||||
call = CallRecord.objects.get(operation_type=CallRecord.OperationType.VISION)
|
||||
self.assertEqual(call.status, CallRecord.Status.FAILED)
|
||||
self.assertEqual(
|
||||
PointsLedger.objects.filter(
|
||||
ref_call=call,
|
||||
change_type=PointsLedger.ChangeType.CONSUME,
|
||||
).count(),
|
||||
1,
|
||||
)
|
||||
self.assertEqual(
|
||||
PointsLedger.objects.filter(
|
||||
ref_call=call,
|
||||
change_type=PointsLedger.ChangeType.REFUND,
|
||||
).count(),
|
||||
1,
|
||||
)
|
||||
|
||||
def test_generate_image_stores_file_returns_url_and_does_not_store_raw_base64(self):
|
||||
encoded = base64.b64encode(b"input-image").decode("ascii")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user