AI创想

标题: Agent之LangGraph [打印本页]

作者: 米落枫    时间: 4 天前
标题: Agent之LangGraph
作者:CSDN博客
一、LangGraph概念

1.1 什么是LangGraph

(, 下载次数: 13)


LangGraph图包含三部分:
LangGraph提供了构建生产级智能体应用的核心能力:
1.2 与LangChain的区别

特性
LangGraph
LangChain
抽象级别
低级,提供细粒度控制
高级,开箱即用
状态管理
内置状态机和检查点
需要自行管理状态
执行模型
基于图的并行执行
线性链式执行
持久化
原生支持
需要额外实现
适用场景
复杂、有状态的智能体应用开发
简单的链式调用
LangGraph在以下场景下,应用更加广泛:
LangGraph官方源码地址及官网地址如下:
二、LangGraph入门案例

2.1 环境准备

创建环境:
  1. conda create –-name=langgraph_venv 3.12
  2. conda activate langgraph_venv
  3. pip install langgraph==1.0.5
  4. # 确认安装成功
  5. pip show langgraph
复制代码
2.2 入门示例

LangGraph代码编写流程

代码示例
  1. from typing import TypedDict
  2. from langgraph.constants import START, END
  3. from langgraph.graph import StateGraph
  4. # 第一步 定义状态  {key:value}
  5. class MyAgentState(TypedDict):
  6.     query:str
  7.     rag_search:str
  8.     web_search:str
  9.     llm_answer:str
  10. # 第二步 定义节点
  11. # rag搜索  {“query”:"....."}
  12. def rag_search_node(state:MyAgentState):
  13.     # state字典根据query这个key获取问题
  14.     query = state["query"]
  15.     # 模拟rag
  16.     rag_search = f"根据用户问题{query},返回rag_search结果"
  17.     return {"rag_search":rag_search}
  18. # web搜索  {“query”:"....."}
  19. def web_search_node(state:MyAgentState):
  20.     query = state["query"]
  21.     web_search = f"根据用户问题{query},返回web_search结果"
  22.     return {"web_search": web_search}
  23. # llm整合数据节点  {“query”:"....."}
  24. def llm_node(state:MyAgentState):
  25.     query = state["query"]
  26.     rag_search = state["rag_search"]
  27.     web_search = state["web_search"]
  28.     # 模拟
  29.     llm_answer = f"用户问题:{query},LLM基于{rag_search}和{web_search}的最终回复"
  30.     return {"llm_answer": llm_answer}
  31. # 第三步 构建builder,添加节点和边
  32. builder = StateGraph(state_schema=MyAgentState)
  33. # 添加节点
  34. builder.add_node(rag_search_node) # "rag_search_node":rag_search_node
  35. builder.add_node(web_search_node)
  36. builder.add_node(llm_node)
  37. # 添加边
  38. builder.add_edge(START, "rag_search_node")
  39. builder.add_edge(START, "web_search_node")
  40. builder.add_edge("rag_search_node", "llm_node")
  41. builder.add_edge("web_search_node", "llm_node")
  42. builder.add_edge("llm_node", END)
  43. # 第四步 编译图
  44. graph = builder.compile()
  45. # 第五步 调用图
  46. res = graph.invoke({"query":"如何使用LangGraph"})
  47. print(res)
复制代码
特殊说明

三、LangGraph状态

        通过LangGraph构建应用程序时,第一步就是定义State,State代表了整个图中 节点 的状态数据,以及应用最终结果和目标。
3.1 状态的定义

State Schema可以通过三种方式定义:
第一种 继承TypedDict类实现(推荐)
  1. # TypedDict是Python提供的一种类型提示工具,用于为字典(Dict)的键和值指定精确的类型信息。
  2. from typing import TypedDict
  3. class MyAgentState(TypedDict):
  4.     query:str
  5.     rag_search:str
  6.     web_search:str
  7.     llm_answer:str
复制代码
第二种 继承BaseModel类实现
  1. # Pydantic BaseModel:Pydantic 提供运行时数据校验,并支持静态类型检查工具进行类型推导。状态类通过继承Pydantic的BaseModel,定义键的类型和reducer函数;
  2. from pydantic import BaseModel
  3. class MyStateFull(BaseModel):
  4.     rag_result:str
  5.     web_search_result:str
  6.     query:str
