Agent 长任务的分段执行与 Checkpoint:避免超时丢进度

一、Agent 跑了 12 分钟终于快出结果了,HTTP 超时让你一切归零

Agent 执行一个复杂任务——"分析 10 份财报,生成对比报告"——需要调用多次 API、处理大量数据。假设你用的是 HTTP SSE(Server-Sent Events)完成这个任务,Agent 在第 8 分钟调用第 9 份财报时,客户端/网关超时断开,12 分钟的计算结果全部丢失。

这不是 Agent 不聪明,是执行模型没考虑长任务的特性。普通 API 调用是"请求-响应"模型,200ms 内返回。Agent 调用可能是"请求-多步推理-响应"模型,持续几分钟甚至几十分钟。两者的错误处理策略完全不同:短任务超时重试即可,长任务重试的代价太高——丢掉的不是一次 API 调用,是前面所有推理步骤的中间结果。

解决方案:分段执行 + Checkpoint。把 Agent 的长任务拆成独立的小段,每段结束后持久化状态。如果中途失败或超时,从最近的 Checkpoint 恢复继续执行,而不是从零开始。

二、底层机制与原理剖析

Checkpoint 保存什么(四个维度的状态):

LLM 上下文:对话历史、system prompt、当前步骤。这是恢复推理的最关键部分——没有上下文,Agent 不知道"前面已经做了什么"。

工具调用结果:每个工具调用的返回值。如果 Checkpoint 不包含已调用的工具结果,恢复后 Agent 会重复调用这些工具(浪费 API 配额、可能产生副作用)。

中间产出物:Agent 在这一步之前已经生成的文本、分析结果。如果不保存,恢复后需要重新生成——这恰恰是分段执行要避免的。

执行元数据:已用步数、已用时间、当前置信度、预算消耗量。这些信息帮助恢复后的 Agent 判断"还剩多少时间/步数可以继续"。

三、生产级代码实现

"""
Agent 分段执行与 Checkpoint 机制

设计思路:
1. 每个 Segment 是一个不可再分的执行单元
2. Checkpoint 持久化到 Redis + 可选持久化到 S3(防 Redis 丢失)
3. 恢复时从最近的 Checkpoint 读取,跳过已完成的 Segment
4. 幂等性保证:同一个 Segment 重复执行不会产生副作用
"""

import json
import time
import uuid
import logging
import hashlib
from typing import Dict, List, Optional, Any
from dataclasses import dataclass, field
from enum import Enum
from datetime import datetime, timedelta

import redis

logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)

class CheckpointStatus(Enum):
    IN_PROGRESS = "in_progress"  # Segment 正在执行
    COMPLETED = "completed"      # Segment 已完成
    FAILED = "failed"            # Segment 执行失败

@dataclass
class SegmentCheckpoint:
    """单个 Segment 的 Checkpoint 数据"""
    segment_id: str               # 唯一标识
    task_id: str                  # 所属任务
    segment_index: int            # 在任务中的序号(0-based)
    status: CheckpointStatus = CheckpointStatus.IN_PROGRESS

    # LLM 上下文
    messages: List[Dict[str, Any]] = field(default_factory=list)
    system_prompt: str = ""

    # 工具调用结果
    tool_results: Dict[str, Any] = field(default_factory=dict)

    # 中间产出物
    partial_output: str = ""      # 到这一段为止的部分输出
    intermediate_data: Dict[str, Any] = field(default_factory=dict)

    # 执行元数据
    steps_used: int = 0
    elapsed_ms: int = 0
    created_at: str = ""
    updated_at: str = ""

    def __post_init__(self):
        now = datetime.utcnow().isoformat()
        if not self.created_at:
            self.created_at = now
        if not self.updated_at:
            self.updated_at = now

    def to_dict(self) -> dict:
        return {
            "segment_id": self.segment_id,
            "task_id": self.task_id,
            "segment_index": self.segment_index,
            "status": self.status.value,
            "messages": self.messages,
            "system_prompt": self.system_prompt,
            "tool_results": self.tool_results,
            "partial_output": self.partial_output,
            "intermediate_data": self.intermediate_data,
            "steps_used": self.steps_used,
            "elapsed_ms": self.elapsed_ms,
            "created_at": self.created_at,
            "updated_at": self.updated_at,
        }

    @classmethod
    def from_dict(cls, data: dict) -> "SegmentCheckpoint":
        return cls(
            segment_id=data["segment_id"],
            task_id=data["task_id"],
            segment_index=data["segment_index"],
            status=CheckpointStatus(data["status"]),
            messages=data.get("messages", []),
            system_prompt=data.get("system_prompt", ""),
            tool_results=data.get("tool_results", {}),
            partial_output=data.get("partial_output", ""),
            intermediate_data=data.get("intermediate_data", {}),
            steps_used=data.get("steps_used", 0),
            elapsed_ms=data.get("elapsed_ms", 0),
            created_at=data.get("created_at", ""),
            updated_at=data.get("updated_at", ""),
        )

