Skip to content

Commit 2c8975e

Browse files
committed
fix(core): split LiteLLM query and document embeddings
Signed-off-by: phernandez <paul@basicmachines.co>
1 parent 187ca1a commit 2c8975e

7 files changed

Lines changed: 314 additions & 9 deletions

File tree

docs/semantic-search.md

Lines changed: 32 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -99,10 +99,12 @@ All settings are fields on `BasicMemoryConfig` and can be set via environment va
9999
| Config Field | Env Var | Default | Description |
100100
|---|---|---|---|
101101
| `semantic_search_enabled` | `BASIC_MEMORY_SEMANTIC_SEARCH_ENABLED` | Auto (`true` when semantic deps are available) | Enable semantic search. Required before vector/hybrid modes work. |
102-
| `semantic_embedding_provider` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_PROVIDER` | `"fastembed"` | Embedding provider: `"fastembed"` (local) or `"openai"` (API). |
102+
| `semantic_embedding_provider` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_PROVIDER` | `"fastembed"` | Embedding provider: `"fastembed"` (local), `"openai"` (API), or `"litellm"` (multi-provider API). |
103103
| `semantic_embedding_model` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_MODEL` | `"bge-small-en-v1.5"` | Model identifier. Auto-adjusted per provider if left at default. |
104-
| `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. |
105-
| `semantic_embedding_batch_size` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_BATCH_SIZE` | `64` | Number of texts to embed per batch. |
104+
| `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. |
105+
| `semantic_embedding_batch_size` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_BATCH_SIZE` | `2` | Number of texts to embed per batch. |
106+
| `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. |
107+
| `semantic_embedding_query_input_type` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_QUERY_INPUT_TYPE` | Auto for known LiteLLM models | Optional LiteLLM `input_type` for search queries. |
106108
| `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. |
107109

108110
## Embedding Providers
@@ -135,7 +137,31 @@ export BASIC_MEMORY_SEMANTIC_EMBEDDING_PROVIDER=openai
135137
export OPENAI_API_KEY=sk-...
136138
```
137139

138-
When switching from FastEmbed to OpenAI (or vice versa), you must rebuild embeddings since the vector dimensions differ:
140+
### LiteLLM
141+
142+
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.
143+
144+
```bash
145+
export BASIC_MEMORY_SEMANTIC_SEARCH_ENABLED=true
146+
export BASIC_MEMORY_SEMANTIC_EMBEDDING_PROVIDER=litellm
147+
export BASIC_MEMORY_SEMANTIC_EMBEDDING_MODEL=cohere/embed-english-v3.0
148+
export BASIC_MEMORY_SEMANTIC_EMBEDDING_DIMENSIONS=1024
149+
export COHERE_API_KEY=...
150+
```
151+
152+
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:
153+
154+
- Cohere v3: documents use `search_document`, queries use `search_query`
155+
- NVIDIA NIM retrieval models: documents use `passage`, queries use `query`
156+
157+
For other asymmetric LiteLLM models, set the input types explicitly:
158+
159+
```bash
160+
export BASIC_MEMORY_SEMANTIC_EMBEDDING_DOCUMENT_INPUT_TYPE=passage
161+
export BASIC_MEMORY_SEMANTIC_EMBEDDING_QUERY_INPUT_TYPE=query
162+
```
163+
164+
When switching providers, models, dimensions, or LiteLLM document/query input types, rebuild embeddings:
139165

140166
```bash
141167
bm reindex --embeddings
@@ -203,9 +229,10 @@ bm reindex -p my-project
203229

204230
- **Upgrade note**: Migration now performs a one-time automatic embedding backfill on upgrade.
205231
- **Manual enable case**: If you explicitly had `semantic_search_enabled=false` and then turn it on
206-
- **Provider change**: After switching between `fastembed` and `openai`
232+
- **Provider change**: After switching between `fastembed`, `openai`, and `litellm`
207233
- **Model change**: After changing `semantic_embedding_model`
208234
- **Dimension change**: After changing `semantic_embedding_dimensions`
235+
- **LiteLLM role change**: After changing `semantic_embedding_document_input_type` or `semantic_embedding_query_input_type`
209236

210237
The reindex command shows progress with embedded/skipped/error counts:
211238

