Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
main.py4315 linesDownload Raw Back to pregel
1from __future__ import annotations2 3import asyncio4import concurrent5import concurrent.futures6import contextlib7import queue8import warnings9import weakref10from collections import defaultdict, deque11from collections.abc import (12    AsyncIterator,13    Awaitable,14    Callable,15    Iterator,16    Mapping,17    Sequence,18)19from dataclasses import is_dataclass, replace20from datetime import timedelta21from functools import partial22from inspect import isclass23from typing import (24    Any,25    Generic,26    Literal,27    cast,28    get_type_hints,29    overload,30)31from uuid import UUID, uuid532 33from langchain_core._api import beta34from langchain_core.globals import get_debug35from langchain_core.runnables import (36    RunnableSequence,37)38from langchain_core.runnables.base import Input, Output39from langchain_core.runnables.config import (40    RunnableConfig,41    get_async_callback_manager_for_config,42    get_callback_manager_for_config,43)44from langchain_core.runnables.graph import Graph45from langchain_core.runnables.schema import StreamEvent46from langgraph.cache.base import BaseCache47from langgraph.checkpoint.base import (48    BaseCheckpointSaver,49    Checkpoint,50    CheckpointTuple,51)52from langgraph.store.base import BaseStore53from pydantic import BaseModel, TypeAdapter54from typing_extensions import Self, Unpack, deprecated, is_typeddict55 56from langgraph._internal import _serde57from langgraph._internal._config import (58    ensure_config,59    merge_configs,60    patch_checkpoint_map,61    patch_config,62    patch_configurable,63    recast_checkpoint_ns,64)65from langgraph._internal._constants import (66    CACHE_NS_WRITES,67    CONF,68    CONFIG_KEY_CACHE,69    CONFIG_KEY_CHECKPOINT_ID,70    CONFIG_KEY_CHECKPOINT_NS,71    CONFIG_KEY_CHECKPOINTER,72    CONFIG_KEY_DURABILITY,73    CONFIG_KEY_NODE_FINISHED,74    CONFIG_KEY_READ,75    CONFIG_KEY_RUNNER_SUBMIT,76    CONFIG_KEY_RUNTIME,77    CONFIG_KEY_SEND,78    CONFIG_KEY_STREAM,79    CONFIG_KEY_STREAM_MESSAGES_V2,80    CONFIG_KEY_TASK_ID,81    CONFIG_KEY_THREAD_ID,82    ERROR,83    INPUT,84    INTERRUPT,85    NS_END,86    NS_SEP,87    NULL_TASK_ID,88    PUSH,89    TASKS,90)91from langgraph._internal._pydantic import create_model92from langgraph._internal._queue import (  # type: ignore[attr-defined]93    AsyncQueue,94    SyncQueue,95)96from langgraph._internal._runnable import (97    Runnable,98    RunnableLike,99    RunnableSeq,100    coerce_to_runnable,101)102from langgraph._internal._timeout import coerce_timeout_policy103from langgraph._internal._typing import MISSING, DeprecatedKwargs104from langgraph.callbacks import (105    GraphInterruptEvent,106    GraphResumeEvent,107    get_async_graph_callback_manager_for_config,108    get_sync_graph_callback_manager_for_config,109)110from langgraph.channels.base import BaseChannel111from langgraph.channels.topic import Topic112from langgraph.config import get_config113from langgraph.constants import END114from langgraph.errors import (115    ErrorCode,116    GraphDrained,117    GraphRecursionError,118    InvalidUpdateError,119    create_error_message,120)121from langgraph.managed.base import ManagedValueSpec122from langgraph.pregel._algo import (123    PregelTaskWrites,124    _scratchpad,125    apply_writes,126    local_read,127    prepare_next_tasks,128)129from langgraph.pregel._call import identifier130from langgraph.pregel._checkpoint import (131    achannels_from_checkpoint,132    channels_from_checkpoint,133    copy_checkpoint,134    create_checkpoint,135    empty_checkpoint,136)137from langgraph.pregel._draw import draw_graph138from langgraph.pregel._io import map_input, read_channels139from langgraph.pregel._loop import (140    AsyncPregelLoop,141    SyncPregelLoop,142)143from langgraph.pregel._messages import (144    StreamMessagesHandler,145    StreamMessagesHandlerV2,146)147from langgraph.pregel._read import DEFAULT_BOUND, PregelNode148from langgraph.pregel._retry import RetryPolicy149from langgraph.pregel._runner import PregelRunner150from langgraph.pregel._tools import StreamToolCallHandler151from langgraph.pregel._utils import (152    get_new_channel_versions,153    validate_timeout_supported,154)155from langgraph.pregel._validate import validate_graph, validate_keys156from langgraph.pregel._write import ChannelWrite, ChannelWriteEntry157from langgraph.pregel.debug import get_bolded_text, get_colored_text, tasks_w_writes158from langgraph.pregel.protocol import PregelProtocol, StreamChunk, StreamProtocol159from langgraph.runtime import (160    DEFAULT_RUNTIME,161    BaseUser,162    RunControl,163    Runtime,164    ServerInfo,165)166from langgraph.stream._mux import StreamMux167from langgraph.stream._types import StreamTransformer168from langgraph.stream.run_stream import AsyncGraphRunStream, GraphRunStream169from langgraph.stream.transformers import (170    LifecycleTransformer,171    MessagesTransformer,172    SubgraphTransformer,173    ValuesTransformer,174)175from langgraph.types import (176    All,177    CachePolicy,178    Checkpointer,179    Command,180    Durability,181    GraphOutput,182    Interrupt,183    Send,184    StateSnapshot,185    StateUpdate,186    StreamMode,187    StreamPart,188    TimeoutPolicy,189    ensure_valid_checkpointer,190)191from langgraph.typing import ContextT, InputT, OutputT, StateT192from langgraph.warnings import LangGraphDeprecatedSinceV10193 194try:195    from langchain_core.tracers._streaming import _StreamingCallbackHandler196except ImportError:197    _StreamingCallbackHandler = None  # type: ignore198 199__all__ = ("NodeBuilder", "Pregel")200 201_WriteValue = Callable[[Input], Output] | Any202 203 204class NodeBuilder:205    __slots__ = (206        "_channels",207        "_triggers",208        "_tags",209        "_metadata",210        "_writes",211        "_bound",212        "_retry_policy",213        "_cache_policy",214        "_timeout",215    )216 217    _channels: str | list[str]218    _triggers: list[str]219    _tags: list[str]220    _metadata: dict[str, Any]221    _writes: list[ChannelWriteEntry]222    _bound: Runnable223    _retry_policy: list[RetryPolicy]224    _cache_policy: CachePolicy | None225    _timeout: TimeoutPolicy | None226 227    def __init__(228        self,229    ) -> None:230        self._channels = []231        self._triggers = []232        self._tags = []233        self._metadata = {}234        self._writes = []235        self._bound = DEFAULT_BOUND236        self._retry_policy = []237        self._cache_policy = None238        self._timeout = None239 240    def subscribe_only(241        self,242        channel: str,243    ) -> Self:244        """Subscribe to a single channel."""245        if not self._channels:246            self._channels = channel247        else:248            raise ValueError(249                "Cannot subscribe to single channels when other channels are already subscribed to"250            )251 252        self._triggers.append(channel)253 254        return self255 256    def subscribe_to(257        self,258        *channels: str,259        read: bool = True,260    ) -> Self:261        """Add channels to subscribe to.262 263        Node will be invoked when any of these channels are updated, with a dict of the264        channel values as input.265 266        Args:267            channels: Channel name(s) to subscribe to268            read: If `True`, the channels will be included in the input to the node.269                Otherwise, they will trigger the node without being sent in input.270 271        Returns:272            Self for chaining273        """274        if isinstance(self._channels, str):275            raise ValueError(276                "Cannot subscribe to channels when subscribed to a single channel"277            )278        if read:279            if not self._channels:280                self._channels = list(channels)281            else:282                self._channels.extend(channels)283 284        if isinstance(channels, str):285            self._triggers.append(channels)286        else:287            self._triggers.extend(channels)288 289        return self290 291    def read_from(292        self,293        *channels: str,294    ) -> Self:295        """Adds the specified channels to read from, without subscribing to them."""296        assert isinstance(self._channels, list), (297            "Cannot read additional channels when subscribed to single channels"298        )299        self._channels.extend(channels)300        return self301 302    def do(303        self,304        node: RunnableLike,305    ) -> Self:306        """Adds the specified node."""307        if self._bound is not DEFAULT_BOUND:308            self._bound = RunnableSeq(309                self._bound, coerce_to_runnable(node, name=None, trace=True)310            )311        else:312            self._bound = coerce_to_runnable(node, name=None, trace=True)313        return self314 315    def write_to(316        self,317        *channels: str | ChannelWriteEntry,318        **kwargs: _WriteValue,319    ) -> Self:320        """Add channel writes.321 322        Args:323            *channels: Channel names to write to.324            **kwargs: Channel name and value mappings.325 326        Returns:327            Self for chaining328        """329        self._writes.extend(330            ChannelWriteEntry(c) if isinstance(c, str) else c for c in channels331        )332        self._writes.extend(333            ChannelWriteEntry(k, mapper=v)334            if callable(v)335            else ChannelWriteEntry(k, value=v)336            for k, v in kwargs.items()337        )338 339        return self340 341    def meta(self, *tags: str, **metadata: Any) -> Self:342        """Add tags or metadata to the node."""343        self._tags.extend(tags)344        self._metadata.update(metadata)345        return self346 347    def add_retry_policies(self, *policies: RetryPolicy) -> Self:348        """Adds retry policies to the node."""349        self._retry_policy.extend(policies)350        return self351 352    def add_cache_policy(self, policy: CachePolicy) -> Self:353        """Adds cache policies to the node."""354        self._cache_policy = policy355        return self356 357    def set_timeout(self, timeout: float | timedelta | TimeoutPolicy | None) -> Self:358        """Set the per-attempt timeout policy for this node."""359        self._timeout = coerce_timeout_policy(timeout)360        return self361 362    def build(self) -> PregelNode:363        """Builds the node."""364        return PregelNode(365            channels=self._channels,366            triggers=self._triggers,367            tags=self._tags,368            metadata=self._metadata,369            writers=[ChannelWrite(self._writes)],370            bound=self._bound,371            retry_policy=self._retry_policy,372            cache_policy=self._cache_policy,373            timeout=self._timeout,374        )375 376 377# Kwargs that ``stream_events(version="v3")`` / ``astream_events(version="v3")``378# manage internally and must not be overridden by callers. ``stream_mode`` is379# derived from the transformer mux; ``subgraphs`` is forced True so nested380# namespaces flow through scoped muxes. Forwarding either to the inner381# ``stream(...)`` would silently break v3's invariants, so we raise instead.382_V3_INVARIANT_KWARGS: tuple[str, ...] = ("stream_mode", "subgraphs")383 384 385def _reject_v3_invariant_kwargs(kwargs: dict[str, Any]) -> None:386    collisions = [k for k in _V3_INVARIANT_KWARGS if k in kwargs]387    if collisions:388        raise TypeError(389            "stream_events(version='v3') / astream_events(version='v3') do "390            f"not accept {', '.join(collisions)}; v3 owns these "391            "(stream_mode is built from the transformer mux, subgraphs is "392            "forced True so nested namespaces flow through scoped muxes)."393        )394 395 396def _collect_stream_modes(mux: Any) -> list[StreamMode]:397    """Return the union of `required_stream_modes` across registered transformers.398 399    Transformers declare the stream modes they need to function, and400    `stream_events(version="v3")` asks the graph for exactly that union — no hardcoded401    default set. If zero transformers declare a given mode, the graph402    does not stream events for it.403    """404    modes: set[StreamMode] = set()405    for transformer in mux._transformers:406        modes.update(407            cast(408                "tuple[StreamMode, ...]",409                getattr(transformer, "required_stream_modes", ()),410            )411        )412    return list(modes)413 414 415def _normalize_stream_transformer_factories(416    specs: Sequence[Callable[[tuple[str, ...]], Any]] | None,417) -> list[Callable[[tuple[str, ...]], Any]]:418    """Normalize stream transformer specs to scoped factories.419 420    A stream transformer spec is a callable that accepts421    `scope: tuple[str, ...]` and returns a fresh `StreamTransformer`.422    Transformer classes work when their constructor follows the same423    shape. Pre-built instances are rejected because they cannot be424    cloned into subgraph scopes.425    """426    factories: list[Callable[[tuple[str, ...]], Any]] = []427    for spec in specs or ():428        if isinstance(spec, StreamTransformer):429            raise TypeError(430                "stream_events(version='v3') transformers must be scope-aware callables, "431                f"got pre-built instance {type(spec).__name__}. Pass the "432                "transformer class or a factory like "433                "`lambda scope: MyTransformer(scope, ...)`."434            )435        if not callable(spec):436            raise TypeError(437                "stream_events(version='v3') transformers must be scope-aware callables, "438                f"got {type(spec).__name__}."439            )440 441        def factory(scope: tuple[str, ...], _spec: Callable[..., Any] = spec) -> Any:442            return _spec(scope)443 444        factories.append(factory)445    return factories446 447 448class Pregel(449    PregelProtocol[StateT, ContextT, InputT, OutputT],450    Generic[StateT, ContextT, InputT, OutputT],451):452    """Pregel manages the runtime behavior for LangGraph applications.453 454    ## Overview455 456    Pregel combines [**actors**](https://en.wikipedia.org/wiki/Actor_model)457    and **channels** into a single application.458    **Actors** read data from channels and write data to channels.459    Pregel organizes the execution of the application into multiple steps,460    following the **Pregel Algorithm**/**Bulk Synchronous Parallel** model.461 462    Each step consists of three phases:463 464    - **Plan**: Determine which **actors** to execute in this step. For example,465        in the first step, select the **actors** that subscribe to the special466        **input** channels; in subsequent steps,467        select the **actors** that subscribe to channels updated in the previous step.468    - **Execution**: Execute all selected **actors** in parallel,469        until all complete, or one fails, or a timeout is reached. During this470        phase, channel updates are invisible to actors until the next step.471    - **Update**: Update the channels with the values written by the **actors**472        in this step.473 474    Repeat until no **actors** are selected for execution, or a maximum number of475    steps is reached.476 477    ## Actors478 479    An **actor** is a `PregelNode`.480    It subscribes to channels, reads data from them, and writes data to them.481    It can be thought of as an **actor** in the Pregel algorithm.482    `PregelNodes` implement LangChain's483    Runnable interface.484 485    ## Channels486 487    Channels are used to communicate between actors (`PregelNodes`).488    Each channel has a value type, an update type, and an update function – which489    takes a sequence of updates and490    modifies the stored value. Channels can be used to send data from one chain to491    another, or to send data from a chain to itself in a future step. LangGraph492    provides a number of built-in channels:493 494    ### Basic channels: LastValue and Topic495 496    - `LastValue`: The default channel, stores the last value sent to the channel,497       useful for input and output values, or for sending data from one step to the next498    - `Topic`: A configurable PubSub Topic, useful for sending multiple values499       between *actors*, or for accumulating output. Can be configured to deduplicate500       values, and/or to accumulate values over the course of multiple steps.501 502    ### Advanced channels: Context and BinaryOperatorAggregate503 504    - `Context`: exposes the value of a context manager, managing its lifecycle.505        Useful for accessing external resources that require setup and/or teardown. e.g.506        `client = Context(httpx.Client)`507    - `BinaryOperatorAggregate`: stores a persistent value, updated by applying508        a binary operator to the current value and each update509        sent to the channel, useful for computing aggregates over multiple steps. e.g.510        `total = BinaryOperatorAggregate(int, operator.add)`511 512    ## Examples513 514    Most users will interact with Pregel via a515    [StateGraph (Graph API)][langgraph.graph.StateGraph] or via an516    [entrypoint (Functional API)][langgraph.func.entrypoint].517 518    However, for **advanced** use cases, Pregel can be used directly. If you're519    not sure whether you need to use Pregel directly, then the answer is probably no520    - you should use the Graph API or Functional API instead. These are higher-level521    interfaces that will compile down to Pregel under the hood.522 523    Here are some examples to give you a sense of how it works:524 525    Example: Single node application526        ```python527        from langgraph.channels import EphemeralValue528        from langgraph.pregel import Pregel, NodeBuilder529 530        node1 = (531            NodeBuilder().subscribe_only("a")532            .do(lambda x: x + x)533            .write_to("b")534        )535 536        app = Pregel(537            nodes={"node1": node1},538            channels={539                "a": EphemeralValue(str),540                "b": EphemeralValue(str),541            },542            input_channels=["a"],543            output_channels=["b"],544        )545 546        app.invoke({"a": "foo"})547        ```548 549        ```con550        {'b': 'foofoo'}551        ```552 553    Example: Using multiple nodes and multiple output channels554        ```python555        from langgraph.channels import LastValue, EphemeralValue556        from langgraph.pregel import Pregel, NodeBuilder557 558        node1 = (559            NodeBuilder().subscribe_only("a")560            .do(lambda x: x + x)561            .write_to("b")562        )563 564        node2 = (565            NodeBuilder().subscribe_to("b")566            .do(lambda x: x["b"] + x["b"])567            .write_to("c")568        )569 570 571        app = Pregel(572            nodes={"node1": node1, "node2": node2},573            channels={574                "a": EphemeralValue(str),575                "b": LastValue(str),576                "c": EphemeralValue(str),577            },578            input_channels=["a"],579            output_channels=["b", "c"],580        )581 582        app.invoke({"a": "foo"})583        ```584 585        ```con586        {'b': 'foofoo', 'c': 'foofoofoofoo'}587        ```588 589    Example: Using a Topic channel590        ```python591        from langgraph.channels import LastValue, EphemeralValue, Topic592        from langgraph.pregel import Pregel, NodeBuilder593 594        node1 = (595            NodeBuilder().subscribe_only("a")596            .do(lambda x: x + x)597            .write_to("b", "c")598        )599 600        node2 = (601            NodeBuilder().subscribe_only("b")602            .do(lambda x: x + x)603            .write_to("c")604        )605 606 607        app = Pregel(608            nodes={"node1": node1, "node2": node2},609            channels={610                "a": EphemeralValue(str),611                "b": EphemeralValue(str),612                "c": Topic(str, accumulate=True),613            },614            input_channels=["a"],615            output_channels=["c"],616        )617 618        app.invoke({"a": "foo"})619        ```620 621        ```pycon622        {"c": ["foofoo", "foofoofoofoo"]}623        ```624 625    Example: Using a `BinaryOperatorAggregate` channel626        ```python627        from langgraph.channels import EphemeralValue, BinaryOperatorAggregate628        from langgraph.pregel import Pregel, NodeBuilder629 630 631        node1 = (632            NodeBuilder().subscribe_only("a")633            .do(lambda x: x + x)634            .write_to("b", "c")635        )636 637        node2 = (638            NodeBuilder().subscribe_only("b")639            .do(lambda x: x + x)640            .write_to("c")641        )642 643 644        def reducer(current, update):645            if current:646                return current + " | " + update647            else:648                return update649 650 651        app = Pregel(652            nodes={"node1": node1, "node2": node2},653            channels={654                "a": EphemeralValue(str),655                "b": EphemeralValue(str),656                "c": BinaryOperatorAggregate(str, operator=reducer),657            },658            input_channels=["a"],659            output_channels=["c"],660        )661 662        app.invoke({"a": "foo"})663        ```664 665        ```con666        {'c': 'foofoo | foofoofoofoo'}667        ```668 669    Example: Introducing a cycle670        This example demonstrates how to introduce a cycle in the graph, by having671        a chain write to a channel it subscribes to.672 673        Execution will continue until a `None` value is written to the channel.674 675        ```python676        from langgraph.channels import EphemeralValue677        from langgraph.pregel import Pregel, NodeBuilder, ChannelWriteEntry678 679        example_node = (680            NodeBuilder()681            .subscribe_only("value")682            .do(lambda x: x + x if len(x) < 10 else None)683            .write_to(ChannelWriteEntry(channel="value", skip_none=True))684        )685 686        app = Pregel(687            nodes={"example_node": example_node},688            channels={689                "value": EphemeralValue(str),690            },691            input_channels=["value"],692            output_channels=["value"],693        )694 695        app.invoke({"value": "a"})696        ```697 698        ```con699        {'value': 'aaaaaaaaaaaaaaaa'}700        ```701    """702 703    nodes: dict[str, PregelNode]704 705    channels: dict[str, BaseChannel | ManagedValueSpec]706 707    stream_mode: StreamMode = "values"708    """Mode to stream output, defaults to 'values'."""709 710    stream_eager: bool = False711    """Whether to force emitting stream events eagerly, automatically turned on712    for stream_mode "messages" and "custom"."""713 714    output_channels: str | Sequence[str]715 716    stream_channels: str | Sequence[str] | None = None717    """Channels to stream, defaults to all channels not in reserved channels"""718 719    interrupt_after_nodes: All | Sequence[str]720 721    interrupt_before_nodes: All | Sequence[str]722 723    input_channels: str | Sequence[str]724 725    step_timeout: float | None = None726    """Maximum time to wait for a step to complete, in seconds."""727 728    debug: bool729    """Whether to print debug information during execution."""730 731    checkpointer: Checkpointer = None732    """`Checkpointer` used to save and load graph state."""733 734    store: BaseStore | None = None735    """Memory store to use for SharedValues."""736 737    cache: BaseCache | None = None738    """Cache to use for storing node results."""739 740    retry_policy: Sequence[RetryPolicy] = ()741    """Retry policies to use when running tasks. Empty set disables retries."""742 743    cache_policy: CachePolicy | None = None744    """Cache policy to use for all nodes. Can be overridden by individual nodes."""745 746    context_schema: type[ContextT] | None = None747    """Specifies the schema for the context object that will be passed to the workflow."""748 749    config: RunnableConfig | None = None750 751    name: str = "LangGraph"752 753    trigger_to_nodes: Mapping[str, Sequence[str]]754    node_error_handler_map: Mapping[str, str]755 756    def __init__(757        self,758        *,759        nodes: dict[str, PregelNode | NodeBuilder],760        channels: dict[str, BaseChannel | ManagedValueSpec] | None,761        auto_validate: bool = True,762        stream_mode: StreamMode = "values",763        stream_eager: bool = False,764        output_channels: str | Sequence[str],765        stream_channels: str | Sequence[str] | None = None,766        interrupt_after_nodes: All | Sequence[str] = (),767        interrupt_before_nodes: All | Sequence[str] = (),768        input_channels: str | Sequence[str],769        step_timeout: float | None = None,770        debug: bool | None = None,771        checkpointer: Checkpointer = None,772        store: BaseStore | None = None,773        cache: BaseCache | None = None,774        retry_policy: RetryPolicy | Sequence[RetryPolicy] = (),775        cache_policy: CachePolicy | None = None,776        context_schema: type[ContextT] | None = None,777        config: RunnableConfig | None = None,778        trigger_to_nodes: Mapping[str, Sequence[str]] | None = None,779        node_error_handler_map: Mapping[str, str] | None = None,780        name: str = "LangGraph",781        stream_transformers: Sequence[Callable[[tuple[str, ...]], Any]] | None = None,782        **deprecated_kwargs: Unpack[DeprecatedKwargs],783    ) -> None:784        if (785            config_type := deprecated_kwargs.get("config_type", MISSING)786        ) is not MISSING:787            warnings.warn(788                "`config_type` is deprecated and will be removed. Please use `context_schema` instead.",789                category=LangGraphDeprecatedSinceV10,790                stacklevel=2,791            )792 793            if context_schema is None:794                context_schema = cast(type[ContextT], config_type)795 796        checkpointer = ensure_valid_checkpointer(checkpointer)797 798        self.nodes = {799            k: v.build() if isinstance(v, NodeBuilder) else v for k, v in nodes.items()800        }801        self.channels = channels or {}802        if TASKS in self.channels and not isinstance(self.channels[TASKS], Topic):803            raise ValueError(804                f"Channel '{TASKS}' is reserved and cannot be used in the graph."805            )806        else:807            self.channels[TASKS] = Topic(Send, accumulate=False)808        self.stream_mode = stream_mode809        self.stream_eager = stream_eager810        self.output_channels = output_channels811        self.stream_channels = stream_channels812        self.interrupt_after_nodes = interrupt_after_nodes813        self.interrupt_before_nodes = interrupt_before_nodes814        self.input_channels = input_channels815        self.step_timeout = step_timeout816        self.debug = debug if debug is not None else get_debug()817        self.checkpointer = checkpointer818        self.store = store819        self.cache = cache820        self.retry_policy = (821            (retry_policy,) if isinstance(retry_policy, RetryPolicy) else retry_policy822        )823        self.cache_policy = cache_policy824        self.context_schema = context_schema825        self.config = config826        self.trigger_to_nodes = trigger_to_nodes or {}827        self.node_error_handler_map = node_error_handler_map or {}828        self.name = name829        self.stream_transformers: tuple[Callable[[tuple[str, ...]], Any], ...] = tuple(830            stream_transformers or ()831        )832        self._serde_allowlist: set[tuple[str, ...]] | None = None833        if auto_validate:834            self.validate()835 836    def _apply_checkpointer_allowlist(837        self, checkpointer: BaseCheckpointSaver | None838    ) -> BaseCheckpointSaver | None:839        if not _serde.STRICT_MSGPACK_ENABLED:840            return checkpointer841        return _serde.apply_checkpointer_allowlist(checkpointer, self._serde_allowlist)842 843    def get_graph(844        self, config: RunnableConfig | None = None, *, xray: int | bool = False845    ) -> Graph:846        """Return a drawable representation of the computation graph."""847        # gather subgraphs848        if xray:849            subgraphs = {850                k: v.get_graph(851                    config,852                    xray=xray if isinstance(xray, bool) or xray <= 0 else xray - 1,853                )854                for k, v in self.get_subgraphs()855            }856        else:857            subgraphs = {}858 859        return draw_graph(860            merge_configs(self.config, config),861            nodes=self.nodes,862            specs=self.channels,863            input_channels=self.input_channels,864            interrupt_after_nodes=self.interrupt_after_nodes,865            interrupt_before_nodes=self.interrupt_before_nodes,866            trigger_to_nodes=self.trigger_to_nodes,867            checkpointer=self.checkpointer,868            subgraphs=subgraphs,869        )870 871    async def aget_graph(872        self, config: RunnableConfig | None = None, *, xray: int | bool = False873    ) -> Graph:874        """Return a drawable representation of the computation graph."""875 876        # gather subgraphs877        if xray:878            subpregels: dict[str, PregelProtocol] = {879                k: v async for k, v in self.aget_subgraphs()880            }881            subgraphs = {882                k: v883                for k, v in zip(884                    subpregels,885                    await asyncio.gather(886                        *(887                            p.aget_graph(888                                config,889                                xray=xray890                                if isinstance(xray, bool) or xray <= 0891                                else xray - 1,892                            )893                            for p in subpregels.values()894                        )895                    ),896                )897            }898        else:899            subgraphs = {}900 901        return draw_graph(902            merge_configs(self.config, config),903            nodes=self.nodes,904            specs=self.channels,905            input_channels=self.input_channels,906            interrupt_after_nodes=self.interrupt_after_nodes,907            interrupt_before_nodes=self.interrupt_before_nodes,908            trigger_to_nodes=self.trigger_to_nodes,909            checkpointer=self.checkpointer,910            subgraphs=subgraphs,911        )912 913    def _repr_mimebundle_(self, **kwargs: Any) -> dict[str, Any]:914        """Mime bundle used by Jupyter to display the graph"""915        return {916            "text/plain": repr(self),917            "image/png": self.get_graph().draw_mermaid_png(),918        }919 920    def copy(self, update: dict[str, Any] | None = None) -> Self:921        attrs = {k: v for k, v in self.__dict__.items() if k != "__orig_class__"}922        attrs.update(update or {})923        return self.__class__(**attrs)924 925    def with_config(self, config: RunnableConfig | None = None, **kwargs: Any) -> Self:926        """Create a copy of the Pregel object with an updated config."""927        return self.copy(928            {"config": merge_configs(self.config, config, cast(RunnableConfig, kwargs))}929        )930 931    def validate(self) -> Self:932        for name, node in self.nodes.items():933            if node.timeout is not None:934                validate_timeout_supported(node.node or node.bound, name=name)935        validate_graph(936            self.nodes,937            {k: v for k, v in self.channels.items() if isinstance(v, BaseChannel)},938            {k: v for k, v in self.channels.items() if not isinstance(v, BaseChannel)},939            self.input_channels,940            self.output_channels,941            self.stream_channels,942            self.interrupt_after_nodes,943            self.interrupt_before_nodes,944        )945        self.trigger_to_nodes = _trigger_to_nodes(self.nodes)946        return self947 948    @deprecated(949        "`config_schema` is deprecated. Use `get_context_jsonschema` for the relevant schema instead.",950        category=None,951    )952    def config_schema(self, *, include: Sequence[str] | None = None) -> type[BaseModel]:953        warnings.warn(954            "`config_schema` is deprecated. Use `get_context_jsonschema` for the relevant schema instead.",955            category=LangGraphDeprecatedSinceV10,956            stacklevel=2,957        )958 959        include = include or []960        fields = {961            **(962                {"configurable": (self.context_schema, None)}963                if self.context_schema964                else {}965            ),966            **{967                field_name: (field_type, None)968                for field_name, field_type in get_type_hints(RunnableConfig).items()969                if field_name in [i for i in include if i != "configurable"]970            },971        }972        return create_model(self.get_name("Config"), field_definitions=fields)973 974    @deprecated(975        "`get_config_jsonschema` is deprecated. Use `get_context_jsonschema` instead.",976        category=None,977    )978    def get_config_jsonschema(979        self, *, include: Sequence[str] | None = None980    ) -> dict[str, Any]:981        warnings.warn(982            "`get_config_jsonschema` is deprecated. Use `get_context_jsonschema` instead.",983            category=LangGraphDeprecatedSinceV10,984            stacklevel=2,985        )986 987        with warnings.catch_warnings():988            warnings.filterwarnings("ignore", category=LangGraphDeprecatedSinceV10)989            schema = self.config_schema(include=include)990        return schema.model_json_schema()991 992    def get_context_jsonschema(self) -> dict[str, Any] | None:993        if (context_schema := self.context_schema) is None:994            return None995 996        if isclass(context_schema) and issubclass(context_schema, BaseModel):997            return context_schema.model_json_schema()998        elif is_typeddict(context_schema) or is_dataclass(context_schema):999            return TypeAdapter(context_schema).json_schema()1000        else:1001            raise ValueError(1002                f"Invalid context schema type: {context_schema}. Must be a BaseModel, TypedDict or dataclass."1003            )1004 1005    @property1006    def InputType(self) -> Any:1007        if isinstance(self.input_channels, str):1008            channel = self.channels[self.input_channels]1009            if isinstance(channel, BaseChannel):1010                return channel.UpdateType1011 1012    def get_input_schema(self, config: RunnableConfig | None = None) -> type[BaseModel]:1013        config = merge_configs(self.config, config)1014        if isinstance(self.input_channels, str):1015            return super().get_input_schema(config)1016        else:1017            return create_model(1018                self.get_name("Input"),1019                field_definitions={1020                    k: (c.UpdateType, None)1021                    for k in self.input_channels or self.channels.keys()1022                    if (c := self.channels[k]) and isinstance(c, BaseChannel)1023                },1024            )1025 1026    def get_input_jsonschema(1027        self, config: RunnableConfig | None = None1028    ) -> dict[str, Any]:1029        schema = self.get_input_schema(config)1030        return schema.model_json_schema()1031 1032    @property1033    def OutputType(self) -> Any:1034        if isinstance(self.output_channels, str):1035            channel = self.channels[self.output_channels]1036            if isinstance(channel, BaseChannel):1037                return channel.ValueType1038 1039    def get_output_schema(1040        self, config: RunnableConfig | None = None1041    ) -> type[BaseModel]:1042        config = merge_configs(self.config, config)1043        if isinstance(self.output_channels, str):1044            return super().get_output_schema(config)1045        else:1046            return create_model(1047                self.get_name("Output"),1048                field_definitions={1049                    k: (c.ValueType, None)1050                    for k in self.output_channels1051                    if (c := self.channels[k]) and isinstance(c, BaseChannel)1052                },1053            )1054 1055    def get_output_jsonschema(1056        self, config: RunnableConfig | None = None1057    ) -> dict[str, Any]:1058        schema = self.get_output_schema(config)1059        return schema.model_json_schema()1060 1061    @property1062    def stream_channels_list(self) -> Sequence[str]:1063        stream_channels = self.stream_channels_asis1064        return (1065            [stream_channels] if isinstance(stream_channels, str) else stream_channels1066        )1067 1068    @property1069    def stream_channels_asis(self) -> str | Sequence[str]:1070        return self.stream_channels or [1071            k for k in self.channels if isinstance(self.channels[k], BaseChannel)1072        ]1073 1074    def get_subgraphs(1075        self, *, namespace: str | None = None, recurse: bool = False1076    ) -> Iterator[tuple[str, PregelProtocol]]:1077        """Get the subgraphs of the graph.1078 1079        Args:1080            namespace: The namespace to filter the subgraphs by.1081            recurse: Whether to recurse into the subgraphs.1082                If `False`, only the immediate subgraphs will be returned.1083 1084        Returns:1085            An iterator of the `(namespace, subgraph)` pairs.1086        """1087        for name, node in self.nodes.items():1088            # filter by prefix1089            if namespace is not None:1090                if not namespace.startswith(name):1091                    continue1092 1093            # find the subgraph, if any1094            graph = node.subgraphs[0] if node.subgraphs else None1095 1096            # if found, yield recursively1097            if graph:1098                if name == namespace:1099                    yield name, graph1100                    return  # we found it, stop searching1101                if namespace is None:1102                    yield name, graph1103                if recurse and isinstance(graph, Pregel):1104                    if namespace is not None:1105                        namespace = namespace[len(name) + 1 :]1106                    yield from (1107                        (f"{name}{NS_SEP}{n}", s)1108                        for n, s in graph.get_subgraphs(1109                            namespace=namespace, recurse=recurse1110                        )1111                    )1112 1113    async def aget_subgraphs(1114        self, *, namespace: str | None = None, recurse: bool = False1115    ) -> AsyncIterator[tuple[str, PregelProtocol]]:1116        """Get the subgraphs of the graph.1117 1118        Args:1119            namespace: The namespace to filter the subgraphs by.1120            recurse: Whether to recurse into the subgraphs.1121                If `False`, only the immediate subgraphs will be returned.1122 1123        Returns:1124            An iterator of the `(namespace, subgraph)` pairs.1125        """1126        for name, node in self.get_subgraphs(namespace=namespace, recurse=recurse):1127            yield name, node1128 1129    # Mappers for v2 stream coercion (pydantic/dataclass).1130    # Set by CompiledStateGraph; None for base Pregel.1131    _output_mapper: Callable[[Any], Any] | None = None1132    _state_mapper: Callable[[Any], Any] | None = None1133 1134    def _migrate_checkpoint(self, checkpoint: Checkpoint) -> None:1135        """Migrate a saved checkpoint to new channel layout."""1136        if checkpoint["v"] < 4 and checkpoint.get("pending_sends"):1137            pending_sends: list[Send] = checkpoint.pop("pending_sends")1138            checkpoint["channel_values"][TASKS] = pending_sends1139            checkpoint["channel_versions"][TASKS] = max(1140                checkpoint["channel_versions"].values()1141            )1142 1143    def _prepare_state_snapshot(1144        self,1145        config: RunnableConfig,1146        saved: CheckpointTuple | None,1147        recurse: BaseCheckpointSaver | None = None,1148        apply_pending_writes: bool = False,1149    ) -> StateSnapshot:1150        if not saved:1151            return StateSnapshot(1152                values={},1153                next=(),1154                config=config,1155                metadata=None,1156                created_at=None,1157                parent_config=None,1158                tasks=(),1159                interrupts=(),1160            )1161 1162        # migrate checkpoint if needed1163        self._migrate_checkpoint(saved.checkpoint)1164 1165        step = saved.metadata.get("step", -1) + 11166        stop = step + 21167        channels, managed = channels_from_checkpoint(1168            self.channels,1169            saved.checkpoint,1170            saver=self.checkpointer1171            if isinstance(self.checkpointer, BaseCheckpointSaver)1172            else None,1173            config=saved.config,1174        )1175        # tasks for this checkpoint1176        next_tasks = prepare_next_tasks(1177            saved.checkpoint,1178            saved.pending_writes or [],1179            self.nodes,1180            channels,1181            managed,1182            saved.config,1183            step,1184            stop,1185            for_execution=True,1186            store=self.store,1187            checkpointer=(1188                self.checkpointer1189                if isinstance(self.checkpointer, BaseCheckpointSaver)1190                else None1191            ),1192            manager=None,1193        )1194        # get the subgraphs1195        subgraphs = dict(self.get_subgraphs())1196        parent_ns = saved.config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "")1197        task_states: dict[str, RunnableConfig | StateSnapshot] = {}1198        for task in next_tasks.values():1199            if task.name not in subgraphs:1200                continue

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

codekingpro/portable-devtools · Team Ai