diff --git a/libs/langgraph/tests/memory_assert.py b/libs/langgraph/tests/memory_assert.py index c0ede2f20..09c9e52af 100644 --- a/libs/langgraph/tests/memory_assert.py +++ b/libs/langgraph/tests/memory_assert.py @@ -7,6 +7,7 @@ from langchain_core.runnables import RunnableConfig from langgraph.checkpoint.base import ( Checkpoint, CheckpointMetadata, + CheckpointTuple, SerializerProtocol, copy_checkpoint, ) @@ -119,3 +120,11 @@ class MemorySaverAssertCheckpointMetadata(MemorySaver): return await asyncio.get_running_loop().run_in_executor( None, self.put, config, checkpoint, metadata ) + + +class MemorySaverNoPending(MemorySaver): + def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]: + result = super().get_tuple(config) + if result: + return CheckpointTuple(result.config, result.checkpoint, result.metadata) + return result diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index db902eca0..6740d9a96 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -66,6 +66,7 @@ from tests.any_str import AnyStr from tests.memory_assert import ( MemorySaverAssertCheckpointMetadata, MemorySaverAssertImmutable, + MemorySaverNoPending, NoopSerializer, ) @@ -1149,10 +1150,32 @@ def test_cond_edge_after_send() -> None: builder.add_conditional_edges("1", send_for_fun) builder.add_conditional_edges("2", route_to_three) graph = builder.compile() - assert graph.invoke(["0"]) == ["0", "1", "2", "3"] +async def test_checkpointer_null_pending_writes() -> None: + class Node: + def __init__(self, name: str): + self.name = name + setattr(self, "__name__", name) + + def __call__(self, state): + return [self.name] + + builder = StateGraph(Annotated[list, operator.add]) + builder.add_node(Node("1")) + builder.add_edge(START, "1") + graph = builder.compile(checkpointer=MemorySaverNoPending()) + assert graph.invoke([], {"configurable": {"thread_id": "foo"}}) == ["1"] + assert graph.invoke([], {"configurable": {"thread_id": "foo"}}) == ["1"] * 2 + assert (await graph.ainvoke([], {"configurable": {"thread_id": "foo"}})) == [ + "1" + ] * 3 + assert (await graph.ainvoke([], {"configurable": {"thread_id": "foo"}})) == [ + "1" + ] * 4 + + def test_invoke_checkpoint_sqlite(mocker: MockerFixture) -> None: adder = mocker.Mock(side_effect=lambda x: x["total"] + x["input"])