diff --git a/docs/docs/concepts/pregel.md b/docs/docs/concepts/pregel.md index eccd469cb..039ba1226 100644 --- a/docs/docs/concepts/pregel.md +++ b/docs/docs/concepts/pregel.md @@ -27,7 +27,7 @@ Repeat until no **actors** are selected for execution, or a maximum number of st ## Actors -An **actor** is a [PregelNode][langgraph.pregel.read.PregelNode]. It subscribes to channels, reads data from them, and writes data to them. It can be thought of as an **actor** in the Pregel algorithm. [PregelNodes][langgraph.pregel.read.PregelNode] implement LangChain's Runnable interface. +An **actor** is a `PregelNode`. It subscribes to channels, reads data from them, and writes data to them. It can be thought of as an **actor** in the Pregel algorithm. `PregelNodes` implement LangChain's Runnable interface. ## Channels diff --git a/docs/docs/reference/agents.md b/docs/docs/reference/agents.md index 29cdc050b..2893d1906 100644 --- a/docs/docs/reference/agents.md +++ b/docs/docs/reference/agents.md @@ -6,19 +6,29 @@ - AgentState - create_react_agent +::: langgraph.prebuilt.tool_node.ToolNode + options: + show_if_no_docstring: true + show_root_heading: true + show_root_full_path: false + inherited_members: false + members: + - inject_tool_args + ::: langgraph.prebuilt.tool_node options: members: - - ToolNode - InjectedState - InjectedStore - tools_condition -::: langgraph.prebuilt.tool_validator +::: langgraph.prebuilt.tool_validator.ValidationNode options: - members: - - ValidationNode - + show_if_no_docstring: true + show_root_heading: true + show_root_full_path: false + inherited_members: false + members: false ::: langgraph.prebuilt.interrupt options: diff --git a/docs/docs/reference/checkpoints.md b/docs/docs/reference/checkpoints.md index dcac4c27a..12ed304fd 100644 --- a/docs/docs/reference/checkpoints.md +++ b/docs/docs/reference/checkpoints.md @@ -25,5 +25,11 @@ ::: langgraph.checkpoint.sqlite.aio ::: langgraph.checkpoint.postgres + options: + members: + - PostgresSaver -::: langgraph.checkpoint.postgres.aio \ No newline at end of file +::: langgraph.checkpoint.postgres.aio + options: + members: + - AsyncPostgresSaver \ No newline at end of file diff --git a/docs/docs/reference/graphs.md b/docs/docs/reference/graphs.md index c67e2136a..53ecc68db 100644 --- a/docs/docs/reference/graphs.md +++ b/docs/docs/reference/graphs.md @@ -1,16 +1,75 @@ # Graph Definitions -::: langgraph.graph.graph +::: langgraph.graph.state.StateGraph options: + show_if_no_docstring: true + show_root_heading: true + show_root_full_path: false members: - - Graph - - CompiledGraph + - add_node + - add_edge + - add_conditional_edges + - add_sequence + - compile -::: langgraph.graph.state +::: langgraph.graph.state.CompiledStateGraph options: + show_if_no_docstring: true + show_root_heading: true + show_root_full_path: false members: - - StateGraph - - CompiledStateGraph + - stream + - astream + - invoke + - ainvoke + - get_state + - aget_state + - get_state_history + - aget_state_history + - update_state + - aupdate_state + - bulk_update_state + - abulk_update_state + - get_graph + - aget_graph + - get_subgraphs + - aget_subgraphs + - with_config + +::: langgraph.graph.graph.Graph + options: + show_if_no_docstring: true + show_root_heading: true + show_root_full_path: false + members: + - add_node + - add_edge + - add_conditional_edges + - compile + +::: langgraph.graph.graph.CompiledGraph + options: + show_if_no_docstring: true + show_root_heading: true + show_root_full_path: false + members: + - stream + - astream + - invoke + - ainvoke + - get_state + - aget_state + - get_state_history + - aget_state_history + - update_state + - aupdate_state + - bulk_update_state + - abulk_update_state + - get_graph + - aget_graph + - get_subgraphs + - aget_subgraphs + - with_config ::: langgraph.graph.message options: diff --git a/docs/docs/reference/pregel.md b/docs/docs/reference/pregel.md index c1896ff20..a0ada895a 100644 --- a/docs/docs/reference/pregel.md +++ b/docs/docs/reference/pregel.md @@ -1,7 +1,25 @@ # Pregel -::: langgraph.pregel +::: langgraph.pregel.Pregel options: + show_if_no_docstring: true + show_root_heading: true + show_root_full_path: false members: - - Pregel - - PregelNode \ No newline at end of file + - stream + - astream + - invoke + - ainvoke + - get_state + - aget_state + - get_state_history + - aget_state_history + - update_state + - aupdate_state + - bulk_update_state + - abulk_update_state + - get_graph + - aget_graph + - get_subgraphs + - aget_subgraphs + - with_config \ No newline at end of file diff --git a/docs/docs/reference/types.md b/docs/docs/reference/types.md index b42b11f35..aac7c4f59 100644 --- a/docs/docs/reference/types.md +++ b/docs/docs/reference/types.md @@ -10,7 +10,6 @@ - CachePolicy - Interrupt - PregelTask - - PregelExecutableTask - StateSnapshot - Send - Command diff --git a/docs/mkdocs.yml b/docs/mkdocs.yml index 1ab2d3d56..3e4845dcb 100644 --- a/docs/mkdocs.yml +++ b/docs/mkdocs.yml @@ -390,13 +390,13 @@ nav: - Graphs: reference/graphs.md - Checkpointing: reference/checkpoints.md - Storage: reference/store.md - - Channels: reference/channels.md - - Errors: reference/errors.md - Types: reference/types.md - - Constants: reference/constants.md - - Pregel: reference/pregel.md - Config: reference/config.md - Functional API: reference/func.md + - Errors: reference/errors.md + - Constants: reference/constants.md + - Pregel: reference/pregel.md + - Channels: reference/channels.md - Prebuilt: - Agents: reference/agents.md - Supervisor: reference/supervisor.md diff --git a/docs/poetry.lock b/docs/poetry.lock index eecaee72c..1e360e335 100644 --- a/docs/poetry.lock +++ b/docs/poetry.lock @@ -3744,7 +3744,7 @@ langgraph-prebuilt = ">=0.1.7,<0.2.0" type = "git" url = "https://github.com/langchain-ai/langgraph-supervisor-py" reference = "HEAD" -resolved_reference = "18c8e7b7ca67828de6b71a83d7393ecb3c0372b0" +resolved_reference = "5cbfa9748ec809c0ee92c6a821fb18851b9ddfde" [[package]] name = "langgraph-swarm" @@ -3764,7 +3764,7 @@ langgraph = ">=0.3.5,<0.4.0" type = "git" url = "https://github.com/langchain-ai/langgraph-swarm-py" reference = "HEAD" -resolved_reference = "aa56a3935b41adc265239c014b1d341f9bbee4ba" +resolved_reference = "472a871aed829dc7fd655819cd5fbc8b0291f507" [[package]] name = "langmem" diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py index f4d8a322a..abacab49c 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py @@ -27,6 +27,8 @@ Conn = _internal.Conn # For backward compatibility class PostgresSaver(BasePostgresSaver): + """Checkpointer that stores checkpoints in a Postgres database.""" + lock: threading.Lock def __init__( diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py index 06354bb8f..70faa906c 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py @@ -27,6 +27,8 @@ Conn = _ainternal.Conn # For backward compatibility class AsyncPostgresSaver(BasePostgresSaver): + """Asynchronous checkpointer that stores checkpoints in a Postgres database.""" + lock: asyncio.Lock def __init__( diff --git a/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py b/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py index 49577b3aa..3a9ab50e2 100644 --- a/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py +++ b/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py @@ -34,6 +34,8 @@ EMPTY_BYTES = b"" class JsonPlusSerializer(SerializerProtocol): + """Serializer that uses ormsgpack, with a fallback to extended JSON serializer.""" + def __init__( self, *, __unpack_ext_hook__: Optional[Callable[[int, bytes], Any]] = None ) -> None: diff --git a/libs/checkpoint/langgraph/store/base/batch.py b/libs/checkpoint/langgraph/store/base/batch.py index 4883a3282..63ea02031 100644 --- a/libs/checkpoint/langgraph/store/base/batch.py +++ b/libs/checkpoint/langgraph/store/base/batch.py @@ -1,3 +1,5 @@ +"""Utilities for batching operations in a background task.""" + import asyncio import functools import weakref diff --git a/libs/langgraph/langgraph/channels/base.py b/libs/langgraph/langgraph/channels/base.py index 230aa4096..f7b53238f 100644 --- a/libs/langgraph/langgraph/channels/base.py +++ b/libs/langgraph/langgraph/channels/base.py @@ -13,6 +13,8 @@ C = TypeVar("C") class BaseChannel(Generic[Value, Update, C], ABC): + """Base class for all channels.""" + __slots__ = ("key", "typ") def __init__(self, typ: Any, key: str = "") -> None: diff --git a/libs/langgraph/langgraph/graph/graph.py b/libs/langgraph/langgraph/graph/graph.py index cbc58e819..f9aad0e2e 100644 --- a/libs/langgraph/langgraph/graph/graph.py +++ b/libs/langgraph/langgraph/graph/graph.py @@ -76,6 +76,15 @@ class Graph: *, metadata: Optional[dict[str, Any]] = None, ) -> Self: + """Add a new node to the graph. + + Args: + node (Union[str, RunnableLike]): The function or runnable this node will run. + If a string is provided, it will be used as the node name, and action will be used as the function or runnable. + action (Optional[RunnableLike]): The action associated with the node. (default: None) + Will be used as the node function or runnable if `node` is a string (node name). + metadata (Optional[dict[str, Any]]): The metadata associated with the node. (default: None) + """ if isinstance(node, str): for character in (NS_SEP, NS_END): if character in node: @@ -110,6 +119,12 @@ class Graph: return self def add_edge(self, start_key: str, end_key: str) -> Self: + """Add a directed edge from the start node to the end node. + + Args: + start_key (str): The key of the start node of the edge. + end_key (str): The key of the end node of the edge. + """ if self.compiled: logger.warning( "Adding an edge to a graph that has already been compiled. This will " @@ -302,6 +317,25 @@ class Graph: debug: bool = False, name: Optional[str] = None, ) -> "CompiledGraph": + """Compiles the graph into a `CompiledGraph` object. + + The compiled graph implements the `Runnable` interface and can be invoked, + streamed, batched, and run asynchronously. + + Args: + checkpointer (Optional[Union[Checkpointer, Literal[False]]]): A checkpoint saver object or flag. + If provided, this Checkpointer serves as a fully versioned "short-term memory" for the graph, + allowing it to be paused, resumed, and replayed from any point. + If None, it may inherit the parent graph's checkpointer when used as a subgraph. + If False, it will not use or inherit any checkpointer. + interrupt_before (Optional[Sequence[str]]): An optional list of node names to interrupt before. + interrupt_after (Optional[Sequence[str]]): An optional list of node names to interrupt after. + debug (bool): A flag indicating whether to enable debug mode. + name (Optional[str]): The name to use for the compiled graph. + + Returns: + CompiledGraph: The compiled graph. + """ # assign default values interrupt_before = interrupt_before or [] interrupt_after = interrupt_after or [] diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 177a78f80..9653e2282 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -123,42 +123,44 @@ class StateGraph(Graph): config_schema (Optional[Type[Any]]): The schema class that defines the configuration. Use this to expose configurable parameters in your API. - Examples: - >>> from langchain_core.runnables import RunnableConfig - >>> from typing_extensions import Annotated, TypedDict - >>> from langgraph.checkpoint.memory import MemorySaver - >>> from langgraph.graph import StateGraph - >>> - >>> def reducer(a: list, b: int | None) -> list: - ... if b is not None: - ... return a + [b] - ... return a - >>> - >>> class State(TypedDict): - ... x: Annotated[list, reducer] - >>> - >>> class ConfigSchema(TypedDict): - ... r: float - >>> - >>> graph = StateGraph(State, config_schema=ConfigSchema) - >>> - >>> def node(state: State, config: RunnableConfig) -> dict: - ... r = config["configurable"].get("r", 1.0) - ... x = state["x"][-1] - ... next_value = x * r * (1 - x) - ... return {"x": next_value} - >>> - >>> graph.add_node("A", node) - >>> graph.set_entry_point("A") - >>> graph.set_finish_point("A") - >>> compiled = graph.compile() - >>> - >>> print(compiled.config_specs) - [ConfigurableFieldSpec(id='r', annotation=, name=None, description=None, default=None, is_shared=False, dependencies=None)] - >>> - >>> step1 = compiled.invoke({"x": 0.5}, {"configurable": {"r": 3.0}}) - >>> print(step1) - {'x': [0.5, 0.75]}""" + Example: + ```python + from langchain_core.runnables import RunnableConfig + from typing_extensions import Annotated, TypedDict + from langgraph.checkpoint.memory import MemorySaver + from langgraph.graph import StateGraph + + def reducer(a: list, b: int | None) -> list: + if b is not None: + return a + [b] + return a + + class State(TypedDict): + x: Annotated[list, reducer] + + class ConfigSchema(TypedDict): + r: float + + graph = StateGraph(State, config_schema=ConfigSchema) + + def node(state: State, config: RunnableConfig) -> dict: + r = config["configurable"].get("r", 1.0) + x = state["x"][-1] + next_value = x * r * (1 - x) + return {"x": next_value} + + graph.add_node("A", node) + graph.set_entry_point("A") + graph.set_finish_point("A") + compiled = graph.compile() + + print(compiled.config_specs) + # [ConfigurableFieldSpec(id='r', annotation=, name=None, description=None, default=None, is_shared=False, dependencies=None)] + + step1 = compiled.invoke({"x": 0.5}, {"configurable": {"r": 3.0}}) + # {'x': [0.5, 0.75]} + ``` + """ nodes: dict[str, StateNodeSpec] # type: ignore[assignment] channels: dict[str, BaseChannel] @@ -251,17 +253,8 @@ class StateGraph(Graph): retry: Optional[Union[RetryPolicy, Sequence[RetryPolicy]]] = None, destinations: Optional[Union[dict[str, str], tuple[str, ...]]] = None, ) -> Self: - """Adds a new node to the state graph. + """Add a new node to the state graph. Will take the name of the function/runnable as the node name. - - Args: - node (RunnableLike): The function or runnable this node will run. - - Raises: - ValueError: If the key is already being used as a state key. - - Returns: - Self: The instance of the state graph, allowing for method chaining. """ ... @@ -276,18 +269,7 @@ class StateGraph(Graph): retry: Optional[Union[RetryPolicy, Sequence[RetryPolicy]]] = None, destinations: Optional[Union[dict[str, str], tuple[str, ...]]] = None, ) -> Self: - """Adds a new node to the state graph. - - Args: - node (str): The key of the node. - action (RunnableLike): The action associated with the node. - - Raises: - ValueError: If the key is already being used as a state key. - - Returns: - Self: The instance of the state graph, allowing for method chaining. - """ + """Add a new node to the state graph.""" ... def add_node( @@ -300,13 +282,13 @@ class StateGraph(Graph): retry: Optional[Union[RetryPolicy, Sequence[RetryPolicy]]] = None, destinations: Optional[Union[dict[str, str], tuple[str, ...]]] = None, ) -> Self: - """Adds a new node to the state graph. - - Will take the name of the function/runnable as the node name. + """Add a new node to the state graph. Args: node (Union[str, RunnableLike]): The function or runnable this node will run. + If a string is provided, it will be used as the node name, and action will be used as the function or runnable. action (Optional[RunnableLike]): The action associated with the node. (default: None) + Will be used as the node function or runnable if `node` is a string (node name). metadata (Optional[dict[str, Any]]): The metadata associated with the node. (default: None) input (Optional[Type[Any]]): The input schema for the node. (default: the graph's input schema) retry (Optional[Union[RetryPolicy, Sequence[RetryPolicy]]]): The policy for retrying the node. (default: None) @@ -319,29 +301,29 @@ class StateGraph(Graph): Raises: ValueError: If the key is already being used as a state key. - Examples: - ```pycon - >>> from langgraph.graph import START, StateGraph - ... - >>> def my_node(state, config): - ... return {"x": state["x"] + 1} - ... - >>> builder = StateGraph(dict) - >>> builder.add_node(my_node) # node name will be 'my_node' - >>> builder.add_edge(START, "my_node") - >>> graph = builder.compile() - >>> graph.invoke({"x": 1}) - {'x': 2} - ``` - Customize the name: + Example: + ```python + from langgraph.graph import START, StateGraph - ```pycon - >>> builder = StateGraph(dict) - >>> builder.add_node("my_fair_node", my_node) - >>> builder.add_edge(START, "my_fair_node") - >>> graph = builder.compile() - >>> graph.invoke({"x": 1}) - {'x': 2} + def my_node(state, config): + return {"x": state["x"] + 1} + + builder = StateGraph(dict) + builder.add_node(my_node) # node name will be 'my_node' + builder.add_edge(START, "my_node") + graph = builder.compile() + graph.invoke({"x": 1}) + # {'x': 2} + ``` + + Example: Customize the name: + ```python + builder = StateGraph(dict) + builder.add_node("my_fair_node", my_node) + builder.add_edge(START, "my_fair_node") + graph = builder.compile() + graph.invoke({"x": 1}) + # {'x': 2} ``` Returns: @@ -444,7 +426,7 @@ class StateGraph(Graph): return self def add_edge(self, start_key: Union[str, list[str]], end_key: str) -> Self: - """Adds a directed edge from the start node (or list of start nodes) to the end node. + """Add a directed edge from the start node (or list of start nodes) to the end node. When a single start node is provided, the graph will wait for that node to complete before executing the end node. When multiple start nodes are provided, @@ -584,7 +566,7 @@ class StateGraph(Graph): debug: bool = False, name: Optional[str] = None, ) -> "CompiledStateGraph": - """Compiles the state graph into a `CompiledGraph` object. + """Compiles the state graph into a `CompiledStateGraph` object. The compiled graph implements the `Runnable` interface and can be invoked, streamed, batched, and run asynchronously. @@ -598,6 +580,7 @@ class StateGraph(Graph): interrupt_before (Optional[Sequence[str]]): An optional list of node names to interrupt before. interrupt_after (Optional[Sequence[str]]): An optional list of node names to interrupt after. debug (bool): A flag indicating whether to enable debug mode. + name (Optional[str]): The name to use for the compiled graph. Returns: CompiledStateGraph: The compiled state graph. diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index e7170616a..9228ed34a 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -230,15 +230,15 @@ class Pregel(PregelProtocol): ## Actors - An **actor** is a [PregelNode][langgraph.pregel.read.PregelNode]. + An **actor** is a `PregelNode`. It subscribes to channels, reads data from them, and writes data to them. It can be thought of as an **actor** in the Pregel algorithm. - [PregelNodes][langgraph.pregel.read.PregelNode] implement LangChain's + `PregelNodes` implement LangChain's Runnable interface. ## Channels - Channels are used to communicate between actors (PregelNodes). + Channels are used to communicate between actors (`PregelNodes`). Each channel has a value type, an update type, and an update function – which takes a sequence of updates and modifies the stored value. Channels can be used to send data from one chain to @@ -560,7 +560,7 @@ class Pregel(PregelProtocol): def get_graph( self, config: RunnableConfig | None = None, *, xray: int | bool = False ) -> Graph: - """Returns a drawable representation of the computation graph.""" + """Return a drawable representation of the computation graph.""" # gather subgraphs if xray: subgraphs = { @@ -588,7 +588,7 @@ class Pregel(PregelProtocol): async def aget_graph( self, config: RunnableConfig | None = None, *, xray: int | bool = False ) -> Graph: - """Returns a drawable representation of the computation graph.""" + """Return a drawable representation of the computation graph.""" # gather subgraphs if xray: @@ -639,6 +639,7 @@ class Pregel(PregelProtocol): return self.__class__(**attrs) def with_config(self, config: RunnableConfig | None = None, **kwargs: Any) -> Self: + """Create a copy of the Pregel object with an updated config.""" return self.copy( {"config": merge_configs(self.config, config, cast(RunnableConfig, kwargs))} ) @@ -801,6 +802,16 @@ class Pregel(PregelProtocol): def get_subgraphs( self, *, namespace: str | None = None, recurse: bool = False ) -> Iterator[tuple[str, PregelProtocol]]: + """Get the subgraphs of the graph. + + Args: + namespace (Optional[str]): The namespace to filter the subgraphs by. + recurse (bool): Whether to recurse into the subgraphs. + If False, only the immediate subgraphs will be returned. + + Returns: + Iterator[tuple[str, PregelProtocol]]: An iterator of the (namespace, subgraph) pairs. + """ for name, node in self.nodes.items(): # filter by prefix if namespace is not None: @@ -830,6 +841,16 @@ class Pregel(PregelProtocol): async def aget_subgraphs( self, *, namespace: str | None = None, recurse: bool = False ) -> AsyncIterator[tuple[str, PregelProtocol]]: + """Get the subgraphs of the graph. + + Args: + namespace (Optional[str]): The namespace to filter the subgraphs by. + recurse (bool): Whether to recurse into the subgraphs. + If False, only the immediate subgraphs will be returned. + + Returns: + AsyncIterator[tuple[str, PregelProtocol]]: An iterator of the (namespace, subgraph) pairs. + """ for name, node in self.get_subgraphs(namespace=namespace, recurse=recurse): yield name, node @@ -1171,8 +1192,8 @@ class Pregel(PregelProtocol): before: RunnableConfig | None = None, limit: int | None = None, ) -> Iterator[StateSnapshot]: - config = ensure_config(config) """Get the history of the state of the graph.""" + config = ensure_config(config) checkpointer: BaseCheckpointSaver | None = ensure_config(config)[CONF].get( CONFIG_KEY_CHECKPOINTER, self.checkpointer ) @@ -1222,8 +1243,8 @@ class Pregel(PregelProtocol): before: RunnableConfig | None = None, limit: int | None = None, ) -> AsyncIterator[StateSnapshot]: + """Asynchronously get the history of the state of the graph.""" config = ensure_config(config) - """Get the history of the state of the graph.""" checkpointer: BaseCheckpointSaver | None = ensure_config(config)[CONF].get( CONFIG_KEY_CHECKPOINTER, self.checkpointer ) @@ -1688,7 +1709,7 @@ class Pregel(PregelProtocol): config: RunnableConfig, supersteps: Sequence[Sequence[StateUpdate]], ) -> RunnableConfig: - """Apply updates to the graph state in bulk. Requires a checkpointer to be set. + """Asynchronously apply updates to the graph state in bulk. Requires a checkpointer to be set. Args: config: The config to apply the updates to. @@ -2113,7 +2134,7 @@ class Pregel(PregelProtocol): values: dict[str, Any] | Any, as_node: str | None = None, ) -> RunnableConfig: - """Update the state of the graph asynchronously with the given values, as if they came from + """Asynchronously update the state of the graph with the given values, as if they came from node `as_node`. If `as_node` is not provided, it will be set to the last node that updated the state, if not ambiguous. """ @@ -2216,101 +2237,100 @@ class Pregel(PregelProtocol): Yields: The output of each step in the graph. The output shape depends on the stream_mode. - Examples: - Using different stream modes with a graph: - ```pycon - >>> import operator - >>> from typing_extensions import Annotated, TypedDict - >>> from langgraph.graph import StateGraph, START - ... - >>> class State(TypedDict): - ... alist: Annotated[list, operator.add] - ... another_list: Annotated[list, operator.add] - ... - >>> builder = StateGraph(State) - >>> builder.add_node("a", lambda _state: {"another_list": ["hi"]}) - >>> builder.add_node("b", lambda _state: {"alist": ["there"]}) - >>> builder.add_edge("a", "b") - >>> builder.add_edge(START, "a") - >>> graph = builder.compile() - ``` - With stream_mode="values": + Example: Using stream_mode="values": + ```python + import operator + from typing_extensions import Annotated, TypedDict + from langgraph.graph import StateGraph, START - ```pycon - >>> for event in graph.stream({"alist": ['Ex for stream_mode="values"']}, stream_mode="values"): - ... print(event) - {'alist': ['Ex for stream_mode="values"'], 'another_list': []} - {'alist': ['Ex for stream_mode="values"'], 'another_list': ['hi']} - {'alist': ['Ex for stream_mode="values"', 'there'], 'another_list': ['hi']} - ``` - With stream_mode="updates": + class State(TypedDict): + alist: Annotated[list, operator.add] + another_list: Annotated[list, operator.add] - ```pycon - >>> for event in graph.stream({"alist": ['Ex for stream_mode="updates"']}, stream_mode="updates"): - ... print(event) - {'a': {'another_list': ['hi']}} - {'b': {'alist': ['there']}} - ``` - With stream_mode="debug": + builder = StateGraph(State) + builder.add_node("a", lambda _state: {"another_list": ["hi"]}) + builder.add_node("b", lambda _state: {"alist": ["there"]}) + builder.add_edge("a", "b") + builder.add_edge(START, "a") + graph = builder.compile() - ```pycon - >>> for event in graph.stream({"alist": ['Ex for stream_mode="debug"']}, stream_mode="debug"): - ... print(event) - {'type': 'task', 'timestamp': '2024-06-23T...+00:00', 'step': 1, 'payload': {'id': '...', 'name': 'a', 'input': {'alist': ['Ex for stream_mode="debug"'], 'another_list': []}, 'triggers': ['start:a']}} - {'type': 'task_result', 'timestamp': '2024-06-23T...+00:00', 'step': 1, 'payload': {'id': '...', 'name': 'a', 'result': [('another_list', ['hi'])]}} - {'type': 'task', 'timestamp': '2024-06-23T...+00:00', 'step': 2, 'payload': {'id': '...', 'name': 'b', 'input': {'alist': ['Ex for stream_mode="debug"'], 'another_list': ['hi']}, 'triggers': ['a']}} - {'type': 'task_result', 'timestamp': '2024-06-23T...+00:00', 'step': 2, 'payload': {'id': '...', 'name': 'b', 'result': [('alist', ['there'])]}} + for event in graph.stream({"alist": ['Ex for stream_mode="values"']}, stream_mode="values"): + print(event) + + # {'alist': ['Ex for stream_mode="values"'], 'another_list': []} + # {'alist': ['Ex for stream_mode="values"'], 'another_list': ['hi']} + # {'alist': ['Ex for stream_mode="values"', 'there'], 'another_list': ['hi']} ``` - With stream_mode="custom": + Example: Using stream_mode="updates": + ```python + for event in graph.stream({"alist": ['Ex for stream_mode="updates"']}, stream_mode="updates"): + print(event) - ```pycon - >>> from langgraph.types import StreamWriter - ... - >>> def node_a(state: State, writer: StreamWriter): - ... writer({"custom_data": "foo"}) - ... return {"alist": ["hi"]} - ... - >>> builder = StateGraph(State) - >>> builder.add_node("a", node_a) - >>> builder.add_edge(START, "a") - >>> graph = builder.compile() - ... - >>> for event in graph.stream({"alist": ['Ex for stream_mode="custom"']}, stream_mode="custom"): - ... print(event) - {'custom_data': 'foo'} + # {'a': {'another_list': ['hi']}} + # {'b': {'alist': ['there']}} ``` - With stream_mode="messages": + Example: Using stream_mode="debug": + ```python + for event in graph.stream({"alist": ['Ex for stream_mode="debug"']}, stream_mode="debug"): + print(event) - ```pycon - >>> from typing_extensions import Annotated, TypedDict - >>> from langgraph.graph import StateGraph, START - >>> from langchain_openai import ChatOpenAI - ... - >>> llm = ChatOpenAI(model="gpt-4o-mini") - ... - >>> class State(TypedDict): - ... question: str - ... answer: str - ... - >>> def node_a(state: State): - ... response = llm.invoke(state["question"]) - ... return {"answer": response.content} - ... - >>> builder = StateGraph(State) - >>> builder.add_node("a", node_a) - >>> builder.add_edge(START, "a") - >>> graph = builder.compile() + # {'type': 'task', 'timestamp': '2024-06-23T...+00:00', 'step': 1, 'payload': {'id': '...', 'name': 'a', 'input': {'alist': ['Ex for stream_mode="debug"'], 'another_list': []}, 'triggers': ['start:a']}} + # {'type': 'task_result', 'timestamp': '2024-06-23T...+00:00', 'step': 1, 'payload': {'id': '...', 'name': 'a', 'result': [('another_list', ['hi'])]}} + # {'type': 'task', 'timestamp': '2024-06-23T...+00:00', 'step': 2, 'payload': {'id': '...', 'name': 'b', 'input': {'alist': ['Ex for stream_mode="debug"'], 'another_list': ['hi']}, 'triggers': ['a']}} + # {'type': 'task_result', 'timestamp': '2024-06-23T...+00:00', 'step': 2, 'payload': {'id': '...', 'name': 'b', 'result': [('alist', ['there'])]}} + ``` - >>> for event in graph.stream({"question": "What is the capital of France?"}, stream_mode="messages"): - ... print(event) - (AIMessageChunk(content='The', additional_kwargs={}, response_metadata={}, id='...'), {'langgraph_step': 1, 'langgraph_node': 'a', 'langgraph_triggers': ['start:a'], 'langgraph_path': ('__pregel_pull', 'a'), 'langgraph_checkpoint_ns': '...', 'checkpoint_ns': '...', 'ls_provider': 'openai', 'ls_model_name': 'gpt-4o-mini', 'ls_model_type': 'chat', 'ls_temperature': 0.7}) - (AIMessageChunk(content=' capital', additional_kwargs={}, response_metadata={}, id='...'), {'langgraph_step': 1, 'langgraph_node': 'a', 'langgraph_triggers': ['start:a'], ...}) - (AIMessageChunk(content=' of', additional_kwargs={}, response_metadata={}, id='...'), {...}) - (AIMessageChunk(content=' France', additional_kwargs={}, response_metadata={}, id='...'), {...}) - (AIMessageChunk(content=' is', additional_kwargs={}, response_metadata={}, id='...'), {...}) - (AIMessageChunk(content=' Paris', additional_kwargs={}, response_metadata={}, id='...'), {...}) + Example: Using stream_mode="custom": + ```python + from langgraph.types import StreamWriter + + def node_a(state: State, writer: StreamWriter): + writer({"custom_data": "foo"}) + return {"alist": ["hi"]} + + builder = StateGraph(State) + builder.add_node("a", node_a) + builder.add_edge(START, "a") + graph = builder.compile() + + for event in graph.stream({"alist": ['Ex for stream_mode="custom"']}, stream_mode="custom"): + print(event) + + # {'custom_data': 'foo'} + ``` + + Example: Using stream_mode="messages": + ```python + from typing_extensions import Annotated, TypedDict + from langgraph.graph import StateGraph, START + from langchain_openai import ChatOpenAI + + llm = ChatOpenAI(model="gpt-4o-mini") + + class State(TypedDict): + question: str + answer: str + + def node_a(state: State): + response = llm.invoke(state["question"]) + return {"answer": response.content} + + builder = StateGraph(State) + builder.add_node("a", node_a) + builder.add_edge(START, "a") + graph = builder.compile() + + for event in graph.stream({"question": "What is the capital of France?"}, stream_mode="messages"): + print(event) + + # (AIMessageChunk(content='The', additional_kwargs={}, response_metadata={}, id='...'), {'langgraph_step': 1, 'langgraph_node': 'a', 'langgraph_triggers': ['start:a'], 'langgraph_path': ('__pregel_pull', 'a'), 'langgraph_checkpoint_ns': '...', 'checkpoint_ns': '...', 'ls_provider': 'openai', 'ls_model_name': 'gpt-4o-mini', 'ls_model_type': 'chat', 'ls_temperature': 0.7}) + # (AIMessageChunk(content=' capital', additional_kwargs={}, response_metadata={}, id='...'), {'langgraph_step': 1, 'langgraph_node': 'a', 'langgraph_triggers': ['start:a'], ...}) + # (AIMessageChunk(content=' of', additional_kwargs={}, response_metadata={}, id='...'), {...}) + # (AIMessageChunk(content=' France', additional_kwargs={}, response_metadata={}, id='...'), {...}) + # (AIMessageChunk(content=' is', additional_kwargs={}, response_metadata={}, id='...'), {...}) + # (AIMessageChunk(content=' Paris', additional_kwargs={}, response_metadata={}, id='...'), {...}) ``` """ @@ -2478,7 +2498,7 @@ class Pregel(PregelProtocol): debug: bool | None = None, subgraphs: bool = False, ) -> AsyncIterator[dict[str, Any] | Any]: - """Stream graph steps for a single input. + """Asynchronously stream graph steps for a single input. Args: input: The input to the graph. @@ -2503,101 +2523,100 @@ class Pregel(PregelProtocol): Yields: The output of each step in the graph. The output shape depends on the stream_mode. - Examples: - Using different stream modes with a graph: - ```pycon - >>> import operator - >>> from typing_extensions import Annotated, TypedDict - >>> from langgraph.graph import StateGraph, START - ... - >>> class State(TypedDict): - ... alist: Annotated[list, operator.add] - ... another_list: Annotated[list, operator.add] - ... - >>> builder = StateGraph(State) - >>> builder.add_node("a", lambda _state: {"another_list": ["hi"]}) - >>> builder.add_node("b", lambda _state: {"alist": ["there"]}) - >>> builder.add_edge("a", "b") - >>> builder.add_edge(START, "a") - >>> graph = builder.compile() - ``` - With stream_mode="values": + Example: Using stream_mode="values": + ```python + import operator + from typing_extensions import Annotated, TypedDict + from langgraph.graph import StateGraph, START - ```pycon - >>> async for event in graph.astream({"alist": ['Ex for stream_mode="values"']}, stream_mode="values"): - ... print(event) - {'alist': ['Ex for stream_mode="values"'], 'another_list': []} - {'alist': ['Ex for stream_mode="values"'], 'another_list': ['hi']} - {'alist': ['Ex for stream_mode="values"', 'there'], 'another_list': ['hi']} - ``` - With stream_mode="updates": + class State(TypedDict): + alist: Annotated[list, operator.add] + another_list: Annotated[list, operator.add] - ```pycon - >>> async for event in graph.astream({"alist": ['Ex for stream_mode="updates"']}, stream_mode="updates"): - ... print(event) - {'a': {'another_list': ['hi']}} - {'b': {'alist': ['there']}} - ``` - With stream_mode="debug": + builder = StateGraph(State) + builder.add_node("a", lambda _state: {"another_list": ["hi"]}) + builder.add_node("b", lambda _state: {"alist": ["there"]}) + builder.add_edge("a", "b") + builder.add_edge(START, "a") + graph = builder.compile() - ```pycon - >>> async for event in graph.astream({"alist": ['Ex for stream_mode="debug"']}, stream_mode="debug"): - ... print(event) - {'type': 'task', 'timestamp': '2024-06-23T...+00:00', 'step': 1, 'payload': {'id': '...', 'name': 'a', 'input': {'alist': ['Ex for stream_mode="debug"'], 'another_list': []}, 'triggers': ['start:a']}} - {'type': 'task_result', 'timestamp': '2024-06-23T...+00:00', 'step': 1, 'payload': {'id': '...', 'name': 'a', 'result': [('another_list', ['hi'])]}} - {'type': 'task', 'timestamp': '2024-06-23T...+00:00', 'step': 2, 'payload': {'id': '...', 'name': 'b', 'input': {'alist': ['Ex for stream_mode="debug"'], 'another_list': ['hi']}, 'triggers': ['a']}} - {'type': 'task_result', 'timestamp': '2024-06-23T...+00:00', 'step': 2, 'payload': {'id': '...', 'name': 'b', 'result': [('alist', ['there'])]}} + async for event in graph.astream({"alist": ['Ex for stream_mode="values"']}, stream_mode="values"): + print(event) + + # {'alist': ['Ex for stream_mode="values"'], 'another_list': []} + # {'alist': ['Ex for stream_mode="values"'], 'another_list': ['hi']} + # {'alist': ['Ex for stream_mode="values"', 'there'], 'another_list': ['hi']} ``` - With stream_mode="custom": + Example: Using stream_mode="updates": + ```python + async for event in graph.astream({"alist": ['Ex for stream_mode="updates"']}, stream_mode="updates"): + print(event) - ```pycon - >>> from langgraph.types import StreamWriter - ... - >>> async def node_a(state: State, writer: StreamWriter): - ... writer({"custom_data": "foo"}) - ... return {"alist": ["hi"]} - ... - >>> builder = StateGraph(State) - >>> builder.add_node("a", node_a) - >>> builder.add_edge(START, "a") - >>> graph = builder.compile() - ... - >>> async for event in graph.astream({"alist": ['Ex for stream_mode="custom"']}, stream_mode="custom"): - ... print(event) - {'custom_data': 'foo'} + # {'a': {'another_list': ['hi']}} + # {'b': {'alist': ['there']}} ``` - With stream_mode="messages": + Example: Using stream_mode="debug": + ```python + async for event in graph.astream({"alist": ['Ex for stream_mode="debug"']}, stream_mode="debug"): + print(event) - ```pycon - >>> from typing_extensions import Annotated, TypedDict - >>> from langgraph.graph import StateGraph, START - >>> from langchain_openai import ChatOpenAI - ... - >>> llm = ChatOpenAI(model="gpt-4o-mini") - ... - >>> class State(TypedDict): - ... question: str - ... answer: str - ... - >>> async def node_a(state: State): - ... response = await llm.ainvoke(state["question"]) - ... return {"answer": response.content} - ... - >>> builder = StateGraph(State) - >>> builder.add_node("a", node_a) - >>> builder.add_edge(START, "a") - >>> graph = builder.compile() + # {'type': 'task', 'timestamp': '2024-06-23T...+00:00', 'step': 1, 'payload': {'id': '...', 'name': 'a', 'input': {'alist': ['Ex for stream_mode="debug"'], 'another_list': []}, 'triggers': ['start:a']}} + # {'type': 'task_result', 'timestamp': '2024-06-23T...+00:00', 'step': 1, 'payload': {'id': '...', 'name': 'a', 'result': [('another_list', ['hi'])]}} + # {'type': 'task', 'timestamp': '2024-06-23T...+00:00', 'step': 2, 'payload': {'id': '...', 'name': 'b', 'input': {'alist': ['Ex for stream_mode="debug"'], 'another_list': ['hi']}, 'triggers': ['a']}} + # {'type': 'task_result', 'timestamp': '2024-06-23T...+00:00', 'step': 2, 'payload': {'id': '...', 'name': 'b', 'result': [('alist', ['there'])]}} + ``` - >>> for event in graph.stream({"question": "What is the capital of France?"}, stream_mode="messages"): - ... print(event) - (AIMessageChunk(content='The', additional_kwargs={}, response_metadata={}, id='...'), {'langgraph_step': 1, 'langgraph_node': 'a', 'langgraph_triggers': ['start:a'], 'langgraph_path': ('__pregel_pull', 'a'), 'langgraph_checkpoint_ns': '...', 'checkpoint_ns': '...', 'ls_provider': 'openai', 'ls_model_name': 'gpt-4o-mini', 'ls_model_type': 'chat', 'ls_temperature': 0.7}) - (AIMessageChunk(content=' capital', additional_kwargs={}, response_metadata={}, id='...'), {'langgraph_step': 1, 'langgraph_node': 'a', 'langgraph_triggers': ['start:a'], ...}) - (AIMessageChunk(content=' of', additional_kwargs={}, response_metadata={}, id='...'), {...}) - (AIMessageChunk(content=' France', additional_kwargs={}, response_metadata={}, id='...'), {...}) - (AIMessageChunk(content=' is', additional_kwargs={}, response_metadata={}, id='...'), {...}) - (AIMessageChunk(content=' Paris', additional_kwargs={}, response_metadata={}, id='...'), {...}) + Example: Using stream_mode="custom": + ```python + from langgraph.types import StreamWriter + + async def node_a(state: State, writer: StreamWriter): + writer({"custom_data": "foo"}) + return {"alist": ["hi"]} + + builder = StateGraph(State) + builder.add_node("a", node_a) + builder.add_edge(START, "a") + graph = builder.compile() + + async for event in graph.astream({"alist": ['Ex for stream_mode="custom"']}, stream_mode="custom"): + print(event) + + # {'custom_data': 'foo'} + ``` + + Example: Using stream_mode="messages": + ```python + from typing_extensions import Annotated, TypedDict + from langgraph.graph import StateGraph, START + from langchain_openai import ChatOpenAI + + llm = ChatOpenAI(model="gpt-4o-mini") + + class State(TypedDict): + question: str + answer: str + + async def node_a(state: State): + response = await llm.ainvoke(state["question"]) + return {"answer": response.content} + + builder = StateGraph(State) + builder.add_node("a", node_a) + builder.add_edge(START, "a") + graph = builder.compile() + + async for event in graph.astream({"question": "What is the capital of France?"}, stream_mode="messages"): + print(event) + + # (AIMessageChunk(content='The', additional_kwargs={}, response_metadata={}, id='...'), {'langgraph_step': 1, 'langgraph_node': 'a', 'langgraph_triggers': ['start:a'], 'langgraph_path': ('__pregel_pull', 'a'), 'langgraph_checkpoint_ns': '...', 'checkpoint_ns': '...', 'ls_provider': 'openai', 'ls_model_name': 'gpt-4o-mini', 'ls_model_type': 'chat', 'ls_temperature': 0.7}) + # (AIMessageChunk(content=' capital', additional_kwargs={}, response_metadata={}, id='...'), {'langgraph_step': 1, 'langgraph_node': 'a', 'langgraph_triggers': ['start:a'], ...}) + # (AIMessageChunk(content=' of', additional_kwargs={}, response_metadata={}, id='...'), {...}) + # (AIMessageChunk(content=' France', additional_kwargs={}, response_metadata={}, id='...'), {...}) + # (AIMessageChunk(content=' is', additional_kwargs={}, response_metadata={}, id='...'), {...}) + # (AIMessageChunk(content=' Paris', additional_kwargs={}, response_metadata={}, id='...'), {...}) ``` """ diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index 86a377684..1cfe02d9d 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -133,7 +133,8 @@ class CachePolicy(NamedTuple): @dataclasses.dataclass(**_DC_KWARGS) class Interrupt: - """ + """Information about an interrupt that occurred in a node. + !!! version-added "Added in version 0.2.24." """ @@ -156,6 +157,8 @@ class StateUpdate(NamedTuple): class PregelTask(NamedTuple): + """A Pregel task.""" + id: str name: str path: tuple[Union[str, int, tuple], ...] diff --git a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py index 70756b780..d945b95eb 100644 --- a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py +++ b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py @@ -398,23 +398,22 @@ def create_react_agent( ``` Example: + ```python + from langgraph.prebuilt import create_react_agent - ```python - from langgraph.prebuilt import create_react_agent + def check_weather(location: str) -> str: + '''Return the weather forecast for the specified location.''' + return f"It's always sunny in {location}" - def check_weather(location: str) -> str: - '''Return the weather forecast for the specified location.''' - return f"It's always sunny in {location}" - - graph = create_react_agent( - "anthropic:claude-3-7-sonnet-latest", - tools=[check_weather], - prompt="You are a helpful assistant", - ) - inputs = {"messages": [{"role": "user", "content": "what is the weather in sf"}]} - for chunk in graph.stream(inputs, stream_mode="updates"): - print(chunk) - ``` + graph = create_react_agent( + "anthropic:claude-3-7-sonnet-latest", + tools=[check_weather], + prompt="You are a helpful assistant", + ) + inputs = {"messages": [{"role": "user", "content": "what is the weather in sf"}]} + for chunk in graph.stream(inputs, stream_mode="updates"): + print(chunk) + ``` """ if version not in ("v1", "v2"): raise ValueError( diff --git a/libs/prebuilt/langgraph/prebuilt/tool_validator.py b/libs/prebuilt/langgraph/prebuilt/tool_validator.py index 70796a8d1..0e58c7d6c 100644 --- a/libs/prebuilt/langgraph/prebuilt/tool_validator.py +++ b/libs/prebuilt/langgraph/prebuilt/tool_validator.py @@ -73,87 +73,57 @@ class ValidationNode(RunnableCallable): Returns: (Union[Dict[str, List[ToolMessage]], Sequence[ToolMessage]]): A list of ToolMessages with the validated content or error messages. - Examples: - Example usage for re-prompting the model to generate a valid response: - >>> from typing import Literal, Annotated - >>> from typing_extensions import TypedDict - ... - >>> from langchain_anthropic import ChatAnthropic - >>> from pydantic import BaseModel, field_validator - ... - >>> from langgraph.graph import END, START, StateGraph - >>> from langgraph.prebuilt import ValidationNode - >>> from langgraph.graph.message import add_messages - ... - ... - >>> class SelectNumber(BaseModel): - ... a: int - ... - ... @field_validator("a") - ... def a_must_be_meaningful(cls, v): - ... if v != 37: - ... raise ValueError("Only 37 is allowed") - ... return v - ... - ... - >>> builder = StateGraph(Annotated[list, add_messages]) - >>> llm = ChatAnthropic(model="claude-3-5-haiku-latest").bind_tools([SelectNumber]) - >>> builder.add_node("model", llm) - >>> builder.add_node("validation", ValidationNode([SelectNumber])) - >>> builder.add_edge(START, "model") - ... - ... - >>> def should_validate(state: list) -> Literal["validation", "__end__"]: - ... if state[-1].tool_calls: - ... return "validation" - ... return END - ... - ... - >>> builder.add_conditional_edges("model", should_validate) - ... - ... - >>> def should_reprompt(state: list) -> Literal["model", "__end__"]: - ... for msg in state[::-1]: - ... # None of the tool calls were errors - ... if msg.type == "ai": - ... return END - ... if msg.additional_kwargs.get("is_error"): - ... return "model" - ... return END - ... - ... - >>> builder.add_conditional_edges("validation", should_reprompt) - ... - ... - >>> graph = builder.compile() - >>> res = graph.invoke(("user", "Select a number, any number")) - >>> # Show the retry logic - >>> for msg in res: - ... msg.pretty_print() - ================================ Human Message ================================= - Select a number, any number - ================================== Ai Message ================================== - [{'id': 'toolu_01JSjT9Pq8hGmTgmMPc6KnvM', 'input': {'a': 42}, 'name': 'SelectNumber', 'type': 'tool_use'}] - Tool Calls: - SelectNumber (toolu_01JSjT9Pq8hGmTgmMPc6KnvM) - Call ID: toolu_01JSjT9Pq8hGmTgmMPc6KnvM - Args: - a: 42 - ================================= Tool Message ================================= - Name: SelectNumber - ValidationError(model='SelectNumber', errors=[{'loc': ('a',), 'msg': 'Only 37 is allowed', 'type': 'value_error'}]) - Respond after fixing all validation errors. - ================================== Ai Message ================================== - [{'id': 'toolu_01PkxSVxNxc5wqwCPW1FiSmV', 'input': {'a': 37}, 'name': 'SelectNumber', 'type': 'tool_use'}] - Tool Calls: - SelectNumber (toolu_01PkxSVxNxc5wqwCPW1FiSmV) - Call ID: toolu_01PkxSVxNxc5wqwCPW1FiSmV - Args: - a: 37 - ================================= Tool Message ================================= - Name: SelectNumber - {"a": 37} + Example: + ```python title="Example usage for re-prompting the model to generate a valid response:" + from typing import Literal, Annotated + from typing_extensions import TypedDict + from langchain_anthropic import ChatAnthropic + from pydantic import BaseModel, field_validator + + from langgraph.graph import END, START, StateGraph + from langgraph.prebuilt import ValidationNode + from langgraph.graph.message import add_messages + + class SelectNumber(BaseModel): + a: int + + @field_validator("a") + def a_must_be_meaningful(cls, v): + if v != 37: + raise ValueError("Only 37 is allowed") + return v + + builder = StateGraph(Annotated[list, add_messages]) + llm = ChatAnthropic(model="claude-3-5-haiku-latest").bind_tools([SelectNumber]) + builder.add_node("model", llm) + builder.add_node("validation", ValidationNode([SelectNumber])) + builder.add_edge(START, "model") + + def should_validate(state: list) -> Literal["validation", "__end__"]: + if state[-1].tool_calls: + return "validation" + return END + + builder.add_conditional_edges("model", should_validate) + + def should_reprompt(state: list) -> Literal["model", "__end__"]: + for msg in state[::-1]: + # None of the tool calls were errors + if msg.type == "ai": + return END + if msg.additional_kwargs.get("is_error"): + return "model" + return END + + builder.add_conditional_edges("validation", should_reprompt) + + graph = builder.compile() + res = graph.invoke(("user", "Select a number, any number")) + # Show the retry logic + for msg in res: + msg.pretty_print() + ``` """ def __init__(