mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
[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:
parent
c3c6255689
commit
be60d12ff7
4 changed files with 244 additions and 10 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue