Bump SDK Py (#1935)

This commit is contained in:
William FH
2024-10-01 08:17:51 +00:00
committed by GitHub
parent 53c2e4d8c2
commit dd88ac6224
5 changed files with 101 additions and 20 deletions
@@ -21,6 +21,7 @@ from langgraph.store.base import GetOp, ListNamespacesOp, Op, PutOp, Result, Sea
from langgraph.store.postgres.base import (
BasePostgresStore,
Row,
_decode_ns_bytes,
_group_ops,
_row_to_item,
)
@@ -127,17 +128,19 @@ class AsyncPostgresStore(BasePostgresStore[AsyncConnection]):
results: list[Result],
) -> None:
queries = self._get_batch_search_queries(search_ops)
cursors: list[tuple[AsyncCursor[Any], int, SearchOp]] = []
cursors: list[tuple[AsyncCursor[Any], int]] = []
for (query, params), (idx, op) in zip(queries, search_ops):
for (query, params), (idx, _) in zip(queries, search_ops):
cur = self.conn.cursor(binary=True)
await cur.execute(query, params)
cursors.append((cur, idx, op))
cursors.append((cur, idx))
for cur, idx, op in cursors:
for cur, idx in cursors:
rows = cast(list[Row], await cur.fetchall())
items = [
_row_to_item(op.namespace_prefix, row, loader=self._deserializer)
_row_to_item(
_decode_ns_bytes(row["prefix"]), row, loader=self._deserializer
)
for row in rows
]
results[idx] = items
@@ -156,9 +159,7 @@ class AsyncPostgresStore(BasePostgresStore[AsyncConnection]):
for cur, idx in cursors:
rows = cast(list[dict], await cur.fetchall())
namespaces = [
tuple(row["truncated_prefix"].decode()[1:].split(".")) for row in rows
]
namespaces = [_decode_ns_bytes(row["truncated_prefix"]) for row in rows]
results[idx] = namespaces
@classmethod
@@ -152,7 +152,7 @@ class BasePostgresStore(BaseStore, Generic[C]):
queries: list[tuple[str, Sequence]] = []
for _, op in search_ops:
query = """
SELECT key, value, created_at, updated_at
SELECT prefix, key, value, created_at, updated_at, prefix
FROM store
WHERE prefix <@ %s
"""
@@ -297,17 +297,19 @@ class PostgresStore(BasePostgresStore[Connection]):
results: list[Result],
) -> None:
queries = self._get_batch_search_queries(search_ops)
cursors: list[tuple[Cursor[Any], int, SearchOp]] = []
cursors: list[tuple[Cursor[Any], int]] = []
for (query, params), (idx, op) in zip(queries, search_ops):
for (query, params), (idx, _) in zip(queries, search_ops):
cur = self.conn.cursor(binary=True)
cur.execute(query, params)
cursors.append((cur, idx, op))
cursors.append((cur, idx))
for cur, idx, op in cursors:
for cur, idx in cursors:
rows = cast(list[Row], cur.fetchall())
items = [
_row_to_item(op.namespace_prefix, row, loader=self._deserializer)
_row_to_item(
_decode_ns_bytes(row["prefix"]), row, loader=self._deserializer
)
for row in rows
]
results[idx] = items
@@ -326,9 +328,7 @@ class PostgresStore(BasePostgresStore[Connection]):
for cur, idx in cursors:
rows = cast(list[dict], cur.fetchall())
namespaces = [
tuple(row["truncated_prefix"].decode()[1:].split(".")) for row in rows
]
namespaces = [_decode_ns_bytes(row["truncated_prefix"]) for row in rows]
results[idx] = namespaces
@classmethod
@@ -433,3 +433,9 @@ def _json_loads(content: Union[bytes, orjson.Fragment]) -> Any:
else:
content = content.contents.encode()
return orjson.loads(cast(bytes, content))
def _decode_ns_bytes(namespace: Union[str, bytes]) -> tuple[str, ...]:
if isinstance(namespace, bytes):
namespace = namespace.decode()[1:]
return tuple(namespace.split("."))
@@ -45,12 +45,14 @@ async def test_abatch_order(store: AsyncPostgresStore) -> None:
"value": '{"data": "value1"}',
"created_at": datetime.now(),
"updated_at": datetime.now(),
"prefix": "test.foo",
},
{
"key": "key2",
"value": '{"data": "value2"}',
"created_at": datetime.now(),
"updated_at": datetime.now(),
"prefix": "test.bar",
},
]
)
@@ -61,6 +63,7 @@ async def test_abatch_order(store: AsyncPostgresStore) -> None:
"value": '{"data": "value1"}',
"created_at": datetime.now(),
"updated_at": datetime.now(),
"prefix": "test.foo",
},
]
)
@@ -151,12 +154,14 @@ async def test_batch_get_ops(store: AsyncPostgresStore) -> None:
"value": '{"data": "value1"}',
"created_at": datetime.now(),
"updated_at": datetime.now(),
"prefix": "test.foo",
},
{
"key": "key2",
"value": '{"data": "value2"}',
"created_at": datetime.now(),
"updated_at": datetime.now(),
"prefix": "test.bar",
},
]
)
@@ -205,12 +210,14 @@ async def test_batch_search_ops(store: AsyncPostgresStore) -> None:
"value": '{"data": "value1"}',
"created_at": datetime.now(),
"updated_at": datetime.now(),
"prefix": "test.foo",
},
{
"key": "key2",
"value": '{"data": "value2"}',
"created_at": datetime.now(),
"updated_at": datetime.now(),
"prefix": "test.bar",
},
]
)
@@ -422,13 +429,26 @@ class TestAsyncPostgresStore:
{"title": "Report A", "author": "John Doe", "tags": ["final"]},
{"title": "Report B", "author": "Alice Johnson", "tags": ["draft"]},
]
empty = await store.asearch(
(
"scoped",
"assistant_id",
"shared",
"6c5356f6-63ab-4158-868d-cd9fd14c736e",
),
limit=10,
offset=0,
)
assert len(empty) == 0
for namespace, item in zip(test_namespaces, test_items):
await store.aput(namespace, f"item_{namespace[-1]}", item)
docs_result = await store.asearch(["test_search", "documents"])
assert len(docs_result) == 2
assert all(item.namespace[1] == "documents" for item in docs_result)
assert all([item.namespace[1] == "documents" for item in docs_result]), [
item.namespace for item in docs_result
]
reports_result = await store.asearch(["test_search", "reports"])
assert len(reports_result) == 2
@@ -460,6 +480,51 @@ class TestAsyncPostgresStore:
all_items = page1 + page2
assert len(all_items) == 4
assert len(set(item.key for item in all_items)) == 4
empty = await store.asearch(
(
"scoped",
"assistant_id",
"shared",
"again",
"maybe",
"some-long",
"6be5cb0e-2eb4-42e6-bb6b-fba3c269db25",
),
limit=10,
offset=0,
)
assert len(empty) == 0
# Test with a namespace beginning with a number (like a UUID)
uuid_namespace = (str(uuid.uuid4()), "documents")
uuid_item_id = "uuid_doc"
uuid_item_value = {
"title": "UUID Document",
"content": "This document has a UUID namespace.",
}
# Insert the item with the UUID namespace
await store.aput(uuid_namespace, uuid_item_id, uuid_item_value)
# Retrieve the item to verify it was stored correctly
retrieved_item = await store.aget(uuid_namespace, uuid_item_id)
assert retrieved_item is not None
assert retrieved_item.namespace == uuid_namespace
assert retrieved_item.key == uuid_item_id
assert retrieved_item.value == uuid_item_value
# Search for the item using the UUID namespace
search_result = await store.asearch([uuid_namespace[0]])
assert len(search_result) == 1
assert search_result[0].key == uuid_item_id
assert search_result[0].value == uuid_item_value
# Clean up: delete the item with the UUID namespace
await store.adelete(uuid_namespace, uuid_item_id)
# Verify the item was deleted
deleted_item = await store.aget(uuid_namespace, uuid_item_id)
assert deleted_item is None
for namespace in test_namespaces:
await store.adelete(namespace, f"item_{namespace[-1]}")
+10 -1
View File
@@ -43,12 +43,14 @@ def test_batch_order(store: PostgresStore) -> None:
"value": '{"data": "value1"}',
"created_at": datetime.now(),
"updated_at": datetime.now(),
"prefix": "test.foo",
},
{
"key": "key2",
"value": '{"data": "value2"}',
"created_at": datetime.now(),
"updated_at": datetime.now(),
"prefix": "test.bar",
},
]
)
@@ -59,6 +61,7 @@ def test_batch_order(store: PostgresStore) -> None:
"value": '{"data": "value1"}',
"created_at": datetime.now(),
"updated_at": datetime.now(),
"prefix": "test.foo",
},
]
)
@@ -149,12 +152,14 @@ def test_batch_get_ops(store: PostgresStore) -> None:
"value": '{"data": "value1"}',
"created_at": datetime.now(),
"updated_at": datetime.now(),
"prefix": "test.foo",
},
{
"key": "key2",
"value": '{"data": "value2"}',
"created_at": datetime.now(),
"updated_at": datetime.now(),
"prefix": "test.bar",
},
]
)
@@ -203,12 +208,14 @@ def test_batch_search_ops(store: PostgresStore) -> None:
"value": '{"data": "value1"}',
"created_at": datetime.now(),
"updated_at": datetime.now(),
"prefix": "test.foo",
},
{
"key": "key2",
"value": '{"data": "value2"}',
"created_at": datetime.now(),
"updated_at": datetime.now(),
"prefix": "test.bar",
},
]
)
@@ -423,7 +430,9 @@ class TestPostgresStore:
docs_result = store.search(["test_search", "documents"])
assert len(docs_result) == 2
assert all(item.namespace[1] == "documents" for item in docs_result)
assert all(
[item.namespace[1] == "documents" for item in docs_result]
), docs_result
reports_result = store.search(["test_search", "reports"])
assert len(reports_result) == 2
+1 -1
View File
@@ -1,6 +1,6 @@
[tool.poetry]
name = "langgraph-sdk"
version = "0.1.31"
version = "0.1.32"
description = "SDK for interacting with LangGraph API"
authors = []
license = "MIT"