zcbot/scripts/smoke_read_document.py

117 lines
4.2 KiB
Python

"""Smoke: read_document(豆包文档理解)端到端走通 + 扫描件 OCR + save_md 落盘验证。
跑法: .venv/Scripts/python.exe scripts/smoke_read_document.py
依赖 .env 里 ARK_API_KEY / ZCBOT_DB_URL。**会真调豆包文档理解,产生 < ¥0.01 费用**。
校验:
1. 合成 3 页扫描件 PDF(复用 probe_ark_doc 的生成器,无文本层)
2. ReadDocumentTool.execute(save_md=...) 返回 banner + saved: + 预览
3. save_md 文件落盘且三页魔术串全命中(页覆盖 + OCR 保真)
4. usage_events 多一行 kind="vision",units 含 document 路径
5. _count_pdf_pages 软探页数 = 3(页数闸门的数据源)
"""
from __future__ import annotations
import sys
import uuid
from pathlib import Path
ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(ROOT))
sys.path.insert(0, str(ROOT / "scripts"))
try:
sys.stdout.reconfigure(encoding="utf-8", errors="replace") # type: ignore[attr-defined]
except Exception:
pass
# probe_ark_doc import 时自带 .env 加载 + sys.path 处理
from probe_ark_doc import _magic, make_scanned_pdf
from sqlalchemy import text
from core.ark_client import ArkConfig
from core.storage import session_scope
from core.storage.models import Task, User
from tools.read_document import ReadDocumentTool, _count_pdf_pages
def main() -> int:
cfg = ArkConfig.load()
if cfg is None:
print("[SKIP] ARK_API_KEY 未设(或 doubao.yaml 缺失),无法测真接口")
return 0
vision_cfg = (cfg.raw.get("vision") or {})
if not vision_cfg:
print("[SKIP] doubao.yaml 无 vision 段")
return 0
variant_key, variant_cfg = next(iter(vision_cfg.items()))
print(f"[setup] variant={variant_key} model={variant_cfg.get('model_id')} "
f"max_pdf_mb={variant_cfg.get('max_pdf_mb')} "
f"max_pdf_pages={variant_cfg.get('max_pdf_pages')}")
uid = uuid.uuid4()
tid = uuid.uuid4()
ws_user = ROOT / "workspace" / "users" / str(uid)
wd = ws_user / "smoke_readdoc"
pdf = wd / "upload" / "scan.pdf"
pdf.parent.mkdir(parents=True, exist_ok=True)
make_scanned_pdf(pdf, 3)
print(f"[setup] 合成扫描件 {pdf.name}({pdf.stat().st_size} bytes, 3 页, 无文本层)")
n = _count_pdf_pages(pdf)
assert n == 3, f"_count_pdf_pages 应 3,实际 {n}"
print(f"[OK] _count_pdf_pages = {n}")
with session_scope() as s:
s.add(User(user_id=uid))
with session_scope() as s:
s.add(Task(task_id=tid, user_id=uid, name="smoke_readdoc", working_dir=str(wd)))
tool = ReadDocumentTool(
ark_cfg=cfg,
vision_variant_cfg=variant_cfg,
variant_key=variant_key,
working_dir=wd,
task_id=tid,
user_id=uid,
base_dir=wd,
user_root=ws_user,
)
print("[call] execute(document='upload/scan.pdf', save_md='source/scan.md')")
result = tool.execute(document="upload/scan.pdf", save_md="source/scan.md")
print(f"[tool result]\n{result}\n")
if result.startswith("[Error]"):
print("[FAIL] tool 返回错误")
return 2
assert "saved:" in result, "返回里缺 saved: 行"
saved = wd / "source" / "scan.md"
assert saved.is_file(), f"save_md 未落盘: {saved}"
body = saved.read_text(encoding="utf-8")
flat = body.replace(" ", "").replace("-", "")
magics = [_magic(i) for i in range(3)]
hits = [m for m in magics if m.replace("-", "") in flat]
assert len(hits) == 3, f"魔术串命中 {len(hits)}/3(漏 {set(magics) - set(hits)})"
print(f"[OK] save_md 落盘 {len(body)} 字符,3/3 魔术串命中")
with session_scope() as s:
rows = s.execute(text(
"SELECT kind, model_profile, units, cost_cny FROM usage_events "
"WHERE task_id = :tid"
), {"tid": str(tid)}).all()
assert len(rows) == 1, f"usage_events 行数应 1,实际 {len(rows)}"
row = rows[0]
assert row.kind == "vision", f"kind 应 vision,实际 {row.kind}"
assert "document" in row.units, f"units 缺 document: {row.units}"
print(f"[OK] usage_events: kind={row.kind} model={row.model_profile} "
f"cost_cny={row.cost_cny} units={row.units}")
print("\n[PASS] smoke_read_document 全部通过")
return 0
if __name__ == "__main__":
sys.exit(main())