[LLM Translation - Redis] fix: redis caching for embedding response models (#12750)

* fix: redis caching for embedding responses

* add helper

* add mypy fixes

* lint fix

* review changes

* remove file

* fix ruff

* add if check

* add if check
This commit is contained in:
Jugal D. Bhatt 2025-07-19 05:01:10 +05:30 • committed by GitHub
parent c3c6255689
commit be60d12ff7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 244 additions and 10 deletions

View file

@ -590,6 +590,38 @@ class Cache:
except Exception as e:
verbose_logger.exception(f"LiteLLM Cache: Excepton add_cache: {str(e)}")
def _convert_to_cached_embedding(self, embedding_response: Any, model: Optional[str]) -> CachedEmbedding:
"""
Convert any embedding response into the standardized CachedEmbedding TypedDict format.
"""
try:
if isinstance(embedding_response, dict):
return {
"embedding": embedding_response.get("embedding"),
"index": embedding_response.get("index"),
"object": embedding_response.get("object"),
"model": model,
}
elif hasattr(embedding_response, 'model_dump'):
data = embedding_response.model_dump()
return {
"embedding": data.get("embedding"),
"index": data.get("index"),
"object": data.get("object"),
"model": model,
}
else:
data = vars(embedding_response)
return {
"embedding": data.get("embedding"),
"index": data.get("index"),
"object": data.get("object"),
"model": model,
}
except KeyError as e:
raise ValueError(f"Missing expected key in embedding response: {e}")
def add_embedding_response_to_cache(
self,
result: EmbeddingResponse,
@ -600,8 +632,13 @@ class Cache:
preset_cache_key = self.get_cache_key(**{**kwargs, "input": input})
kwargs["cache_key"] = preset_cache_key
embedding_response = result.data[idx_in_result_data]
# Always convert to properly typed CachedEmbedding
model_name = result.model
embedding_dict: CachedEmbedding = self._convert_to_cached_embedding(embedding_response, model_name)
cache_key, cached_data, kwargs = self._add_cache_logic(
result=embedding_response,
result=embedding_dict,
**kwargs,
)
return cache_key, cached_data, kwargs

View file

@ -36,6 +36,7 @@ from pydantic import BaseModel
import litellm
from litellm._logging import print_verbose, verbose_logger
from litellm.caching.caching import S3Cache
from litellm.types.caching import CachedEmbedding
from litellm.litellm_core_utils.logging_utils import (
_assemble_complete_response_from_streaming_chunks,
)
@ -305,10 +306,25 @@ class LLMCachingHandler:
else:
raise ValueError("input must be a string or a list")
def _extract_model_from_cached_results(self, non_null_list: List[Tuple[int, CachedEmbedding]]) -> Optional[str]:
"""
Helper method to extract the model name from cached results.
Args:
non_null_list: List of (idx, cr) tuples where cr is the cached result dict
Returns:
Optional[str]: The model name if found, None otherwise
"""
for _, cr in non_null_list:
if isinstance(cr, dict) and cr.get("model"):
return cr["model"]
return None
def _process_async_embedding_cached_response(
self,
final_embedding_cached_response: Optional[EmbeddingResponse],
cached_result: List[Optional[Dict[str, Any]]],
cached_result: List[Optional[CachedEmbedding]],
kwargs: Dict[str, Any],
logging_obj: LiteLLMLoggingObj,
start_time: datetime.datetime,
@ -345,9 +361,12 @@ class LLMCachingHandler:
non_null_list.append((idx, cr))
kwargs["input"] = remaining_list
if len(non_null_list) > 0:
verbose_logger.debug(f"EMBEDDING CACHE HIT! - {len(non_null_list)}")
# Use the model from the first non-null cached result, fallback to kwargs if not present
model_name = self._extract_model_from_cached_results(non_null_list)
if not model_name:
model_name = kwargs.get("model")
final_embedding_cached_response = EmbeddingResponse(
model=kwargs.get("model"),
model=model_name,
data=[None] * len(kwargs_input_as_list),
)
final_embedding_cached_response._hidden_params["cache_hit"] = True
@ -356,11 +375,13 @@ class LLMCachingHandler:
for val in non_null_list:
idx, cr = val # (idx, cr) tuple
if cr is not None:
final_embedding_cached_response.data[idx] = Embedding(
embedding=cr["embedding"],
index=idx,
object="embedding",
)
embedding_data = cr.get("embedding")
if embedding_data is not None:
final_embedding_cached_response.data[idx] = Embedding(
embedding=embedding_data,
index=idx,
object="embedding",
)
if isinstance(kwargs_input_as_list[idx], str):
from litellm.utils import token_counter

View file

@ -1,5 +1,5 @@
from enum import Enum
from typing import Any, Dict, Literal, Optional, TypedDict, Union
from typing import Any, Dict, List, Literal, Optional, TypedDict, Union
from pydantic import BaseModel
@ -87,3 +87,11 @@ class HealthCheckCacheParams(BaseModel):
redis_kwargs: Optional[Dict[str, Any]] = None
namespace: Optional[str] = None
redis_version: Optional[Union[str, int, float]] = None
class CachedEmbedding(TypedDict):
"""Type definition for cached embedding objects"""
embedding: Optional[List[float]]
index: Optional[int]
object: Optional[str]
model: Optional[str]

View file

@ -334,3 +334,171 @@ def test_combine_cached_embedding_response_multiple_missing_values():
assert result.data[1].embedding == [0.4, 0.5, 0.6]
assert result.data[2].embedding == [0.4, 0.5, 0.6]
assert result.data[3].embedding == [0.7, 0.8, 0.9]
@pytest.mark.asyncio
async def test_embedding_cache_model_field_consistency():
"""
Test that the model field is consistently preserved in cached embedding responses.
This ensures that cache hits return the same model field as the original API response.
"""
# Setup cache
setup_cache()
caching_handler = LLMCachingHandler(
original_function=aembedding, request_kwargs={}, start_time=datetime.now()
)
# Create a mock embedding response with a specific model
original_model = "text-embedding-005"
embedding_response = EmbeddingResponse(
model=original_model,
data=[
Embedding(embedding=[0.1, 0.2, 0.3], index=0, object="embedding"),
Embedding(embedding=[0.4, 0.5, 0.6], index=1, object="embedding"),
]
)
# Mock logging object
logging_obj = LiteLLMLogging(
litellm_call_id=str(datetime.now()),
call_type=CallTypes.aembedding.value,
model=original_model,
messages=[], # Not used for embeddings
function_id=str(uuid.uuid4()),
stream=False,
start_time=datetime.now(),
)
# Test parameters
kwargs = {
"model": original_model,
"input": ["test input 1", "test input 2"],
"caching": True
}
# Step 1: Cache the embedding response
await caching_handler.async_set_cache(
result=embedding_response,
original_function=aembedding,
kwargs=kwargs
)
# Step 2: Retrieve from cache
cached_response = await caching_handler._async_get_cache(
model=original_model,
original_function=aembedding,
logging_obj=logging_obj,
start_time=datetime.now(),
call_type=CallTypes.aembedding.value,
kwargs=kwargs,
)
# Step 3: Verify the model field is preserved
assert cached_response.final_embedding_cached_response is not None
assert cached_response.final_embedding_cached_response.model == original_model
assert len(cached_response.final_embedding_cached_response.data) == 2
assert cached_response.final_embedding_cached_response.data[0].embedding == [0.1, 0.2, 0.3]
assert cached_response.final_embedding_cached_response.data[0].index == 0
assert cached_response.final_embedding_cached_response.data[1].embedding == [0.4, 0.5, 0.6]
assert cached_response.final_embedding_cached_response.data[1].index == 1
# Verify cache hit flag is set
assert cached_response.final_embedding_cached_response._hidden_params["cache_hit"] == True
@pytest.mark.asyncio
async def test_embedding_cache_model_field_with_vendor_prefix():
"""
Test that the model field is preserved even when using vendor-prefixed model names.
This simulates the real-world scenario where models might be prefixed with vendor names.
"""
# Setup cache
setup_cache()
caching_handler = LLMCachingHandler(
original_function=aembedding, request_kwargs={}, start_time=datetime.now()
)
# Test with vendor-prefixed model name (like vertex_ai/text-embedding-005)
vendor_model = "vertex_ai/text-embedding-005"
actual_model = "text-embedding-005" # What the provider actually returns
# Create embedding response with the actual model name (as returned by provider)
embedding_response = EmbeddingResponse(
model=actual_model, # Provider returns this
data=[
Embedding(embedding=[0.1, 0.2, 0.3], index=0, object="embedding"),
]
)
# Mock logging object
logging_obj = LiteLLMLogging(
litellm_call_id=str(datetime.now()),
call_type=CallTypes.aembedding.value,
model=vendor_model,
messages=[],
function_id=str(uuid.uuid4()),
stream=False,
start_time=datetime.now(),
)
# Test parameters with vendor-prefixed model
kwargs = {
"model": vendor_model, # Request uses vendor prefix
"input": ["test input"],
"caching": True
}
# Cache the response
await caching_handler.async_set_cache(
result=embedding_response,
original_function=aembedding,
kwargs=kwargs
)
# Retrieve from cache
cached_response = await caching_handler._async_get_cache(
model=vendor_model,
original_function=aembedding,
logging_obj=logging_obj,
start_time=datetime.now(),
call_type=CallTypes.aembedding.value,
kwargs=kwargs,
)
# Verify the model field matches the original provider response, not the request
assert cached_response.final_embedding_cached_response is not None
assert cached_response.final_embedding_cached_response.model == actual_model # Should be the provider's model name
assert cached_response.final_embedding_cached_response.model != vendor_model # Should NOT be the vendor-prefixed name
def test_extract_model_from_cached_results():
"""
Test the helper method that extracts model names from cached results.
"""
caching_handler = LLMCachingHandler(
original_function=aembedding, request_kwargs={}, start_time=datetime.now()
)
# Test with valid cached results
non_null_list = [
(0, {"embedding": [0.1, 0.2], "index": 0, "object": "embedding", "model": "text-embedding-005"}),
(1, {"embedding": [0.3, 0.4], "index": 1, "object": "embedding", "model": "text-embedding-005"}),
]
model_name = caching_handler._extract_model_from_cached_results(non_null_list)
assert model_name == "text-embedding-005"
# Test with missing model field
non_null_list_no_model = [
(0, {"embedding": [0.1, 0.2], "index": 0, "object": "embedding"}),
(1, {"embedding": [0.3, 0.4], "index": 1, "object": "embedding"}),
]
model_name = caching_handler._extract_model_from_cached_results(non_null_list_no_model)
assert model_name is None
# Test with empty list
model_name = caching_handler._extract_model_from_cached_results([])
assert model_name is None