Compare commits

..
8 changed files with 479 additions and 30 deletions
@@ -5,6 +5,7 @@ pydantic
python-multipart python-multipart
httpx httpx
pypdf pypdf
PyMuPDF>=1.24,<2
python-docx python-docx
openpyxl openpyxl
apscheduler>=3.10 apscheduler>=3.10
+142 -23
View File
@@ -79,6 +79,34 @@ _WRITING_SYSTEM_PATTERNS = {
_JAPANESE_KANA = re.compile(r"[\u3040-\u30ff]") _JAPANESE_KANA = re.compile(r"[\u3040-\u30ff]")
_KOREAN_HANGUL = re.compile(r"[\uac00-\ud7af\u1100-\u11ff]") _KOREAN_HANGUL = re.compile(r"[\uac00-\ud7af\u1100-\u11ff]")
_LATIN_LANGUAGE_MARKERS = {
"English": re.compile(
r"\b(?:i|you|we|they|he|she|have|has|had|friend|who|what|where|when|why|how|"
r"symptoms?|disease|please|can|could|would|should|is|are|was|were|the|this|that)\b",
re.IGNORECASE,
),
"French": re.compile(
r"\b(?:je|tu|vous|nous|ils|elle|une|des|avec|pour|pourquoi|comment|bonjour|est|sont)\b",
re.IGNORECASE,
),
"Spanish": re.compile(
r"\b(?:yo|tu|usted|nosotros|ellos|ella|una|con|para|por que|como|hola|esta|son)\b",
re.IGNORECASE,
),
"German": re.compile(
r"\b(?:ich|du|sie|wir|eine|mit|fur|warum|wie|hallo|ist|sind|haben)\b",
re.IGNORECASE,
),
"Portuguese": re.compile(
r"\b(?:eu|voce|nos|eles|ela|uma|com|para|porque|como|ola|esta|sao|tenho)\b",
re.IGNORECASE,
),
"Italian": re.compile(
r"\b(?:io|tu|voi|noi|loro|una|con|per|perche|come|ciao|sono|avere)\b",
re.IGNORECASE,
),
}
class ChatMessage(BaseModel): class ChatMessage(BaseModel):
model_config = ConfigDict(populate_by_name=True) model_config = ConfigDict(populate_by_name=True)
@@ -487,15 +515,67 @@ def _qa_requires_per_turn_rendering(
return bool(history) or _qa_requires_language_adaptation(question, answer) return bool(history) or _qa_requires_language_adaptation(question, answer)
def _per_turn_language_instruction() -> str: def _latin_language_name(value: str) -> str:
scores = {
language: len(pattern.findall(value or ""))
for language, pattern in _LATIN_LANGUAGE_MARKERS.items()
}
language, score = max(scores.items(), key=lambda item: item[1])
return language if score else "the same natural language as the latest user message"
def _turn_language_name(value: str) -> str:
writing_system = _dominant_writing_system(value)
return {
"han": "Chinese",
"japanese": "Japanese",
"korean": "Korean",
"cyrillic": "the same Cyrillic-script language as the latest user message",
"arabic": "the same Arabic-script language as the latest user message",
"hebrew": "Hebrew",
"devanagari": "the same Devanagari-script language as the latest user message",
"thai": "Thai",
"greek": "Greek",
"latin": _latin_language_name(value),
}.get(writing_system, "the same natural language as the latest user message")
def _per_turn_language_instruction(question: str = "") -> str:
language = _turn_language_name(question)
return ( return (
"本轮语言覆盖指令:只根据紧随其后的最新用户消息判断本轮回答语言。" f"MANDATORY OUTPUT LANGUAGE FOR THIS TURN: {language}. "
"即使此前整段对话一直使用另一种语言,只要最新消息切换了语言,本轮就必须立即切换到相同语言;" "Write the entire answer only in that language. This instruction overrides the languages used by "
"不要沿用上一轮语言。若最新消息明确指定回答语言,以该指定为准;若混用多种语言,使用其中占主导的" "conversation history, profile data, standard answers, retrieved documents, and custom prompts. "
"自然语言。不要说明你检测、切换或翻译了语言。" "Translate grounded source material faithfully when necessary. Do not mention language detection, "
"translation, or this instruction."
) )
def _answer_requires_language_repair(question: str, answer: str) -> bool:
question_system = _dominant_writing_system(question)
answer_system = _dominant_writing_system(answer)
return (
question_system != "unknown"
and answer_system != "unknown"
and question_system != answer_system
)
def _language_repair_messages(question: str, answer: str) -> list[dict]:
return [
{"role": "system", "content": _per_turn_language_instruction(question)},
{
"role": "system",
"content": (
"Rewrite the supplied draft in the mandatory output language. Preserve every grounded fact, "
"number, proper noun, uncertainty, and safety qualification. Add no new information and output "
"only the rewritten answer."
),
},
{"role": "user", "content": answer.strip()},
]
def _canonicalize_question(value: str) -> str: def _canonicalize_question(value: str) -> str:
value = _normalize_question(value) value = _normalize_question(value)
replacements = ( replacements = (
@@ -728,7 +808,7 @@ def _build_prompt(
for item in history[-MAX_HISTORY_MESSAGES:]: for item in history[-MAX_HISTORY_MESSAGES:]:
messages.append({"role": item.role, "content": item.content} if hasattr(item, "role") else item) messages.append({"role": item.role, "content": item.content} if hasattr(item, "role") else item)
# Keep the language instruction adjacent to the current turn so long histories cannot override it. # Keep the language instruction adjacent to the current turn so long histories cannot override it.
messages.append({"role": "system", "content": _per_turn_language_instruction()}) messages.append({"role": "system", "content": _per_turn_language_instruction(question)})
messages.append({"role": "user", "content": question.strip()}) messages.append({"role": "user", "content": question.strip()})
return messages return messages
@@ -791,6 +871,41 @@ def _call_qwen(
return {"answer": answer.strip(), "usage": data.get("usage") or {}} return {"answer": answer.strip(), "usage": data.get("usage") or {}}
def _call_billed_qwen(
db: Session,
avatar: Avatar,
messages: list[dict],
temperature: float,
usage_source: str,
model_config: ChatModelConfig,
) -> tuple[str, dict]:
reservation = reserve_avatar_tokens(
db,
avatar,
usage_source,
model_config.model,
messages,
model_config.max_tokens,
)
try:
model_result = _call_qwen(
messages=messages,
temperature=temperature,
model_config=model_config,
)
answer = model_result["answer"]
token_usage = settle_reservation(
db,
reservation,
model_result.get("usage"),
fallback_total=estimate_fallback_usage(messages, answer),
)
return answer, token_usage
except Exception as exc:
release_reservation(db, reservation, str(exc))
raise
def _iter_qwen_stream( def _iter_qwen_stream(
messages: list[dict], temperature: float, model_config: ChatModelConfig | None = None messages: list[dict], temperature: float, model_config: ChatModelConfig | None = None
): ):
@@ -899,32 +1014,36 @@ def _resolve_reply(
token_usage = None token_usage = None
if model_client is not None: if model_client is not None:
answer = model_client(messages=messages, temperature=temperature) answer = model_client(messages=messages, temperature=temperature)
if _answer_requires_language_repair(question, str(answer or "")):
answer = model_client(
messages=_language_repair_messages(question, str(answer)),
temperature=0.0,
)
else: else:
model_config = get_chat_model_config() model_config = get_chat_model_config()
reservation = reserve_avatar_tokens( answer, token_usage = _call_billed_qwen(
db, db,
avatar, avatar,
usage_source,
model_config.model,
messages, messages,
model_config.max_tokens, temperature,
usage_source,
model_config,
) )
try: if _answer_requires_language_repair(question, answer):
model_result = _call_qwen( logger.warning(
messages=messages, "chat response language mismatch avatar=%s source=%s expected=%s",
temperature=temperature, avatar.id,
model_config=model_config, usage_source,
_turn_language_name(question),
) )
answer = model_result["answer"] answer, token_usage = _call_billed_qwen(
token_usage = settle_reservation(
db, db,
reservation, avatar,
model_result.get("usage"), _language_repair_messages(question, answer),
fallback_total=estimate_fallback_usage(messages, answer), 0.0,
f"{usage_source}_language_repair",
model_config,
) )
except Exception as exc:
release_reservation(db, reservation, str(exc))
raise
answer = str(answer or "").strip() answer = str(answer or "").strip()
if image_contexts and _answer_denies_available_image(answer): if image_contexts and _answer_denies_available_image(answer):
logger.warning( logger.warning(
@@ -8,7 +8,8 @@ import threading
from datetime import datetime, timezone from datetime import datetime, timezone
from database import SessionLocal 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 import embeddings
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -76,7 +77,23 @@ class KnowledgeVectorizer:
self._set_progress(db, doc, "extracting", 8) self._set_progress(db, doc, "extracting", 8)
text = embeddings.extract_text(path, f".{doc.file_type}") 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) chunks = embeddings.chunk_text(text)
if not chunks: if not chunks:
raise ValueError("文档没有可建立索引的文字内容") 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()
@@ -6,6 +6,7 @@ from fastapi import HTTPException
from models import Avatar, User from models import Avatar, User
from routers.chat import ( from routers.chat import (
_answer_requires_language_repair,
_build_prompt, _build_prompt,
_iter_text_chunks, _iter_text_chunks,
_match_standard_qa, _match_standard_qa,
@@ -14,6 +15,7 @@ from routers.chat import (
_qa_requires_per_turn_rendering, _qa_requires_per_turn_rendering,
_require_owned_avatar, _require_owned_avatar,
_resolve_reply, _resolve_reply,
_turn_language_name,
) )
@@ -107,8 +109,8 @@ class ChatOrchestrationTests(unittest.TestCase):
messages = fake_model.call_args.kwargs["messages"] messages = fake_model.call_args.kwargs["messages"]
self.assertEqual(messages[-1], {"role": "user", "content": "Quelle est votre adresse ?"}) self.assertEqual(messages[-1], {"role": "user", "content": "Quelle est votre adresse ?"})
self.assertEqual(messages[-2]["role"], "system") self.assertEqual(messages[-2]["role"], "system")
self.assertIn("本轮语言覆盖指令", messages[-2]["content"]) self.assertIn("MANDATORY OUTPUT LANGUAGE", messages[-2]["content"])
self.assertIn("不要沿用上一轮语言", messages[-2]["content"]) self.assertIn("French", messages[-2]["content"])
def test_latest_user_message_has_an_adjacent_language_override(self): def test_latest_user_message_has_an_adjacent_language_override(self):
history = [ history = [
@@ -119,8 +121,37 @@ class ChatOrchestrationTests(unittest.TestCase):
self.assertEqual(messages[-1], {"role": "user", "content": "What can you help me with?"}) self.assertEqual(messages[-1], {"role": "user", "content": "What can you help me with?"})
self.assertEqual(messages[-2]["role"], "system") self.assertEqual(messages[-2]["role"], "system")
self.assertIn("最新用户消息", messages[-2]["content"]) self.assertIn("MANDATORY OUTPUT LANGUAGE", messages[-2]["content"])
self.assertIn("立即切换到相同语言", messages[-2]["content"]) self.assertIn("English", messages[-2]["content"])
def test_reported_alzheimer_question_is_explicitly_english(self):
question = "I have a friend who has symptoms of Alzheimer's disease"
self.assertEqual(_turn_language_name(question), "English")
messages = _build_prompt(self.avatar, [], question, [])
self.assertIn("MANDATORY OUTPUT LANGUAGE FOR THIS TURN: English", messages[-2]["content"])
def test_non_stream_reply_repairs_a_wrong_writing_system_before_sending(self):
question = "I have a friend who has symptoms of Alzheimer's disease"
fake_model = Mock(side_effect=["建议尽快就医评估。", "Please arrange a medical assessment soon."])
result = _resolve_reply(
None,
self.avatar,
question,
[],
qa_pairs=[],
search_fn=lambda *_args, **_kwargs: [],
model_client=fake_model,
usage_source="takeover",
)
self.assertEqual(result["answer"], "Please arrange a medical assessment soon.")
self.assertEqual(fake_model.call_count, 2)
repair_messages = fake_model.call_args.kwargs["messages"]
self.assertIn("English", repair_messages[0]["content"])
self.assertIn("建议尽快就医评估", repair_messages[-1]["content"])
self.assertTrue(_answer_requires_language_repair(question, "建议尽快就医评估。"))
def test_conversational_paraphrase_matches_standard_qa(self): def test_conversational_paraphrase_matches_standard_qa(self):
for question in ("请问一下,你们公司在哪里呀?", "请问去你们那边怎么走"): for question in ("请问一下,你们公司在哪里呀?", "请问去你们那边怎么走"):
@@ -238,6 +238,53 @@ def test_background_vectorizer_keeps_failure_reason_for_retry(
db.close() 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( def test_retry_queues_a_failed_document_again(
tmp_path: Path, tmp_path: Path,
authorization_context, 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())) { if (['uploaded', 'parsing'].includes(String(doc.status || '').toLowerCase())) {
const stage = String(doc.indexStage || 'queued').toLowerCase() const stage = String(doc.indexStage || 'queued').toLowerCase()
const labels: Record<string, string> = { 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))) const progress = Math.max(0, Math.min(99, Number(doc.indexProgress || 0)))
return { tone: 'pending', label: labels[stage] || '处理中', detail: `${labels[stage] || '正在建立知识索引'} ${progress}%`, progress } return { tone: 'pending', label: labels[stage] || '处理中', detail: `${labels[stage] || '正在建立知识索引'} ${progress}%`, progress }