Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_write.py193 linesDownload Raw Back to pregel
1from __future__ import annotations2 3from collections.abc import Callable, Sequence4from typing import (5    Any,6    NamedTuple,7    TypeVar,8    cast,9)10 11from langchain_core.runnables import Runnable, RunnableConfig12 13from langgraph._internal._constants import CONF, CONFIG_KEY_SEND, TASKS14from langgraph._internal._runnable import RunnableCallable15from langgraph._internal._typing import MISSING16from langgraph.errors import InvalidUpdateError17from langgraph.types import Send18 19TYPE_SEND = Callable[[Sequence[tuple[str, Any]]], None]20R = TypeVar("R", bound=Runnable)21 22SKIP_WRITE = object()23PASSTHROUGH = object()24 25 26class ChannelWriteEntry(NamedTuple):27    channel: str28    """Channel name to write to."""29    value: Any = PASSTHROUGH30    """Value to write, or PASSTHROUGH to use the input."""31    skip_none: bool = False32    """Whether to skip writing if the value is None."""33    mapper: Callable | None = None34    """Function to transform the value before writing."""35 36 37class ChannelWriteTupleEntry(NamedTuple):38    mapper: Callable[[Any], Sequence[tuple[str, Any]] | None]39    """Function to extract tuples from value."""40    value: Any = PASSTHROUGH41    """Value to write, or PASSTHROUGH to use the input."""42    static: Sequence[tuple[str, Any, str | None]] | None = None43    """Optional, declared writes for static analysis."""44 45 46class ChannelWrite(RunnableCallable):47    """Implements the logic for sending writes to CONFIG_KEY_SEND.48    Can be used as a runnable or as a static method to call imperatively."""49 50    writes: list[ChannelWriteEntry | ChannelWriteTupleEntry | Send]51    """Sequence of write entries or Send objects to write."""52 53    def __init__(54        self,55        writes: Sequence[ChannelWriteEntry | ChannelWriteTupleEntry | Send],56        *,57        tags: Sequence[str] | None = None,58    ):59        super().__init__(60            func=self._write,61            afunc=self._awrite,62            name=None,63            tags=tags,64            trace=False,65        )66        self.writes = cast(67            list[ChannelWriteEntry | ChannelWriteTupleEntry | Send], writes68        )69 70    def get_name(self, suffix: str | None = None, *, name: str | None = None) -> str:71        if not name:72            name = f"ChannelWrite<{','.join(w.channel if isinstance(w, ChannelWriteEntry) else '...' if isinstance(w, ChannelWriteTupleEntry) else w.node for w in self.writes)}>"73        return super().get_name(suffix, name=name)74 75    def _write(self, input: Any, config: RunnableConfig) -> None:76        writes = [77            ChannelWriteEntry(write.channel, input, write.skip_none, write.mapper)78            if isinstance(write, ChannelWriteEntry) and write.value is PASSTHROUGH79            else ChannelWriteTupleEntry(write.mapper, input)80            if isinstance(write, ChannelWriteTupleEntry) and write.value is PASSTHROUGH81            else write82            for write in self.writes83        ]84        self.do_write(85            config,86            writes,87        )88        return input89 90    async def _awrite(self, input: Any, config: RunnableConfig) -> None:91        writes = [92            ChannelWriteEntry(write.channel, input, write.skip_none, write.mapper)93            if isinstance(write, ChannelWriteEntry) and write.value is PASSTHROUGH94            else ChannelWriteTupleEntry(write.mapper, input)95            if isinstance(write, ChannelWriteTupleEntry) and write.value is PASSTHROUGH96            else write97            for write in self.writes98        ]99        self.do_write(100            config,101            writes,102        )103        return input104 105    @staticmethod106    def do_write(107        config: RunnableConfig,108        writes: Sequence[ChannelWriteEntry | ChannelWriteTupleEntry | Send],109        allow_passthrough: bool = True,110    ) -> None:111        # validate112        for w in writes:113            if isinstance(w, ChannelWriteEntry):114                if w.channel == TASKS:115                    raise InvalidUpdateError(116                        "Cannot write to the reserved channel TASKS"117                    )118                if w.value is PASSTHROUGH and not allow_passthrough:119                    raise InvalidUpdateError("PASSTHROUGH value must be replaced")120            if isinstance(w, ChannelWriteTupleEntry):121                if w.value is PASSTHROUGH and not allow_passthrough:122                    raise InvalidUpdateError("PASSTHROUGH value must be replaced")123        # if we want to persist writes found before hitting a ParentCommand124        # can move this to a finally block125        write: TYPE_SEND = config[CONF][CONFIG_KEY_SEND]126        write(_assemble_writes(writes))127 128    @staticmethod129    def is_writer(runnable: Runnable) -> bool:130        """Used by PregelNode to distinguish between writers and other runnables."""131        return (132            isinstance(runnable, ChannelWrite)133            or getattr(runnable, "_is_channel_writer", MISSING) is not MISSING134        )135 136    @staticmethod137    def get_static_writes(138        runnable: Runnable,139    ) -> Sequence[tuple[str, Any, str | None]] | None:140        """Used to get conditional writes a writer declares for static analysis."""141        if isinstance(runnable, ChannelWrite):142            return [143                w144                for entry in runnable.writes145                if isinstance(entry, ChannelWriteTupleEntry) and entry.static146                for w in entry.static147            ] or None148        elif writes := getattr(runnable, "_is_channel_writer", MISSING):149            if writes is not MISSING:150                writes = cast(151                    Sequence[tuple[ChannelWriteEntry | Send, str | None]],152                    writes,153                )154                entries = [e for e, _ in writes]155                labels = [la for _, la in writes]156                return [(*t, la) for t, la in zip(_assemble_writes(entries), labels)]157 158    @staticmethod159    def register_writer(160        runnable: R,161        static: Sequence[tuple[ChannelWriteEntry | Send, str | None]] | None = None,162    ) -> R:163        """Used to mark a runnable as a writer, so that it can be detected by is_writer.164        Instances of ChannelWrite are automatically marked as writers.165        Optionally, a list of declared writes can be passed for static analysis."""166        # using object.__setattr__ to work around objects that override __setattr__167        # eg. pydantic models and dataclasses168        object.__setattr__(runnable, "_is_channel_writer", static)169        return runnable170 171 172def _assemble_writes(173    writes: Sequence[ChannelWriteEntry | ChannelWriteTupleEntry | Send],174) -> list[tuple[str, Any]]:175    """Assembles the writes into a list of tuples."""176    tuples: list[tuple[str, Any]] = []177    for w in writes:178        if isinstance(w, Send):179            tuples.append((TASKS, w))180        elif isinstance(w, ChannelWriteTupleEntry):181            if ww := w.mapper(w.value):182                tuples.extend(ww)183        elif isinstance(w, ChannelWriteEntry):184            value = w.mapper(w.value) if w.mapper is not None else w.value185            if value is SKIP_WRITE:186                continue187            if w.skip_none and value is None:188                continue189            tuples.append((w.channel, value))190        else:191            raise ValueError(f"Invalid write entry: {w}")192    return tuples193 
codekingpro/portable-devtools · Team Ai