diff --git a/README.md b/README.md index 82d91fa75..e43bd4692 100644 --- a/README.md +++ b/README.md @@ -37,7 +37,8 @@ search. - **Local-first.** Plain text on your disk. Forever. - **Two-way.** AI and humans write to the same files; sync keeps them in step. - **A real knowledge graph.** Observations and wikilinks compound into context. -- **Semantic search.** Find notes by meaning, not just keywords. +- **Semantic search.** Find notes by meaning, not just keywords, with optional + cross-encoder reranking for higher-quality vector and hybrid results. - **MCP-native.** Works with every major AI client and IDE. - **Progressive tool discovery.** Every tool is tagged with behavior hints (read-only, destructive, idempotent) so agents pick the right tool on @@ -357,6 +358,8 @@ Try a prompt: - **Semantic vector search.** Find notes by meaning, not just keywords. Hybrid full-text + vector ranking with FastEmbed embeddings, on SQLite or Postgres. +- **Optional search reranking.** Rescore the strongest vector and hybrid + candidates with a local FastEmbed cross-encoder or a LiteLLM-backed provider. - **Schema system.** Infer, validate, and diff the structure of your knowledge base with `schema_infer`, `schema_validate`, `schema_diff`. - **Per-project cloud routing.** Route individual projects through the cloud @@ -374,6 +377,36 @@ Try a prompt: Full [CHANGELOG](CHANGELOG.md) for v0.18 → v0.20. +## Optional cross-encoder reranking + +Reranking adds a second relevance pass after vector or hybrid retrieval. It is +disabled by default because it adds inference latency and, for the local +provider, a first-run model download. Text, title, and permalink searches keep +their existing ranking. + +Enable the default local FastEmbed reranker: + +```bash +export BASIC_MEMORY_SEMANTIC_SEARCH_ENABLED=true +export BASIC_MEMORY_RERANKER_ENABLED=true +``` + +The default model is `jinaai/jina-reranker-v1-tiny-en`. To use a hosted +reranker through LiteLLM instead: + +```bash +export BASIC_MEMORY_SEMANTIC_SEARCH_ENABLED=true +export BASIC_MEMORY_RERANKER_ENABLED=true +export BASIC_MEMORY_RERANKER_PROVIDER=litellm +export BASIC_MEMORY_RERANKER_MODEL=cohere/rerank-v3.5 +export COHERE_API_KEY=... +``` + +The feature fails fast on invalid configuration and does not silently fall +back to retrieval order when an enabled provider fails. See the +[semantic search guide](docs/semantic-search.md#cross-encoder-reranking) for +provider setup, all settings, tuning, pagination, and failure behavior. + ## Why Basic Memory Most LLM conversations are ephemeral. You ask a question, get an answer, then diff --git a/docs/semantic-search.md b/docs/semantic-search.md index fb54f4d37..e0424657a 100644 --- a/docs/semantic-search.md +++ b/docs/semantic-search.md @@ -322,6 +322,117 @@ Score-based fusion uses the formula `max(vec, fts) + bonus * min(vec, fts)` to p | `vector` | Conceptual queries, paraphrase matching, exploratory searches | | `hybrid` | General-purpose search combining precision and recall | +## Cross-Encoder Reranking + +Reranking is an optional second stage for vector and hybrid search. Initial +retrieval finds a candidate pool efficiently; a cross-encoder then reads each +query and candidate together and replaces the leading candidates' scores with +more precise relevance scores. + +Reranking is disabled by default. It requires semantic search and does not +change `text`, `title`, or `permalink` search behavior. + +### Quick Start + +Use the default local FastEmbed provider: + +```bash +export BASIC_MEMORY_SEMANTIC_SEARCH_ENABLED=true +export BASIC_MEMORY_RERANKER_ENABLED=true +``` + +The first reranked search downloads the model when it is not already cached, +then loads `jinaai/jina-reranker-v1-tiny-en`. Later searches reuse the +process-wide model instance and cache. + +To use a hosted reranker through LiteLLM: + +```bash +export BASIC_MEMORY_SEMANTIC_SEARCH_ENABLED=true +export BASIC_MEMORY_RERANKER_ENABLED=true +export BASIC_MEMORY_RERANKER_PROVIDER=litellm +export BASIC_MEMORY_RERANKER_MODEL=cohere/rerank-v3.5 +export COHERE_API_KEY=... +``` + +LiteLLM model names must use explicit `provider/model` routing. Standard +provider environment variables, such as `COHERE_API_KEY`, work normally. You +can instead set `BASIC_MEMORY_RERANKER_API_KEY` to pass a credential directly +to LiteLLM. + +### Providers + +| Provider | Runs | Default model | Tradeoff | +|---|---|---|---| +| `fastembed` | Locally with ONNX | `jinaai/jina-reranker-v1-tiny-en` | No API key or per-query cost; downloads a model on first use and adds local inference latency. | +| `litellm` | Hosted or self-hosted API | No implicit hosted default | Supports LiteLLM rerank providers such as Cohere, Jina, and Voyage; adds network latency, provider cost, and credential requirements. | + +### Configuration Reference + +All settings use the `BASIC_MEMORY_` environment prefix: + +| Config Field | Environment Variable | Default | Description | +|---|---|---|---| +| `reranker_enabled` | `BASIC_MEMORY_RERANKER_ENABLED` | `false` | Enable reranking for vector and hybrid search. Requires semantic search. | +| `reranker_provider` | `BASIC_MEMORY_RERANKER_PROVIDER` | `fastembed` | `fastembed` for a local ONNX cross-encoder or `litellm` for an API provider. | +| `reranker_model` | `BASIC_MEMORY_RERANKER_MODEL` | `jinaai/jina-reranker-v1-tiny-en` | Model identifier. LiteLLM requires explicit `provider/model` routing. | +| `reranker_candidates` | `BASIC_MEMORY_RERANKER_CANDIDATES` | `20` | Number of leading retrieval results rescored on every page. Larger values can improve recall but increase latency and provider usage. | +| `reranker_max_document_chars` | `BASIC_MEMORY_RERANKER_MAX_DOCUMENT_CHARS` | `0` | Maximum characters sent per candidate. `0` sends the full matched text; a positive cap bounds latency and request size. | +| `reranker_api_base` | `BASIC_MEMORY_RERANKER_API_BASE` | Unset | Optional custom endpoint for the LiteLLM provider. | +| `reranker_api_key` | `BASIC_MEMORY_RERANKER_API_KEY` | Unset | Optional credential passed directly to LiteLLM. When unset, LiteLLM resolves provider credentials from its normal environment variables. | + +Configuration is validated at startup. Basic Memory rejects unsupported +providers, blank models, unavailable FastEmbed models or dependencies, a +LiteLLM model without a provider prefix, and reranking without semantic search. + +### Ranking and Pagination Behavior + +- Basic Memory reranks the same fixed leading candidate window on every + non-empty page. Results outside that window retain retrieval order and are + calibrated at or below the reranked floor. When that floor is zero, stable + prefix-before-tail ordering breaks the tie. +- The matched chunk is placed before optional title context in the reranker + document. A positive document-character cap therefore preserves the passage + that caused the retrieval match. +- Reranker scores replace retrieval scores for the reranked prefix and remain + within the public `[0, 1]` search score range. +- Multi-project MCP search reranks inside each project's search service, then + merges the returned project-owned scores. It does not perform a second global + rerank in the MCP process. + +These rules keep independently fetched pages stable while still allowing +larger requests to expand the untouched retrieval tail. + +### Failure Behavior + +An enabled reranker is part of the requested ranking contract, so Basic Memory +does not silently return retrieval order when it fails: + +- Temporary provider, rate-limit, connection, timeout, or first-download + failures return HTTP 503. Retry the search after the provider recovers. +- Malformed or incomplete provider responses return HTTP 502. +- Authentication, model, dependency, and permanent configuration failures + surface directly instead of being treated as transient. +- In multi-project search, a retryable failure from any project aborts the + aggregate page instead of returning a partial result set whose ordering could + change on retry. + +### Tuning + +Start with the defaults, then tune only if measurements justify it: + +- Increase `reranker_candidates` when relevant results enter the retrieval set + but remain outside the desired cutoff. This increases local inference time or + hosted provider usage. +- Set `reranker_max_document_chars` to a positive value such as `1000` to + bound latency and hosted request size for long notes. The matched chunk comes + first, so a modest cap retains the strongest retrieval signal. +- Keep reranking disabled when retrieval latency matters more than the + additional ranking pass. + +Changing reranker providers, models, candidate counts, or document caps does +not change stored embeddings, so it does not require `bm reindex --embeddings`. + ## The Reindex Command The `bm reindex` command rebuilds search indexes without dropping the database. diff --git a/src/basic_memory/api/v2/routers/search_router.py b/src/basic_memory/api/v2/routers/search_router.py index 902ac4c62..83ff1e2bd 100644 --- a/src/basic_memory/api/v2/routers/search_router.py +++ b/src/basic_memory/api/v2/routers/search_router.py @@ -11,6 +11,8 @@ import logfire from basic_memory.api.v2.utils import to_search_results from basic_memory.repository.semantic_errors import ( + RerankProviderContractError, + RerankTransientError, SemanticDependenciesMissingError, SemanticSearchDisabledError, ) @@ -91,6 +93,15 @@ async def search( raise HTTPException(status_code=400, detail=str(exc)) from exc except SemanticDependenciesMissingError as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc + except RerankTransientError as exc: + # Returning raw retrieval order would make pagination inconsistent with + # earlier reranked pages. Preserve ordering semantics and make the outage + # explicitly retryable instead. + raise HTTPException(status_code=503, detail=str(exc)) from exc + except RerankProviderContractError as exc: + # Upstream reranker returned a malformed response — an upstream fault, not a + # client error and not a transient outage (those map to a retryable 503). + raise HTTPException(status_code=502, detail=str(exc)) from exc except ValueError as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc diff --git a/src/basic_memory/config_models.py b/src/basic_memory/config_models.py index 188039cd2..d39bcfcd3 100644 --- a/src/basic_memory/config_models.py +++ b/src/basic_memory/config_models.py @@ -56,6 +56,11 @@ class DatabaseBackend(str, Enum): POSTGRES = "postgres" +# Default reranker model — a small local fastembed cross-encoder. +# LiteLLM cannot route this name, so the litellm provider requires an explicit override. +DEFAULT_FASTEMBED_RERANK_MODEL = "jinaai/jina-reranker-v1-tiny-en" + + def _default_semantic_search_enabled() -> bool: """Enable semantic search by default when required local semantic dependencies exist.""" required_modules = ("fastembed", "sqlite_vec") @@ -375,6 +380,47 @@ def __init__(self, **data: Any) -> None: ... "When unset, defaults to 'hybrid' if semantic search is enabled, otherwise 'text'.", ) + # Reranker configuration (cross-encoder rescoring of the top vector/hybrid candidates) + reranker_enabled: bool = Field( + default=False, + description="Enable cross-encoder reranking of vector/hybrid search candidates. " + "Off by default: adds latency and a first-run model download. Requires semantic search.", + ) + reranker_provider: str = Field( + default="fastembed", + description="Reranker provider: 'fastembed' (local ONNX cross-encoder) or 'litellm' " + "(Cohere/Jina/Voyage/etc. via API).", + ) + reranker_model: str = Field( + default=DEFAULT_FASTEMBED_RERANK_MODEL, + description="Reranker model identifier. For litellm use the 'provider/model' form, " + "e.g. 'cohere/rerank-v3.5'.", + ) + reranker_max_document_chars: int = Field( + default=0, + description="Max characters of each candidate's text passed to the cross-encoder. " + "0 (default) sends the full matched text — the model still truncates to its own token " + "limit. Set a positive cap (e.g. ~1000) to bound rerank latency on long notes; the " + "most-relevant matched chunk leads the text, so a modest cap keeps most of the signal.", + ge=0, + ) + reranker_api_base: str | None = Field( + default=None, + description="Optional custom API base URL for the litellm reranker provider " + "(self-hosted OpenAI-compatible rerank endpoints).", + ) + reranker_api_key: str | None = Field( + default=None, + description="Optional API key passed to the litellm reranker provider. When unset, " + "litellm resolves credentials from provider environment variables.", + ) + reranker_candidates: int = Field( + default=20, + description="Number of top retrieval candidates to rescore with the reranker before " + "returning the requested page. Larger widens recall at the cost of latency.", + gt=0, + ) + # Database connection pool configuration (Postgres only) db_pool_size: int = Field( default=20, @@ -865,6 +911,66 @@ def project_list(self) -> List[ProjectConfig]: # pragma: no cover for name, entry in self.projects.items() ] + @model_validator(mode="after") + def validate_reranker_config(self) -> "BasicMemoryConfig": + """Fail fast on reranker configs that cannot work. + + - Reranking runs only on vector/hybrid retrieval, so it needs semantic search; + accepting reranker_enabled=True without it would silently never rerank. + - FastEmbed exposes a finite registered model catalog, so reject typos while + loading config instead of deferring them to the first search request. + - The default model is a local fastembed cross-encoder that litellm cannot + route, so the litellm provider needs an explicit provider/model model id. + """ + if not self.reranker_enabled: + return self + if not self.semantic_search_enabled: + raise ValueError( + "reranker_enabled=True requires semantic_search_enabled=True " + "(reranking operates on vector/hybrid search results)." + ) + provider = self.reranker_provider.strip().lower() + if provider not in {"fastembed", "litellm"}: + raise ValueError("reranker_provider must be one of: fastembed, litellm") + model = self.reranker_model.strip() + if not model: + raise ValueError("reranker_model must not be blank when reranker_enabled=True") + self.reranker_provider = provider + self.reranker_model = model + if provider == "fastembed": + try: + from fastembed.rerank.cross_encoder import TextCrossEncoder + except ImportError as exc: + raise ValueError( + "reranker_provider='fastembed' requires the fastembed package. " + "Install/update basic-memory to include semantic dependencies." + ) from exc + + supported_models = { + entry["model"] for entry in TextCrossEncoder.list_supported_models() + } + if model not in supported_models: + supported_names = ", ".join(sorted(supported_models)) + raise ValueError( + f"Unsupported FastEmbed reranker model {model!r}. " + f"Supported models: {supported_names}" + ) + if provider == "litellm" and model == DEFAULT_FASTEMBED_RERANK_MODEL: + raise ValueError( + "reranker_provider='litellm' requires an explicit reranker_model in " + "'provider/model' form (e.g. 'cohere/rerank-v3.5'); the default " + f"{DEFAULT_FASTEMBED_RERANK_MODEL!r} is a local fastembed model " + "litellm cannot route." + ) + if provider == "litellm": + provider_name, separator, model_name = model.partition("/") + if not separator or not provider_name or not model_name: + raise ValueError( + "reranker_provider='litellm' requires an explicit reranker_model in " + "'provider/model' form (e.g. 'cohere/rerank-v3.5')." + ) + return self + @model_validator(mode="after") def ensure_project_paths_exists(self) -> "BasicMemoryConfig": # pragma: no cover """Ensure project paths exist. diff --git a/src/basic_memory/mcp/tools/search.py b/src/basic_memory/mcp/tools/search.py index dd268f30f..41224cc4e 100644 --- a/src/basic_memory/mcp/tools/search.py +++ b/src/basic_memory/mcp/tools/search.py @@ -6,6 +6,7 @@ from uuid import UUID import logfire +from httpx import HTTPStatusError from loguru import logger from fastmcp import Context from pydantic import AliasChoices, BeforeValidator, Field @@ -38,6 +39,8 @@ SearchRetrievalMode, ) +_SERVICE_UNAVAILABLE_HEADING = "# Search Failed - Service Temporarily Unavailable" + def _default_search_type() -> str: """Pick default search mode from config, falling back to auto-detection. @@ -55,6 +58,31 @@ def _default_search_type() -> str: return "hybrid" if config.semantic_search_enabled else "text" +def _is_service_unavailable_error(error: BaseException) -> bool: + """Return whether an explicit HTTP cause marks a retryable service outage.""" + current: BaseException | None = error + while current is not None: + if isinstance(current, HTTPStatusError): + return current.response.status_code == 503 + current = current.__cause__ + return False + + +def _format_service_unavailable_response(project: str, error_message: str, query: str) -> str: + """Keep retryable outages distinct so fan-out never turns them into partial success.""" + return dedent(f""" + {_SERVICE_UNAVAILABLE_HEADING} + + Search for '{query}' in project '{project}' could not complete: {error_message} + + No partial results were returned because retrying with a changing project set + could duplicate or skip results across pages. + + ## Next step + Retry the same search after the service recovers. + """).strip() + + def _format_search_error_response( project: str, error_message: str, query: str, search_type: str = "text" ) -> str: @@ -593,6 +621,8 @@ async def _search_all_projects( continue if isinstance(results, str): + if results.startswith(_SERVICE_UNAVAILABLE_HEADING): + return results if not results.startswith("# Search Failed"): return results logger.warning( @@ -608,6 +638,9 @@ async def _search_all_projects( any_project_has_more = any_project_has_more or results.get("has_more") is True merged_results.extend(_qualify_results_for_project(raw_results, project_ref)) + # Each project owns retrieval and optional reranking behind its typed API client. + # The MCP process only merges returned scores; it must not instantiate repository + # providers with local credentials for content fetched through another route. sorted_results = sorted(merged_results, key=_result_score, reverse=True) start = (requested_page - 1) * requested_page_size end = start + requested_page_size @@ -1179,6 +1212,10 @@ async def search_notes( logger.error( f"Search failed for query '{query or ''}': {e}, project: {active_project.name}" ) + if _is_service_unavailable_error(e): + return _format_service_unavailable_response( + active_project.name, str(e), query or "" + ) # Return formatted error message as string for better user experience return _format_search_error_response( active_project.name, str(e), query or "", effective_search_type diff --git a/src/basic_memory/redaction.py b/src/basic_memory/redaction.py index 2b4e649a6..a81c335d4 100644 --- a/src/basic_memory/redaction.py +++ b/src/basic_memory/redaction.py @@ -9,11 +9,23 @@ from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse # Fields in BasicMemoryConfig that contain secrets and must never be surfaced. -SECRET_FIELDS = frozenset({"cloud_api_key", "semantic_embedding_api_key"}) +SECRET_FIELDS = frozenset( + { + "cloud_api_key", + "reranker_api_key", + "semantic_embedding_api_key", + } +) # Fields whose values are URLs that may embed user:password credentials. # The userinfo component is stripped before surfacing. -URL_FIELDS = frozenset({"database_url", "semantic_embedding_api_base"}) +URL_FIELDS = frozenset( + { + "database_url", + "reranker_api_base", + "semantic_embedding_api_base", + } +) _SECRET_QUERY_KEYS = frozenset( { diff --git a/src/basic_memory/repository/fastembed_rerank_provider.py b/src/basic_memory/repository/fastembed_rerank_provider.py new file mode 100644 index 000000000..d7985dbd7 --- /dev/null +++ b/src/basic_memory/repository/fastembed_rerank_provider.py @@ -0,0 +1,160 @@ +"""FastEmbed-based local cross-encoder reranker provider.""" + +from __future__ import annotations + +import asyncio +import math +import os +from typing import TYPE_CHECKING, Any + +from loguru import logger +from requests import exceptions as requests_exceptions + +from basic_memory.repository.rerank_provider import RerankProvider, validate_rerank_scores +from basic_memory.repository.semantic_errors import ( + RerankProviderContractError, + RerankTransientError, + SemanticDependenciesMissingError, +) + +if TYPE_CHECKING: + from fastembed.rerank.cross_encoder import TextCrossEncoder # pragma: no cover + + +_TRANSIENT_DOWNLOAD_STATUS_CODES = frozenset({408, 425, 429}) +_TRUE_ENV_VALUES = frozenset({"1", "ON", "YES", "TRUE"}) + + +def _consume_model_load_exception(task: asyncio.Task["TextCrossEncoder"]) -> None: + """Retrieve background load failures when the request that started them was cancelled.""" + if not task.cancelled(): + task.exception() + + +def _is_transient_model_load_error( + exc: requests_exceptions.RequestException | ValueError, + model_name: str, +) -> bool: + """Classify recoverable first-download failures without hiding model/config errors.""" + if isinstance(exc, (requests_exceptions.ConnectionError, requests_exceptions.Timeout)): + return True + if isinstance(exc, requests_exceptions.HTTPError): + status_code = exc.response.status_code if exc.response is not None else None + return status_code in _TRANSIENT_DOWNLOAD_STATUS_CODES or ( + status_code is not None and status_code >= 500 + ) + # FastEmbed retries Hugging Face/GCS downloads internally and, after exhausting + # both sources, replaces the transport cause with this model-specific ValueError. + if str(exc) != f"Could not load model {model_name} from any source.": + return False + + # Trigger: Hugging Face offline mode turns an absent cache entry into the same + # exhausted-source ValueError as a temporary download failure. + # Why: an explicitly offline process cannot recover by retrying the model load. + # Outcome: preserve the permanent configuration/cache error for the caller. + offline_mode = any( + os.environ.get(name, "").strip().upper() in _TRUE_ENV_VALUES + for name in ("HF_HUB_OFFLINE", "TRANSFORMERS_OFFLINE") + ) + return not offline_mode + + +class FastEmbedRerankProvider(RerankProvider): + """Local ONNX cross-encoder reranker backed by FastEmbed. + + Shares the FastEmbed model cache with the embedding provider — reranker + models live under a distinct model subdir, so the two never collide. + """ + + def __init__( + self, + model_name: str = "Xenova/ms-marco-MiniLM-L-6-v2", + *, + cache_dir: str | None = None, + threads: int | None = None, + ) -> None: + self.model_name = model_name + self.cache_dir = cache_dir + self.threads = threads + self._model: TextCrossEncoder | None = None + self._model_load_task: asyncio.Task[TextCrossEncoder] | None = None + # Serialize the one-time model load; concurrent first queries must not each + # construct (and download) the ONNX model. + self._model_lock = asyncio.Lock() + + def runtime_log_attrs(self) -> dict[str, Any]: + return {"model_name": self.model_name, "threads": self.threads} + + def _create_model(self) -> "TextCrossEncoder": + try: + from fastembed.rerank.cross_encoder import TextCrossEncoder + except ImportError as exc: # pragma: no cover - exercised via monkeypatched tests + raise SemanticDependenciesMissingError( + "fastembed package is missing. " + "Install/update basic-memory to include semantic dependencies: " + "pip install -U basic-memory" + ) from exc + model_kwargs: dict[str, Any] = {"model_name": self.model_name} + if self.cache_dir is not None: + model_kwargs["cache_dir"] = self.cache_dir + if self.threads is not None: + model_kwargs["threads"] = self.threads + return TextCrossEncoder(**model_kwargs) + + async def _load_model_once(self) -> "TextCrossEncoder": + """Own one model construction independently of any requesting task.""" + try: + try: + model = await asyncio.to_thread(self._create_model) + except (requests_exceptions.RequestException, ValueError) as exc: + if not _is_transient_model_load_error(exc, self.model_name): + raise + raise RerankTransientError( + f"FastEmbed reranker model download failed temporarily: {self.model_name}" + ) from exc + self._model = model + logger.info("FastEmbed reranker loaded: model_name={model}", model=self.model_name) + return model + finally: + # The worker owns cleanup, so cancellation of the request awaiting it + # cannot erase the in-flight state and trigger a duplicate construction. + if self._model_load_task is asyncio.current_task(): + self._model_load_task = None + + async def _load_model(self) -> "TextCrossEncoder": + if self._model is not None: + return self._model + async with self._model_lock: + if self._model is not None: + return self._model + if self._model_load_task is None: + load_task = asyncio.create_task(self._load_model_once()) + load_task.add_done_callback(_consume_model_load_exception) + self._model_load_task = load_task + else: + load_task = self._model_load_task + + # Shield the shared construction from request cancellation. The cancelled + # caller still exits promptly while the next caller joins the same worker. + return await asyncio.shield(load_task) + + async def rerank(self, query: str, documents: list[str]) -> list[float]: + if not documents: + return [] + model = await self._load_model() + # TextCrossEncoder.rerank is sync/CPU-bound and yields one raw logit per doc + # in input order; run it off the event loop and squash to [0, 1] so callers + # get a bounded relevance on the same scale as API-based rerankers. + scores = await asyncio.to_thread( + lambda: [_sigmoid(float(score)) for score in model.rerank(query, documents)] + ) + return validate_rerank_scores(scores, len(documents)) + + +def _sigmoid(x: float) -> float: + """Map a cross-encoder logit to a [0, 1] relevance, clamping to avoid overflow.""" + if not math.isfinite(x): + raise RerankProviderContractError(f"Reranker returned a non-finite logit: {x!r}") + # exp(710+) overflows a float; clamp well before that — the tails are ~0/~1 anyway. + x = max(-30.0, min(30.0, x)) + return 1.0 / (1.0 + math.exp(-x)) diff --git a/src/basic_memory/repository/litellm_rerank_provider.py b/src/basic_memory/repository/litellm_rerank_provider.py new file mode 100644 index 000000000..ff4ccc076 --- /dev/null +++ b/src/basic_memory/repository/litellm_rerank_provider.py @@ -0,0 +1,121 @@ +"""LiteLLM-based reranker provider. + +Routes rerank requests to any provider LiteLLM supports (Cohere, Jina, Voyage, +Together, AWS Bedrock, self-hosted OpenAI-compatible endpoints, ...) via +``litellm.arerank``. Model strings use the ``provider/model`` format, e.g. +``cohere/rerank-v3.5`` or ``jina_ai/jina-reranker-v2-base-multilingual``. + +See https://docs.litellm.ai/docs/rerank for supported rerank models. +""" + +from __future__ import annotations + +from typing import Any + +from pydantic import ValidationError + +from basic_memory.repository.litellm_provider import _import_litellm +from basic_memory.repository.rerank_provider import RerankProvider, validate_rerank_scores +from basic_memory.repository.semantic_errors import ( + RerankProviderContractError, + RerankTransientError, +) + + +class LiteLLMRerankProvider(RerankProvider): + """Reranker provider backed by the litellm SDK.""" + + def __init__( + self, + model_name: str = "cohere/rerank-v3.5", + *, + api_key: str | None = None, + api_base: str | None = None, + timeout: float = 30.0, + ) -> None: + self.model_name = model_name + self._api_key = api_key + self._api_base = api_base + self._timeout = timeout + + def runtime_log_attrs(self) -> dict[str, Any]: + return {"model_name": self.model_name, "api_base_set": self._api_base is not None} + + async def rerank(self, query: str, documents: list[str]) -> list[float]: + if not documents: + return [] + + litellm = _import_litellm() + params: dict[str, Any] = { + "model": self.model_name, + "query": query, + "documents": documents, + # Ask for every candidate back so we can rescore the full pool; the + # caller decides the final cut. + "top_n": len(documents), + "timeout": self._timeout, + } + if self._api_key: + params["api_key"] = self._api_key + if self._api_base is not None: + params["api_base"] = self._api_base + + transient_errors = ( + litellm.Timeout, + litellm.APIConnectionError, + litellm.RateLimitError, + litellm.BadGatewayError, + litellm.ServiceUnavailableError, + litellm.InternalServerError, + ) + try: + response = await litellm.arerank(**params) + except transient_errors as exc: + raise RerankTransientError( + f"Rerank provider is temporarily unavailable for model {self.model_name!r}." + ) from exc + except ValidationError as exc: + # LiteLLM constructs its typed RerankResponse inside the awaited call. + # A validation failure therefore describes malformed upstream response + # data, not an invalid search request from our caller. + raise RerankProviderContractError( + f"Rerank provider returned an invalid response for model {self.model_name!r}." + ) from exc + # litellm.arerank returns RerankResponse (Cohere response format): `results` + # is an optional list of TypedDict items with required `index` and + # `relevance_score` keys. A missing/empty list is a contract break, not a + # transient blip. + results = response.results + if not results: + raise RerankProviderContractError( + f"Rerank response contained no results for {len(documents)} documents." + ) + + # Rerank responses are indexed and may arrive out of order, so rebuild an + # input-aligned score vector. We request top_n == len(documents), so a + # complete response must cover every index exactly once; out-of-range + # indices or gaps are truncated/malformed responses we must not silently + # paper over with a 0.0 floor (fail fast). + scores = [0.0] * len(documents) + seen: set[int] = set() + for item in results: + index = int(item["index"]) + if not 0 <= index < len(documents): + raise RerankProviderContractError( + f"Rerank response index {index} is out of range for {len(documents)} documents." + ) + # A repeated index still leaves `seen` == the full set when every other + # index is present, so the coverage check below can't catch it — reject + # here to hold the "each index exactly once" contract (a duplicate would + # otherwise silently overwrite the earlier score, last-write-wins). + if index in seen: + raise RerankProviderContractError( + f"Rerank response repeated index {index} for {len(documents)} documents." + ) + scores[index] = float(item["relevance_score"]) + seen.add(index) + if seen != set(range(len(documents))): + raise RerankProviderContractError( + f"Rerank response covered {len(seen)} of {len(documents)} documents." + ) + return validate_rerank_scores(scores, len(documents)) diff --git a/src/basic_memory/repository/postgres_search_repository.py b/src/basic_memory/repository/postgres_search_repository.py index 2a17e4758..a0d217259 100644 --- a/src/basic_memory/repository/postgres_search_repository.py +++ b/src/basic_memory/repository/postgres_search_repository.py @@ -15,6 +15,8 @@ from basic_memory.config import BasicMemoryConfig, ConfigManager from basic_memory.repository.embedding_provider import EmbeddingProvider from basic_memory.repository.embedding_provider_factory import create_embedding_provider +from basic_memory.repository.rerank_provider import RerankProvider +from basic_memory.repository.rerank_provider_factory import create_rerank_provider from basic_memory.repository.search_index_row import SearchIndexRow from basic_memory.repository.search_query import relaxed_query_words from basic_memory.repository.semantic_chunking import VectorChunkRecord @@ -58,6 +60,7 @@ def __init__( project_id: int, app_config: BasicMemoryConfig | None = None, embedding_provider: EmbeddingProvider | None = None, + rerank_provider: RerankProvider | None = None, ): super().__init__(session_maker, project_id) self._app_config = app_config or ConfigManager().config @@ -71,12 +74,18 @@ def __init__( self._app_config.semantic_postgres_prepare_concurrency ) self._embedding_provider = embedding_provider + self._rerank_provider = rerank_provider + self._reranker_candidates = self._app_config.reranker_candidates + self._reranker_max_document_chars = self._app_config.reranker_max_document_chars self._vector_dimensions = 384 self._vector_tables_initialized = False self._vector_tables_lock = asyncio.Lock() if self._semantic_enabled and self._embedding_provider is None: self._embedding_provider = create_embedding_provider(self._app_config) + # create_rerank_provider returns None unless reranking is enabled. + if self._semantic_enabled and self._rerank_provider is None: + self._rerank_provider = create_rerank_provider(self._app_config) if self._embedding_provider is not None: self._vector_dimensions = self._embedding_provider.dimensions diff --git a/src/basic_memory/repository/rerank_provider.py b/src/basic_memory/repository/rerank_provider.py new file mode 100644 index 000000000..c3b196aec --- /dev/null +++ b/src/basic_memory/repository/rerank_provider.py @@ -0,0 +1,85 @@ +"""Reranker provider protocol for pluggable cross-encoder rescoring. + +A reranker rescores a small candidate set (query + document together via +cross-attention) after retrieval, recovering relevant results that bi-encoder / +FTS ranking left just below the top-k cutoff. Mirrors ``embedding_provider`` so +the same provider families and config shape apply to a different pipeline stage. +""" + +import math +from collections.abc import Sequence +from typing import Any, Protocol + +from basic_memory.repository.semantic_errors import RerankProviderContractError + + +def build_rerank_document(title: str | None, body: str | None, max_chars: int) -> str: + """Assemble the text handed to the cross-encoder for one candidate. + + Lead with the matched body because it carries the retrieval signal, then append + the title so short or title-only candidates still provide context. Truncate to + ``max_chars`` (when positive) so long notes don't inflate cross-encoder latency. + """ + title = title or "" + body = body or "" + text = f"{body}\n{title}" if (title and body) else (body or title) + if max_chars > 0 and len(text) > max_chars: + text = text[:max_chars] + return text + + +def demote_tail_scores(floor: float, count: int) -> list[float]: + """Return bounded tail scores at or below the reranked floor. + + Reranked candidates carry [0, 1] relevance while un-reranked ones still hold + retrieval scores on a different scale; left as is, a tail candidate could + numerically outrank a reranked one. Positive floors produce descending scores + in ``(0, floor)`` based only on the row's tail rank, so fetching more rows does + not rescale earlier results. A zero floor must remain zero to preserve the public + [0, 1] contract, so callers preserve the reranked-pool-before-tail ordering + explicitly instead of relying on a numerically smaller sentinel. + """ + return [floor / (index + 2) for index in range(count)] + + +def validate_rerank_scores(scores: Sequence[float | str], expected_count: int) -> list[float]: + """Return finite scores in ``[0, 1]`` or raise a provider contract error.""" + if len(scores) != expected_count: + raise RerankProviderContractError( + f"Reranker returned {len(scores)} scores for {expected_count} documents." + ) + + validated: list[float] = [] + for index, score in enumerate(scores): + try: + value = float(score) + except (TypeError, ValueError) as exc: + raise RerankProviderContractError( + f"Reranker score at index {index} is not a number: {score!r}." + ) from exc + if not math.isfinite(value) or not 0.0 <= value <= 1.0: + raise RerankProviderContractError( + f"Reranker score at index {index} must be finite and in [0, 1], got {value!r}." + ) + validated.append(value) + return validated + + +class RerankProvider(Protocol): + """Contract for cross-encoder reranking providers.""" + + model_name: str + + async def rerank(self, query: str, documents: list[str]) -> list[float]: + """Return a relevance score in ``[0, 1]`` per document, aligned to input order. + + Higher is more relevant. Providers whose model emits raw logits (e.g. a + local cross-encoder) must squash to ``[0, 1]`` themselves so the search + pipeline can treat every reranker's output on one comparable scale and + keep the public result score bounded. + """ + ... + + def runtime_log_attrs(self) -> dict[str, Any]: + """Return provider-specific runtime settings suitable for startup logs.""" + ... diff --git a/src/basic_memory/repository/rerank_provider_factory.py b/src/basic_memory/repository/rerank_provider_factory.py new file mode 100644 index 000000000..a9eeceb55 --- /dev/null +++ b/src/basic_memory/repository/rerank_provider_factory.py @@ -0,0 +1,109 @@ +"""Factory for creating configured reranker providers. + +Mirrors ``embedding_provider_factory``: string-dispatch on ``reranker_provider`` +plus a process-wide singleton cache so the cross-encoder model loads once. The +fastembed path reuses the embedding provider's resolved cache dir and CPU-aware +thread budget — reranker and embedder models live in the same cache under +distinct model subdirs. +""" + +from threading import Lock + +from loguru import logger + +from basic_memory.config import BasicMemoryConfig +from basic_memory.repository.embedding_provider_factory import ( + _resolve_cache_dir, + _resolve_fastembed_runtime_knobs, + _sensitive_value_digest, +) +from basic_memory.repository.rerank_provider import RerankProvider + +# Key on the fields that change the loaded provider's identity: provider, model, +# (for the litellm path) the endpoint/key routing, and the resolved cache dir. The +# cache dir matters because the fastembed provider is constructed with it — omitting +# it (as an earlier version did) lets two configs with different cache dirs share one +# singleton pointing at the wrong directory, the #741/#872 class of bug the embedding +# factory guards against. CPU-derived thread counts stay out (they drift per call). +type RerankCacheKey = tuple[str, str, str | None, str | None, str] + +_RERANK_PROVIDER_CACHE: dict[RerankCacheKey, RerankProvider] = {} +_RERANK_PROVIDER_CACHE_LOCK = Lock() + + +def _rerank_cache_key(app_config: BasicMemoryConfig) -> RerankCacheKey: + provider_name = app_config.reranker_provider.strip().lower() + api_base_digest = None + api_key_digest = None + if provider_name == "litellm": + api_base_digest = _sensitive_value_digest(app_config.reranker_api_base) + api_key_digest = _sensitive_value_digest(app_config.reranker_api_key) + return ( + provider_name, + app_config.reranker_model, + api_base_digest, + api_key_digest, + _resolve_cache_dir(app_config), + ) + + +def reset_rerank_provider_cache() -> None: + """Clear the process-level reranker provider cache (used by tests).""" + with _RERANK_PROVIDER_CACHE_LOCK: + _RERANK_PROVIDER_CACHE.clear() + + +def create_rerank_provider(app_config: BasicMemoryConfig) -> RerankProvider | None: + """Create a reranker provider, or ``None`` when reranking is disabled. + + Returning ``None`` (rather than a no-op provider) keeps the disabled path + allocation-free and lets the search pipeline skip reranking with a simple + identity check. + """ + # Trigger: reranking is opt-in and off by default. + # Why: a cross-encoder adds latency and a first-run model download; existing + # users must see zero change until they turn it on. + # Outcome: no provider, no import of fastembed/litellm rerank paths. + if not app_config.reranker_enabled: + return None + + cache_key = _rerank_cache_key(app_config) + with _RERANK_PROVIDER_CACHE_LOCK: + if cached_provider := _RERANK_PROVIDER_CACHE.get(cache_key): + return cached_provider + + provider: RerankProvider + provider_name = app_config.reranker_provider.strip().lower() + if provider_name == "fastembed": + # Deferred import: fastembed (and its onnxruntime dep) may not be installed. + from basic_memory.repository.fastembed_rerank_provider import FastEmbedRerankProvider + + resolved_threads, _ = _resolve_fastembed_runtime_knobs(app_config) + provider = FastEmbedRerankProvider( + model_name=app_config.reranker_model, + cache_dir=_resolve_cache_dir(app_config), + threads=resolved_threads, + ) + elif provider_name == "litellm": + from basic_memory.repository.litellm_rerank_provider import LiteLLMRerankProvider + + provider = LiteLLMRerankProvider( + model_name=app_config.reranker_model, + api_key=app_config.reranker_api_key, + api_base=app_config.reranker_api_base, + ) + else: + raise ValueError(f"Unsupported reranker provider: {provider_name}") + + with _RERANK_PROVIDER_CACHE_LOCK: + if cached_provider := _RERANK_PROVIDER_CACHE.get(cache_key): + return cached_provider + if _RERANK_PROVIDER_CACHE: + logger.warning( + "Creating a second distinct reranker provider in this process; " + "the model will be loaded again. existing_keys={existing} new_key={new}", + existing=list(_RERANK_PROVIDER_CACHE.keys()), + new=cache_key, + ) + _RERANK_PROVIDER_CACHE[cache_key] = provider + return provider diff --git a/src/basic_memory/repository/search_repository.py b/src/basic_memory/repository/search_repository.py index ed6c05729..999402488 100644 --- a/src/basic_memory/repository/search_repository.py +++ b/src/basic_memory/repository/search_repository.py @@ -14,6 +14,7 @@ from basic_memory.config import BasicMemoryConfig, DatabaseBackend from basic_memory.repository.embedding_provider_factory import create_embedding_provider +from basic_memory.repository.rerank_provider_factory import create_rerank_provider from basic_memory.repository.postgres_search_repository import PostgresSearchRepository from basic_memory.repository.search_index_row import SearchIndexRow from basic_memory.repository.search_repository_base import VectorSyncBatchResult @@ -138,8 +139,12 @@ def create_search_repository( # Outcome: resolve the cached singleton here once and inject it, so the provider # is the single source of truth across all callers of this factory. embedding_provider = None + rerank_provider = None if config.semantic_search_enabled: embedding_provider = create_embedding_provider(config) + # Returns None unless reranking is enabled; resolve the cached singleton + # here so both backends share one process-wide reranker model. + rerank_provider = create_rerank_provider(config) if database_backend == DatabaseBackend.POSTGRES: # pragma: no cover return PostgresSearchRepository( # pragma: no cover @@ -147,6 +152,7 @@ def create_search_repository( project_id=project_id, app_config=app_config, embedding_provider=embedding_provider, + rerank_provider=rerank_provider, ) else: return SQLiteSearchRepository( @@ -154,6 +160,7 @@ def create_search_repository( project_id=project_id, app_config=app_config, embedding_provider=embedding_provider, + rerank_provider=rerank_provider, ) diff --git a/src/basic_memory/repository/search_repository_base.py b/src/basic_memory/repository/search_repository_base.py index 5719e7774..094b301a5 100644 --- a/src/basic_memory/repository/search_repository_base.py +++ b/src/basic_memory/repository/search_repository_base.py @@ -19,6 +19,12 @@ EmbeddingProvider, embedding_provider_identity, ) +from basic_memory.repository.rerank_provider import ( + RerankProvider, + build_rerank_document, + demote_tail_scores, + validate_rerank_scores, +) from basic_memory.repository.search_index_row import SearchIndexRow from basic_memory.repository.semantic_chunking import ( SemanticSourceRow, @@ -45,6 +51,10 @@ # --- Semantic search constants --- VECTOR_FILTER_SCAN_LIMIT = 50000 +# Over-fetch factor for the rerank candidate chunk pool: chunks collapse to unique +# (type, id) rows before reranking, so fetch several times reranker_candidates chunks +# to keep enough unique documents in the rerank window. +RERANK_POOL_CHUNK_FANOUT = 4 FUSION_BONUS = 0.3 FTS_GATE_THRESHOLD = 0.0 TOP_CHUNKS_PER_RESULT = 5 @@ -77,6 +87,11 @@ class SearchRepositoryBase(ABC): _semantic_vector_k: int _semantic_min_similarity: float _embedding_provider: Optional[EmbeddingProvider] + # Class-level defaults: a repo with no reranker configured (or a lightweight + # test double that bypasses __init__) safely skips reranking. + _rerank_provider: Optional[RerankProvider] = None + _reranker_candidates: int = 20 + _reranker_max_document_chars: int = 0 _semantic_embedding_sync_batch_size: int _vector_dimensions: int _vector_tables_initialized: bool @@ -833,6 +848,130 @@ def _parse_chunk_key(chunk_key: str) -> SearchIndexKey: parts = chunk_key.split(":") return parts[0], int(parts[1]) + # ------------------------------------------------------------------ + # Shared semantic search: cross-encoder reranking + # ------------------------------------------------------------------ + + def _should_rerank(self, query_text: str) -> bool: + """Return whether a configured reranker should run for this query.""" + return self._rerank_provider is not None and bool(query_text) + + def _rerank_candidate_limit(self) -> int: + """Return the fixed chunk window that owns reranker-prefix membership.""" + return max( + self._semantic_vector_k, + self._reranker_candidates * RERANK_POOL_CHUNK_FANOUT, + ) + + def _candidate_limit(self, limit: int, offset: int, query_text: str) -> int: + """Size the retrieval candidate *chunk* pool for vector/hybrid search. + + ``candidate_limit`` bounds vector chunks, but many chunks of one large note + collapse to a single ``(type, id)`` row before reranking, so a chunk count does + not equal a unique-document count. When reranking is active we over-fetch by + ``RERANK_POOL_CHUNK_FANOUT`` so a few multi-chunk notes can't starve the rerank + window below ``reranker_candidates`` unique rows. This is best-effort headroom, + not a hard guarantee — a single note dominating the entire nearest-neighbour set + can still yield fewer unique rows (a pathological corpus shape). + """ + if self._should_rerank(query_text): + # Trigger: the requested window extends beyond the fixed reranked prefix. + # Why: a bounded prefix alone can under-fill large pages and hide the + # semantic pagination probe even when more matches exist. + # Outcome: keep prefix membership fixed while adding chunk headroom only + # for the untouched tail that this request must return. + rerank_candidate_limit = self._rerank_candidate_limit() + tail_size = max(0, limit + offset - self._reranker_candidates) + return rerank_candidate_limit + tail_size * 10 + return max(self._semantic_vector_k, (limit + offset) * 10) + + def _rerank_document_text(self, row: SearchIndexRow) -> str: + """Build the document text handed to the cross-encoder for one candidate. + + Prefer the matched chunk (the most relevant passage of a large note), + falling back to the stored snippet. + """ + body = row.matched_chunk_text or row.content_snippet or "" + return build_rerank_document(row.title, body, self._reranker_max_document_chars) + + @staticmethod + def _demote_tail(tail: list[SearchIndexRow], floor: float) -> list[SearchIndexRow]: + """Rescore un-reranked tail rows at or below the floor, preserving their order. + + The reranked pool carries [0, 1] relevance scores while the tail still holds + raw retrieval scores on a different scale ([0, 1.3] for fused hybrid). Left as + is, a tail row could outrank a reranked row numerically. Positive floors put + the tail strictly below the pool; a zero floor yields zeroes because no smaller + score exists in the public [0, 1] range. The returned pool-plus-tail sequence, + rather than a later score-only sort, owns that tie-breaking invariant. + """ + return [ + replace(row, score=score) + for row, score in zip(tail, demote_tail_scores(floor, len(tail))) + ] + + async def _rerank_and_paginate( + self, + query_text: str, + rows: list[SearchIndexRow], + *, + offset: int, + limit: int, + stable_rows: list[SearchIndexRow] | None = None, + ) -> list[SearchIndexRow]: + """Rerank the top candidates, then return the requested ``[offset:offset+limit]`` page. + + Trigger: a reranker is configured and there is a real query. + Why: bi-encoder/FTS ranking lands the gold document in the top-N but often + just below the top-k cutoff (#950); a cross-encoder that reads query and + document together recovers those near-misses. + Outcome: the first ``reranker_candidates`` rows are reordered by reranker + relevance (which replaces ``score``); the requested page is sliced from the + reordered list. + + Every non-empty page rescores the same fixed prefix before slicing so the + untouched tail can be demoted onto the reranker's public ``[0, 1]`` scale. + """ + page_end = offset + limit + if self._rerank_provider is None or not query_text: + return rows[offset:page_end] + + # Trigger: pagination needs more rows than the fixed rerank retrieval window. + # Why: an expanded retrieval may introduce or strengthen raw candidates, but + # letting them replace the original prefix causes duplicates and skips. + # Outcome: the fixed window owns prefix membership; the expanded result only + # supplies new, de-duplicated tail rows. + pool_source = stable_rows if stable_rows is not None else rows + pool = pool_source[: self._reranker_candidates] + pool_keys = {(row.type, row.id) for row in pool} + tail = [row for row in rows if (row.type, row.id) not in pool_keys] + ordered_rows = pool + tail + + # Skip only when there is no prefix to calibrate or the requested page is + # empty. Even a singleton prefix or a wholly-tail page needs the prefix's + # relevance floor so raw hybrid scores cannot leak into cross-project sorting. + if not pool or offset >= len(ordered_rows): + return ordered_rows[offset:page_end] + + documents = [self._rerank_document_text(row) for row in pool] + # A transient provider failure must surface instead of switching this page + # back to retrieval order. A prior page may already have returned reranked + # order, so degrading here can duplicate one result and omit another. + scores = validate_rerank_scores( + await self._rerank_provider.rerank(query_text, documents), + len(pool), + ) + + order = sorted(range(len(pool)), key=lambda i: scores[i], reverse=True) + reranked = [replace(pool[i], score=scores[i]) for i in order] + logger.debug( + "Reranked candidates: pool={pool} model={model}", + pool=len(pool), + model=self._rerank_provider.model_name, + ) + demoted_tail = self._demote_tail(tail, floor=reranked[-1].score or 0.0) + return (reranked + demoted_tail)[offset:page_end] + async def _search_vector_only( self, *, @@ -848,19 +987,25 @@ async def _search_vector_only( min_similarity: Optional[float] = None, limit: int, offset: int, + candidate_limit: int | None = None, _emit_observability_log: bool = True, + _apply_rerank: bool = True, ) -> List[SearchIndexRow]: """Run vector-only search returning chunk-level results. Returns individual search_index rows (entities, observations, relations) ranked by vector similarity. Each observation or relation is a first-class result, not collapsed into its parent entity. + + ``candidate_limit`` is supplied only by a composed retrieval stage that + already sized the shared candidate pool. """ self._assert_semantic_available() await self._ensure_vector_tables() assert self._embedding_provider is not None query_text = search_text.strip() - candidate_limit = max(self._semantic_vector_k, (limit + offset) * 10) + if candidate_limit is None: + candidate_limit = self._candidate_limit(limit, offset, query_text) query_start = time.perf_counter() embed_start = time.perf_counter() query_embedding = await self._embedding_provider.embed_query(query_text) @@ -1008,8 +1153,45 @@ def _log_vector_summary() -> None: ranked_rows.sort(key=lambda item: item.score or 0.0, reverse=True) hydrate_ms = (time.perf_counter() - hydrate_start) * 1000 + # Rerank over the wide candidate pool, then slice to the page. Suppressed when + # hybrid calls this internally (_apply_rerank=False) — hybrid reranks its own + # fused result; _rerank_and_paginate no-ops back to a plain slice otherwise. + if _apply_rerank: + stable_rows = ranked_rows + if self._should_rerank(query_text): + stable_candidate_limit = self._rerank_candidate_limit() + if candidate_limit > stable_candidate_limit: + stable_rows = await self._search_vector_only( + search_text=search_text, + permalink=permalink, + permalink_match=permalink_match, + title=title, + note_types=note_types, + after_date=after_date, + search_item_types=search_item_types, + categories=categories, + metadata_filters=metadata_filters, + min_similarity=min_similarity, + limit=stable_candidate_limit, + offset=0, + candidate_limit=stable_candidate_limit, + _emit_observability_log=False, + _apply_rerank=False, + ) + output = await self._rerank_and_paginate( + query_text, + ranked_rows, + offset=offset, + limit=limit, + stable_rows=stable_rows, + ) + else: + output = ranked_rows[offset : offset + limit] + # Vector latency owns the optional rerank stage too. Logging before the + # awaited provider call hides the feature's dominant cost and can suppress + # the slow-query warning entirely. _log_vector_summary() - return ranked_rows[offset : offset + limit] + return output async def _fetch_entity_rows_by_ids(self, entity_ids: list[int]) -> dict[int, SearchIndexRow]: """Fetch entity-type search_index rows by their entity_id values.""" @@ -1088,6 +1270,9 @@ async def _search_hybrid( min_similarity: Optional[float] = None, limit: int, offset: int, + _candidate_limit_override: int | None = None, + _apply_rerank: bool = True, + _emit_observability_log: bool = True, ) -> List[SearchIndexRow]: """Fuse FTS and vector results using score-based fusion. @@ -1097,8 +1282,14 @@ async def _search_hybrid( """ self._assert_semantic_available() query_text = search_text.strip() + rerank_configured = self._should_rerank(query_text) + rerank_enabled = _apply_rerank and rerank_configured query_start = time.perf_counter() - candidate_limit = max(self._semantic_vector_k, (limit + offset) * 10) + candidate_limit = ( + _candidate_limit_override + if _candidate_limit_override is not None + else self._candidate_limit(limit, offset, query_text) + ) fts_start = time.perf_counter() # allow_relaxed: question-form queries rarely AND-match, and a dead FTS # branch silently degrades hybrid to vector-only ranking. Fusion plus @@ -1133,7 +1324,13 @@ async def _search_hybrid( min_similarity=min_similarity, limit=candidate_limit, offset=0, + # Trigger: reranking owns a bounded candidate window shared by both legs. + # Why: the disabled path historically expands the vector leg again to + # preserve recall when many vector chunks collapse into a few search rows. + # Outcome: avoid double expansion only when reranking is actually active. + candidate_limit=candidate_limit if rerank_configured else None, _emit_observability_log=False, + _apply_rerank=False, ) vector_ms = (time.perf_counter() - vector_start) * 1000 fusion_start = time.perf_counter() @@ -1151,25 +1348,31 @@ async def _search_hybrid( fts_max = max(fts_abs) if fts_abs else 1.0 fts_scores: dict[SearchIndexKey, float] = {} - for row in fts_results: + fts_ranks: dict[SearchIndexKey, int] = {} + for rank, row in enumerate(fts_results): if row.id is None: continue + row_key = (row.type, row.id) norm = abs(row.score or 0.0) / fts_max if fts_max > 0 else 0.0 # Gate: FTS scores below threshold contribute zero if norm < FTS_GATE_THRESHOLD: norm = 0.0 - fts_scores[(row.type, row.id)] = norm - rows_by_key[(row.type, row.id)] = row + fts_scores[row_key] = norm + fts_ranks.setdefault(row_key, rank) + rows_by_key[row_key] = row vec_scores: dict[SearchIndexKey, float] = {} - for row in vector_results: + vec_ranks: dict[SearchIndexKey, int] = {} + for rank, row in enumerate(vector_results): if row.id is None: continue + row_key = (row.type, row.id) # Trigger: no re-normalization by vec_max # Why: vector similarity is already calibrated [0, 1]; re-normalizing # inflates weak matches when the entire result set is mediocre - vec_scores[(row.type, row.id)] = row.score or 0.0 - rows_by_key[(row.type, row.id)] = row + vec_scores[row_key] = row.score or 0.0 + vec_ranks.setdefault(row_key, rank) + rows_by_key[row_key] = row # Fuse: max(v, f) + FUSION_BONUS * min(v, f) # Preserves the dominant signal; bonus rewards dual-source agreement. @@ -1181,18 +1384,74 @@ async def _search_hybrid( fused_scores[row_key] = max(v, f) + FUSION_BONUS * min(v, f) ranked = sorted(fused_scores.items(), key=lambda item: item[1], reverse=True) - output: list[SearchIndexRow] = [] - for row_key, fused_score in ranked[offset : offset + limit]: + + def _materialize(entry: tuple[SearchIndexKey, float]) -> SearchIndexRow: + row_key, fused_score = entry row = rows_by_key[row_key] # Trigger: FTS-only results have no matched_chunk_text from vector search. # Why: without chunk text, API falls back to truncated content, losing answer text. # Outcome: FTS-only results get full content_snippet as matched_chunk. if row.matched_chunk_text is None and row.content_snippet: row = replace(row, matched_chunk_text=row.content_snippet) - output.append(replace(row, score=fused_score)) + return replace(row, score=fused_score) + + # Rerank the top fused candidates before paginating. When reranking is active + # we materialize the whole candidate list (cheap next to a cross-encoder call) + # and hand it to the shared paginate helper; the disabled path stays cheap by + # materializing only the requested page. + if rerank_enabled: + candidates = [_materialize(entry) for entry in ranked] + stable_candidates = candidates + stable_candidate_limit = self._rerank_candidate_limit() + if candidate_limit > stable_candidate_limit: + stable_candidates = await self._search_hybrid( + search_text=search_text, + permalink=permalink, + permalink_match=permalink_match, + title=title, + note_types=note_types, + after_date=after_date, + search_item_types=search_item_types, + categories=categories, + metadata_filters=metadata_filters, + min_similarity=min_similarity, + limit=stable_candidate_limit, + offset=0, + _candidate_limit_override=stable_candidate_limit, + _apply_rerank=False, + _emit_observability_log=False, + ) + stable_keys = {(row.type, row.id) for row in stable_candidates} + expanded_tail = [entry for entry in ranked if entry[0] not in stable_keys] + + # Trigger: deeper pages expand the FTS/vector retrieval windows. + # Why: score fusion can strengthen an existing row when its second + # signal appears later, moving it across a page already returned. + # Outcome: freeze the fixed fused universe, then order newly admitted + # rows by their earliest source rank. That rank cannot improve after a + # row first appears, so each larger window only appends to the tail. + expanded_tail.sort( + key=lambda entry: ( + min( + fts_ranks.get(entry[0], candidate_limit), + vec_ranks.get(entry[0], candidate_limit), + ), + entry[0], + ) + ) + candidates = stable_candidates + [_materialize(entry) for entry in expanded_tail] + output = await self._rerank_and_paginate( + query_text, + candidates, + offset=offset, + limit=limit, + stable_rows=stable_candidates, + ) + else: + output = [_materialize(entry) for entry in ranked[offset : offset + limit]] fusion_ms = (time.perf_counter() - fusion_start) * 1000 total_ms = (time.perf_counter() - query_start) * 1000 - if total_ms > 2500: + if _emit_observability_log and total_ms > 2500: logger.warning( "[SEMANTIC_SLOW_QUERY] Semantic query timing: project_id={project_id} " "retrieval_mode={retrieval_mode} query_length={query_length} " diff --git a/src/basic_memory/repository/semantic_errors.py b/src/basic_memory/repository/semantic_errors.py index edc6fe9da..545d5c16d 100644 --- a/src/basic_memory/repository/semantic_errors.py +++ b/src/basic_memory/repository/semantic_errors.py @@ -7,3 +7,16 @@ class SemanticSearchDisabledError(RuntimeError): class SemanticDependenciesMissingError(RuntimeError): """Raised when a semantic search dependency is unavailable or misconfigured.""" + + +class RerankProviderContractError(RuntimeError): + """Raised when a reranker provider violates its response contract. + + A distinct type so the search pipeline can surface this permanent fault (a + provider/config bug) instead of degrading it to un-reranked results the way it + handles transient reranker failures. + """ + + +class RerankTransientError(RuntimeError): + """Raised when a reranker is temporarily unavailable and the request should be retried.""" diff --git a/src/basic_memory/repository/sqlite_search_repository.py b/src/basic_memory/repository/sqlite_search_repository.py index e918e6af6..35d75df0f 100644 --- a/src/basic_memory/repository/sqlite_search_repository.py +++ b/src/basic_memory/repository/sqlite_search_repository.py @@ -24,6 +24,8 @@ ) from basic_memory.repository.embedding_provider import EmbeddingProvider from basic_memory.repository.embedding_provider_factory import create_embedding_provider +from basic_memory.repository.rerank_provider import RerankProvider +from basic_memory.repository.rerank_provider_factory import create_rerank_provider from basic_memory.repository.search_index_row import SearchIndexRow from basic_memory.repository.search_query import relaxed_query_words from basic_memory.repository.search_repository_base import SearchRepositoryBase @@ -48,6 +50,7 @@ def __init__( project_id: int, app_config: BasicMemoryConfig | None = None, embedding_provider: EmbeddingProvider | None = None, + rerank_provider: RerankProvider | None = None, ): super().__init__(session_maker, project_id) self._entity_columns: set[str] | None = None @@ -59,6 +62,9 @@ def __init__( self._app_config.semantic_embedding_sync_batch_size ) self._embedding_provider = embedding_provider + self._rerank_provider = rerank_provider + self._reranker_candidates = self._app_config.reranker_candidates + self._reranker_max_document_chars = self._app_config.reranker_max_document_chars self._sqlite_vec_load_lock = asyncio.Lock() self._sqlite_prepare_write_lock = asyncio.Lock() self._vector_tables_initialized = False @@ -69,6 +75,9 @@ def __init__( # This conversion is correct only for unit-normalized embeddings. # Provider implementations must return normalized vectors. self._embedding_provider = create_embedding_provider(self._app_config) + # create_rerank_provider returns None unless reranking is enabled. + if self._semantic_enabled and self._rerank_provider is None: + self._rerank_provider = create_rerank_provider(self._app_config) if self._embedding_provider is not None: self._vector_dimensions = self._embedding_provider.dimensions diff --git a/test-int/semantic/test_semantic_coverage.py b/test-int/semantic/test_semantic_coverage.py index 5e499a7bd..0e4881089 100644 --- a/test-int/semantic/test_semantic_coverage.py +++ b/test-int/semantic/test_semantic_coverage.py @@ -38,6 +38,23 @@ PG_FASTEMBED = SearchCombo("postgres-fastembed", DatabaseBackend.POSTGRES, "fastembed", 384) +class _RecordingReranker: + """Deterministic reranker that records whether hybrid applies the stage once.""" + + model_name = "recording-reranker" + + def __init__(self) -> None: + self.calls = 0 + + async def rerank(self, query: str, documents: list[str]) -> list[float]: + self.calls += 1 + document_count = len(documents) + return [(document_count - index) / document_count for index in range(document_count)] + + def runtime_log_attrs(self) -> dict[str, object]: + return {} + + @pytest.mark.asyncio @pytest.mark.semantic @pytest.mark.benchmark @@ -116,6 +133,150 @@ async def test_postgres_hybrid_search(postgres_engine_factory, tmp_path): ) +@pytest.mark.asyncio +@pytest.mark.semantic +@pytest.mark.benchmark +async def test_postgres_hybrid_preserves_candidate_windows( + postgres_engine_factory, + tmp_path, + monkeypatch, +): + """Exercise baseline and reranked hybrid candidate sizing through real Postgres queries.""" + skip_if_needed(PG_FASTEMBED) + if postgres_engine_factory is None: + pytest.skip("Postgres engine not available") + + provider = _create_fastembed_provider() + search_service = await create_search_service( + postgres_engine_factory, PG_FASTEMBED, tmp_path, embedding_provider=provider + ) + await seed_benchmark_notes(search_service, note_count=120) + + repo = cast(Any, search_service.repository) + repo._rerank_provider = None + repo._semantic_vector_k = 100 + repo._reranker_candidates = 100 + + candidate_limits: list[int] = [] + run_vector_query = repo._run_vector_query + + async def record_vector_query( + session: Any, + query_embedding: list[float], + candidate_limit: int, + ) -> list[dict[str, Any]]: + candidate_limits.append(candidate_limit) + return await run_vector_query(session, query_embedding, candidate_limit) + + monkeypatch.setattr(repo, "_run_vector_query", record_vector_query) + + baseline_results = await search_service.search( + SearchQuery( + text="database migration schema", + retrieval_mode=SearchRetrievalMode.HYBRID, + entity_types=[SearchItemType.ENTITY], + ), + limit=10, + ) + + assert baseline_results + assert candidate_limits == [1000] + + candidate_limits.clear() + repo._semantic_vector_k = 5 + repo._reranker_candidates = 20 + reranker = _RecordingReranker() + repo._rerank_provider = reranker + reranked_results = await search_service.search( + SearchQuery( + text="database migration schema", + retrieval_mode=SearchRetrievalMode.HYBRID, + entity_types=[SearchItemType.ENTITY], + ), + limit=11, + ) + + assert reranked_results + assert candidate_limits == [80] + + candidate_limits.clear() + growing_prefix_results = await search_service.search( + SearchQuery( + text="database migration schema", + retrieval_mode=SearchRetrievalMode.HYBRID, + entity_types=[SearchItemType.ENTITY], + ), + limit=21, + ) + + assert growing_prefix_results + assert reranker.calls == 2 + assert candidate_limits == [90, 80] + + candidate_limits.clear() + large_page_with_probe = await search_service.search( + SearchQuery( + text="database migration schema", + retrieval_mode=SearchRetrievalMode.HYBRID, + entity_types=[SearchItemType.ENTITY], + min_similarity=0.0, + ), + limit=101, + ) + + assert len(large_page_with_probe) == 101 + assert reranker.calls == 3 + assert candidate_limits == [890, 80] + + # A larger retrieval window must extend the same sequence rather than + # reordering rows already exposed by an earlier deep page. + reranker.calls = 0 + candidate_limits.clear() + stable_window = await search_service.search( + SearchQuery( + text="database migration schema", + retrieval_mode=SearchRetrievalMode.HYBRID, + entity_types=[SearchItemType.ENTITY], + min_similarity=0.0, + ), + limit=40, + ) + + assert len(stable_window) == 40 + assert candidate_limits == [280, 80] + + candidate_limits.clear() + third_page = await search_service.search( + SearchQuery( + text="database migration schema", + retrieval_mode=SearchRetrievalMode.HYBRID, + entity_types=[SearchItemType.ENTITY], + min_similarity=0.0, + ), + limit=10, + offset=20, + ) + assert candidate_limits == [180, 80] + + candidate_limits.clear() + fourth_page = await search_service.search( + SearchQuery( + text="database migration schema", + retrieval_mode=SearchRetrievalMode.HYBRID, + entity_types=[SearchItemType.ENTITY], + min_similarity=0.0, + ), + limit=10, + offset=30, + ) + assert candidate_limits == [280, 80] + + assert [row.permalink for row in third_page] == [row.permalink for row in stable_window[20:30]] + assert [row.permalink for row in fourth_page] == [row.permalink for row in stable_window[30:40]] + assert reranker.calls == 3 + assert all(row.score is not None and row.score < 0.05 for row in third_page + fourth_page) + + @pytest.mark.asyncio @pytest.mark.semantic @pytest.mark.benchmark diff --git a/tests/api/v2/test_search_router.py b/tests/api/v2/test_search_router.py index 5459481ad..8d885296c 100644 --- a/tests/api/v2/test_search_router.py +++ b/tests/api/v2/test_search_router.py @@ -11,6 +11,8 @@ from basic_memory.models import Project from basic_memory.repository.search_index_row import SearchIndexRow from basic_memory.repository.semantic_errors import ( + RerankProviderContractError, + RerankTransientError, SemanticDependenciesMissingError, SemanticSearchDisabledError, ) @@ -436,6 +438,58 @@ async def count(self, *args, **kwargs): assert "Semantic dependencies are missing" in response.json()["detail"] +@pytest.mark.asyncio +async def test_search_router_returns_503_for_transient_reranker_failure( + client: AsyncClient, app, v2_project_url +): + """A transient reranker outage should be retryable, not silently reorder results.""" + + class RaisingSearchService: + async def search(self, *args, **kwargs): + raise RerankTransientError("Reranker is temporarily unavailable.") + + async def count(self, *args, **kwargs): + raise RerankTransientError("Reranker is temporarily unavailable.") + + app.dependency_overrides[get_search_service_v2_external] = lambda: RaisingSearchService() + try: + response = await client.post( + f"{v2_project_url}/search/", + json={"text": "semantic query", "retrieval_mode": "hybrid"}, + ) + finally: + app.dependency_overrides.pop(get_search_service_v2_external, None) + + assert response.status_code == 503 + assert "temporarily unavailable" in response.json()["detail"] + + +@pytest.mark.asyncio +async def test_search_router_returns_502_for_reranker_contract_failure( + client: AsyncClient, app, v2_project_url +): + """A malformed upstream reranker response should map to a provider 502.""" + + class RaisingSearchService: + async def search(self, *args, **kwargs): + raise RerankProviderContractError("Reranker returned malformed scores.") + + async def count(self, *args, **kwargs): + raise RerankProviderContractError("Reranker returned malformed scores.") + + app.dependency_overrides[get_search_service_v2_external] = lambda: RaisingSearchService() + try: + response = await client.post( + f"{v2_project_url}/search/", + json={"text": "semantic query", "retrieval_mode": "hybrid"}, + ) + finally: + app.dependency_overrides.pop(get_search_service_v2_external, None) + + assert response.status_code == 502 + assert response.json()["detail"] == "Reranker returned malformed scores." + + @pytest.mark.asyncio async def test_search_router_returns_400_for_invalid_vector_query( client: AsyncClient, app, v2_project_url diff --git a/tests/cli/test_config_command.py b/tests/cli/test_config_command.py index bba656907..79e291ad5 100644 --- a/tests/cli/test_config_command.py +++ b/tests/cli/test_config_command.py @@ -239,6 +239,16 @@ def test_config_get_never_prints_cloud_api_key(runner, write_config): assert "cloud_api_key = ********" in result.output +def test_config_get_never_prints_reranker_api_key(runner, write_config): + write_config(_base_config(reranker_api_key="reranker-super-secret")) + + result = runner.invoke(app, ["config", "get", "reranker_api_key"]) + + assert result.exit_code == 0, result.output + assert "reranker-super-secret" not in result.output + assert "reranker_api_key = ********" in result.output + + def test_config_list_never_prints_cloud_api_key(runner, write_config): write_config(_base_config(cloud_api_key="bmc_super_secret_token")) @@ -250,6 +260,48 @@ def test_config_list_never_prints_cloud_api_key(runner, write_config): assert rows["cloud_api_key"]["value"] == "********" +def test_config_list_redacts_reranker_credentials(runner, write_config): + write_config( + _base_config( + reranker_api_key="reranker-super-secret", + reranker_api_base=( + "https://provider-user:provider-password@reranker.example.com/v1" + "?api_key=query-secret&timeout=30" + ), + ) + ) + + result = runner.invoke(app, ["config", "list", "--json"]) + + assert result.exit_code == 0, result.output + for secret in ( + "reranker-super-secret", + "provider-user", + "provider-password", + "query-secret", + ): + assert secret not in result.output + rows = {row["key"]: row for row in json.loads(result.output)} + assert rows["reranker_api_key"]["value"] == "********" + assert rows["reranker_api_base"]["value"] == ( + "https://***@reranker.example.com/v1?api_key=%2A%2A%2A&timeout=30" + ) + + +def test_config_set_never_prints_reranker_api_key(runner, write_config): + config_file = write_config(_base_config()) + + result = runner.invoke( + app, + ["config", "set", "reranker_api_key", "reranker-super-secret"], + ) + + assert result.exit_code == 0, result.output + assert "reranker-super-secret" not in result.output + assert "reranker_api_key = ********" in result.output + assert json.loads(config_file.read_text())["reranker_api_key"] == "reranker-super-secret" + + def test_config_get_masks_database_url_credentials(runner, write_config): write_config(_base_config(database_url="postgresql://dbuser:dbpass@host.example.com:5432/bm")) diff --git a/tests/mcp/test_tool_basic_memory_diagnostics.py b/tests/mcp/test_tool_basic_memory_diagnostics.py index 99b0710d6..43a2c29b2 100644 --- a/tests/mcp/test_tool_basic_memory_diagnostics.py +++ b/tests/mcp/test_tool_basic_memory_diagnostics.py @@ -45,6 +45,18 @@ def test_redact_config_removes_semantic_embedding_api_key(): assert result["semantic_embedding_api_base"] == "https://embeddings.example.com/v1" +def test_redact_config_removes_reranker_api_key(): + raw = { + "reranker_api_key": "reranker-secret", + "reranker_api_base": "https://reranker.example.com/v1", + } + + result = _redact_config(raw) + + assert "reranker_api_key" not in result + assert result["reranker_api_base"] == "https://reranker.example.com/v1" + + def test_redact_config_scrubs_semantic_embedding_api_base_credentials(): raw = { "semantic_embedding_api_base": ( @@ -60,6 +72,21 @@ def test_redact_config_scrubs_semantic_embedding_api_base_credentials(): ) +def test_redact_config_scrubs_reranker_api_base_credentials(): + raw = { + "reranker_api_base": ( + "https://provider-user:provider-password@reranker.example.com/v1" + "?api_key=query-secret&timeout=30" + ), + } + + result = _redact_config(raw) + + assert result["reranker_api_base"] == ( + "https://***@reranker.example.com/v1?api_key=%2A%2A%2A&timeout=30" + ) + + def test_redact_config_passes_through_safe_fields(): raw = {"default_project": "main", "log_level": "INFO", "env": "dev"} result = _redact_config(raw) @@ -165,6 +192,33 @@ def test_diagnostics_redacts_semantic_embedding_api_key(tmp_path): assert "timeout=30" in result +def test_diagnostics_redacts_reranker_credentials(tmp_path): + """Reranker credentials must never appear in diagnostic output.""" + config_data = { + "reranker_api_key": "reranker-super-secret", + "reranker_api_base": ( + "https://provider-user:provider-password@reranker.example.com/v1" + "?api_key=query-secret&timeout=30" + ), + "projects": {}, + } + config_file = tmp_path / "config.json" + config_file.write_text(json.dumps(config_data)) + + result = basic_memory_diagnostics() + + for secret in ( + "reranker-super-secret", + "provider-user", + "provider-password", + "query-secret", + ): + assert secret not in result + assert "reranker_api_key" not in result + assert "reranker.example.com" in result + assert "timeout=30" in result + + def test_diagnostics_config_missing(tmp_path): """When config file does not exist, output should say so.""" config_file = tmp_path / "config.json" diff --git a/tests/mcp/tools/test_search_notes_multi_project.py b/tests/mcp/tools/test_search_notes_multi_project.py index bfffebd9a..b088ce238 100644 --- a/tests/mcp/tools/test_search_notes_multi_project.py +++ b/tests/mcp/tools/test_search_notes_multi_project.py @@ -3,6 +3,8 @@ from contextlib import asynccontextmanager import importlib +from httpx import HTTPStatusError, Request, Response +from mcp.server.fastmcp.exceptions import ToolError import pytest from basic_memory.schemas.search import SearchItemType, SearchResponse, SearchResult @@ -289,6 +291,86 @@ async def search(self, payload, page, page_size): assert any("team index unavailable" in warning for warning in warnings) +@pytest.mark.asyncio +async def test_search_notes_search_all_projects_propagates_retryable_service_outage( + monkeypatch, cloud_routing +): + """A retryable project outage must fail the merged page instead of returning a partial one.""" + clients_mod = importlib.import_module("basic_memory.mcp.clients") + search_mod = importlib.import_module("basic_memory.mcp.tools.search") + project_refs = [ + { + "project": "personal/main", + "project_id": "11111111-1111-1111-1111-111111111111", + }, + { + "project": "team-paul/main", + "project_id": "22222222-2222-2222-2222-222222222222", + }, + ] + + async def fake_load_search_project_refs(context=None): + return project_refs + + class StubProject: + def __init__(self, name: str | None, external_id: str | None): + self.name = name or "main" + self.external_id = external_id or "local-main" + + @asynccontextmanager + async def fake_get_project_client(project=None, context=None, project_id=None): + yield object(), StubProject(project, project_id) + + async def fake_resolve_project_and_path(client, identifier, project=None, context=None): + return StubProject(project, None), identifier, False + + class MockSearchClient: + def __init__(self, client, project_id): + self.project_id = project_id + + async def search(self, payload, page, page_size): + if self.project_id == "22222222-2222-2222-2222-222222222222": + request = Request("POST", "https://api.example/search") + response = Response( + 503, + request=request, + json={"detail": "Reranker temporarily unavailable"}, + ) + try: + response.raise_for_status() + except HTTPStatusError as exc: + raise ToolError("Reranker temporarily unavailable") from exc + return SearchResponse( + results=[ + SearchResult( + title="Personal result", + permalink="main/personal-result", + content="MCP content", + type=SearchItemType.ENTITY, + score=0.5, + file_path="/main/personal-result.md", + ) + ], + current_page=page, + page_size=page_size, + total=1, + ) + + monkeypatch.setattr(search_mod, "_load_search_project_refs", fake_load_search_project_refs) + monkeypatch.setattr(search_mod, "get_project_client", fake_get_project_client) + monkeypatch.setattr(search_mod, "resolve_project_and_path", fake_resolve_project_and_path) + monkeypatch.setattr(clients_mod, "SearchClient", MockSearchClient) + + result = await search_mod.search_notes( + query="MCP Test Note", + search_all_projects=True, + output_format="json", + ) + + assert isinstance(result, str) + assert result.startswith("# Search Failed - Service Temporarily Unavailable") + + @pytest.mark.asyncio async def test_search_notes_search_all_projects_local_omits_project_id(monkeypatch, local_routing): """Without a cloud route, fan-out must address each project by name only. diff --git a/tests/repository/test_fastembed_rerank_provider.py b/tests/repository/test_fastembed_rerank_provider.py new file mode 100644 index 000000000..5c9e1ccf1 --- /dev/null +++ b/tests/repository/test_fastembed_rerank_provider.py @@ -0,0 +1,314 @@ +"""Tests for FastEmbedRerankProvider.""" + +import asyncio +import builtins +import sys +import threading + +import pytest +from requests import Response, exceptions as requests_exceptions + +from basic_memory.repository.fastembed_rerank_provider import FastEmbedRerankProvider +from basic_memory.repository.semantic_errors import ( + RerankProviderContractError, + RerankTransientError, + SemanticDependenciesMissingError, +) + + +class _StubCrossEncoder: + init_count = 0 + last_init_kwargs: dict = {} + + def __init__(self, model_name: str, cache_dir: str | None = None, threads: int | None = None): + _StubCrossEncoder.last_init_kwargs = { + "model_name": model_name, + "cache_dir": cache_dir, + "threads": threads, + } + _StubCrossEncoder.init_count += 1 + + def rerank(self, query: str, documents: list[str]): + # Score = count of query-token overlaps, so tests can assert ordering. + tokens = set(query.lower().split()) + for doc in documents: + yield float(len(tokens & set(doc.lower().split()))) + + +def _install_stub(monkeypatch) -> None: + module = type(sys)("fastembed.rerank.cross_encoder") + setattr(module, "TextCrossEncoder", _StubCrossEncoder) + monkeypatch.setitem(sys.modules, "fastembed.rerank.cross_encoder", module) + _StubCrossEncoder.init_count = 0 + + +def _http_error(status_code: int) -> requests_exceptions.HTTPError: + response = Response() + response.status_code = status_code + return requests_exceptions.HTTPError(f"HTTP {status_code}", response=response) + + +@pytest.mark.asyncio +async def test_lazy_loads_once_and_reuses_model(monkeypatch): + _install_stub(monkeypatch) + provider = FastEmbedRerankProvider(model_name="stub-reranker") + assert provider._model is None + + await provider.rerank("auth token", ["auth token doc", "unrelated"]) + await provider.rerank("auth token", ["another auth doc"]) + + assert _StubCrossEncoder.init_count == 1 + assert provider._model is not None + + +@pytest.mark.asyncio +async def test_cancelled_first_waiter_does_not_start_a_second_model_load(monkeypatch): + """Later searches join the worker thread left running by a cancelled request.""" + provider = FastEmbedRerankProvider(model_name="stub-reranker") + loop = asyncio.get_running_loop() + load_started = asyncio.Event() + release_load = threading.Event() + model = object() + load_count = 0 + + def create_model(): + nonlocal load_count + load_count += 1 + loop.call_soon_threadsafe(load_started.set) + assert release_load.wait(timeout=5) + return model + + monkeypatch.setattr(provider, "_create_model", create_model) + + first_waiter = asyncio.create_task(provider._load_model()) + await load_started.wait() + first_waiter.cancel() + with pytest.raises(asyncio.CancelledError): + await first_waiter + + second_waiter = asyncio.create_task(provider._load_model()) + for _ in range(100): + if load_count > 1: + break + await asyncio.sleep(0.001) + release_load.set() + + assert await second_waiter is model + assert load_count == 1 + + +@pytest.mark.asyncio +async def test_waiter_reuses_model_loaded_before_it_acquires_the_lock(monkeypatch): + """The in-lock double check wins if another load finishes while a caller waits.""" + _install_stub(monkeypatch) + provider = FastEmbedRerankProvider(model_name="stub-reranker") + model = provider._create_model() + await provider._model_lock.acquire() + waiter = asyncio.create_task(provider._load_model()) + await asyncio.sleep(0) + + provider._model = model + provider._model_lock.release() + + assert await waiter is model + + +@pytest.mark.asyncio +async def test_rerank_squashes_logits_to_unit_interval_in_input_order(monkeypatch): + """Raw cross-encoder logits are sigmoid-squashed to [0, 1], preserving order.""" + _install_stub(monkeypatch) + provider = FastEmbedRerankProvider(model_name="stub-reranker") + + # Stub logits: 0.0 (no overlap) and 2.0 (two-token overlap). + scores = await provider.rerank("auth token", ["nothing here", "auth token match"]) + + assert scores == [pytest.approx(0.5), pytest.approx(0.8807970779778823)] + assert all(0.0 <= s <= 1.0 for s in scores) + + +@pytest.mark.asyncio +async def test_rerank_sigmoid_handles_extreme_logits(monkeypatch): + """Large-magnitude logits are clamped so exp() never overflows.""" + module = type(sys)("fastembed.rerank.cross_encoder") + + class _Extreme: + def __init__(self, **kwargs): + pass + + def rerank(self, query, documents): + yield -1000.0 + yield 1000.0 + + setattr(module, "TextCrossEncoder", _Extreme) + monkeypatch.setitem(sys.modules, "fastembed.rerank.cross_encoder", module) + + provider = FastEmbedRerankProvider(model_name="stub-reranker") + scores = await provider.rerank("q", ["a", "b"]) + assert scores[0] == pytest.approx(0.0, abs=1e-6) + assert scores[1] == pytest.approx(1.0, abs=1e-6) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("logit", [float("nan"), float("inf"), float("-inf")]) +async def test_rerank_rejects_non_finite_logits(monkeypatch, logit): + module = type(sys)("fastembed.rerank.cross_encoder") + + class _NonFinite: + def __init__(self, **kwargs): + pass + + def rerank(self, query, documents): + yield logit + + setattr(module, "TextCrossEncoder", _NonFinite) + monkeypatch.setitem(sys.modules, "fastembed.rerank.cross_encoder", module) + + provider = FastEmbedRerankProvider(model_name="stub-reranker") + with pytest.raises(RerankProviderContractError, match="non-finite logit"): + await provider.rerank("q", ["a"]) + + +@pytest.mark.asyncio +async def test_rerank_rejects_missing_model_scores(monkeypatch): + module = type(sys)("fastembed.rerank.cross_encoder") + + class _Incomplete: + def __init__(self, **kwargs): + pass + + def rerank(self, query, documents): + yield 0.5 + + setattr(module, "TextCrossEncoder", _Incomplete) + monkeypatch.setitem(sys.modules, "fastembed.rerank.cross_encoder", module) + + provider = FastEmbedRerankProvider(model_name="stub-reranker") + with pytest.raises(RerankProviderContractError, match="1 scores for 2 documents"): + await provider.rerank("q", ["a", "b"]) + + +@pytest.mark.asyncio +async def test_empty_documents_short_circuit(monkeypatch): + _install_stub(monkeypatch) + provider = FastEmbedRerankProvider(model_name="stub-reranker") + assert await provider.rerank("auth", []) == [] + # Empty input must not trigger a model load. + assert provider._model is None + + +@pytest.mark.asyncio +async def test_passes_cache_dir_and_threads_to_model(monkeypatch): + _install_stub(monkeypatch) + provider = FastEmbedRerankProvider( + model_name="stub-reranker", cache_dir="/tmp/rr-cache", threads=3 + ) + await provider.rerank("auth", ["auth doc"]) + assert _StubCrossEncoder.last_init_kwargs == { + "model_name": "stub-reranker", + "cache_dir": "/tmp/rr-cache", + "threads": 3, + } + + +@pytest.mark.asyncio +async def test_missing_dependency_raises_actionable_error(monkeypatch): + monkeypatch.delitem(sys.modules, "fastembed.rerank.cross_encoder", raising=False) + real_import = builtins.__import__ + + def _raising_import(name, *args, **kwargs): + if name.startswith("fastembed"): + raise ImportError("no fastembed") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", _raising_import) + provider = FastEmbedRerankProvider(model_name="stub-reranker") + with pytest.raises(SemanticDependenciesMissingError): + await provider.rerank("auth", ["auth doc"]) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "load_error", + [ + requests_exceptions.Timeout("download timed out"), + requests_exceptions.ConnectionError("connection reset"), + _http_error(503), + ValueError("Could not load model stub-reranker from any source."), + ], +) +async def test_transient_model_download_error_is_classified(monkeypatch, load_error): + monkeypatch.delenv("HF_HUB_OFFLINE", raising=False) + monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising=False) + module = type(sys)("fastembed.rerank.cross_encoder") + + class _DownloadFailure: + def __init__(self, **kwargs): + raise load_error + + setattr(module, "TextCrossEncoder", _DownloadFailure) + monkeypatch.setitem(sys.modules, "fastembed.rerank.cross_encoder", module) + + provider = FastEmbedRerankProvider(model_name="stub-reranker") + with pytest.raises(RerankTransientError, match="model download failed temporarily"): + await provider.rerank("auth", ["auth doc"]) + assert provider._model is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("offline_env", ["HF_HUB_OFFLINE", "TRANSFORMERS_OFFLINE"]) +async def test_offline_cache_miss_remains_permanent(monkeypatch, offline_env): + monkeypatch.delenv("HF_HUB_OFFLINE", raising=False) + monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising=False) + monkeypatch.setenv(offline_env, "1") + module = type(sys)("fastembed.rerank.cross_encoder") + load_error = ValueError("Could not load model stub-reranker from any source.") + + class _OfflineCacheMiss: + def __init__(self, **kwargs): + raise load_error + + setattr(module, "TextCrossEncoder", _OfflineCacheMiss) + monkeypatch.setitem(sys.modules, "fastembed.rerank.cross_encoder", module) + + provider = FastEmbedRerankProvider(model_name="stub-reranker") + with pytest.raises(ValueError, match="Could not load model"): + await provider.rerank("auth", ["auth doc"]) + assert provider._model is None + + +@pytest.mark.asyncio +async def test_unsupported_model_error_remains_permanent(monkeypatch): + module = type(sys)("fastembed.rerank.cross_encoder") + + class _UnsupportedModel: + def __init__(self, **kwargs): + raise ValueError("Model unknown is not supported in TextCrossEncoder.") + + setattr(module, "TextCrossEncoder", _UnsupportedModel) + monkeypatch.setitem(sys.modules, "fastembed.rerank.cross_encoder", module) + + provider = FastEmbedRerankProvider(model_name="unknown") + with pytest.raises(ValueError, match="not supported"): + await provider.rerank("auth", ["auth doc"]) + + +@pytest.mark.asyncio +async def test_download_auth_error_remains_permanent(monkeypatch): + module = type(sys)("fastembed.rerank.cross_encoder") + auth_error = _http_error(401) + + class _Unauthorized: + def __init__(self, **kwargs): + raise auth_error + + setattr(module, "TextCrossEncoder", _Unauthorized) + monkeypatch.setitem(sys.modules, "fastembed.rerank.cross_encoder", module) + + provider = FastEmbedRerankProvider(model_name="private-model") + with pytest.raises(requests_exceptions.HTTPError): + await provider.rerank("auth", ["auth doc"]) + + +def test_runtime_log_attrs(): + provider = FastEmbedRerankProvider(model_name="stub-reranker", threads=2) + assert provider.runtime_log_attrs() == {"model_name": "stub-reranker", "threads": 2} diff --git a/tests/repository/test_litellm_rerank_provider.py b/tests/repository/test_litellm_rerank_provider.py new file mode 100644 index 000000000..a6d89d481 --- /dev/null +++ b/tests/repository/test_litellm_rerank_provider.py @@ -0,0 +1,263 @@ +"""Tests for LiteLLMRerankProvider.""" + +import types + +import pytest +from pydantic import BaseModel, ValidationError + +from basic_memory.repository.litellm_rerank_provider import LiteLLMRerankProvider +from basic_memory.repository.semantic_errors import ( + RerankProviderContractError, + RerankTransientError, +) + + +class _Response: + def __init__(self, results): + self.results = results + + +class _TransientProviderError(RuntimeError): + pass + + +class _BadGatewayError(RuntimeError): + pass + + +class _SDKRerankResponse(BaseModel): + results: list[dict] + + +def _fake_litellm(response, recorder: dict, *, exc: Exception | None = None): + async def arerank(**params): + recorder.update(params) + if exc is not None: + raise exc + return response + + return types.SimpleNamespace( + arerank=arerank, + Timeout=_TransientProviderError, + APIConnectionError=_TransientProviderError, + RateLimitError=_TransientProviderError, + BadGatewayError=_BadGatewayError, + ServiceUnavailableError=_TransientProviderError, + InternalServerError=_TransientProviderError, + ) + + +@pytest.mark.asyncio +async def test_rerank_realigns_out_of_order_indexed_results(monkeypatch): + """Rerank responses are indexed and may arrive out of order; realign to input.""" + recorder: dict = {} + response = _Response( + [ + {"index": 2, "relevance_score": 0.9}, + {"index": 0, "relevance_score": 0.1}, + {"index": 1, "relevance_score": 0.5}, + ] + ) + monkeypatch.setattr( + "basic_memory.repository.litellm_rerank_provider._import_litellm", + lambda: _fake_litellm(response, recorder), + ) + provider = LiteLLMRerankProvider(model_name="cohere/rerank-v3.5") + + scores = await provider.rerank("q", ["doc0", "doc1", "doc2"]) + + assert scores == [0.1, 0.5, 0.9] + assert recorder["model"] == "cohere/rerank-v3.5" + assert recorder["top_n"] == 3 + + +@pytest.mark.asyncio +async def test_rerank_forwards_routing_params(monkeypatch): + recorder: dict = {} + response = _Response( + [{"index": 0, "relevance_score": 0.7}, {"index": 1, "relevance_score": 0.2}] + ) + monkeypatch.setattr( + "basic_memory.repository.litellm_rerank_provider._import_litellm", + lambda: _fake_litellm(response, recorder), + ) + provider = LiteLLMRerankProvider( + model_name="jina_ai/jina-reranker-v2", + api_key="secret", + api_base="https://rr.example", + ) + + scores = await provider.rerank("q", ["a", "b"]) + + assert scores == [0.7, 0.2] + assert recorder["api_key"] == "secret" + assert recorder["api_base"] == "https://rr.example" + + +@pytest.mark.asyncio +async def test_incomplete_response_raises(monkeypatch): + """A response that omits a requested document is a fault, not a silent 0.0.""" + response = _Response([{"index": 1, "relevance_score": 0.8}]) # index 0 missing + monkeypatch.setattr( + "basic_memory.repository.litellm_rerank_provider._import_litellm", + lambda: _fake_litellm(response, {}), + ) + provider = LiteLLMRerankProvider() + with pytest.raises(RerankProviderContractError, match="covered 1 of 2 documents"): + await provider.rerank("q", ["dropped", "kept"]) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("bad_index", [-1, 2]) +async def test_out_of_range_index_raises(monkeypatch, bad_index): + """Indices outside [0, len) are a malformed response, not a valid position. + + A negative index would silently overwrite a wrong slot via Python indexing and + could even satisfy the coverage check; fail fast instead. + """ + response = _Response( + [ + {"index": bad_index, "relevance_score": 0.8}, + {"index": 0, "relevance_score": 0.4}, + ] + ) + monkeypatch.setattr( + "basic_memory.repository.litellm_rerank_provider._import_litellm", + lambda: _fake_litellm(response, {}), + ) + provider = LiteLLMRerankProvider() + with pytest.raises(RerankProviderContractError, match="out of range"): + await provider.rerank("q", ["a", "b"]) + + +@pytest.mark.asyncio +async def test_duplicate_index_raises(monkeypatch): + """A repeated index with otherwise-full coverage slips past the coverage check. + + `[0, 1, 0]` for two documents leaves `seen == {0, 1}` — the count check passes — + while the second index-0 item silently overwrites the first score. Reject the + duplicate so the "each index exactly once" contract holds. + """ + response = _Response( + [ + {"index": 0, "relevance_score": 0.8}, + {"index": 1, "relevance_score": 0.5}, + {"index": 0, "relevance_score": 0.1}, + ] + ) + monkeypatch.setattr( + "basic_memory.repository.litellm_rerank_provider._import_litellm", + lambda: _fake_litellm(response, {}), + ) + provider = LiteLLMRerankProvider() + with pytest.raises(RerankProviderContractError, match="repeated index 0"): + await provider.rerank("q", ["a", "b"]) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("results", [None, []]) +async def test_missing_results_raises(monkeypatch, results): + """litellm's RerankResponse.results is optional; None/empty is a contract break.""" + monkeypatch.setattr( + "basic_memory.repository.litellm_rerank_provider._import_litellm", + lambda: _fake_litellm(_Response(results), {}), + ) + provider = LiteLLMRerankProvider() + with pytest.raises(RerankProviderContractError, match="no results"): + await provider.rerank("q", ["a", "b"]) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("bad_score", [-0.1, 1.1, float("nan"), float("inf")]) +async def test_invalid_relevance_score_raises(monkeypatch, bad_score): + response = _Response([{"index": 0, "relevance_score": bad_score}]) + monkeypatch.setattr( + "basic_memory.repository.litellm_rerank_provider._import_litellm", + lambda: _fake_litellm(response, {}), + ) + provider = LiteLLMRerankProvider() + + with pytest.raises(RerankProviderContractError, match=r"finite and in \[0, 1\]"): + await provider.rerank("q", ["a"]) + + +@pytest.mark.asyncio +async def test_transient_provider_error_is_classified_for_fallback(monkeypatch): + transient_error = _TransientProviderError("provider unavailable") + monkeypatch.setattr( + "basic_memory.repository.litellm_rerank_provider._import_litellm", + lambda: _fake_litellm(None, {}, exc=transient_error), + ) + provider = LiteLLMRerankProvider() + + with pytest.raises(RerankTransientError) as caught: + await provider.rerank("q", ["a"]) + + assert caught.value.__cause__ is transient_error + + +@pytest.mark.asyncio +async def test_bad_gateway_error_is_classified_for_fallback(monkeypatch): + bad_gateway_error = _BadGatewayError("upstream returned 502") + monkeypatch.setattr( + "basic_memory.repository.litellm_rerank_provider._import_litellm", + lambda: _fake_litellm(None, {}, exc=bad_gateway_error), + ) + provider = LiteLLMRerankProvider() + + with pytest.raises(RerankTransientError) as caught: + await provider.rerank("q", ["a"]) + + assert caught.value.__cause__ is bad_gateway_error + + +@pytest.mark.asyncio +async def test_sdk_response_validation_error_is_provider_contract_error(monkeypatch): + with pytest.raises(ValidationError) as invalid_response: + _SDKRerankResponse.model_validate({}) + monkeypatch.setattr( + "basic_memory.repository.litellm_rerank_provider._import_litellm", + lambda: _fake_litellm(None, {}, exc=invalid_response.value), + ) + provider = LiteLLMRerankProvider() + + with pytest.raises(RerankProviderContractError, match="invalid response") as caught: + await provider.rerank("q", ["a"]) + + assert caught.value.__cause__ is invalid_response.value + + +@pytest.mark.asyncio +async def test_unknown_provider_error_surfaces(monkeypatch): + permanent_error = RuntimeError("invalid credentials or model") + monkeypatch.setattr( + "basic_memory.repository.litellm_rerank_provider._import_litellm", + lambda: _fake_litellm(None, {}, exc=permanent_error), + ) + provider = LiteLLMRerankProvider() + + with pytest.raises(RuntimeError, match="invalid credentials or model"): + await provider.rerank("q", ["a"]) + + +@pytest.mark.asyncio +async def test_empty_documents_short_circuit(monkeypatch): + called = False + + def _boom(): + nonlocal called + called = True + raise AssertionError("litellm must not be imported for empty input") + + monkeypatch.setattr("basic_memory.repository.litellm_rerank_provider._import_litellm", _boom) + provider = LiteLLMRerankProvider() + assert await provider.rerank("q", []) == [] + assert called is False + + +def test_runtime_log_attrs(): + provider = LiteLLMRerankProvider(model_name="cohere/rerank-v3.5", api_base="https://rr") + assert provider.runtime_log_attrs() == { + "model_name": "cohere/rerank-v3.5", + "api_base_set": True, + } diff --git a/tests/repository/test_rerank_pipeline.py b/tests/repository/test_rerank_pipeline.py new file mode 100644 index 000000000..80dfe1b3b --- /dev/null +++ b/tests/repository/test_rerank_pipeline.py @@ -0,0 +1,841 @@ +"""Rerank stage wiring in the shared search pipeline (vector + hybrid).""" + +from datetime import datetime, timezone +from typing import Any +from unittest.mock import MagicMock + +import pytest + +from basic_memory.config import BasicMemoryConfig, DatabaseBackend +from basic_memory.repository.search_index_row import SearchIndexRow +from basic_memory.repository.rerank_provider import demote_tail_scores, validate_rerank_scores +from basic_memory.repository.search_repository_base import RERANK_POOL_CHUNK_FANOUT +from basic_memory.repository.semantic_errors import ( + RerankProviderContractError, + RerankTransientError, + SemanticDependenciesMissingError, +) +from basic_memory.repository.sqlite_search_repository import SQLiteSearchRepository +from basic_memory.schemas.search import SearchItemType, SearchRetrievalMode + + +class _StubEmbeddingProvider: + """Deterministic embeddings that give the two auth notes DIFFERENT similarity. + + An "auth" doc containing "deep" is tilted slightly off the query axis, so vector + retrieval ranks the plain-auth note strictly above it. That makes the pre-rerank + baseline a real ordering (not a tie), so a rerank that promotes the lower note is + a genuine "recover the below-cutoff doc" scenario (#950), not a coin flip. + """ + + model_name = "stub" + dimensions = 4 + + async def embed_query(self, text: str) -> list[float]: + return self._vectorize(text) + + async def embed_documents(self, texts: list[str]) -> list[list[float]]: + return [self._vectorize(t) for t in texts] + + def runtime_log_attrs(self) -> dict: + return {} + + @staticmethod + def _vectorize(text: str) -> list[float]: + lowered = text.lower() + if "auth" not in lowered: + return [0.0, 0.0, 0.0, 1.0] + # Unit vectors; cos with the query axis [1,0,0,0] is 1.0 vs 0.9. + if "deep" in lowered: + return [0.9, 0.4358898943540674, 0.0, 0.0] + return [1.0, 0.0, 0.0, 0.0] + + +class _FakeReranker: + """Scores a document by the marker substring it contains; records call count.""" + + model_name = "fake-reranker" + + def __init__(self, score_by_marker: dict[str, float]): + self.score_by_marker = score_by_marker + self.calls = 0 + + async def rerank(self, query: str, documents: list[str]) -> list[float]: + self.calls += 1 + scores = [] + for doc in documents: + score = 0.0 + for marker, value in self.score_by_marker.items(): + if marker in doc: + score = value + scores.append(score) + return scores + + def runtime_log_attrs(self) -> dict: + return {} + + +class _BadReranker: + model_name = "bad" + + async def rerank(self, query: str, documents: list[str]) -> list[float]: + return [] # deliberately misaligned (no exception) + + def runtime_log_attrs(self) -> dict: + return {} + + +class _ExplodingReranker: + """Typed transient failure the pipeline must surface.""" + + model_name = "boom" + + async def rerank(self, query: str, documents: list[str]) -> list[float]: + raise RerankTransientError("cross-encoder backend unreachable") + + def runtime_log_attrs(self) -> dict: + return {} + + +class _SucceedsThenTransientReranker: + """Rerank one page, then model a temporary provider outage.""" + + model_name = "flaky" + + def __init__(self): + self.calls = 0 + + async def rerank(self, query: str, documents: list[str]) -> list[float]: + self.calls += 1 + if self.calls > 1: + raise RerankTransientError("cross-encoder backend unreachable") + return [0.1, 0.9] + + def runtime_log_attrs(self) -> dict: + return {} + + +class _PermanentFaultReranker: + """Permanent fault (bad config/deps): the pipeline must surface, not swallow.""" + + model_name = "permanent" + + def __init__(self, exc: Exception): + self._exc = exc + + async def rerank(self, query: str, documents: list[str]) -> list[float]: + raise self._exc + + def runtime_log_attrs(self) -> dict: + return {} + + +def test_validate_rerank_scores_rejects_non_numeric_value(): + """Provider contract validation rejects values that cannot become finite floats.""" + with pytest.raises(RerankProviderContractError, match="is not a number"): + validate_rerank_scores(["not-a-score"], expected_count=1) + + +def _entity_row(*, project_id: int, row_id: int, title: str, permalink: str, content: str): + now = datetime.now(timezone.utc) + return SearchIndexRow( + project_id=project_id, + id=row_id, + type=SearchItemType.ENTITY.value, + title=title, + permalink=permalink, + file_path=f"{permalink}.md", + metadata={"note_type": "spec"}, + entity_id=row_id, + content_stems=content, + content_snippet=content, + created_at=now, + updated_at=now, + ) + + +def _row(**overrides) -> SearchIndexRow: + now = datetime.now(timezone.utc) + base: dict[str, Any] = dict( + project_id=1, + id=1, + type=SearchItemType.ENTITY.value, + file_path="x.md", + created_at=now, + updated_at=now, + title="Title", + content_snippet="snippet", + score=0.5, + ) + base.update(overrides) + return SearchIndexRow(**base) + + +def _unit_repo() -> SQLiteSearchRepository: + """A repo built without a real DB — for the pure rerank helper methods.""" + config = BasicMemoryConfig( + env="test", + projects={"test-project": "/tmp/test"}, + default_project="test-project", + database_backend=DatabaseBackend.SQLITE, + semantic_search_enabled=True, + ) + return SQLiteSearchRepository( + MagicMock(), + project_id=1, + app_config=config, + embedding_provider=_StubEmbeddingProvider(), + ) + + +# --- Pure helper behavior --- + + +def test_should_rerank_gating(): + repo = _unit_repo() + repo._rerank_provider = None + assert repo._should_rerank("auth") is False + repo._rerank_provider = _FakeReranker({}) + assert repo._should_rerank("") is False + assert repo._should_rerank("auth") is True + + +def test_rerank_document_text_fallbacks(): + repo = _unit_repo() + assert repo._rerank_document_text(_row(title="T", matched_chunk_text="chunk")) == "chunk\nT" + assert ( + repo._rerank_document_text(_row(title="T", matched_chunk_text=None, content_snippet="snip")) + == "snip\nT" + ) + assert ( + repo._rerank_document_text(_row(title=None, matched_chunk_text="only-body")) == "only-body" + ) + assert ( + repo._rerank_document_text(_row(title="only-title", content_snippet=None)) == "only-title" + ) + assert repo._rerank_document_text(_row(title=None, content_snippet=None)) == "" + + +def test_rerank_document_text_truncation(): + repo = _unit_repo() + row = _row(title="T", matched_chunk_text="x" * 500) # full text = 500 + "\nT" = 502 chars + repo._reranker_max_document_chars = 0 # disabled + assert len(repo._rerank_document_text(row)) == 502 + repo._reranker_max_document_chars = 100 # trims to the leading (most-relevant) text + trimmed = repo._rerank_document_text(row) + assert len(trimmed) == 100 and trimmed == "x" * 100 + repo._reranker_max_document_chars = 10_000 # no-op when already under the cap + assert len(repo._rerank_document_text(row)) == 502 + + +def test_rerank_document_text_cap_preserves_matched_body_with_long_title(): + repo = _unit_repo() + repo._reranker_max_document_chars = 8 + row = _row(title="title-" * 20, matched_chunk_text="MATCHED body") + + assert "MATCHED" in repo._rerank_document_text(row) + + +def test_demoted_tail_scores_are_stable_as_the_tail_grows(): + """An existing tail rank keeps its score when later results are fetched.""" + first_page_scores = demote_tail_scores(floor=0.2, count=2) + growing_prefix_scores = demote_tail_scores(floor=0.2, count=5) + + assert growing_prefix_scores[:2] == first_page_scores + assert growing_prefix_scores == sorted(growing_prefix_scores, reverse=True) + assert all(0.0 < score < 0.2 for score in growing_prefix_scores) + + +def test_candidate_limit_over_fetches_chunks_for_rerank_pool(): + """With reranking active, over-fetch chunks so dedup can't starve the rerank window.""" + repo = _unit_repo() + repo._semantic_vector_k = 5 + repo._reranker_candidates = 20 + + repo._rerank_provider = None + assert repo._candidate_limit(limit=1, offset=0, query_text="auth") == 10 # max(5, 10) + + repo._rerank_provider = _FakeReranker({}) + assert ( + repo._candidate_limit(limit=1, offset=0, query_text="auth") == 20 * RERANK_POOL_CHUNK_FANOUT + ) + assert repo._candidate_limit(limit=1, offset=0, query_text="") == 10 # no query → no bump + + +def test_candidate_limit_expands_only_for_results_beyond_rerank_pool(): + """The fixed rerank window grows only enough to supply the requested tail.""" + repo = _unit_repo() + repo._semantic_vector_k = 5 + repo._reranker_candidates = 20 + repo._rerank_provider = _FakeReranker({}) + + first_page_limit = repo._candidate_limit(limit=10, offset=0, query_text="auth") + + assert first_page_limit == 20 * RERANK_POOL_CHUNK_FANOUT + assert repo._candidate_limit(limit=20, offset=0, query_text="auth") == first_page_limit + assert repo._candidate_limit(limit=10, offset=10, query_text="auth") == first_page_limit + assert repo._candidate_limit(limit=21, offset=0, query_text="auth") == 90 + assert repo._candidate_limit(limit=10, offset=19, query_text="auth") == 170 + assert repo._candidate_limit(limit=10, offset=20, query_text="auth") == 180 + # Large first pages still retrieve their untouched tail and pagination probe. + assert repo._candidate_limit(limit=101, offset=0, query_text="auth") == 890 + + +@pytest.mark.asyncio +async def test_rerank_paginate_noop_paths(): + repo = _unit_repo() + rows = [_row(id=1), _row(id=2)] + repo._rerank_provider = None + assert await repo._rerank_and_paginate("auth", rows, offset=0, limit=10) == rows + + reranker = _FakeReranker({}) + repo._rerank_provider = reranker + assert await repo._rerank_and_paginate("", rows, offset=0, limit=10) == rows + assert await repo._rerank_and_paginate("auth", [], offset=0, limit=10) == [] + assert await repo._rerank_and_paginate("auth", rows, offset=2, limit=10) == [] + assert reranker.calls == 0 + + +@pytest.mark.asyncio +async def test_rerank_paginate_reorders_rescore_and_demotes_tail(): + repo = _unit_repo() + repo._rerank_provider = _FakeReranker({"Alpha": 0.1, "Bravo": 0.9, "Charlie": 0.5}) + repo._reranker_candidates = 2 + rows = [ + _row(id=1, title="Alpha"), + _row(id=2, title="Bravo"), + _row(id=3, title="Charlie"), # past the pool + ] + + result = await repo._rerank_and_paginate("auth", rows, offset=0, limit=3) + + assert [r.title for r in result] == ["Bravo", "Alpha", "Charlie"] + assert result[0].score == 0.9 # reranker relevance replaces the prior score + # Tail is demoted strictly below the reranked floor (0.1) so the page stays + # monotonic and in [0, 1] — no scale mixing. + scores = [r.score for r in result if r.score is not None] + assert len(scores) == 3 + assert scores == sorted(scores, reverse=True) + assert all(0.0 <= s <= 1.0 for s in scores) + assert result[2].title == "Charlie" + assert scores[2] < scores[1] + + +@pytest.mark.asyncio +async def test_rerank_paginate_preserves_pool_before_tail_at_zero_floor(): + repo = _unit_repo() + repo._rerank_provider = _FakeReranker({"Alpha": 0.0, "Bravo": 0.9}) + repo._reranker_candidates = 2 + rows = [ + _row(id=1, title="Alpha"), + _row(id=2, title="Bravo"), + _row(id=3, title="Charlie"), + ] + + result = await repo._rerank_and_paginate("auth", rows, offset=0, limit=3) + + assert [row.title for row in result] == ["Bravo", "Alpha", "Charlie"] + assert [row.score for row in result] == [0.9, 0.0, 0.0] + + +@pytest.mark.asyncio +async def test_rerank_paginate_scores_singleton_prefix_and_demotes_tail(): + """A one-row prefix still calibrates scores before cross-project merging.""" + repo = _unit_repo() + reranker = _FakeReranker({"Only": 0.4}) + repo._rerank_provider = reranker + repo._reranker_candidates = 2 + stable_rows = [_row(id=1, title="Only", score=0.5)] + expanded_rows = [ + stable_rows[0], + _row(id=2, title="Tail", score=1.3), + ] + + result = await repo._rerank_and_paginate( + "auth", + expanded_rows, + offset=0, + limit=2, + stable_rows=stable_rows, + ) + + assert reranker.calls == 1 + assert [row.title for row in result] == ["Only", "Tail"] + assert result[0].score == 0.4 + assert result[1].score is not None and result[1].score < 0.4 + + +@pytest.mark.asyncio +async def test_rerank_paginate_calibrates_tail_scores_on_deep_page(): + """Deep pages rescore the fixed prefix before returning its calibrated tail.""" + repo = _unit_repo() + reranker = _FakeReranker({"n1": 0.9, "n2": 0.8}) + repo._rerank_provider = reranker + repo._reranker_candidates = 2 + stable_rows = [_row(id=1, title="n1"), _row(id=2, title="n2")] + expanded_rows = [ + _row(id=3, title="newly strengthened"), + stable_rows[0], + stable_rows[1], + _row(id=4, title="n4"), + _row(id=5, title="n5"), + ] + + result = await repo._rerank_and_paginate( + "auth", + expanded_rows, + offset=2, + limit=2, + stable_rows=stable_rows, + ) + + assert reranker.calls == 1 + assert [r.id for r in result] == [3, 4] + assert all(row.score is not None and row.score < 0.8 for row in result) + + +@pytest.mark.asyncio +async def test_rerank_paginate_keeps_expanded_candidates_out_of_stable_prefix(): + """A larger tail retrieval cannot replace candidates in the reranked prefix.""" + repo = _unit_repo() + reranker = _FakeReranker({"Alpha": 0.1, "Bravo": 0.9, "Charlie": 1.0}) + repo._rerank_provider = reranker + repo._reranker_candidates = 2 + stable_rows = [_row(id=1, title="Alpha"), _row(id=2, title="Bravo")] + expanded_rows = [ + _row(id=3, title="Charlie"), + stable_rows[0], + stable_rows[1], + _row(id=4, title="Delta"), + ] + + result = await repo._rerank_and_paginate( + "auth", + expanded_rows, + offset=0, + limit=3, + stable_rows=stable_rows, + ) + expanded_result = await repo._rerank_and_paginate( + "auth", + expanded_rows + [_row(id=5, title="Echo"), _row(id=6, title="Foxtrot")], + offset=0, + limit=3, + stable_rows=stable_rows, + ) + + assert [row.title for row in result] == ["Bravo", "Alpha", "Charlie"] + assert [row.title for row in expanded_result] == ["Bravo", "Alpha", "Charlie"] + assert [row.score for row in expanded_result] == [row.score for row in result] + assert reranker.calls == 2 + + +@pytest.mark.asyncio +async def test_rerank_paginate_surfaces_transient_provider_error(): + """Transient failures must not silently replace reranked order with retrieval order.""" + repo = _unit_repo() + repo._rerank_provider = _ExplodingReranker() + repo._reranker_candidates = 20 + rows = [_row(id=1, title="A"), _row(id=2, title="B")] + + with pytest.raises(RerankTransientError, match="backend unreachable"): + await repo._rerank_and_paginate("auth", rows, offset=0, limit=10) + + +@pytest.mark.asyncio +async def test_rerank_paginate_does_not_duplicate_results_when_later_page_is_transient(): + """A later page fails instead of changing order and repeating an earlier result.""" + repo = _unit_repo() + repo._rerank_provider = _SucceedsThenTransientReranker() + repo._reranker_candidates = 2 + rows = [_row(id=1, title="A"), _row(id=2, title="B")] + + first_page = await repo._rerank_and_paginate("auth", rows, offset=0, limit=1) + + assert [row.id for row in first_page] == [2] + with pytest.raises(RerankTransientError, match="backend unreachable"): + await repo._rerank_and_paginate("auth", rows, offset=1, limit=1) + + +@pytest.mark.asyncio +async def test_rerank_paginate_misaligned_scores_raise(): + """A length mismatch is a provider bug — fail fast, don't degrade.""" + repo = _unit_repo() + repo._rerank_provider = _BadReranker() + repo._reranker_candidates = 20 + with pytest.raises(RerankProviderContractError, match="Reranker returned 0 scores"): + await repo._rerank_and_paginate("auth", [_row(id=1), _row(id=2)], offset=0, limit=10) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "exc", + [ + RerankProviderContractError("incomplete rerank response"), + SemanticDependenciesMissingError("fastembed missing"), + RuntimeError("unexpected provider failure"), + ], +) +async def test_rerank_paginate_surfaces_permanent_faults(exc): + """Permanent faults (contract break, missing deps) propagate — not silently degraded.""" + repo = _unit_repo() + repo._rerank_provider = _PermanentFaultReranker(exc) + repo._reranker_candidates = 20 + with pytest.raises(type(exc)): + await repo._rerank_and_paginate("auth", [_row(id=1), _row(id=2)], offset=0, limit=10) + + +# --- End-to-end through the SQLite repo --- + + +def _enable_semantic(repo: SQLiteSearchRepository) -> None: + try: + import sqlite_vec # noqa: F401 + except ImportError: # pragma: no cover + pytest.skip("sqlite-vec dependency is required for vector search tests.") + repo._semantic_enabled = True + repo._embedding_provider = _StubEmbeddingProvider() + repo._vector_dimensions = 4 + repo._vector_tables_initialized = False + repo._semantic_min_similarity = 0.0 + + +async def _index_two_auth_notes(repo: SQLiteSearchRepository) -> None: + await repo.init_search_index() + await repo.bulk_index_items( + [ + _entity_row( + project_id=repo.project_id, + row_id=401, + title="Alpha Auth Guide", + permalink="specs/alpha", + content="auth login session token overview", + ), + _entity_row( + project_id=repo.project_id, + row_id=402, + title="Bravo Auth Guide", + permalink="specs/bravo", + content="auth login session token deep dive", + ), + ] + ) + await repo.sync_entity_vectors(401) + await repo.sync_entity_vectors(402) + + +@pytest.mark.asyncio +async def test_vector_search_applies_reranker(search_repository): + if not isinstance(search_repository, SQLiteSearchRepository): + pytest.skip("sqlite-vec repository behavior is local SQLite-only.") + + _enable_semantic(search_repository) + await _index_two_auth_notes(search_repository) + # Promote the note that vector similarity alone leaves tied/second. + reranker = _FakeReranker({"Alpha": 0.1, "Bravo": 0.9}) + search_repository._rerank_provider = reranker + + results = await search_repository.search( + search_text="auth session token", + retrieval_mode=SearchRetrievalMode.VECTOR, + limit=5, + ) + + # Baseline vector order is alpha > bravo (see test_search_without_reranker); + # the reranker promotes bravo — a genuine below-cutoff recovery, not a tie-break. + assert [r.permalink for r in results[:2]] == ["specs/bravo", "specs/alpha"] + assert results[0].score == 0.9 + scores = [r.score for r in results if r.score is not None] + assert scores == sorted(scores, reverse=True) + assert reranker.calls == 1 + + +@pytest.mark.asyncio +async def test_vector_search_expands_tail_from_stable_rerank_pool( + search_repository, + monkeypatch, +): + """Vector retrieval keeps its fixed prefix when a request also needs tail rows.""" + if not isinstance(search_repository, SQLiteSearchRepository): + pytest.skip("sqlite-vec repository behavior is local SQLite-only.") + + _enable_semantic(search_repository) + await _index_two_auth_notes(search_repository) + search_repository._semantic_vector_k = 5 + search_repository._reranker_candidates = 2 + search_repository._rerank_provider = _FakeReranker({"Alpha": 0.1, "Bravo": 0.9}) + + candidate_limits: list[int] = [] + run_vector_query = search_repository._run_vector_query + + async def record_vector_query( + session: Any, + query_embedding: list[float], + candidate_limit: int, + ) -> list[dict[str, Any]]: + candidate_limits.append(candidate_limit) + return await run_vector_query(session, query_embedding, candidate_limit) + + monkeypatch.setattr(search_repository, "_run_vector_query", record_vector_query) + + results = await search_repository.search( + search_text="auth session token", + retrieval_mode=SearchRetrievalMode.VECTOR, + limit=3, + ) + + assert [row.permalink for row in results] == ["specs/bravo", "specs/alpha"] + assert candidate_limits == [18, 8] + + +@pytest.mark.asyncio +async def test_vector_slow_query_timing_includes_reranker(search_repository, monkeypatch): + if not isinstance(search_repository, SQLiteSearchRepository): + pytest.skip("sqlite-vec repository behavior is local SQLite-only.") + + _enable_semantic(search_repository) + await _index_two_auth_notes(search_repository) + clock = {"now": 0.0} + reranker = _FakeReranker({"Alpha": 0.1, "Bravo": 0.9}) + original_rerank = reranker.rerank + + async def slow_rerank(query: str, documents: list[str]) -> list[float]: + clock["now"] = 3.0 + return await original_rerank(query, documents) + + search_repository._rerank_provider = reranker + monkeypatch.setattr(reranker, "rerank", slow_rerank) + monkeypatch.setattr( + "basic_memory.repository.search_repository_base.time.perf_counter", + lambda: clock["now"], + ) + warning = MagicMock() + monkeypatch.setattr( + "basic_memory.repository.search_repository_base.logger.warning", + warning, + ) + + await search_repository.search( + search_text="auth session token", + retrieval_mode=SearchRetrievalMode.VECTOR, + limit=5, + ) + + warning.assert_called_once() + assert warning.call_args.args[0].startswith("[SEMANTIC_SLOW_QUERY]") + assert warning.call_args.kwargs["total_ms"] == 3000.0 + + +@pytest.mark.asyncio +async def test_hybrid_search_reranks_once(search_repository): + """Hybrid reranks the fused result exactly once — not again inside its vector leg.""" + if not isinstance(search_repository, SQLiteSearchRepository): + pytest.skip("sqlite-vec repository behavior is local SQLite-only.") + + _enable_semantic(search_repository) + await _index_two_auth_notes(search_repository) + reranker = _FakeReranker({"Alpha": 0.1, "Bravo": 0.9}) + search_repository._rerank_provider = reranker + + results = await search_repository.search( + search_text="auth session token", + retrieval_mode=SearchRetrievalMode.HYBRID, + limit=5, + ) + + assert results[0].permalink == "specs/bravo" + assert results[0].score == 0.9 + scores = [r.score for r in results if r.score is not None] + assert scores == sorted(scores, reverse=True) + assert reranker.calls == 1 + + +@pytest.mark.asyncio +async def test_hybrid_search_preserves_candidate_windows( + search_repository, + monkeypatch, +): + """Hybrid preserves legacy recall unless reranking owns the shared candidate pool.""" + if not isinstance(search_repository, SQLiteSearchRepository): + pytest.skip("sqlite-vec repository behavior is local SQLite-only.") + + _enable_semantic(search_repository) + await _index_two_auth_notes(search_repository) + search_repository._semantic_vector_k = 100 + search_repository._reranker_candidates = 100 + search_repository._rerank_provider = None + + candidate_limits: list[int] = [] + run_vector_query = search_repository._run_vector_query + + async def record_vector_query( + session: Any, + query_embedding: list[float], + candidate_limit: int, + ) -> list[dict[str, Any]]: + candidate_limits.append(candidate_limit) + return await run_vector_query(session, query_embedding, candidate_limit) + + monkeypatch.setattr(search_repository, "_run_vector_query", record_vector_query) + + baseline_results = await search_repository.search( + search_text="auth session token", + retrieval_mode=SearchRetrievalMode.HYBRID, + limit=10, + ) + + assert baseline_results + assert candidate_limits == [1000] + + candidate_limits.clear() + search_repository._semantic_vector_k = 5 + search_repository._reranker_candidates = 20 + reranker = _FakeReranker({"Alpha": 0.1, "Bravo": 0.9}) + search_repository._rerank_provider = reranker + reranked_results = await search_repository.search( + search_text="auth session token", + retrieval_mode=SearchRetrievalMode.HYBRID, + limit=11, + ) + + assert reranked_results + assert candidate_limits == [80] + + candidate_limits.clear() + growing_prefix_results = await search_repository.search( + search_text="auth session token", + retrieval_mode=SearchRetrievalMode.HYBRID, + limit=21, + ) + + assert growing_prefix_results + assert reranker.calls == 2 + assert candidate_limits == [90, 80] + + +@pytest.mark.asyncio +async def test_hybrid_search_keeps_deep_tail_stable_as_candidate_window_grows(monkeypatch): + """Late dual-source evidence cannot move a row across an earlier tail page.""" + repo = _unit_repo() + repo._semantic_vector_k = 2 + repo._reranker_candidates = 2 + reranker = _FakeReranker({"Alpha": 0.9, "Bravo": 0.8}) + repo._rerank_provider = reranker + + charlie = _row(id=3, title="Charlie") + delta = _row(id=4, title="Delta") + + async def fake_fts_search(*args, limit: int, **kwargs) -> list[SearchIndexRow]: + assert kwargs["retrieval_mode"] == SearchRetrievalMode.FTS + if limit <= 8: + return [ + _row(id=1, title="Alpha", score=10.0), + _row(id=2, title="Bravo", score=9.0), + ] + return [ + _row(id=1, title="Alpha", score=10.0), + _row(id=2, title="Bravo", score=9.0), + _row(id=3, title="Charlie", score=8.0), + _row(id=4, title="Delta", score=7.0), + ] + + async def fake_vector_search(**kwargs) -> list[SearchIndexRow]: + candidate_limit = kwargs["candidate_limit"] + if candidate_limit <= 8: + return [ + _row(id=1, title="Alpha", score=1.0), + _row(id=2, title="Bravo", score=0.9), + ] + if candidate_limit <= 18: + return [ + _row(id=1, title="Alpha", score=1.0), + _row(id=2, title="Bravo", score=0.9), + _row(id=4, title="Delta", score=0.95), + ] + return [ + _row(id=1, title="Alpha", score=1.0), + _row(id=2, title="Bravo", score=0.9), + _row(id=3, title="Charlie", score=1.0), + _row(id=4, title="Delta", score=0.7), + ] + + monkeypatch.setattr(repo, "search", fake_fts_search) + monkeypatch.setattr(repo, "_search_vector_only", fake_vector_search) + + async def deep_page(offset: int) -> list[SearchIndexRow]: + return await repo._search_hybrid( + search_text="auth", + permalink=None, + permalink_match=None, + title=None, + note_types=None, + after_date=None, + search_item_types=None, + categories=None, + metadata_filters=None, + limit=1, + offset=offset, + ) + + first_tail_page = await deep_page(offset=2) + second_tail_page = await deep_page(offset=3) + + assert [row.id for row in first_tail_page] == [charlie.id] + assert [row.id for row in second_tail_page] == [delta.id] + assert reranker.calls == 2 + + +@pytest.mark.asyncio +async def test_hybrid_search_surfaces_transient_reranker_error(search_repository): + """Hybrid search must not replace reranked order with raw order during an outage.""" + if not isinstance(search_repository, SQLiteSearchRepository): + pytest.skip("sqlite-vec repository behavior is local SQLite-only.") + + _enable_semantic(search_repository) + await _index_two_auth_notes(search_repository) + search_repository._rerank_provider = _ExplodingReranker() + + with pytest.raises(RerankTransientError, match="backend unreachable"): + await search_repository.search( + search_text="auth session token", + retrieval_mode=SearchRetrievalMode.HYBRID, + limit=5, + ) + + +@pytest.mark.asyncio +async def test_hybrid_search_propagates_contract_error(search_repository): + """A provider-contract break (e.g. incomplete rerank response) must surface, not hide.""" + if not isinstance(search_repository, SQLiteSearchRepository): + pytest.skip("sqlite-vec repository behavior is local SQLite-only.") + + _enable_semantic(search_repository) + await _index_two_auth_notes(search_repository) + search_repository._rerank_provider = _PermanentFaultReranker( + RerankProviderContractError("incomplete rerank response") + ) + + with pytest.raises(RerankProviderContractError): + await search_repository.search( + search_text="auth session token", + retrieval_mode=SearchRetrievalMode.HYBRID, + limit=5, + ) + + +@pytest.mark.asyncio +async def test_search_without_reranker_keeps_baseline(search_repository): + if not isinstance(search_repository, SQLiteSearchRepository): + pytest.skip("sqlite-vec repository behavior is local SQLite-only.") + + _enable_semantic(search_repository) + await _index_two_auth_notes(search_repository) + search_repository._rerank_provider = None + + results = await search_repository.search( + search_text="auth session token", + retrieval_mode=SearchRetrievalMode.VECTOR, + limit=5, + ) + # Vector similarity ranks the plain-auth note above the "deep"-tilted one. + assert [r.permalink for r in results] == ["specs/alpha", "specs/bravo"] diff --git a/tests/repository/test_rerank_provider_factory.py b/tests/repository/test_rerank_provider_factory.py new file mode 100644 index 000000000..705e5be41 --- /dev/null +++ b/tests/repository/test_rerank_provider_factory.py @@ -0,0 +1,233 @@ +"""Tests for the reranker provider factory.""" + +import builtins + +import pytest +from pydantic import ValidationError + +from basic_memory.config import BasicMemoryConfig, DatabaseBackend +from basic_memory.repository.fastembed_rerank_provider import FastEmbedRerankProvider +from basic_memory.repository.litellm_rerank_provider import LiteLLMRerankProvider +import basic_memory.repository.rerank_provider_factory as factory +from basic_memory.repository.rerank_provider_factory import ( + create_rerank_provider, + reset_rerank_provider_cache, +) + + +def _config(**overrides) -> BasicMemoryConfig: + base = dict( + env="test", + projects={"test-project": "/tmp/test"}, + default_project="test-project", + database_backend=DatabaseBackend.SQLITE, + ) + base.update(overrides) + return BasicMemoryConfig(**base) + + +@pytest.fixture(autouse=True) +def _clear_cache(): + reset_rerank_provider_cache() + yield + reset_rerank_provider_cache() + + +def test_disabled_returns_none(): + """Reranking is off by default; the factory yields no provider.""" + assert create_rerank_provider(_config(reranker_enabled=False)) is None + + +def test_fastembed_provider_selected_with_resolved_cache_dir(): + config = _config( + reranker_enabled=True, + reranker_provider="fastembed", + reranker_model="Xenova/ms-marco-MiniLM-L-6-v2", + semantic_embedding_cache_dir="/tmp/fastembed-cache", + ) + provider = create_rerank_provider(config) + assert isinstance(provider, FastEmbedRerankProvider) + assert provider.model_name == "Xenova/ms-marco-MiniLM-L-6-v2" + # Reranker shares the embedding provider's resolved cache dir. + assert provider.cache_dir == "/tmp/fastembed-cache" + + +def test_fastembed_provider_rejects_unsupported_model_at_startup(): + """An enabled FastEmbed reranker must fail during config loading.""" + with pytest.raises(ValidationError, match="Unsupported FastEmbed reranker model"): + _config( + reranker_enabled=True, + reranker_provider="fastembed", + reranker_model="unsupported/model-typo", + ) + + +def test_fastembed_provider_rejects_missing_dependency_at_startup(monkeypatch): + """An enabled local reranker must report a missing FastEmbed install during config load.""" + real_import = builtins.__import__ + + def import_without_fastembed(name, *args, **kwargs): + if name == "fastembed.rerank.cross_encoder": + raise ImportError("fastembed is unavailable") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", import_without_fastembed) + + with pytest.raises(ValidationError, match="requires the fastembed package"): + _config( + reranker_enabled=True, + reranker_provider="fastembed", + ) + + +def test_litellm_provider_selected_with_routing(): + config = _config( + reranker_enabled=True, + reranker_provider="litellm", + reranker_model="cohere/rerank-v3.5", + reranker_api_key="secret", + reranker_api_base="https://rerank.example", + ) + provider = create_rerank_provider(config) + assert isinstance(provider, LiteLLMRerankProvider) + assert provider.model_name == "cohere/rerank-v3.5" + assert provider._api_key == "secret" + assert provider._api_base == "https://rerank.example" + + +def test_unsupported_provider_raises(): + """Unknown provider names fail at the configuration boundary.""" + with pytest.raises(ValidationError, match="reranker_provider must be one of"): + _config(reranker_enabled=True, reranker_provider="nope") + + +def test_same_config_returns_cached_singleton(): + config = _config(reranker_enabled=True, reranker_provider="fastembed") + first = create_rerank_provider(config) + second = create_rerank_provider(config) + assert first is second + + +def test_distinct_config_creates_second_provider(): + """A second distinct key builds a new provider (and warns about the reload).""" + first = create_rerank_provider( + _config( + reranker_enabled=True, + reranker_model="Xenova/ms-marco-MiniLM-L-6-v2", + ) + ) + second = create_rerank_provider( + _config( + reranker_enabled=True, + reranker_model="Xenova/ms-marco-MiniLM-L-12-v2", + ) + ) + assert first is not second + + +def test_distinct_cache_dir_does_not_collide(): + """Two configs differing only in cache dir must not share one singleton (#741/#872).""" + a = create_rerank_provider( + _config(reranker_enabled=True, semantic_embedding_cache_dir="/tmp/rr-a") + ) + b = create_rerank_provider( + _config(reranker_enabled=True, semantic_embedding_cache_dir="/tmp/rr-b") + ) + assert a is not b + assert isinstance(a, FastEmbedRerankProvider) and isinstance(b, FastEmbedRerankProvider) + assert a.cache_dir == "/tmp/rr-a" + assert b.cache_dir == "/tmp/rr-b" + + +def test_reranker_enabled_requires_semantic_search(): + """Config rejects reranking without semantic search rather than silently no-op'ing.""" + with pytest.raises(ValidationError, match="requires semantic_search_enabled"): + _config(reranker_enabled=True, semantic_search_enabled=False) + + +def test_litellm_provider_rejects_default_fastembed_model(): + """Selecting litellm without overriding the model is a footgun; reject it at config.""" + with pytest.raises(ValidationError, match="requires an explicit reranker_model"): + _config(reranker_enabled=True, reranker_provider="litellm") + + +def test_litellm_provider_accepts_explicit_model(): + """An explicit provider/model model id passes and reaches the provider.""" + provider = create_rerank_provider( + _config( + reranker_enabled=True, + reranker_provider="litellm", + reranker_model="cohere/rerank-v3.5", + ) + ) + assert isinstance(provider, LiteLLMRerankProvider) + assert provider.model_name == "cohere/rerank-v3.5" + + +def test_litellm_provider_normalizes_model_whitespace(): + """Accepted provider/model identifiers are stored in their routable form.""" + config = _config( + reranker_enabled=True, + reranker_provider=" LiteLLM ", + reranker_model=" cohere/rerank-v3.5 ", + ) + + provider = create_rerank_provider(config) + + assert config.reranker_provider == "litellm" + assert config.reranker_model == "cohere/rerank-v3.5" + assert isinstance(provider, LiteLLMRerankProvider) + assert provider.model_name == "cohere/rerank-v3.5" + + +def test_litellm_provider_rejects_unroutable_model_name(): + """LiteLLM model names include the provider routing prefix.""" + with pytest.raises(ValidationError, match="requires an explicit reranker_model"): + _config( + reranker_enabled=True, + reranker_provider="litellm", + reranker_model="rerank-v3.5", + ) + + +@pytest.mark.parametrize("provider", ["fastembed", "litellm"]) +def test_enabled_provider_rejects_blank_model(provider): + """Every enabled provider needs a model before its first search.""" + with pytest.raises(ValidationError, match="reranker_model must not be blank"): + _config( + reranker_enabled=True, + reranker_provider=provider, + reranker_model=" \t ", + ) + + +def test_reset_clears_cache(): + config = _config(reranker_enabled=True, reranker_provider="fastembed") + first = create_rerank_provider(config) + reset_rerank_provider_cache() + second = create_rerank_provider(config) + assert first is not second + + +def test_concurrent_race_returns_winning_provider(monkeypatch): + """The double-checked lock returns the racing writer's provider, not ours. + + Simulates a concurrent caller that populated the cache after our first miss + but before we take the write lock: the second in-lock check must win. + """ + winner = object() + + class _RacyCache(dict): + def __init__(self): + super().__init__() + self._gets = 0 + + def get(self, key, default=None): + # First check (outside lock) misses so we build; second check (in lock) + # finds the winner another thread inserted mid-flight. + self._gets += 1 + return winner if self._gets > 1 else None + + monkeypatch.setattr(factory, "_RERANK_PROVIDER_CACHE", _RacyCache()) + result = create_rerank_provider(_config(reranker_enabled=True, reranker_provider="fastembed")) + assert result is winner