Implement StateGraph

This commit is contained in:
Nuno Campos
2024-01-12 17:55:19 -08:00
parent b1bf68fde4
commit 426a5c1b10
6 changed files with 456 additions and 139 deletions
+6 -2
View File
@@ -1,5 +1,5 @@
from contextlib import contextmanager
from typing import Callable, Generator, Generic, Optional, Sequence, Type
from typing import Annotated, Callable, Generator, Generic, Optional, Sequence, Type
from typing_extensions import Self
@@ -16,7 +16,11 @@ class BinaryOperatorAggregate(Generic[Value], BaseChannel[Value, Value, Value]):
```
"""
def __init__(self, typ: Type[Value], operator: Callable[[Value, Value], Value]):
def __init__(
self,
typ: Type[Value],
operator: Callable[[Value, Value], Value],
):
self.typ = typ
self.operator = operator
try:
+3 -132
View File
@@ -1,133 +1,4 @@
from asyncio import iscoroutinefunction
from collections import defaultdict
from typing import Any, Callable, Dict, NamedTuple
from langgraph.graph.graph import END, Graph
from langgraph.graph.state import StateGraph
from langchain_core.runnables import Runnable
from langchain_core.runnables.base import (
RunnableLambda,
RunnableLike,
coerce_to_runnable,
)
from langgraph.pregel import Channel, Pregel
END = "__end__"
class Branch(NamedTuple):
condition: Callable[..., str]
ends: dict[str, str]
def runnable(self, input: Any) -> Runnable:
result = self.condition(input)
destination = self.ends[result]
return Channel.write_to(f"{destination}:inbox" if destination != END else END)
class Graph:
def __init__(self) -> None:
self.nodes: dict[str, Runnable] = {}
self.edges = set[tuple[str, str]]()
self.branches: defaultdict[str, list[Branch]] = defaultdict(list)
def add_node(self, key: str, action: RunnableLike) -> None:
if key in self.nodes:
raise ValueError(f"Node `{key}` already present.")
if key == END:
raise ValueError(f"Node `{key}` is reserved.")
self.nodes[key] = coerce_to_runnable(action)
def add_edge(self, start_key: str, end_key: str) -> None:
if start_key == END:
raise ValueError("END cannot be a start node")
if start_key not in self.nodes:
raise ValueError(f"Need to add_node `{start_key}` first")
if end_key not in self.nodes and end_key != END:
raise ValueError(f"Need to add_node `{end_key}` first")
# TODO: support multiple message passing
if start_key in set(start for start, _ in self.edges):
raise ValueError(f"Already found path for {start_key}")
self.edges.add((start_key, end_key))
def add_conditional_edges(
self,
start_key: str,
condition: Callable[..., str],
conditional_edge_mapping: Dict[str, str],
) -> None:
if start_key not in self.nodes:
raise ValueError(f"Need to add_node `{start_key}` first")
if iscoroutinefunction(condition):
raise ValueError("Condition cannot be a coroutine function")
for destination in conditional_edge_mapping.values():
if destination not in self.nodes and destination != END:
raise ValueError(f"Need to add_node `{destination}` first")
self.branches[start_key].append(Branch(condition, conditional_edge_mapping))
def set_entry_point(self, key: str) -> None:
if key not in self.nodes:
raise ValueError(f"Need to add_node `{key}` first")
self.entry_point = key
def set_finish_point(self, key: str) -> None:
return self.add_edge(key, END)
def compile(self) -> Pregel:
################################################
# STEP 1: VALIDATE GRAPH STRUCTURE #
################################################
all_starts = {src for src, _ in self.edges} | {src for src in self.branches}
all_ends = (
{end for _, end in self.edges}
| {
end
for branch_list in self.branches.values()
for branch in branch_list
for end in branch.ends.values()
}
| {self.entry_point}
)
for node in self.nodes:
if node not in all_ends:
raise ValueError(f"Node `{node}` is not reachable")
if node not in all_starts:
raise ValueError(f"Node `{node}` is a dead-end")
################################################
# STEP 2: CREATE GRAPH #
################################################
outgoing_edges = defaultdict(list)
for start, end in self.edges:
outgoing_edges[start].append(f"{end}:inbox" if end != END else END)
nodes = {
key: (Channel.subscribe_to(f"{key}:inbox") | node | Channel.write_to(key))
for key, node in self.nodes.items()
}
for key in self.nodes:
outgoing = outgoing_edges[key]
edges_key = f"{key}:edges"
if outgoing or key in self.branches:
nodes[edges_key] = Channel.subscribe_to(key)
if outgoing:
nodes[edges_key] |= Channel.write_to(*[dest for dest in outgoing])
if key in self.branches:
for branch in self.branches[key]:
nodes[edges_key] |= RunnableLambda(
branch.runnable, name=f"{key}_condition"
)
return Pregel(
nodes=nodes,
input=f"{self.entry_point}:inbox",
output=END,
hidden=[f"{node}:inbox" for node in self.nodes],
)
__all__ = ["END", "Graph", "StateGraph"]
+128
View File
@@ -0,0 +1,128 @@
from asyncio import iscoroutinefunction
from collections import defaultdict
from typing import Any, Callable, Dict, NamedTuple
from langchain_core.runnables import Runnable
from langchain_core.runnables.base import (
RunnableLambda,
RunnableLike,
coerce_to_runnable,
)
from langgraph.pregel import Channel, Pregel
END = "__end__"
class Branch(NamedTuple):
condition: Callable[..., str]
ends: dict[str, str]
def runnable(self, input: Any) -> Runnable:
result = self.condition(input)
destination = self.ends[result]
return Channel.write_to(f"{destination}:inbox" if destination != END else END)
class Graph:
def __init__(self) -> None:
self.nodes: dict[str, Runnable] = {}
self.edges = set[tuple[str, str]]()
self.branches: defaultdict[str, list[Branch]] = defaultdict(list)
def add_node(self, key: str, action: RunnableLike) -> None:
if key in self.nodes:
raise ValueError(f"Node `{key}` already present.")
if key == END:
raise ValueError(f"Node `{key}` is reserved.")
self.nodes[key] = coerce_to_runnable(action)
def add_edge(self, start_key: str, end_key: str) -> None:
if start_key == END:
raise ValueError("END cannot be a start node")
if start_key not in self.nodes:
raise ValueError(f"Need to add_node `{start_key}` first")
if end_key not in self.nodes and end_key != END:
raise ValueError(f"Need to add_node `{end_key}` first")
# TODO: support multiple message passing
if start_key in set(start for start, _ in self.edges):
raise ValueError(f"Already found path for {start_key}")
self.edges.add((start_key, end_key))
def add_conditional_edges(
self,
start_key: str,
condition: Callable[..., str],
conditional_edge_mapping: Dict[str, str],
) -> None:
if start_key not in self.nodes:
raise ValueError(f"Need to add_node `{start_key}` first")
if iscoroutinefunction(condition):
raise ValueError("Condition cannot be a coroutine function")
for destination in conditional_edge_mapping.values():
if destination not in self.nodes and destination != END:
raise ValueError(f"Need to add_node `{destination}` first")
self.branches[start_key].append(Branch(condition, conditional_edge_mapping))
def set_entry_point(self, key: str) -> None:
if key not in self.nodes:
raise ValueError(f"Need to add_node `{key}` first")
self.entry_point = key
def set_finish_point(self, key: str) -> None:
return self.add_edge(key, END)
def validate(self) -> None:
all_starts = {src for src, _ in self.edges} | {src for src in self.branches}
all_ends = (
{end for _, end in self.edges}
| {
end
for branch_list in self.branches.values()
for branch in branch_list
for end in branch.ends.values()
}
| {self.entry_point}
)
for node in self.nodes:
if node not in all_ends:
raise ValueError(f"Node `{node}` is not reachable")
if node not in all_starts:
raise ValueError(f"Node `{node}` is a dead-end")
def compile(self) -> Pregel:
self.validate()
outgoing_edges = defaultdict(list)
for start, end in self.edges:
outgoing_edges[start].append(f"{end}:inbox" if end != END else END)
nodes = {
key: (Channel.subscribe_to(f"{key}:inbox") | node | Channel.write_to(key))
for key, node in self.nodes.items()
}
for key in self.nodes:
outgoing = outgoing_edges[key]
edges_key = f"{key}:edges"
if outgoing or key in self.branches:
nodes[edges_key] = Channel.subscribe_to(key)
if outgoing:
nodes[edges_key] |= Channel.write_to(*[dest for dest in outgoing])
if key in self.branches:
for branch in self.branches[key]:
nodes[edges_key] |= RunnableLambda(
branch.runnable, name=f"{key}_condition"
)
return Pregel(
nodes=nodes,
input=f"{self.entry_point}:inbox",
output=END,
hidden=[f"{node}:inbox" for node in self.nodes],
)
+117
View File
@@ -0,0 +1,117 @@
from collections import defaultdict
from functools import partial
from inspect import signature
from typing import Any, Optional, Type
from langchain_core.runnables import RunnableConfig, RunnableLambda
from langgraph.channels.base import BaseChannel
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels.last_value import LastValue
from langgraph.graph.graph import END, Graph
from langgraph.pregel import Channel, Pregel
from langgraph.pregel.read import ChannelRead
from langgraph.pregel.write import ChannelWrite
START = "__start__"
class StateGraph(Graph):
def __init__(self, schema: Type[Any]) -> None:
super().__init__()
self.schema = schema
self.channels = _get_channels(schema)
def compile(self) -> Pregel:
self.validate()
if any(key in self.nodes for key in self.channels):
raise ValueError("Cannot use channel names as node names")
state_keys = list(self.channels)
outgoing_edges = defaultdict(list)
for start, end in self.edges:
outgoing_edges[start].append(f"{end}:inbox" if end != END else END)
nodes = {
key: (
Channel.subscribe_to(f"{key}:inbox")
| partial(_coerce_state, self.schema) # coerce/validate using schema
| node
| _update_state
| Channel.write_to(key)
)
for key, node in self.nodes.items()
}
for key in self.nodes:
outgoing = outgoing_edges[key]
edges_key = f"{key}:edges"
if outgoing or key in self.branches:
nodes[edges_key] = Channel.subscribe_to(key) | ChannelRead(state_keys)
if outgoing:
nodes[edges_key] |= Channel.write_to(*[dest for dest in outgoing])
if key in self.branches:
for branch in self.branches[key]:
nodes[edges_key] |= RunnableLambda(
branch.runnable, name=f"{key}_condition"
)
nodes[START] = (
Channel.subscribe_to(f"{START}:inbox")
| _update_state
| Channel.write_to(START)
)
nodes[f"{START}:edges"] = (
Channel.subscribe_to(START)
| ChannelRead(state_keys)
| Channel.write_to(f"{self.entry_point}:inbox")
)
return Pregel(
nodes=nodes,
channels=self.channels,
input=f"{START}:inbox",
output=END,
hidden=[f"{node}:inbox" for node in self.nodes] + [START] + state_keys,
)
def _coerce_state(schema: Type[Any], input: dict[str, Any]) -> dict[str, Any]:
return schema(**input)
def _update_state(input: dict[str, Any], config: RunnableConfig):
ChannelWrite.do_write(config, **input)
return input
def _get_channels(schema: Type[dict]) -> dict[str, BaseChannel]:
if not hasattr(schema, "__annotations__"):
raise ValueError("Schema must be a class with type annotations")
channels: dict[str, BaseChannel] = {}
for name, typ in schema.__annotations__.items():
if channel := _is_field_binop(typ):
channels[name] = channel
else:
channels[name] = LastValue(typ)
return channels
def _is_field_binop(typ: Type[Any]) -> Optional[BinaryOperatorAggregate]:
if hasattr(typ, "__metadata__"):
meta = typ.__metadata__
if len(meta) == 1 and callable(meta[0]):
sig = signature(meta[0])
params = list(sig.parameters.values())
if len(params) == 2 and len(
[
p
for p in params
if p.kind in (p.POSITIONAL_ONLY, p.POSITIONAL_OR_KEYWORD)
]
):
return BinaryOperatorAggregate(typ, meta[0])
+12 -4
View File
@@ -23,7 +23,7 @@ from langgraph.constants import CONFIG_KEY_READ
class ChannelRead(RunnableLambda):
channel: str
channel: Union[str, list[str]]
@property
def config_specs(self) -> list[ConfigurableFieldSpec]:
@@ -37,7 +37,7 @@ class ChannelRead(RunnableLambda):
),
]
def __init__(self, channel: str) -> None:
def __init__(self, channel: Union[str, list[str]]) -> None:
super().__init__(func=self._read, afunc=self._aread)
self.channel = channel
self.name = f"ChannelRead<{channel}>"
@@ -50,7 +50,11 @@ class ChannelRead(RunnableLambda):
f"Runnable {self} is not configured with a read function"
"Make sure to call in the context of a Pregel process"
)
return read(self.channel)
return (
read(self.channel)
if isinstance(self.channel, str)
else {chan: read(chan) for chan in self.channel}
)
async def _aread(self, _: Any, config: RunnableConfig) -> Any:
try:
@@ -60,7 +64,11 @@ class ChannelRead(RunnableLambda):
f"Runnable {self} is not configured with a read function"
"Make sure to call in the context of a Pregel process"
)
return read(self.channel)
return (
read(self.channel)
if isinstance(self.channel, str)
else {chan: read(chan) for chan in self.channel}
)
default_bound: RunnablePassthrough = RunnablePassthrough()
+190 -1
View File
@@ -2,7 +2,7 @@ import operator
import time
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from typing import Generator
from typing import Annotated, Generator, TypedDict
import pytest
from langchain_core.runnables import RunnablePassthrough
@@ -15,6 +15,7 @@ from langgraph.channels.last_value import LastValue
from langgraph.channels.topic import Topic
from langgraph.checkpoint.memory import MemorySaver
from langgraph.graph import END, Graph
from langgraph.graph.state import StateGraph
from langgraph.pregel import Channel, Pregel
from langgraph.pregel.reserved import ReservedChannels
@@ -771,3 +772,191 @@ def test_conditional_graph() -> None:
}
},
]
def test_conditional_graph_state() -> None:
from copy import deepcopy
from langchain.llms.fake import FakeStreamingListLLM
from langchain_community.tools import tool
from langchain_core.agents import AgentAction, AgentFinish
from langchain_core.prompts import PromptTemplate
class AgentState(TypedDict):
input: str
agent_outcome: AgentAction | AgentFinish | None
intermediate_steps: Annotated[list[tuple[AgentAction, str]], operator.add]
# Assemble the tools
@tool()
def search_api(query: str) -> str:
"""Searches the API for the query."""
return f"result for {query}"
tools = [search_api]
# Construct the agent
prompt = PromptTemplate.from_template("Hello!")
llm = FakeStreamingListLLM(
responses=[
"tool:search_api:query",
"tool:search_api:another",
"finish:answer",
]
)
def agent_parser(input: str) -> AgentFinish | AgentAction:
if input.startswith("finish"):
_, answer = input.split(":")
return {
"agent_outcome": AgentFinish(
return_values={"answer": answer}, log=input
)
}
else:
_, tool_name, tool_input = input.split(":")
return {
"agent_outcome": AgentAction(
tool=tool_name, tool_input=tool_input, log=input
)
}
agent = prompt | llm | agent_parser
# Define tool execution logic
def execute_tools(data: AgentState) -> dict:
agent_action: AgentAction = data.pop("agent_outcome")
observation = {t.name: t for t in tools}[agent_action.tool].invoke(
agent_action.tool_input
)
return {"intermediate_steps": [(agent_action, observation)]}
# Define decision-making logic
def should_continue(data: AgentState) -> str:
# Logic to decide whether to continue in the loop or exit
if isinstance(data["agent_outcome"], AgentFinish):
return "exit"
else:
return "continue"
# Define a new graph
workflow = StateGraph(AgentState)
workflow.add_node("agent", agent)
workflow.add_node("tools", execute_tools)
workflow.set_entry_point("agent")
workflow.add_conditional_edges(
"agent", should_continue, {"continue": "tools", "exit": END}
)
workflow.add_edge("tools", "agent")
app = workflow.compile()
assert app.invoke({"input": "what is weather in sf"}) == {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
"result for query",
),
(
AgentAction(
tool="search_api",
tool_input="another",
log="tool:search_api:another",
),
"result for another",
),
],
"agent_outcome": AgentFinish(
return_values={"answer": "answer"}, log="finish:answer"
),
}
assert [deepcopy(c) for c in app.stream({"input": "what is weather in sf"})] == [
{
"agent": {
"agent_outcome": AgentAction(
tool="search_api", tool_input="query", log="tool:search_api:query"
),
}
},
{
"tools": {
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
"result for query",
)
],
}
},
{
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="another",
log="tool:search_api:another",
),
}
},
{
"tools": {
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="another",
log="tool:search_api:another",
),
"result for another",
),
],
}
},
{
"agent": {
"agent_outcome": AgentFinish(
return_values={"answer": "answer"}, log="finish:answer"
),
}
},
{
"__end__": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
"result for query",
),
(
AgentAction(
tool="search_api",
tool_input="another",
log="tool:search_api:another",
),
"result for another",
),
],
"agent_outcome": AgentFinish(
return_values={"answer": "answer"}, log="finish:answer"
),
}
},
]