zcbot/tools/read_document.py

244 lines
10 KiB
Python

"""read_document: 让纯文本主模型""扫描件 / 图片型 PDF。
markitdown 只能抽 PDF 的文本层,扫描件(纯图页)抽出来是空的 —— 这条死路由本 tool
补上:整份 PDF base64 内联喂豆包 seed-2.0-lite 的文档理解(方舟按页栅格化后视觉
OCR),产出 markdown 全文。格式 / 上限 / 成本均经 scripts/probe_ark_doc.py 实测:
file 内容块 + `data:application/pdf;base64,` 前缀;单页栅格化 3600 万像素上限;
输入约 1300 token/页(100 页约 ¥0.09)。
与 look_at_image 同一 model / key / 记账通道(usage_events kind="vision"),
配置在 config/media/doubao.yaml 的 vision 段(max_pdf_mb / max_pdf_pages)。
"""
from __future__ import annotations
from pathlib import Path
from typing import Any, Optional
from uuid import UUID
from core.ark_client import ArkConfig
from core.storage.usage import record_vision_usage
from .base import Tool
from .output import compact_tool_output
from .image_ref import _CONTAINER_ROOT, load_pdf_as_data_url, resolve_in_root
from .media_common import ark_chat_with_retry, extract_chat_answer, record_usage_safe
_DEFAULT_QUESTION = (
"这是一份多页 PDF 文档。请逐页把其中的文字完整 OCR 成 markdown:"
"每页以「== 第N页 ==」开头;表格转成 markdown 表格;保留标题层级与段落换行;"
"公式尽量用 LaTeX;不要总结、不要遗漏、不要自行补充原文没有的内容。"
)
# 保存到文件时,tool 返回值里带的正文预览长度(全文在文件里,预览只为让模型确认质量)
_PREVIEW_CHARS = 1500
def _count_pdf_pages(pdf: Path) -> Optional[int]:
"""软探页数(pdfminer 随 markitdown[pdf] 已在依赖里);解析不了返 None 不拦路。"""
try:
from pdfminer.pdfdocument import PDFDocument
from pdfminer.pdfpage import PDFPage
from pdfminer.pdfparser import PDFParser
with open(pdf, "rb") as f:
return sum(1 for _ in PDFPage.create_pages(PDFDocument(PDFParser(f))))
except Exception:
return None
class ReadDocumentTool(Tool):
name = "read_document"
description = (
"Read a PDF (including SCANNED/image-only PDFs that markitdown can't extract) using "
"Doubao Seed 2.0 Lite document understanding — OCRs every page into markdown. "
"Use when markitdown output for a PDF is empty/near-empty (scanned document), or to "
"ask a specific question about a PDF's content. Pass the PDF path; optionally "
"`question` (default: full per-page OCR to markdown) and `save_md` (relative path to "
"write the full text, e.g. 'source/xxx.md' — recommended for multi-page docs so the "
"full text lands in a file instead of flooding context). Costs roughly 0.001-0.002 "
"CNY per page; docs over the page limit must be split first."
)
parameters = {
"type": "object",
"properties": {
"document": {
"type": "string",
"description": (
"PDF 相对路径(task_dir 内,如 'source/xxx.pdf',或用户消息里"
"`[用户上传的文件]` 行给的路径)。"
),
},
"question": {
"type": "string",
"description": (
"想从文档里知道什么(可选)。如「第3章的检测指标是什么」。"
"不传则默认逐页完整 OCR 成 markdown。"
),
},
"save_md": {
"type": "string",
"description": (
"把全文写到这个相对路径(可选,如 'source/xxx.md')。多页 OCR 建议必传:"
"全文落文件,tool 只返回开头预览,不撑爆上下文。"
),
},
},
"required": ["document"],
}
def __init__(
self,
*,
ark_cfg: ArkConfig,
vision_variant_cfg: dict,
variant_key: str,
working_dir: Path,
task_id: UUID,
user_id: UUID,
base_dir: Optional[Path] = None,
user_root: Optional[Path] = None,
) -> None:
super().__init__(base_dir, user_root=user_root)
self.ark_cfg = ark_cfg
self.cfg = vision_variant_cfg
self.variant_key = variant_key
self.working_dir = Path(working_dir)
self.task_id = task_id
self.user_id = user_id
def execute(
self,
document: str,
question: Optional[str] = None,
save_md: Optional[str] = None,
) -> str:
if not (document or "").strip():
return "[Error] document(PDF 路径)不能为空"
cfg = self.cfg
max_bytes = int(float(cfg.get("max_pdf_mb", 30)) * 1024 * 1024)
data_url, disp, err = load_pdf_as_data_url(
document.strip(),
working_dir=self.working_dir,
user_root=self.user_root,
display_fn=self._display,
max_bytes=max_bytes,
)
if err:
return err
# 页数软闸:约 1300 token/页,超过上限会撞模型上下文窗口 → 白付一次失败调用。
# pdfminer 解析不了(加密/损坏)不拦,让 API 报错兜底。
max_pages = int(cfg.get("max_pdf_pages", 100))
resolved = resolve_in_root(document.strip(), self.working_dir, self.user_root)
n_pages = _count_pdf_pages(resolved) if resolved is not None else None
if n_pages is not None and n_pages > max_pages:
return (
f"[Error] 文档 {n_pages} 页超过单次 {max_pages} 页上限(上下文约束)。"
f"先把 PDF 拆成不超过 {max_pages} 页的分卷再逐卷调用。"
)
q = (question or "").strip() or _DEFAULT_QUESTION
model_id = cfg["model_id"]
timeout_s = float(cfg.get("doc_request_timeout_s", 600))
endpoint = cfg.get("endpoint", "/chat/completions")
body: dict[str, Any] = {
"model": model_id,
"messages": [
{
"role": "user",
"content": [
{"type": "text", "text": q},
{
"type": "file",
"file": {"filename": Path(disp).name, "file_data": data_url},
},
],
}
],
}
# 超时透明重试(细节见 media_common.ark_chat_with_retry)
resp, api_err = ark_chat_with_retry(
self.ark_cfg, endpoint, body,
timeout_s=timeout_s,
retries=int(cfg.get("timeout_retries", 1)),
tool_name="read_document",
)
if api_err:
return api_err
assert resp is not None
answer, truncated = extract_chat_answer(resp)
if not answer:
return (
"[Error] 文档理解响应缺内容(模型未返回文本)。"
"可能 PDF 损坏 / 页面超单页像素上限,稍后重试或先重存该 PDF。"
)
usage = resp.get("usage") or {}
tin = int(usage.get("prompt_tokens", 0) or 0)
tout = int(usage.get("completion_tokens", 0) or 0)
cost = record_usage_safe(
"read_document", record_vision_usage,
task_id=self.task_id,
user_id=self.user_id,
model_profile=f"doubao.{self.variant_key}",
prompt_tokens=tin,
completion_tokens=tout,
input_cny_per_mtoken=float(cfg.get("price_cny_per_mtoken_input", 0)),
output_cny_per_mtoken=float(cfg.get("price_cny_per_mtoken_output", 0)),
extra_units={"document": disp},
)
cost_cny = float(cost or 0)
banner = (
f"[read_document] model={model_id} · document={disp}"
+ (f" · pages={n_pages}" if n_pages else "")
+ f" · tokens={tin}+{tout} · cost=¥{cost_cny:.4f}"
)
trunc_note = (
"\n[注意] 输出被模型 token 上限截断,末尾页可能缺失 —— "
"用 question 按页段分次问(如「只 OCR 第 51-100 页」)补齐。"
) if truncated else ""
if save_md and save_md.strip():
saved_disp, save_err = self._save_md(save_md.strip(), answer)
if save_err:
# 保存失败不吞掉 OCR 结果:降级为直接返回(截断保护)
return f"{banner}\n{save_err}(全文改为直接返回)\n\n{compact_tool_output(answer)}{trunc_note}"
preview = answer[:_PREVIEW_CHARS]
more = f"\n...(预览截断,全文 {len(answer)} 字符在文件里)" if len(answer) > _PREVIEW_CHARS else ""
return f"{banner}\nsaved: {saved_disp}\n\n{preview}{more}{trunc_note}"
return f"{banner}\n\n{compact_tool_output(answer)}{trunc_note}"
def _save_md(self, rel: str, text: str) -> tuple[str, str]:
"""把全文写到 task 内。返回 (display_path, error)。
与读取侧 resolve_in_root 同款三形态(相对 / 宿主绝对 / 容器 `/workspace/...`),
但目标是写入 —— 不要求文件已存在,只做 user_root 边界校验。
"""
p = Path(rel)
is_container = rel == _CONTAINER_ROOT or rel.startswith(_CONTAINER_ROOT + "/")
if self.user_root is not None and is_container:
target = self.user_root / rel[len(_CONTAINER_ROOT):].lstrip("/")
elif p.is_absolute():
target = p
else:
target = self.working_dir / p
target = target.resolve()
root = (self.user_root or self.working_dir).resolve()
try:
target.relative_to(root)
except ValueError:
return "", f"[Error] save_md 越界: {rel!r}(须在 task 目录内)"
try:
target.parent.mkdir(parents=True, exist_ok=True)
target.write_text(text, encoding="utf-8")
except OSError as e:
return "", f"[Error] 写入 {rel} 失败: {type(e).__name__}: {e}"
return self._display(target), ""