feat: add multi-image vision analysis
This commit is contained in:
+8
-6
@@ -18,15 +18,16 @@ class ModelCapabilityError(AliasResolutionError):
|
||||
|
||||
|
||||
REQUIRED_CAPABILITIES = {
|
||||
ModelAlias.OperationType.TITLE: "text",
|
||||
ModelAlias.OperationType.IMAGE: "image",
|
||||
ModelAlias.OperationType.TITLE: frozenset({"text"}),
|
||||
ModelAlias.OperationType.IMAGE: frozenset({"image"}),
|
||||
ModelAlias.OperationType.VISION: frozenset({"text", "vision"}),
|
||||
}
|
||||
|
||||
|
||||
def resolve_model_alias(operation_type: str, alias: str | None = None) -> ModelAlias:
|
||||
"""Resolve an external capability alias to an active ModelAlias row."""
|
||||
required_capability = REQUIRED_CAPABILITIES.get(operation_type)
|
||||
if required_capability is None:
|
||||
required_capabilities = REQUIRED_CAPABILITIES.get(operation_type)
|
||||
if required_capabilities is None:
|
||||
raise AliasResolutionError(f"unsupported operation_type: {operation_type}")
|
||||
|
||||
queryset = ModelAlias.objects.select_related("ai_model").filter(
|
||||
@@ -47,10 +48,11 @@ def resolve_model_alias(operation_type: str, alias: str | None = None) -> ModelA
|
||||
|
||||
ai_model: AiModel = model_alias.ai_model
|
||||
capabilities = ai_model.capabilities_set()
|
||||
if required_capability not in capabilities:
|
||||
missing_capabilities = required_capabilities - capabilities
|
||||
if missing_capabilities:
|
||||
raise ModelCapabilityError(
|
||||
f"alias {model_alias.alias} maps to model {ai_model.name} without "
|
||||
f"{required_capability} capability"
|
||||
f"{', '.join(sorted(missing_capabilities))} capability"
|
||||
)
|
||||
return model_alias
|
||||
|
||||
|
||||
+5
-3
@@ -14,13 +14,15 @@ PUBLIC_DEFAULT_RESOLUTION = "default"
|
||||
|
||||
|
||||
def _has_required_capability(model_alias: ModelAlias) -> bool:
|
||||
required_capability = REQUIRED_CAPABILITIES.get(model_alias.operation_type)
|
||||
if required_capability is None:
|
||||
required_capabilities = REQUIRED_CAPABILITIES.get(model_alias.operation_type)
|
||||
if required_capabilities is None:
|
||||
return False
|
||||
return required_capability in model_alias.ai_model.capabilities_set()
|
||||
return required_capabilities.issubset(model_alias.ai_model.capabilities_set())
|
||||
|
||||
|
||||
def _requires_image_input(ai_model: AiModel, operation_type: str) -> bool:
|
||||
if operation_type == ModelAlias.OperationType.VISION:
|
||||
return True
|
||||
if operation_type != ModelAlias.OperationType.IMAGE:
|
||||
return False
|
||||
try:
|
||||
|
||||
@@ -0,0 +1,18 @@
|
||||
# Generated by Django 5.2.15 on 2026-07-16 03:52
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('ai', '0004_alter_aiconfigauditlog_action_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='modelalias',
|
||||
name='operation_type',
|
||||
field=models.CharField(choices=[('title', '生成标题'), ('image', '生成图片'), ('vision', '图片理解')], max_length=32, verbose_name='操作类型'),
|
||||
),
|
||||
]
|
||||
@@ -92,6 +92,7 @@ class ModelAlias(models.Model):
|
||||
class OperationType(models.TextChoices):
|
||||
TITLE = "title", "生成标题"
|
||||
IMAGE = "image", "生成图片"
|
||||
VISION = "vision", "图片理解"
|
||||
|
||||
alias = models.SlugField("能力别名", max_length=64)
|
||||
operation_type = models.CharField("操作类型", max_length=32, choices=OperationType.choices)
|
||||
|
||||
@@ -4,6 +4,7 @@ from .base import (
|
||||
AiProviderError,
|
||||
AiResponseParseError,
|
||||
ImageGenerationResult,
|
||||
MultimodalImage,
|
||||
Provider,
|
||||
ResolvedModel,
|
||||
TextGenerationResult,
|
||||
@@ -16,6 +17,7 @@ __all__ = [
|
||||
"AiProviderError",
|
||||
"AiResponseParseError",
|
||||
"ImageGenerationResult",
|
||||
"MultimodalImage",
|
||||
"Provider",
|
||||
"ResolvedModel",
|
||||
"TextGenerationResult",
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Mapping, Protocol
|
||||
from typing import Any, Mapping, Protocol, Sequence
|
||||
|
||||
|
||||
class AiProviderError(RuntimeError):
|
||||
@@ -78,6 +78,12 @@ class ImageGenerationResult:
|
||||
raw: Mapping[str, Any]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MultimodalImage:
|
||||
data: bytes
|
||||
mime_type: str = "image/png"
|
||||
|
||||
|
||||
class Provider(Protocol):
|
||||
def capabilities(self) -> set[str]:
|
||||
...
|
||||
@@ -108,6 +114,16 @@ class Provider(Protocol):
|
||||
) -> ImageGenerationResult:
|
||||
...
|
||||
|
||||
def analyze_images(
|
||||
self,
|
||||
prompt: str,
|
||||
model: ResolvedModel,
|
||||
*,
|
||||
images: Sequence[MultimodalImage],
|
||||
parameters: Mapping[str, Any] | None = None,
|
||||
) -> TextGenerationResult:
|
||||
...
|
||||
|
||||
|
||||
def validate_model_config(model: ResolvedModel) -> None:
|
||||
errors = []
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Mapping
|
||||
from typing import Any, Mapping, Sequence
|
||||
|
||||
import requests
|
||||
|
||||
@@ -8,6 +8,7 @@ from .base import (
|
||||
AiCapabilityError,
|
||||
AiResponseParseError,
|
||||
ImageGenerationResult,
|
||||
MultimodalImage,
|
||||
ResolvedModel,
|
||||
TextGenerationResult,
|
||||
validate_model_config,
|
||||
@@ -18,6 +19,7 @@ from .utils import (
|
||||
API_IMAGES,
|
||||
API_IMAGES_EDITS,
|
||||
extract_image_from_response,
|
||||
extract_raw_text,
|
||||
extract_text_from_response,
|
||||
extract_titles_from_response,
|
||||
image_request_timeout,
|
||||
@@ -166,6 +168,37 @@ class ChatCompletionsProvider(BaseHttpProvider):
|
||||
raise AiResponseParseError("AI response did not contain an image")
|
||||
return ImageGenerationResult(image=image_bytes, model_used=model.model, raw=raw)
|
||||
|
||||
def analyze_images(
|
||||
self,
|
||||
prompt: str,
|
||||
model: ResolvedModel,
|
||||
*,
|
||||
images: Sequence[MultimodalImage],
|
||||
parameters: Mapping[str, Any] | None = None,
|
||||
) -> TextGenerationResult:
|
||||
validate_model_config(model)
|
||||
if not images:
|
||||
raise AiCapabilityError("vision analysis requires at least one image")
|
||||
url = normalize_api_url(model.url, API_CHAT)
|
||||
payload = build_chat_vision_payload(
|
||||
model,
|
||||
prompt,
|
||||
images=images,
|
||||
parameters=parameters,
|
||||
)
|
||||
response = self.session.post(
|
||||
url,
|
||||
headers=self._headers(model, json=True),
|
||||
json=payload,
|
||||
timeout=self._timeout(model, "1K"),
|
||||
)
|
||||
response.raise_for_status()
|
||||
raw = response.json()
|
||||
text = extract_raw_text(raw).strip()
|
||||
if not text:
|
||||
raise AiResponseParseError("AI response did not contain text")
|
||||
return TextGenerationResult(text=text, titles=(), model_used=model.model, raw=raw)
|
||||
|
||||
|
||||
class GeminiProvider(ChatCompletionsProvider):
|
||||
def generate_text(
|
||||
@@ -241,6 +274,37 @@ class GeminiProvider(ChatCompletionsProvider):
|
||||
raise AiResponseParseError("AI response did not contain an image")
|
||||
return ImageGenerationResult(image=image_bytes, model_used=model.model, raw=raw)
|
||||
|
||||
def analyze_images(
|
||||
self,
|
||||
prompt: str,
|
||||
model: ResolvedModel,
|
||||
*,
|
||||
images: Sequence[MultimodalImage],
|
||||
parameters: Mapping[str, Any] | None = None,
|
||||
) -> TextGenerationResult:
|
||||
validate_model_config(model)
|
||||
if not images:
|
||||
raise AiCapabilityError("vision analysis requires at least one image")
|
||||
url = normalize_api_url(model.url, API_GEMINI).replace("{model}", model.model)
|
||||
payload = build_gemini_vision_payload(
|
||||
model,
|
||||
prompt,
|
||||
images=images,
|
||||
parameters=parameters,
|
||||
)
|
||||
response = self.session.post(
|
||||
url,
|
||||
headers=self._headers(model, json=True),
|
||||
json=payload,
|
||||
timeout=self._timeout(model, "1K"),
|
||||
)
|
||||
response.raise_for_status()
|
||||
raw = response.json()
|
||||
text = extract_raw_text(raw).strip()
|
||||
if not text:
|
||||
raise AiResponseParseError("AI response did not contain text")
|
||||
return TextGenerationResult(text=text, titles=(), model_used=model.model, raw=raw)
|
||||
|
||||
|
||||
class ImagesGenerationProvider(BaseHttpProvider):
|
||||
def capabilities(self) -> set[str]:
|
||||
@@ -249,6 +313,9 @@ class ImagesGenerationProvider(BaseHttpProvider):
|
||||
def generate_text(self, *args: Any, **kwargs: Any) -> TextGenerationResult:
|
||||
raise AiCapabilityError("images generation provider cannot generate text")
|
||||
|
||||
def analyze_images(self, *args: Any, **kwargs: Any) -> TextGenerationResult:
|
||||
raise AiCapabilityError("images generation provider cannot analyze images")
|
||||
|
||||
def generate_image(
|
||||
self,
|
||||
prompt: str,
|
||||
@@ -298,6 +365,9 @@ class ImagesEditsProvider(BaseHttpProvider):
|
||||
def generate_text(self, *args: Any, **kwargs: Any) -> TextGenerationResult:
|
||||
raise AiCapabilityError("images edits provider cannot generate text")
|
||||
|
||||
def analyze_images(self, *args: Any, **kwargs: Any) -> TextGenerationResult:
|
||||
raise AiCapabilityError("images edits provider cannot analyze images")
|
||||
|
||||
def generate_image(
|
||||
self,
|
||||
prompt: str,
|
||||
@@ -384,6 +454,32 @@ def build_chat_image_payload(
|
||||
)
|
||||
|
||||
|
||||
def build_chat_vision_payload(
|
||||
model: ResolvedModel,
|
||||
prompt: str,
|
||||
*,
|
||||
images: Sequence[MultimodalImage],
|
||||
parameters: Mapping[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
content: list[dict[str, Any]] = [{"type": "text", "text": prompt}]
|
||||
content.extend(
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": image_bytes_to_data_url(image.data, image.mime_type),
|
||||
},
|
||||
}
|
||||
for image in images
|
||||
)
|
||||
payload: dict[str, Any] = {
|
||||
"model": model.model,
|
||||
"messages": [{"role": "user", "content": content}],
|
||||
"stream": False,
|
||||
}
|
||||
apply_extra_body(payload, model, parameters)
|
||||
return payload
|
||||
|
||||
|
||||
def build_gemini_payload(
|
||||
model: ResolvedModel,
|
||||
prompt: str,
|
||||
@@ -406,6 +502,26 @@ def build_gemini_payload(
|
||||
return payload
|
||||
|
||||
|
||||
def build_gemini_vision_payload(
|
||||
model: ResolvedModel,
|
||||
prompt: str,
|
||||
*,
|
||||
images: Sequence[MultimodalImage],
|
||||
parameters: Mapping[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
parts: list[dict[str, Any]] = [{"text": prompt}]
|
||||
for image in images:
|
||||
data_url = image_bytes_to_data_url(image.data, image.mime_type)
|
||||
mime_type, data = split_data_url(data_url)
|
||||
parts.append({"inlineData": {"mimeType": mime_type, "data": data}})
|
||||
payload: dict[str, Any] = {
|
||||
"contents": [{"parts": parts}],
|
||||
"generationConfig": {"responseModalities": ["TEXT"]},
|
||||
}
|
||||
apply_extra_body(payload, model, parameters)
|
||||
return payload
|
||||
|
||||
|
||||
def apply_extra_body(
|
||||
payload: dict[str, Any],
|
||||
model: ResolvedModel,
|
||||
|
||||
+121
-2
@@ -11,8 +11,18 @@ from apps.ai.admin import AiConfigAuditLogAdmin, AiModelAdmin, ModelAliasAdmin
|
||||
from apps.ai.aliases import AliasNotFoundError, ModelCapabilityError, resolve_alias
|
||||
from apps.ai.importers import import_ai_models_config
|
||||
from apps.ai.models import AiConfigAuditLog, AiModel, ModelAlias
|
||||
from apps.ai.providers import AiCapabilityError, ResolvedModel, get_provider, resolve_api_type
|
||||
from apps.ai.providers.openai_compatible import ChatCompletionsProvider, ImagesEditsProvider
|
||||
from apps.ai.providers import (
|
||||
AiCapabilityError,
|
||||
MultimodalImage,
|
||||
ResolvedModel,
|
||||
get_provider,
|
||||
resolve_api_type,
|
||||
)
|
||||
from apps.ai.providers.openai_compatible import (
|
||||
ChatCompletionsProvider,
|
||||
GeminiProvider,
|
||||
ImagesEditsProvider,
|
||||
)
|
||||
from apps.ai.providers.utils import image_request_timeout, resolution_to_size
|
||||
|
||||
|
||||
@@ -255,6 +265,85 @@ class ChatCompletionsProviderTests(SimpleTestCase):
|
||||
|
||||
self.assertEqual(session.posts[0]["timeout"], (30, 600))
|
||||
|
||||
def test_analyze_images_builds_ordered_payload_and_preserves_full_text(self):
|
||||
session = FakeSession(
|
||||
FakeResponse(
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"content": "第一张是正面图。\n第二张是细节图。"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
)
|
||||
provider = ChatCompletionsProvider(session=session)
|
||||
model = ResolvedModel(
|
||||
name="Vision text",
|
||||
url="https://api.vectorengine.ai/v1",
|
||||
model="vision-model",
|
||||
api_key="test-key",
|
||||
api_type="chat",
|
||||
)
|
||||
|
||||
result = provider.analyze_images(
|
||||
"比较两张商品图",
|
||||
model,
|
||||
images=(
|
||||
MultimodalImage(b"first-image", "image/jpeg"),
|
||||
MultimodalImage(b"second-image", "image/png"),
|
||||
),
|
||||
parameters={"temperature": 0.2, "messages": []},
|
||||
)
|
||||
|
||||
self.assertEqual(result.text, "第一张是正面图。\n第二张是细节图。")
|
||||
self.assertEqual(result.titles, ())
|
||||
content = session.posts[0]["json"]["messages"][0]["content"]
|
||||
self.assertEqual(content[0], {"type": "text", "text": "比较两张商品图"})
|
||||
self.assertTrue(content[1]["image_url"]["url"].startswith("data:image/jpeg;base64,"))
|
||||
self.assertTrue(content[2]["image_url"]["url"].startswith("data:image/png;base64,"))
|
||||
self.assertEqual(session.posts[0]["json"]["temperature"], 0.2)
|
||||
|
||||
def test_gemini_analyze_images_builds_ordered_inline_data(self):
|
||||
session = FakeSession(
|
||||
FakeResponse(
|
||||
{
|
||||
"candidates": [
|
||||
{"content": {"parts": [{"text": "多图分析结果"}]}}
|
||||
]
|
||||
}
|
||||
)
|
||||
)
|
||||
provider = GeminiProvider(session=session)
|
||||
model = ResolvedModel(
|
||||
name="Gemini vision",
|
||||
url="https://gemini.example.com",
|
||||
model="gemini-vision",
|
||||
api_key="test-key",
|
||||
api_type="gemini",
|
||||
)
|
||||
|
||||
result = provider.analyze_images(
|
||||
"理解这些图片",
|
||||
model,
|
||||
images=(
|
||||
MultimodalImage(b"one", "image/webp"),
|
||||
MultimodalImage(b"two", "image/jpeg"),
|
||||
),
|
||||
)
|
||||
|
||||
self.assertEqual(result.text, "多图分析结果")
|
||||
parts = session.posts[0]["json"]["contents"][0]["parts"]
|
||||
self.assertEqual(parts[0], {"text": "理解这些图片"})
|
||||
self.assertEqual(parts[1]["inlineData"]["mimeType"], "image/webp")
|
||||
self.assertEqual(parts[2]["inlineData"]["mimeType"], "image/jpeg")
|
||||
self.assertEqual(
|
||||
session.posts[0]["json"]["generationConfig"]["responseModalities"],
|
||||
["TEXT"],
|
||||
)
|
||||
|
||||
|
||||
class ImagesEditsProviderTests(SimpleTestCase):
|
||||
def test_generate_image_builds_multipart_request_and_parses_base64(self):
|
||||
@@ -344,6 +433,13 @@ class ImagesEditsProviderTests(SimpleTestCase):
|
||||
with self.assertRaises(AiCapabilityError):
|
||||
provider.generate_text("Generate title", model)
|
||||
|
||||
with self.assertRaises(AiCapabilityError):
|
||||
provider.analyze_images(
|
||||
"Analyze image",
|
||||
model,
|
||||
images=(MultimodalImage(b"source-image"),),
|
||||
)
|
||||
|
||||
|
||||
@override_settings(AI_KEY_ENCRYPTION_KEY=TEST_ENCRYPTION_KEY)
|
||||
class AiModelEncryptionTests(TestCase):
|
||||
@@ -440,6 +536,29 @@ class AliasResolutionTests(TestCase):
|
||||
self.assertEqual(resolved.model, "gpt-image-2")
|
||||
self.assertIn("image", resolved.capabilities)
|
||||
|
||||
def test_resolve_vision_alias_requires_text_and_vision_capabilities(self):
|
||||
ModelAlias.objects.create(
|
||||
operation_type=ModelAlias.OperationType.VISION,
|
||||
alias="vision-standard",
|
||||
ai_model=self.text_model,
|
||||
is_default=True,
|
||||
)
|
||||
|
||||
resolved = resolve_alias(ModelAlias.OperationType.VISION)
|
||||
|
||||
self.assertEqual(resolved.model, "gpt-5.5")
|
||||
self.assertTrue({"text", "vision"}.issubset(resolved.capabilities))
|
||||
|
||||
def test_resolve_vision_alias_rejects_vision_model_without_text(self):
|
||||
ModelAlias.objects.create(
|
||||
operation_type=ModelAlias.OperationType.VISION,
|
||||
alias="bad-vision",
|
||||
ai_model=self.image_model,
|
||||
)
|
||||
|
||||
with self.assertRaises(ModelCapabilityError):
|
||||
resolve_alias(ModelAlias.OperationType.VISION, "bad-vision")
|
||||
|
||||
def test_resolve_alias_rejects_capability_mismatch(self):
|
||||
ModelAlias.objects.create(
|
||||
operation_type=ModelAlias.OperationType.TITLE,
|
||||
|
||||
Reference in New Issue
Block a user