谈结点函数的输入输出之前,有必要了解的是State的设计LangGraph-0x04-StateGraph为什么需要State
通过上面,就已经厘清了一件事情,就是每个结点函数的入参、出参可能不同,但是一定都是State的子集。
那么对于函数的入参而言
- 最少的字段就是0个,也就是函数签名如
def fn():
- 最多的字段就是State的全部字段,那么函数签名就是
def fn(state:State):
其实大部分的实际开发中,都是每个结点函数都有自己的特定的字段,为了解耦各个函数的入参,最简单的就是给每个函数都定义自己的参数类型就行了,然后告诉框架参数类型是什么就行
下面是示例代码
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82
| from typing import TypedDict, Annotated
from langgraph.graph import StateGraph, START, END from langgraph.graph.message import add_messages
class Ctx(TypedDict): mark: Annotated[list[str], add_messages] question: str documents: list[str] answer: str
class RetrieveArg(TypedDict): question: str
def retrieve(state: RetrieveArg) -> dict: r""" 结点1 查询知识库 """ return { "mark": ["结点1"], "documents": ["SOP资料", "RAG资料"] }
class GenerateArg(TypedDict): question: str documents: list[str]
def generate(state: GenerateArg) -> dict: r""" 结点2 调用LLM """ return { "mark": ["结点2"], "answer": "LLM的答案" }
class CheckArg(TypedDict): question: str documents: list[str] answer: str
def check(state: CheckArg) -> dict: r""" 结点3 负责检查答案 """ return { "mark": ["结点3"], "answer": state["answer"] + "->已经审核过了" }
builder = StateGraph(Ctx)
builder.add_node(retrieve, input_schema=RetrieveArg) builder.add_node(generate, input_schema=GenerateArg) builder.add_node(check, input_schema=CheckArg)
builder.add_edge(START, retrieve.__name__) builder.add_edge(retrieve.__name__, generate.__name__) builder.add_edge(generate.__name__, check.__name__) builder.add_edge(check.__name__, END)
graph = builder.compile()
init_state: Ctx = {"question": "this is my question"} result = graph.invoke(init_state)
print(result)
|
在注册结点函数的时候builder.add_node(retrieve, input_schema=RetrieveArg)告诉框架这个函数的入参类型是啥,那么如果没显式告诉框架怎么办呢
当用户没有告诉框架结点函数的参数类型,就要框架自己去推断了LangGraph-0x09-LangGraph如何自动推断结点输入