Cache Embeddings support for redis-semantic cache

- Add supports of Cache for embeddings when using redis-semantic
This commit is contained in:
fcenedes 2025-12-02 09:01:20 +01:00
parent 4c7a988454
commit 1be0e0c40b
3 changed files with 169 additions and 25 deletions

View file

@ -98,6 +98,10 @@ class Cache:
gcs_path: Optional[str] = None,
redis_semantic_cache_embedding_model: str = "text-embedding-ada-002",
redis_semantic_cache_index_name: Optional[str] = None,
# Embeddings cache (RedisVL) options for redis-semantic
redis_semantic_embedding_cache_enabled: bool = False,
redis_semantic_embedding_cache_ttl: Optional[int] = None,
redis_semantic_embedding_cache_name: Optional[str] = None,
redis_flush_size: Optional[int] = None,
redis_startup_nodes: Optional[List] = None,
disk_cache_dir: Optional[str] = None,
@ -111,8 +115,7 @@ class Cache:
gcp_ssl_ca_certs: Optional[str] = None,
**kwargs,
):
"""
Initializes the cache based on the given type.
"""Initializes the cache based on the given type.
Args:
type (str, optional): The type of cache to initialize. Can be "local", "redis", "redis-semantic", "qdrant-semantic", "s3" or "disk". Defaults to "local".
@ -195,6 +198,9 @@ class Cache:
similarity_threshold=similarity_threshold,
embedding_model=redis_semantic_cache_embedding_model,
index_name=redis_semantic_cache_index_name,
embedding_cache_enabled=redis_semantic_embedding_cache_enabled,
embedding_cache_ttl=redis_semantic_embedding_cache_ttl,
embedding_cache_name=redis_semantic_embedding_cache_name,
**kwargs,
)
elif type == LiteLLMCacheType.QDRANT_SEMANTIC:

View file

