diff --git a/pyproject.toml b/pyproject.toml index bc8f8df..bba1b50 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -42,6 +42,7 @@ sentence-transformers = ["sentence-transformers>=3.0"] # bedrock needs nothing beyond boto3. # --- Optional capabilities --- +async = ["aioboto3>=13"] rerank = ["sentence-transformers>=3.0"] # cross-encoder reranking redis = ["redis>=5.0"] # RedisCache / AWS ElastiCache mcp = ["mcp>=1.0; python_version >= '3.10'"] # MCP client ingestion @@ -87,7 +88,8 @@ dev = [ "pytest-asyncio>=0.23", "pytest-cov>=5.0", "hypothesis>=6.100", - "moto[dynamodb]>=5.0", + "aioboto3>=13", + "moto[dynamodb,server]>=5.0", "mypy>=1.11,<2.0", "ruff>=0.6", "pre-commit>=3.5", diff --git a/src/dynavec/__init__.py b/src/dynavec/__init__.py index a0aaec8..0d025a7 100644 --- a/src/dynavec/__init__.py +++ b/src/dynavec/__init__.py @@ -20,6 +20,7 @@ from __future__ import annotations +from .async_client import AsyncDynavec from .bm25 import BM25Index from .cache import BaseCache, DynamoDBCache, RedisCache, SemanticCache, warm_cache from .client import Dynavec @@ -91,6 +92,7 @@ __all__ = [ "Dynavec", + "AsyncDynavec", "DynavecConfig", "AWSCredentials", "Document", diff --git a/src/dynavec/async_client.py b/src/dynavec/async_client.py new file mode 100644 index 0000000..f901d45 --- /dev/null +++ b/src/dynavec/async_client.py @@ -0,0 +1,448 @@ +"""Async Dynavec client for high-concurrency workloads.""" + +from __future__ import annotations + +import asyncio +from collections.abc import AsyncIterator, Iterable, Sequence +from contextlib import AsyncExitStack +from types import TracebackType +from typing import TYPE_CHECKING, Any, Union + +if TYPE_CHECKING: + from .cache import BaseCache + +from .client_common import ( + DDBPayload, + HotPayload, + S3Payload, + apply_transform_pipeline, + assign_embeddings, + build_write_payloads, + documents_to_embed, + embedding_texts, + split_key, +) +from .config import DynavecConfig +from .credentials import AWSCredentials, resolve_async_session +from .embeddings.base import Embedder +from .exceptions import ConfigurationError, DimensionMismatchError +from .hot import HotTier +from .metadata import build_s3_filter +from .models import Document, SearchResult, UpsertResult +from .retrieval import distance_to_score +from .stores.async_dynamodb import AsyncDynamoDBStore +from .stores.async_s3vectors import AsyncS3VectorsStore +from .telemetry import TelemetryRecorder +from .transforms import Transform, TransformPipeline, as_pipeline + +TransformSpec = Union[TransformPipeline, Transform, Iterable[Transform]] +Metadata = dict[str, Any] + + +class AsyncDynavec: + """Async Dynavec client backed by aioboto3.""" + + def __init__( + self, + config: DynavecConfig, + embedder: Embedder | None = None, + *, + credentials: AWSCredentials | None = None, + boto_session: Any | None = None, + transform: TransformSpec | None = None, + cache: BaseCache | None = None, + telemetry: TelemetryRecorder | None = None, + ) -> None: + self.config = config + self.embedder = embedder + self._session = resolve_async_session( + credentials, + boto_session, + ) + + self._vectors = AsyncS3VectorsStore( + config, + boto_session=self._session, + ) + self._docs = AsyncDynamoDBStore( + config, + boto_session=self._session, + ) + + self._default_transform = as_pipeline(transform) + self._cache = cache + self._telemetry = telemetry + self._hot = ( + HotTier(config) + if config.hot_tier + else None + ) + + self._exit_stack: AsyncExitStack | None = None + + if ( + embedder is not None + and embedder.dimension != config.dimension + ): + raise ConfigurationError( + f"Embedder dimension ({embedder.dimension}) " + f"!= index dimension ({config.dimension}). " + "Fix the embedder or DynavecConfig.dimension." + ) + + async def __aenter__(self) -> AsyncDynavec: + if self._exit_stack is not None: + return self + + stack = AsyncExitStack() + await stack.__aenter__() + + try: + await stack.enter_async_context(self._vectors) + await stack.enter_async_context(self._docs) + except Exception: + await stack.aclose() + raise + + self._exit_stack = stack + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> None: + stack = self._exit_stack + + if stack is None: + return + + try: + await stack.__aexit__( + exc_type, + exc_value, + traceback, + ) + finally: + self._exit_stack = None + + async def aclose(self) -> None: + stack = self._exit_stack + + if stack is None: + return + + try: + await stack.aclose() + finally: + self._exit_stack = None + + def _require_open(self) -> None: + if self._exit_stack is None: + raise RuntimeError( + "AsyncDynavec is not open. " + "Use it inside 'async with' before performing AWS operations." + ) + + def _invalidate_cache(self, namespace: str) -> None: + if ( + self._cache is not None + and self.config.cache_invalidate_on_write + ): + self._cache.invalidate(namespace) + + async def _resolve_query_vector( + self, + query: str | None, + vector: list[float] | None, + ) -> list[float]: + if vector is not None: + if len(vector) != self.config.dimension: + raise DimensionMismatchError( + f"Query vector dimension {len(vector)} " + f"!= {self.config.dimension}." + ) + return vector + + if query is None: + raise ValueError( + "Provide either 'query' text or a 'vector'." + ) + + if self.embedder is None: + raise ConfigurationError( + "Text query requires an embedder. " + "Pass one to AsyncDynavec(...) or query " + "with a precomputed 'vector'." + ) + + return await self.embedder.aembed_query(query) + + async def asearch( + self, + query: str | None = None, + *, + vector: list[float] | None = None, + top_k: int = 10, + namespace: str = "default", + filter: Metadata | None = None, + ) -> list[SearchResult]: + """Search asynchronously using query text or a precomputed vector.""" + self._require_open() + + query_vector = await self._resolve_query_vector( + query, + vector, + ) + + raw = await self._vectors.query( + query_vector=query_vector, + top_k=top_k, + filter=build_s3_filter( + filter, + namespace, + ), + return_metadata=True, + return_distance=True, + ) + + if not raw: + return [] + + hits = [ + ( + split_key(item["key"])[1], + item.get("distance"), + ) + for item in raw + ] + + ids = [doc_id for doc_id, _ in hits] + + hydrated = await self._docs.get_many( + namespace, + ids, + ) + + results: list[SearchResult] = [] + + for doc_id, distance in hits: + doc = hydrated.get(doc_id, {}) + + results.append( + SearchResult( + id=doc_id, + score=( + distance_to_score( + distance, + self.config.distance_metric, + ) + if distance is not None + else 0.0 + ), + distance=distance, + text=doc.get("text"), + metadata=doc.get("metadata", {}), + ttl=doc.get("ttl"), + ) + ) + + return results + + async def asearch_stream( + self, + query: str | None = None, + *, + vector: list[float] | None = None, + top_k: int = 50, + namespace: str = "default", + filter: Metadata | None = None, + page_size: int | None = None, + ) -> AsyncIterator[SearchResult]: + """Stream search results asynchronously page by page.""" + self._require_open() + + query_vector = await self._resolve_query_vector( + query, + vector, + ) + + yielded = 0 + + effective_page_size = ( + page_size + if page_size is not None + else self.config.top_k_page_size + ) + + async for page in self._vectors.query_pages( + query_vector=query_vector, + top_k=top_k, + filter=build_s3_filter( + filter, + namespace, + ), + return_metadata=True, + return_distance=True, + page_size=effective_page_size, + ): + page_hits = [ + ( + split_key(item["key"])[1], + item.get("distance"), + ) + for item in page + ] + + hydrated = await self._docs.get_many( + namespace, + [doc_id for doc_id, _ in page_hits], + ) + + for doc_id, distance in page_hits: + if yielded >= top_k: + return + + doc = hydrated.get(doc_id, {}) + + yield SearchResult( + id=doc_id, + score=( + distance_to_score( + distance, + self.config.distance_metric, + ) + if distance is not None + else 0.0 + ), + distance=distance, + text=doc.get("text"), + metadata=doc.get("metadata", {}), + ) + + yielded += 1 + + async def _prepare( + self, + docs: list[Document], + namespace: str, + auto_metadata: bool, + transform: TransformSpec | None, + default_ttl_seconds: int | None = None, + ) -> tuple[ + list[S3Payload], + list[DDBPayload], + list[str], + list[HotPayload], + ]: + pipeline = ( + as_pipeline(transform) + or self._default_transform + ) + + apply_transform_pipeline( + docs, + namespace, + pipeline, + ) + + to_embed = documents_to_embed(docs) + + if to_embed: + if self.embedder is None: + raise ConfigurationError( + "Some documents have no vector and no embedder " + "is configured. Pass an embedder to " + "AsyncDynavec(...) or provide precomputed vectors." + ) + + texts = embedding_texts(to_embed) + + vectors = await self.embedder.aembed_documents( + texts + ) + + assign_embeddings( + docs, + to_embed, + vectors, + ) + + return build_write_payloads( + docs, + self.config, + namespace, + auto_metadata, + default_ttl_seconds=default_ttl_seconds, + ) + + async def aupsert( + self, + documents: Sequence[ + Document | dict[str, Any] + ] + | None = None, + *, + namespace: str = "default", + auto_metadata: bool = False, + transform: TransformSpec | None = None, + ttl_seconds: int | None = None, + ) -> UpsertResult: + """Insert or overwrite documents asynchronously.""" + if ttl_seconds is not None and ttl_seconds <= 0: + raise ValueError( + f"ttl_seconds must be positive, got {ttl_seconds}." + ) + + if not documents: + return UpsertResult( + count=0, + ids=[], + ) + + self._require_open() + + docs = [ + document + if isinstance(document, Document) + else Document(**document) + for document in documents + ] + + ( + s3_payload, + ddb_payload, + ids, + hot_payload, + ) = await self._prepare( + docs, + namespace, + auto_metadata, + transform, + default_ttl_seconds=ttl_seconds, + ) + + await asyncio.gather( + self._vectors.put_vectors( + s3_payload, + max_workers=self.config.max_workers, + ), + self._docs.put_many( + namespace, + ddb_payload, + ), + ) + + if self._hot is not None: + self._hot.insert_many( + namespace, + hot_payload, + ) + + self._invalidate_cache(namespace) + + return UpsertResult( + count=len(ids), + ids=ids, + ) diff --git a/src/dynavec/client.py b/src/dynavec/client.py index 0a8f989..9d66d64 100644 --- a/src/dynavec/client.py +++ b/src/dynavec/client.py @@ -32,7 +32,7 @@ from functools import partial from pathlib import Path from types import TracebackType -from typing import TYPE_CHECKING, Any, Literal, Optional, TextIO, Union, overload +from typing import TYPE_CHECKING, Any, Literal, TextIO, Union, overload if TYPE_CHECKING: from .cache import BaseCache @@ -40,6 +40,18 @@ import numpy as np +from .client_common import ( + DDBPayload, + HotPayload, + S3Payload, + apply_transform_pipeline, + assign_embeddings, + build_write_payloads, + documents_to_embed, + embedding_texts, + s3_key, + split_key, +) from .config import NS_METADATA_KEY, TEXT_METADATA_KEY, DynavecConfig from .credentials import AWSCredentials, resolve_session from .embeddings.base import Embedder @@ -52,7 +64,7 @@ ) from .graph import GraphStore from .hot import HotTier -from .metadata import build_s3_filter, generate_auto_metadata, split_metadata +from .metadata import build_s3_filter from .metrics import normalize_scores as normalize_metric_scores from .metrics import rescore as metric_rescore from .metrics import score as metric_score @@ -68,10 +80,9 @@ from .provisioning import provision_all from .retrieval import distance_to_score, maximal_marginal_relevance, reciprocal_rank_fusion from .stores import DynamoDBStore, S3VectorsStore -from .stores.dynamodb import check_item_size from .telemetry import TelemetryRecorder -from .transforms import Transform, TransformContext, TransformPipeline, as_pipeline -from .utils import KEY_SEPARATOR, chunked, decode_key_component, encode_key_component +from .transforms import Transform, TransformPipeline, as_pipeline +from .utils import chunked Metadata = dict[str, Any] _S3_PUT_CHUNK = 500 @@ -79,9 +90,6 @@ RescoreSpec = Union[str, dict[str, float]] TransformSpec = Union[TransformPipeline, Transform, Iterable[Transform]] -S3Payload = tuple[str, list[float], Metadata] -DDBPayload = tuple[str, Optional[str], Metadata] -HotPayload = tuple[str, list[float], Optional[str], Metadata] class Dynavec: @@ -185,11 +193,10 @@ def graph(self) -> GraphStore: # ------------------------------------------------------------- key helpers def _s3_key(self, namespace: str, doc_id: str) -> str: - return f"{encode_key_component(namespace)}{KEY_SEPARATOR}{encode_key_component(doc_id)}" + return s3_key(namespace, doc_id) def _split_key(self, key: str) -> tuple[str, str]: - namespace, _, doc_id = key.partition(KEY_SEPARATOR) - return decode_key_component(namespace), decode_key_component(doc_id) + return split_key(key) def _run_parallel(self, tasks: list[Callable[[], None]]) -> None: """Run zero-arg callables; parallel if enabled, else sequential.""" @@ -211,87 +218,49 @@ def _prepare( auto_metadata: bool, transform: TransformSpec | None, default_ttl_seconds: int | None = None, - ) -> tuple[list[S3Payload], list[DDBPayload], list[str], list[HotPayload]]: + ) -> tuple[ + list[S3Payload], + list[DDBPayload], + list[str], + list[HotPayload], + ]: pipeline = as_pipeline(transform) or self._default_transform - # 1) transforms may set/rewrite text, vector, metadata - if pipeline is not None: - for d in docs: - ctx = pipeline( - TransformContext( - id=d.id, - text=d.text, - vector=d.vector, - metadata=dict(d.metadata), - namespace=namespace, - ) - ) - d.text, d.vector, d.metadata = ctx.text, ctx.vector, ctx.metadata + # 1) shared transform logic + apply_transform_pipeline( + docs, + namespace, + pipeline, + ) + + # 2) sync embedding stays specific to Dynavec + to_embed = documents_to_embed(docs) - # 2) embed anything still missing a vector, in one batched call - to_embed = [(i, d.text) for i, d in enumerate(docs) if d.vector is None] if to_embed: if self.embedder is None: raise ConfigurationError( "Some documents have no vector and no embedder is configured. " "Pass an embedder to Dynavec(...) or provide precomputed vectors." ) - optional_texts = [t for _, t in to_embed] - if any(t is None for t in optional_texts): - raise ConfigurationError("A document has neither text nor vector.") - texts = [t for t in optional_texts if t is not None] + + texts = embedding_texts(to_embed) vectors = self.embedder.embed_documents(texts) - for (idx, _), vec in zip(to_embed, vectors): - docs[idx].vector = vec - - # 3) validate + build payloads - s3_payload: list[S3Payload] = [] - ddb_payload: list[DDBPayload] = [] - ids: list[str] = [] - hot_payload: list[HotPayload] = [] - for d in docs: - vector = d.vector - if vector is None: - raise ConfigurationError(f"Embedder did not return a vector for document {d.id!r}.") - if len(vector) != self.config.dimension: - raise DimensionMismatchError( - f"Document {d.id!r} vector has dimension {len(vector)}, " - f"expected {self.config.dimension}." - ) - meta = dict(d.metadata) - if auto_metadata: - auto = generate_auto_metadata(d.text) - auto.update(meta) - meta = auto - s3_meta, ddb_meta = split_metadata(meta, self.config, namespace, d.text) - doc_ttl_seconds = d.ttl_seconds if d.ttl_seconds is not None else default_ttl_seconds - if doc_ttl_seconds is not None and doc_ttl_seconds <= 0: - raise ValueError( - f"Document {d.id!r} ttl_seconds must be positive, got {doc_ttl_seconds}." - ) - ttl_timestamp = ( - int(time.time() + doc_ttl_seconds) if doc_ttl_seconds is not None else None - ) - if ttl_timestamp is not None: - ddb_meta["_ttl"] = ttl_timestamp - # fail before either store is written, not partway through a batch - check_item_size( - namespace, - d.id, - d.text, - ddb_meta, - self.config.gzip_threshold_bytes, - ttl=ttl_timestamp, - ttl_attribute=self.config.dynamodb_ttl_attribute, + assign_embeddings( + docs, + to_embed, + vectors, ) - s3_payload.append((self._s3_key(namespace, d.id), vector, s3_meta)) - ddb_payload.append((d.id, d.text, ddb_meta)) - ids.append(d.id) - # Hot tier keeps the full (merged) metadata + text so warmed - # namespaces need neither an S3 query nor a DynamoDB read. - hot_payload.append((d.id, vector, d.text, meta)) - return s3_payload, ddb_payload, ids, hot_payload + + # 3) shared validation + payload construction + return build_write_payloads( + docs, + self.config, + namespace, + auto_metadata, + default_ttl_seconds=default_ttl_seconds, + ) + def _write( self, diff --git a/src/dynavec/client_common.py b/src/dynavec/client_common.py new file mode 100644 index 0000000..2ac71fb --- /dev/null +++ b/src/dynavec/client_common.py @@ -0,0 +1,203 @@ +"""Pure helpers shared by sync and async Dynavec clients.""" + +from __future__ import annotations + +import time +from typing import Any, Optional + +from .config import DynavecConfig +from .exceptions import ConfigurationError, DimensionMismatchError +from .metadata import generate_auto_metadata, split_metadata +from .models import Document +from .stores.dynamodb import check_item_size +from .transforms import TransformContext, TransformPipeline +from .utils import KEY_SEPARATOR, decode_key_component, encode_key_component + +Metadata = dict[str, Any] +S3Payload = tuple[str, list[float], Metadata] +DDBPayload = tuple[str, Optional[str], Metadata] +HotPayload = tuple[str, list[float], Optional[str], Metadata] +EmbeddingTarget = tuple[int, Optional[str]] + + +def s3_key(namespace: str, doc_id: str) -> str: + return ( + f"{encode_key_component(namespace)}" + f"{KEY_SEPARATOR}" + f"{encode_key_component(doc_id)}" + ) + + +def split_key(key: str) -> tuple[str, str]: + namespace, _, doc_id = key.partition(KEY_SEPARATOR) + return ( + decode_key_component(namespace), + decode_key_component(doc_id), + ) + + +def apply_transform_pipeline( + docs: list[Document], + namespace: str, + pipeline: TransformPipeline | None, +) -> None: + if pipeline is None: + return + + for doc in docs: + context = pipeline( + TransformContext( + id=doc.id, + text=doc.text, + vector=doc.vector, + metadata=dict(doc.metadata), + namespace=namespace, + ) + ) + + doc.text = context.text + doc.vector = context.vector + doc.metadata = context.metadata + + +def documents_to_embed( + docs: list[Document], +) -> list[EmbeddingTarget]: + return [ + (index, doc.text) + for index, doc in enumerate(docs) + if doc.vector is None + ] + + +def embedding_texts( + targets: list[EmbeddingTarget], +) -> list[str]: + optional_texts = [text for _, text in targets] + + if any(text is None for text in optional_texts): + raise ConfigurationError( + "A document has neither text nor vector." + ) + + return [ + text + for text in optional_texts + if text is not None + ] + + +def assign_embeddings( + docs: list[Document], + targets: list[EmbeddingTarget], + vectors: list[list[float]], +) -> None: + for (index, _), vector in zip(targets, vectors): + docs[index].vector = vector + + +def build_write_payloads( + docs: list[Document], + config: DynavecConfig, + namespace: str, + auto_metadata: bool, + default_ttl_seconds: int | None = None, +) -> tuple[ + list[S3Payload], + list[DDBPayload], + list[str], + list[HotPayload], +]: + s3_payload: list[S3Payload] = [] + ddb_payload: list[DDBPayload] = [] + ids: list[str] = [] + hot_payload: list[HotPayload] = [] + + for doc in docs: + vector = doc.vector + + if vector is None: + raise ConfigurationError( + f"Embedder did not return a vector for document " + f"{doc.id!r}." + ) + + if len(vector) != config.dimension: + raise DimensionMismatchError( + f"Document {doc.id!r} vector has dimension " + f"{len(vector)}, expected {config.dimension}." + ) + + metadata = dict(doc.metadata) + + if auto_metadata: + generated = generate_auto_metadata(doc.text) + generated.update(metadata) + metadata = generated + + s3_metadata, ddb_metadata = split_metadata( + metadata, + config, + namespace, + doc.text, + ) + + doc_ttl_seconds = ( + doc.ttl_seconds + if doc.ttl_seconds is not None + else default_ttl_seconds + ) + + if doc_ttl_seconds is not None and doc_ttl_seconds <= 0: + raise ValueError( + f"Document {doc.id!r} ttl_seconds must be positive, " + f"got {doc_ttl_seconds}." + ) + + ttl_timestamp = ( + int(time.time() + doc_ttl_seconds) + if doc_ttl_seconds is not None + else None + ) + + if ttl_timestamp is not None: + ddb_metadata["_ttl"] = ttl_timestamp + + check_item_size( + namespace, + doc.id, + doc.text, + ddb_metadata, + config.gzip_threshold_bytes, + ttl=ttl_timestamp, + ttl_attribute=config.dynamodb_ttl_attribute, + ) + + s3_payload.append( + ( + s3_key(namespace, doc.id), + vector, + s3_metadata, + ) + ) + + ddb_payload.append( + ( + doc.id, + doc.text, + ddb_metadata, + ) + ) + + ids.append(doc.id) + + hot_payload.append( + ( + doc.id, + vector, + doc.text, + metadata, + ) + ) + + return s3_payload, ddb_payload, ids, hot_payload diff --git a/src/dynavec/credentials.py b/src/dynavec/credentials.py index 794e998..ce52051 100644 --- a/src/dynavec/credentials.py +++ b/src/dynavec/credentials.py @@ -12,9 +12,12 @@ from __future__ import annotations +import importlib from dataclasses import dataclass from typing import Any +from .exceptions import MissingDependencyError + @dataclass(frozen=True) class AWSCredentials: @@ -89,3 +92,52 @@ def resolve_session(credentials: AWSCredentials | None, boto_session: Any | None import boto3 return boto3.Session() + + +def resolve_async_session( + credentials: AWSCredentials | None, + boto_session: Any | None, +) -> Any: + """Pick an aioboto3 session: explicit session > credentials > default chain.""" + if boto_session is not None: + return boto_session + + try: + aioboto3: Any = importlib.import_module("aioboto3") + except ImportError as exc: + raise MissingDependencyError("AsyncDynavec", "aioboto3", "async") from exc + + if credentials is None: + return aioboto3.Session() + + if credentials.assume_role_arn: + # Reuse the existing synchronous STS assume-role flow once during setup. + sync_session = credentials.session() + resolved = sync_session.get_credentials() + if resolved is None: + raise RuntimeError("Unable to resolve assumed-role AWS credentials.") + + frozen = resolved.get_frozen_credentials() + return aioboto3.Session( + aws_access_key_id=frozen.access_key, + aws_secret_access_key=frozen.secret_key, + aws_session_token=frozen.token, + region_name=credentials.region, + ) + + kwargs: dict[str, str] = {} + + if credentials.access_key_id and credentials.secret_access_key: + kwargs["aws_access_key_id"] = credentials.access_key_id + kwargs["aws_secret_access_key"] = credentials.secret_access_key + + if credentials.session_token: + kwargs["aws_session_token"] = credentials.session_token + + if credentials.profile_name: + kwargs["profile_name"] = credentials.profile_name + + if credentials.region: + kwargs["region_name"] = credentials.region + + return aioboto3.Session(**kwargs) diff --git a/src/dynavec/stores/async_dynamodb.py b/src/dynavec/stores/async_dynamodb.py new file mode 100644 index 0000000..496cdcb --- /dev/null +++ b/src/dynavec/stores/async_dynamodb.py @@ -0,0 +1,260 @@ +"""Async DynamoDB document / metadata store backed by aioboto3.""" + +from __future__ import annotations + +import logging +import time +from collections.abc import Sequence +from types import TracebackType +from typing import Any + +from ..config import DynavecConfig +from ..logging import log_store_event +from ..utils import async_retry +from .dynamodb import ( + _BATCH_GET_LIMIT, + Metadata, + _build_item, + _check_built_item, + _from_dynamo, + _pk, + _read_text, +) + + +class AsyncDynamoDBStore: + """Async wrapper around DynamoDB document and metadata I/O.""" + + _logger = logging.getLogger("dynavec.stores.async_dynamodb") + + def __init__(self, config: DynavecConfig, boto_session: Any) -> None: + self._config = config + self._session = boto_session + self._resource_context: Any | None = None + self._ddb: Any | None = None + self._table: Any | None = None + + async def __aenter__(self) -> AsyncDynamoDBStore: + if self._ddb is not None: + return self + + resource_kwargs: dict[str, object] = { + "region_name": self._config.region, + } + + botocore_config = self._config.botocore_config() + if botocore_config is not None: + resource_kwargs["config"] = botocore_config + + self._resource_context = self._session.resource( + "dynamodb", + **resource_kwargs, + ) + + try: + self._ddb = await self._resource_context.__aenter__() + self._table = await self._ddb.Table(self._config.table) + except Exception: + try: + await self._resource_context.__aexit__( + None, + None, + None, + ) + finally: + self._table = None + self._ddb = None + self._resource_context = None + raise + + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> None: + try: + if self._resource_context is not None: + await self._resource_context.__aexit__( + exc_type, + exc_value, + traceback, + ) + finally: + self._table = None + self._ddb = None + self._resource_context = None + + async def aclose(self) -> None: + try: + if self._resource_context is not None: + await self._resource_context.__aexit__( + None, + None, + None, + ) + finally: + self._table = None + self._ddb = None + self._resource_context = None + + @staticmethod + def _pk(namespace: str, doc_id: str) -> str: + return _pk(namespace, doc_id) + + def _require_ddb(self) -> Any: + if self._ddb is None: + raise RuntimeError( + "AsyncDynamoDBStore is not open. " + "Use it inside 'async with' before performing AWS operations." + ) + return self._ddb + + def _require_table(self) -> Any: + if self._table is None: + raise RuntimeError( + "AsyncDynamoDBStore is not open. " + "Use it inside 'async with' before performing AWS operations." + ) + return self._table + + async def put_many( + self, + namespace: str, + items: Sequence[ + tuple[str, str | None, Metadata] + | tuple[str, str | None, Metadata, int | None] + ], + ) -> None: + """Upsert document tuples using DynamoDB's async batch writer.""" + t0 = time.perf_counter() + + threshold = self._config.gzip_threshold_bytes + ttl_attr = self._config.dynamodb_ttl_attribute + + built = [ + _build_item( + namespace, + item[0], + item[1], + item[2], + threshold, + ttl=item[3] if len(item) > 3 else None, + ttl_attribute=ttl_attr, + ) + for item in items + ] + + for item in built: + _check_built_item(item) + + table = self._require_table() + + async with table.batch_writer( + overwrite_by_pkeys=["pk"], + ) as batch: + for item in built: + await batch.put_item(Item=item) + + log_store_event( + self._logger, + "dynamodb.put_many", + self._config.structured_logging, + table=self._config.table, + namespace=namespace, + count=len(items), + duration_ms=round( + (time.perf_counter() - t0) * 1000, + 2, + ), + ) + + @async_retry() + async def get_many( + self, + namespace: str, + ids: list[str], + ) -> dict[str, dict[str, Any]]: + """Hydrate documents by id asynchronously.""" + t0 = time.perf_counter() + + if not ids: + log_store_event( + self._logger, + "dynamodb.get_many", + self._config.structured_logging, + table=self._config.table, + namespace=namespace, + requested_count=0, + returned_count=0, + duration_ms=0.0, + ) + return {} + + ddb = self._require_ddb() + + keys = [ + {"pk": self._pk(namespace, doc_id)} + for doc_id in ids + ] + + out: dict[str, dict[str, Any]] = {} + + for start in range(0, len(keys), _BATCH_GET_LIMIT): + chunk = keys[start : start + _BATCH_GET_LIMIT] + + request: dict[str, Any] | None = { + self._config.table: { + "Keys": chunk, + } + } + + while request: + response = await ddb.batch_get_item( + RequestItems=request, + ) + + for item in response["Responses"].get( + self._config.table, + [], + ): + entry: dict[str, Any] = { + "text": _read_text(item), + "metadata": _from_dynamo( + item.get("metadata", {}) + ), + } + + ttl_attr = self._config.dynamodb_ttl_attribute + if ttl_attr in item: + entry["ttl"] = int(item[ttl_attr]) + + out[item["id"]] = entry + + unprocessed = ( + response.get("UnprocessedKeys") or {} + ) + + request = ( + unprocessed + if unprocessed + else None + ) + + log_store_event( + self._logger, + "dynamodb.get_many", + self._config.structured_logging, + table=self._config.table, + namespace=namespace, + requested_count=len(ids), + returned_count=len(out), + duration_ms=round( + (time.perf_counter() - t0) * 1000, + 2, + ), + ) + + return out diff --git a/src/dynavec/stores/async_s3vectors.py b/src/dynavec/stores/async_s3vectors.py new file mode 100644 index 0000000..d963ded --- /dev/null +++ b/src/dynavec/stores/async_s3vectors.py @@ -0,0 +1,381 @@ +"""Async Amazon S3 Vectors store backed by aioboto3.""" + +from __future__ import annotations + +import asyncio +import logging +import time +from collections.abc import AsyncIterator +from types import TracebackType +from typing import Any + +from ..config import DynavecConfig +from ..logging import log_store_event +from ..utils import TokenBucket, async_retry +from .s3vectors import _GET_LIMIT, _MAX_TOP_K, _PUT_LIMIT, Metadata, _f32 + + +class AsyncS3VectorsStore: + """Async wrapper around Amazon S3 Vectors I/O.""" + + _logger = logging.getLogger("dynavec.stores.async_s3vectors") + + def __init__(self, config: DynavecConfig, boto_session: Any) -> None: + self._config = config + self._session = boto_session + self._put_limiter = ( + TokenBucket(config.put_rps) + if config.put_rps is not None + else None + ) + self._query_limiter = ( + TokenBucket(config.query_rps) + if config.query_rps is not None + else None + ) + self._client_context: Any | None = None + self._client: Any | None = None + + async def __aenter__(self) -> AsyncS3VectorsStore: + if self._client is not None: + return self + + client_kwargs: dict[str, object] = {"region_name": self._config.region} + botocore_config = self._config.botocore_config() + if botocore_config is not None: + client_kwargs["config"] = botocore_config + + self._client_context = self._session.client( + "s3vectors", + **client_kwargs, + ) + self._client = await self._client_context.__aenter__() + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> None: + try: + if self._client_context is not None: + await self._client_context.__aexit__( + exc_type, + exc_value, + traceback, + ) + finally: + self._client = None + self._client_context = None + + async def aclose(self) -> None: + try: + if self._client_context is not None: + await self._client_context.__aexit__( + None, + None, + None, + ) + finally: + self._client = None + self._client_context = None + + def _require_client(self) -> Any: + if self._client is None: + raise RuntimeError( + "AsyncS3VectorsStore is not open. " + "Use it inside 'async with' before performing AWS operations." + ) + return self._client + + @async_retry() + async def _put_batch( + self, + payload: list[dict[str, Any]], + ) -> None: + if self._put_limiter is not None: + await self._put_limiter.acquire_async() + + client = self._require_client() + await client.put_vectors( + vectorBucketName=self._config.vector_bucket, + indexName=self._config.index, + vectors=payload, + ) + + async def put_vectors( + self, + vectors: list[tuple[str, list[float], Metadata]], + max_workers: int = 8, + ) -> None: + """Insert or overwrite vectors using concurrent async batches.""" + if not vectors: + return + + if max_workers <= 0: + raise ValueError("max_workers must be greater than 0.") + + t0 = time.perf_counter() + chunks = [ + vectors[i : i + _PUT_LIMIT] + for i in range(0, len(vectors), _PUT_LIMIT) + ] + + semaphore = asyncio.Semaphore(min(max_workers, len(chunks))) + + async def upload_chunk( + chunk: list[tuple[str, list[float], Metadata]], + ) -> None: + payload = [ + { + "key": key, + "data": {"float32": _f32(vector)}, + "metadata": metadata, + } + for key, vector, metadata in chunk + ] + + async with semaphore: + await self._put_batch(payload) + + tasks = [ + asyncio.create_task(upload_chunk(chunk)) + for chunk in chunks + ] + + try: + await asyncio.gather(*tasks) + except Exception: + for task in tasks: + if not task.done(): + task.cancel() + + await asyncio.gather(*tasks, return_exceptions=True) + raise + + log_store_event( + self._logger, + "s3vectors.put_vectors", + self._config.structured_logging, + bucket=self._config.vector_bucket, + index=self._config.index, + count=len(vectors), + duration_ms=round((time.perf_counter() - t0) * 1000, 2), + ) + + + def _query_kwargs( + self, + query_vector: list[float], + top_k: int, + filter: Metadata | None, + return_metadata: bool, + return_distance: bool, + ) -> dict[str, Any]: + kwargs: dict[str, Any] = { + "vectorBucketName": self._config.vector_bucket, + "indexName": self._config.index, + "queryVector": {"float32": _f32(query_vector)}, + "topK": top_k, + "returnMetadata": return_metadata, + "returnDistance": return_distance, + } + if filter: + kwargs["filter"] = filter + return kwargs + + @async_retry() + async def query( + self, + query_vector: list[float], + top_k: int, + filter: Metadata | None = None, + return_metadata: bool = True, + return_distance: bool = True, + ) -> list[dict[str, Any]]: + """Run an ANN query and drain paginated results up to ``top_k``.""" + t0 = time.perf_counter() + + if top_k <= 0: + log_store_event( + self._logger, + "s3vectors.query", + self._config.structured_logging, + bucket=self._config.vector_bucket, + index=self._config.index, + top_k=top_k, + filtered=filter is not None, + returned_count=0, + duration_ms=0.0, + ) + return [] + + if top_k > _MAX_TOP_K: + raise ValueError( + f"top_k ({top_k}) exceeds Amazon S3 Vectors " + f"maximum limit of {_MAX_TOP_K}." + ) + + results: list[dict[str, Any]] = [] + + async for page in self.query_pages( + query_vector, + top_k, + filter=filter, + return_metadata=return_metadata, + return_distance=return_distance, + ): + results.extend(page) + if len(results) >= top_k: + break + + res = results[:top_k] + + log_store_event( + self._logger, + "s3vectors.query", + self._config.structured_logging, + bucket=self._config.vector_bucket, + index=self._config.index, + top_k=top_k, + filtered=filter is not None, + returned_count=len(res), + duration_ms=round((time.perf_counter() - t0) * 1000, 2), + ) + + return res + + async def query_pages( + self, + query_vector: list[float], + top_k: int, + filter: Metadata | None = None, + return_metadata: bool = True, + return_distance: bool = True, + page_size: int | None = None, + ) -> AsyncIterator[list[dict[str, Any]]]: + """Yield result pages from the async S3 Vectors paginator.""" + if top_k <= 0: + return + + if top_k > _MAX_TOP_K: + raise ValueError( + f"top_k ({top_k}) exceeds Amazon S3 Vectors " + f"maximum limit of {_MAX_TOP_K}." + ) + + effective_page_size = ( + page_size + if page_size is not None + else self._config.top_k_page_size + ) + + if effective_page_size is not None and effective_page_size <= 0: + raise ValueError("page_size must be a positive integer.") + + kwargs = self._query_kwargs( + query_vector, + top_k, + filter, + return_metadata, + return_distance, + ) + + client = self._require_client() + paginator = client.get_paginator("query_vectors") + + yielded = 0 + buffer: list[dict[str, Any]] = [] + + page_iterator = paginator.paginate( + PaginationConfig={"MaxItems": top_k}, + **kwargs, + ).__aiter__() + + while True: + if self._query_limiter is not None: + await self._query_limiter.acquire_async() + + try: + page = await page_iterator.__anext__() + except StopAsyncIteration: + break + + vectors = page.get("vectors", []) + + if vectors: + if effective_page_size is None: + remaining = top_k - yielded + + if len(vectors) > remaining: + vectors = vectors[:remaining] + + yield vectors + yielded += len(vectors) + + if yielded >= top_k: + return + + else: + buffer.extend(vectors) + + while ( + len(buffer) >= effective_page_size + and yielded < top_k + ): + chunk = buffer[:effective_page_size] + buffer = buffer[effective_page_size:] + + remaining = top_k - yielded + + if len(chunk) > remaining: + chunk = chunk[:remaining] + + yield chunk + yielded += len(chunk) + + if yielded >= top_k: + return + + if page.get("NextToken") is None: + break + + if ( + effective_page_size is not None + and buffer + and yielded < top_k + ): + remaining = top_k - yielded + yield buffer[:remaining] + + + @async_retry() + async def get_vectors( + self, + keys: list[str], + return_metadata: bool = False, + ) -> dict[str, dict[str, Any]]: + """Fetch stored vectors by key for MMR reranking.""" + client = self._require_client() + + out: dict[str, dict[str, Any]] = {} + + for start in range(0, len(keys), _GET_LIMIT): + chunk = keys[start : start + _GET_LIMIT] + + response = await client.get_vectors( + vectorBucketName=self._config.vector_bucket, + indexName=self._config.index, + keys=chunk, + returnData=True, + returnMetadata=return_metadata, + ) + + for vector in response.get("vectors", []): + out[vector["key"]] = { + "vector": vector.get("data", {}).get("float32"), + "metadata": vector.get("metadata", {}), + } + + return out diff --git a/src/dynavec/utils.py b/src/dynavec/utils.py index 3e2c03d..70e8715 100644 --- a/src/dynavec/utils.py +++ b/src/dynavec/utils.py @@ -6,11 +6,12 @@ from __future__ import annotations +import asyncio import functools import random import threading import time -from collections.abc import Iterable, Iterator +from collections.abc import Awaitable, Iterable, Iterator from typing import Any, Callable, TypeVar T = TypeVar("T") @@ -85,6 +86,44 @@ def wrapper(*args: Any, **kwargs: Any) -> T: return decorator +def async_retry( + max_attempts: int = 5, + base_delay: float = 0.1, + max_delay: float = 5.0, + retry_on: Callable[[Exception], bool] = is_retryable, + retry_delay: Callable[[Exception], float | None] | None = None, +) -> Callable[ + [Callable[..., Awaitable[T]]], + Callable[..., Awaitable[T]], +]: + """Async retry decorator with exponential backoff and full jitter.""" + + def decorator( + fn: Callable[..., Awaitable[T]], + ) -> Callable[..., Awaitable[T]]: + @functools.wraps(fn) + async def wrapper(*args: Any, **kwargs: Any) -> T: + attempt = 0 + while True: + try: + return await fn(*args, **kwargs) + except Exception as exc: # noqa: BLE001 + attempt += 1 + if attempt >= max_attempts or not retry_on(exc): + raise + + server_delay = retry_delay(exc) if retry_delay is not None else None + if server_delay is not None: + await asyncio.sleep(server_delay) + else: + delay = min(max_delay, base_delay * (2 ** (attempt - 1))) + await asyncio.sleep(random.uniform(0, delay)) + + return wrapper + + return decorator + + def timed( sink: Callable[[str, float], None] | None = None, ) -> Callable[[Callable[..., T]], Callable[..., T]]: @@ -143,22 +182,37 @@ def __init__( self._lock = threading.Lock() + def _acquire_wait_time(self) -> float | None: + with self._lock: + now = time.monotonic() + elapsed = now - self.last_time + + self.tokens = min( + self.capacity, + self.tokens + elapsed * self.rate, + ) + self.last_time = now + + if self.tokens >= 1: + self.tokens -= 1 + return None + + return (1 - self.tokens) / self.rate + def acquire(self) -> None: while True: - with self._lock: - now = time.monotonic() - elapsed = now - self.last_time + wait_time = self._acquire_wait_time() + + if wait_time is None: + return - self.tokens = min( - self.capacity, - self.tokens + elapsed * self.rate, - ) - self.last_time = now + time.sleep(wait_time) - if self.tokens >= 1: - self.tokens -= 1 - return + async def acquire_async(self) -> None: + while True: + wait_time = self._acquire_wait_time() - wait_time = (1 - self.tokens) / self.rate + if wait_time is None: + return - time.sleep(wait_time) \ No newline at end of file + await asyncio.sleep(wait_time) diff --git a/tests/test_async_client.py b/tests/test_async_client.py new file mode 100644 index 0000000..a871731 --- /dev/null +++ b/tests/test_async_client.py @@ -0,0 +1,795 @@ +"""Tests for the high-level AsyncDynavec client.""" + +from __future__ import annotations + +import asyncio +from collections.abc import AsyncIterator +from typing import Any + +import pytest + +from dynavec.async_client import AsyncDynavec +from dynavec.client_common import s3_key +from dynavec.config import NS_METADATA_KEY, DynavecConfig +from dynavec.embeddings.base import Embedder +from dynavec.exceptions import ( + ConfigurationError, + DimensionMismatchError, +) + + +def _config() -> DynavecConfig: + return DynavecConfig( + vector_bucket="bucket", + index="index", + table="docs", + dimension=4, + ) + + +class AsyncOnlyEmbedder(Embedder): + dimension = 4 + + def __init__(self) -> None: + self.async_calls: list[list[str]] = [] + self.async_query_calls: list[str] = [] + + def embed_query( + self, + text: str, + ) -> list[float]: + raise AssertionError( + "AsyncDynavec must not call sync embed_query()." + ) + + async def aembed_query( + self, + text: str, + ) -> list[float]: + self.async_query_calls.append(text) + return [1.0, 2.0, 3.0, 4.0] + + def embed_documents( + self, + texts: list[str], + ) -> list[list[float]]: + raise AssertionError( + "AsyncDynavec must not call sync embed_documents()." + ) + + async def aembed_documents( + self, + texts: list[str], + ) -> list[list[float]]: + self.async_calls.append(texts) + + return [ + [1.0, 2.0, 3.0, 4.0] + for _ in texts + ] + + +class FakeVectorStore: + def __init__(self) -> None: + self.entered = False + self.exited = False + self.put_calls: list[ + tuple[list[tuple[str, list[float], dict[str, Any]]], int] + ] = [] + self.query_results: list[dict[str, Any]] = [] + self.query_calls: list[dict[str, Any]] = [] + self.query_page_results: list[list[dict[str, Any]]] = [] + self.query_page_calls: list[dict[str, Any]] = [] + + async def query( + self, + query_vector: list[float], + top_k: int, + filter: dict[str, Any] | None = None, + return_metadata: bool = True, + return_distance: bool = True, + ) -> list[dict[str, Any]]: + self.query_calls.append( + { + "query_vector": query_vector, + "top_k": top_k, + "filter": filter, + "return_metadata": return_metadata, + "return_distance": return_distance, + } + ) + + return self.query_results + + async def __aenter__(self) -> FakeVectorStore: + self.entered = True + return self + + async def __aexit__( + self, + exc_type: Any, + exc_value: Any, + traceback: Any, + ) -> None: + self.exited = True + + async def put_vectors( + self, + vectors: list[ + tuple[str, list[float], dict[str, Any]] + ], + max_workers: int = 8, + ) -> None: + self.put_calls.append( + (vectors, max_workers) + ) + + async def query_pages( + self, + query_vector: list[float], + top_k: int, + filter: dict[str, Any] | None = None, + return_metadata: bool = True, + return_distance: bool = True, + page_size: int | None = None, + ) -> AsyncIterator[list[dict[str, Any]]]: + self.query_page_calls.append( + { + "query_vector": query_vector, + "top_k": top_k, + "filter": filter, + "return_metadata": return_metadata, + "return_distance": return_distance, + "page_size": page_size, + } + ) + + for page in self.query_page_results: + yield page + + +class FakeDocumentStore: + def __init__(self) -> None: + self.entered = False + self.exited = False + self.put_calls: list[ + tuple[ + str, + list[ + tuple[ + str, + str | None, + dict[str, Any], + ] + ], + ] + ] = [] + self.documents: dict[str, dict[str, Any]] = {} + self.get_calls: list[tuple[str, list[str]]] = [] + + async def get_many( + self, + namespace: str, + ids: list[str], + ) -> dict[str, dict[str, Any]]: + self.get_calls.append( + ( + namespace, + ids, + ) + ) + + return { + doc_id: self.documents[doc_id] + for doc_id in ids + if doc_id in self.documents + } + + async def __aenter__( + self, + ) -> FakeDocumentStore: + self.entered = True + return self + + async def __aexit__( + self, + exc_type: Any, + exc_value: Any, + traceback: Any, + ) -> None: + self.exited = True + + async def put_many( + self, + namespace: str, + items: list[ + tuple[ + str, + str | None, + dict[str, Any], + ] + ], + ) -> None: + self.put_calls.append( + (namespace, items) + ) + + +def _install_fake_stores( + client: AsyncDynavec, + vectors: FakeVectorStore, + documents: FakeDocumentStore, +) -> None: + client._vectors = vectors # type: ignore[assignment] + client._docs = documents # type: ignore[assignment] + + +def _client( + *, + embedder: Embedder | None = None, +) -> tuple[ + AsyncDynavec, + FakeVectorStore, + FakeDocumentStore, +]: + client = AsyncDynavec( + _config(), + embedder=embedder, + boto_session=object(), + ) + + vectors = FakeVectorStore() + documents = FakeDocumentStore() + + _install_fake_stores( + client, + vectors, + documents, + ) + + return client, vectors, documents + + +async def test_async_context_opens_and_closes_stores() -> None: + client, vectors, documents = _client() + + async with client: + assert vectors.entered + assert documents.entered + assert not vectors.exited + assert not documents.exited + + assert vectors.exited + assert documents.exited + + +async def test_aclose_closes_open_stores() -> None: + client, vectors, documents = _client() + + await client.__aenter__() + + assert vectors.entered + assert documents.entered + + await client.aclose() + + assert vectors.exited + assert documents.exited + + # Closing twice should be safe. + await client.aclose() + + +async def test_aupsert_requires_open_client() -> None: + client, _, _ = _client() + + with pytest.raises( + RuntimeError, + match="not open", + ): + await client.aupsert( + [ + { + "id": "a", + "vector": [1.0, 2.0, 3.0, 4.0], + } + ] + ) + + +async def test_aupsert_uses_async_embedding_and_writes_both_stores() -> None: + embedder = AsyncOnlyEmbedder() + client, vectors, documents = _client( + embedder=embedder + ) + + async with client: + result = await client.aupsert( + [ + { + "id": "a", + "text": "hello", + "metadata": { + "topic": "test", + }, + } + ], + namespace="tenant", + ) + + assert result.count == 1 + assert result.ids == ["a"] + + assert embedder.async_calls == [ + ["hello"] + ] + + assert len(vectors.put_calls) == 1 + assert len(documents.put_calls) == 1 + + vector_payload, _ = vectors.put_calls[0] + + assert len(vector_payload) == 1 + assert vector_payload[0][1] == [ + 1.0, + 2.0, + 3.0, + 4.0, + ] + + namespace, document_payload = ( + documents.put_calls[0] + ) + + assert namespace == "tenant" + assert document_payload[0][0] == "a" + assert document_payload[0][1] == "hello" + + +async def test_aupsert_passes_default_ttl( + monkeypatch: pytest.MonkeyPatch, +) -> None: + client, _, documents = _client() + + monkeypatch.setattr( + "dynavec.client_common.time.time", + lambda: 1_700_000_000, + ) + + async with client: + await client.aupsert( + [ + { + "id": "a", + "vector": [1.0, 2.0, 3.0, 4.0], + } + ], + namespace="tenant", + ttl_seconds=60, + ) + + namespace, document_payload = documents.put_calls[0] + + assert namespace == "tenant" + assert document_payload[0][2]["_ttl"] == 1_700_000_060 + + +async def test_aupsert_rejects_non_positive_ttl() -> None: + client, _, _ = _client() + + with pytest.raises( + ValueError, + match="ttl_seconds must be positive", + ): + await client.aupsert( + [ + { + "id": "a", + "vector": [1.0, 2.0, 3.0, 4.0], + } + ], + ttl_seconds=0, + ) + + +async def test_aupsert_runs_s3_and_dynamodb_writes_concurrently() -> None: + s3_started = asyncio.Event() + ddb_started = asyncio.Event() + + class CoordinatedVectors( + FakeVectorStore + ): + async def put_vectors( + self, + vectors: list[ + tuple[ + str, + list[float], + dict[str, Any], + ] + ], + max_workers: int = 8, + ) -> None: + s3_started.set() + + await asyncio.wait_for( + ddb_started.wait(), + timeout=1, + ) + + class CoordinatedDocuments( + FakeDocumentStore + ): + async def put_many( + self, + namespace: str, + items: list[ + tuple[ + str, + str | None, + dict[str, Any], + ] + ], + ) -> None: + ddb_started.set() + + await asyncio.wait_for( + s3_started.wait(), + timeout=1, + ) + + client = AsyncDynavec( + _config(), + boto_session=object(), + ) + + vectors = CoordinatedVectors() + documents = CoordinatedDocuments() + + _install_fake_stores( + client, + vectors, + documents, + ) + + async with client: + await asyncio.wait_for( + client.aupsert( + [ + { + "id": "a", + "vector": [ + 1.0, + 2.0, + 3.0, + 4.0, + ], + } + ] + ), + timeout=1, + ) + + assert s3_started.is_set() + assert ddb_started.is_set() + + +async def test_partial_enter_failure_closes_first_store() -> None: + vectors = FakeVectorStore() + + class FailingDocumentStore( + FakeDocumentStore + ): + async def __aenter__( + self, + ) -> FailingDocumentStore: + raise RuntimeError( + "DynamoDB open failed" + ) + + client = AsyncDynavec( + _config(), + boto_session=object(), + ) + + documents = FailingDocumentStore() + + _install_fake_stores( + client, + vectors, + documents, + ) + + with pytest.raises( + RuntimeError, + match="DynamoDB open failed", + ): + async with client: + pass + + assert vectors.entered + assert vectors.exited + + +async def test_resolve_query_vector_uses_precomputed_vector(): + client, _, _ = _client() + + vector = [0.1, 0.2, 0.3, 0.4] + + result = await client._resolve_query_vector( + None, + vector, + ) + + assert result is vector + + +async def test_resolve_query_vector_rejects_wrong_dimension(): + client, _, _ = _client() + + with pytest.raises( + DimensionMismatchError, + match="Query vector dimension", + ): + await client._resolve_query_vector( + None, + [0.1, 0.2], + ) + + +async def test_resolve_query_vector_requires_query_or_vector(): + client, _, _ = _client() + + with pytest.raises( + ValueError, + match="Provide either 'query' text or a 'vector'", + ): + await client._resolve_query_vector( + None, + None, + ) + + +async def test_resolve_query_vector_requires_embedder(): + client, _, _ = _client() + + with pytest.raises( + ConfigurationError, + match="Text query requires an embedder", + ): + await client._resolve_query_vector( + "hello", + None, + ) + + +async def test_resolve_query_vector_uses_async_embedder(): + embedder = AsyncOnlyEmbedder() + client, _, _ = _client(embedder=embedder) + + result = await client._resolve_query_vector( + "hello", + None, + ) + + assert result == [1.0, 2.0, 3.0, 4.0] + assert embedder.async_query_calls == ["hello"] + + +async def test_asearch_queries_and_hydrates_results(): + client, vectors, documents = _client() + + vectors.query_results = [ + { + "key": s3_key("default", "doc-1"), + "distance": 0.2, + }, + { + "key": s3_key("default", "doc-2"), + "distance": 0.4, + }, + ] + + documents.documents = { + "doc-1": { + "text": "First document", + "metadata": {"topic": "python"}, + "ttl": 1_700_000_060, + }, + "doc-2": { + "text": "Second document", + "metadata": {"topic": "asyncio"}, + }, + } + + async with client: + results = await client.asearch( + vector=[0.1, 0.2, 0.3, 0.4], + top_k=2, + ) + + assert [result.id for result in results] == [ + "doc-1", + "doc-2", + ] + + assert results[0].distance == 0.2 + assert results[0].text == "First document" + assert results[0].metadata == {"topic": "python"} + assert results[0].ttl == 1_700_000_060 + + assert results[1].distance == 0.4 + assert results[1].text == "Second document" + assert results[1].metadata == {"topic": "asyncio"} + assert results[1].ttl is None + + assert documents.get_calls == [ + ( + "default", + ["doc-1", "doc-2"], + ) + ] + + assert len(vectors.query_calls) == 1 + assert vectors.query_calls[0]["query_vector"] == [ + 0.1, + 0.2, + 0.3, + 0.4, + ] + assert vectors.query_calls[0]["top_k"] == 2 + + +async def test_asearch_returns_empty_when_vector_search_has_no_hits(): + client, vectors, documents = _client() + + vectors.query_results = [] + + async with client: + results = await client.asearch( + vector=[0.1, 0.2, 0.3, 0.4], + ) + + assert results == [] + assert documents.get_calls == [] + + +async def test_asearch_passes_namespace_and_filter_to_vector_store(): + client, vectors, _ = _client() + + vectors.query_results = [] + + async with client: + await client.asearch( + vector=[0.1, 0.2, 0.3, 0.4], + namespace="tenant-a", + filter={"topic": "python"}, + ) + + assert vectors.query_calls[0]["filter"] == { + "$and": [ + {"topic": "python"}, + {NS_METADATA_KEY: "tenant-a"}, + ] + } + + +async def test_asearch_stream_yields_results_page_by_page(): + client, vectors, documents = _client() + + vectors.query_page_results = [ + [ + { + "key": s3_key("tenant-a", "doc-1"), + "distance": 0.1, + }, + { + "key": s3_key("tenant-a", "doc-2"), + "distance": 0.2, + }, + ], + [ + { + "key": s3_key("tenant-a", "doc-3"), + "distance": 0.3, + } + ], + ] + + documents.documents = { + "doc-1": { + "text": "First", + "metadata": {"topic": "python"}, + }, + "doc-2": { + "text": "Second", + "metadata": {"topic": "python"}, + }, + "doc-3": { + "text": "Third", + "metadata": {"topic": "python"}, + }, + } + + async with client: + results = [ + result + async for result in client.asearch_stream( + vector=[0.1, 0.2, 0.3, 0.4], + top_k=3, + namespace="tenant-a", + filter={"topic": "python"}, + page_size=2, + ) + ] + + assert [result.id for result in results] == [ + "doc-1", + "doc-2", + "doc-3", + ] + + assert documents.get_calls == [ + ( + "tenant-a", + ["doc-1", "doc-2"], + ), + ( + "tenant-a", + ["doc-3"], + ), + ] + + assert len(vectors.query_page_calls) == 1 + + call = vectors.query_page_calls[0] + + assert call["top_k"] == 3 + assert call["page_size"] == 2 + assert call["filter"] == { + "$and": [ + {"topic": "python"}, + {NS_METADATA_KEY: "tenant-a"}, + ] + } + + +async def test_asearch_stream_stops_at_top_k(): + client, vectors, documents = _client() + + vectors.query_page_results = [ + [ + { + "key": s3_key("default", "doc-1"), + "distance": 0.1, + }, + { + "key": s3_key("default", "doc-2"), + "distance": 0.2, + }, + { + "key": s3_key("default", "doc-3"), + "distance": 0.3, + }, + ] + ] + + documents.documents = { + "doc-1": {"text": "First", "metadata": {}}, + "doc-2": {"text": "Second", "metadata": {}}, + "doc-3": {"text": "Third", "metadata": {}}, + } + + async with client: + results = [ + result + async for result in client.asearch_stream( + vector=[0.1, 0.2, 0.3, 0.4], + top_k=2, + ) + ] + + assert [result.id for result in results] == [ + "doc-1", + "doc-2", + ] diff --git a/tests/test_async_dynamodb.py b/tests/test_async_dynamodb.py new file mode 100644 index 0000000..46b5614 --- /dev/null +++ b/tests/test_async_dynamodb.py @@ -0,0 +1,440 @@ +"""Tests for the async DynamoDB store.""" + +from __future__ import annotations + +from decimal import Decimal +from typing import Any + +import pytest + +from dynavec.config import DynavecConfig +from dynavec.exceptions import ItemTooLargeError +from dynavec.stores.async_dynamodb import AsyncDynamoDBStore +from dynavec.stores.dynamodb import MAX_ITEM_BYTES + + +def _config() -> DynavecConfig: + return DynavecConfig( + vector_bucket="bucket", + index="index", + table="docs", + dimension=8, + ) + + +class FakeBatchWriter: + def __init__(self) -> None: + self.items: list[dict[str, Any]] = [] + self.entered = False + self.exited = False + + async def __aenter__(self) -> FakeBatchWriter: + self.entered = True + return self + + async def __aexit__( + self, + exc_type: Any, + exc_value: Any, + traceback: Any, + ) -> None: + self.exited = True + + async def put_item(self, **kwargs: Any) -> None: + self.items.append(kwargs["Item"]) + + +class FakeTable: + def __init__(self) -> None: + self.writer = FakeBatchWriter() + self.batch_writer_calls: list[dict[str, Any]] = [] + + def batch_writer(self, **kwargs: Any) -> FakeBatchWriter: + self.batch_writer_calls.append(kwargs) + return self.writer + + +class FakeDynamoResource: + def __init__( + self, + table: FakeTable, + responses: list[dict[str, Any]] | None = None, + ) -> None: + self.table = table + self.responses = list(responses or []) + self.batch_get_calls: list[dict[str, Any]] = [] + self.table_calls: list[str] = [] + + async def Table(self, name: str) -> FakeTable: + self.table_calls.append(name) + return self.table + + async def batch_get_item( + self, + **kwargs: Any, + ) -> dict[str, Any]: + self.batch_get_calls.append(kwargs) + + if self.responses: + return self.responses.pop(0) + + return { + "Responses": {"docs": []}, + "UnprocessedKeys": {}, + } + + +class FakeResourceContext: + def __init__(self, resource: FakeDynamoResource) -> None: + self.resource = resource + self.entered = False + self.exited = False + + async def __aenter__(self) -> FakeDynamoResource: + self.entered = True + return self.resource + + async def __aexit__( + self, + exc_type: Any, + exc_value: Any, + traceback: Any, + ) -> None: + self.exited = True + + +class FakeSession: + def __init__(self, resource: FakeDynamoResource) -> None: + self.resource_instance = resource + self.context: FakeResourceContext | None = None + self.calls: list[tuple[str, dict[str, Any]]] = [] + + def resource( + self, + service_name: str, + **kwargs: Any, + ) -> FakeResourceContext: + self.calls.append((service_name, kwargs)) + self.context = FakeResourceContext( + self.resource_instance + ) + return self.context + + +async def test_async_store_enters_and_closes_resource(): + table = FakeTable() + resource = FakeDynamoResource(table) + session = FakeSession(resource) + + async with AsyncDynamoDBStore(_config(), session): + assert session.context is not None + assert session.context.entered + assert resource.table_calls == ["docs"] + + assert session.context is not None + assert session.context.exited + assert session.calls[0][0] == "dynamodb" + + +async def test_put_many_builds_and_writes_items(): + table = FakeTable() + resource = FakeDynamoResource(table) + session = FakeSession(resource) + + async with AsyncDynamoDBStore(_config(), session) as store: + await store.put_many( + "tenant", + [ + ("a", "First", {"score": 0.5}), + ("b", "Second", {"year": 2026}), + ], + ) + + assert table.batch_writer_calls == [ + {"overwrite_by_pkeys": ["pk"]} + ] + + assert len(table.writer.items) == 2 + + assert table.writer.items[0]["pk"] == "tenant#a" + assert table.writer.items[0]["id"] == "a" + assert table.writer.items[0]["text"] == "First" + assert table.writer.items[0]["metadata"]["score"] == Decimal( + "0.5" + ) + + +async def test_put_many_writes_custom_ttl_attribute(): + config = DynavecConfig( + vector_bucket="bucket", + index="index", + table="docs", + dimension=8, + dynamodb_ttl_attribute="expires_at", + ) + + table = FakeTable() + resource = FakeDynamoResource(table) + session = FakeSession(resource) + + async with AsyncDynamoDBStore(config, session) as store: + await store.put_many( + "tenant", + [ + ( + "a", + "First", + { + "topic": "test", + "_ttl": 1_700_000_060, + }, + ) + ], + ) + + item = table.writer.items[0] + + assert item["expires_at"] == 1_700_000_060 + assert "_ttl" not in item["metadata"] + assert item["metadata"]["topic"] == "test" + + +async def test_put_many_rejects_oversized_item_before_writing(): + table = FakeTable() + resource = FakeDynamoResource(table) + session = FakeSession(resource) + + async with AsyncDynamoDBStore(_config(), session) as store: + with pytest.raises(ItemTooLargeError): + await store.put_many( + "tenant", + [ + ("ok", "fits", {}), + ( + "big", + "x" * MAX_ITEM_BYTES, + {}, + ), + ], + ) + + assert table.batch_writer_calls == [] + + +async def test_get_many_retries_only_unprocessed_keys(): + table = FakeTable() + + remaining = { + "docs": { + "Keys": [ + {"pk": "tenant#b"}, + {"pk": "tenant#c"}, + ] + } + } + + resource = FakeDynamoResource( + table, + responses=[ + { + "Responses": { + "docs": [ + { + "id": "a", + "text": "First document", + "metadata": { + "topic": "aws", + }, + } + ] + }, + "UnprocessedKeys": remaining, + }, + { + "Responses": { + "docs": [ + { + "id": "c", + "text": "Third document", + "metadata": { + "score": Decimal("0.1"), + }, + }, + { + "id": "b", + "text": "Second document", + "metadata": { + "year": Decimal("2026"), + }, + }, + ] + }, + "UnprocessedKeys": {}, + }, + ], + ) + + session = FakeSession(resource) + + async with AsyncDynamoDBStore(_config(), session) as store: + documents = await store.get_many( + "tenant", + ["a", "b", "c"], + ) + + assert resource.batch_get_calls == [ + { + "RequestItems": { + "docs": { + "Keys": [ + {"pk": "tenant#a"}, + {"pk": "tenant#b"}, + {"pk": "tenant#c"}, + ] + } + } + }, + { + "RequestItems": remaining, + }, + ] + + assert documents == { + "a": { + "text": "First document", + "metadata": {"topic": "aws"}, + }, + "b": { + "text": "Second document", + "metadata": {"year": 2026}, + }, + "c": { + "text": "Third document", + "metadata": {"score": 0.1}, + }, + } + + +async def test_get_many_batches_at_100_keys(): + table = FakeTable() + + resource = FakeDynamoResource( + table, + responses=[ + { + "Responses": {"docs": []}, + "UnprocessedKeys": {}, + }, + { + "Responses": {"docs": []}, + "UnprocessedKeys": {}, + }, + { + "Responses": {"docs": []}, + "UnprocessedKeys": {}, + }, + ], + ) + + session = FakeSession(resource) + ids = [f"doc-{i}" for i in range(250)] + + async with AsyncDynamoDBStore(_config(), session) as store: + await store.get_many("tenant", ids) + + sizes = [ + len( + call["RequestItems"]["docs"]["Keys"] + ) + for call in resource.batch_get_calls + ] + + assert sizes == [100, 100, 50] + + +async def test_operations_require_open_store(): + store = AsyncDynamoDBStore( + _config(), + FakeSession( + FakeDynamoResource(FakeTable()) + ), + ) + + with pytest.raises(RuntimeError, match="not open"): + await store.put_many( + "tenant", + [("a", "text", {})], + ) + + with pytest.raises(RuntimeError, match="not open"): + await store.get_many( + "tenant", + ["a"], + ) + + +async def test_get_many_returns_ttl(): + config = DynavecConfig( + vector_bucket="bucket", + index="index", + table="docs", + dimension=8, + dynamodb_ttl_attribute="expires_at", + ) + + table = FakeTable() + resource = FakeDynamoResource( + table, + responses=[ + { + "Responses": { + "docs": [ + { + "id": "a", + "text": "First", + "metadata": { + "topic": "test", + }, + "expires_at": Decimal("1700000060"), + } + ] + }, + "UnprocessedKeys": {}, + } + ], + ) + + session = FakeSession(resource) + + async with AsyncDynamoDBStore(config, session) as store: + documents = await store.get_many( + "tenant", + ["a"], + ) + + assert documents["a"]["ttl"] == 1_700_000_060 + + +async def test_enter_closes_resource_if_table_creation_fails(): + table = FakeTable() + resource = FakeDynamoResource(table) + + async def fail_table(name: str): + raise RuntimeError("table creation failed") + + resource.Table = fail_table # type: ignore[method-assign] + session = FakeSession(resource) + + store = AsyncDynamoDBStore(_config(), session) + + with pytest.raises( + RuntimeError, + match="table creation failed", + ): + await store.__aenter__() + + assert session.context is not None + assert session.context.exited + assert store._ddb is None + assert store._table is None + assert store._resource_context is None diff --git a/tests/test_async_dynamodb_moto.py b/tests/test_async_dynamodb_moto.py new file mode 100644 index 0000000..12104d6 --- /dev/null +++ b/tests/test_async_dynamodb_moto.py @@ -0,0 +1,124 @@ +"""Moto-backed integration tests for the async DynamoDB store.""" + +from __future__ import annotations + +from typing import Any + +import aioboto3 +import boto3 +import pytest +from moto.server import ThreadedMotoServer + +from dynavec.config import DynavecConfig +from dynavec.stores.async_dynamodb import AsyncDynamoDBStore + + +class MotoAsyncSession: + """aioboto3 session that routes DynamoDB calls to Moto server.""" + + def __init__(self, endpoint_url: str) -> None: + self._endpoint_url = endpoint_url + self._session = aioboto3.Session( + aws_access_key_id="testing", + aws_secret_access_key="testing", + region_name="us-east-1", + ) + + def resource( + self, + service_name: str, + **kwargs: Any, + ) -> Any: + return self._session.resource( + service_name, + endpoint_url=self._endpoint_url, + **kwargs, + ) + + +@pytest.fixture +def moto_dynamodb_endpoint() -> str: + server = ThreadedMotoServer(port=0) + server.start() + + try: + _, port = server.get_host_and_port() + endpoint_url = f"http://127.0.0.1:{port}" + + client = boto3.client( + "dynamodb", + region_name="us-east-1", + endpoint_url=endpoint_url, + aws_access_key_id="testing", + aws_secret_access_key="testing", + ) + + client.create_table( + TableName="docs", + KeySchema=[ + { + "AttributeName": "pk", + "KeyType": "HASH", + } + ], + AttributeDefinitions=[ + { + "AttributeName": "pk", + "AttributeType": "S", + } + ], + BillingMode="PAY_PER_REQUEST", + ) + + yield endpoint_url + finally: + server.stop() + + +async def test_async_dynamodb_roundtrip_with_moto( + moto_dynamodb_endpoint: str, +) -> None: + config = DynavecConfig( + vector_bucket="test-bucket", + index="test-index", + table="docs", + dimension=4, + region="us-east-1", + ) + + session = MotoAsyncSession( + moto_dynamodb_endpoint, + ) + + async with AsyncDynamoDBStore( + config, + session, + ) as store: + await store.put_many( + "tenant-a", + [ + ( + "doc-1", + "Hello async world", + { + "topic": "async", + "source": "moto", + }, + ) + ], + ) + + documents = await store.get_many( + "tenant-a", + ["doc-1"], + ) + + assert documents == { + "doc-1": { + "text": "Hello async world", + "metadata": { + "topic": "async", + "source": "moto", + }, + } + } diff --git a/tests/test_async_s3vectors.py b/tests/test_async_s3vectors.py new file mode 100644 index 0000000..9a46cba --- /dev/null +++ b/tests/test_async_s3vectors.py @@ -0,0 +1,382 @@ +"""Tests for the async S3 Vectors store.""" + +from __future__ import annotations + +import asyncio +from typing import Any +from unittest.mock import AsyncMock + +import pytest + +from dynavec.config import DynavecConfig +from dynavec.stores.async_s3vectors import AsyncS3VectorsStore + + +class FakeAsyncPaginator: + def __init__(self, pages: list[dict[str, Any]]) -> None: + self.pages = pages + self.calls: list[dict[str, Any]] = [] + + async def paginate(self, **kwargs: Any): + self.calls.append(kwargs) + + for page in self.pages: + yield page + + +def _config() -> DynavecConfig: + return DynavecConfig( + vector_bucket="test-bucket", + index="test-index", + table="test-table", + dimension=4, + ) + + +class FakeClientContext: + def __init__(self, client: Any) -> None: + self.client = client + self.entered = False + self.exited = False + + async def __aenter__(self) -> Any: + self.entered = True + return self.client + + async def __aexit__( + self, + exc_type: Any, + exc_value: Any, + traceback: Any, + ) -> None: + self.exited = True + + +class FakeSession: + def __init__(self, client: Any) -> None: + self.client_instance = client + self.context: FakeClientContext | None = None + self.calls: list[tuple[str, dict[str, Any]]] = [] + + def client(self, service_name: str, **kwargs: Any) -> FakeClientContext: + self.calls.append((service_name, kwargs)) + self.context = FakeClientContext(self.client_instance) + return self.context + + +class FakeS3Client: + def __init__( + self, + pages: list[dict[str, Any]] | None = None, + ) -> None: + self.put_calls: list[dict[str, Any]] = [] + self.get_calls: list[dict[str, Any]] = [] + self.paginator = FakeAsyncPaginator(pages or []) + + async def put_vectors(self, **kwargs: Any) -> None: + self.put_calls.append(kwargs) + + def get_paginator(self, name: str) -> FakeAsyncPaginator: + assert name == "query_vectors" + return self.paginator + + async def get_vectors(self, **kwargs: Any) -> dict[str, Any]: + self.get_calls.append(kwargs) + + return { + "vectors": [ + { + "key": key, + "data": {"float32": [1.0, 2.0, 3.0, 4.0]}, + "metadata": {"source": "test"}, + } + for key in kwargs["keys"] + ] + } + + +async def test_async_store_enters_and_closes_client(): + client = FakeS3Client() + session = FakeSession(client) + + async with AsyncS3VectorsStore(_config(), session): + assert session.context is not None + assert session.context.entered + + assert session.context is not None + assert session.context.exited + assert session.calls[0][0] == "s3vectors" + + +async def test_put_vectors_batches_at_service_limit(): + client = FakeS3Client() + session = FakeSession(client) + + vectors = [ + (f"key-{i}", [0.1, 0.2, 0.3, 0.4], {"tag": "test"}) + for i in range(1200) + ] + + async with AsyncS3VectorsStore(_config(), session) as store: + await store.put_vectors(vectors) + + assert sorted(len(call["vectors"]) for call in client.put_calls) == [ + 200, + 500, + 500, + ] + + +async def test_put_vectors_runs_batches_concurrently(monkeypatch): + session = FakeSession(FakeS3Client()) + + async with AsyncS3VectorsStore(_config(), session) as store: + active = 0 + max_active = 0 + + async def fake_put_batch(payload): + nonlocal active, max_active + + active += 1 + max_active = max(max_active, active) + + await asyncio.sleep(0) + + active -= 1 + + monkeypatch.setattr(store, "_put_batch", fake_put_batch) + + vectors = [ + (f"key-{i}", [0.1, 0.2, 0.3, 0.4], {}) + for i in range(1200) + ] + + await store.put_vectors(vectors) + + assert max_active > 1 + + +async def test_put_vectors_propagates_errors(monkeypatch): + session = FakeSession(FakeS3Client()) + + async with AsyncS3VectorsStore(_config(), session) as store: + + async def fail(payload): + raise RuntimeError("S3 API failure") + + monkeypatch.setattr(store, "_put_batch", fail) + + with pytest.raises(RuntimeError, match="S3 API failure"): + await store.put_vectors( + [("key", [0.1, 0.2, 0.3, 0.4], {})] + ) + + +async def test_put_vectors_requires_open_store(): + store = AsyncS3VectorsStore(_config(), FakeSession(FakeS3Client())) + + with pytest.raises(RuntimeError, match="not open"): + await store.put_vectors( + [("key", [0.1, 0.2, 0.3, 0.4], {})] + ) + + +async def test_query_returns_paginated_results(): + pages = [ + { + "vectors": [ + {"key": f"k{i}", "distance": i * 0.01} + for i in range(100) + ], + "NextToken": "token-1", + }, + { + "vectors": [ + {"key": f"k{i}", "distance": i * 0.01} + for i in range(100, 150) + ] + }, + ] + + client = FakeS3Client(pages) + session = FakeSession(client) + + async with AsyncS3VectorsStore(_config(), session) as store: + results = await store.query( + [0.1, 0.2, 0.3, 0.4], + top_k=150, + ) + + assert len(results) == 150 + assert results[0]["key"] == "k0" + assert results[-1]["key"] == "k149" + + +async def test_query_truncates_at_top_k(): + pages = [ + { + "vectors": [ + {"key": f"k{i}"} + for i in range(100) + ], + "NextToken": "token-1", + }, + { + "vectors": [ + {"key": f"k{i}"} + for i in range(100, 200) + ] + }, + ] + + client = FakeS3Client(pages) + session = FakeSession(client) + + async with AsyncS3VectorsStore(_config(), session) as store: + results = await store.query( + [0.1, 0.2, 0.3, 0.4], + top_k=130, + ) + + assert len(results) == 130 + assert results[-1]["key"] == "k129" + + +async def test_query_pages_respects_page_size(): + pages = [ + { + "vectors": [ + {"key": f"k{i}"} + for i in range(25) + ] + } + ] + + client = FakeS3Client(pages) + session = FakeSession(client) + + async with AsyncS3VectorsStore(_config(), session) as store: + result_pages = [ + page + async for page in store.query_pages( + [0.1, 0.2, 0.3, 0.4], + top_k=25, + page_size=10, + ) + ] + + assert [len(page) for page in result_pages] == [10, 10, 5] + + +async def test_query_rejects_excessive_top_k(): + client = FakeS3Client() + session = FakeSession(client) + + async with AsyncS3VectorsStore(_config(), session) as store: + with pytest.raises( + ValueError, + match="exceeds Amazon S3 Vectors maximum limit", + ): + await store.query( + [0.1, 0.2, 0.3, 0.4], + top_k=10_001, + ) + + +async def test_put_batch_uses_rate_limiter(): + config = DynavecConfig( + vector_bucket="test-bucket", + index="test-index", + table="test-table", + dimension=4, + put_rps=10, + ) + + client = FakeS3Client() + session = FakeSession(client) + + async with AsyncS3VectorsStore(config, session) as store: + limiter = AsyncMock() + store._put_limiter = limiter + + payload = [ + { + "key": "key-1", + "data": { + "float32": [0.1, 0.2, 0.3, 0.4], + }, + "metadata": {}, + } + ] + + await store._put_batch(payload) + + limiter.acquire_async.assert_awaited_once_with() + assert len(client.put_calls) == 1 + + +async def test_query_pages_uses_rate_limiter_per_page(): + pages = [ + { + "vectors": [{"key": "k0"}], + "NextToken": "token-1", + }, + { + "vectors": [{"key": "k1"}], + "NextToken": "token-2", + }, + { + "vectors": [{"key": "k2"}], + }, + ] + + config = DynavecConfig( + vector_bucket="test-bucket", + index="test-index", + table="test-table", + dimension=4, + query_rps=10, + ) + + client = FakeS3Client(pages) + session = FakeSession(client) + + async with AsyncS3VectorsStore(config, session) as store: + limiter = AsyncMock() + store._query_limiter = limiter + + result_pages = [ + page + async for page in store.query_pages( + [0.1, 0.2, 0.3, 0.4], + top_k=3, + ) + ] + + assert [len(page) for page in result_pages] == [1, 1, 1] + assert limiter.acquire_async.await_count == 3 + + +async def test_get_vectors_batches_keys(): + client = FakeS3Client() + session = FakeSession(client) + + keys = [f"key-{i}" for i in range(250)] + + async with AsyncS3VectorsStore(_config(), session) as store: + result = await store.get_vectors( + keys, + return_metadata=True, + ) + + assert [len(call["keys"]) for call in client.get_calls] == [ + 100, + 100, + 50, + ] + + assert all(call["returnData"] is True for call in client.get_calls) + assert all(call["returnMetadata"] is True for call in client.get_calls) + + assert len(result) == 250 + assert set(result) == set(keys) diff --git a/tests/test_credentials.py b/tests/test_credentials.py new file mode 100644 index 0000000..1f26749 --- /dev/null +++ b/tests/test_credentials.py @@ -0,0 +1,155 @@ +"""Tests for AWS credential/session resolution.""" + +from types import SimpleNamespace + +import pytest + +import dynavec.credentials as credentials_module +from dynavec.credentials import AWSCredentials, resolve_async_session +from dynavec.exceptions import MissingDependencyError + + +class _FakeAioboto3: + def __init__(self) -> None: + self.calls = [] + self.session = object() + + def Session(self, **kwargs): + self.calls.append(kwargs) + return self.session + + +def test_resolve_async_session_prefers_explicit_session(monkeypatch): + explicit_session = object() + + def fail_import(name): + raise AssertionError(f"unexpected import: {name}") + + monkeypatch.setattr(credentials_module.importlib, "import_module", fail_import) + + result = resolve_async_session(None, explicit_session) + + assert result is explicit_session + + +def test_resolve_async_session_missing_dependency(monkeypatch): + def missing_import(name): + raise ImportError(name) + + monkeypatch.setattr(credentials_module.importlib, "import_module", missing_import) + + with pytest.raises(MissingDependencyError) as exc_info: + resolve_async_session(None, None) + + assert exc_info.value.feature == "AsyncDynavec" + assert exc_info.value.package == "aioboto3" + assert exc_info.value.extra == "async" + + +def test_resolve_async_session_default_chain(monkeypatch): + fake_aioboto3 = _FakeAioboto3() + monkeypatch.setattr( + credentials_module.importlib, + "import_module", + lambda name: fake_aioboto3, + ) + + result = resolve_async_session(None, None) + + assert result is fake_aioboto3.session + assert fake_aioboto3.calls == [{}] + + +def test_resolve_async_session_static_credentials(monkeypatch): + fake_aioboto3 = _FakeAioboto3() + monkeypatch.setattr( + credentials_module.importlib, + "import_module", + lambda name: fake_aioboto3, + ) + + credentials = AWSCredentials( + access_key_id="access", + secret_access_key="secret", + session_token="token", + region="us-east-1", + ) + + result = resolve_async_session(credentials, None) + + assert result is fake_aioboto3.session + assert fake_aioboto3.calls == [ + { + "aws_access_key_id": "access", + "aws_secret_access_key": "secret", + "aws_session_token": "token", + "region_name": "us-east-1", + } + ] + + +def test_resolve_async_session_profile(monkeypatch): + fake_aioboto3 = _FakeAioboto3() + monkeypatch.setattr( + credentials_module.importlib, + "import_module", + lambda name: fake_aioboto3, + ) + + credentials = AWSCredentials( + profile_name="dev", + region="us-west-2", + ) + + resolve_async_session(credentials, None) + + assert fake_aioboto3.calls == [ + { + "profile_name": "dev", + "region_name": "us-west-2", + } + ] + + +def test_resolve_async_session_assume_role_reuses_sync_resolution(monkeypatch): + fake_aioboto3 = _FakeAioboto3() + monkeypatch.setattr( + credentials_module.importlib, + "import_module", + lambda name: fake_aioboto3, + ) + + frozen = SimpleNamespace( + access_key="assumed-access", + secret_key="assumed-secret", + token="assumed-token", + ) + resolved = SimpleNamespace( + get_frozen_credentials=lambda: frozen, + ) + sync_session = SimpleNamespace( + get_credentials=lambda: resolved, + ) + + monkeypatch.setattr( + AWSCredentials, + "session", + lambda self: sync_session, + ) + + credentials = AWSCredentials( + assume_role_arn="arn:aws:iam::123456789012:role/test", + region="us-east-1", + ) + + result = resolve_async_session(credentials, None) + + assert result is fake_aioboto3.session + assert fake_aioboto3.calls == [ + { + "aws_access_key_id": "assumed-access", + "aws_secret_access_key": "assumed-secret", + "aws_session_token": "assumed-token", + "region_name": "us-east-1", + } + ] diff --git a/tests/test_utils_transforms.py b/tests/test_utils_transforms.py index 58b97ec..b6931a1 100644 --- a/tests/test_utils_transforms.py +++ b/tests/test_utils_transforms.py @@ -5,7 +5,7 @@ import pytest from dynavec.transforms import TransformContext, TransformPipeline, as_pipeline -from dynavec.utils import TokenBucket, chunked, is_retryable, retry +from dynavec.utils import TokenBucket, async_retry, chunked, is_retryable, retry def test_chunked_generator(): @@ -161,4 +161,98 @@ def test_token_bucket_does_not_exceed_capacity(): assert bucket.tokens == 1 bucket.acquire() - assert bucket.tokens == 1 \ No newline at end of file + assert bucket.tokens == 1 + + +async def test_token_bucket_async_waits_for_token( + monkeypatch, +): + delays = [] + + async def fake_sleep(delay): + delays.append(delay) + + monkeypatch.setattr( + "dynavec.utils.asyncio.sleep", + fake_sleep, + ) + + with patch( + "dynavec.utils.time.monotonic", + side_effect=[ + 100.0, + 100.0, + 100.0, + 100.05, + 100.11, + ], + ): + bucket = TokenBucket( + rate=10, + capacity=2, + ) + + await bucket.acquire_async() + await bucket.acquire_async() + await bucket.acquire_async() + + assert delays == [pytest.approx(0.05)] + assert bucket.tokens == pytest.approx(0.1) + + +async def test_async_retry_retries_then_succeeds(): + calls = {"n": 0} + + class Throttle(Exception): + response = {"Error": {"Code": "ThrottlingException"}} + + @async_retry(max_attempts=5, base_delay=0.0) + async def flaky(): + calls["n"] += 1 + if calls["n"] < 3: + raise Throttle() + return "ok" + + assert await flaky() == "ok" + assert calls["n"] == 3 + + +async def test_async_retry_does_not_retry_non_retryable(): + calls = {"n": 0} + + @async_retry(max_attempts=5, base_delay=0.0) + async def boom(): + calls["n"] += 1 + raise ValueError("nope") + + with pytest.raises(ValueError): + await boom() + + assert calls["n"] == 1 + + +async def test_async_retry_uses_retry_delay(monkeypatch): + calls = {"n": 0} + delays = [] + + class Throttle(Exception): + response = {"Error": {"Code": "ThrottlingException"}} + + async def fake_sleep(delay): + delays.append(delay) + + monkeypatch.setattr("dynavec.utils.asyncio.sleep", fake_sleep) + + @async_retry( + max_attempts=3, + retry_delay=lambda exc: 3.0, + ) + async def flaky(): + calls["n"] += 1 + if calls["n"] == 1: + raise Throttle() + return "ok" + + assert await flaky() == "ok" + assert calls["n"] == 2 + assert delays == [3.0]