243 lines
9.9 KiB
Python
243 lines
9.9 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
|
|
|
|
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), ""
|