mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-13 05:07:51 +02:00
Implement search() for SqliteSaver.
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user