From 1be0e0c40b4a39a08263172b4de43eb3b47fe15b Mon Sep 17 00:00:00 2001 From: fcenedes Date: Tue, 2 Dec 2025 09:01:20 +0100 Subject: [PATCH] Cache Embeddings support for redis-semantic cache - Add supports of Cache for embeddings when using redis-semantic --- litellm/caching/caching.py | 10 +- litellm/caching/redis_semantic_cache.py | 115 ++++++++++++++---- .../caching/test_redis_semantic_cache.py | 69 +++++++++++ 3 files changed, 169 insertions(+), 25 deletions(-) diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 82fc37e0cb4..3679a581813 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -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: diff --git a/litellm/caching/redis_semantic_cache.py b/litellm/caching/redis_semantic_cache.py index c76f27377d8..def28586551 100644 --- a/litellm/caching/redis_semantic_cache.py +++ b/litellm/caching/redis_semantic_cache.py @@ -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 diff --git a/tests/test_litellm/caching/test_redis_semantic_cache.py b/tests/test_litellm/caching/test_redis_semantic_cache.py index f9946e266fe..db473ea44a1 100644 --- a/tests/test_litellm/caching/test_redis_semantic_cache.py +++ b/tests/test_litellm/caching/test_redis_semantic_cache.py @@ -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]