复制代码
第三种 使用python装饰器 dataclass
  1. # Dataclass: dataclass是Python标准库中的一个装饰器,用于自动生成常见特殊方法(如__init__、__repr__、__eq__等),从而简化主要用作数据容器的类的定义。状态类通过dataclass装饰器装饰后,定义键的类型和reducer函数即可。
  2. from dataclasses import dataclass
  3. @dataclass
  4. class MyStateFullTwo():
  5.     rag_result:str
  6.     web_search_result:str
  7.     final_answer:str
  8.     query:str
复制代码
3.2 输入输出隔离

在LangGraph 当中,可以精细管理输入到图中的状态键有哪些,以及输出的状态键有哪些。这是通过初始化StateGraph时,分别指定三个参数:state_schema、input_schema 和 output_schema 来实现的。
  1. from typing import TypedDict
  2. from langgraph.constants import START, END
  3. from langgraph.graph import StateGraph
  4. # 1 定义全局状态
  5. class MyAgentState(TypedDict):
  6.     query:str
  7.     rag_search:str
  8.     web_search:str
  9.     llm_answer:str
  10. # 2 定义输入 和 输出状态类
  11. class InputSchema(TypedDict):
  12.     query:str
  13. class OutputSchema(TypedDict):
  14.     llm_answer: str
  15. # 3 定义节点
  16. def rag_search_node(state:MyAgentState):
  17.     # state字典根据query这个key获取问题
  18.     query = state["query"]
  19.     # 模拟rag
  20.     rag_search = f"根据用户问题{query},返回rag_search结果"
  21.     # 返回数据,langgraph封装自动把返回数据更新到传输状态里面
  22.     return {"rag_search":rag_search}
  23. def web_search_node(state:MyAgentState):
  24.     query = state["query"]
  25.     web_search = f"根据用户问题{query},返回web_search结果"
  26.     return {"web_search": web_search}
  27. def llm_node(state:MyAgentState):
  28.     query = state["query"]
  29.     rag_search = state["rag_search"]
  30.     web_search = state["web_search"]
  31.     # 模拟
  32.     llm_answer = f"用户问题:{query},LLM基于{rag_search}和{web_search}的最终回复"
  33.     # 返回
  34.     return {"llm_answer": llm_answer}
  35. # 4 创建builder,添加节点和边
  36. builder = StateGraph(state_schema=MyAgentState,
  37.            input_schema=InputSchema,
  38.            output_schema=OutputSchema)
  39. # 添加节点
  40. builder.add_node(rag_search_node) # "rag_search_node":rag_search_node
  41. builder.add_node(web_search_node)
  42. builder.add_node(llm_node)
  43. # 添加边
  44. builder.add_edge(START, "rag_search_node")
  45. builder.add_edge(START, "web_search_node")
  46. builder.add_edge("rag_search_node", "llm_node")
  47. builder.add_edge("web_search_node", "llm_node")
  48. builder.add_edge("llm_node", END)
  49. # 第四步 编译图
  50. graph = builder.compile()
  51. # 第五步 调用图
  52. res = graph.invoke({"query":"如何使用LangGraph","rag_search":"xxxxxx"})
  53. print(res)
复制代码
3.3 Reducer函数

        Reducer用于进行当前增量状态(节点输出的状态)和全局状态的合并。State中的每个键都有其独立的reducer函数。每个node的返回值中的每个key与全局state_schema中对应的key进行合并更新,具体更新逻辑取决于每个key指定的reducer函数。
