167 lines
6.0 KiB
Python
167 lines
6.0 KiB
Python
"""上传图片 → 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()
|