AI Agent自研框架调度引擎与执行流控制(上)

自研框架调度引擎与执行流控制

专题定位:深入自研 Agent 框架的调度内核,从执行主循环、消息传递、工具调度、状态管理、记忆存取、规划推理到并发控制,逐层拆解七大核心机理的设计原理、工程实现与性能权衡。每章配备 Mermaid 架构图、核心模块源码级实现与量化分析矩阵。

前置阅读:第 08 篇《自研 Agent 框架总体架构设计》(四层架构、核心抽象、事件驱动、插件体系)

代码规范:所有示例代码遵循 Python 3.11+ 类型注解,异步优先(async/await),接口面向抽象而非实现。


目录


第 1 章 执行机理:Agent 主循环的心脏搏动

1.1 主循环的本质

Agent 主循环(Agent Loop)是整个框架的心脏——它驱动着感知、推理、规划、行动、观察、反思六个阶段的持续流转。没有主循环,Agent 只是一堆零散的组件;有了主循环,组件才能协同成一个"会思考、会行动"的智能体。

从控制论的视角看,Agent 主循环本质上是一个感知-行动循环(Perception-Action Loop)的扩展版本,它在经典控制论反馈环的基础上,增加了推理、规划和反思三个认知层:

经典控制论: 感知 → 决策 → 行动 → 反馈
Agent 扩展: 感知 → 推理 → 规划 → 行动 → 观察 → 反思

1.2 六阶段流转全景

未完成

已完成

用户输入

1. Perceive 感知

2. Reason 推理

3. Plan 规划
需要工具?

4. Act 行动
执行工具调用

5. Observe 观察
解析工具结果

6. Reflect 反思
可选

返回最终结果

六阶段详解:

阶段名称核心职责输入输出
1Perceive(感知)接收外部输入,更新状态,写入短期记忆用户输入 / 工具返回 / 环境事件AgentState.messages 更新
2Reason(推理)组装 Prompt,调用 LLM,解析输出system_prompt + memory + messagesLLM 响应(文本 + tool_calls)
3Plan(规划)判断是否需要工具,选择工具+参数,确定执行顺序LLM tool_calls执行计划(串行/并行)
4Act(行动)执行工具调用,沙箱隔离,超时控制,错误捕获工具 + 参数ToolResult
5Observe(观察)解析工具结果,写入记忆,判断终止条件ToolResult更新后的状态 + 是否继续
6Reflect(反思)评估结果质量,总结经验,调整策略最终结果 + 原始目标反思笔记(可选)

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 的扩展点,允许在不修改核心循环代码的前提下,注入横切关注点(如日志、追踪、限流、缓存、安全审计等)。

LLM引擎 AgentLoop核心 中间件2 中间件1 调用方 LLM引擎 AgentLoop核心 中间件2 中间件1 调用方 before(context) before(context) 执行主循环 chat() response result after(context, result) result after(context, result) final_result

中间件接口定义:

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 ms70 – 85%极高Prompt 缓存、模型路由、流式输出
Plan(规划)< 5 ms~0%极低内置于 Reason,无需单独优化
Act(行动)100 – 10,000 ms10 – 25%并行执行、超时控制、连接池
Observe(观察)< 2 ms~0%极低无需优化
Reflect(反思)200 – 1,000 ms5 – 10%可选开关、异步后台执行
单步总计800 – 16,000 ms100%-LLM 调用是核心瓶颈

关键洞察:

  1. LLM 调用占绝对主导:Reason 阶段消耗了 70-85% 的时间,优化 LLM 调用是性能提升的核心
  2. 工具执行波动大:Act 阶段耗时从 100ms 到 10s 不等,取决于工具类型(本地函数 vs 远程 API vs 代码执行)
  3. 反思是可选开销:Reflect 增加 5-10% 开销,建议仅在需要学习/迭代的场景开启
  4. 感知/规划/观察几乎可忽略:这三个阶段加起来不到 1%,优化收益极低

第 2 章 消息传递机理:事件驱动的组件通信

2.1 为什么需要消息传递

在一个复杂的 Agent 框架中,组件之间的通信方式决定了系统的可扩展性、可观测性和解耦程度。如果每个组件都直接调用其他组件的方法,系统会变成一张紧密耦合的网——改动一个组件可能牵一发而动全身。

消息传递机制通过事件总线(EventBus) 实现了发布/订阅模式,让组件之间通过事件而非直接调用来通信:

  • 发布者(Publisher):只负责发布事件,不关心谁来处理
  • 订阅者(Subscriber):只关心自己感兴趣的事件,不关心谁发布的
  • 事件总线(EventBus):负责事件的路由和分发

2.2 三种通信模式

自研框架采用混合通信模式,不同场景使用不同的通信方式:

通信模式

核心路径

横切关注点

管道模式

同步调用

LLM / Memory / Tool

异步事件

Tracing / Logging / Metrics

回调链

中间件 / 过滤器

