197 lines
7.7 KiB
Python
197 lines
7.7 KiB
Python
"""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()]
|