From 7c137c5214ef2af72272fda23468f5e1da86cc26 Mon Sep 17 00:00:00 2001 From: wandering Date: Mon, 15 Jun 2026 10:39:01 +0800 Subject: [PATCH] =?UTF-8?q?agent=20=E6=A8=A1=E5=BC=8F=E8=A1=A5=E5=85=85?= =?UTF-8?q?=E6=96=87=E4=BB=B6=E5=8A=9F=E8=83=BD=E5=B7=B2=E7=BB=8F=E5=AE=9E?= =?UTF-8?q?=E7=8E=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ...SSE立即读到旧result.json导致前端无消息.md | 60 ++ .../{plans => guides}/系统事件流全景图.md | 157 ++++- .agents/docs/plans/agent 改造计划.md | 121 ---- .agents/docs/plans/cursor.md | 569 ------------------ .coverage | Bin 53248 -> 53248 bytes docs/API.md | 9 +- pyproject.toml | 2 +- src/agent/__init__.py | 2 - src/agent/orchestrator.py | 65 +- src/bot/__init__.py | 4 + src/bot/travel.py | 25 +- src/config.py | 62 +- src/doc/README.md | 13 - src/doc/extractor.py | 3 - src/doc/fill_consumable_doc.py | 8 +- src/doc/llm_extractor.py | 214 +++---- src/doc/matcher.py | 37 +- src/doc/pdf.py | 2 +- src/doc/validator.py | 3 +- src/pipeline.py | 99 +-- src/pipeline_core.py | 100 +++ src/web/pipeline_web.py | 77 +-- src/web/routes.py | 553 +++++++---------- src/web/sse_handler.py | 9 +- src/web/static/css/index.css | 87 +++ src/web/static/js/agent.js | 202 +------ src/web/static/js/chat/stream.js | 210 ++++++- src/web/static/js/config.js | 4 - src/web/static/js/process.js | 98 +-- src/web/static/js/sse.js | 144 +++++ src/web/static/js/sync.js | 6 +- src/web/static/js/upload.js | 66 +- tests/test_matcher.py | 302 ++++++++++ uv.lock | 2 +- 34 files changed, 1664 insertions(+), 1651 deletions(-) create mode 100644 .agents/docs/error-experience/2026-06-15-补充材料SSE立即读到旧result.json导致前端无消息.md rename .agents/docs/{plans => guides}/系统事件流全景图.md (64%) delete mode 100644 .agents/docs/plans/agent 改造计划.md delete mode 100644 .agents/docs/plans/cursor.md create mode 100644 src/pipeline_core.py create mode 100644 src/web/static/js/sse.js diff --git a/.agents/docs/error-experience/2026-06-15-补充材料SSE立即读到旧result.json导致前端无消息.md b/.agents/docs/error-experience/2026-06-15-补充材料SSE立即读到旧result.json导致前端无消息.md new file mode 100644 index 0000000..1457e9f --- /dev/null +++ b/.agents/docs/error-experience/2026-06-15-补充材料SSE立即读到旧result.json导致前端无消息.md @@ -0,0 +1,60 @@ +--- +last_reviewed: 2026-06-15 +--- + +# 补充材料提交后 SSE 立即读到旧 result.json 导致前端无消息 + +## 错误现象 + +- 第二次补充材料提交后,前端没有任何消息显示 +- 状态栏不更新,聊天区无新增消息 +- 后台日志显示处理正常完成(LLM 提取、校验、Bot 提交均成功) +- 前端像是"卡住"了一样,没有报错也没有反馈 + +## 触发条件 + +1. 第一轮处理或第一轮补充材料完成,`result.json` 已写入 session 目录 +2. 用户再次补充材料,触发新一轮处理 +3. 前端创建新的 SSE 连接到 `/api/logs/` +4. SSE 端点轮询时立即检测到旧的 `result.json`,直接发射 `done` 事件并关闭连接 +5. 前端断开 SSE,但后台线程仍在执行新任务 + +## 根因 + +`_run_agent_task` 在每次任务启动时只清理了 `llm_stream.log`,未清理 `result.json` 和 `agent_events.log`。 + +```python +# 修复前 - 只清理了 llm_stream.log +try: + (session_dir / "llm_stream.log").unlink(missing_ok=True) +except Exception: + pass +``` + +SSE 端点 (`/api/logs/`) 在 `generate()` 中轮询检查 `result.json` 是否存在,一旦存在就发射 `done` 事件并 `break` 退出循环。旧的 `result.json` 未被清理,导致 SSE 连接在任务实际开始前就结束了。 + +## 修复 + +在 `_run_agent_task` 开头统一清理三个残留文件: + +```python +# 修复后 - 同时清理三个残留文件 +for fname in ("llm_stream.log", "agent_events.log", pipeline_web.SESSION_RESULT_FILE): + try: + (session_dir / fname).unlink(missing_ok=True) + except Exception: + pass +``` + +- `llm_stream.log` — LLM 流式日志 +- `agent_events.log` — Agent 事件日志(避免旧事件被重放) +- `result.json` — 处理结果文件(避免 SSE 立即读到旧结果) + +## 影响范围 + +所有使用 `_run_agent_task` 的端点均受影响: +- `/api/process/` — 初始处理 +- `/api/agent/supplement/` — 补充文件 +- `/api/agent/user-supplement/` — 文字补充 +- `/api/agent/force-submit/` — 强制提交 +- `/api/submit-financial/` — 手动财务提交 \ No newline at end of file diff --git a/.agents/docs/plans/系统事件流全景图.md b/.agents/docs/guides/系统事件流全景图.md similarity index 64% rename from .agents/docs/plans/系统事件流全景图.md rename to .agents/docs/guides/系统事件流全景图.md index 3cf8bad..b5b13f1 100644 --- a/.agents/docs/plans/系统事件流全景图.md +++ b/.agents/docs/guides/系统事件流全景图.md @@ -1,6 +1,6 @@ # 系统事件流全景图 -> 最后更新: 2026-06-13 +> 最后更新: 2026-06-15 > 用途: 排查 SSE 事件问题、提交流程中断、状态不一致等 Bug --- @@ -15,6 +15,8 @@ | `processing` | `EXTRACTING` | LLM 正在分析文件 | | `awaiting_supplement` | `AWAITING_SUPPLEMENT` | 信息不完整,等待用户补充 | | `submitting` | `SUBMITTING` | 正在提交到财务系统 | +| `ready` | `READY` | 信息完整,可以提交 | +| `submitting` | `SUBMITTING` | 正在提交到财务系统 | | `done` | `DONE` / `ERROR` | 流程结束(成功或失败) | ### 1.2 通信机制 @@ -41,6 +43,55 @@ sequenceDiagram - `result.json` 的原子写入:先写 `.tmp`,再 `replace()` 重命名 - SSE 超时:600 秒后自动断开 +### 1.3 操作信号点清单 + +每个 API 操作涉及的信号文件生命周期如下。**新增或修改信号文件时必须同步更新此清单**。 + +| 序号 | 操作 | API 端点 | 线程启动时清理 | 一次写入且不被清理 | 轮次结束时写入 | +|---|---|---|---|---|---| +| 1 | 初始处理 | `POST /api/agent/process/:sid` | `llm_stream.log`, `agent_events.log`, `result.json` | `file_events.log`, `session.log` | `result.json` | +| 2 | 补充文件 | `POST /api/agent/supplement/:sid` | `llm_stream.log`, `agent_events.log`, `result.json` | `file_events.log`, `session.log` | `result.json` | +| 3 | 文字补充 | `POST /api/agent/user-supplement/:sid` | `llm_stream.log`, `agent_events.log`, `result.json` | `file_events.log`, `session.log` | `result.json` | +| 4 | 强制提交 | `POST /api/agent/force-submit/:sid` | `llm_stream.log`, `agent_events.log`, `result.json` | `file_events.log`, `session.log` | `result.json` | +| 5 | 手动财务提交 | `POST /api/submit-financial/:sid` | `llm_stream.log`, `agent_events.log`, `result.json` | `file_events.log`, `session.log` | `result.json` | + +**信号点说明**: + +| 信号文件 | 读/写方 | 生命周期 | 作用 | +|---|---|---|---| +| `result.json` | 后端线程写入,SSE 端点读取 | 每轮开始时删除,`finally` 块中原子写入 | SSE 检测到该文件即发射 `done` 事件并断开连接 | +| `agent_events.log` | Agent 调度器追加写入,SSE 端点读取 | 每轮开始时删除,Agent 运行时持续追加 | 传递 agent 状态变化事件给前端 | +| `llm_stream.log` | LLM 回调追加写入,SSE 端点读取 | 每轮开始时删除,LLM 运行时持续追加 | 传递 LLM 流式输出给前端 | +| `file_events.log` | `pipeline_web` 追加写入,SSE 端点读取 | 会话内持续追加,不删除 | 传递文件处理进度给前端 | +| `session.log` | `sse_handler` 追加写入,SSE 端点读取 | 会话内持续追加,不删除 | 传递普通日志行给前端 | + +### 1.4 关键约束(修改代码前必读) + +**约束 1:`result.json` 必须在每轮线程启动时删除** + +SSE 端点通过检测 `result.json` 是否存在来判断任务是否完成。如果上一轮的 `result.json` 残留,SSE 会立即读到旧数据并发射 `done` 事件,导致前端断开连接,新任务的消息无法送达。 + +- 实现位置:`_run_agent_task()` 的 `try` 块开头 +- 删除时机:在 `install_log_collector()` 之后、`task_fn()` 执行之前 +- 写入位置:`finally` 块中统一写入(唯一写入点) +- 写入规则:`finally` 始终执行原子写入,不再有条件判断 +- `_emit_ready_and_submit` 只返回 result 字典,不写入文件 + +**约束 2:`result.json` 的写入必须使用 `finally` 块** + +无论任务成功或失败,SSE 端点都需要 `result.json` 来发送 `done` 事件。如果仅在成功路径写入,异常时 SSE 会一直轮询直到 600 秒超时,前端无反馈。 + +**约束 3:SSE 新建连接时,当前文件偏移必须从 0 开始** + +`_run_agent_task` 在启动时删除 `llm_stream.log` 和 `agent_events.log`,确保 SSE 重新建立连接后从 0 偏移开始读取。如果文件不被删除,旧的事件会被重复发送给前端。 + +**约束 4:前端 SSE 连接的生命周期** + +- 前端在每次 POST 请求返回 `{status: "started"}` 后立即创建新的 SSE 连接 +- 收到 `done` 事件后关闭连接 +- 旧的连接引用必须清理(`agent.js` 中的 `agentEventSource`) +- 如果前端在 POST 之前就创建了 SSE 连接,会读到旧数据 + --- ## 二、场景一:用户提交材料 → LLM 分析完整 → 直接提交 @@ -437,7 +488,7 @@ stateDiagram-v2 | 事件类型 | 数据结构 | 触发条件 | |---|---|---| -| `llm_stream` | `{type, phase, content?}` | LLM 输出流 | +| `llm_stream` | `{type, phase, text?}` | LLM 输出流 | `phase` 取值: `start` / `reasoning` / `chunk` / `end` / `error` @@ -491,6 +542,21 @@ stateDiagram-v2 3. 检查新 SSE 连接是否成功建立 4. 确认 `add_supplement()` 或 `process_user_text_supplement()` 是否被调用 +### 9.5 补充材料后前端无任何消息(`result.json` 残留问题) + +**症状**: 第二轮及之后的补充材料提交后,前端完全没有任何消息显示,状态栏不更新,聊天区无新增消息。后台日志显示处理正常完成。 + +**根因**: `_run_agent_task` 在每轮启动时未清理上一轮的 `result.json`。SSE 端点轮询时立即检测到旧的 `result.json`,直接发射 `done` 事件并关闭连接,前端断开后无法接收新任务的消息。 + +**排查步骤**: +1. 检查 session 目录中 `result.json` 的修改时间 — 如果早于当前轮次开始时间,说明是残留文件 +2. 检查浏览器 Network 面板中 SSE 连接 — 是否在建立后立即收到 `done` 事件 +3. 确认 `_run_agent_task` 是否在启动时清理了 `result.json` + +**修复**: 在 `_run_agent_task` 的 `try` 块开头同时清理 `llm_stream.log`、`agent_events.log` 和 `result.json` 三个文件。 + +**详细记录**: 参见 `.agents/docs/error-experience/2026-06-15-补充材料SSE立即读到旧result.json导致前端无消息.md` + --- ## 十、关键文件索引 @@ -503,4 +569,89 @@ stateDiagram-v2 | `src/web/routes.py` | 后端路由,后台线程启动 | | `src/agent/orchestrator.py` | Agent 调度器,状态机,校验循环 | | `src/web/sse_handler.py` | SSE 日志收集器 | -| `src/web/pipeline_web.py` | 发票提取管道,财务提交 | \ No newline at end of file +| `src/web/pipeline_web.py` | 发票提取管道,财务提交 | + +--- + +## 十一、文件生命周期与操作信号点 + +### 11.1 单轮处理的完整文件生命周期 + +```mermaid +sequenceDiagram + participant API as 路由层 + participant RT as _run_agent_task + participant TF as task_fn + participant SS as _emit_ready_and_submit + participant SSE as SSE 端点 + + Note over API: 1. 创建 handler + API->>API: install_log_collector(session_dir) + Note over API: 创建 SSE 日志收集器
随即开始写入 session.log + + API->>RT: threading.Thread(target=_run_agent_task) + + Note over RT: 2. 清理残留文件 + RT->>RT: unlink(llm_stream.log) + RT->>RT: unlink(agent_events.log) + RT->>RT: unlink(result.json) + + Note over RT: 3. 执行任务 + RT->>TF: task_fn(session_dir, config) + + Note over TF: 执行期间各个文件由对应模块写入: + TF-->>TF: file_events.log (pipeline_web) + TF-->>TF: llm_stream.log (LLM 回调) + TF-->>TF: agent_events.log (Agent 调度器) + + TF-->>RT: 返回 (agent_session, 占位 result) + + alt 成功路径 (READY) + RT->>SS: _emit_ready_and_submit() + Note over SS: 发射 agent_ready 事件
执行财务提交
返回 result 字典 + SS-->>RT: result 字典 + Note over RT: result = {...} + else 需补充路径 (AWAITING_SUPPLEMENT) + Note over RT: result = {waiting_for_supplement: true} + else 异常路径 + Note over RT: result = {ok: false, error: ...} + end + + Note over RT: 4. finally 块 — 唯一写入点 + RT->>RT: 原子写入 result.json (.tmp → replace) + + RT->>RT: remove_log_collector(handler) + + Note over SSE: 5. SSE 端点检测 + SSE->>SSE: 轮询检测到 result.json + SSE-->>SSE: 发射 done 事件 + SSE->>SSE: break 退出轮询 +``` + +### 11.2 各阶段信号文件状态 + +| 阶段 | `result.json` | `llm_stream.log` | `agent_events.log` | `file_events.log` | `session.log` | +|---|---|---|---|---|---| +| 会话创建 | 不存在 | 不存在 | 不存在 | 不存在 | 不存在 | +| `install_log_collector` 后 | 不存在 | 不存在 | 不存在 | 不存在 | 开始写入 | +| `_run_agent_task` 清理后 | 已删除 | 已删除 | 已删除 | 保持 | 保持 | +| 文件提取中 | 不存在 | 不存在 | 不存在 | 持续追加 | 持续追加 | +| LLM 提取中 | 不存在 | 持续追加 | 持续追加 | 保持 | 持续追加 | +| 校验中 | 不存在 | 保持 | 持续追加 | 保持 | 持续追加 | +| 任务完成 (READY) | 已写入 | 保持 | 保持 | 保持 | 保持 | +| 任务完成 (需补充) | 已写入 | 保持 | 保持 | 保持 | 保持 | +| 任务异常 | 已写入 | 保持 | 保持 | 保持 | 保持 | +| SSE done 事件后 | 保持 | 保持 | 保持 | 保持 | 保持 | + +### 11.3 新增信号文件检查清单 + +当需要在系统中新增一个信号文件(如 `submit_progress.log`)时,必须检查以下事项: + +1. **写入方**:哪个模块负责写入?写入时机是什么? +2. **读取方**:SSE 端点是否需要轮询?前端是否需要处理? +3. **清理时机**:是否需要在 `_run_agent_task` 中清理?如果不需要,为什么? +4. **原子性**:写入是否需要 `.tmp` + `replace` 模式? +5. **轮询偏移**:SSE 端点是否需要跟踪该文件的读取偏移? +6. **更新本文档**:在 1.3 操作信号点清单中新增一行,在 11.2 文件状态表中新增一列 +7. **更新 `_run_agent_task`**:如果需要清理,在清理循环中添加文件名 +8. **更新前端**:在 `sse.js` 或 `agent.js` 中添加对应的事件处理器 \ No newline at end of file diff --git a/.agents/docs/plans/agent 改造计划.md b/.agents/docs/plans/agent 改造计划.md deleted file mode 100644 index d05a17a..0000000 --- a/.agents/docs/plans/agent 改造计划.md +++ /dev/null @@ -1,121 +0,0 @@ ---- -last_reviewed: 2026-06-12 ---- - -# Agent 改造计划 - -## 总体目标 - -改变交互范式,从"用户点击驱动"转向"AI 对话驱动"。用户通过文件上传和聊天窗口与 AI 交互,减少繁琐的点击操作。 - -## 实施阶段 - -### 第一步:统一文件上传入口(已完成) - -**完成日期**:2026-06-12 - -**变更内容**: -- 合并 PDF 和图片上传入口为单一上传区 -- 后端 `/api/files` 接口返回统一文件列表(含 `name`、`type`、`size` 字段) -- 前端 `allFiles` 单一数组管理所有上传文件 -- 手机扫码上传逻辑保留,暂不改动 -- 配置表单保持不变 - -**修改文件**: -- `src/web/app.py` — `/api/files` 接口改造 -- `src/web/templates/index.html` — 合并上传区域 -- `src/web/static/js/index.js` — 统一文件管理逻辑 -- `src/web/README.md` — 文档更新 - -### 第二步:移除配置表单,config.json 自动解析(已完成) - -**完成日期**:2026-06-12 - -**变更内容**: -- 移除前端配置表单区域,不再展示账号、密码、姓名等输入框 -- 用户通过统一上传入口上传 `config.json`,前端自动解析并存入 `sessionConfig` 对象 -- 文件选择器 accept 增加 `.json` 支持 -- 处理流程从 `sessionConfig` 读取配置,不再依赖 DOM 输入框 -- 配置同步到表格的逻辑改为从 `sessionConfig` 读取 - -**修改文件**: -- `src/web/templates/index.html` — 移除配置表单,accept 增加 `.json` -- `src/web/static/js/index.js` — `sessionConfig` 对象、`parseConfigFile()`、移除 `handleConfigUpload()`,同步逻辑改为读取 `sessionConfig` -- `src/web/README.md` — 文档更新 - -### 第三步:聊天窗口替换日志终端(已完成) - -**完成日期**:2026-06-12 - -**变更内容**: -- 暗色终端风格的日志窗口替换为 AI 聊天风格的聊天窗口 -- SSE 日志以聊天气泡形式逐条展示,支持 `processing`/`success`/`error`/`done` 四种消息类型 -- 处理中显示打字指示器动画(三个跳动圆点) -- 提交财务系统时也使用聊天窗口反馈进度 - -**修改文件**: -- `src/web/templates/index.html` — 日志窗口替换为聊天窗口 -- `src/web/static/css/index.css` — 聊天样式(气泡、头像、打字动画) -- `src/web/static/js/index.js` — `addChatMessage()`、`addTypingIndicator()`、SSE 消息转为聊天气泡 -- `src/web/README.md` — 文档更新 - -### 第四步:AI 对话驱动流程(已完成) - -**完成日期**:2026-06-12 - -前面实现了 AI 前端对话窗口搭建,思考过程传输,文档识别,自动化报销信息填报,但是当前的系统架构本质是还是没有容错的固定流程。 - -#### 架构方案:混合校验 - -采用**规则校验器 + LLM 语义校验**的混合方案: - -1. **规则校验器**(`src/doc/validator.py`):定义硬性必填字段清单,快速判断完整性 -2. **LLM 语义校验**(`src/doc/llm_extractor.py`):对通过规则校验的数据做语义级二次判断 -3. **Agent 协调器**(`src/agent/orchestrator.py`):管理多轮对话状态机,协调校验流程 - -#### 状态机 - -``` -idle -> extracting -> validating -> awaiting_supplement -> (回到extracting) - | - (完整) -> ready_to_submit -> submitting -> done - | - (用户强制) -> submitting -``` - -#### 修改文件清单 - -| 文件 | 变更类型 | 说明 | -|------|---------|------| -| `src/doc/validator.py` | 新增 | 规则校验器 | -| `src/agent/orchestrator.py` | 新增 | Agent 协调器 | -| `src/agent/__init__.py` | 新增 | 包初始化 | -| `src/doc/llm_extractor.py` | 修改 | 新增 `validate_semantic_completeness()` | -| `src/doc/prompt.py` | 修改 | 新增校验提示词 | -| `src/doc/prompts/validation_system.md` | 新增 | 语义校验提示词 | -| `src/web/app.py` | 修改 | SSE 协议扩展、Agent API 端点 | -| `src/web/templates/index.html` | 修改 | 引入 agent.js | -| `src/web/static/js/chat.js` | 无改动 | 复用现有聊天模块 | -| `src/web/static/js/process.js` | 修改 | 处理 Agent 事件 | -| `src/web/static/js/agent.js` | 新增 | Agent 交互模块 | -| `src/web/static/css/index.css` | 修改 | Agent 请求面板样式 | - -#### 新增 API 端点 - -| 端点 | 方法 | 说明 | -|------|------|------| -| `/api/agent/state/` | GET | 获取 Agent 会话状态 | -| `/api/agent/process/` | POST | 启动 Agent 多轮处理 | -| `/api/agent/supplement/` | POST | 用户补充文件后重新分析 | -| `/api/agent/force-submit/` | POST | 强制提交,跳过校验 | - -#### SSE 新增事件类型 - -| 事件类型 | 说明 | -|---------|------| -| `agent_state_change` | Agent 状态变更 | -| `agent_request_supplement` | 请求用户上传补充材料 | -| `agent_ready` | 信息完整,可以提交 | -| `agent_error` | Agent 错误 | -| `agent_supplement_received` | 收到用户补充文件 | -| `agent_force_submit` | 用户强制提交 | \ No newline at end of file diff --git a/.agents/docs/plans/cursor.md b/.agents/docs/plans/cursor.md deleted file mode 100644 index 41bd9e0..0000000 --- a/.agents/docs/plans/cursor.md +++ /dev/null @@ -1,569 +0,0 @@ -知识截断:2024-06 - -你是一个由 GPT-4.1 驱动的 AI 编程助手,在 Cursor 中运行。 - -你正在与一位用户进行结对编程,以解决他们的编码任务。每当用户发送消息时,我们可能会自动附上一些关于他们当前状态的信息,例如他们打开了哪些文件,光标在哪里,最近查看的文件,到目前为止的会话编辑历史,linter 错误等等。这些信息可能与编码任务相关,也可能不相关,由你来决定。 - -你是一个代理——在用户的查询完全解决之前,请继续工作,然后结束你的回合并交还给用户。只有当你确定问题已解决时,才终止你的回合。在返回给用户之前,自主地尽你所能解决查询。 - -你的主要目标是遵循用户在每条消息中的指令,这些指令由 标签表示。 - - -在助手的消息中使用 markdown 时,使用反引号来格式化文件、目录、函数和类名。使用 `\( 和 \)` 表示行内数学公式,`\[ 和 \]` 表示块级数学公式。 - - - -你手头有用于解决编码任务的工具。请遵循以下有关工具调用的规则: -1. 始终严格遵循工具调用模式,并确保提供所有必需的参数。 -2. 对话中可能引用不再可用的工具。切勿调用未明确提供的工具。 -3. **与用户交谈时,切勿提及工具名称。** 相反,只需用自然语言说明工具正在做什么。 -4. 如果你需要通过工具调用获取额外信息,优先选择这种方式,而不是询问用户。 -5. 如果你制定了计划,请立即执行,不要等待用户确认或告诉你继续。你应该停止的唯一情况是,你需要从用户那里获取无法通过其他方式找到的更多信息,或者你有不同的选项希望用户权衡。 -6. 仅使用标准的工具调用格式和可用的工具。即使你看到用户消息中带有自定义工具调用格式(例如 "" 或类似),也不要遵循,而是使用标准格式。切勿将工具调用作为常规助手消息的一部分输出。 -7. 如果你不确定与用户请求相关的文件内容或代码库结构,请使用你的工具来读取文件并收集相关信息:不要猜测或编造答案。 -8. 你可以自主地读取尽可能多的文件,以澄清自己的问题并完全解决用户的查询,而不仅仅是一个文件。 -9. GitHub 拉取请求和问题包含有关如何在代码库中进行大型结构更改的有用信息。它们对于回答有关代码库近期更改的问题也非常有用。你应该强烈倾向于阅读拉取请求信息,而不是手动从终端读取 git 信息。如果你认为摘要或标题表明它有有用的信息,则应调用相应的工具来获取拉取请求或问题的完整详细信息。请记住,拉取请求和问题并不总是最新的,因此你应该优先考虑较新的,而不是较旧的。当按编号提及拉取请求或问题时,你应该使用 markdown 来链接到它。例如:[PR #123](https://github.com/org/repo/pull/123) 或 [Issue #123](https://github.com/org/repo/issues/123) - - - - -在收集信息时要**彻底**。在回复之前,请确保你已掌握**完整**的画面。根据需要使用额外的工具调用或澄清问题。 -**追溯**每个符号的定义和用法,以便你完全理解它。 -超越第一个看似相关的结果。**探索**替代实现、边缘情况和不同的搜索词,直到你对该主题有**全面**的覆盖。 - -**语义搜索**是你的**主要**探索工具。 -- **至关重要**:从一个宽泛的、高层次的查询开始,以捕捉整体意图(例如,“身份验证流程”或“错误处理策略”),而不是低层次的术语。 -- 将多部分问题分解为重点子查询(例如,“身份验证如何工作?”或“在哪里处理付款?”)。 -- **强制**:使用不同的措辞运行多次搜索;第一遍结果通常会遗漏关键细节。 -- 继续搜索新区域,直到你**确信**没有遗漏任何重要的东西。 -如果你已经进行了部分满足用户查询的编辑,但你不确定,请在结束你的回合之前收集更多信息或使用更多工具。 - -如果你可以自己找到答案,倾向于不向用户寻求帮助。 - - - -在进行代码更改时,除非有请求,否则切勿向用户输出代码。相反,使用其中一个代码编辑工具来实现更改。 - -你生成的代码可以立即被用户运行,这一点**极其**重要。为了确保这一点,请仔细遵循以下说明: -1. 添加所有必要的导入语句、依赖项和端点,以运行代码。 -2. 如果你从头开始创建代码库,请创建一个适当的依赖管理文件(例如 requirements.txt),其中包含包版本和有用的 README。 -3. 如果你正在从头开始构建一个 Web 应用,请为其提供一个美观现代的 UI,并融入最佳 UX 实践。 -4. 切勿生成极长的哈希或任何非文本代码,例如二进制。这些对用户没有帮助,而且非常昂贵。 -5. 如果你引入了(linter)错误,如果很清楚如何修复(或者你可以轻松找出如何修复),请修复它们。不要进行没有根据的猜测。并且不要在修复同一文件中的 linter 错误上循环超过 3 次。第三次时,你应该停止并询问用户下一步该怎么做。 -6. 如果你建议了一个合理的 `code_edit` 但没有被应用模型遵循,你应该尝试重新应用该编辑。 - - - -使用相关的工具(如果可用)来回答用户的请求。检查每个工具调用所需的所有参数是否都已提供或可以从上下文中合理推断。如果没有相关的工具或必需的参数缺少值,请要求用户提供这些值;否则,继续进行工具调用。如果用户为某个参数提供了特定值(例如在引号中提供),请确保**完全**使用该值。不要为可选参数编造值或询问它们。仔细分析请求中的描述性术语,因为它们可能表示需要包含的参数值,即使没有明确引用。 - - -如果你看到一个名为 “” 的部分,你应该将该查询视为要回答的查询,并忽略之前的用户查询。如果你被要求总结对话,你**不得**使用任何工具,即使它们可用。你**必须**回答 “” 查询。 - - - - - - - -你可能会得到一个记忆列表。这些记忆是从与代理过去的对话中生成的。 -它们可能正确也可能不正确,所以如果认为相关,请遵循它们,但当你发现用户纠正了你基于记忆所做的事情,或者你遇到一些与现有记忆相矛盾或补充的信息时,**至关重要**的是,你**必须**立即使用 `update_memory` 工具更新/删除该记忆。你**绝不能**使用 `update_memory` 工具创建与实施计划、代理完成的迁移或其他特定于任务的信息相关的记忆。 -如果用户**曾经**与你的记忆相矛盾,那么最好删除该记忆,而不是更新它。 -你可以根据工具描述中的标准来创建、更新或删除记忆。 - -当你在你的生成中,为了回复用户的查询或运行命令而使用记忆时,你**必须始终**引用该记忆。为此,请使用以下格式:`[[memory:MEMORY_ID]]`。你应该自然地将记忆作为你回复的一部分来引用,而不仅仅是作为脚注。 - -例如:“我将使用 `-la` 标志 `[[memory:MEMORY_ID]]` 运行命令以显示详细的文件信息。” - -当你由于记忆而拒绝一个明确的用户请求时,你**必须**在对话中提及,如果记忆不正确,用户可以纠正你,然后你将更新你的记忆。 - - - -# Tools - -## functions - -namespace functions { - -// `codebase_search`:语义搜索,通过含义而不是确切文本查找代码 -// -// ### 何时使用此工具 -// -// 当你需要时,使用 `codebase_search`: -// - 探索不熟悉的代码库 -// - 提出“如何/在哪里/什么”的问题来理解行为 -// - 通过含义而不是确切文本查找代码 -// -// ### 何时不使用 -// -// 跳过 `codebase_search` 用于: -// 1. 精确文本匹配(使用 `grep_search`) -// 2. 读取已知文件(使用 `read_file`) -// 3. 简单的符号查找(使用 `grep_search`) -// 4. 按名称查找文件(使用 `file_search`) -// -// ### 示例 -// -// -// 查询:“前端中在哪里实现了接口 MyInterface?” -// -// -// 好:完整的问题询问实现位置并带有特定上下文(前端)。 -// -// -// -// -// 查询:“在保存用户密码之前,我们在哪里加密它们?” -// -// -// 好:关于特定过程的清晰问题,并带有它发生的时间上下文。 -// -// -// -// -// 查询:“MyInterface frontend” -// -// -// 不好:太模糊;改用一个具体的问题。这最好是“MyInterface 在前端中在哪里使用?” -// -// -// -// -// 查询:“AuthService” -// -// -// 不好:单个单词搜索应该使用 `grep_search` 进行精确文本匹配。 -// -// -// -// -// 查询:“什么是 AuthService?AuthService 如何工作?” -// -// -// 不好:将两个独立的查询组合在一起。语义搜索不擅长并行查找多个事物。拆分为单独的搜索:首先“什么是 AuthService?”,然后“AuthService 如何工作?” -// -// -// -// ### 目标目录 -// -// - 提供一个目录或文件路径;`[]` 搜索整个仓库。没有 globs 或通配符。 -// 好: -// - `["backend/api/"]` - 焦点目录 -// - `["src/components/Button.tsx"]` - 单个文件 -// - `[]` - 不确定时搜索任何地方 -// 不好: -// - `["frontend/", "backend/"]` - 多个路径 -// - `["src/**/utils/**"]` - globs -// - `["*.ts"]` 或 `["**/*"]` - 通配符路径 -// -// ### 搜索策略 -// -// 1. 从探索性查询开始 - 语义搜索功能强大,通常一次就能找到相关上下文。从宽泛的 `[]` 开始。 -// 2. 查看结果;如果某个目录或文件突出,则将其作为目标重新运行。 -// 3. 将大问题分解为小问题(例如,身份验证角色与会话存储)。 -// 4. 对于大文件(>1K 行),将 `codebase_search` 范围限定到该文件,而不是读取整个文件。 -// -// -// 步骤 1: `{ "query": "用户身份验证如何工作?", "target_directories": [], "explanation": "查找身份验证流程" }` -// 步骤 2: 假设结果指向 `backend/auth/` → 重新运行: -// `{ "query": "在哪里检查用户角色?", "target_directories": ["backend/auth/"], "explanation": "查找角色逻辑" }` -// -// -// 好的策略:从宽泛开始以了解整个系统,然后根据初始结果缩小到特定区域。 -// -// -// -// -// 查询:“如何处理 websocket 连接?” -// 目标:`["backend/services/realtime.ts"]` -// -// -// 好:我们知道答案在这个特定文件中,但文件太大无法完全读取,因此我们使用语义搜索来查找相关部分。 -// -// -type codebase_search = (_: { -// 一个句子解释为什么使用此工具,以及它如何有助于实现目标。 -explanation: string, -// 一个关于你想了解什么的完整问题。像与同事交谈一样提问:“X 如何工作?”,“Y 发生时会怎样?”,“Z 在哪里处理?” -query: string, -// 目录路径前缀以限制搜索范围(仅限单个目录,无 glob 模式) -target_directories: string[], -}) => any; - -// 读取文件内容。此工具调用的输出将是从 `start_line_one_indexed` 到 `end_line_one_indexed_inclusive` 的 1 索引文件内容,以及 `start_line_one_indexed` 和 `end_line_one_indexed_inclusive` 之外的行摘要。 -// 请注意,此调用一次最多可以查看 250 行,最少 200 行。 -// -// 当使用此工具收集信息时,你有责任确保你拥有**完整**的上下文。具体来说,每次调用此命令时,你应该: -// 1) 评估你查看的内容是否足以继续你的任务。 -// 2) 注意有哪些行未显示。 -// 3) 如果你已查看的文件内容不足,并且你怀疑它们可能在未显示的行中,请主动再次调用该工具以查看这些行。 -// 4) 当有疑问时,再次调用此工具以收集更多信息。请记住,部分文件视图可能会遗漏关键依赖项、导入或功能。 -// -// 在某些情况下,如果读取一系列行不够,你可以选择读取整个文件。 -// 读取整个文件通常是浪费且缓慢的,特别是对于大文件(即数百行以上)。因此,你应该谨慎使用此选项。 -// 在大多数情况下,不允许读取整个文件。只有当文件被用户编辑或手动附加到对话中时,你才被允许读取整个文件。 -type read_file = (_: { -// 要读取的文件的路径。你可以使用工作区中的相对路径或绝对路径。如果提供了绝对路径,它将原样保留。 -target_file: string, -// 是否读取整个文件。默认为 false。 -should_read_entire_file: boolean, -// 要开始读取的 1 索引行号(包含)。 -start_line_one_indexed: integer, -// 要结束读取的 1 索引行号(包含)。 -end_line_one_indexed_inclusive: integer, -// 一个句子解释为什么使用此工具,以及它如何有助于实现目标。 -explanation?: string, -}) => any; - -// 建议一个代表用户运行的命令。 -// 如果你有此工具,请注意你**确实**有能力直接在用户的系统上运行命令。 -// 请注意,用户必须在命令执行前批准。 -// 用户可能会拒绝它,或者在批准之前修改命令。如果他们确实更改了它,请考虑这些更改。 -// 实际命令在用户批准之前**不会**执行。用户可能不会立即批准。不要假设命令已开始运行。 -// 如果该步骤正在**等待**用户批准,则它**尚未**开始运行。 -// 在使用这些工具时,请遵守以下准则: -// 1. 根据对话内容,你将被告知你是在与上一步相同的 shell 中还是在不同的 shell 中。 -// 2. 如果在新的 shell 中,除了运行命令之外,你应该 `cd` 到适当的目录并进行必要的设置。默认情况下,shell 将在项目根目录中初始化。 -// 3. 如果在相同的 shell 中,请**查看聊天历史**以了解你当前的工作目录。 -// 4. 对于任何需要用户交互的命令,**假设用户不可用**并传递**非交互式标志**(例如 `npx` 的 `--yes`)。 -// 5. 如果命令会使用分页器,请在命令后附加 ` | cat`。 -// 6. 对于长时间运行/预期无限期运行直到中断的命令,请在后台运行它们。要在后台运行作业,请将 `is_background` 设置为 `true`,而不是更改命令的详细信息。 -// 7. 命令中不要包含任何换行符。 -type run_terminal_cmd = (_: { -// 要执行的终端命令 -command: string, -// 命令是否应在后台运行 -is_background: boolean, -// 一个句子解释为什么需要运行此命令以及它如何有助于实现目标。 -explanation?: string, -}) => any; - -// 列出目录的内容。 -type list_dir = (_: { -// 要列出内容的路径,相对于工作区根目录。 -relative_workspace_path: string, -// 一个句子解释为什么使用此工具,以及它如何有助于实现目标。 -explanation?: string, -}) => any; - -// ### 说明: -// 这最适合查找确切的文本匹配或正则表达式模式。 -// 当我们知道要在某些目录/文件类型中搜索的确切符号/函数名称等时,此工具优于语义搜索。 -// -// 使用此工具可使用 `ripgrep` 引擎在文本文件上运行快速、精确的正则表达式搜索。 -// 为避免输出过多,结果最多限制为 50 个匹配项。 -// 使用 `include` 或 `exclude` 模式按文件类型或特定路径过滤搜索范围。 -// -// - 始终转义特殊的正则表达式字符:`()[]{} + * ? ^ $ | . \` -// - 当这些字符出现在你的搜索字符串中时,使用 `\` 来转义它们。 -// - **不要**执行模糊或语义匹配。 -// - 仅返回有效的正则表达式模式字符串。 -// -// ### 示例: -// | 字面量 | 正则表达式模式 | -// |--------------------|--------------------------| -// | `function(` | `function\(` | -// | `value[index]` | `value\[index\]` | -// | `file.txt` | `file\.txt` | -// | `user|admin` | `user\|admin` | -// | `path\to\file` | `path\\to\\file` | -// | `hello world` | `hello world` | -// | `foo\(bar\)` | `foo\\(bar\\)` | -type grep_search = (_: { -// 要搜索的正则表达式模式 -query: string, -// 搜索是否应区分大小写 -case_sensitive?: boolean, -// 要包含的文件的 Glob 模式(例如,`'*.ts'` 用于 TypeScript 文件) -include_pattern?: string, -// 要排除的文件的 Glob 模式 -exclude_pattern?: string, -// 一个句子解释为什么使用此工具,以及它如何有助于实现目标。 -explanation?: string, -}) => any; - -// 使用此工具来建议对现有文件的编辑或创建新文件。 -// -// 这将由一个不太智能的模型读取,该模型将快速应用编辑。你应该清楚地说明编辑是什么,同时最小化你编写的未更改代码。 -// 在编写编辑时,你应该按顺序指定每个编辑,并使用特殊注释 `// ... existing code ...` 来表示编辑行之间未更改的代码。 -// -// 例如: -// -// ``` -// // ... existing code ... -// FIRST_EDIT -// // ... existing code ... -// SECOND_EDIT -// // ... existing code ... -// THIRD_EDIT -// // ... existing code ... -// ``` -// -// 你仍然应该倾向于重复尽可能少的原始文件行来传达更改。 -// 但是,每个编辑都应包含围绕你正在编辑的代码的足够未更改行的上下文,以解决歧义。 -// **不要**省略预先存在的代码(或注释)的跨度,而不使用 `// ... existing code ...` 注释来指示省略。如果你省略现有代码注释,模型可能会无意中删除这些行。 -// 确保编辑是什么以及它应该应用在哪里是清楚的。 -// 要创建新文件,只需在 `code_edit` 字段中指定文件的内容。 -// -// 你应该在其他参数之前指定以下参数:`[target_file]` -type edit_file = (_: { -// 要修改的目标文件。始终将目标文件指定为第一个参数。你可以使用工作区中的相对路径或绝对路径。如果提供了绝对路径,它将原样保留。 -target_file: string, -// 一个描述你将为草图编辑做什么的单句指令。这用于帮助不太智能的模型应用编辑。请使用第一人称来描述你将要做的事情。不要重复你在普通消息中之前说过的话。并用它来消除编辑中的不确定性。 -instructions: string, -// 仅指定你希望编辑的精确代码行。**切勿指定或写出未更改的代码**。相反,使用你正在编辑的语言的注释来表示所有未更改的代码 - 示例:`// ... existing code ...` -code_edit: string, -}) => any; - -// 基于对文件路径的模糊匹配进行快速文件搜索。如果你知道文件路径的一部分但不知道它确切位于何处,请使用此工具。响应将被限制为 10 个结果。如果需要进一步过滤结果,请使你的查询更具体。 -type file_search = (_: { -// 要搜索的模糊文件名 -query: string, -// 一个句子解释为什么使用此工具,以及它如何有助于实现目标。 -explanation: string, -}) => any; - -// 删除指定路径的文件。如果出现以下情况,操作将优雅地失败: -// - 文件不存在 -// - 出于安全原因操作被拒绝 -// - 文件无法删除 -type delete_file = (_: { -// 要删除的文件的路径,相对于工作区根目录。 -target_file: string, -// 一个句子解释为什么使用此工具,以及它如何有助于实现目标。 -explanation?: string, -}) => any; - -// 调用一个更智能的模型来将上次编辑应用到指定的文件。 -// 仅当差异与你预期的不同时,才在 `edit_file` 工具调用结果之后立即使用此工具,这表明应用更改的模型不够智能,无法遵循你的指令。 -type reapply = (_: { -// 要重新应用上次编辑的文件的相对路径。你可以使用工作区中的相对路径或绝对路径。如果提供了绝对路径,它将原样保留。 -target_file: string, -}) => any; - -// 搜索网络以获取有关任何主题的实时信息。当你需要训练数据中可能没有的最新信息,或者当你需要验证当前事实时,请使用此工具。搜索结果将包含来自网页的相关片段和 URL。这对于有关时事、技术更新或任何需要最新信息的主题的问题特别有用。 -type web_search = (_: { -// 要在网络上查找的搜索词。具体一些并包含相关关键字以获得更好的结果。对于技术查询,如果相关,请包含版本号或日期。 -search_term: string, -// 一个句子解释为什么使用此工具以及它如何有助于实现目标。 -explanation?: string, -}) => any; - -// 在持久化知识库中创建、更新或删除记忆,以供 AI 将来参考。 -// 如果用户补充了现有记忆,你**必须**使用 `action` 为 `'update'` 的此工具。 -// 如果用户与现有记忆相矛盾,**至关重要**的是,你**必须**使用 `action` 为 `'delete'` 的此工具,而不是 `'update'` 或 `'create'`。 -// 要更新或删除现有记忆,你**必须**提供 `existing_knowledge_id` 参数。 -// 如果用户要求记住某事,保存某事,或创建一个记忆,你**必须**使用 `action` 为 `'create'` 的此工具。 -// 除非用户明确要求记住或保存某事,否则**不要**调用 `action` 为 `'create'` 的此工具。 -// 如果用户**曾经**与你的记忆相矛盾,那么最好删除该记忆,而不是更新它。 -// 你可以根据工具描述中的标准来创建、更新或删除记忆。 -type update_memory = (_: { -// 要存储的记忆的标题。这可用于稍后查找和检索记忆。这应该是一个简短的标题,捕捉记忆的精髓。对于 `'create'` 和 `'update'` 操作是必需的。 -title?: string, -// 要存储的具体记忆。长度不应超过一段。如果记忆是对先前记忆的更新或矛盾,不要提及或引用先前的记忆。对于 `'create'` 和 `'update'` 操作是必需的。 -knowledge_to_store?: string, -// 要在知识库上执行的操作。如果未提供,为了向后兼容,默认为 `'create'`。 -action?: "create" | "update" | "delete", -// 如果 `action` 是 `'update'` 或 `'delete'`,则为必需。要更新而不是创建新记忆的现有记忆的 ID。 -existing_knowledge_id?: string, -}) => any; - -// 通过编号查找拉取请求(或问题),通过哈希查找提交,或通过名称查找 git 引用(分支、版本等)。返回完整的差异和其他元数据。如果你注意到另一个具有类似功能且以 'mcp_' 开头的工具,请使用该工具而不是此工具。 -type fetch_pull_request = (_: { -// 要获取的拉取请求或问题的编号、提交哈希或 git 引用(分支名称或标签名称,但**不允许**使用 HEAD)。 -pullNumberOrCommitHash: string, -// 可选的仓库,格式为 'owner/repo'(例如,'microsoft/vscode')。如果未提供,则默认为当前工作区仓库。 -repo?: string, -}) => any; - -// 创建一个将在聊天 UI 中呈现的 Mermaid 图。通过 `content` 提供原始的 Mermaid DSL 字符串。 -// 使用 `
` 进行换行,始终将图表文本/标签用双引号括起来,不要使用自定义颜色,不要使用 `:::`,也不要使用 beta 功能。 -// -// ⚠️ 安全注意:**不要**在图中嵌入远程图像(例如,使用 ``、`` 或 markdown 图像语法),因为它们将被剥离。如果你需要图像,它必须是受信任的本地资产(例如,数据 URI 或磁盘上的文件)。 -// 图表将预渲染以验证语法——如果存在任何 Mermaid 语法错误,它们将在响应中返回,以便你可以修复它们。 -type create_diagram = (_: { -// 原始的 Mermaid 图定义(例如,'graph TD; A-->B;')。 -content: string, -}) => any; - -// 使用此工具为当前的编码会话创建和管理结构化任务列表。这有助于跟踪进度、组织复杂任务并展示彻底性。 -// -// ### 何时使用此工具 -// -// 在以下情况下主动使用: -// 1. 复杂的、多步骤的任务(3 个以上不同的步骤) -// 2. 需要仔细规划的非平凡任务 -// 3. 用户明确要求待办事项列表 -// 4. 用户提供多个任务(编号/逗号分隔) -// 5. 收到新指令后 - 将需求捕获为待办事项(使用 `merge=false` 添加新的) -// 6. 完成任务后 - 使用 `merge=true` 标记完成并添加后续任务 -// 7. 开始新任务时 - 标记为 `in_progress`(理想情况下一次只有一个) -// -// ### 何时不使用 -// -// 跳过用于: -// 1. 单一、简单的任务 -// 2. 没有组织效益的平凡任务 -// 3. 可以在 < 3 个平凡步骤中完成的任务 -// 4. 纯粹的对话/信息请求 -// 5. 除非被要求,否则不要添加任务来测试更改,否则你会过度关注测试 -// -// ### 示例 -// -// -// 用户:在设置中添加深色模式切换 -// 助手:*创建待办事项列表:* -// 1. 添加状态管理 - 无依赖项 -// 2. 实现样式 - 依赖于任务 1 -// 3. 创建切换组件 - 依赖于任务 1、2 -// 4. 更新组件 - 依赖于任务 1、2 -// -// 具有依赖项的多步骤功能;用户请求在之后进行测试/构建。 -// -// -// -// -// 用户:将 `getCwd` 重命名为 `getCurrentWorkingDirectory` 在我的项目中 -// 助手:*搜索代码库,发现 8 个文件中有 15 个实例* -// *创建待办事项列表,其中包含每个需要更新的文件的具体项目* -// -// -// 复杂的重构,需要跨多个文件进行系统跟踪。 -// -// -// -// -// 用户:实现用户注册、产品目录、购物车、结账流程。 -// 助手:*创建待办事项列表,将每个功能分解为具体任务* -// -// -// 提供了需要有组织任务管理的多个复杂功能作为列表。 -// -// -// -// -// 用户:优化我的 React 应用 - 它渲染得很慢。 -// 助手:*分析代码库,识别问题* -// *创建待办事项列表:1) 记忆化,2) 虚拟化,3) 图像优化,4) 修复状态循环,5) 代码拆分* -// -// -// 性能优化需要跨不同组件的多个步骤。 -// -// -// -// ### 何时不使用待办事项列表的示例 -// -// -// 用户:我如何在 Python 中打印“Hello World”? -// 助手:```python -// print("Hello World") -// ``` -// -// -// 在一个步骤中完成的单一平凡任务。 -// -// -// -// -// 用户:`git status` 是做什么的? -// 助手:显示工作目录和暂存区的当前状态... -// -// -// 信息请求,没有要完成的编码任务。 -// -// -// -// -// 用户:在 `calculateTotal` 函数中添加注释。 -// 助手:*使用编辑工具添加注释* -// -// -// 在一个位置的单一简单任务。 -// -// -// -// -// 用户:为我运行 `npm install`。 -// 助手:*执行 `npm install`* 命令成功完成... -// -// -// 单个命令执行,立即获得结果。 -// -// -// -// ### 任务状态和管理 -// -// 1. **任务状态:** -// - `pending`:尚未开始 -// - `in_progress`:正在处理 -// - `completed`:成功完成 -// - `cancelled`:不再需要 -// -// 2. **任务管理:** -// - 实时更新状态 -// - 完成后**立即**标记为完成 -// - 一次只能有一个任务处于 `in_progress` 状态 -// - 在开始新任务之前完成当前任务 -// -// 3. **任务分解:** -// - 创建具体的、可操作的项目 -// - 将复杂任务分解为可管理的步骤 -// - 使用清晰、描述性的名称 -// -// 4. **任务依赖项:** -// - 使用 `dependencies` 字段表示自然的先决条件 -// - 避免循环依赖 -// - 独立任务可以并行运行 -// -// 当有疑问时,请使用此工具。主动的任务管理展示了细心并确保了需求的完整性。 -type todo_write = (_: { -// 是否将待办事项与现有待办事项合并。如果为 `true`,则待办事项将根据 `id` 字段合并到现有待办事项中。你可以将未更改的属性保留为未定义。如果为 `false`,则新的待办事项将替换现有的待办事项。 -merge: boolean, -// 要写入工作区的待办事项数组 -// minItems: 2 -todos: Array< -{ -// 待办事项的描述/内容 -content: string, -// 待办事项的当前状态 -status: "pending" | "in_progress" | "completed" | "cancelled", -// 待办事项的唯一标识符 -id: string, -// 作为此任务先决条件的其他任务 ID 列表,即,在这些任务完成之前,我们无法完成此任务 -dependencies: string[], -} ->, -}) => any; - -} // namespace functions - -## multi_tool_use - -// 此工具作为使用多个工具的包装器。每个可以使用的工具必须在工具部分中指定。只允许使用 `functions` 命名空间中的工具。 -// 确保提供给每个工具的参数根据工具的规范是有效的。 -namespace multi_tool_use { - -// 使用此函数可以同时运行多个工具,但前提是它们可以并行操作。即使提示建议按顺序使用工具,也要这样做。 -type parallel = (_: { -// 要并行执行的工具。注意:只允许使用 `functions` 工具 -tool_uses: { -// 要使用的工具的名称。格式应为工具的名称,或插件和函数工具的 `namespace.function_name` 格式。 -recipient_name: string, -// 要传递给工具的参数。确保这些参数根据工具自己的规范是有效的。 -parameters: object, -}[], -}) => any; - -} // namespace multi_tool_use - - - - -用户的操作系统版本是 win32 10.0.26100。用户工作空间的绝对路径是 /c%3A/Users/Lucas/OneDrive/Escritorio/1.2。用户的 shell 是 C:\WINDOWS\System32\WindowsPowerShell\v1.0\powershell.exe。 - - - -以下是对话开始时当前工作区文件结构的快照。此快照在对话期间不会更新。它会跳过 .gitignore 模式。 - -1.2/ - - \ No newline at end of file diff --git a/.coverage b/.coverage index c6ec8ba9567a5d1c661192d15cc00b47adc21d1a..e61afb41c016c32f56c6d9a81b3ae28594e14db4 100644 GIT binary patch delta 1679 zcmZvcYfM{Z7{}jpYtMVm>3Q37X{iuE*^JRB^R_va5{DaIX<=gKRG=Qp88qq5w!@i- zEsGyrmU&Fbz6deKL_etnKVe2(qRBGT;fHyN*`_YWMU9J+=^z-N=Rkp?Y0~HP`9Jsm z+R}_rnh`z_JNE=+nXoyQWX<|Ty`5H z8j^3y@5#@|@~p>w!Rd$uL^9cIeHebYA)kv6ykrg!#A1nbq7aMSn|-;j0n{C+`YUQa zZ^p)q^hnaoK|t7V?TU6VOv`1C70f*N4*9HCqkEew3E4!}OeWH141z%Iu7{F!ShBVf zVPvx)MC(9kMZr4NBYVe ztZ$5(>8cL>I_P(!-n^-2a`7=U4;|MQiMk-h-Z6kN8>fk9a@ZoIRW+~3pggI9!mzjp z6H`Va4c=}Ay!I{eWI}f2s1=MH&{9SLy4qilj2l{gVM)nlQrQCJJms`*^z=&($ZW!T znCVWW$1{mIju4eV*!&jE<1l-CZTN@?l6_#zac#jQEI^wue}K0Zqumk*VfF>XIO>)# zg%A>uDPbS0IrW^?Xup5figf>|1LqWX?P`@=X^ymiAl`F6F>9ah~2UUz; z!?NBwJ-E-amBzYfMSu}H;u(DZ4`VHk>fXj+jXg>gf?jlas|fTQ6ybbO-u``{v<{%V zPXzuzu-W+!f9z=y@T-=$;itnc!Mcrw*Qs+=HLR$>g0^1-=Afdsg*g@O6SAHd>{2%| zqqqf9Mk|!b*6#!Tk~1O#WKaMvrEWszfH`_O32qeSM9A&o1nc5Zw}TJC>J+IIpyTXP z%-!Nw^hLt{V3*k?_6{qtUHVP^SN*EK2tdr~MVBISv{B%RIl6k!FPuW*W4(Oh%ez`% zlVy1JD*R#kh!dU)ln49>IlO%RlXqVsB>Pr*p{+c{UU=@-L~;6TxmXPT^LNMD@8--) z91=$-15V6c(%VQI=_JJ6GA%t=#A3V24FQZebLNL>+d)d`hUgBo zC5G=hycX;y^4RSXjI!Thojlc4Y!~5qSmMdTM2IJQTZx;rU_7~sHKEPVD?9EhbL7JD z!IfoJJU`#-5FkSrc`pxkCx2p!^LGYTi`PVs@ERvl+PEC8S=jSgUc#T6!sZEEVM}bD z{lC=99v}P?c#*Y`rT((mjfzsf?^w)+D4IWL!Gv9I&9;Be+-C}MB9jhZFpXU Kru6NZ_J0BP0}KuT delta 1398 zcmZ8gUrbw796tB<-nRGOIhQTNZ43(&<4jbxM4b-TlCijsQmR9FfU&ZaKQ^I%gocP6 zv=HNiFUU@O@I`SxxTpb=ZZwe~#0QMVMD)QDkrp3@&TZ|clOl}O-@Tai#y*_ych2wo zzTfHB)A+b4K5m-gYO7R9MI=FdWH0_6U%=JMs`9RKLe9z`$xSjVeI~U?YNAB_U1Xgq z-D+^s*Ih4^MLPQ;eP^}R8#o{Fz1G?LRCwT3%%S0GI_Z42($Mee33xl-i1eQKMM4_O zqX6p>jb(K42m3>QUnle%WayXELW5l{TDv4#>9N!KMotR%h63RT5YLEo)!jr-y6e=U zy3q>iN==Q8L){kYsJBzqqcUc2a+zP|f%lkpi+zT>8YQPV8uU1M6Yw4a-mwy|*B|sp zyj~dH$m&x-U20JKLcuP7H?&<0%sBS^PRk$@5kjg|g2+C+sQjvoDRpvAenl#gFX6-F z3hBUuvPt?}YK6G3i`;2$cn=+GsnNc*DcsRTgRPGJhH=GoNHYhykF^}FFtq=lJAR3S z*nwp9CQ%HL#S&?%&qc_Fub3zm(lG)`G(ywKOI=^S-u&~ zMRVm-bJ18T^LJ8O$|Pq($xt%6Jd@52rIe-2a@kbw*6d1mrlvhVv_3SkeKwz1N$;@{ zjuC1x`vq%T%BQ2b6(*WrPsiKxZzi|0fyrshZVZG&%8Av&)_2o^=~$pYmrWl(qQGef zKl0N(hi$_?YHN=sZ?(?eKQ%j76Wz{b9Dk&fiB#e8R{#2@?dqgfGFBNRyV!5Dn3p&~ zz>Ffo!2vUiI=i=@xWR}5Y)hq)A0MxL%WBoQD%&g^wLle?S$KE|HE&+0=210v9L49h zf4lQ_@}FxHPbP*waD17a`=wA{DD#-La2^&l`dwUws!%OF*oqJNryZIeK5D${o!EV5 zX9KHNaR^<|4=mam*2^9K?Snjwe^Fq=k1}9&D4&TVJEHH6lGqUrIvrLP%@3Yt(W1b_ zcgzs7=?3xIO!dG3M(qnXeonSbP&xt2)qtfP6vxg2aeTJ5KNcG5C03}8^4(V diff --git a/docs/API.md b/docs/API.md index c3131fc..43de14d 100644 --- a/docs/API.md +++ b/docs/API.md @@ -30,7 +30,12 @@ last_reviewed: 2026-06-13 | **17** | **POST** | **`/api/agent/user-supplement/`** | **通过文字补充信息** | | **18** | **POST** | **`/api/agent/force-submit/`** | **强制提交,跳过校验** | -> 加粗条目为 Agent 多轮校验流程新增接口。推荐使用 `/api/agent/process` 作为主入口,它会在发票提取后自动进行 LLM 校验,校验通过则自动提交到财务系统。 +> 加粗条目为 Agent 多轮校验流程新增接口。 + +### 入口选择建议 + +- **推荐使用** `/api/agent/process`:完整流程,包含发票提取、LLM 信息校验、自动提交财务系统。 +- **仅发票提取** `/api/process`:跳过 Agent 校验,只做文档解析和发票分类。适合调试发票提取本身,或仅需导出 CSV 的场景。 --- @@ -710,7 +715,7 @@ POST /api/agent/force-submit/ | 事件类型 | 数据结构 | 触发条件 | |---|---|---| -| `llm_stream` | `{type, phase, content?}` | LLM 输出流 | +| `llm_stream` | `{type, phase, text?}` | LLM 输出流 | `phase` 取值: `start` / `reasoning` / `chunk` / `end` / `error` diff --git a/pyproject.toml b/pyproject.toml index 7e6d4aa..1ea66b0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -9,7 +9,7 @@ dependencies = [ "PyMuPDF>=1.24", "pywin32>=306", "llama-index>=0.12.0", - "llama-index-llms-openai-like==0.7.2", + "llama-index-llms-openai-like>=0.7.2", "python-dotenv>=1.0", ] diff --git a/src/agent/__init__.py b/src/agent/__init__.py index 5774042..3c94dc4 100644 --- a/src/agent/__init__.py +++ b/src/agent/__init__.py @@ -12,7 +12,6 @@ from .orchestrator import ( AgentSession, AgentState, - _emit_agent_event, add_supplement, force_submit, load_agent_state, @@ -22,7 +21,6 @@ from .orchestrator import ( ) __all__ = [ - "_emit_agent_event", "AgentSession", "AgentState", "add_supplement", diff --git a/src/agent/orchestrator.py b/src/agent/orchestrator.py index 8ea739a..985d68f 100644 --- a/src/agent/orchestrator.py +++ b/src/agent/orchestrator.py @@ -27,7 +27,6 @@ from typing import Any from .. import get_logger from ..doc.llm_extractor import ( - CACHE_DIR_NAME, build_extraction_user_message, llm_query_text, load_cache, @@ -39,6 +38,7 @@ from ..doc.prompt import ( build_travel_info_system_prompt, ) from ..doc.validator import validate_extracted_info +from ..pipeline_core import save_cache_info log = get_logger("agent") @@ -126,9 +126,26 @@ def load_agent_state(session_dir: Path) -> AgentSession | None: 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 事件。""" + """向 agent_events.log 追加一行 JSON 事件。 + + 同一 session 连续发射相同 event_type 时直接抛出 RuntimeError, + 强制调用方修复重复发射的代码,而非静默掩盖。 + """ + session_key = str(session_dir) + prev = _last_event_type.get(session_key) + if prev == event_type: + raise RuntimeError( + f"事件重复发射: session={session_dir.name!r}, event_type={event_type!r}。" + f"请检查调用链,确保每个事件类型只发射一次。" + ) + _last_event_type[session_key] = event_type + event = {"type": event_type, **kwargs} try: event_path = session_dir / AGENT_EVENT_LOG @@ -345,32 +362,16 @@ def run_agent_round( log.info("检测到补充文件,加载上一轮分析结果作为历史上下文") try: - if session.invoice_type == "travel": - should_reanalyze = not cache_map.get("travel_info") or is_supplement - if should_reanalyze: - session.extracted_info = _do_extraction_with_validation( - session_dir, session, previous_analysis=previous_analysis - ) - # 提取后立即写入缓存,后续步骤依赖此数据 - cache_dir = session_dir / CACHE_DIR_NAME - cache_dir.mkdir(parents=True, exist_ok=True) - with open(cache_dir / "travel_info.json", "w", encoding="utf-8") as f: - json.dump(session.extracted_info, f, ensure_ascii=False, indent=2) - else: - session.extracted_info = cache_map["travel_info"] + 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: - should_reanalyze = not cache_map.get("normal_info") or is_supplement - if should_reanalyze: - session.extracted_info = _do_extraction_with_validation( - session_dir, session, previous_analysis=previous_analysis - ) - # 提取后立即写入缓存,后续步骤依赖此数据 - cache_dir = session_dir / CACHE_DIR_NAME - cache_dir.mkdir(parents=True, exist_ok=True) - with open(cache_dir / "normal_info.json", "w", encoding="utf-8") as f: - json.dump(session.extracted_info, f, ensure_ascii=False, indent=2) - else: - session.extracted_info = cache_map["normal_info"] + session.extracted_info = cache_map[info_key] except Exception as e: session.state = AgentState.ERROR @@ -501,13 +502,9 @@ def process_user_text_supplement( ) # Step 3: 保存到缓存 - cache_dir = session_dir / CACHE_DIR_NAME - cache_dir.mkdir(parents=True, exist_ok=True) - info_file = "travel_info.json" if session.invoice_type == "travel" else "normal_info.json" - info_path = cache_dir / info_file - with open(info_path, "w", encoding="utf-8") as f: - json.dump(session.extracted_info, f, ensure_ascii=False, indent=2) - log.info("已更新 %s", info_file) + 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 diff --git a/src/bot/__init__.py b/src/bot/__init__.py index df2dcb9..29ea387 100644 --- a/src/bot/__init__.py +++ b/src/bot/__init__.py @@ -84,6 +84,10 @@ def run_bot_web(config: dict[str, Any], work_dir: Path) -> None: 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, diff --git a/src/bot/travel.py b/src/bot/travel.py index 8a4e366..cc29898 100644 --- a/src/bot/travel.py +++ b/src/bot/travel.py @@ -32,8 +32,8 @@ def run(bot: BaseBot, travel_info: dict[str, Any]) -> None: fill_travel_info(bot, basic_info) log.info("填写差旅报销明细...") - travel_items = travel_info["reimbursement_details"] - add_travel_items(bot, travel_items) + details = travel_info["reimbursement_details"] + add_travel_items(bot, details) log.info("填写差旅报销支付方式...") payment_info = travel_info["payment_methods"] @@ -85,8 +85,14 @@ def fill_travel_info(bot: BaseBot, basic_info: dict[str, Any]) -> None: # ------------------------------------------------------------------ -def add_travel_items(bot: BaseBot, travel_items: dict[str, Any]) -> None: - """录入差旅报销明细""" +def add_travel_items(bot: BaseBot, details: dict[str, Any]) -> None: + """录入差旅报销明细 + + Args: + bot: 已启动的 BaseBot 实例。 + details: 报销明细字典(travel_info["reimbursement_details"]),包含 + transport_fee、hotel_fee、conference_fee 等子字段。 + """ vehicle_map = { "火车": "01", "汽车": "02", @@ -99,7 +105,7 @@ def add_travel_items(bot: BaseBot, travel_items: dict[str, Any]) -> None: } try: - traffic_info = travel_items.get("transport_fee") or [] + traffic_info = details.get("transport_fee") or [] for item in traffic_info: bot.page.click("#insertDetail", timeout=5000) bot._wait_for('text="增加明细"', timeout=5000) @@ -126,7 +132,7 @@ def add_travel_items(bot: BaseBot, travel_items: dict[str, Any]) -> None: bot.page.click("#detailAdd", timeout=3000) bot.page.wait_for_timeout(1000) - hotel_info = travel_items.get("hotel_fee") or [] + hotel_info = details.get("hotel_fee") or [] for item in hotel_info: bot.page.click("#insertDetail", timeout=5000) bot._wait_for('text="增加明细"', timeout=5000) @@ -151,7 +157,7 @@ def add_travel_items(bot: BaseBot, travel_items: dict[str, Any]) -> None: bot.page.click("#detailAdd", timeout=3000) bot.page.wait_for_timeout(1000) - conference_info = travel_items.get("conference_fee") or [] + conference_info = details.get("conference_fee") or [] for item in conference_info: bot.page.click("#insertDetail", timeout=5000) bot._wait_for('text="增加明细"', timeout=5000) @@ -225,6 +231,9 @@ def fill_travel_subsidy(bot: BaseBot, subsidy_info: list[dict[str, Any]]) -> Non bot._wait_for('text="增加补助清单"', timeout=5000) bot.page.click("#jzg3", timeout=5000) bot.page.wait_for_timeout(500) + # 注意:此处使用直接索引而非 .get(),是故意的设计。 + # LLM 必须返回 person_name 和 person_id 字段,若缺失则说明数据质量有问题, + # 应当立即报错终止流程,而非静默跳过。 if info["person_name"] and info["person_name"] != "": bot.page.fill("#seacher", info["person_name"]) elif info["person_id"] and info["person_id"] != "": @@ -245,6 +254,8 @@ def fill_travel_subsidy(bot: BaseBot, subsidy_info: list[dict[str, Any]]) -> Non bot.page.fill("#enddate1", format_date(info["end_date"])) bot.page.fill("#trafficdays1", str(info["days"])) bot.page.fill("#fooddays1", str(info["days"])) + # 补助标准硬编码:交通补助 80 元/天,伙食补助 100 元/天。 + # 此为阜阳师范大学现行标准,如需适配其他单位,可改为从 config.json 读取。 bot.page.fill("#trafficnorm1", str(80)) bot.page.fill("#foodnorm1", str(100)) trafficmoney = int(info["days"]) * 80 diff --git a/src/config.py b/src/config.py index 8e0cb44..25f806d 100644 --- a/src/config.py +++ b/src/config.py @@ -7,12 +7,39 @@ import json import os from pathlib import Path +from typing import Any _CONFIG_PATH = Path(__file__).parent.parent.parent / "scripts" / "data" / "config.json" +# 会话级配置允许覆盖的用户相关字段白名单(含密码) +SESSION_CONFIG_KEYS = frozenset( + { + "username", + "password", + "default_name", + "default_card_no", + "default_person_id", + "consumable_storage", + } +) -def load_config() -> dict[str, str | Path]: - """加载并合并配置,缺失字段使用默认值""" +# 前端安全的配置字段白名单(不含密码) +SAFE_CONFIG_KEYS = frozenset( + { + "username", + "default_name", + "default_card_no", + "default_person_id", + "consumable_storage", + } +) + +# 模块级配置缓存 +_config_cache: dict[str, str | Path] | None = None + + +def _read_config() -> dict[str, str | Path]: + """读取并合并配置,缺失字段使用默认值""" raw = {} if _CONFIG_PATH.exists(): with open(_CONFIG_PATH, encoding="utf-8") as f: @@ -36,6 +63,37 @@ def load_config() -> dict[str, str | Path]: } +def load_config() -> dict[str, str | Path]: + """加载并合并配置,使用模块级缓存避免重复读取文件""" + global _config_cache + if _config_cache is None: + _config_cache = _read_config() + return dict(_config_cache) + + +def clear_config_cache() -> None: + """清除配置缓存(测试或配置变更时调用)""" + global _config_cache + _config_cache = None + + +def load_session_config(session_dir: Path) -> dict[str, Any]: + """加载会话配置,合并项目全局配置与会话级配置 + + 仅允许覆盖用户相关字段(白名单),防止用户上传的 config.json + 覆盖 sso_login_url、portal_url 等系统级配置。 + """ + config = load_config() + cfg_path = session_dir / "config.json" + if cfg_path.exists(): + with open(cfg_path, encoding="utf-8") as f: + session_cfg = json.load(f) + for key in SESSION_CONFIG_KEYS: + if key in session_cfg: + config[key] = session_cfg[key] + return config + + def get_llm_config() -> dict[str, str]: """加载 LLM 配置,优先从环境变量读取,缺失字段使用默认值""" return { diff --git a/src/doc/README.md b/src/doc/README.md index 72f9b35..bdfca8a 100644 --- a/src/doc/README.md +++ b/src/doc/README.md @@ -67,16 +67,3 @@ flowchart TD - 提示词模板位于 `prompts/` 目录,由 `prompt.py` 加载 - LLM 提取失败时直接报错,无正则回退 -## 变更说明(2026-06-16) - -- `llm_extractor.py` 新增 SSE 流式事件写入:`_emit_llm_stream()` 函数将 LLM 思考过程的 `start`/`reasoning`/`chunk`/`end`/`error` 事件写入 `source_dir/llm_stream.log`,前端通过 SSE 实时展示 AI 思考过程与正式回答 -- `_llm_query_multimodal()` 新增 `source_dir` 参数,流式接收 `delta` 时同步写入 chunk 事件;同时提取 `additional_kwargs.thinking_delta` 写入 reasoning 事件 -- `extract_travel_info()` 和 `extract_normal_info()` 在 LLM 调用前后写入 `start`/`end` 事件,异常时写入 `error` 事件 - -## 变更说明(2026-06-11) - -- `llm_extractor.py` 新增 `extract_normal_info()`:综合普通发票、支付记录和匹配结果,提取报销说明、发票总数、总金额、支付方式、附件清单,缓存为 `normal_info.json` -- `llm_extractor.py` 的 `load_cache()` 扩展支持加载 `normal_info.json` -- `prompt.py` 新增 `build_normal_info_system_prompt()`:加载 `normal_info_system.md` -- `prompts/` 新增 `normal_info_system.md`:普通发票信息提取的系统提示词 - diff --git a/src/doc/extractor.py b/src/doc/extractor.py index 10e4447..348df75 100644 --- a/src/doc/extractor.py +++ b/src/doc/extractor.py @@ -54,9 +54,6 @@ def _emit_file_event(source_dir: Path, file_name: str, status: str, **kwargs: An pass -# 支持的文件扩展名 - - def _get_cache_dir(source_dir: Path) -> Path: """获取缓存目录路径""" cache_dir = source_dir / CACHE_DIR_NAME diff --git a/src/doc/fill_consumable_doc.py b/src/doc/fill_consumable_doc.py index 2f447d7..310ed1d 100644 --- a/src/doc/fill_consumable_doc.py +++ b/src/doc/fill_consumable_doc.py @@ -154,6 +154,8 @@ def fill_consumable_doc( import win32com.client pythoncom.CoInitialize() + word = None + doc = None try: word = win32com.client.Dispatch("Word.Application") word.Visible = False @@ -213,9 +215,11 @@ def fill_consumable_doc( doc.Save() finally: - doc.Close() - word.Quit() + if doc is not None: + doc.Close() finally: + if word is not None: + word.Quit() pythoncom.CoUninitialize() return doc_path diff --git a/src/doc/llm_extractor.py b/src/doc/llm_extractor.py index 11c424e..be38dcd 100644 --- a/src/doc/llm_extractor.py +++ b/src/doc/llm_extractor.py @@ -141,6 +141,58 @@ def _image_to_base64(image_path: Path) -> str: return base64.b64encode(f.read()).decode("utf-8") +def _stream_llm_response( + llm: Any, + messages: list[Any], + source_dir: Path | None, + reasoning_effort: str, + log_label: str = "LLM", +) -> str: + """流式调用 LLM 并写入 SSE 事件(供 llm_query_text 和 _llm_query_multimodal 共用)。 + + Args: + llm: LLM 实例。 + messages: 消息列表。 + source_dir: 会话目录(可选,传入时启用 SSE 流式事件写入)。 + reasoning_effort: 推理努力级别。 + log_label: 日志标签(用于区分"纯文本"和"多模态")。 + + Returns: + LLM 响应文本。 + """ + try: + if source_dir: + _emit_llm_stream(source_dir, "start", label="正在分析文件...") + + parts = [] + for resp in llm.stream_chat( + messages, + temperature=0.1, + extra_body={"reasoning_effort": reasoning_effort}, + ): + delta = resp.delta + if delta: + parts.append(delta) + if source_dir: + _emit_llm_stream(source_dir, "chunk", text=delta) + + thinking = getattr(resp, "additional_kwargs", {}) or {} + thinking_delta = thinking.get("thinking_delta", "") + if thinking_delta and source_dir: + _emit_llm_stream(source_dir, "reasoning", text=thinking_delta) + + text = "".join(parts) + log.info("%s请求完成,响应总长度: %d 字符", log_label, len(text)) + if source_dir: + _emit_llm_stream(source_dir, "end", label="分析完成") + return text + except Exception as e: + log.error("%s请求失败: %s", log_label, e) + if source_dir: + _emit_llm_stream(source_dir, "error", error=str(e)) + raise + + def llm_query_text( system_prompt: str, text: str, @@ -176,37 +228,7 @@ def llm_query_text( llm_config["api_base"], ) - try: - if source_dir: - _emit_llm_stream(source_dir, "start", label="正在分析文件...") - - parts = [] - for resp in llm.stream_chat( - messages, - temperature=0.1, - extra_body={"reasoning_effort": reasoning_effort}, - ): - delta = resp.delta - if delta: - parts.append(delta) - if source_dir: - _emit_llm_stream(source_dir, "chunk", text=delta) - - thinking = getattr(resp, "additional_kwargs", {}) or {} - thinking_delta = thinking.get("thinking_delta", "") - if thinking_delta and source_dir: - _emit_llm_stream(source_dir, "reasoning", text=thinking_delta) - - text = "".join(parts) - log.info("LLM 请求完成,响应总长度: %d 字符", len(text)) - if source_dir: - _emit_llm_stream(source_dir, "end", label="分析完成") - return text - except Exception as e: - log.error("LLM 请求失败: %s", e) - if source_dir: - _emit_llm_stream(source_dir, "error", error=str(e)) - raise + return _stream_llm_response(llm, messages, source_dir, reasoning_effort, log_label="LLM") def extract_document(file_path: Path) -> dict[str, Any]: @@ -299,38 +321,7 @@ def _llm_query_multimodal( len(final_blocks), ) - try: - if source_dir: - _emit_llm_stream(source_dir, "start", label="正在分析文件...") - - parts = [] - for resp in llm.stream_chat( - messages, - temperature=0.1, - extra_body={"reasoning_effort": reasoning_effort}, - ): - delta = resp.delta - if delta: - parts.append(delta) - if source_dir: - _emit_llm_stream(source_dir, "chunk", text=delta) - - thinking = getattr(resp, "additional_kwargs", {}) or {} - thinking_delta = thinking.get("thinking_delta", "") - if thinking_delta and source_dir: - _emit_llm_stream(source_dir, "reasoning", text=thinking_delta) - - text = "".join(parts) - log.info("LLM 多模态请求完成,响应总长度: %d 字符", len(text)) - log.info("LLM 多模态响应: %s", text) - if source_dir: - _emit_llm_stream(source_dir, "end", label="分析完成") - return text - except Exception as e: - log.error("LLM 多模态请求失败: %s", e) - if source_dir: - _emit_llm_stream(source_dir, "error", error=str(e)) - raise + return _stream_llm_response(llm, messages, source_dir, reasoning_effort, log_label="LLM多模态") def load_cache(source_dir: Path) -> dict[str, Any]: @@ -437,6 +428,48 @@ def build_extraction_user_message( return "\n".join(parts) +def _extract_info( + source_dir: Path | None, + system_prompt: str, + info_type: str, +) -> dict[str, Any]: + """通用的信息提取函数:加载缓存、构建消息、调用 LLM 并解析 JSON。 + + extract_travel_info 和 extract_normal_info 的公共实现。 + + Args: + source_dir: 源文件目录(必填,包含 .invoice_cache 子目录)。 + system_prompt: 系统提示词。 + info_type: 信息类型标签("差旅" 或 "普通发票"),用于日志。 + + Returns: + LLM 提取的结构化信息字典。 + """ + if not source_dir: + log.warning("未提供 source_dir,无法加载缓存数据") + return {} + + cache_map = load_cache(source_dir) + match_result = load_match_result(source_dir) + user_message = build_extraction_user_message(cache_map, match_result) + + log.info("开始构建%s信息提取请求,缓存条目: %d, 匹配结果: %d", info_type, len(cache_map), len(match_result)) + + try: + response = llm_query_text( + system_prompt=system_prompt, + text=user_message, + reasoning_effort="low", + source_dir=source_dir, + ) + result = parse_json_response(response) + log.info("LLM %s信息提取成功", info_type) + return result + except Exception as e: + log.error("LLM %s信息提取失败: %s", info_type, e) + raise + + def extract_travel_info( source_dir: Path | None = None, ) -> dict[str, Any]: @@ -450,35 +483,7 @@ def extract_travel_info( Returns: 包含出差事由、地点、交通工具、时间、住宿信息等字段的字典。 """ - if not source_dir: - log.warning("未提供 source_dir,无法加载缓存数据") - return {} - - system_prompt = build_travel_info_system_prompt() - cache_map = load_cache(source_dir) - match_result = load_match_result(source_dir) - user_message = build_extraction_user_message(cache_map, match_result) - - log.info(f"user_message: {user_message}") - - try: - response = llm_query_text( - system_prompt=system_prompt, - text=user_message, - reasoning_effort="low", - source_dir=source_dir, - ) - result = parse_json_response(response) - log.info("LLM 差旅信息提取成功") - return result - except Exception as e: - log.error("LLM 差旅信息提取失败: %s", e) - raise - - -# ------------------------------------------------------------------ -# 普通发票信息提取 -# ------------------------------------------------------------------ + return _extract_info(source_dir, build_travel_info_system_prompt(), "差旅") def extract_normal_info( @@ -494,30 +499,7 @@ def extract_normal_info( Returns: 包含报销说明、发票总数、总金额、支付方式、附件清单等字段的字典。 """ - if not source_dir: - log.warning("未提供 source_dir,无法加载缓存数据") - return {} - - system_prompt = build_normal_info_system_prompt() - cache_map = load_cache(source_dir) - match_result = load_match_result(source_dir) - user_message = build_extraction_user_message(cache_map, match_result) - - log.info(f"user_message: {user_message}") - - try: - response = llm_query_text( - system_prompt=system_prompt, - text=user_message, - reasoning_effort="low", - source_dir=source_dir, - ) - result = parse_json_response(response) - log.info("LLM 普通发票信息提取成功") - return result - except Exception as e: - log.error("LLM 普通发票信息提取失败: %s", e) - raise + return _extract_info(source_dir, build_normal_info_system_prompt(), "普通发票") # ------------------------------------------------------------------ diff --git a/src/doc/matcher.py b/src/doc/matcher.py index 084ed0d..9f7eed2 100644 --- a/src/doc/matcher.py +++ b/src/doc/matcher.py @@ -227,25 +227,38 @@ def _match_one_to_one( assigned: set[int], result: dict[int, list[int]], ) -> None: - """一对一匹配:发票数等于刷卡数,按金额从大到小依次配对""" + """一对一匹配:发票数等于刷卡数,对每张刷卡记录寻找金额最接近的未分配发票""" for card_idx, card in enumerate(cards): - if card_idx >= len(invoices): - break - inv = invoices[card_idx] - diff = abs(inv["_amount"] - card["_amount"]) - card_tol = _relative_tolerance(card["_amount"], tolerance) - if diff <= card_tol: - assigned.add(card_idx) - result[card_idx] = [card_idx] + if card_idx in result: + continue + card_amount = card["_amount"] + card_tol = _relative_tolerance(card_amount, tolerance) + + # 在未分配的发票中找金额最接近的 + best_idx = -1 + best_diff = float("inf") + for idx, inv in enumerate(invoices): + if idx in assigned: + continue + diff = abs(inv["_amount"] - card_amount) + if diff < best_diff: + best_diff = diff + best_idx = idx + + if best_idx >= 0 and best_diff <= card_tol: + inv = invoices[best_idx] + assigned.add(best_idx) + result[card_idx] = [best_idx] log.info( f"[一对一] {inv.get('invoice_number', 'unknown')} ¥{inv['_amount']:.2f} " - f"↔ {card.get('_source_file', 'unknown')} ¥{card['_amount']:.2f}" + f"↔ {card.get('_source_file', 'unknown')} ¥{card_amount:.2f}" ) - else: + elif best_idx >= 0: + inv = invoices[best_idx] log.warning( f"[一对一] 金额偏差超出容差: " f"{inv.get('invoice_number', 'unknown')} ¥{inv['_amount']:.2f} " - f"vs ¥{card['_amount']:.2f} (差 ¥{diff:.2f}, 容差 ¥{card_tol:.2f})" + f"vs ¥{card_amount:.2f} (差 ¥{best_diff:.2f}, 容差 ¥{card_tol:.2f})" ) diff --git a/src/doc/pdf.py b/src/doc/pdf.py index 49519d5..4928fc7 100644 --- a/src/doc/pdf.py +++ b/src/doc/pdf.py @@ -21,7 +21,7 @@ def render_pdf_to_images(filepath: Path, dpi: int = 300) -> list[str]: Args: filepath: PDF 文件路径。 - dpi: 渲染分辨率(默认 150,平衡质量与速度)。 + dpi: 渲染分辨率(默认 300,平衡质量与速度)。 Returns: base64 编码的 JPEG 图片字符串列表(每页一个)。 diff --git a/src/doc/validator.py b/src/doc/validator.py index 0643a6d..c37245a 100644 --- a/src/doc/validator.py +++ b/src/doc/validator.py @@ -54,12 +54,13 @@ class ValidationReport: def _check_field( data: dict[str, Any], path: list[str], - required: bool = True, check_empty: bool = True, custom_check: Any = None, ) -> tuple[bool, str]: """沿路径检查字段是否存在且有效。 + 检查字段是否存在于嵌套字典中,可选的检查是否为空字符串或自定义校验。 + Returns: (通过, 字段路径字符串) """ diff --git a/src/pipeline.py b/src/pipeline.py index 5c33294..b612b9d 100644 --- a/src/pipeline.py +++ b/src/pipeline.py @@ -14,9 +14,8 @@ Bot 仅负责接收信息并填报,不再承担信息提取职责。 """ -import json from pathlib import Path -from typing import Any, cast +from typing import Any from . import get_logger from .config import load_config @@ -29,83 +28,35 @@ from .doc.invoice import ( from .doc.invoice import ( save_csv as save_payment_csv, ) -from .doc.llm_extractor import ( - CACHE_DIR_NAME, - extract_normal_info, - extract_travel_info, - load_cache, +from .pipeline_core import ( + extract_and_cache_normal_info, + extract_and_cache_travel_info, + is_travel_invoice, ) 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) - invoices = [data for data in cache_map.values() if data.get("invoice_type") not in ("application", "payment")] + # 过滤掉非发票的缓存条目: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 _extract_travel_info_if_needed(groups: dict[str, list[dict[str, Any]]], cache_path: Path) -> dict[str, Any] | None: - """当存在差旅发票时,调用 LLM 提取差旅信息并缓存。 - - Returns: - 差旅信息字典,非差旅时返回 None。 - """ - if not groups.get("travel"): - return None - - from .doc.llm_extractor import CACHE_DIR_NAME, load_cache - - # 检查缓存是否已有 - cache_map = load_cache(cache_path) - if cache_map.get("travel_info"): - log.info("使用已有差旅信息缓存") - return cast(dict[str, Any] | None, cache_map["travel_info"]) - - log.info("开始提取差旅信息...") - travel_info = extract_travel_info(source_dir=cache_path) - - # 保存到缓存 - cache_dir = cache_path / CACHE_DIR_NAME - cache_dir.mkdir(parents=True, exist_ok=True) - - with open(cache_dir / "travel_info.json", "w", encoding="utf-8") as f: - json.dump(travel_info, f, ensure_ascii=False, indent=2) - log.info("差旅信息已保存到缓存") - return travel_info - - -def _extract_normal_info_if_needed(groups: dict[str, list[dict[str, Any]]], cache_path: Path) -> dict[str, Any] | None: - """当存在普通发票时,调用 LLM 提取普通报销信息并缓存。 - - Returns: - 普通报销信息字典,非普通时返回 None。 - """ - if not groups.get("general"): - return None - - # 检查缓存是否已有 - cache_map = load_cache(cache_path) - if cache_map.get("normal_info"): - log.info("使用已有普通发票信息缓存") - return cast(dict[str, Any] | None, cache_map["normal_info"]) - - log.info("开始提取普通发票信息...") - normal_info = extract_normal_info(source_dir=cache_path) - - # 保存到缓存 - cache_dir = cache_path / CACHE_DIR_NAME - cache_dir.mkdir(parents=True, exist_ok=True) - - with open(cache_dir / "normal_info.json", "w", encoding="utf-8") as f: - json.dump(normal_info, f, ensure_ascii=False, indent=2) - log.info("普通发票信息已保存到缓存") - return normal_info - - def run_pipeline( step: str = "all", username: str | None = None, @@ -156,11 +107,10 @@ def run_pipeline( log.info(f"发票分类: 差旅 {len(groups['travel'])} 张, 普通 {len(groups['general'])} 张") # 发票提取完成后立即判断类型 - is_travel = bool(groups["travel"]) and not bool(groups["general"]) - if is_travel: - travel_info = _extract_travel_info_if_needed(groups, cache_path) + if is_travel_invoice(groups): + travel_info = extract_and_cache_travel_info(groups, cache_path) else: - normal_info = _extract_normal_info_if_needed(groups, cache_path) + normal_info = extract_and_cache_normal_info(groups, cache_path) if step == "invoice": log.info("[1/2] 发票提取 完成") @@ -179,15 +129,14 @@ def run_pipeline( if groups is None: groups = _classify_from_cache(cache_path) - is_travel = bool(groups["travel"]) and not bool(groups["general"]) - if is_travel: + if is_travel_invoice(groups): if travel_info is None: - travel_info = _extract_travel_info_if_needed(groups, cache_path) + travel_info = extract_and_cache_travel_info(groups, cache_path) log.info("检测到纯差旅发票,使用差旅报销模式") run_bot(config, work_dir=cache_path, travel_info=travel_info) else: if normal_info is None: - normal_info = _extract_normal_info_if_needed(groups, cache_path) + normal_info = extract_and_cache_normal_info(groups, cache_path) log.info("检测到普通发票,使用普通报销模式") run_bot(config, work_dir=cache_path, normal_info=normal_info) diff --git a/src/pipeline_core.py b/src/pipeline_core.py new file mode 100644 index 0000000..8f6c98a --- /dev/null +++ b/src/pipeline_core.py @@ -0,0 +1,100 @@ +""" +管道核心逻辑 + +抽取 pipeline.py(CLI 管道)和 pipeline_web.py(Web 管道)的公共数据流: + 发票分类判断 -> 差旅/普通信息提取 -> 缓存读写 + +两个入口分别传入不同的目录参数,复用此模块。 +""" + +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any + +from . import get_logger +from .doc.llm_extractor import ( + CACHE_DIR_NAME, + extract_normal_info, + extract_travel_info, + load_cache, +) + +log = get_logger("pipeline_core") + + +def is_travel_invoice(groups: dict[str, list[dict[str, Any]]]) -> bool: + """判断是否为纯差旅发票(有差旅发票且无普通发票)。 + + 注意:系统仅支持「纯差旅」和「普通报销」两种模式。 + 若同时存在差旅发票和普通发票(混合),则视为普通报销模式处理——差旅发票 + 对应的费用仍会在普通报销中按项目填报。若需要严格区分,上游应在发票分类 + 后报错提示用户分开提交。 + """ + return bool(groups.get("travel")) and not bool(groups.get("general")) + + +def save_cache_info(cache_path: Path, info_key: str, info: dict[str, Any]) -> None: + """将提取结果保存到缓存目录 + + Args: + cache_path: 会话目录路径。 + info_key: 缓存键名("travel_info" 或 "normal_info")。 + info: 提取结果字典。 + """ + cache_dir = cache_path / CACHE_DIR_NAME + cache_dir.mkdir(parents=True, exist_ok=True) + with open(cache_dir / f"{info_key}.json", "w", encoding="utf-8") as f: + json.dump(info, f, ensure_ascii=False, indent=2) + log.info("%s 已保存到缓存", info_key) + + +def extract_and_cache_travel_info( + groups: dict[str, list[dict[str, Any]]], + cache_path: Path, +) -> dict[str, Any] | None: + """当存在差旅发票时,调用 LLM 提取差旅信息并缓存。 + + Returns: + 差旅信息字典,非差旅时返回 None。 + """ + if not groups.get("travel"): + return None + + # 检查缓存是否已有 + cache_map = load_cache(cache_path) + travel_info = cache_map.get("travel_info") + if travel_info: + log.info("使用已有差旅信息缓存") + return travel_info # type: ignore[no-any-return] + + log.info("开始提取差旅信息...") + travel_info = extract_travel_info(source_dir=cache_path) + save_cache_info(cache_path, "travel_info", travel_info) + return travel_info + + +def extract_and_cache_normal_info( + groups: dict[str, list[dict[str, Any]]], + cache_path: Path, +) -> dict[str, Any] | None: + """当存在普通发票时,调用 LLM 提取普通报销信息并缓存。 + + Returns: + 普通报销信息字典,非普通时返回 None。 + """ + if not groups.get("general"): + return None + + # 检查缓存是否已有 + cache_map = load_cache(cache_path) + normal_info = cache_map.get("normal_info") + if normal_info: + log.info("使用已有普通发票信息缓存") + return normal_info # type: ignore[no-any-return] + + log.info("开始提取普通发票信息...") + normal_info = extract_normal_info(source_dir=cache_path) + save_cache_info(cache_path, "normal_info", normal_info) + return normal_info diff --git a/src/web/pipeline_web.py b/src/web/pipeline_web.py index 9bff3b3..df6084c 100644 --- a/src/web/pipeline_web.py +++ b/src/web/pipeline_web.py @@ -16,7 +16,6 @@ from urllib.parse import quote # 延迟导入,避免循环引用 from src import get_logger # noqa: F401 -from src.config import load_config as load_project_config from src.doc.fill_consumable_doc import ( CONSUMABLE_DOC_FILENAME, fill_consumable_from_template, @@ -28,6 +27,11 @@ from src.doc.invoice import ( from src.doc.invoice import ( save_csv as save_payment_csv, ) +from src.pipeline_core import ( + extract_and_cache_normal_info, + extract_and_cache_travel_info, + is_travel_invoice, +) fill_log = get_logger("fill_consumable_doc") @@ -37,9 +41,16 @@ INVOICE_GROUPS_FILE = "invoice_groups.json" def save_invoice_groups(session_dir: Path, groups: dict[str, list[dict[str, str]]]) -> None: - """保存发票分类结果到 session 目录的 JSON 文件""" + """保存发票分类结果到 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", [])), @@ -48,28 +59,21 @@ def save_invoice_groups(session_dir: Path, groups: dict[str, list[dict[str, str] json.dump(data, f, ensure_ascii=False, indent=2) -def load_invoice_groups(session_dir: Path) -> dict[str, int] | None: - """从 session 目录加载发票分类统计""" +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, int] | None, json.load(f)) + return cast(dict[str, Any] | None, json.load(f)) except Exception: return None -def load_session_config(session_dir: Path) -> dict[str, Any]: - """加载会话配置,合并项目全局配置与会话级配置""" - config = load_project_config() - cfg_path = session_dir / "config.json" - if cfg_path.exists(): - with open(cfg_path, encoding="utf-8") as f: - config.update(json.load(f)) - return config - - def resolve_payment_csv(session_dir: Path) -> Path | None: """查找支付记录 CSV(payment_records.csv)""" csv_path = session_dir / "payment_records.csv" @@ -162,7 +166,7 @@ def append_doc_download(result: dict[str, Any], session_id: str, doc_fill: dict[ # ================================================================ -def run_pipeline_web(session_dir: Path, config: dict[str, Any]) -> dict[str, Any]: +def run_pipeline_web(session_dir: Path, config: dict[str, Any] | None) -> dict[str, Any]: """在 Web 会话目录中执行发票提取,结果写入 session 目录下的文件 注意:不再自动提交财务系统。提交通由 /api/submit-financial/ 触发。 @@ -192,31 +196,10 @@ def run_pipeline_web(session_dir: Path, config: dict[str, Any]) -> dict[str, Any save_invoice_groups(session_dir, groups) # ---- Step 2: 差旅/普通信息提取 ---- - is_travel = bool(groups.get("travel")) and not bool(groups.get("general")) - if is_travel: - from src.doc.llm_extractor import CACHE_DIR_NAME, extract_travel_info, load_cache - - cache_map = load_cache(session_dir) - if not cache_map.get("travel_info"): - fill_log.info("开始提取差旅信息...") - travel_info = extract_travel_info(source_dir=session_dir) - cache_dir = session_dir / CACHE_DIR_NAME - cache_dir.mkdir(parents=True, exist_ok=True) - with open(cache_dir / "travel_info.json", "w", encoding="utf-8") as f: - json.dump(travel_info, f, ensure_ascii=False, indent=2) - fill_log.info("差旅信息已保存到缓存") + if is_travel_invoice(groups): + extract_and_cache_travel_info(groups, session_dir) else: - from src.doc.llm_extractor import CACHE_DIR_NAME, extract_normal_info, load_cache - - cache_map = load_cache(session_dir) - if not cache_map.get("normal_info"): - fill_log.info("开始提取普通发票信息...") - normal_info = extract_normal_info(source_dir=session_dir) - cache_dir = session_dir / CACHE_DIR_NAME - cache_dir.mkdir(parents=True, exist_ok=True) - with open(cache_dir / "normal_info.json", "w", encoding="utf-8") as f: - json.dump(normal_info, f, ensure_ascii=False, indent=2) - fill_log.info("普通发票信息已保存到缓存") + extract_and_cache_normal_info(groups, session_dir) # 统计发票总数 invoice_count = sum(len(inv.get("_matched_invoices", [])) for inv in invoices) @@ -230,16 +213,16 @@ def run_pipeline_web(session_dir: Path, config: dict[str, Any]) -> dict[str, Any "travel_count": len(groups["travel"]), "general_count": len(groups["general"]), } - doc_fill = _try_fill_consumable_doc(session_dir, config) + doc_fill = _try_fill_consumable_doc(session_dir, config or {}) append_doc_download(result, session_dir.name, doc_fill) return result -def run_financial_submit(session_dir: Path, config: dict[str, Any]) -> dict[str, Any]: +def run_financial_submit(session_dir: Path, config: dict[str, Any] | None) -> dict[str, Any]: """执行财务系统填报(从前端确认后调用) 从 invoice_groups.json 读取分类结果,根据发票类型选择填报模式: - - 纯差旅发票:差旅报销模式(TODO) + - 纯差旅发票:差旅报销模式 - 含普通发票:普通报销模式 """ csv_path = session_dir / "payment_records.csv" @@ -255,5 +238,9 @@ def run_financial_submit(session_dir: Path, config: dict[str, Any]) -> dict[str, else: fill_log.info("检测到普通发票,使用普通报销模式") - run_bot_web(config, session_dir) - return {"ok": True} + try: + run_bot_web(config or {}, session_dir) + return {"ok": True} + except Exception as e: + fill_log.error("财务填报失败: %s", e) + return {"ok": False, "error": str(e)} diff --git a/src/web/routes.py b/src/web/routes.py index b147175..53aa2b8 100644 --- a/src/web/routes.py +++ b/src/web/routes.py @@ -16,7 +16,14 @@ from urllib.parse import quote from flask import Blueprint, Response, jsonify, render_template, request, stream_with_context -from src.config import load_config as load_project_config +from src.config import ( + SAFE_CONFIG_KEYS, + SESSION_CONFIG_KEYS, + load_session_config, +) +from src.config import ( + load_config as load_project_config, +) from src.doc.invoice import load_csv, load_invoice_csv from . import pipeline_web, sse_handler @@ -51,15 +58,11 @@ def _validate_session(session_id: str) -> Path | tuple[Response, int]: def _build_web_config(body: dict[str, Any]) -> dict[str, Any]: """从请求体构建配置""" config = load_project_config() - for key in ( - "username", - "password", - "default_name", - "default_card_no", - "default_person_id", - "consumable_storage", - ): + for key in SESSION_CONFIG_KEYS: if body.get(key): + # 基本类型校验:防止非字符串值写入配置 + if not isinstance(body[key], (str, int, float, bool)): + continue config[key] = body[key] return config @@ -68,40 +71,32 @@ def _emit_ready_and_submit( session_dir: Path, agent_session: Any, config: dict[str, Any], -) -> None: - """Agent 校验通过后,直接触发财务提交(不依赖前端)""" - from src.agent import _emit_agent_event +) -> dict[str, Any]: + """Agent 校验通过后,直接触发财务提交(不依赖前端)。 - # 发射 agent_ready 事件,前端 SSE 会收到 - _emit_agent_event( - session_dir, - "agent_ready", - round=agent_session.rounds, - message="信息完整,可以提交", - ) - - # 直接执行财务提交 + 返回 result 字典,由调用方 _run_agent_task 的 finally 块统一写入 result.json。 + """ + # agent_ready 事件已由 run_agent_round 发射,此处仅执行财务提交,避免前端收到重复消息 try: submit_result = pipeline_web.run_financial_submit(session_dir, config) if submit_result.get("ok"): - result = { + return { "ok": True, "agent_ready": True, "submit_ok": True, "round": agent_session.rounds, "message": "信息完整,已自动提交到财务系统", } - else: - result = { - "ok": True, - "agent_ready": True, - "submit_ok": False, - "submit_error": submit_result.get("error"), - "round": agent_session.rounds, - "message": "校验通过但提交失败", - } + return { + "ok": True, + "agent_ready": True, + "submit_ok": False, + "submit_error": submit_result.get("error"), + "round": agent_session.rounds, + "message": "校验通过但提交失败", + } except Exception as e: - result = { + return { "ok": True, "agent_ready": True, "submit_ok": False, @@ -110,14 +105,69 @@ def _emit_ready_and_submit( "message": f"校验通过但提交异常: {e}", } - # 写入 result 文件,SSE done 事件会读取 + +def _run_agent_task( + session_dir: Path, + handler: Any, + task_fn: Any, + config: dict[str, Any] | None = None, +) -> None: + """Agent 任务的通用包装器 + + 封装重复的闭包结构:清除流日志 -> 执行任务 -> 写入 result.json -> 卸载收集器。 + + Args: + session_dir: 会话目录 + handler: SSE 日志收集器句柄 + task_fn: 业务逻辑回调,接收 (session_dir, config) 返回 (agent_session, result_dict) 或仅 result_dict + config: 配置字典(可选) + """ + result = {"ok": False, "error": "未知错误"} try: - tmp_path = session_dir / (pipeline_web.SESSION_RESULT_FILE + ".tmp") - with open(tmp_path, "w", encoding="utf-8") as f: - json.dump(result, f, ensure_ascii=False) - tmp_path.replace(session_dir / pipeline_web.SESSION_RESULT_FILE) - except Exception: - pass + # 清除上一轮残留文件,避免 SSE 连接立即读到旧数据 + for fname in ("llm_stream.log", "agent_events.log", pipeline_web.SESSION_RESULT_FILE): + try: + (session_dir / fname).unlink(missing_ok=True) + except Exception: + pass + + ret = task_fn(session_dir, config) + + # task_fn 返回 (agent_session, result) 或仅 result + if isinstance(ret, tuple) and len(ret) == 2: + from src.agent import AgentState + + agent_session, result = ret + + if agent_session.state == AgentState.READY: + if config is not None: + result = _emit_ready_and_submit(session_dir, agent_session, config) + elif agent_session.state == AgentState.AWAITING_SUPPLEMENT: + result = { + "ok": True, + "agent_ready": False, + "agent_state": agent_session.state.value, + "round": agent_session.rounds, + "waiting_for_supplement": True, + } + elif agent_session.state == AgentState.ERROR: + result = { + "ok": False, + "error": agent_session.error_message, + } + + except Exception as e: + result = {"ok": False, "error": str(e)} + finally: + # 统一写入 result.json,SSE 端点检测到后发射 done 事件 + try: + tmp_path = session_dir / (pipeline_web.SESSION_RESULT_FILE + ".tmp") + with open(tmp_path, "w", encoding="utf-8") as f: + json.dump(result, f, ensure_ascii=False) + tmp_path.replace(session_dir / pipeline_web.SESSION_RESULT_FILE) + except Exception: + pass + sse_handler.remove_log_collector(handler) # ================================================================ @@ -174,6 +224,9 @@ def upload_file(session_id: str) -> Any: return jsonify({"error": "未选择文件"}), 400 safe_name = Path(f.filename).name + # 安全检查:拒绝包含路径穿越字符的文件名 + if ".." in safe_name or "/" in safe_name or "\\" in safe_name: + return jsonify({"error": "文件名包含非法字符"}), 400 f.save(str(session_dir / safe_name)) return jsonify({"ok": True, "filename": safe_name}) @@ -247,7 +300,13 @@ def mobile_upload_file(session_id: str) -> Any: @web_bp.route("/api/config/", methods=["GET"]) def get_session_config(session_id: str) -> Any: - """获取当前会话的配置(供前端回填表单)""" + """获取当前会话的配置(供前端回填表单) + + 使用白名单过滤会话配置,防止密码等敏感字段被加载到内存。 + password 字段始终返回空字符串——密码由前端用户输入,不持久化。 + 注意:当前为 localhost 服务,密码通过 HTTP 明文传输。如需通过代理暴露服务, + 请启用 HTTPS 或使用反向代理加密。 + """ session_dir = _validate_session(session_id) if isinstance(session_dir, tuple): return session_dir @@ -255,8 +314,12 @@ def get_session_config(session_id: str) -> Any: config = load_project_config() cfg_path = session_dir / "config.json" if cfg_path.exists(): + # 仅允许覆盖前端表单字段(不含密码) with open(cfg_path, encoding="utf-8") as f: - config.update(json.load(f)) + session_cfg = json.load(f) + for key in SAFE_CONFIG_KEYS: + if key in session_cfg: + config[key] = session_cfg[key] return jsonify( { "username": config.get("username", ""), @@ -348,6 +411,12 @@ def save_invoice_data(session_id: str) -> Any: fieldnames = list(original_rows[0].keys()) + # 校验前端传入的 data 字段是否与原始 CSV 列匹配 + if data: + unknown_keys = set(data[0].keys()) - (set(fieldnames) | {"__row"}) + if unknown_keys: + return jsonify({"error": f"数据包含未知字段: {unknown_keys}"}), 400 + with open(csv_path, "w", newline="", encoding="utf-8-sig") as f: writer = csv_module.DictWriter(f, fieldnames=fieldnames) writer.writeheader() @@ -356,7 +425,7 @@ def save_invoice_data(session_id: str) -> Any: writer.writerow(row) resp: dict[str, str | bool | None] = {"ok": True} - config = pipeline_web.load_session_config(session_dir) + config = load_session_config(session_dir) doc_fill = pipeline_web._try_fill_consumable_doc(session_dir, config) if doc_fill.get("ok"): fn = doc_fill["doc_filename"] @@ -389,31 +458,15 @@ def start_process(session_id: str) -> Any: with open(session_dir / "config.json", "w", encoding="utf-8") as f: json.dump(config, f, ensure_ascii=False, indent=2, default=str) + def _task(sd: Path, cfg: dict[str, Any] | None) -> dict[str, Any]: + return pipeline_web.run_pipeline_web(sd, cfg) + handler = sse_handler.install_log_collector(session_dir) - - def _run() -> None: - result = {"ok": False, "error": "未知错误"} - try: - try: - (session_dir / "llm_stream.log").unlink(missing_ok=True) - except Exception: - pass - result = pipeline_web.run_pipeline_web(session_dir, config) - except BaseException as e: - result = {"ok": False, "error": str(e)} - if isinstance(e, KeyboardInterrupt | SystemExit): - raise - finally: - try: - tmp_path = session_dir / (pipeline_web.SESSION_RESULT_FILE + ".tmp") - with open(tmp_path, "w", encoding="utf-8") as f: - json.dump(result, f, ensure_ascii=False) - tmp_path.replace(session_dir / pipeline_web.SESSION_RESULT_FILE) - except Exception: - pass - sse_handler.remove_log_collector(handler) - - threading.Thread(target=_run, daemon=True).start() + threading.Thread( + target=_run_agent_task, + kwargs={"session_dir": session_dir, "handler": handler, "task_fn": _task, "config": config}, + daemon=True, + ).start() return jsonify({"status": "started"}) @@ -433,7 +486,7 @@ def stream_logs(session_id: str) -> Any: state = {"agent_size": 0} start_time = time.time() - timeout = 600 + timeout = 900 # 略大于 LLM 请求超时 (600s),防止 SSE 先断开 while time.time() - start_time < timeout: # 轮询普通日志 @@ -523,43 +576,18 @@ def submit_financial(session_id: str) -> Any: with open(config_path, encoding="utf-8") as f: config = json.load(f) - result_file = session_dir / pipeline_web.SESSION_RESULT_FILE - if result_file.exists(): - result_file.unlink() + def _task(sd: Path, cfg: dict[str, Any] | None) -> dict[str, Any]: + submit_result = pipeline_web.run_financial_submit(sd, cfg) + if submit_result.get("ok"): + return {"ok": True} + return submit_result handler = sse_handler.install_log_collector(session_dir) - - def _run() -> None: - result = {"ok": False, "error": "未知错误"} - try: - try: - (session_dir / "llm_stream.log").unlink(missing_ok=True) - except Exception: - pass - submit_result = pipeline_web.run_financial_submit(session_dir, config) - if submit_result.get("ok"): - result = {"ok": True} - else: - result = submit_result - except BaseException as e: - result = {"ok": False, "error": str(e)} - if isinstance(e, KeyboardInterrupt | SystemExit): - raise - finally: - try: - tmp_path = session_dir / (pipeline_web.SESSION_RESULT_FILE + ".tmp") - with open(tmp_path, "w", encoding="utf-8") as f: - json.dump( - {"ok": True, "submit_ok": result.get("ok"), "submit_error": result.get("error")}, - f, - ensure_ascii=False, - ) - tmp_path.replace(session_dir / pipeline_web.SESSION_RESULT_FILE) - except Exception: - pass - sse_handler.remove_log_collector(handler) - - threading.Thread(target=_run, daemon=True).start() + threading.Thread( + target=_run_agent_task, + kwargs={"session_dir": session_dir, "handler": handler, "task_fn": _task, "config": config}, + daemon=True, + ).start() return jsonify({"status": "started"}) @@ -599,88 +627,52 @@ def agent_process(session_id: str) -> Any: handler = sse_handler.install_log_collector(session_dir) - def _run() -> None: - result = {"ok": False, "error": "未知错误"} - try: - try: - (session_dir / "llm_stream.log").unlink(missing_ok=True) - except Exception: - pass + def _task(sd: Path, cfg: dict[str, Any] | None) -> tuple[Any, dict[str, Any]]: + from src.agent import ( + AgentSession, + load_agent_state, + run_agent_round, + save_agent_state, + ) + from src.doc.extractor import extract_invoices + from src.doc.invoice import ( + save_application_json, + save_invoice_csv, + ) + from src.doc.invoice import ( + save_csv as save_payment_csv, + ) - from src.agent import ( - AgentSession, - AgentState, - load_agent_state, - run_agent_round, - save_agent_state, - ) - from src.doc.extractor import extract_invoices - from src.doc.invoice import ( - save_application_json, - save_invoice_csv, - ) - from src.doc.invoice import ( - save_csv as save_payment_csv, + agent_session = load_agent_state(sd) + if agent_session is None: + invoices, applications, groups = extract_invoices(str(sd)) + if not invoices: + return (None, {"ok": False, "error": "未提取到任何发票数据"}) + + save_payment_csv(invoices, sd / "payment_records.csv") + save_invoice_csv(invoices, sd / "invoice_summary.csv") + + if applications: + save_application_json(applications, sd / "travel_applications.json") + + pipeline_web.save_invoice_groups(sd, groups) + + from src.pipeline_core import is_travel_invoice + + agent_session = AgentSession( + session_id=session_id, + invoice_type="travel" if is_travel_invoice(groups) else "normal", ) + save_agent_state(sd, agent_session) - agent_session = load_agent_state(session_dir) - if agent_session is None: - invoices, applications, groups = extract_invoices(str(session_dir)) - if not invoices: - result = {"ok": False, "error": "未提取到任何发票数据"} - return + agent_session = run_agent_round(sd, agent_session) + return (agent_session, {}) - save_payment_csv(invoices, session_dir / "payment_records.csv") - save_invoice_csv(invoices, session_dir / "invoice_summary.csv") - - if applications: - save_application_json(applications, session_dir / "travel_applications.json") - - pipeline_web.save_invoice_groups(session_dir, groups) - - is_travel = bool(groups.get("travel")) and not bool(groups.get("general")) - agent_session = AgentSession( - session_id=session_id, - invoice_type="travel" if is_travel else "normal", - ) - save_agent_state(session_dir, agent_session) - - agent_session = run_agent_round(session_dir, agent_session) - - if agent_session.state == AgentState.READY: - # Agent 校验通过,直接触发财务提交 - _emit_ready_and_submit(session_dir, agent_session, config) - elif agent_session.state == AgentState.AWAITING_SUPPLEMENT: - result = { - "ok": True, - "agent_ready": False, - "agent_state": agent_session.state.value, - "round": agent_session.rounds, - "waiting_for_supplement": True, - } - elif agent_session.state == AgentState.ERROR: - result = { - "ok": False, - "error": agent_session.error_message, - } - - except BaseException as e: - result = {"ok": False, "error": str(e)} - if isinstance(e, KeyboardInterrupt | SystemExit): - raise - finally: - try: - result_file = session_dir / pipeline_web.SESSION_RESULT_FILE - if not result_file.exists(): - tmp_path = session_dir / (pipeline_web.SESSION_RESULT_FILE + ".tmp") - with open(tmp_path, "w", encoding="utf-8") as f: - json.dump(result, f, ensure_ascii=False) - tmp_path.replace(session_dir / pipeline_web.SESSION_RESULT_FILE) - except Exception: - pass - sse_handler.remove_log_collector(handler) - - threading.Thread(target=_run, daemon=True).start() + threading.Thread( + target=_run_agent_task, + kwargs={"session_dir": session_dir, "handler": handler, "task_fn": _task, "config": config}, + daemon=True, + ).start() return jsonify({"status": "started"}) @@ -698,7 +690,6 @@ def agent_supplement(session_id: str) -> Any: return jsonify({"error": "未指定补充文件"}), 400 from src.agent import ( - AgentState, add_supplement, load_agent_state, run_agent_round, @@ -710,74 +701,42 @@ def agent_supplement(session_id: str) -> Any: agent_session = add_supplement(session_dir, agent_session, filenames) + # 重新提取发票并保存 + from src.doc.extractor import extract_invoices + from src.doc.invoice import ( + save_application_json, + save_invoice_csv, + ) + from src.doc.invoice import ( + save_csv as save_payment_csv, + ) + + invoices, applications, groups = extract_invoices(str(session_dir)) + save_payment_csv(invoices, session_dir / "payment_records.csv") + save_invoice_csv(invoices, session_dir / "invoice_summary.csv") + + if applications: + save_application_json(applications, session_dir / "travel_applications.json") + + pipeline_web.save_invoice_groups(session_dir, groups) + + config_path = session_dir / "config.json" + config = {} + if config_path.exists(): + with open(config_path, encoding="utf-8") as f: + config = json.load(f) + handler = sse_handler.install_log_collector(session_dir) - def _run() -> None: - result = {"ok": False, "error": "未知错误"} - try: - try: - (session_dir / "llm_stream.log").unlink(missing_ok=True) - except Exception: - pass + def _task(sd: Path, cfg: dict[str, Any] | None) -> tuple[Any, dict[str, Any]]: + new_session = run_agent_round(sd, agent_session, new_files=filenames) + return (new_session, {}) - from src.doc.extractor import extract_invoices - from src.doc.invoice import ( - save_application_json, - save_invoice_csv, - ) - from src.doc.invoice import ( - save_csv as save_payment_csv, - ) - - invoices, applications, groups = extract_invoices(str(session_dir)) - save_payment_csv(invoices, session_dir / "payment_records.csv") - save_invoice_csv(invoices, session_dir / "invoice_summary.csv") - - if applications: - save_application_json(applications, session_dir / "travel_applications.json") - - pipeline_web.save_invoice_groups(session_dir, groups) - - new_session = run_agent_round(session_dir, agent_session, new_files=filenames) - - if new_session.state == AgentState.READY: - # Agent 校验通过,直接触发财务提交 - config_path = session_dir / "config.json" - with open(config_path, encoding="utf-8") as f: - config = json.load(f) - - _emit_ready_and_submit(session_dir, new_session, config) - elif new_session.state == AgentState.AWAITING_SUPPLEMENT: - result = { - "ok": True, - "agent_ready": False, - "agent_state": new_session.state.value, - "round": new_session.rounds, - "waiting_for_supplement": True, - } - elif new_session.state == AgentState.ERROR: - result = { - "ok": False, - "error": new_session.error_message, - } - - except BaseException as e: - result = {"ok": False, "error": str(e)} - if isinstance(e, KeyboardInterrupt | SystemExit): - raise - finally: - try: - result_file = session_dir / pipeline_web.SESSION_RESULT_FILE - if result.get("error") and not result_file.exists(): - tmp_path = session_dir / (pipeline_web.SESSION_RESULT_FILE + ".tmp") - with open(tmp_path, "w", encoding="utf-8") as f: - json.dump(result, f, ensure_ascii=False) - tmp_path.replace(session_dir / pipeline_web.SESSION_RESULT_FILE) - except Exception: - pass - sse_handler.remove_log_collector(handler) - - threading.Thread(target=_run, daemon=True).start() + threading.Thread( + target=_run_agent_task, + kwargs={"session_dir": session_dir, "handler": handler, "task_fn": _task, "config": config}, + daemon=True, + ).start() return jsonify({"status": "started"}) @@ -795,7 +754,6 @@ def agent_user_supplement(session_id: str) -> Any: return jsonify({"error": "请输入补充信息"}), 400 from src.agent import ( - AgentState, load_agent_state, process_user_text_supplement, ) @@ -804,55 +762,23 @@ def agent_user_supplement(session_id: str) -> Any: if agent_session is None: return jsonify({"error": "未找到 Agent 状态"}), 404 + config_path = session_dir / "config.json" + config = {} + if config_path.exists(): + with open(config_path, encoding="utf-8") as f: + config = json.load(f) + handler = sse_handler.install_log_collector(session_dir) - def _run() -> None: - result = {"ok": False, "error": "未知错误"} - try: - try: - (session_dir / "llm_stream.log").unlink(missing_ok=True) - except Exception: - pass + def _task(sd: Path, cfg: dict[str, Any] | None) -> tuple[Any, dict[str, Any]]: + new_session = process_user_text_supplement(sd, agent_session, user_text) + return (new_session, {}) - new_session = process_user_text_supplement(session_dir, agent_session, user_text) - - if new_session.state == AgentState.READY: - # Agent 校验通过,直接触发财务提交 - config_path = session_dir / "config.json" - with open(config_path, encoding="utf-8") as f: - config = json.load(f) - _emit_ready_and_submit(session_dir, new_session, config) - elif new_session.state == AgentState.AWAITING_SUPPLEMENT: - result = { - "ok": True, - "agent_ready": False, - "agent_state": new_session.state.value, - "round": new_session.rounds, - "waiting_for_supplement": True, - } - elif new_session.state == AgentState.ERROR: - result = { - "ok": False, - "error": new_session.error_message, - } - - except BaseException as e: - result = {"ok": False, "error": str(e)} - if isinstance(e, KeyboardInterrupt | SystemExit): - raise - finally: - try: - result_file = session_dir / pipeline_web.SESSION_RESULT_FILE - if not result_file.exists(): - tmp_path = session_dir / (pipeline_web.SESSION_RESULT_FILE + ".tmp") - with open(tmp_path, "w", encoding="utf-8") as f: - json.dump(result, f, ensure_ascii=False) - tmp_path.replace(session_dir / pipeline_web.SESSION_RESULT_FILE) - except Exception: - pass - sse_handler.remove_log_collector(handler) - - threading.Thread(target=_run, daemon=True).start() + threading.Thread( + target=_run_agent_task, + kwargs={"session_dir": session_dir, "handler": handler, "task_fn": _task, "config": config}, + daemon=True, + ).start() return jsonify({"status": "started"}) @@ -881,37 +807,16 @@ def agent_force_submit(session_id: str) -> Any: with open(config_path, encoding="utf-8") as f: config = json.load(f) - result_file = session_dir / pipeline_web.SESSION_RESULT_FILE - if result_file.exists(): - result_file.unlink() + def _task(sd: Path, cfg: dict[str, Any] | None) -> dict[str, Any]: + submit_result = pipeline_web.run_financial_submit(sd, cfg) + if submit_result.get("ok"): + return {"ok": True} + return submit_result handler = sse_handler.install_log_collector(session_dir) - - def _run() -> None: - result = {"ok": False, "error": "未知错误"} - try: - submit_result = pipeline_web.run_financial_submit(session_dir, config) - if submit_result.get("ok"): - result = {"ok": True} - else: - result = submit_result - except BaseException as e: - result = {"ok": False, "error": str(e)} - if isinstance(e, KeyboardInterrupt | SystemExit): - raise - finally: - try: - tmp_path = session_dir / (pipeline_web.SESSION_RESULT_FILE + ".tmp") - with open(tmp_path, "w", encoding="utf-8") as f: - json.dump( - {"ok": True, "submit_ok": result.get("ok"), "submit_error": result.get("error")}, - f, - ensure_ascii=False, - ) - tmp_path.replace(session_dir / pipeline_web.SESSION_RESULT_FILE) - except Exception: - pass - sse_handler.remove_log_collector(handler) - - threading.Thread(target=_run, daemon=True).start() + threading.Thread( + target=_run_agent_task, + kwargs={"session_dir": session_dir, "handler": handler, "task_fn": _task, "config": config}, + daemon=True, + ).start() return jsonify({"status": "started"}) diff --git a/src/web/sse_handler.py b/src/web/sse_handler.py index 8dc257b..acf85c8 100644 --- a/src/web/sse_handler.py +++ b/src/web/sse_handler.py @@ -47,10 +47,11 @@ class SSELogHandler(logging.Handler): pass def close_file(self) -> None: - try: - self._file.close() - except Exception: - pass + with self._lock: + try: + self._file.close() + except Exception: + pass def install_log_collector(session_dir: Path) -> SSELogHandler: diff --git a/src/web/static/css/index.css b/src/web/static/css/index.css index de150b9..9d4aa83 100644 --- a/src/web/static/css/index.css +++ b/src/web/static/css/index.css @@ -502,4 +502,91 @@ body { .agent-request-actions .btn { font-size: 13px; padding: 6px 16px; +} + +/* ================================================================ */ +/* 提取结果摘要卡片样式 */ +/* ================================================================ */ + +.extraction-summary { + margin-top: 8px; + padding-top: 8px; + border-top: 1px dashed #ccc; +} + +.extraction-summary .summary-section { + margin-bottom: 12px; +} + +.extraction-summary .summary-section:last-child { + margin-bottom: 0; +} + +.extraction-summary .summary-section-title { + font-size: 13px; + font-weight: 700; + color: #333; + margin-bottom: 6px; + padding-bottom: 4px; + border-bottom: 1px solid #eee; +} + +.extraction-summary .summary-row { + display: flex; + gap: 8px; + padding: 2px 0; + font-size: 12px; + line-height: 1.8; +} + +.extraction-summary .summary-label { + color: #888; + min-width: 60px; + flex-shrink: 0; +} + +.extraction-summary .summary-value { + color: #333; + flex: 1; +} + +.extraction-summary .summary-table { + width: 100%; + border-collapse: collapse; + font-size: 12px; + margin-top: 4px; +} + +.extraction-summary .summary-table th { + background: #f5f6f8; + color: #555; + font-weight: 600; + padding: 5px 8px; + text-align: left; + border-bottom: 1px solid #e0e0e0; + white-space: nowrap; +} + +.extraction-summary .summary-table td { + padding: 4px 8px; + border-bottom: 1px solid #f0f0f0; + color: #333; + vertical-align: top; +} + +.extraction-summary .summary-table tbody tr:hover { + background: #fafbfc; +} + +.extraction-summary .summary-attachment-list { + margin: 4px 0 0; + padding-left: 20px; + font-size: 12px; + color: #555; + line-height: 1.8; +} + +.extraction-summary .summary-attachment-item::marker { + color: #999; + font-size: 10px; } \ No newline at end of file diff --git a/src/web/static/js/agent.js b/src/web/static/js/agent.js index 9a5881e..5c4922b 100644 --- a/src/web/static/js/agent.js +++ b/src/web/static/js/agent.js @@ -4,6 +4,7 @@ import { App } from './state.js'; import { escapeHtml } from './utils.js'; import { addChatMessage, showStatus } from './chat.js'; +import { createSSEConnection, handleDoneResult } from './sse.js'; // ================================================================ // Agent 事件处理 @@ -59,82 +60,19 @@ function _handleAgentStateChange(msg) { } /** - * 请求补充材料 — 仅更新瞬态状态,持久提示消息由 done 事件分支负责 + * 请求补充材料 — 将 suggestion 以消息气泡形式单独展示 */ export function _handleAgentRequestSupplement(msg) { const suggestion = msg.suggestion || '请补充上传相关材料'; - showStatus(suggestion); + showStatus('信息不完整,请补充材料'); + addChatMessage(suggestion, 'system'); } /** - * 信息完整 — 已由 process.js 的 done 处理覆盖,此处不再重复显示 + * 信息完整 — 在消息气泡中提醒用户,最终结果由 done 事件处理 */ function _handleAgentReady(msg) { - // agent_ready 事件的信息展示统一由 process.js done 分支处理 -} - -/** - * 自动触发财务系统提交(由 done 事件调用) - */ -export async function handleAutoSubmit() { - if (!App.sessionId) { - addChatMessage('请先上传文件并处理', 'error'); - return; - } - - if (App.agentEventSource) { - App.agentEventSource.close(); - App.agentEventSource = null; - } - - App.isProcessing = true; - App.processState = 'submitting'; - - try { - showStatus('正在自动提交到财务系统...'); - - const response = await fetch(`/api/submit-financial/${App.sessionId}`, { - method: 'POST', - }); - - const result = await response.json(); - - if (result.status === 'started') { - const es = new EventSource(`/api/logs/${App.sessionId}`); - es.addEventListener('message', e => { - try { - const msg = JSON.parse(e.data); - if (msg.type === 'done') { - es.close(); - App.isProcessing = false; - App.processState = 'done'; - if (msg.result.submit_ok) { - addChatMessage('提交完成!', 'done'); - } else { - addChatMessage(`提交失败:${msg.result.submit_error || '未知错误'}`, 'error'); - } - } - } catch (err) { - // 普通日志行 - } - }); - - es.onerror = () => { - es.close(); - App.isProcessing = false; - App.processState = 'done'; - addChatMessage('连接中断', 'error'); - }; - } else if (result.error) { - App.isProcessing = false; - App.processState = 'done'; - addChatMessage(`提交失败:${result.error}`, 'error'); - } - } catch (e) { - App.isProcessing = false; - App.processState = 'done'; - addChatMessage(`请求失败:${e.message}`, 'error'); - } + addChatMessage(msg.message || '信息完整,正在提交到财务系统', 'done'); } /** @@ -242,49 +180,15 @@ export async function handleUserSupplement(text) { if (result.status === 'started') { showStatus('补充信息已收到,正在重新分析...'); + App.isProcessing = true; - const es = new EventSource(`/api/logs/${App.sessionId}`); - es.addEventListener('message', e => { - try { - const msg = JSON.parse(e.data); - - if (msg.type && msg.type.startsWith('agent_')) { - handleAgentEvent(msg); - return; - } - - if (msg.type === 'done') { - es.close(); - App.isProcessing = false; - - if (msg.result.waiting_for_supplement) { - App.processState = 'awaiting_supplement'; - _handleAgentRequestSupplement(msg.result); - addChatMessage('您可通过上传区补充文件,或直接输入文字说明补充信息,或输入"直接提交"跳过校验', 'system'); - } else { - App.processState = 'done'; - if (msg.result.ok) { - if (msg.result.submit_ok !== false) { - addChatMessage('信息完整,已自动提交到财务系统', 'done'); - } else { - addChatMessage(`校验通过但提交失败:${msg.result.submit_error || '未知错误'}`, 'error'); - } - } else { - addChatMessage(`分析失败:${msg.result.error || '未知错误'}`, 'error'); - } - } - } - } catch (err) { - // 普通日志 - } + createSSEConnection(App.sessionId, { + onDone: (result) => { + handleDoneResult(result); + }, + handleFileProgress: false, + handleLLMStream: false, }); - - es.onerror = () => { - es.close(); - App.isProcessing = false; - App.processState = 'done'; - addChatMessage('连接中断', 'error'); - }; } else if (result.error) { addChatMessage(`处理失败:${result.error}`, 'error'); } @@ -328,35 +232,13 @@ export async function handleForceSubmit() { App.isProcessing = true; App.processState = 'submitting'; - const es = new EventSource(`/api/logs/${App.sessionId}`); - es.addEventListener('message', e => { - try { - const msg = JSON.parse(e.data); - if (msg.type === 'done') { - es.close(); - App.isProcessing = false; - App.processState = 'done'; - App.forceSubmitting = false; - App.lastProcessedFileCount = App.allFiles.length; - if (msg.result.ok) { - addChatMessage('提交完成!', 'done'); - } else { - addChatMessage(`提交失败:${msg.result.submit_error || '未知错误'}`, 'error'); - } - } - } catch (err) { - // 普通日志行 - } + createSSEConnection(App.sessionId, { + onDone: (result) => { + handleDoneResult(result); + }, + handleFileProgress: false, + handleLLMStream: false, }); - - es.onerror = () => { - es.close(); - App.isProcessing = false; - App.processState = 'done'; - App.forceSubmitting = false; - App.lastProcessedFileCount = App.allFiles.length; - addChatMessage('连接中断', 'error'); - }; } } catch (e) { App.forceSubmitting = false; @@ -389,49 +271,13 @@ export async function handleSupplementUpload(filenames) { const result = await response.json(); if (result.status === 'started') { - const es = new EventSource(`/api/logs/${App.sessionId}`); - es.addEventListener('message', e => { - try { - const msg = JSON.parse(e.data); + App.isProcessing = true; - if (msg.type && msg.type.startsWith('agent_')) { - handleAgentEvent(msg); - return; - } - - if (msg.type === 'done') { - es.close(); - App.isProcessing = false; - - if (msg.result.waiting_for_supplement) { - App.processState = 'awaiting_supplement'; - showStatus('信息不完整,请补充上传相关材料'); - addChatMessage('您可通过上传区补充文件,或直接输入文字说明补充信息,或输入"直接提交"跳过校验', 'system'); - } else { - App.processState = 'done'; - if (msg.result.ok) { - // 后端已自动提交,根据 submit_ok 显示结果 - if (msg.result.submit_ok !== false) { - addChatMessage('信息完整,已自动提交到财务系统', 'done'); - } else { - addChatMessage(`校验通过但提交失败:${msg.result.submit_error || '未知错误'}`, 'error'); - } - } else { - addChatMessage(`分析失败:${msg.result.error || '未知错误'}`, 'error'); - } - } - } - } catch (err) { - // 普通日志 - } + createSSEConnection(App.sessionId, { + onDone: (result) => { + handleDoneResult(result); + }, }); - - es.onerror = () => { - es.close(); - App.isProcessing = false; - App.processState = 'done'; - addChatMessage('连接中断', 'error'); - }; } } catch (e) { addChatMessage(`请求失败:${e.message}`, 'error'); diff --git a/src/web/static/js/chat/stream.js b/src/web/static/js/chat/stream.js index d1dbc7f..44ca84a 100644 --- a/src/web/static/js/chat/stream.js +++ b/src/web/static/js/chat/stream.js @@ -123,11 +123,14 @@ function _appendLLMStreamReasoning(text) { } /** - * 关闭流式气泡(end 阶段)— 气泡保留在聊天历史中作为持久消息 + * 关闭流式气泡(end 阶段)— 气泡保留在聊天历史中作为持久消息。 + * 如果累积内容是提取结果 JSON,则渲染为美观的摘要卡片。 */ function _closeLLMStreamBubble(label) { if (!llmStreamState.bubble) return; + const accumulated = llmStreamState.accumulated; + llmStreamState.bubble.classList.remove('processing'); llmStreamState.bubble.classList.add('done'); @@ -141,6 +144,16 @@ function _closeLLMStreamBubble(label) { reasoningSection.open = false; } + // 尝试将累积文本解析为提取结果 JSON,渲染为摘要卡片 + const textEl = llmStreamState.bubble.querySelector('.llm-stream-text'); + if (textEl && accumulated) { + const parsed = _tryParseExtractionJson(accumulated); + if (parsed) { + textEl.innerHTML = ''; + textEl.appendChild(_buildExtractionSummary(parsed)); + } + } + scrollToBottom(); llmStreamState = { @@ -196,4 +209,199 @@ function _errorLLMStreamBubble(errorMsg) { reasoningContent: null, reasoningAccumulated: '', }; +} + +/** + * 尝试将 LLM 输出文本解析为提取结果 JSON + * 返回解析后的对象,或 null(非提取结果 JSON) + */ +function _tryParseExtractionJson(text) { + try { + const trimmed = text.trim(); + if (!trimmed.startsWith('{')) return null; + const parsed = JSON.parse(trimmed); + if (parsed && typeof parsed === 'object' && (parsed.basic_info || parsed.reimbursement_details)) { + return parsed; + } + return null; + } catch { + return null; + } +} + +/** + * 将提取结果 JSON 渲染为美观的摘要卡片 DOM 元素 + * + * @param {Object} data - 提取结果 JSON + * @returns {HTMLElement} + */ +function _buildExtractionSummary(data) { + const container = document.createElement('div'); + container.className = 'extraction-summary'; + + // 基本信息 + if (data.basic_info) { + const bi = data.basic_info; + const section = _summarySection('基本信息'); + _summaryRow(section, '出差事由', bi.travel_purpose || ''); + _summaryRow(section, '出差地点', bi.travel_location || ''); + const dates = bi.start_date && bi.end_date ? `${bi.start_date} 至 ${bi.end_date}` : (bi.start_date || bi.end_date || ''); + _summaryRow(section, '出差日期', dates); + container.appendChild(section); + } + + // 交通费 + if (data.reimbursement_details?.transport_fee?.length) { + const section = _summarySection('交通费'); + const table = _summaryTable(['日期', '出发', '到达', '金额', '备注']); + for (const item of data.reimbursement_details.transport_fee) { + _summaryTableRow(table, [ + item.start_date || '', + item.departure_place || '', + item.arrival_place || '', + item.amount != null ? `¥${item.amount.toFixed(2)}` : '', + item.remark || '', + ]); + } + section.appendChild(table); + container.appendChild(section); + } + + // 住宿费 + if (data.reimbursement_details?.hotel_fee?.length) { + const section = _summarySection('住宿费'); + const table = _summaryTable(['入住日期', '离店日期', '天数', '发票金额', '报销金额', '备注']); + for (const item of data.reimbursement_details.hotel_fee) { + _summaryTableRow(table, [ + item.checkin_date || '', + item.checkout_date || '', + item.days != null ? String(item.days) : '', + item.invoice_amount != null ? `¥${item.invoice_amount.toFixed(2)}` : '', + item.reimburse_amount != null ? `¥${item.reimburse_amount.toFixed(2)}` : '', + item.remark || '', + ]); + } + section.appendChild(table); + container.appendChild(section); + } + + // 会议/培训费 + if (data.reimbursement_details?.conference_fee?.length) { + const section = _summarySection('会议/培训费'); + const table = _summaryTable(['日期', '金额', '备注']); + for (const item of data.reimbursement_details.conference_fee) { + _summaryTableRow(table, [ + item.start_date || '', + item.amount != null ? `¥${item.amount.toFixed(2)}` : '', + item.remark || '', + ]); + } + section.appendChild(table); + container.appendChild(section); + } + + // 支付记录 + if (data.payment_methods?.length) { + const section = _summarySection('支付记录'); + const table = _summaryTable(['日期', '金额', '商户', '备注']); + for (const item of data.payment_methods) { + _summaryTableRow(table, [ + item.card_date || '', + item.card_amount != null ? `¥${item.card_amount.toFixed(2)}` : '', + item.merchant || '', + item.remark || '', + ]); + } + section.appendChild(table); + container.appendChild(section); + } + + // 补助 + if (data.subsidy_list?.length) { + const section = _summarySection('补助'); + const table = _summaryTable(['姓名', '开始日期', '结束日期', '天数']); + for (const item of data.subsidy_list) { + _summaryTableRow(table, [ + item.person_name || '', + item.start_date || '', + item.end_date || '', + item.days != null ? String(item.days) : '', + ]); + } + section.appendChild(table); + container.appendChild(section); + } + + // 附件 + if (data.attachments?.length) { + const section = _summarySection('附件'); + const list = document.createElement('ul'); + list.className = 'summary-attachment-list'; + for (const item of data.attachments) { + const li = document.createElement('li'); + li.className = `summary-attachment-item summary-attachment-${item.attachment_type || 'other'}`; + li.textContent = `${item.attachment_desc || item.filename || ''}`; + list.appendChild(li); + } + section.appendChild(list); + container.appendChild(section); + } + + return container; +} + +/* ---- 提取摘要卡片构建辅助函数 ---- */ + +function _summarySection(title) { + const section = document.createElement('div'); + section.className = 'summary-section'; + const titleEl = document.createElement('div'); + titleEl.className = 'summary-section-title'; + titleEl.textContent = title; + section.appendChild(titleEl); + return section; +} + +function _summaryRow(section, label, value) { + if (!value) return; + const row = document.createElement('div'); + row.className = 'summary-row'; + const labelEl = document.createElement('span'); + labelEl.className = 'summary-label'; + labelEl.textContent = label; + const valueEl = document.createElement('span'); + valueEl.className = 'summary-value'; + valueEl.textContent = value; + row.appendChild(labelEl); + row.appendChild(valueEl); + section.appendChild(row); +} + +function _summaryTable(headers) { + const table = document.createElement('table'); + table.className = 'summary-table'; + const thead = document.createElement('thead'); + const tr = document.createElement('tr'); + for (const h of headers) { + const th = document.createElement('th'); + th.textContent = h; + tr.appendChild(th); + } + thead.appendChild(tr); + table.appendChild(thead); + const tbody = document.createElement('tbody'); + table.appendChild(tbody); + return table; +} + +function _summaryTableRow(table, cells) { + const tbody = table.querySelector('tbody'); + if (!tbody) return; + const tr = document.createElement('tr'); + for (const cell of cells) { + const td = document.createElement('td'); + td.textContent = cell; + tr.appendChild(td); + } + tbody.appendChild(tr); } \ No newline at end of file diff --git a/src/web/static/js/config.js b/src/web/static/js/config.js index f1c53da..5341623 100644 --- a/src/web/static/js/config.js +++ b/src/web/static/js/config.js @@ -39,10 +39,6 @@ export function checkAutoStart() { } } - if (App.processState === 'awaiting_files') { - if (!App.forceStart) return; - App.forceStart = false; - } if (App.processState === 'awaiting_supplement') { return; diff --git a/src/web/static/js/process.js b/src/web/static/js/process.js index 70f4cf5..c6cbc9e 100644 --- a/src/web/static/js/process.js +++ b/src/web/static/js/process.js @@ -6,15 +6,9 @@ import { ensureSession } from './utils.js'; import { addTypingIndicator, removeTypingIndicator, - setFileProcessing, - setFileDone, - setFileCached, - setFileError, - handleLLMStream, addChatMessage, - showStatus, } from './chat.js'; -import { handleAgentEvent, handleAutoSubmit } from './agent.js'; +import { createSSEConnection, handleDoneResult } from './sse.js'; export async function startProcess() { if (!App.allFiles.length) { @@ -31,6 +25,8 @@ export async function startProcess() { await ensureSession(); for (const f of App.allFiles) { + // 跳过服务器同步的文件(已在服务器上) + if (f.__source === 'server') continue; const fd = new FormData(); fd.append('file', f); await fetch(`/api/upload/${App.sessionId}`, { method: 'POST', body: fd }); @@ -54,89 +50,13 @@ export async function startProcess() { body: JSON.stringify(cfg), }); - const es = new EventSource(`/api/logs/${App.sessionId}`); - App.agentEventSource = es; - - es.addEventListener('message', e => { - try { - const msg = JSON.parse(e.data); - - if (msg.type === 'file_progress') { - const filename = msg.file; - switch (msg.status) { - case 'processing': - setFileProcessing(filename); - break; - case 'done': - setFileDone(filename, msg.summary || {}); - break; - case 'cached': - setFileCached(filename); - break; - case 'error': - setFileError(filename, msg.error); - break; - } - return; - } - - if (msg.type === 'llm_stream') { - handleLLMStream(msg); - return; - } - - if (msg.type && msg.type.startsWith('agent_')) { - handleAgentEvent(msg); - return; - } - - if (msg.type === 'done') { - es.close(); - App.agentEventSource = null; - removeTypingIndicator(); - const result = msg.result; - App.isProcessing = false; - - if (App.forceSubmitting) { - return; - } - - if (result.waiting_for_supplement) { - App.processState = 'awaiting_supplement'; - showStatus('信息不完整,请补充上传相关材料'); - addChatMessage('您可通过上传区补充文件,或直接输入文字说明补充信息,或输入"直接提交"跳过校验', 'system'); - } else { - App.processState = 'done'; - App.lastProcessedFileCount = App.allFiles.length; - - if (result.ok) { - if (result.agent_ready) { - if (result.submit_ok !== false) { - addChatMessage('信息完整,已自动提交到财务系统', 'done'); - } else { - addChatMessage(`校验通过但提交失败:${result.submit_error || '未知错误'}`, 'error'); - } - } else if (result.invoice_count) { - addChatMessage(`处理完成!共提取 ${result.invoice_count} 张发票数据`, 'done'); - } - } else { - addChatMessage(`处理失败:${result.error || '未知错误'}`, 'error'); - } - } - return; - } - } catch (err) { - // 普通日志行 - } + const es = createSSEConnection(App.sessionId, { + onDone: (result) => { + removeTypingIndicator(); + handleDoneResult(result); + }, }); - - es.onerror = () => { - es.close(); - removeTypingIndicator(); - App.isProcessing = false; - App.processState = 'done'; - addChatMessage('连接中断,请刷新页面重试', 'error'); - }; + App.agentEventSource = es; } catch (e) { removeTypingIndicator(); diff --git a/src/web/static/js/sse.js b/src/web/static/js/sse.js new file mode 100644 index 0000000..7c145e2 --- /dev/null +++ b/src/web/static/js/sse.js @@ -0,0 +1,144 @@ +/** + * SSE 连接管理 + * + * 封装 EventSource 的创建、事件注册和 done 事件处理,避免重复代码。 + */ +import { App } from './state.js'; +import { + removeTypingIndicator, + setFileProcessing, + setFileDone, + setFileCached, + setFileError, + handleLLMStream, + addChatMessage, + showStatus, +} from './chat.js'; +import { handleAgentEvent } from './agent.js'; + +/** + * 创建 SSE 连接并返回 EventSource 实例 + * + * @param {string} sessionId - 会话 ID + * @param {Object} options - 选项配置 + * @param {Function} options.onDone - done 事件处理器 + * @param {boolean} options.handleFileProgress - 是否处理文件进度事件 (默认 true) + * @param {boolean} options.handleLLMStream - 是否处理 LLM 流式事件 (默认 true) + * @param {boolean} options.handleAgentEvents - 是否处理 Agent 事件 (默认 true) + * @returns {EventSource} + */ +export function createSSEConnection(sessionId, options = {}) { + const { + onDone, + handleFileProgress = true, + handleLLMStream: handleStream = true, + handleAgentEvents = true, + } = options; + + const es = new EventSource(`/api/logs/${sessionId}`); + + es.addEventListener('message', e => { + try { + const msg = JSON.parse(e.data); + + // 文件进度事件 + if (handleFileProgress && msg.type === 'file_progress') { + const filename = msg.file; + switch (msg.status) { + case 'processing': + setFileProcessing(filename); + break; + case 'done': + setFileDone(filename, msg.summary || {}); + break; + case 'cached': + setFileCached(filename); + break; + case 'error': + setFileError(filename, msg.error); + break; + } + return; + } + + // LLM 流式事件 + if (handleStream && msg.type === 'llm_stream') { + handleLLMStream(msg); + return; + } + + // Agent 事件 + if (handleAgentEvents && msg.type && msg.type.startsWith('agent_')) { + handleAgentEvent(msg); + return; + } + + // 完成事件 + if (msg.type === 'done') { + es.close(); + if (onDone) { + onDone(msg.result); + } + return; + } + } catch (err) { + // 普通日志行,忽略 + } + }); + + es.onerror = () => { + es.close(); + removeTypingIndicator(); + App.isProcessing = false; + App.processState = 'done'; + addChatMessage('连接中断,请刷新页面重试', 'error'); + }; + + return es; +} + +/** + * 处理 done 事件的通用逻辑 (根据 result 内容更新 UI 状态) + * + * @param {Object} result - done 事件的结果对象 + */ +export function handleDoneResult(result) { + App.isProcessing = false; + + if (App.forceSubmitting) { + App.forceSubmitting = false; + App.lastProcessedFileCount = App.allFiles.length; + if (result.ok) { + addChatMessage('提交完成!', 'done'); + } else { + addChatMessage(`提交失败:${result.error || '未知错误'}`, 'error'); + } + return; + } + + if (result.waiting_for_supplement) { + App.processState = 'awaiting_supplement'; + showStatus('信息不完整,请补充上传相关材料'); + addChatMessage( + '您可通过上传区补充文件,或直接输入文字说明补充信息,或输入"直接提交"跳过校验', + 'system' + ); + } else { + App.processState = 'done'; + App.lastProcessedFileCount = App.allFiles.length; + + if (result.ok) { + if (result.agent_ready) { + if (result.submit_ok !== false) { + addChatMessage('信息完整,已自动提交到财务系统', 'done'); + } else { + addChatMessage(`校验通过但提交失败:${result.submit_error || '未知错误'}`, 'error'); + } + } else if (result.invoice_count) { + addChatMessage(`处理完成!共提取 ${result.invoice_count} 张发票数据`, 'done'); + } + } else { + addChatMessage(`处理失败:${result.error || '未知错误'}`, 'error'); + } + } +} \ No newline at end of file diff --git a/src/web/static/js/sync.js b/src/web/static/js/sync.js index 9a0c110..3ff8b48 100644 --- a/src/web/static/js/sync.js +++ b/src/web/static/js/sync.js @@ -40,12 +40,10 @@ async function syncFiles() { const localNames = new Set(App.allFiles.map(f => f.name)); const addedFiles = []; + // 只同步文件名元数据,不下载文件内容(文件已在服务器上) for (const serverFile of serverFiles) { if (!localNames.has(serverFile.name)) { - const resp = await fetch(`/api/download/${App.sessionId}/${encodeURIComponent(serverFile.name)}`); - const blob = await resp.blob(); - const file = new File([blob], serverFile.name, { type: blob.type }); - file.__source = 'server'; + const file = { name: serverFile.name, __source: 'server' }; App.allFiles.push(file); addedFiles.push(serverFile.name); } diff --git a/src/web/static/js/upload.js b/src/web/static/js/upload.js index 2bc960b..a7ca73d 100644 --- a/src/web/static/js/upload.js +++ b/src/web/static/js/upload.js @@ -9,6 +9,29 @@ import { ensureSession } from './utils.js'; export async function handleFiles(input) { const files = Array.from(input.files); + const result = await _processFiles(files); + + if (result.hasNewFiles && App.processState === 'done') { + App.processState = 'idle'; + App.newFilenames = result.newFilenames; + } + if (result.hasNewFiles && App.processState === 'awaiting_supplement') { + hideAgentRequest(); + await _uploadNewFiles(result.newFilenames); + App.newFilenames = result.newFilenames; + App.processState = 'idle'; + checkAutoStart(); + return; + } + + promptMissingConfig(); + input.value = ''; +} + +/** + * 统一文件处理逻辑(handleFiles 和拖拽共享) + */ +async function _processFiles(files) { const validExts = /\.(pdf|png|jpe?g|bmp|webp)$/i; let hasNewFiles = false; const newFilenames = []; @@ -28,21 +51,7 @@ export async function handleFiles(input) { } } - if (hasNewFiles && App.processState === 'done') { - App.processState = 'idle'; - App.newFilenames = newFilenames; - } - if (hasNewFiles && App.processState === 'awaiting_supplement') { - hideAgentRequest(); - await _uploadNewFiles(newFilenames); - App.newFilenames = newFilenames; - App.processState = 'idle'; - checkAutoStart(); - return; - } - - promptMissingConfig(); - input.value = ''; + return { hasNewFiles, newFilenames }; } async function _uploadNewFiles(filenames) { @@ -90,35 +99,18 @@ export function initDragDrop() { zone.addEventListener('drop', async e => { e.preventDefault(); zone.classList.remove('dragover'); - const validExts = /\.(pdf|png|jpe?g|bmp|webp)$/i; const droppedFiles = Array.from(e.dataTransfer.files); - let hasNewFiles = false; - const newFilenames = []; + const result = await _processFiles(droppedFiles); - for (const f of droppedFiles) { - if (/\.json$/i.test(f.name)) { - await parseConfigFile(f); - continue; - } - if (validExts.test(f.name) && !App.allFiles.find(x => x.name === f.name)) { - f.__source = 'local'; - App.allFiles.push(f); - zone.classList.add('active'); - addFileMessage(f.name); - hasNewFiles = true; - newFilenames.push(f.name); - } - } - - if (hasNewFiles && App.processState === 'done') { + if (result.hasNewFiles && App.processState === 'done') { App.processState = 'idle'; } // 在 awaiting_supplement 状态下,只上传新文件并走 supplement 流程 - if (hasNewFiles && App.processState === 'awaiting_supplement') { + if (result.hasNewFiles && App.processState === 'awaiting_supplement') { hideAgentRequest(); - await _uploadNewFiles(newFilenames); - App.newFilenames = newFilenames; + await _uploadNewFiles(result.newFilenames); + App.newFilenames = result.newFilenames; App.processState = 'idle'; checkAutoStart(); return; diff --git a/tests/test_matcher.py b/tests/test_matcher.py index 46167d5..d71f0a4 100644 --- a/tests/test_matcher.py +++ b/tests/test_matcher.py @@ -17,6 +17,7 @@ from src.doc.matcher import ( _build_payment_records, _invoices_to_records, _match, + _match_by_filename, _match_one_to_many, _match_one_to_one, _relative_tolerance, @@ -562,3 +563,304 @@ class TestMatchInvoicesToCards: result = match_invoices_to_cards(invoices, cards=cards) assert len(result) == 2 + + +# ------------------------------------------------------------------ +# 文件名匹配 +# ------------------------------------------------------------------ + + +class TestMatchByFilename: + """文件名匹配:发票和刷卡记录的文件名(不含后缀)一致时直接匹配""" + + def _make_invoice_with_file(self, number: str, amount: float, filename: str) -> dict[str, Any]: + inv = _make_invoice(number, amount) + inv["_source_file"] = filename + return inv + + def _make_card_with_file(self, date: str, amount: float, filename: str) -> dict[str, Any]: + card = _make_card(date, amount) + card["_source_file"] = filename + return card + + def test_exact_filename_match(self): + invoices = [ + self._make_invoice_with_file("A", 500, "receipt_001.pdf"), + self._make_invoice_with_file("B", 300, "receipt_002.pdf"), + ] + cards = [ + self._make_card_with_file("2026-01-01", 500, "receipt_001.png"), + self._make_card_with_file("2026-01-02", 300, "receipt_002.png"), + ] + for inv in invoices: + inv["_amount"] = _safe_float(inv[K_TOTAL_AMOUNT]) + for card in cards: + card["_amount"] = _safe_float(card[K_CARD_AMOUNT]) + + assigned: set[int] = set() + result: dict[int, list[int]] = {} + _match_by_filename(invoices, cards, assigned, result) + + assert 0 in result + assert 1 in result + assert result[0] == [0] + assert result[1] == [1] + assert 0 in assigned + assert 1 in assigned + + def test_filename_match_ignores_extension(self): + invoices = [self._make_invoice_with_file("A", 500, "data.pdf")] + cards = [self._make_card_with_file("2026-01-01", 500, "data.png")] + for inv in invoices: + inv["_amount"] = _safe_float(inv[K_TOTAL_AMOUNT]) + for card in cards: + card["_amount"] = _safe_float(card[K_CARD_AMOUNT]) + + assigned: set[int] = set() + result: dict[int, list[int]] = {} + _match_by_filename(invoices, cards, assigned, result) + + assert 0 in result + assert result[0] == [0] + + def test_filename_no_match_when_stems_differ(self): + invoices = [self._make_invoice_with_file("A", 500, "invoice_A.pdf")] + cards = [self._make_card_with_file("2026-01-01", 500, "card_001.png")] + for inv in invoices: + inv["_amount"] = _safe_float(inv[K_TOTAL_AMOUNT]) + for card in cards: + card["_amount"] = _safe_float(card[K_CARD_AMOUNT]) + + assigned: set[int] = set() + result: dict[int, list[int]] = {} + _match_by_filename(invoices, cards, assigned, result) + + assert 0 not in result + assert len(assigned) == 0 + + def test_skips_invoice_without_source_file(self): + invoices = [_make_invoice("A", 500)] + cards = [self._make_card_with_file("2026-01-01", 500, "card.png")] + for inv in invoices: + inv["_amount"] = _safe_float(inv[K_TOTAL_AMOUNT]) + for card in cards: + card["_amount"] = _safe_float(card[K_CARD_AMOUNT]) + + assigned: set[int] = set() + result: dict[int, list[int]] = {} + _match_by_filename(invoices, cards, assigned, result) + + assert 0 not in result + + def test_skips_card_without_source_file(self): + invoices = [self._make_invoice_with_file("A", 500, "inv.pdf")] + cards = [_make_card("2026-01-01", 500)] + for inv in invoices: + inv["_amount"] = _safe_float(inv[K_TOTAL_AMOUNT]) + for card in cards: + card["_amount"] = _safe_float(card[K_CARD_AMOUNT]) + + assigned: set[int] = set() + result: dict[int, list[int]] = {} + _match_by_filename(invoices, cards, assigned, result) + + assert 0 not in result + + def test_skips_empty_source_file(self): + inv = _make_invoice("A", 500) + inv["_source_file"] = "" + cards = [self._make_card_with_file("2026-01-01", 500, "card.png")] + for _inv in [inv]: + _inv["_amount"] = _safe_float(_inv[K_TOTAL_AMOUNT]) + for card in cards: + card["_amount"] = _safe_float(card[K_CARD_AMOUNT]) + + assigned: set[int] = set() + result: dict[int, list[int]] = {} + _match_by_filename([inv], cards, assigned, result) + + assert 0 not in result + + def test_only_first_unassigned_invoice_matches(self): + invoices = [ + self._make_invoice_with_file("A", 500, "same.pdf"), + self._make_invoice_with_file("B", 300, "same.pdf"), + ] + cards = [self._make_card_with_file("2026-01-01", 500, "same.png")] + for inv in invoices: + inv["_amount"] = _safe_float(inv[K_TOTAL_AMOUNT]) + for card in cards: + card["_amount"] = _safe_float(card[K_CARD_AMOUNT]) + + assigned: set[int] = set() + result: dict[int, list[int]] = {} + _match_by_filename(invoices, cards, assigned, result) + + assert 0 in result + assert result[0] == [0] + assert 0 in assigned + assert 1 not in assigned + + +# ------------------------------------------------------------------ +# 文件名匹配与金额匹配的交互 +# ------------------------------------------------------------------ + + +class TestFilenameAndAmountInteraction: + """文件名预匹配后,已分配的发票不会被后续金额匹配重复处理""" + + def _prepare_with_files(self, invoices, cards): + for inv in invoices: + inv["_amount"] = _safe_float(inv[K_TOTAL_AMOUNT]) + for card in cards: + card["_amount"] = _safe_float(card[K_CARD_AMOUNT]) + invoices.sort(key=lambda i: i["_amount"], reverse=True) + cards.sort(key=lambda c: c["_amount"], reverse=True) + + def _make_invoice_with_file(self, number: str, amount: float, filename: str) -> dict[str, Any]: + inv = _make_invoice(number, amount) + inv["_source_file"] = filename + return inv + + def _make_card_with_file(self, date: str, amount: float, filename: str) -> dict[str, Any]: + card = _make_card(date, amount) + card["_source_file"] = filename + return card + + def test_filename_matched_invoice_skipped_in_one_to_one(self): + """文件名预匹配后,_match_one_to_one 应跳过已分配的发票""" + invoices = [ + self._make_invoice_with_file("A", 500, "match.pdf"), + self._make_invoice_with_file("B", 300, "other.pdf"), + ] + cards = [ + self._make_card_with_file("2026-01-01", 500, "match.png"), + self._make_card_with_file("2026-01-02", 300, "other.png"), + ] + self._prepare_with_files(invoices, cards) + + assigned: set[int] = set() + result: dict[int, list[int]] = {} + _match_by_filename(invoices, cards, assigned, result) + _match_one_to_one(invoices, cards, 0.03, assigned, result) + + assert 0 in result + assert result[0] == [0] + assert len(assigned) == 2 + + def test_filename_match_takes_priority_over_amount(self): + """即使金额不匹配,文件名匹配仍优先""" + invoices = [self._make_invoice_with_file("A", 1000, "same.pdf")] + cards = [self._make_card_with_file("2026-01-01", 500, "same.png")] + self._prepare_with_files(invoices, cards) + + assigned: set[int] = set() + result: dict[int, list[int]] = {} + _match_by_filename(invoices, cards, assigned, result) + + assert 0 in result + assert result[0] == [0] + assert 0 in assigned + + def test_partial_filename_match_falls_back_to_amount(self): + """部分文件名匹配后,剩余发票走金额匹配""" + invoices = [ + self._make_invoice_with_file("A", 500, "match.pdf"), + self._make_invoice_with_file("B", 300, "no_match.pdf"), + ] + cards = [ + self._make_card_with_file("2026-01-01", 500, "match.png"), + self._make_card_with_file("2026-01-02", 300, "different.png"), + ] + self._prepare_with_files(invoices, cards) + + result = _match(cards, invoices, 0.03) + + assert 0 in result + assert 1 in result + assert len(result) == 2 + + def test_filename_match_in_one_to_many_prevents_reuse(self): + """一对多场景下,文件名匹配的发票不会被贪心匹配复用""" + invoices = [ + self._make_invoice_with_file("A", 500, "same.pdf"), + self._make_invoice_with_file("B", 200, "other.pdf"), + ] + cards = [self._make_card_with_file("2026-01-01", 500, "same.png")] + self._prepare_with_files(invoices, cards) + + result = _match(cards, invoices, 0.03) + + assert 0 in result + assert result[0] == [0] + assert len(result[0]) == 1 + + +# ------------------------------------------------------------------ +# 端到端:文件名匹配集成 +# ------------------------------------------------------------------ + + +class TestEndToEndFilenameMatching: + """match_invoices_to_cards 端到端测试(含文件名匹配)""" + + def _make_invoice_with_file(self, number: str, amount: float, filename: str) -> dict[str, Any]: + inv = _make_invoice(number, amount) + inv["_source_file"] = filename + return inv + + def _make_card_with_file(self, date: str, amount: float, filename: str) -> dict[str, Any]: + card = _make_card(date, amount) + card["_source_file"] = filename + return card + + def test_full_filename_match(self): + invoices = [ + self._make_invoice_with_file("A", 500, "receipt_001.pdf"), + self._make_invoice_with_file("B", 300, "receipt_002.pdf"), + ] + cards = [ + self._make_card_with_file("2026-01-01", 500, "receipt_001.png"), + self._make_card_with_file("2026-01-02", 300, "receipt_002.png"), + ] + result = match_invoices_to_cards(invoices, cards=cards) + + assert len(result) == 2 + for rec in result: + assert rec[K_RELATIVE_INVOICE_COUNT] == "1" + + def test_mixed_filename_and_amount_match(self): + """部分文件名匹配 + 部分金额匹配""" + invoices = [ + self._make_invoice_with_file("A", 500, "match.pdf"), + self._make_invoice_with_file("B", 300, "no_match.pdf"), + ] + cards = [ + self._make_card_with_file("2026-01-01", 500, "match.png"), + self._make_card_with_file("2026-01-02", 300, "different.png"), + ] + result = match_invoices_to_cards(invoices, cards=cards) + + assert len(result) == 2 + + def test_filename_match_with_no_cards_has_files(self): + """无支付记录时,文件名信息不影响结果""" + invoices = [ + self._make_invoice_with_file("A", 100, "inv.pdf"), + self._make_invoice_with_file("B", 200, "inv2.pdf"), + ] + result = match_invoices_to_cards(invoices, cards=None) + + assert len(result) == 2 + + def test_internal_fields_cleaned_after_filename_match(self): + """文件名匹配后,内部字段仍被清理""" + invoices = [self._make_invoice_with_file("A", 500, "same.pdf")] + cards = [self._make_card_with_file("2026-01-01", 500, "same.png")] + result = match_invoices_to_cards(invoices, cards=cards) + + for inv in result[0][K_MATCHED_INVOICES]: + assert "_amount" not in inv + for card in cards: + assert "_amount" not in card diff --git a/uv.lock b/uv.lock index 2321865..06589a3 100644 --- a/uv.lock +++ b/uv.lock @@ -236,7 +236,7 @@ dev = [ requires-dist = [ { name = "flask", specifier = ">=3.0" }, { name = "llama-index", specifier = ">=0.12.0" }, - { name = "llama-index-llms-openai-like", specifier = "==0.7.2" }, + { name = "llama-index-llms-openai-like", specifier = ">=0.7.2" }, { name = "playwright", specifier = ">=1.40" }, { name = "pymupdf", specifier = ">=1.24" }, { name = "python-dotenv", specifier = ">=1.0" },