自研框架调度引擎与执行流控制
专题定位:深入自研 Agent 框架的调度内核,从执行主循环、消息传递、工具调度、状态管理、记忆存取、规划推理到并发控制,逐层拆解七大核心机理的设计原理、工程实现与性能权衡。每章配备 Mermaid 架构图、核心模块源码级实现与量化分析矩阵。
前置阅读:第 08 篇《自研 Agent 框架总体架构设计》(四层架构、核心抽象、事件驱动、插件体系)
代码规范:所有示例代码遵循 Python 3.11+ 类型注解,异步优先(async/await),接口面向抽象而非实现。
目录
- 第 1 章 执行机理:Agent 主循环的心脏搏动
- 第 2 章 消息传递机理:事件驱动的组件通信
- 第 3 章 工具调度机理:从发现到执行的五步链路
- 第 4 章 状态管理机理:检查点快照与时间旅行
- 第 5 章 记忆存取机理:编码-存储-检索-衰减全链路
- 第 6 章 规划推理机理:DAG 分解与动态重规划
- 第 7 章 并发控制机理:协程-信号量-令牌桶三级体系
- 第 8 章 七大机理协同工作:完整执行流全景
- 第 9 章 性能基准测试与优化策略
- 第 10 章 与主流框架调度机制对比
第 1 章 执行机理:Agent 主循环的心脏搏动
1.1 主循环的本质
Agent 主循环(Agent Loop)是整个框架的心脏——它驱动着感知、推理、规划、行动、观察、反思六个阶段的持续流转。没有主循环,Agent 只是一堆零散的组件;有了主循环,组件才能协同成一个"会思考、会行动"的智能体。
从控制论的视角看,Agent 主循环本质上是一个感知-行动循环(Perception-Action Loop)的扩展版本,它在经典控制论反馈环的基础上,增加了推理、规划和反思三个认知层:
经典控制论: 感知 → 决策 → 行动 → 反馈
Agent 扩展: 感知 → 推理 → 规划 → 行动 → 观察 → 反思
1.2 六阶段流转全景
六阶段详解:
| 阶段 | 名称 | 核心职责 | 输入 | 输出 |
|---|---|---|---|---|
| 1 | Perceive(感知) | 接收外部输入,更新状态,写入短期记忆 | 用户输入 / 工具返回 / 环境事件 | AgentState.messages 更新 |
| 2 | Reason(推理) | 组装 Prompt,调用 LLM,解析输出 | system_prompt + memory + messages | LLM 响应(文本 + tool_calls) |
| 3 | Plan(规划) | 判断是否需要工具,选择工具+参数,确定执行顺序 | LLM tool_calls | 执行计划(串行/并行) |
| 4 | Act(行动) | 执行工具调用,沙箱隔离,超时控制,错误捕获 | 工具 + 参数 | ToolResult |
| 5 | Observe(观察) | 解析工具结果,写入记忆,判断终止条件 | ToolResult | 更新后的状态 + 是否继续 |
| 6 | Reflect(反思) | 评估结果质量,总结经验,调整策略 | 最终结果 + 原始目标 | 反思笔记(可选) |
1.3 核心实现:AgentLoop 类
"""
执行机理核心实现 - AgentLoop
"""
from __future__ import annotations
import asyncio
import json
import logging
from dataclasses import dataclass, field
from datetime import datetime
from typing import Any, Optional
logger = logging.getLogger(__name__)
@dataclass
class AgentState:
"""Agent 运行时状态"""
session_id: str
max_steps: int = 20
step: int = 0
tool_call_count: int = 0
total_tokens: int = 0
start_time: datetime = field(default_factory=datetime.now)
status: str = "initialized" # initialized/running/paused/completed/error
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass
class StepMetrics:
"""单步性能指标"""
step: int
perceive_ms: float = 0.0
reason_ms: float = 0.0
plan_ms: float = 0.0
act_ms: float = 0.0
observe_ms: float = 0.0
reflect_ms: float = 0.0
tool_calls: int = 0
tokens_used: int = 0
class AgentLoop:
"""
Agent 主循环实现
驱动 Perceive → Reason → Plan → Act → Observe → Reflect 的完整流转。
支持最大步数限制、中间件拦截、事件通知、检查点恢复。
"""
def __init__(
self,
agent: "Agent",
max_steps: int = 20,
enable_reflection: bool = False,
middlewares: Optional[list["IMiddleware"]] = None,
):
self.agent = agent
self.max_steps = max_steps
self.enable_reflection = enable_reflection
self.middlewares = middlewares or []
self.metrics: list[StepMetrics] = []
async def run(self, query: str) -> str:
"""
执行一个完整的 Agent 循环
Args:
query: 用户输入查询
Returns:
Agent 最终响应文本
"""
state = self.agent.get_state()
state.max_steps = self.max_steps
state.status = "running"
logger.info(f"[AgentLoop] 开始执行 session={state.session_id}")
# 执行中间件 before 链
context = {"query": query, "agent": self.agent}
for mw in self.middlewares:
context = await mw.before(context)
if context.get("_blocked"):
return context.get("_block_reason", "请求被拦截")
# 主循环
result = ""
for step in range(self.max_steps):
state.step = step
step_metrics = StepMetrics(step=step)
# 1. Perceive - 感知
t0 = asyncio.get_event_loop().time()
await self._perceive(query if step == 0 else "", step)
step_metrics.perceive_ms = (asyncio.get_event_loop().time() - t0) * 1000
# 2. Reason - 推理
t0 = asyncio.get_event_loop().time()
response = await self._reason()
step_metrics.reason_ms = (asyncio.get_event_loop().time() - t0) * 1000
step_metrics.tokens_used = getattr(response, "usage", {}).get("total_tokens", 0)
state.total_tokens += step_metrics.tokens_used
# 3. Plan - 规划(判断是否需要工具调用)
t0 = asyncio.get_event_loop().time()
tool_calls = self._plan(response)
step_metrics.plan_ms = (asyncio.get_event_loop().time() - t0) * 1000
step_metrics.tool_calls = len(tool_calls)
if not tool_calls:
# 无工具调用 = 最终答案
result = response.content or ""
# 6. Reflect - 反思(可选)
if self.enable_reflection:
t0 = asyncio.get_event_loop().time()
await self._reflect(query, result)
step_metrics.reflect_ms = (asyncio.get_event_loop().time() - t0) * 1000
self.metrics.append(step_metrics)
break
# 4. Act - 行动(执行工具)
t0 = asyncio.get_event_loop().time()
tool_results = await self._act(tool_calls)
step_metrics.act_ms = (asyncio.get_event_loop().time() - t0) * 1000
state.tool_call_count += len(tool_results)
# 5. Observe - 观察
t0 = asyncio.get_event_loop().time()
await self._observe(tool_calls, tool_results)
step_metrics.observe_ms = (asyncio.get_event_loop().time() - t0) * 1000
self.metrics.append(step_metrics)
# 保存检查点(每步自动保存)
if self.agent.state_manager:
await self.agent.state_manager.save_checkpoint()
# 触发 step_complete 事件
await self.agent._emit("step_complete", {
"step": step,
"metrics": step_metrics.__dict__,
})
else:
# 达到最大步数
result = f"达到最大步数限制 ({self.max_steps} 步),任务未完成。"
state.status = "error"
state.status = "completed"
# 执行中间件 after 链(逆序)
for mw in reversed(self.middlewares):
result = await mw.after(context, result)
logger.info(f"[AgentLoop] 执行完成: {state.step+1} 步, "
f"{state.tool_call_count} 次工具调用, "
f"{state.total_tokens} tokens")
return result
async def _perceive(self, content: str, step: int) -> None:
"""感知阶段:接收输入,更新状态,写入记忆"""
if step == 0:
# 第一步:用户输入
await self.agent.memory.add_message(
role="user", content=content
)
# 非第一步:工具返回结果已在 observe 阶段写入
await self.agent._emit("on_perceive", {"step": step})
async def _reason(self) -> "LLMResponse":
"""推理阶段:组装 Prompt,调用 LLM"""
messages = await self._build_messages()
response = await self.agent.llm.chat(
messages=messages,
tools=self.agent.get_tool_schemas(),
temperature=self.agent.config.temperature,
)
# 写入助手消息
await self.agent.memory.add_message(
role="assistant",
content=response.content or "",
tool_calls=getattr(response, "tool_calls", None),
)
await self.agent._emit("on_think", {
"thought": response.content or "",
"tool_calls": len(getattr(response, "tool_calls", []) or []),
})
return response
def _plan(self, response: "LLMResponse") -> list:
"""规划阶段:解析 LLM 输出,确定执行计划"""
tool_calls = getattr(response, "tool_calls", None) or []
return tool_calls
async def _act(self, tool_calls: list) -> list["ToolResult"]:
"""行动阶段:执行工具调用"""
if not tool_calls:
return []
await self.agent._emit("on_act", {
"tools": [tc.function.name for tc in tool_calls],
"count": len(tool_calls),
})
# 使用工具调度器执行
results = await self.agent.tool_scheduler.schedule(tool_calls)
return results
async def _observe(self, tool_calls: list, tool_results: list) -> None:
"""观察阶段:解析工具结果,写入记忆"""
for tc, tr in zip(tool_calls, tool_results):
tool_name = tc.function.name
output = tr.output if tr.success else f"Error: {tr.error}"
await self.agent.memory.add_message(
role="tool",
content=output,
name=tool_name,
tool_call_id=tc.id,
)
await self.agent._emit("on_observe", {
"results": [
{"tool": tc.function.name, "success": tr.success}
for tc, tr in zip(tool_calls, tool_results)
]
})
async def _reflect(self, query: str, result: str) -> None:
"""反思阶段:评估结果质量,总结经验"""
reflection_prompt = f"""请对以下问答过程进行反思:
原始问题: {query}
最终答案: {result}
请从以下角度评估:
1. 答案是否准确、完整地回答了问题?
2. 推理过程是否高效?有无冗余步骤?
3. 工具使用是否合理?有无更好的选择?
4. 有哪些可以改进的地方?
输出反思总结。"""
try:
reflection = await self.agent.llm.chat([
{"role": "user", "content": reflection_prompt}
])
# 写入长期记忆
await self.agent.memory.add_reflection(
query=query,
result=result,
reflection=reflection.content or "",
)
await self.agent._emit("on_reflect", {
"reflection": reflection.content,
})
except Exception as e:
logger.warning(f"反思阶段出错: {e}")
async def _build_messages(self) -> list[dict]:
"""构建发送给 LLM 的消息列表"""
messages = []
# System Prompt
if self.agent.config.system_prompt:
messages.append({
"role": "system",
"content": self.agent.config.system_prompt
})
# 记忆注入(短期 + 长期检索)
memory_context = await self.agent.memory.get_context()
if memory_context:
messages.append({
"role": "system",
"content": f"相关记忆:\n{memory_context}"
})
# 对话历史
messages.extend(await self.agent.memory.get_messages())
return messages
def get_summary_metrics(self) -> dict:
"""获取汇总性能指标"""
if not self.metrics:
return {}
total_steps = len(self.metrics)
total_perceive = sum(m.perceive_ms for m in self.metrics)
total_reason = sum(m.reason_ms for m in self.metrics)
total_plan = sum(m.plan_ms for m in self.metrics)
total_act = sum(m.act_ms for m in self.metrics)
total_observe = sum(m.observe_ms for m in self.metrics)
total_reflect = sum(m.reflect_ms for m in self.metrics)
total_tool_calls = sum(m.tool_calls for m in self.metrics)
total_tokens = sum(m.tokens_used for m in self.metrics)
total_ms = total_perceive + total_reason + total_plan + total_act + total_observe + total_reflect
return {
"total_steps": total_steps,
"total_tool_calls": total_tool_calls,
"total_tokens": total_tokens,
"total_time_ms": round(total_ms, 2),
"avg_step_time_ms": round(total_ms / total_steps, 2),
"breakdown": {
"perceive": {"ms": round(total_perceive, 2), "pct": round(total_perceive/total_ms*100, 1)},
"reason": {"ms": round(total_reason, 2), "pct": round(total_reason/total_ms*100, 1)},
"plan": {"ms": round(total_plan, 2), "pct": round(total_plan/total_ms*100, 1)},
"act": {"ms": round(total_act, 2), "pct": round(total_act/total_ms*100, 1)},
"observe": {"ms": round(total_observe, 2), "pct": round(total_observe/total_ms*100, 1)},
"reflect": {"ms": round(total_reflect, 2), "pct": round(total_reflect/total_ms*100, 1)},
}
}
1.4 中间件管道设计
中间件(Middleware)是 AgentLoop 的扩展点,允许在不修改核心循环代码的前提下,注入横切关注点(如日志、追踪、限流、缓存、安全审计等)。
中间件接口定义:
from abc import ABC, abstractmethod
from typing import Any
class IMiddleware(ABC):
"""中间件接口"""
@abstractmethod
async def before(self, context: dict[str, Any]) -> dict[str, Any]:
"""
核心逻辑执行前调用
返回修改后的 context。
若 context["_blocked"] = True,则中断执行,
并返回 context["_block_reason"] 作为结果。
"""
...
@abstractmethod
async def after(self, context: dict[str, Any], result: str) -> str:
"""
核心逻辑执行后调用(逆序执行)
返回修改后的 result。
"""
...
class LoggingMiddleware(IMiddleware):
"""日志中间件"""
async def before(self, context: dict) -> dict:
query = context.get("query", "")
logger.info(f"[Middleware] 请求开始: query={query[:50]}...")
context["_start_time"] = asyncio.get_event_loop().time()
return context
async def after(self, context: dict, result: str) -> str:
elapsed = (asyncio.get_event_loop().time() - context.get("_start_time", 0)) * 1000
logger.info(f"[Middleware] 请求完成: 耗时 {elapsed:.1f}ms, "
f"结果长度 {len(result)}")
return result
class RateLimitMiddleware(IMiddleware):
"""限流中间件"""
def __init__(self, rate_limiter: "RateLimiter"):
self.rate_limiter = rate_limiter
async def before(self, context: dict) -> dict:
try:
await asyncio.wait_for(self.rate_limiter.acquire(), timeout=5.0)
except asyncio.TimeoutError:
context["_blocked"] = True
context["_block_reason"] = "请求过于频繁,请稍后再试。"
return context
async def after(self, context: dict, result: str) -> str:
return result
class PromptCacheMiddleware(IMiddleware):
"""Prompt 缓存中间件 - 对完全相同的查询直接返回缓存结果"""
def __init__(self, cache: "Cache", ttl: int = 3600):
self.cache = cache
self.ttl = ttl
async def before(self, context: dict) -> dict:
query = context.get("query", "")
cache_key = f"prompt_cache:{hash(query)}"
cached = await self.cache.get(cache_key)
if cached:
context["_blocked"] = True
context["_block_reason"] = cached
context["_from_cache"] = True
else:
context["_cache_key"] = cache_key
return context
async def after(self, context: dict, result: str) -> str:
cache_key = context.get("_cache_key")
if cache_key and not context.get("_from_cache"):
await self.cache.set(cache_key, result, ttl=self.ttl)
return result
1.5 量化分析与性能瓶颈
| 阶段 | 单步耗时范围 | 占比 | 优化潜力 | 主要优化手段 |
|---|---|---|---|---|
| Perceive(感知) | < 1 ms | ~0% | 极低 | 无需优化 |
| Reason(推理) | 500 – 5,000 ms | 70 – 85% | 极高 | Prompt 缓存、模型路由、流式输出 |
| Plan(规划) | < 5 ms | ~0% | 极低 | 内置于 Reason,无需单独优化 |
| Act(行动) | 100 – 10,000 ms | 10 – 25% | 中 | 并行执行、超时控制、连接池 |
| Observe(观察) | < 2 ms | ~0% | 极低 | 无需优化 |
| Reflect(反思) | 200 – 1,000 ms | 5 – 10% | 中 | 可选开关、异步后台执行 |
| 单步总计 | 800 – 16,000 ms | 100% | - | LLM 调用是核心瓶颈 |
关键洞察:
- LLM 调用占绝对主导:Reason 阶段消耗了 70-85% 的时间,优化 LLM 调用是性能提升的核心
- 工具执行波动大:Act 阶段耗时从 100ms 到 10s 不等,取决于工具类型(本地函数 vs 远程 API vs 代码执行)
- 反思是可选开销:Reflect 增加 5-10% 开销,建议仅在需要学习/迭代的场景开启
- 感知/规划/观察几乎可忽略:这三个阶段加起来不到 1%,优化收益极低
第 2 章 消息传递机理:事件驱动的组件通信
2.1 为什么需要消息传递
在一个复杂的 Agent 框架中,组件之间的通信方式决定了系统的可扩展性、可观测性和解耦程度。如果每个组件都直接调用其他组件的方法,系统会变成一张紧密耦合的网——改动一个组件可能牵一发而动全身。
消息传递机制通过事件总线(EventBus) 实现了发布/订阅模式,让组件之间通过事件而非直接调用来通信:
- 发布者(Publisher):只负责发布事件,不关心谁来处理
- 订阅者(Subscriber):只关心自己感兴趣的事件,不关心谁发布的
- 事件总线(EventBus):负责事件的路由和分发
2.2 三种通信模式
自研框架采用混合通信模式,不同场景使用不同的通信方式:
| 通信模式 | 延迟 | 吞吐量 | 适用场景 | 实现方式 |
|---|---|---|---|---|
| 同步调用 | 0.1 – 1 ms | 低 | LLM、Memory、Tool 等核心路径的直接调用 | 直接方法调用 |
| 异步事件 | 1 – 10 ms | 高 | 追踪、日志、指标收集等横切关注点 | EventBus 发布/订阅 |
| 回调链 | 0.5 – 5 ms | 中 | 中间件管道、请求过滤器 | Middleware Chain |
| 发布/订阅 | 1 – 10 ms | 高 | 多组件并行通知 | EventBus + Handler 列表 |
2.3 事件总线架构
2.4 核心实现:EventBus
"""
消息传递机理核心实现 - EventBus
"""
from __future__ import annotations
import asyncio
import logging
from collections import defaultdict
from typing import Any, Callable, Optional
logger = logging.getLogger(__name__)
EventHandler = Callable[[dict[str, Any]], Any]
class EventBus:
"""
事件总线 - 发布/订阅模式
支持同步发布(阻塞主流程)和异步发布(不阻塞主流程)两种模式。
- 同步发布:用于需要确保处理完成后再继续的场景
- 异步发布:用于日志、指标等不影响主流程的场景
"""
def __init__(self):
self._handlers: dict[str, list[EventHandler]] = defaultdict(list)
self._async_queue: Optional[asyncio.Queue] = None
self._async_task: Optional[asyncio.Task] = None
self._running: bool = False
self._wildcard_handlers: list[EventHandler] = []
def subscribe(self, event: str, handler: EventHandler) -> None:
"""
订阅指定事件
Args:
event: 事件名称,使用 "*" 订阅所有事件
handler: 事件处理函数
"""
if event == "*":
self._wildcard_handlers.append(handler)
else:
self._handlers[event].append(handler)
def unsubscribe(self, event: str, handler: EventHandler) -> None:
"""取消订阅"""
if event == "*":
self._wildcard_handlers.remove(handler)
else:
if handler in self._handlers.get(event, []):
self._handlers[event].remove(handler)
async def publish(self, event: str, data: dict[str, Any]) -> None:
"""
同步发布事件
按顺序调用所有订阅者,阻塞直到全部处理完成。
单个 handler 异常不影响其他 handler。
Args:
event: 事件名称
data: 事件数据
"""
# 特定事件的 handler
for handler in self._handlers.get(event, []):
try:
if asyncio.iscoroutinefunction(handler):
await handler(data)
else:
handler(data)
except Exception as e:
logger.error(f"[EventBus] 同步 handler 出错 "
f"(event={event}): {e}", exc_info=True)
# 通配符 handler
for handler in self._wildcard_handlers:
try:
if asyncio.iscoroutinefunction(handler):
await handler({"event": event, **data})
else:
handler({"event": event, **data})
except Exception as e:
logger.error(f"[EventBus] 通配符 handler 出错 "
f"(event={event}): {e}", exc_info=True)
async def publish_async(self, event: str, data: dict[str, Any]) -> None:
"""
异步发布事件(不阻塞主流程)
事件放入队列,由后台协程处理。
适用于日志、指标等不影响主流程的场景。
Args:
event: 事件名称
data: 事件数据
"""
if not self._running:
await self.start_async_processor()
if self._async_queue is not None:
try:
self._async_queue.put_nowait((event, data))
except asyncio.QueueFull:
logger.warning(f"[EventBus] 异步队列已满,丢弃事件: {event}")
async def start_async_processor(self, max_queue_size: int = 1000) -> None:
"""启动异步事件处理器"""
if self._running:
return
self._async_queue = asyncio.Queue(maxsize=max_queue_size)
self._running = True
self._async_task = asyncio.create_task(self._async_worker())
logger.info("[EventBus] 异步事件处理器已启动")
async def stop(self) -> None:
"""停止事件总线"""
self._running = False
if self._async_task:
self._async_task.cancel()
try:
await self._async_task
except asyncio.CancelledError:
pass
logger.info("[EventBus] 事件总线已停止")
async def _async_worker(self) -> None:
"""异步事件处理协程"""
while self._running:
try:
event, data = await self._async_queue.get()
# 处理特定事件 handler
for handler in self._handlers.get(event, []):
try:
if asyncio.iscoroutinefunction(handler):
await handler(data)
else:
handler(data)
except Exception as e:
logger.error(f"[EventBus] 异步 handler 出错 "
f"(event={event}): {e}", exc_info=True)
# 处理通配符 handler
for handler in self._wildcard_handlers:
try:
if asyncio.iscoroutinefunction(handler):
await handler({"event": event, **data})
else:
handler({"event": event, **data})
except Exception as e:
logger.error(f"[EventBus] 异步通配符 handler 出错 "
f"(event={event}): {e}", exc_info=True)
self._async_queue.task_done()
except asyncio.CancelledError:
break
except Exception as e:
logger.error(f"[EventBus] 异步 worker 异常: {e}", exc_info=True)
def get_handler_count(self, event: Optional[str] = None) -> int:
"""获取指定事件的订阅者数量"""
if event is None:
return sum(len(hs) for hs in self._handlers.values()) + len(self._wildcard_handlers)
return len(self._handlers.get(event, []))
def list_events(self) -> list[str]:
"""列出所有有订阅者的事件"""
return list(self._handlers.keys())
2.5 标准事件列表
框架定义了以下标准事件,所有核心组件都会在关键节点发布:
| 事件名称 | 触发时机 | 数据字段 | 建议模式 |
|---|---|---|---|
on_perceive | 感知阶段开始 | step, query | 异步 |
on_think | LLM 返回推理结果 | step, thought, tool_calls | 异步 |
on_act | 工具执行前 | tools, count | 异步 |
on_observe | 工具结果返回后 | results | 异步 |
step_complete | 单步执行完成 | step, metrics | 异步 |
on_reflect | 反思完成后 | reflection | 异步 |
on_error | 发生错误时 | error, step, context | 同步 |
on_complete | 整个任务完成 | result, metrics | 异步 |
tool_registered | 新工具注册时 | tool_name, tool_schema | 异步 |
memory_stored | 记忆写入后 | memory_id, memory_type | 异步 |
checkpoint_saved | 检查点保存后 | checkpoint_id, step | 异步 |
2.6 典型订阅者实现
"""
典型事件订阅者实现
"""
import time
from typing import Any
class LoggingHandler:
"""日志订阅者 - 将关键事件输出到日志"""
def __init__(self, level: str = "INFO"):
self.level = level
async def __call__(self, data: dict[str, Any]) -> None:
event = data.get("event", "unknown")
step = data.get("step", "?")
if event == "on_think":
thought = data.get("thought", "")[:100]
tool_calls = data.get("tool_calls", 0)
logger.info(f"[Step {step}] 思考: {thought}... "
f"(工具调用: {tool_calls})")
elif event == "on_act":
tools = data.get("tools", [])
logger.info(f"[Step {step}] 执行工具: {tools}")
elif event == "on_error":
error = data.get("error", "")
logger.error(f"[Step {step}] 错误: {error}")
class MetricsHandler:
"""指标订阅者 - 收集性能指标"""
def __init__(self):
self.step_times: list[float] = []
self.tool_call_count: int = 0
self.total_tokens: int = 0
async def __call__(self, data: dict[str, Any]) -> None:
event = data.get("event", "")
if event == "step_complete":
metrics = data.get("metrics", {})
total_ms = sum(metrics.get(k, 0) for k in
["perceive_ms", "reason_ms", "plan_ms",
"act_ms", "observe_ms", "reflect_ms"])
self.step_times.append(total_ms)
self.tool_call_count += metrics.get("tool_calls", 0)
self.total_tokens += metrics.get("tokens_used", 0)
def get_summary(self) -> dict:
return {
"total_steps": len(self.step_times),
"total_tool_calls": self.tool_call_count,
"total_tokens": self.total_tokens,
"avg_step_time_ms": (sum(self.step_times) / len(self.step_times)
if self.step_times else 0),
}
class TracingHandler:
"""追踪订阅者 - 生成完整执行轨迹,用于调试"""
def __init__(self):
self.trace: list[dict] = []
async def __call__(self, data: dict[str, Any]) -> None:
event = data.get("event", "")
self.trace.append({
"timestamp": time.time(),
"event": event,
"data": {k: v for k, v in data.items() if k != "event"},
})
def get_trace(self) -> list[dict]:
"""获取完整追踪记录"""
return self.trace
def print_trace(self) -> None:
"""打印追踪记录(调试用)"""
for entry in self.trace:
ts = entry["timestamp"]
event = entry["event"]
data_str = str(entry["data"])[:80]
print(f"[{ts:.3f}] {event:20s} {data_str}")
class WebhookNotifier:
"""Webhook 通知订阅者 - 将事件推送到外部 Webhook"""
def __init__(self, webhook_url: str, events: Optional[list[str]] = None):
self.webhook_url = webhook_url
self.events = events # None = 所有事件
async def __call__(self, data: dict[str, Any]) -> None:
event = data.get("event", "")
if self.events and event not in self.events:
return
# 异步 HTTP 调用(不阻塞)
asyncio.create_task(self._send_webhook(event, data))
async def _send_webhook(self, event: str, data: dict) -> None:
try:
import aiohttp
async with aiohttp.ClientSession() as session:
async with session.post(self.webhook_url, json={
"event": event,
"data": data,
}) as resp:
if resp.status != 200:
logger.warning(f"Webhook 返回 {resp.status}: {event}")
except Exception as e:
logger.warning(f"Webhook 发送失败 ({event}): {e}")
2.7 回调链(中间件管道)
回调链(Callback Chain)是另一种消息传递模式,采用管道-过滤器架构,将多个处理步骤串联成一条链。每个中间件都可以:
- 修改输入(before)
- 修改输出(after)
- 中断流程(返回 _blocked)
"""
回调链实现 - CallbackChain
"""
from typing import Any
class CallbackChain:
"""
回调链 - 中间件管道模式
before 链正序执行,after 链逆序执行(洋葱模型)。
任一中间件可通过设置 _blocked 中断执行。
"""
def __init__(self, middlewares: list[IMiddleware]):
self.middlewares = middlewares
async def execute(self, context: dict[str, Any],
core_fn: Callable) -> Any:
"""
执行回调链
Args:
context: 上下文(会在中间件间传递)
core_fn: 核心逻辑函数
Returns:
最终结果
"""
# before 链(正序)
for mw in self.middlewares:
context = await mw.before(context)
if context.get("_blocked"):
return context.get("_block_reason", "")
# 核心逻辑
result = await core_fn(context)
# after 链(逆序)
for mw in reversed(self.middlewares):
result = await mw.after(context, result)
return result
第 3 章 工具调度机理:从发现到执行的五步链路
3.1 工具调度的挑战
工具调用(Tool Calling)是 Agent 与外部世界交互的桥梁。看似简单的"调用一个函数"背后,隐藏着一系列工程挑战:
- 工具发现:LLM 输出的工具名可能不存在、拼写错误
- 参数校验:LLM 生成的参数可能格式错误、缺少必填字段、类型不匹配
- 执行安全:工具可能执行危险操作、超时、抛出异常
- 结果处理:工具输出可能超长、格式不规范、需要二次解析
- 错误恢复:工具失败后如何重试、降级、告知 LLM 调整策略
自研框架的工具调度器(ToolScheduler)用一条五步链路系统性地解决这些问题。
3.2 五步链路全景
3.3 核心实现:ToolScheduler
"""
工具调度机理核心实现 - ToolScheduler
"""
from __future__ import annotations
import asyncio
import json
import logging
from dataclasses import dataclass, field
from typing import Any, Optional
logger = logging.getLogger(__name__)
@dataclass
class ToolResult:
"""工具执行结果"""
success: bool
output: str = ""
error: Optional[str] = None
execution_time_ms: float = 0.0
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass
class RetryPolicy:
"""重试策略"""
max_retries: int = 3
base_delay: float = 1.0 # 初始延迟(秒)
backoff_factor: float = 2.0 # 指数退避因子
retry_on_exceptions: tuple[type[Exception], ...] = (
TimeoutError,
ConnectionError,
RuntimeError,
)
class ToolScheduler:
"""
工具调度器
五步完整链路:发现 → 校验 → 调用 → 解析 → 错误处理
支持并行执行、超时控制、指数退避重试、结果截断。
"""
def __init__(
self,
registry: "ToolRegistry",
executor: "IExecutor",
retry_policy: Optional[RetryPolicy] = None,
max_output_length: int = 10000,
default_timeout: float = 30.0,
):
self.registry = registry
self.executor = executor
self.retry_policy = retry_policy or RetryPolicy()
self.max_output_length = max_output_length
self.default_timeout = default_timeout
async def schedule(self, tool_calls: list[Any]) -> list[ToolResult]:
"""
调度多个工具调用
单个工具串行执行,多个工具判断是否可并行。
所有工具调用始终返回 ToolResult,不会抛出异常。
Args:
tool_calls: LLM 返回的工具调用列表
Returns:
工具执行结果列表(与输入顺序一致)
"""
if not tool_calls:
return []
if len(tool_calls) == 1:
# 单个工具:直接执行
return [await self._execute_single(tool_calls[0])]
# 多个工具:并行执行
tasks = [self._execute_single(tc) for tc in tool_calls]
results = await asyncio.gather(*tasks, return_exceptions=True)
# 确保全部返回 ToolResult
return [
r if isinstance(r, ToolResult)
else ToolResult(success=False, output="", error=str(r))
for r in results
]
async def _execute_single(self, tool_call: Any) -> ToolResult:
"""
执行单个工具调用(完整五步链路 + 重试)
Args:
tool_call: 单个工具调用对象
Returns:
ToolResult
"""
tool_name = tool_call.function.name
start_time = asyncio.get_event_loop().time()
# Step 1: 工具发现
tool = self.registry.get(tool_name)
if not tool:
return ToolResult(
success=False,
output="",
error=f"工具 '{tool_name}' 不存在。可用工具: {self.registry.list_names()}",
execution_time_ms=(asyncio.get_event_loop().time() - start_time) * 1000,
)
# Step 2: 参数校验
try:
args = json.loads(tool_call.function.arguments)
except json.JSONDecodeError as e:
return ToolResult(
success=False,
output="",
error=f"参数 JSON 解析失败: {e}",
execution_time_ms=(asyncio.get_event_loop().time() - start_time) * 1000,
)
# 校验参数 schema
validation_error = self._validate_args(tool, args)
if validation_error:
return ToolResult(
success=False,
output="",
error=f"参数校验失败: {validation_error}",
execution_time_ms=(asyncio.get_event_loop().time() - start_time) * 1000,
)
# Step 3+5: 调用执行(带重试)
last_result: Optional[ToolResult] = None
for attempt in range(self.retry_policy.max_retries):
result = await self._execute_tool(tool, args)
last_result = result
if result.success:
# Step 4: 结果解析(成功时)
result = self._parse_result(result)
result.execution_time_ms = (
asyncio.get_event_loop().time() - start_time
) * 1000
return result
# 判断是否可重试
if not self._is_retryable(result.error):
break
if attempt < self.retry_policy.max_retries - 1:
delay = self.retry_policy.base_delay * (
self.retry_policy.backoff_factor ** attempt
)
logger.info(f"[ToolScheduler] 工具 '{tool_name}' 第 {attempt+1} 次失败,"
f"{delay:.1f}s 后重试...")
await asyncio.sleep(delay)
# 所有重试都失败
if last_result:
last_result.execution_time_ms = (
asyncio.get_event_loop().time() - start_time
) * 1000
return last_result
return ToolResult(
success=False,
output="",
error="未知错误",
execution_time_ms=(asyncio.get_event_loop().time() - start_time) * 1000,
)
async def _execute_tool(self, tool: "ITool", args: dict) -> ToolResult:
"""实际执行工具(单次)"""
timeout = getattr(tool, "timeout", self.default_timeout)
try:
result = await asyncio.wait_for(
self.executor.execute_tool(tool, args),
timeout=timeout,
)
return result
except asyncio.TimeoutError:
return ToolResult(
success=False,
output="",
error=f"工具执行超时({timeout}s)",
)
except Exception as e:
return ToolResult(
success=False,
output="",
error=f"工具执行异常: {type(e).__name__}: {e}",
)
def _validate_args(self, tool: "ITool", args: dict) -> Optional[str]:
"""
参数校验
Returns:
错误信息(None 表示校验通过)
"""
schema = getattr(tool, "parameters", {})
if not schema:
return None # 无 schema,跳过校验
required = schema.get("required", [])
properties = schema.get("properties", {})
# 检查必填字段
for field_name in required:
if field_name not in args:
return f"缺少必填字段: {field_name}"
# 类型检查(简化版)
for field_name, value in args.items():
if field_name not in properties:
continue # 额外字段不报错,LLM 可能传多了
prop = properties[field_name]
expected_type = prop.get("type", "")
if expected_type == "string" and not isinstance(value, str):
return f"字段 '{field_name}' 类型错误: 期望 string, 实际 {type(value).__name__}"
elif expected_type == "integer" and not isinstance(value, int):
return f"字段 '{field_name}' 类型错误: 期望 integer, 实际 {type(value).__name__}"
elif expected_type == "number" and not isinstance(value, (int, float)):
return f"字段 '{field_name}' 类型错误: 期望 number, 实际 {type(value).__name__}"
elif expected_type == "boolean" and not isinstance(value, bool):
return f"字段 '{field_name}' 类型错误: 期望 boolean, 实际 {type(value).__name__}"
elif expected_type == "array" and not isinstance(value, list):
return f"字段 '{field_name}' 类型错误: 期望 array, 实际 {type(value).__name__}"
elif expected_type == "object" and not isinstance(value, dict):
return f"字段 '{field_name}' 类型错误: 期望 object, 实际 {type(value).__name__}"
return None
def _parse_result(self, result: ToolResult) -> ToolResult:
"""结果解析:截断、标准化"""
# 输出截断(防止超长输出撑爆 context)
if len(result.output) > self.max_output_length:
truncated = result.output[:self.max_output_length]
result.output = truncated + f"\n... [输出已截断,原长度 {len(result.output)} 字符]"
result.metadata["truncated"] = True
result.metadata["original_length"] = len(result.output)
return result
def _is_retryable(self, error: Optional[str]) -> bool:
"""判断错误是否可重试"""
if not error:
return False
retryable_keywords = [
"超时", "timeout",
"连接", "connection",
"网络", "network",
"500", "502", "503", "504",
"临时", "temporary",
"限流", "rate limit",
]
error_lower = error.lower()
return any(kw.lower() in error_lower for kw in retryable_keywords)
3.4 工具注册表(ToolRegistry)
工具注册表是工具发现的基础设施,负责工具的注册、查找和元数据管理。
"""
工具注册表实现
"""
from __future__ import annotations
import logging
from typing import Any, Optional
logger = logging.getLogger(__name__)
class ToolRegistry:
"""
工具注册表
负责工具的注册、注销、查找和元数据管理。
支持按名称查找、按标签过滤、动态注册/注销。
"""
def __init__(self):
self._tools: dict[str, "ITool"] = {}
self._categories: dict[str, list[str]] = {}
def register(self, tool: "ITool") -> None:
"""注册工具"""
name = tool.name
if name in self._tools:
logger.warning(f"[ToolRegistry] 工具 '{name}' 已存在,将被覆盖")
self._tools[name] = tool
# 按分类索引
category = getattr(tool, "category", "general")
if category not in self._categories:
self._categories[category] = []
if name not in self._categories[category]:
self._categories[category].append(name)
logger.debug(f"[ToolRegistry] 已注册工具: {name}")
def unregister(self, name: str) -> bool:
"""注销工具"""
if name not in self._tools:
return False
tool = self._tools.pop(name)
category = getattr(tool, "category", "general")
if category in self._categories and name in self._categories[category]:
self._categories[category].remove(name)
logger.debug(f"[ToolRegistry] 已注销工具: {name}")
return True
def get(self, name: str) -> Optional["ITool"]:
"""按名称获取工具"""
return self._tools.get(name)
def list_names(self) -> list[str]:
"""列出所有工具名称"""
return list(self._tools.keys())
def get_schemas(self) -> list[dict]:
"""获取所有工具的 OpenAI function calling schema"""
return [tool.get_schema() for tool in self._tools.values()]
def get_by_category(self, category: str) -> list["ITool"]:
"""按分类获取工具列表"""
names = self._categories.get(category, [])
return [self._tools[name] for name in names if name in self._tools]
def search(self, query: str) -> list["ITool"]:
"""
搜索工具(按名称和描述模糊匹配)
用于工具选择阶段,帮助 LLM 缩小工具范围。
"""
query_lower = query.lower()
results = []
for tool in self._tools.values():
name_match = query_lower in tool.name.lower()
desc = getattr(tool, "description", "")
desc_match = query_lower in desc.lower()
if name_match or desc_match:
results.append(tool)
return results
def __len__(self) -> int:
return len(self._tools)
def __contains__(self, name: str) -> bool:
return name in self._tools
3.5 执行器(Executor)
执行器负责实际调用工具,是工具调度的执行层。支持多种执行模式:
"""
执行器接口与实现
"""
from __future__ import annotations
import asyncio
import logging
from abc import ABC, abstractmethod
from typing import Any
logger = logging.getLogger(__name__)
class IExecutor(ABC):
"""执行器接口"""
@abstractmethod
async def execute_tool(self, tool: "ITool", args: dict[str, Any]) -> ToolResult:
"""执行工具"""
...
class LocalExecutor(IExecutor):
"""
本地函数执行器
直接调用 Python 函数,适用于本地工具。
"""
async def execute_tool(self, tool: "ITool", args: dict[str, Any]) -> ToolResult:
start = asyncio.get_event_loop().time()
try:
# 支持同步和异步工具
if asyncio.iscoroutinefunction(tool.execute):
output = await tool.execute(**args)
else:
output = tool.execute(**args)
elapsed = (asyncio.get_event_loop().time() - start) * 1000
if isinstance(output, ToolResult):
output.execution_time_ms = elapsed
return output
return ToolResult(
success=True,
output=str(output),
execution_time_ms=elapsed,
)
except Exception as e:
elapsed = (asyncio.get_event_loop().time() - start) * 1000
return ToolResult(
success=False,
output="",
error=f"{type(e).__name__}: {e}",
execution_time_ms=elapsed,
)
class RemoteExecutor(IExecutor):
"""
远程 API 执行器
通过 HTTP 调用远程工具,适用于外部 API。
"""
def __init__(self, base_url: str = "", timeout: float = 30.0):
self.base_url = base_url
self.timeout = timeout
async def execute_tool(self, tool: "ITool", args: dict[str, Any]) -> ToolResult:
import aiohttp
start = asyncio.get_event_loop().time()
url = getattr(tool, "endpoint", "")
if not url.startswith("http"):
url = self.base_url + url
try:
async with aiohttp.ClientSession() as session:
async with session.post(
url,
json=args,
timeout=aiohttp.ClientTimeout(total=self.timeout),
) as resp:
data = await resp.json()
elapsed = (asyncio.get_event_loop().time() - start) * 1000
if resp.status == 200:
return ToolResult(
success=True,
output=str(data.get("result", data)),
execution_time_ms=elapsed,
metadata={"status_code": resp.status},
)
else:
return ToolResult(
success=False,
output="",
error=f"HTTP {resp.status}: {data.get('error', data)}",
execution_time_ms=elapsed,
)
except Exception as e:
elapsed = (asyncio.get_event_loop().time() - start) * 1000
return ToolResult(
success=False,
output="",
error=f"远程调用失败: {type(e).__name__}: {e}",
execution_time_ms=elapsed,
)
class SandboxExecutor(IExecutor):
"""
沙箱代码执行器
在隔离环境中执行代码,适用于代码解释器类工具。
使用 subprocess 隔离,支持超时控制。
"""
def __init__(self, timeout: float = 30.0, memory_limit_mb: int = 512):
self.timeout = timeout
self.memory_limit_mb = memory_limit_mb
async def execute_tool(self, tool: "ITool", args: dict[str, Any]) -> ToolResult:
code = args.get("code", "")
language = args.get("language", "python")
start = asyncio.get_event_loop().time()
try:
if language == "python":
result = await self._execute_python(code)
elif language == "bash":
result = await self._execute_bash(code)
else:
return ToolResult(
success=False,
output="",
error=f"不支持的语言: {language}",
)
elapsed = (asyncio.get_event_loop().time() - start) * 1000
result.execution_time_ms = elapsed
return result
except asyncio.TimeoutError:
elapsed = (asyncio.get_event_loop().time() - start) * 1000
return ToolResult(
success=False,
output="",
error=f"代码执行超时({self.timeout}s)",
execution_time_ms=elapsed,
)
except Exception as e:
elapsed = (asyncio.get_event_loop().time() - start) * 1000
return ToolResult(
success=False,
output="",
error=f"执行异常: {type(e).__name__}: {e}",
execution_time_ms=elapsed,
)
async def _execute_python(self, code: str) -> ToolResult:
"""执行 Python 代码"""
proc = await asyncio.create_subprocess_exec(
"python3", "-c", code,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
)
try:
stdout, stderr = await asyncio.wait_for(
proc.communicate(),
timeout=self.timeout,
)
except asyncio.TimeoutError:
proc.kill()
await proc.wait()
raise
if proc.returncode == 0:
return ToolResult(
success=True,
output=stdout.decode("utf-8", errors="replace"),
)
else:
return ToolResult(
success=False,
output=stdout.decode("utf-8", errors="replace"),
error=stderr.decode("utf-8", errors="replace"),
)
async def _execute_bash(self, code: str) -> ToolResult:
"""执行 Bash 命令"""
proc = await asyncio.create_subprocess_exec(
"bash", "-c", code,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
)
try:
stdout, stderr = await asyncio.wait_for(
proc.communicate(),
timeout=self.timeout,
)
except asyncio.TimeoutError:
proc.kill()
await proc.wait()
raise
if proc.returncode == 0:
return ToolResult(
success=True,
output=stdout.decode("utf-8", errors="replace"),
)
else:
return ToolResult(
success=False,
output=stdout.decode("utf-8", errors="replace"),
error=stderr.decode("utf-8", errors="replace"),
)
3.6 量化分析
| 步骤 | 平均耗时 | 失败率 | 重试策略 | 优化手段 |
|---|---|---|---|---|
| 工具发现 | < 1 ms | 0.1% | 直接报错 | 哈希表 O(1) 查找 |
| 参数校验 | 1 – 5 ms | 2 – 5% | 返回 LLM 修正 | Schema 缓存、预编译 |
| 工具调用 | 100 – 10,000 ms | 5 – 15% | 指数退避 × 3 | 连接池、并行执行、超时控制 |
| 结果解析 | < 2 ms | < 1% | 截断处理 | 流式解析、增量截断 |
| 错误处理 | < 5 ms | - | 返回错误信息 | 错误分类、重试决策 |
关键洞察:
- 工具调用是第二大耗时源:仅次于 LLM 调用,且波动极大(100ms → 10s)
- 参数校验失败率不可忽视:2-5% 的失败率意味着 LLM 经常生成错误参数,需要良好的错误反馈机制
- 重试策略要精准:不是所有错误都该重试,网络/超时类错误重试才有意义
- 结果截断是必要的:不加截断的工具输出可能轻易撑爆 LLM context window
第 4 章 状态管理机理:检查点快照与时间旅行
4.1 为什么需要状态管理
Agent 的执行过程是有状态的——每一步的推理、工具调用、记忆写入都会改变 Agent 的状态。如果没有良好的状态管理机制,会面临以下问题:
- 故障恢复:执行中途崩溃后,无法从断点继续,只能从头开始
- 调试困难:出了问题无法回溯到出错前的状态进行分析
- 分支探索:无法对同一问题尝试不同策略并对比结果
- 长任务中断:长时间运行的任务无法暂停和恢复
自研框架的状态管理机制通过检查点快照(Checkpoint) 解决这些问题,支持状态保存、恢复、历史回溯和时间旅行调试。
4.2 状态生命周期
状态说明:
| 状态 | 含义 | 可转换到 |
|---|---|---|
| Created | 刚创建,尚未初始化 | Initialized |
| Initialized | 配置加载完成,等待执行 | Running |
| Running | 正在执行主循环 | Paused, Completed, Error |
| Paused | 已暂停,状态已持久化 | Resumed |
| Resumed | 从暂停恢复,即将继续 | Running |
| Completed | 任务正常完成 | -(终态) |
| Error | 执行出错 | Running(重试) |
4.3 检查点机制
检查点是状态管理的核心——每执行一步,自动保存当前状态的快照。
检查点存储内容:
| 数据项 | 说明 | 大小估算 |
|---|---|---|
| session_id | 会话唯一标识 | 36 bytes (UUID) |
| step | 当前步数 | 4 bytes |
| messages | 对话历史列表 | 2 – 50 KB |
| tool_call_count | 工具调用次数 | 4 bytes |
| total_tokens | 累计 token 数 | 4 bytes |
| memory_state | 记忆系统状态 | 1 – 100 KB(可选) |
| metadata | 自定义元数据 | 可变 |
| timestamp | 保存时间戳 | 8 bytes |
| 总计 | - | ~3 – 150 KB/步 |
4.4 核心实现:StateManager
"""
状态管理机理核心实现 - StateManager
"""
from __future__ import annotations
import json
import logging
import pickle
from abc import ABC, abstractmethod
from dataclasses import dataclass, field, asdict
from datetime import datetime
from typing import Any, Optional
from uuid import uuid4
logger = logging.getLogger(__name__)
@dataclass
class Checkpoint:
"""检查点数据结构"""
checkpoint_id: str
session_id: str
step: int
state_data: dict[str, Any]
timestamp: datetime = field(default_factory=datetime.now)
version: str = "1.0"
def to_dict(self) -> dict:
return {
"checkpoint_id": self.checkpoint_id,
"session_id": self.session_id,
"step": self.step,
"state_data": self.state_data,
"timestamp": self.timestamp.isoformat(),
"version": self.version,
}
@classmethod
def from_dict(cls, data: dict) -> "Checkpoint":
return cls(
checkpoint_id=data["checkpoint_id"],
session_id=data["session_id"],
step=data["step"],
state_data=data["state_data"],
timestamp=datetime.fromisoformat(data["timestamp"]),
version=data.get("version", "1.0"),
)
class CheckpointStore(ABC):
"""检查点存储接口"""
@abstractmethod
async def save(self, checkpoint: Checkpoint) -> str:
"""保存检查点,返回 checkpoint_id"""
...
@abstractmethod
async def load(self, session_id: str,
checkpoint_id: Optional[str] = None) -> Optional[Checkpoint]:
"""加载检查点(不传 checkpoint_id 则加载最新的)"""
...
@abstractmethod
async def list_checkpoints(self, session_id: str) -> list[dict]:
"""列出会话的所有检查点元数据"""
...
@abstractmethod
async def delete(self, session_id: str, checkpoint_id: str) -> bool:
"""删除指定检查点"""
...
@abstractmethod
async def clear(self, session_id: str) -> int:
"""清除会话的所有检查点,返回删除数量"""
...
class MemoryCheckpointStore(CheckpointStore):
"""
内存检查点存储
适用于开发测试,进程退出后数据丢失。
"""
def __init__(self, max_checkpoints_per_session: int = 100):
self._store: dict[str, list[Checkpoint]] = {}
self.max_checkpoints = max_checkpoints_per_session
async def save(self, checkpoint: Checkpoint) -> str:
session_id = checkpoint.session_id
if session_id not in self._store:
self._store[session_id] = []
self._store[session_id].append(checkpoint)
# 限制数量
if len(self._store[session_id]) > self.max_checkpoints:
self._store[session_id] = self._store[session_id][-self.max_checkpoints:]
return checkpoint.checkpoint_id
async def load(self, session_id: str,
checkpoint_id: Optional[str] = None) -> Optional[Checkpoint]:
checkpoints = self._store.get(session_id, [])
if not checkpoints:
return None
if checkpoint_id:
for cp in checkpoints:
if cp.checkpoint_id == checkpoint_id:
return cp
return None
else:
return checkpoints[-1] # 最新的
async def list_checkpoints(self, session_id: str) -> list[dict]:
checkpoints = self._store.get(session_id, [])
return [
{
"checkpoint_id": cp.checkpoint_id,
"step": cp.step,
"timestamp": cp.timestamp.isoformat(),
}
for cp in checkpoints
]
async def delete(self, session_id: str, checkpoint_id: str) -> bool:
checkpoints = self._store.get(session_id, [])
for i, cp in enumerate(checkpoints):
if cp.checkpoint_id == checkpoint_id:
checkpoints.pop(i)
return True
return False
async def clear(self, session_id: str) -> int:
count = len(self._store.get(session_id, []))
if session_id in self._store:
del self._store[session_id]
return count
class StateManager:
"""
状态管理器
负责状态的保存、恢复、历史查询和迁移。
每步自动保存检查点,支持从任意检查点恢复。
"""
def __init__(self, store: CheckpointStore,
auto_save: bool = True,
max_checkpoints: int = 50):
self.store = store
self.auto_save = auto_save
self.max_checkpoints = max_checkpoints
self._current_state: Optional[dict] = None
self._session_id: Optional[str] = None
def init_session(self, session_id: Optional[str] = None,
initial_state: Optional[dict] = None) -> str:
"""
初始化新会话
Args:
session_id: 会话 ID(不传则自动生成)
initial_state: 初始状态
Returns:
session_id
"""
self._session_id = session_id or str(uuid4())
self._current_state = initial_state or {
"session_id": self._session_id,
"step": 0,
"messages": [],
"tool_call_count": 0,
"total_tokens": 0,
"status": "initialized",
"metadata": {},
}
return self._session_id
@property
def current_state(self) -> Optional[dict]:
"""获取当前状态"""
return self._current_state
@property
def session_id(self) -> Optional[str]:
"""获取当前会话 ID"""
return self._session_id
async def save_checkpoint(self, step: Optional[int] = None) -> str:
"""
保存当前状态为检查点
Args:
step: 当前步数(不传则从状态中取)
Returns:
checkpoint_id
"""
if not self._current_state or not self._session_id:
raise ValueError("无活动会话,请先调用 init_session()")
if step is not None:
self._current_state["step"] = step
checkpoint = Checkpoint(
checkpoint_id=str(uuid4()),
session_id=self._session_id,
step=self._current_state.get("step", 0),
state_data=self._serialize_state(self._current_state),
)
cp_id = await self.store.save(checkpoint)
logger.debug(f"[StateManager] 已保存检查点: {cp_id} (step {checkpoint.step})")
return cp_id
async def restore(self, session_id: str,
checkpoint_id: Optional[str] = None) -> dict:
"""
从检查点恢复状态
Args:
session_id: 会话 ID
checkpoint_id: 检查点 ID(不传则恢复最新的)
Returns:
恢复后的状态
"""
checkpoint = await self.store.load(session_id, checkpoint_id)
if not checkpoint:
raise ValueError(
f"检查点不存在: session={session_id}, "
f"checkpoint={checkpoint_id or 'latest'}"
)
self._session_id = session_id
self._current_state = self._deserialize_state(checkpoint.state_data)
self._current_state["status"] = "paused"
logger.info(f"[StateManager] 已从检查点恢复: {checkpoint.checkpoint_id} "
f"(step {checkpoint.step})")
return self._current_state
async def list_history(self, session_id: Optional[str] = None) -> list[dict]:
"""
列出所有检查点(用于时间旅行)
Args:
session_id: 会话 ID(不传则使用当前会话)
Returns:
检查点列表(按时间正序)
"""
sid = session_id or self._session_id
if not sid:
return []
return await self.store.list_checkpoints(sid)
async def time_travel(self, checkpoint_id: str) -> dict:
"""
时间旅行:跳转到指定检查点
与 restore 不同的是,time_travel 会保留后续检查点,
方便在不同时间点间跳转。
Args:
checkpoint_id: 目标检查点 ID
Returns:
跳转后的状态
"""
if not self._session_id:
raise ValueError("无活动会话")
return await self.restore(self._session_id, checkpoint_id)
async def migrate_state(self, old_state: dict,
new_schema: dict[str, Any]) -> dict:
"""
状态迁移(版本升级时使用)
当状态 schema 发生变化时,将旧状态迁移到新格式。
Args:
old_state: 旧状态数据
new_schema: 新 schema 定义(字段名 → 默认值)
Returns:
迁移后的状态
"""
migrated = old_state.copy()
for key, default in new_schema.items():
if key not in migrated:
migrated[key] = default
# 标记迁移版本
migrated["_migrated_from"] = old_state.get("version", "unknown")
migrated["version"] = "migrated"
logger.info(f"[StateManager] 状态迁移完成,新增 {len(new_schema)} 个字段")
return migrated
async def clear_history(self, session_id: Optional[str] = None) -> int:
"""清除会话的所有检查点"""
sid = session_id or self._session_id
if not sid:
return 0
count = await self.store.clear(sid)
logger.info(f"[StateManager] 已清除 {count} 个检查点 (session={sid})")
return count
def _serialize_state(self, state: dict) -> dict:
"""序列化状态(默认直接返回 dict,子类可覆写)"""
# 深拷贝防止外部修改影响存储
import copy
return copy.deepcopy(state)
def _deserialize_state(self, data: dict) -> dict:
"""反序列化状态"""
return data.copy()
def update_state(self, **kwargs) -> None:
"""更新当前状态的部分字段"""
if self._current_state is None:
raise ValueError("无活动状态")
self._current_state.update(kwargs)
4.5 持久化存储实现
内存存储只适用于开发测试。生产环境需要持久化存储,以下是两种常见实现:
"""
持久化检查点存储实现
"""
import json
import os
from typing import Optional
class FileCheckpointStore(CheckpointStore):
"""
文件系统检查点存储
目录结构:
{base_dir}/
{session_id}/
checkpoint_{step}_{id}.json
metadata.json
"""
def __init__(self, base_dir: str = "./checkpoints"):
self.base_dir = base_dir
os.makedirs(base_dir, exist_ok=True)
def _session_dir(self, session_id: str) -> str:
return os.path.join(self.base_dir, session_id)
async def save(self, checkpoint: Checkpoint) -> str:
session_dir = self._session_dir(checkpoint.session_id)
os.makedirs(session_dir, exist_ok=True)
filename = f"cp_{checkpoint.step:04d}_{checkpoint.checkpoint_id[:8]}.json"
filepath = os.path.join(session_dir, filename)
with open(filepath, "w", encoding="utf-8") as f:
json.dump(checkpoint.to_dict(), f, ensure_ascii=False, indent=2)
return checkpoint.checkpoint_id
async def load(self, session_id: str,
checkpoint_id: Optional[str] = None) -> Optional[Checkpoint]:
session_dir = self._session_dir(session_id)
if not os.path.isdir(session_dir):
return None
if checkpoint_id:
# 按 ID 查找
for filename in os.listdir(session_dir):
if filename.endswith(".json") and checkpoint_id[:8] in filename:
filepath = os.path.join(session_dir, filename)
with open(filepath, "r", encoding="utf-8") as f:
return Checkpoint.from_dict(json.load(f))
return None
else:
# 加载最新的(按文件名排序)
files = sorted([f for f in os.listdir(session_dir) if f.endswith(".json")])
if not files:
return None
filepath = os.path.join(session_dir, files[-1])
with open(filepath, "r", encoding="utf-8") as f:
return Checkpoint.from_dict(json.load(f))
async def list_checkpoints(self, session_id: str) -> list[dict]:
session_dir = self._session_dir(session_id)
if not os.path.isdir(session_dir):
return []
result = []
for filename in sorted(os.listdir(session_dir)):
if not filename.endswith(".json"):
continue
filepath = os.path.join(session_dir, filename)
with open(filepath, "r", encoding="utf-8") as f:
data = json.load(f)
result.append({
"checkpoint_id": data["checkpoint_id"],
"step": data["step"],
"timestamp": data["timestamp"],
})
return result
async def delete(self, session_id: str, checkpoint_id: str) -> bool:
session_dir = self._session_dir(session_id)
if not os.path.isdir(session_dir):
return False
for filename in os.listdir(session_dir):
if checkpoint_id[:8] in filename:
os.remove(os.path.join(session_dir, filename))
return True
return False
async def clear(self, session_id: str) -> int:
import shutil
session_dir = self._session_dir(session_id)
if not os.path.isdir(session_dir):
return 0
count = len([f for f in os.listdir(session_dir) if f.endswith(".json")])
shutil.rmtree(session_dir)
return count
4.6 量化分析
| 操作 | 耗时范围 | 存储开销 | 触发频率 |
|---|---|---|---|
| 保存检查点(内存) | 1 – 5 ms | 2 – 50 KB/步 | 每步 1 次 |
| 保存检查点(文件) | 5 – 50 ms | 2 – 50 KB/步 | 每步 1 次 |
| 恢复检查点(内存) | 1 – 3 ms | - | 故障时 / 调试时 |
| 恢复检查点(文件) | 5 – 30 ms | - | 故障时 / 调试时 |
| 列出历史 | 5 – 100 ms | - | 调试时 |
| 状态迁移 | 1 – 10 ms | - | 版本升级时 |
| 时间旅行跳转 | 5 – 50 ms | - | 调试时 |
存储成本估算(20 步对话):
- 内存存储:~100 KB – 1 MB
- 文件存储:~100 KB – 1 MB
- 1000 个会话:~100 MB – 1 GB
1307

被折叠的 条评论
为什么被折叠?



