Rename Packet to Send

This commit is contained in:
Nuno Campos
2024-05-31 15:17:22 -07:00
parent 5f3c61da98
commit 5b5323b94f
9 changed files with 26 additions and 35 deletions
+2 -2
View File
@@ -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."""
+1 -1
View File
@@ -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
+5 -7
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.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
]
+4 -7
View File
@@ -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}
),
)
)
+2 -2
View File
@@ -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}"
)
+5 -9
View File
@@ -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)
]
+2 -2
View File
@@ -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"
+2 -2
View File
@@ -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
+3 -3
View File
@@ -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