Reducer常用函数有以下几种:
  1. from operator import add
  2. from typing import TypedDict, Annotated, List
  3. from langgraph.constants import START, END
  4. from langgraph.graph import StateGraph
  5. # 1 定义全局状态
  6. class MyAgentState(TypedDict):
  7.     """
  8.         每个状态类key有默认reducer,函数默认效果是覆盖,新值覆盖旧值
  9.         效果:通过一个key改变reducer逻辑,把覆盖效果-》合并效果
  10.     """
  11.     query:str
  12.     rag_search:str
  13.     web_search:str
  14.     llm_answer:str
  15.     # reducer   operator.add(1,2)=3   operator.add([k1],[k2])=[k1,k2]合并
  16.     test_key:Annotated[List[str],add]
  17. # 第二步 定义节点
  18. def rag_search_node(state:MyAgentState):
  19.     query = state["query"]
  20.     # 模拟rag
  21.     rag_search = f"根据用户问题{query},返回rag_search结果"
  22.     # 返回数据,langgraph封装自动把返回数据更新到传输状态里面
  23.     return {"rag_search":rag_search,"test_key":["test_key_rag"]}
  24. def web_search_node(state:MyAgentState):
  25.     query = state["query"]
  26.     web_search = f"根据用户问题{query},返回web_search结果"
  27.     return {"web_search": web_search,"test_key":["test_key_web"]}
  28. def llm_node(state:MyAgentState):
  29.     query = state["query"]
  30.     rag_search = state["rag_search"]
  31.     web_search = state["web_search"]
  32.     llm_answer = f"用户问题:{query},LLM基于{rag_search}和{web_search}的最终回复"
  33.     return {"llm_answer": llm_answer}
  34. # 第三步 构建builder,添加节点和边
  35. builder = StateGraph(state_schema=MyAgentState)
  36. # 添加节点
  37. builder.add_node(rag_search_node) # "rag_search_node":rag_search_node
  38. builder.add_node(web_search_node)
  39. builder.add_node(llm_node)
  40. # 添加边
  41. builder.add_edge(START, "rag_search_node")
  42. builder.add_edge(START, "web_search_node")
  43. builder.add_edge("rag_search_node", "llm_node")
  44. builder.add_edge("web_search_node", "llm_node")
  45. builder.add_edge("llm_node", END)
  46. # 第四步 编译图
  47. graph = builder.compile()
  48. # 第五步 调用图
  49. res = graph.invoke({"query":"如何使用LangGraph"})
  50. print(res)
复制代码
自定义Reducer函数:
  1. from typing import List, TypedDict, Annotated
  2. from langchain_core.messages import AnyMessage, AIMessage, ToolMessage, HumanMessage
  3. from langgraph.constants import START
  4. from langgraph.graph import StateGraph
  5. # 创建操作方法:合并  AIMessage  HuManMessage
  6. def add_message(message_list_left:List[AnyMessage],
  7.               message_list_right:List[AnyMessage])->List[AnyMessage]:
  8.     print("my_add_message_reducer函数的调用:")
  9.     print("message_list_left:", message_list_left)
  10.     print("message_list_right:", message_list_right)
  11.     new_message_list = message_list_left + message_list_right
  12.     return new_message_list
  13. # 创建状态类
  14. class MyAgentState(TypedDict):
  15.     messages:Annotated[List[AnyMessage],add_message]
  16.    
  17. # 创建节点
  18. def llm_node(state:MyAgentState):
  19.     ai_message = AIMessage(content="xxxx")
  20.     return {"messages":ai_message}
  21. def tool_node(state:MyAgentState):
  22.     tool_message = ToolMessage(content="xxxx", tool_call_id='xx')
  23.     return {"messages":tool_message}
  24. # 创建builder
  25. builder = StateGraph(state_schema=MyAgentState)
  26. builder.add_node(llm_node)
  27. builder.add_node(tool_node)
  28. builder.add_edge(START,'llm_node')
  29. builder.add_edge('llm_node',
  30.                    'tool_node')
  31. graph = builder.compile()
  32. res = graph.invoke(
  33.     {"messages":
  34.      [HumanMessage(content = "你好")]})
  35. print(res)
复制代码
3.4 状态存储

       在前面的例子当中,我们看到,每次用户invoke时,LangGraph都会初始化一个空状态,然后将用户传入的初始状态合并进来,再继续往下执行。这在一些一次性、简单任务过程中,没有什么问题。但是对于一些复杂任务,就会出现一些问题,考虑以下两个场景:
