Fix bug with WHERE clause being incorrect when passing before param.

This commit is contained in:
Andrew Nguonly
2024-05-14 10:38:26 -07:00
parent c80090465f
commit d6d33f9551
3 changed files with 39 additions and 24 deletions
+13 -10
View File
@@ -278,17 +278,20 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
Iterator[CheckpointTuple]: An iterator of checkpoint tuples.
"""
await self.setup()
query = (
f"SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints {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 {search_where(metadata_query)}AND thread_ts < ? ORDER BY thread_ts DESC"
# construct query
SELECT = "SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints "
WHERE = search_where(
metadata_query, [] if before is None else ["thread_ts < ?"]
)
if limit:
query += f" LIMIT {limit}"
async with self.conn.execute(
query,
(() if before is None else (str(before["configurable"]["thread_ts"]),)),
) as cursor:
ORDER_BY = "ORDER BY thread_ts DESC "
LIMIT = f"LIMIT {limit}" if limit else ""
query = f"{SELECT}{WHERE}{ORDER_BY}{LIMIT}"
params = () if before is None else (str(before["configurable"]["thread_ts"]),)
# execute query
async with self.conn.execute(query, params) as cursor:
async for thread_id, thread_ts, parent_ts, value, metadata in cursor:
yield CheckpointTuple(
{"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}},
+24 -14
View File
@@ -4,7 +4,7 @@ import sqlite3
import threading
from contextlib import AbstractContextManager, contextmanager
from types import TracebackType
from typing import Any, Iterator, Optional
from typing import Any, Iterator, List, Optional
from langchain_core.runnables import RunnableConfig
from typing_extensions import Self
@@ -360,18 +360,21 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
Yields:
Iterator[CheckpointTuple]: An iterator of checkpoint tuples.
"""
query = (
f"SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints {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 {search_where(metadata_query)}AND thread_ts < ? ORDER BY thread_ts DESC"
# construct query
SELECT = "SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints "
WHERE = search_where(
metadata_query, [] if before is None else ["thread_ts < ?"]
)
if limit:
query += f" LIMIT {limit}"
ORDER_BY = "ORDER BY thread_ts DESC "
LIMIT = f"LIMIT {limit}" if limit else ""
query = f"{SELECT}{WHERE}{ORDER_BY}{LIMIT}"
params = () if before is None else (before["configurable"]["thread_ts"],)
# execute query
with self.cursor(transaction=False) as cur:
cur.execute(
query,
(() if before is None else (before["configurable"]["thread_ts"],)),
)
cur.execute(query, params)
for thread_id, thread_ts, parent_ts, value, metadata in cur:
yield CheckpointTuple(
{"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}},
@@ -436,8 +439,9 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
}
def search_where(metadata_query: CheckpointMetadata) -> str:
"""Return WHERE clause for (a)search() given metadata query.
def search_where(metadata_query: CheckpointMetadata, predicates: List[str] = []) -> str:
"""Return WHERE clause for (a)search() given metadata query and
predicates.
This method returns the operator as well (=, IS).
"""
@@ -459,11 +463,17 @@ def search_where(metadata_query: CheckpointMetadata) -> str:
return f"= '{str(query_value)}'"
where = "WHERE "
# process metadata query
for query_key, query_value in metadata_query.items():
where += f"json_extract(CAST(metadata AS TEXT), '$.{query_key}') {_where_value(query_value)} AND "
# process predicates
for predicate in predicates:
where += f"{predicate} AND "
if where == "WHERE ":
# there are no query key/value pairs
# there are no query key/value pairs or predicates
return ""
else:
# remove trailing AND
+2
View File
@@ -80,8 +80,10 @@ class TestMemorySaver:
def test_create_where(self):
# call method / assertions
expected_where_1 = "WHERE json_extract(CAST(metadata AS TEXT), '$.source') = 'input' AND json_extract(CAST(metadata AS TEXT), '$.step') = 2 AND json_extract(CAST(metadata AS TEXT), '$.writes') = '{}' AND json_extract(CAST(metadata AS TEXT), '$.score') = 1 AND thread_ts < ? "
expected_where_2 = "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 "
expected_where_3 = ""
assert search_where(self.metadata_1, ["thread_ts < ?"]) == expected_where_1
assert search_where(self.metadata_2) == expected_where_2
assert search_where(self.metadata_3) == expected_where_3