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