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