mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-23 16:12:25 +02:00
Merge pull request #263 from langchain-ai/nc/1apr/rename-channel-invoke
Rename ChannelInvoke to PregelNode
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
@@ -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=(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user