@ -45,10 +45,13 @@ class RedisSemanticCache(BaseCache):
similarity_threshold: Optional[float] = None,
embedding_model: str = "text-embedding-ada-002",
index_name: Optional[str] = None,
# Embeddings cache (RedisVL) configuration
embedding_cache_enabled: bool = False,
embedding_cache_ttl: Optional[int] = None,
embedding_cache_name: Optional[str] = None,
**kwargs,
):
"""
Initialize the Redis Semantic Cache.
"""Initialize the Redis Semantic Cache.
Args:
host: Redis host address
@ -59,6 +62,9 @@ class RedisSemanticCache(BaseCache):
where 1.0 requires exact matches and 0.0 accepts any match
embedding_model: Model to use for generating embeddings
index_name: Name for the Redis index
embedding_cache_enabled: Whether to enable RedisVL EmbeddingsCache
embedding_cache_ttl: Default TTL for embeddings cache entries in seconds
embedding_cache_name: Optional name prefix for embeddings cache keys
ttl: Default time-to-live for cache entries in seconds
**kwargs: Additional arguments passed to the Redis client
@ -66,6 +72,8 @@ class RedisSemanticCache(BaseCache):
Exception: If similarity_threshold is not provided or required Redis
connection information is missing
"""
# Import RedisVL components lazily to avoid hard dependency when not used
from redisvl.extensions.cache.embeddings import EmbeddingsCache
from redisvl.extensions.llmcache import SemanticCache
from redisvl.utils.vectorize import CustomTextVectorizer
@ -87,6 +95,14 @@ class RedisSemanticCache(BaseCache):
self.distance_threshold = 1 - similarity_threshold
self.embedding_model = embedding_model
# Embeddings cache configuration
self.embedding_cache_enabled: bool = embedding_cache_enabled
self.embedding_cache_ttl: Optional[int] = embedding_cache_ttl
self.embedding_cache_name: str = (
embedding_cache_name or "litellm_redis_semantic_embeddings_cache"
)
self._embeddings_cache: Optional[EmbeddingsCache] = None
# Set up Redis connection
if redis_url is None:
try:
@ -106,8 +122,26 @@ class RedisSemanticCache(BaseCache):
print_verbose(f"Redis semantic-cache redis_url: {redis_url}")
# Initialize the embeddings cache if enabled
if self.embedding_cache_enabled:
try:
self._embeddings_cache = EmbeddingsCache(
name=self.embedding_cache_name,
redis_url=redis_url,
ttl=self.embedding_cache_ttl,
)
except Exception as e: # pragma: no cover - defensive, treat as non-fatal
print_verbose(
f"Redis semantic-cache: failed to initialize EmbeddingsCache, "
f"disabling embedding cache. Error: {str(e)}"
)
self.embedding_cache_enabled = False
self._embeddings_cache = None
# Initialize the Redis vectorizer and cache
cache_vectorizer = CustomTextVectorizer(self._get_embedding)
cache_vectorizer = CustomTextVectorizer(
self._get_embedding, cache=self._embeddings_cache
)
self.llmcache = SemanticCache(
name=index_name,
@ -117,6 +151,14 @@ class RedisSemanticCache(BaseCache):
overwrite=False,
)
@property
def embeddings_cache(self):
"""Expose the underlying EmbeddingsCache instance (if any).
This is primarily for tests and potential reuse by other components.
"""
return self._embeddings_cache
def _get_ttl(self, **kwargs) -> Optional[int]:
"""
Get the TTL (time-to-live) value for cache entries.
@ -133,16 +175,16 @@ class RedisSemanticCache(BaseCache):
return ttl
def _get_embedding(self, prompt: str) -> List[float]:
"""
Generate an embedding vector for the given prompt using the configured embedding model.
"""Generate an embedding vector for the given prompt.
Args:
prompt: The text to generate an embedding for
Returns:
List[float]: The embedding vector
This is the sync embedding function used by RedisVL's CustomTextVectorizer.
It deliberately bypasses LiteLLM's high-level cache (cache={"no-store": True})
because EmbeddingsCache is responsible for caching at this layer.
"""
# Create an embedding from prompt
# NOTE: EmbeddingsCache is already wired into CustomTextVectorizer via
# the `cache` parameter in __init__, so this method only needs to
# compute the embedding when there is an EmbeddingsCache miss.
embedding_response = cast(
EmbeddingResponse,
litellm.embedding(
@ -151,8 +193,7 @@ class RedisSemanticCache(BaseCache):
cache={"no-store": True, "no-cache": True},
),
)
embedding = embedding_response["data"][0]["embedding"]
return embedding
return embedding_response["data"][0]["embedding"]
def _get_cache_logic(self, cached_response: Any) -> Any:
"""
@ -269,16 +310,28 @@ class RedisSemanticCache(BaseCache):
print_verbose(f"Error retrieving from Redis semantic cache: {str(e)}")
async def _get_async_embedding(self, prompt: str, **kwargs) -> List[float]:
"""
Asynchronously generate an embedding for the given prompt.
"""Asynchronously generate an embedding for the given prompt.
Args:
prompt: The text to generate an embedding for
**kwargs: Additional arguments that may contain metadata
Returns:
List[float]: The embedding vector
This is used by the async semantic cache paths. It first checks the
RedisVL EmbeddingsCache (if enabled) before falling back to the
underlying embedding provider.
"""
# Fast path: check embeddings cache if available
if self.embedding_cache_enabled and self._embeddings_cache is not None:
try:
cached = await self._embeddings_cache.aget(
text=prompt,
model_name=self.embedding_model,
)
if cached is not None:
return cached["embedding"] # type: ignore[index]
except Exception as e: # pragma: no cover - defensive
print_verbose(
f"Redis semantic-cache: EmbeddingsCache.aget failed, "
f"falling back to provider. Error: {str(e)}"
)
from litellm.proxy.proxy_server import llm_model_list, llm_router
# Route the embedding request through the proxy if appropriate
@ -310,8 +363,24 @@ class RedisSemanticCache(BaseCache):
cache={"no-store": True, "no-cache": True},
)
# Extract and return the embedding vector
return embedding_response["data"][0]["embedding"]
embedding_vec = embedding_response["data"][0]["embedding"]
# Store in embeddings cache (best effort)
if self.embedding_cache_enabled and self._embeddings_cache is not None:
try:
await self._embeddings_cache.aset(
text=prompt,
model_name=self.embedding_model,
embedding=embedding_vec,
ttl=self.embedding_cache_ttl,
)
except Exception as e: # pragma: no cover - defensive
print_verbose(
"Redis semantic-cache: EmbeddingsCache.aset failed; "
f"continuing without cache. Error: {str(e)}"
)
return embedding_vec
except Exception as e:
print_verbose(f"Error generating async embedding: {str(e)}")
raise ValueError(f"Failed to generate embedding: {str(e)}") from e

View file

@ -144,3 +144,72 @@ async def test_redis_semantic_cache_async_get_cache(monkeypatch):
# Verify methods were called
redis_semantic_cache._get_async_embedding.assert_called_once()
redis_semantic_cache.llmcache.acheck.assert_called_once()
@patch.dict(
"sys.modules",
{
"redisvl.extensions.llmcache": MagicMock(),
"redisvl.utils.vectorize": MagicMock(),
"redisvl.extensions.cache.embeddings": MagicMock(),
},
)
def test_redis_semantic_cache_embeddings_cache_enabled(monkeypatch):
from litellm.caching.redis_semantic_cache import RedisSemanticCache
# Set environment variables
monkeypatch.setenv("REDIS_HOST", "localhost")
monkeypatch.setenv("REDIS_PORT", "6379")
monkeypatch.setenv("REDIS_PASSWORD", "test_password")
cache = RedisSemanticCache(
similarity_threshold=0.8,
embedding_cache_enabled=True,
embedding_cache_name="test_embed_cache",
embedding_cache_ttl=123,
)
# Ensure embeddings cache is initialized when enabled
assert cache.embedding_cache_enabled is True
assert cache.embeddings_cache is not None
@pytest.mark.asyncio
async def test_redis_semantic_cache_async_embedding_uses_cache(monkeypatch):
# Patch redisvl modules and EmbeddingsCache class specifically
embeddings_cache_mock_cls = MagicMock()
embeddings_cache_instance = AsyncMock()
embeddings_cache_mock_cls.return_value = embeddings_cache_instance
with patch.dict(
"sys.modules",
{
"redisvl.extensions.llmcache": MagicMock(),
"redisvl.utils.vectorize": MagicMock(),
"redisvl.extensions.cache.embeddings": MagicMock(
EmbeddingsCache=embeddings_cache_mock_cls
),
},
):
from litellm.caching.redis_semantic_cache import RedisSemanticCache
# Set environment variables
monkeypatch.setenv("REDIS_HOST", "localhost")
monkeypatch.setenv("REDIS_PORT", "6379")
monkeypatch.setenv("REDIS_PASSWORD", "test_password")
# Embeddings cache returns a cached vector
embeddings_cache_instance.aget.return_value = {"embedding": [0.1, 0.2, 0.3]}
cache = RedisSemanticCache(
similarity_threshold=0.8,
embedding_cache_enabled=True,
)
# Call internal async embedding helper
result = await cache._get_async_embedding("hello world")
# Should have used the embeddings cache and not fallen back to provider
embeddings_cache_instance.aget.assert_awaited_once()
assert result == [0.1, 0.2, 0.3]