Implement search() for SqliteSaver.

This commit is contained in:
Andrew Nguonly
2024-05-13 16:59:59 -07:00
parent dcb7278c56
commit 9b1efe9698
3 changed files with 168 additions and 7 deletions
+91
View File
@@ -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
-7
View File
@@ -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
+77
View File
@@ -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