Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_read.py299 linesDownload Raw Back to pregel
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 
codekingpro/portable-devtools · Team Ai