fix(debug): add failing test for self-referencing

This commit is contained in:
Tat Dat Duong
2024-10-10 02:21:17 +02:00
parent db0f508269
commit ac8b51f1f2
+65 -3
View File
@@ -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
)