Expose filter arg in get_state_history

- Combine list and search methods in Checkpointer
This commit is contained in:
Nuno Campos
2024-06-05 17:14:35 -07:00
parent 052284f21d
commit d366bacf03
11 changed files with 154 additions and 348 deletions
+12 -4
View File
@@ -51,18 +51,26 @@ class TestAsyncSqliteSaver:
query_4: CheckpointMetadata = {"source": "update", "step": 1} # no match
async with self.sqlite_saver as sqlite_saver:
search_results_1 = [c async for c in sqlite_saver.asearch(query_1)]
search_results_1 = [
c async for c in sqlite_saver.alist(None, filter=query_1)
]
assert len(search_results_1) == 1
assert search_results_1[0].metadata == self.metadata_1
search_results_2 = [c async for c in sqlite_saver.asearch(query_2)]
search_results_2 = [
c async for c in sqlite_saver.alist(None, filter=query_2)
]
assert len(search_results_2) == 1
assert search_results_2[0].metadata == self.metadata_2
search_results_3 = [c async for c in sqlite_saver.asearch(query_3)]
search_results_3 = [
c async for c in sqlite_saver.alist(None, filter=query_3)
]
assert len(search_results_3) == 2
search_results_4 = [c async for c in sqlite_saver.asearch(query_4)]
search_results_4 = [
c async for c in sqlite_saver.alist(None, filter=query_4)
]
assert len(search_results_4) == 0
# TODO: test before and limit params
+16 -8
View File
@@ -50,18 +50,18 @@ class TestMemorySaver:
query_3: CheckpointMetadata = {} # search by no keys, return all checkpoints
query_4: CheckpointMetadata = {"source": "update", "step": 1} # no match
search_results_1 = list(self.memory_saver.search(query_1))
search_results_1 = list(self.memory_saver.list(None, filter=query_1))
assert len(search_results_1) == 1
assert search_results_1[0].metadata == self.metadata_1
search_results_2 = list(self.memory_saver.search(query_2))
search_results_2 = list(self.memory_saver.list(None, filter=query_2))
assert len(search_results_2) == 1
assert search_results_2[0].metadata == self.metadata_2
search_results_3 = list(self.memory_saver.search(query_3))
search_results_3 = list(self.memory_saver.list(None, filter=query_3))
assert len(search_results_3) == 2
search_results_4 = list(self.memory_saver.search(query_4))
search_results_4 = list(self.memory_saver.list(None, filter=query_4))
assert len(search_results_4) == 0
# TODO: test before and limit params
@@ -81,16 +81,24 @@ 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)]
search_results_1 = [
c async for c in self.memory_saver.alist(None, filter=query_1)
]
assert len(search_results_1) == 1
assert search_results_1[0].metadata == self.metadata_1
search_results_2 = [c async for c in self.memory_saver.asearch(query_2)]
search_results_2 = [
c async for c in self.memory_saver.alist(None, filter=query_2)
]
assert len(search_results_2) == 1
assert search_results_2[0].metadata == self.metadata_2
search_results_3 = [c async for c in self.memory_saver.asearch(query_3)]
search_results_3 = [
c async for c in self.memory_saver.alist(None, filter=query_3)
]
assert len(search_results_3) == 2
search_results_4 = [c async for c in self.memory_saver.asearch(query_4)]
search_results_4 = [
c async for c in self.memory_saver.alist(None, filter=query_4)
]
assert len(search_results_4) == 0
+23 -13
View File
@@ -56,40 +56,50 @@ class TestSqliteSaver:
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))
search_results_1 = list(self.sqlite_saver.list(None, filter=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))
search_results_2 = list(self.sqlite_saver.list(None, filter=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))
search_results_3 = list(self.sqlite_saver.list(None, filter=query_3))
assert len(search_results_3) == 2
search_results_4 = list(self.sqlite_saver.search(query_4))
search_results_4 = list(self.sqlite_saver.list(None, filter=query_4))
assert len(search_results_4) == 0
# TODO: test before and limit params
def test_search_where(self):
# call method / assertions
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 = "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(None, self.metadata_1, self.config_1) == (
expected_predicate_1,
expected_param_values_1,
)
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_predicate_1 = [
"json_extract(CAST(metadata AS TEXT), '$.source') = ?",
"json_extract(CAST(metadata AS TEXT), '$.step') = ?",
"json_extract(CAST(metadata AS TEXT), '$.writes') = ?",
"json_extract(CAST(metadata AS TEXT), '$.score') = ?",
]
expected_predicate_2 = [
"json_extract(CAST(metadata AS TEXT), '$.source') = ?",
"json_extract(CAST(metadata AS TEXT), '$.step') = ?",
"json_extract(CAST(metadata AS TEXT), '$.writes') = ?",
"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 = ()
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,
+4 -4
View File
@@ -2870,7 +2870,7 @@ Some examples of past conversations:
metadata = chkpnt_tuple_1.metadata
# not needed in application code, only for testing
hiscored = list(saver.search({"score": 1}))
hiscored = list(saver.list(None, filter={"score": 1}))
assert hiscored == []
# mark as "good"
@@ -2878,7 +2878,7 @@ Some examples of past conversations:
saver.put(config, checkpoint, metadata)
# not needed in application code, only for testing
hiscored = list(saver.search({"score": 1}))
hiscored = list(saver.list(None, filter={"score": 1}))
assert len(hiscored) == 1
assert hiscored[0].checkpoint["channel_values"]["messages"] == first_messages
@@ -2921,14 +2921,14 @@ Some examples of past conversations:
metadata = chkpnt_tuple_2.metadata
# not needed in application code, only for testing
hiscored = list(saver.search({"score": 1}))
hiscored = list(saver.list(None, filter={"score": 1}))
assert len(hiscored) == 1
# mark as "good"
metadata["score"] = 1
saver.put(config, checkpoint, metadata)
hiscored = list(saver.search({"score": 1}))
hiscored = list(saver.list(None, filter={"score": 1}))
assert len(hiscored) == 2
assert app.invoke(
+2 -2
View File
@@ -2643,14 +2643,14 @@ Some examples of past conversations:
metadata = chkpnt_tuple_1.metadata
# not needed in application code, only for testing
assert [c async for c in saver.asearch({"score": 1})] == []
assert [c async for c in saver.alist(None, filter={"score": 1})] == []
# mark as "good"
metadata["score"] = 1
await saver.aput(config, checkpoint, metadata)
# not needed in application code, only for testing
hiscored = [c async for c in saver.asearch({"score": 1})]
hiscored = [c async for c in saver.alist(None, filter={"score": 1})]
assert len(hiscored) == 1
assert hiscored[0].checkpoint["channel_values"]["messages"] == first_messages