feat(avatar): add private vision chat support
This commit is contained in:
@@ -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
|
||||
Reference in New Issue
Block a user