1)场景一:在agent的一个会话里,需要在多次调用当中保持上下文
2)场景二:图执行过程当中报错,不想重复已执行完节点,想要实现断点续传
        (1)invoke时传递None作为初始参数;
        (2)传入相同的thread_id。
  1. # 从状态中恢复执行
  2. # 某个节点执行出错,从出错节点恢复执行,已经执行过不再执行
  3. import sqlite3
  4. from typing import TypedDict
  5. from langgraph.checkpoint.sqlite import SqliteSaver
  6. from langgraph.constants import START
  7. from langgraph.graph import StateGraph
  8. # 创建状态
  9. class ResumeState(TypedDict):
  10.     key_1:str
  11.     key_2:str
  12.     key_3:str
  13. # 创建节点
  14. def node_1(state:ResumeState):
  15.     print(state)
  16.     print("node1节点调用了")
  17.     return {"key_1":"value_1"}
  18. def node_2(state:ResumeState):
  19.     print(state)
  20.     print("node2节点调用了")
  21.     #raise Exception("node2节点出错了....")
  22.     return {"key_2":"value_2"}
  23. def node_3(state:ResumeState):
  24.     print(state)
  25.     print("node3节点调用了")
  26.     return {"key_3":"value_3"}
  27. # 创建builder,添加节点和边
  28. builder = StateGraph(state_schema=ResumeState)
  29. builder.add_node(node_1)
  30. builder.add_node(node_2)
  31. builder.add_node(node_3)
  32. builder.add_edge(START, "node_1")
  33. builder.add_edge("node_1", "node_2")
  34. builder.add_edge("node_2", "node_3")
  35. # 编译和执行
  36. checkpointer = SqliteSaver(
  37.     conn=sqlite3.connect("./resume_demo.db",
  38.                          check_same_thread=False),
  39. )
  40. graph = builder.compile(checkpointer=checkpointer)
  41. # 第一次调用:node2节点出错
  42. # res = graph.invoke({},
  43. #    config={"configurable":
  44. #         {"thread_id":"2"}})
  45. # print(res)
  46. # 第二次调用:从node2恢复执行
  47. # 要想从故障中恢复,需要以下两点:
  48. # (1)invoke时传递None作为初始参数;
  49. # (2)传入相同的thread_id。
  50. res = graph.invoke(None,
  51.    config={"configurable":
  52.         {"thread_id":"2"}})
  53. print(res)
复制代码
3.5 LangGraph底层运行算法

        在了解如何获取历史状态之前,首先我们需要了解一下LangGraph底层算法—Pregel。
        Pregel 是LangGraph底层的一个类,管理 LangGraph 应用程序的运行时(runtime)行为。也就是说,整个图结构,从开始到结束的迭代执行过程,是由Pregel控制管理的。
在这个类当中,有两大主要组件:
  1. from typing import TypedDict,Annotated,List
  2. from operator import add
  3. from langgraph.graph import StateGraph
  4. from langgraph.constants import START
  5. class MyState(TypedDict):
  6.     aggregates:Annotated[List[str],add]
  7. def node_a(state:MyState):
  8.     return {"aggregates":["a"]}
  9. def node_b(state:MyState):
  10.     return {"aggregates":["b"]}
  11. def node_c(state:MyState):
  12.     return {"aggregates":["c"]}
  13. def node_b_2(state:MyState):
  14.     return {"aggregates":["b_2"]}
  15. def node_d(state:MyState):
  16.     return {"aggregates":["d"]}
  17. builder = StateGraph(MyState)
  18. builder.add_node(node_a)
  19. builder.add_node(node_b)
  20. builder.add_node(node_c)
  21. builder.add_node(node_b_2)
  22. builder.add_node(node_d)
  23. builder.add_edge(START,"node_a")
  24. builder.add_edge("node_a","node_b")
  25. builder.add_edge("node_a","node_c")
  26. builder.add_edge("node_b","node_b_2")
  27. builder.add_edge("node_b_2","node_d")
  28. builder.add_edge("node_c","node_d")
  29. graph = builder.compile()
  30. print('graph当中的nodes',graph.nodes)
  31. print('graph当中的通道channels',graph.channels)
  32. res = graph.invoke({})
  33. print(res)
