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.token_billing import InsufficientTokensError 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_image_upload_preserves_insufficient_points_response(authorization_context): context = authorization_context with patch( "routers.chat._analyze_image_bytes", side_effect=InsufficientTokensError("积分余额不足"), ): response = client.post( f"/api/avatar/{context['avatar'].id}/chat/images", headers=context["owner_headers"], files={"file": ("private.png", b"image-bytes", "image/png")}, ) assert response.status_code == 402 assert response.json()["detail"] == "积分余额不足" 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