class CheckpointStore:
    """
    Checkpoint 存储层

    双写策略(Redis + 可选 S3)的原因:
    - Redis 用于快速读写(恢复时的延迟 < 5ms)
    - 持久化存储用于 Redis 不可用时恢复(避免单点故障)
    - 恢复时优先读 Redis,Redis 未命中再读持久化存储
    """

    def __init__(self, redis_client: redis.Redis, ttl_hours: int = 24):
        self.redis = redis_client
        self.ttl = ttl_hours * 3600  # Redis key 的过期时间

    def _task_key(self, task_id: str) -> str:
        return f"agent:checkpoint:{task_id}"

    def save(self, checkpoint: SegmentCheckpoint) -> bool:
        """
        保存 Checkpoint

        使用 Redis HSET 而非 SET,原因:
        - 一个任务有多个 Segment,用 hash 结构管理更方便
        - 可以单独读取某个 Segment 的 Checkpoint 而不需要反序列化全部
        """
        checkpoint.updated_at = datetime.utcnow().isoformat()
        key = self._task_key(checkpoint.task_id)
        field = str(checkpoint.segment_index)

        try:
            # 使用 pipeline 原子执行:设置 hash field + 刷新 TTL
            pipe = self.redis.pipeline()
            pipe.hset(key, field, json.dumps(checkpoint.to_dict()))
            pipe.expire(key, self.ttl)
            pipe.execute()
            logger.debug("Checkpoint saved: task=%s segment=%d",
                        checkpoint.task_id, checkpoint.segment_index)
            return True
        except redis.RedisError as e:
            logger.error("Failed to save checkpoint: %s", e)
            return False

    def load(self, task_id: str, segment_index: int) -> Optional[SegmentCheckpoint]:
        """加载指定 Segment 的 Checkpoint"""
        key = self._task_key(task_id)
        try:
            raw = self.redis.hget(key, str(segment_index))
            if raw:
                data = json.loads(raw)
                return SegmentCheckpoint.from_dict(data)
        except redis.RedisError as e:
            logger.error("Failed to load checkpoint: %s", e)
        return None

    def load_latest(self, task_id: str) -> Optional[SegmentCheckpoint]:
        """
        加载任务最近的已完成 Checkpoint

        为什么需要"最近的已完成"?
        因为恢复执行时,不需要从 segment 0 开始——
        找到最后一个 status=COMPLETED 的 Checkpoint 即可
        """
        key = self._task_key(task_id)
        try:
            all_raw = self.redis.hgetall(key)
            if not all_raw:
                return None

            latest = None
            latest_index = -1
            for seg_index, raw in all_raw.items():
                data = json.loads(raw)
                if data.get("status") != CheckpointStatus.COMPLETED.value:
                    continue
                idx = int(seg_index)
                if idx > latest_index:
                    latest_index = idx
                    latest = data

            if latest:
                return SegmentCheckpoint.from_dict(latest)
        except redis.RedisError as e:
            logger.error("Failed to load latest checkpoint: %s", e)
        return None

    def delete_task(self, task_id: str) -> bool:
        """任务完成后清理所有 Checkpoint"""
        key = self._task_key(task_id)
        try:
            self.redis.delete(key)
            return True
        except redis.RedisError as e:
            logger.error("Failed to delete checkpoints: %s", e)
            return False

