Lint and fix

This commit is contained in:
William Fu-Hinthorn
2024-12-06 10:05:29 -08:00
parent 4406828c39
commit 054ca2f78d
5 changed files with 215 additions and 129 deletions
@@ -21,8 +21,6 @@ from langgraph.store.base import (
)
from langgraph.store.base.batch import (
BatchedBaseStore,
AsyncBatchedBaseStore,
SyncBatchedBaseStore,
)
from langgraph.store.postgres.base import (
_PLACEHOLDER,
@@ -40,7 +38,7 @@ from langgraph.store.postgres.base import (
logger = logging.getLogger(__name__)
class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Conn]):
class AsyncPostgresStore(BatchedBaseStore, BasePostgresStore[_ainternal.Conn]):
"""Asynchronous Postgres-backed store with optional vector search using pgvector.
!!! example "Examples"
@@ -160,15 +158,8 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
return results
# def batch(self, ops: Iterable[Op]) -> list[Result]:
# return asyncio.run_coroutine_threadsafe(self.abatch(ops), self.loop).result()
def batch(self, ops: Iterable[Op]) -> list[Result]:
futures = []
for op in ops:
fut = self._loop.create_future()
self._aqueue[fut] = op
futures.append(fut)
return [fut.result() for fut in asyncio.as_completed(futures)]
return asyncio.run_coroutine_threadsafe(self.abatch(ops), self.loop).result()
@classmethod
@asynccontextmanager
@@ -30,7 +30,6 @@ from typing_extensions import TypedDict
from langgraph.checkpoint.postgres import _ainternal as _ainternal
from langgraph.checkpoint.postgres import _internal as _pg_internal
from langgraph.store.base import (
BaseStore,
GetOp,
IndexConfig,
Item,
@@ -44,6 +43,7 @@ from langgraph.store.base import (
get_text_at_path,
tokenize_path,
)
from langgraph.store.base.batch import SyncBatchedBaseStore
if TYPE_CHECKING:
from langchain_core.embeddings import Embeddings
@@ -533,7 +533,7 @@ class BasePostgresStore(Generic[C]):
raise ValueError(f"Unsupported operator: {op}")
class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
class PostgresStore(SyncBatchedBaseStore, BasePostgresStore[_pg_internal.Conn]):
"""Postgres-backed store with optional vector search using pgvector.
!!! example "Examples"
@@ -4,6 +4,7 @@ import itertools
import sys
import uuid
from collections.abc import AsyncIterator
from concurrent.futures import ThreadPoolExecutor
from contextlib import asynccontextmanager
from typing import Any, Optional
@@ -64,104 +65,94 @@ async def store(request) -> AsyncIterator[AsyncPostgresStore]:
await conn.execute(f"DROP DATABASE {database}")
async def test_large_batches(store: AsyncPostgresStore) -> None:
N = 100
def test_large_batches(store: AsyncPostgresStore) -> None:
N = 1000
M = 10
coros = []
from concurrent.futures import ThreadPoolExecutor
with ThreadPoolExecutor(max_workers=10) as executor:
for m in range(M):
ops = []
for i in range(N):
for i in range(N):
coros.append(
executor.submit(
store.put,
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
value={"foo": "bar" + str(i)},
)
)
coros.append(
executor.submit(
store.get,
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
)
)
coros.append(
executor.submit(
store.list_namespaces,
prefix=None,
max_depth=m + 1,
)
)
coros.append(
executor.submit(
store.search,
("test",),
)
)
coros.append(
executor.submit(
store.put,
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
value={"foo": "bar" + str(i)},
)
)
coros.append(
executor.submit(
store.put,
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
None,
)
)
# ops.extend(
# [
# PutOp(
# namespace=("test", "foo", "bar", "baz", str(m % 2)),
# key=f"key{i}", # {m}",
# value=None,
# ),
# GetOp(namespace=("test",), key=f"key{i}{m}"),
# ListNamespacesOp(
# match_conditions=None, max_depth=i + 1, limit=m, offset=0
# ),
# SearchOp(
# namespace_prefix=("test",),
# filter=None,
# limit=10,
# offset=0,
# ),
# ]
#
# ops.extend(
# [
# # PutOp(
# # namespace=("test", "foo", "bar", "baz", str(m % 2)),
# # key=f"key{i}", # {m}",
# # value={"data": f"value{i}{m}"},
# # ),
# # GetOp(namespace=("test",), key=f"key{i}{m}"),
# # ListNamespacesOp(
# # match_conditions=None, max_depth=i + 1, limit=m, offset=0
# # ),
# # SearchOp(
# # namespace_prefix=("test",),
# # filter={"data": f"value{i}{m}"},
# # limit=10,
# # offset=0,
# # ),
# ]
# )
_ = [
executor.submit(
store.put,
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
value={"foo": "bar" + str(i)},
),
executor.submit(
store.get,
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
),
executor.submit(
store.list_namespaces,
prefix=None,
max_depth=m + 1,
),
executor.submit(
store.search,
("test",),
),
executor.submit(
store.put,
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
value={"foo": "bar" + str(i)},
),
executor.submit(
store.put,
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
None,
),
]
# coros.extend(ops)
# executor.map(store.batch, [[op] for op in coros])
# await asyncio.gather(*coros) # *[store.abatch(ops) for ops in coros])
async def test_large_batches_async(store: AsyncPostgresStore) -> None:
N = 1000
M = 10
coros = []
for m in range(M):
for i in range(N):
coros.append(
store.aput(
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
value={"foo": "bar" + str(i)},
)
)
coros.append(
store.aget(
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
)
)
coros.append(
store.alist_namespaces(
prefix=None,
max_depth=m + 1,
)
)
coros.append(
store.asearch(
("test",),
)
)
coros.append(
store.aput(
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
value={"foo": "bar" + str(i)},
)
)
coros.append(
store.adelete(
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
)
)
await asyncio.gather(*coros)
async def test_abatch_order(store: AsyncPostgresStore) -> None:
+93 -5
View File
@@ -1,5 +1,7 @@
# type: ignore
import asyncio
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from typing import Any, Optional
from uuid import uuid4
@@ -17,11 +19,7 @@ from langgraph.store.base import (
SearchOp,
)
from langgraph.store.postgres import PostgresStore
from tests.conftest import (
DEFAULT_URI,
VECTOR_TYPES,
CharacterEmbeddings,
)
from tests.conftest import DEFAULT_URI, VECTOR_TYPES, CharacterEmbeddings
@pytest.fixture(scope="function", params=["default", "pipe", "pool"])
@@ -59,6 +57,96 @@ def store(request) -> PostgresStore:
conn.execute(f"DROP DATABASE {database}")
def test_large_batches(store: PostgresStore) -> None:
N = 1000
M = 10
with ThreadPoolExecutor(max_workers=10) as executor:
for m in range(M):
for i in range(N):
_ = [
executor.submit(
store.put,
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
value={"foo": "bar" + str(i)},
),
executor.submit(
store.get,
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
),
executor.submit(
store.list_namespaces,
prefix=None,
max_depth=m + 1,
),
executor.submit(
store.search,
("test",),
),
executor.submit(
store.put,
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
value={"foo": "bar" + str(i)},
),
executor.submit(
store.put,
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
None,
),
]
async def test_large_batches_async(store: PostgresStore) -> None:
N = 1000
M = 10
coros = []
for m in range(M):
for i in range(N):
coros.append(
store.aput(
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
value={"foo": "bar" + str(i)},
)
)
coros.append(
store.aget(
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
)
)
coros.append(
store.alist_namespaces(
prefix=None,
max_depth=m + 1,
)
)
coros.append(
store.asearch(
("test",),
)
)
coros.append(
store.aput(
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
value={"foo": "bar" + str(i)},
)
)
coros.append(
store.adelete(
("test", "foo", "bar", "baz", str(m % 2)),
f"key{i}",
)
)
await asyncio.gather(*coros)
def test_batch_order(store: PostgresStore) -> None:
# Setup test data
store.put(("test", "foo"), "key1", {"data": "value1"})
+36 -20
View File
@@ -14,6 +14,7 @@ from langgraph.store.base import (
NamespacePath,
Op,
PutOp,
Result,
SearchItem,
SearchOp,
_validate_namespace,
@@ -23,6 +24,10 @@ from langgraph.store.base import (
class AsyncBatchedBaseStoreMixin:
"""Efficiently batch operations in a background task."""
_loop: asyncio.AbstractEventLoop
_aqueue: dict[asyncio.Future, Op]
_task: asyncio.Task
async def aget(
self,
namespace: tuple[str, ...],
@@ -92,14 +97,6 @@ class AsyncBatchedBaseStoreMixin:
self._aqueue[fut] = op
return await fut
# def batch(self, ops: Iterable[Op]) -> list[Result]:
# futures = []
# for op in ops:
# fut = self._loop.create_future()
# self._aqueue[fut] = op
# futures.append(fut)
# return [fut.result() for fut in asyncio.as_completed(futures)]
class AsyncBatchedBaseStore(AsyncBatchedBaseStoreMixin, BaseStore):
"""Efficiently batch operations in a background task."""
@@ -115,6 +112,14 @@ class AsyncBatchedBaseStore(AsyncBatchedBaseStoreMixin, BaseStore):
def __del__(self) -> None:
self._task.cancel()
def batch(self, ops: Iterable[Op]) -> list[Result]:
futures = []
for op in ops:
fut = self._loop.create_future()
self._aqueue[fut] = op
futures.append(fut)
return [fut.result() for fut in asyncio.as_completed(futures)]
def _dedupe_ops(values: list[Op]) -> tuple[Optional[list[int]], list[Op]]:
"""Dedupe operations while preserving order for results.
@@ -195,13 +200,16 @@ async def _run(
class SyncBatchedBaseStoreMixin(BaseStore):
"""Efficiently batch operations in a background thread."""
_sync_queue: dict[Future, Op]
_sync_thread: threading.Thread
def get(
self,
namespace: tuple[str, ...],
key: str,
) -> Optional[Item]:
fut = Future()
self._queue[fut] = GetOp(namespace, key)
fut: Future[Optional[Item]] = Future()
self._sync_queue[fut] = GetOp(namespace, key)
return fut.result()
def search(
@@ -214,8 +222,8 @@ class SyncBatchedBaseStoreMixin(BaseStore):
limit: int = 10,
offset: int = 0,
) -> list[SearchItem]:
fut = Future()
self._queue[fut] = SearchOp(namespace_prefix, filter, limit, offset, query)
fut: Future[list[SearchItem]] = Future()
self._sync_queue[fut] = SearchOp(namespace_prefix, filter, limit, offset, query)
return fut.result()
def put(
@@ -226,8 +234,8 @@ class SyncBatchedBaseStoreMixin(BaseStore):
index: Optional[Union[Literal[False], list[str]]] = None,
) -> None:
_validate_namespace(namespace)
fut = Future()
self._queue[fut] = PutOp(namespace, key, value, index)
fut: Future[None] = Future()
self._sync_queue[fut] = PutOp(namespace, key, value, index)
return fut.result()
def delete(
@@ -235,8 +243,8 @@ class SyncBatchedBaseStoreMixin(BaseStore):
namespace: tuple[str, ...],
key: str,
) -> None:
fut = Future()
self._queue[fut] = PutOp(namespace, key, None)
fut: Future[None] = Future()
self._sync_queue[fut] = PutOp(namespace, key, None)
return fut.result()
def list_namespaces(
@@ -248,7 +256,7 @@ class SyncBatchedBaseStoreMixin(BaseStore):
limit: int = 100,
offset: int = 0,
) -> list[tuple[str, ...]]:
fut = Future()
fut: Future[list[tuple[str, ...]]] = Future()
match_conditions = []
if prefix:
match_conditions.append(MatchCondition(match_type="prefix", path=prefix))
@@ -261,7 +269,7 @@ class SyncBatchedBaseStoreMixin(BaseStore):
limit=limit,
offset=offset,
)
self._queue[fut] = op
self._sync_queue[fut] = op
return fut.result()
@@ -283,10 +291,18 @@ class SyncBatchedBaseStore(SyncBatchedBaseStoreMixin, BaseStore):
def __del__(self) -> None:
# Signal the thread to stop
if self._sync_thread.is_alive():
empty_future = Future()
empty_future: Future = Future()
self._sync_queue[empty_future] = None # type: ignore
self._sync_thread.join()
async def abatch(self, ops: Iterable[Op]) -> list[Result]:
futures = []
for op in ops:
fut: Future[Result] = Future()
self._sync_queue[fut] = op
futures.append(fut)
return [fut.result() for fut in futures]
class BatchedBaseStore(
AsyncBatchedBaseStoreMixin, SyncBatchedBaseStoreMixin, BaseStore
@@ -317,7 +333,7 @@ class BatchedBaseStore(
def __del__(self) -> None:
# Signal the thread to stop
if self._sync_thread.is_alive():
empty_future = Future()
empty_future: Future[None] = Future()
self._sync_queue[empty_future] = None # type: ignore
self._sync_thread.join()