mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
2df965513e
commit
654664795f
2 changed files with 18 additions and 5 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue