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
+56
View File
@@ -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()
+38
View File
@@ -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)
@@ -8,3 +8,4 @@ pypdf
python-docx
openpyxl
apscheduler>=3.10
Pillow>=10.4
+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
@@ -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
@@ -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",
)
@@ -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)
@@ -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()]
@@ -5,6 +5,7 @@ from database import init_db, SessionLocal
from models import (
Authorization,
Avatar,
ChatAttachment,
TakeoverCursor,
TakeoverMessage,
TakeoverReplyTask,
@@ -95,6 +96,9 @@ def authorization_context():
finally:
db.rollback()
avatar_ids = [avatar.id, other_avatar.id]
db.query(ChatAttachment).filter(
ChatAttachment.avatar_id.in_(avatar_ids)
).delete(synchronize_session=False)
db.query(TakeoverReplyTask).filter(
TakeoverReplyTask.avatar_id.in_(avatar_ids)
).delete(synchronize_session=False)
@@ -0,0 +1,273 @@
import json
from datetime import datetime, timedelta
from types import SimpleNamespace
from unittest.mock import Mock, patch
import pytest
from fastapi import HTTPException
from fastapi.testclient import TestClient
from database import SessionLocal
from main import app
from models import ChatAttachment
from routers.chat import (
ChatIn,
_attachment_contexts,
_load_chat_attachments,
_resolve_reply,
)
from services.chat_attachment_service import purge_expired_chat_attachments
from services.vision_service import PreparedImage
client = TestClient(app)
GENERAL_RESULT = {
"content": json.dumps({
"category": "general_image",
"summary": "一张包含产品路线图的截图",
"visible_text": "产品路线图",
"key_facts": ["包含三个阶段"],
"uncertainties": [],
"medical": {},
}, ensure_ascii=False),
"usage": {"total_tokens": 120},
}
def test_owner_can_upload_and_cache_image_analysis(authorization_context):
context = authorization_context
prepared = PreparedImage(b"jpeg", "image/jpeg", 100, 80)
with (
patch("routers.chat.prepare_image", return_value=prepared),
patch("routers.chat._run_billed_vision_call", return_value=GENERAL_RESULT),
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/chat/images",
headers=context["owner_headers"],
files={"file": ("roadmap.png", b"image-bytes", "image/png")},
)
assert response.status_code == 200
payload = response.json()["data"]
assert payload["status"] == "ready"
assert payload["category"] == "general_image"
assert payload["summary"] == "一张包含产品路线图的截图"
db = SessionLocal()
try:
stored = db.query(ChatAttachment).filter(ChatAttachment.id == payload["id"]).one()
assert stored.avatar_id == context["avatar"].id
assert stored.extracted_text == "产品路线图"
assert stored.structured_data["key_facts"] == ["包含三个阶段"]
finally:
db.close()
def test_non_owner_cannot_upload_chat_image(authorization_context):
context = authorization_context
response = client.post(
f"/api/avatar/{context['avatar'].id}/chat/images",
headers=context["other_headers"],
files={"file": ("private.png", b"image-bytes", "image/png")},
)
assert response.status_code == 403
def test_public_share_can_upload_without_exposing_analysis_details(authorization_context):
context = authorization_context
db = SessionLocal()
try:
avatar = db.get(type(context["avatar"]), context["avatar"].id)
avatar.share_token = f"share-{context['suffix']}"
db.commit()
share_token = avatar.share_token
finally:
db.close()
with (
patch(
"routers.chat.prepare_image",
return_value=PreparedImage(b"jpeg", "image/jpeg", 100, 80),
),
patch("routers.chat._run_billed_vision_call", return_value=GENERAL_RESULT),
):
response = client.post(
f"/api/public/avatar/{share_token}/chat/images",
files={"file": ("visitor.png", b"image-bytes", "image/png")},
)
assert response.status_code == 200
payload = response.json()["data"]
assert payload["status"] == "ready"
assert "structuredData" not in payload
assert "extractedText" not in payload
assert "visionModel" not in payload
assert "ocrModel" not in payload
db = SessionLocal()
try:
stored = db.get(ChatAttachment, payload["id"])
assert stored.uploader_kind == "public"
assert stored.avatar_id == context["avatar"].id
finally:
db.close()
def test_medical_document_uses_ocr_result(authorization_context):
context = authorization_context
general = {
"content": json.dumps({
"category": "medical_document",
"summary": "血常规报告",
"visible_text": "初步文字",
"key_facts": [],
"uncertainties": [],
"medical": {"document_type": "检验报告"},
}, ensure_ascii=False),
"usage": {},
}
ocr = {"content": "白细胞 11.2 x10^9/L", "usage": {}}
with (
patch("routers.chat.prepare_image", return_value=PreparedImage(b"jpeg", "image/jpeg", 100, 80)),
patch("routers.chat._run_billed_vision_call", side_effect=[general, ocr]) as model,
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/chat/images",
headers=context["owner_headers"],
files={"file": ("report.jpg", b"image-bytes", "image/jpeg")},
)
assert response.status_code == 200
attachment_id = response.json()["data"]["id"]
assert model.call_count == 2
assert model.call_args_list[1].kwargs["source"] == "vision_medical_ocr"
db = SessionLocal()
try:
stored = db.query(ChatAttachment).filter(ChatAttachment.id == attachment_id).one()
assert stored.extracted_text == "白细胞 11.2 x10^9/L"
assert stored.ocr_model == "qwen-vl-ocr"
assert "不能替代医生诊断" in stored.warning
finally:
db.close()
def test_attachment_cannot_cross_avatar_boundary(authorization_context):
context = authorization_context
db = SessionLocal()
try:
attachment = ChatAttachment(
avatar_id=context["avatar"].id,
filename="private.jpg",
status="ready",
expires_at=datetime.utcnow() + timedelta(hours=1),
)
db.add(attachment)
db.commit()
body = ChatIn(message="看看图片", attachmentIds=[attachment.id])
with pytest.raises(HTTPException, match="不属于当前分身") as caught:
_load_chat_attachments(db, context["other_avatar"].id, body)
assert caught.value.status_code == 400
finally:
db.close()
def test_expired_attachment_is_removed(authorization_context):
context = authorization_context
db = SessionLocal()
try:
attachment = ChatAttachment(
avatar_id=context["avatar"].id,
filename="expired.jpg",
status="ready",
expires_at=datetime.utcnow() - timedelta(seconds=1),
)
db.add(attachment)
db.commit()
attachment_id = attachment.id
body = ChatIn(message="看看图片", attachmentIds=[attachment_id])
with pytest.raises(HTTPException):
_load_chat_attachments(db, context["avatar"].id, body)
assert db.query(ChatAttachment).filter(ChatAttachment.id == attachment_id).first() is None
finally:
db.close()
def test_cleanup_keeps_unexpired_attachment(authorization_context):
context = authorization_context
now = datetime.utcnow()
db = SessionLocal()
try:
expired = ChatAttachment(
avatar_id=context["avatar"].id,
filename="expired.jpg",
status="ready",
expires_at=now - timedelta(seconds=1),
)
active = ChatAttachment(
avatar_id=context["avatar"].id,
filename="active.jpg",
status="ready",
expires_at=now + timedelta(hours=1),
)
db.add_all([expired, active])
db.commit()
expired_id, active_id = expired.id, active.id
assert purge_expired_chat_attachments(db, now=now) == 1
assert db.get(ChatAttachment, expired_id) is None
assert db.get(ChatAttachment, active_id) is not None
finally:
db.close()
def test_image_context_keeps_standard_answer_authoritative():
avatar = SimpleNamespace(
id="avatar-vision",
name="测试分身",
description="产品顾问",
config={},
)
model = Mock(return_value="标准退款期限是七天;图片显示的是商品包装。")
result = _resolve_reply(
None,
avatar,
"退款期限是多少?",
[],
qa_pairs=[SimpleNamespace(question="退款期限是多少?", answer="七天", enabled=True)],
search_fn=Mock(return_value=[]),
model_client=model,
image_contexts=[{
"id": "attachment",
"filename": "product.jpg",
"category": "general_image",
"summary": "商品包装",
"extractedText": "",
"structuredData": {},
"warning": "",
}],
)
assert result["source"] == "qa"
system = model.call_args.kwargs["messages"][0]["content"]
assert "已确认标准答案" in system
assert "七天" in system
assert "商品包装" in system
assert "标准答题对中的事实优先级高于图片资料" in system
def test_attachment_context_does_not_expose_internal_fields():
row = SimpleNamespace(
id="attachment",
filename="case.jpg",
category="medical_document",
summary="门诊病例",
extracted_text="主诉:咳嗽",
structured_data={"medical": {"chief_complaint": "咳嗽"}},
warning="请核对原文",
)
context = _attachment_contexts([row])[0]
assert context["filename"] == "case.jpg"
assert "avatar_id" not in context
assert "vision_model" not in context
@@ -26,6 +26,8 @@ def test_admin_runtime_config_takes_priority(monkeypatch):
"api_base_url": "https://model.test/v1/",
"api_key": "runtime-key",
"model": "avatar-model",
"vision_model": "avatar-vision-model",
"ocr_model": "avatar-ocr-model",
"max_tokens": 2048,
"timeout_seconds": 42,
}
@@ -37,6 +39,8 @@ def test_admin_runtime_config_takes_priority(monkeypatch):
assert config.source == "admin"
assert config.api_base_url == "https://model.test/v1"
assert config.model == "avatar-model"
assert config.vision_model == "avatar-vision-model"
assert config.ocr_model == "avatar-ocr-model"
assert config.max_tokens == 2048
request.assert_called_once_with(
"http://config.test/runtime",
@@ -51,6 +55,8 @@ def test_runtime_failure_falls_back_to_environment(monkeypatch):
monkeypatch.setenv("CHAT_API_URL", "https://fallback.test/v1/")
monkeypatch.setenv("CHAT_API_KEY", "fallback-key")
monkeypatch.setenv("CHAT_MODEL", "fallback-model")
monkeypatch.setenv("VISION_MODEL", "fallback-vision")
monkeypatch.setenv("VISION_OCR_MODEL", "fallback-ocr")
monkeypatch.setenv("CHAT_MAX_OUTPUT_TOKENS", "1536")
request = httpx.Request("GET", "http://config.test/runtime")
@@ -64,6 +70,8 @@ def test_runtime_failure_falls_back_to_environment(monkeypatch):
assert config.api_base_url == "https://fallback.test/v1"
assert config.api_key == "fallback-key"
assert config.model == "fallback-model"
assert config.vision_model == "fallback-vision"
assert config.ocr_model == "fallback-ocr"
assert config.max_tokens == 1536
@@ -20,8 +20,9 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
):
import main
maintenance_scheduler = MagicMock()
scheduler = MagicMock()
mock_scheduler_class.return_value = scheduler
mock_scheduler_class.side_effect = [maintenance_scheduler, scheduler]
boxim = MagicMock()
mock_boxim_class.return_value = boxim
takeover = MagicMock()
@@ -47,6 +48,10 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
assert config["BOXIM_API_BASE_URL"] == "https://im.example/api"
mock_takeover_class.assert_called_once_with(main.SessionLocal, boxim)
maintenance_scheduler.add_job.assert_called_once()
assert maintenance_scheduler.add_job.call_args.kwargs["id"] == "chat_attachment_cleanup"
maintenance_scheduler.start.assert_called_once_with()
assert scheduler.add_job.call_count == 2
poll_call, process_call = scheduler.add_job.call_args_list
assert poll_call.args[0] is takeover.poll_messages
@@ -62,6 +67,7 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
scheduler.start.assert_called_once_with()
main.takeover_scheduler = None
main.maintenance_scheduler = None
@patch("main.AsyncIOScheduler")
@@ -73,6 +79,7 @@ def test_scheduler_failure_does_not_stop_the_api(mock_scheduler_class):
main.on_startup()
assert main.takeover_scheduler is None
assert main.maintenance_scheduler is None
def test_shutdown_stops_only_the_scheduler():
@@ -80,9 +87,14 @@ def test_shutdown_stops_only_the_scheduler():
scheduler = MagicMock()
scheduler.running = True
maintenance_scheduler = MagicMock()
maintenance_scheduler.running = True
main.takeover_scheduler = scheduler
main.maintenance_scheduler = maintenance_scheduler
main.on_shutdown()
scheduler.shutdown.assert_called_once_with(wait=False)
maintenance_scheduler.shutdown.assert_called_once_with(wait=False)
assert main.takeover_scheduler is None
assert main.maintenance_scheduler is None
@@ -0,0 +1,98 @@
import io
import json
from unittest.mock import Mock, patch
import pytest
from PIL import Image
from services.chat_model_config import ChatModelConfig
from services.vision_service import (
ImageValidationError,
build_attachment_warning,
call_vision_model,
parse_vision_analysis,
prepare_image,
)
def _image_bytes(fmt="PNG", size=(120, 80)):
output = io.BytesIO()
Image.new("RGB", size, "#f97316").save(output, format=fmt)
return output.getvalue()
def _config():
return ChatModelConfig(
api_base_url="https://model.test/v1",
api_key="secret-key",
model="chat-model",
max_tokens=1024,
timeout_seconds=30,
vision_model="vision-model",
ocr_model="ocr-model",
vision_max_tokens=2048,
vision_timeout_seconds=90,
source="test",
)
def test_prepare_image_validates_and_reencodes_without_metadata():
prepared = prepare_image(_image_bytes())
assert prepared.mime_type == "image/jpeg"
assert prepared.width == 120
assert prepared.height == 80
with Image.open(io.BytesIO(prepared.data)) as image:
assert image.format == "JPEG"
assert not image.getexif()
def test_prepare_image_rejects_non_image_content():
with pytest.raises(ImageValidationError, match="格式无效"):
prepare_image(b"not-an-image")
def test_vision_request_uses_openai_compatible_image_content():
response = Mock()
response.raise_for_status.return_value = None
response.json.return_value = {
"choices": [{"message": {"content": '{"category":"general_image"}'}}],
"usage": {"total_tokens": 88},
}
prepared = prepare_image(_image_bytes())
with patch("services.vision_service.httpx.post", return_value=response) as request:
result = call_vision_model(
prepared,
_config(),
model="vision-model",
prompt="describe",
json_output=True,
)
payload = request.call_args.kwargs["json"]
content = payload["messages"][0]["content"]
assert payload["model"] == "vision-model"
assert payload["response_format"] == {"type": "json_object"}
assert content[0]["type"] == "image_url"
assert content[0]["image_url"]["url"].startswith("data:image/jpeg;base64,")
assert content[1] == {"type": "text", "text": "describe"}
assert result["usage"]["total_tokens"] == 88
def test_parse_medical_analysis_and_build_warning():
analysis = parse_vision_analysis(json.dumps({
"category": "medical_document",
"summary": "血常规报告",
"visible_text": "白细胞 11.2",
"key_facts": ["白细胞偏高"],
"uncertainties": ["日期模糊"],
"medical": {"document_type": "检验报告"},
}, ensure_ascii=False))
assert analysis["category"] == "medical_document"
assert analysis["medical"]["document_type"] == "检验报告"
warning = build_attachment_warning(analysis, ocr_failed=True)
assert "日期模糊" in warning
assert "人工核对" in warning
assert "不能替代医生诊断" in warning