- src/doc/ 拆分为 src/core/extraction/, matching/, validation/(核心业务逻辑) - src/bot/ 重命名为 src/infra/browser/(浏览器自动化基础设施) - fill_consumable_doc.py → src/infra/documents/consumable.py - 新增 Agent 调度模块:coordinator.py, events.py, session.py,重构 orchestrator.py - 更新 AGENTS.md、README.md 及所有子目录 README
209 lines
6.6 KiB
Python
209 lines
6.6 KiB
Python
#!/usr/bin/env python3
|
||
"""单独测试事前申请单的信息提取功能
|
||
|
||
用于调试 LLM 对事前申请单的提取准确率。
|
||
|
||
用法:
|
||
# 测试项目根目录的事前申请单.pdf
|
||
python scripts/test_application_extract.py
|
||
|
||
# 测试指定文件
|
||
python scripts/test_application_extract.py --file path/to/file.pdf
|
||
|
||
# 测试 scripts/data 目录下的事前申请单
|
||
python scripts/test_application_extract.py --dir scripts/data
|
||
"""
|
||
|
||
import argparse
|
||
import io
|
||
import json
|
||
import sys
|
||
from pathlib import Path
|
||
|
||
# 加载 .env 环境变量(在导入 src 模块之前)
|
||
from dotenv import load_dotenv
|
||
|
||
ROOT = Path(__file__).resolve().parent.parent
|
||
load_dotenv(ROOT / ".env")
|
||
|
||
# Windows 终端强制 UTF-8
|
||
sys.stdout = io.TextIOWrapper(sys.stdout.buffer, encoding="utf-8", errors="replace")
|
||
|
||
sys.path.insert(0, str(ROOT)) # noqa: E402
|
||
|
||
from src.core.extraction import extract_document # noqa: E402
|
||
|
||
|
||
def test_single_file(file_path: Path) -> None:
|
||
"""测试单个文件的提取效果"""
|
||
print("=" * 60)
|
||
print(f"测试文件: {file_path.name}")
|
||
print("=" * 60)
|
||
|
||
if not file_path.exists():
|
||
print(f"文件不存在: {file_path}")
|
||
return
|
||
|
||
try:
|
||
# 调用统一提取接口
|
||
result = extract_document(file_path)
|
||
|
||
if not result:
|
||
print("[ERROR] 提取返回空结果")
|
||
return
|
||
|
||
# 检查类型判断
|
||
inv_type = result.get("invoice_type", "")
|
||
print(f"\n[类型判断] invoice_type = {inv_type}")
|
||
|
||
if inv_type != "application":
|
||
print(f"[WARNING] 类型判断错误!期望 'application',实际得到 '{inv_type}'")
|
||
else:
|
||
print("[OK] 类型判断正确")
|
||
|
||
# 展示提取结果
|
||
print("\n[提取结果]")
|
||
print(json.dumps(result, ensure_ascii=False, indent=2))
|
||
|
||
# 字段完整性检查
|
||
print("\n[字段检查]")
|
||
expected_fields = {
|
||
"invoice_type": "类型标识",
|
||
"project_name": "项目名称",
|
||
"purpose": "出差事由",
|
||
"start_date": "开始日期",
|
||
"end_date": "结束日期",
|
||
"person_info": "人员信息",
|
||
}
|
||
|
||
for field, desc in expected_fields.items():
|
||
value = result.get(field)
|
||
if value is None:
|
||
print(f" [MISSING] {desc} ({field}) - 字段缺失")
|
||
elif value == "" or value == []:
|
||
print(f" [EMPTY] {desc} ({field}) - 字段为空")
|
||
else:
|
||
print(f" [OK] {desc} ({field})")
|
||
|
||
# 日期格式检查
|
||
for date_field in ["start_date", "end_date"]:
|
||
date_value = result.get(date_field, "")
|
||
if date_value and len(date_value) == 10:
|
||
try:
|
||
parts = date_value.split("-")
|
||
if len(parts) == 3:
|
||
int(parts[0]) # year
|
||
int(parts[1]) # month
|
||
int(parts[2]) # day
|
||
print(f" [OK] {date_field} 格式正确 (YYYY-MM-DD)")
|
||
except (ValueError, IndexError):
|
||
print(f" [ERROR] {date_field} 格式错误: {date_value}")
|
||
elif date_value:
|
||
print(f" [ERROR] {date_field} 格式错误: {date_value}")
|
||
|
||
# 人员信息结构检查
|
||
person_info = result.get("person_info")
|
||
if person_info:
|
||
if isinstance(person_info, list):
|
||
print(f"\n[人员信息] 共 {len(person_info)} 人")
|
||
for i, person in enumerate(person_info, 1):
|
||
pid = person.get("person_id", "")
|
||
pname = person.get("person_name", "")
|
||
print(f" 人员 {i}: {pname} ({pid})")
|
||
elif isinstance(person_info, dict):
|
||
print("\n[人员信息] 单人格式")
|
||
print(f" 姓名: {person_info.get('person_name', '')}")
|
||
print(f" 编号: {person_info.get('person_id', '')}")
|
||
|
||
except Exception as e:
|
||
print(f"\n[EXCEPTION] 提取失败: {e}")
|
||
import traceback
|
||
|
||
traceback.print_exc()
|
||
|
||
|
||
def test_cache_comparison(file_path: Path) -> None:
|
||
"""对比原始提取和缓存结果"""
|
||
cache_dir = file_path.parent / ".invoice_cache"
|
||
cache_file = cache_dir / f"{file_path.stem}{file_path.suffix}.json"
|
||
|
||
if not cache_file.exists():
|
||
print(f"\n[INFO] 无缓存文件对比: {cache_file}")
|
||
return
|
||
|
||
print("\n" + "=" * 60)
|
||
print("[缓存对比]")
|
||
print("=" * 60)
|
||
|
||
try:
|
||
with open(cache_file, encoding="utf-8") as f:
|
||
cache_data = json.load(f)
|
||
|
||
cached_result = cache_data.get("extracted_data", {})
|
||
print("\n[缓存数据]")
|
||
print(json.dumps(cached_result, ensure_ascii=False, indent=2))
|
||
|
||
# 对比关键字段
|
||
print("\n[字段对比]")
|
||
for key in ["invoice_type", "project_name", "purpose", "start_date", "end_date"]:
|
||
cached_value = cached_result.get(key, "<缺失>")
|
||
print(f" {key}: {cached_value}")
|
||
|
||
except Exception as e:
|
||
print(f"[ERROR] 读取缓存失败: {e}")
|
||
|
||
|
||
def main() -> None:
|
||
parser = argparse.ArgumentParser(description="测试事前申请单信息提取")
|
||
parser.add_argument(
|
||
"--file",
|
||
type=str,
|
||
default=None,
|
||
help="要测试的文件路径 (默认: 扫描 scripts/data 目录)",
|
||
)
|
||
parser.add_argument(
|
||
"--dir",
|
||
type=str,
|
||
default=str(ROOT / "scripts" / "data"),
|
||
help="扫描目录下所有事前申请单文件 (默认: scripts/data)",
|
||
)
|
||
parser.add_argument(
|
||
"--no-cache-compare",
|
||
action="store_true",
|
||
help="不执行缓存对比",
|
||
)
|
||
|
||
args = parser.parse_args()
|
||
|
||
if args.dir:
|
||
# 目录模式:扫描所有PDF文件
|
||
dir_path = Path(args.dir)
|
||
if not dir_path.exists():
|
||
print(f"目录不存在: {dir_path}")
|
||
sys.exit(1)
|
||
|
||
pdf_files = sorted(dir_path.glob("*.pdf"))
|
||
if not pdf_files:
|
||
print(f"目录下没有找到PDF文件: {dir_path}")
|
||
sys.exit(1)
|
||
|
||
print(f"发现 {len(pdf_files)} 个PDF文件,开始逐个测试...\n")
|
||
for pdf in pdf_files:
|
||
if "事前申请" in pdf.name or "申请单" in pdf.name:
|
||
test_single_file(pdf)
|
||
if not args.no_cache_compare:
|
||
test_cache_comparison(pdf)
|
||
print()
|
||
else:
|
||
# 单文件模式
|
||
file_path = Path(args.file)
|
||
test_single_file(file_path)
|
||
if not args.no_cache_compare:
|
||
test_cache_comparison(file_path)
|
||
|
||
print("\n测试完成!")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|