codekingpro/portable-devtools
114k
1"""Replay state for subgraph checkpoint loading during time-travel."""2 3from __future__ import annotations4 5from typing import TYPE_CHECKING6 7from langgraph._internal._constants import NS_END8 9if TYPE_CHECKING:10 from langchain_core.runnables import RunnableConfig11 from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointTuple12 13 14class ReplayState:15 """Tracks which subgraphs have already loaded their pre-replay checkpoint.16 17 During a parent replay, each subgraph's first invocation should restore the18 checkpoint from before the replay point. Subsequent invocations of the same19 subgraph (e.g. in a loop) should use normal checkpoint loading so they pick20 up freshly created checkpoints.21 22 The single `ReplayState` instance is shared by reference across all derived23 configs within one parent execution.24 """25 26 __slots__ = ("checkpoint_id", "_visited_ns")27 28 def __init__(self, checkpoint_id: str) -> None:29 self.checkpoint_id = checkpoint_id30 # DO NOT CHANGE THIS VARIABLE – it may need to be rehydrated31 # in other runtimes32 self._visited_ns: set[str] = set()33 34 def _is_first_visit(self, checkpoint_ns: str) -> bool:35 """Return True the first time a subgraph namespace is seen.36 37 The task-id suffix is stripped so that the same logical subgraph38 (e.g. ``"sub_node"``) is recognized across loop iterations even39 though each iteration has a different task id.40 """41 # "sub_node:task_id" -> "sub_node"42 stable_ns = (43 checkpoint_ns.rsplit(NS_END, 1)[0]44 if NS_END in checkpoint_ns45 else checkpoint_ns46 )47 if stable_ns in self._visited_ns:48 return False49 self._visited_ns.add(stable_ns)50 return True51 52 def get_checkpoint(53 self,54 checkpoint_ns: str,55 checkpointer: BaseCheckpointSaver,56 checkpoint_config: RunnableConfig,57 ) -> CheckpointTuple | None:58 """Load the right checkpoint for a subgraph during replay.59 60 On the first call for a given subgraph namespace, returns the latest61 checkpoint created *before* the replay point. On subsequent calls62 (e.g. the same subgraph in a later loop iteration), falls back to63 normal latest-checkpoint loading.64 """65 if self._is_first_visit(checkpoint_ns):66 for saved in checkpointer.list(67 checkpoint_config,68 before={"configurable": {"checkpoint_id": self.checkpoint_id}},69 limit=1,70 ):71 return saved72 return None73 return checkpointer.get_tuple(checkpoint_config)74 75 async def aget_checkpoint(76 self,77 checkpoint_ns: str,78 checkpointer: BaseCheckpointSaver,79 checkpoint_config: RunnableConfig,80 ) -> CheckpointTuple | None:81 """Async version of `get_checkpoint`."""82 if self._is_first_visit(checkpoint_ns):83 async for saved in checkpointer.alist(84 checkpoint_config,85 before={"configurable": {"checkpoint_id": self.checkpoint_id}},86 limit=1,87 ):88 return saved89 return None90 return await checkpointer.aget_tuple(checkpoint_config)91 