mirror of
https://github.com/basicmachines-co/basic-memory
synced 2026-06-21 13:47:35 +00:00
fix(core): split LiteLLM query and document embeddings
Signed-off-by: phernandez <paul@basicmachines.co>
This commit is contained in:
+32
-5
@@ -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:
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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])
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user