AI大模型之LangGraph:线程与检查点

一、基础概念

LangGraph中的持久化,使用持久化的核心组件线程(thread)和检查点(checkpoint)

1.1 什么是持久化?

在LangGraph图执行过程中,每一个节点执行后会将图状态生成一份快照保存起来,这些快照在图执行结束后也可以访问,这种机制被称为LangGraph的持久化

1.2 持久化的作用?

通过LangGraph的持久化,可以实现人在环路、时间旅行、记忆组件等功能

1.3 什么是检查点与线程?

在LangGraph中,持久化数据是由检查点管理器(checkpointer)来实现的,当图中的节点执行完一个超级步骤(super-step)时,会保存一份图状态快照信息(StateSnapshot)到线程(thread)中。 每一个线程都有一个唯一标识thread ID,可以通过thread ID找到thread,再从thread读取中的快照数据,这些保存的图状态数据快照就被称为检查点(checkpoint),在检查点中就包含了图状态数据(State)等信息。
其中超级步骤(super-step)是指:
对于顺序节点:一个顺序节点执行完后,就完成了一个超级步骤,此时会保存检查点。
对于并行节点:所有并行节点都执行完后,才共同完成一个超级步骤,此时保存一个检查点。

在这里插入图片描述

1.4 使用步骤
1.4.1 首先要在创建图时指定检查点管理器(checkpointer),LangGraph的持久化功能才能生效
# 可以使用redis作缓存
checkpointer = InMemorySaver()
agent = graph.compile(checkpointer=checkpointer)
1.4.2 指定thread_id
config = {"configurable": {"thread_id": "1"}}
state = agent.invoke({"result": []}, config)
1.5 检查点的使用
1.5.1 检查点结构

检查点作为图状态数据快照,它包含以下关键属性:
config:Config配置信息
metadata:元数据信息
values: 图状态数据
tasks:有关下一步要执行的任务信息PregelTask对象元组(在执行完成的最终检查点中,此字段通常为空)

1.5.2 获取最新检查点 get_state()

在图完成编译和执行之后,可以调用get_state()方法获取最新的检查点,用法如下,需要传入配置信息,配置信息中需要包含thread_id。

config = {"configurable": {"thread_id": "1"}}
print("==============获取最新的检查点===============")
latest_state = agent.get_state(config)
print(latest_state)

输出:

StateSnapshot(values={'result': ['你好,我是a节点', '你好,我是b节点']}, next=(), config={'configurable': {'thread_id': '1', 'checkpoint_ns': '', 'checkpoint_id': '1f0de5ff-1c74-6c92-8002-404168dee320'}}, metadata={'source': 'loop', 'step': 2, 'parents': {}}, created_at='2025-12-21T11:26:50.811499+00:00', parent_config={'configurable': {'thread_id': '1', 'checkpoint_ns': '', 'checkpoint_id': '1f0de5ff-1c73-6626-8001-38a35205281b'}}, tasks=(), interrupts=())
1.5.3 获取历史检查点列表 get_state_history()
config = {"configurable": {"thread_id": "1"}}
history_states = agent.get_state_history(config)
print("==============获取历史检查点===============")
for state in history_states:
    print(state)

执行结果,包含四个检查点信息,并且顺序是从最新的检查点开始排列
这四个检查点分别在四个super-step进行保存:
1、第一个检查点是一个空的检查点,下一个要执行的节点是START节点。
2、第二个检查点是执行完START节点,下一个要执行的是a_node节点,并且图状态数据values还是初始化状态。
3、第三个检查点是执行完a_node节点,下一个要执行的是b_node节点,并且图状态数据为’result’: [‘你好,我是a节点’]
4、第四个检查点是执行完b_node节点,没有下一个要执行的节点,并且图状态数据为{‘result’: [‘你好,我是a节点’, ‘你好,我是b节点’]}

1.5.4 回放检查点

检查点回放是指:通过传递的thread_id和checkpoint_id,对指定检查点之前的步骤进行重放,不会重新执行,在检查点之后的步骤全部都会重新执行

1.5.5 更新图状态 update_state()

执行更新后,系统会创建一个新的检查点,并添加到在当前thread中

二、代码示例

