diff --git a/pyproject.toml b/pyproject.toml index e1e73f4..b5b1ad8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -62,6 +62,7 @@ typecheck = [ "llama-index-core>=0.11", "haystack-ai>=2.0", "dspy>=3.3; python_version >= '3.10'", + "semantic-kernel>=1.44; python_version >= '3.10'", ] all = [ "openai>=1.40", @@ -172,6 +173,8 @@ module = [ "redis.*", "requests", "requests.*", + "semantic_kernel", + "semantic_kernel.*", "sentence_transformers", "sentence_transformers.*", "voyageai", diff --git a/src/dynavec/__init__.py b/src/dynavec/__init__.py index 3618fd0..a0aaec8 100644 --- a/src/dynavec/__init__.py +++ b/src/dynavec/__init__.py @@ -20,6 +20,7 @@ from __future__ import annotations +from .bm25 import BM25Index from .cache import BaseCache, DynamoDBCache, RedisCache, SemanticCache, warm_cache from .client import Dynavec from .config import DynavecConfig @@ -72,6 +73,8 @@ reciprocal_rank_fusion, ) from .retrievers import ( + BM25HybridRetriever, + BM25Retriever, HyDERetriever, MultiQueryRetriever, QueryExpansionRetriever, @@ -112,6 +115,9 @@ "QueryExpansionRetriever", "MultiQueryRetriever", "HyDERetriever", + "BM25Index", + "BM25Retriever", + "BM25HybridRetriever", "FitResult", "RRFWeightFitter", "HotTier", diff --git a/src/dynavec/bm25.py b/src/dynavec/bm25.py new file mode 100644 index 0000000..ae0d4e8 --- /dev/null +++ b/src/dynavec/bm25.py @@ -0,0 +1,316 @@ +"""Lightweight Okapi BM25 index and tokenizer for lexical retrieval. + +Implements pure-Python Okapi BM25 with zero third-party dependencies. +Computes Robertson-Spärck Jones / Lucene IDF and BM25 term-frequency saturation +with length normalization, supporting exact keyword lookups, part numbers, SKUs, +and acronyms. +""" + +from __future__ import annotations + +import math +import re +from collections import Counter +from collections.abc import Callable, Sequence +from typing import Any + +from .models import Document, Metadata, SearchResult + +_TOKEN_RE = re.compile(r"[a-zA-Z0-9]+(?:[-_.:][a-zA-Z0-9]+)*") +_SUBTOKEN_RE = re.compile(r"[a-zA-Z0-9]+") + +DEFAULT_STOPWORDS: frozenset[str] = frozenset({ + "a", "an", "and", "are", "as", "at", "be", "but", "by", "for", "if", "in", + "into", "is", "it", "no", "not", "of", "on", "or", "such", "that", "the", + "their", "then", "there", "these", "they", "this", "to", "was", "will", "with" +}) + + +def default_tokenize( + text: str, + *, + remove_stopwords: bool = True, + stopwords: frozenset[str] | set[str] | None = None, +) -> list[str]: + """Tokenize text into search terms, preserving compound identifiers. + + Extracts compound identifiers (e.g. 'SKU-892-XZ', 'XPS-13-9310', + 'ConnectionResetError') as both full compound tokens and individual + sub-tokens for maximal recall on both exact matches and partial phrases. + """ + if not text: + return [] + + lowered = text.lower() + stop_set = stopwords if stopwords is not None else DEFAULT_STOPWORDS + + tokens: list[str] = [] + # 1) Extract compound patterns (e.g. sku-892-xz, 10.0.0.1) + compounds = _TOKEN_RE.findall(lowered) + for comp in compounds: + if remove_stopwords and comp in stop_set: + continue + tokens.append(comp) + # 2) If compound has delimiters, also emit individual sub-tokens + if any(c in comp for c in "-_.:"): + subtokens = _SUBTOKEN_RE.findall(comp) + for sub in subtokens: + if remove_stopwords and sub in stop_set: + continue + tokens.append(sub) + + return tokens + + +def _matches_filter(doc_meta: Metadata, filter_spec: Metadata) -> bool: + """Check if document metadata matches key-value filter conditions.""" + for key, expected in filter_spec.items(): + actual = doc_meta.get(key) + if isinstance(expected, dict): + # Operator support: e.g. {"$in": [...]}, {"$eq": val} + for op, val in expected.items(): + if op == "$eq" and actual != val: + return False + if op == "$ne" and actual == val: + return False + if op == "$in" and (not isinstance(val, (list, tuple, set)) or actual not in val): + return False + if op == "$nin" and isinstance(val, (list, tuple, set)) and actual in val: + return False + if op == "$gt" and (actual is None or actual <= val): + return False + if op == "$gte" and (actual is None or actual < val): + return False + if op == "$lt" and (actual is None or actual >= val): + return False + if op == "$lte" and (actual is None or actual > val): + return False + elif actual != expected: + return False + return True + + +class BM25Index: + """In-memory Okapi BM25 inverted index with zero third-party dependencies.""" + + def __init__( + self, + *, + k1: float = 1.5, + b: float = 0.75, + tokenizer: Callable[[str], list[str]] | None = None, + remove_stopwords: bool = True, + stopwords: frozenset[str] | set[str] | None = None, + ) -> None: + if k1 < 0: + raise ValueError("k1 must be non-negative") + if not (0.0 <= b <= 1.0): + raise ValueError("b must be in [0, 1]") + + self.k1 = k1 + self.b = b + self.tokenizer = tokenizer or default_tokenize + self.remove_stopwords = remove_stopwords + self.stopwords = stopwords + + # Stored documents: id -> {text, metadata, len, tf} + self._docs: dict[str, dict[str, Any]] = {} + # Term frequencies across corpus: term -> document count + self._doc_freqs: Counter[str] = Counter() + # Inverted index: term -> set(doc_id) + self._inverted_index: dict[str, set[str]] = {} + + self._total_docs: int = 0 + self._total_length: int = 0 + self._avg_doc_len: float = 0.0 + self._idf_cache: dict[str, float] = {} + + def __len__(self) -> int: + return self._total_docs + + def _tokenize(self, text: str) -> list[str]: + if self.tokenizer is default_tokenize: + return default_tokenize( + text, + remove_stopwords=self.remove_stopwords, + stopwords=self.stopwords, + ) + return self.tokenizer(text) + + def add_document( + self, + doc_id: str, + text: str | None, + metadata: Metadata | None = None, + ) -> None: + """Add or overwrite a document in the index.""" + if not doc_id: + raise ValueError("doc_id cannot be empty") + + raw_text = text or "" + tokens = self._tokenize(raw_text) + tf = Counter(tokens) + doc_len = len(tokens) + + # If overwriting existing document, remove it first + if doc_id in self._docs: + self.remove_document(doc_id) + + meta = dict(metadata or {}) + self._docs[doc_id] = { + "text": raw_text, + "metadata": meta, + "len": doc_len, + "tf": tf, + } + + for term in tf: + self._doc_freqs[term] += 1 + if term not in self._inverted_index: + self._inverted_index[term] = set() + self._inverted_index[term].add(doc_id) + + self._total_docs += 1 + self._total_length += doc_len + self._avg_doc_len = self._total_length / self._total_docs if self._total_docs > 0 else 0.0 + self._idf_cache.clear() + + def add_documents( + self, + documents: Sequence[Document | dict[str, Any] | tuple[str, str | None]], + ) -> None: + """Batch-insert multiple documents into the index.""" + for d in documents: + if isinstance(d, Document): + self.add_document(d.id, d.text, d.metadata) + elif isinstance(d, dict): + self.add_document( + str(d["id"]), + d.get("text"), + d.get("metadata"), + ) + elif isinstance(d, tuple): + self.add_document(d[0], d[1]) + else: + raise TypeError(f"Unsupported document type: {type(d)}") + + def remove_document(self, doc_id: str) -> bool: + """Remove a document from the index. Returns True if removed, False if absent.""" + doc = self._docs.pop(doc_id, None) + if doc is None: + return False + + doc_len = doc["len"] + tf: Counter[str] = doc["tf"] + + for term in tf: + self._doc_freqs[term] -= 1 + if self._doc_freqs[term] <= 0: + del self._doc_freqs[term] + if term in self._inverted_index: + self._inverted_index[term].discard(doc_id) + if not self._inverted_index[term]: + del self._inverted_index[term] + + self._total_docs -= 1 + self._total_length -= doc_len + self._avg_doc_len = self._total_length / self._total_docs if self._total_docs > 0 else 0.0 + self._idf_cache.clear() + return True + + def clear(self) -> None: + """Clear all indexed documents.""" + self._docs.clear() + self._doc_freqs.clear() + self._inverted_index.clear() + self._total_docs = 0 + self._total_length = 0 + self._avg_doc_len = 0.0 + self._idf_cache.clear() + + def idf(self, term: str) -> float: + """Calculate Robertson-Spärck Jones / Lucene smoothed inverse document frequency.""" + cached = self._idf_cache.get(term) + if cached is not None: + return cached + + df = self._doc_freqs.get(term, 0) + # Smoothing prevents negative weights on common terms + score = math.log(1.0 + (self._total_docs - df + 0.5) / (df + 0.5)) + self._idf_cache[term] = score + return score + + def score(self, query_tokens: list[str], doc_id: str) -> float: + """Compute the BM25 score of a single document for tokenized query terms.""" + doc = self._docs.get(doc_id) + if doc is None: + return 0.0 + + tf = doc["tf"] + doc_len = doc["len"] + avg_len = self._avg_doc_len or 1.0 + + score = 0.0 + for term in query_tokens: + count = tf.get(term, 0) + if count == 0: + continue + term_idf = self.idf(term) + # Okapi BM25 TF saturation formula + denom = count + self.k1 * (1.0 - self.b + self.b * (doc_len / avg_len)) + score += term_idf * (count * (self.k1 + 1.0)) / denom + + return score + + def search( + self, + query: str, + *, + top_k: int = 10, + filter: Metadata | None = None, + ) -> list[SearchResult]: + """Search the BM25 index and return ranked SearchResult items.""" + if top_k < 1: + raise ValueError("top_k must be >= 1") + if not query or not query.strip() or self._total_docs == 0: + return [] + + tokens = self._tokenize(query) + if not tokens: + return [] + + # Find candidate documents containing at least one term + candidate_ids: set[str] = set() + for term in tokens: + doc_ids = self._inverted_index.get(term) + if doc_ids: + candidate_ids.update(doc_ids) + + if not candidate_ids: + return [] + + scored: list[tuple[float, str]] = [] + for doc_id in candidate_ids: + doc_entry = self._docs[doc_id] + if filter and not _matches_filter(doc_entry["metadata"], filter): + continue + bm25_score = self.score(tokens, doc_id) + if bm25_score > 0.0: + scored.append((bm25_score, doc_id)) + + # Sort descending by score; tie-break deterministically by doc_id + scored.sort(key=lambda item: (-item[0], item[1])) + + results: list[SearchResult] = [] + for score, doc_id in scored[:top_k]: + doc_entry = self._docs[doc_id] + results.append( + SearchResult( + id=doc_id, + score=round(score, 6), + text=doc_entry["text"], + metadata=doc_entry["metadata"], + ) + ) + + return results diff --git a/src/dynavec/client.py b/src/dynavec/client.py index 093c0b7..cafa209 100644 --- a/src/dynavec/client.py +++ b/src/dynavec/client.py @@ -36,6 +36,7 @@ if TYPE_CHECKING: from .cache import BaseCache + from .retrievers import BM25HybridRetriever, BM25Retriever import numpy as np @@ -1050,6 +1051,61 @@ def as_hyde_retriever( **kw, ) + def as_bm25_retriever( + self, + namespace: str = "default", + **kw: Any, + ) -> BM25Retriever: + """Create a :class:`~dynavec.retrievers.BM25Retriever` bound to this client.""" + from .retrievers import BM25Retriever + + return BM25Retriever(self, namespace=namespace, **kw) + + def as_hybrid_retriever( + self, + namespace: str = "default", + *, + dense_weight: float = 1.0, + sparse_weight: float = 0.8, + **kw: Any, + ) -> BM25HybridRetriever: + """Create a :class:`~dynavec.retrievers.BM25HybridRetriever` bound to this client.""" + from .retrievers import BM25HybridRetriever + + return BM25HybridRetriever( + self, + namespace=namespace, + dense_weight=dense_weight, + sparse_weight=sparse_weight, + **kw, + ) + + def hybrid_search( + self, + query: str, + *, + top_k: int = 10, + namespace: str = "default", + dense_weight: float = 1.0, + sparse_weight: float = 0.8, + rrf_k: int = 60, + bm25_retriever: Any = None, + filter: Metadata | None = None, + use_cache: bool | None = None, + **kw: Any, + ) -> list[SearchResult]: + """Execute hybrid search combining dense ANN vector search and sparse BM25 lexical search.""" + retriever = self.as_hybrid_retriever( + namespace=namespace, + dense_weight=dense_weight, + sparse_weight=sparse_weight, + rrf_k=rrf_k, + top_k=top_k, + bm25_retriever=bm25_retriever, + **kw, + ) + return retriever.search(query, top_k=top_k, filter=filter, use_cache=use_cache) + def _resolve_query_vector(self, query: str | None, vector: list[float] | None) -> list[float]: if vector is not None: if len(vector) != self.config.dimension: diff --git a/src/dynavec/namespace.py b/src/dynavec/namespace.py index b085a09..81a9110 100644 --- a/src/dynavec/namespace.py +++ b/src/dynavec/namespace.py @@ -11,10 +11,11 @@ from collections.abc import Iterator, Sequence from typing import TYPE_CHECKING, Any, Literal, cast, overload -from .models import Document, ExplainedSearchResult, SearchResult, UpsertResult +from .models import Document, ExplainedSearchResult, Metadata, SearchResult, UpsertResult if TYPE_CHECKING: from .client import Dynavec + from .retrievers import BM25HybridRetriever, BM25Retriever class NamespaceView: @@ -141,6 +142,57 @@ def as_hyde_retriever( **kw, ) + def as_bm25_retriever(self, **kw: Any) -> BM25Retriever: + """Create a :class:`~dynavec.retrievers.BM25Retriever` pinned to this namespace.""" + from .retrievers import BM25Retriever + + return BM25Retriever(self._db, namespace=self._ns, **kw) + + def as_hybrid_retriever( + self, + *, + dense_weight: float = 1.0, + sparse_weight: float = 0.8, + **kw: Any, + ) -> BM25HybridRetriever: + """Create a :class:`~dynavec.retrievers.BM25HybridRetriever` pinned to this namespace.""" + from .retrievers import BM25HybridRetriever + + return BM25HybridRetriever( + self._db, + namespace=self._ns, + dense_weight=dense_weight, + sparse_weight=sparse_weight, + **kw, + ) + + def hybrid_search( + self, + query: str, + *, + top_k: int = 10, + dense_weight: float = 1.0, + sparse_weight: float = 0.8, + rrf_k: int = 60, + bm25_retriever: Any = None, + filter: Metadata | None = None, + use_cache: bool | None = None, + **kw: Any, + ) -> list[SearchResult]: + """Execute hybrid search combining dense ANN vector search and sparse BM25 lexical search.""" + return self._db.hybrid_search( + query, + top_k=top_k, + namespace=self._ns, + dense_weight=dense_weight, + sparse_weight=sparse_weight, + rrf_k=rrf_k, + bm25_retriever=bm25_retriever, + filter=filter, + use_cache=use_cache, + **kw, + ) + def export_namespace(self, output: Any, **kw: Any) -> int: return self._db.export_namespace(output, namespace=self._ns, **kw) diff --git a/src/dynavec/retrievers.py b/src/dynavec/retrievers.py index fed2960..79005cb 100644 --- a/src/dynavec/retrievers.py +++ b/src/dynavec/retrievers.py @@ -45,8 +45,9 @@ import numpy as np +from .bm25 import BM25Index from .exceptions import ConfigurationError -from .models import Metadata, SearchResult +from .models import Document, Metadata, SearchResult from .namespace import NamespaceView from .retrieval import reciprocal_rank_fusion @@ -58,7 +59,14 @@ OnGenerateError = Literal["fallback", "raise"] HyDEStrategy = Literal["average", "fuse"] -__all__ = ["QueryExpansionRetriever", "MultiQueryRetriever", "HyDERetriever"] +__all__ = [ + "QueryExpansionRetriever", + "MultiQueryRetriever", + "HyDERetriever", + "BM25Index", + "BM25Retriever", + "BM25HybridRetriever", +] def _clean_texts(original: str, candidates: object, limit: int) -> list[str]: @@ -469,3 +477,241 @@ async def _async_plan( ) -> 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) + + +class BM25Retriever: + """Sparse lexical retriever powered by an in-memory Okapi BM25 index.""" + + def __init__( + self, + source: Dynavec | NamespaceView, + *, + namespace: str = "default", + top_k: int = 10, + k1: float = 1.5, + b: float = 0.75, + index: BM25Index | None = None, + documents: Sequence[Document | dict[str, Any]] | None = None, + populate_from_store: bool = False, + ) -> None: + if top_k < 1: + raise ValueError("top_k must be >= 1") + + if isinstance(source, NamespaceView): + self._db = source._db + self._namespace = source.namespace + else: + self._db = source + self._namespace = namespace + + self.top_k = top_k + self.index = index or BM25Index(k1=k1, b=b) + + if documents: + self.index_documents(documents) + elif populate_from_store: + self.populate_from_store() + + @property + def namespace(self) -> str: + return self._namespace + + def index_documents( + self, + documents: Sequence[Document | dict[str, Any] | tuple[str, str | None]], + ) -> None: + """Add documents to the underlying BM25 index.""" + self.index.add_documents(documents) + + def populate_from_store(self, max_docs: int | None = None) -> int: + """Hydrate and index documents from the underlying DynamoDB / S3 store.""" + count = 0 + try: + for hit in self._db.list_vectors(self._namespace, hydrate=True): + if hit.text: + self.index.add_document(hit.id, hit.text, hit.metadata) + count += 1 + if max_docs is not None and count >= max_docs: + break + except Exception as exc: + logger.warning("Could not auto-populate BM25 index from store: %s", exc) + return count + + def search( + self, + query: str, + *, + top_k: int | None = None, + filter: Metadata | None = None, + ) -> list[SearchResult]: + """Perform lexical BM25 search.""" + if not isinstance(query, str) or not query.strip(): + raise ValueError("query must be a non-empty string") + k = top_k if top_k is not None else self.top_k + if k < 1: + raise ValueError("top_k must be >= 1") + return self.index.search(query, top_k=k, filter=filter) + + async def asearch( + self, + query: str, + *, + top_k: int | None = None, + filter: Metadata | None = None, + ) -> list[SearchResult]: + """Async variant of lexical BM25 search.""" + return await asyncio.to_thread(self.search, query, top_k=top_k, filter=filter) + + +class BM25HybridRetriever: + """Hybrid dense-sparse fusion retriever combining S3 Vectors ANN and BM25 via RRF.""" + + def __init__( + self, + source: Dynavec | NamespaceView, + *, + namespace: str = "default", + dense_weight: float = 1.0, + sparse_weight: float = 0.8, + rrf_k: int = 60, + top_k: int = 10, + per_query_k: int | None = None, + bm25_retriever: BM25Retriever | None = None, + k1: float = 1.5, + b: float = 0.75, + documents: Sequence[Document | dict[str, Any]] | None = None, + weights: Sequence[float] | Any | None = None, + ) -> None: + if top_k < 1: + raise ValueError("top_k must be >= 1") + if rrf_k < 1: + raise ValueError("rrf_k must be >= 1") + if per_query_k is not None and per_query_k < 1: + raise ValueError("per_query_k must be >= 1") + + if isinstance(source, NamespaceView): + self._db = source._db + self._namespace = source.namespace + else: + self._db = source + self._namespace = namespace + + if weights is not None: + if hasattr(weights, "weights"): + weights = weights.weights + if len(weights) != 2: + raise ValueError("weights must contain exactly 2 elements [dense_weight, sparse_weight]") + dense_weight, sparse_weight = float(weights[0]), float(weights[1]) + + if dense_weight <= 0 or sparse_weight <= 0: + raise ValueError("weights must be positive") + + self.dense_weight = dense_weight + self.sparse_weight = sparse_weight + self.rrf_k = rrf_k + self.top_k = top_k + self.per_query_k = per_query_k + + if bm25_retriever is not None: + self.bm25_retriever = bm25_retriever + else: + self.bm25_retriever = BM25Retriever( + source, + namespace=self._namespace, + k1=k1, + b=b, + documents=documents, + top_k=top_k, + ) + + self._local_executor: ThreadPoolExecutor | None = None + + @property + def namespace(self) -> str: + return self._namespace + + @property + def _executor(self) -> ThreadPoolExecutor: + if hasattr(self._db, "_executor") and self._db._executor is not None: + return self._db._executor + if self._local_executor is None: + self._local_executor = ThreadPoolExecutor(max_workers=8) + return self._local_executor + + def index_documents( + self, + documents: Sequence[Document | dict[str, Any] | tuple[str, str | None]], + ) -> None: + """Add documents to the BM25 index.""" + self.bm25_retriever.index_documents(documents) + + def search( + self, + query: str, + *, + top_k: int | None = None, + filter: Metadata | None = None, + use_cache: bool | None = None, + weights: Sequence[float] | Any | None = None, + ) -> list[SearchResult]: + """Execute parallel dense ANN and sparse BM25 retrieval and fuse via RRF.""" + if not isinstance(query, str) or not query.strip(): + raise ValueError("query must be a non-empty string") + effective_k = top_k if top_k is not None else self.top_k + if effective_k < 1: + raise ValueError("top_k must be >= 1") + depth = self.per_query_k or max(effective_k * 2, 10) + + w_dense = self.dense_weight + w_sparse = self.sparse_weight + if weights is not None: + if hasattr(weights, "weights"): + weights = weights.weights + if len(weights) != 2: + raise ValueError("weights must contain exactly 2 elements [dense_weight, sparse_weight]") + w_dense, w_sparse = float(weights[0]), float(weights[1]) + + # Run dense vector search and sparse BM25 search in parallel + f_dense = self._executor.submit( + self._db.search, + query, + top_k=depth, + namespace=self._namespace, + filter=filter, + use_cache=use_cache, + ) + f_sparse = self._executor.submit( + self.bm25_retriever.search, + query, + top_k=depth, + filter=filter, + ) + + dense_hits = f_dense.result() + sparse_hits = f_sparse.result() + + fused = reciprocal_rank_fusion( + [dense_hits, sparse_hits], + k=self.rrf_k, + weights=[w_dense, w_sparse], + ) + return fused[:effective_k] + + async def asearch( + self, + query: str, + *, + top_k: int | None = None, + filter: Metadata | None = None, + use_cache: bool | None = None, + weights: Sequence[float] | Any | None = None, + ) -> list[SearchResult]: + """Async variant of hybrid search.""" + return await asyncio.to_thread( + self.search, + query, + top_k=top_k, + filter=filter, + use_cache=use_cache, + weights=weights, + ) diff --git a/tests/test_bm25_retriever.py b/tests/test_bm25_retriever.py new file mode 100644 index 0000000..666d16e --- /dev/null +++ b/tests/test_bm25_retriever.py @@ -0,0 +1,383 @@ +import pytest +from test_client_inmemory import FakeDDB, FakeGraph, FakeS3 + +import dynavec.client as client_mod +from dynavec import ( + BM25HybridRetriever, + BM25Index, + BM25Retriever, + Document, + Dynavec, + DynavecConfig, + FitResult, +) +from dynavec.bm25 import default_tokenize +from dynavec.embeddings.base import Embedder + +# Test embeddings in 4D +E0 = [1.0, 0.0, 0.0, 0.0] +E1 = [0.0, 1.0, 0.0, 0.0] +E2 = [0.0, 0.0, 1.0, 0.0] +E3 = [0.0, 0.0, 0.0, 1.0] + + +class FakeEmbedder(Embedder): + dimension = 4 + + def __init__(self): + self.query_calls: list[str] = [] + + def embed_query(self, text: str) -> list[float]: + self.query_calls.append(text) + if "laptop" in text.lower(): + return E1 + if "server" in text.lower(): + return E2 + return E0 + + def embed_documents(self, texts: list[str]) -> list[list[float]]: + res = [] + for t in texts: + tl = t.lower() + if "laptop" in tl: + res.append(E1) + elif "server" in tl: + res.append(E2) + else: + res.append(E0) + return res + + +class FakeS3WithPages(FakeS3): + def list_pages(self, return_data=False, return_metadata=True, page_size=None): + keys = list(self._store.keys()) + chunk_size = page_size or 2 + for i in range(0, len(keys), chunk_size): + page = [] + for k in keys[i : i + chunk_size]: + vec, meta = self._store[k] + item = {"key": k} + if return_data: + item["data"] = {"float32": vec} + if return_metadata: + item["metadata"] = meta + page.append(item) + yield page + + +@pytest.fixture +def db(monkeypatch): + monkeypatch.setattr(client_mod, "S3VectorsStore", FakeS3WithPages) + monkeypatch.setattr(client_mod, "DynamoDBStore", FakeDDB) + monkeypatch.setattr(client_mod, "GraphStore", FakeGraph) + + cfg = DynavecConfig(vector_bucket="b", index="i", table="t", dimension=4, region="us-east-1") + client = Dynavec(cfg, embedder=FakeEmbedder()) + + docs = [ + Document( + id="doc1", + text="Dell XPS-13-9310 ultrabook laptop with Intel Core i7", + vector=E1, + metadata={"category": "hardware", "sku": "SKU-9310-XPS"}, + ), + Document( + id="doc2", + text="Lenovo ThinkPad X1 Carbon laptop for business travel", + vector=E1, + metadata={"category": "hardware", "sku": "SKU-X1-CARB"}, + ), + Document( + id="doc3", + text="High performance rack server Dell PowerEdge R750 with dual Xeon", + vector=E2, + metadata={"category": "server", "sku": "SKU-R750-PWR"}, + ), + Document( + id="doc4", + text="Python troubleshooting guide: handling ConnectionResetError during TLS handshake", + vector=E0, + metadata={"category": "software", "code": "ERR-104"}, + ), + ] + client.upsert(docs) + return client + + +# --------------------------------------------------------------------------- +# Tokenizer tests +# --------------------------------------------------------------------------- + +def test_default_tokenize_basic(): + tokens = default_tokenize("The quick brown fox jumps over the lazy dog.") + # Stopwords ("the", "over") filtered out + assert "quick" in tokens + assert "brown" in tokens + assert "fox" in tokens + assert "dog" in tokens + assert "the" not in tokens + + +def test_default_tokenize_compound_identifiers(): + tokens = default_tokenize("Order item SKU-892-XZ and Dell XPS-13-9310") + # Exact compound preservation + assert "sku-892-xz" in tokens + assert "xps-13-9310" in tokens + # Sub-token preservation + assert "sku" in tokens + assert "892" in tokens + assert "xz" in tokens + assert "xps" in tokens + assert "13" in tokens + assert "9310" in tokens + + +def test_default_tokenize_custom_stopwords(): + tokens = default_tokenize("apple banana cherry", stopwords=frozenset({"banana"})) + assert tokens == ["apple", "cherry"] + + tokens_no_removal = default_tokenize("apple banana", remove_stopwords=False) + assert "banana" in tokens_no_removal + + +def test_default_tokenize_empty_and_special(): + assert default_tokenize("") == [] + assert default_tokenize(" ") == [] + assert default_tokenize("!@#$%^&*()") == [] + + +# --------------------------------------------------------------------------- +# BM25Index tests +# --------------------------------------------------------------------------- + +def test_bm25_index_validation(): + with pytest.raises(ValueError, match="k1 must be non-negative"): + BM25Index(k1=-1.0) + with pytest.raises(ValueError, match="b must be in"): + BM25Index(b=1.5) + with pytest.raises(ValueError, match="b must be in"): + BM25Index(b=-0.1) + + +def test_bm25_index_crud(): + index = BM25Index() + assert len(index) == 0 + + with pytest.raises(ValueError, match="doc_id cannot be empty"): + index.add_document("", "some text") + + index.add_document("doc1", "first document with keyword python") + index.add_document("doc2", "second document with keyword rust") + assert len(index) == 2 + + # Search keyword + res = index.search("python") + assert len(res) == 1 + assert res[0].id == "doc1" + assert res[0].score > 0 + + # Overwrite doc1 + index.add_document("doc1", "updated document without keyword") + assert len(index) == 2 + res_after = index.search("python") + assert len(res_after) == 0 + + # Remove document + assert index.remove_document("doc1") is True + assert index.remove_document("nonexistent") is False + assert len(index) == 1 + + # Clear + index.clear() + assert len(index) == 0 + assert index.search("rust") == [] + + +def test_bm25_index_batch_add(): + index = BM25Index() + docs = [ + Document(id="d1", text="alpha beta", metadata={"cat": "a"}), + {"id": "d2", "text": "beta gamma", "metadata": {"cat": "b"}}, + ("d3", "gamma delta"), + ] + index.add_documents(docs) + assert len(index) == 3 + + with pytest.raises(TypeError, match="Unsupported document type"): + index.add_documents([123]) # type: ignore + + +def test_bm25_index_filtering(): + index = BM25Index() + index.add_document("d1", "database query optimization", {"type": "db", "priority": 1}) + index.add_document("d2", "database index design", {"type": "db", "priority": 5}) + index.add_document("d3", "database backup procedures", {"type": "ops", "priority": 3}) + + # Exact filter + hits = index.search("database", filter={"type": "db"}) + hit_ids = {h.id for h in hits} + assert hit_ids == {"d1", "d2"} + + # Operator filter: $in + hits_in = index.search("database", filter={"type": {"$in": ["ops"]}}) + assert [h.id for h in hits_in] == ["d3"] + + # Operator filter: $gte + hits_gte = index.search("database", filter={"priority": {"$gte": 3}}) + hit_gte_ids = {h.id for h in hits_gte} + assert hit_gte_ids == {"d2", "d3"} + + # Operator filter: $eq and $ne + hits_ne = index.search("database", filter={"type": {"$ne": "db"}}) + assert [h.id for h in hits_ne] == ["d3"] + + +def test_bm25_index_empty_query(): + index = BM25Index() + index.add_document("d1", "sample text") + assert index.search("") == [] + assert index.search(" ") == [] + with pytest.raises(ValueError, match="top_k must be >= 1"): + index.search("sample", top_k=0) + + +# --------------------------------------------------------------------------- +# BM25Retriever tests +# --------------------------------------------------------------------------- + +def test_bm25_retriever_basic(db): + retriever = BM25Retriever(db) + retriever.index_documents([ + Document(id="doc1", text="Dell XPS-13-9310 laptop"), + Document(id="doc2", text="Lenovo ThinkPad laptop"), + ]) + + results = retriever.search("XPS-13-9310") + assert len(results) >= 1 + assert results[0].id == "doc1" + assert "XPS-13-9310" in results[0].text + + with pytest.raises(ValueError, match="query must be a non-empty string"): + retriever.search("") + with pytest.raises(ValueError, match="top_k must be >= 1"): + retriever.search("test", top_k=-1) + + +@pytest.mark.asyncio +async def test_bm25_retriever_async(db): + retriever = BM25Retriever(db) + retriever.index_documents([ + Document(id="d1", text="ConnectionResetError encountered"), + Document(id="d2", text="HTTP 500 internal server error"), + ]) + + results = await retriever.asearch("ConnectionResetError") + assert len(results) == 1 + assert results[0].id == "d1" + + +def test_bm25_retriever_populate_from_store(db): + retriever = BM25Retriever(db, populate_from_store=True) + # The fake store has 4 documents seeded + assert len(retriever.index) == 4 + + results = retriever.search("XPS-13-9310") + assert len(results) >= 1 + assert results[0].id == "doc1" + + +# --------------------------------------------------------------------------- +# BM25HybridRetriever tests +# --------------------------------------------------------------------------- + +def test_hybrid_retriever_initialization_and_weights(db): + retriever = BM25HybridRetriever(db, dense_weight=1.2, sparse_weight=0.6) + assert retriever.dense_weight == 1.2 + assert retriever.sparse_weight == 0.6 + + with pytest.raises(ValueError, match="weights must contain exactly 2 elements"): + BM25HybridRetriever(db, weights=[1.0]) + + with pytest.raises(ValueError, match="weights must be positive"): + BM25HybridRetriever(db, weights=[-1.0, 1.0]) + + # FitResult compatibility + fit_res = FitResult(weights=[0.8, 0.4], score=0.95, method="grid", n_evaluations=10) + retriever_fitted = BM25HybridRetriever(db, weights=fit_res) + assert retriever_fitted.dense_weight == 0.8 + assert retriever_fitted.sparse_weight == 0.4 + + +def test_hybrid_retriever_search_lexical_boost(db): + # Seed BM25 retriever with same docs + hybrid = BM25HybridRetriever(db, dense_weight=1.0, sparse_weight=1.5) + hybrid.bm25_retriever.populate_from_store() + + # Query for exact model: "XPS-13-9310" + # S3 vector search maps "XPS-13-9310" to E0 (ranking doc4 top), + # but BM25 strongly ranks doc1 top. Fused RRF elevates doc1 to rank 1. + results = hybrid.search("XPS-13-9310", top_k=2) + assert len(results) > 0 + assert results[0].id == "doc1" + + +def test_hybrid_retriever_weights_override_and_fit_result(db): + hybrid = BM25HybridRetriever(db, dense_weight=1.0, sparse_weight=0.5) + hybrid.bm25_retriever.populate_from_store() + + fit_res = FitResult(weights=[1.5, 0.5], score=0.92, method="grid", n_evaluations=5) + results = hybrid.search("laptop", weights=fit_res, top_k=3) + assert len(results) > 0 + + with pytest.raises(ValueError, match="weights must contain exactly 2 elements"): + hybrid.search("laptop", weights=[1.0, 2.0, 3.0]) + + +@pytest.mark.asyncio +async def test_hybrid_retriever_asearch(db): + hybrid = BM25HybridRetriever(db) + hybrid.bm25_retriever.populate_from_store() + + results = await hybrid.asearch("server Dell PowerEdge", top_k=2) + assert len(results) > 0 + assert results[0].id == "doc3" + + +# --------------------------------------------------------------------------- +# Client & NamespaceView Ergonomics tests +# --------------------------------------------------------------------------- + +def test_client_retriever_ergonomics(db): + # as_bm25_retriever + bm25 = db.as_bm25_retriever(populate_from_store=True) + assert isinstance(bm25, BM25Retriever) + assert len(bm25.index) == 4 + + # as_hybrid_retriever + hybrid = db.as_hybrid_retriever(dense_weight=1.0, sparse_weight=0.9, bm25_retriever=bm25) + assert isinstance(hybrid, BM25HybridRetriever) + assert hybrid.sparse_weight == 0.9 + + # hybrid_search directly on client + res = db.hybrid_search("laptop Dell XPS", bm25_retriever=bm25, top_k=2) + assert len(res) <= 2 + assert res[0].id in ("doc1", "doc2") + + +def test_namespace_view_retriever_ergonomics(db): + ns = db.namespace("default") + + # as_bm25_retriever on namespace + bm25 = ns.as_bm25_retriever(populate_from_store=True) + assert isinstance(bm25, BM25Retriever) + assert bm25.namespace == "default" + + # as_hybrid_retriever on namespace + hybrid = ns.as_hybrid_retriever(dense_weight=1.0, sparse_weight=0.7, bm25_retriever=bm25) + assert isinstance(hybrid, BM25HybridRetriever) + assert hybrid.namespace == "default" + + # hybrid_search on namespace + res = ns.hybrid_search("ConnectionResetError", bm25_retriever=bm25, top_k=1) + assert len(res) == 1 + assert res[0].id == "doc4"