复制代码
3.6 获取图执行的历史状态

        构建好的graph实例,可以通过get_state() / get_state_history()方法,传入想要获取的历史状态的thread_id,即可拿到图执行过程中的历史状态,get_state()方法获取到的是最近一个时间步的状态,get_state_history()获取到的是图执行当中所有时间步的历史状态,其输出值为一个迭代器,按照时间步倒序排列。
  1. # ....................
  2. # 图编译
  3. graph = builder.compile(checkpointer=InMemorySaver())
  4. # 图调用
  5. res = graph.invoke({},
  6.              config={"configurable":
  7.         {"thread_id":"1"}})
  8. print(res)
  9. print("="*50)
  10. # 获取图执行的历史状态
  11. history_states = graph.get_state_history(config={'configurable':{'thread_id':'1'}})
  12. for state in history_states:
  13.     print(state)
  14. print("="*50)
  15. # 获取最近一次
  16. state = graph.get_state(config={'configurable':{'thread_id':'1'}})
  17. print(state)
复制代码
四、节点

4.1 节点的输入输出

        在LangGraph中,一般来讲,节点都是Python函数(可以是同步的,也可以是异步的),它们接受以下参数:
       以上参数,会在运行过程中,自动被LangGraph运行时注入。
       节点的输出为当前节点对状态的增量更新,而不能将接收到的整个状态实例返回出去。
        LangGraph底层会将所有节点输出的状态,都作为增量状态,并尝试和当前的全局状态做一次合并操作。如果将整个状态都输出,对于没有配置任何reducer的状态键,langgraph底层无法完成合并,会抛出异常;对于配置了reducer的状态键,如果节点输出了不属于该节点更新的状态,也会导致数据产生问题。
  1. from typing import List, TypedDict
  2. from langchain_core.runnables import RunnableConfig
  3. from langgraph.constants import START
  4. from langgraph.graph import StateGraph
  5. from langgraph.runtime import Runtime
  6. # 定义状态
  7. class CustomerState(TypedDict):
  8.     query: str          # 用户问题
  9.     response: str       # 客服回复
  10.     log: List[str]      # 处理日志
  11. # 定义节点
  12. # state:状态 业务数据
  13. # runtime:上下文对象,比如连接对象db
  14. # config:配置信息,比如thread_id
  15. def node_customer_service(state:CustomerState,
  16.                           runtime: Runtime,
  17.                           config: RunnableConfig):
  18.     print("传入的状态为,", state)
  19.     print("传入的配置为,", config)
  20.     print("传入的runtime为,", runtime)
  21. # 创建buildler
  22. builder = StateGraph(CustomerState)
  23. builder.add_node(node_customer_service)
  24. builder.add_edge(START,"node_customer_service")
  25. # 图编译
  26. graph = builder.compile()
  27. # 图执行
  28. # 三个参数: input   config   context
  29. config = {
  30.     "user_id":"123_vip",
  31.     "configurable":{"thread_id":"1234"}
  32. }
  33. context={
  34.     "db":"mysqldb",
  35.     "llm":"llm_result"
  36. }
  37. res = graph.invoke({"query":"hello"},config=config,context=context)
  38. print(res)
复制代码
4.2 特殊节点

        START和END节点都是LangGraph当中的特殊节点。START它代表着将用户输入发送到图中的节点。引用此节点的主要目的是确定应首先调用哪些节点。END节点是一个特殊节点,代表终止节点。当想表示哪些边在完成后没有动作时,会引用这个节点。
START和END节点本质上是一个字符串,如下所示:
  1. END = sys.intern("__end__")
  2. """The last (maybe virtual) node in graph-style Pregel."""
  3. START = sys.intern("__start__")
  4. """The first (maybe virtual) node in graph-style Pregel."""
复制代码
START作为一个特殊节点,在图执行完START节点之后,也会创建一个checkpoint。
注:END节点可不写,LangGraph默认会自动添加。START节点对应的边必须进行指定。
4.3 节点缓存

        LangGraph支持基于节点输入对节点进行缓存。对于配置了缓存的节点,且缓存结果没有过期,以相同的输入再次调用节点时,可直接从缓存当中读取结果,不需要再进行节点计算。
