diff --git a/src/dynavec/stores/s3vectors.py b/src/dynavec/stores/s3vectors.py index 6ca25a7..79a0232 100644 --- a/src/dynavec/stores/s3vectors.py +++ b/src/dynavec/stores/s3vectors.py @@ -11,6 +11,7 @@ import logging import time from collections.abc import Iterator +from concurrent.futures import ThreadPoolExecutor, as_completed from typing import Any import numpy as np @@ -65,19 +66,31 @@ def _put_batch(self, payload: list[dict]) -> None: vectors=payload, ) + def put_vectors( self, vectors: list[tuple[str, list[float], Metadata]], + max_workers: int = 8, ) -> None: - """Insert/overwrite (key, vector, filterable_metadata) triples.""" + """Insert/overwrite (key, vector, filterable_metadata) triples in parallelized batches.""" + if not vectors: + return + t0 = time.perf_counter() - for start in range(0, len(vectors), _PUT_LIMIT): - chunk = vectors[start : start + _PUT_LIMIT] + chunks = [vectors[i : i + _PUT_LIMIT] for i in range(0, len(vectors), _PUT_LIMIT)] + + def _upload_chunk(chunk: list[tuple[str, list[float], Metadata]]) -> None: payload = [ {"key": key, "data": {"float32": _f32(vec)}, "metadata": meta} for key, vec, meta in chunk ] self._put_batch(payload) + + with ThreadPoolExecutor(max_workers=min(max_workers, len(chunks))) as executor: + futures = [executor.submit(_upload_chunk, chunk) for chunk in chunks] + for future in as_completed(futures): + future.result() # Ensures any exception raised in thread is re-raised + log_store_event( self._logger, "s3vectors.put_vectors", diff --git a/tests/test_s3vectors.py b/tests/test_s3vectors.py index e4cf0e2..3d48a7d 100644 --- a/tests/test_s3vectors.py +++ b/tests/test_s3vectors.py @@ -1,3 +1,4 @@ +import time from unittest.mock import MagicMock import pytest @@ -164,3 +165,52 @@ def test_query_pages_invalid_page_size(): def test_config_rejects_non_positive_top_k_page_size(): with pytest.raises(ValueError, match="top_k_page_size must be a positive integer"): _config(top_k_page_size=0) + + + + + + + +def test_put_vectors_parallelization(): + store, _ = _store_with_pages([]) + + + vectors = [ + (f"key_{i}", [0.1] * 128, {"tag": "test"}) + for i in range(1500) + ] + + def mock_put_batch(payload): + time.sleep(0.1) + + store._put_batch = MagicMock(side_effect=mock_put_batch) + + t0 = time.perf_counter() + store.put_vectors(vectors) + duration = time.perf_counter() - t0 + + + assert duration < 0.5 + assert store._put_batch.call_count == 3 + + +def test_put_vectors_error_propagation(): + store, _ = _store_with_pages([]) + + vectors = [ + (f"key_{i}", [0.1] * 128, {"tag": "test"}) + for i in range(1000) + ] + + def mock_put_batch_with_error(payload): + if payload[0]["key"] == "key_0": + raise RuntimeError("S3 API Failure") + + store._put_batch = MagicMock(side_effect=mock_put_batch_with_error) + + with pytest.raises(RuntimeError, match="S3 API Failure"): + store.put_vectors(vectors) + + +