From 5b5323b94fb0acb9983c0443f47de942d8232c95 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 31 May 2024 15:17:22 -0700 Subject: [PATCH] Rename Packet to Send --- langgraph/checkpoint/base.py | 4 ++-- langgraph/constants.py | 2 +- langgraph/graph/graph.py | 12 +++++------- langgraph/graph/state.py | 11 ++++------- langgraph/pregel/__init__.py | 4 ++-- langgraph/pregel/write.py | 14 +++++--------- langgraph/serde/jsonplus.py | 4 ++-- tests/test_pregel.py | 4 ++-- tests/test_pregel_async.py | 6 +++--- 9 files changed, 26 insertions(+), 35 deletions(-) diff --git a/langgraph/checkpoint/base.py b/langgraph/checkpoint/base.py index 673a83278..19d550af0 100644 --- a/langgraph/checkpoint/base.py +++ b/langgraph/checkpoint/base.py @@ -15,7 +15,7 @@ from typing import ( from langchain_core.runnables import ConfigurableFieldSpec, RunnableConfig from langgraph.checkpoint.id import uuid6 -from langgraph.constants import Packet +from langgraph.constants import Send from langgraph.serde.base import SerializerProtocol from langgraph.serde.jsonplus import JsonPlusSerializer @@ -74,7 +74,7 @@ class Checkpoint(TypedDict): Used to determine which nodes to execute next. """ - pending_packets: List[Packet] + pending_packets: List[Send] """List of packets sent to nodes but not yet processed. Cleared by the next checkpoint.""" diff --git a/langgraph/constants.py b/langgraph/constants.py index f4c372e71..35e4bbd1a 100644 --- a/langgraph/constants.py +++ b/langgraph/constants.py @@ -11,6 +11,6 @@ RESERVED = {INTERRUPT, TASKS, CONFIG_KEY_SEND, CONFIG_KEY_READ} TAG_HIDDEN = "langsmith:hidden" -class Packet(NamedTuple): +class Send(NamedTuple): node: str arg: Any diff --git a/langgraph/graph/graph.py b/langgraph/graph/graph.py index 8acc8f9b4..af9b7bc6c 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.constants import TAG_HIDDEN, Packet +from langgraph.constants import TAG_HIDDEN, Send from langgraph.errors import InvalidUpdateError from langgraph.pregel import Channel, Pregel from langgraph.pregel.read import PregelNode @@ -105,14 +105,12 @@ class Branch(NamedTuple): if not isinstance(result, list): result = [result] if self.ends: - destinations = [ - r if isinstance(r, Packet) else self.ends[r] for r in result - ] + destinations = [r if isinstance(r, Send) else self.ends[r] for r in result] else: destinations = result if any(dest is None or dest == START for dest in destinations): raise ValueError("Branch did not return a valid destination") - if any(p.node == END for p in destinations if isinstance(p, Packet)): + if any(p.node == END for p in destinations if isinstance(p, Send)): raise InvalidUpdateError("Cannot send a packet to the END node") return writer(destinations) or input @@ -404,10 +402,10 @@ class CompiledGraph(Pregel): self.nodes[end].channels.append(start) def attach_branch(self, start: str, name: str, branch: Branch) -> None: - def branch_writer(packets: list[Union[str, Packet]]) -> Optional[ChannelWrite]: + def branch_writer(packets: list[Union[str, Send]]) -> Optional[ChannelWrite]: writes = [ ChannelWriteEntry(f"branch:{start}:{name}:{p}" if p != END else END) - if not isinstance(p, Packet) + if not isinstance(p, Send) else p for p in packets ] diff --git a/langgraph/graph/state.py b/langgraph/graph/state.py index 49bd34e37..1af5d3ae3 100644 --- a/langgraph/graph/state.py +++ b/langgraph/graph/state.py @@ -27,7 +27,7 @@ from langgraph.channels.named_barrier_value import NamedBarrierValue from langgraph.checkpoint import BaseCheckpointSaver from langgraph.constants import TAG_HIDDEN from langgraph.errors import InvalidUpdateError -from langgraph.graph.graph import END, START, Branch, CompiledGraph, Graph, Packet +from langgraph.graph.graph import END, START, Branch, CompiledGraph, Graph, Send from langgraph.managed.base import ManagedValue, is_managed_value from langgraph.pregel.read import ChannelRead, PregelNode from langgraph.pregel.types import All @@ -382,11 +382,11 @@ class CompiledStateGraph(CompiledGraph): ) def attach_branch(self, start: str, name: str, branch: Branch) -> None: - def branch_writer(packets: list[Union[str, Packet]]) -> Optional[ChannelWrite]: + def branch_writer(packets: list[Union[str, Send]]) -> Optional[ChannelWrite]: if filtered := [p for p in packets if p != END]: writes = [ ChannelWriteEntry(f"branch:{start}:{name}:{p}", start) - if not isinstance(p, Packet) + if not isinstance(p, Send) else p for p in filtered ] @@ -395,10 +395,7 @@ class CompiledStateGraph(CompiledGraph): ChannelWriteEntry( f"branch:{start}:{name}:then", WaitForNames( - { - p.node if isinstance(p, Packet) else p - for p in filtered - } + {p.node if isinstance(p, Send) else p for p in filtered} ), ) ) diff --git a/langgraph/pregel/__init__.py b/langgraph/pregel/__init__.py index ca7f787c3..2c6f7a306 100644 --- a/langgraph/pregel/__init__.py +++ b/langgraph/pregel/__init__.py @@ -69,7 +69,7 @@ from langgraph.constants import ( INTERRUPT, TAG_HIDDEN, TASKS, - Packet, + Send, ) from langgraph.errors import GraphRecursionError, InvalidUpdateError from langgraph.managed.base import ( @@ -1532,7 +1532,7 @@ def _local_write( ) -> None: for chan, value in writes: if chan == TASKS: - if not isinstance(value, Packet): + if not isinstance(value, Send): raise InvalidUpdateError( f"Invalid packet type, expected Packet, got {value}" ) diff --git a/langgraph/pregel/write.py b/langgraph/pregel/write.py index 82610c322..e85885cc9 100644 --- a/langgraph/pregel/write.py +++ b/langgraph/pregel/write.py @@ -16,7 +16,7 @@ from typing import ( from langchain_core.runnables import Runnable, RunnableConfig from langchain_core.runnables.utils import ConfigurableFieldSpec -from langgraph.constants import CONFIG_KEY_SEND, TASKS, Packet +from langgraph.constants import CONFIG_KEY_SEND, TASKS, Send from langgraph.errors import InvalidUpdateError from langgraph.utils import RunnableCallable @@ -36,7 +36,7 @@ class ChannelWriteEntry(NamedTuple): class ChannelWrite(RunnableCallable): - writes: Sequence[Union[ChannelWriteEntry, Packet]] + writes: Sequence[Union[ChannelWriteEntry, Send]] """ Sequence of write entries, each of which is a tuple of: - channel name @@ -50,7 +50,7 @@ class ChannelWrite(RunnableCallable): def __init__( self, - writes: Sequence[Union[ChannelWriteEntry, Packet]], + writes: Sequence[Union[ChannelWriteEntry, Send]], *, tags: Optional[list[str]] = None, require_at_least_one_of: Optional[Sequence[str]] = None, @@ -83,9 +83,7 @@ class ChannelWrite(RunnableCallable): def _write(self, input: Any, config: RunnableConfig) -> None: # split packets and entries - writes = [ - (TASKS, packet) for packet in self.writes if isinstance(packet, Packet) - ] + writes = [(TASKS, packet) for packet in self.writes if isinstance(packet, Send)] entries = [ write for write in self.writes if isinstance(write, ChannelWriteEntry) ] @@ -115,9 +113,7 @@ class ChannelWrite(RunnableCallable): async def _awrite(self, input: Any, config: RunnableConfig) -> None: # split packets and entries - writes = [ - (TASKS, packet) for packet in self.writes if isinstance(packet, Packet) - ] + writes = [(TASKS, packet) for packet in self.writes if isinstance(packet, Send)] entries = [ write for write in self.writes if isinstance(write, ChannelWriteEntry) ] diff --git a/langgraph/serde/jsonplus.py b/langgraph/serde/jsonplus.py index f4f49a7a4..73b6d7413 100644 --- a/langgraph/serde/jsonplus.py +++ b/langgraph/serde/jsonplus.py @@ -9,7 +9,7 @@ from uuid import UUID from langchain_core.load.load import Reviver from langchain_core.load.serializable import Serializable -from langgraph.constants import Packet +from langgraph.constants import Send from langgraph.serde.base import SerializerProtocol LC_REVIVER = Reviver() @@ -65,7 +65,7 @@ class JsonPlusSerializer(SerializerProtocol): elif isinstance(obj, Enum): return self._encode_constructor_args(obj.__class__, args=[obj.value]) elif isinstance(obj, NamedTuple): - return self._encode_constructor_args(Packet, args=[*obj]) + return self._encode_constructor_args(Send, args=[*obj]) else: raise TypeError( f"Object of type {obj.__class__.__name__} is not JSON serializable" diff --git a/tests/test_pregel.py b/tests/test_pregel.py index 4997b62fe..5b7cbcff0 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -34,7 +34,7 @@ from langgraph.channels.context import Context from langgraph.channels.last_value import LastValue from langgraph.channels.topic import Topic from langgraph.checkpoint.sqlite import SqliteSaver -from langgraph.constants import Packet +from langgraph.constants import Send from langgraph.errors import InvalidUpdateError from langgraph.graph import END, Graph from langgraph.graph.graph import START @@ -3503,7 +3503,7 @@ def test_state_graph_packets() -> None: ), "nodes can pass extra data to their cond edges, which isn't saved in state" # Logic to decide whether to continue in the loop or exit if tool_calls := data["messages"][-1].tool_calls: - return [Packet("tools", tool_call) for tool_call in tool_calls] + return [Send("tools", tool_call) for tool_call in tool_calls] else: return END diff --git a/tests/test_pregel_async.py b/tests/test_pregel_async.py index 944a6bef8..dabdf964f 100644 --- a/tests/test_pregel_async.py +++ b/tests/test_pregel_async.py @@ -32,7 +32,7 @@ from langgraph.channels.context import Context from langgraph.channels.last_value import LastValue from langgraph.channels.topic import Topic from langgraph.checkpoint.aiosqlite import AsyncSqliteSaver -from langgraph.constants import Packet +from langgraph.constants import Send from langgraph.errors import InvalidUpdateError from langgraph.graph import END, Graph, StateGraph from langgraph.graph.graph import START @@ -2584,7 +2584,7 @@ Some examples of past conversations: def should_continue(data: AgentState) -> str: # Logic to decide whether to continue in the loop or exit if tool_calls := data["messages"][-1].tool_calls: - return [Packet("tools", tool_call) for tool_call in tool_calls] + return [Send("tools", tool_call) for tool_call in tool_calls] else: return "exit" @@ -3148,7 +3148,7 @@ async def test_state_graph_packets() -> None: def should_continue(data: AgentState) -> str: # Logic to decide whether to continue in the loop or exit if tool_calls := data["messages"][-1].tool_calls: - return [Packet("tools", tool_call) for tool_call in tool_calls] + return [Send("tools", tool_call) for tool_call in tool_calls] else: return END