pyproject.toml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -85,6 +85,7 @@ markers = [
8585
"windows: Windows-specific tests (deselect with '-m \"not windows\"')",
8686
"smoke: Fast end-to-end smoke tests for MCP flows",
8787
"semantic: Tests requiring semantic dependencies (fastembed, sqlite-vec, openai)",
88+
"live: Tests that call external provider APIs and require explicit opt-in",
8889
]
8990

9091
[tool.ruff]

src/basic_memory/config.py

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -263,6 +263,20 @@ def __init__(self, **data: Any) -> None: ...
263263
description="Maximum number of concurrent provider requests for batched embedding generation when the active provider supports request-level concurrency.",
264264
gt=0,
265265
)
266+
semantic_embedding_document_input_type: str | None = Field(
267+
default=None,
268+
description=(
269+
"Optional LiteLLM input_type for indexed document/passages. "
270+
"Use with asymmetric embedding models such as Cohere or NVIDIA retrieval models."
271+
),
272+
)
273+
semantic_embedding_query_input_type: str | None = Field(
274+
default=None,
275+
description=(
276+
"Optional LiteLLM input_type for search queries. "
277+
"Use with asymmetric embedding models such as Cohere or NVIDIA retrieval models."
278+
),
279+
)
266280
semantic_embedding_sync_batch_size: int = Field(
267281
default=2,
268282
description="Batch size for vector sync orchestration flushes.",

src/basic_memory/repository/embedding_provider_factory.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,8 @@
1212
int | None,
1313
int,
1414
int,
15+
str | None,
16+
str | None,
1517
str,
1618
int | None,
1719
int | None,
@@ -88,6 +90,8 @@ def _provider_cache_key(app_config: BasicMemoryConfig) -> ProviderCacheKey:
8890
app_config.semantic_embedding_dimensions,
8991
app_config.semantic_embedding_batch_size,
9092
app_config.semantic_embedding_request_concurrency,
93+
app_config.semantic_embedding_document_input_type,
94+
app_config.semantic_embedding_query_input_type,
9195
_resolve_cache_dir(app_config),
9296
resolved_threads,
9397
resolved_parallel,
@@ -161,6 +165,8 @@ def create_embedding_provider(app_config: BasicMemoryConfig) -> EmbeddingProvide
161165
model_name=model_name,
162166
batch_size=app_config.semantic_embedding_batch_size,
163167
request_concurrency=app_config.semantic_embedding_request_concurrency,
168+
document_input_type=app_config.semantic_embedding_document_input_type,
169+
query_input_type=app_config.semantic_embedding_query_input_type,
164170
**extra_kwargs,
165171
)
166172
else:

src/basic_memory/repository/litellm_provider.py

Lines changed: 43 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,30 @@
2121
from basic_memory.repository.semantic_errors import SemanticDependenciesMissingError
2222

2323

24+
def _default_input_types(model_name: str) -> tuple[str | None, str | None]:
25+
"""Return role-specific LiteLLM input_type defaults for known asymmetric models."""
26+
normalized = model_name.strip().lower()
27+
28+
# Cohere v3 embeddings require search_document/search_query to distinguish
29+
# index-time passages from retrieval-time queries. LiteLLM supports both
30+
# direct Cohere model names and provider-prefixed forms.
31+
cohere_v3 = (
32+
normalized.startswith("cohere/")
33+
or normalized.startswith("bedrock/cohere.")
34+
or normalized.startswith("cohere.")
35+
or normalized.startswith("embed-")
36+
) and "-v3" in normalized
37+
if cohere_v3:
38+
return "search_document", "search_query"
39+
40+
# NVIDIA retrieval embeddings use passage/query roles. The provider prefix
41+
# is part of LiteLLM's model routing, so this stays narrowly scoped.
42+
if normalized.startswith("nvidia_nim/"):
43+
return "passage", "query"
44+
45+
return None, None
46+
47+
2448
class LiteLLMEmbeddingProvider(EmbeddingProvider):
2549
"""Embedding provider backed by the litellm SDK."""
2650

