diff --git a/langgraph/channels/binop.py b/langgraph/channels/binop.py index 6ec62fcb9..962ce7140 100644 --- a/langgraph/channels/binop.py +++ b/langgraph/channels/binop.py @@ -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: diff --git a/langgraph/graph/__init__.py b/langgraph/graph/__init__.py index 4d5393c5a..1173035b5 100644 --- a/langgraph/graph/__init__.py +++ b/langgraph/graph/__init__.py @@ -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"] diff --git a/langgraph/graph/graph.py b/langgraph/graph/graph.py new file mode 100644 index 000000000..4caac7a16 --- /dev/null +++ b/langgraph/graph/graph.py @@ -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], + ) diff --git a/langgraph/graph/state.py b/langgraph/graph/state.py new file mode 100644 index 000000000..c3f58edb6 --- /dev/null +++ b/langgraph/graph/state.py @@ -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]) diff --git a/langgraph/pregel/read.py b/langgraph/pregel/read.py index 33fadf46c..11c1a050b 100644 --- a/langgraph/pregel/read.py +++ b/langgraph/pregel/read.py @@ -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() diff --git a/tests/test_pregel.py b/tests/test_pregel.py index cccb64fce..a56e641f2 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -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" + ), + } + }, + ]