第一次提交
This commit is contained in:
@@ -0,0 +1,24 @@
|
||||
"""用户 AI api_key 的对称加密(Fernet)。密钥来自 settings.fernet_key。"""
|
||||
from cryptography.fernet import Fernet, InvalidToken
|
||||
|
||||
from app.config import settings
|
||||
|
||||
|
||||
def _fernet() -> Fernet:
|
||||
if not settings.fernet_key:
|
||||
raise RuntimeError(
|
||||
"FERNET_KEY 未配置,无法加解密用户 api_key。"
|
||||
"生成:python -c \"from cryptography.fernet import Fernet; print(Fernet.generate_key().decode())\""
|
||||
)
|
||||
return Fernet(settings.fernet_key.encode("utf-8"))
|
||||
|
||||
|
||||
def encrypt(plaintext: str) -> str:
|
||||
return _fernet().encrypt(plaintext.encode("utf-8")).decode("utf-8")
|
||||
|
||||
|
||||
def decrypt(token: str) -> str | None:
|
||||
try:
|
||||
return _fernet().decrypt(token.encode("utf-8")).decode("utf-8")
|
||||
except (InvalidToken, ValueError):
|
||||
return None
|
||||
@@ -0,0 +1,166 @@
|
||||
"""上传图片 → OCR → 结构化提取 → 生成草稿题目。作为 BackgroundTask 运行。
|
||||
|
||||
注意:BackgroundTask 在请求返回后执行,需要自己开新的 DB session。
|
||||
"""
|
||||
import json
|
||||
import logging
|
||||
|
||||
from pydantic import ValidationError
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.db import SessionLocal
|
||||
from app.models.image import Image
|
||||
from app.models.job import JobStatus, ProcessingJob
|
||||
from app.models.question import Question, QuestionType
|
||||
from app.models.user import UserSettings
|
||||
from app.schemas.question import QuestionCreate
|
||||
from app.services.llm.client import LLMClient, LLMConfigError
|
||||
from app.services.llm.prompts import EXTRACT_SYSTEM, EXTRACT_USER_TEMPLATE
|
||||
from app.services.ocr.pipeline import run_ocr
|
||||
from app.services.storage import download_image_async
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _parse_extract_json(raw: str) -> dict:
|
||||
"""从模型输出里解析 JSON(容忍代码块包裹)。"""
|
||||
text = raw.strip()
|
||||
if text.startswith("```"):
|
||||
text = text.split("```", 2)[1]
|
||||
if text.startswith("json"):
|
||||
text = text[4:]
|
||||
return json.loads(text.strip())
|
||||
|
||||
|
||||
async def run_extract_job(job_id: str, user_id: str, object_key: str, mime: str):
|
||||
async with SessionLocal() as db:
|
||||
job = await db.get(ProcessingJob, job_id)
|
||||
if job is None:
|
||||
logger.error("job %s 不存在", job_id)
|
||||
return
|
||||
try:
|
||||
# 加载用户 AI 配置
|
||||
settings_row = await db.scalar(
|
||||
select(UserSettings).where(UserSettings.user_id == user_id)
|
||||
)
|
||||
llm = LLMClient(settings_row)
|
||||
|
||||
# 1. OCR
|
||||
job.status = JobStatus.ocr_running
|
||||
await db.commit()
|
||||
image_bytes = await download_image_async(object_key)
|
||||
ocr = await run_ocr(image_bytes, mime, llm)
|
||||
|
||||
# 2. 结构化提取
|
||||
job.status = JobStatus.ai_running
|
||||
job.engine_used = ocr.engine
|
||||
await db.commit()
|
||||
extract = await llm.chat_text(
|
||||
EXTRACT_SYSTEM,
|
||||
EXTRACT_USER_TEMPLATE.format(content=ocr.markdown),
|
||||
json_mode=True,
|
||||
)
|
||||
questions = _build_questions(user_id, extract.content, ocr, job.image_id)
|
||||
|
||||
# 缓存原始 OCR 文本,便于日后追溯或重新拆分
|
||||
if job.image_id:
|
||||
image = await db.get(Image, job.image_id)
|
||||
if image is not None:
|
||||
image.ocr_markdown = ocr.markdown
|
||||
|
||||
db.add_all(questions)
|
||||
await db.flush()
|
||||
job.result_question_ids = [q.id for q in questions]
|
||||
job.status = JobStatus.done
|
||||
await db.commit()
|
||||
logger.info(
|
||||
"job %s 完成,engine=%s,拆出 %d 条笔记",
|
||||
job_id,
|
||||
ocr.engine,
|
||||
len(questions),
|
||||
)
|
||||
except LLMConfigError as e:
|
||||
await _fail(db, job, f"AI 未配置:{e}")
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.exception("job %s 失败", job_id)
|
||||
await _fail(db, job, str(e))
|
||||
|
||||
|
||||
def _build_questions(
|
||||
user_id: str, extract_raw: str, ocr, image_id: str | None
|
||||
) -> list[Question]:
|
||||
"""把提取结果拆成多条笔记。
|
||||
|
||||
一张照片常含多道题,每道题存成独立笔记。整体解析失败时不丢内容,
|
||||
降级为一条 unclassified 笔记(原文进题干),用户可自行整理。
|
||||
"""
|
||||
try:
|
||||
data = _parse_extract_json(extract_raw)
|
||||
raw_items = data.get("questions")
|
||||
if not isinstance(raw_items, list) or not raw_items:
|
||||
raise ValueError("提取结果里没有 questions 数组")
|
||||
except (json.JSONDecodeError, KeyError, ValueError) as e:
|
||||
logger.warning("提取结果解析失败,降级为一条未分类笔记:%s", e)
|
||||
return [_fallback_question(user_id, ocr, image_id)]
|
||||
|
||||
questions: list[Question] = []
|
||||
for idx, item in enumerate(raw_items):
|
||||
try:
|
||||
questions.append(_build_one(user_id, item, ocr, image_id))
|
||||
except (KeyError, ValueError, ValidationError, TypeError) as e:
|
||||
# 单条坏掉不影响其他条:这条退化为未分类,原文留在题干里
|
||||
logger.warning("第 %d 条解析失败,存为未分类:%s", idx + 1, e)
|
||||
stem = ""
|
||||
if isinstance(item, dict):
|
||||
stem = str(item.get("stem_markdown") or "")
|
||||
questions.append(
|
||||
_fallback_question(user_id, ocr, image_id, stem_override=stem or None)
|
||||
)
|
||||
|
||||
return questions or [_fallback_question(user_id, ocr, image_id)]
|
||||
|
||||
|
||||
def _build_one(user_id: str, item: dict, ocr, image_id: str | None) -> Question:
|
||||
parsed = QuestionCreate(
|
||||
# 模型没给或给了不认识的题型,就当未分类,别硬猜
|
||||
type=_coerce_type(item.get("type")),
|
||||
stem_markdown=item.get("stem_markdown") or ocr.markdown,
|
||||
options=item.get("options"),
|
||||
correct_answer=item.get("correct_answer"),
|
||||
difficulty=item.get("difficulty"),
|
||||
)
|
||||
return Question(
|
||||
user_id=user_id,
|
||||
type=parsed.type,
|
||||
stem_markdown=parsed.stem_markdown,
|
||||
options=[o.model_dump() for o in parsed.options] if parsed.options else None,
|
||||
correct_answer=parsed.correct_answer,
|
||||
difficulty=parsed.difficulty,
|
||||
source_image_id=image_id,
|
||||
ocr_engine=ocr.engine,
|
||||
)
|
||||
|
||||
|
||||
def _coerce_type(raw) -> QuestionType:
|
||||
try:
|
||||
return QuestionType(raw)
|
||||
except ValueError:
|
||||
return QuestionType.unclassified
|
||||
|
||||
|
||||
def _fallback_question(
|
||||
user_id: str, ocr, image_id: str | None, stem_override: str | None = None
|
||||
) -> Question:
|
||||
return Question(
|
||||
user_id=user_id,
|
||||
type=QuestionType.unclassified,
|
||||
stem_markdown=stem_override or ocr.markdown,
|
||||
source_image_id=image_id,
|
||||
ocr_engine=ocr.engine,
|
||||
)
|
||||
|
||||
|
||||
async def _fail(db, job: ProcessingJob, msg: str):
|
||||
job.status = JobStatus.failed
|
||||
job.error_message = msg[:1000]
|
||||
await db.commit()
|
||||
@@ -0,0 +1,24 @@
|
||||
"""客观题本地判分:精确比对答案集合。"""
|
||||
from app.models.question import OBJECTIVE_TYPES, Question
|
||||
|
||||
|
||||
def is_objective(q: Question) -> bool:
|
||||
return q.type in OBJECTIVE_TYPES
|
||||
|
||||
|
||||
def can_grade_locally(q: Question) -> bool:
|
||||
"""客观题且录了标准答案,才能本地判。"""
|
||||
return is_objective(q) and bool(q.correct_answer)
|
||||
|
||||
|
||||
def grade_objective(q: Question, user_answer: list[str]) -> bool | None:
|
||||
"""选择/判断题:忽略顺序与大小写,比对答案集合。
|
||||
|
||||
没录标准答案时返回 None(无法判定),而不是 False —— 否则一条只是
|
||||
随手记下、还没填答案的笔记会被静默算错,并污染复习清单。
|
||||
"""
|
||||
correct = {str(a).strip().upper() for a in (q.correct_answer or [])}
|
||||
if not correct:
|
||||
return None
|
||||
given = {str(a).strip().upper() for a in (user_answer or [])}
|
||||
return correct == given
|
||||
@@ -0,0 +1,111 @@
|
||||
"""AI 任务:判分、讲解、知识点总结。统一记录到 ai_generations。"""
|
||||
import json
|
||||
import logging
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models.practice import AiGeneration
|
||||
from app.models.question import Question
|
||||
from app.models.user import UserSettings
|
||||
from app.services.llm.client import LLMClient
|
||||
from app.services.llm.prompts import (
|
||||
EXPLAIN_SYSTEM,
|
||||
EXPLAIN_USER_TEMPLATE,
|
||||
JUDGE_SYSTEM,
|
||||
JUDGE_USER_TEMPLATE,
|
||||
TAGS_SYSTEM,
|
||||
TAGS_USER_TEMPLATE,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def load_llm(db: AsyncSession, user_id: str) -> LLMClient:
|
||||
"""加载用户的 LLMClient(未配置会抛 LLMConfigError)。"""
|
||||
row = await db.scalar(select(UserSettings).where(UserSettings.user_id == user_id))
|
||||
return LLMClient(row)
|
||||
|
||||
|
||||
def _parse_json(raw: str) -> dict:
|
||||
text = raw.strip()
|
||||
if text.startswith("```"):
|
||||
text = text.split("```", 2)[1]
|
||||
if text.startswith("json"):
|
||||
text = text[4:]
|
||||
return json.loads(text.strip())
|
||||
|
||||
|
||||
async def _record(
|
||||
db: AsyncSession, user_id: str, question_id: str | None, task: str, result, content: str
|
||||
):
|
||||
db.add(
|
||||
AiGeneration(
|
||||
user_id=user_id,
|
||||
question_id=question_id,
|
||||
task=task,
|
||||
model=result.model,
|
||||
request_tokens=result.prompt_tokens,
|
||||
response_tokens=result.completion_tokens,
|
||||
content_markdown=content,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
async def judge_answer(
|
||||
db: AsyncSession, llm: LLMClient, user_id: str, q: Question, user_answer: list[str]
|
||||
) -> dict:
|
||||
"""主观题判分。返回 {is_correct, feedback_markdown, key_points_missed}。"""
|
||||
answer = "、".join(q.correct_answer) if q.correct_answer else "(未提供标准答案)"
|
||||
result = await llm.chat_text(
|
||||
JUDGE_SYSTEM,
|
||||
JUDGE_USER_TEMPLATE.format(
|
||||
stem=q.stem_markdown, answer=answer, user_answer="\n".join(user_answer)
|
||||
),
|
||||
json_mode=True,
|
||||
)
|
||||
await _record(db, user_id, q.id, "judge", result, result.content)
|
||||
try:
|
||||
data = _parse_json(result.content)
|
||||
return {
|
||||
"is_correct": bool(data.get("is_correct")),
|
||||
"feedback_markdown": data.get("feedback_markdown", ""),
|
||||
"key_points_missed": data.get("key_points_missed", []),
|
||||
}
|
||||
except (json.JSONDecodeError, KeyError) as e:
|
||||
logger.warning("judge 解析失败:%s", e)
|
||||
return {
|
||||
"is_correct": False,
|
||||
"feedback_markdown": result.content,
|
||||
"key_points_missed": [],
|
||||
}
|
||||
|
||||
|
||||
async def explain_question(
|
||||
db: AsyncSession, llm: LLMClient, user_id: str, q: Question
|
||||
) -> str:
|
||||
answer_hint = ""
|
||||
if q.correct_answer:
|
||||
answer_hint = f"参考答案:{'、'.join(q.correct_answer)}\n\n"
|
||||
result = await llm.chat_text(
|
||||
EXPLAIN_SYSTEM,
|
||||
EXPLAIN_USER_TEMPLATE.format(stem=q.stem_markdown, answer_hint=answer_hint),
|
||||
)
|
||||
await _record(db, user_id, q.id, "explain", result, result.content)
|
||||
return result.content
|
||||
|
||||
|
||||
async def summarize_tags(
|
||||
db: AsyncSession, llm: LLMClient, user_id: str, q: Question
|
||||
) -> list[str]:
|
||||
result = await llm.chat_text(
|
||||
TAGS_SYSTEM, TAGS_USER_TEMPLATE.format(stem=q.stem_markdown), json_mode=True
|
||||
)
|
||||
await _record(db, user_id, q.id, "summarize", result, result.content)
|
||||
try:
|
||||
data = _parse_json(result.content)
|
||||
tags = data.get("tags", [])
|
||||
return [str(t).strip() for t in tags if str(t).strip()][:4]
|
||||
except (json.JSONDecodeError, KeyError) as e:
|
||||
logger.warning("tags 解析失败:%s", e)
|
||||
return []
|
||||
@@ -0,0 +1,103 @@
|
||||
"""统一的 OpenAI 兼容 LLM client:文本 + 多模态。
|
||||
|
||||
配置来自用户的 UserSettings(base_url / api_key / 模型名)。api_key 在库里
|
||||
是 Fernet 加密的,这里解密后使用。
|
||||
"""
|
||||
import base64
|
||||
from dataclasses import dataclass
|
||||
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
from app.models.user import UserSettings
|
||||
from app.services.crypto import decrypt
|
||||
|
||||
|
||||
class LLMConfigError(RuntimeError):
|
||||
"""用户尚未正确配置 AI 接口。"""
|
||||
|
||||
|
||||
@dataclass
|
||||
class LLMResult:
|
||||
content: str
|
||||
model: str | None
|
||||
prompt_tokens: int | None
|
||||
completion_tokens: int | None
|
||||
|
||||
|
||||
class LLMClient:
|
||||
def __init__(self, settings: UserSettings):
|
||||
if not settings or not settings.llm_base_url or not settings.llm_api_key_encrypted:
|
||||
raise LLMConfigError("请先在设置里配置 AI 接口(base_url 与 api_key)")
|
||||
api_key = decrypt(settings.llm_api_key_encrypted)
|
||||
if not api_key:
|
||||
raise LLMConfigError("api_key 解密失败,请在设置里重新填写")
|
||||
self._client = AsyncOpenAI(base_url=settings.llm_base_url, api_key=api_key)
|
||||
self._text_model = settings.llm_text_model
|
||||
self._vision_model = settings.llm_vision_model or settings.llm_text_model
|
||||
|
||||
async def chat_text(
|
||||
self,
|
||||
system: str,
|
||||
user: str,
|
||||
*,
|
||||
json_mode: bool = False,
|
||||
model: str | None = None,
|
||||
) -> LLMResult:
|
||||
m = model or self._text_model
|
||||
if not m:
|
||||
raise LLMConfigError("未配置文本模型名")
|
||||
kwargs = {}
|
||||
if json_mode:
|
||||
kwargs["response_format"] = {"type": "json_object"}
|
||||
resp = await self._client.chat.completions.create(
|
||||
model=m,
|
||||
messages=[
|
||||
{"role": "system", "content": system},
|
||||
{"role": "user", "content": user},
|
||||
],
|
||||
**kwargs,
|
||||
)
|
||||
return self._to_result(resp, m)
|
||||
|
||||
async def chat_vision(
|
||||
self,
|
||||
system: str,
|
||||
user_text: str,
|
||||
image_bytes: bytes,
|
||||
mime: str = "image/jpeg",
|
||||
*,
|
||||
json_mode: bool = False,
|
||||
model: str | None = None,
|
||||
) -> LLMResult:
|
||||
m = model or self._vision_model
|
||||
if not m:
|
||||
raise LLMConfigError("未配置多模态模型名")
|
||||
data_url = f"data:{mime};base64,{base64.b64encode(image_bytes).decode()}"
|
||||
kwargs = {}
|
||||
if json_mode:
|
||||
kwargs["response_format"] = {"type": "json_object"}
|
||||
resp = await self._client.chat.completions.create(
|
||||
model=m,
|
||||
messages=[
|
||||
{"role": "system", "content": system},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": user_text},
|
||||
{"type": "image_url", "image_url": {"url": data_url}},
|
||||
],
|
||||
},
|
||||
],
|
||||
**kwargs,
|
||||
)
|
||||
return self._to_result(resp, m)
|
||||
|
||||
@staticmethod
|
||||
def _to_result(resp, model: str) -> LLMResult:
|
||||
usage = getattr(resp, "usage", None)
|
||||
return LLMResult(
|
||||
content=resp.choices[0].message.content or "",
|
||||
model=model,
|
||||
prompt_tokens=getattr(usage, "prompt_tokens", None),
|
||||
completion_tokens=getattr(usage, "completion_tokens", None),
|
||||
)
|
||||
@@ -0,0 +1,75 @@
|
||||
"""各 AI 任务的 prompt 模板。"""
|
||||
|
||||
# ---- OCR:多模态直接看图,转成含 LaTeX 的 Markdown(不作答)----
|
||||
VLM_OCR_SYSTEM = (
|
||||
"你是一个精准的题目识别助手。用户会给你一张包含题目的图片,"
|
||||
"请把图片中的题目内容原样转写为 Markdown 文本。要求:\n"
|
||||
"1. 数学公式用 LaTeX 表示,行内用 $...$,独立公式用 $$...$$。\n"
|
||||
"2. 保留题目的选项(如 A/B/C/D)、题号、结构。\n"
|
||||
"3. 只转写题目本身,不要作答、不要解释、不要添加任何额外内容。\n"
|
||||
"4. 如果有多道题,全部转写。"
|
||||
)
|
||||
VLM_OCR_USER = "请转写这张图片里的题目。"
|
||||
|
||||
# ---- 结构化提取:把 OCR/图片内容解析成结构化题目 JSON ----
|
||||
EXTRACT_SYSTEM = """你是一个题目结构化助手。用户会给你一段题目文本(Markdown,可能含 LaTeX 公式),
|
||||
其中可能包含一道或多道题。请把**每一道题**解析成一个对象,输出 JSON:
|
||||
{
|
||||
"questions": [
|
||||
{
|
||||
"type": "unclassified | single_choice | multiple_choice | true_false | fill_blank | short_answer",
|
||||
"stem_markdown": "题干(保留原始 LaTeX)",
|
||||
"options": [{"key": "A", "text_markdown": "..."}] // 仅选择/判断题;否则为 null,
|
||||
"correct_answer": ["A"] 或 ["文本答案"] 或 null, // 原文未给出答案则为 null
|
||||
"difficulty": null 或 1-5
|
||||
}
|
||||
]
|
||||
}
|
||||
题型判断规则:
|
||||
- 有 ABCD 等多个选项且只选一个 → single_choice;可多选 → multiple_choice
|
||||
- 判断对错(正确/错误、对/错、True/False)→ true_false,options 用 [{"key":"T","text_markdown":"正确"},{"key":"F","text_markdown":"错误"}]
|
||||
- 有下划线、括号、填空 → fill_blank
|
||||
- 其余开放性问答 → short_answer
|
||||
- 判断不了、或者这段内容不像一道完整的题 → unclassified(这是可以接受的,不要硬猜)
|
||||
|
||||
重要:
|
||||
- 有几道题就输出几个对象,不要合并成一条,也不要漏掉。
|
||||
- 只有一道题时也要放在 questions 数组里。
|
||||
- 原文没给答案就填 null,不要自己解题填答案。
|
||||
只输出 JSON,不要额外文字。"""
|
||||
|
||||
EXTRACT_USER_TEMPLATE = "题目文本:\n\n{content}"
|
||||
|
||||
# ---- 判分(主观题)----
|
||||
JUDGE_SYSTEM = """你是一个阅卷助手。根据题目、标准答案(可能没有)和学生答案,判断对错并给出反馈。
|
||||
输出 JSON:
|
||||
{
|
||||
"is_correct": true/false,
|
||||
"feedback_markdown": "简短反馈,指出对错原因(可含 LaTeX)",
|
||||
"key_points_missed": ["遗漏的要点", ...]
|
||||
}
|
||||
只输出 JSON。"""
|
||||
|
||||
JUDGE_USER_TEMPLATE = """题目:
|
||||
{stem}
|
||||
|
||||
标准答案:{answer}
|
||||
|
||||
学生答案:
|
||||
{user_answer}"""
|
||||
|
||||
# ---- 讲解 ----
|
||||
EXPLAIN_SYSTEM = (
|
||||
"你是一位耐心的老师。请针对给定题目给出面向自学者的分步讲解,"
|
||||
"解释解题思路而非只给结论。使用 Markdown,公式用 LaTeX($...$ / $$...$$)。"
|
||||
)
|
||||
EXPLAIN_USER_TEMPLATE = """题目:
|
||||
{stem}
|
||||
|
||||
{answer_hint}请给出详细讲解。"""
|
||||
|
||||
# ---- 知识点标签 ----
|
||||
TAGS_SYSTEM = """你是一个知识点归类助手。根据题目内容,输出 1-4 个简短的知识点标签。
|
||||
输出 JSON:{"tags": ["标签1", "标签2"]}
|
||||
标签要简洁(2-8 字),是学科知识点,不要句子。只输出 JSON。"""
|
||||
TAGS_USER_TEMPLATE = "题目:\n{stem}"
|
||||
@@ -0,0 +1,14 @@
|
||||
"""OCR 引擎抽象接口。"""
|
||||
from dataclasses import dataclass
|
||||
from typing import Protocol
|
||||
|
||||
|
||||
@dataclass
|
||||
class OcrResult:
|
||||
markdown: str
|
||||
confidence: float # 0-1,未知时给 1.0
|
||||
engine: str # "pix2text" / "vlm"
|
||||
|
||||
|
||||
class OcrEngine(Protocol):
|
||||
async def recognize(self, image_bytes: bytes, mime: str) -> OcrResult: ...
|
||||
@@ -0,0 +1,35 @@
|
||||
"""OCR 编排:Pix2Text 优先,失败/低置信/未部署时降级到多模态 VLM。"""
|
||||
import logging
|
||||
|
||||
from app.config import settings
|
||||
from app.services.llm.client import LLMClient
|
||||
from app.services.ocr.base import OcrResult
|
||||
from app.services.ocr.pix2text_client import Pix2TextClient
|
||||
from app.services.ocr.vlm_engine import VlmOcrEngine
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Pix2Text 结果置信度低于此值则降级到 VLM
|
||||
_CONFIDENCE_THRESHOLD = 0.5
|
||||
|
||||
|
||||
async def run_ocr(image_bytes: bytes, mime: str, llm: LLMClient) -> OcrResult:
|
||||
"""返回识别结果,engine 字段标明最终使用的引擎。"""
|
||||
# 1. 若配置了 Pix2Text 服务且健康,优先用
|
||||
if settings.ocr_service_url:
|
||||
client = Pix2TextClient(settings.ocr_service_url)
|
||||
if await client.healthy():
|
||||
try:
|
||||
result = await client.recognize(image_bytes, mime)
|
||||
if result.markdown.strip() and result.confidence >= _CONFIDENCE_THRESHOLD:
|
||||
return result
|
||||
logger.info(
|
||||
"Pix2Text 置信度低(%.2f)或结果为空,降级 VLM", result.confidence
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("Pix2Text 识别失败,降级 VLM:%s", e)
|
||||
else:
|
||||
logger.info("Pix2Text 服务不健康,降级 VLM")
|
||||
|
||||
# 2. 降级 / 默认:多模态大模型
|
||||
return await VlmOcrEngine(llm).recognize(image_bytes, mime)
|
||||
@@ -0,0 +1,32 @@
|
||||
"""Pix2Text 独立微服务的 HTTP 客户端(后续阶段部署 OCR 服务时启用)。"""
|
||||
import httpx
|
||||
|
||||
from app.services.ocr.base import OcrResult
|
||||
|
||||
|
||||
class Pix2TextClient:
|
||||
def __init__(self, service_url: str, timeout: float = 60.0):
|
||||
self._url = service_url.rstrip("/")
|
||||
self._timeout = timeout
|
||||
|
||||
async def healthy(self) -> bool:
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=5.0) as c:
|
||||
r = await c.get(f"{self._url}/health")
|
||||
return r.status_code == 200
|
||||
except httpx.HTTPError:
|
||||
return False
|
||||
|
||||
async def recognize(self, image_bytes: bytes, mime: str) -> OcrResult:
|
||||
async with httpx.AsyncClient(timeout=self._timeout) as c:
|
||||
r = await c.post(
|
||||
f"{self._url}/ocr",
|
||||
files={"file": ("image", image_bytes, mime)},
|
||||
)
|
||||
r.raise_for_status()
|
||||
data = r.json()
|
||||
return OcrResult(
|
||||
markdown=data["markdown"],
|
||||
confidence=float(data.get("confidence", 1.0)),
|
||||
engine="pix2text",
|
||||
)
|
||||
@@ -0,0 +1,15 @@
|
||||
"""多模态大模型 OCR 引擎:直接看图转写题目为 Markdown。"""
|
||||
from app.services.llm.client import LLMClient
|
||||
from app.services.llm.prompts import VLM_OCR_SYSTEM, VLM_OCR_USER
|
||||
from app.services.ocr.base import OcrResult
|
||||
|
||||
|
||||
class VlmOcrEngine:
|
||||
def __init__(self, llm: LLMClient):
|
||||
self._llm = llm
|
||||
|
||||
async def recognize(self, image_bytes: bytes, mime: str) -> OcrResult:
|
||||
result = await self._llm.chat_vision(
|
||||
VLM_OCR_SYSTEM, VLM_OCR_USER, image_bytes, mime=mime
|
||||
)
|
||||
return OcrResult(markdown=result.content, confidence=1.0, engine="vlm")
|
||||
@@ -0,0 +1,43 @@
|
||||
"""密码哈希(bcrypt)与 JWT 签发/校验。"""
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import bcrypt
|
||||
import jwt
|
||||
|
||||
from app.config import settings
|
||||
|
||||
_BCRYPT_ROUNDS = 12
|
||||
|
||||
|
||||
def hash_password(password: str) -> str:
|
||||
hashed = bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt(rounds=_BCRYPT_ROUNDS))
|
||||
return hashed.decode("utf-8")
|
||||
|
||||
|
||||
def verify_password(password: str, password_hash: str) -> bool:
|
||||
try:
|
||||
return bcrypt.checkpw(password.encode("utf-8"), password_hash.encode("utf-8"))
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
def create_access_token(subject: str) -> str:
|
||||
"""subject 为 user id。"""
|
||||
now = datetime.now(timezone.utc)
|
||||
payload = {
|
||||
"sub": subject,
|
||||
"iat": now,
|
||||
"exp": now + timedelta(minutes=settings.access_token_expire_minutes),
|
||||
}
|
||||
return jwt.encode(payload, settings.jwt_secret, algorithm=settings.jwt_algorithm)
|
||||
|
||||
|
||||
def decode_access_token(token: str) -> str | None:
|
||||
"""返回 user id,失败返回 None。"""
|
||||
try:
|
||||
payload = jwt.decode(
|
||||
token, settings.jwt_secret, algorithms=[settings.jwt_algorithm]
|
||||
)
|
||||
return payload.get("sub")
|
||||
except jwt.PyJWTError:
|
||||
return None
|
||||
@@ -0,0 +1,38 @@
|
||||
"""启动时的清理工作。"""
|
||||
import logging
|
||||
|
||||
from sqlalchemy import update
|
||||
|
||||
from app.db import SessionLocal
|
||||
from app.models.job import JobStatus, ProcessingJob
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 进程重启后不可能再继续的中间状态
|
||||
_STALE_STATUSES = (
|
||||
JobStatus.pending,
|
||||
JobStatus.ocr_running,
|
||||
JobStatus.ai_running,
|
||||
)
|
||||
|
||||
|
||||
async def fail_stale_jobs() -> int:
|
||||
"""把残留的进行中任务标记为失败。
|
||||
|
||||
任务跑在 BackgroundTasks 里,进程重启就没了。不清理的话这些任务会永远
|
||||
停在 pending/*_running,前端会一直轮询下去。
|
||||
"""
|
||||
async with SessionLocal() as db:
|
||||
result = await db.execute(
|
||||
update(ProcessingJob)
|
||||
.where(ProcessingJob.status.in_(_STALE_STATUSES))
|
||||
.values(
|
||||
status=JobStatus.failed,
|
||||
error_message="服务重启,任务中断,请重新上传",
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
count = result.rowcount or 0
|
||||
if count:
|
||||
logger.info("已清理 %d 个残留任务", count)
|
||||
return count
|
||||
@@ -0,0 +1,99 @@
|
||||
"""对象存储(S3 兼容):上传/下载图片,含校验与去重。
|
||||
|
||||
boto3 是同步库,直接在 async 路径里调用会阻塞事件循环(传大图时整个服务卡住),
|
||||
所以对外暴露 async 版本(to_thread 包装),同步版本保留给测试和脚本用。
|
||||
"""
|
||||
import asyncio
|
||||
import hashlib
|
||||
import io
|
||||
import uuid
|
||||
|
||||
import boto3
|
||||
from botocore.config import Config
|
||||
from PIL import Image as PILImage
|
||||
|
||||
from app.config import settings
|
||||
|
||||
# Pillow 格式 → 扩展名
|
||||
_FORMAT_EXT = {"JPEG": "jpg", "PNG": "png", "WEBP": "webp"}
|
||||
# 解压炸弹防护:单边像素上限
|
||||
_MAX_DIMENSION = 12000
|
||||
|
||||
|
||||
class StorageError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class ImageValidationError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
def _s3_client():
|
||||
if not settings.s3_bucket:
|
||||
raise StorageError("S3 未配置(S3_BUCKET 为空)")
|
||||
return boto3.client(
|
||||
"s3",
|
||||
endpoint_url=settings.s3_endpoint_url,
|
||||
region_name=settings.s3_region,
|
||||
aws_access_key_id=settings.s3_access_key,
|
||||
aws_secret_access_key=settings.s3_secret_key,
|
||||
config=Config(signature_version="s3v4"),
|
||||
)
|
||||
|
||||
|
||||
def validate_image(data: bytes) -> tuple[str, str]:
|
||||
"""校验是真实图片且格式在白名单内。返回 (mime, ext)。"""
|
||||
if len(data) > settings.max_upload_bytes:
|
||||
raise ImageValidationError("文件超过大小上限")
|
||||
try:
|
||||
img = PILImage.open(io.BytesIO(data))
|
||||
img.verify() # 校验完整性(魔数级)
|
||||
# verify 后需重新打开才能读属性
|
||||
img = PILImage.open(io.BytesIO(data))
|
||||
except Exception as e: # noqa: BLE001
|
||||
raise ImageValidationError(f"不是有效的图片文件:{e}") from e
|
||||
|
||||
fmt = img.format or ""
|
||||
if fmt not in _FORMAT_EXT:
|
||||
raise ImageValidationError(f"不支持的图片格式:{fmt}")
|
||||
if img.width > _MAX_DIMENSION or img.height > _MAX_DIMENSION:
|
||||
raise ImageValidationError("图片尺寸过大")
|
||||
|
||||
mime = f"image/{'jpeg' if fmt == 'JPEG' else fmt.lower()}"
|
||||
if mime not in settings.allowed_image_mimes:
|
||||
raise ImageValidationError(f"不支持的 MIME:{mime}")
|
||||
return mime, _FORMAT_EXT[fmt]
|
||||
|
||||
|
||||
def upload_image(user_id: str, data: bytes, ext: str, mime: str) -> str:
|
||||
"""上传到 S3,返回 object_key。"""
|
||||
key = f"{user_id}/{uuid.uuid4().hex}.{ext}"
|
||||
client = _s3_client()
|
||||
client.put_object(
|
||||
Bucket=settings.s3_bucket,
|
||||
Key=key,
|
||||
Body=data,
|
||||
ContentType=mime,
|
||||
)
|
||||
return key
|
||||
|
||||
|
||||
def download_image(object_key: str) -> bytes:
|
||||
client = _s3_client()
|
||||
resp = client.get_object(Bucket=settings.s3_bucket, Key=object_key)
|
||||
return resp["Body"].read()
|
||||
|
||||
|
||||
def sha256_hex(data: bytes) -> str:
|
||||
return hashlib.sha256(data).hexdigest()
|
||||
|
||||
|
||||
# ---- async 包装:在 async 路径里请用这两个 ----
|
||||
|
||||
|
||||
async def upload_image_async(user_id: str, data: bytes, ext: str, mime: str) -> str:
|
||||
return await asyncio.to_thread(upload_image, user_id, data, ext, mime)
|
||||
|
||||
|
||||
async def download_image_async(object_key: str) -> bytes:
|
||||
return await asyncio.to_thread(download_image, object_key)
|
||||
@@ -0,0 +1,36 @@
|
||||
"""标签 get-or-create(按用户隔离)。"""
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models.question import Tag
|
||||
|
||||
|
||||
async def get_or_create_tags(
|
||||
db: AsyncSession, user_id: str, names: list[str]
|
||||
) -> list[Tag]:
|
||||
"""返回该用户下对应名字的 Tag,不存在则创建。不提交(由调用方 commit)。"""
|
||||
cleaned = [n.strip() for n in names if n and n.strip()]
|
||||
if not cleaned:
|
||||
return []
|
||||
# 去重保序
|
||||
seen: dict[str, None] = {}
|
||||
for n in cleaned:
|
||||
seen.setdefault(n, None)
|
||||
unique = list(seen.keys())
|
||||
|
||||
existing = (
|
||||
await db.scalars(
|
||||
select(Tag).where(Tag.user_id == user_id, Tag.name.in_(unique))
|
||||
)
|
||||
).all()
|
||||
by_name = {t.name: t for t in existing}
|
||||
|
||||
result: list[Tag] = []
|
||||
for name in unique:
|
||||
tag = by_name.get(name)
|
||||
if tag is None:
|
||||
tag = Tag(user_id=user_id, name=name)
|
||||
db.add(tag)
|
||||
by_name[name] = tag
|
||||
result.append(tag)
|
||||
return result
|
||||
Reference in New Issue
Block a user