mirror of
https://github.com/basicmachines-co/basic-memory
synced 2026-06-21 13:47:35 +00:00
fix(core): harden LiteLLM provider configuration
Signed-off-by: phernandez <paul@basicmachines.co>
This commit is contained in:
@@ -101,7 +101,7 @@ All settings are fields on `BasicMemoryConfig` and can be set via environment va
|
||||
| `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), `"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/LiteLLM OpenAI. Override when using a non-default model. |
|
||||
| `semantic_embedding_dimensions` | `BASIC_MEMORY_SEMANTIC_EMBEDDING_DIMENSIONS` | Provider default | Vector dimensions. 384 for FastEmbed, 1536 for OpenAI/LiteLLM OpenAI. Required when using a non-default LiteLLM 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. |
|
||||
@@ -149,6 +149,8 @@ export BASIC_MEMORY_SEMANTIC_EMBEDDING_DIMENSIONS=1024
|
||||
export COHERE_API_KEY=...
|
||||
```
|
||||
|
||||
Basic Memory creates vector tables before the first embedding call, so non-default LiteLLM models must set `BASIC_MEMORY_SEMANTIC_EMBEDDING_DIMENSIONS`. The LiteLLM OpenAI default (`openai/text-embedding-3-small`) uses 1536 dimensions automatically.
|
||||
|
||||
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`
|
||||
|
||||
@@ -245,7 +245,11 @@ class BasicMemoryConfig(BaseSettings):
|
||||
)
|
||||
semantic_embedding_dimensions: int | None = Field(
|
||||
default=None,
|
||||
description="Embedding vector dimensions. Auto-detected from provider if not set (384 for FastEmbed, 1536 for OpenAI).",
|
||||
description=(
|
||||
"Embedding vector dimensions. Uses provider defaults when unset "
|
||||
"(384 for FastEmbed, 1536 for OpenAI and LiteLLM OpenAI default); "
|
||||
"required for custom LiteLLM models."
|
||||
),
|
||||
)
|
||||
# Trigger: full local rebuilds spend most of their time waiting behind shared
|
||||
# embed flushes, not constructing vectors themselves.
|
||||
|
||||
@@ -107,8 +107,11 @@ def reset_embedding_provider_cache() -> None:
|
||||
def create_embedding_provider(app_config: BasicMemoryConfig) -> EmbeddingProvider:
|
||||
"""Create an embedding provider based on semantic config.
|
||||
|
||||
When semantic_embedding_dimensions is set in config, it overrides
|
||||
the provider's default dimensions (384 for FastEmbed, 1536 for OpenAI).
|
||||
When semantic_embedding_dimensions is set in config, it overrides the
|
||||
provider's default dimensions (384 for FastEmbed, 1536 for OpenAI and
|
||||
the LiteLLM OpenAI default). Custom LiteLLM models require an explicit
|
||||
dimension because the vector table schema is created before the first
|
||||
embedding response is available.
|
||||
"""
|
||||
cache_key = _provider_cache_key(app_config)
|
||||
with _EMBEDDING_PROVIDER_CACHE_LOCK:
|
||||
@@ -161,6 +164,15 @@ def create_embedding_provider(app_config: BasicMemoryConfig) -> EmbeddingProvide
|
||||
model_name = app_config.semantic_embedding_model or "openai/text-embedding-3-small"
|
||||
if model_name == "bge-small-en-v1.5":
|
||||
model_name = "openai/text-embedding-3-small"
|
||||
if (
|
||||
app_config.semantic_embedding_dimensions is None
|
||||
and model_name != "openai/text-embedding-3-small"
|
||||
):
|
||||
raise ValueError(
|
||||
"semantic_embedding_dimensions must be set when "
|
||||
"semantic_embedding_provider='litellm' uses a non-default model. "
|
||||
f"Configured model: {model_name!r}."
|
||||
)
|
||||
provider = LiteLLMEmbeddingProvider(
|
||||
model_name=model_name,
|
||||
batch_size=app_config.semantic_embedding_batch_size,
|
||||
|
||||
@@ -15,6 +15,7 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import math
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from basic_memory.repository.embedding_provider import EmbeddingProvider
|
||||
@@ -45,6 +46,24 @@ def _default_input_types(model_name: str) -> tuple[str | None, str | None]:
|
||||
return None, None
|
||||
|
||||
|
||||
def _import_litellm() -> Any:
|
||||
"""Import LiteLLM without letting its import-time dotenv hook read cwd secrets."""
|
||||
# Constraint: LiteLLM 1.85.0 loads .env files at import time when
|
||||
# LITELLM_MODE defaults to DEV. Basic Memory intentionally does not load
|
||||
# arbitrary cwd .env files, so set the production mode before importing
|
||||
# unless the caller already made an explicit LiteLLM choice.
|
||||
os.environ.setdefault("LITELLM_MODE", "PRODUCTION")
|
||||
|
||||
try:
|
||||
import litellm
|
||||
except ImportError as exc:
|
||||
raise SemanticDependenciesMissingError(
|
||||
"litellm dependency is missing. Install with: pip install litellm"
|
||||
) from exc
|
||||
|
||||
return litellm
|
||||
|
||||
|
||||
class LiteLLMEmbeddingProvider(EmbeddingProvider):
|
||||
"""Embedding provider backed by the litellm SDK."""
|
||||
|
||||
@@ -86,12 +105,7 @@ class LiteLLMEmbeddingProvider(EmbeddingProvider):
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
try:
|
||||
import litellm
|
||||
except ImportError as exc:
|
||||
raise SemanticDependenciesMissingError(
|
||||
"litellm dependency is missing. Install with: pip install litellm"
|
||||
) from exc
|
||||
litellm = _import_litellm()
|
||||
|
||||
batches = [
|
||||
texts[start : start + self.batch_size]
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
import builtins
|
||||
import math
|
||||
import os
|
||||
import sys
|
||||
from types import SimpleNamespace
|
||||
|
||||
@@ -193,6 +194,47 @@ async def test_litellm_provider_missing_dependency_raises_actionable_error(monke
|
||||
await provider.embed_query("test")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_provider_sets_production_mode_before_import(monkeypatch):
|
||||
"""Unset LiteLLM mode should not let LiteLLM import load cwd .env files."""
|
||||
monkeypatch.delitem(sys.modules, "litellm", raising=False)
|
||||
monkeypatch.delenv("LITELLM_MODE", raising=False)
|
||||
observed_modes: list[str | None] = []
|
||||
original_import = builtins.__import__
|
||||
|
||||
async def _aembedding(**kwargs):
|
||||
return _make_embedding_response(kwargs["input"])
|
||||
|
||||
def _observing_import(name, globals=None, locals=None, fromlist=(), level=0):
|
||||
if name == "litellm":
|
||||
observed_modes.append(os.environ.get("LITELLM_MODE"))
|
||||
module = type(sys)("litellm")
|
||||
setattr(module, "aembedding", _aembedding)
|
||||
monkeypatch.setitem(sys.modules, "litellm", module)
|
||||
return module
|
||||
return original_import(name, globals, locals, fromlist, level)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", _observing_import)
|
||||
|
||||
provider = LiteLLMEmbeddingProvider(dimensions=3)
|
||||
await provider.embed_query("test")
|
||||
|
||||
assert observed_modes == ["PRODUCTION"]
|
||||
assert os.environ["LITELLM_MODE"] == "PRODUCTION"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_provider_preserves_explicit_litellm_mode(monkeypatch):
|
||||
"""An explicit LiteLLM mode should stay under the caller's control."""
|
||||
monkeypatch.setenv("LITELLM_MODE", "CUSTOM")
|
||||
_install_litellm_stub(monkeypatch)
|
||||
|
||||
provider = LiteLLMEmbeddingProvider(dimensions=3)
|
||||
await provider.embed_query("test")
|
||||
|
||||
assert os.environ["LITELLM_MODE"] == "CUSTOM"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_provider_output_ordering(monkeypatch):
|
||||
"""Vectors should be returned in the same order as input texts.
|
||||
@@ -256,16 +298,33 @@ def test_factory_forwards_litellm_document_and_query_input_types():
|
||||
semantic_search_enabled=True,
|
||||
semantic_embedding_provider="litellm",
|
||||
semantic_embedding_model="nvidia_nim/nvidia/embed-qa-4",
|
||||
semantic_embedding_dimensions=1024,
|
||||
semantic_embedding_document_input_type="passage",
|
||||
semantic_embedding_query_input_type="query",
|
||||
)
|
||||
provider = create_embedding_provider(config)
|
||||
|
||||
assert isinstance(provider, LiteLLMEmbeddingProvider)
|
||||
assert provider.dimensions == 1024
|
||||
assert provider.document_input_type == "passage"
|
||||
assert provider.query_input_type == "query"
|
||||
|
||||
|
||||
def test_factory_requires_litellm_dimensions_for_custom_models():
|
||||
"""Custom LiteLLM models need explicit dimensions before vector tables are created."""
|
||||
config = BasicMemoryConfig(
|
||||
env="test",
|
||||
projects={"test": "/tmp/basic-memory-test"},
|
||||
default_project="test",
|
||||
semantic_search_enabled=True,
|
||||
semantic_embedding_provider="litellm",
|
||||
semantic_embedding_model="cohere/embed-english-v3.0",
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="semantic_embedding_dimensions"):
|
||||
create_embedding_provider(config)
|
||||
|
||||
|
||||
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