feat(avatar): add private vision chat support

This commit is contained in:
stefanfeng
2026-08-31 15:10:50 +08:00
parent 094f8cd40f
commit 016bc22c05
22 changed files with 1616 additions and 59 deletions
+398 -16
View File
@@ -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