Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
remote.py1192 linesDownload Raw Back to pregel
1from __future__ import annotations2 3import logging4from collections.abc import AsyncIterator, Iterator, Sequence5from dataclasses import asdict6from typing import (7    Any,8    Literal,9    cast,10    overload,11)12from uuid import UUID13 14import langsmith as ls15from langchain_core.runnables import RunnableConfig16from langchain_core.runnables.graph import (17    Edge as DrawableEdge,18)19from langchain_core.runnables.graph import (20    Graph as DrawableGraph,21)22from langchain_core.runnables.graph import (23    Node as DrawableNode,24)25from langgraph.checkpoint.base import CheckpointMetadata26from langgraph_sdk.client import (27    LangGraphClient,28    SyncLangGraphClient,29    get_client,30    get_sync_client,31)32from langgraph_sdk.schema import (33    Checkpoint,34    Context,35    QueryParamTypes,36    ThreadState,37)38from langgraph_sdk.schema import (39    Command as CommandSDK,40)41from langgraph_sdk.schema import (42    StreamMode as StreamModeSDK,43)44from typing_extensions import Self45 46from langgraph._internal._config import merge_configs47from langgraph._internal._constants import (48    CONF,49    CONFIG_KEY_CHECKPOINT_ID,50    CONFIG_KEY_CHECKPOINT_MAP,51    CONFIG_KEY_CHECKPOINT_NS,52    CONFIG_KEY_STREAM,53    CONFIG_KEY_TASK_ID,54    INTERRUPT,55    NS_SEP,56)57from langgraph.errors import GraphInterrupt, ParentCommand58from langgraph.pregel.protocol import PregelProtocol, StreamProtocol59from langgraph.types import (60    All,61    Command,62    GraphOutput,63    Interrupt,64    PregelTask,65    StateSnapshot,66    StreamMode,67    StreamPart,68)69 70logger = logging.getLogger(__name__)71 72__all__ = ("RemoteGraph", "RemoteException")73 74_CONF_DROPLIST = frozenset(75    (76        CONFIG_KEY_CHECKPOINT_MAP,77        CONFIG_KEY_CHECKPOINT_ID,78        CONFIG_KEY_CHECKPOINT_NS,79        CONFIG_KEY_TASK_ID,80    ),81)82 83 84def _sanitize_config_value(v: Any) -> Any:85    """Recursively sanitize a config value to ensure it contains only primitives."""86    if isinstance(v, (str, int, float, bool, UUID)):87        return v88    elif isinstance(v, dict):89        sanitized_dict = {}90        for k, val in v.items():91            if isinstance(k, str):92                sanitized_value = _sanitize_config_value(val)93                if sanitized_value is not None:94                    sanitized_dict[k] = sanitized_value95        return sanitized_dict96    elif isinstance(v, (list, tuple)):97        sanitized_list = []98        for item in v:99            sanitized_item = _sanitize_config_value(item)100            if sanitized_item is not None:101                sanitized_list.append(sanitized_item)102        return sanitized_list103    return None104 105 106class RemoteException(Exception):107    """Exception raised when an error occurs in the remote graph."""108 109    pass110 111 112class RemoteGraph(PregelProtocol):113    """The `RemoteGraph` class is a client implementation for calling remote114    APIs that implement the LangGraph Server API specification.115 116    For example, the `RemoteGraph` class can be used to call APIs from deployments117    on LangSmith Deployment.118 119    `RemoteGraph` behaves the same way as a `Graph` and can be used directly as120    a node in another `Graph`.121    """122 123    assistant_id: str124    name: str | None125 126    def __init__(127        self,128        assistant_id: str,  # graph_id129        /,130        *,131        url: str | None = None,132        api_key: str | None = None,133        headers: dict[str, str] | None = None,134        client: LangGraphClient | None = None,135        sync_client: SyncLangGraphClient | None = None,136        config: RunnableConfig | None = None,137        name: str | None = None,138        distributed_tracing: bool = False,139    ):140        """Specify `url`, `api_key`, and/or `headers` to create default sync and async clients.141 142        If `client` or `sync_client` are provided, they will be used instead of the default clients.143        See `LangGraphClient` and `SyncLangGraphClient` for details on the default clients. At least144        one of `url`, `client`, or `sync_client` must be provided.145 146        Args:147            assistant_id: The assistant ID or graph name of the remote graph to use.148            url: The URL of the remote API.149            api_key: The API key to use for authentication. If not provided, it will be read from the environment (`LANGGRAPH_API_KEY`, `LANGSMITH_API_KEY`, or `LANGCHAIN_API_KEY`).150            headers: Additional headers to include in the requests.151            client: A `LangGraphClient` instance to use instead of creating a default client.152            sync_client: A `SyncLangGraphClient` instance to use instead of creating a default client.153            config: An optional `RunnableConfig` instance with additional configuration.154            name: Human-readable name to attach to the RemoteGraph instance.155                This is useful for adding `RemoteGraph` as a subgraph via `graph.add_node(remote_graph)`.156                If not provided, defaults to the assistant ID.157            distributed_tracing: Whether to enable sending LangSmith distributed tracing headers.158        """159        self.assistant_id = assistant_id160        if name is None:161            self.name = assistant_id162        else:163            self.name = name164        self.config = config165        self.distributed_tracing = distributed_tracing166 167        if client is None and url is not None:168            client = get_client(url=url, api_key=api_key, headers=headers)169        self.client = client170 171        if sync_client is None and url is not None:172            sync_client = get_sync_client(url=url, api_key=api_key, headers=headers)173        self.sync_client = sync_client174 175    def _validate_client(self) -> LangGraphClient:176        if self.client is None:177            raise ValueError(178                "Async client is not initialized: please provide `url` or `client` when initializing `RemoteGraph`."179            )180        return self.client181 182    def _validate_sync_client(self) -> SyncLangGraphClient:183        if self.sync_client is None:184            raise ValueError(185                "Sync client is not initialized: please provide `url` or `sync_client` when initializing `RemoteGraph`."186            )187        return self.sync_client188 189    def copy(self, update: dict[str, Any]) -> Self:190        attrs = {**self.__dict__, **update}191        return self.__class__(attrs.pop("assistant_id"), **attrs)192 193    def with_config(self, config: RunnableConfig | None = None, **kwargs: Any) -> Self:194        return self.copy(195            {"config": merge_configs(self.config, config, cast(RunnableConfig, kwargs))}196        )197 198    def _get_drawable_nodes(199        self, graph: dict[str, list[dict[str, Any]]]200    ) -> dict[str, DrawableNode]:201        nodes = {}202        for node in graph["nodes"]:203            node_id = str(node["id"])204            node_data = node.get("data", {})205 206            # Get node name from node_data if available. If not, use node_id.207            node_name = node.get("name")208            if node_name is None:209                if isinstance(node_data, dict):210                    node_name = node_data.get("name", node_id)211                else:212                    node_name = node_id213 214            nodes[node_id] = DrawableNode(215                id=node_id,216                name=node_name,217                data=node_data,218                metadata=node.get("metadata"),219            )220        return nodes221 222    def get_graph(223        self,224        config: RunnableConfig | None = None,225        *,226        xray: int | bool = False,227        headers: dict[str, str] | None = None,228        params: QueryParamTypes | None = None,229    ) -> DrawableGraph:230        """Get graph by graph name.231 232        This method calls `GET /assistants/{assistant_id}/graph`.233 234        Args:235            config: This parameter is not used.236            xray: Include graph representation of subgraphs. If an integer237                value is provided, only subgraphs with a depth less than or238                equal to the value will be included.239 240        Returns:241            The graph information for the assistant in JSON format.242        """243        sync_client = self._validate_sync_client()244        graph = sync_client.assistants.get_graph(245            assistant_id=self.assistant_id,246            xray=xray,247            headers=headers,248            params=params,249        )250        return DrawableGraph(251            nodes=self._get_drawable_nodes(graph),252            edges=[DrawableEdge(**edge) for edge in graph["edges"]],253        )254 255    async def aget_graph(256        self,257        config: RunnableConfig | None = None,258        *,259        xray: int | bool = False,260        headers: dict[str, str] | None = None,261        params: QueryParamTypes | None = None,262    ) -> DrawableGraph:263        """Get graph by graph name.264 265        This method calls `GET /assistants/{assistant_id}/graph`.266 267        Args:268            config: This parameter is not used.269            xray: Include graph representation of subgraphs. If an integer270                value is provided, only subgraphs with a depth less than or271                equal to the value will be included.272 273        Returns:274            The graph information for the assistant in JSON format.275        """276        client = self._validate_client()277        graph = await client.assistants.get_graph(278            assistant_id=self.assistant_id,279            xray=xray,280            headers=headers,281            params=params,282        )283        return DrawableGraph(284            nodes=self._get_drawable_nodes(graph),285            edges=[DrawableEdge(**edge) for edge in graph["edges"]],286        )287 288    def _create_state_snapshot(self, state: ThreadState) -> StateSnapshot:289        tasks: list[PregelTask] = []290        for task in state["tasks"]:291            interrupts = tuple(292                Interrupt(**interrupt) for interrupt in task["interrupts"]293            )294 295            tasks.append(296                PregelTask(297                    id=task["id"],298                    name=task["name"],299                    path=tuple(),300                    error=Exception(task["error"]) if task["error"] else None,301                    interrupts=interrupts,302                    state=(303                        self._create_state_snapshot(task["state"])304                        if task["state"]305                        else (306                            cast(RunnableConfig, {"configurable": task["checkpoint"]})307                            if task["checkpoint"]308                            else None309                        )310                    ),311                    result=task.get("result"),312                )313            )314 315        return StateSnapshot(316            values=state["values"],317            next=tuple(state["next"]) if state["next"] else tuple(),318            config={319                "configurable": {320                    "thread_id": state["checkpoint"]["thread_id"],321                    "checkpoint_ns": state["checkpoint"]["checkpoint_ns"],322                    "checkpoint_id": state["checkpoint"]["checkpoint_id"],323                    "checkpoint_map": state["checkpoint"].get("checkpoint_map", {}),324                }325            },326            metadata=CheckpointMetadata(**state["metadata"]),327            created_at=state["created_at"],328            parent_config=(329                {330                    "configurable": {331                        "thread_id": state["parent_checkpoint"]["thread_id"],332                        "checkpoint_ns": state["parent_checkpoint"]["checkpoint_ns"],333                        "checkpoint_id": state["parent_checkpoint"]["checkpoint_id"],334                        "checkpoint_map": state["parent_checkpoint"].get(335                            "checkpoint_map", {}336                        ),337                    }338                }339                if state["parent_checkpoint"]340                else None341            ),342            tasks=tuple(tasks),343            interrupts=tuple([i for task in tasks for i in task.interrupts]),344        )345 346    def _get_checkpoint(self, config: RunnableConfig | None) -> Checkpoint | None:347        if config is None:348            return None349 350        checkpoint = {}351 352        if "thread_id" in config["configurable"]:353            checkpoint["thread_id"] = config["configurable"]["thread_id"]354        if "checkpoint_ns" in config["configurable"]:355            checkpoint["checkpoint_ns"] = config["configurable"]["checkpoint_ns"]356        if "checkpoint_id" in config["configurable"]:357            checkpoint["checkpoint_id"] = config["configurable"]["checkpoint_id"]358        if "checkpoint_map" in config["configurable"]:359            checkpoint["checkpoint_map"] = config["configurable"]["checkpoint_map"]360 361        return checkpoint if checkpoint else None362 363    def _get_config(self, checkpoint: Checkpoint) -> RunnableConfig:364        return {365            "configurable": {366                "thread_id": checkpoint["thread_id"],367                "checkpoint_ns": checkpoint["checkpoint_ns"],368                "checkpoint_id": checkpoint["checkpoint_id"],369                "checkpoint_map": checkpoint.get("checkpoint_map", {}),370            }371        }372 373    def _sanitize_config(self, config: RunnableConfig) -> RunnableConfig:374        """Sanitize the config to remove non-serializable fields."""375        sanitized: RunnableConfig = {}376        if "recursion_limit" in config:377            sanitized["recursion_limit"] = config["recursion_limit"]378        if "tags" in config:379            sanitized["tags"] = [tag for tag in config["tags"] if isinstance(tag, str)]380 381        if "metadata" in config:382            sanitized["metadata"] = {}383            for k, v in config["metadata"].items():384                if (385                    isinstance(k, str)386                    and (sanitized_value := _sanitize_config_value(v)) is not None387                ):388                    sanitized["metadata"][k] = sanitized_value389 390        if "configurable" in config:391            sanitized["configurable"] = {}392            for k, v in config["configurable"].items():393                if (394                    isinstance(k, str)395                    and k not in _CONF_DROPLIST396                    and (sanitized_value := _sanitize_config_value(v)) is not None397                ):398                    sanitized["configurable"][k] = sanitized_value399 400        return sanitized401 402    def get_state(403        self,404        config: RunnableConfig,405        *,406        subgraphs: bool = False,407        headers: dict[str, str] | None = None,408        params: QueryParamTypes | None = None,409    ) -> StateSnapshot:410        """Get the state of a thread.411 412        This method calls `POST /threads/{thread_id}/state/checkpoint` if a413        checkpoint is specified in the config or `GET /threads/{thread_id}/state`414        if no checkpoint is specified.415 416        Args:417            config: A `RunnableConfig` that includes `thread_id` in the418                `configurable` field.419            subgraphs: Include subgraphs in the state.420            headers: Optional custom headers to include with the request.421            params: Optional query parameters to include with the request.422 423        Returns:424            The latest state of the thread.425        """426        sync_client = self._validate_sync_client()427        merged_config = merge_configs(self.config, config)428 429        state = sync_client.threads.get_state(430            thread_id=merged_config["configurable"]["thread_id"],431            checkpoint=self._get_checkpoint(merged_config),432            subgraphs=subgraphs,433            headers=headers,434            params=params,435        )436        return self._create_state_snapshot(state)437 438    async def aget_state(439        self,440        config: RunnableConfig,441        *,442        subgraphs: bool = False,443        headers: dict[str, str] | None = None,444        params: QueryParamTypes | None = None,445    ) -> StateSnapshot:446        """Get the state of a thread.447 448        This method calls `POST /threads/{thread_id}/state/checkpoint` if a449        checkpoint is specified in the config or `GET /threads/{thread_id}/state`450        if no checkpoint is specified.451 452        Args:453            config: A `RunnableConfig` that includes `thread_id` in the454                `configurable` field.455            subgraphs: Include subgraphs in the state.456            headers: Optional custom headers to include with the request.457            params: Optional query parameters to include with the request.458 459        Returns:460            The latest state of the thread.461        """462        client = self._validate_client()463        merged_config = merge_configs(self.config, config)464 465        state = await client.threads.get_state(466            thread_id=merged_config["configurable"]["thread_id"],467            checkpoint=self._get_checkpoint(merged_config),468            subgraphs=subgraphs,469            headers=headers,470            params=params,471        )472        return self._create_state_snapshot(state)473 474    def get_state_history(475        self,476        config: RunnableConfig,477        *,478        filter: dict[str, Any] | None = None,479        before: RunnableConfig | None = None,480        limit: int | None = None,481        headers: dict[str, str] | None = None,482        params: QueryParamTypes | None = None,483    ) -> Iterator[StateSnapshot]:484        """Get the state history of a thread.485 486        This method calls `POST /threads/{thread_id}/history`.487 488        Args:489            config: A `RunnableConfig` that includes `thread_id` in the490                `configurable` field.491            filter: Metadata to filter on.492            before: A `RunnableConfig` that includes checkpoint metadata.493            limit: Max number of states to return.494 495        Returns:496            States of the thread.497        """498        sync_client = self._validate_sync_client()499        merged_config = merge_configs(self.config, config)500 501        states = sync_client.threads.get_history(502            thread_id=merged_config["configurable"]["thread_id"],503            limit=limit if limit else 10,504            before=self._get_checkpoint(before),505            metadata=filter,506            checkpoint=self._get_checkpoint(merged_config),507            headers=headers,508            params=params,509        )510        for state in states:511            yield self._create_state_snapshot(state)512 513    async def aget_state_history(514        self,515        config: RunnableConfig,516        *,517        filter: dict[str, Any] | None = None,518        before: RunnableConfig | None = None,519        limit: int | None = None,520        headers: dict[str, str] | None = None,521        params: QueryParamTypes | None = None,522    ) -> AsyncIterator[StateSnapshot]:523        """Get the state history of a thread.524 525        This method calls `POST /threads/{thread_id}/history`.526 527        Args:528            config: A `RunnableConfig` that includes `thread_id` in the529                `configurable` field.530            filter: Metadata to filter on.531            before: A `RunnableConfig` that includes checkpoint metadata.532            limit: Max number of states to return.533            headers: Optional custom headers to include with the request.534            params: Optional query parameters to include with the request.535 536        Returns:537            States of the thread.538        """539        client = self._validate_client()540        merged_config = merge_configs(self.config, config)541 542        states = await client.threads.get_history(543            thread_id=merged_config["configurable"]["thread_id"],544            limit=limit if limit else 10,545            before=self._get_checkpoint(before),546            metadata=filter,547            checkpoint=self._get_checkpoint(merged_config),548            headers=headers,549            params=params,550        )551        for state in states:552            yield self._create_state_snapshot(state)553 554    def bulk_update_state(555        self,556        config: RunnableConfig,557        updates: list[tuple[dict[str, Any] | None, str | None]],558    ) -> RunnableConfig:559        raise NotImplementedError560 561    async def abulk_update_state(562        self,563        config: RunnableConfig,564        updates: list[tuple[dict[str, Any] | None, str | None]],565    ) -> RunnableConfig:566        raise NotImplementedError567 568    def update_state(569        self,570        config: RunnableConfig,571        values: dict[str, Any] | Any | None,572        as_node: str | None = None,573        *,574        headers: dict[str, str] | None = None,575        params: QueryParamTypes | None = None,576    ) -> RunnableConfig:577        """Update the state of a thread.578 579        This method calls `POST /threads/{thread_id}/state`.580 581        Args:582            config: A `RunnableConfig` that includes `thread_id` in the583                `configurable` field.584            values: Values to update to the state.585            as_node: Update the state as if this node had just executed.586 587        Returns:588            `RunnableConfig` for the updated thread.589        """590        sync_client = self._validate_sync_client()591        merged_config = merge_configs(self.config, config)592 593        response: dict = sync_client.threads.update_state(  # type: ignore594            thread_id=merged_config["configurable"]["thread_id"],595            values=values,596            as_node=as_node,597            checkpoint=self._get_checkpoint(merged_config),598            headers=headers,599            params=params,600        )601        return self._get_config(response["checkpoint"])602 603    async def aupdate_state(604        self,605        config: RunnableConfig,606        values: dict[str, Any] | Any | None,607        as_node: str | None = None,608        *,609        headers: dict[str, str] | None = None,610        params: QueryParamTypes | None = None,611    ) -> RunnableConfig:612        """Update the state of a thread.613 614        This method calls `POST /threads/{thread_id}/state`.615 616        Args:617            config: A `RunnableConfig` that includes `thread_id` in the618                `configurable` field.619            values: Values to update to the state.620            as_node: Update the state as if this node had just executed.621 622        Returns:623            `RunnableConfig` for the updated thread.624        """625        client = self._validate_client()626        merged_config = merge_configs(self.config, config)627 628        response: dict = await client.threads.update_state(  # type: ignore629            thread_id=merged_config["configurable"]["thread_id"],630            values=values,631            as_node=as_node,632            checkpoint=self._get_checkpoint(merged_config),633            headers=headers,634            params=params,635        )636        return self._get_config(response["checkpoint"])637 638    def _get_stream_modes(639        self,640        stream_mode: StreamMode | list[StreamMode] | None,641        config: RunnableConfig | None,642        default: StreamMode = "updates",643    ) -> tuple[list[StreamModeSDK], list[StreamModeSDK], bool, StreamProtocol | None]:644        """Return a tuple of the final list of stream modes sent to the645        remote graph and a boolean flag indicating if stream mode 'updates'646        was present in the original list of stream modes.647 648        'updates' mode is added to the list of stream modes so that interrupts649        can be detected in the remote graph.650        """651        updated_stream_modes: list[StreamModeSDK] = []652        req_single = True653        # coerce to list, or add default stream mode654        if stream_mode:655            if isinstance(stream_mode, str):656                updated_stream_modes.append(stream_mode)657            else:658                req_single = False659                updated_stream_modes.extend(stream_mode)660        else:661            updated_stream_modes.append(default)662        requested_stream_modes = updated_stream_modes.copy()663        # add any from parent graph664        stream: StreamProtocol | None = (665            (config or {}).get(CONF, {}).get(CONFIG_KEY_STREAM)666        )667        if stream:668            updated_stream_modes.extend(stream.modes)669        # map "messages" to "messages-tuple"670        if "messages" in updated_stream_modes:671            updated_stream_modes.remove("messages")672            updated_stream_modes.append("messages-tuple")673 674        # if requested "messages-tuple",675        # map to "messages" in requested_stream_modes676        if "messages-tuple" in requested_stream_modes:677            requested_stream_modes.remove("messages-tuple")678            requested_stream_modes.append("messages")679 680        # add 'updates' mode if not present681        if "updates" not in updated_stream_modes:682            updated_stream_modes.append("updates")683 684        # remove 'events', as it's not supported in Pregel685        if "events" in updated_stream_modes:686            updated_stream_modes.remove("events")687        return (updated_stream_modes, requested_stream_modes, req_single, stream)688 689    @overload690    def stream(691        self,692        input: dict[str, Any] | Any,693        config: RunnableConfig | None = None,694        *,695        context: Context | None = None,696        stream_mode: StreamMode | list[StreamMode] | None = None,697        interrupt_before: All | Sequence[str] | None = None,698        interrupt_after: All | Sequence[str] | None = None,699        subgraphs: bool = False,700        headers: dict[str, str] | None = None,701        params: QueryParamTypes | None = None,702        version: Literal["v2"],703        **kwargs: Any,704    ) -> Iterator[StreamPart]: ...705 706    @overload707    def stream(708        self,709        input: dict[str, Any] | Any,710        config: RunnableConfig | None = None,711        *,712        context: Context | None = None,713        stream_mode: StreamMode | list[StreamMode] | None = None,714        interrupt_before: All | Sequence[str] | None = None,715        interrupt_after: All | Sequence[str] | None = None,716        subgraphs: bool = False,717        headers: dict[str, str] | None = None,718        params: QueryParamTypes | None = None,719        version: Literal["v1"] = ...,720        **kwargs: Any,721    ) -> Iterator[dict[str, Any] | Any]: ...722 723    def stream(724        self,725        input: dict[str, Any] | Any,726        config: RunnableConfig | None = None,727        *,728        context: Context | None = None,729        stream_mode: StreamMode | list[StreamMode] | None = None,730        interrupt_before: All | Sequence[str] | None = None,731        interrupt_after: All | Sequence[str] | None = None,732        subgraphs: bool = False,733        headers: dict[str, str] | None = None,734        params: QueryParamTypes | None = None,735        version: Literal["v1", "v2"] = "v1",736        **kwargs: Any,737    ) -> Iterator[dict[str, Any] | Any]:738        """Create a run and stream the results.739 740        This method calls `POST /threads/{thread_id}/runs/stream` if a `thread_id`741        is specified in the `configurable` field of the config or742        `POST /runs/stream` otherwise.743 744        Args:745            input: Input to the graph.746            config: A `RunnableConfig` for graph invocation.747            stream_mode: Stream mode(s) to use.748            interrupt_before: Interrupt the graph before these nodes.749            interrupt_after: Interrupt the graph after these nodes.750            subgraphs: Stream from subgraphs.751            headers: Additional headers to pass to the request.752            **kwargs: Additional params to pass to client.runs.stream.753 754        Yields:755            The output of the graph.756        """757        sync_client = self._validate_sync_client()758        merged_config = merge_configs(self.config, config)759        sanitized_config = self._sanitize_config(merged_config)760        stream_modes, requested, req_single, stream = self._get_stream_modes(761            stream_mode, config762        )763        if isinstance(input, Command):764            command: CommandSDK | None = cast(CommandSDK, asdict(input))765            input = None766        else:767            command = None768        thread_id = sanitized_config.get("configurable", {}).pop("thread_id", None)769 770        for chunk in sync_client.runs.stream(771            thread_id=thread_id,772            assistant_id=self.assistant_id,773            input=input,774            command=command,775            config=sanitized_config,776            context=context,777            stream_mode=stream_modes,778            interrupt_before=interrupt_before,779            interrupt_after=interrupt_after,780            stream_subgraphs=subgraphs or stream is not None,781            if_not_exists="create",782            headers=(783                _merge_tracing_headers(headers) if self.distributed_tracing else headers784            ),785            params=params,786            **kwargs,787        ):788            # split mode and ns789            if NS_SEP in chunk.event:790                mode, ns_ = chunk.event.split(NS_SEP, 1)791                ns = tuple(ns_.split(NS_SEP))792            else:793                mode, ns = chunk.event, ()794            # raise ParentCommand exception for command events795            if mode == "command" and chunk.data.get("graph") == Command.PARENT:796                raise ParentCommand(Command(**chunk.data))797            # prepend caller ns (as it is not passed to remote graph)798            if caller_ns := (config or {}).get(CONF, {}).get(CONFIG_KEY_CHECKPOINT_NS):799                caller_ns = tuple(caller_ns.split(NS_SEP))800                ns = caller_ns + ns801            # stream to parent stream802            if stream is not None and mode in stream.modes:803                stream((ns, mode, chunk.data))804            # raise interrupt or errors805            if chunk.event.startswith("updates"):806                if isinstance(chunk.data, dict) and INTERRUPT in chunk.data:807                    if caller_ns:808                        raise GraphInterrupt(809                            [Interrupt(**i) for i in chunk.data[INTERRUPT]]810                        )811            elif chunk.event.startswith("error"):812                raise RemoteException(chunk.data)813            # filter for what was actually requested814            if mode not in requested:815                continue816 817            if chunk.event.startswith("messages"):818                chunk = chunk._replace(data=tuple(chunk.data))819 820            # emit chunk821            if version == "v2":822                ints: tuple[Interrupt, ...] = ()823                if mode == "values" and isinstance(chunk.data, dict):824                    ints = tuple(825                        Interrupt(**i) if isinstance(i, dict) else i826                        for i in chunk.data.pop(INTERRUPT, ())827                    )828                yield {"type": mode, "ns": ns, "data": chunk.data, "interrupts": ints}829            elif subgraphs:830                if NS_SEP in chunk.event:831                    mode, ns_ = chunk.event.split(NS_SEP, 1)832                    ns = tuple(ns_.split(NS_SEP))833                else:834                    mode, ns = chunk.event, ()835                if req_single:836                    yield ns, chunk.data837                else:838                    yield ns, mode, chunk.data839            elif req_single:840                yield chunk.data841            else:842                yield chunk843 844    @overload845    def astream(846        self,847        input: dict[str, Any] | Any,848        config: RunnableConfig | None = None,849        *,850        context: Context | None = None,851        stream_mode: StreamMode | list[StreamMode] | None = None,852        interrupt_before: All | Sequence[str] | None = None,853        interrupt_after: All | Sequence[str] | None = None,854        subgraphs: bool = False,855        headers: dict[str, str] | None = None,856        params: QueryParamTypes | None = None,857        version: Literal["v2"],858        **kwargs: Any,859    ) -> AsyncIterator[StreamPart]: ...860 861    @overload862    def astream(863        self,864        input: dict[str, Any] | Any,865        config: RunnableConfig | None = None,866        *,867        context: Context | None = None,868        stream_mode: StreamMode | list[StreamMode] | None = None,869        interrupt_before: All | Sequence[str] | None = None,870        interrupt_after: All | Sequence[str] | None = None,871        subgraphs: bool = False,872        headers: dict[str, str] | None = None,873        params: QueryParamTypes | None = None,874        version: Literal["v1"] = ...,875        **kwargs: Any,876    ) -> AsyncIterator[dict[str, Any] | Any]: ...877 878    async def astream(879        self,880        input: dict[str, Any] | Any,881        config: RunnableConfig | None = None,882        *,883        context: Context | None = None,884        stream_mode: StreamMode | list[StreamMode] | None = None,885        interrupt_before: All | Sequence[str] | None = None,886        interrupt_after: All | Sequence[str] | None = None,887        subgraphs: bool = False,888        headers: dict[str, str] | None = None,889        params: QueryParamTypes | None = None,890        version: Literal["v1", "v2"] = "v1",891        **kwargs: Any,892    ) -> AsyncIterator[dict[str, Any] | Any]:893        """Create a run and stream the results.894 895        This method calls `POST /threads/{thread_id}/runs/stream` if a `thread_id`896        is specified in the `configurable` field of the config or897        `POST /runs/stream` otherwise.898 899        Args:900            input: Input to the graph.901            config: A `RunnableConfig` for graph invocation.902            stream_mode: Stream mode(s) to use.903            interrupt_before: Interrupt the graph before these nodes.904            interrupt_after: Interrupt the graph after these nodes.905            subgraphs: Stream from subgraphs.906            headers: Additional headers to pass to the request.907            **kwargs: Additional params to pass to client.runs.stream.908 909        Yields:910            The output of the graph.911        """912        client = self._validate_client()913        merged_config = merge_configs(self.config, config)914        sanitized_config = self._sanitize_config(merged_config)915        stream_modes, requested, req_single, stream = self._get_stream_modes(916            stream_mode, config917        )918        if isinstance(input, Command):919            command: CommandSDK | None = cast(CommandSDK, asdict(input))920            input = None921        else:922            command = None923        thread_id = sanitized_config.get("configurable", {}).pop("thread_id", None)924 925        async for chunk in client.runs.stream(926            thread_id=thread_id,927            assistant_id=self.assistant_id,928            input=input,929            command=command,930            config=sanitized_config,931            context=context,932            stream_mode=stream_modes,933            interrupt_before=interrupt_before,934            interrupt_after=interrupt_after,935            stream_subgraphs=subgraphs or stream is not None,936            if_not_exists="create",937            headers=(938                _merge_tracing_headers(headers) if self.distributed_tracing else headers939            ),940            params=params,941            **kwargs,942        ):943            # split mode and ns944            if NS_SEP in chunk.event:945                mode, ns_ = chunk.event.split(NS_SEP, 1)946                ns = tuple(ns_.split(NS_SEP))947            else:948                mode, ns = chunk.event, ()949            # raise ParentCommand exception for command events950            if mode == "command" and chunk.data.get("graph") == Command.PARENT:951                raise ParentCommand(Command(**chunk.data))952            # prepend caller ns (as it is not passed to remote graph)953            if caller_ns := (config or {}).get(CONF, {}).get(CONFIG_KEY_CHECKPOINT_NS):954                caller_ns = tuple(caller_ns.split(NS_SEP))955                ns = caller_ns + ns956            # stream to parent stream957            if stream is not None and mode in stream.modes:958                stream((ns, mode, chunk.data))959            # raise interrupt or errors960            if chunk.event.startswith("updates"):961                if isinstance(chunk.data, dict) and INTERRUPT in chunk.data:962                    if caller_ns:963                        raise GraphInterrupt(964                            [Interrupt(**i) for i in chunk.data[INTERRUPT]]965                        )966            elif chunk.event.startswith("error"):967                raise RemoteException(chunk.data)968            # filter for what was actually requested969            if mode not in requested:970                continue971 972            if chunk.event.startswith("messages"):973                chunk = chunk._replace(data=tuple(chunk.data))974 975            # emit chunk976            if version == "v2":977                ints: tuple[Interrupt, ...] = ()978                if mode == "values" and isinstance(chunk.data, dict):979                    ints = tuple(980                        Interrupt(**i) if isinstance(i, dict) else i981                        for i in chunk.data.pop(INTERRUPT, ())982                    )983                yield {"type": mode, "ns": ns, "data": chunk.data, "interrupts": ints}984            elif subgraphs:985                if NS_SEP in chunk.event:986                    mode, ns_ = chunk.event.split(NS_SEP, 1)987                    ns = tuple(ns_.split(NS_SEP))988                else:989                    mode, ns = chunk.event, ()990                if req_single:991                    yield ns, chunk.data992                else:993                    yield ns, mode, chunk.data994            elif req_single:995                yield chunk.data996            else:997                yield chunk998 999    async def astream_events(1000        self,1001        input: Any,1002        config: RunnableConfig | None = None,1003        *,1004        version: Literal["v1", "v2"],1005        include_names: Sequence[All] | None = None,1006        include_types: Sequence[All] | None = None,1007        include_tags: Sequence[All] | None = None,1008        exclude_names: Sequence[All] | None = None,1009        exclude_types: Sequence[All] | None = None,1010        exclude_tags: Sequence[All] | None = None,1011        **kwargs: Any,1012    ) -> AsyncIterator[dict[str, Any]]:1013        raise NotImplementedError1014 1015    @overload1016    def invoke(1017        self,1018        input: dict[str, Any] | Any,1019        config: RunnableConfig | None = None,1020        *,1021        context: Context | None = None,1022        interrupt_before: All | Sequence[str] | None = None,1023        interrupt_after: All | Sequence[str] | None = None,1024        headers: dict[str, str] | None = None,1025        params: QueryParamTypes | None = None,1026        version: Literal["v2"],1027        **kwargs: Any,1028    ) -> GraphOutput[dict[str, Any]]: ...1029 1030    @overload1031    def invoke(1032        self,1033        input: dict[str, Any] | Any,1034        config: RunnableConfig | None = None,1035        *,1036        context: Context | None = None,1037        interrupt_before: All | Sequence[str] | None = None,1038        interrupt_after: All | Sequence[str] | None = None,1039        headers: dict[str, str] | None = None,1040        params: QueryParamTypes | None = None,1041        version: Literal["v1"] = ...,1042        **kwargs: Any,1043    ) -> dict[str, Any] | Any: ...1044 1045    def invoke(1046        self,1047        input: dict[str, Any] | Any,1048        config: RunnableConfig | None = None,1049        *,1050        context: Context | None = None,1051        interrupt_before: All | Sequence[str] | None = None,1052        interrupt_after: All | Sequence[str] | None = None,1053        headers: dict[str, str] | None = None,1054        params: QueryParamTypes | None = None,1055        version: Literal["v1", "v2"] = "v1",1056        **kwargs: Any,1057    ) -> dict[str, Any] | Any:1058        """Create a run, wait until it finishes and return the final state.1059 1060        Args:1061            input: Input to the graph.1062            config: A `RunnableConfig` for graph invocation.1063            interrupt_before: Interrupt the graph before these nodes.1064            interrupt_after: Interrupt the graph after these nodes.1065            headers: Additional headers to pass to the request.1066            version: The streaming format version. `"v1"` (default) returns the1067                traditional format, `"v2"` returns `StreamPart` typed dicts.1068            **kwargs: Additional params to pass to RemoteGraph.stream.1069 1070        Returns:1071            The output of the graph.1072        """1073        for chunk in self.stream(  # type: ignore[misc, call-overload]1074            input,1075            config=config,1076            context=context,1077            interrupt_before=interrupt_before,1078            interrupt_after=interrupt_after,1079            headers=headers,1080            stream_mode="values",1081            params=params,1082            version=version,1083            **kwargs,1084        ):1085            pass1086        try:1087            if version == "v2":1088                return GraphOutput(1089                    value=chunk["data"],1090                    interrupts=tuple(chunk.get("interrupts", ())),1091                )1092            return chunk1093        except UnboundLocalError:1094            logger.warning("No events received from remote graph")1095            return None1096 1097    @overload1098    async def ainvoke(1099        self,1100        input: dict[str, Any] | Any,1101        config: RunnableConfig | None = None,1102        *,1103        context: Context | None = None,1104        interrupt_before: All | Sequence[str] | None = None,1105        interrupt_after: All | Sequence[str] | None = None,1106        headers: dict[str, str] | None = None,1107        params: QueryParamTypes | None = None,1108        version: Literal["v2"],1109        **kwargs: Any,1110    ) -> GraphOutput[dict[str, Any]]: ...1111 1112    @overload1113    async def ainvoke(1114        self,1115        input: dict[str, Any] | Any,1116        config: RunnableConfig | None = None,1117        *,1118        context: Context | None = None,1119        interrupt_before: All | Sequence[str] | None = None,1120        interrupt_after: All | Sequence[str] | None = None,1121        headers: dict[str, str] | None = None,1122        params: QueryParamTypes | None = None,1123        version: Literal["v1"] = ...,1124        **kwargs: Any,1125    ) -> dict[str, Any] | Any: ...1126 1127    async def ainvoke(1128        self,1129        input: dict[str, Any] | Any,1130        config: RunnableConfig | None = None,1131        *,1132        context: Context | None = None,1133        interrupt_before: All | Sequence[str] | None = None,1134        interrupt_after: All | Sequence[str] | None = None,1135        headers: dict[str, str] | None = None,1136        params: QueryParamTypes | None = None,1137        version: Literal["v1", "v2"] = "v1",1138        **kwargs: Any,1139    ) -> dict[str, Any] | Any:1140        """Create a run, wait until it finishes and return the final state.1141 1142        Args:1143            input: Input to the graph.1144            config: A `RunnableConfig` for graph invocation.1145            interrupt_before: Interrupt the graph before these nodes.1146            interrupt_after: Interrupt the graph after these nodes.1147            headers: Additional headers to pass to the request.1148            version: The streaming format version. `"v1"` (default) returns the1149                traditional format, `"v2"` returns `StreamPart` typed dicts.1150            **kwargs: Additional params to pass to RemoteGraph.astream.1151 1152        Returns:1153            The output of the graph.1154        """1155        async for chunk in self.astream(  # type: ignore[misc, call-overload]1156            input,1157            config=config,1158            context=context,1159            interrupt_before=interrupt_before,1160            interrupt_after=interrupt_after,1161            headers=headers,1162            stream_mode="values",1163            params=params,1164            version=version,1165            **kwargs,1166        ):1167            pass1168        try:1169            if version == "v2":1170                return GraphOutput(1171                    value=chunk["data"],1172                    interrupts=tuple(chunk.get("interrupts", ())),1173                )1174            return chunk1175        except UnboundLocalError:1176            logger.warning("No events received from remote graph")1177            return None1178 1179 1180def _merge_tracing_headers(headers: dict[str, str] | None) -> dict[str, str] | None:1181    if rt := ls.get_current_run_tree():1182        tracing_headers = rt.to_headers()1183        if headers:1184            if "baggage" in headers:1185                tracing_headers["baggage"] = (1186                    f"{headers['baggage']},{tracing_headers['baggage']}"1187                )1188            headers.update(tracing_headers)1189        else:1190            headers = tracing_headers1191    return headers1192 
codekingpro/portable-devtools · Team Ai