diff --git a/langgraph/checkpoint/sqlite.py b/langgraph/checkpoint/sqlite.py index 84dd96a82..e3bd33efb 100644 --- a/langgraph/checkpoint/sqlite.py +++ b/langgraph/checkpoint/sqlite.py @@ -1,3 +1,4 @@ +import json import pickle import sqlite3 import threading @@ -337,6 +338,62 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): ), ) + def search( + self, + metadata_query: CheckpointMetadata, + *, + before: Optional[RunnableConfig] = None, + limit: Optional[int] = None, + ) -> Iterator[CheckpointTuple]: + """Search for checkpoints by metadata. + + This method retrieves a list of checkpoint tuples from the SQLite + database based on the provided metadata query. The metadata query does + not need to contain all keys defined in the CheckpointMetadata class. + The checkpoints are ordered by timestamp in descending order. + + Args: + metadata_query (CheckpointMetadata): The metadata query to use for searching the checkpoints. + before (Optional[RunnableConfig]): If provided, only checkpoints before the specified timestamp are returned. Defaults to None. + limit (Optional[int]): The maximum number of checkpoints to return. Defaults to None. + + Yields: + Iterator[CheckpointTuple]: An iterator of checkpoint tuples. + """ + query = ( + f"SELECT json_extract(CAST(metadata AS TEXT), '$.writes'), thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints {self.search_where(metadata_query)}ORDER BY thread_ts DESC" + if before is None + else f"SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints {self.search_where(metadata_query)}AND thread_ts < ? ORDER BY thread_ts DESC" + ) + if limit: + query += f" LIMIT {limit}" + + print("final query", query) + with self.cursor(transaction=False) as cur: + cur.execute( + query, + ( + () if before is None else (before["configurable"]["thread_ts"],) + ), + ) + for writes, thread_id, thread_ts, parent_ts, value, metadata in cur: + print("writes after json extract", writes) + yield CheckpointTuple( + {"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}}, + self.serde.loads(value), + self.serde.loads(metadata) if metadata is not None else {}, + ( + { + "configurable": { + "thread_id": thread_id, + "thread_ts": parent_ts, + } + } + if parent_ts + else None + ), + ) + def put( self, config: RunnableConfig, @@ -382,3 +439,37 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): "thread_ts": checkpoint["ts"], } } + + def search_where(self, metadata_query: CheckpointMetadata) -> str: + """Return WHERE clause for (a)search() given metadata query. + + This method returns the operator as well (=, IS). + """ + def _where_value(query_value: Any) -> str: + if query_value is None: + return "IS NULL" + elif isinstance(query_value, str): + return f"= '{query_value}'" + elif isinstance(query_value, int) or isinstance(query_value, float): + return f"= {query_value}" + elif isinstance(query_value, bool): + return f"= {1 if query_value else 0}" + elif isinstance(query_value, dict) or isinstance(query_value, list): + # query value for JSON object cannot have trailing space after separators (, :) + # SQLite json_extract() returns JSON string without whitespace + return f"= '{json.dumps(query_value, separators=(',', ':'))}'" + else: + return f"= '{str(query_value)}'" + + where = "WHERE " + for query_key, query_value in metadata_query.items(): + where += f"json_extract(CAST(metadata AS TEXT), '$.{query_key}') {_where_value(query_value)} AND " + + if where == "WHERE ": + # there are no query key/value pairs + return "" + else: + # remove trailing AND + where = where[:-4] + # where clause contains an extra trailing space + return where diff --git a/tests/checkpoint/test_memory.py b/tests/checkpoint/test_memory.py index 9b5bf5660..a271f2703 100644 --- a/tests/checkpoint/test_memory.py +++ b/tests/checkpoint/test_memory.py @@ -44,12 +44,6 @@ class TestMemorySaver: "score": None, } - async def _async_iterator_to_list(self, async_iterator: AsyncIterator): - result = [] - async for item in async_iterator: - result.append(item) - return result - async def test_search(self): # set up test # save checkpoints @@ -90,7 +84,6 @@ class TestMemorySaver: query_3: CheckpointMetadata = {} # search by no keys, return all checkpoints query_4: CheckpointMetadata = {"source": "update", "step": 1} # no match - search_results_1 = [c async for c in self.memory_saver.asearch(query_1)] assert len(search_results_1) == 1 assert search_results_1[0].metadata == self.metadata_1 diff --git a/tests/checkpoint/test_sqlite.py b/tests/checkpoint/test_sqlite.py new file mode 100644 index 000000000..2f66b6ee9 --- /dev/null +++ b/tests/checkpoint/test_sqlite.py @@ -0,0 +1,77 @@ +import pytest + +from langchain_core.runnables import RunnableConfig + +from langgraph.checkpoint.base import Checkpoint, CheckpointMetadata +from langgraph.checkpoint.sqlite import SqliteSaver + + +class TestMemorySaver: + @pytest.fixture(autouse=True) + def setup(self): + self.sqlite_saver = SqliteSaver.from_conn_string(":memory:") + + # objects for test setup + self.config_1: RunnableConfig = {"configurable": {"thread_id": "thread-1", "thread_ts": "1"}} + self.config_2: RunnableConfig = {"configurable": {"thread_id": "thread-2", "thread_ts": "2"}} + + self.chkpnt_1: Checkpoint = { + "v": 1, + "ts": "1", + "channel_values": {}, + "channel_versions": {}, + "versions_seen": {} + } + self.chkpnt_2: Checkpoint = { + "v": 2, + "ts": "2", + "channel_values": {}, + "channel_versions": {}, + "versions_seen": {} + } + + self.metadata_1: CheckpointMetadata = { + "source": "input", + "step": 2, + "writes": {}, + "score": 1, + } + self.metadata_2: CheckpointMetadata = { + "source": "loop", + "step": 1, + "writes": {"foo": "bar"}, + "score": None, + } + + def test_search(self): + # set up test + # save checkpoints + self.sqlite_saver.put(self.config_1, self.chkpnt_1, self.metadata_1) + self.sqlite_saver.put(self.config_2, self.chkpnt_2, self.metadata_2) + + # call method / assertions + query_1: CheckpointMetadata = {"source": "input"} # search by 1 key + query_2: CheckpointMetadata = {"step": 1, "writes": {"foo": "bar"}} # search by multiple keys + query_3: CheckpointMetadata = {} # search by no keys, return all checkpoints + query_4: CheckpointMetadata = {"source": "update", "step": 1} # no match + + search_results_1 = list(self.sqlite_saver.search(query_1)) + assert len(search_results_1) == 1 + assert search_results_1[0].metadata == self.metadata_1 + + search_results_2 = list(self.sqlite_saver.search(query_2)) + assert len(search_results_2) == 1 + assert search_results_2[0].metadata == self.metadata_2 + + search_results_3 = list(self.sqlite_saver.search(query_3)) + assert len(search_results_3) == 2 + + search_results_4 = list(self.sqlite_saver.search(query_4)) + assert len(search_results_4) == 0 + + # TODO: test before and limit params + + def test_create_where(self): + # call method / assertions + expected_where = "WHERE json_extract(CAST(metadata AS TEXT), '$.source') = 'loop' AND json_extract(CAST(metadata AS TEXT), '$.step') = 1 AND json_extract(CAST(metadata AS TEXT), '$.writes') = '{\"foo\":\"bar\"}' AND json_extract(CAST(metadata AS TEXT), '$.score') IS NULL " + assert self.sqlite_saver.search_where(self.metadata_2) == expected_where