codekingpro/portable-devtools
114k
1from __future__ import annotations2 3from collections.abc import AsyncIterator, Callable, Iterator, Mapping, Sequence4from datetime import timedelta5from functools import cached_property6from typing import (7 Any,8)9 10from langchain_core.runnables import Runnable, RunnableConfig11 12from langgraph._internal._config import merge_configs13from langgraph._internal._constants import CONF, CONFIG_KEY_READ14from langgraph._internal._runnable import RunnableCallable, RunnableSeq15from langgraph._internal._timeout import coerce_timeout_policy16from langgraph.pregel._utils import find_subgraph_pregel17from langgraph.pregel._write import ChannelWrite18from langgraph.pregel.protocol import PregelProtocol19from langgraph.types import CachePolicy, RetryPolicy, TimeoutPolicy20 21READ_TYPE = Callable[[str | Sequence[str], bool], Any | dict[str, Any]]22INPUT_CACHE_KEY_TYPE = tuple[Callable[..., Any], tuple[str, ...]]23 24 25class ChannelRead(RunnableCallable):26 """Implements the logic for reading state from CONFIG_KEY_READ.27 Usable both as a runnable as well as a static method to call imperatively."""28 29 channel: str | list[str]30 31 fresh: bool = False32 33 mapper: Callable[[Any], Any] | None = None34 35 def __init__(36 self,37 channel: str | list[str],38 *,39 fresh: bool = False,40 mapper: Callable[[Any], Any] | None = None,41 tags: list[str] | None = None,42 ) -> None:43 super().__init__(44 func=self._read,45 afunc=self._aread,46 tags=tags,47 name=None,48 trace=False,49 )50 self.fresh = fresh51 self.mapper = mapper52 self.channel = channel53 54 def get_name(self, suffix: str | None = None, *, name: str | None = None) -> str:55 if name:56 pass57 elif isinstance(self.channel, str):58 name = f"ChannelRead<{self.channel}>"59 else:60 name = f"ChannelRead<{','.join(self.channel)}>"61 return super().get_name(suffix, name=name)62 63 def _read(self, _: Any, config: RunnableConfig) -> Any:64 return self.do_read(65 config, select=self.channel, fresh=self.fresh, mapper=self.mapper66 )67 68 async def _aread(self, _: Any, config: RunnableConfig) -> Any:69 return self.do_read(70 config, select=self.channel, fresh=self.fresh, mapper=self.mapper71 )72 73 @staticmethod74 def do_read(75 config: RunnableConfig,76 *,77 select: str | list[str],78 fresh: bool = False,79 mapper: Callable[[Any], Any] | None = None,80 ) -> Any:81 try:82 read: READ_TYPE = config[CONF][CONFIG_KEY_READ]83 except KeyError:84 raise RuntimeError(85 "Not configured with a read function"86 "Make sure to call in the context of a Pregel process"87 )88 if mapper:89 return mapper(read(select, fresh))90 else:91 return read(select, fresh)92 93 94DEFAULT_BOUND = RunnableCallable(lambda input: input)95 96 97class PregelNode:98 """A node in a Pregel graph. This won't be invoked as a runnable by the graph99 itself, but instead acts as a container for the components necessary to make100 a PregelExecutableTask for a node."""101 102 channels: str | list[str]103 """The channels that will be passed as input to `bound`.104 If a str, the node will be invoked with its value if it isn't empty.105 If a list, the node will be invoked with a dict of those channels' values."""106 107 triggers: list[str]108 """If any of these channels is written to, this node will be triggered in109 the next step."""110 111 mapper: Callable[[Any], Any] | None112 """A function to transform the input before passing it to `bound`."""113 114 writers: list[Runnable]115 """A list of writers that will be executed after `bound`, responsible for116 taking the output of `bound` and writing it to the appropriate channels."""117 118 bound: Runnable[Any, Any]119 """The main logic of the node. This will be invoked with the input from 120 `channels`."""121 122 retry_policy: Sequence[RetryPolicy] | None123 """The retry policies to use when invoking the node."""124 125 cache_policy: CachePolicy | None126 """The cache policy to use when invoking the node."""127 128 timeout: TimeoutPolicy | None129 """Timeout policy for a single invocation.130 131 If exceeded, `NodeTimeoutError` is raised and the retry policy (if any)132 decides whether to retry. Supported only for async nodes.133 """134 135 tags: Sequence[str] | None136 """Tags to attach to the node for tracing."""137 138 metadata: Mapping[str, Any] | None139 """Metadata to attach to the node for tracing."""140 141 is_error_handler: bool142 """Whether this node is registered as an error handler node."""143 144 error_handler_node: str | None145 """Optional handler node name for failures from this node."""146 147 subgraphs: Sequence[PregelProtocol]148 """Subgraphs used by the node."""149 150 def __init__(151 self,152 *,153 channels: str | list[str],154 triggers: Sequence[str],155 mapper: Callable[[Any], Any] | None = None,156 writers: list[Runnable] | None = None,157 tags: list[str] | None = None,158 metadata: Mapping[str, Any] | None = None,159 bound: Runnable[Any, Any] | None = None,160 retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,161 cache_policy: CachePolicy | None = None,162 is_error_handler: bool = False,163 error_handler_node: str | None = None,164 subgraphs: Sequence[PregelProtocol] | None = None,165 timeout: float | timedelta | TimeoutPolicy | None = None,166 ) -> None:167 self.channels = channels168 self.triggers = list(triggers)169 self.mapper = mapper170 self.writers = writers or []171 self.bound = bound if bound is not None else DEFAULT_BOUND172 self.cache_policy = cache_policy173 if isinstance(retry_policy, RetryPolicy):174 self.retry_policy = (retry_policy,)175 else:176 self.retry_policy = retry_policy177 self.timeout = coerce_timeout_policy(timeout)178 self.tags = tags179 self.metadata = metadata180 self.is_error_handler = is_error_handler181 self.error_handler_node = error_handler_node182 if subgraphs is not None:183 self.subgraphs = subgraphs184 elif self.bound is not DEFAULT_BOUND:185 try:186 subgraph = find_subgraph_pregel(self.bound)187 except Exception:188 subgraph = None189 if subgraph:190 self.subgraphs = [subgraph]191 else:192 self.subgraphs = []193 else:194 self.subgraphs = []195 196 def copy(self, update: dict[str, Any]) -> PregelNode:197 attrs = {**self.__dict__, **update}198 # Drop the cached properties199 attrs.pop("flat_writers", None)200 attrs.pop("node", None)201 attrs.pop("input_cache_key", None)202 return PregelNode(**attrs)203 204 @cached_property205 def flat_writers(self) -> list[Runnable]:206 """Get writers with optimizations applied. Dedupes consecutive ChannelWrites."""207 writers = self.writers.copy()208 while (209 len(writers) > 1210 and isinstance(writers[-1], ChannelWrite)211 and isinstance(writers[-2], ChannelWrite)212 ):213 # we can combine writes if they are consecutive214 # careful to not modify the original writers list or ChannelWrite215 writers[-2] = ChannelWrite(216 writes=writers[-2].writes + writers[-1].writes,217 )218 writers.pop()219 return writers220 221 @cached_property222 def node(self) -> Runnable[Any, Any] | None:223 """Get a runnable that combines `bound` and `writers`."""224 writers = self.flat_writers225 if self.bound is DEFAULT_BOUND and not writers:226 return None227 elif self.bound is DEFAULT_BOUND and len(writers) == 1:228 return writers[0]229 elif self.bound is DEFAULT_BOUND:230 return RunnableSeq(*writers)231 elif writers:232 return RunnableSeq(self.bound, *writers)233 else:234 return self.bound235 236 @cached_property237 def input_cache_key(self) -> INPUT_CACHE_KEY_TYPE:238 """Get a cache key for the input to the node.239 This is used to avoid calculating the same input multiple times."""240 return (241 self.mapper,242 tuple(self.channels)243 if isinstance(self.channels, list)244 else (self.channels,),245 )246 247 def invoke(248 self,249 input: Any,250 config: RunnableConfig | None = None,251 **kwargs: Any | None,252 ) -> Any:253 self_config: RunnableConfig = {"metadata": self.metadata, "tags": self.tags}254 return self.bound.invoke(255 input,256 merge_configs(self_config, config),257 **kwargs,258 )259 260 async def ainvoke(261 self,262 input: Any,263 config: RunnableConfig | None = None,264 **kwargs: Any | None,265 ) -> Any:266 self_config: RunnableConfig = {"metadata": self.metadata, "tags": self.tags}267 return await self.bound.ainvoke(268 input,269 merge_configs(self_config, config),270 **kwargs,271 )272 273 def stream(274 self,275 input: Any,276 config: RunnableConfig | None = None,277 **kwargs: Any | None,278 ) -> Iterator[Any]:279 self_config: RunnableConfig = {"metadata": self.metadata, "tags": self.tags}280 yield from self.bound.stream(281 input,282 merge_configs(self_config, config),283 **kwargs,284 )285 286 async def astream(287 self,288 input: Any,289 config: RunnableConfig | None = None,290 **kwargs: Any | None,291 ) -> AsyncIterator[Any]:292 self_config: RunnableConfig = {"metadata": self.metadata, "tags": self.tags}293 async for item in self.bound.astream(294 input,295 merge_configs(self_config, config),296 **kwargs,297 ):298 yield item299 