274 lines
9.1 KiB
Python
274 lines
9.1 KiB
Python
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
|