通信模式延迟吞吐量适用场景实现方式
同步调用0.1 – 1 msLLM、Memory、Tool 等核心路径的直接调用直接方法调用
异步事件1 – 10 ms追踪、日志、指标收集等横切关注点EventBus 发布/订阅
回调链0.5 – 5 ms中间件管道、请求过滤器Middleware Chain
发布/订阅1 – 10 ms多组件并行通知EventBus + Handler 列表

2.3 事件总线架构

发布事件

分发

订阅者

日志处理器
LoggingHandler

追踪处理器
TracingHandler

指标处理器
MetricsHandler

回调处理器
CallbackHandler

Webhook 通知
WebhookNotifier

事件总线 EventBus

on_think

on_act

on_observe

step_complete

on_reflect

on_error

Agent 核心

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_thinkLLM 返回推理结果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)

输入

Middleware 1
before

Middleware 2
before

核心逻辑

Middleware 2
after

Middleware 1
after

输出

"""
回调链实现 - 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 与外部世界交互的桥梁。看似简单的"调用一个函数"背后,隐藏着一系列工程挑战:

  1. 工具发现:LLM 输出的工具名可能不存在、拼写错误
  2. 参数校验:LLM 生成的参数可能格式错误、缺少必填字段、类型不匹配
  3. 执行安全:工具可能执行危险操作、超时、抛出异常
  4. 结果处理:工具输出可能超长、格式不规范、需要二次解析
  5. 错误恢复:工具失败后如何重试、降级、告知 LLM 调整策略

自研框架的工具调度器(ToolScheduler)用一条五步链路系统性地解决这些问题。

3.2 五步链路全景

工具存在

工具不存在

校验通过

校验失败

成功

失败

可重试

不可重试

LLM tool_calls

Step 1
工具发现

Step 2
参数校验

返回错误
工具不存在

Step 3
调用执行

返回错误
参数错误

Step 4
结果解析

Step 5
错误处理

ToolResult

指数退避重试
最多 3 次

返回错误结果

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)

执行器负责实际调用工具,是工具调度的执行层。支持多种执行模式:

执行器 IExecutor

LocalExecutor
本地函数执行

RemoteExecutor
远程 API 执行

SandboxExecutor
沙箱代码执行

MCPExecutor
MCP 工具执行

ToolScheduler

工具实现

"""
执行器接口与实现
"""
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 ms0.1%直接报错哈希表 O(1) 查找
参数校验1 – 5 ms2 – 5%返回 LLM 修正Schema 缓存、预编译
工具调用100 – 10,000 ms5 – 15%指数退避 × 3连接池、并行执行、超时控制
结果解析< 2 ms< 1%截断处理流式解析、增量截断
错误处理< 5 ms-返回错误信息错误分类、重试决策

关键洞察:

  1. 工具调用是第二大耗时源:仅次于 LLM 调用,且波动极大(100ms → 10s)
  2. 参数校验失败率不可忽视:2-5% 的失败率意味着 LLM 经常生成错误参数,需要良好的错误反馈机制
  3. 重试策略要精准:不是所有错误都该重试,网络/超时类错误重试才有意义
  4. 结果截断是必要的:不加截断的工具输出可能轻易撑爆 LLM context window

第 4 章 状态管理机理:检查点快照与时间旅行

4.1 为什么需要状态管理

Agent 的执行过程是有状态的——每一步的推理、工具调用、记忆写入都会改变 Agent 的状态。如果没有良好的状态管理机制,会面临以下问题:

  1. 故障恢复:执行中途崩溃后,无法从断点继续,只能从头开始
  2. 调试困难:出了问题无法回溯到出错前的状态进行分析
  3. 分支探索:无法对同一问题尝试不同策略并对比结果
  4. 长任务中断:长时间运行的任务无法暂停和恢复

自研框架的状态管理机制通过检查点快照(Checkpoint) 解决这些问题,支持状态保存、恢复、历史回溯和时间旅行调试。

4.2 状态生命周期

初始化配置

run()

pause()

resume()

任务完成

发生错误

重试/恢复

持久化退出

Created

Initialized

Running

Paused

Resumed

Completed

Error

状态说明:

状态含义可转换到
Created刚创建,尚未初始化Initialized
Initialized配置加载完成,等待执行Running
Running正在执行主循环Paused, Completed, Error
Paused已暂停,状态已持久化Resumed
Resumed从暂停恢复,即将继续Running
Completed任务正常完成-(终态)
Error执行出错Running(重试)

4.3 检查点机制

检查点是状态管理的核心——每执行一步,自动保存当前状态的快照。

检查点

执行流

Step 0

Step 1

Step 2

Step 3

Step 4

CP #0

CP #1

CP #2

CP #3

CP #4

检查点存储内容:

数据项说明大小估算
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 ms2 – 50 KB/步每步 1 次
保存检查点(文件)5 – 50 ms2 – 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

评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包

打赏作者

千江明月

你的鼓励将是我创作的最大动力

¥1 ¥2 ¥4 ¥6 ¥10 ¥20
扫码支付:¥1
获取中
扫码支付

您的余额不足,请更换扫码支付或充值

打赏作者

实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值