From 1b35f07fd7614c05417e0a26acee7a92a917937b Mon Sep 17 00:00:00 2001 From: wandering Date: Thu, 2 Jul 2026 18:36:19 +0800 Subject: [PATCH] =?UTF-8?q?refactor:=20=E6=9E=B6=E6=9E=84=E9=87=8D?= =?UTF-8?q?=E7=BB=84=20=E2=80=94=20doc/bot=20=E2=86=92=20core/infra?= =?UTF-8?q?=EF=BC=8C=E6=96=B0=E5=A2=9E=20Agent=20=E8=B0=83=E5=BA=A6?= =?UTF-8?q?=E6=A8=A1=E5=9D=97?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 --- .agents/docs/guides/项目架构全景图.md | 401 +++++++++++++ .agents/docs/plans/README.md | 20 + .agents/docs/plans/架构分析-2026-06-15.md | 225 +++++++ .coverage | Bin 53248 -> 53248 bytes AGENTS.md | 93 ++- README.md | 128 +++- config/README.md | 24 + config/validation_rules.json | 135 +++++ docs/API.md | 5 +- docs/报销操作指南.md | 2 +- pyproject.toml | 4 - scripts/README.md | 20 + scripts/debug_stream_fields.py | 2 +- scripts/test_application_extract.py | 2 +- scripts/test_multimodal.py | 4 +- scripts/test_travel_info.py | 2 +- src/README.md | 79 +-- src/agent/coordinator.py | 410 +++++++++++++ src/agent/events.py | 75 +++ src/agent/orchestrator.py | 551 ++---------------- src/agent/session.py | 101 ++++ src/bot/README.md | 33 -- src/bot/__init__.py | 97 --- src/core/README.md | 21 + src/core/__init__.py | 4 + src/core/extraction/README.md | 29 + src/core/extraction/__init__.py | 42 ++ src/{doc => core/extraction}/extractor.py | 62 +- src/{doc => core/extraction}/llm_extractor.py | 14 +- src/core/matching/README.md | 27 + src/core/matching/__init__.py | 8 + src/{doc => core/matching}/matcher.py | 2 +- src/core/validation/README.md | 34 ++ src/core/validation/__init__.py | 28 + src/core/validation/validator.py | 537 +++++++++++++++++ src/doc/README.md | 69 --- src/doc/__init__.py | 4 - src/doc/prompts/validation_system.md | 83 --- src/doc/validator.py | 430 -------------- src/infra/README.md | 21 + src/infra/__init__.py | 4 + src/infra/browser/README.md | 44 ++ src/infra/browser/__init__.py | 100 ++++ src/{bot => infra/browser}/base.py | 5 +- src/{bot => infra/browser}/normal.py | 5 +- src/{bot => infra/browser}/travel.py | 8 +- src/infra/documents/README.md | 36 ++ src/infra/documents/__init__.py | 32 + .../documents/consumable.py} | 9 +- src/{doc => infra/documents}/invoice.py | 2 +- src/{doc => infra/documents}/pdf.py | 2 +- src/infra/llm/README.md | 32 + src/infra/llm/__init__.py | 18 + src/{doc => infra/llm}/prompt.py | 10 +- src/{doc => infra/llm}/prompts/README.md | 6 +- .../llm}/prompts/invoice_system.md | 0 .../llm}/prompts/normal_info_system.md | 0 .../llm}/prompts/supplement_system.md | 0 .../llm}/prompts/travel_info_system.md | 0 src/pipeline.py | 53 +- src/pipeline_core.py | 124 +++- src/web/app.py | 9 +- src/web/pipeline_web.py | 82 +-- src/web/routes.py | 14 +- src/web/static/js/agent.js | 6 + tests/test_extractor.py | 72 +-- tests/test_invoice.py | 2 +- tests/test_llm_extractor.py | 8 +- tests/test_matcher.py | 2 +- 69 files changed, 2965 insertions(+), 1548 deletions(-) create mode 100644 .agents/docs/guides/项目架构全景图.md create mode 100644 .agents/docs/plans/README.md create mode 100644 .agents/docs/plans/架构分析-2026-06-15.md create mode 100644 config/README.md create mode 100644 config/validation_rules.json create mode 100644 scripts/README.md create mode 100644 src/agent/coordinator.py create mode 100644 src/agent/events.py create mode 100644 src/agent/session.py delete mode 100644 src/bot/README.md delete mode 100644 src/bot/__init__.py create mode 100644 src/core/README.md create mode 100644 src/core/__init__.py create mode 100644 src/core/extraction/README.md create mode 100644 src/core/extraction/__init__.py rename src/{doc => core/extraction}/extractor.py (86%) rename src/{doc => core/extraction}/llm_extractor.py (98%) create mode 100644 src/core/matching/README.md create mode 100644 src/core/matching/__init__.py rename src/{doc => core/matching}/matcher.py (99%) create mode 100644 src/core/validation/README.md create mode 100644 src/core/validation/__init__.py create mode 100644 src/core/validation/validator.py delete mode 100644 src/doc/README.md delete mode 100644 src/doc/__init__.py delete mode 100644 src/doc/prompts/validation_system.md delete mode 100644 src/doc/validator.py create mode 100644 src/infra/README.md create mode 100644 src/infra/__init__.py create mode 100644 src/infra/browser/README.md create mode 100644 src/infra/browser/__init__.py rename src/{bot => infra/browser}/base.py (98%) rename src/{bot => infra/browser}/normal.py (99%) rename src/{bot => infra/browser}/travel.py (98%) create mode 100644 src/infra/documents/README.md create mode 100644 src/infra/documents/__init__.py rename src/{doc/fill_consumable_doc.py => infra/documents/consumable.py} (97%) rename src/{doc => infra/documents}/invoice.py (99%) rename src/{doc => infra/documents}/pdf.py (97%) create mode 100644 src/infra/llm/README.md create mode 100644 src/infra/llm/__init__.py rename src/{doc => infra/llm}/prompt.py (79%) rename src/{doc => infra/llm}/prompts/README.md (74%) rename src/{doc => infra/llm}/prompts/invoice_system.md (100%) rename src/{doc => infra/llm}/prompts/normal_info_system.md (100%) rename src/{doc => infra/llm}/prompts/supplement_system.md (100%) rename src/{doc => infra/llm}/prompts/travel_info_system.md (100%) diff --git a/.agents/docs/guides/项目架构全景图.md b/.agents/docs/guides/项目架构全景图.md new file mode 100644 index 0000000..649ca57 --- /dev/null +++ b/.agents/docs/guides/项目架构全景图.md @@ -0,0 +1,401 @@ +# 项目架构全景图 + +> 最后更新: 2026-06-15 +> 用途: 理解项目整体结构、模块职责、依赖关系和数据流 + +--- + +## 一、分层架构总览 + +``` +src/ +├── agent/ Agent 调度层(协调提取-校验-修正循环,状态机管理) +├── core/ 核心业务层(纯逻辑,零框架依赖) +├── infra/ 基础设施层(浏览器、文档、LLM 提示词) +├── web/ Web 界面层(Flask + SSE) +├── pipeline.py CLI 流程编排 +├── pipeline_core.py CLI/Web 公共管道逻辑 +├── main.py CLI 入口 +├── config.py 配置加载 +└── exceptions.py 异常定义 +``` + +### 依赖方向 + +```mermaid +graph TD + classDef entry fill:#e8eaf6,stroke:#3f51b5,color:#1a237e + classDef orchestrate fill:#e0f2f1,stroke:#00897b,color:#004d40 + classDef agent fill:#fff8e1,stroke:#ff8f00,color:#3e2723 + classDef core fill:#e3f2fd,stroke:#1565c0,color:#0d47a1 + classDef infra fill:#e8f5e9,stroke:#2e7d32,color:#1b5e20 + + subgraph 入口层 + CLI["main.py"]:::entry + WEB["web/app.py"]:::entry + end + + subgraph 编排层 + PIPE["pipeline.py"]:::orchestrate + PIPE_WEB["web/pipeline_web.py"]:::orchestrate + PIPE_CORE["pipeline_core.py"]:::orchestrate + end + + subgraph Agent调度层 + AGENT["agent/orchestrator.py"]:::agent + SESSION["agent/session.py"]:::agent + EVENTS["agent/events.py"]:::agent + end + + subgraph 核心业务层 + EXTRACT["core/extraction/"]:::core + MATCH["core/matching/"]:::core + VALID["core/validation/"]:::core + end + + subgraph 基础设施层 + BROWSER["infra/browser/"]:::infra + DOCS["infra/documents/"]:::infra + LLM["infra/llm/"]:::infra + end + + CLI --> PIPE + WEB --> PIPE_WEB + PIPE --> PIPE_CORE + PIPE --> EXTRACT + PIPE --> BROWSER + PIPE_WEB --> AGENT + PIPE_WEB --> PIPE_CORE + PIPE_WEB --> EXTRACT + AGENT --> EXTRACT + AGENT --> VALID + AGENT --> LLM + AGENT --> PIPE_CORE + EXTRACT --> MATCH + EXTRACT --> DOCS + EXTRACT --> LLM + MATCH --> DOCS + BROWSER --> DOCS +``` + +**关键约束**: +- `infra` 不依赖 `core` 和 `agent`,只提供工具能力 +- `core` 零外部依赖,不依赖 Flask、Playwright 等框架 +- `agent` 依赖 `core` 和 `infra`,作为调度中枢编排各模块 +- 所有跨层调用均通过 `__init__.py` 导出的稳定接口 + +--- + +## 二、模块清单 + +### 2.1 Agent 调度层 (`src/agent/`) + +| 文件 | 职责 | +|------|------| +| `coordinator.py` | 核心协调逻辑:提取-校验-修正循环(最多 3 次重试)、用户补充处理、强制提交 | +| `session.py` | 会话状态:`AgentState` 枚举、`AgentSession` 数据类、状态持久化(原子写入) | +| `events.py` | SSE 事件发射:事件去重、事件日志追加、事件读取 | +| `orchestrator.py` | 兼容层:从子模块重新导出所有符号,保持旧导入路径可用 | + +**对外接口**:`AgentSession`, `AgentState`, `run_agent_round()`, `force_submit()`, `add_supplement()`, `process_user_text_supplement()`, `load_agent_state()`, `save_agent_state()` + +### 2.2 核心业务层 (`src/core/`) + +| 子模块 | 职责 | 对外接口 | +|------|------|------| +| `extraction/extractor.py` | 编排入口:扫描目录 → 逐文件提取 → 分类 → 金额匹配 | `extract_invoices()`, `extract_document()` | +| `extraction/llm_extractor.py` | LLM 多模态提取核心:统一文档提取、差旅/普通信息提取、缓存管理、SSE 流式事件 | `llm_query_text()`, `extract_travel_info()`, `extract_normal_info()`, `load_cache()` | +| `matching/matcher.py` | 发票与支付记录按金额匹配(一对一 / 一对多贪心,相对容差 3%) | `match_invoices_to_cards()` | +| `validation/validator.py` | 声明式规则校验引擎,规则从 JSON 配置文件加载 | `validate_extracted_info()`, `ValidationReport` | + +### 2.3 基础设施层 (`src/infra/`) + +| 子模块 | 职责 | 对外接口 | +|------|------|------| +| `browser/base.py` | `BaseBot` 基类:Playwright 浏览器生命周期、登录、导航、截图 | 内部基类 | +| `browser/travel.py` | 差旅报销填报:基本信息 → 明细 → 支付 → 补助 → 附件上传 | 内部流程 | +| `browser/normal.py` | 普通报销填报:基本信息 → 总明细 → 支付 → 附件上传 | 内部流程 | +| `browser/__init__.py` | 浏览器入口:类型路由和流程调度 | `run_bot()`, `run_bot_web()` | +| `documents/invoice.py` | 发票数据模型、CSV/JSON 读写、发票分类 | `load_csv()`, `save_csv()`, `save_invoice_csv()`, `classify_invoice_batch()` | +| `documents/pdf.py` | PDF 渲染为图片(PyMuPDF) | `render_pdf_to_images()` | +| `documents/consumable.py` | 易耗品出库单填写:CSV → Word 模板 | `fill_consumable_doc()` | +| `llm/prompt.py` | LLM 提示词加载 | `build_invoice_system_prompt()`, `build_travel_info_system_prompt()`, `build_normal_info_system_prompt()` | + +### 2.4 Web 界面层 (`src/web/`) + +| 文件/目录 | 职责 | +|------|------| +| `app.py` | Flask 应用入口,注册蓝图和模板 | +| `routes.py` | 路由定义:会话管理、文件上传、配置、SSE 日志流、Agent 交互 API | +| `pipeline_web.py` | Web 管道逻辑:发票提取 + 出库单生成 + 财务提交 | +| `sse_handler.py` | SSE 日志收集器、日志转义、文件轮询 | +| `templates/` | `index.html`(PC 端主界面)、`mobile_upload.html`(移动端上传) | +| `static/js/` | 前端逻辑(按加载顺序):`state.js` → `utils.js` → `chat.js` → `upload.js` → `config.js` → `process.js` → `sync.js` → `index.js` | + +--- + +## 三、CLI 模式数据流 + +```mermaid +graph TD + CLI_ENTRY["main.py --step all"] --> PIPE["pipeline.py run_pipeline()"] + + subgraph Step1["Step 1: 发票提取"] + PIPE --> EXT["core/extraction/extractor.py extract_invoices()"] + EXT --> DOC["逐文件提取"] + DOC --> LLM["LLM 多模态识别 (infra/llm)"] + LLM --> CLASS["分类: train/hotel/general/payment/application"] + CLASS --> MATCH["core/matching/matcher.py 金额匹配"] + MATCH --> SAVE["infra/documents/ CSV/JSON 保存"] + end + + subgraph Step2["Step 2: 信息提取"] + SAVE --> TYPE{"判断报销类型"} + TYPE -->|差旅| TRAVEL["提取差旅信息 → travel_info.json"] + TYPE -->|普通| NORMAL["提取普通发票信息 → normal_info.json"] + end + + subgraph Step3["Step 3: 浏览器填报"] + TRAVEL --> BOT["infra/browser/ 填报"] + NORMAL --> BOT + BOT -->|差旅| BOT_T["browser/travel.py"] + BOT -->|普通| BOT_N["browser/normal.py"] + end +``` + +**关键文件输出**: + +| 文件 | 来源 | 说明 | +|------|------|------| +| `payment_records.csv` | Step 1 | 支付记录级别(每笔刷卡记录一行) | +| `invoice_summary.csv` | Step 1 | 发票级别(每张发票一行) | +| `travel_applications.json` | Step 1 | 出差事前申请单 | +| `invoice_groups.json` | Step 1 | 发票分类结果 | +| `travel_info.json` | Step 2 | 差旅信息:交通/住宿明细、补贴、附件清单 | +| `normal_info.json` | Step 2 | 普通发票信息:报销说明、发票总数、总金额、附件清单 | + +--- + +## 四、Web 模式数据流 + +```mermaid +sequenceDiagram + participant F as 前端 (浏览器) + participant API as routes.py + participant PW as pipeline_web.py + participant AG as agent/orchestrator.py + participant EX as core/extraction/ + participant VA as core/validation/ + participant SSE as SSE 轮询 + + F->>API: POST /api/session → 创建 session + F->>API: POST /api/upload/:sid → 上传文件 + F->>API: POST /api/agent/process/:sid + API-->>F: {status: "started"} + F->>SSE: GET /api/logs/:sid (SSE 长连接) + + Note over API: 后台 daemon 线程启动 + + API->>PW: extract_invoices(session_dir) + PW->>EX: 发票提取 + 分类 + 匹配 + EX-->>PW: payment_records, applications, groups + + API->>AG: run_agent_round(session_dir, session) + loop 校验-修正循环 (最多 3 次) + AG->>EX: llm_query_text() 提取信息 + AG->>VA: validate_extracted_info() 规则校验 + alt 校验失败 + AG->>AG: 构建修正提示 + end + end + AG-->>API: session (READY 或 AWAITING_SUPPLEMENT) + + SSE-->>F: file_progress, llm_stream, agent_state_change, agent_ready/agent_request_supplement + + API->>API: 写入 result.json + SSE-->>F: done (携带 result) + F->>F: 关闭 SSE, 展示结果 +``` + +### Web 模式特有的 Agent 调度 + +CLI 模式中 `pipeline.py` 直接调用 `extract_invoices()` → `infra/browser/`,不经过 Agent 层。 + +Web 模式中 `routes.py` 启动后台线程,调用 `agent/orchestrator.py` 作为调度中枢: + +``` +run_agent_round() +├── 1. load_cache() — 检查缓存 +├── 2. _do_extraction_with_validation() — 提取-校验-修正循环 +│ ├── llm_query_text() — LLM 提取结构化信息 +│ ├── validate_extracted_info() — 规则校验 +│ └── 校验失败 → 构建修正提示 → 再次调用 LLM (最多 3 次) +├── 3. 判断 can_submit 字段 +│ ├── true → READY → 自动触发财务提交 +│ └── false → AWAITING_SUPPLEMENT → 等待用户补充 +├── 4. 用户补充处理 +│ ├── add_supplement() — 记录补充文件 +│ └── process_user_text_supplement() — LLM 解析文字补充 +└── 5. save_agent_state() — 持久化状态 +``` + +--- + +## 五、Agent 状态机 + +```mermaid +stateDiagram-v2 + [*] --> IDLE: 会话创建 + + IDLE --> EXTRACTING: POST /api/agent/process + IDLE --> EXTRACTING: POST /api/agent/supplement + IDLE --> EXTRACTING: POST /api/agent/user-supplement + + EXTRACTING --> READY: can_submit == true + EXTRACTING --> AWAITING_SUPPLEMENT: can_submit == false + EXTRACTING --> ERROR: 异常 / 轮次超限 + + READY --> SUBMITTING: _emit_ready_and_submit() + SUBMITTING --> DONE: 财务提交完成 + + AWAITING_SUPPLEMENT --> EXTRACTING: 用户补充文件/文字 + AWAITING_SUPPLEMENT --> READY: 用户强制提交 + + note right of EXTRACTING + LLM 提取 + validator 校验 + 最多 3 次重试 + end note +``` + +### 终态保护 + +以下状态为终态,再次触发 `run_agent_round()` 会被跳过: +- `DONE` — 提交完成 +- `SUBMITTING` — 提交中 +- `READY` — 准备提交 + +### 轮次保护 + +默认最多 5 轮(`AgentSession.max_rounds`),超限后进入 `ERROR` 状态,用户可选择强制提交。 + +--- + +## 六、SSE 事件通信机制 + +```mermaid +graph LR + subgraph 后端写入 + AGENT[agent/orchestrator.py] -->|追加写入| AE[agent_events.log] + LLM[LLM 回调] -->|追加写入| LS[llm_stream.log] + PW[pipeline_web.py] -->|追加写入| FE[file_events.log] + SH[sse_handler.py] -->|追加写入| SL[session.log] + RT[_run_agent_task] -->|finally 原子写入| RJ[result.json] + end + + subgraph SSE 轮询 (0.5s) + POLL[SSE 端点] -->|读取| AE + POLL -->|读取| LS + POLL -->|读取| FE + POLL -->|读取| SL + POLL -->|检测| RJ + end + + POLL -->|event: agent_*| FRONT[前端 agent.js] + POLL -->|event: llm_stream| FRONT + POLL -->|event: file_progress| FRONT + POLL -->|event: done| FRONT +``` + +### 信号文件生命周期 + +| 阶段 | `result.json` | `llm_stream.log` | `agent_events.log` | `file_events.log` | `session.log` | +|------|:--:|:--:|:--:|:--:|:--:| +| 会话创建 | 不存在 | 不存在 | 不存在 | 不存在 | 不存在 | +| 后台线程启动 | 已删除 | 已删除 | 已删除 | 保持 | 保持 | +| 文件提取中 | 不存在 | 不存在 | 不存在 | 持续追加 | 持续追加 | +| LLM 提取中 | 不存在 | 持续追加 | 持续追加 | 保持 | 持续追加 | +| 校验中 | 不存在 | 保持 | 持续追加 | 保持 | 持续追加 | +| 任务完成 | 已写入 | 保持 | 保持 | 保持 | 保持 | +| SSE done 事件 | 保持 | 保持 | 保持 | 保持 | 保持 | + +--- + +## 七、发票类型路由 + +```mermaid +graph TD + INPUT["上传文件 (PDF/图片)"] --> EXT["LLM 多模态识别"] + EXT --> TYPE{"invoice_type?"} + + TYPE -->|train| TRAVEL["差旅报销流程"] + TYPE -->|hotel| TRAVEL + TYPE -->|general| NORMAL["普通报销流程"] + TYPE -->|payment| MATCH["参与金额匹配"] + TYPE -->|application| APP["存储为 JSON"] + + TRAVEL --> TRAVEL_INFO["提取差旅信息
travel_info.json"] + TRAVEL_INFO --> TRAVEL_BOT["browser/travel.py
填报差旅报销单"] + + NORMAL --> NORMAL_INFO["提取普通发票信息
normal_info.json"] + NORMAL_INFO --> NORMAL_BOT["browser/normal.py
填报普通报销单"] + NORMAL_INFO --> CONSUMABLE["生成易耗品出库单
(仅普通报销)"] + + MATCH --> MERGE["合并到对应发票组"] + + style TRAVEL fill:#cfe2ff,stroke:#0d6efd + style NORMAL fill:#f8d7da,stroke:#dc3545 + style MATCH fill:#d1e7dd,stroke:#198754 + style APP fill:#fff3cd,stroke:#ffc107 +``` + +| 发票类型 | `invoice_type` | 报销流程 | 生成出库单 | +|----------|---------------|---------|:--:| +| 高铁票/火车票 | `train` | 差旅报销 | 否 | +| 酒店住宿 | `hotel` | 差旅报销 | 否 | +| 普通发票 | `general` | 普通报销 | 是 | +| 支付记录 | `payment` | 参与匹配 | 否 | +| 出差申请单 | `application` | 单独存储 | 否 | + +> 差旅发票和普通发票不支持混报,混合时系统按普通报销处理。 + +--- + +## 八、设计原则 + +| 原则 | 说明 | +|------|------| +| **Agent 是调度中枢** | 校验-修正循环由 Agent 编排,不内嵌在 `llm_extractor` 中 | +| **模块职责单一** | `llm_extractor` 只管提取,`validator` 只管校验,Agent 负责编排 | +| **core 零外部依赖** | 不依赖 Flask、Playwright 等框架 | +| **infra 不依赖业务** | 基础设施层只提供工具能力,不包含业务逻辑 | +| **缓存优先** | 信息提取优先读取 `.invoice_cache`,避免重复调用 LLM | +| **轮次保护** | 默认 5 轮上限,校验-修正循环最多重试 3 次 | +| **终态保护** | `DONE`/`SUBMITTING`/`READY` 状态下不再重复处理 | +| **容错降级** | 规则校验 3 次重试后返回最佳结果,不阻断流程 | +| **原子写入** | 状态文件先写 `.tmp` 再 `rename()`,防止读取不完整数据 | + +--- + +## 九、关键文件索引 + +| 文件 | 职责 | +|------|------| +| `src/main.py` | CLI 入口 | +| `src/web/app.py` | Web 入口 | +| `src/pipeline.py` | CLI 流程编排 | +| `src/pipeline_core.py` | CLI/Web 公共管道逻辑 | +| `src/web/pipeline_web.py` | Web 管道逻辑 + 财务提交 | +| `src/web/routes.py` | Web 路由 + 后台线程启动 | +| `src/agent/coordinator.py` | Agent 核心协调逻辑 | +| `src/agent/session.py` | 会话状态定义与持久化 | +| `src/agent/events.py` | SSE 事件发射 | +| `src/core/extraction/extractor.py` | 发票提取编排入口 | +| `src/core/extraction/llm_extractor.py` | LLM 多模态提取核心 | +| `src/core/matching/matcher.py` | 金额匹配 | +| `src/core/validation/validator.py` | 声明式规则校验 | +| `src/infra/browser/base.py` | 浏览器自动化基类 | +| `src/infra/documents/invoice.py` | 发票数据模型 | +| `src/web/sse_handler.py` | SSE 日志收集器 | +| `src/web/static/js/process.js` | 前端主提交流程 | +| `src/web/static/js/agent.js` | 前端 Agent 交互处理 | +| `config.json` | 项目配置 | diff --git a/.agents/docs/plans/README.md b/.agents/docs/plans/README.md new file mode 100644 index 0000000..7b88abd --- /dev/null +++ b/.agents/docs/plans/README.md @@ -0,0 +1,20 @@ +--- +last_reviewed: 2026-06-15 +--- + +# .agents/docs/plans — 实施方案与工作交接 + +存放项目实施方案、架构分析报告、重构计划等规划类文档。 + +## 文件 + +| 文件 | 说明 | +|------|------| +| `架构分析-2026-06-15.md` | 项目架构分析与重构建议(模块拆分、分层设计、接口契约) | + +## 用途 + +- 架构决策记录 +- 重构实施方案 +- 工作交接说明 +- 技术选型论证 diff --git a/.agents/docs/plans/架构分析-2026-06-15.md b/.agents/docs/plans/架构分析-2026-06-15.md new file mode 100644 index 0000000..24e87d6 --- /dev/null +++ b/.agents/docs/plans/架构分析-2026-06-15.md @@ -0,0 +1,225 @@ +# 项目架构分析与重构建议 + +## 一、当前架构总览 + +``` +src/ +├── main.py # CLI 入口 +├── pipeline.py # CLI 管道编排 +├── pipeline_core.py # CLI/Web 公共管道逻辑 +├── config.py # 配置加载 +├── exceptions.py # 异常定义 +│ +├── doc/ # 文档处理模块(职责过重) +│ ├── extractor.py # 发票提取编排 +│ ├── llm_extractor.py # LLM 提取核心 +│ ├── invoice.py # 发票数据模型 + CSV 工具 +│ ├── matcher.py # 发票匹配逻辑 +│ ├── validator.py # 信息校验规则 +│ ├── prompt.py # 提示词加载 +│ ├── pdf.py # PDF 渲染 +│ ├── fill_consumable_doc.py # 出库单填写 +│ └── prompts/ # LLM 提示词模板 +│ +├── agent/ # Agent 调度模块 +│ └── orchestrator.py # 校验-修正循环调度 +│ +├── bot/ # 浏览器自动化模块 +│ ├── base.py # 浏览器基类 +│ ├── travel.py # 差旅填报 +│ └── normal.py # 普通报销填报 +│ +└── web/ # Web 界面模块 + ├── app.py # Flask 应用 + ├── routes.py # 路由定义 + ├── pipeline_web.py # Web 管道逻辑(与 pipeline_core 重复) + ├── sse_handler.py # SSE 日志流处理 + └── static/templates/ # 前端资源 +``` + +--- + +## 二、问题分析 + +### 2.1 职责不清(高耦合) + +| 问题 | 位置 | 说明 | +|------|------|------| +| **doc 模块职责过重** | `src/doc/` | 同时负责:提取、匹配、校验、提示词、PDF渲染、出库单填写、CSV操作 | +| **Web 层重复逻辑** | `pipeline_web.py` vs `pipeline_core.py` | 两者的 `is_travel_invoice`、`extract_and_cache_*` 逻辑重复 | +| **提示词与校验耦合** | `validator.py` | 校验规则直接引用提示词相关函数,缺乏分层 | +| **bot 模块位置** | `src/bot/` | 浏览器自动化属于基础设施,却被放在 src 根目录而非独立模块 | + +### 2.2 逻辑混乱 + +1. **`src/doc/validator.py`** 的问题: + - 校验规则(`TRAVEL_VALIDATION_RULES`)硬编码在模块中,修改需改代码 + - `FieldRule` 和 `ArrayRule` 类与校验逻辑紧耦合 + - 数组元素字段支持简单格式和详细格式两种配置,增加了理解成本 + +2. **`src/doc/prompt.py`** 的问题: + - 简单的文件读取包装,但调用方分散 + - `build_invoice_system_prompt()` 和 `build_travel_info_system_prompt()` 分别调用,但结构相似 + +3. **`src/agent/orchestrator.py`** 的问题: + - 校验循环与提取逻辑混合在 `_do_extraction_with_validation` + - SSE 事件发射逻辑(`_emit_agent_event`)与业务逻辑混杂 + - 状态机转换逻辑分散 + +### 2.3 分层不合理 + +``` +当前分层(按目录): + main.py → pipeline.py → doc/ + bot/ + ↓ + pipeline_web.py → web/ + +建议分层(按职责): + 应用层: main.py, pipeline.py, pipeline_web.py + 业务层: agent/orchestrator.py, doc/validator.py, doc/matcher.py + 提取层: doc/extractor.py, doc/llm_extractor.py + 基础设施层: bot/, web/, doc/pdf.py, doc/fill_consumable_doc.py +``` + +--- + +## 三、重构建议 + +### 3.1 目录重组 + +``` +src/ +├── main.py # CLI 入口 +├── config.py # 配置加载 +├── exceptions.py # 异常定义 +│ +├── apps/ # 应用层(管道编排) +│ ├── cli/ # CLI 应用 +│ │ └── pipeline.py +│ └── web/ # Web 应用 +│ ├── app.py +│ ├── routes.py +│ ├── pipeline.py # Web 专用管道 +│ └── sse.py +│ +├── core/ # 核心业务逻辑 +│ ├── agent/ # Agent 调度 +│ │ ├── orchestrator.py +│ │ └── session.py +│ ├── validation/ # 校验模块 +│ │ ├── validator.py +│ │ └── rules/ # 校验规则(可配置化) +│ ├── matching/ # 匹配模块 +│ │ └── matcher.py +│ └── extraction/ # 提取模块 +│ ├── extractor.py +│ └── llm.py +│ +├── infra/ # 基础设施层 +│ ├── browser/ # 浏览器自动化 +│ │ ├── base.py +│ │ ├── travel.py +│ │ └── normal.py +│ ├── documents/ # 文档处理 +│ │ ├── invoice.py +│ │ ├── pdf.py +│ │ └── consumable.py +│ └── llm/ # LLM 接口 +│ └── prompts/ # 提示词模板 +│ +└── shared/ # 共享工具 + ├── logging.py + └── cache.py +``` + +### 3.2 关键重构点 + +#### 3.2.1 doc 模块拆分 + +| 职责 | 建议移动位置 | +|------|-------------| +| `validator.py` | `core/validation/` | +| `matcher.py` | `core/matching/` | +| `llm_extractor.py` | `core/extraction/` | +| `extractor.py` | `core/extraction/` | +| `invoice.py` | `infra/documents/` | +| `pdf.py` | `infra/documents/` | +| `fill_consumable_doc.py` | `infra/documents/` | +| `prompt.py` + `prompts/` | `infra/llm/` | + +#### 3.2.2 消除重复逻辑 + +**问题**: `pipeline_web.py` 和 `pipeline_core.py` 都有相似逻辑: +- `is_travel_invoice()` +- `extract_and_cache_travel_info()` +- `extract_and_cache_normal_info()` + +**建议**: 将这些公共逻辑统一到 `core/pipeline/` 目录,两个入口调用同一模块。 + +#### 3.2.3 Validator 重构 + +**当前问题**: +- 校验规则硬编码 +- `FieldRule` 和 `ArrayRule` 类过于复杂 + +**建议**: +- 将校验规则外部化为 JSON/YAML 配置文件 +- 简化 `FieldRule` 为单一数据结构 +- 统一顶层字段和数组元素字段的校验方式 + +#### 3.2.4 Agent 拆分 + +**当前问题**: +- `orchestrator.py` 包含:状态机、SSE 事件、校验循环、提取逻辑 + +**建议**: +``` +agent/ +├── session.py # 状态机定义 + 会话数据模型 +├── coordinator.py # 校验-修正循环 +├── events.py # SSE 事件发射 +└── orchestrator.py # 总调度入口 +``` + +### 3.3 接口契约强化 + +| 模块 | 依赖关系 | 接口契约 | +|------|----------|----------| +| `core/extraction` | 被 `apps/*` 调用 | 返回 `(payment_records, applications, groups)` | +| `core/validation` | 被 `agent/*` 调用 | `validate(info, rules) -> ValidationReport` | +| `core/matching` | 被 `extraction` 调用 | `match(invoices, cards) -> List[Dict]` | +| `infra/browser` | 被 `apps/*` 调用 | `run(bot, info) -> None` | +| `infra/llm` | 被 `core/extraction` 调用 | `extract_document(file) -> dict` | + +--- + +## 四、优先重构顺序 + +### 第一阶段(降低耦合) +1. 将 `doc/` 拆分为 `core/` + `infra/` +2. 消除 `pipeline_web.py` 和 `pipeline_core.py` 的重复逻辑 +3. 将 `bot/` 移动到 `infra/browser/` + +### 第二阶段(职责清晰化) +4. 拆分 `agent/orchestrator.py` 为多个模块 +5. 外部化 `validator.py` 的校验规则为配置文件 +6. 统一 SSE 事件处理接口 + +### 第三阶段(可维护性) +7. 完善 `__init__.py` 的接口导出 +8. 添加模块间依赖注入机制 +9. 建立跨模块调用规范 + +--- + +## 五、当前项目优点 + +1. **日志规范**: 统一的 `get_logger()` 方式,全局日志管理 +2. **异常体系**: 清晰的 `ReimbursementError` 异常层次 +3. **SSE 事件协议**: 良好的实时反馈机制 +4. **缓存设计**: `llm_extractor.py` 的缓存加载逻辑完善 +5. **声明式校验**: `validator.py` 的规则配置思路正确 + +--- + +*生成时间: 2026-06-15* diff --git a/.coverage b/.coverage index e61afb41c016c32f56c6d9a81b3ae28594e14db4..bd92bc28c8da951fa49c52f346339f0d853b4e73 100644 GIT binary patch delta 100 zcmZozz}&Eac>`O6*f|FNPyF}zukkPBujSX{`^@)}?=IiD&4L1(_$FWMyKHLA!otXz z!OFyN;Jd0PgM&H?0}}%a0|?ZBNd~6c{ZF1+GBil8fBpHu_4*&d&uW=AU+dR$005Q! BB5nWx delta 97 zcmZozz}&Eac>`O6*hL2ZPyF}zukkP8Z{RoN`^NW%?*ZS%&4L13`6i$1yKH2@!otXz x#LC2Qpw^R#!Ag-qfI)!)1RgMhDS;QiRrSv^GMuT1{r!u-{?s1-&DZ+18~`U2AVL5D diff --git a/AGENTS.md b/AGENTS.md index 2773b6b..e0e4234 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -4,26 +4,93 @@ alwaysApply: true --- --- -last_reviewed: 2026-06-09 +last_reviewed: 2026-07-02 --- -# AGENTS 索引 +# AGENTS — 项目操作指南 -本文件是规则的入口。详细策略文本位于 `.agents/docs/standards/*.md`。 +本文件为 Agent 提供高信号量的项目操作知识,避免重复探索。 ## 文档边界 -* 一定不要用**表情文字**输出任何内容,禁止!!!!! -* `docs/` 目录专门存放面向开源用户、外部贡献者的项目公开文档及说明文件。 -* 维护规范、实施方案、经验总结、拉取请求佐证材料与各类内部记录资料,均统一放置在 `.agents/` 目录下,避免内部自动化流程相关内容混入公开文档目录。 -* 每个文件夹下都有一个 `README.md` 文件用来交代这个文件夹的作用以及重要的信息。 +* **禁止使用表情文字**输出任何内容。 +* `docs/` 目录存放面向开源用户、外部贡献者的公开文档。 +* `.agents/` 目录存放维护规范、实施方案、经验总结等内部资料。 +* 每个文件夹下都有 `README.md` 说明该文件夹的作用和重要信息。 -## 标准目录 +## 开发命令(必须使用 uv) -* 标准文档元数据:`.agents/docs/standards/README.md` -* 调试规范:`.agents/docs/standards/调试规范.md` -* 复利式工程实践:`.agents/docs/standards/复利式工程实践.md` +项目使用 `uv` 管理依赖,所有包版本锁定在 `uv.lock` 中。 -## 项目的架构思想 +| 操作 | Makefile (跨平台) | tasks.py (Windows) | +|------|-------------------|---------------------| +| 安装依赖 + pre-commit | `make install` | `python tasks.py install` | +| 代码检查(lint+format+typecheck+deptry) | `make check` | `python tasks.py check` | +| 运行测试(含覆盖率报告) | `make test` | `python tasks.py test` | +| 运行 CLI 全流程 | `make run` | `python tasks.py run` | +| 清理缓存和虚拟环境 | `make clean` | `python tasks.py clean` | -* Agent 是负责调度的中枢,负责调度各个模块 +**注意:** `tasks.py` 中的 `check` 命令使用 `&&` 连接,Windows PowerShell 不支持 `&&`,但 `tasks.py` 内部已处理为单行字符串。 + +## 代码质量工具链(执行顺序) + +1. **Ruff lint** — `uv run ruff check .` (select: E, F, W, I, N, UP, B; ignore: E501) +2. **Ruff format** — `uv run ruff format --check .` (line-length: 120) +3. **MyPy strict mode** — `uv run mypy src/main.py` (strict=true, warn_return_any, ignore_missing_imports) +4. **deptry** — `uv run deptry .` (检测未声明、未使用、过时依赖) + +### pre-commit 钩子(仅 Ruff) + +`.pre-commit-config.yaml` 配置了两个 hook: +- `ruff --fix` — lint 并自动修复 +- `ruff-format` — 格式化 + +**注意:** MyPy 和 deptry **不在** pre-commit 中,需要手动运行 `make check`。 + +## 项目架构(Agent 调度模式) + +核心入口:`src/agent/orchestrator.py` — Agent 是负责调度的中枢,协调以下模块: +- `extraction/extractor.py` — 文件扫描 → LLM 多模态提取 → JSON 缓存 +- `matching/matcher.py` — 支付记录与发票金额匹配 +- `validation/validator.py` — 声明式校验器(规则配置与引擎分离) +- `infra/browser/travel.py` / `normal.py` — 浏览器自动化填报 + +### 数据流关键产物 + +| 文件 | 生成阶段 | 作用 | +|------|---------|------| +| `.invoice_cache/*.json` | extractor 提取 | 单张发票/支付记录的结构化数据 | +| `match_result.json` | matcher 匹配 | 支付截图与发票的关联关系 | +| `travel_info.json` / `normal_info.json` | LLM 综合提取 | 差旅/普通报销所需的全部结构化数据 | +| `invoice_summary.csv` | extractor 提取 | 普通发票汇总(用于生成易耗品出库单) | + +### 缓存机制 + +CLI 模式:`scripts/data/.invoice_cache/` +Web 模式:`src/web/uploads//.invoice_cache/` + +缓存文件与源文件同名(如 `发票1.pdf` → `.invoice_cache/发票1.json`),后续步骤均从缓存读取。删除缓存后下次处理会重新提取。 + +## 重要约束 + +* **Windows-only**:易耗品出库单填写依赖 Microsoft Word + COM (`pywin32`),仅 Windows 可用 +* **浏览器自动化**:使用 Playwright,填报时会打开 Chromium,请勿手动干扰 +* **敏感信息**:`scripts/config.json` 含登录凭据,勿提交到公开仓库 +* **发票类型区分**:差旅发票(高铁票/酒店住宿)不生成易耗品出库单,走差旅报销流程;普通发票生成出库单 + +## Web 服务 + +```bash +uv run python src/web/app.py +# 访问 http://localhost:5000 +``` + +Web 端浏览器填报以无头模式运行。会话产物存放在 `src/web/uploads//`,每次上传生成独立会话。 + +## 测试 + +```bash +make test # pytest + coverage report (term-missing) +``` + +测试目录:`tests/`,配置在 `pyproject.toml` 中 (`testpaths = ["tests"]`, `pythonpath = ["."]`)。 diff --git a/README.md b/README.md index 53b79a4..c6d64e9 100644 --- a/README.md +++ b/README.md @@ -18,20 +18,40 @@ ├── src/ │ ├── __init__.py # 包初始化 / 日志器 │ ├── config.py # 配置加载 -│ ├── bot.py # 浏览器自动填报 +│ ├── exceptions.py # 异常定义 │ ├── pipeline.py # CLI 流程编排 +│ ├── pipeline_core.py # CLI/Web 公共管道逻辑 │ ├── main.py # CLI 入口 -│ ├── doc/ # 文档处理模块 -│ │ ├── extractor.py # 编排入口:串联 PDF 读取 → LLM 提取 → 分类 -│ │ ├── pdf.py # PDF 图片渲染(PyMuPDF,供多模态 LLM 使用) -│ │ ├── llm_extractor.py # LLM 信息提取 -│ │ ├── matcher.py # 数据匹配与校验 -│ │ ├── invoice.py # 发票类型常量、分类逻辑、CSV 读写工具 -│ │ ├── fill_consumable_doc.py # 将 CSV 填入易耗品出库单(Word COM) -│ │ ├── prompt.py # LLM 提示词模板 -│ │ └── prompts/ # 提示词模板文件 -│ └── web/ -│ ├── app.py # Web 服务入口 +│ ├── agent/ # Agent 调度模块 +│ │ ├── orchestrator.py # 总调度入口 +│ │ ├── coordinator.py # 校验-修正循环 +│ │ ├── session.py # 状态机与会话数据 +│ │ └── events.py # SSE 事件发射 +│ ├── core/ # 核心业务逻辑 +│ │ ├── extraction/ # 信息提取 +│ │ │ ├── extractor.py # 编排入口:串联文件扫描 → 提取 → 分类 +│ │ │ └── llm_extractor.py # LLM 多模态信息提取 +│ │ ├── matching/ # 金额匹配 +│ │ │ └── matcher.py # 支付记录与发票关联 +│ │ └── validation/ # 校验模块 +│ │ └── validator.py # 声明式校验器 +│ ├── infra/ # 基础设施层 +│ │ ├── browser/ # 浏览器自动化 +│ │ │ ├── base.py # BaseBot 基类 +│ │ │ ├── travel.py # 差旅报销填报流程 +│ │ │ └── normal.py # 普通报销填报流程 +│ │ ├── documents/ # 文档处理 +│ │ │ ├── invoice.py # 发票数据模型 + CSV 工具 +│ │ │ ├── pdf.py # PDF 图片渲染 +│ │ │ └── consumable.py # 易耗品出库单填写(Word COM) +│ │ └── llm/ # LLM 接口 +│ │ ├── prompt.py # 提示词加载 +│ │ └── prompts/ # 提示词模板文件 +│ └── web/ # Web 界面模块 +│ ├── app.py # Flask 应用入口 +│ ├── routes.py # 路由定义 +│ ├── pipeline_web.py # Web 管道逻辑 +│ ├── sse_handler.py # SSE 日志流处理 │ ├── templates/ │ │ ├── index.html # PC 端主页 │ │ └── mobile_upload.html # 移动端扫码上传 @@ -48,6 +68,66 @@ └── *.pdf / *.jpg / *.png # 发票 PDF 或图片(CLI 模式,放在 scripts/data/) ``` +## 声明式校验器 + +`src/core/validation/validator.py` 采用**规则配置与校验引擎分离**的设计模式,支持声明式定义校验规则: + +### 设计特点 + +| 特性 | 说明 | +|------|------| +| **声明式配置** | 校验规则以数据结构形式定义,无需编写代码 | +| **统一路径定位** | 使用 `path` 统一定位字段,如 `["basic_info", "travel_purpose"]` | +| **自定义校验函数** | 支持为字段定义自定义校验逻辑(日期格式、正数检查等) | +| **数组元素校验** | 支持校验数组字段的最小元素数量及每个元素的必填字段 | +| **向后兼容** | 支持简单格式 `["field1", "field2"]` 和详细格式 `{"path": [...], "custom_check": ...}` | + +### 规则配置示例 + +```python +# 差旅报销校验规则 +TRAVEL_VALIDATION_RULES = { + "fields": [ + {"path": ["basic_info", "travel_purpose"], "description": "出差事由"}, + {"path": ["basic_info", "start_date"], "custom_check": _is_valid_date}, + ], + "arrays": [ + { + "path": ["payment_methods"], + "min_items": 1, # 至少1条支付记录 + "element_fields": [ + {"path": ["card_date"], "description": "刷卡日期"}, + {"path": ["card_amount"], "custom_check": _is_positive_number}, + ], + }, + ], +} +``` + +### 校验规则类型 + +| 规则类型 | 用途 | 关键字段 | +|----------|------|----------| +| `fields` | 顶层单值字段校验 | `path`, `custom_check`, `check_empty` | +| `arrays` | 数组字段校验 | `path`, `min_items`, `element_fields` | + +### 内置校验函数 + +- `_is_valid_date(value)` — 检查日期格式是否为 `YYYY-MM-DD` +- `_is_positive_number(value)` — 检查值是否为正数 + +### 扩展自定义校验 + +```python +# 定义自定义校验函数 +def check_vehicle_type(value): + valid_types = ["飞机", "火车", "汽车", "打车"] + return isinstance(value, str) and value.strip() in valid_types + +# 在规则中使用 +{"path": ["vehicle_type"], "custom_check": check_vehicle_type} +``` + ## 数据流 ```mermaid @@ -73,21 +153,21 @@ flowchart TB MatchResult --> NormalLLM NormalLLM --> NormalInfo[(normal_info.json)] - TravelInfo -->|差旅基本信息| Bot_T[bot/travel.py
差旅填报流程] + TravelInfo -->|差旅基本信息| Bot_T[infra/browser/travel.py
差旅填报流程] TravelInfo -->|报销明细| Bot_T TravelInfo -->|支付方式| Bot_T TravelInfo -->|补助清单| Bot_T TravelInfo -->|附件清单| Bot_T Bot_T --> Submit_T[差旅报销提交] - NormalInfo -->|报销说明| Bot_G[bot/normal.py
普通填报流程] + NormalInfo -->|报销说明| Bot_G[infra/browser/normal.py
普通填报流程] NormalInfo -->|发票总数/金额| Bot_G NormalInfo -->|支付方式| Bot_G NormalInfo -->|附件清单| Bot_G Bot_G --> Submit_G[普通报销提交] General --> CSV[(invoice_summary.csv)] - CSV --> Fill[fill_consumable_doc] + CSV --> Fill[consumable.py] Fill --> Doc[易耗品、出库单.doc] ``` @@ -103,14 +183,14 @@ flowchart TB ### bot 模块架构 -`bot/` 包负责浏览器自动化填报,仅接收已提取的信息并执行填报操作,不承担信息提取职责: +`infra/browser/` 包负责浏览器自动化填报,仅接收已提取的信息并执行填报操作,不承担信息提取职责: | 模块 | 职责 | |------|------| -| `bot/base.py` | `BaseBot` 基类:浏览器生命周期、登录、导航、截图 | -| `bot/travel.py` | 差旅填报流程:基本信息 → 差旅明细 → 支付方式 → 补助清单 → 附件上传 | -| `bot/normal.py` | 普通填报流程:基本信息 → 总明细 → 支付方式 → 附件上传 | -| `bot/__init__.py` | 入口函数:`run_bot()` / `run_bot_web()`,负责类型判断和流程路由 | +| `infra/browser/base.py` | `BaseBot` 基类:浏览器生命周期、登录、导航、截图 | +| `infra/browser/travel.py` | 差旅填报流程:基本信息 → 差旅明细 → 支付方式 → 补助清单 → 附件上传 | +| `infra/browser/normal.py` | 普通填报流程:基本信息 → 总明细 → 支付方式 → 附件上传 | +| `infra/browser/__init__.py` | 入口函数:`run_bot()` / `run_bot_web()`,负责类型判断和流程路由 | ## 环境要求 @@ -200,10 +280,10 @@ uv run python src/main.py -u 工号 -p 密码 需已生成 `invoice_summary.csv`,且本机已安装 **Microsoft Word**: ```bash -uv run python -m src.doc.fill_consumable_doc -uv run python -m src.doc.fill_consumable_doc --csv invoice_summary.csv --doc "易耗品、出库单.doc" -uv run python -m src.doc.fill_consumable_doc --config scripts/config.json # 指定配置文件 -uv run python -m src.doc.fill_consumable_doc --no-backup # 不生成 .doc.bak 备份 +uv run python -m src.infra.documents.consumable +uv run python -m src.infra.documents.consumable --csv invoice_summary.csv --doc "易耗品、出库单.doc" +uv run python -m src.infra.documents.consumable --config scripts/config.json # 指定配置文件 +uv run python -m src.infra.documents.consumable --no-backup # 不生成 .doc.bak 备份 ``` 填写规则概要: diff --git a/config/README.md b/config/README.md new file mode 100644 index 0000000..bd2864a --- /dev/null +++ b/config/README.md @@ -0,0 +1,24 @@ +--- +last_reviewed: 2026-06-15 +--- + +# config — 配置文件目录 + +## 文件 + +| 文件 | 说明 | +|------|------| +| `validation_rules.json` | 声明式校验规则配置:定义差旅和普通报销的必填字段、数组元素校验规则和自定义校验函数 | + +## validation_rules.json 结构 + +```json +{ + "version": "1.0", + "custom_checks": { ... }, + "travel": { "fields": [...], "arrays": [...] }, + "normal": { "fields": [...], "arrays": [...] } +} +``` + +校验引擎 `src/core/validation/validator.py` 在启动时读取此文件,若文件不存在则使用内置默认规则。 diff --git a/config/validation_rules.json b/config/validation_rules.json new file mode 100644 index 0000000..b801364 --- /dev/null +++ b/config/validation_rules.json @@ -0,0 +1,135 @@ +{ + "version": "1.0", + "custom_checks": { + "is_valid_date": "检查日期格式是否为 YYYY-MM-DD", + "is_positive_number": "检查是否为正数(整数或浮点数)", + "is_positive_integer": "检查是否为正整数" + }, + "travel": { + "description": "差旅报销校验规则", + "fields": [ + { + "path": ["basic_info", "travel_purpose"], + "required": true, + "check_empty": true, + "description": "出差事由" + }, + { + "path": ["basic_info", "travel_location"], + "required": true, + "check_empty": true, + "description": "出差地点" + }, + { + "path": ["basic_info", "start_date"], + "required": true, + "check_empty": true, + "custom_check": "is_valid_date", + "description": "出差开始日期" + }, + { + "path": ["basic_info", "end_date"], + "required": true, + "check_empty": true, + "custom_check": "is_valid_date", + "description": "出差结束日期" + } + ], + "arrays": [ + { + "path": ["reimbursement_details", "transport_fee"], + "min_items": 1, + "description": "交通费用明细", + "element_fields": [ + {"path": ["vehicle_type"], "required": true, "check_empty": true, "description": "交通工具类型"}, + {"path": ["start_date"], "required": true, "check_empty": true, "custom_check": "is_valid_date", "description": "出发日期"}, + {"path": ["end_date"], "required": true, "check_empty": true, "custom_check": "is_valid_date", "description": "到达日期"}, + {"path": ["departure_place"], "required": true, "check_empty": true, "description": "出发地"}, + {"path": ["arrival_place"], "required": true, "check_empty": true, "description": "目的地"}, + {"path": ["amount"], "required": true, "check_empty": true, "custom_check": "is_positive_number", "description": "金额"}, + {"path": ["bill_count"], "required": true, "check_empty": true, "custom_check": "is_positive_integer", "description": "票据张数"}, + {"path": ["remark"], "required": true, "check_empty": false, "description": "备注说明"} + ] + }, + { + "path": ["payment_methods"], + "min_items": 1, + "description": "支付方式记录", + "element_fields": [ + {"path": ["card_date"], "required": true, "check_empty": true, "custom_check": "is_valid_date", "description": "刷卡日期"}, + {"path": ["card_amount"], "required": true, "check_empty": true, "custom_check": "is_positive_number", "description": "支付金额"}, + {"path": ["merchant"], "required": true, "check_empty": true, "description": "商户名称"}, + {"path": ["remark"], "required": true, "check_empty": false, "description": "备注"} + ] + }, + { + "path": ["subsidy_list"], + "min_items": 1, + "description": "补助清单", + "element_fields": [ + {"path": ["person_id"], "required": true, "check_empty": true, "description": "人员工号"}, + {"path": ["person_name"], "required": true, "check_empty": true, "description": "人员姓名"}, + {"path": ["start_date"], "required": true, "check_empty": true, "custom_check": "is_valid_date", "description": "补助开始日期"}, + {"path": ["end_date"], "required": true, "check_empty": true, "custom_check": "is_valid_date", "description": "补助结束日期"}, + {"path": ["days"], "required": true, "check_empty": true, "custom_check": "is_positive_integer", "description": "补助天数"} + ] + }, + { + "path": ["attachments"], + "min_items": 0, + "description": "附件列表", + "element_fields": [ + {"path": ["filename"], "required": true, "check_empty": true, "description": "文件名"}, + {"path": ["attachment_type"], "required": true, "check_empty": true, "description": "附件类型"} + ] + } + ] + }, + "normal": { + "description": "普通报销校验规则", + "fields": [ + { + "path": ["basic_info", "reimbursement_description"], + "required": true, + "check_empty": true, + "description": "报销事由" + }, + { + "path": ["reimbursement_details", "total_invoices"], + "required": true, + "check_empty": true, + "custom_check": "is_positive_integer", + "description": "发票总数" + }, + { + "path": ["reimbursement_details", "total_amount"], + "required": true, + "check_empty": true, + "custom_check": "is_positive_number", + "description": "总金额" + } + ], + "arrays": [ + { + "path": ["payment_methods"], + "min_items": 1, + "description": "支付方式记录", + "element_fields": [ + {"path": ["card_date"], "required": true, "check_empty": true, "custom_check": "is_valid_date", "description": "刷卡日期"}, + {"path": ["card_amount"], "required": true, "check_empty": true, "custom_check": "is_positive_number", "description": "支付金额"}, + {"path": ["merchant"], "required": true, "check_empty": true, "description": "商户名称"}, + {"path": ["remark"], "required": true, "check_empty": false, "description": "备注"} + ] + }, + { + "path": ["attachments"], + "min_items": 0, + "description": "附件列表", + "element_fields": [ + {"path": ["filename"], "required": true, "check_empty": true, "description": "文件名"}, + {"path": ["attachment_type"], "required": true, "check_empty": true, "description": "附件类型"} + ] + } + ] + } +} \ No newline at end of file diff --git a/docs/API.md b/docs/API.md index 43de14d..0e8616d 100644 --- a/docs/API.md +++ b/docs/API.md @@ -371,7 +371,7 @@ Accept: text/event-stream 检测到 `result.json` 存在时,读取后发送 `done` 事件并断开连接。 -**SSE 超时:** 600 秒。 +**SSE 超时:** 900 秒。 --- @@ -696,6 +696,7 @@ POST /api/agent/force-submit/ | `agent_request_supplement` | `{type, round, missing_fields, missing_materials, semantic_issues, suggestion}` | 校验未通过 | | `agent_supplement_received` | `{type, files}` | 收到用户补充 | | `agent_force_submit` | `{type, message}` | 用户强制提交 | +| `agent_extract_status` | `{type, state, round, attempt, message}` | 校验-修正循环中的每次尝试结果 | | `agent_error` | `{type, message}` | 提取失败 | | `agent_max_rounds` | `{type, message}` | 达到最大轮次 | @@ -849,7 +850,7 @@ sequenceDiagram 不经过 Web、在本地直接填写出库单: ```bash -uv run python -m src.doc.fill_consumable_doc --csv invoice_summary.csv --doc "易耗品、出库单.doc" +uv run python -m src.infra.documents.consumable --csv invoice_summary.csv --doc "易耗品、出库单.doc" ``` 详见 [README.md](./README.md)。 \ No newline at end of file diff --git a/docs/报销操作指南.md b/docs/报销操作指南.md index 68eb79f..78fb4ed 100644 --- a/docs/报销操作指南.md +++ b/docs/报销操作指南.md @@ -427,7 +427,7 @@ PC 端生成二维码指向移动端上传页面,手机端上传的图片通 ### 单独填写出库单 ```bash -uv run python -m src.doc.fill_consumable_doc --csv invoice_summary.csv --doc "易耗品、出库单.doc" +uv run python -m src.infra.documents.consumable --csv invoice_summary.csv --doc "易耗品、出库单.doc" ``` ### 分步执行管道 diff --git a/pyproject.toml b/pyproject.toml index 1ea66b0..7d12c8b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -38,10 +38,6 @@ warn_return_any = true warn_unused_configs = true ignore_missing_imports = true -[[tool.mypy.overrides]] -module = "tests.*" -ignore_errors = true - [tool.deptry] ignore_notebooks = true diff --git a/scripts/README.md b/scripts/README.md new file mode 100644 index 0000000..4a0c8c1 --- /dev/null +++ b/scripts/README.md @@ -0,0 +1,20 @@ +--- +last_reviewed: 2026-06-15 +--- + +# scripts — 调试脚本与数据目录 + +## 子目录 + +| 目录 | 说明 | +|------|------| +| `data/` | CLI 模式的数据目录:发票源文件、`config.json`、`.invoice_cache` 缓存 | + +## 脚本 + +| 文件 | 说明 | +|------|------| +| `debug_stream_fields.py` | 诊断 stream_chat 返回对象的字段结构 | +| `test_application_extract.py` | 测试出差申请单提取 | +| `test_multimodal.py` | 测试多模态 LLM 识别 | +| `test_travel_info.py` | 测试差旅信息提取 | diff --git a/scripts/debug_stream_fields.py b/scripts/debug_stream_fields.py index 1fb229d..20ef8e9 100644 --- a/scripts/debug_stream_fields.py +++ b/scripts/debug_stream_fields.py @@ -10,7 +10,7 @@ from dotenv import load_dotenv # noqa: E402 from llama_index.core.llms import ChatMessage # noqa: E402 from src.config import get_llm_config # noqa: E402 -from src.doc.llm_extractor import _create_llm # noqa: E402 +from src.core.extraction import _create_llm # noqa: E402 load_dotenv(Path(__file__).parent / ".env") diff --git a/scripts/test_application_extract.py b/scripts/test_application_extract.py index a4f2000..e4b5ae0 100644 --- a/scripts/test_application_extract.py +++ b/scripts/test_application_extract.py @@ -31,7 +31,7 @@ sys.stdout = io.TextIOWrapper(sys.stdout.buffer, encoding="utf-8", errors="repla sys.path.insert(0, str(ROOT)) # noqa: E402 -from src.doc.llm_extractor import extract_document # noqa: E402 +from src.core.extraction import extract_document # noqa: E402 def test_single_file(file_path: Path) -> None: diff --git a/scripts/test_multimodal.py b/scripts/test_multimodal.py index a02bcea..34a5a68 100644 --- a/scripts/test_multimodal.py +++ b/scripts/test_multimodal.py @@ -15,8 +15,8 @@ sys.stdout = io.TextIOWrapper(sys.stdout.buffer, encoding="utf-8", errors="repla ROOT = Path(__file__).resolve().parent.parent sys.path.insert(0, str(ROOT)) # noqa: E402 -from src.doc.llm_extractor import extract_document # noqa: E402 -from src.doc.pdf import render_pdf_to_images # noqa: E402 +from src.core.extraction import extract_document # noqa: E402 +from src.infra.documents.pdf import render_pdf_to_images # noqa: E402 def test_render() -> None: diff --git a/scripts/test_travel_info.py b/scripts/test_travel_info.py index 7af2c5d..cf1049c 100644 --- a/scripts/test_travel_info.py +++ b/scripts/test_travel_info.py @@ -15,7 +15,7 @@ from pathlib import Path ROOT = Path(__file__).resolve().parent.parent sys.path.insert(0, str(ROOT)) # noqa: E402 -from src.doc.llm_extractor import extract_travel_info # noqa: E402 +from src.core.extraction import extract_travel_info # noqa: E402 def main() -> None: diff --git a/src/README.md b/src/README.md index 3c59e0a..8167751 100644 --- a/src/README.md +++ b/src/README.md @@ -1,74 +1,77 @@ --- -last_reviewed: 2026-06-12 +last_reviewed: 2026-06-15 --- # src — 主源码目录 -包含财务报销自动化系统的核心模块。 +包含财务报销自动化系统的全部源码模块。 + +## 架构分层 + +``` +src/ +├── agent/ Agent 调度层(协调提取-校验-修正循环) +├── core/ 核心业务层(提取、匹配、校验) +├── infra/ 基础设施层(浏览器、文档、LLM 提示词) +├── web/ Web 界面层(Flask + SSE) +├── pipeline.py CLI 流程编排 +├── pipeline_core.py CLI/Web 公共管道逻辑 +├── main.py CLI 入口 +├── config.py 配置加载 +└── exceptions.py 异常定义 +``` ## 模块清单 | 文件/目录 | 说明 | |-----------|------| -| `__init__.py` | 包初始化:提供 `get_logger()` 日志工厂(支持终端 + 文件双输出,按日期自动分文件) | -| `config.py` | 配置加载:从 `config.json` 读取用户凭据和默认值,从环境变量读取服务端配置(SSO 地址、LLM 参数) | -| `pipeline.py` | 流程编排:串联发票提取 → 类型判断 → 差旅/普通信息提取 → 浏览器填报,支持分步执行 | -| `main.py` | CLI 入口:支持 `--step` 分步执行、`-u/-p` 覆盖凭据、`--cache-dir` 指定缓存目录 | -| `bot/` | 浏览器自动化:Playwright 驱动的财务系统填报机器人(仅负责接收信息并填报) | -| `doc/` | 文档处理模块:PDF 渲染、LLM 提取、支付匹配、发票分类、出库单生成 | -| `web/` | Web 界面模块:Flask 应用、SSE 日志、可编辑表格、移动端上传、会话隔离 | +| `agent/` | Agent 调度:校验-修正循环、状态机管理、SSE 事件发射 | +| `core/` | 核心业务逻辑:信息提取、金额匹配、信息校验 | +| `infra/` | 基础设施:浏览器自动填报、文档处理、LLM 提示词管理 | +| `web/` | Web 界面:Flask 应用、SSE 日志流、可编辑表格、移动端上传、会话隔离 | +| `pipeline.py` | CLI 流程编排:串联提取 → 类型判断 → 信息提取 → 浏览器填报 | +| `pipeline_core.py` | CLI/Web 公共管道逻辑:发票类型判断、缓存提取 | +| `main.py` | CLI 入口:`--step` 分步执行、`-u/-p` 覆盖凭据 | +| `config.py` | 配置加载:`config.json` + 环境变量 | +| `exceptions.py` | 异常层次定义 | ## 数据流 ```mermaid graph TD A[CLI/Web 入口] --> B["pipeline.py (编排)"] - B --> C["doc/extractor.py (统一提取入口)"] - C --> D["doc/pdf.py (PDF 渲染为图片)"] - C --> E["doc/llm_extractor.py (多模态 LLM 识别)"] + B --> C["core/extraction/extractor.py (统一提取入口)"] + C --> D["infra/documents/pdf.py (PDF 渲染为图片)"] + C --> E["core/extraction/llm_extractor.py (多模态 LLM 识别)"] E --> F["发票 invoice_type=train/hotel/general"] E --> G["支付记录 invoice_type=payment"] E --> H["出差事前申请单 invoice_type=application"] - C --> I["doc/matcher.py (发票与支付记录按金额匹配)"] + C --> I["core/matching/matcher.py (发票与支付记录按金额匹配)"] I --> J["一对一匹配 发票数 == 刷卡数"] I --> K["一对多匹配 贪心算法 相对容差 3%"] - C --> L["doc/invoice.py (CSV/JSON 读写)"] + C --> L["infra/documents/invoice.py (CSV/JSON 读写)"] L --> M["payment_records.csv (支付记录级别)"] L --> N["invoice_summary.csv (发票级别)"] L --> O["travel_applications.json (出差申请单)"] B --> R{"判断报销类型"} - R -->|差旅| T["doc/llm_extractor.py (差旅信息提取)"] - R -->|普通| V["doc/llm_extractor.py (普通发票信息提取)"] + R -->|差旅| T["core/extraction/llm_extractor.py (差旅信息提取)"] + R -->|普通| V["core/extraction/llm_extractor.py (普通发票信息提取)"] T --> W["travel_info.json (差旅信息: 交通/住宿明细、补贴、附件清单)"] V --> X["normal_info.json (普通发票信息: 报销说明、发票总数、总金额、支付方式、附件清单)"] - W --> P["bot/ (浏览器填报 - 仅接收信息并填报)"] + W --> P["infra/browser/ (浏览器填报 - 仅接收信息并填报)"] X --> P P --> Q["差旅模式: travel_info.json → 填报差旅单 → 上传差旅附件"] P --> S["普通模式: 基本信息 → 录入明细 → 支付信息 → 上传附件"] ``` -**数据流变更(2026-06-11):** 差旅信息提取从 `bot.py` 提升到 `pipeline.py` 编排层。在发票提取和匹配完成后立即判断报销类型,差旅发票调用 LLM 提取 `travel_info.json`,普通发票调用 LLM 提取 `normal_info.json`。Bot 仅负责接收信息并填报,不再承担信息提取职责。 +## 子模块文档 -## 文档处理子模块 (`doc/`) - -详见 [`doc/README.md`](doc/README.md) - -核心能力: -- **统一文档提取**:LLM 自行判断文档类型(发票/支付记录/出差事前申请单),无需正则回退 -- **JSON 缓存**:提取结果缓存于 `.invoice_cache/`,避免重复处理 -- **金额匹配**:支持一对多匹配,相对容差 3%,未匹配发票单独列为记录 -- **差旅信息提取**:综合发票、支付记录和匹配结果,提取出差事由、地点、时间等 -- **普通发票信息提取**:综合普通发票、支付记录和匹配结果,提取报销说明、发票总数、总金额、支付方式、附件清单 -- **出库单生成**:将 CSV 数据填入 Word 模板(pywin32 COM,仅 Windows) - -## Web 界面子模块 (`web/`) - -详见 [`web/README.md`](web/README.md) - -核心能力: -- **会话隔离**:每次上传生成独立 `session_id`,文件/日志/配置/结果各自隔离 -- **移动端同步**:PC 端生成二维码指向 `/mobile/`,跨设备协作上传 -- **可编辑表格**:前端加载 CSV 数据,支持在线编辑后保存 +| 目录 | 文档 | +|------|------| +| `agent/` | [`agent/README.md`](agent/README.md) | +| `core/` | [`core/README.md`](core/README.md) | +| `infra/` | [`infra/README.md`](infra/README.md) | +| `web/` | [`web/README.md`](web/README.md) | ## 启动方式 diff --git a/src/agent/coordinator.py b/src/agent/coordinator.py new file mode 100644 index 0000000..bade9f4 --- /dev/null +++ b/src/agent/coordinator.py @@ -0,0 +1,410 @@ +"""Agent 协调器 + +作为调度中枢,编排信息提取、规则校验的完整流程。 + +校验-修正循环由 Agent 层调度: + 1. Agent 调用 LLM 提取信息 + 2. Agent 调用 validator.py 校验 + 3. 校验失败则构建修正提示,再次调用 LLM + 4. 重复直到校验通过或达到最大重试次数 + 5. LLM 在输出中包含 can_submit 和 suggestion 字段,用于判断信息完整性 +""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any + +from .. import get_logger +from ..core.extraction import ( + build_extraction_user_message, + llm_query_text, + load_cache, + load_match_result, + merge_supplement_into_info, + parse_json_response, + process_user_supplement, +) +from ..core.validation import validate_extracted_info +from ..infra.llm import ( + build_normal_info_system_prompt, + build_travel_info_system_prompt, +) +from ..pipeline_core import save_cache_info +from .events import emit_agent_event +from .session import AgentSession, AgentState, save_agent_state + +log = get_logger("agent.coordinator") + +# 规则校验-修正循环的最大重试次数 +MAX_VALIDATION_RETRIES = 3 + + +# ------------------------------------------------------------------ +# 辅助:构建修正提示 +# ------------------------------------------------------------------ + + +def _build_correction_prompt( + base_message: str, + report: Any, +) -> str: + """根据校验报告构建修正提示,追加到原始用户消息后。""" + error_feedback = ( + f"\n\n=== 上一次输出的校验结果 ===\n" + f"校验未通过,发现以下问题:\n" + f"缺失字段 ({len(report.missing_fields)} 个):{', '.join(report.missing_fields)}\n" + ) + if report.missing_materials: + error_feedback += f"可能需要补充的材料:{', '.join(report.missing_materials)}\n" + if report.suggestion: + error_feedback += f"建议:{report.suggestion}\n" + error_feedback += ( + "\n请根据以上校验结果修正你的输出,确保所有必填字段都有值。" + "如果某个字段确实没有数据,请给出合理的猜测值。" + "再次返回完整的 JSON 结果。" + ) + return base_message + error_feedback + + +# ------------------------------------------------------------------ +# 核心协调逻辑 +# ------------------------------------------------------------------ + + +def _do_extraction_with_validation( + session_dir: Path, + session: AgentSession, + previous_analysis: dict[str, Any] | None = None, + cache_map: dict[str, Any] | None = None, +) -> dict[str, Any]: + """Agent 调度的提取-校验-修正循环。 + + 流程: + 1. 加载缓存数据和匹配结果 + 2. 构建用户消息 + 3. 调用 LLM 提取 + 4. 调用 validator 校验 + 5. 校验失败则构建修正提示,回到步骤 3 + 6. 最多重试 MAX_VALIDATION_RETRIES 次 + + LLM 输出的 JSON 中额外包含 can_submit 和 suggestion 字段, + 用于判断信息是否完整可提交。 + + Args: + session_dir: 会话目录。 + session: 当前 Agent 会话。 + previous_analysis: 上一轮 LLM 分析结果(可选,补充文件时传入作为历史上下文)。 + + Returns: + 校验通过的结构化数据(或达到重试上限后的最佳结果)。 + """ + cache_map = cache_map or load_cache(session_dir) + match_result = load_match_result(session_dir) + + # 选择系统提示词 + if session.invoice_type == "travel": + system_prompt = build_travel_info_system_prompt() + else: + system_prompt = build_normal_info_system_prompt() + + base_message = build_extraction_user_message(cache_map, match_result, previous_analysis=previous_analysis) + current_message = base_message + + for attempt in range(1, MAX_VALIDATION_RETRIES + 1): + log.info( + "LLM 提取第 %d/%d 次尝试 (%s)", + attempt, + MAX_VALIDATION_RETRIES, + session.invoice_type, + ) + emit_agent_event( + session_dir, + "agent_state_change", + state=AgentState.EXTRACTING, + round=session.rounds, + attempt=attempt, + message=f"正在分析文件... (第{attempt}次)", + ) + + # Step 1: 调用 LLM 提取 + try: + response = llm_query_text( + system_prompt=system_prompt, + text=current_message, + reasoning_effort="low", + source_dir=session_dir, + ) + result = parse_json_response(response) + except Exception as e: + log.error("LLM 提取失败: %s", e) + emit_agent_event( + session_dir, + "agent_error", + message=f"LLM 提取失败: {e}", + ) + raise + + # Step 2: 调用 validator 校验 + report = validate_extracted_info(result, invoice_type=session.invoice_type) + + if report.valid: + log.info("规则校验通过 (第 %d 次尝试)", attempt) + emit_agent_event( + session_dir, + "agent_extract_status", + state=AgentState.EXTRACTING, + round=session.rounds, + attempt=attempt, + message=f"规则校验通过 (第{attempt}次)", + ) + return result + + # Step 3: 校验失败,构建修正提示 + log.warning( + "规则校验未通过 (第 %d/%d 次): 缺失 %d 个字段 - %s", + attempt, + MAX_VALIDATION_RETRIES, + len(report.missing_fields), + report.missing_fields, + ) + emit_agent_event( + session_dir, + "agent_extract_status", + state=AgentState.EXTRACTING, + round=session.rounds, + attempt=attempt, + message=f"规则校验未通过,缺失 {len(report.missing_fields)} 个字段,正在请求 LLM 修正...", + ) + current_message = _build_correction_prompt(current_message, report) + + # 所有重试都失败,返回最后一次结果 + log.error( + "LLM 提取经过 %d 次尝试仍未通过规则校验,返回最后一次结果 (置信度: %.0f%%)", + MAX_VALIDATION_RETRIES, + report.confidence * 100, + ) + return result + + +def run_agent_round( + session_dir: Path, + session: AgentSession, + new_files: list[str] | None = None, +) -> AgentSession: + """执行一轮 Agent 处理:提取-校验-修正循环。 + + Args: + session_dir: 会话目录。 + session: 当前 Agent 会话。 + new_files: 新增的文件列表(可选,补充文件时传入)。 + + Returns: + 更新后的 Agent 会话。 + + 注意: + - 信息提取优先从 .invoice_cache 缓存读取,避免重复调用 LLM。 + - Agent 调度提取-校验-修正循环:LLM 提取 -> validator 校验 -> 失败则反馈修正。 + - 若缓存缺失则执行提取后立即写回缓存(travel_info.json / normal_info.json)。 + - 补充文件时(new_files 非空),加载上一轮分析结果作为历史上下文,强制重新分析。 + """ + # 终态保护:会话已提交或已完成时不再重复处理 + if session.is_terminal(): + log.info("Agent 会话已处于终态 (%s),跳过重复处理", session.state.value) + return session + + if session.rounds >= session.max_rounds: + session.state = AgentState.ERROR + session.error_message = f"已达到最大轮次 ({session.max_rounds}),请检查信息或强制提交" + log.warning("Agent 达到最大轮次限制") + emit_agent_event( + session_dir, + "agent_max_rounds", + message=session.error_message, + ) + return session + + session.rounds += 1 + log.info("开始第 %d 轮 Agent 处理", session.rounds) + + # ---- Step 1: 信息提取(Agent 调度校验-修正循环) ---- + session.state = AgentState.EXTRACTING + + # 判断是否为补充文件场景:有新文件传入时,加载上一轮分析结果作为上下文 + cache_map = load_cache(session_dir) + is_supplement = bool(new_files) + previous_analysis = None + if is_supplement: + info_key = "travel_info" if session.invoice_type == "travel" else "normal_info" + previous_analysis = cache_map.get(info_key) + if previous_analysis: + log.info("检测到补充文件,加载上一轮分析结果作为历史上下文") + + try: + info_key = "travel_info" if session.invoice_type == "travel" else "normal_info" + should_reanalyze = not cache_map.get(info_key) or is_supplement + if should_reanalyze: + session.extracted_info = _do_extraction_with_validation( + session_dir, session, previous_analysis=previous_analysis, cache_map=cache_map + ) + # 提取后立即写入缓存,后续步骤依赖此数据 + save_cache_info(session_dir, info_key, session.extracted_info) + else: + session.extracted_info = cache_map[info_key] + emit_agent_event( + session_dir, + "agent_state_change", + state=session.state, + round=session.rounds, + message="使用缓存数据,无需重新分析", + ) + + except Exception as e: + session.state = AgentState.ERROR + session.error_message = f"信息提取失败: {e}" + log.error("Agent 信息提取失败: %s", e) + emit_agent_event( + session_dir, + "agent_error", + message=session.error_message, + ) + return session + + # ---- 判断结果(从 LLM 提取结果中读 can_submit) ---- + can_submit = session.extracted_info.get("can_submit", True) + suggestion = session.extracted_info.get("suggestion", "") + + if can_submit: + session.state = AgentState.READY + emit_agent_event( + session_dir, + "agent_ready", + round=session.rounds, + message="信息完整,可以提交", + ) + log.info("Agent 校验通过,信息完整") + else: + session.state = AgentState.AWAITING_SUPPLEMENT + + combined_suggestion = suggestion or "信息不完整,请补充材料" + + emit_agent_event( + session_dir, + "agent_request_supplement", + round=session.rounds, + missing_fields=[], + missing_materials=[], + semantic_issues=[], + suggestion=combined_suggestion, + ) + log.info("Agent 请求补充: %s", combined_suggestion) + + save_agent_state(session_dir, session) + return session + + +def force_submit( + session_dir: Path, + session: AgentSession, +) -> AgentSession: + """用户强制提交,跳过校验。""" + session.state = AgentState.READY + log.info("用户强制提交,跳过校验") + emit_agent_event( + session_dir, + "agent_force_submit", + message="用户选择强制提交", + ) + save_agent_state(session_dir, session) + return session + + +def add_supplement( + session_dir: Path, + session: AgentSession, + filenames: list[str], +) -> AgentSession: + """记录用户补充的文件。""" + session.user_supplements.extend(filenames) + emit_agent_event( + session_dir, + "agent_supplement_received", + files=filenames, + ) + log.info("收到用户补充文件: %s", filenames) + save_agent_state(session_dir, session) + return session + + +def process_user_text_supplement( + session_dir: Path, + session: AgentSession, + user_text: str, +) -> AgentSession: + """处理用户通过文字补充的信息。 + + 流程: + 1. LLM 分析用户文字,提取需要更新的字段 + 2. 合并到已提取的信息中 + 3. 保存到缓存 + 4. 重新执行一轮 Agent 校验 + + Args: + session_dir: 会话目录。 + session: 当前 Agent 会话。 + user_text: 用户输入的文字。 + + Returns: + 更新后的 Agent 会话。 + """ + log.info("收到用户文字补充: %s", user_text) + emit_agent_event( + session_dir, + "agent_supplement_received", + files=[user_text[:50]], # 简短显示 + ) + + # Step 1: LLM 分析用户文字 + supplement_result = process_user_supplement( + user_text=user_text, + extracted_info=session.extracted_info, + invoice_type=session.invoice_type, + source_dir=session_dir, + ) + + updated_fields = supplement_result.get("updated_fields", {}) + unparsed = supplement_result.get("unparsed_info", "") + + if updated_fields: + # Step 2: 合并到已提取信息 + session.extracted_info = merge_supplement_into_info( + session.extracted_info, + updated_fields, + ) + + # Step 3: 保存到缓存 + info_key = "travel_info" if session.invoice_type == "travel" else "normal_info" + save_cache_info(session_dir, info_key, session.extracted_info) + log.info("已更新 %s", info_key) + + # Step 4: 重新执行 Agent 校验 + session.state = AgentState.EXTRACTING + emit_agent_event( + session_dir, + "agent_state_change", + state=session.state, + round=session.rounds, + message="正在重新校验...", + ) + session = run_agent_round(session_dir, session) + else: + # 没有可更新的字段 + msg = unparsed or "未识别到可更新的报销信息" + emit_agent_event( + session_dir, + "agent_supplement_received", + files=[msg], + ) + log.info("用户补充未识别到有效信息: %s", msg) + + return session diff --git a/src/agent/events.py b/src/agent/events.py new file mode 100644 index 0000000..2efc12f --- /dev/null +++ b/src/agent/events.py @@ -0,0 +1,75 @@ +"""Agent 事件系统 + +负责 SSE 事件的发射和管理,用于实时通知前端状态变化。 +""" + +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any + +from .. import get_logger + +log = get_logger("agent.events") + +AGENT_EVENT_LOG = "agent_events.log" + +# 去重守卫:记录每个 session 上一次发射的事件类型,防止连续重复发射 +# key: str(session_dir), value: 上一次的 event_type +_last_event_type: dict[str, str] = {} + + +def emit_agent_event(session_dir: Path, event_type: str, **kwargs: Any) -> None: + """向 agent_events.log 追加一行 JSON 事件。 + + 同一 session 连续发射相同 event_type 时直接抛出 RuntimeError, + 强制调用方修复重复发射的代码,而非静默掩盖。 + + 注意:agent_state_change 会在同一轮提取中多次发射不同消息 + ("正在分析"、"校验通过"、"校验失败,请求修正"),这是合法行为。 + 去重守卫仅检查 event_type 字符串是否完全相同,不检查 kwargs。 + 因此不要在同一个 event_type 下连续发射不同消息,应使用不同的事件类型。 + """ + session_key = str(session_dir) + prev = _last_event_type.get(session_key) + if prev == event_type: + raise RuntimeError( + f"事件重复发射: session={session_dir.name!r}, event_type={event_type!r}。" + f"请检查调用链,确保每个事件类型只发射一次。" + ) + _last_event_type[session_key] = event_type + + event = {"type": event_type, **kwargs} + try: + event_path = session_dir / AGENT_EVENT_LOG + with open(event_path, "a", encoding="utf-8") as f: + f.write(json.dumps(event, ensure_ascii=False) + "\n") + except Exception: + pass + + +def clear_event_history(session_dir: Path) -> None: + """清除指定会话的事件历史。""" + session_key = str(session_dir) + _last_event_type.pop(session_key, None) + event_path = session_dir / AGENT_EVENT_LOG + if event_path.exists(): + event_path.unlink() + + +def read_events(session_dir: Path) -> list[dict[str, Any]]: + """读取指定会话的所有事件记录。""" + event_path = session_dir / AGENT_EVENT_LOG + if not event_path.exists(): + return [] + events = [] + try: + with open(event_path, encoding="utf-8") as f: + for line in f: + line = line.strip() + if line: + events.append(json.loads(line)) + except Exception as e: + log.warning("读取事件日志失败: %s", e) + return events diff --git a/src/agent/orchestrator.py b/src/agent/orchestrator.py index 985d68f..ee0346e 100644 --- a/src/agent/orchestrator.py +++ b/src/agent/orchestrator.py @@ -1,529 +1,46 @@ -"""Agent 协调器 +"""Agent 协调器(兼容层) -作为调度中枢,编排信息提取、规则校验的完整流程。 +此文件为向后兼容而保留,所有功能已迁移到以下子模块: +- session.py: 会话状态定义和持久化 +- events.py: SSE 事件发射系统 +- coordinator.py: 核心协调逻辑 -校验-修正循环由 Agent 层调度: - 1. Agent 调用 LLM 提取信息 - 2. Agent 调用 validator.py 校验 - 3. 校验失败则构建修正提示,再次调用 LLM - 4. 重复直到校验通过或达到最大重试次数 - 5. LLM 在输出中包含 can_submit 和 suggestion 字段,用于判断信息完整性 - -状态机: - idle -> extracting -> awaiting_supplement -> (回到extracting) - | - (完整) -> ready -> submitting -> done - | - (用户强制) -> ready +原有导入路径保持可用。 """ -from __future__ import annotations - -import json -from dataclasses import asdict, dataclass, field -from enum import StrEnum -from pathlib import Path -from typing import Any - -from .. import get_logger -from ..doc.llm_extractor import ( - build_extraction_user_message, - llm_query_text, - load_cache, - load_match_result, - parse_json_response, +# 从新模块重新导出所有符号,保持向后兼容 +from .coordinator import ( + add_supplement, + force_submit, + process_user_text_supplement, + run_agent_round, ) -from ..doc.prompt import ( - build_normal_info_system_prompt, - build_travel_info_system_prompt, +from .events import emit_agent_event as _emit_agent_event +from .session import ( + AgentSession, + AgentState, + load_agent_state, + save_agent_state, ) -from ..doc.validator import validate_extracted_info -from ..pipeline_core import save_cache_info -log = get_logger("agent") +# 为旧代码提供兼容的私有函数 +_emit_agent_event = _emit_agent_event -# 规则校验-修正循环的最大重试次数 +# 常量保持不变 MAX_VALIDATION_RETRIES = 3 - - -# ------------------------------------------------------------------ -# 状态枚举 -# ------------------------------------------------------------------ - - -class AgentState(StrEnum): - IDLE = "idle" - EXTRACTING = "extracting" - AWAITING_SUPPLEMENT = "awaiting_supplement" - READY = "ready" - SUBMITTING = "submitting" - DONE = "done" - ERROR = "error" - - -# ------------------------------------------------------------------ -# Agent 会话数据模型 -# ------------------------------------------------------------------ - - -@dataclass -class AgentSession: - """Agent 会话状态""" - - session_id: str - state: AgentState = AgentState.IDLE - rounds: int = 0 - max_rounds: int = 5 - invoice_type: str = "travel" # "travel" 或 "normal" - extracted_info: dict[str, Any] = field(default_factory=dict) - validation_reports: list[dict[str, Any]] = field(default_factory=list) - user_supplements: list[str] = field(default_factory=list) - error_message: str = "" - - def to_dict(self) -> dict[str, Any]: - return asdict(self) - - @classmethod - def from_dict(cls, data: dict[str, Any]) -> AgentSession: - # 兼容旧版本:state 可能是字符串 - if "state" in data and isinstance(data["state"], str): - data["state"] = AgentState(data["state"]) - return cls(**data) - - -# ------------------------------------------------------------------ -# 持久化 -# ------------------------------------------------------------------ - AGENT_STATE_FILE = "agent_state.json" - - -def save_agent_state(session_dir: Path, session: AgentSession) -> None: - """将 Agent 会话状态持久化到 session 目录。""" - state_path = session_dir / AGENT_STATE_FILE - tmp_path = session_dir / (AGENT_STATE_FILE + ".tmp") - with open(tmp_path, "w", encoding="utf-8") as f: - json.dump(session.to_dict(), f, ensure_ascii=False, indent=2) - tmp_path.replace(state_path) - - -def load_agent_state(session_dir: Path) -> AgentSession | None: - """从 session 目录加载 Agent 会话状态。""" - state_path = session_dir / AGENT_STATE_FILE - if not state_path.exists(): - return None - try: - with open(state_path, encoding="utf-8") as f: - return AgentSession.from_dict(json.load(f)) - except Exception as e: - log.warning("加载 Agent 状态失败: %s", e) - return None - - -# ------------------------------------------------------------------ -# SSE 事件发射 -# ------------------------------------------------------------------ - AGENT_EVENT_LOG = "agent_events.log" -# 去重守卫:记录每个 session 上一次发射的事件类型,防止连续重复发射 -# key: str(session_dir), value: 上一次的 event_type -_last_event_type: dict[str, str] = {} - - -def _emit_agent_event(session_dir: Path, event_type: str, **kwargs: Any) -> None: - """向 agent_events.log 追加一行 JSON 事件。 - - 同一 session 连续发射相同 event_type 时直接抛出 RuntimeError, - 强制调用方修复重复发射的代码,而非静默掩盖。 - """ - session_key = str(session_dir) - prev = _last_event_type.get(session_key) - if prev == event_type: - raise RuntimeError( - f"事件重复发射: session={session_dir.name!r}, event_type={event_type!r}。" - f"请检查调用链,确保每个事件类型只发射一次。" - ) - _last_event_type[session_key] = event_type - - event = {"type": event_type, **kwargs} - try: - event_path = session_dir / AGENT_EVENT_LOG - with open(event_path, "a", encoding="utf-8") as f: - f.write(json.dumps(event, ensure_ascii=False) + "\n") - except Exception: - pass - - -# ------------------------------------------------------------------ -# 辅助:构建修正提示 -# ------------------------------------------------------------------ - - -def _build_correction_prompt( - base_message: str, - report: Any, -) -> str: - """根据校验报告构建修正提示,追加到原始用户消息后。""" - error_feedback = ( - f"\n\n=== 上一次输出的校验结果 ===\n" - f"校验未通过,发现以下问题:\n" - f"缺失字段 ({len(report.missing_fields)} 个):{', '.join(report.missing_fields)}\n" - ) - if report.missing_materials: - error_feedback += f"可能需要补充的材料:{', '.join(report.missing_materials)}\n" - if report.suggestion: - error_feedback += f"建议:{report.suggestion}\n" - error_feedback += ( - "\n请根据以上校验结果修正你的输出,确保所有必填字段都有值。" - "如果某个字段确实没有数据,请给出合理的猜测值。" - "再次返回完整的 JSON 结果。" - ) - return base_message + error_feedback - - -# ------------------------------------------------------------------ -# 核心协调逻辑 -# ------------------------------------------------------------------ - - -def _do_extraction_with_validation( - session_dir: Path, - session: AgentSession, - previous_analysis: dict[str, Any] | None = None, -) -> dict[str, Any]: - """Agent 调度的提取-校验-修正循环。 - - 流程: - 1. 加载缓存数据和匹配结果 - 2. 构建用户消息 - 3. 调用 LLM 提取 - 4. 调用 validator 校验 - 5. 校验失败则构建修正提示,回到步骤 3 - 6. 最多重试 MAX_VALIDATION_RETRIES 次 - - LLM 输出的 JSON 中额外包含 can_submit 和 suggestion 字段, - 用于判断信息是否完整可提交。 - - Args: - session_dir: 会话目录。 - session: 当前 Agent 会话。 - previous_analysis: 上一轮 LLM 分析结果(可选,补充文件时传入作为历史上下文)。 - - Returns: - 校验通过的结构化数据(或达到重试上限后的最佳结果)。 - """ - cache_map = load_cache(session_dir) - match_result = load_match_result(session_dir) - - # 选择系统提示词 - if session.invoice_type == "travel": - system_prompt = build_travel_info_system_prompt() - else: - system_prompt = build_normal_info_system_prompt() - - base_message = build_extraction_user_message(cache_map, match_result, previous_analysis=previous_analysis) - current_message = base_message - - for attempt in range(1, MAX_VALIDATION_RETRIES + 1): - log.info( - "LLM 提取第 %d/%d 次尝试 (%s)", - attempt, - MAX_VALIDATION_RETRIES, - session.invoice_type, - ) - _emit_agent_event( - session_dir, - "agent_state_change", - state=AgentState.EXTRACTING, - round=session.rounds, - attempt=attempt, - message=f"正在分析文件... (第{attempt}次)", - ) - - # Step 1: 调用 LLM 提取 - try: - response = llm_query_text( - system_prompt=system_prompt, - text=current_message, - reasoning_effort="low", - source_dir=session_dir, - ) - result = parse_json_response(response) - except Exception as e: - log.error("LLM 提取失败: %s", e) - _emit_agent_event( - session_dir, - "agent_error", - message=f"LLM 提取失败: {e}", - ) - raise - - # Step 2: 调用 validator 校验 - report = validate_extracted_info(result, invoice_type=session.invoice_type) - - if report.valid: - log.info("规则校验通过 (第 %d 次尝试)", attempt) - _emit_agent_event( - session_dir, - "agent_state_change", - state=AgentState.EXTRACTING, - round=session.rounds, - attempt=attempt, - message=f"规则校验通过 (第{attempt}次)", - ) - return result - - # Step 3: 校验失败,构建修正提示 - log.warning( - "规则校验未通过 (第 %d/%d 次): 缺失 %d 个字段 - %s", - attempt, - MAX_VALIDATION_RETRIES, - len(report.missing_fields), - report.missing_fields, - ) - _emit_agent_event( - session_dir, - "agent_state_change", - state=AgentState.EXTRACTING, - round=session.rounds, - attempt=attempt, - message=f"规则校验未通过,缺失 {len(report.missing_fields)} 个字段,正在请求 LLM 修正...", - ) - current_message = _build_correction_prompt(current_message, report) - - # 所有重试都失败,返回最后一次结果 - log.error( - "LLM 提取经过 %d 次尝试仍未通过规则校验,返回最后一次结果 (置信度: %.0f%%)", - MAX_VALIDATION_RETRIES, - report.confidence * 100, - ) - return result - - -def run_agent_round( - session_dir: Path, - session: AgentSession, - new_files: list[str] | None = None, -) -> AgentSession: - """执行一轮 Agent 处理:提取-校验-修正循环。 - - Args: - session_dir: 会话目录。 - session: 当前 Agent 会话。 - new_files: 新增的文件列表(可选,补充文件时传入)。 - - Returns: - 更新后的 Agent 会话。 - - 注意: - - 信息提取优先从 .invoice_cache 缓存读取,避免重复调用 LLM。 - - Agent 调度提取-校验-修正循环:LLM 提取 -> validator 校验 -> 失败则反馈修正。 - - 若缓存缺失则执行提取后立即写回缓存(travel_info.json / normal_info.json)。 - - 补充文件时(new_files 非空),加载上一轮分析结果作为历史上下文,强制重新分析。 - """ - # 终态保护:会话已提交或已完成时不再重复处理 - if session.state in (AgentState.DONE, AgentState.SUBMITTING, AgentState.READY): - log.info("Agent 会话已处于终态 (%s),跳过重复处理", session.state.value) - return session - - if session.rounds >= session.max_rounds: - session.state = AgentState.ERROR - session.error_message = f"已达到最大轮次 ({session.max_rounds}),请检查信息或强制提交" - log.warning("Agent 达到最大轮次限制") - _emit_agent_event( - session_dir, - "agent_max_rounds", - message=session.error_message, - ) - return session - - session.rounds += 1 - log.info("开始第 %d 轮 Agent 处理", session.rounds) - - # ---- Step 1: 信息提取(Agent 调度校验-修正循环) ---- - session.state = AgentState.EXTRACTING - _emit_agent_event( - session_dir, - "agent_state_change", - state=session.state, - round=session.rounds, - message="正在分析文件...", - ) - - # 判断是否为补充文件场景:有新文件传入时,加载上一轮分析结果作为上下文 - cache_map = load_cache(session_dir) - is_supplement = bool(new_files) - previous_analysis = None - if is_supplement: - info_key = "travel_info" if session.invoice_type == "travel" else "normal_info" - previous_analysis = cache_map.get(info_key) - if previous_analysis: - log.info("检测到补充文件,加载上一轮分析结果作为历史上下文") - - try: - info_key = "travel_info" if session.invoice_type == "travel" else "normal_info" - should_reanalyze = not cache_map.get(info_key) or is_supplement - if should_reanalyze: - session.extracted_info = _do_extraction_with_validation( - session_dir, session, previous_analysis=previous_analysis - ) - # 提取后立即写入缓存,后续步骤依赖此数据 - save_cache_info(session_dir, info_key, session.extracted_info) - else: - session.extracted_info = cache_map[info_key] - - except Exception as e: - session.state = AgentState.ERROR - session.error_message = f"信息提取失败: {e}" - log.error("Agent 信息提取失败: %s", e) - _emit_agent_event( - session_dir, - "agent_error", - message=session.error_message, - ) - return session - - # ---- 判断结果(从 LLM 提取结果中读 can_submit) ---- - can_submit = session.extracted_info.get("can_submit", True) - suggestion = session.extracted_info.get("suggestion", "") - - if can_submit: - session.state = AgentState.READY - _emit_agent_event( - session_dir, - "agent_ready", - round=session.rounds, - message="信息完整,可以提交", - ) - log.info("Agent 校验通过,信息完整") - else: - session.state = AgentState.AWAITING_SUPPLEMENT - - combined_suggestion = suggestion or "信息不完整,请补充材料" - - _emit_agent_event( - session_dir, - "agent_request_supplement", - round=session.rounds, - missing_fields=[], - missing_materials=[], - semantic_issues=[], - suggestion=combined_suggestion, - ) - log.info("Agent 请求补充: %s", combined_suggestion) - - save_agent_state(session_dir, session) - return session - - -def force_submit( - session_dir: Path, - session: AgentSession, -) -> AgentSession: - """用户强制提交,跳过校验。""" - session.state = AgentState.READY - log.info("用户强制提交,跳过校验") - _emit_agent_event( - session_dir, - "agent_force_submit", - message="用户选择强制提交", - ) - save_agent_state(session_dir, session) - return session - - -def add_supplement( - session_dir: Path, - session: AgentSession, - filenames: list[str], -) -> AgentSession: - """记录用户补充的文件。""" - session.user_supplements.extend(filenames) - _emit_agent_event( - session_dir, - "agent_supplement_received", - files=filenames, - ) - log.info("收到用户补充文件: %s", filenames) - save_agent_state(session_dir, session) - return session - - -def process_user_text_supplement( - session_dir: Path, - session: AgentSession, - user_text: str, -) -> AgentSession: - """处理用户通过文字补充的信息。 - - 流程: - 1. LLM 分析用户文字,提取需要更新的字段 - 2. 合并到已提取的信息中 - 3. 保存到缓存 - 4. 重新执行一轮 Agent 校验 - - Args: - session_dir: 会话目录。 - session: 当前 Agent 会话。 - user_text: 用户输入的文字。 - - Returns: - 更新后的 Agent 会话。 - """ - from ..doc.llm_extractor import ( - merge_supplement_into_info, - process_user_supplement, - ) - - log.info("收到用户文字补充: %s", user_text) - _emit_agent_event( - session_dir, - "agent_supplement_received", - files=[user_text[:50]], # 简短显示 - ) - - # Step 1: LLM 分析用户文字 - supplement_result = process_user_supplement( - user_text=user_text, - extracted_info=session.extracted_info, - invoice_type=session.invoice_type, - source_dir=session_dir, - ) - - updated_fields = supplement_result.get("updated_fields", {}) - unparsed = supplement_result.get("unparsed_info", "") - - if updated_fields: - # Step 2: 合并到已提取信息 - session.extracted_info = merge_supplement_into_info( - session.extracted_info, - updated_fields, - ) - - # Step 3: 保存到缓存 - info_key = "travel_info" if session.invoice_type == "travel" else "normal_info" - save_cache_info(session_dir, info_key, session.extracted_info) - log.info("已更新 %s", info_key) - - # Step 4: 重新执行 Agent 校验 - session.state = AgentState.EXTRACTING - _emit_agent_event( - session_dir, - "agent_state_change", - state=session.state, - round=session.rounds, - message="正在重新校验...", - ) - session = run_agent_round(session_dir, session) - else: - # 没有可更新的字段 - msg = unparsed or "未识别到可更新的报销信息" - _emit_agent_event( - session_dir, - "agent_supplement_received", - files=[msg], - ) - log.info("用户补充未识别到有效信息: %s", msg) - - return session +__all__ = [ + "AgentSession", + "AgentState", + "save_agent_state", + "load_agent_state", + "run_agent_round", + "force_submit", + "add_supplement", + "process_user_text_supplement", + "MAX_VALIDATION_RETRIES", + "AGENT_STATE_FILE", + "AGENT_EVENT_LOG", +] diff --git a/src/agent/session.py b/src/agent/session.py new file mode 100644 index 0000000..ee79c29 --- /dev/null +++ b/src/agent/session.py @@ -0,0 +1,101 @@ +"""Agent 会话管理 + +负责会话状态的定义、序列化和持久化。 +""" + +from __future__ import annotations + +import json +from dataclasses import asdict, dataclass, field +from enum import StrEnum +from pathlib import Path +from typing import Any + +from .. import get_logger + +log = get_logger("agent.session") + +# ------------------------------------------------------------------ +# 状态枚举 +# ------------------------------------------------------------------ + + +class AgentState(StrEnum): + IDLE = "idle" + EXTRACTING = "extracting" + AWAITING_SUPPLEMENT = "awaiting_supplement" + READY = "ready" + SUBMITTING = "submitting" + DONE = "done" + ERROR = "error" + + +# ------------------------------------------------------------------ +# Agent 会话数据模型 +# ------------------------------------------------------------------ + + +@dataclass +class AgentSession: + """Agent 会话状态""" + + session_id: str + state: AgentState = AgentState.IDLE + rounds: int = 0 + max_rounds: int = 5 + invoice_type: str = "travel" # "travel" 或 "normal" + extracted_info: dict[str, Any] = field(default_factory=dict) + validation_reports: list[dict[str, Any]] = field(default_factory=list) + user_supplements: list[str] = field(default_factory=list) + error_message: str = "" + + def to_dict(self) -> dict[str, Any]: + return asdict(self) + + @classmethod + def from_dict(cls, data: dict[str, Any]) -> AgentSession: + # 兼容旧版本:state 可能是字符串 + if "state" in data and isinstance(data["state"], str): + data["state"] = AgentState(data["state"]) + return cls(**data) + + def is_terminal(self) -> bool: + """判断会话是否处于终态(已提交或已完成)""" + return self.state in (AgentState.DONE, AgentState.SUBMITTING, AgentState.READY) + + +# ------------------------------------------------------------------ +# 持久化 +# ------------------------------------------------------------------ + +AGENT_STATE_FILE = "agent_state.json" + + +def save_agent_state(session_dir: Path, session: AgentSession) -> None: + """将 Agent 会话状态持久化到 session 目录。""" + state_path = session_dir / AGENT_STATE_FILE + tmp_path = session_dir / (AGENT_STATE_FILE + ".tmp") + with open(tmp_path, "w", encoding="utf-8") as f: + json.dump(session.to_dict(), f, ensure_ascii=False, indent=2) + tmp_path.replace(state_path) + + +def load_agent_state(session_dir: Path) -> AgentSession | None: + """从 session 目录加载 Agent 会话状态。""" + state_path = session_dir / AGENT_STATE_FILE + if not state_path.exists(): + return None + try: + with open(state_path, encoding="utf-8") as f: + return AgentSession.from_dict(json.load(f)) + except Exception as e: + log.warning("加载 Agent 状态失败: %s", e) + return None + + +def create_agent_session(session_id: str, invoice_type: str = "travel") -> AgentSession: + """创建新的 Agent 会话。""" + return AgentSession( + session_id=session_id, + invoice_type=invoice_type, + ) diff --git a/src/bot/README.md b/src/bot/README.md deleted file mode 100644 index 97ea030..0000000 --- a/src/bot/README.md +++ /dev/null @@ -1,33 +0,0 @@ ---- -last_reviewed: 2026-06-12 ---- - -# bot — 浏览器自动化填报模块 - -使用 Playwright 操作财务报销系统,自动完成登录、填单、上传附件等操作。 - -## 模块清单 - -| 文件 | 说明 | -|------|------| -| `__init__.py` | 对外入口:`run_bot()` 和 `run_bot_web()`,负责类型判断和流程路由 | -| `base.py` | `BaseBot` 基类:浏览器生命周期、登录、导航、截图、日期格式化 | -| `travel.py` | 差旅报销填报流程:基本信息 → 差旅明细 → 支付方式 → 补助清单 → 附件上传 | -| `normal.py` | 普通发票报销填报流程:基本信息 → 总明细 → 支付方式 → 附件上传 | - -## 架构设计 - -``` -run_bot(config, travel_info, normal_info) - ├── 创建 BaseBot,启动浏览器,登录门户 - ├── travel_info 存在 → travel.run(bot, travel_info) - └── normal_info 存在 → normal.run(bot, normal_info) -``` - -- **`BaseBot`** 只保留公共操作(launch、login、navigate、create_new_form、close、screenshot) -- **差旅/普通流程** 作为独立函数接受 `bot: BaseBot` 参数,符合函数式编程偏好 -- **`__init__.py`** 仅做路由分发,不包含具体填报逻辑 - -## 变更历史 - -- **2026-06-12**:从 `bot.py` 单文件重构为 `bot/` 包,分离差旅和普通报销逻辑 \ No newline at end of file diff --git a/src/bot/__init__.py b/src/bot/__init__.py deleted file mode 100644 index 29ea387..0000000 --- a/src/bot/__init__.py +++ /dev/null @@ -1,97 +0,0 @@ -""" -浏览器自动化填报 - -使用 Playwright 操作财务报销系统,自动完成登录、填单、上传附件等操作。 - -对外接口: - run_bot(config, travel_info, normal_info) 启动浏览器并执行填报流程 - run_bot_web(config, work_dir) Web 模式填报(从缓存加载信息) -""" - -from pathlib import Path -from typing import Any - -from .. import get_logger -from .base import BaseBot - -log = get_logger("bot") - - -def run_bot( - config: dict[str, Any], - headless: bool = False, - work_dir: Path | None = None, - travel_info: dict[str, Any] | None = None, - normal_info: dict[str, Any] | None = None, -) -> None: - """执行完整的浏览器填报流程,根据发票类型自动路由 - - Args: - config: 配置字典。 - headless: 是否无头模式。 - work_dir: 工作目录。 - travel_info: 差旅信息(由 pipeline 层提前提取并传入,非差旅时传 None)。 - normal_info: 普通发票信息(由 pipeline 层提前提取并传入,非普通时传 None)。 - """ - if not config["username"] or not config["password"]: - raise ValueError("缺少用户名或密码") - - if not work_dir: - raise ValueError("缺少工作目录") - - bot = BaseBot(config, headless=headless) - bot.work_dir = work_dir - - try: - bot.launch() - bot.login_portal() - - if travel_info is not None: - log.info("处理差旅发票...") - bot.navigate_to_reimburse(page_key="travel_page") - bot.create_new_form() - from . import travel - - travel.run(bot, travel_info) - elif normal_info is not None: - log.info("处理普通发票...") - bot.navigate_to_reimburse(page_key="reimburse_page") - bot.create_new_form() - from . import normal - - normal.run(bot, normal_info) - else: - raise ValueError("缺少差旅信息(travel_info)和普通发票信息(normal_info),无法继续填报") - - except Exception as e: - log.error(f"操作失败: {e}") - try: - bot._screenshot("error") - except Exception: - pass - raise - finally: - bot.close() - - -def run_bot_web(config: dict[str, Any], work_dir: Path) -> None: - """Web 模式填报 — headless,附件从指定目录读取 - - Web 端的信息提取由 app.py 的管道负责,此处从缓存加载。 - """ - from ..doc.llm_extractor import load_cache - - cache_map = load_cache(work_dir) - travel_info = cache_map.get("travel_info") - normal_info = cache_map.get("normal_info") - - if travel_info is None and normal_info is None: - raise ValueError(f"缓存中未找到 travel_info 或 normal_info,请先执行发票提取流程。工作目录: {work_dir}") - - run_bot( - config, - headless=True, - work_dir=work_dir, - travel_info=travel_info, - normal_info=normal_info, - ) diff --git a/src/core/README.md b/src/core/README.md new file mode 100644 index 0000000..38d555d --- /dev/null +++ b/src/core/README.md @@ -0,0 +1,21 @@ +--- +last_reviewed: 2026-06-15 +--- + +# src/core — 核心业务逻辑 + +项目的核心业务层,负责信息提取、金额匹配和信息校验。此层不依赖 Web 框架或浏览器自动化等基础设施。 + +## 子模块 + +| 目录 | 说明 | +|------|------| +| `extraction/` | 文档信息提取:PDF/图片 → LLM 多模态识别 → 结构化数据 | +| `matching/` | 发票与支付记录按金额匹配(一对一 / 一对多) | +| `validation/` | 声明式信息完整性校验,规则从 JSON 配置文件加载 | + +## 设计原则 + +- **零外部依赖**:不依赖 Flask、Playwright 等框架 +- **接口契约**:每个子模块通过 `__init__.py` 导出稳定的对外接口 +- **错误传播**:明确的异常层次,便于上层统一处理 diff --git a/src/core/__init__.py b/src/core/__init__.py new file mode 100644 index 0000000..6fbfeb5 --- /dev/null +++ b/src/core/__init__.py @@ -0,0 +1,4 @@ +"""核心业务逻辑模块 + +提供信息提取、规则校验和发票匹配功能。 +""" diff --git a/src/core/extraction/README.md b/src/core/extraction/README.md new file mode 100644 index 0000000..302558d --- /dev/null +++ b/src/core/extraction/README.md @@ -0,0 +1,29 @@ +--- +last_reviewed: 2026-06-15 +--- + +# src/core/extraction — 信息提取 + +从 PDF 发票和图片中提取结构化数据,是系统数据流的起点。 + +## 文件 + +| 文件 | 职责 | +|------|------| +| `extractor.py` | 编排入口:扫描目录 → 逐文件提取 → 分类(发票/支付记录/申请单)→ 金额匹配 | +| `llm_extractor.py` | LLM 多模态提取核心:统一文档提取、差旅/普通信息提取、缓存管理、SSE 流式事件 | + +## 对外接口 + +| 函数 | 说明 | +|------|------| +| `extract_invoices(directory)` | 统一提取入口,返回 `(payment_records, applications, groups)` | +| `extract_document(file_path)` | 从单个图片/PDF 提取信息 | +| `extract_travel_info(source_dir)` | 综合发票和匹配结果提取差旅信息 | +| `extract_normal_info(source_dir)` | 提取普通发票报销信息 | +| `load_cache(source_dir)` | 加载缓存的结构化数据 | +| `llm_query_text(...)` | 纯文本 LLM 查询(供 Agent 调度使用) | + +## 缓存机制 + +提取结果缓存在 `.invoice_cache/` 目录中,文件名与源文件同名(`发票1.pdf` → `.invoice_cache/发票1.json`),避免重复调用 LLM。 diff --git a/src/core/extraction/__init__.py b/src/core/extraction/__init__.py new file mode 100644 index 0000000..8cbd4ae --- /dev/null +++ b/src/core/extraction/__init__.py @@ -0,0 +1,42 @@ +"""信息提取模块 + +提供发票/文档结构化提取、LLM 辅助提取等功能。 +""" + +from .extractor import ( + EXTRACTION_PARALLEL_COUNT, + FILE_EVENTS_LOG, + SUPPORTED_EXTENSIONS, + extract_invoices, +) +from .llm_extractor import ( + CACHE_DIR_NAME, + build_extraction_user_message, + extract_document, + extract_normal_info, + extract_travel_info, + llm_query_text, + load_cache, + load_match_result, + merge_supplement_into_info, + parse_json_response, + process_user_supplement, +) + +__all__ = [ + "CACHE_DIR_NAME", + "build_extraction_user_message", + "extract_document", + "extract_normal_info", + "extract_travel_info", + "load_cache", + "load_match_result", + "llm_query_text", + "merge_supplement_into_info", + "parse_json_response", + "process_user_supplement", + "EXTRACTION_PARALLEL_COUNT", + "FILE_EVENTS_LOG", + "SUPPORTED_EXTENSIONS", + "extract_invoices", +] diff --git a/src/doc/extractor.py b/src/core/extraction/extractor.py similarity index 86% rename from src/doc/extractor.py rename to src/core/extraction/extractor.py index 348df75..6b9b109 100644 --- a/src/doc/extractor.py +++ b/src/core/extraction/extractor.py @@ -20,20 +20,25 @@ SSE 文件进度事件: """ import json +import os +from concurrent.futures import ThreadPoolExecutor, as_completed from pathlib import Path from typing import Any -from .. import get_logger -from ..exceptions import ExtractionError -from .invoice import CACHE_DIR_NAME, classify_invoice_batch +from ... import get_logger +from ...core.matching import match_invoices_to_cards +from ...exceptions import ExtractionError +from ...infra.documents.invoice import CACHE_DIR_NAME, classify_invoice_batch from .llm_extractor import extract_document -from .matcher import match_invoices_to_cards log = get_logger("extractor") # SSE 文件进度事件日志文件名 FILE_EVENTS_LOG = "file_events.log" +# 并行提取文件数,可通过环境变量 EXTRACTION_PARALLEL_COUNT 配置 +EXTRACTION_PARALLEL_COUNT = int(os.environ.get("EXTRACTION_PARALLEL_COUNT", "3")) + # 支持的文件扩展名 SUPPORTED_EXTENSIONS = {".pdf", ".png", ".jpg", ".jpeg", ".bmp", ".webp"} @@ -285,28 +290,39 @@ def extract_invoices( # 记录失败文件及其错误信息 failed_files: list[tuple[str, str]] = [] - for file_path in all_files: - result, err = _extract_document(file_path, cache_dir, source_dir) + max_workers = max(1, EXTRACTION_PARALLEL_COUNT) + with ThreadPoolExecutor(max_workers=max_workers) as executor: + future_to_file = {executor.submit(_extract_document, fp, cache_dir, source_dir): fp for fp in all_files} - if not result: - if err: - failed_files.append((file_path.name, err)) - log.warning(f"未能解析: {file_path.name}") - continue + for future in as_completed(future_to_file): + file_path = future_to_file[future] + try: + result, err = future.result() + except Exception as e: + err_msg = str(e) + failed_files.append((file_path.name, err_msg)) + log.warning(f"未能解析: {file_path.name} ({err_msg})") + continue - inv_type = result.get("invoice_type", "") + if not result: + if err: + failed_files.append((file_path.name, err)) + log.warning(f"未能解析: {file_path.name}") + continue - if inv_type == "application": - applications.append(result) - log.info(f"[{inv_type}] 已解析: {file_path.name}") - elif inv_type == "payment": - all_cards.append(result) - log.info(f"[{inv_type}] 已解析: {file_path.name}") - elif result.get("invoice_number"): - all_invoices.append(result) - log.info(f"[{inv_type}] 已解析: {file_path.name}") - else: - log.warning(f"无法分类: {file_path.name} (invoice_type={inv_type})") + inv_type = result.get("invoice_type", "") + + if inv_type == "application": + applications.append(result) + log.info(f"[{inv_type}] 已解析: {file_path.name}") + elif inv_type == "payment": + all_cards.append(result) + log.info(f"[{inv_type}] 已解析: {file_path.name}") + elif result.get("invoice_number"): + all_invoices.append(result) + log.info(f"[{inv_type}] 已解析: {file_path.name}") + else: + log.warning(f"无法分类: {file_path.name} (invoice_type={inv_type})") # 全部文件提取失败时抛出异常,携带原始错误信息 if failed_files and not all_invoices and not all_cards and not applications: diff --git a/src/doc/llm_extractor.py b/src/core/extraction/llm_extractor.py similarity index 98% rename from src/doc/llm_extractor.py rename to src/core/extraction/llm_extractor.py index be38dcd..a200198 100644 --- a/src/doc/llm_extractor.py +++ b/src/core/extraction/llm_extractor.py @@ -38,9 +38,9 @@ import json from pathlib import Path from typing import Any, cast -from .. import get_logger -from .invoice import CACHE_DIR_NAME -from .prompt import ( +from ... import get_logger +from ...infra.documents.invoice import CACHE_DIR_NAME +from ...infra.llm.prompt import ( build_invoice_system_prompt, build_normal_info_system_prompt, build_supplement_system_prompt, @@ -91,7 +91,7 @@ def _create_llm() -> Any: log.error("缺少 llama-index-llms-openai-like,请执行: uv pip install llama-index-llms-openai-like") raise - from ..config import get_llm_config + from ...config import get_llm_config llm_config = get_llm_config() return OpenAILike( @@ -213,7 +213,7 @@ def llm_query_text( from llama_index.core.base.llms.types import TextBlock from llama_index.core.llms import ChatMessage - from ..config import get_llm_config + from ...config import get_llm_config messages = [ ChatMessage(role="system", content=system_prompt), @@ -242,7 +242,7 @@ def extract_document(file_path: Path) -> dict[str, Any]: Returns: 包含提取字段的字典。 """ - from .pdf import render_pdf_to_images + from ...infra.documents.pdf import render_pdf_to_images system_prompt = build_invoice_system_prompt() user_text = f"请分析以下财务文档并提取信息:\n\n文件名: {file_path.name}" @@ -291,7 +291,7 @@ def _llm_query_multimodal( from llama_index.core.base.llms.types import ImageBlock, TextBlock from llama_index.core.llms import ChatMessage - from ..config import get_llm_config + from ...config import get_llm_config if blocks is not None: final_blocks = blocks diff --git a/src/core/matching/README.md b/src/core/matching/README.md new file mode 100644 index 0000000..ae31fe1 --- /dev/null +++ b/src/core/matching/README.md @@ -0,0 +1,27 @@ +--- +last_reviewed: 2026-06-15 +--- + +# src/core/matching — 金额匹配 + +将提取到的发票数据与支付记录(刷卡截图)按金额进行匹配。 + +## 文件 + +| 文件 | 职责 | +|------|------| +| `matcher.py` | 匹配引擎:一对一匹配、一对多贪心匹配、未匹配发票处理 | + +## 匹配策略 + +| 场景 | 策略 | +|------|------| +| 发票数 == 支付记录数 | 一对一匹配:按金额降序配对,相对容差内即匹配 | +| 发票数 > 支付记录数 | 一对多匹配:贪心算法凑金额,相对容差 3% | +| 文件名匹配 | 最高优先级:文件名(不含后缀)一致时直接匹配 | +| 未匹配发票 | 单独列为一条支付记录,`remark` 标记为 `"unmatched"` | + +## 业务约束 + +- 发票总金额 >= 支付总金额 +- 输出以支付记录为主键的结果列表 diff --git a/src/core/matching/__init__.py b/src/core/matching/__init__.py new file mode 100644 index 0000000..1f83b01 --- /dev/null +++ b/src/core/matching/__init__.py @@ -0,0 +1,8 @@ +"""匹配模块 + +提供发票与支付记录的金额匹配功能。 +""" + +from .matcher import match_invoices_to_cards + +__all__ = ["match_invoices_to_cards"] diff --git a/src/doc/matcher.py b/src/core/matching/matcher.py similarity index 99% rename from src/doc/matcher.py rename to src/core/matching/matcher.py index 9f7eed2..5b31c52 100644 --- a/src/doc/matcher.py +++ b/src/core/matching/matcher.py @@ -49,7 +49,7 @@ from pathlib import Path from typing import Any -from .. import get_logger +from ... import get_logger log = get_logger("matcher") diff --git a/src/core/validation/README.md b/src/core/validation/README.md new file mode 100644 index 0000000..65c5fad --- /dev/null +++ b/src/core/validation/README.md @@ -0,0 +1,34 @@ +--- +last_reviewed: 2026-06-15 +--- + +# src/core/validation — 信息校验 + +对 LLM 提取的报销信息进行声明式规则校验,判断是否满足填报要求。 + +## 文件 + +| 文件 | 职责 | +|------|------| +| `validator.py` | 校验引擎:加载 JSON 规则配置 → 遍历字段/数组 → 输出校验报告 | + +## 设计特点 + +- **规则与引擎分离**:校验规则存储在 `config/validation_rules.json`,引擎只负责执行 +- **统一路径定位**:使用 `path` 列表定位嵌套字段,如 `["basic_info", "travel_purpose"]` +- **自定义校验**:支持 `custom_check` 函数(日期格式、正数检查等) +- **数组元素校验**:支持 `min_items` 最小数量 + 每个元素的必填字段 + +## 校验规则类型 + +| 类型 | 用途 | 配置项 | +|------|------|--------| +| `fields` | 顶层单值字段 | `path`, `required`, `custom_check`, `check_empty` | +| `arrays` | 数组字段 | `path`, `min_items`, `element_fields` | + +## 对外接口 + +| 函数 | 说明 | +|------|------| +| `validate(info, invoice_type)` | 执行校验,返回 `ValidationReport` | +| `get_missing_fields(report)` | 提取缺失字段列表 | diff --git a/src/core/validation/__init__.py b/src/core/validation/__init__.py new file mode 100644 index 0000000..debfb8a --- /dev/null +++ b/src/core/validation/__init__.py @@ -0,0 +1,28 @@ +"""校验模块 + +提供报销信息的规则级校验功能。 +""" + +from .validator import ( + ArrayRule, + FieldRule, + ValidationReport, + ValidationRules, + get_validation_rules, + reload_validation_rules, + validate_extracted_info, + validate_normal_info, + validate_travel_info, +) + +__all__ = [ + "validate_extracted_info", + "validate_travel_info", + "validate_normal_info", + "ValidationReport", + "FieldRule", + "ArrayRule", + "ValidationRules", + "get_validation_rules", + "reload_validation_rules", +] diff --git a/src/core/validation/validator.py b/src/core/validation/validator.py new file mode 100644 index 0000000..6e82088 --- /dev/null +++ b/src/core/validation/validator.py @@ -0,0 +1,537 @@ +"""信息完整性校验器 + +对 LLM 提取的报销信息进行规则级校验,判断是否满足填报要求。 + +校验规则从 JSON 配置文件加载,支持声明式配置。 + +设计理念: + 使用声明式规则配置,将校验规则与校验逻辑分离,提高可读性和可维护性。 +""" + +from __future__ import annotations + +import json +import re +from collections.abc import Callable +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, TypedDict + +from ... import get_logger + +log = get_logger("validator") + +# ------------------------------------------------------------------ +# 日期格式 +# ------------------------------------------------------------------ + +DATE_PATTERN = re.compile(r"^\d{4}-\d{2}-\d{2}$") + + +def _is_valid_date(value: str) -> bool: + """检查日期格式是否为 YYYY-MM-DD。""" + return bool(DATE_PATTERN.match(value)) + + +def _is_positive_number(value: Any) -> bool: + """检查值是否为正数(整数或浮点数)。""" + return isinstance(value, int | float) and value > 0 + + +def _is_positive_integer(value: Any) -> bool: + """检查值是否为正整数。""" + return isinstance(value, int) and value > 0 + + +# 自定义校验函数注册表 +_CUSTOM_CHECKS: dict[str, Callable[[Any], bool]] = { + "is_valid_date": _is_valid_date, + "is_positive_number": _is_positive_number, + "is_positive_integer": _is_positive_integer, +} + + +# ------------------------------------------------------------------ +# 规则定义 +# ------------------------------------------------------------------ + + +class FieldRule(TypedDict, total=False): + """字段校验规则(统一使用 path 定位)""" + + path: list[str] # 字段路径(统一定位方式) + required: bool = True # 是否必填(默认必填) + check_empty: bool = True # 是否检查空字符串(默认检查) + custom_check: str | Callable[[Any], bool] | None = None # 自定义校验函数(名称或函数) + description: str = "" # 字段描述(用于生成友好提示) + + +class ArrayRule(TypedDict, total=False): + """数组校验规则""" + + path: list[str] # 数组路径 + min_items: int = 1 # 最小元素数量 + element_fields: list[str | FieldRule] = [] # 元素字段规则 + description: str = "" # 数组描述 + + +class ValidationRules(TypedDict): + """校验规则集合""" + + fields: list[FieldRule] # 字段规则列表 + arrays: list[ArrayRule] # 数组规则列表 + + +class ValidationConfig(TypedDict): + """校验配置结构""" + + version: str + custom_checks: dict[str, str] + travel: ValidationRules + normal: ValidationRules + + +# ------------------------------------------------------------------ +# 配置加载 +# ------------------------------------------------------------------ + +_CONFIG_PATH = Path(__file__).parent.parent.parent / "config" / "validation_rules.json" +_cached_rules: ValidationConfig | None = None + + +def _load_validation_config() -> ValidationConfig: + """加载校验规则配置文件。""" + global _cached_rules + if _cached_rules is not None: + return _cached_rules + + if not _CONFIG_PATH.exists(): + log.warning("校验规则配置文件不存在: %s,使用内置默认规则", _CONFIG_PATH) + return _load_default_rules() + + try: + with open(_CONFIG_PATH, encoding="utf-8") as f: + config = json.load(f) + _cached_rules = _resolve_custom_checks(config) + log.info("校验规则配置加载成功") + return _cached_rules + except Exception as e: + log.error("加载校验规则配置失败: %s,使用内置默认规则", e) + return _load_default_rules() + + +def _resolve_custom_checks(config: dict[str, Any]) -> ValidationConfig: + """解析配置中的自定义校验函数名称,替换为实际函数引用。""" + + def resolve_rule(rule: dict[str, Any]) -> dict[str, Any]: + if "custom_check" in rule and isinstance(rule["custom_check"], str): + check_name = rule["custom_check"] + if check_name in _CUSTOM_CHECKS: + rule["custom_check"] = _CUSTOM_CHECKS[check_name] + else: + log.warning("未知的自定义校验函数: %s", check_name) + rule["custom_check"] = None + return rule + + # 解析 travel 规则的 fields + for field_rule in config.get("travel", {}).get("fields", []): + resolve_rule(field_rule) + # 解析 element_fields + for array_rule in config.get("travel", {}).get("arrays", []): + for elem_field in array_rule.get("element_fields", []): + if isinstance(elem_field, dict): + resolve_rule(elem_field) + + # 解析 normal 规则的 fields + for field_rule in config.get("normal", {}).get("fields", []): + resolve_rule(field_rule) + # 解析 element_fields + for array_rule in config.get("normal", {}).get("arrays", []): + for elem_field in array_rule.get("element_fields", []): + if isinstance(elem_field, dict): + resolve_rule(elem_field) + + return config # type: ignore[return-value] + + +def _load_default_rules() -> ValidationConfig: + """返回内置的默认校验规则(当配置文件不存在时使用)。""" + return { + "version": "1.0", + "custom_checks": {}, + "travel": { + "fields": [ + {"path": ["basic_info", "travel_purpose"], "description": "出差事由"}, + {"path": ["basic_info", "travel_location"], "description": "出差地点"}, + {"path": ["basic_info", "start_date"], "custom_check": _is_valid_date, "description": "出差开始日期"}, + {"path": ["basic_info", "end_date"], "custom_check": _is_valid_date, "description": "出差结束日期"}, + ], + "arrays": [ + { + "path": ["reimbursement_details", "transport_fee"], + "min_items": 1, + "element_fields": [ + {"path": ["vehicle_type"], "description": "交通工具类型"}, + {"path": ["start_date"], "custom_check": _is_valid_date, "description": "出发日期"}, + {"path": ["end_date"], "custom_check": _is_valid_date, "description": "到达日期"}, + {"path": ["departure_place"], "description": "出发地"}, + {"path": ["arrival_place"], "description": "目的地"}, + {"path": ["amount"], "custom_check": _is_positive_number, "description": "金额"}, + {"path": ["bill_count"], "custom_check": _is_positive_integer, "description": "票据张数"}, + {"path": ["remark"], "check_empty": False, "description": "备注说明"}, + ], + "description": "交通费用明细", + }, + { + "path": ["payment_methods"], + "min_items": 1, + "element_fields": [ + {"path": ["card_date"], "custom_check": _is_valid_date, "description": "刷卡日期"}, + {"path": ["card_amount"], "custom_check": _is_positive_number, "description": "支付金额"}, + {"path": ["merchant"], "description": "商户名称"}, + {"path": ["remark"], "check_empty": False, "description": "备注"}, + ], + "description": "支付方式记录", + }, + { + "path": ["subsidy_list"], + "min_items": 1, + "element_fields": [ + {"path": ["person_id"], "description": "人员工号"}, + {"path": ["person_name"], "description": "人员姓名"}, + {"path": ["start_date"], "custom_check": _is_valid_date, "description": "补助开始日期"}, + {"path": ["end_date"], "custom_check": _is_valid_date, "description": "补助结束日期"}, + {"path": ["days"], "custom_check": _is_positive_integer, "description": "补助天数"}, + ], + "description": "补助清单", + }, + { + "path": ["attachments"], + "min_items": 0, + "element_fields": [ + {"path": ["filename"], "description": "文件名"}, + {"path": ["attachment_type"], "description": "附件类型"}, + ], + "description": "附件列表", + }, + ], + }, + "normal": { + "fields": [ + {"path": ["basic_info", "reimbursement_description"], "description": "报销事由"}, + { + "path": ["reimbursement_details", "total_invoices"], + "custom_check": _is_positive_integer, + "description": "发票总数", + }, + { + "path": ["reimbursement_details", "total_amount"], + "custom_check": _is_positive_number, + "description": "总金额", + }, + ], + "arrays": [ + { + "path": ["payment_methods"], + "min_items": 1, + "element_fields": [ + {"path": ["card_date"], "custom_check": _is_valid_date, "description": "刷卡日期"}, + {"path": ["card_amount"], "custom_check": _is_positive_number, "description": "支付金额"}, + {"path": ["merchant"], "description": "商户名称"}, + {"path": ["remark"], "check_empty": False, "description": "备注"}, + ], + "description": "支付方式记录", + }, + { + "path": ["attachments"], + "min_items": 0, + "element_fields": [ + {"path": ["filename"], "description": "文件名"}, + {"path": ["attachment_type"], "description": "附件类型"}, + ], + "description": "附件列表", + }, + ], + }, + } + + +def get_validation_rules(invoice_type: str) -> ValidationRules: + """获取指定发票类型的校验规则。 + + Args: + invoice_type: 发票类型,'travel' 或 'normal'。 + + Returns: + 对应的校验规则。 + """ + config = _load_validation_config() + return config.get(invoice_type, config.get("travel", {})) # type: ignore[return-value] + + +# ------------------------------------------------------------------ +# 数据模型 +# ------------------------------------------------------------------ + + +@dataclass +class ValidationReport: + """校验结果报告""" + + valid: bool + missing_fields: list[str] = field(default_factory=list) + missing_materials: list[str] = field(default_factory=list) + confidence: float = 0.0 + suggestion: str = "" + + +# ------------------------------------------------------------------ +# 通用校验引擎 +# ------------------------------------------------------------------ + + +def _check_field( + data: dict[str, Any], + rule: FieldRule, +) -> tuple[bool, str]: + """根据字段规则检查字段。""" + path = rule["path"] + check_empty = rule.get("check_empty", True) + custom_check = rule.get("custom_check") + + current = data + for key in path: + if not isinstance(current, dict): + return (False, ".".join(path)) + if key not in current: + return (False, ".".join(path)) + current = current[key] + + if check_empty and isinstance(current, str) and not current.strip(): + return (False, ".".join(path)) + + if custom_check is not None and not custom_check(current): + return (False, ".".join(path)) + + return (True, ".".join(path)) + + +def _check_array( + data: dict[str, Any], + rule: ArrayRule, +) -> tuple[list[str], int, int]: + """根据数组规则检查数组。 + + Returns: + (缺失字段列表, 总检查数, 通过检查数) + """ + missing: list[str] = [] + path = rule["path"] + min_items = rule.get("min_items", 1) + element_fields = rule.get("element_fields", []) + path_str = ".".join(path) + + total_checks = 1 # 数组存在性和最小数量检查 + passed_checks = 0 + + # 遍历路径获取数组 + current = data + for key in path: + if not isinstance(current, dict) or key not in current: + return ([path_str], total_checks, passed_checks) + current = current[key] + + # 检查数组是否满足最小数量要求 + if not isinstance(current, list) or len(current) < min_items: + return ([path_str], total_checks, passed_checks) + + passed_checks += 1 # 数组检查通过 + + # 检查数组元素的字段 + if element_fields: + for i, item in enumerate(current): + if not isinstance(item, dict): + missing.append(f"{path_str}[{i}]") + total_checks += len(element_fields) + continue + + for field_rule in element_fields: + total_checks += 1 + # 支持两种格式:简单字符串格式 和 详细规则格式 + if isinstance(field_rule, str): + field_rule_dict: FieldRule = {"path": [field_rule]} + else: + field_rule_dict = field_rule + + # 复用 _check_field 函数检查元素字段 + ok, _ = _check_field(item, field_rule_dict) + if ok: + passed_checks += 1 + else: + field_path_str = ".".join(field_rule_dict["path"]) + missing.append(f"{path_str}[{i}].{field_path_str}") + + return missing, total_checks, passed_checks + + +def _validate_with_rules(data: dict[str, Any], rules: ValidationRules) -> tuple[list[str], int, int]: + """使用规则配置进行校验。""" + missing: list[str] = [] + total_checks = 0 + passed_checks = 0 + + # 校验字段规则 + for rule in rules.get("fields", []): + total_checks += 1 + ok, field_path = _check_field(data, rule) + if ok: + passed_checks += 1 + else: + missing.append(field_path) + + # 校验数组规则 + for rule in rules.get("arrays", []): + array_missing, array_total, array_passed = _check_array(data, rule) + total_checks += array_total + passed_checks += array_passed + missing.extend(array_missing) + + return missing, total_checks, passed_checks + + +# ------------------------------------------------------------------ +# 校验入口 +# ------------------------------------------------------------------ + + +def validate_travel_info(data: dict[str, Any]) -> ValidationReport: + """校验差旅报销信息的完整性。""" + rules = get_validation_rules("travel") + missing, total_checks, passed_checks = _validate_with_rules(data, rules) + + confidence = passed_checks / total_checks if total_checks > 0 else 0.0 + missing_materials = _infer_missing_materials(missing, data) + suggestion = _build_suggestion(missing, missing_materials) + + return ValidationReport( + valid=len(missing) == 0, + missing_fields=missing, + missing_materials=missing_materials, + confidence=round(confidence, 2), + suggestion=suggestion, + ) + + +def validate_normal_info(data: dict[str, Any]) -> ValidationReport: + """校验普通报销信息的完整性。""" + rules = get_validation_rules("normal") + missing, total_checks, passed_checks = _validate_with_rules(data, rules) + + confidence = passed_checks / total_checks if total_checks > 0 else 0.0 + missing_materials = _infer_missing_materials(missing, data) + suggestion = _build_suggestion(missing, missing_materials) + + return ValidationReport( + valid=len(missing) == 0, + missing_fields=missing, + missing_materials=missing_materials, + confidence=round(confidence, 2), + suggestion=suggestion, + ) + + +# ------------------------------------------------------------------ +# 缺失材料推断 +# ------------------------------------------------------------------ + + +def _infer_missing_materials( + missing_fields: list[str], + data: dict[str, Any], +) -> list[str]: + """根据缺失字段推断可能需要补充的材料类型。""" + materials: list[str] = [] + field_set = set(missing_fields) + + if any("start_date" in f or "end_date" in f for f in field_set): + if "basic_info.start_date" in field_set or "basic_info.end_date" in field_set: + materials.append("出差事前申请单") + + if "basic_info.travel_purpose" in field_set: + materials.append("出差事前申请单") + + if "basic_info.travel_location" in field_set: + materials.append("交通工具发票") + + if "payment_methods" in field_set or any("payment_methods[" in f for f in field_set): + materials.append("支付记录截图") + + if any("transport_fee" in f for f in field_set): + materials.append("交通工具发票") + + if any("subsidy_list" in f for f in field_set): + materials.append("出差事前申请单") + + if "basic_info.reimbursement_description" in field_set: + materials.append("发票或支付记录") + + return list(dict.fromkeys(materials)) + + +def _build_suggestion( + missing_fields: list[str], + missing_materials: list[str], +) -> str: + """生成用户友好的建议信息。""" + if not missing_fields: + return "" + + if missing_materials: + material_names = "、".join(missing_materials) + return f"信息不完整,请补充上传:{material_names}" + + return f"信息不完整,缺少 {len(missing_fields)} 个字段" + + +# ------------------------------------------------------------------ +# 统一入口 +# ------------------------------------------------------------------ + + +def validate_extracted_info( + data: dict[str, Any], + invoice_type: str = "travel", +) -> ValidationReport: + """校验提取信息的完整性。 + + Args: + data: LLM 提取的结构化信息。 + invoice_type: 发票类型,'travel' 或 'normal'。 + + Returns: + 校验报告。 + """ + log.info("开始校验 %s 报销信息完整性", invoice_type) + + if invoice_type == "travel": + report = validate_travel_info(data) + else: + report = validate_normal_info(data) + + status = "通过" if report.valid else "未通过" + log.info( + "校验结果: %s (置信度: %.0f%%, 缺失字段: %d)", + status, + report.confidence * 100, + len(report.missing_fields), + ) + + return report + + +def reload_validation_rules() -> None: + """重新加载校验规则配置(用于运行时热更新)。""" + global _cached_rules + _cached_rules = None + _load_validation_config() + log.info("校验规则已重新加载") diff --git a/src/doc/README.md b/src/doc/README.md deleted file mode 100644 index bdfca8a..0000000 --- a/src/doc/README.md +++ /dev/null @@ -1,69 +0,0 @@ ---- - -## last_reviewed: 2026-06-12 - -# src/doc — 文档处理模块 - -负责发票信息提取、基于 LLM 的支付截图信息识别、差旅/普通报销信息提取、以及将数据填入 Word 出库单模板。 - -## 模块清单 - - -| 文件 | 作用 | -| ------------------------ | ---------------------------------------------------------------------------- | -| `extractor.py` | 编排入口:串联 PDF 读取 → LLM 提取 → 支付截图匹配 → 分类 | -| `pdf.py` | PDF 图片渲染(PyMuPDF) | -| `llm_extractor.py` | 基于 LLM 的信息提取(发票文本 + 支付截图多模态 + 差旅/普通报销信息综合提取) | -| `matcher.py` | 发票与支付截图按金额匹配,回填刷卡信息至发票记录 | -| `invoice.py` | 发票类型常量、分类逻辑、CSV 读写工具 | -| `fill_consumable_doc.py` | 将 CSV 数据填入易耗品出库单 Word 模板(pywin32 COM) | -| `prompt.py` | LLM 提示词模板加载 | -| `prompts/` | 提示词模板文件(`invoice_system.md`、`travel_info_system.md`、`normal_info_system.md`) | - - -## 数据流 - -```mermaid -flowchart TD - A["PDF 发票"] --> B["pdf.py
PDF 图片渲染"] - B --> C["llm_extractor.py
发票文本提取"] - C --> D["(发票列表)"] - - E["支付截图"] --> F["llm_extractor.py
多模态提取"] - F --> G["(刷卡记录)"] - - D --> H["matcher.py
金额贪心匹配 / 相对容差 3%"] - G --> H - - H --> I["invoice.py
分类 + CSV 回填"] - I --> J["CSV
刷卡日期/卡号/金额"] - - J --> K["fill_consumable_doc
易耗品出库单"] - K --> L["易耗品出库单.doc"] - - D --> M["llm_extractor.py
差旅/普通信息综合提取"] - H --> M - - M --> N{"发票类型判断"} - N -->|差旅发票| O["extract_travel_info()"] - N -->|普通发票| P["extract_normal_info()"] - - O --> Q["travel_info.json"] - P --> R["normal_info.json"] -``` - - - -## 依赖说明 - -- **PyMuPDF (pymupdf)** — PDF 图片渲染 -- **pywin32** — Word COM 自动化(仅 Windows) -- **llama-index** — LLM 信息提取 - -## 注意事项 - -- `fill_consumable_doc.py` 依赖 Microsoft Word + COM,仅 Windows 可用 -- LLM 提取不会覆盖 CSV 中已有非空字段 -- 提示词模板位于 `prompts/` 目录,由 `prompt.py` 加载 -- LLM 提取失败时直接报错,无正则回退 - diff --git a/src/doc/__init__.py b/src/doc/__init__.py deleted file mode 100644 index 97668d1..0000000 --- a/src/doc/__init__.py +++ /dev/null @@ -1,4 +0,0 @@ -"""文档处理模块 - -包含发票提取、LLM 信息提取、出库单填写等功能。 -""" diff --git a/src/doc/prompts/validation_system.md b/src/doc/prompts/validation_system.md deleted file mode 100644 index 6b455f0..0000000 --- a/src/doc/prompts/validation_system.md +++ /dev/null @@ -1,83 +0,0 @@ -#性格 - -你是财务报销信息完整性校验助手。你的任务是检查已提取的报销信息在语义上是否足够支撑完成报销系统填报。 - -**核心原则**:不仅要检查字段是否存在,还要判断信息在逻辑上是否自洽、是否足以完成填报。 - ---- - -## 输入数据说明 - -你会收到以下数据: -1. **已提取的报销信息**:包含基本信息、报销明细、支付方式、附件清单等 -2. **发票缓存数据**:原始发票提取的结构化数据 -3. **匹配结果**:发票与支付记录的关联关系 - ---- - -## 校验维度 - -### 1. 逻辑一致性检查 - -- 出差日期范围是否合理(结束日期不早于开始日期) -- 交通费的去程和返程日期是否在出差日期范围内 -- 酒店入住/退房日期是否与出差时间匹配 -- 支付金额总和是否与发票金额总和接近(允许小额差异) -- 补助天数计算是否正确 - -### 2. 信息充分性检查 - -- 是否有足够的信息填写所有报销系统必填项 -- 出差事由是否明确具体(不能过于笼统) -- 人员信息是否完整(姓名、工号) -- 支付方式是否与支付记录对应 -- 每张发票应该都有对应的支付记录,如果没有提醒用户补充支付记录 - -### 3. 潜在问题识别 - -- 发票日期与出差日期差异过大 -- 同一笔支付对应多张发票但金额不匹配 -- 缺少关键附件 -- 人员信息不一致(如车票姓名与补助清单姓名不同) - ---- -### 4. 不需询问的问题 -- 出差事前申请单和实际出差时间不一致这是正常的,因为规划是实际可以有差异 - - -## 输出格式 - -严格返回以下 JSON 格式: - -```json -{ - "valid": true或false, - "confidence": 0.0到1.0之间的数字, - "issues": ["问题描述1", "问题描述2"], - "missing_info": ["缺失信息1", "缺失信息2"], - "suggestion": "补充建议" -} -``` - -### 字段说明 - -| 字段 | 类型 | 说明 | -|------|------|------| -| `valid` | boolean | 信息在语义上是否足够支撑填报 | -| `confidence` | number | 置信度,1.0 表示完全确定,0.0 表示完全不确定 | -| `issues` | array | 发现的逻辑问题列表,无问题则为空数组 | -| `missing_info` | array | 语义上缺失的关键信息列表 | -| `suggestion` | string | 补充建议,说明需要用户上传什么材料 | - -### 判定标准 - -- **valid = true**:信息完整且逻辑自洽,可以直接填报 -- **valid = false**:存在信息缺失或逻辑矛盾,需要补充材料 - ---- - -## 最终输出要求 - -- 严格只输出 JSON 字符串 -- JSON 语法必须正确 -- 不要包含任何思考过程或解释文字 \ No newline at end of file diff --git a/src/doc/validator.py b/src/doc/validator.py deleted file mode 100644 index c37245a..0000000 --- a/src/doc/validator.py +++ /dev/null @@ -1,430 +0,0 @@ -"""信息完整性校验器 - -对 LLM 提取的报销信息进行规则级校验,判断是否满足填报要求。 - -校验规则基于两个 schema: - - 差旅报销:basic_info + reimbursement_details + payment_methods + subsidy_list + attachments - - 普通报销:basic_info + reimbursement_details + payment_methods + attachments - -校验结果包含缺失字段列表、缺失材料推断和建议信息。 -""" - -from __future__ import annotations - -import re -from dataclasses import dataclass, field -from typing import Any - -from .. import get_logger - -log = get_logger("validator") - -# ------------------------------------------------------------------ -# 日期格式 -# ------------------------------------------------------------------ - -DATE_PATTERN = re.compile(r"^\d{4}-\d{2}-\d{2}$") - - -def _is_valid_date(value: str) -> bool: - return bool(DATE_PATTERN.match(value)) - - -# ------------------------------------------------------------------ -# 数据模型 -# ------------------------------------------------------------------ - - -@dataclass -class ValidationReport: - """校验结果报告""" - - valid: bool - missing_fields: list[str] = field(default_factory=list) - missing_materials: list[str] = field(default_factory=list) - confidence: float = 0.0 - suggestion: str = "" - - -# ------------------------------------------------------------------ -# 校验辅助 -# ------------------------------------------------------------------ - - -def _check_field( - data: dict[str, Any], - path: list[str], - check_empty: bool = True, - custom_check: Any = None, -) -> tuple[bool, str]: - """沿路径检查字段是否存在且有效。 - - 检查字段是否存在于嵌套字典中,可选的检查是否为空字符串或自定义校验。 - - Returns: - (通过, 字段路径字符串) - """ - current = data - for key in path: - if not isinstance(current, dict): - return (False, ".".join(path)) - if key not in current: - return (False, ".".join(path)) - current = current[key] - - if check_empty and isinstance(current, str) and not current.strip(): - return (False, ".".join(path)) - - if custom_check is not None and not custom_check(current): - return (False, ".".join(path)) - - return (True, ".".join(path)) - - -def _check_array_min( - data: dict[str, Any], - path: list[str], - min_items: int = 1, -) -> tuple[bool, str]: - """检查数组字段是否存在且至少有 min_items 项。""" - current = data - for key in path: - if not isinstance(current, dict): - return (False, ".".join(path)) - if key not in current: - return (False, ".".join(path)) - current = current[key] - - if not isinstance(current, list) or len(current) < min_items: - return (False, ".".join(path)) - - return (True, ".".join(path)) - - -def _check_array_element_fields( - items: Any, - required_fields: list[str], - field_path_prefix: str, -) -> list[str]: - """检查数组中每个元素是否包含必填字段。""" - missing: list[str] = [] - if not isinstance(items, list): - return [field_path_prefix] - - for i, item in enumerate(items): - if not isinstance(item, dict): - missing.append(f"{field_path_prefix}[{i}]") - continue - for fld in required_fields: - if fld not in item or (isinstance(item[fld], str) and not item[fld].strip()): - missing.append(f"{field_path_prefix}[{i}].{fld}") - - return missing - - -def _is_positive_number(value: Any) -> bool: - """检查值是否为正数。""" - return isinstance(value, int | float) and value > 0 - - -# ------------------------------------------------------------------ -# 差旅报销校验 -# ------------------------------------------------------------------ - - -def validate_travel_info(data: dict[str, Any]) -> ValidationReport: - """校验差旅报销信息的完整性。""" - missing: list[str] = [] - total_checks = 0 - passed_checks = 0 - - # basic_info 必填字段 - basic_fields = [ - "travel_purpose", - "travel_location", - "start_date", - "end_date", - ] - - for fld in basic_fields: - total_checks += 1 - path = ["basic_info", fld] - if fld in ("start_date", "end_date"): - ok, _ = _check_field(data, path, custom_check=_is_valid_date) - else: - ok, _ = _check_field(data, path) - if ok: - passed_checks += 1 - else: - missing.append(".".join(path)) - - # reimbursement_details.transport_fee (至少一条) - total_checks += 1 - ok, _ = _check_array_min(data, ["reimbursement_details", "transport_fee"], min_items=1) - if ok: - passed_checks += 1 - else: - missing.append("reimbursement_details.transport_fee") - - # 检查 transport_fee 元素字段 - transport = data.get("reimbursement_details", {}).get("transport_fee", []) - transport_fields = [ - "vehicle_type", - "start_date", - "end_date", - "departure_place", - "arrival_place", - "amount", - "bill_count", - "remark", - ] - missing.extend( - _check_array_element_fields( - transport, - transport_fields, - "reimbursement_details.transport_fee", - ) - ) - - # payment_methods (至少一条) - total_checks += 1 - ok, _ = _check_array_min(data, ["payment_methods"], min_items=1) - if ok: - passed_checks += 1 - else: - missing.append("payment_methods") - - # 检查 payment_methods 元素 - payments = data.get("payment_methods", []) - payment_fields = ["card_date", "card_amount", "merchant", "remark"] - missing.extend( - _check_array_element_fields( - payments, - payment_fields, - "payment_methods", - ) - ) - - # subsidy_list (至少一条) - total_checks += 1 - ok, _ = _check_array_min(data, ["subsidy_list"], min_items=1) - if ok: - passed_checks += 1 - else: - missing.append("subsidy_list") - - # 检查 subsidy_list 元素 - subsidies = data.get("subsidy_list", []) - subsidy_fields = ["person_id", "person_name", "start_date", "end_date", "days"] - missing.extend( - _check_array_element_fields( - subsidies, - subsidy_fields, - "subsidy_list", - ) - ) - - # attachments (可选,有数据时检查) - attachments = data.get("attachments", []) - if attachments: - attach_fields = ["filename", "attachment_type"] - missing.extend( - _check_array_element_fields( - attachments, - attach_fields, - "attachments", - ) - ) - - # 计算置信度 - confidence = passed_checks / total_checks if total_checks > 0 else 0.0 - - # 推断缺失材料 - missing_materials = _infer_missing_materials(missing, data) - - # 生成建议 - suggestion = _build_suggestion(missing, missing_materials) - - return ValidationReport( - valid=len(missing) == 0, - missing_fields=missing, - missing_materials=missing_materials, - confidence=round(confidence, 2), - suggestion=suggestion, - ) - - -# ------------------------------------------------------------------ -# 普通报销校验 -# ------------------------------------------------------------------ - - -def validate_normal_info(data: dict[str, Any]) -> ValidationReport: - """校验普通报销信息的完整性。""" - missing: list[str] = [] - total_checks = 0 - passed_checks = 0 - - # basic_info.reimbursement_description - total_checks += 1 - ok, _ = _check_field(data, ["basic_info", "reimbursement_description"]) - if ok: - passed_checks += 1 - else: - missing.append("basic_info.reimbursement_description") - - # reimbursement_details.total_invoices - total_checks += 1 - ok, _ = _check_field( - data, - ["reimbursement_details", "total_invoices"], - custom_check=lambda v: isinstance(v, int) and v > 0, - ) - if ok: - passed_checks += 1 - else: - missing.append("reimbursement_details.total_invoices") - - # reimbursement_details.total_amount - total_checks += 1 - ok, _ = _check_field( - data, - ["reimbursement_details", "total_amount"], - custom_check=_is_positive_number, - ) - if ok: - passed_checks += 1 - else: - missing.append("reimbursement_details.total_amount") - - # payment_methods (至少一条) - total_checks += 1 - ok, _ = _check_array_min(data, ["payment_methods"], min_items=1) - if ok: - passed_checks += 1 - else: - missing.append("payment_methods") - - payments = data.get("payment_methods", []) - payment_fields = ["card_date", "card_amount", "merchant", "remark"] - missing.extend( - _check_array_element_fields( - payments, - payment_fields, - "payment_methods", - ) - ) - - # attachments - attachments = data.get("attachments", []) - if attachments: - attach_fields = ["filename", "attachment_type"] - missing.extend( - _check_array_element_fields( - attachments, - attach_fields, - "attachments", - ) - ) - - confidence = passed_checks / total_checks if total_checks > 0 else 0.0 - missing_materials = _infer_missing_materials(missing, data) - suggestion = _build_suggestion(missing, missing_materials) - - return ValidationReport( - valid=len(missing) == 0, - missing_fields=missing, - missing_materials=missing_materials, - confidence=round(confidence, 2), - suggestion=suggestion, - ) - - -# ------------------------------------------------------------------ -# 缺失材料推断 -# ------------------------------------------------------------------ - - -def _infer_missing_materials( - missing_fields: list[str], - data: dict[str, Any], -) -> list[str]: - """根据缺失字段推断可能需要补充的材料类型。""" - materials: list[str] = [] - field_set = set(missing_fields) - - if any("start_date" in f or "end_date" in f for f in field_set): - if "basic_info.start_date" in field_set or "basic_info.end_date" in field_set: - materials.append("出差事前申请单") - - if "basic_info.travel_purpose" in field_set: - materials.append("出差事前申请单") - - if "basic_info.travel_location" in field_set: - materials.append("交通工具发票") - - if "payment_methods" in field_set or any("payment_methods[" in f for f in field_set): - materials.append("支付记录截图") - - if any("transport_fee" in f for f in field_set): - materials.append("交通工具发票") - - if any("subsidy_list" in f for f in field_set): - materials.append("出差事前申请单") - - if "basic_info.reimbursement_description" in field_set: - materials.append("发票或支付记录") - - # 去重 - return list(dict.fromkeys(materials)) - - -def _build_suggestion( - missing_fields: list[str], - missing_materials: list[str], -) -> str: - """生成用户友好的建议信息。""" - if not missing_fields: - return "" - - if missing_materials: - material_names = "、".join(missing_materials) - return f"信息不完整,请补充上传:{material_names}" - - return f"信息不完整,缺少 {len(missing_fields)} 个字段" - - -# ------------------------------------------------------------------ -# 统一入口 -# ------------------------------------------------------------------ - - -def validate_extracted_info( - data: dict[str, Any], - invoice_type: str = "travel", -) -> ValidationReport: - """校验提取信息的完整性。 - - Args: - data: LLM 提取的结构化信息。 - invoice_type: 发票类型,'travel' 或 'normal'。 - - Returns: - 校验报告。 - """ - log.info("开始校验 %s 报销信息完整性", invoice_type) - - if invoice_type == "travel": - report = validate_travel_info(data) - else: - report = validate_normal_info(data) - - status = "通过" if report.valid else "未通过" - log.info( - "校验结果: %s (置信度: %.0f%%, 缺失字段: %d)", - status, - report.confidence * 100, - len(report.missing_fields), - ) - - return report diff --git a/src/infra/README.md b/src/infra/README.md new file mode 100644 index 0000000..0627e75 --- /dev/null +++ b/src/infra/README.md @@ -0,0 +1,21 @@ +--- +last_reviewed: 2026-06-15 +--- + +# src/infra — 基础设施层 + +提供浏览器自动化、文档处理和 LLM 接口等底层能力。此层不包含业务逻辑,只提供工具和平台能力。 + +## 子模块 + +| 目录 | 说明 | +|------|------| +| `browser/` | Playwright 驱动的财务系统自动填报 | +| `documents/` | 发票数据模型、PDF 渲染、Word 出库单填写 | +| `llm/` | LLM 提示词模板加载与管理 | + +## 设计原则 + +- **无业务逻辑**:只提供工具能力,不包含业务流程判断 +- **可替换性**:每个子模块通过 `__init__.py` 导出接口,便于替换实现 +- **与 core 层解耦**:infra 不依赖 core,core 可通过接口调用 infra diff --git a/src/infra/__init__.py b/src/infra/__init__.py new file mode 100644 index 0000000..171d146 --- /dev/null +++ b/src/infra/__init__.py @@ -0,0 +1,4 @@ +"""基础设施模块 + +提供浏览器自动化、文档处理和 LLM 接口功能。 +""" diff --git a/src/infra/browser/README.md b/src/infra/browser/README.md new file mode 100644 index 0000000..1e9e16d --- /dev/null +++ b/src/infra/browser/README.md @@ -0,0 +1,44 @@ +--- +last_reviewed: 2026-06-15 +--- + +# src/infra/browser — 浏览器自动化 + +使用 Playwright 操作财务报销系统,自动完成登录、填单、上传附件等操作。 + +## 文件 + +| 文件 | 职责 | +|------|------| +| `base.py` | `BaseBot` 基类:浏览器生命周期、登录信息门户、导航到报销系统、创建新单据、截图 | +| `travel.py` | 差旅报销填报流程:基本信息 → 差旅明细 → 支付方式 → 补助清单 → 附件上传 | +| `normal.py` | 普通报销填报流程:基本信息 → 总明细 → 支付方式 → 附件上传 | +| `__init__.py` | 入口函数:`run_bot()` / `run_bot_web()`,负责类型路由和流程调度 | + +## 对外接口 + +| 函数 | 说明 | +|------|------| +| `run_bot(config, travel_info, normal_info)` | CLI 模式:根据传入信息判断差旅/普通报销 | +| `run_bot_web(config, work_dir)` | Web 模式:从缓存加载信息后执行填报 | + +## 填报流程 + +### 差旅报销(travel) +1. 填写基本信息(事由、地点、日期、项目编号) +2. 添加差旅明细(交通费用逐条录入) +3. 填写支付方式(公务卡刷卡记录) +4. 填写补助清单(按天计算交通补助 + 伙食补助) +5. 上传附件(发票、申请单等) + +### 普通报销(normal) +1. 填写基本信息(报销事由、金额) +2. 填写发票明细(总数、总金额) +3. 填写支付方式 +4. 上传附件 + +## 注意事项 + +- 浏览器填报会启动 Chromium,请勿手动干扰自动化流程 +- 调试截图保存在 `images/` 目录 +- Web 模式以无头模式运行 diff --git a/src/infra/browser/__init__.py b/src/infra/browser/__init__.py new file mode 100644 index 0000000..2cd3de8 --- /dev/null +++ b/src/infra/browser/__init__.py @@ -0,0 +1,100 @@ +"""浏览器自动化填报 + +使用 Playwright 操作财务报销系统,自动完成登录、填单、上传附件等操作。 + +对外接口: + run_bot(config, travel_info, normal_info) 启动浏览器并执行填报流程 + run_bot_web(config, work_dir) Web 模式填报(从缓存加载信息) +""" + +from pathlib import Path +from typing import Any + +from ... import get_logger +from .base import BaseBot + +log = get_logger("bot") + + +def run_bot( + config: dict[str, Any], + headless: bool = False, + work_dir: Path | None = None, + travel_info: dict[str, Any] | None = None, + normal_info: dict[str, Any] | None = None, +) -> None: + """启动浏览器并执行填报流程。 + + 根据传入的报销信息判断执行差旅报销还是普通报销流程。 + + Args: + config: 财务系统配置(含 URL、账号密码等)。 + headless: 是否无头模式。 + work_dir: 工作目录。 + travel_info: 差旅报销信息(可选)。 + normal_info: 普通报销信息(可选)。 + """ + invoice_type = "" + if travel_info and normal_info: + invoice_type = "mixed" + elif travel_info: + invoice_type = "travel" + elif normal_info: + invoice_type = "normal" + else: + log.error("未提供任何报销信息") + return + + log.info(f"启动填报流程: {invoice_type}") + + if invoice_type == "travel": + from .travel import run as run_travel + + bot = BaseBot(config, headless=headless) + bot.work_dir = work_dir + + try: + bot.launch() + bot.login_portal() + bot.navigate_to_reimburse(page_key="travel_page") + bot.create_new_form() + run_travel(bot, travel_info) + finally: + bot.close() + elif invoice_type == "normal": + from .normal import run as run_normal + + bot = BaseBot(config, headless=headless) + bot.work_dir = work_dir + + try: + bot.launch() + bot.login_portal() + bot.navigate_to_reimburse(page_key="reimburse_page") + bot.create_new_form() + run_normal(bot, normal_info) + finally: + bot.close() + else: + log.warning("暂不支持混合报销流程") + + +def run_bot_web(config: dict[str, Any], work_dir: str | Path) -> None: + """Web 模式填报(从缓存加载信息)。 + + 根据 work_dir 下的 .invoice_cache 目录中已提取的信息, + 自动判断执行差旅报销还是普通报销流程。 + + Args: + config: 财务系统配置(含 URL、账号密码等)。 + work_dir: 工作目录(包含 .invoice_cache 子目录)。 + """ + from ...core.extraction import load_cache + + work_dir = Path(work_dir) + cache = load_cache(work_dir) + + travel_info = cache.get("travel_info") + normal_info = cache.get("normal_info") + + run_bot(config, headless=True, work_dir=work_dir, travel_info=travel_info, normal_info=normal_info) diff --git a/src/bot/base.py b/src/infra/browser/base.py similarity index 98% rename from src/bot/base.py rename to src/infra/browser/base.py index bfd81e4..927e3d5 100644 --- a/src/bot/base.py +++ b/src/infra/browser/base.py @@ -1,5 +1,4 @@ -""" -浏览器自动化填报 — 公共基类 +"""浏览器自动化填报 — 公共基类 提供浏览器生命周期管理、登录、导航、截图等公共操作。 """ @@ -7,7 +6,7 @@ from pathlib import Path from typing import Any -from .. import get_logger +from ... import get_logger log = get_logger("bot") diff --git a/src/bot/normal.py b/src/infra/browser/normal.py similarity index 99% rename from src/bot/normal.py rename to src/infra/browser/normal.py index a37d162..c8259cb 100644 --- a/src/bot/normal.py +++ b/src/infra/browser/normal.py @@ -1,5 +1,4 @@ -""" -普通报销填报流程 +"""普通报销填报流程 负责普通发票报销的完整填报步骤: 基本信息 → 总明细 → 支付方式 → 附件上传 @@ -7,7 +6,7 @@ from typing import Any -from .. import get_logger +from ... import get_logger from .base import BaseBot, format_date log = get_logger("bot.normal") diff --git a/src/bot/travel.py b/src/infra/browser/travel.py similarity index 98% rename from src/bot/travel.py rename to src/infra/browser/travel.py index cc29898..0c3a34e 100644 --- a/src/bot/travel.py +++ b/src/infra/browser/travel.py @@ -1,5 +1,4 @@ -""" -差旅报销填报流程 +"""差旅报销填报流程 负责差旅报销的完整填报步骤: 基本信息 → 差旅明细 → 支付方式 → 补助清单 → 附件上传 @@ -7,7 +6,7 @@ from typing import Any -from .. import get_logger +from ... import get_logger from .base import BaseBot, format_date log = get_logger("bot.travel") @@ -294,7 +293,7 @@ def upload_travel_attachments(bot: BaseBot, attachment_info: list[dict[str, Any] bot.page.select_option("#fjlx", "1") else: bot.page.select_option("#fjlx", "2") - bot.page.fill("#fpsmxx", info["attachment_desc"]) + bot.page.fill("#fpsmxx", info.get("attachment_desc", "")) if attachment_file and attachment_file.exists(): bot.page.set_input_files("#file", str(attachment_file)) bot.page.wait_for_timeout(1000) @@ -304,4 +303,5 @@ def upload_travel_attachments(bot: BaseBot, attachment_info: list[dict[str, Any] log.error(f"差旅附件上传失败: {e}") bot._screenshot("travel_attachment_error") raise + bot._screenshot("travel_attachment_done") diff --git a/src/infra/documents/README.md b/src/infra/documents/README.md new file mode 100644 index 0000000..0bc8e66 --- /dev/null +++ b/src/infra/documents/README.md @@ -0,0 +1,36 @@ +--- +last_reviewed: 2026-06-15 +--- + +# src/infra/documents — 文档处理 + +提供发票数据模型、PDF 渲染和 Word 出库单填写功能。 + +## 文件 + +| 文件 | 职责 | +|------|------| +| `invoice.py` | 发票数据模型:类型常量、CSV 列定义、CSV/JSON 读写工具、发票分类 | +| `pdf.py` | PDF 渲染为图片(PyMuPDF),供多模态 LLM 识别使用 | +| `consumable.py` | 易耗品出库单填写:读取 CSV → 填入 Word 模板(pywin32 COM,仅 Windows) | + +## 对外接口 + +| 函数 | 说明 | +|------|------| +| `load_csv(path)` | 读取支付记录 CSV | +| `save_csv(payment_records, path)` | 保存支付记录 CSV | +| `save_invoice_csv(payment_records, path)` | 保存发票级别 CSV | +| `classify_invoice_batch(cache_map)` | 按类型批量分类发票 | +| `render_pdf_to_images(pdf_path)` | PDF → 图片列表 | +| `fill_consumable_doc(csv_path, doc_path)` | 将 CSV 数据填入 Word 模板 | + +## 缓存目录 + +`.invoice_cache/` 是系统级缓存目录名常量,定义在 `invoice.py` 中,被提取和匹配模块统一引用。 + +## 易耗品出库单 + +- 需要 **Windows + Microsoft Word + pywin32** +- 模板文件为项目根目录的 `易耗品、出库单.doc` +- 填写规则:日期用当天日期,品名/规格/数量/单价从 CSV 解析,字体统一宋体五号 diff --git a/src/infra/documents/__init__.py b/src/infra/documents/__init__.py new file mode 100644 index 0000000..1eb5675 --- /dev/null +++ b/src/infra/documents/__init__.py @@ -0,0 +1,32 @@ +"""文档处理基础设施 + +提供发票数据模型、PDF 渲染、出库单填写等功能。 +""" + +from .consumable import ( + CONSUMABLE_DOC_FILENAME, + fill_consumable_doc, + fill_consumable_from_template, +) +from .invoice import ( + CACHE_DIR_NAME, + classify_invoice_batch, + load_csv, + load_invoice_csv, + save_application_json, + save_csv, + save_invoice_csv, +) + +__all__ = [ + "CACHE_DIR_NAME", + "classify_invoice_batch", + "load_csv", + "load_invoice_csv", + "save_csv", + "save_invoice_csv", + "save_application_json", + "CONSUMABLE_DOC_FILENAME", + "fill_consumable_doc", + "fill_consumable_from_template", +] diff --git a/src/doc/fill_consumable_doc.py b/src/infra/documents/consumable.py similarity index 97% rename from src/doc/fill_consumable_doc.py rename to src/infra/documents/consumable.py index 310ed1d..40895e3 100644 --- a/src/doc/fill_consumable_doc.py +++ b/src/infra/documents/consumable.py @@ -1,5 +1,4 @@ -""" -将 invoice_summary.csv 填入「易耗品、出库单.doc」表格。 +"""将 invoice_summary.csv 填入「易耗品、出库单.doc」表格。 仅写入表格数据单元格,保留原模板字体、边框与版式。 """ @@ -13,9 +12,9 @@ from datetime import date from pathlib import Path from typing import Any -from .. import get_logger -from ..config import load_config -from ..doc.invoice import load_invoice_csv +from ... import get_logger +from ...config import load_config +from .invoice import load_invoice_csv log = get_logger("fill_consumable_doc") diff --git a/src/doc/invoice.py b/src/infra/documents/invoice.py similarity index 99% rename from src/doc/invoice.py rename to src/infra/documents/invoice.py index a466d24..a7bfa3a 100644 --- a/src/doc/invoice.py +++ b/src/infra/documents/invoice.py @@ -13,7 +13,7 @@ import csv import json from pathlib import Path -from .. import get_logger +from ... import get_logger log = get_logger("invoice") diff --git a/src/doc/pdf.py b/src/infra/documents/pdf.py similarity index 97% rename from src/doc/pdf.py rename to src/infra/documents/pdf.py index 4928fc7..17a7698 100644 --- a/src/doc/pdf.py +++ b/src/infra/documents/pdf.py @@ -11,7 +11,7 @@ from pathlib import Path import fitz -from .. import get_logger +from ... import get_logger log = get_logger("pdf") diff --git a/src/infra/llm/README.md b/src/infra/llm/README.md new file mode 100644 index 0000000..5eb0bfd --- /dev/null +++ b/src/infra/llm/README.md @@ -0,0 +1,32 @@ +--- +last_reviewed: 2026-06-15 +--- + +# src/infra/llm — LLM 提示词管理 + +管理 LLM 提示词模板的加载,供 `core/extraction/llm_extractor.py` 调用。 + +## 文件 + +| 文件 | 职责 | +|------|------| +| `prompt.py` | 提示词加载:从 `prompts/` 目录读取 `.md` 模板文件 | +| `prompts/` | 提示词模板目录(Markdown 格式) | + +## 提示词模板 + +| 文件 | 用途 | +|------|------| +| `invoice_system.md` | 发票提取系统提示词 | +| `travel_info_system.md` | 差旅信息提取系统提示词 | +| `normal_info_system.md` | 普通发票信息提取系统提示词 | +| `supplement_system.md` | 用户补充信息后的二次提取提示词 | +| `validation_system.md` | 校验修正提示词 | + +## 对外接口 + +| 函数 | 说明 | +|------|------| +| `build_invoice_system_prompt()` | 构建发票提取系统提示词 | +| `build_travel_info_system_prompt()` | 构建差旅信息提取系统提示词 | +| `build_normal_info_system_prompt()` | 构建普通发票信息提取系统提示词 | diff --git a/src/infra/llm/__init__.py b/src/infra/llm/__init__.py new file mode 100644 index 0000000..f435c51 --- /dev/null +++ b/src/infra/llm/__init__.py @@ -0,0 +1,18 @@ +"""LLM 接口模块 + +提供 LLM 提示词模板加载功能。 +""" + +from .prompt import ( + build_invoice_system_prompt, + build_normal_info_system_prompt, + build_supplement_system_prompt, + build_travel_info_system_prompt, +) + +__all__ = [ + "build_invoice_system_prompt", + "build_normal_info_system_prompt", + "build_supplement_system_prompt", + "build_travel_info_system_prompt", +] diff --git a/src/doc/prompt.py b/src/infra/llm/prompt.py similarity index 79% rename from src/doc/prompt.py rename to src/infra/llm/prompt.py index 026fe90..eb87f2d 100644 --- a/src/doc/prompt.py +++ b/src/infra/llm/prompt.py @@ -1,7 +1,6 @@ -""" -LLM 提示词模板 +"""LLM 提示词模板 -从 src/prompts/ 目录加载 .md 文件作为提示词模板。 +从 infra/llm/prompts/ 目录加载 .md 文件作为提示词模板。 """ import os @@ -34,8 +33,3 @@ def build_normal_info_system_prompt() -> str: def build_supplement_system_prompt() -> str: """构建用户补充信息分析系统提示词。""" return _load_prompt("supplement_system.md") - - -def build_validation_system_prompt() -> str: - """构建语义校验系统提示词。""" - return _load_prompt("validation_system.md") diff --git a/src/doc/prompts/README.md b/src/infra/llm/prompts/README.md similarity index 74% rename from src/doc/prompts/README.md rename to src/infra/llm/prompts/README.md index 511de36..68f0a0e 100644 --- a/src/doc/prompts/README.md +++ b/src/infra/llm/prompts/README.md @@ -2,9 +2,9 @@ last_reviewed: 2026-06-12 --- -# src/doc/prompts — LLM 提示词模板 +# src/infra/llm/prompts — LLM 提示词模板 -存放 LLM 信息提取使用的系统提示词模板文件,由 `src/doc/prompt.py` 动态加载。 +存放 LLM 信息提取使用的系统提示词模板文件,由 `src/infra/llm/prompt.py` 动态加载。 ## 模板清单 @@ -16,5 +16,5 @@ last_reviewed: 2026-06-12 ## 加载方式 ```python -from src.doc.prompt import build_invoice_system_prompt, build_travel_info_system_prompt +from src.infra.llm import build_invoice_system_prompt, build_travel_info_system_prompt ``` \ No newline at end of file diff --git a/src/doc/prompts/invoice_system.md b/src/infra/llm/prompts/invoice_system.md similarity index 100% rename from src/doc/prompts/invoice_system.md rename to src/infra/llm/prompts/invoice_system.md diff --git a/src/doc/prompts/normal_info_system.md b/src/infra/llm/prompts/normal_info_system.md similarity index 100% rename from src/doc/prompts/normal_info_system.md rename to src/infra/llm/prompts/normal_info_system.md diff --git a/src/doc/prompts/supplement_system.md b/src/infra/llm/prompts/supplement_system.md similarity index 100% rename from src/doc/prompts/supplement_system.md rename to src/infra/llm/prompts/supplement_system.md diff --git a/src/doc/prompts/travel_info_system.md b/src/infra/llm/prompts/travel_info_system.md similarity index 100% rename from src/doc/prompts/travel_info_system.md rename to src/infra/llm/prompts/travel_info_system.md diff --git a/src/pipeline.py b/src/pipeline.py index b612b9d..10b06b3 100644 --- a/src/pipeline.py +++ b/src/pipeline.py @@ -19,44 +19,18 @@ from typing import Any from . import get_logger from .config import load_config -from .doc.extractor import extract_invoices -from .doc.invoice import ( - classify_invoice_batch, - save_application_json, - save_invoice_csv, -) -from .doc.invoice import ( - save_csv as save_payment_csv, -) +from .core.extraction import extract_invoices from .pipeline_core import ( extract_and_cache_normal_info, extract_and_cache_travel_info, + extract_info_by_type, is_travel_invoice, + process_invoices, ) log = get_logger("pipeline") -def _classify_from_cache(cache_path: Path) -> dict[str, list[dict[str, Any]]]: - """从缓存目录读取发票数据并按类型分组 - - 注意:load_cache 会加载所有 .invoice_cache/*.json,包括 travel_info 和 normal_info - 等提取结果缓存(它们没有 invoice_type 字段),需要显式过滤掉。 - """ - from .doc.llm_extractor import load_cache - - cache_map = load_cache(cache_path) - # 过滤掉非发票的缓存条目:travel_info、normal_info 等提取结果 - invoices = [ - data - for key, data in cache_map.items() - if key not in ("travel_info", "normal_info") - and isinstance(data, dict) - and data.get("invoice_type") not in ("application", "payment") - ] - return classify_invoice_batch(invoices) - - def run_pipeline( step: str = "all", username: str | None = None, @@ -98,19 +72,11 @@ def run_pipeline( log.error("未提取到任何发票数据") return 1 - save_payment_csv(payment_records, cache_path / "payment_records.csv") - save_invoice_csv(payment_records, cache_path / "invoice_summary.csv") + # 使用公共函数处理发票数据 + process_invoices(payment_records, applications, groups, cache_path) - if applications: - save_application_json(applications, cache_path / "travel_applications.json") - - log.info(f"发票分类: 差旅 {len(groups['travel'])} 张, 普通 {len(groups['general'])} 张") - - # 发票提取完成后立即判断类型 - if is_travel_invoice(groups): - travel_info = extract_and_cache_travel_info(groups, cache_path) - else: - normal_info = extract_and_cache_normal_info(groups, cache_path) + # 发票提取完成后立即判断类型并提取信息 + travel_info, normal_info = extract_info_by_type(groups, cache_path) if step == "invoice": log.info("[1/2] 发票提取 完成") @@ -124,9 +90,12 @@ def run_pipeline( log.info("[2/2] 报销提交") log.info("=" * 60) - from .bot import run_bot + from .infra.browser import run_bot if groups is None: + # 从缓存重新分类(仅 submit 阶段需要) + from .pipeline_core import _classify_from_cache + groups = _classify_from_cache(cache_path) if is_travel_invoice(groups): diff --git a/src/pipeline_core.py b/src/pipeline_core.py index 8f6c98a..64ba4fd 100644 --- a/src/pipeline_core.py +++ b/src/pipeline_core.py @@ -11,15 +11,21 @@ from __future__ import annotations import json from pathlib import Path -from typing import Any +from typing import Any, cast from . import get_logger -from .doc.llm_extractor import ( +from .core.extraction import ( CACHE_DIR_NAME, extract_normal_info, extract_travel_info, load_cache, ) +from .infra.documents import ( + classify_invoice_batch, + save_application_json, + save_invoice_csv, +) +from .infra.documents import save_csv as save_payment_csv log = get_logger("pipeline_core") @@ -98,3 +104,117 @@ def extract_and_cache_normal_info( normal_info = extract_normal_info(source_dir=cache_path) save_cache_info(cache_path, "normal_info", normal_info) return normal_info + + +def _classify_from_cache(cache_path: Path) -> dict[str, list[dict[str, Any]]]: + """从缓存目录读取发票数据并按类型分组 + + 注意:load_cache 会加载所有 .invoice_cache/*.json,包括 travel_info 和 normal_info + 等提取结果缓存(它们没有 invoice_type 字段),需要显式过滤掉。 + """ + cache_map = load_cache(cache_path) + # 过滤掉非发票的缓存条目:travel_info、normal_info 等提取结果 + invoices = [ + data + for key, data in cache_map.items() + if key not in ("travel_info", "normal_info") + and isinstance(data, dict) + and data.get("invoice_type") not in ("application", "payment") + ] + return classify_invoice_batch(invoices) + + +def save_invoice_groups(session_dir: Path, groups: dict[str, list[dict[str, str]]]) -> None: + """保存发票分类结果到目录的 JSON 文件 + + 保存完整发票分组数据(供 is_travel_invoice/extract_and_cache_* 使用), + 同时保留计数字段(供快速统计使用)。 + """ + groups_path = session_dir / "invoice_groups.json" + data = { + "travel": groups.get("travel", []), + "general": groups.get("general", []), + "application": groups.get("application", []), + "travel_count": len(groups.get("travel", [])), + "general_count": len(groups.get("general", [])), + "application_count": len(groups.get("application", [])), + } + with open(groups_path, "w", encoding="utf-8") as f: + json.dump(data, f, ensure_ascii=False, indent=2) + + +def load_invoice_groups(session_dir: Path) -> dict[str, Any] | None: + """从目录加载发票分类结果 + + 返回包含完整发票分组数据和计数字段的字典。 + """ + groups_path = session_dir / "invoice_groups.json" + if not groups_path.exists(): + return None + try: + with open(groups_path, encoding="utf-8") as f: + return cast(dict[str, Any] | None, json.load(f)) + except Exception: + return None + + +def process_invoices( + payment_records: list[dict[str, Any]], + applications: list[dict[str, Any]], + groups: dict[str, list[dict[str, Any]]], + output_dir: Path, +) -> dict[str, Any]: + """处理提取的发票数据:保存 CSV、申请单和分类结果 + + Args: + payment_records: 支付记录列表。 + applications: 申请单列表。 + groups: 发票分类结果。 + output_dir: 输出目录。 + + Returns: + 包含发票统计信息的字典。 + """ + # 保存 CSV:支付记录级别(供 bot/出库单使用)和发票级别(供人工参考) + save_payment_csv(payment_records, output_dir / "payment_records.csv") + save_invoice_csv(payment_records, output_dir / "invoice_summary.csv") + + # 出差申请单单独保存 + if applications: + save_application_json(applications, output_dir / "travel_applications.json") + + # 保存分类结果(供后续步骤统一读取) + save_invoice_groups(output_dir, groups) + + # 统计发票总数 + invoice_count = sum(len(inv.get("_matched_invoices", [])) for inv in payment_records) + + log.info(f"发票分类: 差旅 {len(groups['travel'])} 张, 普通 {len(groups['general'])} 张") + + return { + "invoice_count": invoice_count, + "travel_count": len(groups["travel"]), + "general_count": len(groups["general"]), + "application_count": len(groups.get("application", [])), + } + + +def extract_info_by_type( + groups: dict[str, list[dict[str, Any]]], + cache_path: Path, +) -> tuple[dict[str, Any] | None, dict[str, Any] | None]: + """根据发票类型提取差旅或普通报销信息 + + Args: + groups: 发票分类结果。 + cache_path: 缓存目录路径。 + + Returns: + (travel_info, normal_info) 元组,根据发票类型返回对应信息。 + """ + if is_travel_invoice(groups): + travel_info = extract_and_cache_travel_info(groups, cache_path) + return travel_info, None + else: + normal_info = extract_and_cache_normal_info(groups, cache_path) + return None, normal_info diff --git a/src/web/app.py b/src/web/app.py index ac5def2..b1b8396 100644 --- a/src/web/app.py +++ b/src/web/app.py @@ -22,15 +22,22 @@ sys.path.insert(0, str(PROJECT_ROOT)) from flask import Flask # noqa: E402, I001 from src.web import pipeline_web, routes # noqa: E402, I001 -from src.doc.fill_consumable_doc import CONSUMABLE_DOC_FILENAME # noqa: E402, I001 +from src.infra.documents import CONSUMABLE_DOC_FILENAME # noqa: E402, I001 app = Flask(__name__, template_folder="templates") UPLOAD_BASE = PROJECT_ROOT / "src" / "web" / "uploads" +_app_initialized = False + def create_app() -> Flask: """应用工厂:初始化配置并注册路由""" + global _app_initialized + if _app_initialized: + return app + _app_initialized = True + # 设置出库单模板路径 pipeline_web.set_consumable_template(PROJECT_ROOT / CONSUMABLE_DOC_FILENAME) diff --git a/src/web/pipeline_web.py b/src/web/pipeline_web.py index df6084c..1b19cc6 100644 --- a/src/web/pipeline_web.py +++ b/src/web/pipeline_web.py @@ -8,29 +8,21 @@ Web 管道逻辑 - 发票分类数据持久化 """ -import json import time from pathlib import Path -from typing import Any, cast +from typing import Any from urllib.parse import quote # 延迟导入,避免循环引用 from src import get_logger # noqa: F401 -from src.doc.fill_consumable_doc import ( +from src.infra.documents import ( CONSUMABLE_DOC_FILENAME, fill_consumable_from_template, ) -from src.doc.invoice import ( - save_application_json, - save_invoice_csv, -) -from src.doc.invoice import ( - save_csv as save_payment_csv, -) from src.pipeline_core import ( - extract_and_cache_normal_info, - extract_and_cache_travel_info, - is_travel_invoice, + extract_info_by_type, + load_invoice_groups, + process_invoices, ) fill_log = get_logger("fill_consumable_doc") @@ -40,40 +32,6 @@ SESSION_RESULT_FILE = "result.json" INVOICE_GROUPS_FILE = "invoice_groups.json" -def save_invoice_groups(session_dir: Path, groups: dict[str, list[dict[str, str]]]) -> None: - """保存发票分类结果到 session 目录的 JSON 文件 - - 保存完整发票分组数据(供 is_travel_invoice/extract_and_cache_* 使用), - 同时保留计数字段(供快速统计使用)。 - """ - groups_path = session_dir / INVOICE_GROUPS_FILE - data = { - "travel": groups.get("travel", []), - "general": groups.get("general", []), - "application": groups.get("application", []), - "travel_count": len(groups.get("travel", [])), - "general_count": len(groups.get("general", [])), - "application_count": len(groups.get("application", [])), - } - with open(groups_path, "w", encoding="utf-8") as f: - json.dump(data, f, ensure_ascii=False, indent=2) - - -def load_invoice_groups(session_dir: Path) -> dict[str, Any] | None: - """从 session 目录加载发票分类结果 - - 返回包含完整发票分组数据和计数字段的字典。 - """ - groups_path = session_dir / INVOICE_GROUPS_FILE - if not groups_path.exists(): - return None - try: - with open(groups_path, encoding="utf-8") as f: - return cast(dict[str, Any] | None, json.load(f)) - except Exception: - return None - - def resolve_payment_csv(session_dir: Path) -> Path | None: """查找支付记录 CSV(payment_records.csv)""" csv_path = session_dir / "payment_records.csv" @@ -175,7 +133,7 @@ def run_pipeline_web(session_dir: Path, config: dict[str, Any] | None) -> dict[s - 差旅发票:调用 LLM 提取差旅信息并缓存到 travel_info.json - 普通发票:无需额外提取(normal_info.json 待实现) """ - from src.doc.extractor import extract_invoices + from src.core.extraction import extract_invoices start = time.time() @@ -184,34 +142,20 @@ def run_pipeline_web(session_dir: Path, config: dict[str, Any] | None) -> dict[s if not invoices: return {"ok": False, "error": "未提取到任何发票数据"} - # 保存 CSV:支付记录级别(供 bot/出库单使用)和发票级别(供人工参考) - save_payment_csv(invoices, session_dir / "payment_records.csv") - save_invoice_csv(invoices, session_dir / "invoice_summary.csv") - - # 出差申请单单独保存 - if applications: - save_application_json(applications, session_dir / "travel_applications.json") - - # 保存分类结果(供后续步骤统一读取) - save_invoice_groups(session_dir, groups) + # 使用公共函数处理发票数据 + stats = process_invoices(invoices, applications, groups, session_dir) # ---- Step 2: 差旅/普通信息提取 ---- - if is_travel_invoice(groups): - extract_and_cache_travel_info(groups, session_dir) - else: - extract_and_cache_normal_info(groups, session_dir) - - # 统计发票总数 - invoice_count = sum(len(inv.get("_matched_invoices", [])) for inv in invoices) + extract_info_by_type(groups, session_dir) elapsed = time.time() - start result = { "ok": True, "elapsed": f"{elapsed:.1f}s", - "invoice_count": invoice_count, + "invoice_count": stats["invoice_count"], "csv_url": f"/api/download/{session_dir.name}/invoice_summary.csv", - "travel_count": len(groups["travel"]), - "general_count": len(groups["general"]), + "travel_count": stats["travel_count"], + "general_count": stats["general_count"], } doc_fill = _try_fill_consumable_doc(session_dir, config or {}) append_doc_download(result, session_dir.name, doc_fill) @@ -229,7 +173,7 @@ def run_financial_submit(session_dir: Path, config: dict[str, Any] | None) -> di if not csv_path.exists(): return {"ok": False, "error": "未找到发票数据,请先处理"} - from src.bot import run_bot_web + from src.infra.browser import run_bot_web groups = load_invoice_groups(session_dir) if groups: diff --git a/src/web/routes.py b/src/web/routes.py index 53aa2b8..03f86cf 100644 --- a/src/web/routes.py +++ b/src/web/routes.py @@ -24,7 +24,7 @@ from src.config import ( from src.config import ( load_config as load_project_config, ) -from src.doc.invoice import load_csv, load_invoice_csv +from src.infra.documents import load_csv, load_invoice_csv from . import pipeline_web, sse_handler @@ -634,12 +634,12 @@ def agent_process(session_id: str) -> Any: run_agent_round, save_agent_state, ) - from src.doc.extractor import extract_invoices - from src.doc.invoice import ( + from src.core.extraction import extract_invoices + from src.infra.documents import ( save_application_json, save_invoice_csv, ) - from src.doc.invoice import ( + from src.infra.documents import ( save_csv as save_payment_csv, ) @@ -702,12 +702,12 @@ def agent_supplement(session_id: str) -> Any: agent_session = add_supplement(session_dir, agent_session, filenames) # 重新提取发票并保存 - from src.doc.extractor import extract_invoices - from src.doc.invoice import ( + from src.core.extraction import extract_invoices + from src.infra.documents import ( save_application_json, save_invoice_csv, ) - from src.doc.invoice import ( + from src.infra.documents import ( save_csv as save_payment_csv, ) diff --git a/src/web/static/js/agent.js b/src/web/static/js/agent.js index 5c4922b..891df55 100644 --- a/src/web/static/js/agent.js +++ b/src/web/static/js/agent.js @@ -35,6 +35,12 @@ export function handleAgentEvent(msg) { case 'agent_force_submit': _handleForceSubmit(msg); break; + case 'agent_max_rounds': + _handleAgentError(msg); + break; + case 'agent_extract_status': + _handleAgentStateChange(msg); + break; } } diff --git a/tests/test_extractor.py b/tests/test_extractor.py index 87c0f02..f5eaa02 100644 --- a/tests/test_extractor.py +++ b/tests/test_extractor.py @@ -12,7 +12,7 @@ from typing import Any import pytest from src import exceptions -from src.doc.extractor import extract_invoices +from src.core.extraction import extract_invoices # 字段键名(与源码中的字符串字面量保持一致) K_INVOICE_TYPE = "invoice_type" @@ -78,7 +78,7 @@ class TestExtractInvoices: """extract_invoices 编排函数""" def test_empty_directory(self, tmp_path: Path, monkeypatch): - monkeypatch.setattr("src.doc.extractor._find_all_files", lambda d: []) + monkeypatch.setattr("src.core.extraction.extractor._find_all_files", lambda d: []) records, apps, groups = extract_invoices(str(tmp_path)) assert records == [] @@ -89,8 +89,8 @@ class TestExtractInvoices: pdf = tmp_path / "broken.pdf" pdf.touch() - monkeypatch.setattr("src.doc.extractor._find_all_files", lambda d: [pdf]) - monkeypatch.setattr("src.doc.extractor._extract_document", lambda p, c, s: (None, "parse error")) + monkeypatch.setattr("src.core.extraction.extractor._find_all_files", lambda d: [pdf]) + monkeypatch.setattr("src.core.extraction.extractor._extract_document", lambda p, c, s: (None, "parse error")) with pytest.raises(exceptions.ExtractionError) as exc_info: extract_invoices(str(tmp_path)) @@ -107,16 +107,14 @@ class TestExtractInvoices: inv1 = _make_invoice("INV001", 300.0, INVOICE_TYPE_GENERAL) inv2 = _make_invoice("INV002", 200.0, INVOICE_TYPE_GENERAL) - monkeypatch.setattr("src.doc.extractor._find_all_files", lambda d: [pdf1, pdf2]) + monkeypatch.setattr("src.core.extraction.extractor._find_all_files", lambda d: [pdf1, pdf2]) - call_index = [0] + path_to_result = {pdf1: inv1, pdf2: inv2} def fake_extract(path, cache_dir, source_dir): - idx = call_index[0] - call_index[0] += 1 - return (inv1, None) if idx == 0 else (inv2, None) + return (path_to_result[path], None) - monkeypatch.setattr("src.doc.extractor._extract_document", fake_extract) + monkeypatch.setattr("src.core.extraction.extractor._extract_document", fake_extract) def fake_match(invoices, cards): return [ @@ -131,7 +129,7 @@ class TestExtractInvoices: } ] - monkeypatch.setattr("src.doc.extractor.match_invoices_to_cards", fake_match) + monkeypatch.setattr("src.core.matching.matcher.match_invoices_to_cards", fake_match) records, apps, groups = extract_invoices(str(tmp_path)) @@ -153,19 +151,17 @@ class TestExtractInvoices: inv_general = _make_invoice("GEN001", 150.0, INVOICE_TYPE_GENERAL) monkeypatch.setattr( - "src.doc.extractor._find_all_files", + "src.core.extraction.extractor._find_all_files", lambda d: [pdf1, pdf2, pdf3], ) invoices_list = [inv_train, inv_hotel, inv_general] - call_index = [0] + path_to_result = dict(zip([pdf1, pdf2, pdf3], invoices_list, strict=True)) def fake_extract(path, cache_dir, source_dir): - idx = call_index[0] - call_index[0] += 1 - return (invoices_list[idx], None) + return (path_to_result[path], None) - monkeypatch.setattr("src.doc.extractor._extract_document", fake_extract) + monkeypatch.setattr("src.core.extraction.extractor._extract_document", fake_extract) def fake_match(invoices, cards): return [ @@ -198,7 +194,7 @@ class TestExtractInvoices: }, ] - monkeypatch.setattr("src.doc.extractor.match_invoices_to_cards", fake_match) + monkeypatch.setattr("src.core.matching.matcher.match_invoices_to_cards", fake_match) records, apps, groups = extract_invoices(str(tmp_path)) @@ -221,19 +217,16 @@ class TestExtractInvoices: app = _make_application() monkeypatch.setattr( - "src.doc.extractor._find_all_files", + "src.core.extraction.extractor._find_all_files", lambda d: [pdf1, pdf2], ) - results = [inv, app] - call_index = [0] + path_to_result = {pdf1: inv, pdf2: app} def fake_extract(path, cache_dir, source_dir): - idx = call_index[0] - call_index[0] += 1 - return (results[idx], None) + return (path_to_result[path], None) - monkeypatch.setattr("src.doc.extractor._extract_document", fake_extract) + monkeypatch.setattr("src.core.extraction.extractor._extract_document", fake_extract) def fake_match(invoices, cards): return [ @@ -248,7 +241,7 @@ class TestExtractInvoices: } ] - monkeypatch.setattr("src.doc.extractor.match_invoices_to_cards", fake_match) + monkeypatch.setattr("src.core.matching.matcher.match_invoices_to_cards", fake_match) records, apps, groups = extract_invoices(str(tmp_path)) @@ -266,19 +259,16 @@ class TestExtractInvoices: card = _make_card(300.0) monkeypatch.setattr( - "src.doc.extractor._find_all_files", + "src.core.extraction.extractor._find_all_files", lambda d: [pdf1, pdf2], ) - results = [inv, card] - call_index = [0] + path_to_result = {pdf1: inv, pdf2: card} def fake_extract(path, cache_dir, source_dir): - idx = call_index[0] - call_index[0] += 1 - return (results[idx], None) + return (path_to_result[path], None) - monkeypatch.setattr("src.doc.extractor._extract_document", fake_extract) + monkeypatch.setattr("src.core.extraction.extractor._extract_document", fake_extract) def fake_match(invoices, cards): return [ @@ -293,7 +283,7 @@ class TestExtractInvoices: } ] - monkeypatch.setattr("src.doc.extractor.match_invoices_to_cards", fake_match) + monkeypatch.setattr("src.core.matching.matcher.match_invoices_to_cards", fake_match) records, apps, groups = extract_invoices(str(tmp_path)) @@ -309,19 +299,17 @@ class TestExtractInvoices: inv = _make_invoice("INV001", 300.0, INVOICE_TYPE_GENERAL) monkeypatch.setattr( - "src.doc.extractor._find_all_files", + "src.core.extraction.extractor._find_all_files", lambda d: [pdf1, pdf2], ) - results = [inv, None] - call_index = [0] + path_to_result = {pdf1: inv, pdf2: None} def fake_extract(path, cache_dir, source_dir): - idx = call_index[0] - call_index[0] += 1 - return (results[idx], "parse error") if results[idx] is None else (results[idx], None) + result = path_to_result[path] + return (result, "parse error") if result is None else (result, None) - monkeypatch.setattr("src.doc.extractor._extract_document", fake_extract) + monkeypatch.setattr("src.core.extraction.extractor._extract_document", fake_extract) def fake_match(invoices, cards): return [ @@ -336,7 +324,7 @@ class TestExtractInvoices: } ] - monkeypatch.setattr("src.doc.extractor.match_invoices_to_cards", fake_match) + monkeypatch.setattr("src.core.matching.matcher.match_invoices_to_cards", fake_match) records, apps, groups = extract_invoices(str(tmp_path)) diff --git a/tests/test_invoice.py b/tests/test_invoice.py index 9d05523..4888ed2 100644 --- a/tests/test_invoice.py +++ b/tests/test_invoice.py @@ -3,7 +3,7 @@ 覆盖发票分类、CSV 列定义、常量校验。 """ -from src.doc.invoice import ( +from src.infra.documents import ( INVOICE_LEVEL_COLUMNS, PAYMENT_RECORD_COLUMNS, classify_invoice_batch, diff --git a/tests/test_llm_extractor.py b/tests/test_llm_extractor.py index 1d0ca33..da6cd24 100644 --- a/tests/test_llm_extractor.py +++ b/tests/test_llm_extractor.py @@ -14,7 +14,7 @@ from pathlib import Path import pytest -from src.doc.llm_extractor import ( +from src.core.extraction import ( _image_to_base64, extract_document, parse_json_response, @@ -131,7 +131,7 @@ class TestExtractDocument: def fake_query(system_prompt, text, image_b64, max_tokens=4096): return response_text - monkeypatch.setattr("src.doc.llm_extractor._llm_query_multimodal", fake_query) + monkeypatch.setattr("src.core.extraction.llm_extractor._llm_query_multimodal", fake_query) def test_success(self, tmp_path: Path, monkeypatch): img_path = tmp_path / "card.png" @@ -156,7 +156,7 @@ class TestExtractDocument: def fake_query(system_prompt, text, image_b64, max_tokens=4096): raise RuntimeError("模型不可用") - monkeypatch.setattr("src.doc.llm_extractor._llm_query_multimodal", fake_query) + monkeypatch.setattr("src.core.extraction.llm_extractor._llm_query_multimodal", fake_query) with pytest.raises(RuntimeError, match="模型不可用"): extract_document(img_path) @@ -193,7 +193,7 @@ class TestExtractDocument: received_b64 = image_b64s return json.dumps({K_CARD_DATE: "2026-01-01", K_CARD_NO: "0000", K_CARD_AMOUNT: "100"}) - monkeypatch.setattr("src.doc.llm_extractor._llm_query_multimodal", capture_b64) + monkeypatch.setattr("src.core.extraction.llm_extractor._llm_query_multimodal", capture_b64) extract_document(img_path) assert received_b64 is not None diff --git a/tests/test_matcher.py b/tests/test_matcher.py index d71f0a4..0690baa 100644 --- a/tests/test_matcher.py +++ b/tests/test_matcher.py @@ -12,7 +12,7 @@ from __future__ import annotations from typing import Any -from src.doc.matcher import ( +from src.core.matching import ( _build_invoice_summary, _build_payment_records, _invoices_to_records,