fix(core): split LiteLLM query and document embeddings

Signed-off-by: phernandez <paul@basicmachines.co>
This commit is contained in:
phernandez
2026-05-28 17:59:46 -05:00
parent 187ca1a160
commit 2c8975e9bb
7 changed files with 314 additions and 9 deletions
+32 -5
View File
@@ -99,10 +99,12 @@ All settings are fields on `BasicMemoryConfig` and can be set via environment va
| Config Field | Env Var | Default | Description |
|---|---|---|---|
| `semantic_search_enabled` | `BASIC_MEMORY_SEMANTIC_SEARCH_ENABLED` | Auto (`true` when semantic deps are available) | Enable semantic search. Required before vector/hybrid modes work. |
| `semantic_embedding_provider` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_PROVIDER` | `"fastembed"` | Embedding provider: `"fastembed"` (local) or `"openai"` (API). |
| `semantic_embedding_provider` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_PROVIDER` | `"fastembed"` | Embedding provider: `"fastembed"` (local), `"openai"` (API), or `"litellm"` (multi-provider API). |
| `semantic_embedding_model` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_MODEL` | `"bge-small-en-v1.5"` | Model identifier. Auto-adjusted per provider if left at default. |
| `semantic_embedding_dimensions` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_DIMENSIONS` | Auto-detected | Vector dimensions. 384 for FastEmbed, 1536 for OpenAI. Override only if using a non-default model. |
| `semantic_embedding_batch_size` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_BATCH_SIZE` | `64` | Number of texts to embed per batch. |
| `semantic_embedding_dimensions` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_DIMENSIONS` | Auto-detected | Vector dimensions. 384 for FastEmbed, 1536 for OpenAI/LiteLLM OpenAI. Override when using a non-default model. |
| `semantic_embedding_batch_size` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_BATCH_SIZE` | `2` | Number of texts to embed per batch. |
| `semantic_embedding_document_input_type` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_DOCUMENT_INPUT_TYPE` | Auto for known LiteLLM models | Optional LiteLLM `input_type` for indexed document/passages. |
| `semantic_embedding_query_input_type` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_QUERY_INPUT_TYPE` | Auto for known LiteLLM models | Optional LiteLLM `input_type` for search queries. |
| `semantic_vector_k` | `BASIC_MEMORY_SEMANTIC_VECTOR_K` | `100` | Candidate count for vector nearest-neighbour retrieval. Higher values improve recall at the cost of latency. |
## Embedding Providers
@@ -135,7 +137,31 @@ export BASIC_MEMORY_SEMANTIC_EMBEDDING_PROVIDER=openai
export OPENAI_API_KEY=sk-...
```
When switching from FastEmbed to OpenAI (or vice versa), you must rebuild embeddings since the vector dimensions differ:
### LiteLLM
Uses the LiteLLM SDK to call embedding models from providers such as OpenAI, Cohere, Azure, Bedrock, NVIDIA NIM, and other LiteLLM-supported backends. Requires the provider's API credentials.
```bash
export BASIC_MEMORY_SEMANTIC_SEARCH_ENABLED=true
export BASIC_MEMORY_SEMANTIC_EMBEDDING_PROVIDER=litellm
export BASIC_MEMORY_SEMANTIC_EMBEDDING_MODEL=cohere/embed-english-v3.0
export BASIC_MEMORY_SEMANTIC_EMBEDDING_DIMENSIONS=1024
export COHERE_API_KEY=...
```
Some retrieval models are asymmetric: indexed passages and search queries must be embedded with different provider parameters. Basic Memory automatically sets LiteLLM `input_type` for known asymmetric model families:
- Cohere v3: documents use `search_document`, queries use `search_query`
- NVIDIA NIM retrieval models: documents use `passage`, queries use `query`
For other asymmetric LiteLLM models, set the input types explicitly:
```bash
export BASIC_MEMORY_SEMANTIC_EMBEDDING_DOCUMENT_INPUT_TYPE=passage
export BASIC_MEMORY_SEMANTIC_EMBEDDING_QUERY_INPUT_TYPE=query
```
When switching providers, models, dimensions, or LiteLLM document/query input types, rebuild embeddings:
```bash
bm reindex --embeddings
@@ -203,9 +229,10 @@ bm reindex -p my-project
- **Upgrade note**: Migration now performs a one-time automatic embedding backfill on upgrade.
- **Manual enable case**: If you explicitly had `semantic_search_enabled=false` and then turn it on
- **Provider change**: After switching between `fastembed` and `openai`
- **Provider change**: After switching between `fastembed`, `openai`, and `litellm`
- **Model change**: After changing `semantic_embedding_model`
- **Dimension change**: After changing `semantic_embedding_dimensions`
- **LiteLLM role change**: After changing `semantic_embedding_document_input_type` or `semantic_embedding_query_input_type`
The reindex command shows progress with embedded/skipped/error counts:
+1
View File
@@ -85,6 +85,7 @@ markers = [
"windows: Windows-specific tests (deselect with '-m \"not windows\"')",
"smoke: Fast end-to-end smoke tests for MCP flows",
"semantic: Tests requiring semantic dependencies (fastembed, sqlite-vec, openai)",
"live: Tests that call external provider APIs and require explicit opt-in",
]
[tool.ruff]
+14
View File
@@ -263,6 +263,20 @@ class BasicMemoryConfig(BaseSettings):
description="Maximum number of concurrent provider requests for batched embedding generation when the active provider supports request-level concurrency.",
gt=0,
)
semantic_embedding_document_input_type: str | None = Field(
default=None,
description=(
"Optional LiteLLM input_type for indexed document/passages. "
"Use with asymmetric embedding models such as Cohere or NVIDIA retrieval models."
),
)
semantic_embedding_query_input_type: str | None = Field(
default=None,
description=(
"Optional LiteLLM input_type for search queries. "
"Use with asymmetric embedding models such as Cohere or NVIDIA retrieval models."
),
)
semantic_embedding_sync_batch_size: int = Field(
default=2,
description="Batch size for vector sync orchestration flushes.",
@@ -12,6 +12,8 @@ type ProviderCacheKey = tuple[
int | None,
int,
int,
str | None,
str | None,
str,
int | None,
int | None,
@@ -88,6 +90,8 @@ def _provider_cache_key(app_config: BasicMemoryConfig) -> ProviderCacheKey:
app_config.semantic_embedding_dimensions,
app_config.semantic_embedding_batch_size,
app_config.semantic_embedding_request_concurrency,
app_config.semantic_embedding_document_input_type,
app_config.semantic_embedding_query_input_type,
_resolve_cache_dir(app_config),
resolved_threads,
resolved_parallel,
@@ -161,6 +165,8 @@ def create_embedding_provider(app_config: BasicMemoryConfig) -> EmbeddingProvide
model_name=model_name,
batch_size=app_config.semantic_embedding_batch_size,
request_concurrency=app_config.semantic_embedding_request_concurrency,
document_input_type=app_config.semantic_embedding_document_input_type,
query_input_type=app_config.semantic_embedding_query_input_type,
**extra_kwargs,
)
else:
@@ -21,6 +21,30 @@ from basic_memory.repository.embedding_provider import EmbeddingProvider
from basic_memory.repository.semantic_errors import SemanticDependenciesMissingError
def _default_input_types(model_name: str) -> tuple[str | None, str | None]:
"""Return role-specific LiteLLM input_type defaults for known asymmetric models."""
normalized = model_name.strip().lower()
# Cohere v3 embeddings require search_document/search_query to distinguish
# index-time passages from retrieval-time queries. LiteLLM supports both
# direct Cohere model names and provider-prefixed forms.
cohere_v3 = (
normalized.startswith("cohere/")
or normalized.startswith("bedrock/cohere.")
or normalized.startswith("cohere.")
or normalized.startswith("embed-")
) and "-v3" in normalized
if cohere_v3:
return "search_document", "search_query"
# NVIDIA retrieval embeddings use passage/query roles. The provider prefix
# is part of LiteLLM's model routing, so this stays narrowly scoped.
if normalized.startswith("nvidia_nim/"):
return "passage", "query"
return None, None
class LiteLLMEmbeddingProvider(EmbeddingProvider):
"""Embedding provider backed by the litellm SDK."""
@@ -33,6 +57,8 @@ class LiteLLMEmbeddingProvider(EmbeddingProvider):
dimensions: int = 1536,
api_key: str | None = None,
timeout: float = 30.0,
document_input_type: str | None = None,
query_input_type: str | None = None,
) -> None:
self.model_name = model_name
self.dimensions = dimensions
@@ -40,15 +66,23 @@ class LiteLLMEmbeddingProvider(EmbeddingProvider):
self.request_concurrency = request_concurrency
self._api_key = api_key
self._timeout = timeout
default_document_input_type, default_query_input_type = _default_input_types(model_name)
self.document_input_type = document_input_type or default_document_input_type
self.query_input_type = query_input_type or default_query_input_type
def runtime_log_attrs(self) -> dict[str, int]:
def runtime_log_attrs(self) -> dict[str, Any]:
"""Return provider-specific runtime settings suitable for startup logs."""
return {
attrs: dict[str, Any] = {
"provider_batch_size": self.batch_size,
"request_concurrency": self.request_concurrency,
}
if self.document_input_type:
attrs["document_input_type"] = self.document_input_type
if self.query_input_type:
attrs["query_input_type"] = self.query_input_type
return attrs
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
async def _embed(self, texts: list[str], *, input_type: str | None) -> list[list[float]]:
if not texts:
return []
@@ -76,6 +110,8 @@ class LiteLLMEmbeddingProvider(EmbeddingProvider):
}
if self._api_key:
params["api_key"] = self._api_key
if input_type:
params["input_type"] = input_type
response = await litellm.aembedding(**params)
@@ -129,6 +165,9 @@ class LiteLLMEmbeddingProvider(EmbeddingProvider):
)
return normalized
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
return await self._embed(texts, input_type=self.document_input_type)
async def embed_query(self, text: str) -> list[float]:
vectors = await self.embed_documents([text])
vectors = await self._embed([text], input_type=self.query_input_type)
return vectors[0] if vectors else [0.0] * self.dimensions
@@ -0,0 +1,163 @@
"""Opt-in live LiteLLM provider checks against real embedding APIs.
These tests intentionally do not run in normal CI. Enable them with
``BASIC_MEMORY_RUN_LITELLM_INTEGRATION=1`` and provider API keys when validating
new LiteLLM model support before merging or releasing.
"""
from __future__ import annotations
import json
import math
import os
from dataclasses import dataclass
from typing import Any
import pytest
from basic_memory.repository.litellm_provider import LiteLLMEmbeddingProvider
pytestmark = [
pytest.mark.semantic,
pytest.mark.slow,
pytest.mark.live,
pytest.mark.skipif(
os.getenv("BASIC_MEMORY_RUN_LITELLM_INTEGRATION") != "1",
reason="Set BASIC_MEMORY_RUN_LITELLM_INTEGRATION=1 to run live LiteLLM tests",
),
]
@dataclass(frozen=True)
class LiteLLMLiveCase:
"""A real LiteLLM embedding model to exercise end-to-end."""
name: str
model: str
dimensions: int
api_key_env: str | None = None
document_input_type: str | None = None
query_input_type: str | None = None
def _custom_cases() -> list[LiteLLMLiveCase]:
"""Load additional live model cases from BASIC_MEMORY_TEST_LITELLM_CASES."""
raw = os.getenv("BASIC_MEMORY_TEST_LITELLM_CASES")
if not raw:
return []
values = json.loads(raw)
if not isinstance(values, list):
raise ValueError("BASIC_MEMORY_TEST_LITELLM_CASES must be a JSON array")
cases: list[LiteLLMLiveCase] = []
for value in values:
if not isinstance(value, dict):
raise ValueError("Each LiteLLM live case must be a JSON object")
case_data: dict[str, Any] = value
cases.append(
LiteLLMLiveCase(
name=str(case_data["name"]),
model=str(case_data["model"]),
dimensions=int(case_data["dimensions"]),
api_key_env=case_data.get("api_key_env"),
document_input_type=case_data.get("document_input_type"),
query_input_type=case_data.get("query_input_type"),
)
)
return cases
def _live_cases() -> list[LiteLLMLiveCase | Any]:
"""Return built-in and user-supplied live cases whose credentials are available."""
cases: list[LiteLLMLiveCase] = []
if os.getenv("OPENAI_API_KEY"):
cases.append(
LiteLLMLiveCase(
name="openai-text-embedding-3-small",
model="openai/text-embedding-3-small",
dimensions=1536,
api_key_env="OPENAI_API_KEY",
)
)
if os.getenv("COHERE_API_KEY"):
cases.append(
LiteLLMLiveCase(
name="cohere-embed-english-v3",
model="cohere/embed-english-v3.0",
dimensions=1024,
api_key_env="COHERE_API_KEY",
)
)
cases.extend(_custom_cases())
if cases:
return cases
return [
pytest.param(
None,
marks=pytest.mark.skip(
reason=(
"No LiteLLM live cases configured. Set OPENAI_API_KEY, "
"COHERE_API_KEY, or BASIC_MEMORY_TEST_LITELLM_CASES."
)
),
)
]
def _cosine(a: list[float], b: list[float]) -> float:
"""Compute cosine similarity for live ranking sanity checks."""
dot = sum(x * y for x, y in zip(a, b, strict=True))
norm_a = math.sqrt(sum(x * x for x in a))
norm_b = math.sqrt(sum(y * y for y in b))
if norm_a == 0 or norm_b == 0:
return 0.0
return dot / (norm_a * norm_b)
def _assert_valid_vector(vector: list[float], dimensions: int) -> None:
"""Assert provider output is a usable normalized vector."""
assert len(vector) == dimensions
assert all(math.isfinite(value) for value in vector)
norm = math.sqrt(sum(value * value for value in vector))
assert norm == pytest.approx(1.0, abs=1e-6)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"case",
_live_cases(),
ids=lambda case: case.name if isinstance(case, LiteLLMLiveCase) else "no-live-cases",
)
async def test_litellm_live_model_embeds_documents_and_queries(
case: LiteLLMLiveCase,
) -> None:
"""A live LiteLLM model should embed documents and rank a related query higher."""
api_key = os.getenv(case.api_key_env) if case.api_key_env else None
provider = LiteLLMEmbeddingProvider(
model_name=case.model,
dimensions=case.dimensions,
batch_size=2,
api_key=api_key,
timeout=60.0,
document_input_type=case.document_input_type,
query_input_type=case.query_input_type,
)
documents = [
"OAuth login refresh tokens keep an authenticated web session active.",
"A sourdough starter ferments flour and water before bread baking.",
]
vectors = await provider.embed_documents(documents)
query_vector = await provider.embed_query("authentication login token flow")
assert len(vectors) == 2
for vector in [*vectors, query_vector]:
_assert_valid_vector(vector, case.dimensions)
assert _cosine(query_vector, vectors[0]) > _cosine(query_vector, vectors[1])
+55
View File
@@ -130,6 +130,42 @@ async def test_litellm_provider_drop_params_always_set(monkeypatch):
assert calls[0]["drop_params"] is True
@pytest.mark.asyncio
async def test_litellm_provider_uses_cohere_document_and_query_input_types(monkeypatch):
"""Cohere v3 embeddings require different input_type values per embedding role."""
calls = _install_litellm_stub(monkeypatch)
provider = LiteLLMEmbeddingProvider(
model_name="cohere/embed-english-v3.0",
batch_size=2,
dimensions=3,
)
await provider.embed_documents(["indexed passage"])
await provider.embed_query("retrieval query")
assert calls[0]["input_type"] == "search_document"
assert calls[1]["input_type"] == "search_query"
@pytest.mark.asyncio
async def test_litellm_provider_uses_explicit_document_and_query_input_types(monkeypatch):
"""Explicit input_type overrides should support asymmetric providers beyond Cohere."""
calls = _install_litellm_stub(monkeypatch)
provider = LiteLLMEmbeddingProvider(
model_name="nvidia_nim/nvidia/embed-qa-4",
batch_size=2,
dimensions=3,
document_input_type="passage",
query_input_type="query",
)
await provider.embed_documents(["indexed passage"])
await provider.embed_query("retrieval query")
assert calls[0]["input_type"] == "passage"
assert calls[1]["input_type"] == "query"
@pytest.mark.asyncio
async def test_litellm_provider_dimension_mismatch_raises_error(monkeypatch):
"""Provider should fail fast when response dimensions differ from configured."""
@@ -211,6 +247,25 @@ def test_factory_maps_default_model_for_litellm():
assert provider.model_name == "openai/text-embedding-3-small"
def test_factory_forwards_litellm_document_and_query_input_types():
"""Factory should pass role-specific LiteLLM input_type config to the provider."""
config = BasicMemoryConfig(
env="test",
projects={"test": "/tmp/basic-memory-test"},
default_project="test",
semantic_search_enabled=True,
semantic_embedding_provider="litellm",
semantic_embedding_model="nvidia_nim/nvidia/embed-qa-4",
semantic_embedding_document_input_type="passage",
semantic_embedding_query_input_type="query",
)
provider = create_embedding_provider(config)
assert isinstance(provider, LiteLLMEmbeddingProvider)
assert provider.document_input_type == "passage"
assert provider.query_input_type == "query"
def test_runtime_log_attrs():
"""runtime_log_attrs should return batch_size and concurrency."""
provider = LiteLLMEmbeddingProvider(batch_size=32, request_concurrency=8)