Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
857d6f2562 | ||
|
|
9afc2d5a6c | ||
|
|
8585d101d5 | ||
|
|
62eb9578fd | ||
|
|
848657219e |
@@ -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
|
||||||
|
|||||||
@@ -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 }
|
||||||
|
|||||||
Reference in New Issue
Block a user