From dd88ac622485b9a73155062ad81dacafbecd5fdb Mon Sep 17 00:00:00 2001 From: William FH <13333726+hinthornw@users.noreply.github.com> Date: Tue, 1 Oct 2024 01:17:51 -0700 Subject: [PATCH] Bump SDK Py (#1935) --- .../langgraph/store/postgres/aio.py | 17 ++--- .../langgraph/store/postgres/base.py | 24 ++++--- .../tests/test_async_store.py | 67 ++++++++++++++++++- libs/checkpoint-postgres/tests/test_store.py | 11 ++- libs/sdk-py/pyproject.toml | 2 +- 5 files changed, 101 insertions(+), 20 deletions(-) diff --git a/libs/checkpoint-postgres/langgraph/store/postgres/aio.py b/libs/checkpoint-postgres/langgraph/store/postgres/aio.py index ef73482ea..981cc91b4 100644 --- a/libs/checkpoint-postgres/langgraph/store/postgres/aio.py +++ b/libs/checkpoint-postgres/langgraph/store/postgres/aio.py @@ -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 diff --git a/libs/checkpoint-postgres/langgraph/store/postgres/base.py b/libs/checkpoint-postgres/langgraph/store/postgres/base.py index ba4a39d1f..f5ad3408d 100644 --- a/libs/checkpoint-postgres/langgraph/store/postgres/base.py +++ b/libs/checkpoint-postgres/langgraph/store/postgres/base.py @@ -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(".")) diff --git a/libs/checkpoint-postgres/tests/test_async_store.py b/libs/checkpoint-postgres/tests/test_async_store.py index 244ad299e..ad6504312 100644 --- a/libs/checkpoint-postgres/tests/test_async_store.py +++ b/libs/checkpoint-postgres/tests/test_async_store.py @@ -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]}") diff --git a/libs/checkpoint-postgres/tests/test_store.py b/libs/checkpoint-postgres/tests/test_store.py index 47ffce88f..bc952a862 100644 --- a/libs/checkpoint-postgres/tests/test_store.py +++ b/libs/checkpoint-postgres/tests/test_store.py @@ -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 diff --git a/libs/sdk-py/pyproject.toml b/libs/sdk-py/pyproject.toml index df727e781..9d7750314 100644 --- a/libs/sdk-py/pyproject.toml +++ b/libs/sdk-py/pyproject.toml @@ -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"