From bce464db56c0374c9dcbc198a8e477e3713c6a6f Mon Sep 17 00:00:00 2001 From: eunwoo song Date: Thu, 17 Sep 2026 12:52:26 +0900 Subject: [PATCH 1/8] ci: add mypy type checking --- .github/workflows/ci.yml | 13 +++++++ CONTRIBUTING.md | 24 +++++++----- Makefile | 12 ++++-- pyproject.toml | 61 +++++++++++++++++++++++++++++++ src/dynavec/cache.py | 2 +- src/dynavec/client.py | 24 +++++++----- src/dynavec/config.py | 7 +++- src/dynavec/embeddings/bedrock.py | 3 +- src/dynavec/embeddings/mistral.py | 4 +- src/dynavec/embeddings/openai.py | 4 +- src/dynavec/embeddings/voyage.py | 4 +- src/dynavec/eval/runner.py | 2 +- src/dynavec/graph.py | 2 +- src/dynavec/ingest.py | 4 +- src/dynavec/metrics.py | 6 ++- src/dynavec/provisioning.py | 10 +++-- src/dynavec/quantization.py | 52 ++++++++++++++------------ src/dynavec/stores/dynamodb.py | 4 +- src/dynavec/stores/s3vectors.py | 2 +- tests/test_client_inmemory.py | 17 +++++++++ 20 files changed, 190 insertions(+), 67 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 7213a86..4d0f7bb 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -6,6 +6,19 @@ on: pull_request: jobs: + typecheck: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - name: Install uv + uses: astral-sh/setup-uv@v5 + with: + python-version: "3.12" + - name: Install (with dev extras) + run: uv pip install -e ".[dev]" + - name: Type check + run: uv run --no-sync mypy + test: runs-on: ubuntu-latest strategy: diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 9274341..e975674 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -35,7 +35,7 @@ source .venv/bin/activate # Windows: .venv\Scripts\Activate.ps1 make install # editable install with dev + ingest extras # 4. Verify everything works -make check # lint +make check # lint + static type checks make test # run the offline test suite # 5. See every available command @@ -52,9 +52,10 @@ AI agents working in this repo should use the standardized `make` targets: ```bash make help # See all available targets make install # Install dev environment (editable) -make check # Quick health check (lint) +make check # Quick health check (lint + static type checks) +make typecheck # Run mypy across the package make test # Run all unit tests (offline) -make run-ci # Full CI pipeline locally (lint + test) +make run-ci # Full CI pipeline locally (lint + types + test) make format # Auto-format and fix lint issues make docs # Regenerate the static docs site make clean # Remove caches and build artifacts @@ -139,10 +140,12 @@ pre-commit run --all-files # run across the whole tree once | `make install` | Editable install with dev + ingest extras | | `make install-all` | Editable install with **all** extras | | `make format` | Auto-format (`ruff format`) and auto-fix lint (`ruff --fix`) | -| `make lint` / `make check` | Lint with ruff — mirrors CI exactly | +| `make lint` | Lint with ruff | +| `make typecheck` | Type-check the package with mypy | +| `make check` | Run lint and static type checks | | `make test` | Run the offline unit suite (`pytest -q`) | | `make test-live` | Opt-in end-to-end test against **real AWS** (costs money) | -| `make run-ci` | The full CI pipeline locally: lint + test | +| `make run-ci` | The full CI pipeline locally: lint + types + test | | `make docs` | Regenerate the static docs site | | `make clean` | Remove caches and build artifacts | @@ -160,7 +163,7 @@ git checkout -b feat/your-feature # branch off development # ... make changes ... make format # tidy up -make check # lint +make check # lint + static type checks make test # verify ``` @@ -205,9 +208,10 @@ uv run --no-sync pytest tests/test_cache.py -k "jitter" -v - **Style/linting:** [ruff](https://docs.astral.sh/ruff/) (config in `pyproject.toml`, rule sets `E, F, I, UP, B`, line length 100). `make format` fixes most issues automatically. -- **Type hints:** dynavec ships a `py.typed` marker — please add type hints to new public APIs. -- **CI** (`.github/workflows/ci.yml`) runs on every push/PR: **ruff + pytest across Python - 3.9, 3.11, and 3.12**. `make run-ci` reproduces it locally. +- **Type hints:** dynavec ships a `py.typed` marker. Mypy checks the package in a dedicated + CI job; run it locally with `make typecheck`. +- **CI** (`.github/workflows/ci.yml`) runs on every push/PR: **mypy**, plus ruff and + pytest across Python 3.9, 3.11, and 3.12. `make run-ci` reproduces it locally. --- @@ -233,7 +237,7 @@ ci: add Python 3.13 to the matrix **Review checklist:** - [ ] Tests pass (`make test`) -- [ ] Lint passes (`make check`) +- [ ] Lint and static type checks pass (`make check`) - [ ] New/changed behavior has tests - [ ] Docs updated where relevant diff --git a/Makefile b/Makefile index 674e73d..35d29f1 100644 --- a/Makefile +++ b/Makefile @@ -1,10 +1,10 @@ # dynavec developer commands. Run `make help` to see everything. # -# These wrap the exact tools CI uses (uv + ruff + pytest), so `make run-ci` +# These wrap the exact tools CI uses (uv + ruff + mypy + pytest), so `make run-ci` # locally is the same pipeline that runs on your PR. .DEFAULT_GOAL := help -.PHONY: help install install-all format lint check test test-live docs clean run-ci +.PHONY: help install install-all format lint typecheck check test test-live docs clean run-ci PY_DIRS := src benchmarks tests CI_DIRS := src benchmarks # what CI lints (keep in sync with .github/workflows/ci.yml) @@ -26,7 +26,10 @@ format: ## Auto-format and fix lint issues (ruff format + ruff --fix) lint: ## Lint without changing files (mirrors CI exactly) uv run --no-sync ruff check $(CI_DIRS) -check: lint ## Quick health check (lint, no tests) +typecheck: ## Check package annotations with mypy + uv run --no-sync mypy + +check: lint typecheck ## Quick health check (lint + types, no tests) test: ## Run the unit test suite (offline, no AWS needed) uv run --no-sync pytest -q @@ -37,8 +40,9 @@ test-live: ## Run the opt-in end-to-end test against real AWS (costs money) docs: ## Regenerate the static docs site into opensource/dynavec/docs/ uv run --no-sync python tools/build_docs.py -run-ci: ## Run the full CI pipeline locally (lint + test) +run-ci: ## Run the full CI pipeline locally (lint + types + test) $(MAKE) lint + $(MAKE) typecheck $(MAKE) test clean: ## Remove caches and build artifacts diff --git a/pyproject.toml b/pyproject.toml index 5f4d0de..70c0c46 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -75,6 +75,7 @@ dev = [ "pytest-cov>=5.0", "hypothesis>=6.100", "moto[dynamodb]>=5.0", + "mypy>=1.11,<2.0", "ruff>=0.6", "pre-commit>=3.5", ] @@ -109,3 +110,63 @@ target-version = "py39" [tool.ruff.lint] select = ["E", "F", "I", "UP", "B"] ignore = ["E501"] + +[tool.mypy] +python_version = "3.12" +files = ["src/dynavec"] +check_untyped_defs = true +no_implicit_optional = true +pretty = true +show_error_codes = true +warn_redundant_casts = true +warn_unreachable = true +warn_unused_ignores = true + +# Optional integrations deliberately remain importable only when their extras are +# installed. Keep those dependency boundaries explicit while still checking all +# dynavec code that uses them. +[[tool.mypy.overrides]] +module = [ + "bs4", + "bs4.*", + "boto3", + "boto3.*", + "botocore", + "botocore.*", + "crewai", + "crewai.*", + "docx", + "docx.*", + "dspy", + "dspy.*", + "google", + "google.*", + "langchain_core", + "langchain_core.*", + "llama_index", + "llama_index.*", + "matplotlib", + "matplotlib.*", + "mcp", + "mcp.*", + "mistralai", + "mistralai.*", + "openai", + "openai.*", + "openpyxl", + "openpyxl.*", + "pptx", + "pptx.*", + "pypdf", + "pypdf.*", + "redis", + "redis.*", + "requests", + "requests.*", + "sentence_transformers", + "sentence_transformers.*", + "voyageai", + "voyageai.*", + "yaml", +] +ignore_missing_imports = true diff --git a/src/dynavec/cache.py b/src/dynavec/cache.py index 9d725f6..8938ea3 100644 --- a/src/dynavec/cache.py +++ b/src/dynavec/cache.py @@ -222,7 +222,7 @@ def __init__( botocore_config = config.botocore_config() if botocore_config is not None: resource_kwargs["config"] = botocore_config - self._table = session.resource("dynamodb", **resource_kwargs).Table(config.table) # type: ignore[arg-type] + self._table = session.resource("dynamodb", **resource_kwargs).Table(config.table) self.ttl_seconds = ttl_seconds self.ttl_jitter_seconds = ttl_jitter_seconds diff --git a/src/dynavec/client.py b/src/dynavec/client.py index 9dfce97..1bd74b1 100644 --- a/src/dynavec/client.py +++ b/src/dynavec/client.py @@ -27,7 +27,7 @@ import json import time -from collections.abc import Iterator +from collections.abc import Iterator, Sequence from concurrent.futures import ThreadPoolExecutor from pathlib import Path from typing import TYPE_CHECKING, Any, TextIO, Union @@ -211,9 +211,10 @@ def _prepare( "Some documents have no vector and no embedder is configured. " "Pass an embedder to Dynavec(...) or provide precomputed vectors." ) - texts = [t for _, t in to_embed] - if any(t is None for t in texts): + 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] vectors = self.embedder.embed_documents(texts) for (idx, _), vec in zip(to_embed, vectors): docs[idx].vector = vec @@ -221,9 +222,14 @@ def _prepare( # 3) validate + build payloads s3_payload, ddb_payload, ids, hot_payload = [], [], [], [] for d in docs: - if len(d.vector) != self.config.dimension: + 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(d.vector)}, " + f"Document {d.id!r} vector has dimension {len(vector)}, " f"expected {self.config.dimension}." ) meta = dict(d.metadata) @@ -234,12 +240,12 @@ def _prepare( s3_meta, ddb_meta = split_metadata(meta, self.config, namespace, d.text) # fail before either store is written, not partway through a batch check_item_size(namespace, d.id, d.text, ddb_meta) - s3_payload.append((self._s3_key(namespace, d.id), d.vector, s3_meta)) + 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, d.vector, d.text, meta)) + hot_payload.append((d.id, vector, d.text, meta)) return s3_payload, ddb_payload, ids, hot_payload def _write(self, namespace: str, s3_payload: list, ddb_payload: list) -> None: @@ -252,7 +258,7 @@ def _write(self, namespace: str, s3_payload: list, ddb_payload: list) -> None: def upsert( self, - documents: list[Document | dict] | None = None, + documents: Sequence[Document | dict[str, Any]] | None = None, *, namespace: str = "default", auto_metadata: bool = False, @@ -299,7 +305,7 @@ def update( new_vector = vector if new_vector is None: if text is not None and self.embedder is not None: - new_vector = self.embedder.embed_documents([new_text])[0] + new_vector = self.embedder.embed_documents([text])[0] else: fetched = self._vectors.get_vectors([self._s3_key(namespace, id)]) got = fetched.get(self._s3_key(namespace, id)) diff --git a/src/dynavec/config.py b/src/dynavec/config.py index e743e1f..5d52cfd 100644 --- a/src/dynavec/config.py +++ b/src/dynavec/config.py @@ -8,7 +8,10 @@ from __future__ import annotations from dataclasses import dataclass, field -from typing import Literal +from typing import TYPE_CHECKING, Literal + +if TYPE_CHECKING: + from botocore.config import Config DistanceMetric = Literal["cosine", "euclidean"] @@ -98,7 +101,7 @@ class DynavecConfig: structured_logging: bool = False log_level: str = "INFO" - def botocore_config(self): # type: ignore[no-untyped-def] + def botocore_config(self) -> Config | None: """Return a botocore Config with pool tuning, or None for defaults. Local import keeps the base package cheap (boto3/botocore stay diff --git a/src/dynavec/embeddings/bedrock.py b/src/dynavec/embeddings/bedrock.py index 523fdc5..5c99641 100644 --- a/src/dynavec/embeddings/bedrock.py +++ b/src/dynavec/embeddings/bedrock.py @@ -9,7 +9,7 @@ import base64 import json from pathlib import Path -from typing import BinaryIO +from typing import Any, BinaryIO from .base import Embedder, Vector @@ -52,6 +52,7 @@ def __init__( self._is_cohere = model_id.startswith("cohere.") def _invoke(self, text: str, input_type: str) -> Vector: + body: dict[str, Any] if self._is_titan: body = {"inputText": text} if self._requested_dim is not None: diff --git a/src/dynavec/embeddings/mistral.py b/src/dynavec/embeddings/mistral.py index 117c37b..51e4bd5 100644 --- a/src/dynavec/embeddings/mistral.py +++ b/src/dynavec/embeddings/mistral.py @@ -2,6 +2,8 @@ from __future__ import annotations +from typing import Any + from ..exceptions import MissingDependencyError from .base import Embedder, Vector @@ -53,7 +55,7 @@ def embed_documents(self, texts: list[str]) -> list[Vector]: out: list[Vector] = [] for i in range(0, len(texts), self.batch_size): chunk = texts[i : i + self.batch_size] - kwargs = {"model": self.model, "inputs": chunk} + kwargs: dict[str, Any] = {"model": self.model, "inputs": chunk} if self._requested_dim is not None: kwargs["output_dimension"] = self._requested_dim resp = self._client.embeddings.create(**kwargs) diff --git a/src/dynavec/embeddings/openai.py b/src/dynavec/embeddings/openai.py index 03bb2e6..6f728bd 100644 --- a/src/dynavec/embeddings/openai.py +++ b/src/dynavec/embeddings/openai.py @@ -2,6 +2,8 @@ from __future__ import annotations +from typing import Any + from ..exceptions import MissingDependencyError from .base import Embedder, Vector @@ -50,7 +52,7 @@ def embed_documents(self, texts: list[str]) -> list[Vector]: out: list[Vector] = [] for i in range(0, len(texts), self.batch_size): chunk = texts[i : i + self.batch_size] - kwargs = {"model": self.model, "input": chunk} + kwargs: dict[str, Any] = {"model": self.model, "input": chunk} if self._requested_dim is not None: kwargs["dimensions"] = self._requested_dim resp = self._client.embeddings.create(**kwargs) diff --git a/src/dynavec/embeddings/voyage.py b/src/dynavec/embeddings/voyage.py index 97ffab1..f67747c 100644 --- a/src/dynavec/embeddings/voyage.py +++ b/src/dynavec/embeddings/voyage.py @@ -2,6 +2,8 @@ from __future__ import annotations +from typing import Any + from ..exceptions import MissingDependencyError from .base import Embedder, Vector @@ -65,7 +67,7 @@ def _embed(self, texts: list[str], input_type: str) -> list[Vector]: out: list[Vector] = [] for i in range(0, len(texts), self.batch_size): chunk = texts[i : i + self.batch_size] - kwargs = {"model": self.model, "input_type": input_type} + kwargs: dict[str, Any] = {"model": self.model, "input_type": input_type} if self._requested_dim is not None: kwargs["output_dimension"] = self._requested_dim resp = self._client.embed(chunk, **kwargs) diff --git a/src/dynavec/eval/runner.py b/src/dynavec/eval/runner.py index 787e276..a91c492 100644 --- a/src/dynavec/eval/runner.py +++ b/src/dynavec/eval/runner.py @@ -63,7 +63,7 @@ def evaluate_sample( def run( self, - dataset: Sequence[dict[str, Any] | tuple[str, str | Sequence[str], str]], + dataset: Sequence[dict[str, Any] | Sequence[Any]], run_faithfulness: bool = True, run_answer_relevance: bool = True, ) -> EvalSummary: diff --git a/src/dynavec/graph.py b/src/dynavec/graph.py index 012808c..726e8f7 100644 --- a/src/dynavec/graph.py +++ b/src/dynavec/graph.py @@ -171,7 +171,7 @@ def __init__(self, config: DynavecConfig, boto_session=None) -> None: botocore_config = config.botocore_config() if botocore_config is not None: resource_kwargs["config"] = botocore_config - self._ddb = session.resource("dynamodb", **resource_kwargs) # type: ignore[arg-type] + self._ddb = session.resource("dynamodb", **resource_kwargs) self._table = self._ddb.Table(config.table) @staticmethod diff --git a/src/dynavec/ingest.py b/src/dynavec/ingest.py index 297c67a..15fc781 100644 --- a/src/dynavec/ingest.py +++ b/src/dynavec/ingest.py @@ -122,7 +122,7 @@ def __init__(self, path: str | Path) -> None: self._document_cls = docx.Document def __iter__(self) -> Iterator[Record]: - doc = self._document_cls(self._path) + doc = self._document_cls(str(self._path)) paragraphs = [p.text.strip() for p in doc.paragraphs if p.text and p.text.strip()] if not paragraphs: return @@ -156,7 +156,7 @@ def __init__(self, path: str | Path) -> None: self._presentation_cls = Presentation def __iter__(self) -> Iterator[Record]: - prs = self._presentation_cls(self._path) + prs = self._presentation_cls(str(self._path)) path_str = self._path.as_posix() for slide_num, slide in enumerate(prs.slides, start=1): diff --git a/src/dynavec/metrics.py b/src/dynavec/metrics.py index cb90549..827aff7 100644 --- a/src/dynavec/metrics.py +++ b/src/dynavec/metrics.py @@ -70,11 +70,13 @@ def composite_score( if not weights: raise ValueError("weights must be non-empty") total = 0.0 - acc = None + acc: np.ndarray | None = None for metric, w in weights.items(): s = normalize_scores(score(query, mat, metric)) acc = s * w if acc is None else acc + s * w total += w + if acc is None: # Defensive guard for unusual mapping implementations. + raise ValueError("weights must be non-empty") return acc / (total or 1.0) @@ -84,7 +86,7 @@ def rescore( spec: Metric | dict[str, float], *, normalize: bool = False, -) -> np.ndarray: +) -> tuple[np.ndarray, np.ndarray]: """Return an ordering (indices, best first) for the candidates under ``spec``. ``spec`` is a metric name or a ``{metric: weight}`` combination. diff --git a/src/dynavec/provisioning.py b/src/dynavec/provisioning.py index 61356e4..86c5b45 100644 --- a/src/dynavec/provisioning.py +++ b/src/dynavec/provisioning.py @@ -10,6 +10,8 @@ from __future__ import annotations +from typing import Any + from .config import TEXT_METADATA_KEY, DynavecConfig from .exceptions import ProvisioningError @@ -26,7 +28,7 @@ def ensure_vector_bucket(config: DynavecConfig, boto_session=None) -> None: botocore_config = config.botocore_config() if botocore_config is not None: client_kwargs["config"] = botocore_config - s3v = session.client("s3vectors", **client_kwargs) # type: ignore[arg-type] + s3v = session.client("s3vectors", **client_kwargs) try: s3v.create_vector_bucket(vectorBucketName=config.vector_bucket) except Exception as exc: # noqa: BLE001 @@ -43,7 +45,7 @@ def ensure_index(config: DynavecConfig, boto_session=None) -> None: botocore_config = config.botocore_config() if botocore_config is not None: client_kwargs["config"] = botocore_config - s3v = session.client("s3vectors", **client_kwargs) # type: ignore[arg-type] + s3v = session.client("s3vectors", **client_kwargs) non_filterable = list(config.non_filterable_keys) if config.store_text_in_s3vectors and TEXT_METADATA_KEY not in non_filterable: @@ -75,10 +77,10 @@ def ensure_table(config: DynavecConfig, boto_session=None) -> None: botocore_config = config.botocore_config() if botocore_config is not None: client_kwargs["config"] = botocore_config - ddb = session.client("dynamodb", **client_kwargs) # type: ignore[arg-type] + ddb = session.client("dynamodb", **client_kwargs) try: - create_kwargs = { + create_kwargs: dict[str, Any] = { "TableName": config.table, "AttributeDefinitions": [{"AttributeName": "pk", "AttributeType": "S"}], "KeySchema": [{"AttributeName": "pk", "KeyType": "HASH"}], diff --git a/src/dynavec/quantization.py b/src/dynavec/quantization.py index 9ac1db3..ee05dd7 100644 --- a/src/dynavec/quantization.py +++ b/src/dynavec/quantization.py @@ -92,24 +92,24 @@ def code_size_bytes(self) -> int: # ----------------------------------------------------------------- encode def encode(self, vectors: np.ndarray) -> np.ndarray: - self._check_fitted() + codebooks, dsub = self._fitted_state() x = np.asarray(vectors, dtype=np.float32) if x.ndim == 1: x = x.reshape(1, -1) codes = np.empty((x.shape[0], self.m), dtype=np.uint8) for j in range(self.m): - sub = x[:, j * self._dsub : (j + 1) * self._dsub] - d = ((sub[:, None, :] - self._codebooks[j][None, :, :]) ** 2).sum(axis=2) + sub = x[:, j * dsub : (j + 1) * dsub] + d = ((sub[:, None, :] - codebooks[j][None, :, :]) ** 2).sum(axis=2) codes[:, j] = d.argmin(axis=1) return codes def decode(self, codes: np.ndarray) -> np.ndarray: """Approximate reconstruction from codes.""" - self._check_fitted() + codebooks, dsub = self._fitted_state() codes = np.atleast_2d(codes) - out = np.empty((codes.shape[0], self.m * self._dsub), dtype=np.float32) + out = np.empty((codes.shape[0], self.m * dsub), dtype=np.float32) for j in range(self.m): - out[:, j * self._dsub : (j + 1) * self._dsub] = self._codebooks[j][codes[:, j]] + out[:, j * dsub : (j + 1) * dsub] = codebooks[j][codes[:, j]] return out # --------------------------------------------------------------- distance @@ -119,14 +119,14 @@ def asymmetric_distances(self, query: np.ndarray, codes: np.ndarray) -> np.ndarr Precomputes a per-subspace distance table so scoring N codes is a few table lookups — the reason PQ is fast at scale. """ - self._check_fitted() + codebooks, dsub = self._fitted_state() q = np.asarray(query, dtype=np.float32).reshape(-1) codes = np.atleast_2d(codes) # distance table: (m, ksub) table = np.empty((self.m, self.ksub), dtype=np.float32) for j in range(self.m): - qsub = q[j * self._dsub : (j + 1) * self._dsub] - table[j] = ((self._codebooks[j] - qsub) ** 2).sum(axis=1) + qsub = q[j * dsub : (j + 1) * dsub] + table[j] = ((codebooks[j] - qsub) ** 2).sum(axis=1) # sum table lookups across subspaces dists = np.zeros(codes.shape[0], dtype=np.float32) for j in range(self.m): @@ -151,21 +151,25 @@ def save(self, file: str | Path | BinaryIO) -> None: RuntimeError If the quantizer has not been .fit() yet. """ - self._check_fitted() - payload = { - "format_version": np.array(FORMAT_VERSION, dtype=np.int32), - "m": np.array(self.m, dtype=np.int32), - "nbits": np.array(self.nbits, dtype=np.int32), - "iters": np.array(self.iters, dtype=np.int32), - "seed": np.array(self.seed, dtype=np.int32), - "dsub": np.array(self._dsub, dtype=np.int32), - "codebooks": self._codebooks, - } + codebooks, dsub = self._fitted_state() + + def write_payload(target: BinaryIO) -> None: + np.savez( + target, + format_version=np.array(FORMAT_VERSION, dtype=np.int32), + m=np.array(self.m, dtype=np.int32), + nbits=np.array(self.nbits, dtype=np.int32), + iters=np.array(self.iters, dtype=np.int32), + seed=np.array(self.seed, dtype=np.int32), + dsub=np.array(dsub, dtype=np.int32), + codebooks=codebooks, + ) + if isinstance(file, (str, Path)): with open(file, "wb") as f: - np.savez(f, **payload) + write_payload(f) else: - np.savez(file, **payload) + write_payload(file) @classmethod def load(cls, file: str | Path | BinaryIO) -> ProductQuantizer: @@ -212,10 +216,10 @@ def load(cls, file: str | Path | BinaryIO) -> ProductQuantizer: except Exception as exc: raise ValueError(f"Failed to load ProductQuantizer: {exc}") from exc - def _check_fitted(self) -> None: - if self._codebooks is None: + def _fitted_state(self) -> tuple[np.ndarray, int]: + if self._codebooks is None or self._dsub is None: raise RuntimeError("ProductQuantizer must be .fit() before use") - + return self._codebooks, self._dsub @dataclass diff --git a/src/dynavec/stores/dynamodb.py b/src/dynavec/stores/dynamodb.py index 0db5529..103cfdf 100644 --- a/src/dynavec/stores/dynamodb.py +++ b/src/dynavec/stores/dynamodb.py @@ -125,7 +125,7 @@ def __init__(self, config: DynavecConfig, boto_session=None) -> None: botocore_config = config.botocore_config() if botocore_config is not None: resource_kwargs["config"] = botocore_config - self._ddb = session.resource("dynamodb", **resource_kwargs) # type: ignore[arg-type] + self._ddb = session.resource("dynamodb", **resource_kwargs) self._table = self._ddb.Table(config.table) self._logger = logging.getLogger("dynavec.stores.dynamodb") @@ -182,7 +182,7 @@ def get_many(self, namespace: str, ids: list[str]) -> dict[str, dict[str, Any]]: for start in range(0, len(keys), _BATCH_GET_LIMIT): chunk = keys[start : start + _BATCH_GET_LIMIT] - request = {self._config.table: {"Keys": chunk}} + request: dict[str, Any] | None = {self._config.table: {"Keys": chunk}} while request: resp = self._ddb.batch_get_item(RequestItems=request) for item in resp["Responses"].get(self._config.table, []): diff --git a/src/dynavec/stores/s3vectors.py b/src/dynavec/stores/s3vectors.py index 6ca25a7..50b611f 100644 --- a/src/dynavec/stores/s3vectors.py +++ b/src/dynavec/stores/s3vectors.py @@ -48,7 +48,7 @@ def __init__(self, config: DynavecConfig, boto_session=None) -> None: botocore_config = config.botocore_config() if botocore_config is not None: client_kwargs["config"] = botocore_config - self._client = session.client("s3vectors", **client_kwargs) # type: ignore[arg-type] + self._client = session.client("s3vectors", **client_kwargs) self._logger = logging.getLogger("dynavec.stores.s3vectors") def get_index(self) -> dict: diff --git a/tests/test_client_inmemory.py b/tests/test_client_inmemory.py index f8ec948..009c0ab 100644 --- a/tests/test_client_inmemory.py +++ b/tests/test_client_inmemory.py @@ -249,6 +249,11 @@ def embed_documents(self, texts): return [[0.1, 0.2] for _ in texts] +class MissingOutputEmbedder(HashEmbedder): + def embed_documents(self, texts): + return [] + + def test_embedder_output_dimension_mismatch_raises(db): from dynavec.exceptions import DimensionMismatchError @@ -261,6 +266,18 @@ def test_embedder_output_dimension_mismatch_raises(db): db.upsert([Document(id="x", text="hello")]) +def test_embedder_missing_output_raises_configuration_error(db): + from dynavec.exceptions import ConfigurationError + + db.embedder = MissingOutputEmbedder(8) + + with pytest.raises( + ConfigurationError, + match=r"Embedder did not return a vector for document 'x'", + ): + db.upsert([Document(id="x", text="hello")]) + + def test_auto_metadata_switch(db): db.upsert([Document(id="1", text="hello world")], auto_metadata=True) got = db.get(["1"])[0] From 58d25ec0db6b217f3f0e8d0ed669ce50cc35064a Mon Sep 17 00:00:00 2001 From: eunwoo song Date: Thu, 17 Sep 2026 13:09:45 +0900 Subject: [PATCH 2/8] ci: enforce strict mypy checks --- CONTRIBUTING.md | 6 +- pyproject.toml | 6 +- src/dynavec/cache.py | 110 +++++++++++++++--- src/dynavec/cli.py | 4 +- src/dynavec/client.py | 6 +- src/dynavec/credentials.py | 11 +- src/dynavec/dashboard.py | 37 ++++-- src/dynavec/embeddings/__init__.py | 4 +- src/dynavec/embeddings/bedrock.py | 14 +-- src/dynavec/embeddings/gemini.py | 6 +- src/dynavec/embeddings/ollama.py | 4 +- .../embeddings/sentence_transformers.py | 4 +- src/dynavec/eval/base.py | 22 +++- src/dynavec/eval/judges.py | 2 +- src/dynavec/eval/runner.py | 10 +- src/dynavec/graph.py | 9 +- src/dynavec/hot.py | 3 +- src/dynavec/ingest.py | 17 ++- src/dynavec/integrations/dspy.py | 24 +++- src/dynavec/integrations/langchain.py | 47 +++++--- src/dynavec/integrations/llamaindex.py | 38 +++--- src/dynavec/integrations/tools.py | 11 +- src/dynavec/mcp/server.py | 20 ++-- src/dynavec/metrics.py | 8 +- src/dynavec/namespace.py | 25 ++-- src/dynavec/provisioning.py | 13 ++- src/dynavec/retrieval.py | 4 +- src/dynavec/spfresh.py | 4 +- src/dynavec/stores/dynamodb.py | 2 +- src/dynavec/stores/s3vectors.py | 11 +- src/dynavec/transforms.py | 15 ++- src/dynavec/utils.py | 4 +- 32 files changed, 336 insertions(+), 165 deletions(-) diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index e975674..dd4f332 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -141,7 +141,7 @@ pre-commit run --all-files # run across the whole tree once | `make install-all` | Editable install with **all** extras | | `make format` | Auto-format (`ruff format`) and auto-fix lint (`ruff --fix`) | | `make lint` | Lint with ruff | -| `make typecheck` | Type-check the package with mypy | +| `make typecheck` | Type-check the package with mypy's strict mode | | `make check` | Run lint and static type checks | | `make test` | Run the offline unit suite (`pytest -q`) | | `make test-live` | Opt-in end-to-end test against **real AWS** (costs money) | @@ -208,8 +208,8 @@ uv run --no-sync pytest tests/test_cache.py -k "jitter" -v - **Style/linting:** [ruff](https://docs.astral.sh/ruff/) (config in `pyproject.toml`, rule sets `E, F, I, UP, B`, line length 100). `make format` fixes most issues automatically. -- **Type hints:** dynavec ships a `py.typed` marker. Mypy checks the package in a dedicated - CI job; run it locally with `make typecheck`. +- **Type hints:** dynavec ships a `py.typed` marker. Mypy checks the package in strict mode + in a dedicated CI job; run it locally with `make typecheck`. - **CI** (`.github/workflows/ci.yml`) runs on every push/PR: **mypy**, plus ruff and pytest across Python 3.9, 3.11, and 3.12. `make run-ci` reproduces it locally. diff --git a/pyproject.toml b/pyproject.toml index 70c0c46..6267bab 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -114,13 +114,9 @@ ignore = ["E501"] [tool.mypy] python_version = "3.12" files = ["src/dynavec"] -check_untyped_defs = true -no_implicit_optional = true +strict = true pretty = true show_error_codes = true -warn_redundant_casts = true -warn_unreachable = true -warn_unused_ignores = true # Optional integrations deliberately remain importable only when their extras are # installed. Keep those dependency boundaries explicit while still checking all diff --git a/src/dynavec/cache.py b/src/dynavec/cache.py index 8938ea3..5a36220 100644 --- a/src/dynavec/cache.py +++ b/src/dynavec/cache.py @@ -23,25 +23,30 @@ import time from abc import ABC, abstractmethod from collections import OrderedDict -from typing import TYPE_CHECKING +from collections.abc import Sequence +from typing import TYPE_CHECKING, Any, cast import numpy as np +from .config import DynavecConfig from .exceptions import ConfigurationError, MissingDependencyError from .models import SearchResult if TYPE_CHECKING: from .client import Dynavec +MetadataFilter = dict[str, Any] +CacheEntry = tuple[np.ndarray, list[SearchResult], int] -def _signature(namespace: str, top_k: int, filter: dict | None) -> str: + +def _signature(namespace: str, top_k: int, filter: MetadataFilter | None) -> str: payload = json.dumps( {"ns": namespace, "k": top_k, "f": filter or {}}, sort_keys=True ) return hashlib.sha256(payload.encode()).hexdigest()[:24] -def _vec_key(query_vector, ndigits: int = 3) -> str: +def _vec_key(query_vector: Sequence[float], ndigits: int = 3) -> str: rounded = [round(float(x), ndigits) for x in query_vector] return hashlib.sha256(json.dumps(rounded).encode()).hexdigest()[:24] @@ -60,7 +65,7 @@ def _deserialize(blob: str) -> list[SearchResult]: ] -def _deep_size(value, seen: set[int] | None = None) -> int: +def _deep_size(value: object, seen: set[int] | None = None) -> int: """Estimate an object graph's resident size without requiring serialization.""" if seen is None: seen = set() @@ -88,11 +93,24 @@ def __init__(self) -> None: self.misses: int = 0 @abstractmethod - def get(self, namespace, query_vector, top_k, filter) -> list[SearchResult] | None: + def get( + self, + namespace: str, + query_vector: Sequence[float], + top_k: int, + filter: MetadataFilter | None, + ) -> list[SearchResult] | None: ... @abstractmethod - def put(self, namespace, query_vector, top_k, filter, results) -> None: + def put( + self, + namespace: str, + query_vector: Sequence[float], + top_k: int, + filter: MetadataFilter | None, + results: list[SearchResult], + ) -> None: ... def stats(self) -> dict[str, int | float]: @@ -136,7 +154,7 @@ def __init__( self._size_bytes = 0 self._lru: OrderedDict[tuple[str, str], None] = OrderedDict() # key: signature -> OrderedDict[vec_key -> (unit_vec, results, size_bytes)] - self._buckets: dict[str, OrderedDict[str, tuple]] = {} + self._buckets: dict[str, OrderedDict[str, CacheEntry]] = {} @property def size_bytes(self) -> int: @@ -145,9 +163,15 @@ def size_bytes(self) -> int: @staticmethod def _unit(v: np.ndarray) -> np.ndarray: - return v / (np.linalg.norm(v) + 1e-12) + return cast(np.ndarray, v / (np.linalg.norm(v) + 1e-12)) - def get(self, namespace, query_vector, top_k, filter): + def get( + self, + namespace: str, + query_vector: Sequence[float], + top_k: int, + filter: MetadataFilter | None, + ) -> list[SearchResult] | None: sig = _signature(namespace, top_k, filter) bucket = self._buckets.get(sig) if not bucket: @@ -168,7 +192,14 @@ def get(self, namespace, query_vector, top_k, filter): self.misses += 1 return None - def put(self, namespace, query_vector, top_k, filter, results): + def put( + self, + namespace: str, + query_vector: Sequence[float], + top_k: int, + filter: MetadataFilter | None, + results: list[SearchResult], + ) -> None: sig = _signature(namespace, top_k, filter) bucket = self._buckets.setdefault(sig, OrderedDict()) vk = _vec_key(query_vector) @@ -206,8 +237,8 @@ class DynamoDBCache(BaseCache): def __init__( self, - config, - boto_session=None, + config: DynavecConfig, + boto_session: Any | None = None, ttl_seconds: int = 3600, ttl_jitter_seconds: int = 0, ) -> None: @@ -227,10 +258,21 @@ def __init__( self.ttl_jitter_seconds = ttl_jitter_seconds @staticmethod - def _pk(namespace, query_vector, top_k, filter) -> str: + def _pk( + namespace: str, + query_vector: Sequence[float], + top_k: int, + filter: MetadataFilter | None, + ) -> str: return f"__cache__#{_signature(namespace, top_k, filter)}#{_vec_key(query_vector)}" - def get(self, namespace, query_vector, top_k, filter): + def get( + self, + namespace: str, + query_vector: Sequence[float], + top_k: int, + filter: MetadataFilter | None, + ) -> list[SearchResult] | None: resp = self._table.get_item( Key={"pk": self._pk(namespace, query_vector, top_k, filter)} ) @@ -244,7 +286,14 @@ def get(self, namespace, query_vector, top_k, filter): self.hits += 1 return _deserialize(item["results"]) - def put(self, namespace, query_vector, top_k, filter, results): + def put( + self, + namespace: str, + query_vector: Sequence[float], + top_k: int, + filter: MetadataFilter | None, + results: list[SearchResult], + ) -> None: jitter = random.randint(0, self.ttl_jitter_seconds) self._table.put_item( Item={ @@ -259,7 +308,12 @@ def put(self, namespace, query_vector, top_k, filter, results): class RedisCache(BaseCache): """Shared exact-match cache on Redis / AWS ElastiCache.""" - def __init__(self, url: str = "redis://localhost:6379/0", ttl_seconds: int = 3600, client=None): + def __init__( + self, + url: str = "redis://localhost:6379/0", + ttl_seconds: int = 3600, + client: Any | None = None, + ) -> None: super().__init__() if client is not None: self._r = client @@ -272,10 +326,21 @@ def __init__(self, url: str = "redis://localhost:6379/0", ttl_seconds: int = 360 self.ttl_seconds = ttl_seconds @staticmethod - def _key(namespace, query_vector, top_k, filter) -> str: + def _key( + namespace: str, + query_vector: Sequence[float], + top_k: int, + filter: MetadataFilter | None, + ) -> str: return f"dynavec:{_signature(namespace, top_k, filter)}:{_vec_key(query_vector)}" - def get(self, namespace, query_vector, top_k, filter): + def get( + self, + namespace: str, + query_vector: Sequence[float], + top_k: int, + filter: MetadataFilter | None, + ) -> list[SearchResult] | None: blob = self._r.get(self._key(namespace, query_vector, top_k, filter)) if blob: self.hits += 1 @@ -283,7 +348,14 @@ def get(self, namespace, query_vector, top_k, filter): self.misses += 1 return None - def put(self, namespace, query_vector, top_k, filter, results): + def put( + self, + namespace: str, + query_vector: Sequence[float], + top_k: int, + filter: MetadataFilter | None, + results: list[SearchResult], + ) -> None: self._r.set( self._key(namespace, query_vector, top_k, filter), _serialize(results), diff --git a/src/dynavec/cli.py b/src/dynavec/cli.py index c689cc0..b5b3438 100644 --- a/src/dynavec/cli.py +++ b/src/dynavec/cli.py @@ -4,6 +4,8 @@ import argparse import sys +from collections.abc import Callable +from typing import Any from .client import Dynavec from .config import DynavecConfig @@ -212,7 +214,7 @@ def _import(args: argparse.Namespace) -> int: return 1 -def _check(label: str, callback) -> bool: +def _check(label: str, callback: Callable[[], str]) -> bool: try: detail = callback() except Exception as exc: # noqa: BLE001 diff --git a/src/dynavec/client.py b/src/dynavec/client.py index 1bd74b1..0fb2a01 100644 --- a/src/dynavec/client.py +++ b/src/dynavec/client.py @@ -27,10 +27,12 @@ import json import time -from collections.abc import Iterator, Sequence +from collections.abc import Callable, Iterable, Iterator, Sequence from concurrent.futures import ThreadPoolExecutor +from functools import partial from pathlib import Path -from typing import TYPE_CHECKING, Any, TextIO, Union +from types import TracebackType +from typing import TYPE_CHECKING, Any, Optional, TextIO, Union if TYPE_CHECKING: from .cache import BaseCache diff --git a/src/dynavec/credentials.py b/src/dynavec/credentials.py index ce0cdc4..794e998 100644 --- a/src/dynavec/credentials.py +++ b/src/dynavec/credentials.py @@ -13,6 +13,7 @@ from __future__ import annotations from dataclasses import dataclass +from typing import Any @dataclass(frozen=True) @@ -28,14 +29,14 @@ class AWSCredentials: role_session_name: str = "dynavec" external_id: str | None = None - def session(self): + def session(self) -> Any: """Build a ``boto3.Session`` from this credential description.""" import boto3 if self.assume_role_arn: return self._assume_role_session(boto3) - kwargs = {} + kwargs: dict[str, str] = {} if self.access_key_id and self.secret_access_key: kwargs["aws_access_key_id"] = self.access_key_id kwargs["aws_secret_access_key"] = self.secret_access_key @@ -47,9 +48,9 @@ def session(self): kwargs["region_name"] = self.region return boto3.Session(**kwargs) - def _assume_role_session(self, boto3): + def _assume_role_session(self, boto3: Any) -> Any: # Base session used only to call STS. - base_kwargs = {} + base_kwargs: dict[str, str] = {} if self.access_key_id and self.secret_access_key: base_kwargs["aws_access_key_id"] = self.access_key_id base_kwargs["aws_secret_access_key"] = self.secret_access_key @@ -79,7 +80,7 @@ def _assume_role_session(self, boto3): ) -def resolve_session(credentials: AWSCredentials | None, boto_session): +def resolve_session(credentials: AWSCredentials | None, boto_session: Any | None) -> Any: """Pick a boto3 session: explicit session > credentials > default chain.""" if boto_session is not None: return boto_session diff --git a/src/dynavec/dashboard.py b/src/dynavec/dashboard.py index 0853096..da32bb1 100644 --- a/src/dynavec/dashboard.py +++ b/src/dynavec/dashboard.py @@ -23,6 +23,7 @@ import json from http.server import BaseHTTPRequestHandler, HTTPServer +from typing import Any from urllib.parse import parse_qs, urlparse from .telemetry import TelemetryRecorder, aggregate, aggregate_eval @@ -208,12 +209,17 @@ """ -def _make_handler(recorder: TelemetryRecorder): +def _make_handler(recorder: TelemetryRecorder) -> type[BaseHTTPRequestHandler]: class Handler(BaseHTTPRequestHandler): - def log_message(self, *a): # quiet + def log_message(self, format: str, *args: Any) -> None: # quiet pass - def _send(self, code, body, ctype="application/json"): + def _send( + self, + code: int, + body: str | bytes, + ctype: str = "application/json", + ) -> None: data = body.encode() if isinstance(body, str) else body self.send_response(code) self.send_header("Content-Type", ctype) @@ -221,14 +227,16 @@ def _send(self, code, body, ctype="application/json"): self.end_headers() self.wfile.write(data) - def do_GET(self): + def do_GET(self) -> None: parsed = urlparse(self.path) path, qs = parsed.path, parse_qs(parsed.query) if path == "/" or path == "/index.html": - return self._send(200, _INDEX_HTML, "text/html; charset=utf-8") + self._send(200, _INDEX_HTML, "text/html; charset=utf-8") + return if path == "/api/metrics": window = int(qs.get("window", ["3600"])[0]) - return self._send(200, json.dumps(aggregate(recorder.snapshot(), window))) + self._send(200, json.dumps(aggregate(recorder.snapshot(), window))) + return if path == "/api/traces": evs = recorder.events( limit=int(qs.get("limit", ["100"])[0]), @@ -236,15 +244,19 @@ def do_GET(self): namespace=(qs.get("namespace", [None])[0] or None), status=(qs.get("status", [None])[0] or None), ) - return self._send(200, json.dumps([e.to_dict() for e in evs])) + self._send(200, json.dumps([e.to_dict() for e in evs])) + return if path.startswith("/api/trace/"): ev = recorder.get(path.rsplit("/", 1)[-1]) if ev is None: - return self._send(404, json.dumps({"error": "not found"})) - return self._send(200, json.dumps(ev.to_dict())) + self._send(404, json.dumps({"error": "not found"})) + return + self._send(200, json.dumps(ev.to_dict())) + return if path == "/api/eval/summary": window = int(qs.get("window", ["86400"])[0]) - return self._send(200, json.dumps(aggregate_eval(recorder.snapshot(), window))) + self._send(200, json.dumps(aggregate_eval(recorder.snapshot(), window))) + return if path == "/api/eval/runs": limit = int(qs.get("limit", ["50"])[0]) all_events = recorder.events(limit=recorder._events.maxlen or 10000) @@ -257,8 +269,9 @@ def do_GET(self): or e.eval_mrr is not None or e.eval_ndcg is not None ][:limit] - return self._send(200, json.dumps(eval_runs)) - return self._send(404, json.dumps({"error": "not found"})) + self._send(200, json.dumps(eval_runs)) + return + self._send(404, json.dumps({"error": "not found"})) return Handler diff --git a/src/dynavec/embeddings/__init__.py b/src/dynavec/embeddings/__init__.py index 05ab99d..29efe2b 100644 --- a/src/dynavec/embeddings/__init__.py +++ b/src/dynavec/embeddings/__init__.py @@ -7,7 +7,7 @@ from __future__ import annotations -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any from .base import Embedder, Vector from .cache import EmbeddingCache, InMemoryCache @@ -59,7 +59,7 @@ } -def __getattr__(name: str): # PEP 562 lazy submodule attribute access +def __getattr__(name: str) -> Any: # PEP 562 lazy submodule attribute access if name in _LAZY: import importlib diff --git a/src/dynavec/embeddings/bedrock.py b/src/dynavec/embeddings/bedrock.py index 5c99641..3c73caf 100644 --- a/src/dynavec/embeddings/bedrock.py +++ b/src/dynavec/embeddings/bedrock.py @@ -9,7 +9,7 @@ import base64 import json from pathlib import Path -from typing import Any, BinaryIO +from typing import Any, BinaryIO, cast from .base import Embedder, Vector @@ -39,7 +39,7 @@ def __init__( model_id: str = "amazon.titan-embed-text-v2:0", region: str | None = None, dimension: int | None = None, - boto_session=None, + boto_session: Any | None = None, ) -> None: import boto3 # local import keeps base import cheap @@ -65,8 +65,8 @@ def _invoke(self, text: str, input_type: str) -> Vector: resp = self._client.invoke_model(modelId=self.model_id, body=json.dumps(body)) payload = json.loads(resp["body"].read()) if self._is_cohere: - return payload["embeddings"][0] - return payload["embedding"] + return cast(Vector, payload["embeddings"][0]) + return cast(Vector, payload["embedding"]) def embed_documents(self, texts: list[str]) -> list[Vector]: return [self._invoke(t, "search_document") for t in texts] @@ -100,7 +100,7 @@ def __init__( self, dimension: int = 1024, region: str | None = None, - boto_session=None, + boto_session: Any | None = None, ) -> None: import boto3 @@ -141,7 +141,7 @@ def _invoke( if text is None and image_base64 is None: raise ValueError("At least one of 'text' or 'image' must be provided.") - body: dict = { + body: dict[str, Any] = { "embeddingConfig": { "outputEmbeddingLength": self.dimension, } @@ -156,7 +156,7 @@ def _invoke( body=json.dumps(body), ) payload = json.loads(resp["body"].read()) - return payload["embedding"] + return cast(Vector, payload["embedding"]) def embed_documents(self, texts: list[str]) -> list[Vector]: """Embed a batch of text documents.""" diff --git a/src/dynavec/embeddings/gemini.py b/src/dynavec/embeddings/gemini.py index d80ade2..e14f0d4 100644 --- a/src/dynavec/embeddings/gemini.py +++ b/src/dynavec/embeddings/gemini.py @@ -2,6 +2,8 @@ from __future__ import annotations +from typing import cast + from ..exceptions import MissingDependencyError from .base import Embedder, Vector @@ -46,11 +48,11 @@ def embed_documents(self, texts: list[str]) -> list[Vector]: resp = self._genai.embed_content( model=self.model, content=text, task_type="retrieval_document" ) - out.append(resp["embedding"]) + out.append(cast(Vector, resp["embedding"])) return out def embed_query(self, text: str) -> Vector: resp = self._genai.embed_content( model=self.model, content=text, task_type="retrieval_query" ) - return resp["embedding"] + return cast(Vector, resp["embedding"]) diff --git a/src/dynavec/embeddings/ollama.py b/src/dynavec/embeddings/ollama.py index 1937db0..954b6b3 100644 --- a/src/dynavec/embeddings/ollama.py +++ b/src/dynavec/embeddings/ollama.py @@ -6,7 +6,7 @@ import os import urllib.error import urllib.request -from typing import Any +from typing import Any, cast from .base import Embedder, Vector @@ -66,7 +66,7 @@ def _post(self, endpoint: str, payload: dict[str, Any]) -> dict[str, Any]: ) try: with urllib.request.urlopen(req, timeout=self.timeout) as resp: - return json.loads(resp.read().decode("utf-8")) + return cast(dict[str, Any], json.loads(resp.read().decode("utf-8"))) except urllib.error.URLError as exc: raise RuntimeError(f"Ollama embedding request failed ({url}): {exc}") from exc diff --git a/src/dynavec/embeddings/sentence_transformers.py b/src/dynavec/embeddings/sentence_transformers.py index ca9174e..30bf319 100644 --- a/src/dynavec/embeddings/sentence_transformers.py +++ b/src/dynavec/embeddings/sentence_transformers.py @@ -2,6 +2,8 @@ from __future__ import annotations +from typing import cast + from ..exceptions import MissingDependencyError from .base import Embedder, Vector @@ -48,4 +50,4 @@ def embed_documents(self, texts: list[str]) -> list[Vector]: normalize_embeddings=self.normalize, convert_to_numpy=True, ) - return arr.tolist() + return cast(list[Vector], arr.tolist()) diff --git a/src/dynavec/eval/base.py b/src/dynavec/eval/base.py index b8c076d..60ac10d 100644 --- a/src/dynavec/eval/base.py +++ b/src/dynavec/eval/base.py @@ -6,7 +6,15 @@ import re from abc import ABC, abstractmethod from dataclasses import asdict, dataclass, field -from typing import Any +from typing import Any, cast + + +def _load_json_object(candidate: str) -> dict[str, Any] | None: + parsed: object = json.loads(candidate) + if isinstance(parsed, dict): + # JSON object keys are always strings. + return cast(dict[str, Any], parsed) + return None def extract_json(text: str) -> dict[str, Any]: @@ -23,8 +31,8 @@ def extract_json(text: str) -> dict[str, Any]: # 1. Try direct json parsing try: - data = json.loads(clean) - if isinstance(data, dict): + data = _load_json_object(clean) + if data is not None: return data except json.JSONDecodeError: pass @@ -33,7 +41,9 @@ def extract_json(text: str) -> dict[str, Any]: fenced_match = re.search(r"```(?:json)?\s*(\{.*?\})\s*```", clean, re.DOTALL) if fenced_match: try: - return json.loads(fenced_match.group(1)) + data = _load_json_object(fenced_match.group(1)) + if data is not None: + return data except json.JSONDecodeError: pass @@ -43,7 +53,9 @@ def extract_json(text: str) -> dict[str, Any]: if first_brace != -1 and last_brace != -1 and last_brace > first_brace: candidate = clean[first_brace : last_brace + 1] try: - return json.loads(candidate) + data = _load_json_object(candidate) + if data is not None: + return data except json.JSONDecodeError: pass diff --git a/src/dynavec/eval/judges.py b/src/dynavec/eval/judges.py index 335b9d6..4a59501 100644 --- a/src/dynavec/eval/judges.py +++ b/src/dynavec/eval/judges.py @@ -72,7 +72,7 @@ def __init__( model_id: str = "anthropic.claude-3-haiku-20240307-v1:0", region: str | None = None, temperature: float = 0.0, - boto_session=None, + boto_session: Any | None = None, ) -> None: session = resolve_session(None, boto_session) kwargs: dict[str, Any] = {} diff --git a/src/dynavec/eval/runner.py b/src/dynavec/eval/runner.py index a91c492..e2454d6 100644 --- a/src/dynavec/eval/runner.py +++ b/src/dynavec/eval/runner.py @@ -4,11 +4,13 @@ from collections.abc import Sequence from dataclasses import asdict, dataclass, field -from typing import Any +from typing import Any, Union from .base import BaseJudge, RAGEvalResult from .metrics import evaluate_rag +EvalTriplet = tuple[str, Union[str, Sequence[str]], str] + @dataclass class EvalSummary: @@ -63,7 +65,7 @@ def evaluate_sample( def run( self, - dataset: Sequence[dict[str, Any] | Sequence[Any]], + dataset: Sequence[dict[str, Any] | EvalTriplet], run_faithfulness: bool = True, run_answer_relevance: bool = True, ) -> EvalSummary: @@ -84,10 +86,8 @@ def run( q = str(item.get("query", "")) c = item.get("context", []) a = str(item.get("answer", "")) - elif isinstance(item, (tuple, list)) and len(item) >= 3: - q, c, a = item[0], item[1], item[2] else: - continue + q, c, a = item res = self.evaluate_sample( query=q, diff --git a/src/dynavec/graph.py b/src/dynavec/graph.py index 726e8f7..189365b 100644 --- a/src/dynavec/graph.py +++ b/src/dynavec/graph.py @@ -25,7 +25,7 @@ import re from collections import deque -from typing import Any +from typing import Any, cast from .config import DynavecConfig from .utils import ( @@ -162,7 +162,7 @@ def _render_dot(nodes: list[str], edges: list[tuple[str, str, str]]) -> str: class GraphStore: """DynamoDB-backed property graph sharing the dynavec document table.""" - def __init__(self, config: DynavecConfig, boto_session=None) -> None: + def __init__(self, config: DynavecConfig, boto_session: Any | None = None) -> None: import boto3 session = boto_session or boto3.Session() @@ -289,9 +289,10 @@ def _drop_edges(self, ns: str, src: str, dst: str, relation: str | None = None) # ------------------------------------------------------------------ reads @retry() - def get_node(self, ns: str, entity_id: str) -> dict | None: + def get_node(self, ns: str, entity_id: str) -> dict[str, Any] | None: resp = self._table.get_item(Key={"pk": self._node_pk(ns, entity_id)}) - return resp.get("Item") + item = resp.get("Item") + return cast(dict[str, Any], item) if item is not None else None def neighbors(self, ns: str, entity_id: str, relation: str | None = None) -> list[str]: node = self.get_node(ns, entity_id) diff --git a/src/dynavec/hot.py b/src/dynavec/hot.py index 7d84f3b..44a5b23 100644 --- a/src/dynavec/hot.py +++ b/src/dynavec/hot.py @@ -19,6 +19,7 @@ import threading from collections import OrderedDict +from collections.abc import Callable from typing import Any from .config import DynavecConfig @@ -27,7 +28,7 @@ # --- MongoDB-style metadata matcher (mirrors the S3 Vectors filter dialect) --- -_COMPARATORS = { +_COMPARATORS: dict[str, Callable[[Any, Any], bool]] = { "$eq": lambda a, b: a == b, "$ne": lambda a, b: a != b, "$gt": lambda a, b: a is not None and a > b, diff --git a/src/dynavec/ingest.py b/src/dynavec/ingest.py index 15fc781..79b4304 100644 --- a/src/dynavec/ingest.py +++ b/src/dynavec/ingest.py @@ -17,7 +17,7 @@ from __future__ import annotations import hashlib -from collections.abc import Iterable, Iterator +from collections.abc import Callable, Iterable, Iterator from dataclasses import dataclass, field from pathlib import Path from typing import Any @@ -25,6 +25,7 @@ from .client import Dynavec from .exceptions import MissingDependencyError from .models import Document +from .transforms import Transform, TransformPipeline from .utils import chunked Metadata = dict[str, Any] @@ -60,7 +61,7 @@ def chunk_text(text: str, chunk_size: int = 1000, overlap: int = 150) -> Iterato class IterableSource: """Wrap a list/iterable of records or dicts as a Source.""" - def __init__(self, records: Iterable) -> None: + def __init__(self, records: Iterable[Record | dict[str, Any]]) -> None: self._records = records def __iter__(self) -> Iterator[Record]: @@ -391,12 +392,16 @@ class MCPResourceSource: Optional predicate ``(uri) -> bool`` to select which resources to pull. """ - def __init__(self, session, uri_filter=None) -> None: + def __init__( + self, + session: Any, + uri_filter: Callable[[str], bool] | None = None, + ) -> None: self._session = session self._uri_filter = uri_filter @staticmethod - def _extract_text(contents) -> str: + def _extract_text(contents: Any) -> str: # MCP read_resource returns an object/list of content parts; grab text. parts = getattr(contents, "contents", contents) if isinstance(parts, (list, tuple)): @@ -435,14 +440,14 @@ def __iter__(self) -> Iterator[Record]: def ingest( db: Dynavec, - source: Iterable, + source: Iterable[Record | dict[str, Any]], *, namespace: str = "default", chunk_size: int = 1000, overlap: int = 150, batch_size: int = 256, auto_metadata: bool = True, - transform=None, + transform: TransformPipeline | Transform | Iterable[Transform] | None = None, ) -> int: """Pull records from ``source``, chunk, embed, and upsert. Returns #chunks. diff --git a/src/dynavec/integrations/dspy.py b/src/dynavec/integrations/dspy.py index 5d0a473..c415da9 100644 --- a/src/dynavec/integrations/dspy.py +++ b/src/dynavec/integrations/dspy.py @@ -12,19 +12,31 @@ from __future__ import annotations -from typing import Any +from typing import TYPE_CHECKING, Any from ..client import Dynavec from ..exceptions import MissingDependencyError -try: - import dspy +if TYPE_CHECKING: from dspy.dsp.utils import dotdict -except ImportError as exc: # pragma: no cover - import guard - raise MissingDependencyError("DynavecRM", "dspy", "dspy") from exc + + class _RetrieveBase: + """Typing boundary for the optional DSPy base class.""" + + k: int + + def __init__(self, k: int) -> None: ... + +else: + try: + import dspy + from dspy.dsp.utils import dotdict + except ImportError as exc: # pragma: no cover - import guard + raise MissingDependencyError("DynavecRM", "dspy", "dspy") from exc + _RetrieveBase = dspy.Retrieve -class DynavecRM(dspy.Retrieve): +class DynavecRM(_RetrieveBase): """DSPy retrieval module backed by a :class:`Dynavec` client.""" def __init__( diff --git a/src/dynavec/integrations/langchain.py b/src/dynavec/integrations/langchain.py index f5f079a..f4127c0 100644 --- a/src/dynavec/integrations/langchain.py +++ b/src/dynavec/integrations/langchain.py @@ -15,24 +15,37 @@ import asyncio import uuid from collections.abc import Iterable -from typing import Any +from typing import TYPE_CHECKING, Any, Protocol from ..client import Dynavec from ..embeddings.base import Embedder from ..exceptions import MissingDependencyError from ..models import Document as DVDocument -try: +if TYPE_CHECKING: from langchain_core.documents import Document as LCDocument - from langchain_core.vectorstores import VectorStore -except ImportError as exc: # pragma: no cover - import guard - raise MissingDependencyError("DynavecVectorStore", "langchain-core", "langchain") from exc + + class _VectorStoreBase: + """Typing boundary for the optional LangChain base class.""" + +else: + try: + from langchain_core.documents import Document as LCDocument + from langchain_core.vectorstores import VectorStore as _VectorStoreBase + except ImportError as exc: # pragma: no cover - import guard + raise MissingDependencyError("DynavecVectorStore", "langchain-core", "langchain") from exc + + +class _LCEmbeddings(Protocol): + def embed_documents(self, texts: list[str]) -> list[list[float]]: ... + + def embed_query(self, text: str) -> list[float]: ... class _LCEmbeddingsAdapter(Embedder): """Wrap a LangChain ``Embeddings`` object as a dynavec ``Embedder``.""" - def __init__(self, lc_embeddings, dimension: int) -> None: + def __init__(self, lc_embeddings: _LCEmbeddings, dimension: int) -> None: self._lc = lc_embeddings self.dimension = dimension @@ -43,7 +56,7 @@ def embed_query(self, text: str) -> list[float]: return self._lc.embed_query(text) -class DynavecVectorStore(VectorStore): +class DynavecVectorStore(_VectorStoreBase): """A thin LangChain VectorStore backed by a :class:`Dynavec` client.""" def __init__(self, client: Dynavec, namespace: str = "default") -> None: @@ -51,13 +64,13 @@ def __init__(self, client: Dynavec, namespace: str = "default") -> None: self._namespace = namespace @property - def embeddings(self): # LangChain introspects this + def embeddings(self) -> Any: # LangChain introspects this return self._client.embedder def add_texts( self, texts: Iterable[str], - metadatas: list[dict] | None = None, + metadatas: list[dict[str, Any]] | None = None, ids: list[str] | None = None, **kwargs: Any, ) -> list[str]: @@ -78,7 +91,7 @@ def delete(self, ids: list[str] | None = None, **kwargs: Any) -> bool | None: return True def similarity_search( - self, query: str, k: int = 4, filter: dict | None = None, **kwargs: Any + self, query: str, k: int = 4, filter: dict[str, Any] | None = None, **kwargs: Any ) -> list[LCDocument]: results = self._client.search( query, top_k=k, namespace=self._namespace, filter=filter @@ -92,7 +105,7 @@ async def asimilarity_search( self, query: str, k: int = 4, - filter: dict | None = None, + filter: dict[str, Any] | None = None, **kwargs: Any, ) -> list[LCDocument]: results = await asyncio.to_thread( @@ -111,7 +124,7 @@ async def asimilarity_search( ] def similarity_search_with_score( - self, query: str, k: int = 4, filter: dict | None = None, **kwargs: Any + self, query: str, k: int = 4, filter: dict[str, Any] | None = None, **kwargs: Any ) -> list[tuple[LCDocument, float]]: results = self._client.search( query, top_k=k, namespace=self._namespace, filter=filter @@ -128,7 +141,7 @@ async def asimilarity_search_with_score( self, query: str, k: int = 4, - filter: dict | None = None, + filter: dict[str, Any] | None = None, **kwargs: Any, ) -> list[tuple[LCDocument, float]]: results = await asyncio.to_thread( @@ -155,7 +168,7 @@ def max_marginal_relevance_search( k: int = 4, fetch_k: int = 20, lambda_mult: float = 0.5, - filter: dict | None = None, + filter: dict[str, Any] | None = None, **kwargs: Any, ) -> list[LCDocument]: results = self._client.search( @@ -177,7 +190,7 @@ async def amax_marginal_relevance_search( k: int = 4, fetch_k: int = 20, lambda_mult: float = 0.5, - filter: dict | None = None, + filter: dict[str, Any] | None = None, **kwargs: Any, ) -> list[LCDocument]: results = await asyncio.to_thread( @@ -201,8 +214,8 @@ async def amax_marginal_relevance_search( def from_texts( cls, texts: list[str], - embedding, - metadatas: list[dict] | None = None, + embedding: _LCEmbeddings, + metadatas: list[dict[str, Any]] | None = None, *, client: Dynavec | None = None, namespace: str = "default", diff --git a/src/dynavec/integrations/llamaindex.py b/src/dynavec/integrations/llamaindex.py index c2ca8d1..920556b 100644 --- a/src/dynavec/integrations/llamaindex.py +++ b/src/dynavec/integrations/llamaindex.py @@ -14,26 +14,36 @@ from __future__ import annotations -from typing import Any +from typing import TYPE_CHECKING, Any from ..client import Dynavec from ..exceptions import MissingDependencyError from ..models import Document as DVDocument -try: +if TYPE_CHECKING: from llama_index.core.schema import BaseNode, TextNode - from llama_index.core.vector_stores.types import ( - BasePydanticVectorStore, - VectorStoreQuery, - VectorStoreQueryResult, - ) -except ImportError as exc: # pragma: no cover - import guard - raise MissingDependencyError( - "DynavecLlamaStore", "llama-index-core", "all" - ) from exc - - -class DynavecLlamaStore(BasePydanticVectorStore): + from llama_index.core.vector_stores.types import VectorStoreQuery, VectorStoreQueryResult + + class _BasePydanticVectorStore: + """Typing boundary for the optional LlamaIndex base class.""" + +else: + try: + from llama_index.core.schema import BaseNode, TextNode + from llama_index.core.vector_stores.types import ( + BasePydanticVectorStore as _BasePydanticVectorStore, + ) + from llama_index.core.vector_stores.types import ( + VectorStoreQuery, + VectorStoreQueryResult, + ) + except ImportError as exc: # pragma: no cover - import guard + raise MissingDependencyError( + "DynavecLlamaStore", "llama-index-core", "all" + ) from exc + + +class DynavecLlamaStore(_BasePydanticVectorStore): """Minimal LlamaIndex vector store backed by a :class:`Dynavec` client.""" stores_text: bool = True diff --git a/src/dynavec/integrations/tools.py b/src/dynavec/integrations/tools.py index 1bd0f74..a57960c 100644 --- a/src/dynavec/integrations/tools.py +++ b/src/dynavec/integrations/tools.py @@ -9,10 +9,11 @@ from __future__ import annotations import json -from collections.abc import Sequence -from typing import Any, Callable +from collections.abc import Callable, Sequence +from typing import Any from ..client import Dynavec +from ..models import SearchResult from ..namespace import NamespaceView from ..retrievers import QueryExpansionRetriever @@ -22,8 +23,8 @@ def make_retriever_fn( *, top_k: int = 4, namespace: str = "default", - filter: dict | None = None, - rescore=None, + filter: dict[str, Any] | None = None, + rescore: str | dict[str, float] | None = None, join: str = "\n\n", include_scores: bool = False, ) -> Callable[[str], str]: @@ -35,7 +36,7 @@ def make_retriever_fn( if isinstance(source, QueryExpansionRetriever) and rescore is not None: raise ValueError("rescore is not supported with query-expansion retrievers") - def _search(query: str): + def _search(query: str) -> list[SearchResult]: if isinstance(source, QueryExpansionRetriever): return source.search(query, top_k=top_k, filter=filter) if isinstance(source, NamespaceView): diff --git a/src/dynavec/mcp/server.py b/src/dynavec/mcp/server.py index 30bbe2c..af2d9df 100644 --- a/src/dynavec/mcp/server.py +++ b/src/dynavec/mcp/server.py @@ -7,14 +7,15 @@ import os import sys from collections.abc import Mapping -from typing import Any +from typing import Any, cast from ..client import Dynavec -from ..config import DynavecConfig +from ..config import DistanceMetric, DynavecConfig +from ..embeddings.base import Embedder from ..exceptions import ConfigurationError, MissingDependencyError -def _resolve_embedder(env: Mapping[str, str]): +def _resolve_embedder(env: Mapping[str, str]) -> Embedder | None: """Instantiate an embedder from environment variables or return None.""" embedder_type = (env.get("DYNAVEC_EMBEDDER") or "").strip().lower() model = env.get("DYNAVEC_EMBEDDER_MODEL") @@ -104,7 +105,7 @@ def client_from_env(env: Mapping[str, str] | None = None) -> Dynavec: index=str(index), table=str(table), dimension=dimension, - distance_metric=distance_metric, # type: ignore[arg-type] + distance_metric=cast(DistanceMetric, distance_metric), region=region, filterable_keys=filterable_keys, ) @@ -123,7 +124,7 @@ def _format_hits(hits: list[Any], query: str, context_label: str = "") -> str: return "\n".join(lines).strip() -def create_mcp_server(db: Dynavec | None = None, name: str = "dynavec"): +def create_mcp_server(db: Dynavec | None = None, name: str = "dynavec") -> Any: """Create and configure a FastMCP server exposing dynavec search and graph search tools.""" try: from mcp.server.fastmcp import FastMCP @@ -135,7 +136,6 @@ def create_mcp_server(db: Dynavec | None = None, name: str = "dynavec"): def _get_db() -> Dynavec: return db if db is not None else client_from_env() - @mcp.tool() def dynavec_search( query: str, top_k: int = 5, @@ -164,7 +164,10 @@ def dynavec_search( filter_dict = None if filter_json: try: - filter_dict = json.loads(filter_json) + parsed: object = json.loads(filter_json) + if not isinstance(parsed, dict): + return "Error parsing filter_json: expected a JSON object" + filter_dict = cast(dict[str, Any], parsed) except Exception as exc: return f"Error parsing filter_json: {exc}" @@ -179,7 +182,6 @@ def dynavec_search( ) return _format_hits(hits, query, f" in namespace '{namespace}'") - @mcp.tool() def dynavec_graph_search( query: str, seed_entities: list[str], @@ -212,6 +214,8 @@ def dynavec_graph_search( ) return _format_hits(hits, query, f" (seeds: {seed_entities}, hops: {hops})") + mcp.tool()(dynavec_search) + mcp.tool()(dynavec_graph_search) return mcp diff --git a/src/dynavec/metrics.py b/src/dynavec/metrics.py index 827aff7..67b68de 100644 --- a/src/dynavec/metrics.py +++ b/src/dynavec/metrics.py @@ -14,6 +14,8 @@ from __future__ import annotations +from typing import cast + import numpy as np Metric = str # "cosine" | "dot" | "euclidean" | "manhattan" @@ -34,13 +36,13 @@ def score(query: np.ndarray, mat: np.ndarray, metric: Metric) -> np.ndarray: if metric == "cosine": qn = q / (np.linalg.norm(q) + 1e-12) mn = m / (np.linalg.norm(m, axis=1, keepdims=True) + 1e-12) - return mn @ qn + return cast(np.ndarray, mn @ qn) if metric == "euclidean": d = np.linalg.norm(m - q, axis=1) - return 1.0 / (1.0 + d) + return cast(np.ndarray, 1.0 / (1.0 + d)) if metric == "manhattan": d = np.abs(m - q).sum(axis=1) - return 1.0 / (1.0 + d) + return cast(np.ndarray, 1.0 / (1.0 + d)) raise ValueError(f"unknown metric {metric!r}; expected one of {_VALID}") diff --git a/src/dynavec/namespace.py b/src/dynavec/namespace.py index 22b3a2f..88ad4db 100644 --- a/src/dynavec/namespace.py +++ b/src/dynavec/namespace.py @@ -8,9 +8,10 @@ from __future__ import annotations +from collections.abc import Iterator, Sequence from typing import TYPE_CHECKING, Any -from .models import SearchResult +from .models import Document, SearchResult, UpsertResult if TYPE_CHECKING: from .client import Dynavec @@ -29,23 +30,29 @@ def __init__(self, db: Dynavec, namespace: str) -> None: def namespace(self) -> str: return self._ns - def upsert(self, documents, **kw) -> Any: + def upsert( + self, + documents: Sequence[Document | dict[str, Any]] | None, + **kw: Any, + ) -> UpsertResult: return self._db.upsert(documents, namespace=self._ns, **kw) - def update(self, *args, **kw) -> Any: - return self._db.update(*args, namespace=self._ns, **kw) + def update(self, id: str, **kw: Any) -> UpsertResult: + return self._db.update(id, namespace=self._ns, **kw) - def search(self, query: str | None = None, **kw) -> list[SearchResult]: + def search(self, query: str | None = None, **kw: Any) -> list[SearchResult]: return self._db.search(query, namespace=self._ns, **kw) - def search_stream(self, query: str | None = None, **kw): + def search_stream( + self, query: str | None = None, **kw: Any + ) -> Iterator[SearchResult]: yield from self._db.search_stream(query, namespace=self._ns, **kw) - def get(self, ids, **kw) -> list[SearchResult]: + def get(self, ids: list[str], **kw: Any) -> list[SearchResult]: return self._db.get(ids, namespace=self._ns, **kw) - def delete(self, ids, **kw) -> None: - return self._db.delete(ids, namespace=self._ns, **kw) + def delete(self, ids: list[str], **kw: Any) -> None: + self._db.delete(ids, namespace=self._ns, **kw) def as_multiquery_retriever( self, generate_queries=None, *, llm_generate_queries=None, **kw diff --git a/src/dynavec/provisioning.py b/src/dynavec/provisioning.py index 86c5b45..2caebdf 100644 --- a/src/dynavec/provisioning.py +++ b/src/dynavec/provisioning.py @@ -10,17 +10,18 @@ from __future__ import annotations -from typing import Any +from typing import Any, cast from .config import TEXT_METADATA_KEY, DynavecConfig from .exceptions import ProvisioningError def _client_error_code(exc: Exception) -> str: - return getattr(exc, "response", {}).get("Error", {}).get("Code", "") + response = cast(dict[str, Any], getattr(exc, "response", {})) + return str(response.get("Error", {}).get("Code", "")) -def ensure_vector_bucket(config: DynavecConfig, boto_session=None) -> None: +def ensure_vector_bucket(config: DynavecConfig, boto_session: Any | None = None) -> None: import boto3 session = boto_session or boto3.Session() @@ -37,7 +38,7 @@ def ensure_vector_bucket(config: DynavecConfig, boto_session=None) -> None: raise ProvisioningError(f"Failed to create vector bucket: {exc}") from exc -def ensure_index(config: DynavecConfig, boto_session=None) -> None: +def ensure_index(config: DynavecConfig, boto_session: Any | None = None) -> None: import boto3 session = boto_session or boto3.Session() @@ -69,7 +70,7 @@ def ensure_index(config: DynavecConfig, boto_session=None) -> None: raise ProvisioningError(f"Failed to create vector index: {exc}") from exc -def ensure_table(config: DynavecConfig, boto_session=None) -> None: +def ensure_table(config: DynavecConfig, boto_session: Any | None = None) -> None: import boto3 session = boto_session or boto3.Session() @@ -101,7 +102,7 @@ def ensure_table(config: DynavecConfig, boto_session=None) -> None: ddb.get_waiter("table_exists").wait(TableName=config.table) -def provision_all(config: DynavecConfig, boto_session=None) -> None: +def provision_all(config: DynavecConfig, boto_session: Any | None = None) -> None: """Create every resource dynavec needs. Idempotent.""" ensure_vector_bucket(config, boto_session) ensure_index(config, boto_session) diff --git a/src/dynavec/retrieval.py b/src/dynavec/retrieval.py index fefe722..0b65f6e 100644 --- a/src/dynavec/retrieval.py +++ b/src/dynavec/retrieval.py @@ -7,6 +7,8 @@ from __future__ import annotations +from typing import cast + import numpy as np from .config import DistanceMetric @@ -139,7 +141,7 @@ def maximal_marginal_relevance( def _norm(x: np.ndarray) -> np.ndarray: n = np.linalg.norm(x, axis=-1, keepdims=True) - return x / np.clip(n, 1e-12, None) + return cast(np.ndarray, x / np.clip(n, 1e-12, None)) qn = _norm(q.reshape(1, -1))[0] mn = _norm(mat) diff --git a/src/dynavec/spfresh.py b/src/dynavec/spfresh.py index 550c211..9bacd3e 100644 --- a/src/dynavec/spfresh.py +++ b/src/dynavec/spfresh.py @@ -14,7 +14,7 @@ import threading import uuid from dataclasses import dataclass -from typing import Any, Callable +from typing import Any, Callable, cast import numpy as np @@ -101,7 +101,7 @@ def _normalize_if_needed(self, vec: np.ndarray) -> np.ndarray: if self.metric == "cosine": norm = np.linalg.norm(vec) if norm > 1e-12: - return (vec / norm).astype(np.float32) + return cast(np.ndarray, (vec / norm).astype(np.float32)) return vec.astype(np.float32) # ------------------------------------------------------------------ diff --git a/src/dynavec/stores/dynamodb.py b/src/dynavec/stores/dynamodb.py index 103cfdf..7e48ec6 100644 --- a/src/dynavec/stores/dynamodb.py +++ b/src/dynavec/stores/dynamodb.py @@ -116,7 +116,7 @@ class DynamoDBStore: _logger = logging.getLogger("dynavec.stores.dynamodb") - def __init__(self, config: DynavecConfig, boto_session=None) -> None: + def __init__(self, config: DynavecConfig, boto_session: Any | None = None) -> None: import boto3 # local import: base import stays cheap session = boto_session or boto3.Session() diff --git a/src/dynavec/stores/s3vectors.py b/src/dynavec/stores/s3vectors.py index 50b611f..a76fbc4 100644 --- a/src/dynavec/stores/s3vectors.py +++ b/src/dynavec/stores/s3vectors.py @@ -39,7 +39,7 @@ def _f32(vector: list[float]) -> list[float]: class S3VectorsStore: _logger = logging.getLogger("dynavec.stores.s3vectors") - def __init__(self, config: DynavecConfig, boto_session=None) -> None: + def __init__(self, config: DynavecConfig, boto_session: Any | None = None) -> None: import boto3 # local import: base import stays cheap session = boto_session or boto3.Session() @@ -58,7 +58,7 @@ def get_index(self) -> dict: ) @retry() - def _put_batch(self, payload: list[dict]) -> None: + def _put_batch(self, payload: list[dict[str, Any]]) -> None: self._client.put_vectors( vectorBucketName=self._config.vector_bucket, indexName=self._config.index, @@ -89,7 +89,12 @@ def put_vectors( ) def _query_kwargs( - self, query_vector, top_k, filter, return_metadata, return_distance + 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, diff --git a/src/dynavec/transforms.py b/src/dynavec/transforms.py index 8b4c541..82c3676 100644 --- a/src/dynavec/transforms.py +++ b/src/dynavec/transforms.py @@ -13,8 +13,9 @@ from __future__ import annotations import json +from collections.abc import Iterable, Iterator from dataclasses import dataclass, field -from typing import Any, Callable +from typing import Any, Callable, cast Vector = list[float] Metadata = dict[str, Any] @@ -45,7 +46,7 @@ def add(self, transform: Transform) -> TransformPipeline: self._transforms.append(transform) return self - def __iter__(self): + def __iter__(self) -> Iterator[Transform]: return iter(self._transforms) def __len__(self) -> int: @@ -57,7 +58,9 @@ def __call__(self, ctx: TransformContext) -> TransformContext: return ctx -def as_pipeline(spec) -> TransformPipeline | None: +def as_pipeline( + spec: TransformPipeline | Transform | Iterable[Transform] | None, +) -> TransformPipeline | None: """Coerce ``None`` / a single callable / a list into a pipeline.""" if spec is None: return None @@ -75,7 +78,7 @@ class LambdaTransform: as JSON and must return the same shape (any subset it wants to change). """ - def __init__(self, function_name: str, session, qualifier: str | None = None) -> None: + def __init__(self, function_name: str, session: Any, qualifier: str | None = None) -> None: self._client = session.client("lambda") self._function_name = function_name self._qualifier = qualifier @@ -97,10 +100,10 @@ def __call__(self, ctx: TransformContext) -> TransformContext: if self._qualifier: kwargs["Qualifier"] = self._qualifier resp = self._client.invoke(**kwargs) - body = json.loads(resp["Payload"].read() or b"{}") + body = cast(dict[str, Any], json.loads(resp["Payload"].read() or b"{}")) # Lambda may return a subset; only overwrite what it provides. ctx.text = body.get("text", ctx.text) ctx.vector = body.get("vector", ctx.vector) if "metadata" in body and body["metadata"] is not None: - ctx.metadata = body["metadata"] + ctx.metadata = cast(Metadata, body["metadata"]) return ctx diff --git a/src/dynavec/utils.py b/src/dynavec/utils.py index 350e132..6914da9 100644 --- a/src/dynavec/utils.py +++ b/src/dynavec/utils.py @@ -79,7 +79,9 @@ def wrapper(*args: Any, **kwargs: Any) -> T: return decorator -def timed(sink: Callable[[str, float], None] | None = None): +def timed( + sink: Callable[[str, float], None] | None = None, +) -> Callable[[Callable[..., T]], Callable[..., T]]: """Decorator: report wall-clock seconds to ``sink(name, seconds)``. Handy for wiring dynavec latency into an agent's tracing/telemetry. From 2a6355475df66e763e090882ebcef22865bbb0e5 Mon Sep 17 00:00:00 2001 From: eunwoo song Date: Thu, 17 Sep 2026 20:52:55 +0900 Subject: [PATCH 3/8] fix: preserve integration types and eval compatibility --- .github/workflows/ci.yml | 4 +- CONTRIBUTING.md | 4 +- pyproject.toml | 3 +- src/dynavec/eval/runner.py | 4 +- src/dynavec/integrations/dspy.py | 24 +++--------- src/dynavec/integrations/langchain.py | 19 +++------ src/dynavec/integrations/llamaindex.py | 51 ++++++++++++------------- tests/test_eval.py | 27 +++++++++++++ tests/typing/integration_inheritance.py | 20 ++++++++++ typings/dspy/__init__.pyi | 1 + typings/dspy/dsp/__init__.pyi | 0 typings/dspy/dsp/utils/__init__.pyi | 6 +++ typings/dspy/retrievers/__init__.pyi | 1 + typings/dspy/retrievers/retrieve.pyi | 9 +++++ 14 files changed, 110 insertions(+), 63 deletions(-) create mode 100644 tests/typing/integration_inheritance.py create mode 100644 typings/dspy/__init__.pyi create mode 100644 typings/dspy/dsp/__init__.pyi create mode 100644 typings/dspy/dsp/utils/__init__.pyi create mode 100644 typings/dspy/retrievers/__init__.pyi create mode 100644 typings/dspy/retrievers/retrieve.pyi diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 4d0f7bb..899f19c 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -14,8 +14,8 @@ jobs: uses: astral-sh/setup-uv@v5 with: python-version: "3.12" - - name: Install (with dev extras) - run: uv pip install -e ".[dev]" + - name: Install (with dev and typed integration extras) + run: uv pip install -e ".[dev,langchain,llamaindex,dspy]" - name: Type check run: uv run --no-sync mypy diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index dd4f332..61cc245 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -65,7 +65,9 @@ make clean # Remove caches and build artifacts - Prefer `make` targets over invoking tools directly — they match CI exactly. - For direct tool calls, use the `uv run --no-sync` prefix (plain `uv run` can trigger a - universal resolve that pulls yanked optional deps). + universal resolve that pulls yanked optional deps). Install the `langchain`, + `llamaindex`, and `dspy` extras before type-checking so adapter inheritance is checked + against the real framework APIs. - Run `make run-ci` before declaring a change complete; it is the same pipeline CI runs. - Prefer editing existing files over creating new ones, and follow the conventions in neighboring modules. diff --git a/pyproject.toml b/pyproject.toml index 6267bab..4c0eca9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -113,7 +113,8 @@ ignore = ["E501"] [tool.mypy] python_version = "3.12" -files = ["src/dynavec"] +files = ["src/dynavec", "tests/typing"] +mypy_path = "typings" strict = true pretty = true show_error_codes = true diff --git a/src/dynavec/eval/runner.py b/src/dynavec/eval/runner.py index e2454d6..3bb083b 100644 --- a/src/dynavec/eval/runner.py +++ b/src/dynavec/eval/runner.py @@ -86,8 +86,10 @@ def run( q = str(item.get("query", "")) c = item.get("context", []) a = str(item.get("answer", "")) + elif isinstance(item, (tuple, list)) and len(item) >= 3: + q, c, a = item[0], item[1], item[2] else: - q, c, a = item + continue res = self.evaluate_sample( query=q, diff --git a/src/dynavec/integrations/dspy.py b/src/dynavec/integrations/dspy.py index c415da9..5d0a473 100644 --- a/src/dynavec/integrations/dspy.py +++ b/src/dynavec/integrations/dspy.py @@ -12,31 +12,19 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any +from typing import Any from ..client import Dynavec from ..exceptions import MissingDependencyError -if TYPE_CHECKING: +try: + import dspy from dspy.dsp.utils import dotdict - - class _RetrieveBase: - """Typing boundary for the optional DSPy base class.""" - - k: int - - def __init__(self, k: int) -> None: ... - -else: - try: - import dspy - from dspy.dsp.utils import dotdict - except ImportError as exc: # pragma: no cover - import guard - raise MissingDependencyError("DynavecRM", "dspy", "dspy") from exc - _RetrieveBase = dspy.Retrieve +except ImportError as exc: # pragma: no cover - import guard + raise MissingDependencyError("DynavecRM", "dspy", "dspy") from exc -class DynavecRM(_RetrieveBase): +class DynavecRM(dspy.Retrieve): """DSPy retrieval module backed by a :class:`Dynavec` client.""" def __init__( diff --git a/src/dynavec/integrations/langchain.py b/src/dynavec/integrations/langchain.py index f4127c0..bafaf2e 100644 --- a/src/dynavec/integrations/langchain.py +++ b/src/dynavec/integrations/langchain.py @@ -15,25 +15,18 @@ import asyncio import uuid from collections.abc import Iterable -from typing import TYPE_CHECKING, Any, Protocol +from typing import Any, Protocol from ..client import Dynavec from ..embeddings.base import Embedder from ..exceptions import MissingDependencyError from ..models import Document as DVDocument -if TYPE_CHECKING: +try: from langchain_core.documents import Document as LCDocument - - class _VectorStoreBase: - """Typing boundary for the optional LangChain base class.""" - -else: - try: - from langchain_core.documents import Document as LCDocument - from langchain_core.vectorstores import VectorStore as _VectorStoreBase - except ImportError as exc: # pragma: no cover - import guard - raise MissingDependencyError("DynavecVectorStore", "langchain-core", "langchain") from exc + from langchain_core.vectorstores import VectorStore +except ImportError as exc: # pragma: no cover - import guard + raise MissingDependencyError("DynavecVectorStore", "langchain-core", "langchain") from exc class _LCEmbeddings(Protocol): @@ -56,7 +49,7 @@ def embed_query(self, text: str) -> list[float]: return self._lc.embed_query(text) -class DynavecVectorStore(_VectorStoreBase): +class DynavecVectorStore(VectorStore): """A thin LangChain VectorStore backed by a :class:`Dynavec` client.""" def __init__(self, client: Dynavec, namespace: str = "default") -> None: diff --git a/src/dynavec/integrations/llamaindex.py b/src/dynavec/integrations/llamaindex.py index 920556b..64682ac 100644 --- a/src/dynavec/integrations/llamaindex.py +++ b/src/dynavec/integrations/llamaindex.py @@ -14,36 +14,28 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any +from collections.abc import Sequence +from typing import Any from ..client import Dynavec from ..exceptions import MissingDependencyError from ..models import Document as DVDocument -if TYPE_CHECKING: +try: from llama_index.core.schema import BaseNode, TextNode - from llama_index.core.vector_stores.types import VectorStoreQuery, VectorStoreQueryResult - - class _BasePydanticVectorStore: - """Typing boundary for the optional LlamaIndex base class.""" - -else: - try: - from llama_index.core.schema import BaseNode, TextNode - from llama_index.core.vector_stores.types import ( - BasePydanticVectorStore as _BasePydanticVectorStore, - ) - from llama_index.core.vector_stores.types import ( - VectorStoreQuery, - VectorStoreQueryResult, - ) - except ImportError as exc: # pragma: no cover - import guard - raise MissingDependencyError( - "DynavecLlamaStore", "llama-index-core", "all" - ) from exc - - -class DynavecLlamaStore(_BasePydanticVectorStore): + from llama_index.core.vector_stores.types import ( + BasePydanticVectorStore, + MetadataFilter, + VectorStoreQuery, + VectorStoreQueryResult, + ) +except ImportError as exc: # pragma: no cover - import guard + raise MissingDependencyError( + "DynavecLlamaStore", "llama-index-core", "all" + ) from exc + + +class DynavecLlamaStore(BasePydanticVectorStore): """Minimal LlamaIndex vector store backed by a :class:`Dynavec` client.""" stores_text: bool = True @@ -53,7 +45,7 @@ class DynavecLlamaStore(_BasePydanticVectorStore): _namespace: str def __init__(self, client: Dynavec, namespace: str = "default") -> None: - super().__init__() + super().__init__(stores_text=True) self._client = client self._namespace = namespace @@ -61,7 +53,7 @@ def __init__(self, client: Dynavec, namespace: str = "default") -> None: def client(self) -> Any: return self._client - def add(self, nodes: list[BaseNode], **kwargs: Any) -> list[str]: + def add(self, nodes: Sequence[BaseNode], **kwargs: Any) -> list[str]: docs = [] for node in nodes: meta = node.metadata or {} @@ -83,7 +75,12 @@ def delete(self, ref_doc_id: str, **kwargs: Any) -> None: def query(self, query: VectorStoreQuery, **kwargs: Any) -> VectorStoreQueryResult: flt = None if query.filters is not None: - flt = {f.key: f.value for f in query.filters.filters} + simple_filters: list[MetadataFilter] = [] + for metadata_filter in query.filters.filters: + if not isinstance(metadata_filter, MetadataFilter): + raise ValueError("Nested LlamaIndex metadata filters are not supported.") + simple_filters.append(metadata_filter) + flt = {f.key: f.value for f in simple_filters} results = self._client.search( vector=query.query_embedding, diff --git a/tests/test_eval.py b/tests/test_eval.py index 0631412..621c0d3 100644 --- a/tests/test_eval.py +++ b/tests/test_eval.py @@ -3,6 +3,7 @@ from __future__ import annotations import json +from typing import Any, cast from unittest.mock import MagicMock, patch import pytest @@ -350,6 +351,32 @@ def test_eval_runner_batch_dataset(self) -> None: assert summary_dict["total_samples"] == 2 assert len(summary_dict["results"]) == 2 + def test_eval_runner_preserves_legacy_sequence_handling(self) -> None: + judge = MockJudge( + responses=[ + {"claims": [{"claim": "tuple", "supported": True}]}, + {"score": 0.9}, + {"claims": [{"claim": "list", "supported": True}]}, + {"score": 0.8}, + ] + ) + runner = EvalRunner(judge=judge) + dataset = cast( + Any, + [ + ("Q tuple", ["C tuple"], "A tuple", "ignored"), + ["Q list", ["C list"], "A list", "ignored"], + ("too", "short"), + None, + ], + ) + + summary = runner.run(dataset) + + assert summary.total_samples == 2 + assert [result.query for result in summary.results] == ["Q tuple", "Q list"] + assert [result.answer for result in summary.results] == ["A tuple", "A list"] + # =========================================================================== # 6. Telemetry & Dashboard API Integration Tests diff --git a/tests/typing/integration_inheritance.py b/tests/typing/integration_inheritance.py new file mode 100644 index 0000000..75cc576 --- /dev/null +++ b/tests/typing/integration_inheritance.py @@ -0,0 +1,20 @@ +import dspy +from langchain_core.vectorstores import VectorStore +from llama_index.core.vector_stores.types import BasePydanticVectorStore + +from dynavec.integrations.dspy import DynavecRM +from dynavec.integrations.langchain import DynavecVectorStore +from dynavec.integrations.llamaindex import DynavecLlamaStore + + +def use_langchain(store: DynavecVectorStore) -> VectorStore: + store.as_retriever() + return store + + +def use_dspy(retriever: DynavecRM) -> dspy.Retrieve: + return retriever + + +def use_llamaindex(store: DynavecLlamaStore) -> BasePydanticVectorStore: + return store diff --git a/typings/dspy/__init__.pyi b/typings/dspy/__init__.pyi new file mode 100644 index 0000000..6f896c7 --- /dev/null +++ b/typings/dspy/__init__.pyi @@ -0,0 +1 @@ +from dspy.retrievers.retrieve import Retrieve as Retrieve diff --git a/typings/dspy/dsp/__init__.pyi b/typings/dspy/dsp/__init__.pyi new file mode 100644 index 0000000..e69de29 diff --git a/typings/dspy/dsp/utils/__init__.pyi b/typings/dspy/dsp/utils/__init__.pyi new file mode 100644 index 0000000..6033e84 --- /dev/null +++ b/typings/dspy/dsp/utils/__init__.pyi @@ -0,0 +1,6 @@ +from typing import Any + +class dotdict(dict[str, Any]): + def __getattr__(self, key: str) -> Any: ... + def __setattr__(self, key: str, value: Any) -> None: ... + def __delattr__(self, key: str) -> None: ... diff --git a/typings/dspy/retrievers/__init__.pyi b/typings/dspy/retrievers/__init__.pyi new file mode 100644 index 0000000..6f896c7 --- /dev/null +++ b/typings/dspy/retrievers/__init__.pyi @@ -0,0 +1 @@ +from dspy.retrievers.retrieve import Retrieve as Retrieve diff --git a/typings/dspy/retrievers/retrieve.pyi b/typings/dspy/retrievers/retrieve.pyi new file mode 100644 index 0000000..58d4492 --- /dev/null +++ b/typings/dspy/retrievers/retrieve.pyi @@ -0,0 +1,9 @@ +from typing import Any + +class Retrieve: + k: int + callbacks: list[Any] + + def __init__(self, k: int = ..., callbacks: list[Any] | None = ...) -> None: ... + def __call__(self, *args: Any, **kwargs: Any) -> Any: ... + def forward(self, query: str, k: int | None = ..., **kwargs: Any) -> Any: ... From 5fbf7d12ea338d1e2990f8567f5fd702e3f7ef5f Mon Sep 17 00:00:00 2001 From: eunwoo song Date: Thu, 17 Sep 2026 20:59:14 +0900 Subject: [PATCH 4/8] ci: align local typecheck dependencies --- .github/workflows/ci.yml | 4 ++-- CONTRIBUTING.md | 17 +++++++++-------- Makefile | 4 ++-- pyproject.toml | 7 +++++++ tests/typing/integration_inheritance.py | 7 +++++++ typings/dspy/__init__.pyi | 1 + typings/dspy/predict/__init__.pyi | 1 + typings/dspy/predict/parameter.pyi | 1 + typings/dspy/retrievers/retrieve.pyi | 12 +++++++++++- 9 files changed, 41 insertions(+), 13 deletions(-) create mode 100644 typings/dspy/predict/__init__.pyi create mode 100644 typings/dspy/predict/parameter.pyi diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 899f19c..008c4df 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -14,8 +14,8 @@ jobs: uses: astral-sh/setup-uv@v5 with: python-version: "3.12" - - name: Install (with dev and typed integration extras) - run: uv pip install -e ".[dev,langchain,llamaindex,dspy]" + - name: Install (with dev and type-check extras) + run: uv pip install -e ".[dev,typecheck]" - name: Type check run: uv run --no-sync mypy diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 61cc245..df752c0 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -32,7 +32,7 @@ curl -LsSf https://astral.sh/uv/install.sh | sh # 3. Create a virtual environment and install dev dependencies uv venv source .venv/bin/activate # Windows: .venv\Scripts\Activate.ps1 -make install # editable install with dev + ingest extras +make install # editable install with dev, ingest, and type-check extras # 4. Verify everything works make check # lint + static type checks @@ -65,9 +65,9 @@ make clean # Remove caches and build artifacts - Prefer `make` targets over invoking tools directly — they match CI exactly. - For direct tool calls, use the `uv run --no-sync` prefix (plain `uv run` can trigger a - universal resolve that pulls yanked optional deps). Install the `langchain`, - `llamaindex`, and `dspy` extras before type-checking so adapter inheritance is checked - against the real framework APIs. + universal resolve that pulls yanked optional deps). `make install` includes the + `typecheck` extra so adapter inheritance is checked against the real LangChain, + LlamaIndex, and DSPy APIs. - Run `make run-ci` before declaring a change complete; it is the same pipeline CI runs. - Prefer editing existing files over creating new ones, and follow the conventions in neighboring modules. @@ -111,7 +111,7 @@ dynavec is a single Python package (not a monorepo): **Using Make (recommended):** ```bash -make install # uv pip install -e ".[dev,ingest]" +make install # uv pip install -e ".[dev,ingest,typecheck]" make install-all # everything: all embedders + adapters + dev tools ``` @@ -119,7 +119,7 @@ make install-all # everything: all embedders + adapters + dev tools ```bash uv venv && source .venv/bin/activate -uv pip install -e ".[dev]" # add ,ingest / ,all as needed +uv pip install -e ".[dev,ingest,typecheck]" # same environment as make install ``` ### Pre-commit hooks (optional but encouraged) @@ -139,7 +139,7 @@ pre-commit run --all-files # run across the whole tree once | Command | What it does | |---------|--------------| | `make help` | List all targets | -| `make install` | Editable install with dev + ingest extras | +| `make install` | Editable install with dev, ingest, and strict type-check dependencies | | `make install-all` | Editable install with **all** extras | | `make format` | Auto-format (`ruff format`) and auto-fix lint (`ruff --fix`) | | `make lint` | Lint with ruff | @@ -211,7 +211,8 @@ uv run --no-sync pytest tests/test_cache.py -k "jitter" -v - **Style/linting:** [ruff](https://docs.astral.sh/ruff/) (config in `pyproject.toml`, rule sets `E, F, I, UP, B`, line length 100). `make format` fixes most issues automatically. - **Type hints:** dynavec ships a `py.typed` marker. Mypy checks the package in strict mode - in a dedicated CI job; run it locally with `make typecheck`. + against the optional framework APIs in a dedicated CI job. Run `make install` once, + then `make typecheck` locally. - **CI** (`.github/workflows/ci.yml`) runs on every push/PR: **mypy**, plus ruff and pytest across Python 3.9, 3.11, and 3.12. `make run-ci` reproduces it locally. diff --git a/Makefile b/Makefile index 35d29f1..c52b29d 100644 --- a/Makefile +++ b/Makefile @@ -13,8 +13,8 @@ help: ## Show this help @grep -E '^[a-zA-Z_-]+:.*?## .*$$' $(MAKEFILE_LIST) \ | awk 'BEGIN {FS = ":.*?## "}; {printf " \033[36m%-14s\033[0m %s\n", $$1, $$2}' -install: ## Install dynavec + dev tools in editable mode (recommended) - uv pip install -e ".[dev,ingest]" +install: ## Install dynavec + dev, ingest, and type-check dependencies (recommended) + uv pip install -e ".[dev,ingest,typecheck]" install-all: ## Install everything: all embedders, adapters, and dev tools uv pip install -e ".[all,ingest,dev]" diff --git a/pyproject.toml b/pyproject.toml index 4c0eca9..cb655dc 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -53,6 +53,13 @@ langchain = ["langchain-core>=0.3"] llamaindex = ["llama-index-core>=0.11"] crewai = ["crewai>=0.70,!=1.14.0; python_version >= '3.10'"] dspy = ["dspy>=3.3; python_version >= '3.10'"] +# Framework APIs imported by strict mypy checks. Keep this in sync with the +# integration adapters covered by tests/typing. +typecheck = [ + "langchain-core>=0.3", + "llama-index-core>=0.11", + "dspy>=3.3; python_version >= '3.10'", +] all = [ "openai>=1.40", "google-generativeai>=0.8", diff --git a/tests/typing/integration_inheritance.py b/tests/typing/integration_inheritance.py index 75cc576..a3324e2 100644 --- a/tests/typing/integration_inheritance.py +++ b/tests/typing/integration_inheritance.py @@ -1,4 +1,5 @@ import dspy +from dspy.predict.parameter import Parameter from langchain_core.vectorstores import VectorStore from llama_index.core.vector_stores.types import BasePydanticVectorStore @@ -13,6 +14,12 @@ def use_langchain(store: DynavecVectorStore) -> VectorStore: def use_dspy(retriever: DynavecRM) -> dspy.Retrieve: + retriever.reset() + retriever.load_state(retriever.dump_state()) + return retriever + + +def use_dspy_parameter(retriever: DynavecRM) -> Parameter: return retriever diff --git a/typings/dspy/__init__.pyi b/typings/dspy/__init__.pyi index 6f896c7..8c47bce 100644 --- a/typings/dspy/__init__.pyi +++ b/typings/dspy/__init__.pyi @@ -1 +1,2 @@ +from dspy.predict.parameter import Parameter as Parameter from dspy.retrievers.retrieve import Retrieve as Retrieve diff --git a/typings/dspy/predict/__init__.pyi b/typings/dspy/predict/__init__.pyi new file mode 100644 index 0000000..d3551f6 --- /dev/null +++ b/typings/dspy/predict/__init__.pyi @@ -0,0 +1 @@ +from dspy.predict.parameter import Parameter as Parameter diff --git a/typings/dspy/predict/parameter.pyi b/typings/dspy/predict/parameter.pyi new file mode 100644 index 0000000..ed72368 --- /dev/null +++ b/typings/dspy/predict/parameter.pyi @@ -0,0 +1 @@ +class Parameter: ... diff --git a/typings/dspy/retrievers/retrieve.pyi b/typings/dspy/retrievers/retrieve.pyi index 58d4492..2156c62 100644 --- a/typings/dspy/retrievers/retrieve.pyi +++ b/typings/dspy/retrievers/retrieve.pyi @@ -1,9 +1,19 @@ +from collections.abc import Mapping from typing import Any -class Retrieve: +from dspy.predict.parameter import Parameter + +class Retrieve(Parameter): + name: str + input_variable: str + desc: str + stage: str k: int callbacks: list[Any] def __init__(self, k: int = ..., callbacks: list[Any] | None = ...) -> None: ... + def reset(self) -> None: ... + def dump_state(self) -> dict[str, int]: ... + def load_state(self, state: Mapping[str, Any]) -> None: ... def __call__(self, *args: Any, **kwargs: Any) -> Any: ... def forward(self, query: str, k: int | None = ..., **kwargs: Any) -> Any: ... From 65d410a44cfaebf5e567b614178d9750a7bee172 Mon Sep 17 00:00:00 2001 From: eunwoo song Date: Thu, 17 Sep 2026 21:15:52 +0900 Subject: [PATCH 5/8] fix: preserve DSPy retrieval return contract --- src/dynavec/integrations/dspy.py | 6 +++--- tests/test_dspy.py | 2 ++ tests/typing/integration_inheritance.py | 12 ++++++++++++ typings/dspy/__init__.pyi | 1 + typings/dspy/dsp/__init__.pyi | 0 typings/dspy/dsp/utils/__init__.pyi | 6 ------ typings/dspy/primitives/__init__.pyi | 1 + typings/dspy/primitives/prediction.pyi | 6 ++++++ typings/dspy/retrievers/retrieve.pyi | 9 +++++++-- 9 files changed, 32 insertions(+), 11 deletions(-) delete mode 100644 typings/dspy/dsp/__init__.pyi delete mode 100644 typings/dspy/dsp/utils/__init__.pyi create mode 100644 typings/dspy/primitives/__init__.pyi create mode 100644 typings/dspy/primitives/prediction.pyi diff --git a/src/dynavec/integrations/dspy.py b/src/dynavec/integrations/dspy.py index 5d0a473..821ed83 100644 --- a/src/dynavec/integrations/dspy.py +++ b/src/dynavec/integrations/dspy.py @@ -19,7 +19,7 @@ try: import dspy - from dspy.dsp.utils import dotdict + from dspy.primitives.prediction import Prediction except ImportError as exc: # pragma: no cover - import guard raise MissingDependencyError("DynavecRM", "dspy", "dspy") from exc @@ -42,7 +42,7 @@ def forward( query: str, k: int | None = None, **kwargs: Any, - ) -> list[dotdict]: + ) -> list[Prediction]: k = k if k is not None else self.k results = self._client.search( @@ -53,7 +53,7 @@ def forward( ) return [ - dotdict( + Prediction( long_text=result.text or "", id=result.id, score=result.score, diff --git a/tests/test_dspy.py b/tests/test_dspy.py index 93e187f..de0452b 100644 --- a/tests/test_dspy.py +++ b/tests/test_dspy.py @@ -70,6 +70,7 @@ def test_dynavec_rm_returns_dspy_passages(): passages = rm("dynavec") assert len(passages) == 2 + assert isinstance(passages[0], dspy.Prediction) assert passages[0].long_text == ( "Retrieval-augmented generation combines retrieval with generation." @@ -77,6 +78,7 @@ def test_dynavec_rm_returns_dspy_passages(): assert passages[0].id == "doc-1" assert passages[0].score == 0.95 assert passages[0].metadata == {"topic": "rag"} + assert passages[0]["id"] == "doc-1" assert client.calls == [ ( diff --git a/tests/typing/integration_inheritance.py b/tests/typing/integration_inheritance.py index a3324e2..633c437 100644 --- a/tests/typing/integration_inheritance.py +++ b/tests/typing/integration_inheritance.py @@ -1,5 +1,8 @@ +from typing import assert_type + import dspy from dspy.predict.parameter import Parameter +from dspy.primitives.prediction import Prediction from langchain_core.vectorstores import VectorStore from llama_index.core.vector_stores.types import BasePydanticVectorStore @@ -23,5 +26,14 @@ def use_dspy_parameter(retriever: DynavecRM) -> Parameter: return retriever +def check_dspy_return_contract(retriever: DynavecRM) -> None: + assert_type(retriever.forward("query"), list[Prediction]) + base: dspy.Retrieve = retriever + assert_type( + base.forward("query"), + list[str] | Prediction | list[Prediction], + ) + + def use_llamaindex(store: DynavecLlamaStore) -> BasePydanticVectorStore: return store diff --git a/typings/dspy/__init__.pyi b/typings/dspy/__init__.pyi index 8c47bce..ae2b253 100644 --- a/typings/dspy/__init__.pyi +++ b/typings/dspy/__init__.pyi @@ -1,2 +1,3 @@ from dspy.predict.parameter import Parameter as Parameter +from dspy.primitives.prediction import Prediction as Prediction from dspy.retrievers.retrieve import Retrieve as Retrieve diff --git a/typings/dspy/dsp/__init__.pyi b/typings/dspy/dsp/__init__.pyi deleted file mode 100644 index e69de29..0000000 diff --git a/typings/dspy/dsp/utils/__init__.pyi b/typings/dspy/dsp/utils/__init__.pyi deleted file mode 100644 index 6033e84..0000000 --- a/typings/dspy/dsp/utils/__init__.pyi +++ /dev/null @@ -1,6 +0,0 @@ -from typing import Any - -class dotdict(dict[str, Any]): - def __getattr__(self, key: str) -> Any: ... - def __setattr__(self, key: str, value: Any) -> None: ... - def __delattr__(self, key: str) -> None: ... diff --git a/typings/dspy/primitives/__init__.pyi b/typings/dspy/primitives/__init__.pyi new file mode 100644 index 0000000..6106b53 --- /dev/null +++ b/typings/dspy/primitives/__init__.pyi @@ -0,0 +1 @@ +from dspy.primitives.prediction import Prediction as Prediction diff --git a/typings/dspy/primitives/prediction.pyi b/typings/dspy/primitives/prediction.pyi new file mode 100644 index 0000000..a13fe45 --- /dev/null +++ b/typings/dspy/primitives/prediction.pyi @@ -0,0 +1,6 @@ +from typing import Any + +class Prediction: + def __init__(self, *args: Any, **kwargs: Any) -> None: ... + def __getattr__(self, key: str) -> Any: ... + def __getitem__(self, key: str) -> Any: ... diff --git a/typings/dspy/retrievers/retrieve.pyi b/typings/dspy/retrievers/retrieve.pyi index 2156c62..d82741f 100644 --- a/typings/dspy/retrievers/retrieve.pyi +++ b/typings/dspy/retrievers/retrieve.pyi @@ -2,6 +2,7 @@ from collections.abc import Mapping from typing import Any from dspy.predict.parameter import Parameter +from dspy.primitives.prediction import Prediction class Retrieve(Parameter): name: str @@ -15,5 +16,9 @@ class Retrieve(Parameter): def reset(self) -> None: ... def dump_state(self) -> dict[str, int]: ... def load_state(self, state: Mapping[str, Any]) -> None: ... - def __call__(self, *args: Any, **kwargs: Any) -> Any: ... - def forward(self, query: str, k: int | None = ..., **kwargs: Any) -> Any: ... + def __call__( + self, *args: Any, **kwargs: Any + ) -> list[str] | Prediction | list[Prediction]: ... + def forward( + self, query: str, k: int | None = ..., **kwargs: Any + ) -> list[str] | Prediction | list[Prediction]: ... From 9afe28e3016595d7864dec890831c02433c02212 Mon Sep 17 00:00:00 2001 From: eunwoo song Date: Fri, 18 Sep 2026 07:52:44 +0900 Subject: [PATCH 6/8] fix: type current development additions --- src/dynavec/cache.py | 2 +- src/dynavec/client.py | 12 +++-- src/dynavec/embeddings/cached.py | 5 +- src/dynavec/graph.py | 4 +- src/dynavec/quantization.py | 87 +++++++++++++++++--------------- src/dynavec/stores/dynamodb.py | 11 ++-- 6 files changed, 64 insertions(+), 57 deletions(-) diff --git a/src/dynavec/cache.py b/src/dynavec/cache.py index 5a36220..27a5d5b 100644 --- a/src/dynavec/cache.py +++ b/src/dynavec/cache.py @@ -369,7 +369,7 @@ def warm_cache( *, namespace: str = "default", top_k: int = 10, - **search_kwargs, + **search_kwargs: Any, ) -> int: """Pre-populate the query cache from a list of common queries. diff --git a/src/dynavec/client.py b/src/dynavec/client.py index 0fb2a01..5cff327 100644 --- a/src/dynavec/client.py +++ b/src/dynavec/client.py @@ -95,7 +95,7 @@ def __init__( self._graph_store: GraphStore | None = None self._pool: ThreadPoolExecutor | None = None self._hot: HotTier | None = HotTier(config) if config.hot_tier else None - self._cross_encoder = None + self._cross_encoder: Any | None = None if embedder is not None and embedder.dimension != config.dimension: raise ConfigurationError( @@ -486,8 +486,8 @@ def _search_core( results = self._apply_rescore(query_vector, results, rescore) if rerank == "mmr": results = maximal_marginal_relevance( - results, query_vector, + results, top_k=top_k, lambda_mult=mmr_lambda, ) @@ -580,11 +580,13 @@ def _cross_encoder_rerank( "rerank", ) from exc - if self._cross_encoder is None: - self._cross_encoder = CrossEncoder(self.config.cross_encoder_model) + cross_encoder = self._cross_encoder + if cross_encoder is None: + cross_encoder = CrossEncoder(self.config.cross_encoder_model) + self._cross_encoder = cross_encoder pairs = [(query, result.text) for result in results] - scores = self._cross_encoder.predict(pairs) + scores = cross_encoder.predict(pairs) reranked = sorted( zip(results, scores), diff --git a/src/dynavec/embeddings/cached.py b/src/dynavec/embeddings/cached.py index 2ae6289..8e25104 100644 --- a/src/dynavec/embeddings/cached.py +++ b/src/dynavec/embeddings/cached.py @@ -47,10 +47,7 @@ def __init__( self._inner = inner self._cache = cache if cache is not None else InMemoryCache() self._model_key = model_key or type(inner).__name__ - - @property - def dimension(self) -> int: - return self._inner.dimension + self.dimension = inner.dimension def embed_documents(self, texts: list[str]) -> list[Vector]: if not texts: diff --git a/src/dynavec/graph.py b/src/dynavec/graph.py index 189365b..2d50616 100644 --- a/src/dynavec/graph.py +++ b/src/dynavec/graph.py @@ -333,7 +333,7 @@ def list_node_ids(self, ns: str) -> list[str]: return sorted(entity_id for entity_id, _ in self._scan_nodes(ns, "pk, entity_id")) @retry() - def _scan_nodes(self, ns: str, projection: str) -> list[tuple[str, dict]]: + def _scan_nodes(self, ns: str, projection: str) -> list[tuple[str, dict[str, Any]]]: """``(entity_id, item)`` for every node in ``ns``, following pagination.""" prefix = f"{encode_key_component(ns)}{KEY_SEPARATOR}node{KEY_SEPARATOR}" params: dict[str, Any] = { @@ -341,7 +341,7 @@ def _scan_nodes(self, ns: str, projection: str) -> list[tuple[str, dict]]: "ExpressionAttributeValues": {":prefix": prefix}, "ProjectionExpression": projection, } - nodes: list[tuple[str, dict]] = [] + nodes: list[tuple[str, dict[str, Any]]] = [] while True: resp = self._table.scan(**params) for item in resp.get("Items", []): diff --git a/src/dynavec/quantization.py b/src/dynavec/quantization.py index ee05dd7..2ba26a1 100644 --- a/src/dynavec/quantization.py +++ b/src/dynavec/quantization.py @@ -226,21 +226,20 @@ def _fitted_state(self) -> tuple[np.ndarray, int]: class ScalarQuantizer: """Per-dimension INT8 scalar quantization.""" - def __post_init__(self): - self._mins = None - self._scales = None + def __post_init__(self) -> None: + self._mins: np.ndarray | None = None + self._scales: np.ndarray | None = None @property - def is_fitted(self): - return self._mins is not None + def is_fitted(self) -> bool: + return self._mins is not None and self._scales is not None @property - def code_size_bytes(self): - if self._mins is None: - raise RuntimeError("ScalarQuantizer is not fitted") - return self._mins.shape[0] + def code_size_bytes(self) -> int: + mins, _ = self._fitted_state() + return int(mins.shape[0]) - def fit(self, vectors): + def fit(self, vectors: np.ndarray) -> ScalarQuantizer: vectors = np.asarray(vectors, dtype=np.float32) if vectors.ndim != 2: @@ -257,8 +256,8 @@ def fit(self, vectors): return self - def encode(self, vectors): - self._check_fitted() + def encode(self, vectors: np.ndarray) -> np.ndarray: + mins, scales = self._fitted_state() vectors = np.asarray(vectors, dtype=np.float32) @@ -266,30 +265,32 @@ def encode(self, vectors): raise ValueError("vectors must be a 2D array") codes = np.round( - (vectors - self._mins) / self._scales - 128 + (vectors - mins) / scales - 128 ) - return np.clip(codes, -128, 127).astype(np.int8) + return np.asarray(np.clip(codes, -128, 127), dtype=np.int8) - def decode(self, codes): - self._check_fitted() + def decode(self, codes: np.ndarray) -> np.ndarray: + mins, scales = self._fitted_state() codes = np.asarray(codes, dtype=np.int8) - return ( - (codes.astype(np.float32) + 128) * self._scales - + self._mins - ).astype(np.float32) + return np.asarray( + (codes.astype(np.float32) + 128) * scales + + mins, + dtype=np.float32, + ) - def reconstruction_error(self, vectors): + def reconstruction_error(self, vectors: np.ndarray) -> float: vectors = np.asarray(vectors, dtype=np.float32) reconstructed = self.decode(self.encode(vectors)) return float(np.mean((vectors - reconstructed) ** 2)) - def _check_fitted(self): - if not self.is_fitted: + def _fitted_state(self) -> tuple[np.ndarray, np.ndarray]: + if self._mins is None or self._scales is None: raise RuntimeError("ScalarQuantizer must be .fit() before use") + return self._mins, self._scales @dataclass @@ -332,25 +333,26 @@ def fit(self, vectors: np.ndarray) -> OPQRotation: def transform(self, vectors: np.ndarray) -> np.ndarray: """Apply the learned rotation.""" - self._check_fitted() + rotation = self._fitted_rotation() x = np.asarray(vectors, dtype=np.float32) - return x @ self._rotation + return np.asarray(x @ rotation, dtype=np.float32) def inverse_transform(self, vectors: np.ndarray) -> np.ndarray: """Apply the inverse rotation.""" - self._check_fitted() + rotation = self._fitted_rotation() x = np.asarray(vectors, dtype=np.float32) - return x @ self._rotation.T + return np.asarray(x @ rotation.T, dtype=np.float32) - def _check_fitted(self) -> None: - if not self.is_fitted: + def _fitted_rotation(self) -> np.ndarray: + if self._rotation is None: raise RuntimeError( "OPQRotation must be .fit() before use" ) + return self._rotation def _update_rotation( @@ -388,8 +390,8 @@ def is_fitted(self) -> bool: @property def code_size_bytes(self) -> int: - self._check_fitted() - return self._pq.code_size_bytes + pq, _ = self._fitted_components() + return pq.code_size_bytes def fit(self, vectors: np.ndarray) -> OptimizedProductQuantizer: x = np.asarray(vectors, dtype=np.float32) @@ -442,31 +444,31 @@ def fit(self, vectors: np.ndarray) -> OptimizedProductQuantizer: return self def encode(self, vectors: np.ndarray) -> np.ndarray: - self._check_fitted() + pq, opq = self._fitted_components() - rotated = self._opq.transform(vectors) + rotated = opq.transform(vectors) - return self._pq.encode(rotated) + return pq.encode(rotated) def decode(self, codes: np.ndarray) -> np.ndarray: - self._check_fitted() + pq, opq = self._fitted_components() - rotated = self._pq.decode(codes) + rotated = pq.decode(codes) - return self._opq.inverse_transform(rotated) + return opq.inverse_transform(rotated) def asymmetric_distances( self, query: np.ndarray, codes: np.ndarray, ) -> np.ndarray: - self._check_fitted() + pq, opq = self._fitted_components() - rotated_query = self._opq.transform( + rotated_query = opq.transform( np.asarray(query, dtype=np.float32) ) - return self._pq.asymmetric_distances( + return pq.asymmetric_distances( rotated_query, codes, ) @@ -489,8 +491,9 @@ def training_errors(self) -> list[float]: """PQ reconstruction error after each OPQ iteration.""" return self._training_errors.copy() - def _check_fitted(self) -> None: - if not self.is_fitted: + def _fitted_components(self) -> tuple[ProductQuantizer, OPQRotation]: + if self._pq is None or self._opq is None: raise RuntimeError( "OptimizedProductQuantizer must be .fit() before use" ) + return self._pq, self._opq diff --git a/src/dynavec/stores/dynamodb.py b/src/dynavec/stores/dynamodb.py index 7e48ec6..cadfcae 100644 --- a/src/dynavec/stores/dynamodb.py +++ b/src/dynavec/stores/dynamodb.py @@ -57,7 +57,12 @@ def _pk(namespace: str, doc_id: str) -> str: return f"{encode_key_component(namespace)}{KEY_SEPARATOR}{encode_key_component(doc_id)}" -def _build_item(namespace: str, doc_id: str, text: str | None, metadata: Metadata) -> dict: +def _build_item( + namespace: str, + doc_id: str, + text: str | None, + metadata: Metadata, +) -> dict[str, Any]: item = { "pk": _pk(namespace, doc_id), "ns": namespace, @@ -91,7 +96,7 @@ def _value_size(value: Any) -> int: return len(str(value).encode("utf-8")) -def item_size_bytes(item: dict) -> int: +def item_size_bytes(item: dict[str, Any]) -> int: """Approximate DynamoDB size of ``item``: attribute name bytes plus value sizes.""" return sum(len(name.encode("utf-8")) + _value_size(value) for name, value in item.items()) @@ -105,7 +110,7 @@ def check_item_size(namespace: str, doc_id: str, text: str | None, metadata: Met _check_built_item(_build_item(namespace, doc_id, text, metadata)) -def _check_built_item(item: dict) -> None: +def _check_built_item(item: dict[str, Any]) -> None: size = item_size_bytes(item) if size > MAX_ITEM_BYTES: raise ItemTooLargeError(item["id"], item["ns"], size, MAX_ITEM_BYTES) From 27b90137e66d160a0cf8e0680c9f185d84b85473 Mon Sep 17 00:00:00 2001 From: eunwoo song Date: Sat, 19 Sep 2026 19:38:26 +0900 Subject: [PATCH 7/8] fix: type current development additions --- src/dynavec/cli.py | 14 +-- src/dynavec/client.py | 167 ++++++++++++++++++++++-------- src/dynavec/fusion.py | 2 +- src/dynavec/hot.py | 18 ++-- src/dynavec/integrations/tools.py | 29 ++++-- src/dynavec/namespace.py | 19 ++-- src/dynavec/retrievers.py | 16 +-- src/dynavec/stores/s3vectors.py | 5 +- 8 files changed, 179 insertions(+), 91 deletions(-) diff --git a/src/dynavec/cli.py b/src/dynavec/cli.py index b5b3438..3021af6 100644 --- a/src/dynavec/cli.py +++ b/src/dynavec/cli.py @@ -87,10 +87,10 @@ def _parser() -> argparse.ArgumentParser: return parser -def _session(profile: str | None, region: str | None): +def _session(profile: str | None, region: str | None) -> Any: import boto3 - kwargs = {} + kwargs: dict[str, str] = {} if profile: kwargs["profile_name"] = profile if region: @@ -118,7 +118,7 @@ def _resolve_resources(args: argparse.Namespace) -> tuple[str, str, str]: f"Missing required resource configuration: {', '.join(missing)} " "(provide via flags or DYNAVEC_* environment variables)." ) - return bucket, index, table + return str(bucket), str(index), str(table) def _export(args: argparse.Namespace) -> int: @@ -231,7 +231,7 @@ def _doctor(args: argparse.Namespace) -> int: session = None checks_passed = True - def get_session(): + def get_session() -> Any: nonlocal session if session is None: session = _session(args.profile, args.region) @@ -262,18 +262,18 @@ def get_session(): return 0 if checks_passed else 1 -def _identity(session) -> str: +def _identity(session: Any) -> str: identity = session.client("sts").get_caller_identity() return f"Account: {identity.get('Account', 'unknown')}" -def _check_s3vectors(session, bucket: str, index: str, region: str | None) -> str: +def _check_s3vectors(session: Any, bucket: str, index: str, region: str | None) -> str: client = session.client("s3vectors", region_name=region) client.get_index(vectorBucketName=bucket, indexName=index) return f"Bucket: {bucket}" -def _check_dynamodb(session, table: str, region: str | None) -> str: +def _check_dynamodb(session: Any, table: str, region: str | None) -> str: session.client("dynamodb", region_name=region).describe_table(TableName=table) return "Accessible" diff --git a/src/dynavec/client.py b/src/dynavec/client.py index 5cff327..fe42101 100644 --- a/src/dynavec/client.py +++ b/src/dynavec/client.py @@ -60,14 +60,19 @@ 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 .transforms import TransformContext, as_pipeline +from .telemetry import TelemetryRecorder +from .transforms import Transform, TransformContext, TransformPipeline, as_pipeline from .utils import KEY_SEPARATOR, chunked, decode_key_component, encode_key_component Metadata = dict[str, Any] _S3_PUT_CHUNK = 500 _DDB_CHUNK = 500 -RescoreSpec = Union[str, dict] # "cosine" | "manhattan" | {"cosine":0.7,"dot":0.3} +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: @@ -79,10 +84,10 @@ def __init__( embedder: Embedder | None = None, *, credentials: AWSCredentials | None = None, - boto_session=None, - transform=None, - cache=None, - telemetry=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 @@ -126,7 +131,12 @@ def close(self) -> None: def __enter__(self) -> Dynavec: return self - def __exit__(self, *exc) -> None: + def __exit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> None: self.close() # ------------------------------------------------------------------ setup @@ -172,7 +182,7 @@ 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) - def _run_parallel(self, tasks: list) -> None: + def _run_parallel(self, tasks: list[Callable[[], None]]) -> None: """Run zero-arg callables; parallel if enabled, else sequential.""" if not tasks: return @@ -190,8 +200,8 @@ def _prepare( docs: list[Document], namespace: str, auto_metadata: bool, - transform, - ) -> tuple[list[tuple], list[tuple], list[str], list[tuple]]: + transform: TransformSpec | None, + ) -> 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 @@ -222,7 +232,10 @@ def _prepare( docs[idx].vector = vec # 3) validate + build payloads - s3_payload, ddb_payload, ids, hot_payload = [], [], [], [] + 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: @@ -250,12 +263,17 @@ def _prepare( hot_payload.append((d.id, vector, d.text, meta)) return s3_payload, ddb_payload, ids, hot_payload - def _write(self, namespace: str, s3_payload: list, ddb_payload: list) -> None: - tasks = [] - for chunk in chunked(s3_payload, _S3_PUT_CHUNK): - tasks.append(lambda c=chunk: self._vectors.put_vectors(c)) - for chunk in chunked(ddb_payload, _DDB_CHUNK): - tasks.append(lambda c=chunk: self._docs.put_many(namespace, c)) + def _write( + self, + namespace: str, + s3_payload: list[S3Payload], + ddb_payload: list[DDBPayload], + ) -> None: + tasks: list[Callable[[], None]] = [] + for s3_chunk in chunked(s3_payload, _S3_PUT_CHUNK): + tasks.append(partial(self._vectors.put_vectors, s3_chunk)) + for ddb_chunk in chunked(ddb_payload, _DDB_CHUNK): + tasks.append(partial(self._docs.put_many, namespace, ddb_chunk)) self._run_parallel(tasks) def upsert( @@ -264,7 +282,7 @@ def upsert( *, namespace: str = "default", auto_metadata: bool = False, - transform=None, + transform: TransformSpec | None = None, ) -> UpsertResult: """Insert or overwrite documents (each a :class:`Document` or dict).""" if not documents: @@ -287,7 +305,7 @@ def update( vector: list[float] | None = None, metadata: Metadata | None = None, merge_metadata: bool = True, - transform=None, + transform: TransformSpec | None = None, upsert_if_missing: bool = False, ) -> UpsertResult: """Update an existing document's text, vector, and/or metadata. @@ -419,17 +437,17 @@ def search( def _search_core( self, - query_vector, + query_vector: list[float], *, query: str | None = None, - top_k, - namespace, - filter, - rescore, - rerank, - mmr_lambda, - include_vectors, - normalize_scores, + top_k: int, + namespace: str, + filter: Metadata | None, + rescore: RescoreSpec | None, + rerank: str | None, + mmr_lambda: float, + include_vectors: bool, + normalize_scores: bool, ) -> list[SearchResult]: needs_vectors = rerank == "mmr" or rescore is not None or include_vectors fetch_k = top_k * self.config.over_fetch if (rerank or rescore) else top_k @@ -517,8 +535,19 @@ def _rescore_label(rescore: RescoreSpec | None) -> str | None: return None return rescore if isinstance(rescore, str) else "composite" - def _record_search(self, tel, t0, namespace, top_k, results, cache_hit, - filter, rescore, rerank, query) -> None: + def _record_search( + self, + tel: TelemetryRecorder | None, + t0: float, + namespace: str, + top_k: int, + results: list[SearchResult], + cache_hit: bool | None, + filter: Metadata | None, + rescore: RescoreSpec | None, + rerank: str | None, + query: str | None, + ) -> None: if tel is None: return scores = [r.score for r in results] if results else [] @@ -651,7 +680,12 @@ def search_stream( yielded += 1 def search_many( - self, queries: list[str], *, top_k: int = 10, namespace: str = "default", **kw + self, + queries: list[str], + *, + top_k: int = 10, + namespace: str = "default", + **kw: Any, ) -> list[list[SearchResult]]: """Run several queries concurrently (thread pool over I/O-bound calls).""" futures = [ @@ -662,12 +696,12 @@ def search_many( def as_multiquery_retriever( self, - generate_queries=None, + generate_queries: Any = None, *, - llm_generate_queries=None, + llm_generate_queries: Any = None, namespace: str = "default", - **kw, - ): + **kw: Any, + ) -> Any: """Create a :class:`~dynavec.retrievers.MultiQueryRetriever` bound to this client.""" from .retrievers import MultiQueryRetriever @@ -681,12 +715,12 @@ def as_multiquery_retriever( def as_hyde_retriever( self, - generate_hypothetical=None, + generate_hypothetical: Any = None, *, - llm_generate_hypothetical=None, + llm_generate_hypothetical: Any = None, namespace: str = "default", - **kw, - ): + **kw: Any, + ) -> Any: """Create a :class:`~dynavec.retrievers.HyDERetriever` bound to this client.""" from .retrievers import HyDERetriever @@ -805,17 +839,37 @@ def list_vectors( ) # -------------------------------------------------------------- graph / ER - def graph_add_node(self, entity_id, *, namespace="default", ntype=None, props=None): + def graph_add_node( + self, + entity_id: str, + *, + namespace: str = "default", + ntype: str | None = None, + props: Metadata | None = None, + ) -> None: """Create/update a graph entity (a 'meaning' node).""" self.graph.add_node(namespace, entity_id, ntype, props) - def graph_add_edge(self, src, relation, dst, *, namespace="default", bidirectional=False): + def graph_add_edge( + self, + src: str, + relation: str, + dst: str, + *, + namespace: str = "default", + bidirectional: bool = False, + ) -> None: """Relate two entities: ``(src) -[relation]-> (dst)``.""" self.graph.add_edge(namespace, src, relation, dst) if bidirectional: self.graph.add_edge(namespace, dst, relation, src) - def graph_delete_node(self, entity_id, *, namespace="default"): + def graph_delete_node( + self, + entity_id: str, + *, + namespace: str = "default", + ) -> int: """Delete an entity with its outbound and inbound edges (idempotent). Linked documents and their embeddings are left untouched. Returns the @@ -823,23 +877,44 @@ def graph_delete_node(self, entity_id, *, namespace="default"): """ return self.graph.delete_node(namespace, entity_id) - def graph_delete_edge(self, src, relation, dst, *, namespace="default", bidirectional=False): + def graph_delete_edge( + self, + src: str, + relation: str, + dst: str, + *, + namespace: str = "default", + bidirectional: bool = False, + ) -> int: """Remove ``(src) -[relation]-> (dst)`` (idempotent); return edges removed.""" removed = self.graph.delete_edge(namespace, src, relation, dst) if bidirectional: removed += self.graph.delete_edge(namespace, dst, relation, src) return removed - def graph_link(self, entity_id, doc_ids, *, namespace="default"): + def graph_link( + self, + entity_id: str, + doc_ids: Sequence[str], + *, + namespace: str = "default", + ) -> None: """Attach documents (their S3 Vectors embeddings) to an entity.""" self.graph.link_docs(namespace, entity_id, list(doc_ids)) - def graph_neighbors(self, entity_id, *, namespace="default", relation=None, hops=1): + def graph_neighbors( + self, + entity_id: str, + *, + namespace: str = "default", + relation: str | None = None, + hops: int = 1, + ) -> list[str]: """Breadth-first traversal returning reachable entity ids (excl. seed).""" visited = {entity_id} frontier = [entity_id] for _ in range(hops): - nxt = [] + nxt: list[str] = [] for node in frontier: for nb in self.graph.neighbors(namespace, node, relation): if nb not in visited: diff --git a/src/dynavec/fusion.py b/src/dynavec/fusion.py index e3ab17d..5e73b99 100644 --- a/src/dynavec/fusion.py +++ b/src/dynavec/fusion.py @@ -223,7 +223,7 @@ def _fit_bayesian(self, n_points: int) -> FitResult: def _softmax(x: np.ndarray) -> list[float]: e = np.exp(x - x.max()) - return (e / e.sum()).tolist() + return [float(value) for value in (e / e.sum()).tolist()] def _neg_ndcg(x: np.ndarray) -> float: nonlocal eval_count diff --git a/src/dynavec/hot.py b/src/dynavec/hot.py index 44a5b23..ec8ffcd 100644 --- a/src/dynavec/hot.py +++ b/src/dynavec/hot.py @@ -29,14 +29,14 @@ # --- MongoDB-style metadata matcher (mirrors the S3 Vectors filter dialect) --- _COMPARATORS: dict[str, Callable[[Any, Any], bool]] = { - "$eq": lambda a, b: a == b, - "$ne": lambda a, b: a != b, - "$gt": lambda a, b: a is not None and a > b, - "$gte": lambda a, b: a is not None and a >= b, - "$lt": lambda a, b: a is not None and a < b, - "$lte": lambda a, b: a is not None and a <= b, - "$in": lambda a, b: a in b, - "$nin": lambda a, b: a not in b, + "$eq": lambda a, b: bool(a == b), + "$ne": lambda a, b: bool(a != b), + "$gt": lambda a, b: bool(a is not None and a > b), + "$gte": lambda a, b: bool(a is not None and a >= b), + "$lt": lambda a, b: bool(a is not None and a < b), + "$lte": lambda a, b: bool(a is not None and a <= b), + "$in": lambda a, b: bool(a in b), + "$nin": lambda a, b: bool(a not in b), } @@ -58,7 +58,7 @@ def _match_field(value: Any, condition: Any) -> bool: if not cmp(value, operand): return False return True - return value == condition # bare {k: v} means equality + return bool(value == condition) # bare {k: v} means equality def matches(metadata: dict[str, Any], flt: dict[str, Any] | None) -> bool: diff --git a/src/dynavec/integrations/tools.py b/src/dynavec/integrations/tools.py index a57960c..e5e2644 100644 --- a/src/dynavec/integrations/tools.py +++ b/src/dynavec/integrations/tools.py @@ -156,15 +156,15 @@ def handle_tool_call(self, tool_call: Any) -> dict[str, str]: if parsed_args is None: query = "" elif isinstance(parsed_args, dict): - query = parsed_args.get("query") - if query is None: + query_value: Any = parsed_args.get("query") + if query_value is None: for fallback_key in ("q", "input", "search_query", "text", "prompt"): if fallback_key in parsed_args: - query = parsed_args[fallback_key] + query_value = parsed_args[fallback_key] break - if query is None and parsed_args: - query = next((v for v in parsed_args.values() if isinstance(v, str)), "") - query = str(query or "") + if query_value is None and parsed_args: + query_value = next((v for v in parsed_args.values() if isinstance(v, str)), "") + query = str(query_value or "") elif isinstance(parsed_args, (list, tuple, set)): query = " ".join(str(x) for x in parsed_args) else: @@ -221,7 +221,12 @@ def as_openai_tool( return OpenAIAssistantTool(fn, name=name, description=description) -def as_langchain_tool(source, *, name: str = "dynavec_search", **kw) -> Any: +def as_langchain_tool( + source: Dynavec | NamespaceView | QueryExpansionRetriever, + *, + name: str = "dynavec_search", + **kw: Any, +) -> Any: """Wrap the retriever as a LangChain ``StructuredTool`` (requires langchain-core).""" from ..exceptions import MissingDependencyError @@ -234,7 +239,12 @@ def as_langchain_tool(source, *, name: str = "dynavec_search", **kw) -> Any: return StructuredTool.from_function(func=fn, name=name, description=fn.__doc__) -def as_crewai_tool(source, *, name: str = "dynavec_search", **kw) -> Any: +def as_crewai_tool( + source: Dynavec | NamespaceView, + *, + name: str = "dynavec_search", + **kw: Any, +) -> Any: """Wrap the retriever as a CrewAI tool (requires crewai).""" from ..exceptions import MissingDependencyError @@ -245,9 +255,8 @@ def as_crewai_tool(source, *, name: str = "dynavec_search", **kw) -> Any: fn = make_retriever_fn(source, **kw) - @crewai_tool(name) def _tool(query: str) -> str: """Search the dynavec knowledge base for relevant passages.""" return fn(query) - return _tool + return crewai_tool(name)(_tool) diff --git a/src/dynavec/namespace.py b/src/dynavec/namespace.py index 88ad4db..821fda4 100644 --- a/src/dynavec/namespace.py +++ b/src/dynavec/namespace.py @@ -43,9 +43,7 @@ def update(self, id: str, **kw: Any) -> UpsertResult: def search(self, query: str | None = None, **kw: Any) -> list[SearchResult]: return self._db.search(query, namespace=self._ns, **kw) - def search_stream( - self, query: str | None = None, **kw: Any - ) -> Iterator[SearchResult]: + def search_stream(self, query: str | None = None, **kw: Any) -> Iterator[SearchResult]: yield from self._db.search_stream(query, namespace=self._ns, **kw) def get(self, ids: list[str], **kw: Any) -> list[SearchResult]: @@ -55,8 +53,8 @@ def delete(self, ids: list[str], **kw: Any) -> None: self._db.delete(ids, namespace=self._ns, **kw) def as_multiquery_retriever( - self, generate_queries=None, *, llm_generate_queries=None, **kw - ): + self, generate_queries: Any = None, *, llm_generate_queries: Any = None, **kw: Any + ) -> Any: """Create a :class:`~dynavec.retrievers.MultiQueryRetriever` pinned to this namespace.""" from .retrievers import MultiQueryRetriever @@ -68,8 +66,8 @@ def as_multiquery_retriever( ) def as_hyde_retriever( - self, generate_hypothetical=None, *, llm_generate_hypothetical=None, **kw - ): + self, generate_hypothetical: Any = None, *, llm_generate_hypothetical: Any = None, **kw: Any + ) -> Any: """Create a :class:`~dynavec.retrievers.HyDERetriever` pinned to this namespace.""" from .retrievers import HyDERetriever @@ -79,13 +77,14 @@ def as_hyde_retriever( llm_generate_hypothetical=llm_generate_hypothetical, **kw, ) - def export_namespace(self, output, **kw) -> int: + + def export_namespace(self, output: Any, **kw: Any) -> int: return self._db.export_namespace(output, namespace=self._ns, **kw) - def import_namespace(self, input, **kw) -> int: + def import_namespace(self, input: Any, **kw: Any) -> int: return self._db.import_namespace(input, namespace=self._ns, **kw) - def __iter__(self): + def __iter__(self) -> Iterator[dict[str, Any]]: return self._db.iter_namespace(namespace=self._ns) def __repr__(self) -> str: diff --git a/src/dynavec/retrievers.py b/src/dynavec/retrievers.py index 3ae9167..fed2960 100644 --- a/src/dynavec/retrievers.py +++ b/src/dynavec/retrievers.py @@ -332,7 +332,7 @@ def __init__( *, llm_generate_queries: Callable[[str], Sequence[str]] | None = None, n_queries: int = 3, - **kwargs, + **kwargs: Any, ) -> None: gen = generate_queries if generate_queries is not None else llm_generate_queries if gen is None or not callable(gen): @@ -354,13 +354,15 @@ def _build_plan( plan.append((1.0, self._text_search(text, depth, filter, use_cache))) return plan - def _plan(self, query: str, depth: int, filter: Metadata | None, use_cache: bool | None): + def _plan( + self, query: str, depth: int, filter: Metadata | None, use_cache: bool | None + ) -> list[tuple[float, Callable[[], list[SearchResult]]]]: raw = self._invoke_generator(self._generate_queries, query) return self._build_plan(raw, query, depth, filter, use_cache) async def _async_plan( self, query: str, depth: int, filter: Metadata | None, use_cache: bool | None - ): + ) -> list[tuple[float, Callable[[], list[SearchResult]]]]: raw = await self._async_invoke_generator(self._generate_queries, query) return self._build_plan(raw, query, depth, filter, use_cache) @@ -403,7 +405,7 @@ def __init__( llm_generate_hypothetical: Callable[[str], str | Sequence[str]] | None = None, strategy: HyDEStrategy = "average", max_passages: int = 5, - **kwargs, + **kwargs: Any, ) -> None: gen = ( generate_hypothetical @@ -456,12 +458,14 @@ def _build_plan( return plan - def _plan(self, query: str, depth: int, filter: Metadata | None, use_cache: bool | None): + def _plan( + self, query: str, depth: int, filter: Metadata | None, use_cache: bool | None + ) -> list[tuple[float, Callable[[], list[SearchResult]]]]: raw = self._invoke_generator(self._generate_hypothetical, query) return self._build_plan(raw, query, depth, filter, use_cache) async def _async_plan( self, query: str, depth: int, filter: Metadata | None, use_cache: bool | None - ): + ) -> list[tuple[float, Callable[[], list[SearchResult]]]]: raw = await self._async_invoke_generator(self._generate_hypothetical, query) return self._build_plan(raw, query, depth, filter, use_cache) diff --git a/src/dynavec/stores/s3vectors.py b/src/dynavec/stores/s3vectors.py index a76fbc4..ac31293 100644 --- a/src/dynavec/stores/s3vectors.py +++ b/src/dynavec/stores/s3vectors.py @@ -51,11 +51,12 @@ def __init__(self, config: DynavecConfig, boto_session: Any | None = None) -> No self._client = session.client("s3vectors", **client_kwargs) self._logger = logging.getLogger("dynavec.stores.s3vectors") - def get_index(self) -> dict: - return self._client.get_index( + def get_index(self) -> dict[str, Any]: + result: dict[str, Any] = self._client.get_index( vectorBucketName=self._config.vector_bucket, indexName=self._config.index, ) + return result @retry() def _put_batch(self, payload: list[dict[str, Any]]) -> None: From 852c6802021908027374fdfbc4d144a3e85f506f Mon Sep 17 00:00:00 2001 From: eunwoo song Date: Tue, 22 Sep 2026 09:45:58 +0900 Subject: [PATCH 8/8] fix(types): preserve explained search and live embedder contracts --- CONTRIBUTING.md | 9 ++- src/dynavec/client.py | 99 +++++++++++++++++++++++++++++--- src/dynavec/embeddings/cached.py | 9 ++- src/dynavec/namespace.py | 5 ++ tests/test_client_inmemory.py | 20 +++++++ tests/test_embedding_cache.py | 13 +++++ tests/typing/search_contracts.py | 26 +++++++++ 7 files changed, 170 insertions(+), 11 deletions(-) create mode 100644 tests/typing/search_contracts.py diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index df752c0..874503b 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -212,7 +212,14 @@ uv run --no-sync pytest tests/test_cache.py -k "jitter" -v sets `E, F, I, UP, B`, line length 100). `make format` fixes most issues automatically. - **Type hints:** dynavec ships a `py.typed` marker. Mypy checks the package in strict mode against the optional framework APIs in a dedicated CI job. Run `make install` once, - then `make typecheck` locally. + then `make typecheck` locally. Consumer fixtures in `tests/typing` protect + integration inheritance and the `explain` return contract for client, namespace, + and batch searches. When adding a return-shape option, cover both literal values + and a runtime `bool` across forwarding APIs; do not cast an overloaded callable + to a single result shape to satisfy an executor's type inference. Wrapper + annotations must also preserve live delegation: an embedder such as Ollama can + infer its dimension on the first request, so a cached wrapper must not snapshot + that property during construction. - **CI** (`.github/workflows/ci.yml`) runs on every push/PR: **mypy**, plus ruff and pytest across Python 3.9, 3.11, and 3.12. `make run-ci` reproduces it locally. diff --git a/src/dynavec/client.py b/src/dynavec/client.py index 890e54d..ede5601 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, cast, overload +from typing import TYPE_CHECKING, Any, Literal, Optional, TextIO, Union, overload if TYPE_CHECKING: from .cache import BaseCache @@ -871,20 +871,101 @@ def search_stream( ) yielded += 1 + @overload def search_many( self, queries: list[str], *, top_k: int = 10, namespace: str = "default", - **kw: Any, - ) -> list[list[SearchResult]]: - """Run several queries concurrently (thread pool over I/O-bound calls).""" - search = cast(Callable[..., list[SearchResult]], self.search) - futures = [ - self._executor.submit(search, q, top_k=top_k, namespace=namespace, **kw) - for q in queries - ] + explain: Literal[False] = False, + vector: list[float] | None = None, + filter: Metadata | None = None, + rescore: RescoreSpec | None = None, + rerank: str | None = None, + mmr_lambda: float = 0.5, + include_vectors: bool = False, + use_cache: bool | None = None, + normalize_scores: bool = False, + ) -> list[list[SearchResult]]: ... + + @overload + def search_many( + self, + queries: list[str], + *, + top_k: int = 10, + namespace: str = "default", + explain: Literal[True], + vector: list[float] | None = None, + filter: Metadata | None = None, + rescore: RescoreSpec | None = None, + rerank: str | None = None, + mmr_lambda: float = 0.5, + include_vectors: bool = False, + use_cache: bool | None = None, + normalize_scores: bool = False, + ) -> list[ExplainedSearchResult]: ... + + @overload + def search_many( + self, + queries: list[str], + *, + top_k: int = 10, + namespace: str = "default", + explain: bool, + vector: list[float] | None = None, + filter: Metadata | None = None, + rescore: RescoreSpec | None = None, + rerank: str | None = None, + mmr_lambda: float = 0.5, + include_vectors: bool = False, + use_cache: bool | None = None, + normalize_scores: bool = False, + ) -> list[list[SearchResult]] | list[ExplainedSearchResult]: ... + + def search_many( + self, + queries: list[str], + *, + top_k: int = 10, + namespace: str = "default", + explain: bool = False, + vector: list[float] | None = None, + filter: Metadata | None = None, + rescore: RescoreSpec | None = None, + rerank: str | None = None, + mmr_lambda: float = 0.5, + include_vectors: bool = False, + use_cache: bool | None = None, + normalize_scores: bool = False, + ) -> list[list[SearchResult]] | list[ExplainedSearchResult]: + """Run queries concurrently, returning one search result per input query. + + With ``explain=True``, each item is an ``ExplainedSearchResult``. + """ + if explain: + def explained_search(query: str) -> ExplainedSearchResult: + return self.search( + query, top_k=top_k, namespace=namespace, explain=True, + vector=vector, filter=filter, rescore=rescore, rerank=rerank, + mmr_lambda=mmr_lambda, include_vectors=include_vectors, + use_cache=use_cache, normalize_scores=normalize_scores, + ) + + explained_futures = [self._executor.submit(explained_search, q) for q in queries] + return [f.result() for f in explained_futures] + + def plain_search(query: str) -> list[SearchResult]: + return self.search( + query, top_k=top_k, namespace=namespace, explain=False, + vector=vector, filter=filter, rescore=rescore, rerank=rerank, + mmr_lambda=mmr_lambda, include_vectors=include_vectors, + use_cache=use_cache, normalize_scores=normalize_scores, + ) + + futures = [self._executor.submit(plain_search, q) for q in queries] return [f.result() for f in futures] def as_multiquery_retriever( diff --git a/src/dynavec/embeddings/cached.py b/src/dynavec/embeddings/cached.py index 8e25104..2357c66 100644 --- a/src/dynavec/embeddings/cached.py +++ b/src/dynavec/embeddings/cached.py @@ -47,7 +47,14 @@ def __init__( self._inner = inner self._cache = cache if cache is not None else InMemoryCache() self._model_key = model_key or type(inner).__name__ - self.dimension = inner.dimension + + @property + def dimension(self) -> int: + return self._inner.dimension + + @dimension.setter + def dimension(self, value: int) -> None: + self._inner.dimension = value def embed_documents(self, texts: list[str]) -> list[Vector]: if not texts: diff --git a/src/dynavec/namespace.py b/src/dynavec/namespace.py index 0051b27..727bf88 100644 --- a/src/dynavec/namespace.py +++ b/src/dynavec/namespace.py @@ -50,6 +50,11 @@ def search( self, query: str | None = None, *, explain: Literal[True], **kw: Any ) -> ExplainedSearchResult: ... + @overload + def search( + self, query: str | None = None, *, explain: bool, **kw: Any + ) -> list[SearchResult] | ExplainedSearchResult: ... + def search( self, query: str | None = None, *, explain: bool = False, **kw: Any ) -> list[SearchResult] | ExplainedSearchResult: diff --git a/tests/test_client_inmemory.py b/tests/test_client_inmemory.py index 5f5796c..2fb99cf 100644 --- a/tests/test_client_inmemory.py +++ b/tests/test_client_inmemory.py @@ -497,6 +497,26 @@ def test_search_many_parallel(db): assert all(len(r) == 1 for r in results) +@pytest.mark.parametrize("explain", [False, True]) +def test_search_many_explain_preserves_order_and_namespace(db, explain): + ns = db.namespace("kb") + ns.upsert([Document(id="1", text="apple"), Document(id="2", text="rocket")]) + batches = db.search_many( + ["rocket", "apple"], namespace="kb", top_k=1, explain=explain, + normalize_scores=True, + ) + assert len(batches) == 2 + for result, expected_id in zip(batches, ["2", "1"]): + if explain: + assert isinstance(result, ExplainedSearchResult) + result = result.results + assert isinstance(result, list) + assert [hit.id for hit in result] == [expected_id] + single = ns.search("apple", explain=explain) + assert isinstance(single, ExplainedSearchResult if explain else list) + assert db.search_many([], explain=explain) == [] + + def test_context_manager_closes_pool(db): with db as d: d.upsert([Document(id="1", text="hi")]) diff --git a/tests/test_embedding_cache.py b/tests/test_embedding_cache.py index 5cb5a75..f17d8f3 100644 --- a/tests/test_embedding_cache.py +++ b/tests/test_embedding_cache.py @@ -296,3 +296,16 @@ def test_custom_cache_backend_is_used(): def test_default_cache_is_inmemory(): embedder = CachedEmbedder(_CountingEmbedder()) assert isinstance(embedder.cache, InMemoryCache) + + +@pytest.mark.parametrize("method", ["embed_documents", "embed_query"]) +def test_cached_dimension_tracks_ollama_inference(monkeypatch, method): + from dynavec.embeddings.ollama import OllamaEmbedder + + inner = OllamaEmbedder(model="custom-model") + monkeypatch.setattr(inner, "_post", lambda *args: {"embeddings": [[0.1, 0.2, 0.3]]}) + cached = CachedEmbedder(inner) + assert cached.dimension == 768 + getattr(cached, method)(["hello"] if method == "embed_documents" else "hello") + assert inner.dimension == 3 + assert cached.dimension == 3 diff --git a/tests/typing/search_contracts.py b/tests/typing/search_contracts.py new file mode 100644 index 0000000..1670738 --- /dev/null +++ b/tests/typing/search_contracts.py @@ -0,0 +1,26 @@ +"""Consumer return types must follow the explain flag across search entry points.""" + +from typing import assert_type + +from dynavec import Dynavec, ExplainedSearchResult, SearchResult + + +def check_search_contracts(db: Dynavec, explain: bool) -> None: + assert_type(db.search("q"), list[SearchResult]) + assert_type(db.search("q", explain=False), list[SearchResult]) + assert_type(db.search("q", explain=True), ExplainedSearchResult) + assert_type(db.search("q", explain=explain), list[SearchResult] | ExplainedSearchResult) + + ns = db.namespace("kb") + assert_type(ns.search("q"), list[SearchResult]) + assert_type(ns.search("q", explain=False), list[SearchResult]) + assert_type(ns.search("q", explain=True), ExplainedSearchResult) + assert_type(ns.search("q", explain=explain), list[SearchResult] | ExplainedSearchResult) + + assert_type(db.search_many(["q"]), list[list[SearchResult]]) + assert_type(db.search_many(["q"], explain=False), list[list[SearchResult]]) + assert_type(db.search_many(["q"], explain=True), list[ExplainedSearchResult]) + assert_type( + db.search_many(["q"], explain=explain), + list[list[SearchResult]] | list[ExplainedSearchResult], + )