fix(avatar): OCR image-only knowledge PDFs
This commit is contained in:
@@ -5,6 +5,7 @@ pydantic
|
||||
python-multipart
|
||||
httpx
|
||||
pypdf
|
||||
PyMuPDF>=1.24,<2
|
||||
python-docx
|
||||
openpyxl
|
||||
apscheduler>=3.10
|
||||
|
||||
@@ -8,7 +8,8 @@ import threading
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from database import SessionLocal
|
||||
from models import KnowledgeChunk, KnowledgeDoc
|
||||
from models import Avatar, KnowledgeChunk, KnowledgeDoc
|
||||
from services.pdf_ocr_service import extract_scanned_pdf_text
|
||||
import embeddings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -76,7 +77,23 @@ class KnowledgeVectorizer:
|
||||
|
||||
self._set_progress(db, doc, "extracting", 8)
|
||||
text = embeddings.extract_text(path, f".{doc.file_type}")
|
||||
self._set_progress(db, doc, "chunking", 22)
|
||||
if doc.file_type == "pdf" and not text.strip():
|
||||
avatar = db.get(Avatar, doc.avatar_id)
|
||||
if not avatar:
|
||||
raise ValueError("文档所属分身不存在")
|
||||
|
||||
def ocr_progress(done: int, total: int):
|
||||
percent = 8 + int((done / max(1, total)) * 20)
|
||||
self._set_progress(db, doc, "ocr", min(percent, 28))
|
||||
|
||||
self._set_progress(db, doc, "ocr", 8)
|
||||
text = extract_scanned_pdf_text(
|
||||
db,
|
||||
avatar,
|
||||
path,
|
||||
on_progress=ocr_progress,
|
||||
)
|
||||
self._set_progress(db, doc, "chunking", 29)
|
||||
chunks = embeddings.chunk_text(text)
|
||||
if not chunks:
|
||||
raise ValueError("文档没有可建立索引的文字内容")
|
||||
|
||||
@@ -0,0 +1,130 @@
|
||||
"""OCR fallback for image-only PDF knowledge documents."""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Callable
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from models import Avatar
|
||||
from services.chat_model_config import get_chat_model_config
|
||||
from services.token_billing import (
|
||||
estimate_fallback_usage,
|
||||
release_reservation,
|
||||
reserve_avatar_tokens,
|
||||
settle_reservation,
|
||||
)
|
||||
from services.vision_service import call_vision_model, prepare_image
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
PDF_OCR_PROMPT = (
|
||||
"请逐字转录这一页扫描文档中的全部可见文字和表格,只输出转录内容,不要解释,不要使用 Markdown 代码块。"
|
||||
"保留标题、段落、项目编号、数值和自然换行;看不清的内容写作[无法辨认],不要猜测、纠错或补全。"
|
||||
)
|
||||
|
||||
|
||||
def _positive_int(name: str, default: int, minimum: int, maximum: int) -> int:
|
||||
try:
|
||||
value = int(os.getenv(name, str(default)))
|
||||
except ValueError:
|
||||
value = default
|
||||
return max(minimum, min(maximum, value))
|
||||
|
||||
|
||||
def extract_scanned_pdf_text(
|
||||
db: Session,
|
||||
avatar: Avatar,
|
||||
path: str,
|
||||
*,
|
||||
on_progress: Callable[[int, int], None] | None = None,
|
||||
) -> str:
|
||||
"""Render and OCR an image-only PDF while preserving page order."""
|
||||
try:
|
||||
import pymupdf
|
||||
except ImportError as exc:
|
||||
raise RuntimeError("扫描型 PDF 识别组件未安装") from exc
|
||||
|
||||
max_pages = _positive_int("KNOWLEDGE_PDF_OCR_MAX_PAGES", 80, 1, 300)
|
||||
render_dpi = _positive_int("KNOWLEDGE_PDF_OCR_DPI", 144, 96, 200)
|
||||
max_attempts = _positive_int("KNOWLEDGE_PDF_OCR_ATTEMPTS", 3, 1, 5)
|
||||
model_config = get_chat_model_config()
|
||||
model = model_config.ocr_model or model_config.vision_model
|
||||
if not model_config.api_key or not model:
|
||||
raise RuntimeError("扫描型 PDF 需要配置视觉 OCR 模型")
|
||||
|
||||
texts: list[str] = []
|
||||
with pymupdf.open(path) as document:
|
||||
total_pages = document.page_count
|
||||
if total_pages <= 0:
|
||||
raise ValueError("PDF 没有可识别页面")
|
||||
if total_pages > max_pages:
|
||||
raise ValueError(
|
||||
f"扫描型 PDF 共 {total_pages} 页,超过单次 OCR 上限 {max_pages} 页,请拆分后上传"
|
||||
)
|
||||
|
||||
scale = render_dpi / 72
|
||||
for page_index in range(total_pages):
|
||||
page = document.load_page(page_index)
|
||||
pixmap = page.get_pixmap(
|
||||
matrix=pymupdf.Matrix(scale, scale),
|
||||
colorspace=pymupdf.csRGB,
|
||||
alpha=False,
|
||||
)
|
||||
prepared = prepare_image(pixmap.tobytes("jpeg", jpg_quality=88))
|
||||
estimate_messages = [{
|
||||
"role": "user",
|
||||
"content": f"[扫描 PDF 第 {page_index + 1}/{total_pages} 页]\n{PDF_OCR_PROMPT}",
|
||||
}]
|
||||
reservation = reserve_avatar_tokens(
|
||||
db,
|
||||
avatar,
|
||||
"knowledge_pdf_ocr",
|
||||
model,
|
||||
estimate_messages,
|
||||
model_config.vision_max_tokens,
|
||||
)
|
||||
try:
|
||||
result = None
|
||||
for attempt in range(1, max_attempts + 1):
|
||||
try:
|
||||
result = call_vision_model(
|
||||
prepared,
|
||||
model_config,
|
||||
model=model,
|
||||
prompt=PDF_OCR_PROMPT,
|
||||
json_output=False,
|
||||
)
|
||||
break
|
||||
except RuntimeError:
|
||||
if attempt == max_attempts:
|
||||
raise
|
||||
time.sleep(min(4, attempt))
|
||||
content = str((result or {}).get("content") or "").strip()
|
||||
if not content:
|
||||
raise RuntimeError("扫描型 PDF 页面识别结果为空")
|
||||
settle_reservation(
|
||||
db,
|
||||
reservation,
|
||||
(result or {}).get("usage"),
|
||||
fallback_total=estimate_fallback_usage(estimate_messages, content),
|
||||
)
|
||||
except Exception as exc:
|
||||
release_reservation(db, reservation, str(exc))
|
||||
raise RuntimeError(
|
||||
f"扫描型 PDF 第 {page_index + 1}/{total_pages} 页识别失败:{exc}"
|
||||
) from exc
|
||||
|
||||
texts.append(f"[第 {page_index + 1} 页]\n{content}")
|
||||
if on_progress:
|
||||
on_progress(page_index + 1, total_pages)
|
||||
logger.info(
|
||||
"Scanned PDF OCR completed avatar=%s page=%s/%s",
|
||||
avatar.id,
|
||||
page_index + 1,
|
||||
total_pages,
|
||||
)
|
||||
|
||||
return "\n\n".join(texts).strip()
|
||||
@@ -238,6 +238,53 @@ def test_background_vectorizer_keeps_failure_reason_for_retry(
|
||||
db.close()
|
||||
|
||||
|
||||
def test_background_vectorizer_uses_ocr_for_image_only_pdf(
|
||||
tmp_path: Path,
|
||||
authorization_context,
|
||||
):
|
||||
context = authorization_context
|
||||
with (
|
||||
patch("routers.knowledge.UPLOAD_DIR", str(tmp_path)),
|
||||
patch("routers.knowledge.knowledge_vectorizer.enqueue"),
|
||||
):
|
||||
response = client.post(
|
||||
f"/api/avatar/{context['avatar'].id}/knowledge/docs",
|
||||
headers=context["owner_headers"],
|
||||
files={"file": ("scanned.pdf", b"image-only-pdf", "application/pdf")},
|
||||
)
|
||||
|
||||
payload = response.json()["data"]
|
||||
progress = []
|
||||
with (
|
||||
patch("services.knowledge_vectorizer.UPLOAD_DIR", str(tmp_path)),
|
||||
patch("services.knowledge_vectorizer.embeddings.extract_text", return_value=""),
|
||||
patch(
|
||||
"services.knowledge_vectorizer.extract_scanned_pdf_text",
|
||||
side_effect=lambda _db, _avatar, _path, on_progress: (
|
||||
on_progress(1, 2), on_progress(2, 2), "扫描页文字"
|
||||
)[-1],
|
||||
) as ocr,
|
||||
patch("services.knowledge_vectorizer.embeddings.embed", return_value=[[1.0, 0.0]]),
|
||||
patch.object(knowledge_vectorizer, "_set_progress", wraps=knowledge_vectorizer._set_progress) as set_progress,
|
||||
):
|
||||
knowledge_vectorizer.vectorize_document(payload["id"])
|
||||
progress = [(call.args[2], call.args[3]) for call in set_progress.call_args_list]
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
stored = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == payload["id"]).one()
|
||||
assert stored.status == "ready"
|
||||
assert stored.chunk_count == 1
|
||||
assert ("ocr", 18) in progress
|
||||
assert ("ocr", 28) in progress
|
||||
ocr.assert_called_once()
|
||||
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == stored.id).delete()
|
||||
db.delete(stored)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_retry_queues_a_failed_document_again(
|
||||
tmp_path: Path,
|
||||
authorization_context,
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from services.pdf_ocr_service import extract_scanned_pdf_text
|
||||
|
||||
|
||||
class FakePixmap:
|
||||
def tobytes(self, *_args, **_kwargs):
|
||||
return b"jpeg-page"
|
||||
|
||||
|
||||
class FakePage:
|
||||
def get_pixmap(self, **_kwargs):
|
||||
return FakePixmap()
|
||||
|
||||
|
||||
class FakeDocument:
|
||||
page_count = 2
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *_args):
|
||||
return None
|
||||
|
||||
def load_page(self, _index):
|
||||
return FakePage()
|
||||
|
||||
|
||||
def test_scanned_pdf_ocr_preserves_page_order_and_reports_progress(monkeypatch):
|
||||
fake_pymupdf = SimpleNamespace(
|
||||
open=lambda _path: FakeDocument(),
|
||||
Matrix=lambda x, y: (x, y),
|
||||
csRGB="rgb",
|
||||
)
|
||||
monkeypatch.setitem(__import__("sys").modules, "pymupdf", fake_pymupdf)
|
||||
progress = []
|
||||
reservation = SimpleNamespace()
|
||||
config = SimpleNamespace(
|
||||
api_key="configured",
|
||||
ocr_model="qwen-vl-ocr",
|
||||
vision_model="vision",
|
||||
vision_max_tokens=2048,
|
||||
)
|
||||
|
||||
with (
|
||||
patch("services.pdf_ocr_service.get_chat_model_config", return_value=config),
|
||||
patch("services.pdf_ocr_service.prepare_image", return_value=SimpleNamespace()),
|
||||
patch(
|
||||
"services.pdf_ocr_service.call_vision_model",
|
||||
side_effect=[
|
||||
{"content": "第一页文字", "usage": {"total_tokens": 10}},
|
||||
{"content": "第二页文字", "usage": {"total_tokens": 12}},
|
||||
],
|
||||
),
|
||||
patch("services.pdf_ocr_service.reserve_avatar_tokens", return_value=reservation) as reserve,
|
||||
patch("services.pdf_ocr_service.settle_reservation") as settle,
|
||||
):
|
||||
text = extract_scanned_pdf_text(
|
||||
MagicMock(),
|
||||
SimpleNamespace(id="avatar-1"),
|
||||
"/tmp/scanned.pdf",
|
||||
on_progress=lambda done, total: progress.append((done, total)),
|
||||
)
|
||||
|
||||
assert text == "[第 1 页]\n第一页文字\n\n[第 2 页]\n第二页文字"
|
||||
assert progress == [(1, 2), (2, 2)]
|
||||
assert reserve.call_count == 2
|
||||
assert settle.call_count == 2
|
||||
|
||||
|
||||
def test_scanned_pdf_ocr_releases_tokens_after_retries_fail(monkeypatch):
|
||||
fake_document = FakeDocument()
|
||||
fake_document.page_count = 1
|
||||
fake_pymupdf = SimpleNamespace(
|
||||
open=lambda _path: fake_document,
|
||||
Matrix=lambda x, y: (x, y),
|
||||
csRGB="rgb",
|
||||
)
|
||||
monkeypatch.setitem(__import__("sys").modules, "pymupdf", fake_pymupdf)
|
||||
monkeypatch.setenv("KNOWLEDGE_PDF_OCR_ATTEMPTS", "2")
|
||||
reservation = SimpleNamespace()
|
||||
config = SimpleNamespace(
|
||||
api_key="configured",
|
||||
ocr_model="qwen-vl-ocr",
|
||||
vision_model="vision",
|
||||
vision_max_tokens=2048,
|
||||
)
|
||||
|
||||
with (
|
||||
patch("services.pdf_ocr_service.get_chat_model_config", return_value=config),
|
||||
patch("services.pdf_ocr_service.prepare_image", return_value=SimpleNamespace()),
|
||||
patch("services.pdf_ocr_service.call_vision_model", side_effect=RuntimeError("timeout")) as call,
|
||||
patch("services.pdf_ocr_service.reserve_avatar_tokens", return_value=reservation),
|
||||
patch("services.pdf_ocr_service.release_reservation") as release,
|
||||
patch("services.pdf_ocr_service.time.sleep"),
|
||||
):
|
||||
with pytest.raises(RuntimeError, match="第 1/1 页识别失败"):
|
||||
extract_scanned_pdf_text(MagicMock(), SimpleNamespace(id="avatar-1"), "/tmp/scanned.pdf")
|
||||
|
||||
assert call.call_count == 2
|
||||
release.assert_called_once()
|
||||
@@ -147,7 +147,7 @@ const documentState = (doc: any) => {
|
||||
if (['uploaded', 'parsing'].includes(String(doc.status || '').toLowerCase())) {
|
||||
const stage = String(doc.indexStage || 'queued').toLowerCase()
|
||||
const labels: Record<string, string> = {
|
||||
queued: '等待处理', extracting: '解析文档', chunking: '切分文本', embedding: '向量化中'
|
||||
queued: '等待处理', extracting: '解析文档', ocr: '扫描件识别', chunking: '切分文本', embedding: '向量化中'
|
||||
}
|
||||
const progress = Math.max(0, Math.min(99, Number(doc.indexProgress || 0)))
|
||||
return { tone: 'pending', label: labels[stage] || '处理中', detail: `${labels[stage] || '正在建立知识索引'} ${progress}%`, progress }
|
||||
|
||||
Reference in New Issue
Block a user