code review

This commit is contained in:
vbarda
2024-08-14 12:10:07 -04:00
parent d9618880a3
commit 409b915a3f
2 changed files with 181 additions and 154 deletions
+1 -1
View File
@@ -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})"
+180 -153
View File
@@ -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(