mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-26 01:22:24 +02:00
code review
This commit is contained in:
@@ -86,7 +86,7 @@ class Send:
|
||||
self.id = id or str(uuid4())
|
||||
|
||||
def __hash__(self) -> int:
|
||||
return hash((self.node, self.arg))
|
||||
return hash((self.node, self.arg, self.id))
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"Send(node={self.node!r}, arg={self.arg!r}, id={self.id!r})"
|
||||
|
||||
@@ -195,6 +195,145 @@ class Channel:
|
||||
)
|
||||
|
||||
|
||||
def _get_nodes_and_channels(
|
||||
nodes: Mapping[str, PregelNode],
|
||||
channels: Mapping[str, BaseChannel],
|
||||
checkpoint_ns: str,
|
||||
) -> tuple[Mapping[str, PregelNode], Mapping[str, BaseChannel]]:
|
||||
if checkpoint_ns == "":
|
||||
return nodes, channels
|
||||
|
||||
path = checkpoint_ns.split(CHECKPOINT_NAMESPACE_SEPARATOR)
|
||||
for subgraph_node_name in path:
|
||||
# if we have this separator it means we have a node that was triggered by Send
|
||||
if SEND_CHECKPOINT_NAMESPACE_SEPARATOR in subgraph_node_name:
|
||||
name_parts = subgraph_node_name.split(SEND_CHECKPOINT_NAMESPACE_SEPARATOR)
|
||||
if len(name_parts) != 2:
|
||||
raise ValueError(f"Malformed node name '{subgraph_node_name}'")
|
||||
|
||||
subgraph_node_name = name_parts[0]
|
||||
if subgraph_node_name not in nodes:
|
||||
raise ValueError(f"Couldn't find node '{subgraph_node_name}'.")
|
||||
|
||||
subgraph_node = nodes[subgraph_node_name].get_node()
|
||||
|
||||
if not isinstance(subgraph_node, RunnableSequence):
|
||||
break
|
||||
|
||||
first_step = subgraph_node.steps[0]
|
||||
if isinstance(first_step, Pregel):
|
||||
nodes = first_step.nodes
|
||||
channels = first_step.channels
|
||||
|
||||
return nodes, channels
|
||||
|
||||
|
||||
def _assemble_state_snapshot_hierarchy(
|
||||
root_checkpoint_ns: str,
|
||||
checkpoint_ns_to_state_snapshots: dict[str, StateSnapshot],
|
||||
) -> StateSnapshot:
|
||||
checkpoint_ns_list_to_visit = sorted(
|
||||
checkpoint_ns_to_state_snapshots.keys(),
|
||||
key=lambda x: len(x.split(CHECKPOINT_NAMESPACE_SEPARATOR)),
|
||||
)
|
||||
while checkpoint_ns_list_to_visit:
|
||||
checkpoint_ns = checkpoint_ns_list_to_visit.pop()
|
||||
state_snapshot = checkpoint_ns_to_state_snapshots[checkpoint_ns]
|
||||
*path, subgraph_node = checkpoint_ns.split(CHECKPOINT_NAMESPACE_SEPARATOR)
|
||||
parent_checkpoint_ns = CHECKPOINT_NAMESPACE_SEPARATOR.join(path)
|
||||
if subgraph_node and (
|
||||
parent_state_snapshot := checkpoint_ns_to_state_snapshots.get(
|
||||
parent_checkpoint_ns
|
||||
)
|
||||
):
|
||||
parent_subgraph_snapshots = {
|
||||
**(parent_state_snapshot.subgraph_state_snapshots or {}),
|
||||
subgraph_node: state_snapshot,
|
||||
}
|
||||
checkpoint_ns_to_state_snapshots[
|
||||
parent_checkpoint_ns
|
||||
] = checkpoint_ns_to_state_snapshots[parent_checkpoint_ns]._replace(
|
||||
subgraph_state_snapshots=parent_subgraph_snapshots
|
||||
)
|
||||
|
||||
state_snapshot = checkpoint_ns_to_state_snapshots.pop(root_checkpoint_ns, None)
|
||||
if state_snapshot is None:
|
||||
raise ValueError(f"Missing checkpoint for checkpoint NS '{root_checkpoint_ns}'")
|
||||
return state_snapshot
|
||||
|
||||
|
||||
def _prepare_state_snapshot(
|
||||
saved: CheckpointTuple,
|
||||
nodes: Mapping[str, PregelNode],
|
||||
channels: Mapping[str, BaseChannel],
|
||||
managed_values_dict: dict[str, ManagedValueSpec],
|
||||
select_channels: str | list[str],
|
||||
) -> StateSnapshot:
|
||||
with ChannelsManager(
|
||||
{
|
||||
k: LastValue(None) if isinstance(c, Context) else c
|
||||
for k, c in channels.items()
|
||||
},
|
||||
saved.checkpoint,
|
||||
saved.config,
|
||||
) as channels, ManagedValuesManager(
|
||||
managed_values_dict, ensure_config(saved.config)
|
||||
) as managed:
|
||||
next_tasks = prepare_next_tasks(
|
||||
saved.checkpoint,
|
||||
nodes,
|
||||
channels,
|
||||
managed,
|
||||
saved.config,
|
||||
-1,
|
||||
for_execution=False,
|
||||
)
|
||||
return StateSnapshot(
|
||||
values=read_channels(channels, select_channels),
|
||||
next=tuple(t.name for t in next_tasks),
|
||||
config=saved.config,
|
||||
metadata=saved.metadata,
|
||||
created_at=saved.checkpoint["ts"],
|
||||
parent_config=saved.parent_config,
|
||||
)
|
||||
|
||||
|
||||
async def _prepare_state_snapshot_async(
|
||||
saved: CheckpointTuple,
|
||||
nodes: Mapping[str, PregelNode],
|
||||
channels: Mapping[str, BaseChannel],
|
||||
managed_values_dict: dict[str, ManagedValueSpec],
|
||||
select_channels: str | list[str],
|
||||
) -> StateSnapshot:
|
||||
async with AsyncChannelsManager(
|
||||
{
|
||||
k: LastValue(None) if isinstance(c, Context) else c
|
||||
for k, c in channels.items()
|
||||
},
|
||||
saved.checkpoint,
|
||||
saved.config,
|
||||
) as channels, AsyncManagedValuesManager(
|
||||
managed_values_dict, ensure_config(saved.config)
|
||||
) as managed:
|
||||
next_tasks = prepare_next_tasks(
|
||||
saved.checkpoint,
|
||||
nodes,
|
||||
channels,
|
||||
managed,
|
||||
saved.config,
|
||||
-1,
|
||||
for_execution=False,
|
||||
)
|
||||
return StateSnapshot(
|
||||
values=read_channels(channels, select_channels),
|
||||
next=tuple(t.name for t in next_tasks),
|
||||
config=saved.config,
|
||||
metadata=saved.metadata,
|
||||
created_at=saved.checkpoint["ts"],
|
||||
parent_config=saved.parent_config,
|
||||
)
|
||||
|
||||
|
||||
class Pregel(
|
||||
RunnableSerializable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]
|
||||
):
|
||||
@@ -361,144 +500,6 @@ class Pregel(
|
||||
if is_managed_value(v)
|
||||
}
|
||||
|
||||
def _get_nodes_and_channels(
|
||||
self, checkpoint_ns: str
|
||||
) -> tuple[Mapping[str, PregelNode], Mapping[str, BaseChannel]]:
|
||||
if checkpoint_ns == "":
|
||||
return self.nodes, self.channels
|
||||
|
||||
path = checkpoint_ns.split(CHECKPOINT_NAMESPACE_SEPARATOR)
|
||||
nodes = self.nodes
|
||||
channels = self.channels
|
||||
for subgraph_node_name in path:
|
||||
# if we have this separator it means we have a node that was triggered by Send
|
||||
if SEND_CHECKPOINT_NAMESPACE_SEPARATOR in subgraph_node_name:
|
||||
name_parts = subgraph_node_name.split(
|
||||
SEND_CHECKPOINT_NAMESPACE_SEPARATOR
|
||||
)
|
||||
if len(name_parts) != 2:
|
||||
raise ValueError(f"Malformed node name '{subgraph_node_name}'")
|
||||
|
||||
subgraph_node_name = name_parts[0]
|
||||
if subgraph_node_name not in nodes:
|
||||
raise ValueError(f"Couldn't find node '{subgraph_node_name}'.")
|
||||
|
||||
subgraph_node = nodes[subgraph_node_name].get_node()
|
||||
|
||||
if not isinstance(subgraph_node, RunnableSequence):
|
||||
break
|
||||
|
||||
first_step = subgraph_node.steps[0]
|
||||
if isinstance(first_step, Pregel):
|
||||
nodes = first_step.nodes
|
||||
channels = first_step.channels
|
||||
|
||||
return nodes, channels
|
||||
|
||||
def _prepare_state_snapshot(
|
||||
self,
|
||||
saved: CheckpointTuple,
|
||||
nodes: Mapping[str, PregelNode],
|
||||
channels: Mapping[str, BaseChannel],
|
||||
) -> StateSnapshot:
|
||||
with ChannelsManager(
|
||||
{
|
||||
k: LastValue(None) if isinstance(c, Context) else c
|
||||
for k, c in channels.items()
|
||||
},
|
||||
saved.checkpoint,
|
||||
saved.config,
|
||||
) as channels, ManagedValuesManager(
|
||||
self.managed_values_dict, ensure_config(saved.config)
|
||||
) as managed:
|
||||
next_tasks = prepare_next_tasks(
|
||||
saved.checkpoint,
|
||||
nodes,
|
||||
channels,
|
||||
managed,
|
||||
saved.config,
|
||||
-1,
|
||||
for_execution=False,
|
||||
)
|
||||
return StateSnapshot(
|
||||
values=read_channels(channels, self.stream_channels_asis),
|
||||
next=tuple(t.name for t in next_tasks),
|
||||
config=saved.config,
|
||||
metadata=saved.metadata,
|
||||
created_at=saved.checkpoint["ts"],
|
||||
parent_config=saved.parent_config,
|
||||
)
|
||||
|
||||
async def _prepare_state_snapshot_async(
|
||||
self,
|
||||
saved: CheckpointTuple,
|
||||
nodes: Mapping[str, PregelNode],
|
||||
channels: Mapping[str, BaseChannel],
|
||||
) -> StateSnapshot:
|
||||
async with AsyncChannelsManager(
|
||||
{
|
||||
k: LastValue(None) if isinstance(c, Context) else c
|
||||
for k, c in channels.items()
|
||||
},
|
||||
saved.checkpoint,
|
||||
saved.config,
|
||||
) as channels, AsyncManagedValuesManager(
|
||||
self.managed_values_dict, ensure_config(saved.config)
|
||||
) as managed:
|
||||
next_tasks = prepare_next_tasks(
|
||||
saved.checkpoint,
|
||||
nodes,
|
||||
channels,
|
||||
managed,
|
||||
saved.config,
|
||||
-1,
|
||||
for_execution=False,
|
||||
)
|
||||
return StateSnapshot(
|
||||
values=read_channels(channels, self.stream_channels_asis),
|
||||
next=tuple(t.name for t in next_tasks),
|
||||
config=saved.config,
|
||||
metadata=saved.metadata,
|
||||
created_at=saved.checkpoint["ts"],
|
||||
parent_config=saved.parent_config,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _assemble_state_snapshot_hierarchy(
|
||||
root_checkpoint_ns: str,
|
||||
checkpoint_ns_to_state_snapshots: dict[str, StateSnapshot],
|
||||
) -> StateSnapshot:
|
||||
checkpoint_ns_list_to_visit = sorted(
|
||||
checkpoint_ns_to_state_snapshots.keys(),
|
||||
key=lambda x: len(x.split(CHECKPOINT_NAMESPACE_SEPARATOR)),
|
||||
)
|
||||
while checkpoint_ns_list_to_visit:
|
||||
checkpoint_ns = checkpoint_ns_list_to_visit.pop()
|
||||
state_snapshot = checkpoint_ns_to_state_snapshots[checkpoint_ns]
|
||||
*path, subgraph_node = checkpoint_ns.split(CHECKPOINT_NAMESPACE_SEPARATOR)
|
||||
parent_checkpoint_ns = CHECKPOINT_NAMESPACE_SEPARATOR.join(path)
|
||||
if subgraph_node and (
|
||||
parent_state_snapshot := checkpoint_ns_to_state_snapshots.get(
|
||||
parent_checkpoint_ns
|
||||
)
|
||||
):
|
||||
parent_subgraph_snapshots = {
|
||||
**(parent_state_snapshot.subgraph_state_snapshots or {}),
|
||||
subgraph_node: state_snapshot,
|
||||
}
|
||||
checkpoint_ns_to_state_snapshots[
|
||||
parent_checkpoint_ns
|
||||
] = checkpoint_ns_to_state_snapshots[parent_checkpoint_ns]._replace(
|
||||
subgraph_state_snapshots=parent_subgraph_snapshots
|
||||
)
|
||||
|
||||
state_snapshot = checkpoint_ns_to_state_snapshots.pop(root_checkpoint_ns, None)
|
||||
if state_snapshot is None:
|
||||
raise ValueError(
|
||||
f"Missing checkpoint for checkpoint NS '{root_checkpoint_ns}'"
|
||||
)
|
||||
return state_snapshot
|
||||
|
||||
def get_state(
|
||||
self, config: RunnableConfig, *, include_subgraph_state: bool = False
|
||||
) -> StateSnapshot:
|
||||
@@ -541,13 +542,19 @@ class Pregel(
|
||||
if saved_checkpoint_ns not in checkpoint_ns_to_nodes_and_channels:
|
||||
checkpoint_ns_to_nodes_and_channels[
|
||||
saved_checkpoint_ns
|
||||
] = self._get_nodes_and_channels(saved_checkpoint_ns)
|
||||
] = _get_nodes_and_channels(
|
||||
self.nodes, self.channels, saved_checkpoint_ns
|
||||
)
|
||||
|
||||
nodes, channels = checkpoint_ns_to_nodes_and_channels[
|
||||
saved_checkpoint_ns
|
||||
]
|
||||
state_snapshot = self._prepare_state_snapshot(
|
||||
checkpoint_tuple, nodes, channels
|
||||
state_snapshot = _prepare_state_snapshot(
|
||||
checkpoint_tuple,
|
||||
nodes,
|
||||
channels,
|
||||
self.managed_values_dict,
|
||||
self.stream_channels_asis,
|
||||
)
|
||||
checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot
|
||||
checkpoint_ns_to_checkpoint_id[
|
||||
@@ -559,7 +566,7 @@ class Pregel(
|
||||
values={}, next=(), config=config, metadata=None, created_at=None
|
||||
)
|
||||
|
||||
state_snapshot = self._assemble_state_snapshot_hierarchy(
|
||||
state_snapshot = _assemble_state_snapshot_hierarchy(
|
||||
checkpoint_ns, checkpoint_ns_to_state_snapshots
|
||||
)
|
||||
return state_snapshot
|
||||
@@ -611,13 +618,19 @@ class Pregel(
|
||||
if saved_checkpoint_ns not in checkpoint_ns_to_nodes_and_channels:
|
||||
checkpoint_ns_to_nodes_and_channels[
|
||||
saved_checkpoint_ns
|
||||
] = self._get_nodes_and_channels(saved_checkpoint_ns)
|
||||
] = _get_nodes_and_channels(
|
||||
self.nodes, self.channels, saved_checkpoint_ns
|
||||
)
|
||||
|
||||
nodes, channels = checkpoint_ns_to_nodes_and_channels[
|
||||
saved_checkpoint_ns
|
||||
]
|
||||
state_snapshot = await self._prepare_state_snapshot_async(
|
||||
checkpoint_tuple, nodes, channels
|
||||
state_snapshot = await _prepare_state_snapshot_async(
|
||||
checkpoint_tuple,
|
||||
nodes,
|
||||
channels,
|
||||
self.managed_values_dict,
|
||||
self.stream_channels_asis,
|
||||
)
|
||||
checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot
|
||||
checkpoint_ns_to_checkpoint_id[
|
||||
@@ -629,7 +642,7 @@ class Pregel(
|
||||
values={}, next=(), config=config, metadata=None, created_at=None
|
||||
)
|
||||
|
||||
state_snapshot = self._assemble_state_snapshot_hierarchy(
|
||||
state_snapshot = _assemble_state_snapshot_hierarchy(
|
||||
checkpoint_ns, checkpoint_ns_to_state_snapshots
|
||||
)
|
||||
return state_snapshot
|
||||
@@ -669,10 +682,18 @@ class Pregel(
|
||||
)
|
||||
yield state_snapshot
|
||||
else:
|
||||
nodes, channels = self._get_nodes_and_channels(
|
||||
checkpoint_tuple.config["configurable"]["checkpoint_ns"]
|
||||
nodes, channels = _get_nodes_and_channels(
|
||||
self.nodes,
|
||||
self.channels,
|
||||
checkpoint_tuple.config["configurable"]["checkpoint_ns"],
|
||||
)
|
||||
yield _prepare_state_snapshot(
|
||||
checkpoint_tuple,
|
||||
nodes,
|
||||
channels,
|
||||
self.managed_values_dict,
|
||||
self.stream_channels_asis,
|
||||
)
|
||||
yield self._prepare_state_snapshot(checkpoint_tuple, nodes, channels)
|
||||
|
||||
async def aget_state_history(
|
||||
self,
|
||||
@@ -709,11 +730,17 @@ class Pregel(
|
||||
)
|
||||
yield state_snapshot
|
||||
else:
|
||||
nodes, channels = self._get_nodes_and_channels(
|
||||
checkpoint_tuple.config["configurable"]["checkpoint_ns"]
|
||||
nodes, channels = _get_nodes_and_channels(
|
||||
self.nodes,
|
||||
self.channels,
|
||||
checkpoint_tuple.config["configurable"]["checkpoint_ns"],
|
||||
)
|
||||
yield await self._prepare_state_snapshot_async(
|
||||
checkpoint_tuple, nodes, channels
|
||||
yield await _prepare_state_snapshot_async(
|
||||
checkpoint_tuple,
|
||||
nodes,
|
||||
channels,
|
||||
self.managed_values_dict,
|
||||
self.stream_channels_asis,
|
||||
)
|
||||
|
||||
def update_state(
|
||||
|
||||
Reference in New Issue
Block a user