feat: add multi-image vision analysis

This commit is contained in:
QiuSW
2026-07-16 14:13:19 +08:00
parent 0ea65d4df4
commit 5325eacdfe
29 changed files with 1013 additions and 51 deletions
+328 -1
View File
@@ -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")