Files
Auto-Finance/src/agent/orchestrator.py
2026-06-15 10:39:01 +08:00

530 lines
17 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Agent 协调器
作为调度中枢,编排信息提取、规则校验的完整流程。
校验-修正循环由 Agent 层调度:
1. Agent 调用 LLM 提取信息
2. Agent 调用 validator.py 校验
3. 校验失败则构建修正提示,再次调用 LLM
4. 重复直到校验通过或达到最大重试次数
5. LLM 在输出中包含 can_submit 和 suggestion 字段,用于判断信息完整性
状态机:
idle -> extracting -> awaiting_supplement -> (回到extracting)
|
(完整) -> ready -> submitting -> done
|
(用户强制) -> ready
"""
from __future__ import annotations
import json
from dataclasses import asdict, dataclass, field
from enum import StrEnum
from pathlib import Path
from typing import Any
from .. import get_logger
from ..doc.llm_extractor import (
build_extraction_user_message,
llm_query_text,
load_cache,
load_match_result,
parse_json_response,
)
from ..doc.prompt import (
build_normal_info_system_prompt,
build_travel_info_system_prompt,
)
from ..doc.validator import validate_extracted_info
from ..pipeline_core import save_cache_info
log = get_logger("agent")
# 规则校验-修正循环的最大重试次数
MAX_VALIDATION_RETRIES = 3
# ------------------------------------------------------------------
# 状态枚举
# ------------------------------------------------------------------
class AgentState(StrEnum):
IDLE = "idle"
EXTRACTING = "extracting"
AWAITING_SUPPLEMENT = "awaiting_supplement"
READY = "ready"
SUBMITTING = "submitting"
DONE = "done"
ERROR = "error"
# ------------------------------------------------------------------
# Agent 会话数据模型
# ------------------------------------------------------------------
@dataclass
class AgentSession:
"""Agent 会话状态"""
session_id: str
state: AgentState = AgentState.IDLE
rounds: int = 0
max_rounds: int = 5
invoice_type: str = "travel" # "travel" 或 "normal"
extracted_info: dict[str, Any] = field(default_factory=dict)
validation_reports: list[dict[str, Any]] = field(default_factory=list)
user_supplements: list[str] = field(default_factory=list)
error_message: str = ""
def to_dict(self) -> dict[str, Any]:
return asdict(self)
@classmethod
def from_dict(cls, data: dict[str, Any]) -> AgentSession:
# 兼容旧版本state 可能是字符串
if "state" in data and isinstance(data["state"], str):
data["state"] = AgentState(data["state"])
return cls(**data)
# ------------------------------------------------------------------
# 持久化
# ------------------------------------------------------------------
AGENT_STATE_FILE = "agent_state.json"
def save_agent_state(session_dir: Path, session: AgentSession) -> None:
"""将 Agent 会话状态持久化到 session 目录。"""
state_path = session_dir / AGENT_STATE_FILE
tmp_path = session_dir / (AGENT_STATE_FILE + ".tmp")
with open(tmp_path, "w", encoding="utf-8") as f:
json.dump(session.to_dict(), f, ensure_ascii=False, indent=2)
tmp_path.replace(state_path)
def load_agent_state(session_dir: Path) -> AgentSession | None:
"""从 session 目录加载 Agent 会话状态。"""
state_path = session_dir / AGENT_STATE_FILE
if not state_path.exists():
return None
try:
with open(state_path, encoding="utf-8") as f:
return AgentSession.from_dict(json.load(f))
except Exception as e:
log.warning("加载 Agent 状态失败: %s", e)
return None
# ------------------------------------------------------------------
# SSE 事件发射
# ------------------------------------------------------------------
AGENT_EVENT_LOG = "agent_events.log"
# 去重守卫:记录每个 session 上一次发射的事件类型,防止连续重复发射
# key: str(session_dir), value: 上一次的 event_type
_last_event_type: dict[str, str] = {}
def _emit_agent_event(session_dir: Path, event_type: str, **kwargs: Any) -> None:
"""向 agent_events.log 追加一行 JSON 事件。
同一 session 连续发射相同 event_type 时直接抛出 RuntimeError
强制调用方修复重复发射的代码,而非静默掩盖。
"""
session_key = str(session_dir)
prev = _last_event_type.get(session_key)
if prev == event_type:
raise RuntimeError(
f"事件重复发射: session={session_dir.name!r}, event_type={event_type!r}"
f"请检查调用链,确保每个事件类型只发射一次。"
)
_last_event_type[session_key] = event_type
event = {"type": event_type, **kwargs}
try:
event_path = session_dir / AGENT_EVENT_LOG
with open(event_path, "a", encoding="utf-8") as f:
f.write(json.dumps(event, ensure_ascii=False) + "\n")
except Exception:
pass
# ------------------------------------------------------------------
# 辅助:构建修正提示
# ------------------------------------------------------------------
def _build_correction_prompt(
base_message: str,
report: Any,
) -> str:
"""根据校验报告构建修正提示,追加到原始用户消息后。"""
error_feedback = (
f"\n\n=== 上一次输出的校验结果 ===\n"
f"校验未通过,发现以下问题:\n"
f"缺失字段 ({len(report.missing_fields)} 个){', '.join(report.missing_fields)}\n"
)
if report.missing_materials:
error_feedback += f"可能需要补充的材料:{', '.join(report.missing_materials)}\n"
if report.suggestion:
error_feedback += f"建议:{report.suggestion}\n"
error_feedback += (
"\n请根据以上校验结果修正你的输出,确保所有必填字段都有值。"
"如果某个字段确实没有数据,请给出合理的猜测值。"
"再次返回完整的 JSON 结果。"
)
return base_message + error_feedback
# ------------------------------------------------------------------
# 核心协调逻辑
# ------------------------------------------------------------------
def _do_extraction_with_validation(
session_dir: Path,
session: AgentSession,
previous_analysis: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""Agent 调度的提取-校验-修正循环。
流程:
1. 加载缓存数据和匹配结果
2. 构建用户消息
3. 调用 LLM 提取
4. 调用 validator 校验
5. 校验失败则构建修正提示,回到步骤 3
6. 最多重试 MAX_VALIDATION_RETRIES 次
LLM 输出的 JSON 中额外包含 can_submit 和 suggestion 字段,
用于判断信息是否完整可提交。
Args:
session_dir: 会话目录。
session: 当前 Agent 会话。
previous_analysis: 上一轮 LLM 分析结果(可选,补充文件时传入作为历史上下文)。
Returns:
校验通过的结构化数据(或达到重试上限后的最佳结果)。
"""
cache_map = load_cache(session_dir)
match_result = load_match_result(session_dir)
# 选择系统提示词
if session.invoice_type == "travel":
system_prompt = build_travel_info_system_prompt()
else:
system_prompt = build_normal_info_system_prompt()
base_message = build_extraction_user_message(cache_map, match_result, previous_analysis=previous_analysis)
current_message = base_message
for attempt in range(1, MAX_VALIDATION_RETRIES + 1):
log.info(
"LLM 提取第 %d/%d 次尝试 (%s)",
attempt,
MAX_VALIDATION_RETRIES,
session.invoice_type,
)
_emit_agent_event(
session_dir,
"agent_state_change",
state=AgentState.EXTRACTING,
round=session.rounds,
attempt=attempt,
message=f"正在分析文件... (第{attempt}次)",
)
# Step 1: 调用 LLM 提取
try:
response = llm_query_text(
system_prompt=system_prompt,
text=current_message,
reasoning_effort="low",
source_dir=session_dir,
)
result = parse_json_response(response)
except Exception as e:
log.error("LLM 提取失败: %s", e)
_emit_agent_event(
session_dir,
"agent_error",
message=f"LLM 提取失败: {e}",
)
raise
# Step 2: 调用 validator 校验
report = validate_extracted_info(result, invoice_type=session.invoice_type)
if report.valid:
log.info("规则校验通过 (第 %d 次尝试)", attempt)
_emit_agent_event(
session_dir,
"agent_state_change",
state=AgentState.EXTRACTING,
round=session.rounds,
attempt=attempt,
message=f"规则校验通过 (第{attempt}次)",
)
return result
# Step 3: 校验失败,构建修正提示
log.warning(
"规则校验未通过 (第 %d/%d 次): 缺失 %d 个字段 - %s",
attempt,
MAX_VALIDATION_RETRIES,
len(report.missing_fields),
report.missing_fields,
)
_emit_agent_event(
session_dir,
"agent_state_change",
state=AgentState.EXTRACTING,
round=session.rounds,
attempt=attempt,
message=f"规则校验未通过,缺失 {len(report.missing_fields)} 个字段,正在请求 LLM 修正...",
)
current_message = _build_correction_prompt(current_message, report)
# 所有重试都失败,返回最后一次结果
log.error(
"LLM 提取经过 %d 次尝试仍未通过规则校验,返回最后一次结果 (置信度: %.0f%%)",
MAX_VALIDATION_RETRIES,
report.confidence * 100,
)
return result
def run_agent_round(
session_dir: Path,
session: AgentSession,
new_files: list[str] | None = None,
) -> AgentSession:
"""执行一轮 Agent 处理:提取-校验-修正循环。
Args:
session_dir: 会话目录。
session: 当前 Agent 会话。
new_files: 新增的文件列表(可选,补充文件时传入)。
Returns:
更新后的 Agent 会话。
注意:
- 信息提取优先从 .invoice_cache 缓存读取,避免重复调用 LLM。
- Agent 调度提取-校验-修正循环LLM 提取 -> validator 校验 -> 失败则反馈修正。
- 若缓存缺失则执行提取后立即写回缓存travel_info.json / normal_info.json
- 补充文件时new_files 非空),加载上一轮分析结果作为历史上下文,强制重新分析。
"""
# 终态保护:会话已提交或已完成时不再重复处理
if session.state in (AgentState.DONE, AgentState.SUBMITTING, AgentState.READY):
log.info("Agent 会话已处于终态 (%s),跳过重复处理", session.state.value)
return session
if session.rounds >= session.max_rounds:
session.state = AgentState.ERROR
session.error_message = f"已达到最大轮次 ({session.max_rounds}),请检查信息或强制提交"
log.warning("Agent 达到最大轮次限制")
_emit_agent_event(
session_dir,
"agent_max_rounds",
message=session.error_message,
)
return session
session.rounds += 1
log.info("开始第 %d 轮 Agent 处理", session.rounds)
# ---- Step 1: 信息提取Agent 调度校验-修正循环) ----
session.state = AgentState.EXTRACTING
_emit_agent_event(
session_dir,
"agent_state_change",
state=session.state,
round=session.rounds,
message="正在分析文件...",
)
# 判断是否为补充文件场景:有新文件传入时,加载上一轮分析结果作为上下文
cache_map = load_cache(session_dir)
is_supplement = bool(new_files)
previous_analysis = None
if is_supplement:
info_key = "travel_info" if session.invoice_type == "travel" else "normal_info"
previous_analysis = cache_map.get(info_key)
if previous_analysis:
log.info("检测到补充文件,加载上一轮分析结果作为历史上下文")
try:
info_key = "travel_info" if session.invoice_type == "travel" else "normal_info"
should_reanalyze = not cache_map.get(info_key) or is_supplement
if should_reanalyze:
session.extracted_info = _do_extraction_with_validation(
session_dir, session, previous_analysis=previous_analysis
)
# 提取后立即写入缓存,后续步骤依赖此数据
save_cache_info(session_dir, info_key, session.extracted_info)
else:
session.extracted_info = cache_map[info_key]
except Exception as e:
session.state = AgentState.ERROR
session.error_message = f"信息提取失败: {e}"
log.error("Agent 信息提取失败: %s", e)
_emit_agent_event(
session_dir,
"agent_error",
message=session.error_message,
)
return session
# ---- 判断结果(从 LLM 提取结果中读 can_submit ----
can_submit = session.extracted_info.get("can_submit", True)
suggestion = session.extracted_info.get("suggestion", "")
if can_submit:
session.state = AgentState.READY
_emit_agent_event(
session_dir,
"agent_ready",
round=session.rounds,
message="信息完整,可以提交",
)
log.info("Agent 校验通过,信息完整")
else:
session.state = AgentState.AWAITING_SUPPLEMENT
combined_suggestion = suggestion or "信息不完整,请补充材料"
_emit_agent_event(
session_dir,
"agent_request_supplement",
round=session.rounds,
missing_fields=[],
missing_materials=[],
semantic_issues=[],
suggestion=combined_suggestion,
)
log.info("Agent 请求补充: %s", combined_suggestion)
save_agent_state(session_dir, session)
return session
def force_submit(
session_dir: Path,
session: AgentSession,
) -> AgentSession:
"""用户强制提交,跳过校验。"""
session.state = AgentState.READY
log.info("用户强制提交,跳过校验")
_emit_agent_event(
session_dir,
"agent_force_submit",
message="用户选择强制提交",
)
save_agent_state(session_dir, session)
return session
def add_supplement(
session_dir: Path,
session: AgentSession,
filenames: list[str],
) -> AgentSession:
"""记录用户补充的文件。"""
session.user_supplements.extend(filenames)
_emit_agent_event(
session_dir,
"agent_supplement_received",
files=filenames,
)
log.info("收到用户补充文件: %s", filenames)
save_agent_state(session_dir, session)
return session
def process_user_text_supplement(
session_dir: Path,
session: AgentSession,
user_text: str,
) -> AgentSession:
"""处理用户通过文字补充的信息。
流程:
1. LLM 分析用户文字,提取需要更新的字段
2. 合并到已提取的信息中
3. 保存到缓存
4. 重新执行一轮 Agent 校验
Args:
session_dir: 会话目录。
session: 当前 Agent 会话。
user_text: 用户输入的文字。
Returns:
更新后的 Agent 会话。
"""
from ..doc.llm_extractor import (
merge_supplement_into_info,
process_user_supplement,
)
log.info("收到用户文字补充: %s", user_text)
_emit_agent_event(
session_dir,
"agent_supplement_received",
files=[user_text[:50]], # 简短显示
)
# Step 1: LLM 分析用户文字
supplement_result = process_user_supplement(
user_text=user_text,
extracted_info=session.extracted_info,
invoice_type=session.invoice_type,
source_dir=session_dir,
)
updated_fields = supplement_result.get("updated_fields", {})
unparsed = supplement_result.get("unparsed_info", "")
if updated_fields:
# Step 2: 合并到已提取信息
session.extracted_info = merge_supplement_into_info(
session.extracted_info,
updated_fields,
)
# Step 3: 保存到缓存
info_key = "travel_info" if session.invoice_type == "travel" else "normal_info"
save_cache_info(session_dir, info_key, session.extracted_info)
log.info("已更新 %s", info_key)
# Step 4: 重新执行 Agent 校验
session.state = AgentState.EXTRACTING
_emit_agent_event(
session_dir,
"agent_state_change",
state=session.state,
round=session.rounds,
message="正在重新校验...",
)
session = run_agent_round(session_dir, session)
else:
# 没有可更新的字段
msg = unparsed or "未识别到可更新的报销信息"
_emit_agent_event(
session_dir,
"agent_supplement_received",
files=[msg],
)
log.info("用户补充未识别到有效信息: %s", msg)
return session