Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
state.py1965 linesDownload Raw Back to graph
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

Showing the first 1,200 of 1965 lines. Download the file for the rest.

codekingpro/portable-devtools · Team Ai