diff --git a/langgraph/graph/graph.py b/langgraph/graph/graph.py index 37ec6be78..50c633f4d 100644 --- a/langgraph/graph/graph.py +++ b/langgraph/graph/graph.py @@ -26,7 +26,7 @@ from langchain_core.runnables.graph import ( from langgraph.channels.ephemeral_value import EphemeralValue from langgraph.checkpoint import BaseCheckpointSaver from langgraph.pregel import Channel, Pregel -from langgraph.pregel.read import ChannelInvoke +from langgraph.pregel.read import PregelNode from langgraph.pregel.write import ChannelWrite logger = logging.getLogger(__name__) @@ -293,7 +293,7 @@ class CompiledGraph(Pregel): def attach_node(self, key: str, node: Runnable) -> None: self.channels[key] = EphemeralValue(Any) self.nodes[key] = ( - ChannelInvoke(channels=[], triggers=[]) | node | Channel.write_to(key) + PregelNode(channels=[], triggers=[]) | node | Channel.write_to(key) ) self.stream_channels.append(key) diff --git a/langgraph/graph/state.py b/langgraph/graph/state.py index fe92c6a99..cc064783f 100644 --- a/langgraph/graph/state.py +++ b/langgraph/graph/state.py @@ -14,7 +14,7 @@ from langgraph.channels.named_barrier_value import NamedBarrierValue from langgraph.checkpoint import BaseCheckpointSaver from langgraph.graph.graph import END, START, Branch, CompiledGraph, Graph from langgraph.pregel import Channel -from langgraph.pregel.read import ChannelInvoke, ChannelRead +from langgraph.pregel.read import ChannelRead, PregelNode from langgraph.pregel.write import SKIP_WRITE, ChannelWrite, ChannelWriteEntry logger = logging.getLogger(__name__) @@ -147,7 +147,7 @@ class CompiledStateGraph(CompiledGraph): ).pipe(ChannelWrite(state_write_entries)) else: self.channels[key] = EphemeralValue(Any) - self.nodes[key] = ChannelInvoke( + self.nodes[key] = PregelNode( triggers=[], # read state keys channels=( diff --git a/langgraph/pregel/__init__.py b/langgraph/pregel/__init__.py index dfe2ac00f..5bda279b2 100644 --- a/langgraph/pregel/__init__.py +++ b/langgraph/pregel/__init__.py @@ -76,7 +76,7 @@ from langgraph.pregel.io import ( read_channels, ) from langgraph.pregel.log import logger -from langgraph.pregel.read import ChannelInvoke +from langgraph.pregel.read import PregelNode from langgraph.pregel.types import ( PregelExecutableTask, PregelTaskDescription, @@ -112,7 +112,7 @@ class Channel: *, key: Optional[str] = None, tags: Optional[list[str]] = None, - ) -> ChannelInvoke: + ) -> PregelNode: ... @overload @@ -123,7 +123,7 @@ class Channel: *, key: None = None, tags: Optional[list[str]] = None, - ) -> ChannelInvoke: + ) -> PregelNode: ... @classmethod @@ -133,14 +133,14 @@ class Channel: *, key: Optional[str] = None, tags: Optional[list[str]] = None, - ) -> ChannelInvoke: + ) -> PregelNode: """Runs process.invoke() each time channels are updated, with a dict of the channel values as input.""" if not isinstance(channels, str) and key is not None: raise ValueError( "Can't specify a key when subscribing to multiple channels" ) - return ChannelInvoke( + return PregelNode( channels=cast( Union[Mapping[None, str], Mapping[str, str]], {key: channels} @@ -175,7 +175,7 @@ StreamMode = Literal["values", "updates"] class Pregel( RunnableSerializable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]] ): - nodes: Mapping[str, ChannelInvoke] + nodes: Mapping[str, PregelNode] channels: Mapping[str, BaseChannel] = Field(default_factory=dict) @@ -1114,7 +1114,7 @@ def _apply_writes( @overload def _prepare_next_tasks( checkpoint: Checkpoint, - processes: Mapping[str, ChannelInvoke], + processes: Mapping[str, PregelNode], channels: Mapping[str, BaseChannel], for_execution: Literal[False], ) -> tuple[Checkpoint, list[PregelTaskDescription]]: @@ -1124,7 +1124,7 @@ def _prepare_next_tasks( @overload def _prepare_next_tasks( checkpoint: Checkpoint, - processes: Mapping[str, ChannelInvoke], + processes: Mapping[str, PregelNode], channels: Mapping[str, BaseChannel], for_execution: Literal[True], ) -> tuple[Checkpoint, list[PregelExecutableTask]]: @@ -1133,7 +1133,7 @@ def _prepare_next_tasks( def _prepare_next_tasks( checkpoint: Checkpoint, - processes: Mapping[str, ChannelInvoke], + processes: Mapping[str, PregelNode], channels: Mapping[str, BaseChannel], *, for_execution: bool, diff --git a/langgraph/pregel/read.py b/langgraph/pregel/read.py index 578d22073..f1afd30f0 100644 --- a/langgraph/pregel/read.py +++ b/langgraph/pregel/read.py @@ -83,7 +83,7 @@ class ChannelRead(RunnableLambda): default_bound: RunnablePassthrough = RunnablePassthrough() -class ChannelInvoke(RunnableBindingBase): +class PregelNode(RunnableBindingBase): channels: Union[list[str], Mapping[str, str]] triggers: list[str] = Field(default_factory=list) @@ -163,14 +163,14 @@ class ChannelInvoke(RunnableBindingBase): def __repr_args__(self) -> Any: return [(k, v) for k, v in super().__repr_args__() if k != "bound"] - def join(self, channels: Sequence[str]) -> ChannelInvoke: + def join(self, channels: Sequence[str]) -> PregelNode: assert isinstance(channels, list) or isinstance( channels, tuple ), "channels must be a list or tuple" assert isinstance( self.channels, dict ), "all channels must be named when using .join()" - return ChannelInvoke( + return PregelNode( channels={ **self.channels, **{chan: chan for chan in channels}, @@ -190,9 +190,9 @@ class ChannelInvoke(RunnableBindingBase): Callable[[Any], Other], Mapping[str, Runnable[Any, Other] | Callable[[Any], Other]], ], - ) -> ChannelInvoke: + ) -> PregelNode: if ChannelWrite.is_writer(other): - return ChannelInvoke( + return PregelNode( channels=self.channels, triggers=self.triggers, mapper=self.mapper, @@ -202,7 +202,7 @@ class ChannelInvoke(RunnableBindingBase): config=self.config, ) elif self.bound is default_bound: - return ChannelInvoke( + return PregelNode( channels=self.channels, triggers=self.triggers, mapper=self.mapper, @@ -212,7 +212,7 @@ class ChannelInvoke(RunnableBindingBase): config=self.config, ) else: - return ChannelInvoke( + return PregelNode( channels=self.channels, triggers=self.triggers, mapper=self.mapper, diff --git a/langgraph/pregel/validate.py b/langgraph/pregel/validate.py index 2c62c75e5..69f29cfd4 100644 --- a/langgraph/pregel/validate.py +++ b/langgraph/pregel/validate.py @@ -2,11 +2,11 @@ from typing import Any, Mapping, Optional, Sequence, Type, Union from langgraph.channels.base import BaseChannel from langgraph.constants import INTERRUPT -from langgraph.pregel.read import ChannelInvoke +from langgraph.pregel.read import PregelNode def validate_graph( - nodes: Mapping[str, ChannelInvoke], + nodes: Mapping[str, PregelNode], channels: dict[str, BaseChannel], input_channels: Union[str, Sequence[str]], output_channels: Union[str, Sequence[str]], @@ -19,7 +19,7 @@ def validate_graph( for name, node in nodes.items(): if name == INTERRUPT: raise ValueError(f"Node name {INTERRUPT} is reserved") - if isinstance(node, ChannelInvoke): + if isinstance(node, PregelNode): subscribed_channels.update(node.triggers) else: raise TypeError(