mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-21 15:12:26 +02:00
Add tests
This commit is contained in:
@@ -186,14 +186,14 @@ def map_debug_checkpoint(
|
||||
"id": t.id,
|
||||
"name": t.name,
|
||||
"error": t.error,
|
||||
"state": task_states.get(t.id) if task_states else None,
|
||||
"state": t.state,
|
||||
}
|
||||
if t.error
|
||||
else {
|
||||
"id": t.id,
|
||||
"name": t.name,
|
||||
"interrupts": tuple(asdict(i) for i in t.interrupts),
|
||||
"state": task_states.get(t.id) if task_states else None,
|
||||
"state": t.state,
|
||||
}
|
||||
for t in tasks_w_writes(tasks, pending_writes, task_states)
|
||||
],
|
||||
|
||||
@@ -6996,13 +6996,13 @@ def test_branch_then(
|
||||
# test stream_mode=debug
|
||||
tool_two = tool_two_graph.compile(checkpointer=checkpointer)
|
||||
thread10 = {"configurable": {"thread_id": "10"}}
|
||||
|
||||
|
||||
res = [
|
||||
*tool_two.stream(
|
||||
{"my_key": "value", "market": "DE"}, thread10, stream_mode="debug"
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
assert res == [
|
||||
{
|
||||
"type": "checkpoint",
|
||||
@@ -7029,7 +7029,14 @@ def test_branch_then(
|
||||
},
|
||||
"parent_config": None,
|
||||
"next": ["__start__"],
|
||||
"tasks": [{"id": AnyStr(), "name": "__start__", "interrupts": (), "state": None}],
|
||||
"tasks": [
|
||||
{
|
||||
"id": AnyStr(),
|
||||
"name": "__start__",
|
||||
"interrupts": (),
|
||||
"state": None,
|
||||
}
|
||||
],
|
||||
},
|
||||
},
|
||||
{
|
||||
@@ -7070,7 +7077,9 @@ def test_branch_then(
|
||||
},
|
||||
},
|
||||
"next": ["prepare"],
|
||||
"tasks": [{"id": AnyStr(), "name": "prepare", "interrupts": (), "state": None}],
|
||||
"tasks": [
|
||||
{"id": AnyStr(), "name": "prepare", "interrupts": (), "state": None}
|
||||
],
|
||||
},
|
||||
},
|
||||
{
|
||||
@@ -7134,7 +7143,14 @@ def test_branch_then(
|
||||
},
|
||||
},
|
||||
"next": ["tool_two_slow"],
|
||||
"tasks": [{"id": AnyStr(), "name": "tool_two_slow", "interrupts": (), "state": None}],
|
||||
"tasks": [
|
||||
{
|
||||
"id": AnyStr(),
|
||||
"name": "tool_two_slow",
|
||||
"interrupts": (),
|
||||
"state": None,
|
||||
}
|
||||
],
|
||||
},
|
||||
},
|
||||
{
|
||||
@@ -7198,7 +7214,9 @@ def test_branch_then(
|
||||
},
|
||||
},
|
||||
"next": ["finish"],
|
||||
"tasks": [{"id": AnyStr(), "name": "finish", "interrupts": (), "state": None}],
|
||||
"tasks": [
|
||||
{"id": AnyStr(), "name": "finish", "interrupts": (), "state": None}
|
||||
],
|
||||
},
|
||||
},
|
||||
{
|
||||
@@ -11564,3 +11582,75 @@ def test_enum_node_names():
|
||||
graph = graph.compile()
|
||||
|
||||
assert graph.invoke({"foo": "hello"}) == {"foo": "hello", "bar": "hello!"}
|
||||
|
||||
|
||||
def test_debug_subgraphs():
|
||||
class State(TypedDict):
|
||||
messages: Annotated[list[str], operator.add]
|
||||
|
||||
def node(name):
|
||||
def _node(state: State):
|
||||
return {"messages": [f"entered {name} node"]}
|
||||
|
||||
return _node
|
||||
|
||||
grand_parent = StateGraph(State)
|
||||
parent = StateGraph(State)
|
||||
child = StateGraph(State)
|
||||
|
||||
child.add_node("c_one", node("c_one"))
|
||||
child.add_node("c_two", node("c_two"))
|
||||
child.add_edge(START, "c_one")
|
||||
child.add_edge("c_one", "c_two")
|
||||
child.add_edge("c_two", END)
|
||||
|
||||
parent.add_node("p_one", node("p_one"))
|
||||
parent.add_node("p_two", child.compile())
|
||||
parent.add_edge(START, "p_one")
|
||||
parent.add_edge("p_one", "p_two")
|
||||
parent.add_edge("p_two", END)
|
||||
|
||||
grand_parent.add_node("gp_one", node("gp_one"))
|
||||
grand_parent.add_node("gp_two", parent.compile())
|
||||
grand_parent.add_edge(START, "gp_one")
|
||||
grand_parent.add_edge("gp_one", "gp_two")
|
||||
grand_parent.add_edge("gp_two", END)
|
||||
|
||||
graph = grand_parent.compile(checkpointer=MemorySaver())
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
events = [
|
||||
*graph.stream(
|
||||
{"messages": []},
|
||||
config=config,
|
||||
stream_mode="debug",
|
||||
)
|
||||
]
|
||||
|
||||
checkpoint_events = list(
|
||||
reversed([e["payload"] for e in events if e["type"] == "checkpoint"])
|
||||
)
|
||||
checkpoint_history = list(graph.get_state_history(config))
|
||||
|
||||
assert len(checkpoint_events) == len(checkpoint_history)
|
||||
|
||||
def normalize_config(config: dict | None) -> dict | None:
|
||||
if config is None:
|
||||
return None
|
||||
return config["configurable"]
|
||||
|
||||
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(
|
||||
history.parent_config
|
||||
)
|
||||
|
||||
assert len(stream["tasks"]) == len(history.tasks)
|
||||
for stream_task, history_task in zip(stream["tasks"], history.tasks):
|
||||
assert stream_task["id"] == history_task.id
|
||||
assert stream_task["name"] == history_task.name
|
||||
assert stream_task["interrupts"] == history_task.interrupts
|
||||
assert stream_task.get("error") == history_task.error
|
||||
assert stream_task.get("state") == history_task.state
|
||||
|
||||
Reference in New Issue
Block a user