@@ -33,22 +57,32 @@ def __init__(
3357
dimensions: int = 1536,
3458
api_key: str | None = None,
3559
timeout: float = 30.0,
60+
document_input_type: str | None = None,
61+
query_input_type: str | None = None,
3662
) -> None:
3763
self.model_name = model_name
3864
self.dimensions = dimensions
3965
self.batch_size = batch_size
4066
self.request_concurrency = request_concurrency
4167
self._api_key = api_key
4268
self._timeout = timeout
69+
default_document_input_type, default_query_input_type = _default_input_types(model_name)
70+
self.document_input_type = document_input_type or default_document_input_type
71+
self.query_input_type = query_input_type or default_query_input_type
4372

44-
def runtime_log_attrs(self) -> dict[str, int]:
73+
def runtime_log_attrs(self) -> dict[str, Any]:
4574
"""Return provider-specific runtime settings suitable for startup logs."""
46-
return {
75+
attrs: dict[str, Any] = {
4776
"provider_batch_size": self.batch_size,
4877
"request_concurrency": self.request_concurrency,
4978
}
79+
if self.document_input_type:
80+
attrs["document_input_type"] = self.document_input_type
81+
if self.query_input_type:
82+
attrs["query_input_type"] = self.query_input_type
83+
return attrs
5084

51-
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
85+
async def _embed(self, texts: list[str], *, input_type: str | None) -> list[list[float]]:
5286
if not texts:
5387
return []
5488

@@ -76,6 +110,8 @@ async def embed_batch(batch_index: int, batch: list[str]) -> None:
76110
}
77111
if self._api_key:
78112
params["api_key"] = self._api_key
113+
if input_type:
114+
params["input_type"] = input_type
79115

80116
response = await litellm.aembedding(**params)
81117

@@ -129,6 +165,9 @@ async def embed_batch(batch_index: int, batch: list[str]) -> None:
129165
)
130166
return normalized
131167

