diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 9165fec1e3f..6959467cddd 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -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 diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 43e4fb7c3d9..dcc59b20714 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -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 diff --git a/litellm/types/caching.py b/litellm/types/caching.py index e457fe8a127..2531444ae81 100644 --- a/litellm/types/caching.py +++ b/litellm/types/caching.py @@ -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] diff --git a/tests/local_testing/test_caching_handler.py b/tests/local_testing/test_caching_handler.py index c0e3b8d306a..c969c9d4bea 100644 --- a/tests/local_testing/test_caching_handler.py +++ b/tests/local_testing/test_caching_handler.py @@ -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