codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import inspect4import logging5import typing6import warnings7from collections import defaultdict8from collections.abc import Awaitable, Callable, Hashable, Sequence9from dataclasses import dataclass, is_dataclass10from datetime import timedelta11from functools import partial12from inspect import isclass, isfunction, ismethod, signature13from types import FunctionType14from types import NoneType as NoneType15from typing import (16 Any,17 Generic,18 Literal,19 TypeVar,20 Union,21 cast,22 get_args,23 get_origin,24 get_type_hints,25 overload,26)27 28from langchain_core.runnables import Runnable, RunnableConfig29from langgraph.cache.base import BaseCache30from langgraph.checkpoint.base import Checkpoint31from langgraph.store.base import BaseStore32from pydantic import BaseModel, TypeAdapter33from typing_extensions import NotRequired, Required, Self, Unpack, is_typeddict34 35from langgraph._internal import _serde36from langgraph._internal._constants import (37 INTERRUPT,38 NS_END,39 NS_SEP,40 TASKS,41)42from langgraph._internal._fields import (43 get_cached_annotated_keys,44 get_field_default,45 get_update_as_tuples,46)47from langgraph._internal._pydantic import create_model48from langgraph._internal._runnable import coerce_to_runnable49from langgraph._internal._timeout import coerce_timeout_policy50from langgraph._internal._typing import EMPTY_SEQ, MISSING, DeprecatedKwargs51from langgraph.channels.base import BaseChannel52from langgraph.channels.binop import BinaryOperatorAggregate53from langgraph.channels.delta import DeltaChannel54from langgraph.channels.ephemeral_value import EphemeralValue55from langgraph.channels.last_value import LastValue, LastValueAfterFinish56from langgraph.channels.named_barrier_value import (57 NamedBarrierValue,58 NamedBarrierValueAfterFinish,59)60from langgraph.constants import END, START, TAG_HIDDEN61from langgraph.errors import (62 ErrorCode,63 InvalidUpdateError,64 ParentCommand,65 create_error_message,66)67from langgraph.graph._branch import BranchSpec68from langgraph.graph._node import StateNode, StateNodeSpec69from langgraph.managed.base import (70 ManagedValueSpec,71 is_managed_value,72)73from langgraph.pregel import Pregel74from langgraph.pregel._read import ChannelRead, PregelNode75from langgraph.pregel._write import (76 ChannelWrite,77 ChannelWriteEntry,78 ChannelWriteTupleEntry,79)80from langgraph.types import (81 All,82 CachePolicy,83 Checkpointer,84 Command,85 RetryPolicy,86 Send,87 TimeoutPolicy,88 ensure_valid_checkpointer,89)90from langgraph.typing import ContextT, InputT, NodeInputT, OutputT, StateT91from langgraph.warnings import LangGraphDeprecatedSinceV05, LangGraphDeprecatedSinceV1092 93__all__ = ("StateGraph", "CompiledStateGraph")94 95logger = logging.getLogger(__name__)96 97_CHANNEL_BRANCH_TO = "branch:to:{}"98_DEFAULT_ERROR_HANDLER_NODE = "__default_error_handler__"99 100 101@dataclass(slots=True)102class _NodeDefaults:103 """Default node policies applied to every node at compile time."""104 105 retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None106 cache_policy: CachePolicy | None = None107 error_handler: StateNode[Any, Any] | None = None108 timeout: TimeoutPolicy | None = None109 110 111def _warn_invalid_state_schema(schema: type[Any] | Any) -> None:112 if isinstance(schema, type):113 return114 if typing.get_args(schema):115 return116 warnings.warn(117 f"Invalid state_schema: {schema}. Expected a type or Annotated[type, reducer]. "118 "Please provide a valid schema to ensure correct updates.\n"119 " See: https://langchain-ai.github.io/langgraph/reference/graphs/#stategraph"120 )121 122 123def _get_node_name(node: StateNode[Any, ContextT]) -> str:124 try:125 return getattr(node, "__name__", node.__class__.__name__)126 except AttributeError:127 raise TypeError(f"Unsupported node type: {type(node)}")128 129 130class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):131 """A graph whose nodes communicate by reading and writing to a shared state.132 133 The signature of each node is `State -> Partial<State>`.134 135 Each state key can optionally be annotated with a reducer function that136 will be used to aggregate the values of that key received from multiple nodes.137 The signature of a reducer function is `(Value, Value) -> Value`.138 139 !!! warning140 141 `StateGraph` is a builder class and cannot be used directly for execution.142 You must first call `.compile()` to create an executable graph that supports143 methods like `invoke()`, `stream()`, `astream()`, and `ainvoke()`. See the144 `CompiledStateGraph` documentation for more details.145 146 Args:147 state_schema: The schema class that defines the state.148 context_schema: The schema class that defines the runtime context.149 150 Use this to expose immutable context data to your nodes, like `user_id`, `db_conn`, etc.151 input_schema: The schema class that defines the input to the graph.152 output_schema: The schema class that defines the output from the graph.153 154 !!! warning "`config_schema` Deprecated"155 The `config_schema` parameter is deprecated in v0.6.0 and support will be removed in v2.0.0.156 Please use `context_schema` instead to specify the schema for run-scoped context.157 158 Example:159 ```python160 from langchain_core.runnables import RunnableConfig161 from typing_extensions import Annotated, TypedDict162 from langgraph.checkpoint.memory import InMemorySaver163 from langgraph.graph import StateGraph164 from langgraph.runtime import Runtime165 166 167 def reducer(a: list, b: int | None) -> list:168 if b is not None:169 return a + [b]170 return a171 172 173 class State(TypedDict):174 x: Annotated[list, reducer]175 176 177 class Context(TypedDict):178 r: float179 180 181 graph = StateGraph(state_schema=State, context_schema=Context)182 183 184 def node(state: State, runtime: Runtime[Context]) -> dict:185 r = runtime.context.get("r", 1.0)186 x = state["x"][-1]187 next_value = x * r * (1 - x)188 return {"x": next_value}189 190 191 graph.add_node("A", node)192 graph.set_entry_point("A")193 graph.set_finish_point("A")194 compiled = graph.compile()195 196 step1 = compiled.invoke({"x": 0.5}, context={"r": 3.0})197 # {'x': [0.5, 0.75]}198 ```199 """200 201 edges: set[tuple[str, str]]202 nodes: dict[str, StateNodeSpec[Any, ContextT]]203 branches: defaultdict[str, dict[str, BranchSpec]]204 channels: dict[str, BaseChannel]205 managed: dict[str, ManagedValueSpec]206 schemas: dict[type[Any], dict[str, BaseChannel | ManagedValueSpec]]207 waiting_edges: set[tuple[tuple[str, ...], str]]208 209 compiled: bool210 state_schema: type[StateT]211 context_schema: type[ContextT] | None212 input_schema: type[InputT]213 output_schema: type[OutputT]214 215 def __init__(216 self,217 state_schema: type[StateT],218 context_schema: type[ContextT] | None = None,219 *,220 input_schema: type[InputT] | None = None,221 output_schema: type[OutputT] | None = None,222 **kwargs: Unpack[DeprecatedKwargs],223 ) -> None:224 if (config_schema := kwargs.get("config_schema", MISSING)) is not MISSING:225 warnings.warn(226 "`config_schema` is deprecated and will be removed. Please use `context_schema` instead.",227 category=LangGraphDeprecatedSinceV10,228 stacklevel=2,229 )230 if context_schema is None:231 context_schema = cast(type[ContextT], config_schema)232 233 if (input_ := kwargs.get("input", MISSING)) is not MISSING:234 warnings.warn(235 "`input` is deprecated and will be removed. Please use `input_schema` instead.",236 category=LangGraphDeprecatedSinceV05,237 stacklevel=2,238 )239 if input_schema is None:240 input_schema = cast(type[InputT], input_)241 242 if (output := kwargs.get("output", MISSING)) is not MISSING:243 warnings.warn(244 "`output` is deprecated and will be removed. Please use `output_schema` instead.",245 category=LangGraphDeprecatedSinceV05,246 stacklevel=2,247 )248 if output_schema is None:249 output_schema = cast(type[OutputT], output)250 251 self.nodes = {}252 self.edges = set()253 self.branches = defaultdict(dict)254 self.schemas = {}255 self.channels = {}256 self.managed = {}257 self.compiled = False258 self.waiting_edges = set()259 260 self.state_schema = state_schema261 self.input_schema = cast(type[InputT], input_schema or state_schema)262 self.output_schema = cast(type[OutputT], output_schema or state_schema)263 self.context_schema = context_schema264 265 self._node_defaults: _NodeDefaults = _NodeDefaults()266 267 self._add_schema(self.state_schema)268 self._add_schema(self.input_schema, allow_managed=False)269 self._add_schema(self.output_schema, allow_managed=False)270 271 def set_node_defaults(272 self,273 *,274 retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,275 cache_policy: CachePolicy | None = None,276 error_handler: StateNode[Any, ContextT] | None = None,277 timeout: float | timedelta | TimeoutPolicy | None = None,278 ) -> Self:279 """Set default node policies that apply to every node in this graph.280 281 Per-node values passed to `add_node` always take precedence over these282 defaults. Defaults are applied at `compile()` time. Policies set here283 are **not** inherited by subgraphs.284 285 `retry_policy` and `timeout` defaults apply to **all** nodes,286 including error-handler nodes. `cache_policy` and `error_handler`287 defaults only apply to regular nodes -- caching error-handler results288 is unsafe, and handlers must never catch themselves.289 290 Args:291 retry_policy: Default retry policy for nodes that don't specify292 their own via `add_node(..., retry_policy=...)`. Also applies293 to error-handler nodes.294 cache_policy: Default cache policy for nodes that don't specify295 their own via `add_node(..., cache_policy=...)`. Does **not**296 apply to error-handler nodes.297 error_handler: Default error handler invoked when any regular node298 raises and does not have its own `error_handler` set via299 `add_node`. The handler is **not** invoked when an300 error-handler node itself raises -- handler failures fail the301 run.302 timeout: Default timeout policy for nodes that don't specify their303 own via `add_node(..., timeout=...)`. Also applies to304 error-handler nodes. Accepts a `TimeoutPolicy`, a number of305 seconds (`float`), or a `timedelta`.306 307 Returns:308 Self: The builder instance, for chaining.309 310 Example:311 ```python312 graph = (313 StateGraph(State)314 .set_node_defaults(315 retry_policy=RetryPolicy(max_attempts=3),316 error_handler=my_fallback_handler,317 )318 .add_node("a", node_a)319 .add_node("b", node_b, retry_policy=custom_retry) # overrides default320 .add_edge(START, "a")321 .compile()322 )323 ```324 """325 defaults = self._node_defaults326 if retry_policy is not None:327 defaults.retry_policy = retry_policy328 if cache_policy is not None:329 defaults.cache_policy = cache_policy330 if error_handler is not None:331 defaults.error_handler = error_handler332 if timeout is not None:333 defaults.timeout = coerce_timeout_policy(timeout)334 return self335 336 @property337 def _all_edges(self) -> set[tuple[str, str]]:338 return self.edges | {339 (start, end) for starts, end in self.waiting_edges for start in starts340 }341 342 def _add_schema(self, schema: type[Any], /, allow_managed: bool = True) -> None:343 if schema not in self.schemas:344 _warn_invalid_state_schema(schema)345 channels, managed, type_hints = _get_channels(schema)346 if managed and not allow_managed:347 names = ", ".join(managed)348 schema_name = getattr(schema, "__name__", "")349 raise ValueError(350 f"Invalid managed channels detected in {schema_name}: {names}."351 " Managed channels are not permitted in Input/Output schema."352 )353 self.schemas[schema] = {**channels, **managed}354 for key, channel in channels.items():355 if key in self.channels:356 if self.channels[key] != channel:357 if isinstance(channel, LastValue):358 pass359 else:360 raise ValueError(361 f"Channel '{key}' already exists with a different type"362 )363 else:364 self.channels[key] = channel365 for key, managed in managed.items():366 if key in self.managed:367 if self.managed[key] != managed:368 raise ValueError(369 f"Managed value '{key}' already exists with a different type"370 )371 else:372 self.managed[key] = managed373 374 @overload375 def add_node(376 self,377 node: StateNode[NodeInputT, ContextT],378 *,379 defer: bool = False,380 metadata: dict[str, Any] | None = None,381 input_schema: None = None,382 retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,383 cache_policy: CachePolicy | None = None,384 error_handler: StateNode[Any, ContextT] | None = None,385 destinations: dict[str, str] | tuple[str, ...] | None = None,386 timeout: float | timedelta | TimeoutPolicy | None = None,387 **kwargs: Unpack[DeprecatedKwargs],388 ) -> Self:389 """Add a new node to the `StateGraph`, input schema is inferred as the state schema.390 391 Will take the name of the function/runnable as the node name.392 393 Args:394 node: The function or runnable this node will run.395 defer: Whether to defer the execution of the node until the run is about to end.396 metadata: The metadata associated with the node.397 input_schema: The input schema for the node. (Default: the graph's state schema)398 retry_policy: The retry policy for the node.399 400 If a sequence is provided, the first matching policy will be applied.401 cache_policy: The cache policy for the node.402 destinations: Destinations that indicate where a node can route to.403 404 Useful for edgeless graphs with nodes that return `Command` objects.405 406 If a `dict` is provided, the keys will be used as the target node names and the values will be used as the labels for the edges.407 408 If a `tuple` is provided, the values will be used as the target node names.409 410 !!! warning411 412 This is only used for graph rendering and doesn't have any effect on the graph execution.413 414 Example:415 ```python416 from typing_extensions import TypedDict417 418 from langchain_core.runnables import RunnableConfig419 from langgraph.graph import START, StateGraph420 421 422 class State(TypedDict):423 x: int424 425 426 def my_node(state: State, config: RunnableConfig) -> State:427 return {"x": state["x"] + 1}428 429 430 builder = StateGraph(State)431 builder.add_node(my_node) # node name will be 'my_node'432 builder.add_edge(START, "my_node")433 graph = builder.compile()434 graph.invoke({"x": 1})435 # {'x': 2}436 ```437 438 Returns:439 Self: The instance of the `StateGraph`, allowing for method chaining.440 """441 ...442 443 @overload444 def add_node(445 self,446 node: StateNode[NodeInputT, ContextT],447 *,448 defer: bool = False,449 metadata: dict[str, Any] | None = None,450 input_schema: type[NodeInputT],451 retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,452 cache_policy: CachePolicy | None = None,453 error_handler: StateNode[Any, ContextT] | None = None,454 destinations: dict[str, str] | tuple[str, ...] | None = None,455 timeout: float | timedelta | TimeoutPolicy | None = None,456 **kwargs: Unpack[DeprecatedKwargs],457 ) -> Self:458 """Add a new node to the `StateGraph` where input schema is specified.459 460 Will take the name of the function/runnable as the node name.461 462 Args:463 node: The function or runnable this node will run.464 defer: Whether to defer the execution of the node until the run is about to end.465 metadata: The metadata associated with the node.466 input_schema: The input schema for the node.467 retry_policy: The retry policy for the node.468 469 If a sequence is provided, the first matching policy will be applied.470 cache_policy: The cache policy for the node.471 destinations: Destinations that indicate where a node can route to.472 473 Useful for edgeless graphs with nodes that return `Command` objects.474 475 If a `dict` is provided, the keys will be used as the target node names and the values will be used as the labels for the edges.476 477 If a `tuple` is provided, the values will be used as the target node names.478 479 !!! warning480 481 This is only used for graph rendering and doesn't have any effect on the graph execution.482 483 Example:484 ```python485 from typing_extensions import TypedDict486 487 from langchain_core.runnables import RunnableConfig488 from langgraph.graph import START, StateGraph489 490 491 class State(TypedDict):492 x: int493 494 495 class NodeInput(TypedDict):496 x: int497 498 499 def my_node(state: NodeInput, config: RunnableConfig) -> State:500 return {"x": state["x"] + 1}501 502 503 builder = StateGraph(State)504 builder.add_node(my_node, input_schema=NodeInput) # node name will be 'my_node'505 builder.add_edge(START, "my_node")506 graph = builder.compile()507 graph.invoke({"x": 1})508 # {'x': 2}509 ```510 511 Returns:512 Self: The instance of the `StateGraph`, allowing for method chaining.513 """514 ...515 516 @overload517 def add_node(518 self,519 node: str,520 action: StateNode[NodeInputT, ContextT],521 *,522 defer: bool = False,523 metadata: dict[str, Any] | None = None,524 input_schema: None = None,525 retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,526 cache_policy: CachePolicy | None = None,527 error_handler: StateNode[Any, ContextT] | None = None,528 destinations: dict[str, str] | tuple[str, ...] | None = None,529 timeout: float | timedelta | TimeoutPolicy | None = None,530 **kwargs: Unpack[DeprecatedKwargs],531 ) -> Self:532 """Add a new node to the `StateGraph`, input schema is inferred as the state schema.533 534 Args:535 node: The name of the node.536 action: The function or runnable this node will run.537 defer: Whether to defer the execution of the node until the run is about to end.538 metadata: The metadata associated with the node.539 input_schema: The input schema for the node. (Default: the graph's state schema)540 retry_policy: The retry policy for the node.541 542 If a sequence is provided, the first matching policy will be applied.543 cache_policy: The cache policy for the node.544 destinations: Destinations that indicate where a node can route to.545 546 Useful for edgeless graphs with nodes that return `Command` objects.547 548 If a `dict` is provided, the keys will be used as the target node names and the values will be used as the labels for the edges.549 550 If a `tuple` is provided, the values will be used as the target node names.551 552 !!! warning553 554 This is only used for graph rendering and doesn't have any effect on the graph execution.555 556 Example:557 ```python558 from typing_extensions import TypedDict559 560 from langchain_core.runnables import RunnableConfig561 from langgraph.graph import START, StateGraph562 563 564 class State(TypedDict):565 x: int566 567 568 def my_node(state: State, config: RunnableConfig) -> State:569 return {"x": state["x"] + 1}570 571 572 builder = StateGraph(State)573 builder.add_node("my_fair_node", my_node)574 builder.add_edge(START, "my_fair_node")575 graph = builder.compile()576 graph.invoke({"x": 1})577 # {'x': 2}578 ```579 580 Returns:581 Self: The instance of the `StateGraph`, allowing for method chaining.582 """583 ...584 585 @overload586 def add_node(587 self,588 node: str | StateNode[NodeInputT, ContextT],589 action: StateNode[NodeInputT, ContextT] | None = None,590 *,591 defer: bool = False,592 metadata: dict[str, Any] | None = None,593 input_schema: type[NodeInputT],594 retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,595 cache_policy: CachePolicy | None = None,596 error_handler: StateNode[Any, ContextT] | None = None,597 destinations: dict[str, str] | tuple[str, ...] | None = None,598 timeout: float | timedelta | TimeoutPolicy | None = None,599 **kwargs: Unpack[DeprecatedKwargs],600 ) -> Self:601 """Add a new node to the `StateGraph`, input schema is specified.602 603 Args:604 node: The function or runnable this node will run.605 606 If a string is provided, it will be used as the node name, and action will be used as the function or runnable.607 action: The action associated with the node.608 609 Will be used as the node function or runnable if `node` is a string (node name).610 defer: Whether to defer the execution of the node until the run is about to end.611 metadata: The metadata associated with the node.612 input_schema: The input schema for the node.613 retry_policy: The retry policy for the node.614 615 If a sequence is provided, the first matching policy will be applied.616 cache_policy: The cache policy for the node.617 destinations: Destinations that indicate where a node can route to.618 619 Useful for edgeless graphs with nodes that return `Command` objects.620 621 If a `dict` is provided, the keys will be used as the target node names and the values will be used as the labels for the edges.622 623 If a `tuple` is provided, the values will be used as the target node names.624 625 !!! warning626 627 This is only used for graph rendering and doesn't have any effect on the graph execution.628 629 Example:630 ```python631 from typing_extensions import TypedDict632 633 from langchain_core.runnables import RunnableConfig634 from langgraph.graph import START, StateGraph635 636 637 class State(TypedDict):638 x: int639 640 641 class NodeInput(TypedDict):642 x: int643 644 645 def my_node(state: NodeInput, config: RunnableConfig) -> State:646 return {"x": state["x"] + 1}647 648 649 builder = StateGraph(State)650 builder.add_node("my_fair_node", my_node, input_schema=NodeInput)651 builder.add_edge(START, "my_fair_node")652 graph = builder.compile()653 graph.invoke({"x": 1})654 # {'x': 2}655 ```656 657 Returns:658 Self: The instance of the `StateGraph`, allowing for method chaining.659 """660 ...661 662 def add_node(663 self,664 node: str | StateNode[NodeInputT, ContextT],665 action: StateNode[NodeInputT, ContextT] | None = None,666 *,667 defer: bool = False,668 metadata: dict[str, Any] | None = None,669 input_schema: type[NodeInputT] | None = None,670 retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,671 cache_policy: CachePolicy | None = None,672 error_handler: StateNode[Any, ContextT] | None = None,673 destinations: dict[str, str] | tuple[str, ...] | None = None,674 timeout: float | timedelta | TimeoutPolicy | None = None,675 **kwargs: Unpack[DeprecatedKwargs],676 ) -> Self:677 """Add a new node to the `StateGraph`.678 679 Args:680 node: The function or runnable this node will run.681 682 If a string is provided, it will be used as the node name, and action will be used as the function or runnable.683 action: The action associated with the node.684 685 Will be used as the node function or runnable if `node` is a string (node name).686 defer: Whether to defer the execution of the node until the run is about to end.687 metadata: The metadata associated with the node.688 input_schema: The input schema for the node. (Default: the graph's state schema)689 retry_policy: The retry policy for the node.690 691 If a sequence is provided, the first matching policy will be applied.692 cache_policy: The cache policy for the node.693 error_handler: Optional node-level error handler callable for this node.694 destinations: Destinations that indicate where a node can route to.695 696 Useful for edgeless graphs with nodes that return `Command` objects.697 698 If a `dict` is provided, the keys will be used as the target node names and the values will be used as the labels for the edges.699 700 If a `tuple` is provided, the values will be used as the target node names.701 702 !!! warning703 704 This is only used for graph rendering and doesn't have any effect on the graph execution.705 timeout: Timeout for each node attempt. A number or `timedelta` is706 a hard wall-clock cap and is not refreshed. Use `TimeoutPolicy`707 to configure both a wall-clock `run_timeout` and an708 `idle_timeout` refreshed by progress signals. When exceeded, a709 [`NodeTimeoutError`][langgraph.errors.NodeTimeoutError] is raised710 and the retry policy (if any) decides whether to retry. Timeouts711 are supported only for async nodes; sync nodes cannot be safely712 cancelled in-process.713 714 Example:715 ```python716 from typing_extensions import TypedDict717 718 from langchain_core.runnables import RunnableConfig719 from langgraph.graph import START, StateGraph720 721 722 class State(TypedDict):723 x: int724 725 726 def my_node(state: State, config: RunnableConfig) -> State:727 return {"x": state["x"] + 1}728 729 730 builder = StateGraph(State)731 builder.add_node(my_node) # node name will be 'my_node'732 builder.add_edge(START, "my_node")733 graph = builder.compile()734 graph.invoke({"x": 1})735 # {'x': 2}736 ```737 738 Example: Customize the name:739 ```python740 builder = StateGraph(State)741 builder.add_node("my_fair_node", my_node)742 builder.add_edge(START, "my_fair_node")743 graph = builder.compile()744 graph.invoke({"x": 1})745 # {'x': 2}746 ```747 748 Returns:749 Self: The instance of the `StateGraph`, allowing for method chaining.750 """751 if (retry := kwargs.get("retry", MISSING)) is not MISSING:752 warnings.warn(753 "`retry` is deprecated and will be removed. Please use `retry_policy` instead.",754 category=LangGraphDeprecatedSinceV05,755 )756 if retry_policy is None:757 retry_policy = retry # type: ignore[assignment]758 759 if (input_ := kwargs.get("input", MISSING)) is not MISSING:760 warnings.warn(761 "`input` is deprecated and will be removed. Please use `input_schema` instead.",762 category=LangGraphDeprecatedSinceV05,763 )764 if input_schema is None:765 input_schema = cast(type[NodeInputT] | None, input_)766 timeout = coerce_timeout_policy(timeout)767 768 if not isinstance(node, str):769 action = node770 if isinstance(action, Runnable):771 node = action.get_name()772 else:773 node = getattr(action, "__name__", action.__class__.__name__)774 if node is None:775 raise ValueError(776 "Node name must be provided if action is not a function"777 )778 if self.compiled:779 logger.warning(780 "Adding a node to a graph that has already been compiled. This will "781 "not be reflected in the compiled graph."782 )783 if not isinstance(node, str):784 action = node785 node = cast(str, getattr(action, "name", getattr(action, "__name__", None)))786 if node is None:787 raise ValueError(788 "Node name must be provided if action is not a function"789 )790 if action is None:791 raise RuntimeError792 if node in self.nodes:793 raise ValueError(f"Node `{node}` already present.")794 if node == END or node == START:795 raise ValueError(f"Node `{node}` is reserved.")796 797 for character in (NS_SEP, NS_END):798 if character in node:799 raise ValueError(800 f"'{character}' is a reserved character and is not allowed in the node names."801 )802 803 inferred_input_schema = None804 805 ends: tuple[str, ...] | dict[str, str] = EMPTY_SEQ806 try:807 if (808 isfunction(action)809 or ismethod(action)810 or ismethod(getattr(action, "__call__", None))811 ) and (812 hints := get_type_hints(getattr(action, "__call__"))813 or get_type_hints(action)814 ):815 if input_schema is None:816 first_parameter_name = next(817 iter(818 inspect.signature(819 cast(FunctionType, action)820 ).parameters.keys()821 )822 )823 if input_hint := hints.get(first_parameter_name):824 if isinstance(input_hint, type) and get_type_hints(input_hint):825 inferred_input_schema = input_hint826 if rtn := hints.get("return"):827 # Handle Union types828 rtn_origin = get_origin(rtn)829 if rtn_origin is Union:830 rtn_args = get_args(rtn)831 # Look for Command in the union832 for arg in rtn_args:833 arg_origin = get_origin(arg)834 if arg_origin is Command:835 rtn = arg836 rtn_origin = arg_origin837 break838 839 # Check if it's a Command type840 if (841 rtn_origin is Command842 and (rargs := get_args(rtn))843 and get_origin(rargs[0]) is Literal844 and (vals := get_args(rargs[0]))845 ):846 ends = vals847 except (NameError, TypeError, StopIteration):848 pass849 850 if destinations is not None:851 ends = destinations852 853 resolved_input_schema: type[Any] = (854 input_schema or inferred_input_schema or self.state_schema855 )856 handler_node_name: str | None = None857 if error_handler is not None:858 handler_node_name = f"__error_handler__{node}"859 if handler_node_name in self.nodes:860 raise ValueError(861 f"Auto-generated error handler node `{handler_node_name}` already exists."862 )863 self.nodes[handler_node_name] = StateNodeSpec[Any, ContextT](864 coerce_to_runnable(error_handler, name=handler_node_name, trace=False), # type: ignore[arg-type]865 metadata=None,866 input_schema=resolved_input_schema,867 retry_policy=None,868 cache_policy=None,869 is_error_handler=True,870 )871 872 if input_schema is not None:873 self.nodes[node] = StateNodeSpec[NodeInputT, ContextT](874 coerce_to_runnable(action, name=node, trace=False), # type: ignore[arg-type]875 metadata,876 input_schema=input_schema,877 retry_policy=retry_policy,878 cache_policy=cache_policy,879 error_handler_node=handler_node_name,880 ends=ends,881 defer=defer,882 timeout=timeout,883 )884 elif inferred_input_schema is not None:885 self.nodes[node] = StateNodeSpec(886 coerce_to_runnable(action, name=node, trace=False), # type: ignore[arg-type]887 metadata,888 input_schema=inferred_input_schema,889 retry_policy=retry_policy,890 cache_policy=cache_policy,891 error_handler_node=handler_node_name,892 ends=ends,893 defer=defer,894 timeout=timeout,895 )896 else:897 self.nodes[node] = StateNodeSpec[StateT, ContextT](898 coerce_to_runnable(action, name=node, trace=False), # type: ignore[arg-type]899 metadata,900 input_schema=self.state_schema,901 retry_policy=retry_policy,902 cache_policy=cache_policy,903 error_handler_node=handler_node_name,904 ends=ends,905 defer=defer,906 timeout=timeout,907 )908 909 input_schema = input_schema or inferred_input_schema910 if input_schema is not None:911 self._add_schema(input_schema)912 913 return self914 915 def add_edge(self, start_key: str | list[str], end_key: str) -> Self:916 """Add a directed edge from the start node (or list of start nodes) to the end node.917 918 When a single start node is provided, the graph will wait for that node to complete919 before executing the end node. When multiple start nodes are provided,920 the graph will wait for ALL of the start nodes to complete before executing the end node.921 922 Args:923 start_key: The key(s) of the start node(s) of the edge.924 end_key: The key of the end node of the edge.925 926 Raises:927 ValueError: If the start key is `'END'` or if the start key or end key is not present in the graph.928 929 Returns:930 Self: The instance of the `StateGraph`, allowing for method chaining.931 """932 if self.compiled:933 logger.warning(934 "Adding an edge to a graph that has already been compiled. This will "935 "not be reflected in the compiled graph."936 )937 938 if isinstance(start_key, str):939 if start_key == END:940 raise ValueError("END cannot be a start node")941 if end_key == START:942 raise ValueError("START cannot be an end node")943 944 # run this validation only for non-StateGraph graphs945 if not hasattr(self, "channels") and start_key in set(946 start for start, _ in self.edges947 ):948 raise ValueError(949 f"Already found path for node '{start_key}'.\n"950 "For multiple edges, use StateGraph with an Annotated state key."951 )952 953 self.edges.add((start_key, end_key))954 return self955 956 for start in start_key:957 if start == END:958 raise ValueError("END cannot be a start node")959 if start not in self.nodes:960 raise ValueError(f"Need to add_node `{start}` first")961 if end_key == START:962 raise ValueError("START cannot be an end node")963 if end_key != END and end_key not in self.nodes:964 raise ValueError(f"Need to add_node `{end_key}` first")965 966 self.waiting_edges.add((tuple(start_key), end_key))967 return self968 969 def add_conditional_edges(970 self,971 source: str,972 path: Callable[..., Hashable | Sequence[Hashable]]973 | Callable[..., Awaitable[Hashable | Sequence[Hashable]]]974 | Runnable[Any, Hashable | Sequence[Hashable]],975 path_map: dict[Hashable, str] | list[str] | None = None,976 ) -> Self:977 """Add a conditional edge from the starting node to any number of destination nodes.978 979 Args:980 source: The starting node. This conditional edge will run when981 exiting this node.982 path: The callable that determines the next node or nodes.983 984 If not specifying `path_map` it should return one or more nodes.985 986 If it returns `'END'`, the graph will stop execution.987 path_map: Optional mapping of paths to node names.988 989 If omitted the paths returned by `path` should be node names.990 991 Returns:992 Self: The instance of the graph, allowing for method chaining.993 994 !!! warning995 Without type hints on the `path` function's return value (e.g., `-> Literal["foo", "__end__"]:`)996 or a path_map, the graph visualization assumes the edge could transition to any node in the graph.997 998 """ # noqa: E501999 if self.compiled:1000 logger.warning(1001 "Adding an edge to a graph that has already been compiled. This will "1002 "not be reflected in the compiled graph."1003 )1004 1005 # find a name for the condition1006 path = coerce_to_runnable(path, name=None, trace=True)1007 name = path.name or "condition"1008 # validate the condition1009 if name in self.branches[source]:1010 raise ValueError(1011 f"Branch with name `{path.name}` already exists for node `{source}`"1012 )1013 # save it1014 self.branches[source][name] = BranchSpec.from_path(path, path_map, True)1015 if schema := self.branches[source][name].input_schema:1016 self._add_schema(schema)1017 return self1018 1019 def add_sequence(1020 self,1021 nodes: Sequence[1022 StateNode[NodeInputT, ContextT]1023 | tuple[str, StateNode[NodeInputT, ContextT]]1024 ],1025 ) -> Self:1026 """Add a sequence of nodes that will be executed in the provided order.1027 1028 Args:1029 nodes: A sequence of `StateNode` (callables that accept a `state` arg) or `(name, StateNode)` tuples.1030 1031 If no names are provided, the name will be inferred from the node object (e.g. a `Runnable` or a `Callable` name).1032 1033 Each node will be executed in the order provided.1034 1035 Raises:1036 ValueError: If the sequence is empty.1037 ValueError: If the sequence contains duplicate node names.1038 1039 Returns:1040 Self: The instance of the `StateGraph`, allowing for method chaining.1041 """1042 if len(nodes) < 1:1043 raise ValueError("Sequence requires at least one node.")1044 1045 previous_name: str | None = None1046 for node in nodes:1047 if isinstance(node, tuple) and len(node) == 2:1048 name, node = node1049 else:1050 name = _get_node_name(node)1051 1052 if name in self.nodes:1053 raise ValueError(1054 f"Node names must be unique: node with the name '{name}' already exists. "1055 "If you need to use two different runnables/callables with the same name (for example, using `lambda`), please provide them as tuples (name, runnable/callable)."1056 )1057 1058 self.add_node(name, node)1059 if previous_name is not None:1060 self.add_edge(previous_name, name)1061 1062 previous_name = name1063 1064 return self1065 1066 def set_entry_point(self, key: str) -> Self:1067 """Specifies the first node to be called in the graph.1068 1069 Equivalent to calling `add_edge(START, key)`.1070 1071 Parameters:1072 key (str): The key of the node to set as the entry point.1073 1074 Returns:1075 Self: The instance of the graph, allowing for method chaining.1076 """1077 return self.add_edge(START, key)1078 1079 def set_conditional_entry_point(1080 self,1081 path: Callable[..., Hashable | Sequence[Hashable]]1082 | Callable[..., Awaitable[Hashable | Sequence[Hashable]]]1083 | Runnable[Any, Hashable | Sequence[Hashable]],1084 path_map: dict[Hashable, str] | list[str] | None = None,1085 ) -> Self:1086 """Sets a conditional entry point in the graph.1087 1088 Args:1089 path: The callable that determines the next node or nodes.1090 1091 If not specifying `path_map` it should return one or more nodes.1092 1093 If it returns END, the graph will stop execution.1094 path_map: Optional mapping of paths to node names.1095 1096 If omitted the paths returned by `path` should be node names.1097 1098 Returns:1099 Self: The instance of the graph, allowing for method chaining.1100 """1101 return self.add_conditional_edges(START, path, path_map)1102 1103 def set_finish_point(self, key: str) -> Self:1104 """Marks a node as a finish point of the graph.1105 1106 If the graph reaches this node, it will cease execution.1107 1108 Parameters:1109 key (str): The key of the node to set as the finish point.1110 1111 Returns:1112 Self: The instance of the graph, allowing for method chaining.1113 """1114 return self.add_edge(key, END)1115 1116 def validate(self, interrupt: Sequence[str] | None = None) -> Self:1117 # assemble sources1118 all_sources = {src for src, _ in self._all_edges}1119 for start, branches in self.branches.items():1120 all_sources.add(start)1121 for name, spec in self.nodes.items():1122 if spec.ends:1123 all_sources.add(name)1124 # validate sources1125 for source in all_sources:1126 if source not in self.nodes and source != START:1127 raise ValueError(f"Found edge starting at unknown node '{source}'")1128 1129 if START not in all_sources:1130 raise ValueError(1131 "Graph must have an entrypoint: add at least one edge from START to another node"1132 )1133 1134 # assemble targets1135 all_targets = {end for _, end in self._all_edges}1136 for start, branches in self.branches.items():1137 for cond, branch in branches.items():1138 if branch.ends is not None:1139 for end in branch.ends.values():1140 if end not in self.nodes and end != END:1141 raise ValueError(1142 f"At '{start}' node, '{cond}' branch found unknown target '{end}'"1143 )1144 all_targets.add(end)1145 else:1146 all_targets.add(END)1147 for node in self.nodes:1148 if node != start:1149 all_targets.add(node)1150 for name, spec in self.nodes.items():1151 if spec.ends:1152 all_targets.update(spec.ends)1153 for target in all_targets:1154 if target not in self.nodes and target != END:1155 raise ValueError(f"Found edge ending at unknown node `{target}`")1156 # validate interrupts1157 if interrupt:1158 for node in interrupt:1159 if node not in self.nodes:1160 raise ValueError(f"Interrupt node `{node}` not found")1161 self.compiled = True1162 return self1163 1164 def compile(1165 self,1166 checkpointer: Checkpointer = None,1167 *,1168 cache: BaseCache | None = None,1169 store: BaseStore | None = None,1170 interrupt_before: All | list[str] | None = None,1171 interrupt_after: All | list[str] | None = None,1172 debug: bool = False,1173 name: str | None = None,1174 transformers: Sequence[Callable[[tuple[str, ...]], Any]] | None = None,1175 ) -> CompiledStateGraph[StateT, ContextT, InputT, OutputT]:1176 """Compiles the `StateGraph` into a `CompiledStateGraph` object.1177 1178 The compiled graph implements the `Runnable` interface and can be invoked,1179 streamed, batched, and run asynchronously.1180 1181 Args:1182 checkpointer: A checkpoint saver object or flag.1183 1184 If provided, this `Checkpointer` serves as a fully versioned "short-term memory" for the graph,1185 allowing it to be paused, resumed, and replayed from any point.1186 1187 If `None`, it may inherit the parent graph's checkpointer when used as a subgraph.1188 1189 If `False`, it will not use or inherit any checkpointer.1190 1191 **Important**: When a checkpointer is enabled, you should pass a `thread_id`1192 in the config when invoking the graph:1193 1194 ```python1195 config = {"configurable": {"thread_id": "my-thread"}}1196 graph.invoke(inputs, config)1197 ```1198 1199 The `thread_id` is the key used to store and retrieve checkpoints. Use a1200 unique ID for independent runs, or reuse the same ID to accumulate state