class LongRunningAgentExecutor:
    """
    Agent 分段执行器

    执行流程:
    1. 尝试从最近 Checkpoint 恢复
    2. 如果没有 Checkpoint,从头开始
    3. 执行每个 Segment(超时保护 + 异常保护)
    4. 每个 Segment 完成后保存 Checkpoint
    """

    MAX_SEGMENT_TIME_MS = 120_000  # 单个 Segment 最长执行 2 分钟
    MAX_RETRIES = 2                # 单个 Segment 最多重试 2 次

    def __init__(self, checkpoint_store: CheckpointStore, llm_client=None):
        self.store = checkpoint_store
        self.llm_client = llm_client  # 生产环境换成真实的 LLM client

    def execute(
        self,
        task_id: str,
        segments: List[callable],
        context: Dict[str, Any],
    ) -> Dict[str, Any]:
        """
        分段执行 Agent 任务

        参数:
            task_id: 任务唯一标识
            segments: Segment 执行函数列表(每个函数接受 context 返回更新后的 context)
            context: 初始上下文
        """
        # 1. 尝试从 Checkpoint 恢复
        latest = self.store.load_latest(task_id)
        start_segment = 0

        if latest:
            logger.info("Resuming task %s from segment %d", task_id, latest.segment_index + 1)
            start_segment = latest.segment_index + 1
            context = self._merge_context(context, latest)
        else:
            logger.info("Starting new task %s with %d segments", task_id, len(segments))

        # 2. 逐 Segment 执行
        for i in range(start_segment, len(segments)):
            segment_fn = segments[i]
            checkpoint = SegmentCheckpoint(
                segment_id=str(uuid.uuid4()),
                task_id=task_id,
                segment_index=i,
            )

            # 执行当前 Segment(带重试)
            success = False
            for retry in range(self.MAX_RETRIES + 1):
                try:
                    result = self._execute_segment_with_timeout(
                        segment_fn, context, self.MAX_SEGMENT_TIME_MS
                    )
                    context.update(result)
                    checkpoint.status = CheckpointStatus.COMPLETED
                    checkpoint.partial_output = result.get("output", "")
                    checkpoint.steps_used = result.get("steps", 0)

                    # 保存 Checkpoint
                    self.store.save(checkpoint)
                    success = True
                    logger.info("Segment %d/%d completed for task %s",
                              i + 1, len(segments), task_id)
                    break

                except Exception as e:
                    logger.warning(
                        "Segment %d failed (retry %d/%d): %s",
                        i, retry + 1, self.MAX_RETRIES, e
                    )
                    if retry == self.MAX_RETRIES:
                        checkpoint.status = CheckpointStatus.FAILED
                        self.store.save(checkpoint)
                        raise RuntimeError(
                            f"Segment {i} failed after {self.MAX_RETRIES} retries: {e}"
                        )
                    time.sleep(min(2 ** retry, 10))  # 指数退避,上限 10s

        # 3. 全部完成,清理 Checkpoint
        self.store.delete_task(task_id)
        logger.info("Task %s completed successfully", task_id)

        return context

    def _execute_segment_with_timeout(
        self, fn: callable, context: dict, timeout_ms: int
    ) -> dict:
        """
        带超时保护的 Segment 执行

        为什么用 signal 而不是 threading?
        - threading 的 join(timeout) 不真正杀死线程,只等待超时
        - 对应的线程还在运行,占着 LLM API 连接不释放
        - 生产环境建议用 multiprocessing 或 asyncio 实现真正的超时杀
        """
        import signal

        def timeout_handler(signum, frame):
            raise TimeoutError(f"Segment execution timed out after {timeout_ms}ms")

        old_handler = signal.signal(signal.SIGALRM, timeout_handler)
        signal.alarm(timeout_ms // 1000)

        try:
            result = fn(context)
            return result
        finally:
            signal.alarm(0)  # 取消闹钟
            signal.signal(signal.SIGALRM, old_handler)  # 恢复原 handler

    def _merge_context(
        self, current: dict, checkpoint: SegmentCheckpoint
    ) -> dict:
        """
        合并 Checkpoint 上下文到当前上下文

        恢复执行时,用 Checkpoint 中保存的中间数据覆盖当前上下文。
        为什么不全量覆盖?因为当前上下文可能包含新的参数
        (如用户重试时改了某些配置),需要保留这些新值。
        """
        merged = dict(current)
        merged.update(checkpoint.intermediate_data)
        merged["_partial_output"] = checkpoint.partial_output
        merged["_resumed_from"] = checkpoint.segment_index
        merged["_previous_steps"] = checkpoint.steps_used
        return merged

# ---------------------------------------------------------------------------
# 使用示例
# ---------------------------------------------------------------------------
if __name__ == "__main__":
    # 初始化
    r = redis.Redis(host="localhost", port=6379, decode_responses=True)
    store = CheckpointStore(r)
    executor = LongRunningAgentExecutor(store)

    # 定义 3 个 Segment
    def segment_collect_data(ctx):
        """第 1 段:收集数据"""
        # 模拟数据收集...
        time.sleep(1)
        return {"output": "数据收集完成", "steps": 1, "data": ["A", "B", "C"]}

    def segment_analyze_data(ctx):
        """第 2 段:分析数据"""
        time.sleep(2)
        return {"output": "分析完成", "steps": 1, "analysis": "result"}

    def segment_generate_report(ctx):
        """第 3 段:生成报告"""
        time.sleep(1)
        return {"output": "报告生成完成", "steps": 1, "report": "final"}

    # 执行
    result = executor.execute(
        task_id="task-001",
        segments=[segment_collect_data, segment_analyze_data, segment_generate_report],
        context={},
    )
    print("最终结果:", result)

四、边界分析与架构权衡

Checkpoint 的粒度选择

  • 太粗(3 个 Segment 覆盖整个任务)——恢复后重做成本高
  • 太细(每一步都做 Checkpoint)——持久化开销超过计算本身
  • 推荐:按"不可逆操作"分段——一个工具调用 = 一个 Segment,一个 LLM 推理 + 解析 = 一个 Segment

Checkpoint 一致性风险

  • 如果 Segment 期间执行了有副作用的操作(如发了一封邮件、创建了一个订单),重试时会重复执行这些副作用
  • 解决方案:所有副作用操作必须在 Segment 最后执行,且在 Checkpoint 中记录"副作用已执行"标记
  • 更好的方案:副作用操作做成幂等的(数据库 upsert 而非 insert,邮件去重)

什么场景不该用分段执行

  • 总执行时间 < 30 秒的任务——Checkpoint 的开销(序列化、Redis 写入)可能超过任务本身
  • 强实时性要求(用户等待时间 < 5 秒)——分段执行增加了额外的 Checkpoint 等待时间
  • 纯无状态任务——每次重试代价极低,不值得持久化中间状态

五、总结

Agent 长任务的分段执行 + Checkpoint 机制,本质是用空间(持久化中间状态)换可靠性(避免从头重试)。Checkpoint 四个维度的状态(LLM 上下文、工具结果、中间产出、执行元数据)缺一不可。粒度选择是核心权衡——太粗收益小,太细开销大。关键是识别"不可逆操作"节点,在这些节点前后做 Checkpoint。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