168+
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
169+
return await self._embed(texts, input_type=self.document_input_type)
170+
132171
async def embed_query(self, text: str) -> list[float]:
133-
vectors = await self.embed_documents([text])
172+
vectors = await self._embed([text], input_type=self.query_input_type)
134173
return vectors[0] if vectors else [0.0] * self.dimensions
Lines changed: 163 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,163 @@
1+
"""Opt-in live LiteLLM provider checks against real embedding APIs.
2+
3+
These tests intentionally do not run in normal CI. Enable them with
4+
``BASIC_MEMORY_RUN_LITELLM_INTEGRATION=1`` and provider API keys when validating
5+
new LiteLLM model support before merging or releasing.
6+
"""
7+
8+
from __future__ import annotations
9+
10+
import json
11+
import math
12+
import os
13+
from dataclasses import dataclass
14+
from typing import Any
15+
16+
import pytest
17+
18+
from basic_memory.repository.litellm_provider import LiteLLMEmbeddingProvider
19+
20+
21+
pytestmark = [
22+
pytest.mark.semantic,
23+
pytest.mark.slow,
24+
pytest.mark.live,
25+
pytest.mark.skipif(
26+
os.getenv("BASIC_MEMORY_RUN_LITELLM_INTEGRATION") != "1",
27+
reason="Set BASIC_MEMORY_RUN_LITELLM_INTEGRATION=1 to run live LiteLLM tests",
28+
),
29+
]
30+
31+
32+
@dataclass(frozen=True)
33+
class LiteLLMLiveCase:
34+
"""A real LiteLLM embedding model to exercise end-to-end."""
35+
36+
name: str
37+
model: str
38+
dimensions: int
39+
api_key_env: str | None = None
40+
document_input_type: str | None = None
41+
query_input_type: str | None = None
42+
43+
44+
def _custom_cases() -> list[LiteLLMLiveCase]:
45+
"""Load additional live model cases from BASIC_MEMORY_TEST_LITELLM_CASES."""
46+
raw = os.getenv("BASIC_MEMORY_TEST_LITELLM_CASES")
47+
if not raw:
48+
return []
49+
50+
values = json.loads(raw)
51+
if not isinstance(values, list):
52+
raise ValueError("BASIC_MEMORY_TEST_LITELLM_CASES must be a JSON array")
53+
54+
cases: list[LiteLLMLiveCase] = []
55+
for value in values:
56+
if not isinstance(value, dict):
57+
raise ValueError("Each LiteLLM live case must be a JSON object")
58+
case_data: dict[str, Any] = value
59+
cases.append(
60+
LiteLLMLiveCase(
61+
name=str(case_data["name"]),
62+
model=str(case_data["model"]),
63+
dimensions=int(case_data["dimensions"]),
64+
api_key_env=case_data.get("api_key_env"),
65+
document_input_type=case_data.get("document_input_type"),
66+
query_input_type=case_data.get("query_input_type"),
67+
)
68+
)
69+
return cases
70+
71+
72+
def _live_cases() -> list[LiteLLMLiveCase | Any]:
73+
"""Return built-in and user-supplied live cases whose credentials are available."""
74+
cases: list[LiteLLMLiveCase] = []
75+
76+
if os.getenv("OPENAI_API_KEY"):
77+
cases.append(
78+
LiteLLMLiveCase(
79+
name="openai-text-embedding-3-small",
80+
model="openai/text-embedding-3-small",
81+
dimensions=1536,
82+
api_key_env="OPENAI_API_KEY",
83+
)
84+
)
85+
86+
if os.getenv("COHERE_API_KEY"):
87+
cases.append(
88+
LiteLLMLiveCase(
89+
name="cohere-embed-english-v3",
90+
model="cohere/embed-english-v3.0",
91+
dimensions=1024,
92+
api_key_env="COHERE_API_KEY",
93+
)
94+
)
95+
96+
cases.extend(_custom_cases())
97+
if cases:
98+
return cases
99+
100+
return [
101+
pytest.param(
102+
None,
103+
marks=pytest.mark.skip(
104+
reason=(
105+
"No LiteLLM live cases configured. Set OPENAI_API_KEY, "
106+
"COHERE_API_KEY, or BASIC_MEMORY_TEST_LITELLM_CASES."
107+
)
108+
),
109+
)
110+
]
111+
112+
113+
def _cosine(a: list[float], b: list[float]) -> float:
114+
"""Compute cosine similarity for live ranking sanity checks."""
115+
dot = sum(x * y for x, y in zip(a, b, strict=True))
116+
norm_a = math.sqrt(sum(x * x for x in a))
117+
norm_b = math.sqrt(sum(y * y for y in b))
118+
if norm_a == 0 or norm_b == 0:
119+
return 0.0
120+
return dot / (norm_a * norm_b)
121+
122+
123+
def _assert_valid_vector(vector: list[float], dimensions: int) -> None:
124+
"""Assert provider output is a usable normalized vector."""
125+
assert len(vector) == dimensions
126+
assert all(math.isfinite(value) for value in vector)
127+
norm = math.sqrt(sum(value * value for value in vector))
128+
assert norm == pytest.approx(1.0, abs=1e-6)
129+
130+
131+
@pytest.mark.asyncio
132+
@pytest.mark.parametrize(
133+
"case",
134+
_live_cases(),
135+
ids=lambda case: case.name if isinstance(case, LiteLLMLiveCase) else "no-live-cases",
136+
)
137+
async def test_litellm_live_model_embeds_documents_and_queries(
138+
case: LiteLLMLiveCase,
139+
) -> None:
140+
"""A live LiteLLM model should embed documents and rank a related query higher."""
141+
api_key = os.getenv(case.api_key_env) if case.api_key_env else None
142+
provider = LiteLLMEmbeddingProvider(
143+
model_name=case.model,
144+
dimensions=case.dimensions,
145+
batch_size=2,
146+
api_key=api_key,
147+
timeout=60.0,
148+
document_input_type=case.document_input_type,
149+
query_input_type=case.query_input_type,
150+
)
151+
152+
documents = [
153+
"OAuth login refresh tokens keep an authenticated web session active.",
154+
"A sourdough starter ferments flour and water before bread baking.",
155+
]
156+
vectors = await provider.embed_documents(documents)
157+
query_vector = await provider.embed_query("authentication login token flow")
158+
159+
assert len(vectors) == 2
160+
for vector in [*vectors, query_vector]:
161+
_assert_valid_vector(vector, case.dimensions)
162+
163+
assert _cosine(query_vector, vectors[0]) > _cosine(query_vector, vectors[1])

0 commit comments

Comments
 (0)