Merge pull request #263 from langchain-ai/nc/1apr/rename-channel-invoke

Rename ChannelInvoke to PregelNode
This commit is contained in:
Nuno Campos
2024-04-01 19:37:12 -07:00
committed by GitHub
5 changed files with 23 additions and 23 deletions
+2 -2
View File
@@ -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)
+2 -2
View File
@@ -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=(
+9 -9
View File
@@ -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,
+7 -7
View File
@@ -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,
+3 -3
View File
@@ -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(