code review

This commit is contained in:
vbarda
2024-08-22 17:33:17 -04:00
parent 4935cf52bf
commit 72893d9abb
3 changed files with 47 additions and 48 deletions
+10 -47
View File
@@ -88,7 +88,10 @@ from langgraph.pregel.types import (
StateSnapshot,
StreamMode,
)
from langgraph.pregel.utils import get_new_channel_versions
from langgraph.pregel.utils import (
assemble_state_snapshot_hierarchy,
get_new_channel_versions,
)
from langgraph.pregel.validate import validate_graph, validate_keys
from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry
from langgraph.store.base import BaseStore
@@ -205,40 +208,6 @@ def _get_subgraph(graph: Pregel, checkpoint_ns: str) -> Pregel:
return subgraph_node.bound
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 _has_nested_interrupts(
graph: Pregel,
) -> bool:
@@ -497,7 +466,7 @@ class Pregel(
tasks=(),
)
state_snapshot = _assemble_state_snapshot_hierarchy(
state_snapshot = assemble_state_snapshot_hierarchy(
checkpoint_ns, checkpoint_ns_to_state_snapshots
)
return state_snapshot
@@ -573,7 +542,7 @@ class Pregel(
tasks=(),
)
state_snapshot = _assemble_state_snapshot_hierarchy(
state_snapshot = assemble_state_snapshot_hierarchy(
checkpoint_ns, checkpoint_ns_to_state_snapshots
)
return state_snapshot
@@ -1113,7 +1082,8 @@ class Pregel(
)
if not done:
break # timed out
for fut, task in zip(done, [futures.pop(fut) for fut in done]):
for fut in done:
task = futures.pop(fut)
if exc := _exception(fut):
# save error to checkpointer
if isinstance(exc, GraphInterrupt):
@@ -1364,7 +1334,8 @@ class Pregel(
if not done:
break # timed out
for fut, task in zip(done, [futures.pop(fut) for fut in done]):
for fut in done:
task = futures.pop(fut)
if exc := _exception(fut):
# save error to checkpointer
if isinstance(exc, GraphInterrupt):
@@ -1581,11 +1552,3 @@ def _panic_or_proceed(
inflight.pop().cancel()
# raise timeout error
raise timeout_exc_cls(f"Timed out at step {step}")
def _with_mode(mode: StreamMode, on: bool, iter: Iterator[Any]) -> Iterator[Any]:
if on:
for chunk in iter:
yield (mode, chunk)
else:
yield from iter
+1 -1
View File
@@ -93,7 +93,7 @@ class StateSnapshot(NamedTuple):
tasks: tuple[PregelTask, ...]
"""Tasks to execute in this step. If already attempted, may contain an error."""
subgraph_state_snapshots: Optional[dict[str, "StateSnapshot"]] = None
"""State snapshots of subgraphs represented as a mapping from thread ID suffix to snapshot."""
"""State snapshots of subgraphs represented as a mapping from checkpoint namespace (`checkpoint_ns`) to snapshot."""
All = Literal["*"]
+36
View File
@@ -1,4 +1,6 @@
from langgraph.checkpoint.base import ChannelVersions
from langgraph.constants import CHECKPOINT_NAMESPACE_SEPARATOR
from langgraph.pregel.types import StateSnapshot
def get_new_channel_versions(
@@ -17,3 +19,37 @@ def get_new_channel_versions(
new_versions = current_versions
return new_versions
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