feat(avatar): add private vision chat support

This commit is contained in:
stefanfeng
2026-08-31 15:10:50 +08:00
parent 094f8cd40f
commit 016bc22c05
22 changed files with 1616 additions and 59 deletions
@@ -5,6 +5,7 @@ from database import init_db, SessionLocal
from models import (
Authorization,
Avatar,
ChatAttachment,
TakeoverCursor,
TakeoverMessage,
TakeoverReplyTask,
@@ -95,6 +96,9 @@ def authorization_context():
finally:
db.rollback()
avatar_ids = [avatar.id, other_avatar.id]
db.query(ChatAttachment).filter(
ChatAttachment.avatar_id.in_(avatar_ids)
).delete(synchronize_session=False)
db.query(TakeoverReplyTask).filter(
TakeoverReplyTask.avatar_id.in_(avatar_ids)
).delete(synchronize_session=False)
@@ -0,0 +1,273 @@
import json
from datetime import datetime, timedelta
from types import SimpleNamespace
from unittest.mock import Mock, patch
import pytest
from fastapi import HTTPException
from fastapi.testclient import TestClient
from database import SessionLocal
from main import app
from models import ChatAttachment
from routers.chat import (
ChatIn,
_attachment_contexts,
_load_chat_attachments,
_resolve_reply,
)
from services.chat_attachment_service import purge_expired_chat_attachments
from services.vision_service import PreparedImage
client = TestClient(app)
GENERAL_RESULT = {
"content": json.dumps({
"category": "general_image",
"summary": "一张包含产品路线图的截图",
"visible_text": "产品路线图",
"key_facts": ["包含三个阶段"],
"uncertainties": [],
"medical": {},
}, ensure_ascii=False),
"usage": {"total_tokens": 120},
}
def test_owner_can_upload_and_cache_image_analysis(authorization_context):
context = authorization_context
prepared = PreparedImage(b"jpeg", "image/jpeg", 100, 80)
with (
patch("routers.chat.prepare_image", return_value=prepared),
patch("routers.chat._run_billed_vision_call", return_value=GENERAL_RESULT),
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/chat/images",
headers=context["owner_headers"],
files={"file": ("roadmap.png", b"image-bytes", "image/png")},
)
assert response.status_code == 200
payload = response.json()["data"]
assert payload["status"] == "ready"
assert payload["category"] == "general_image"
assert payload["summary"] == "一张包含产品路线图的截图"
db = SessionLocal()
try:
stored = db.query(ChatAttachment).filter(ChatAttachment.id == payload["id"]).one()
assert stored.avatar_id == context["avatar"].id
assert stored.extracted_text == "产品路线图"
assert stored.structured_data["key_facts"] == ["包含三个阶段"]
finally:
db.close()
def test_non_owner_cannot_upload_chat_image(authorization_context):
context = authorization_context
response = client.post(
f"/api/avatar/{context['avatar'].id}/chat/images",
headers=context["other_headers"],
files={"file": ("private.png", b"image-bytes", "image/png")},
)
assert response.status_code == 403
def test_public_share_can_upload_without_exposing_analysis_details(authorization_context):
context = authorization_context
db = SessionLocal()
try:
avatar = db.get(type(context["avatar"]), context["avatar"].id)
avatar.share_token = f"share-{context['suffix']}"
db.commit()
share_token = avatar.share_token
finally:
db.close()
with (
patch(
"routers.chat.prepare_image",
return_value=PreparedImage(b"jpeg", "image/jpeg", 100, 80),
),
patch("routers.chat._run_billed_vision_call", return_value=GENERAL_RESULT),
):
response = client.post(
f"/api/public/avatar/{share_token}/chat/images",
files={"file": ("visitor.png", b"image-bytes", "image/png")},
)
assert response.status_code == 200
payload = response.json()["data"]
assert payload["status"] == "ready"
assert "structuredData" not in payload
assert "extractedText" not in payload
assert "visionModel" not in payload
assert "ocrModel" not in payload
db = SessionLocal()
try:
stored = db.get(ChatAttachment, payload["id"])
assert stored.uploader_kind == "public"
assert stored.avatar_id == context["avatar"].id
finally:
db.close()
def test_medical_document_uses_ocr_result(authorization_context):
context = authorization_context
general = {
"content": json.dumps({
"category": "medical_document",
"summary": "血常规报告",
"visible_text": "初步文字",
"key_facts": [],
"uncertainties": [],
"medical": {"document_type": "检验报告"},
}, ensure_ascii=False),
"usage": {},
}
ocr = {"content": "白细胞 11.2 x10^9/L", "usage": {}}
with (
patch("routers.chat.prepare_image", return_value=PreparedImage(b"jpeg", "image/jpeg", 100, 80)),
patch("routers.chat._run_billed_vision_call", side_effect=[general, ocr]) as model,
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/chat/images",
headers=context["owner_headers"],
files={"file": ("report.jpg", b"image-bytes", "image/jpeg")},
)
assert response.status_code == 200
attachment_id = response.json()["data"]["id"]
assert model.call_count == 2
assert model.call_args_list[1].kwargs["source"] == "vision_medical_ocr"
db = SessionLocal()
try:
stored = db.query(ChatAttachment).filter(ChatAttachment.id == attachment_id).one()
assert stored.extracted_text == "白细胞 11.2 x10^9/L"
assert stored.ocr_model == "qwen-vl-ocr"
assert "不能替代医生诊断" in stored.warning
finally:
db.close()
def test_attachment_cannot_cross_avatar_boundary(authorization_context):
context = authorization_context
db = SessionLocal()
try:
attachment = ChatAttachment(
avatar_id=context["avatar"].id,
filename="private.jpg",
status="ready",
expires_at=datetime.utcnow() + timedelta(hours=1),
)
db.add(attachment)
db.commit()
body = ChatIn(message="看看图片", attachmentIds=[attachment.id])
with pytest.raises(HTTPException, match="不属于当前分身") as caught:
_load_chat_attachments(db, context["other_avatar"].id, body)
assert caught.value.status_code == 400
finally:
db.close()
def test_expired_attachment_is_removed(authorization_context):
context = authorization_context
db = SessionLocal()
try:
attachment = ChatAttachment(
avatar_id=context["avatar"].id,
filename="expired.jpg",
status="ready",
expires_at=datetime.utcnow() - timedelta(seconds=1),
)
db.add(attachment)
db.commit()
attachment_id = attachment.id
body = ChatIn(message="看看图片", attachmentIds=[attachment_id])
with pytest.raises(HTTPException):
_load_chat_attachments(db, context["avatar"].id, body)
assert db.query(ChatAttachment).filter(ChatAttachment.id == attachment_id).first() is None
finally:
db.close()
def test_cleanup_keeps_unexpired_attachment(authorization_context):
context = authorization_context
now = datetime.utcnow()
db = SessionLocal()
try:
expired = ChatAttachment(
avatar_id=context["avatar"].id,
filename="expired.jpg",
status="ready",
expires_at=now - timedelta(seconds=1),
)
active = ChatAttachment(
avatar_id=context["avatar"].id,
filename="active.jpg",
status="ready",
expires_at=now + timedelta(hours=1),
)
db.add_all([expired, active])
db.commit()
expired_id, active_id = expired.id, active.id
assert purge_expired_chat_attachments(db, now=now) == 1
assert db.get(ChatAttachment, expired_id) is None
assert db.get(ChatAttachment, active_id) is not None
finally:
db.close()
def test_image_context_keeps_standard_answer_authoritative():
avatar = SimpleNamespace(
id="avatar-vision",
name="测试分身",
description="产品顾问",
config={},
)
model = Mock(return_value="标准退款期限是七天;图片显示的是商品包装。")
result = _resolve_reply(
None,
avatar,
"退款期限是多少?",
[],
qa_pairs=[SimpleNamespace(question="退款期限是多少?", answer="七天", enabled=True)],
search_fn=Mock(return_value=[]),
model_client=model,
image_contexts=[{
"id": "attachment",
"filename": "product.jpg",
"category": "general_image",
"summary": "商品包装",
"extractedText": "",
"structuredData": {},
"warning": "",
}],
)
assert result["source"] == "qa"
system = model.call_args.kwargs["messages"][0]["content"]
assert "已确认标准答案" in system
assert "七天" in system
assert "商品包装" in system
assert "标准答题对中的事实优先级高于图片资料" in system
def test_attachment_context_does_not_expose_internal_fields():
row = SimpleNamespace(
id="attachment",
filename="case.jpg",
category="medical_document",
summary="门诊病例",
extracted_text="主诉:咳嗽",
structured_data={"medical": {"chief_complaint": "咳嗽"}},
warning="请核对原文",
)
context = _attachment_contexts([row])[0]
assert context["filename"] == "case.jpg"
assert "avatar_id" not in context
assert "vision_model" not in context
@@ -26,6 +26,8 @@ def test_admin_runtime_config_takes_priority(monkeypatch):
"api_base_url": "https://model.test/v1/",
"api_key": "runtime-key",
"model": "avatar-model",
"vision_model": "avatar-vision-model",
"ocr_model": "avatar-ocr-model",
"max_tokens": 2048,
"timeout_seconds": 42,
}
@@ -37,6 +39,8 @@ def test_admin_runtime_config_takes_priority(monkeypatch):
assert config.source == "admin"
assert config.api_base_url == "https://model.test/v1"
assert config.model == "avatar-model"
assert config.vision_model == "avatar-vision-model"
assert config.ocr_model == "avatar-ocr-model"
assert config.max_tokens == 2048
request.assert_called_once_with(
"http://config.test/runtime",
@@ -51,6 +55,8 @@ def test_runtime_failure_falls_back_to_environment(monkeypatch):
monkeypatch.setenv("CHAT_API_URL", "https://fallback.test/v1/")
monkeypatch.setenv("CHAT_API_KEY", "fallback-key")
monkeypatch.setenv("CHAT_MODEL", "fallback-model")
monkeypatch.setenv("VISION_MODEL", "fallback-vision")
monkeypatch.setenv("VISION_OCR_MODEL", "fallback-ocr")
monkeypatch.setenv("CHAT_MAX_OUTPUT_TOKENS", "1536")
request = httpx.Request("GET", "http://config.test/runtime")
@@ -64,6 +70,8 @@ def test_runtime_failure_falls_back_to_environment(monkeypatch):
assert config.api_base_url == "https://fallback.test/v1"
assert config.api_key == "fallback-key"
assert config.model == "fallback-model"
assert config.vision_model == "fallback-vision"
assert config.ocr_model == "fallback-ocr"
assert config.max_tokens == 1536
@@ -20,8 +20,9 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
):
import main
maintenance_scheduler = MagicMock()
scheduler = MagicMock()
mock_scheduler_class.return_value = scheduler
mock_scheduler_class.side_effect = [maintenance_scheduler, scheduler]
boxim = MagicMock()
mock_boxim_class.return_value = boxim
takeover = MagicMock()
@@ -47,6 +48,10 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
assert config["BOXIM_API_BASE_URL"] == "https://im.example/api"
mock_takeover_class.assert_called_once_with(main.SessionLocal, boxim)
maintenance_scheduler.add_job.assert_called_once()
assert maintenance_scheduler.add_job.call_args.kwargs["id"] == "chat_attachment_cleanup"
maintenance_scheduler.start.assert_called_once_with()
assert scheduler.add_job.call_count == 2
poll_call, process_call = scheduler.add_job.call_args_list
assert poll_call.args[0] is takeover.poll_messages
@@ -62,6 +67,7 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
scheduler.start.assert_called_once_with()
main.takeover_scheduler = None
main.maintenance_scheduler = None
@patch("main.AsyncIOScheduler")
@@ -73,6 +79,7 @@ def test_scheduler_failure_does_not_stop_the_api(mock_scheduler_class):
main.on_startup()
assert main.takeover_scheduler is None
assert main.maintenance_scheduler is None
def test_shutdown_stops_only_the_scheduler():
@@ -80,9 +87,14 @@ def test_shutdown_stops_only_the_scheduler():
scheduler = MagicMock()
scheduler.running = True
maintenance_scheduler = MagicMock()
maintenance_scheduler.running = True
main.takeover_scheduler = scheduler
main.maintenance_scheduler = maintenance_scheduler
main.on_shutdown()
scheduler.shutdown.assert_called_once_with(wait=False)
maintenance_scheduler.shutdown.assert_called_once_with(wait=False)
assert main.takeover_scheduler is None
assert main.maintenance_scheduler is None
@@ -0,0 +1,98 @@
import io
import json
from unittest.mock import Mock, patch
import pytest
from PIL import Image
from services.chat_model_config import ChatModelConfig
from services.vision_service import (
ImageValidationError,
build_attachment_warning,
call_vision_model,
parse_vision_analysis,
prepare_image,
)
def _image_bytes(fmt="PNG", size=(120, 80)):
output = io.BytesIO()
Image.new("RGB", size, "#f97316").save(output, format=fmt)
return output.getvalue()
def _config():
return ChatModelConfig(
api_base_url="https://model.test/v1",
api_key="secret-key",
model="chat-model",
max_tokens=1024,
timeout_seconds=30,
vision_model="vision-model",
ocr_model="ocr-model",
vision_max_tokens=2048,
vision_timeout_seconds=90,
source="test",
)
def test_prepare_image_validates_and_reencodes_without_metadata():
prepared = prepare_image(_image_bytes())
assert prepared.mime_type == "image/jpeg"
assert prepared.width == 120
assert prepared.height == 80
with Image.open(io.BytesIO(prepared.data)) as image:
assert image.format == "JPEG"
assert not image.getexif()
def test_prepare_image_rejects_non_image_content():
with pytest.raises(ImageValidationError, match="格式无效"):
prepare_image(b"not-an-image")
def test_vision_request_uses_openai_compatible_image_content():
response = Mock()
response.raise_for_status.return_value = None
response.json.return_value = {
"choices": [{"message": {"content": '{"category":"general_image"}'}}],
"usage": {"total_tokens": 88},
}
prepared = prepare_image(_image_bytes())
with patch("services.vision_service.httpx.post", return_value=response) as request:
result = call_vision_model(
prepared,
_config(),
model="vision-model",
prompt="describe",
json_output=True,
)
payload = request.call_args.kwargs["json"]
content = payload["messages"][0]["content"]
assert payload["model"] == "vision-model"
assert payload["response_format"] == {"type": "json_object"}
assert content[0]["type"] == "image_url"
assert content[0]["image_url"]["url"].startswith("data:image/jpeg;base64,")
assert content[1] == {"type": "text", "text": "describe"}
assert result["usage"]["total_tokens"] == 88
def test_parse_medical_analysis_and_build_warning():
analysis = parse_vision_analysis(json.dumps({
"category": "medical_document",
"summary": "血常规报告",
"visible_text": "白细胞 11.2",
"key_facts": ["白细胞偏高"],
"uncertainties": ["日期模糊"],
"medical": {"document_type": "检验报告"},
}, ensure_ascii=False))
assert analysis["category"] == "medical_document"
assert analysis["medical"]["document_type"] == "检验报告"
warning = build_attachment_warning(analysis, ocr_failed=True)
assert "日期模糊" in warning
assert "人工核对" in warning
assert "不能替代医生诊断" in warning