refactor: 架构重组 — doc/bot → core/infra,新增 Agent 调度模块

- 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
This commit is contained in:
wandering
2026-07-02 18:36:19 +08:00
parent 7c137c5214
commit 1b35f07fd7
69 changed files with 2965 additions and 1548 deletions

View File

@@ -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["提取差旅信息<br/>travel_info.json"]
TRAVEL_INFO --> TRAVEL_BOT["browser/travel.py<br/>填报差旅报销单"]
NORMAL --> NORMAL_INFO["提取普通发票信息<br/>normal_info.json"]
NORMAL_INFO --> NORMAL_BOT["browser/normal.py<br/>填报普通报销单"]
NORMAL_INFO --> CONSUMABLE["生成易耗品出库单<br/>(仅普通报销)"]
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` | 项目配置 |

View File

@@ -0,0 +1,20 @@
---
last_reviewed: 2026-06-15
---
# .agents/docs/plans — 实施方案与工作交接
存放项目实施方案、架构分析报告、重构计划等规划类文档。
## 文件
| 文件 | 说明 |
|------|------|
| `架构分析-2026-06-15.md` | 项目架构分析与重构建议(模块拆分、分层设计、接口契约) |
## 用途
- 架构决策记录
- 重构实施方案
- 工作交接说明
- 技术选型论证

View File

@@ -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*

BIN
.coverage

Binary file not shown.

View File

@@ -4,26 +4,93 @@ alwaysApply: true
--- ---
--- ---
last_reviewed: 2026-06-09 last_reviewed: 2026-07-02
--- ---
# AGENTS 索引 # AGENTS — 项目操作指南
本文件是规则的入口。详细策略文本位于 `.agents/docs/standards/*.md` 本文件为 Agent 提供高信号量的项目操作知识,避免重复探索
## 文档边界 ## 文档边界
* 一定不要用**表情文字**输出任何内容,禁止!!!!! * **禁止使用表情文字**输出任何内容
* `docs/` 目录专门存放面向开源用户、外部贡献者的项目公开文档及说明文件 * `docs/` 目录存放面向开源用户、外部贡献者的公开文档。
* 维护规范、实施方案、经验总结、拉取请求佐证材料与各类内部记录资料,均统一放置在 `.agents/` 目录下,避免内部自动化流程相关内容混入公开文档目录 * `.agents/` 目录存放维护规范、实施方案、经验总结等内部资料
* 每个文件夹下都有一个 `README.md` 文件用来交代这个文件夹的作用以及重要信息。 * 每个文件夹下都有 `README.md` 说明该文件夹的作用重要信息。
## 标准目录 ## 开发命令(必须使用 uv
* 标准文档元数据:`.agents/docs/standards/README.md` 项目使用 `uv` 管理依赖,所有包版本锁定在 `uv.lock` 中。
* 调试规范:`.agents/docs/standards/调试规范.md`
* 复利式工程实践:`.agents/docs/standards/复利式工程实践.md`
## 项目的架构思想 | 操作 | 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/<session_id>/.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/<session_id>/`,每次上传生成独立会话。
## 测试
```bash
make test # pytest + coverage report (term-missing)
```
测试目录:`tests/`,配置在 `pyproject.toml` 中 (`testpaths = ["tests"]`, `pythonpath = ["."]`)。

126
README.md
View File

@@ -18,20 +18,40 @@
├── src/ ├── src/
│ ├── __init__.py # 包初始化 / 日志器 │ ├── __init__.py # 包初始化 / 日志器
│ ├── config.py # 配置加载 │ ├── config.py # 配置加载
│ ├── bot.py # 浏览器自动填报 │ ├── exceptions.py # 异常定义
│ ├── pipeline.py # CLI 流程编排 │ ├── pipeline.py # CLI 流程编排
│ ├── pipeline_core.py # CLI/Web 公共管道逻辑
│ ├── main.py # CLI 入口 │ ├── main.py # CLI 入口
│ ├── doc/ # 文档处理模块 │ ├── agent/ # Agent 调度模块
│ │ ├── extractor.py # 编排入口:串联 PDF 读取 → LLM 提取 → 分类 │ │ ├── orchestrator.py # 总调度入口
│ │ ├── pdf.py # PDF 图片渲染PyMuPDF供多模态 LLM 使用) │ │ ├── coordinator.py # 校验-修正循环
│ │ ├── llm_extractor.py # LLM 信息提取 │ │ ├── session.py # 状态机与会话数据
│ │ ── matcher.py # 数据匹配与校验 │ │ ── events.py # SSE 事件发射
│ ├── invoice.py # 发票类型常量、分类逻辑、CSV 读写工具 │ ├── core/ # 核心业务逻辑
│ │ ├── fill_consumable_doc.py # 将 CSV 填入易耗品出库单Word COM │ │ ├── extraction/ # 信息提取
│ │ ├── prompt.py # LLM 提示词模板 │ │ │ ├── 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/ # 提示词模板文件 │ │ └── prompts/ # 提示词模板文件
│ └── web/ │ └── web/ # Web 界面模块
│ ├── app.py # Web 服务入口 │ ├── app.py # Flask 应用入口
│ ├── routes.py # 路由定义
│ ├── pipeline_web.py # Web 管道逻辑
│ ├── sse_handler.py # SSE 日志流处理
│ ├── templates/ │ ├── templates/
│ │ ├── index.html # PC 端主页 │ │ ├── index.html # PC 端主页
│ │ └── mobile_upload.html # 移动端扫码上传 │ │ └── mobile_upload.html # 移动端扫码上传
@@ -48,6 +68,66 @@
└── *.pdf / *.jpg / *.png # 发票 PDF 或图片CLI 模式,放在 scripts/data/ └── *.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 ```mermaid
@@ -73,21 +153,21 @@ flowchart TB
MatchResult --> NormalLLM MatchResult --> NormalLLM
NormalLLM --> NormalInfo[(normal_info.json)] NormalLLM --> NormalInfo[(normal_info.json)]
TravelInfo -->|差旅基本信息| Bot_T[bot/travel.py<br/>差旅填报流程] TravelInfo -->|差旅基本信息| Bot_T[infra/browser/travel.py<br/>差旅填报流程]
TravelInfo -->|报销明细| Bot_T TravelInfo -->|报销明细| Bot_T
TravelInfo -->|支付方式| Bot_T TravelInfo -->|支付方式| Bot_T
TravelInfo -->|补助清单| Bot_T TravelInfo -->|补助清单| Bot_T
TravelInfo -->|附件清单| Bot_T TravelInfo -->|附件清单| Bot_T
Bot_T --> Submit_T[差旅报销提交] Bot_T --> Submit_T[差旅报销提交]
NormalInfo -->|报销说明| Bot_G[bot/normal.py<br/>普通填报流程] NormalInfo -->|报销说明| Bot_G[infra/browser/normal.py<br/>普通填报流程]
NormalInfo -->|发票总数/金额| Bot_G NormalInfo -->|发票总数/金额| Bot_G
NormalInfo -->|支付方式| Bot_G NormalInfo -->|支付方式| Bot_G
NormalInfo -->|附件清单| Bot_G NormalInfo -->|附件清单| Bot_G
Bot_G --> Submit_G[普通报销提交] Bot_G --> Submit_G[普通报销提交]
General --> CSV[(invoice_summary.csv)] General --> CSV[(invoice_summary.csv)]
CSV --> Fill[fill_consumable_doc] CSV --> Fill[consumable.py]
Fill --> Doc[易耗品、出库单.doc] Fill --> Doc[易耗品、出库单.doc]
``` ```
@@ -103,14 +183,14 @@ flowchart TB
### bot 模块架构 ### bot 模块架构
`bot/` 包负责浏览器自动化填报,仅接收已提取的信息并执行填报操作,不承担信息提取职责: `infra/browser/` 包负责浏览器自动化填报,仅接收已提取的信息并执行填报操作,不承担信息提取职责:
| 模块 | 职责 | | 模块 | 职责 |
|------|------| |------|------|
| `bot/base.py` | `BaseBot` 基类:浏览器生命周期、登录、导航、截图 | | `infra/browser/base.py` | `BaseBot` 基类:浏览器生命周期、登录、导航、截图 |
| `bot/travel.py` | 差旅填报流程:基本信息 → 差旅明细 → 支付方式 → 补助清单 → 附件上传 | | `infra/browser/travel.py` | 差旅填报流程:基本信息 → 差旅明细 → 支付方式 → 补助清单 → 附件上传 |
| `bot/normal.py` | 普通填报流程:基本信息 → 总明细 → 支付方式 → 附件上传 | | `infra/browser/normal.py` | 普通填报流程:基本信息 → 总明细 → 支付方式 → 附件上传 |
| `bot/__init__.py` | 入口函数:`run_bot()` / `run_bot_web()`,负责类型判断和流程路由 | | `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** 需已生成 `invoice_summary.csv`,且本机已安装 **Microsoft Word**
```bash ```bash
uv run python -m src.doc.fill_consumable_doc uv run python -m src.infra.documents.consumable
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"
uv run python -m src.doc.fill_consumable_doc --config scripts/config.json # 指定配置文件 uv run python -m src.infra.documents.consumable --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 --no-backup # 不生成 .doc.bak 备份
``` ```
填写规则概要: 填写规则概要:

24
config/README.md Normal file
View File

@@ -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` 在启动时读取此文件,若文件不存在则使用内置默认规则。

View File

@@ -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": "附件类型"}
]
}
]
}
}

View File

@@ -371,7 +371,7 @@ Accept: text/event-stream
检测到 `result.json` 存在时,读取后发送 `done` 事件并断开连接。 检测到 `result.json` 存在时,读取后发送 `done` 事件并断开连接。
**SSE 超时:** 600 秒。 **SSE 超时:** 900 秒。
--- ---
@@ -696,6 +696,7 @@ POST /api/agent/force-submit/<session_id>
| `agent_request_supplement` | `{type, round, missing_fields, missing_materials, semantic_issues, suggestion}` | 校验未通过 | | `agent_request_supplement` | `{type, round, missing_fields, missing_materials, semantic_issues, suggestion}` | 校验未通过 |
| `agent_supplement_received` | `{type, files}` | 收到用户补充 | | `agent_supplement_received` | `{type, files}` | 收到用户补充 |
| `agent_force_submit` | `{type, message}` | 用户强制提交 | | `agent_force_submit` | `{type, message}` | 用户强制提交 |
| `agent_extract_status` | `{type, state, round, attempt, message}` | 校验-修正循环中的每次尝试结果 |
| `agent_error` | `{type, message}` | 提取失败 | | `agent_error` | `{type, message}` | 提取失败 |
| `agent_max_rounds` | `{type, message}` | 达到最大轮次 | | `agent_max_rounds` | `{type, message}` | 达到最大轮次 |
@@ -849,7 +850,7 @@ sequenceDiagram
不经过 Web、在本地直接填写出库单 不经过 Web、在本地直接填写出库单
```bash ```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)。 详见 [README.md](./README.md)。

View File

@@ -427,7 +427,7 @@ PC 端生成二维码指向移动端上传页面,手机端上传的图片通
### 单独填写出库单 ### 单独填写出库单
```bash ```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"
``` ```
### 分步执行管道 ### 分步执行管道

View File

@@ -38,10 +38,6 @@ warn_return_any = true
warn_unused_configs = true warn_unused_configs = true
ignore_missing_imports = true ignore_missing_imports = true
[[tool.mypy.overrides]]
module = "tests.*"
ignore_errors = true
[tool.deptry] [tool.deptry]
ignore_notebooks = true ignore_notebooks = true

20
scripts/README.md Normal file
View File

@@ -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` | 测试差旅信息提取 |

View File

@@ -10,7 +10,7 @@ from dotenv import load_dotenv # noqa: E402
from llama_index.core.llms import ChatMessage # noqa: E402 from llama_index.core.llms import ChatMessage # noqa: E402
from src.config import get_llm_config # 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") load_dotenv(Path(__file__).parent / ".env")

View File

@@ -31,7 +31,7 @@ sys.stdout = io.TextIOWrapper(sys.stdout.buffer, encoding="utf-8", errors="repla
sys.path.insert(0, str(ROOT)) # noqa: E402 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: def test_single_file(file_path: Path) -> None:

View File

@@ -15,8 +15,8 @@ sys.stdout = io.TextIOWrapper(sys.stdout.buffer, encoding="utf-8", errors="repla
ROOT = Path(__file__).resolve().parent.parent ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(ROOT)) # noqa: E402 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
from src.doc.pdf import render_pdf_to_images # noqa: E402 from src.infra.documents.pdf import render_pdf_to_images # noqa: E402
def test_render() -> None: def test_render() -> None:

View File

@@ -15,7 +15,7 @@ from pathlib import Path
ROOT = Path(__file__).resolve().parent.parent ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(ROOT)) # noqa: E402 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: def main() -> None:

View File

@@ -1,74 +1,77 @@
--- ---
last_reviewed: 2026-06-12 last_reviewed: 2026-06-15
--- ---
# src — 主源码目录 # 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()` 日志工厂(支持终端 + 文件双输出,按日期自动分文件) | | `agent/` | Agent 调度:校验-修正循环、状态机管理、SSE 事件发射 |
| `config.py` | 配置加载:从 `config.json` 读取用户凭据和默认值从环境变量读取服务端配置SSO 地址、LLM 参数) | | `core/` | 核心业务逻辑:信息提取、金额匹配、信息校验 |
| `pipeline.py` | 流程编排:串联发票提取 → 类型判断 → 差旅/普通信息提取 → 浏览器填报,支持分步执行 | | `infra/` | 基础设施浏览器自动填报、文档处理、LLM 提示词管理 |
| `main.py` | CLI 入口:支持 `--step` 分步执行、`-u/-p` 覆盖凭据、`--cache-dir` 指定缓存目录 | | `web/` | Web 界面Flask 应用、SSE 日志流、可编辑表格、移动端上传、会话隔离 |
| `bot/` | 浏览器自动化Playwright 驱动的财务系统填报机器人(仅负责接收信息并填报 | | `pipeline.py` | CLI 流程编排:串联提取 → 类型判断 → 信息提取 → 浏览器填报 |
| `doc/` | 文档处理模块PDF 渲染、LLM 提取、支付匹配、发票分类、出库单生成 | | `pipeline_core.py` | CLI/Web 公共管道逻辑:发票类型判断、缓存提取 |
| `web/` | Web 界面模块Flask 应用、SSE 日志、可编辑表格、移动端上传、会话隔离 | | `main.py` | CLI 入口:`--step` 分步执行、`-u/-p` 覆盖凭据 |
| `config.py` | 配置加载:`config.json` + 环境变量 |
| `exceptions.py` | 异常层次定义 |
## 数据流 ## 数据流
```mermaid ```mermaid
graph TD graph TD
A[CLI/Web 入口] --> B["pipeline.py (编排)"] A[CLI/Web 入口] --> B["pipeline.py (编排)"]
B --> C["doc/extractor.py (统一提取入口)"] B --> C["core/extraction/extractor.py (统一提取入口)"]
C --> D["doc/pdf.py (PDF 渲染为图片)"] C --> D["infra/documents/pdf.py (PDF 渲染为图片)"]
C --> E["doc/llm_extractor.py (多模态 LLM 识别)"] C --> E["core/extraction/llm_extractor.py (多模态 LLM 识别)"]
E --> F["发票 invoice_type=train/hotel/general"] E --> F["发票 invoice_type=train/hotel/general"]
E --> G["支付记录 invoice_type=payment"] E --> G["支付记录 invoice_type=payment"]
E --> H["出差事前申请单 invoice_type=application"] E --> H["出差事前申请单 invoice_type=application"]
C --> I["doc/matcher.py (发票与支付记录按金额匹配)"] C --> I["core/matching/matcher.py (发票与支付记录按金额匹配)"]
I --> J["一对一匹配 发票数 == 刷卡数"] I --> J["一对一匹配 发票数 == 刷卡数"]
I --> K["一对多匹配 贪心算法 相对容差 3%"] I --> K["一对多匹配 贪心算法 相对容差 3%"]
C --> L["doc/invoice.py (CSV/JSON 读写)"] C --> L["infra/documents/invoice.py (CSV/JSON 读写)"]
L --> M["payment_records.csv (支付记录级别)"] L --> M["payment_records.csv (支付记录级别)"]
L --> N["invoice_summary.csv (发票级别)"] L --> N["invoice_summary.csv (发票级别)"]
L --> O["travel_applications.json (出差申请单)"] L --> O["travel_applications.json (出差申请单)"]
B --> R{"判断报销类型"} B --> R{"判断报销类型"}
R -->|差旅| T["doc/llm_extractor.py (差旅信息提取)"] R -->|差旅| T["core/extraction/llm_extractor.py (差旅信息提取)"]
R -->|普通| V["doc/llm_extractor.py (普通发票信息提取)"] R -->|普通| V["core/extraction/llm_extractor.py (普通发票信息提取)"]
T --> W["travel_info.json (差旅信息: 交通/住宿明细、补贴、附件清单)"] T --> W["travel_info.json (差旅信息: 交通/住宿明细、补贴、附件清单)"]
V --> X["normal_info.json (普通发票信息: 报销说明、发票总数、总金额、支付方式、附件清单)"] V --> X["normal_info.json (普通发票信息: 报销说明、发票总数、总金额、支付方式、附件清单)"]
W --> P["bot/ (浏览器填报 - 仅接收信息并填报)"] W --> P["infra/browser/ (浏览器填报 - 仅接收信息并填报)"]
X --> P X --> P
P --> Q["差旅模式: travel_info.json → 填报差旅单 → 上传差旅附件"] P --> Q["差旅模式: travel_info.json → 填报差旅单 → 上传差旅附件"]
P --> S["普通模式: 基本信息 → 录入明细 → 支付信息 → 上传附件"] 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) | `agent/` | [`agent/README.md`](agent/README.md) |
| `core/` | [`core/README.md`](core/README.md) |
核心能力: | `infra/` | [`infra/README.md`](infra/README.md) |
- **统一文档提取**LLM 自行判断文档类型(发票/支付记录/出差事前申请单),无需正则回退 | `web/` | [`web/README.md`](web/README.md) |
- **JSON 缓存**:提取结果缓存于 `.invoice_cache/`,避免重复处理
- **金额匹配**:支持一对多匹配,相对容差 3%,未匹配发票单独列为记录
- **差旅信息提取**:综合发票、支付记录和匹配结果,提取出差事由、地点、时间等
- **普通发票信息提取**:综合普通发票、支付记录和匹配结果,提取报销说明、发票总数、总金额、支付方式、附件清单
- **出库单生成**:将 CSV 数据填入 Word 模板pywin32 COM仅 Windows
## Web 界面子模块 (`web/`)
详见 [`web/README.md`](web/README.md)
核心能力:
- **会话隔离**:每次上传生成独立 `session_id`,文件/日志/配置/结果各自隔离
- **移动端同步**PC 端生成二维码指向 `/mobile/<sid>`,跨设备协作上传
- **可编辑表格**:前端加载 CSV 数据,支持在线编辑后保存
## 启动方式 ## 启动方式

410
src/agent/coordinator.py Normal file
View File

@@ -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

75
src/agent/events.py Normal file
View File

@@ -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

View File

@@ -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 # 从新模块重新导出所有符号,保持向后兼容
from .coordinator import (
import json add_supplement,
from dataclasses import asdict, dataclass, field force_submit,
from enum import StrEnum process_user_text_supplement,
from pathlib import Path run_agent_round,
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 ..doc.prompt import ( from .events import emit_agent_event as _emit_agent_event
build_normal_info_system_prompt, from .session import (
build_travel_info_system_prompt, 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 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" 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" AGENT_EVENT_LOG = "agent_events.log"
# 去重守卫:记录每个 session 上一次发射的事件类型,防止连续重复发射 __all__ = [
# key: str(session_dir), value: 上一次的 event_type "AgentSession",
_last_event_type: dict[str, str] = {} "AgentState",
"save_agent_state",
"load_agent_state",
def _emit_agent_event(session_dir: Path, event_type: str, **kwargs: Any) -> None: "run_agent_round",
"""向 agent_events.log 追加一行 JSON 事件。 "force_submit",
"add_supplement",
同一 session 连续发射相同 event_type 时直接抛出 RuntimeError "process_user_text_supplement",
强制调用方修复重复发射的代码,而非静默掩盖。 "MAX_VALIDATION_RETRIES",
""" "AGENT_STATE_FILE",
session_key = str(session_dir) "AGENT_EVENT_LOG",
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

101
src/agent/session.py Normal file
View File

@@ -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,
)

View File

@@ -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/` 包,分离差旅和普通报销逻辑

View File

@@ -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,
)

21
src/core/README.md Normal file
View File

@@ -0,0 +1,21 @@
---
last_reviewed: 2026-06-15
---
# src/core — 核心业务逻辑
项目的核心业务层,负责信息提取、金额匹配和信息校验。此层不依赖 Web 框架或浏览器自动化等基础设施。
## 子模块
| 目录 | 说明 |
|------|------|
| `extraction/` | 文档信息提取PDF/图片 → LLM 多模态识别 → 结构化数据 |
| `matching/` | 发票与支付记录按金额匹配(一对一 / 一对多) |
| `validation/` | 声明式信息完整性校验,规则从 JSON 配置文件加载 |
## 设计原则
- **零外部依赖**:不依赖 Flask、Playwright 等框架
- **接口契约**:每个子模块通过 `__init__.py` 导出稳定的对外接口
- **错误传播**:明确的异常层次,便于上层统一处理

4
src/core/__init__.py Normal file
View File

@@ -0,0 +1,4 @@
"""核心业务逻辑模块
提供信息提取、规则校验和发票匹配功能。
"""

View File

@@ -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。

View File

@@ -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",
]

View File

@@ -20,20 +20,25 @@ SSE 文件进度事件:
""" """
import json import json
import os
from concurrent.futures import ThreadPoolExecutor, as_completed
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
from .. import get_logger from ... import get_logger
from ..exceptions import ExtractionError from ...core.matching import match_invoices_to_cards
from .invoice import CACHE_DIR_NAME, classify_invoice_batch from ...exceptions import ExtractionError
from ...infra.documents.invoice import CACHE_DIR_NAME, classify_invoice_batch
from .llm_extractor import extract_document from .llm_extractor import extract_document
from .matcher import match_invoices_to_cards
log = get_logger("extractor") log = get_logger("extractor")
# SSE 文件进度事件日志文件名 # SSE 文件进度事件日志文件名
FILE_EVENTS_LOG = "file_events.log" 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"} SUPPORTED_EXTENSIONS = {".pdf", ".png", ".jpg", ".jpeg", ".bmp", ".webp"}
@@ -285,8 +290,19 @@ def extract_invoices(
# 记录失败文件及其错误信息 # 记录失败文件及其错误信息
failed_files: list[tuple[str, str]] = [] failed_files: list[tuple[str, str]] = []
for file_path in all_files: max_workers = max(1, EXTRACTION_PARALLEL_COUNT)
result, err = _extract_document(file_path, cache_dir, source_dir) 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}
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
if not result: if not result:
if err: if err:

View File

@@ -38,9 +38,9 @@ import json
from pathlib import Path from pathlib import Path
from typing import Any, cast from typing import Any, cast
from .. import get_logger from ... import get_logger
from .invoice import CACHE_DIR_NAME from ...infra.documents.invoice import CACHE_DIR_NAME
from .prompt import ( from ...infra.llm.prompt import (
build_invoice_system_prompt, build_invoice_system_prompt,
build_normal_info_system_prompt, build_normal_info_system_prompt,
build_supplement_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") log.error("缺少 llama-index-llms-openai-like请执行: uv pip install llama-index-llms-openai-like")
raise raise
from ..config import get_llm_config from ...config import get_llm_config
llm_config = get_llm_config() llm_config = get_llm_config()
return OpenAILike( return OpenAILike(
@@ -213,7 +213,7 @@ def llm_query_text(
from llama_index.core.base.llms.types import TextBlock from llama_index.core.base.llms.types import TextBlock
from llama_index.core.llms import ChatMessage from llama_index.core.llms import ChatMessage
from ..config import get_llm_config from ...config import get_llm_config
messages = [ messages = [
ChatMessage(role="system", content=system_prompt), ChatMessage(role="system", content=system_prompt),
@@ -242,7 +242,7 @@ def extract_document(file_path: Path) -> dict[str, Any]:
Returns: Returns:
包含提取字段的字典 包含提取字段的字典
""" """
from .pdf import render_pdf_to_images from ...infra.documents.pdf import render_pdf_to_images
system_prompt = build_invoice_system_prompt() system_prompt = build_invoice_system_prompt()
user_text = f"请分析以下财务文档并提取信息:\n\n文件名: {file_path.name}" 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.base.llms.types import ImageBlock, TextBlock
from llama_index.core.llms import ChatMessage from llama_index.core.llms import ChatMessage
from ..config import get_llm_config from ...config import get_llm_config
if blocks is not None: if blocks is not None:
final_blocks = blocks final_blocks = blocks

View File

@@ -0,0 +1,27 @@
---
last_reviewed: 2026-06-15
---
# src/core/matching — 金额匹配
将提取到的发票数据与支付记录(刷卡截图)按金额进行匹配。
## 文件
| 文件 | 职责 |
|------|------|
| `matcher.py` | 匹配引擎:一对一匹配、一对多贪心匹配、未匹配发票处理 |
## 匹配策略
| 场景 | 策略 |
|------|------|
| 发票数 == 支付记录数 | 一对一匹配:按金额降序配对,相对容差内即匹配 |
| 发票数 > 支付记录数 | 一对多匹配:贪心算法凑金额,相对容差 3% |
| 文件名匹配 | 最高优先级:文件名(不含后缀)一致时直接匹配 |
| 未匹配发票 | 单独列为一条支付记录,`remark` 标记为 `"unmatched"` |
## 业务约束
- 发票总金额 >= 支付总金额
- 输出以支付记录为主键的结果列表

View File

@@ -0,0 +1,8 @@
"""匹配模块
提供发票与支付记录的金额匹配功能。
"""
from .matcher import match_invoices_to_cards
__all__ = ["match_invoices_to_cards"]

View File

@@ -49,7 +49,7 @@
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
from .. import get_logger from ... import get_logger
log = get_logger("matcher") log = get_logger("matcher")

View File

@@ -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)` | 提取缺失字段列表 |

View File

@@ -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",
]

View File

@@ -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("校验规则已重新加载")

View File

@@ -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<br/>PDF 图片渲染"]
B --> C["llm_extractor.py<br/>发票文本提取"]
C --> D["(发票列表)"]
E["支付截图"] --> F["llm_extractor.py<br/>多模态提取"]
F --> G["(刷卡记录)"]
D --> H["matcher.py<br/>金额贪心匹配 / 相对容差 3%"]
G --> H
H --> I["invoice.py<br/>分类 + CSV 回填"]
I --> J["CSV<br/>刷卡日期/卡号/金额"]
J --> K["fill_consumable_doc<br/>易耗品出库单"]
K --> L["易耗品出库单.doc"]
D --> M["llm_extractor.py<br/>差旅/普通信息综合提取"]
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 提取失败时直接报错,无正则回退

View File

@@ -1,4 +0,0 @@
"""文档处理模块
包含发票提取、LLM 信息提取、出库单填写等功能。
"""

View File

@@ -1,83 +0,0 @@
#性格
你是财务报销信息完整性校验助手。你的任务是检查已提取的报销信息在语义上是否足够支撑完成报销系统填报。
**核心原则**:不仅要检查字段是否存在,还要判断信息在逻辑上是否自洽、是否足以完成填报。
---
## 输入数据说明
你会收到以下数据:
1. **已提取的报销信息**:包含基本信息、报销明细、支付方式、附件清单等
2. **发票缓存数据**:原始发票提取的结构化数据
3. **匹配结果**:发票与支付记录的关联关系
---
## 校验维度
### 1. 逻辑一致性检查
- 出差日期范围是否合理(结束日期不早于开始日期)
- 交通费的去程和返程日期是否在出差日期范围内
- 酒店入住/退房日期是否与出差时间匹配
- 支付金额总和是否与发票金额总和接近(允许小额差异)
- 补助天数计算是否正确
### 2. 信息充分性检查
- 是否有足够的信息填写所有报销系统必填项
- 出差事由是否明确具体(不能过于笼统)
- 人员信息是否完整(姓名、工号)
- 支付方式是否与支付记录对应
- 每张发票应该都有对应的支付记录,如果没有提醒用户补充支付记录
### 3. 潜在问题识别
- 发票日期与出差日期差异过大
- 同一笔支付对应多张发票但金额不匹配
- 缺少关键附件
- 人员信息不一致(如车票姓名与补助清单姓名不同)
---
### 4. 不需询问的问题
- 出差事前申请单和实际出差时间不一致这是正常的,因为规划是实际可以有差异
## 输出格式
严格返回以下 JSON 格式:
```json
{
"valid": truefalse,
"confidence": 0.01.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 语法必须正确
- 不要包含任何思考过程或解释文字

View File

@@ -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

21
src/infra/README.md Normal file
View File

@@ -0,0 +1,21 @@
---
last_reviewed: 2026-06-15
---
# src/infra — 基础设施层
提供浏览器自动化、文档处理和 LLM 接口等底层能力。此层不包含业务逻辑,只提供工具和平台能力。
## 子模块
| 目录 | 说明 |
|------|------|
| `browser/` | Playwright 驱动的财务系统自动填报 |
| `documents/` | 发票数据模型、PDF 渲染、Word 出库单填写 |
| `llm/` | LLM 提示词模板加载与管理 |
## 设计原则
- **无业务逻辑**:只提供工具能力,不包含业务流程判断
- **可替换性**:每个子模块通过 `__init__.py` 导出接口,便于替换实现
- **与 core 层解耦**infra 不依赖 corecore 可通过接口调用 infra

4
src/infra/__init__.py Normal file
View File

@@ -0,0 +1,4 @@
"""基础设施模块
提供浏览器自动化、文档处理和 LLM 接口功能。
"""

View File

@@ -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 模式以无头模式运行

View File

@@ -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)

View File

@@ -1,5 +1,4 @@
""" """浏览器自动化填报 — 公共基类
浏览器自动化填报 公共基类
提供浏览器生命周期管理登录导航截图等公共操作 提供浏览器生命周期管理登录导航截图等公共操作
""" """
@@ -7,7 +6,7 @@
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
from .. import get_logger from ... import get_logger
log = get_logger("bot") log = get_logger("bot")

View File

@@ -1,5 +1,4 @@
""" """普通报销填报流程
普通报销填报流程
负责普通发票报销的完整填报步骤 负责普通发票报销的完整填报步骤
基本信息 总明细 支付方式 附件上传 基本信息 总明细 支付方式 附件上传
@@ -7,7 +6,7 @@
from typing import Any from typing import Any
from .. import get_logger from ... import get_logger
from .base import BaseBot, format_date from .base import BaseBot, format_date
log = get_logger("bot.normal") log = get_logger("bot.normal")

View File

@@ -1,5 +1,4 @@
""" """差旅报销填报流程
差旅报销填报流程
负责差旅报销的完整填报步骤 负责差旅报销的完整填报步骤
基本信息 差旅明细 支付方式 补助清单 附件上传 基本信息 差旅明细 支付方式 补助清单 附件上传
@@ -7,7 +6,7 @@
from typing import Any from typing import Any
from .. import get_logger from ... import get_logger
from .base import BaseBot, format_date from .base import BaseBot, format_date
log = get_logger("bot.travel") 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") bot.page.select_option("#fjlx", "1")
else: else:
bot.page.select_option("#fjlx", "2") 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(): if attachment_file and attachment_file.exists():
bot.page.set_input_files("#file", str(attachment_file)) bot.page.set_input_files("#file", str(attachment_file))
bot.page.wait_for_timeout(1000) 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}") log.error(f"差旅附件上传失败: {e}")
bot._screenshot("travel_attachment_error") bot._screenshot("travel_attachment_error")
raise raise
bot._screenshot("travel_attachment_done") bot._screenshot("travel_attachment_done")

View File

@@ -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 解析,字体统一宋体五号

View File

@@ -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",
]

View File

@@ -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 pathlib import Path
from typing import Any from typing import Any
from .. import get_logger from ... import get_logger
from ..config import load_config from ...config import load_config
from ..doc.invoice import load_invoice_csv from .invoice import load_invoice_csv
log = get_logger("fill_consumable_doc") log = get_logger("fill_consumable_doc")

View File

@@ -13,7 +13,7 @@ import csv
import json import json
from pathlib import Path from pathlib import Path
from .. import get_logger from ... import get_logger
log = get_logger("invoice") log = get_logger("invoice")

View File

@@ -11,7 +11,7 @@ from pathlib import Path
import fitz import fitz
from .. import get_logger from ... import get_logger
log = get_logger("pdf") log = get_logger("pdf")

32
src/infra/llm/README.md Normal file
View File

@@ -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()` | 构建普通发票信息提取系统提示词 |

18
src/infra/llm/__init__.py Normal file
View File

@@ -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",
]

View File

@@ -1,7 +1,6 @@
""" """LLM 提示词模板
LLM 提示词模板
src/prompts/ 目录加载 .md 文件作为提示词模板 infra/llm/prompts/ 目录加载 .md 文件作为提示词模板
""" """
import os import os
@@ -34,8 +33,3 @@ def build_normal_info_system_prompt() -> str:
def build_supplement_system_prompt() -> str: def build_supplement_system_prompt() -> str:
"""构建用户补充信息分析系统提示词。""" """构建用户补充信息分析系统提示词。"""
return _load_prompt("supplement_system.md") return _load_prompt("supplement_system.md")
def build_validation_system_prompt() -> str:
"""构建语义校验系统提示词。"""
return _load_prompt("validation_system.md")

View File

@@ -2,9 +2,9 @@
last_reviewed: 2026-06-12 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 ```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
``` ```

View File

@@ -19,44 +19,18 @@ from typing import Any
from . import get_logger from . import get_logger
from .config import load_config from .config import load_config
from .doc.extractor import extract_invoices from .core.extraction 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 .pipeline_core import ( from .pipeline_core import (
extract_and_cache_normal_info, extract_and_cache_normal_info,
extract_and_cache_travel_info, extract_and_cache_travel_info,
extract_info_by_type,
is_travel_invoice, is_travel_invoice,
process_invoices,
) )
log = get_logger("pipeline") 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( def run_pipeline(
step: str = "all", step: str = "all",
username: str | None = None, username: str | None = None,
@@ -98,19 +72,11 @@ def run_pipeline(
log.error("未提取到任何发票数据") log.error("未提取到任何发票数据")
return 1 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") travel_info, normal_info = extract_info_by_type(groups, cache_path)
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)
if step == "invoice": if step == "invoice":
log.info("[1/2] 发票提取 完成") log.info("[1/2] 发票提取 完成")
@@ -124,9 +90,12 @@ def run_pipeline(
log.info("[2/2] 报销提交") log.info("[2/2] 报销提交")
log.info("=" * 60) log.info("=" * 60)
from .bot import run_bot from .infra.browser import run_bot
if groups is None: if groups is None:
# 从缓存重新分类(仅 submit 阶段需要)
from .pipeline_core import _classify_from_cache
groups = _classify_from_cache(cache_path) groups = _classify_from_cache(cache_path)
if is_travel_invoice(groups): if is_travel_invoice(groups):

View File

@@ -11,15 +11,21 @@ from __future__ import annotations
import json import json
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any, cast
from . import get_logger from . import get_logger
from .doc.llm_extractor import ( from .core.extraction import (
CACHE_DIR_NAME, CACHE_DIR_NAME,
extract_normal_info, extract_normal_info,
extract_travel_info, extract_travel_info,
load_cache, 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") log = get_logger("pipeline_core")
@@ -98,3 +104,117 @@ def extract_and_cache_normal_info(
normal_info = extract_normal_info(source_dir=cache_path) normal_info = extract_normal_info(source_dir=cache_path)
save_cache_info(cache_path, "normal_info", normal_info) save_cache_info(cache_path, "normal_info", normal_info)
return 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

View File

@@ -22,15 +22,22 @@ sys.path.insert(0, str(PROJECT_ROOT))
from flask import Flask # noqa: E402, I001 from flask import Flask # noqa: E402, I001
from src.web import pipeline_web, routes # 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") app = Flask(__name__, template_folder="templates")
UPLOAD_BASE = PROJECT_ROOT / "src" / "web" / "uploads" UPLOAD_BASE = PROJECT_ROOT / "src" / "web" / "uploads"
_app_initialized = False
def create_app() -> Flask: 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) pipeline_web.set_consumable_template(PROJECT_ROOT / CONSUMABLE_DOC_FILENAME)

View File

@@ -8,29 +8,21 @@ Web 管道逻辑
- 发票分类数据持久化 - 发票分类数据持久化
""" """
import json
import time import time
from pathlib import Path from pathlib import Path
from typing import Any, cast from typing import Any
from urllib.parse import quote from urllib.parse import quote
# 延迟导入,避免循环引用 # 延迟导入,避免循环引用
from src import get_logger # noqa: F401 from src import get_logger # noqa: F401
from src.doc.fill_consumable_doc import ( from src.infra.documents import (
CONSUMABLE_DOC_FILENAME, CONSUMABLE_DOC_FILENAME,
fill_consumable_from_template, 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 ( from src.pipeline_core import (
extract_and_cache_normal_info, extract_info_by_type,
extract_and_cache_travel_info, load_invoice_groups,
is_travel_invoice, process_invoices,
) )
fill_log = get_logger("fill_consumable_doc") fill_log = get_logger("fill_consumable_doc")
@@ -40,40 +32,6 @@ SESSION_RESULT_FILE = "result.json"
INVOICE_GROUPS_FILE = "invoice_groups.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: def resolve_payment_csv(session_dir: Path) -> Path | None:
"""查找支付记录 CSVpayment_records.csv""" """查找支付记录 CSVpayment_records.csv"""
csv_path = session_dir / "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 - 差旅发票:调用 LLM 提取差旅信息并缓存到 travel_info.json
- 普通发票无需额外提取normal_info.json 待实现) - 普通发票无需额外提取normal_info.json 待实现)
""" """
from src.doc.extractor import extract_invoices from src.core.extraction import extract_invoices
start = time.time() start = time.time()
@@ -184,34 +142,20 @@ def run_pipeline_web(session_dir: Path, config: dict[str, Any] | None) -> dict[s
if not invoices: if not invoices:
return {"ok": False, "error": "未提取到任何发票数据"} return {"ok": False, "error": "未提取到任何发票数据"}
# 保存 CSV支付记录级别供 bot/出库单使用)和发票级别(供人工参考) # 使用公共函数处理发票数据
save_payment_csv(invoices, session_dir / "payment_records.csv") stats = process_invoices(invoices, applications, groups, session_dir)
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)
# ---- Step 2: 差旅/普通信息提取 ---- # ---- Step 2: 差旅/普通信息提取 ----
if is_travel_invoice(groups): extract_info_by_type(groups, session_dir)
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)
elapsed = time.time() - start elapsed = time.time() - start
result = { result = {
"ok": True, "ok": True,
"elapsed": f"{elapsed:.1f}s", "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", "csv_url": f"/api/download/{session_dir.name}/invoice_summary.csv",
"travel_count": len(groups["travel"]), "travel_count": stats["travel_count"],
"general_count": len(groups["general"]), "general_count": stats["general_count"],
} }
doc_fill = _try_fill_consumable_doc(session_dir, config or {}) doc_fill = _try_fill_consumable_doc(session_dir, config or {})
append_doc_download(result, session_dir.name, doc_fill) 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(): if not csv_path.exists():
return {"ok": False, "error": "未找到发票数据,请先处理"} 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) groups = load_invoice_groups(session_dir)
if groups: if groups:

View File

@@ -24,7 +24,7 @@ from src.config import (
from src.config import ( from src.config import (
load_config as load_project_config, 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 from . import pipeline_web, sse_handler
@@ -634,12 +634,12 @@ def agent_process(session_id: str) -> Any:
run_agent_round, run_agent_round,
save_agent_state, save_agent_state,
) )
from src.doc.extractor import extract_invoices from src.core.extraction import extract_invoices
from src.doc.invoice import ( from src.infra.documents import (
save_application_json, save_application_json,
save_invoice_csv, save_invoice_csv,
) )
from src.doc.invoice import ( from src.infra.documents import (
save_csv as save_payment_csv, 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) agent_session = add_supplement(session_dir, agent_session, filenames)
# 重新提取发票并保存 # 重新提取发票并保存
from src.doc.extractor import extract_invoices from src.core.extraction import extract_invoices
from src.doc.invoice import ( from src.infra.documents import (
save_application_json, save_application_json,
save_invoice_csv, save_invoice_csv,
) )
from src.doc.invoice import ( from src.infra.documents import (
save_csv as save_payment_csv, save_csv as save_payment_csv,
) )

View File

@@ -35,6 +35,12 @@ export function handleAgentEvent(msg) {
case 'agent_force_submit': case 'agent_force_submit':
_handleForceSubmit(msg); _handleForceSubmit(msg);
break; break;
case 'agent_max_rounds':
_handleAgentError(msg);
break;
case 'agent_extract_status':
_handleAgentStateChange(msg);
break;
} }
} }

View File

@@ -12,7 +12,7 @@ from typing import Any
import pytest import pytest
from src import exceptions from src import exceptions
from src.doc.extractor import extract_invoices from src.core.extraction import extract_invoices
# 字段键名(与源码中的字符串字面量保持一致) # 字段键名(与源码中的字符串字面量保持一致)
K_INVOICE_TYPE = "invoice_type" K_INVOICE_TYPE = "invoice_type"
@@ -78,7 +78,7 @@ class TestExtractInvoices:
"""extract_invoices 编排函数""" """extract_invoices 编排函数"""
def test_empty_directory(self, tmp_path: Path, monkeypatch): 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)) records, apps, groups = extract_invoices(str(tmp_path))
assert records == [] assert records == []
@@ -89,8 +89,8 @@ class TestExtractInvoices:
pdf = tmp_path / "broken.pdf" pdf = tmp_path / "broken.pdf"
pdf.touch() pdf.touch()
monkeypatch.setattr("src.doc.extractor._find_all_files", lambda d: [pdf]) monkeypatch.setattr("src.core.extraction.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._extract_document", lambda p, c, s: (None, "parse error"))
with pytest.raises(exceptions.ExtractionError) as exc_info: with pytest.raises(exceptions.ExtractionError) as exc_info:
extract_invoices(str(tmp_path)) extract_invoices(str(tmp_path))
@@ -107,16 +107,14 @@ class TestExtractInvoices:
inv1 = _make_invoice("INV001", 300.0, INVOICE_TYPE_GENERAL) inv1 = _make_invoice("INV001", 300.0, INVOICE_TYPE_GENERAL)
inv2 = _make_invoice("INV002", 200.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): def fake_extract(path, cache_dir, source_dir):
idx = call_index[0] return (path_to_result[path], None)
call_index[0] += 1
return (inv1, None) if idx == 0 else (inv2, 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): def fake_match(invoices, cards):
return [ 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)) records, apps, groups = extract_invoices(str(tmp_path))
@@ -153,19 +151,17 @@ class TestExtractInvoices:
inv_general = _make_invoice("GEN001", 150.0, INVOICE_TYPE_GENERAL) inv_general = _make_invoice("GEN001", 150.0, INVOICE_TYPE_GENERAL)
monkeypatch.setattr( monkeypatch.setattr(
"src.doc.extractor._find_all_files", "src.core.extraction.extractor._find_all_files",
lambda d: [pdf1, pdf2, pdf3], lambda d: [pdf1, pdf2, pdf3],
) )
invoices_list = [inv_train, inv_hotel, inv_general] 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): def fake_extract(path, cache_dir, source_dir):
idx = call_index[0] return (path_to_result[path], None)
call_index[0] += 1
return (invoices_list[idx], 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): def fake_match(invoices, cards):
return [ 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)) records, apps, groups = extract_invoices(str(tmp_path))
@@ -221,19 +217,16 @@ class TestExtractInvoices:
app = _make_application() app = _make_application()
monkeypatch.setattr( monkeypatch.setattr(
"src.doc.extractor._find_all_files", "src.core.extraction.extractor._find_all_files",
lambda d: [pdf1, pdf2], lambda d: [pdf1, pdf2],
) )
results = [inv, app] path_to_result = {pdf1: inv, pdf2: app}
call_index = [0]
def fake_extract(path, cache_dir, source_dir): def fake_extract(path, cache_dir, source_dir):
idx = call_index[0] return (path_to_result[path], None)
call_index[0] += 1
return (results[idx], 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): def fake_match(invoices, cards):
return [ 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)) records, apps, groups = extract_invoices(str(tmp_path))
@@ -266,19 +259,16 @@ class TestExtractInvoices:
card = _make_card(300.0) card = _make_card(300.0)
monkeypatch.setattr( monkeypatch.setattr(
"src.doc.extractor._find_all_files", "src.core.extraction.extractor._find_all_files",
lambda d: [pdf1, pdf2], lambda d: [pdf1, pdf2],
) )
results = [inv, card] path_to_result = {pdf1: inv, pdf2: card}
call_index = [0]
def fake_extract(path, cache_dir, source_dir): def fake_extract(path, cache_dir, source_dir):
idx = call_index[0] return (path_to_result[path], None)
call_index[0] += 1
return (results[idx], 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): def fake_match(invoices, cards):
return [ 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)) records, apps, groups = extract_invoices(str(tmp_path))
@@ -309,19 +299,17 @@ class TestExtractInvoices:
inv = _make_invoice("INV001", 300.0, INVOICE_TYPE_GENERAL) inv = _make_invoice("INV001", 300.0, INVOICE_TYPE_GENERAL)
monkeypatch.setattr( monkeypatch.setattr(
"src.doc.extractor._find_all_files", "src.core.extraction.extractor._find_all_files",
lambda d: [pdf1, pdf2], lambda d: [pdf1, pdf2],
) )
results = [inv, None] path_to_result = {pdf1: inv, pdf2: None}
call_index = [0]
def fake_extract(path, cache_dir, source_dir): def fake_extract(path, cache_dir, source_dir):
idx = call_index[0] result = path_to_result[path]
call_index[0] += 1 return (result, "parse error") if result is None else (result, None)
return (results[idx], "parse error") if results[idx] is None else (results[idx], 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): def fake_match(invoices, cards):
return [ 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)) records, apps, groups = extract_invoices(str(tmp_path))

View File

@@ -3,7 +3,7 @@
覆盖发票分类、CSV 列定义、常量校验。 覆盖发票分类、CSV 列定义、常量校验。
""" """
from src.doc.invoice import ( from src.infra.documents import (
INVOICE_LEVEL_COLUMNS, INVOICE_LEVEL_COLUMNS,
PAYMENT_RECORD_COLUMNS, PAYMENT_RECORD_COLUMNS,
classify_invoice_batch, classify_invoice_batch,

View File

@@ -14,7 +14,7 @@ from pathlib import Path
import pytest import pytest
from src.doc.llm_extractor import ( from src.core.extraction import (
_image_to_base64, _image_to_base64,
extract_document, extract_document,
parse_json_response, parse_json_response,
@@ -131,7 +131,7 @@ class TestExtractDocument:
def fake_query(system_prompt, text, image_b64, max_tokens=4096): def fake_query(system_prompt, text, image_b64, max_tokens=4096):
return response_text 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): def test_success(self, tmp_path: Path, monkeypatch):
img_path = tmp_path / "card.png" img_path = tmp_path / "card.png"
@@ -156,7 +156,7 @@ class TestExtractDocument:
def fake_query(system_prompt, text, image_b64, max_tokens=4096): def fake_query(system_prompt, text, image_b64, max_tokens=4096):
raise RuntimeError("模型不可用") 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="模型不可用"): with pytest.raises(RuntimeError, match="模型不可用"):
extract_document(img_path) extract_document(img_path)
@@ -193,7 +193,7 @@ class TestExtractDocument:
received_b64 = image_b64s received_b64 = image_b64s
return json.dumps({K_CARD_DATE: "2026-01-01", K_CARD_NO: "0000", K_CARD_AMOUNT: "100"}) 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) extract_document(img_path)
assert received_b64 is not None assert received_b64 is not None

View File

@@ -12,7 +12,7 @@ from __future__ import annotations
from typing import Any from typing import Any
from src.doc.matcher import ( from src.core.matching import (
_build_invoice_summary, _build_invoice_summary,
_build_payment_records, _build_payment_records,
_invoices_to_records, _invoices_to_records,