feat(avatar): add private vision chat support
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user