Agent 长任务的分段执行与 Checkpoint:避免超时丢进度
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。
更多推荐


所有评论(0)