使用缓存的方法如下:
  1. import time
  2. from typing import TypedDict
  3. from langgraph.cache.memory import InMemoryCache
  4. from langgraph.constants import START
  5. from langgraph.graph import StateGraph
  6. from langgraph.types import CachePolicy
  7. # 需求: 根据user_id生成order_id
  8. # 定义状态
  9. class MyState(TypedDict):
  10.     user_id:str
  11.     order_id:str
  12. # 定义节点
  13. def get_order_id(state:MyState):
  14.     print("调用get_order_id节点.....")
  15.     # 获取user_id
  16.     user_id = state["user_id"]
  17.     # 根据用户id生成order_id
  18.     order_id = f"order_id_{user_id}"
  19.     return {"order_id": order_id}
  20. # 创建builder
  21. builder = StateGraph(state_schema=MyState)
  22. # 添加节点
  23. # 设置缓存策略
  24. builder.add_node(get_order_id,
  25.            cache_policy=CachePolicy(ttl=3))
  26. # 添加边
  27. builder.add_edge(START,"get_order_id")
  28. # 图编译
  29. # 缓存存储方式
  30. cache = InMemoryCache()
  31. graph = builder.compile(cache=cache)
  32. # 图执行
  33. res = graph.invoke({"user_id":"123"})
  34. print(res)
  35. print("=="*50)
  36. res = graph.invoke({"user_id":"123"})
  37. print(res)
  38. print("=="*50)
  39. time.sleep(5)
  40. res = graph.invoke({"user_id":"123"})
  41. print(res)
复制代码
4.4 节点重试

        在很多使用场景中,有些节点由于客观原因限制,导致其执行过程是不稳定的。因此,我们可能希望节点拥有自定义的重试策略,例如在调用API、查询数据库或调用大语言模型(LLM)等情况下。
        为节点添加重试策略,需要在add_node中设置retry_policy参数。retry_policy参数接受一个RetryPolicy命名元组对象。RetryPolicy对象有两个属性:max_attempts和retry_on参数,前者定义了总共重试次数,后者定义了对于哪些异常类型进行重试。
         重试策略为:初始重试时间为0.5秒,幂指数级别更新重试间隔时长,直到达到最大重试次数(默认3次),默认开启1s内的重试时长扰动。
  1. import datetime
  2. from typing import TypedDict
  3. from langgraph.constants import START
  4. from langgraph.graph import StateGraph
  5. from langgraph.types import RetryPolicy
  6. # 重试
  7. attempt = 0
  8. class MyAgentState(TypedDict):
  9.     llm_message: str
  10. def llm_node(state: MyAgentState):
  11.     global attempt
  12.     print(datetime.datetime.now())
  13.     attempt += 1
  14.     if attempt < 10:
  15.         print(f'第{attempt}次调用llm失败')
  16.         raise ConnectionError('调用llm失败')
  17.     print(f'第{attempt}次调用llm成功')
  18.     return {'llm_message': '调用llm成功'}
  19. builder = StateGraph(state_schema=MyAgentState)
  20. builder.add_node(llm_node, retry_policy=RetryPolicy(
  21.     max_attempts=10,        # 最多尝试 10 次(含首次)
  22.     # initial_interval=0.5,   # 第一次重试前等待 0.5 秒
  23.     # backoff_factor=2.0,     # 每次间隔翻倍:0.5 -> 1 -> 2 -> 4 ...
  24.     # max_interval=128.0,      # 间隔上限 10 秒
  25.     # jitter=False,           # 关闭随机抖动,方便观察间隔规律
  26.     # retry_on=(ConnectionError,),  # 只对 ConnectionError 重试
  27. ))
  28. builder.add_edge(START, 'llm_node')
  29. graph = builder.compile()
  30. res = graph.invoke({})
  31. print(res)
复制代码
4.5 图内外数据传递

        默认情况下,通过graph.invoke调用图时,仅在整个图的执行过程都结束之后,我们才能够拿到最终的状态,那么如果想要在图的执行过程当中,想要获取到图内部所产生的数据,应该如何实现?
        考虑Agent的一个场景:首先,我们希望在LLM输出时,就能够拿到LLM所产生的token,在前端进行展示;其次,如果Agent的流程较长,我们希望能够拿到,当前正在执行的节点或者流程是什么。
        由于LangGraph实现了langchain_core当中的Runnable接口,其为我们提供了stream和astream方法,也即流式输出方法,通过流式输出,就能解决前面所说到的问题。
流式输出有下面常见模式:
1. values:每一个执行后,流式输出完整的状态
2. messages:在任何调用了LLM的节点当中,流式输出两元组数据:(LLM Token,metadata)
messages模式代码实现:
  1. load_dotenv()
  2. llm = init_chat_model(
  3.     model="qwen-plus",
  4.     model_provider="openai",
  5.     base_url="https://dashscope.aliyuncs.com/compatible-mode/v1",
  6.     # 千问API Key
  7.     api_key=os.getenv("OPENAI_API_KEY"),
  8. )
  9. def node_stream_messages():
  10.     # 状态
  11.     class MyAgentState(TypedDict):
  12.         query:str
  13.         llm_message:str
  14.     # 节点
  15.     def llm_node(state:MyAgentState):
  16.         query = state["query"]
  17.         response = llm.invoke(query)
  18.         return {"llm_message": response.content}
  19.     # 创建builder
  20.     builder = StateGraph(MyAgentState)
  21.     builder.add_node(llm_node)
  22.     builder.add_edge(START, "llm_node")
  23.     # 编译
  24.     graph = builder.compile()
  25.     # 流式调用
  26.     res = graph.stream(
  27.         {"query": "什么是langgraph"},
  28.         stream_mode="messages",
  29.     )
  30.     # 流式输出两元组数据:(LLM Token,metadata)
  31.     for chunk,metadata in res:
  32.         # print(metadata)
  33.         print(chunk.content,end="")
  34. node_stream_messages()
复制代码
4.6 人工审核节点

        Agent在完成用户设定的任务时,有时候我们希望用户参与当中部分重要决策过程。LangGraph为此提供了一个非常方便的原语:interrupt。
  1. from typing import TypedDict
  2. from langgraph.checkpoint.memory import InMemorySaver
  3. from langgraph.graph import StateGraph
  4. from langgraph.types import interrupt, Command
  5. # 状态
  6. class MyAgentState(TypedDict):
  7.     query:str
  8.     llm_message:str
  9.     rag_message:str
  10. # 节点
  11. def rag_node(state:MyAgentState):
  12.     print("rag node调用....")
  13.     query = state["query"]
  14.     return {"query":query}
  15. def llm_node(state:MyAgentState):
  16.     print("llm node调用....")
  17.     query = state["query"]
  18.     # 拦截中断
  19.     result = interrupt({
  20.         "query": query,
  21.         "msg": "是否调用大语言模型",
  22.     })
  23.     # 模拟判断
  24.     if result:
  25.         rag_message ="llm大语言模型调用了..."
  26.     else:
  27.         rag_message ="不可用调用大语言模型llm..."
  28.     return {"rag_message": rag_message}
  29. builder = StateGraph(MyAgentState)
  30. builder.add_node(rag_node)
  31. builder.add_node(llm_node)
  32. builder.add_edge("__start__",
  33.                  "rag_node")
  34. builder.add_edge("rag_node",
  35.                  "llm_node")
  36. checkpointer = InMemorySaver()
  37. # 使用interrupt,需要给graph,配置一个checkpoint
  38. graph=builder.compile(checkpointer=checkpointer)
  39. res = graph.invoke({"query":"什么是langgraph"},
  40.          config={"configurable":{"thread_id":"1"}})
  41. print(res)
  42. # {
  43. #     'query': '什么是langgraph',
  44. #     '__interrupt__': [
  45. #         Interrupt(value={
  46. #             'query': '什么是langgraph',
  47. #             'msg': '是否调用大语言模型'
  48. #         },
  49. #         id='e070c7acab89ba244cd4eae235e4be01')
  50. #     ]
  51. # }
  52. # 从中断信息获取数据
  53. to_review_data=res["__interrupt__"][0].value
  54. print(to_review_data)
  55. # 业务处理
  56. # 审核通过,继续往后执行
  57. res = graph.invoke(Command(resume=True),
  58.      config={"configurable":{"thread_id":"1"}})
  59. print(res)
复制代码
五、边

5.1 边的本质

        定义边,本质上是定义了,Pregel当中节点订阅的状态通道和节点执行之后,需要更新的状态通道。
