第一次提交
This commit is contained in:
@@ -0,0 +1,383 @@
|
||||
"""OCR 流水线:多模态识别、Pix2Text 降级、提取失败兜底。"""
|
||||
import io
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from PIL import Image as PILImage
|
||||
|
||||
from app.services import storage
|
||||
from app.services.llm.client import LLMResult
|
||||
|
||||
|
||||
def _png_bytes(size=(40, 40)) -> bytes:
|
||||
buf = io.BytesIO()
|
||||
PILImage.new("RGB", size, color="white").save(buf, "PNG")
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
class TestImageValidation:
|
||||
def test_accepts_png(self):
|
||||
mime, ext = storage.validate_image(_png_bytes())
|
||||
assert (mime, ext) == ("image/png", "png")
|
||||
|
||||
def test_rejects_non_image(self):
|
||||
with pytest.raises(storage.ImageValidationError):
|
||||
storage.validate_image(b"this is not an image")
|
||||
|
||||
def test_rejects_oversized(self, monkeypatch):
|
||||
monkeypatch.setattr(storage.settings, "max_upload_bytes", 10)
|
||||
with pytest.raises(storage.ImageValidationError, match="大小上限"):
|
||||
storage.validate_image(_png_bytes())
|
||||
|
||||
def test_rejects_huge_dimensions(self, monkeypatch):
|
||||
"""防解压炸弹:像素维度上限。"""
|
||||
monkeypatch.setattr(storage, "_MAX_DIMENSION", 20)
|
||||
with pytest.raises(storage.ImageValidationError, match="尺寸过大"):
|
||||
storage.validate_image(_png_bytes(size=(50, 50)))
|
||||
|
||||
|
||||
async def _seed_job(client, headers):
|
||||
"""建一张图 + 一个任务,返回 (job_id, user_id, object_key)。"""
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.db import SessionLocal
|
||||
from app.models.image import Image
|
||||
from app.models.job import ProcessingJob
|
||||
from app.models.user import User, UserSettings
|
||||
from app.services import crypto
|
||||
|
||||
async with SessionLocal() as db:
|
||||
user = (await db.scalars(select(User))).first()
|
||||
db.add(
|
||||
UserSettings(
|
||||
user_id=user.id,
|
||||
llm_base_url="http://fake",
|
||||
llm_api_key_encrypted=crypto.encrypt("sk-test"),
|
||||
llm_text_model="text-model",
|
||||
llm_vision_model="vision-model",
|
||||
)
|
||||
)
|
||||
image = Image(
|
||||
user_id=user.id,
|
||||
object_key="k/1.png",
|
||||
mime="image/png",
|
||||
size_bytes=100,
|
||||
sha256="deadbeef",
|
||||
)
|
||||
db.add(image)
|
||||
await db.flush()
|
||||
job = ProcessingJob(user_id=user.id, image_id=image.id, kind="ocr_extract")
|
||||
db.add(job)
|
||||
await db.commit()
|
||||
return job.id, user.id, image.object_key
|
||||
|
||||
|
||||
async def _run_job(client, auth_headers, vision_text: str, extract_text: str):
|
||||
"""跑一次提取流水线,返回 (job, questions)。"""
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.db import SessionLocal
|
||||
from app.models.job import ProcessingJob
|
||||
from app.models.question import Question
|
||||
from app.services.extract_pipeline import run_extract_job
|
||||
|
||||
h = await auth_headers()
|
||||
job_id, user_id, key = await _seed_job(client, h)
|
||||
|
||||
async def fake_vision(self, system, user_text, image_bytes, mime="image/jpeg", **kw):
|
||||
return LLMResult(vision_text, "vision-model", 10, 5)
|
||||
|
||||
async def fake_text(self, system, user, json_mode=False, model=None):
|
||||
return LLMResult(extract_text, "text-model", 8, 4)
|
||||
|
||||
with patch(
|
||||
"app.services.storage.download_image", return_value=_png_bytes()
|
||||
), patch(
|
||||
"app.services.llm.client.LLMClient.chat_vision", fake_vision
|
||||
), patch(
|
||||
"app.services.llm.client.LLMClient.chat_text", fake_text
|
||||
):
|
||||
await run_extract_job(job_id, user_id, key, "image/png")
|
||||
|
||||
async with SessionLocal() as db:
|
||||
job = await db.get(ProcessingJob, job_id)
|
||||
ids = job.result_question_ids or []
|
||||
questions = []
|
||||
if ids:
|
||||
rows = (
|
||||
await db.scalars(select(Question).where(Question.id.in_(ids)))
|
||||
).all()
|
||||
by_id = {q.id: q for q in rows}
|
||||
# 保持流水线产出的顺序
|
||||
questions = [by_id[i] for i in ids if i in by_id]
|
||||
return job, questions
|
||||
|
||||
|
||||
class TestExtractPipeline:
|
||||
async def test_single_question_extraction(self, client, auth_headers):
|
||||
"""未配置 Pix2Text 时走多模态,成功提取出结构化笔记。"""
|
||||
extract = json.dumps(
|
||||
{
|
||||
"questions": [
|
||||
{
|
||||
"type": "single_choice",
|
||||
"stem_markdown": "求 $x^2=9$ 的正根",
|
||||
"options": [
|
||||
{"key": "A", "text_markdown": "2"},
|
||||
{"key": "B", "text_markdown": "3"},
|
||||
],
|
||||
"correct_answer": ["B"],
|
||||
"difficulty": 2,
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
job, qs = await _run_job(client, auth_headers, "求 $x^2=9$ 的正根", extract)
|
||||
|
||||
assert job.status.value == "done"
|
||||
assert job.engine_used == "vlm" # 降级到多模态
|
||||
assert len(qs) == 1
|
||||
assert qs[0].type.value == "single_choice"
|
||||
assert qs[0].correct_answer == ["B"]
|
||||
assert qs[0].ocr_engine == "vlm"
|
||||
|
||||
async def test_one_image_splits_into_many_notes(self, client, auth_headers):
|
||||
"""一张照片常含多道题,应拆成多条独立笔记并共享原图。"""
|
||||
extract = json.dumps(
|
||||
{
|
||||
"questions": [
|
||||
{
|
||||
"type": "single_choice",
|
||||
"stem_markdown": "求 $x^2=9$ 的正根",
|
||||
"options": [{"key": "A", "text_markdown": "3"}],
|
||||
"correct_answer": ["A"],
|
||||
},
|
||||
{
|
||||
"type": "true_false",
|
||||
"stem_markdown": "偶函数关于 y 轴对称",
|
||||
"options": [
|
||||
{"key": "T", "text_markdown": "正确"},
|
||||
{"key": "F", "text_markdown": "错误"},
|
||||
],
|
||||
"correct_answer": ["T"],
|
||||
},
|
||||
{
|
||||
"type": "short_answer",
|
||||
"stem_markdown": "简述牛顿第一定律",
|
||||
"correct_answer": None,
|
||||
},
|
||||
]
|
||||
}
|
||||
)
|
||||
job, qs = await _run_job(client, auth_headers, "三道题的 OCR 文本", extract)
|
||||
|
||||
assert job.status.value == "done"
|
||||
assert len(job.result_question_ids) == 3
|
||||
assert len(qs) == 3
|
||||
assert [q.type.value for q in qs] == [
|
||||
"single_choice",
|
||||
"true_false",
|
||||
"short_answer",
|
||||
]
|
||||
# 都指向同一张原图
|
||||
image_ids = {q.source_image_id for q in qs}
|
||||
assert len(image_ids) == 1 and None not in image_ids
|
||||
|
||||
async def test_original_ocr_text_is_cached(self, client, auth_headers):
|
||||
"""原始 OCR 文本要留在 image 上,便于日后追溯/重新拆分。"""
|
||||
from app.db import SessionLocal
|
||||
from app.models.image import Image as ImageModel
|
||||
|
||||
extract = json.dumps(
|
||||
{"questions": [{"type": "short_answer", "stem_markdown": "第一题"}]}
|
||||
)
|
||||
job, _ = await _run_job(client, auth_headers, "完整的 OCR 原文", extract)
|
||||
|
||||
async with SessionLocal() as db:
|
||||
image = await db.get(ImageModel, job.image_id)
|
||||
assert image.ocr_markdown == "完整的 OCR 原文"
|
||||
|
||||
async def test_malformed_json_falls_back_to_unclassified(
|
||||
self, client, auth_headers
|
||||
):
|
||||
"""整体无法解析时不丢内容:存成一条未分类笔记,原文进题干。"""
|
||||
job, qs = await _run_job(
|
||||
client, auth_headers, "一些无法结构化的题目文字", "这不是 JSON"
|
||||
)
|
||||
|
||||
assert job.status.value == "done"
|
||||
assert len(qs) == 1
|
||||
assert qs[0].type.value == "unclassified"
|
||||
assert qs[0].stem_markdown == "一些无法结构化的题目文字"
|
||||
|
||||
async def test_empty_questions_array_falls_back(self, client, auth_headers):
|
||||
job, qs = await _run_job(
|
||||
client, auth_headers, "原始文本", json.dumps({"questions": []})
|
||||
)
|
||||
assert len(qs) == 1
|
||||
assert qs[0].type.value == "unclassified"
|
||||
assert qs[0].stem_markdown == "原始文本"
|
||||
|
||||
async def test_one_bad_item_does_not_lose_the_others(self, client, auth_headers):
|
||||
"""某一条解析失败,不该拖垮同一张图里的其他题。"""
|
||||
extract = json.dumps(
|
||||
{
|
||||
"questions": [
|
||||
{
|
||||
"type": "single_choice",
|
||||
"stem_markdown": "好的题",
|
||||
"options": [{"key": "A", "text_markdown": "1"}],
|
||||
"correct_answer": ["A"],
|
||||
},
|
||||
{"type": "不存在的题型", "stem_markdown": "题型认不出"},
|
||||
]
|
||||
}
|
||||
)
|
||||
job, qs = await _run_job(client, auth_headers, "OCR 原文", extract)
|
||||
|
||||
assert job.status.value == "done"
|
||||
assert len(qs) == 2
|
||||
assert qs[0].type.value == "single_choice"
|
||||
# 认不出的题型退化为未分类,但内容保住了
|
||||
assert qs[1].type.value == "unclassified"
|
||||
assert qs[1].stem_markdown == "题型认不出"
|
||||
|
||||
async def test_json_in_code_fence_is_parsed(self, client, auth_headers):
|
||||
"""模型常把 JSON 包在 ```json 代码块里,要能容忍。"""
|
||||
fenced = (
|
||||
'```json\n{"questions": [{"type": "short_answer", '
|
||||
'"stem_markdown": "简述牛顿第一定律"}]}\n```'
|
||||
)
|
||||
job, qs = await _run_job(client, auth_headers, "简述牛顿第一定律", fenced)
|
||||
|
||||
assert job.status.value == "done"
|
||||
assert len(qs) == 1
|
||||
assert qs[0].stem_markdown == "简述牛顿第一定律"
|
||||
|
||||
async def test_extracted_notes_need_no_answer(self, client, auth_headers):
|
||||
"""图上没答案是常态,不该因此失败或硬填。"""
|
||||
extract = json.dumps(
|
||||
{
|
||||
"questions": [
|
||||
{
|
||||
"type": "short_answer",
|
||||
"stem_markdown": "证明勾股定理",
|
||||
"correct_answer": None,
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
_, qs = await _run_job(client, auth_headers, "证明勾股定理", extract)
|
||||
assert qs[0].correct_answer is None
|
||||
|
||||
async def test_missing_ai_config_fails_job_gracefully(self, client, auth_headers):
|
||||
"""用户没配 AI 时任务应标记失败并给出可读原因,而不是卡住。"""
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.db import SessionLocal
|
||||
from app.models.image import Image
|
||||
from app.models.job import ProcessingJob
|
||||
from app.models.user import User
|
||||
from app.services.extract_pipeline import run_extract_job
|
||||
|
||||
await auth_headers()
|
||||
async with SessionLocal() as db:
|
||||
user = (await db.scalars(select(User))).first()
|
||||
image = Image(
|
||||
user_id=user.id,
|
||||
object_key="k/2.png",
|
||||
mime="image/png",
|
||||
size_bytes=1,
|
||||
sha256="beef",
|
||||
)
|
||||
db.add(image)
|
||||
await db.flush()
|
||||
job = ProcessingJob(user_id=user.id, image_id=image.id)
|
||||
db.add(job)
|
||||
await db.commit()
|
||||
job_id, user_id = job.id, user.id
|
||||
|
||||
# 没有 UserSettings → LLMConfigError
|
||||
await run_extract_job(job_id, user_id, "k/2.png", "image/png")
|
||||
|
||||
async with SessionLocal() as db:
|
||||
job = await db.get(ProcessingJob, job_id)
|
||||
assert job.status.value == "failed"
|
||||
assert "AI" in job.error_message
|
||||
|
||||
|
||||
class TestOcrFallback:
|
||||
async def test_pix2text_unhealthy_falls_back_to_vlm(self, monkeypatch):
|
||||
"""配了 Pix2Text 但服务不健康时,应自动降级到多模态。"""
|
||||
from app.services.ocr.pipeline import run_ocr
|
||||
|
||||
monkeypatch.setattr(
|
||||
"app.services.ocr.pipeline.settings.ocr_service_url", "http://ocr:8001"
|
||||
)
|
||||
|
||||
class FakeLLM:
|
||||
async def chat_vision(self, *a, **kw):
|
||||
return LLMResult("vlm output", "vision-model", 1, 1)
|
||||
|
||||
async def unhealthy(self):
|
||||
return False
|
||||
|
||||
with patch(
|
||||
"app.services.ocr.pix2text_client.Pix2TextClient.healthy", unhealthy
|
||||
):
|
||||
result = await run_ocr(b"fake", "image/png", FakeLLM())
|
||||
|
||||
assert result.engine == "vlm"
|
||||
assert result.markdown == "vlm output"
|
||||
|
||||
async def test_pix2text_used_when_healthy(self, monkeypatch):
|
||||
from app.services.ocr.base import OcrResult
|
||||
from app.services.ocr.pipeline import run_ocr
|
||||
|
||||
monkeypatch.setattr(
|
||||
"app.services.ocr.pipeline.settings.ocr_service_url", "http://ocr:8001"
|
||||
)
|
||||
|
||||
async def healthy(self):
|
||||
return True
|
||||
|
||||
async def recognize(self, image_bytes, mime):
|
||||
return OcrResult(markdown="pix2text output", confidence=0.9, engine="pix2text")
|
||||
|
||||
with patch(
|
||||
"app.services.ocr.pix2text_client.Pix2TextClient.healthy", healthy
|
||||
), patch(
|
||||
"app.services.ocr.pix2text_client.Pix2TextClient.recognize", recognize
|
||||
):
|
||||
result = await run_ocr(b"fake", "image/png", None)
|
||||
|
||||
assert result.engine == "pix2text"
|
||||
|
||||
async def test_low_confidence_falls_back_to_vlm(self, monkeypatch):
|
||||
"""Pix2Text 置信度过低时也要降级。"""
|
||||
from app.services.ocr.base import OcrResult
|
||||
from app.services.ocr.pipeline import run_ocr
|
||||
|
||||
monkeypatch.setattr(
|
||||
"app.services.ocr.pipeline.settings.ocr_service_url", "http://ocr:8001"
|
||||
)
|
||||
|
||||
class FakeLLM:
|
||||
async def chat_vision(self, *a, **kw):
|
||||
return LLMResult("vlm output", "vision-model", 1, 1)
|
||||
|
||||
async def healthy(self):
|
||||
return True
|
||||
|
||||
async def low_conf(self, image_bytes, mime):
|
||||
return OcrResult(markdown="garbled", confidence=0.1, engine="pix2text")
|
||||
|
||||
with patch(
|
||||
"app.services.ocr.pix2text_client.Pix2TextClient.healthy", healthy
|
||||
), patch(
|
||||
"app.services.ocr.pix2text_client.Pix2TextClient.recognize", low_conf
|
||||
):
|
||||
result = await run_ocr(b"fake", "image/png", FakeLLM())
|
||||
|
||||
assert result.engine == "vlm"
|
||||
Reference in New Issue
Block a user