codekingpro/portable-devtools
114k
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 