add delta overwrite graph coverage

This commit is contained in:
Sydney Runkle
2026-06-17 14:11:23 -04:00
parent 4210feccd9
commit af62aff8ac
+106
View File
@@ -9281,6 +9281,13 @@ def test_send_with_untracked_value_overlapping_keys(
assert state.values.get("dictionary") == {"session_resource": "legal_value"}
def _delta_list_reducer(state: list, writes: Sequence[list]) -> list:
out = list(state)
for write in writes:
out.extend(write)
return out
@pytest.mark.parametrize("as_json", [False, True])
def test_overwrite_sequential(
sync_checkpointer: BaseCheckpointSaver, as_json: bool
@@ -9388,6 +9395,105 @@ def test_overwrite_parallel_error(
graph.invoke({"messages": ["START"]}, config)
@pytest.mark.parametrize("as_json", [False, True])
def test_delta_channel_overwrite_sequential(
sync_checkpointer: BaseCheckpointSaver, as_json: bool
) -> None:
class State(TypedDict):
messages: Annotated[list, DeltaChannel(_delta_list_reducer)]
def node_a(state: State):
return {"messages": ["a"]}
def node_b(state: State):
overwrite = {"__overwrite__": ["b"]} if as_json else Overwrite(["b"])
return {"messages": overwrite}
builder = StateGraph(State)
builder.add_node("node_a", node_a)
builder.add_node("node_b", node_b)
builder.add_edge(START, "node_a")
builder.add_edge("node_a", "node_b")
graph = builder.compile(checkpointer=sync_checkpointer)
config = {"configurable": {"thread_id": "delta-overwrite-sequential"}}
result = graph.invoke({"messages": ["START"]}, config)
assert result == {"messages": ["b"]}
@pytest.mark.parametrize("as_json", [False, True])
def test_delta_channel_overwrite_parallel(
sync_checkpointer: BaseCheckpointSaver, as_json: bool
) -> None:
class State(TypedDict):
messages: Annotated[list, DeltaChannel(_delta_list_reducer)]
def node_a(state: State):
return {"messages": ["a"]}
def node_b(state: State):
overwrite = {"__overwrite__": ["b"]} if as_json else Overwrite(["b"])
return {"messages": overwrite}
def node_c(state: State):
return {"messages": ["c"]}
def node_d(state: State):
return {"messages": ["d"]}
builder = StateGraph(State)
builder.add_node("node_a", node_a)
builder.add_node("node_b", node_b)
builder.add_node("node_c", node_c)
builder.add_node("node_d", node_d)
builder.add_edge(START, "node_a")
builder.add_edge("node_a", "node_b")
builder.add_edge("node_a", "node_c")
builder.add_edge("node_b", "node_d")
builder.add_edge("node_c", "node_d")
graph = builder.compile(checkpointer=sync_checkpointer)
config = {"configurable": {"thread_id": "delta-overwrite-parallel"}}
result = graph.invoke({"messages": ["START"]}, config)
assert result == {"messages": ["b", "d"]}
@pytest.mark.parametrize("as_json", [False, True])
def test_delta_channel_overwrite_parallel_error(
sync_checkpointer: BaseCheckpointSaver, as_json: bool
) -> None:
class State(TypedDict):
messages: Annotated[list, DeltaChannel(_delta_list_reducer)]
def node_a(state: State):
return {"messages": ["a"]}
def node_b(state: State):
overwrite = {"__overwrite__": ["b"]} if as_json else Overwrite(["b"])
return {"messages": overwrite}
def node_c(state: State):
overwrite = {"__overwrite__": ["c"]} if as_json else Overwrite(["c"])
return {"messages": overwrite}
builder = StateGraph(State)
builder.add_node("node_a", node_a)
builder.add_node("node_b", node_b)
builder.add_node("node_c", node_c)
builder.add_edge(START, "node_a")
builder.add_edge("node_a", "node_b")
builder.add_edge("node_a", "node_c")
builder.add_edge("node_b", END)
builder.add_edge("node_c", END)
graph = builder.compile(checkpointer=sync_checkpointer)
config = {"configurable": {"thread_id": "delta-overwrite-parallel-error"}}
with pytest.raises(
InvalidUpdateError, match="Can receive only one Overwrite value per super-step."
):
graph.invoke({"messages": ["START"]}, config)
def test_fork_does_not_apply_pending_writes(
sync_checkpointer: BaseCheckpointSaver,
) -> None: