From 016bc22c05008a7e1830477fcff82771fea40dea Mon Sep 17 00:00:00 2001 From: stefanfeng Date: Mon, 31 Aug 2026 15:10:50 +0800 Subject: [PATCH] feat(avatar): add private vision chat support --- backend/app/api/endpoints/ai_models.py | 6 + backend/app/core/database.py | 39 +- backend/app/models/__init__.py | 2 + backend/app/schemas/__init__.py | 6 + digital-avatar-app/backend/main.py | 56 +++ digital-avatar-app/backend/models.py | 38 ++ digital-avatar-app/backend/requirements.txt | 1 + digital-avatar-app/backend/routers/chat.py | 414 +++++++++++++++++- .../services/chat_attachment_service.py | 20 + .../backend/services/chat_model_config.py | 22 + .../backend/services/token_billing.py | 7 +- .../backend/services/vision_service.py | 196 +++++++++ digital-avatar-app/backend/tests/conftest.py | 4 + .../backend/tests/test_chat_images.py | 273 ++++++++++++ .../backend/tests/test_chat_model_config.py | 8 + .../backend/tests/test_takeover_scheduler.py | 14 +- .../backend/tests/test_vision_service.py | 98 +++++ .../docs/H5_PRODUCTION_DEPLOYMENT.md | 17 +- .../IMAGE_MEDICAL_UNDERSTANDING_DESIGN.md | 184 ++++++++ digital-avatar-app/src/api/index.ts | 56 ++- digital-avatar-app/src/views/AvatarChat.vue | 197 ++++++++- frontend/src/views/AIModels.vue | 17 +- 22 files changed, 1616 insertions(+), 59 deletions(-) create mode 100644 digital-avatar-app/backend/services/chat_attachment_service.py create mode 100644 digital-avatar-app/backend/services/vision_service.py create mode 100644 digital-avatar-app/backend/tests/test_chat_images.py create mode 100644 digital-avatar-app/backend/tests/test_vision_service.py create mode 100644 digital-avatar-app/docs/IMAGE_MEDICAL_UNDERSTANDING_DESIGN.md diff --git a/backend/app/api/endpoints/ai_models.py b/backend/app/api/endpoints/ai_models.py index b2d3c34..cb702f2 100755 --- a/backend/app/api/endpoints/ai_models.py +++ b/backend/app/api/endpoints/ai_models.py @@ -37,6 +37,8 @@ async def create_model(req: AIModelCreateRequest, db=Depends(get_db)): api_base_url=req.api_base_url, api_key_enc=encrypt(req.api_key) if req.api_key else None, model_version=req.model_version, + vision_model_version=req.vision_model_version, + ocr_model_version=req.ocr_model_version, temperature=req.temperature, max_tokens=req.max_tokens, timeout_seconds=req.timeout_seconds, @@ -100,6 +102,8 @@ async def get_digital_avatar_runtime_model( "api_base_url": model.api_base_url or "https://api.openai.com/v1", "api_key": decrypt(model.api_key_enc) if model.api_key_enc else "", "model": model.model_version or model.model_name, + "vision_model": model.vision_model_version or "qwen3.6-flash", + "ocr_model": model.ocr_model_version or "qwen-vl-ocr", "temperature": model.temperature, "max_tokens": model.max_tokens, "timeout_seconds": model.timeout_seconds, @@ -129,6 +133,8 @@ def _format_model(m: AIModelConfig) -> dict: "usage_scope": m.usage_scope, "api_base_url": m.api_base_url, "has_api_key": bool(m.api_key_enc), "model_version": m.model_version, "temperature": m.temperature, + "vision_model_version": m.vision_model_version, + "ocr_model_version": m.ocr_model_version, "max_tokens": m.max_tokens, "timeout_seconds": m.timeout_seconds, "is_default": m.is_default, "is_enabled": m.is_enabled, "created_at": m.created_at.isoformat(), diff --git a/backend/app/core/database.py b/backend/app/core/database.py index 80343ca..2024d04 100755 --- a/backend/app/core/database.py +++ b/backend/app/core/database.py @@ -66,21 +66,36 @@ async def init_db(): PendingReplyTask, TokenStat, AIModelConfig, SystemConfig, LoginLog ) async with engine.begin() as conn: - await conn.execute(text("SELECT GET_LOCK('ai_model_usage_scope_migration', 30)")) + await conn.execute(text("SELECT GET_LOCK('ai_model_config_migration', 30)")) try: - result = await conn.execute(text( - "SELECT COUNT(*) FROM information_schema.COLUMNS " - "WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'ai_model_configs' " - "AND COLUMN_NAME = 'usage_scope'" - )) - if result.scalar_one() == 0: - await conn.execute(text( + columns = ( + ( + "usage_scope", "ALTER TABLE ai_model_configs ADD COLUMN usage_scope " - "VARCHAR(16) NOT NULL DEFAULT 'general' AFTER provider" - )) - logger.info("AI模型配置表已增加 usage_scope 字段") + "VARCHAR(16) NOT NULL DEFAULT 'general' AFTER provider", + ), + ( + "vision_model_version", + "ALTER TABLE ai_model_configs ADD COLUMN vision_model_version " + "VARCHAR(64) NULL AFTER model_version", + ), + ( + "ocr_model_version", + "ALTER TABLE ai_model_configs ADD COLUMN ocr_model_version " + "VARCHAR(64) NULL AFTER vision_model_version", + ), + ) + for column_name, ddl in columns: + result = await conn.execute(text( + "SELECT COUNT(*) FROM information_schema.COLUMNS " + "WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'ai_model_configs' " + "AND COLUMN_NAME = :column_name" + ), {"column_name": column_name}) + if result.scalar_one() == 0: + await conn.execute(text(ddl)) + logger.info("AI模型配置表已增加 %s 字段", column_name) finally: - await conn.execute(text("SELECT RELEASE_LOCK('ai_model_usage_scope_migration')")) + await conn.execute(text("SELECT RELEASE_LOCK('ai_model_config_migration')")) logger.info("✅ 数据库模型注册成功") logger.info("✅ 数据库初始化完成") diff --git a/backend/app/models/__init__.py b/backend/app/models/__init__.py index aa33965..a98bdee 100755 --- a/backend/app/models/__init__.py +++ b/backend/app/models/__init__.py @@ -126,6 +126,8 @@ class AIModelConfig(Base): api_base_url: Mapped[str | None] = mapped_column(String(256)) api_key_enc: Mapped[str | None] = mapped_column(String(512)) model_version: Mapped[str | None] = mapped_column(String(64)) + vision_model_version: Mapped[str | None] = mapped_column(String(64)) + ocr_model_version: Mapped[str | None] = mapped_column(String(64)) temperature: Mapped[float] = mapped_column(Float, default=0.7) max_tokens: Mapped[int] = mapped_column(Integer, default=1000) timeout_seconds: Mapped[int] = mapped_column(Integer, default=30) diff --git a/backend/app/schemas/__init__.py b/backend/app/schemas/__init__.py index bc76534..8def8c4 100755 --- a/backend/app/schemas/__init__.py +++ b/backend/app/schemas/__init__.py @@ -158,6 +158,8 @@ class AIModelCreateRequest(BaseModel): api_base_url: Optional[str] = None api_key: Optional[str] = None model_version: Optional[str] = None + vision_model_version: Optional[str] = Field(None, max_length=64) + ocr_model_version: Optional[str] = Field(None, max_length=64) temperature: float = Field(default=0.7, ge=0.0, le=2.0) max_tokens: int = Field(default=1000, ge=1, le=32000) timeout_seconds: int = Field(default=30, ge=5, le=300) @@ -171,6 +173,8 @@ class AIModelUpdateRequest(BaseModel): api_base_url: Optional[str] = None api_key: Optional[str] = None model_version: Optional[str] = None + vision_model_version: Optional[str] = Field(None, max_length=64) + ocr_model_version: Optional[str] = Field(None, max_length=64) temperature: Optional[float] = Field(None, ge=0.0, le=2.0) max_tokens: Optional[int] = Field(None, ge=1, le=32000) timeout_seconds: Optional[int] = Field(None, ge=5, le=300) @@ -186,6 +190,8 @@ class AIModelResponse(BaseModel): api_base_url: Optional[str] has_api_key: bool model_version: Optional[str] + vision_model_version: Optional[str] + ocr_model_version: Optional[str] temperature: float max_tokens: int timeout_seconds: int diff --git a/digital-avatar-app/backend/main.py b/digital-avatar-app/backend/main.py index 5b8349c..e97a85f 100644 --- a/digital-avatar-app/backend/main.py +++ b/digital-avatar-app/backend/main.py @@ -19,11 +19,13 @@ import routers.huihui_auth import routers.chat import routers.takeover from responses import ok +from services.chat_attachment_service import purge_expired_chat_attachments from services.token_billing import DEFAULT_TOKEN_GRANT, release_stale_reservations logger = logging.getLogger(__name__) takeover_scheduler = None +maintenance_scheduler = None app = FastAPI(title="会会数字分身 API", version="1.0.0") @@ -132,6 +134,15 @@ def on_startup(): # Release stale resources when startup is invoked again by a reload/test. stop_takeover_scheduler() + stop_maintenance_scheduler() + try: + start_maintenance_scheduler() + except Exception as exc: + stop_maintenance_scheduler() + logger.warning( + "Failed to initialize chat attachment cleanup, app will continue: %s", + exc, + ) # --- Takeover scheduler --- try: @@ -196,6 +207,51 @@ def stop_takeover_scheduler(): finally: takeover_scheduler = None + +def purge_expired_chat_attachments_job(): + db = SessionLocal() + try: + count = purge_expired_chat_attachments(db) + if count: + logger.info("Purged %s expired chat image attachment(s)", count) + except Exception as exc: + db.rollback() + logger.warning("Failed to purge expired chat image attachments: %s", exc) + finally: + db.close() + + +def start_maintenance_scheduler(): + global maintenance_scheduler + + purge_expired_chat_attachments_job() + interval_minutes = max( + 5, min(1440, int(os.getenv("CHAT_ATTACHMENT_CLEANUP_MINUTES", "60"))) + ) + maintenance_scheduler = AsyncIOScheduler() + maintenance_scheduler.add_job( + purge_expired_chat_attachments_job, + trigger=IntervalTrigger(minutes=interval_minutes), + id="chat_attachment_cleanup", + max_instances=1, + coalesce=True, + ) + maintenance_scheduler.start() + + +def stop_maintenance_scheduler(): + global maintenance_scheduler + + if maintenance_scheduler is not None: + try: + if maintenance_scheduler.running: + maintenance_scheduler.shutdown(wait=False) + except Exception as exc: + logger.warning("Failed to stop maintenance scheduler cleanly: %s", exc) + finally: + maintenance_scheduler = None + @app.on_event("shutdown") def on_shutdown(): stop_takeover_scheduler() + stop_maintenance_scheduler() diff --git a/digital-avatar-app/backend/models.py b/digital-avatar-app/backend/models.py index 19717fe..828fc29 100644 --- a/digital-avatar-app/backend/models.py +++ b/digital-avatar-app/backend/models.py @@ -257,6 +257,44 @@ class KnowledgeChunk(Base): } +class ChatAttachment(Base): + """Private, avatar-scoped result of one chat image analysis.""" + + __tablename__ = "chat_attachments" + id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex) + avatar_id = Column(String, nullable=False, default="", index=True) + uploader_kind = Column(String, default="owner") # owner | public + filename = Column(String, default="") + mime_type = Column(String, default="") + file_size = Column(Integer, default=0) + status = Column(String, default="processing") # processing | ready | failed + category = Column(String, default="general_image") + summary = Column(Text, default="") + extracted_text = Column(Text, default="") + structured_data = Column(JSON, default=dict) + warning = Column(Text, default="") + vision_model = Column(String, default="") + ocr_model = Column(String, default="") + used_at = Column(DateTime) + expires_at = Column(DateTime, nullable=False) + created_at = Column(DateTime, server_default=func.now()) + + def to_dict(self): + return { + "id": self.id, + "avatarId": self.avatar_id, + "filename": self.filename, + "mimeType": self.mime_type, + "fileSize": self.file_size, + "status": self.status, + "category": self.category, + "summary": self.summary, + "warning": self.warning, + "expiresAt": _iso(self.expires_at), + "createdAt": _iso(self.created_at), + } + + class TokenAccount(Base): __tablename__ = "token_account" id = Column(Integer, primary_key=True) diff --git a/digital-avatar-app/backend/requirements.txt b/digital-avatar-app/backend/requirements.txt index f744392..57d28df 100644 --- a/digital-avatar-app/backend/requirements.txt +++ b/digital-avatar-app/backend/requirements.txt @@ -8,3 +8,4 @@ pypdf python-docx openpyxl apscheduler>=3.10 +Pillow>=10.4 diff --git a/digital-avatar-app/backend/routers/chat.py b/digital-avatar-app/backend/routers/chat.py index fefdee8..83b4c2b 100644 --- a/digital-avatar-app/backend/routers/chat.py +++ b/digital-avatar-app/backend/routers/chat.py @@ -1,21 +1,32 @@ import difflib import json +import logging import os import re import secrets import string +from datetime import datetime, timedelta from typing import Any, Callable import httpx -from fastapi import APIRouter, Body, Depends, Header, HTTPException +from fastapi import APIRouter, Body, Depends, File, Header, HTTPException, UploadFile from fastapi.responses import StreamingResponse -from pydantic import BaseModel, Field +from pydantic import BaseModel, ConfigDict, Field, model_validator from sqlalchemy.orm import Session import embeddings from database import get_db -from models import Avatar, KnowledgeChunk, KnowledgeDoc, QAPair, User +from models import Avatar, ChatAttachment, KnowledgeChunk, KnowledgeDoc, QAPair, User from responses import ok, fail +from services.vision_service import ( + GENERAL_VISION_PROMPT, + MEDICAL_OCR_PROMPT, + ImageValidationError, + build_attachment_warning, + call_vision_model, + parse_vision_analysis, + prepare_image, +) from services.token_billing import ( InsufficientTokensError, estimate_fallback_usage, @@ -24,8 +35,10 @@ from services.token_billing import ( settle_reservation, ) from services.chat_model_config import ChatModelConfig, get_chat_model_config +from services.chat_attachment_service import purge_expired_chat_attachments router = APIRouter(tags=["数字分身聊天"]) +logger = logging.getLogger(__name__) MAX_MESSAGE_LENGTH = 4000 MAX_HISTORY_MESSAGES = 10 @@ -49,14 +62,35 @@ _KOREAN_HANGUL = re.compile(r"[\uac00-\ud7af\u1100-\u11ff]") class ChatMessage(BaseModel): + model_config = ConfigDict(populate_by_name=True) + role: str = Field(pattern="^(user|assistant)$") content: str = Field(min_length=1, max_length=MAX_MESSAGE_LENGTH) + attachment_ids: list[str] = Field( + default_factory=list, + alias="attachmentIds", + max_length=3, + ) class ChatIn(BaseModel): - message: str = Field(min_length=1, max_length=MAX_MESSAGE_LENGTH) + model_config = ConfigDict(populate_by_name=True) + + message: str = Field(default="", max_length=MAX_MESSAGE_LENGTH) + attachment_ids: list[str] = Field( + default_factory=list, + alias="attachmentIds", + max_length=3, + ) history: list[ChatMessage] = Field(default_factory=list, max_length=MAX_HISTORY_MESSAGES) + @model_validator(mode="after") + def require_message_or_image(self): + self.message = self.message.strip() + if not self.message and not self.attachment_ids: + raise ValueError("请输入消息或选择图片") + return self + def _resolve_user(authorization: str | None, db: Session): if not authorization: @@ -77,6 +111,228 @@ def _require_owned_avatar(db: Session, avatar_id: str, authorization: str | None return avatar +def _attachment_expiry() -> datetime: + retention_hours = max( + 1, min(168, int(os.getenv("CHAT_ATTACHMENT_RETENTION_HOURS", "24"))) + ) + return datetime.utcnow() + timedelta(hours=retention_hours) + + +def _chat_attachment_ids(body: ChatIn) -> list[str]: + values = list(body.attachment_ids) + for message in body.history[-MAX_HISTORY_MESSAGES:]: + values.extend(message.attachment_ids) + unique = list(dict.fromkeys(str(value).strip() for value in values if str(value).strip())) + if len(unique) > 3: + raise HTTPException(status_code=400, detail="一次会话最多引用 3 张图片") + return unique + + +def _load_chat_attachments(db: Session, avatar_id: str, body: ChatIn) -> list[ChatAttachment]: + attachment_ids = _chat_attachment_ids(body) + if not attachment_ids: + return [] + purge_expired_chat_attachments(db) + rows = db.query(ChatAttachment).filter( + ChatAttachment.avatar_id == avatar_id, + ChatAttachment.id.in_(attachment_ids), + ).all() + by_id = {row.id: row for row in rows} + if len(by_id) != len(attachment_ids): + raise HTTPException(status_code=400, detail="图片资料不存在、已过期或不属于当前分身") + ordered = [by_id[attachment_id] for attachment_id in attachment_ids] + if any(row.status != "ready" for row in ordered): + raise HTTPException(status_code=409, detail="图片尚未识别完成,请稍后重试") + now = datetime.utcnow() + for row in ordered: + row.used_at = now + db.commit() + return ordered + + +def _attachment_contexts(rows: list[ChatAttachment]) -> list[dict]: + contexts = [] + remaining_text = 12000 + for row in rows: + extracted = (row.extracted_text or "")[:remaining_text] + remaining_text = max(0, remaining_text - len(extracted)) + contexts.append({ + "id": row.id, + "filename": row.filename, + "category": row.category, + "summary": row.summary, + "extractedText": extracted, + "structuredData": row.structured_data or {}, + "warning": row.warning, + }) + return contexts + + +def _image_retrieval_question(question: str, image_contexts: list[dict]) -> str: + parts = [question.strip()] + for context in image_contexts: + parts.extend([ + str(context.get("summary") or "")[:600], + str(context.get("extractedText") or "")[:1200], + ]) + return "\n".join(part for part in parts if part).strip() + + +def _run_billed_vision_call( + db: Session, + avatar: Avatar, + prepared, + *, + model: str, + prompt: str, + source: str, + json_output: bool, + model_config: ChatModelConfig, +) -> dict: + estimate_messages = [{ + "role": "user", + "content": f"[一张待识别图片]\n{prompt}", + }] + reservation = reserve_avatar_tokens( + db, + avatar, + source, + model, + estimate_messages, + model_config.vision_max_tokens, + minimum_reserve_tokens=max( + 1000, int(os.getenv("VISION_TOKEN_RESERVE", "12000")) + ), + ) + try: + result = call_vision_model( + prepared, + model_config, + model=model, + prompt=prompt, + json_output=json_output, + ) + settle_reservation( + db, + reservation, + result.get("usage"), + fallback_total=estimate_fallback_usage( + estimate_messages, result.get("content") or "" + ), + ) + return result + except Exception as exc: + release_reservation(db, reservation, str(exc)) + raise + + +async def _analyze_uploaded_image( + db: Session, + avatar: Avatar, + file: UploadFile, + *, + uploader_kind: str, +) -> ChatAttachment: + max_bytes = max(1024, int(os.getenv("CHAT_IMAGE_MAX_BYTES", str(8 * 1024 * 1024)))) + content = await file.read(max_bytes + 1) + filename = os.path.basename(file.filename or "图片")[:255] + attachment = ChatAttachment( + avatar_id=avatar.id, + uploader_kind=uploader_kind, + filename=filename, + mime_type=(file.content_type or "")[:100], + file_size=len(content), + status="processing", + expires_at=_attachment_expiry(), + ) + db.add(attachment) + db.commit() + db.refresh(attachment) + + try: + prepared = prepare_image(content) + model_config = get_chat_model_config() + vision_result = _run_billed_vision_call( + db, + avatar, + prepared, + model=model_config.vision_model, + prompt=GENERAL_VISION_PROMPT, + source="vision_image", + json_output=True, + model_config=model_config, + ) + analysis = parse_vision_analysis(vision_result["content"]) + extracted_text = analysis.get("visible_text") or "" + ocr_model = "" + ocr_failed = False + if analysis["category"] == "medical_document" and model_config.ocr_model: + try: + ocr_result = _run_billed_vision_call( + db, + avatar, + prepared, + model=model_config.ocr_model, + prompt=MEDICAL_OCR_PROMPT, + source="vision_medical_ocr", + json_output=False, + model_config=model_config, + ) + extracted_text = ocr_result["content"] + ocr_model = model_config.ocr_model + except (RuntimeError, InsufficientTokensError): + ocr_failed = True + logger.warning( + "medical OCR degraded for attachment %s avatar %s", + attachment.id, + avatar.id, + ) + + attachment.mime_type = prepared.mime_type + attachment.status = "ready" + attachment.category = analysis["category"] + attachment.summary = analysis.get("summary") or "图片内容已识别" + attachment.extracted_text = extracted_text + attachment.structured_data = analysis + attachment.warning = build_attachment_warning(analysis, ocr_failed=ocr_failed) + attachment.vision_model = model_config.vision_model + attachment.ocr_model = ocr_model + db.commit() + db.refresh(attachment) + logger.info( + "chat image ready attachment=%s avatar=%s category=%s model=%s ocr=%s", + attachment.id, + avatar.id, + attachment.category, + attachment.vision_model, + bool(attachment.ocr_model), + ) + return attachment + except ImageValidationError as exc: + attachment.status = "failed" + attachment.warning = str(exc) + db.commit() + raise HTTPException(status_code=400, detail=str(exc)) from exc + except InsufficientTokensError: + attachment.status = "failed" + attachment.warning = "积分余额不足" + db.commit() + raise + except RuntimeError as exc: + attachment.status = "failed" + attachment.warning = str(exc) + db.commit() + logger.warning( + "chat image failed attachment=%s avatar=%s error=%s", + attachment.id, + avatar.id, + type(exc).__name__, + ) + raise HTTPException(status_code=502, detail=str(exc)) from exc + finally: + content = b"" + + def _normalize_question(value: str) -> str: value = (value or "").strip().lower() value = re.sub(r"\s+", "", value) @@ -233,6 +489,7 @@ def _build_prompt( knowledge_hits: list[dict], *, standard_answer: str = "", + image_contexts: list[dict] | None = None, ) -> list[dict]: config = _config(avatar) description = (getattr(avatar, "description", "") or "").strip() @@ -241,6 +498,7 @@ def _build_prompt( for hit in knowledge_hits if hit.get("snippet") ) + image_contexts = image_contexts or [] profile_items = [ (label, config[key]) for label, key in ( @@ -265,6 +523,24 @@ def _build_prompt( ) if config["systemPrompt"]: system += f"\n额外系统提示词:{config['systemPrompt']}" + if image_contexts: + image_material = json.dumps(image_contexts, ensure_ascii=False, default=str) + system += ( + "\n以下是当前会话图片经过视觉识别后得到的资料:\n" + f"{image_material}" + "\n图片资料可能包含 OCR 错字、模糊内容或用户尚未确认的信息,只能按可见内容谨慎表达。" + "标准答题对中的事实优先级高于图片资料,知识库事实优先级高于模型推测;发生冲突时遵循更高优先级资料," + "并自然提醒对方核对原图。不得声称看到了图片中不存在的内容。" + ) + if any( + context.get("category") in {"medical_document", "medical_image"} + for context in image_contexts + ): + system += ( + "\n本次包含医疗资料。可以整理病例原文、解释指标含义和提示需要关注的异常,但不能仅凭图片作出" + "确定诊断、疾病分期、处方、停药或治疗决定。医学影像只能客观描述,并提醒结合正规报告和医生意见。" + "回答结尾用与用户相同的语言简短说明图片识别结果仅供辅助,不能替代医生诊断。" + ) if standard_answer: system += ( f"\n以下是本次问题命中的已确认标准答案:\n{standard_answer.strip()}" @@ -277,6 +553,11 @@ def _build_prompt( "\n涉及事实、专业判断、地址、流程、数据或建议时,只能依据本人资料、标准问答形成的上下文" "和以上可靠资料作答,不要补充资料之外的通用知识或自行推测。" ) + elif image_contexts: + system += ( + "\n本次没有命中标准答题对或文件知识库,但已提供图片识别资料。只能围绕图片中的可确认内容、" + "本人资料和当前对话作答;不要补充图片之外的事实、专业判断或具体建议。" + ) else: system += ( "\n本次问题没有检索到可靠资料。除自然寒暄和基于本人资料的回答外,不要凭通用知识给出事实、" @@ -439,14 +720,17 @@ def _resolve_reply( search_fn: Callable[..., list[dict]] | None = None, model_client: Callable[..., str] | None = None, usage_source: str = "chat", + image_contexts: list[dict] | None = None, ) -> dict: + image_contexts = image_contexts or [] + question = question.strip() or "请根据这张图片说明可确认的内容。" if qa_pairs is None: qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all() matched = _match_standard_qa(question, qa_pairs) adapt_qa_language = bool( matched and _qa_requires_language_adaptation(question, matched.answer) ) - if matched and not adapt_qa_language: + if matched and not adapt_qa_language and not image_contexts: return {"answer": matched.answer, "source": "qa", "references": []} if matched: @@ -457,11 +741,19 @@ def _resolve_reply( question, hits, standard_answer=matched.answer, + image_contexts=image_contexts, ) else: search_fn = search_fn or (lambda query, avatar_id: _search_knowledge(db, avatar_id, query)) - hits = search_fn(question, avatar.id) - messages = _build_prompt(avatar, history, question, hits) + retrieval_question = _image_retrieval_question(question, image_contexts) + hits = search_fn(retrieval_question, avatar.id) + messages = _build_prompt( + avatar, + history, + question, + hits, + image_contexts=image_contexts, + ) config = _config(avatar) temperature = 0.0 if matched else min( 0.45 if hits else 0.25, @@ -498,7 +790,9 @@ def _resolve_reply( raise result = { "answer": answer, - "source": "qa" if matched else ("knowledge" if hits else "qwen"), + "source": "qa" if matched else ( + "knowledge" if hits else ("vision" if image_contexts else "qwen") + ), "references": hits, } if token_usage: @@ -514,14 +808,17 @@ def _stream_reply( *, public: bool = False, usage_source: str = "chat_stream", + image_contexts: list[dict] | None = None, ): + image_contexts = image_contexts or [] + question = question.strip() or "请根据这张图片说明可确认的内容。" qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all() matched = _match_standard_qa(question, qa_pairs) adapt_qa_language = bool( matched and _qa_requires_language_adaptation(question, matched.answer) ) messages, reservation = [], None - if matched and not adapt_qa_language: + if matched and not adapt_qa_language and not image_contexts: source, references, chunks = "qa", [], _iter_text_chunks(matched.answer) else: if matched: @@ -533,11 +830,21 @@ def _stream_reply( question, references, standard_answer=matched.answer, + image_contexts=image_contexts, ) else: - references = _search_knowledge(db, avatar.id, question) - source = "knowledge" if references else "qwen" - messages = _build_prompt(avatar, history, question, references) + retrieval_question = _image_retrieval_question(question, image_contexts) + references = _search_knowledge(db, avatar.id, retrieval_question) + source = "knowledge" if references else ( + "vision" if image_contexts else "qwen" + ) + messages = _build_prompt( + avatar, + history, + question, + references, + image_contexts=image_contexts, + ) config = _config(avatar) temperature = 0.0 if matched else min( 0.45 if references else 0.25, @@ -641,11 +948,62 @@ def get_shared_avatar(share_token: str, db: Session = Depends(get_db)): return ok(_public_avatar_payload(_require_shared_avatar(db, share_token))) +@router.post("/avatar/{avatar_id}/chat/images") +async def upload_chat_image( + avatar_id: str, + file: UploadFile = File(...), + authorization: str = Header(None), + db: Session = Depends(get_db), +): + avatar = _require_owned_avatar(db, avatar_id, authorization) + purge_expired_chat_attachments(db) + try: + attachment = await _analyze_uploaded_image( + db, + avatar, + file, + uploader_kind="owner", + ) + return ok(attachment.to_dict()) + except InsufficientTokensError as exc: + raise HTTPException(status_code=402, detail=str(exc)) from exc + + +@router.post("/public/avatar/{share_token}/chat/images") +async def upload_public_chat_image( + share_token: str, + file: UploadFile = File(...), + db: Session = Depends(get_db), +): + avatar = _require_shared_avatar(db, share_token) + purge_expired_chat_attachments(db) + try: + attachment = await _analyze_uploaded_image( + db, + avatar, + file, + uploader_kind="public", + ) + return ok(attachment.to_dict()) + except InsufficientTokensError as exc: + raise HTTPException(status_code=402, detail=str(exc)) from exc + + @router.post("/public/avatar/{share_token}/chat") def public_chat(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)): avatar = _require_shared_avatar(db, share_token) + image_contexts = _attachment_contexts( + _load_chat_attachments(db, avatar.id, body) + ) try: - result = _resolve_reply(db, avatar, body.message, body.history, usage_source="public_chat") + result = _resolve_reply( + db, + avatar, + body.message, + body.history, + usage_source="public_chat", + image_contexts=image_contexts, + ) # 公开访客无需获知知识文件名、检索分数或内部答复来源。 result["references"] = [] result["source"] = "public" @@ -660,8 +1018,17 @@ def public_chat(share_token: str, body: ChatIn = Body(...), db: Session = Depend @router.post("/avatar/{avatar_id}/chat") def chat(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)): avatar = _require_owned_avatar(db, avatar_id, authorization) + image_contexts = _attachment_contexts( + _load_chat_attachments(db, avatar.id, body) + ) try: - return ok(_resolve_reply(db, avatar, body.message, body.history)) + return ok(_resolve_reply( + db, + avatar, + body.message, + body.history, + image_contexts=image_contexts, + )) except InsufficientTokensError as exc: return fail(str(exc), code=402) except RuntimeError as exc: @@ -671,7 +1038,17 @@ def chat(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(N @router.post("/avatar/{avatar_id}/chat/stream") def chat_stream(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)): try: - return _stream_reply(db, _require_owned_avatar(db, avatar_id, authorization), body.message, body.history) + avatar = _require_owned_avatar(db, avatar_id, authorization) + image_contexts = _attachment_contexts( + _load_chat_attachments(db, avatar.id, body) + ) + return _stream_reply( + db, + avatar, + body.message, + body.history, + image_contexts=image_contexts, + ) except InsufficientTokensError as exc: raise HTTPException(status_code=402, detail=str(exc)) from exc @@ -679,13 +1056,18 @@ def chat_stream(avatar_id: str, body: ChatIn = Body(...), authorization: str = H @router.post("/public/avatar/{share_token}/chat/stream") def public_chat_stream(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)): try: + avatar = _require_shared_avatar(db, share_token) + image_contexts = _attachment_contexts( + _load_chat_attachments(db, avatar.id, body) + ) return _stream_reply( db, - _require_shared_avatar(db, share_token), + avatar, body.message, body.history, public=True, usage_source="public_chat_stream", + image_contexts=image_contexts, ) except InsufficientTokensError as exc: raise HTTPException(status_code=402, detail=str(exc)) from exc diff --git a/digital-avatar-app/backend/services/chat_attachment_service.py b/digital-avatar-app/backend/services/chat_attachment_service.py new file mode 100644 index 0000000..501dd5b --- /dev/null +++ b/digital-avatar-app/backend/services/chat_attachment_service.py @@ -0,0 +1,20 @@ +from datetime import datetime + +from sqlalchemy.orm import Session + +from models import ChatAttachment + + +def purge_expired_chat_attachments( + db: Session, + *, + now: datetime | None = None, +) -> int: + """Remove expired derived image data; raw image bytes are never persisted.""" + count = db.query(ChatAttachment).filter( + ChatAttachment.expires_at < (now or datetime.utcnow()) + ).delete(synchronize_session=False) + if count: + db.commit() + db.expire_all() + return count diff --git a/digital-avatar-app/backend/services/chat_model_config.py b/digital-avatar-app/backend/services/chat_model_config.py index a2c71b1..40a4ae9 100644 --- a/digital-avatar-app/backend/services/chat_model_config.py +++ b/digital-avatar-app/backend/services/chat_model_config.py @@ -16,6 +16,10 @@ class ChatModelConfig: model: str max_tokens: int timeout_seconds: float + vision_model: str + ocr_model: str + vision_max_tokens: int + vision_timeout_seconds: float source: str @@ -33,6 +37,10 @@ def _environment_config() -> ChatModelConfig: model=os.getenv("CHAT_MODEL", "qwen-plus"), max_tokens=max(128, int(os.getenv("CHAT_MAX_OUTPUT_TOKENS", "1024"))), timeout_seconds=max(5.0, float(os.getenv("CHAT_TIMEOUT_SECONDS", "30"))), + vision_model=os.getenv("VISION_MODEL", "qwen3.6-flash"), + ocr_model=os.getenv("VISION_OCR_MODEL", "qwen-vl-ocr"), + vision_max_tokens=max(256, int(os.getenv("VISION_MAX_OUTPUT_TOKENS", "2048"))), + vision_timeout_seconds=max(10.0, float(os.getenv("VISION_TIMEOUT_SECONDS", "90"))), source="environment", ) @@ -60,6 +68,20 @@ def _fetch_runtime_config() -> ChatModelConfig | None: model=model, max_tokens=max(128, int(payload.get("max_tokens") or 1024)), timeout_seconds=max(5.0, float(payload.get("timeout_seconds") or 30)), + vision_model=str( + payload.get("vision_model") + or os.getenv("VISION_MODEL", "qwen3.6-flash") + ), + ocr_model=str( + payload.get("ocr_model") + or os.getenv("VISION_OCR_MODEL", "qwen-vl-ocr") + ), + vision_max_tokens=max( + 256, int(os.getenv("VISION_MAX_OUTPUT_TOKENS", "2048")) + ), + vision_timeout_seconds=max( + 10.0, float(os.getenv("VISION_TIMEOUT_SECONDS", "90")) + ), source="admin", ) diff --git a/digital-avatar-app/backend/services/token_billing.py b/digital-avatar-app/backend/services/token_billing.py index 2d483c3..28fb5ce 100644 --- a/digital-avatar-app/backend/services/token_billing.py +++ b/digital-avatar-app/backend/services/token_billing.py @@ -79,12 +79,17 @@ def reserve_avatar_tokens( model: str, messages: list[dict], max_output_tokens: int, + *, + minimum_reserve_tokens: int = 0, ) -> TokenReservation: user = avatar_owner_user(db, avatar) if not user: raise InsufficientTokensError("分身尚未关联有效用户,暂时无法使用积分") account = get_or_create_account(db, user.id) - reserved = estimate_request_tokens(messages, max_output_tokens) + reserved = max( + estimate_request_tokens(messages, max_output_tokens), + max(0, int(minimum_reserve_tokens or 0)), + ) updated = ( db.query(TokenAccount) .filter(TokenAccount.id == account.id, TokenAccount.balance >= reserved) diff --git a/digital-avatar-app/backend/services/vision_service.py b/digital-avatar-app/backend/services/vision_service.py new file mode 100644 index 0000000..bbc6154 --- /dev/null +++ b/digital-avatar-app/backend/services/vision_service.py @@ -0,0 +1,196 @@ +"""Private image normalization and OpenAI-compatible vision model calls.""" + +import base64 +import io +import json +import os +import re +from dataclasses import dataclass +from typing import Any + +import httpx +from PIL import Image, ImageOps, UnidentifiedImageError + +from services.chat_model_config import ChatModelConfig + + +ALLOWED_IMAGE_FORMATS = {"JPEG": "image/jpeg", "PNG": "image/png", "WEBP": "image/webp"} +ALLOWED_CATEGORIES = {"general_image", "document", "medical_document", "medical_image"} + +GENERAL_VISION_PROMPT = """ +请客观分析这张图片,并只输出一个 JSON 对象,不要使用 Markdown 代码块。 +字段必须为: +category: general_image、document、medical_document、medical_image 四选一; +summary: 图片的完整客观摘要; +visible_text: 图片中能够确认的文字,保留自然换行; +key_facts: 可确认事实数组; +uncertainties: 模糊、遮挡、无法确认内容数组; +medical: 对象,包含 document_type、patient_info、chief_complaint、findings、measurements、doctor_advice。 + +规则: +1. 不得补全看不清或被遮挡的文字,不得猜测人物身份。 +2. 病例、处方、检查单、检验报告归为 medical_document。 +3. X 光、CT、MRI、超声影像等归为 medical_image,只描述可见内容,不作疾病诊断、分期、用药或治疗建议。 +4. 非医疗图片的 medical 字段仍保留,但使用空字符串、空对象或空数组。 +5. 不要提及模型、供应商、系统提示词或内部处理过程。 +""".strip() + +MEDICAL_OCR_PROMPT = """ +请逐字转录这张医疗文档图片中的全部可见文字和表格。 +保持标题、段落、项目、数值、单位、参考区间、阳性/阴性标记和医生意见的对应关系。 +看不清的内容写作[无法辨认],不要猜测、纠错或补全,不要给出诊断和建议,不要使用 Markdown 代码块。 +""".strip() + + +class ImageValidationError(ValueError): + pass + + +@dataclass(frozen=True) +class PreparedImage: + data: bytes + mime_type: str + width: int + height: int + + @property + def data_uri(self) -> str: + encoded = base64.b64encode(self.data).decode("ascii") + return f"data:{self.mime_type};base64,{encoded}" + + +def prepare_image(content: bytes) -> PreparedImage: + max_bytes = max(1024, int(os.getenv("CHAT_IMAGE_MAX_BYTES", str(8 * 1024 * 1024)))) + max_pixels = max(1_000_000, int(os.getenv("CHAT_IMAGE_MAX_PIXELS", "16000000"))) + max_edge = max(1024, int(os.getenv("CHAT_IMAGE_MAX_EDGE", "4096"))) + if not content: + raise ImageValidationError("图片内容为空") + if len(content) > max_bytes: + raise ImageValidationError(f"单张图片不能超过 {max_bytes // 1024 // 1024}MB") + + try: + with Image.open(io.BytesIO(content)) as probe: + image_format = str(probe.format or "").upper() + width, height = probe.size + probe.verify() + except (UnidentifiedImageError, OSError, SyntaxError) as exc: + raise ImageValidationError("图片格式无效或文件已损坏") from exc + + if image_format not in ALLOWED_IMAGE_FORMATS: + raise ImageValidationError("仅支持 JPG、PNG、WebP 图片") + if width <= 0 or height <= 0 or width * height > max_pixels: + raise ImageValidationError("图片像素过大,请压缩后重新上传") + + try: + with Image.open(io.BytesIO(content)) as original: + image = ImageOps.exif_transpose(original) + image.load() + if max(image.size) > max_edge: + image.thumbnail((max_edge, max_edge), Image.Resampling.LANCZOS) + if image.mode in {"RGBA", "LA"}: + canvas = Image.new("RGB", image.size, "white") + alpha = image.getchannel("A") + canvas.paste(image.convert("RGB"), mask=alpha) + image = canvas + elif image.mode != "RGB": + image = image.convert("RGB") + output = io.BytesIO() + image.save(output, format="JPEG", quality=92, optimize=True) + normalized = output.getvalue() + normalized_width, normalized_height = image.size + except (OSError, ValueError) as exc: + raise ImageValidationError("图片解码失败,请重新选择图片") from exc + + return PreparedImage( + data=normalized, + mime_type="image/jpeg", + width=normalized_width, + height=normalized_height, + ) + + +def call_vision_model( + prepared: PreparedImage, + model_config: ChatModelConfig, + *, + model: str, + prompt: str, + json_output: bool, +) -> dict: + if not model_config.api_key: + raise RuntimeError("视觉模型服务未配置") + payload: dict[str, Any] = { + "model": model, + "messages": [ + { + "role": "user", + "content": [ + {"type": "image_url", "image_url": {"url": prepared.data_uri}}, + {"type": "text", "text": prompt}, + ], + } + ], + "temperature": 0, + "max_tokens": model_config.vision_max_tokens, + } + if json_output: + payload["response_format"] = {"type": "json_object"} + try: + response = httpx.post( + f"{model_config.api_base_url}/chat/completions", + headers={"Authorization": f"Bearer {model_config.api_key}"}, + json=payload, + timeout=model_config.vision_timeout_seconds, + ) + response.raise_for_status() + data = response.json() + content = data.get("choices", [{}])[0].get("message", {}).get("content", "") + except (httpx.HTTPError, ValueError, KeyError, IndexError) as exc: + raise RuntimeError("图片识别服务暂时不可用") from exc + if not isinstance(content, str) or not content.strip(): + raise RuntimeError("图片识别服务没有返回有效结果") + return {"content": content.strip(), "usage": data.get("usage") or {}} + + +def parse_vision_analysis(content: str) -> dict: + value = (content or "").strip() + fenced = re.match(r"^```(?:json)?\s*(.*?)\s*```$", value, re.DOTALL | re.IGNORECASE) + if fenced: + value = fenced.group(1).strip() + try: + payload = json.loads(value) + except (TypeError, ValueError) as exc: + raise RuntimeError("图片识别结果格式无效") from exc + if not isinstance(payload, dict): + raise RuntimeError("图片识别结果格式无效") + + category = str(payload.get("category") or "general_image").strip().lower() + if category not in ALLOWED_CATEGORIES: + category = "general_image" + medical = payload.get("medical") if isinstance(payload.get("medical"), dict) else {} + return { + "category": category, + "summary": str(payload.get("summary") or "").strip(), + "visible_text": str(payload.get("visible_text") or "").strip(), + "key_facts": _string_list(payload.get("key_facts")), + "uncertainties": _string_list(payload.get("uncertainties")), + "medical": medical, + } + + +def build_attachment_warning(analysis: dict, *, ocr_failed: bool = False) -> str: + warnings = list(analysis.get("uncertainties") or []) + category = analysis.get("category") + if ocr_failed: + warnings.append("精确文字识别暂时不可用,请人工核对图片原文") + if category == "medical_document": + warnings.append("病例识别结果仅供辅助,不能替代医生诊断,请核对原始文档") + elif category == "medical_image": + warnings.append("医学影像仅作客观描述,不能替代影像报告和医生诊断") + return ";".join(dict.fromkeys(item for item in warnings if item)) + + +def _string_list(value: Any) -> list[str]: + if not isinstance(value, list): + return [] + return [str(item).strip() for item in value if str(item).strip()] diff --git a/digital-avatar-app/backend/tests/conftest.py b/digital-avatar-app/backend/tests/conftest.py index ab45d48..1b62971 100644 --- a/digital-avatar-app/backend/tests/conftest.py +++ b/digital-avatar-app/backend/tests/conftest.py @@ -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) diff --git a/digital-avatar-app/backend/tests/test_chat_images.py b/digital-avatar-app/backend/tests/test_chat_images.py new file mode 100644 index 0000000..016391a --- /dev/null +++ b/digital-avatar-app/backend/tests/test_chat_images.py @@ -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 diff --git a/digital-avatar-app/backend/tests/test_chat_model_config.py b/digital-avatar-app/backend/tests/test_chat_model_config.py index 43bd77b..024a581 100644 --- a/digital-avatar-app/backend/tests/test_chat_model_config.py +++ b/digital-avatar-app/backend/tests/test_chat_model_config.py @@ -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 diff --git a/digital-avatar-app/backend/tests/test_takeover_scheduler.py b/digital-avatar-app/backend/tests/test_takeover_scheduler.py index fff0bce..671a745 100644 --- a/digital-avatar-app/backend/tests/test_takeover_scheduler.py +++ b/digital-avatar-app/backend/tests/test_takeover_scheduler.py @@ -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 diff --git a/digital-avatar-app/backend/tests/test_vision_service.py b/digital-avatar-app/backend/tests/test_vision_service.py new file mode 100644 index 0000000..175e858 --- /dev/null +++ b/digital-avatar-app/backend/tests/test_vision_service.py @@ -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 diff --git a/digital-avatar-app/docs/H5_PRODUCTION_DEPLOYMENT.md b/digital-avatar-app/docs/H5_PRODUCTION_DEPLOYMENT.md index cce121a..73f90d9 100644 --- a/digital-avatar-app/docs/H5_PRODUCTION_DEPLOYMENT.md +++ b/digital-avatar-app/docs/H5_PRODUCTION_DEPLOYMENT.md @@ -51,6 +51,16 @@ EMBEDDING_API_URL=https://dashscope.aliyuncs.com/compatible-mode/v1 EMBEDDING_API_KEY= EMBEDDING_MODEL=text-embedding-v3 EMBEDDING_BATCH_SIZE=10 + +VISION_MODEL=qwen3.6-flash +VISION_OCR_MODEL=qwen-vl-ocr +VISION_MAX_OUTPUT_TOKENS=2048 +VISION_TIMEOUT_SECONDS=90 +VISION_TOKEN_RESERVE=12000 +CHAT_IMAGE_MAX_BYTES=8388608 +CHAT_IMAGE_MAX_PIXELS=16000000 +CHAT_ATTACHMENT_RETENTION_HOURS=24 +CHAT_ATTACHMENT_CLEANUP_MINUTES=60 ``` 如生产 AI 配置中心不可用,还应提供当前项目支持的 `OPENAI_API_KEY`、`OPENAI_BASE_URL`、`CHAT_MODEL` 等兜底配置。`/data` 必须挂载持久卷,数据库与知识库文件不可存放在容器临时层。 @@ -109,7 +119,9 @@ location /api/ { } ``` -`proxy_buffering off` 用于数字分身 SSE 流式吐字,`client_max_body_size` 用于知识库文件上传。网关和应用日志必须关闭完整 URL 查询参数记录,任何异常日志都不得输出 token、Authorization 或平台密钥。建议同时设置严格的 `Referrer-Policy: no-referrer`。 +`proxy_buffering off` 用于数字分身 SSE 流式吐字,`client_max_body_size` 同时用于知识库文件和聊天图片上传。应用只保存图片识别结果,不保存原图;识别结果 24 小时失效,后台默认每小时清理一次。公开分享图片识别会消耗分身所有者积分,生产网关应针对 `/api/public/avatar/*/chat/images` 设置每 IP 和每分享令牌的上传频率限制,防止恶意消耗。 + +网关和应用日志必须关闭完整 URL 查询参数记录,任何异常日志都不得输出 token、Authorization、图片 Base64、病例正文或平台密钥。建议同时设置严格的 `Referrer-Policy: no-referrer`。 ## 5. 发布验收 @@ -123,6 +135,9 @@ location /api/ { 8. 重建容器后数据库、头像、知识库文档仍存在,`/api/health` 返回成功。 9. `https://digital.99hui.com/api/health` 可访问,证书域名和有效期正确,HTTP 自动跳转 HTTPS。 10. 微信和支付宝各创建一笔最小套餐订单,未付款时积分不变;支付成功后回调到账一次,重复回调积分不重复增加。 +11. 私聊和公开分享各上传 JPG、PNG、WebP 图片并完成追问;上传非图片、超过 8MB 或跨分身附件时必须拒绝。 +12. 病例图片可以提取可见文字并标记待核对内容,医学影像不作确定诊断;视觉与 OCR 调用分别扣减积分。 +13. 检查服务器上传目录不残留聊天原图,数据库过期图片识别记录在清理周期后删除,日志不出现 Base64 或病例正文。 ## 6. 回滚 diff --git a/digital-avatar-app/docs/IMAGE_MEDICAL_UNDERSTANDING_DESIGN.md b/digital-avatar-app/docs/IMAGE_MEDICAL_UNDERSTANDING_DESIGN.md new file mode 100644 index 0000000..f264152 --- /dev/null +++ b/digital-avatar-app/docs/IMAGE_MEDICAL_UNDERSTANDING_DESIGN.md @@ -0,0 +1,184 @@ +# 数字分身图片与病例理解详细设计 + +## 1. 目标与边界 + +本功能让数字分身在私聊和公开分享聊天中接收图片,并围绕图片内容继续使用现有的“标准答题对 -> 分身独立知识库 -> Qwen 兼容模型”链路回答。 + +第一期支持 JPEG、PNG、WebP,覆盖以下场景: + +1. 普通照片、截图、图表和界面图片的内容理解。 +2. 病例、处方、检查单、检验报告等图片文档的文字和表格提取。 +3. X 光、CT、MRI 等医学影像的客观可见内容描述。 + +第一期不把通用视觉模型的输出当作医学诊断,不自动把图片或病例写入知识库,不保存原图供长期访问,也不支持 DICOM 原始影像。 + +## 2. 核心原则 + +- **资料优先级不变**:标准答题对最高,分身独立知识库其次,图片识别结果属于待核对的会话资料,最后才由模型组织表达。 +- **病例最小留存**:应用不把原图写入业务存储,上传内容在内存中归一化并调用视觉服务;数据库只保存结构化结果和必要元数据。 +- **严格隔离**:每条图片记录必须绑定 `avatar_id`,私聊校验分身所有者,公开聊天校验分享令牌对应的分身。 +- **不确定性显式化**:OCR 看不清、表格列错位、医学影像无法确认时必须指出待核对项,不允许补齐缺失内容。 +- **可计量**:视觉理解和病例 OCR 分别计入分身所有者的积分消耗,失败时释放预留积分。 +- **可降级**:OCR 失败但通用视觉结果有效时仍可回答;视觉主调用失败则不进入聊天发送。 + +## 3. 总体流程 + +```text +用户选择图片 + -> 前端本地预览 + -> 私聊/公开图片上传接口 + -> 文件大小、MIME、真实格式、像素数校验 + -> 自动旋转、缩放、去 EXIF、统一 JPEG + -> 通用视觉模型分类并输出结构化 JSON + -> 若为病例/检查单,再调用 OCR 模型精确转录 + -> 保存结构化结果,不持久化原图 + -> 返回 attachmentId + -> 用户发送文字 + attachmentIds + -> 标准答题对匹配 + -> 用文字 + 图片提取结果检索独立知识库 + -> 把标准答案、知识片段、图片资料注入系统上下文 + -> Qwen SSE 流式回答 +``` + +## 4. 模型编排 + +### 4.1 通用视觉模型 + +默认 `qwen3.6-flash`,可在后台数字分身专用模型配置中修改。输入为归一化后的 Base64 Data URL,要求返回 JSON: + +```json +{ + "category": "general_image|document|medical_document|medical_image", + "summary": "客观、完整的图片描述", + "visible_text": "图片中可确认的文字", + "key_facts": ["事实1", "事实2"], + "uncertainties": ["无法确认的内容"], + "medical": { + "document_type": "", + "patient_info": {}, + "chief_complaint": "", + "findings": [], + "measurements": [], + "doctor_advice": "" + } +} +``` + +模型提示词禁止诊断、补全被遮挡文字、猜测患者身份和输出模型信息。 + +### 4.2 病例 OCR + +当 `category=medical_document` 时追加调用 `qwen-vl-ocr`,按原布局转录文字和表格。OCR 文本优先替换通用视觉输出中的 `visible_text`,但保留通用视觉模型提供的分类、摘要和不确定项。 + +### 4.3 医学影像 + +当 `category=medical_image` 时只保存客观描述,不输出疾病结论、分期、用药或治疗方案。聊天提示词必须要求结合正规影像报告和医生意见,并显示“图片识别结果仅供辅助,不能替代医生诊断”。 + +## 5. 数据模型 + +新增 `chat_attachments`: + +| 字段 | 说明 | +|---|---| +| `id` | 不可猜测的附件 ID | +| `avatar_id` | 所属数字分身,强制隔离 | +| `filename` | 原文件名,去除路径 | +| `mime_type` / `file_size` | 上传元数据 | +| `status` | `processing / ready / failed` | +| `category` | 图片分类 | +| `summary` | 通用视觉摘要 | +| `extracted_text` | 可确认文字/OCR 结果 | +| `structured_data` | 结构化 JSON | +| `warning` | 不确定项和医学提示 | +| `vision_model` / `ocr_model` | 实际调用模型 | +| `created_at` / `used_at` | 创建和最近使用时间 | + +不保存公开原图 URL。应用层不落盘原图;框架上传缓冲在请求结束时关闭,处理结果在 24 小时后自动清理。 + +## 6. API 设计 + +### 6.1 上传并解析 + +- `POST /api/avatar/{avatar_id}/chat/images` +- `POST /api/public/avatar/{share_token}/chat/images` +- `multipart/form-data: file` + +成功返回: + +```json +{ + "id": "attachment-id", + "filename": "病例.jpg", + "status": "ready", + "category": "medical_document", + "summary": "门诊检查单", + "warning": "部分手写内容需要人工核对" +} +``` + +### 6.2 聊天 + +原聊天接口增加: + +```json +{ + "message": "请帮我看看异常指标", + "attachmentIds": ["attachment-id"], + "history": [ + {"role": "user", "content": "上一条问题", "attachmentIds": ["attachment-id"]} + ] +} +``` + +当前消息最多 3 张图,历史最多引用最近 3 个不同附件。后端只读取与当前 `avatar_id` 相同且状态为 `ready` 的记录。 + +## 7. 安全与隐私 + +- 单图最大 8MB,解码后最大 1600 万像素,最长边归一化到 4096 像素以内。 +- 使用 Pillow 验证真实图片格式并防止解压炸弹;重新编码时清除 EXIF、GPS 和其他元数据。 +- 图片不会写入 FastAPI `StaticFiles` 或知识库目录,模型请求和日志不得输出 Base64 内容。 +- 日志只记录附件 ID、分身 ID、状态、耗时和模型,不记录图片 Base64、OCR 全文、病例内容或 API Key。 +- 公开分享上传仍消耗分身所有者积分;余额不足时拒绝视觉调用。 +- 生产环境需要补充用户授权、数据处理协议、存储地域和模型供应商留存策略确认。 + +## 8. 前端交互 + +- 输入框左侧增加图片按钮,支持相册选择和移动端拍照。 +- 选择后显示本地缩略图和“正在识别图片”,识别完成前禁止发送。 +- 用户可删除待发送图片;发送后图片保留在当前会话气泡中,但刷新页面后不恢复原图。 +- 病例和医学影像在输入区及回答下方显示辅助提示,不使用恐吓式红色告警。 +- 上传或识别失败时保留文字输入,明确提示重新选择图片,不产生空白消息。 + +## 9. 配置 + +数字分身专用模型配置新增: + +- `vision_model_version`,默认 `qwen3.6-flash` +- `ocr_model_version`,默认 `qwen-vl-ocr` + +环境变量兜底: + +```dotenv +VISION_MODEL=qwen3.6-flash +VISION_OCR_MODEL=qwen-vl-ocr +VISION_MAX_OUTPUT_TOKENS=2048 +VISION_TIMEOUT_SECONDS=90 +VISION_TOKEN_RESERVE=12000 +CHAT_IMAGE_MAX_BYTES=8388608 +CHAT_IMAGE_MAX_PIXELS=16000000 +CHAT_ATTACHMENT_RETENTION_HOURS=24 +CHAT_ATTACHMENT_CLEANUP_MINUTES=60 +``` + +视觉调用复用数字分身专用配置的 `api_base_url` 和 `api_key`,不额外复制密钥。 + +## 10. 验收标准 + +1. 普通照片、截图和图表能够返回与图片一致的描述并支持追问。 +2. 病例图片可以提取标题、患者字段、检查结果、异常指标和医生意见,模糊内容明确标记待核对。 +3. 上传后服务器业务目录不残留原图,响应和日志不包含 Base64 或完整病例正文。 +4. A 分身无法引用 B 分身附件;公开分享令牌无法访问其他分身附件。 +5. 有图片时标准答题对仍作为最高优先级事实,知识库命中次之。 +6. 视觉与 OCR 积分分别结算,失败调用释放预留积分。 +7. SSE 打字效果、Markdown、用户头像、公开分享和纯文本聊天均无回归。 +8. CT、MRI、X 光回答不作确定诊断,并显示人工复核提示。 diff --git a/digital-avatar-app/src/api/index.ts b/digital-avatar-app/src/api/index.ts index 8e1b325..be7a4a8 100644 --- a/digital-avatar-app/src/api/index.ts +++ b/digital-avatar-app/src/api/index.ts @@ -374,15 +374,35 @@ export const searchKnowledge = (avatarId: string, q: string, topK = 5) => export interface ChatMessage { role: 'user' | 'assistant' content: string + attachmentIds?: string[] } export interface ChatResponse { answer: string - source: 'qa' | 'knowledge' | 'qwen' + source: 'qa' | 'knowledge' | 'vision' | 'qwen' references?: Array<{ docId?: string; filename?: string; fileType?: string; snippet?: string; score?: number }> } -export const sendAvatarChat = (avatarId: string, payload: { message: string; history?: ChatMessage[] }) => +export interface ChatAttachment { + id: string + avatarId: string + filename: string + mimeType: string + fileSize: number + status: 'processing' | 'ready' | 'failed' + category: 'general_image' | 'document' | 'medical_document' | 'medical_image' + summary: string + warning: string + expiresAt: string +} + +export interface ChatPayload { + message: string + attachmentIds?: string[] + history?: ChatMessage[] +} + +export const sendAvatarChat = (avatarId: string, payload: ChatPayload) => request.post(`/avatar/${avatarId}/chat`, payload) export interface PublicAvatar { @@ -401,19 +421,41 @@ export const createAvatarShareLink = (avatarId: string) => export const getPublicAvatar = (shareToken: string) => request.get(`/public/avatar/${shareToken}`) -export const sendPublicAvatarChat = (shareToken: string, payload: { message: string; history?: ChatMessage[] }) => +export const sendPublicAvatarChat = (shareToken: string, payload: ChatPayload) => request.post(`/public/avatar/${shareToken}/chat`, payload) +const imageForm = (file: File) => { + const form = new FormData() + form.append('file', file) + return form +} + +export const uploadAvatarChatImage = (avatarId: string, file: File) => + request.post(`/avatar/${avatarId}/chat/images`, imageForm(file), { + headers: { 'Content-Type': 'multipart/form-data' }, + timeout: 120000 + }) + +export const uploadPublicAvatarChatImage = (shareToken: string, file: File) => + request.post(`/public/avatar/${shareToken}/chat/images`, imageForm(file), { + headers: { 'Content-Type': 'multipart/form-data' }, + timeout: 120000 + }) + type ChatStreamHandlers = { onMeta: (meta: Pick) => void onDelta: (content: string) => void } -const streamChat = async (path: string, payload: { message: string; history?: ChatMessage[] }, handlers: ChatStreamHandlers) => { +const streamChat = async (path: string, payload: ChatPayload, handlers: ChatStreamHandlers) => { const headers: Record = { 'Content-Type': 'application/json', Accept: 'text/event-stream' } if (_authToken) headers.Authorization = `Bearer ${_authToken}` const response = await fetch(`${resolveBaseURL()}${path}`, { method: 'POST', headers, body: JSON.stringify(payload) }) - if (!response.ok || !response.body) throw new Error(`对话请求失败(${response.status})`) + if (!response.ok) { + const errorBody = await response.json().catch(() => null) + throw new Error(errorBody?.detail || errorBody?.message || `对话请求失败(${response.status})`) + } + if (!response.body) throw new Error('对话响应为空,请稍后重试') const reader = response.body.getReader() const decoder = new TextDecoder() @@ -436,10 +478,10 @@ const streamChat = async (path: string, payload: { message: string; history?: Ch } } -export const streamAvatarChat = (avatarId: string, payload: { message: string; history?: ChatMessage[] }, handlers: ChatStreamHandlers) => +export const streamAvatarChat = (avatarId: string, payload: ChatPayload, handlers: ChatStreamHandlers) => streamChat(`/avatar/${avatarId}/chat/stream`, payload, handlers) -export const streamPublicAvatarChat = (shareToken: string, payload: { message: string; history?: ChatMessage[] }, handlers: ChatStreamHandlers) => +export const streamPublicAvatarChat = (shareToken: string, payload: ChatPayload, handlers: ChatStreamHandlers) => streamChat(`/public/avatar/${shareToken}/chat/stream`, payload, handlers) // ==================== 会会用户资料 API ==================== diff --git a/digital-avatar-app/src/views/AvatarChat.vue b/digital-avatar-app/src/views/AvatarChat.vue index 2ac4838..20c996a 100644 --- a/digital-avatar-app/src/views/AvatarChat.vue +++ b/digital-avatar-app/src/views/AvatarChat.vue @@ -1,5 +1,5 @@