import operator
from itertools import islice
from typing import TypedDict, Annotated
from langgraph.graph import StateGraph
from langgraph.constants import START, END
from langgraph.checkpoint.redis import RedisSaver


# =====================================================
# ⭐ 自定义 TTLRedisSaver 用于处理TTL
# =====================================================
class TTLRedisSaver(RedisSaver):
    def __init__(self, *args, ttl_seconds=3600, **kwargs):
        super().__init__(*args, **kwargs)
        self.ttl_seconds = ttl_seconds

    # 带TTL支持的上下文管理器
    @classmethod
    def from_conn_string(cls, conn_string, ttl_seconds=3600):
        parent_cm = super().from_conn_string(conn_string)

        class _TTLContext:
            def __enter__(self):
                saver = parent_cm.__enter__()
                saver.__class__ = cls
                saver.ttl_seconds = ttl_seconds
                return saver

            def __exit__(self, exc_type, exc, tb):
                return parent_cm.__exit__(exc_type, exc, tb)

        return _TTLContext()

    # 重写put方法,包含TTL逻辑
    def put(self, config, checkpoint, metadata, new_versions):
        result = super().put(config, checkpoint, metadata, new_versions)
        self._apply_ttl(config)
        return result

    # 重写get方法,包含TTL逻辑
    def get(self, config):
        data = super().get(config)
        self._apply_ttl(config)
        return data

    # 应用TTL
    def _apply_ttl(self, config):
        try:
            thread_id = config["configurable"]["thread_id"]
            for key in self._redis.scan_iter(f"*{thread_id}*"):
                self._redis.expire(key, self.ttl_seconds)
        except Exception as e:
            print(f"TTL设置失败: {e}")


# =====================================================
# 1️⃣ 定义State(图状态)
# =====================================================
class State(TypedDict):
    result: Annotated[list[str], operator.add]


# =====================================================
# 2️⃣ 定义节点(Node)
# =====================================================
def a_node(state: State):
    return {"result": ["我是大A"]}

def b_node(state: State):
    return {"result": ["我是大B"]}

def c_node(state: State):
    return {"result": ["我是大C"]}


# =====================================================
# 4️⃣ 构建LangGraph
# =====================================================
def build_graph(checkpointer):
    graph = StateGraph(State)
    graph.add_node("a_node", a_node)
    graph.add_node("b_node", b_node)
    graph.add_node("c_node", c_node)
    graph.add_edge(START, "a_node")
    graph.add_edge("a_node", "b_node")
    graph.add_edge("b_node", "c_node")
    graph.add_edge("c_node", END)
    return graph.compile(checkpointer=checkpointer)


# =====================================================
# 5️⃣ 执行Graph(执行节点)
# =====================================================
def run_graph(agent, thread_id):
    # 保留config结构,确保返回时不会丢失配置
    config = {
        "configurable": {  # 确保config中有'configurable'键,且包含'thread_id'
            "thread_id": thread_id  # 确保thread_id传递
        }
    }
    result = agent.invoke({"result": []}, config)  # 返回result并保留原config结构
    config["result"] = result["result"]  # 将result结果存入config中
    print(f"运行图结果: {config}")  # 添加调试输出,查看config结构
    return config


# =====================================================
# 6️⃣ 查看最新的Checkpoint(检查点)
# =====================================================
def show_latest(agent, config):
    print("\n最新的Checkpoint:")
    try:
        print(agent.get_state(config))
    except KeyError as e:
        print(f"获取状态失败: {e}")


# =====================================================
# 7️⃣ 查看历史Checkpoint
# =====================================================
def show_history(agent, config):
    print("\n历史Checkpoint(显示最后3个):")
    print(f"调试: 当前config: {config}")  # 调试输出config

    # 确保config结构包含'thread_id'键
    if "configurable" not in config or "thread_id" not in config["configurable"]:
        print("错误: config中缺少'thread_id'键。")
        return

    history = list(islice(agent.get_state_history(config), 3))  # 只显示最后3个历史
    for state in history:
        print(state)


# =====================================================
# 8️⃣ 时间回放(Replaying Checkpoints)
# =====================================================
def replay_checkpoint(agent, config, index):
    print("\n回放Checkpoint:")
    print(f"调试: 当前config回放前: {config}")  # 调试输出config

    target_state = next(islice(agent.get_state_history(config), index, None))
    checkpoint_id = target_state.config["configurable"]["checkpoint_id"]
    replay_config = {"configurable": {"thread_id": config["configurable"]["thread_id"], "checkpoint_id": checkpoint_id}}
    replay_state = agent.invoke(None, replay_config)
    print("回放结果:", replay_state["result"])


# =====================================================
# 9️⃣ 人工修改Checkpoint(人工干预)
# =====================================================
def update_checkpoint(agent, config):
    print("\n人工修改Checkpoint:")
    agent.update_state(config, {"result": ["人工修改到c节点"]})


# =====================================================
# 🔟 主程序
# =====================================================
def main():
    # ⭐ 使用带TTL支持的RedisSaver
    with TTLRedisSaver.from_conn_string(
            "redis://:123456@192.168.174.198:6379/0",
            ttl_seconds=60  # 设置TTL为1分钟
    ) as checkpointer:

        # 初始化检查点
        try:
            checkpointer.setup()
        except Exception:
            pass

        # 构建LangGraph
        agent = build_graph(checkpointer)

        # 执行一次图
        config = run_graph(agent, "thread-lc")

        # 查看最新Checkpoint
        show_latest(agent, config)

        # 查看历史Checkpoint(显示最后3个)
        show_history(agent, config)

        # 回放Checkpoint(回放第2个)
        replay_checkpoint(agent, config, 2)

        # 修改Checkpoint(人工干预)
        update_checkpoint(agent, config)

        # 查看修改后的历史Checkpoint
        show_history(agent, config)


if __name__ == "__main__":
    main()

注意:这里使用redis做缓存并设有过期时间,也可以使用内存作缓存:

checkpointer = InMemorySaver()
agent = graph.compile(checkpointer=checkpointer)

三、记忆存储 Memory Store

3.1 基础概念

检查点保存的数据与特定线程(Thread)绑定,因此只能在同一个线程内访问。例如,使用 thread_id_a 无法获取到
thread_id_b 的检查点数据。如果我们想让持久化的数据能在不同的线程间共享,那就需要用到记忆存储了

存储方式

in_memory_store = InMemoryStore(),
LangGraph提供了RedisStore、PostgresStore等持久化存储后端

  • put()存储数据

  • get()获取数据

  • search()检索数据

LangGraph的三个节点

load_memory_node:从记忆存储中加载用户聊天历史
通过config获取user_id并拼接好namespace,进行记忆检索,将检索到的记忆保存到State中的chat_history

llm_node:调用llm
首先将之前节点读取的chat_history格式化成字符串,并构造一个SystemMessage。然后将SystemMessage与状态中的最新HumanMessage一同作为参数调用LLM。LLM返回AIMessage,该消息会自动添加到State中的messages列表中

save_memory_node:保存本次对话到Memory Store
将State消息列表messages中的AI消息存储到Memory Store中。下次同一用户调用图时,便可读取到本次对话的历史信息

3.2 代码示例
import json
import operator
from typing import TypedDict, Annotated
import dotenv
from langgraph.store.base import BaseStore
from langgraph.graph import StateGraph
from langgraph.constants import START, END
import redis
from langchain_core.messages import (
    AnyMessage,
    SystemMessage,
    AIMessage,
    HumanMessage,
)
from langchain_core.runnables import RunnableConfig
from langchain_openai import ChatOpenAI

dotenv.load_dotenv()

# =====================================================
# 1️⃣ LLM
# =====================================================
llm = ChatOpenAI(model="gpt-4o-mini", temperature=0)


