fix(redis-semantic-cache): support custom embedding api_base and api_key

RedisSemanticCache previously hardcoded the embedding call to always use
the default OpenAI endpoint, making it impossible to use self-hosted or
alternative embedding providers (e.g. LM Studio, Ollama, Azure).

Changes:
- Add embedding_api_base and embedding_api_key params to RedisSemanticCache.__init__
- Store on self and pass through to litellm.embedding() in _get_embedding()
- Expose as redis_semantic_cache_embedding_api_base and
  redis_semantic_cache_embedding_api_key on the Cache factory class

Backward compatible: both params default to None; existing deployments
using text-embedding-ada-002 via OpenAI are unaffected.

Fixes the startup crash where CustomTextVectorizer.__init__ calls
_get_embedding() synchronously and fails with a 404/auth error when the
embedding model is not accessible at the default OpenAI endpoint.
This commit is contained in:
Petie Clark 2026-03-19 17:46:30 -04:00
parent 2df965513e
commit 654664795f
2 changed files with 18 additions and 5 deletions

View file

@ -99,6 +99,8 @@ class Cache:
gcs_path_service_account: Optional[str] = None,
gcs_path: Optional[str] = None,
redis_semantic_cache_embedding_model: str = "text-embedding-ada-002",
redis_semantic_cache_embedding_api_base: Optional[str] = None,
redis_semantic_cache_embedding_api_key: Optional[str] = None,
redis_semantic_cache_index_name: Optional[str] = None,
redis_flush_size: Optional[int] = None,
redis_startup_nodes: Optional[List] = None,
@ -205,6 +207,8 @@ class Cache:
password=password,
similarity_threshold=similarity_threshold,
embedding_model=redis_semantic_cache_embedding_model,
embedding_api_base=redis_semantic_cache_embedding_api_base,
embedding_api_key=redis_semantic_cache_embedding_api_key,
index_name=redis_semantic_cache_index_name,
**kwargs,
)

View file

@ -44,6 +44,8 @@ class RedisSemanticCache(BaseCache):
redis_url: Optional[str] = None,
similarity_threshold: Optional[float] = None,
embedding_model: str = "text-embedding-ada-002",
embedding_api_base: Optional[str] = None,
embedding_api_key: Optional[str] = None,
index_name: Optional[str] = None,
**kwargs,
):
@ -86,6 +88,8 @@ class RedisSemanticCache(BaseCache):
# While similarity: 1 = most similar, 0 = least similar
self.distance_threshold = 1 - similarity_threshold
self.embedding_model = embedding_model
self.embedding_api_base = embedding_api_base
self.embedding_api_key = embedding_api_key
# Set up Redis connection
if redis_url is None:
@ -143,13 +147,18 @@ class RedisSemanticCache(BaseCache):
List[float]: The embedding vector
"""
# Create an embedding from prompt
embed_kwargs: dict = {
"model": self.embedding_model,
"input": prompt,
"cache": {"no-store": True, "no-cache": True},
}
if self.embedding_api_base is not None:
embed_kwargs["api_base"] = self.embedding_api_base
if self.embedding_api_key is not None:
embed_kwargs["api_key"] = self.embedding_api_key
embedding_response = cast(
EmbeddingResponse,
litellm.embedding(
model=self.embedding_model,
input=prompt,
cache={"no-store": True, "no-cache": True},
),
litellm.embedding(**embed_kwargs),
)
embedding = embedding_response["data"][0]["embedding"]
return embedding