diff --git a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/_delta.py b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/_delta.py index a28313bdb..21251c0e0 100644 --- a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/_delta.py +++ b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/_delta.py @@ -36,6 +36,10 @@ from langgraph.checkpoint.base import DeltaChannelHistory, PendingWrite # descendant's config makes the chain a loop. `step_walk_with_row` stops on a # repeated id; sqlite yields recursive rows lazily, so abandoning the cursor # ends the recursion. +# +# `CROSS JOIN` pins `ancestors` as the outer loop, so each step is one primary +# key lookup. With a plain `JOIN` and no `ANALYZE` stats, sqlite can put +# `checkpoints` outside and scan the whole thread per step. DELTA_STAGE1_SQL = ( "WITH RECURSIVE ancestors(checkpoint_id, parent_checkpoint_id, type, " "checkpoint) AS (" @@ -44,7 +48,7 @@ DELTA_STAGE1_SQL = ( "WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? " "UNION ALL " "SELECT c.checkpoint_id, c.parent_checkpoint_id, c.type, c.checkpoint " - "FROM checkpoints c JOIN ancestors a " + "FROM ancestors a CROSS JOIN checkpoints c " "ON c.checkpoint_id = a.parent_checkpoint_id " "WHERE c.thread_id = ? AND c.checkpoint_ns = ?" ") " diff --git a/libs/checkpoint-sqlite/tests/test_delta_parent_walk.py b/libs/checkpoint-sqlite/tests/test_delta_parent_walk.py index bdd513366..849752a5d 100644 --- a/libs/checkpoint-sqlite/tests/test_delta_parent_walk.py +++ b/libs/checkpoint-sqlite/tests/test_delta_parent_walk.py @@ -11,6 +11,7 @@ from langgraph.checkpoint.base import ( ) from langgraph.checkpoint.sqlite import SqliteSaver +from langgraph.checkpoint.sqlite._delta import DELTA_STAGE1_SQL from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver CHANNEL = "ch" @@ -87,3 +88,17 @@ def test_walk_terminates_when_put_makes_the_parent_chain_cycle() -> None: got = saver.get_delta_channel_history(config=b, channels=[CHANNEL]) assert got[CHANNEL] == {"writes": []} + + +def test_walk_step_looks_up_the_parent_by_primary_key() -> None: + with SqliteSaver.from_conn_string(":memory:") as saver: + saver.setup() + plan = [ + row[3] + for row in saver.conn.execute( + f"EXPLAIN QUERY PLAN {DELTA_STAGE1_SQL}", ("t", "", "id", "t", "") + ) + ] + assert any( + step.startswith("SEARCH c ") and "checkpoint_id=?" in step for step in plan + ), f"recursive step should look up the parent by key, got {plan}"