# =====================================================
# 2️⃣ Redis TTL Store
# =====================================================
class RedisTTLStore(BaseStore):

    def __init__(
            self,
            redis_client,
            ttl: int = 3600,
            max_history: int = 20,
    ):
        """
        redis_client: 外部传入 Redis 连接
        """
        self.redis = redis_client
        self.ttl = ttl
        self.max_history = max_history

    # =============================
    # key
    # =============================
    def _key(self, namespace):
        return ":".join(namespace)

    # =============================
    # put
    # =============================
    def put(self, namespace, key, value, **kwargs):

        k = self._key(namespace)

        self.redis.lpush(k, json.dumps(value))
        self.redis.ltrim(k, 0, self.max_history - 1)
        self.redis.expire(k, self.ttl)

    # =============================
    # search
    # =============================
    def search(self, namespace, **kwargs):

        k = self._key(namespace)
        items = self.redis.lrange(k, 0, -1)

        return [
            {
                "key": str(i),
                "value": json.loads(v),
                "namespace": namespace,
            }
            for i, v in enumerate(reversed(items))
        ]

    # =============================
    # batch
    # =============================
    def batch(self, operations, **kwargs):

        results = []

        for op in operations:

            if op["op"] == "put":
                self.put(
                    op["namespace"],
                    op.get("key"),
                    op["value"],
                )
                results.append(None)

            elif op["op"] == "search":
                results.append(
                    self.search(op["namespace"])
                )

            else:
                raise ValueError(f"Unsupported op: {op['op']}")

        return results

    # =============================
    # async batch
    # =============================
    async def abatch(self, operations, **kwargs):
        return self.batch(operations, **kwargs)


# =====================================================
# 3️⃣ State
# =====================================================
class MessagesState(TypedDict):
    messages: Annotated[list[AnyMessage], operator.add]
    chat_history: Annotated[list[dict], operator.add]


# =====================================================
# 4️⃣ 工具函数
# =====================================================
def get_namespace(config: RunnableConfig):
    user_id = config["configurable"]["user_id"]
    return ("chat_history", user_id)


# =====================================================
# 5️⃣ Load Memory Node
# =====================================================
def load_memory_node(state, config, store):
    memories = store.search(get_namespace(config))

    history = [
        {
            "role": m["value"]["role"],
            "content": m["value"]["content"],
        }
        for m in memories
    ]

    return {"chat_history": history}


# =====================================================
# 6️⃣ LLM Node
# =====================================================
def llm_node(
        state: MessagesState,
        config: RunnableConfig,
        store: BaseStore,
):
    system_message = SystemMessage(
        content="你是一个智能助手,请结合历史记忆回答用户问题"
    )

    # 历史转换为Message对象
    history_msgs = [
        HumanMessage(content=m["content"])
        if m["role"] == "human"
        else AIMessage(content=m["content"])
        for m in state["chat_history"]
    ]

    ai_message = llm.invoke(
        [system_message, *history_msgs, *state["messages"]]
    )

    return {"messages": [ai_message]}


# =====================================================
# 7️⃣ Save Memory Node
# =====================================================
def save_memory_node(
        state: MessagesState,
        config: RunnableConfig,
        store: BaseStore,
):
    namespace = get_namespace(config)

    for msg in state["messages"]:
        if isinstance(msg, (HumanMessage, AIMessage)):
            store.put(
                namespace,
                "",
                {
                    "role": msg.type,
                    "content": msg.content,
                },
            )

    return {}


# =====================================================
# 8️⃣ 构建 Graph
# =====================================================
graph = StateGraph(MessagesState)

graph.add_node("load_memory", load_memory_node)
graph.add_node("llm", llm_node)
graph.add_node("save_memory", save_memory_node)

graph.add_edge(START, "load_memory")
graph.add_edge("load_memory", "llm")
graph.add_edge("llm", "save_memory")
graph.add_edge("save_memory", END)

# =====================================================
# 9️⃣ 使用 Redis Store
# =====================================================
redis_client = redis.from_url(
    "redis://:123456@192.168.174.198:6379/0",
    decode_responses=True,
)
store = RedisTTLStore(
    redis_client=redis_client,
    ttl=60,  # 1小时过期
    max_history=20  # 最大20轮
)

agent = graph.compile(store=store)

# =====================================================
# 🔟 测试运行
# =====================================================
# =====================================================
# 🔟 Human Chat Loop(真正聊天模式)
# =====================================================

config = {"configurable": {"user_id": "lc-325"}}

print("\n===== LangGraph Memory Chat =====")
print("输入 exit 退出\n")

while True:

    user_input = input("👤 你:")

    if user_input.lower() in ["exit", "quit"]:
        print("👋 对话结束")
        break

    state = agent.invoke(
        {"messages": [HumanMessage(content=user_input)]},
        config,
    )

    ai_msg = state["messages"][-1]

    print(f"🤖 AI:{ai_msg.content}\n")
Logo

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

更多推荐