边有几种关键类型:
5.2 条件边

        条件边的本质是一个路由函数,其根据当前状态动态决定下一个要执行的节点。
  1. from typing import TypedDict
  2. from langgraph.constants import END
  3. from langgraph.graph import StateGraph
  4. class MyAgentState(TypedDict):
  5.     query: str
  6.     number:int
  7.     node:str
  8.     rag_message: str
  9. def node_a(state:MyAgentState):
  10.     print("正在执行node_a")
  11.     return {"node":"node_a"}
  12. def node_b(state:MyAgentState):
  13.     print("正在执行node_b")
  14.     return {"node":"node_b"}
  15. #函数,结合当前状态,判断下一个节点
  16. def condtional_function(state:MyAgentState):
  17.     number = state["number"]
  18.     if number % 2 == 0:
  19.         return "node_b"
  20.     else:
  21.         return END
  22. builder = StateGraph(MyAgentState)
  23. builder.add_node(node_a)
  24. builder.add_node(node_b)
  25. # START -> node_a  (条件)->  node_b
  26. builder.add_edge("__start__",
  27.                  "node_a")
  28. # node_a -> node_b 添加条件边
  29. builder.add_conditional_edges(
  30.     source="node_a",
  31.     path=condtional_function
  32. )
  33. graph = builder.compile()
  34. graph.invoke({"number":4})
复制代码
5.3 可控循环边

        通过条件边,我们可以构建带有循环结构的图,例如在典型的REACT模式下,工具调用和大模型总结生成结果形成了一个循环,如下图所示:
(, 下载次数: 15)


        需要注意的是,这种带循环的图结构,有一个隐藏的问题:图执行过程当中,可能因为某些原因,导致一直在循环内循环往复执行,因此LangGraph也提供了一个强制使图的执行终止的递归限制参数。
        递归限制设定了图在抛出错误之前允许执行的超级步骤数量,默认值25,在graph.invoke的config参数中指定。在经过指定数量的超级步骤后,图还没有自然停止执行时,LangGraph会抛出异常GraphRecursionError。
       示例代码如下所示:
  1. from typing import Annotated, Dict, Literal
  2. from typing_extensions import TypedDict
  3. from langgraph.graph import StateGraph, START, END
  4. from langgraph.errors import GraphRecursionError
  5. class LoopState(TypedDict):
  6.     count: int
  7.     result: str
  8.     max_count: int
  9. def node_a(state: LoopState) -> dict:
  10.     """节点a:处理逻辑并更新计数"""
  11.     print(f"执行节点a,当前计数: {state['count']}")
  12.     return {
  13.         'count': state['count'] + 1,
  14.         'result': f"已处理{state['count']}次"
  15.     }
  16. def node_b(state: LoopState) -> dict:
  17.     """节点b:辅助处理"""
  18.     print(f"执行节点b,当前计数: {state['count']}")
  19.     return {
  20.         'result': f"已处理{state['count']}次 - 辅助处理"
  21.     }
  22. def route(state: LoopState) -> Literal["b", END]:
  23.     """条件路由函数:决定是继续循环还是终止"""
  24.     # 终止条件:当计数达到最大值时终止
  25.     if state['count'] >= state['max_count']:
  26.         print(f"满足终止条件,计数 {state['count']} >= {state['max_count']},返回END")
  27.         return END
  28.     else:
  29.         print(f"未满足终止条件,计数 {state['count']} < {state['max_count']},返回b")
  30.         return "b"
  31. # 创建图
  32. builder = StateGraph(LoopState)
  33. # 添加节点
  34. builder.add_node("a", node_a)
  35. builder.add_node("b", node_b)
  36. # 添加边
  37. builder.add_edge(START, "a")
  38. builder.add_conditional_edges("a", route)
  39. builder.add_edge("b", "a")
  40. # 编译图
  41. graph = builder.compile()
  42. # 执行图
  43. print("=== 开始执行工作流 ===")
  44. try:
  45.     result = graph.invoke(input={
  46.         'count': 0,
  47.         'result': '',
  48.         'max_count': 300
  49.     }, config={
  50.         'recursion_limit': 6  # 设置递归限制
  51.     })
  52.     print("=== 执行结果 ===")
  53.     print(result)
  54. except GraphRecursionError as e:
  55. print(f"递归错
复制代码
原文地址:https://blog.csdn.net/weixin_42796403/article/details/162883458




欢迎光临 AI创想 (https://llms-ai.com/) Powered by Discuz! X3.4