mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-08 02:37:52 +02:00
Merge pull request #2070 from langchain-ai/dqbd/debug-self-referencing-checkpoint
fix(debug): self-referencing checkpoints when resuming streaming mid-thread
This commit is contained in:
@@ -250,13 +250,7 @@ class PregelLoop:
|
||||
if self.config[CONF].get(CONFIG_KEY_CHECKPOINT_NS)
|
||||
else ()
|
||||
)
|
||||
self.prev_checkpoint_config = (
|
||||
self.checkpoint_config
|
||||
if self.checkpoint_config
|
||||
and CONF in self.checkpoint_config
|
||||
and CONFIG_KEY_CHECKPOINT_ID in self.checkpoint_config[CONF]
|
||||
else None
|
||||
)
|
||||
self.prev_checkpoint_config = None
|
||||
|
||||
def put_writes(self, task_id: str, writes: Sequence[tuple[str, Any]]) -> None:
|
||||
"""Put writes for a task, to be read by the next tick."""
|
||||
@@ -740,6 +734,7 @@ class SyncPregelLoop(PregelLoop, ContextManager):
|
||||
**saved.config.get(CONF, {}),
|
||||
},
|
||||
}
|
||||
self.prev_checkpoint_config = saved.parent_config
|
||||
self.checkpoint = saved.checkpoint
|
||||
self.checkpoint_metadata = saved.metadata
|
||||
self.checkpoint_pending_writes = (
|
||||
@@ -867,6 +862,7 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
|
||||
**saved.config.get(CONF, {}),
|
||||
},
|
||||
}
|
||||
self.prev_checkpoint_config = saved.parent_config
|
||||
self.checkpoint = saved.checkpoint
|
||||
self.checkpoint_metadata = saved.metadata
|
||||
self.checkpoint_pending_writes = (
|
||||
|
||||
@@ -11584,6 +11584,66 @@ def test_enum_node_names():
|
||||
assert graph.invoke({"foo": "hello"}) == {"foo": "hello", "bar": "hello!"}
|
||||
|
||||
|
||||
def test_debug_retry():
|
||||
class State(TypedDict):
|
||||
messages: Annotated[list[str], operator.add]
|
||||
|
||||
def node(name):
|
||||
def _node(state: State):
|
||||
return {"messages": [f"entered {name} node"]}
|
||||
|
||||
return _node
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("one", node("one"))
|
||||
builder.add_node("two", node("two"))
|
||||
builder.add_edge(START, "one")
|
||||
builder.add_edge("one", "two")
|
||||
builder.add_edge("two", END)
|
||||
|
||||
saver = MemorySaver()
|
||||
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
graph.invoke({"messages": []}, config=config)
|
||||
|
||||
# re-run step: 1
|
||||
target_config = next(
|
||||
c.parent_config for c in saver.list(config) if c.metadata["step"] == 1
|
||||
)
|
||||
update_config = graph.update_state(target_config, values=None)
|
||||
|
||||
events = [*graph.stream(None, config=update_config, stream_mode="debug")]
|
||||
|
||||
checkpoint_events = list(
|
||||
reversed([e["payload"] for e in events if e["type"] == "checkpoint"])
|
||||
)
|
||||
|
||||
checkpoint_history = {
|
||||
c.config["configurable"]["checkpoint_id"]: c
|
||||
for c in graph.get_state_history(config)
|
||||
}
|
||||
|
||||
def lax_normalize_config(config: Optional[dict]) -> Optional[dict]:
|
||||
if config is None:
|
||||
return None
|
||||
return config["configurable"]
|
||||
|
||||
for stream in checkpoint_events:
|
||||
stream_conf = lax_normalize_config(stream["config"])
|
||||
stream_parent_conf = lax_normalize_config(stream["parent_config"])
|
||||
assert stream_conf != stream_parent_conf
|
||||
|
||||
# ensure the streamed checkpoint == checkpoint from checkpointer.list()
|
||||
history = checkpoint_history[stream["config"]["configurable"]["checkpoint_id"]]
|
||||
history_conf = lax_normalize_config(history.config)
|
||||
assert stream_conf == history_conf
|
||||
|
||||
history_parent_conf = lax_normalize_config(history.parent_config)
|
||||
assert stream_parent_conf == history_parent_conf
|
||||
|
||||
|
||||
def test_debug_subgraphs():
|
||||
class State(TypedDict):
|
||||
messages: Annotated[list[str], operator.add]
|
||||
@@ -11627,7 +11687,7 @@ def test_debug_subgraphs():
|
||||
|
||||
assert len(checkpoint_events) == len(checkpoint_history)
|
||||
|
||||
def normalize_config(config: Optional[dict]) -> Optional[dict]:
|
||||
def lax_normalize_config(config: Optional[dict]) -> Optional[dict]:
|
||||
if config is None:
|
||||
return None
|
||||
return config["configurable"]
|
||||
@@ -11635,8 +11695,10 @@ def test_debug_subgraphs():
|
||||
for stream, history in zip(checkpoint_events, checkpoint_history):
|
||||
assert stream["values"] == history.values
|
||||
assert stream["next"] == list(history.next)
|
||||
assert normalize_config(stream["config"]) == normalize_config(history.config)
|
||||
assert normalize_config(stream["parent_config"]) == normalize_config(
|
||||
assert lax_normalize_config(stream["config"]) == lax_normalize_config(
|
||||
history.config
|
||||
)
|
||||
assert lax_normalize_config(stream["parent_config"]) == lax_normalize_config(
|
||||
history.parent_config
|
||||
)
|
||||
|
||||
|
||||
@@ -9820,6 +9820,71 @@ async def test_store_injected_async(checkpointer_name: str, store_name: str) ->
|
||||
) # still overwriting the same one
|
||||
|
||||
|
||||
async def test_debug_retry():
|
||||
class State(TypedDict):
|
||||
messages: Annotated[list[str], operator.add]
|
||||
|
||||
def node(name):
|
||||
async def _node(state: State):
|
||||
return {"messages": [f"entered {name} node"]}
|
||||
|
||||
return _node
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("one", node("one"))
|
||||
builder.add_node("two", node("two"))
|
||||
builder.add_edge(START, "one")
|
||||
builder.add_edge("one", "two")
|
||||
builder.add_edge("two", END)
|
||||
|
||||
saver = MemorySaver()
|
||||
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
await graph.ainvoke({"messages": []}, config=config)
|
||||
|
||||
# re-run step: 1
|
||||
async for c in saver.alist(config):
|
||||
if c.metadata["step"] == 1:
|
||||
target_config = c.parent_config
|
||||
break
|
||||
assert target_config is not None
|
||||
|
||||
update_config = await graph.aupdate_state(target_config, values=None)
|
||||
|
||||
events = [
|
||||
c async for c in graph.astream(None, config=update_config, stream_mode="debug")
|
||||
]
|
||||
|
||||
checkpoint_events = list(
|
||||
reversed([e["payload"] for e in events if e["type"] == "checkpoint"])
|
||||
)
|
||||
|
||||
checkpoint_history = {
|
||||
c.config["configurable"]["checkpoint_id"]: c
|
||||
async for c in graph.aget_state_history(config)
|
||||
}
|
||||
|
||||
def lax_normalize_config(config: Optional[dict]) -> Optional[dict]:
|
||||
if config is None:
|
||||
return None
|
||||
return config["configurable"]
|
||||
|
||||
for stream in checkpoint_events:
|
||||
stream_conf = lax_normalize_config(stream["config"])
|
||||
stream_parent_conf = lax_normalize_config(stream["parent_config"])
|
||||
assert stream_conf != stream_parent_conf
|
||||
|
||||
# ensure the streamed checkpoint == checkpoint from checkpointer.list()
|
||||
history = checkpoint_history[stream["config"]["configurable"]["checkpoint_id"]]
|
||||
history_conf = lax_normalize_config(history.config)
|
||||
assert stream_conf == history_conf
|
||||
|
||||
history_parent_conf = lax_normalize_config(history.parent_config)
|
||||
assert stream_parent_conf == history_parent_conf
|
||||
|
||||
|
||||
async def test_debug_subgraphs():
|
||||
class State(TypedDict):
|
||||
messages: Annotated[list[str], operator.add]
|
||||
|
||||
Reference in New Issue
Block a user