Update implementation for constructing WHERE clause for SqliteSaver so that parameter values are not hardcoded, but bound instead.

This commit is contained in:
Andrew Nguonly
2024-05-14 14:44:59 -07:00
parent 6ce9a2860b
commit 78e3a240f1
4 changed files with 106 additions and 46 deletions
+31 -9
View File
@@ -2,10 +2,10 @@ import pytest
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import Checkpoint, CheckpointMetadata
from langgraph.checkpoint.sqlite import SqliteSaver, search_where
from langgraph.checkpoint.sqlite import SqliteSaver, _metadata_predicate, search_where
class TestMemorySaver:
class TestSqliteSaver:
@pytest.fixture(autouse=True)
def setup(self):
self.sqlite_saver = SqliteSaver.from_conn_string(":memory:")
@@ -78,12 +78,34 @@ class TestMemorySaver:
# TODO: test before and limit params
def test_create_where(self):
def test_search_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 = ""
expected_predicate_1 = "WHERE json_extract(CAST(metadata AS TEXT), '$.source') = ? AND json_extract(CAST(metadata AS TEXT), '$.step') = ? AND json_extract(CAST(metadata AS TEXT), '$.writes') = ? AND json_extract(CAST(metadata AS TEXT), '$.score') = ? AND thread_ts < ? "
expected_param_values_1 = ("input", 2, "{}", 1, "1")
assert search_where(self.metadata_1, self.config_1) == (
expected_predicate_1,
expected_param_values_1,
)
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
def test_metadata_predicate(self):
# call method / assertions
expected_predicate_1 = "json_extract(CAST(metadata AS TEXT), '$.source') = ? AND json_extract(CAST(metadata AS TEXT), '$.step') = ? AND json_extract(CAST(metadata AS TEXT), '$.writes') = ? AND json_extract(CAST(metadata AS TEXT), '$.score') = ? "
expected_predicate_2 = "json_extract(CAST(metadata AS TEXT), '$.source') = ? AND json_extract(CAST(metadata AS TEXT), '$.step') = ? AND json_extract(CAST(metadata AS TEXT), '$.writes') = ? AND json_extract(CAST(metadata AS TEXT), '$.score') IS ? "
expected_predicate_3 = ""
expected_param_values_1 = ("input", 2, "{}", 1)
expected_param_values_2 = ("loop", 1, '{"foo":"bar"}', None)
expected_param_values_3 = ()
assert _metadata_predicate(self.metadata_1) == (
expected_predicate_1,
expected_param_values_1,
)
assert _metadata_predicate(self.metadata_2) == (
expected_predicate_2,
expected_param_values_2,
)
assert _metadata_predicate(self.metadata_3) == (
expected_predicate_3,
expected_param_values_3,
)