diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index a03bff60686..dbd3d8a0c97 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -708,6 +708,20 @@ class Cache: if self.ttl is not None: kwargs["ttl"] = self.ttl + # Sort result.data by index so that positional access in + # add_embedding_response_to_cache correctly maps input[i] to + # the embedding for input[i]. Some providers (e.g. vLLM) may + # return embeddings out of order. See #20456. + if ( + isinstance(result, EmbeddingResponse) + and result.data + and hasattr(result.data[0], "index") + ): + result.data = sorted( + result.data, + key=lambda e: getattr(e, "index", 0), + ) + cache_list = [] if isinstance(kwargs["input"], list): for idx, i in enumerate(kwargs["input"]): diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 4e97197a9de..00a1cfc0a64 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -503,6 +503,16 @@ class LLMCachingHandler: if _caching_handler_response.final_embedding_cached_response is None: return embedding_response + # Sort the API response by the ``index`` field before filling + # None slots. Some providers (e.g. vLLM) may return embedding + # results in a different order than the input, and the sequential + # counter used below assumes sorted order. See #20456. + if embedding_response.data is not None: + embedding_response.data = sorted( + embedding_response.data, + key=lambda e: getattr(e, "index", 0), + ) + idx = 0 final_data_list = [] for item in _caching_handler_response.final_embedding_cached_response.data: @@ -512,6 +522,15 @@ class LLMCachingHandler: else: final_data_list.append(item) + # Correct the ``index`` field on every item so that it matches + # the item's position in the final merged list. Without this, + # API result items retain their provider-relative indices (0, 1, + # 2, …) which do not match the original input positions when + # there were cache hits. See #20456. + for pos, item in enumerate(final_data_list): + if item is not None and hasattr(item, "index"): + item.index = pos + _caching_handler_response.final_embedding_cached_response.data = final_data_list _caching_handler_response.final_embedding_cached_response._hidden_params[ "cache_hit" diff --git a/litellm/proxy/policy_engine/policy_resolve_endpoints.py b/litellm/proxy/policy_engine/policy_resolve_endpoints.py index a81d4405678..eb4d3fc5845 100644 --- a/litellm/proxy/policy_engine/policy_resolve_endpoints.py +++ b/litellm/proxy/policy_engine/policy_resolve_endpoints.py @@ -12,7 +12,6 @@ from fastapi import APIRouter, Depends, HTTPException, Query from litellm._logging import verbose_proxy_logger from litellm.constants import MAX_POLICY_ESTIMATE_IMPACT_ROWS from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry from litellm.proxy.policy_engine.policy_registry import get_policy_registry diff --git a/tests/local_testing/test_caching_handler.py b/tests/local_testing/test_caching_handler.py index 83822b5fcad..6eefee409f2 100644 --- a/tests/local_testing/test_caching_handler.py +++ b/tests/local_testing/test_caching_handler.py @@ -338,6 +338,112 @@ def test_combine_cached_embedding_response_multiple_missing_values(): assert result.data[3].embedding == [0.7, 0.8, 0.9] +def test_combine_cached_embedding_response_out_of_order_api_results(): + """ + Regression test for #20456 — when the provider returns embedding results + out of order (e.g. vLLM), the combine function must still place the + correct embeddings into the correct positions and fix the index fields. + + Cached: [cache_hit, None, None, cache_hit, None] + API returns (out of order): [idx=2, idx=0, idx=1] + Expected final: [cache_hit, api[idx=0], api[idx=1], cache_hit, api[idx=2]] + with all index fields matching their position (0,1,2,3,4). + """ + caching_handler = LLMCachingHandler( + original_function=lambda: None, request_kwargs={}, start_time=datetime.now() + ) + + start_time = datetime.now() + end_time = start_time + timedelta(seconds=1) + + cached_response = EmbeddingResponse( + data=[ + Embedding(embedding=[0.1, 0.2, 0.3], index=0, object="embedding"), + None, + None, + Embedding(embedding=[0.7, 0.8, 0.9], index=3, object="embedding"), + None, + ] + ) + caching_handler_response = CachingHandlerResponse( + final_embedding_cached_response=cached_response + ) + + # API returns results OUT OF ORDER — index field does not match position + api_response = EmbeddingResponse( + data=[ + Embedding(embedding=[1.0, 1.1, 1.2], index=2, object="embedding"), + Embedding(embedding=[0.4, 0.5, 0.6], index=0, object="embedding"), + Embedding(embedding=[2.0, 2.1, 2.2], index=1, object="embedding"), + ] + ) + + result = caching_handler._combine_cached_embedding_response_with_api_result( + _caching_handler_response=caching_handler_response, + embedding_response=api_response, + start_time=start_time, + end_time=end_time, + ) + + assert isinstance(result, EmbeddingResponse) + assert len(result.data) == 5 + # Cached items stay in place + assert result.data[0].embedding == [0.1, 0.2, 0.3] + assert result.data[3].embedding == [0.7, 0.8, 0.9] + # API items placed in correct positions (sorted by their original index) + assert result.data[1].embedding == [0.4, 0.5, 0.6] # was api index=0 + assert result.data[2].embedding == [2.0, 2.1, 2.2] # was api index=1 + assert result.data[4].embedding == [1.0, 1.1, 1.2] # was api index=2 + # All index fields must match their final position + for i, item in enumerate(result.data): + assert item.index == i, f"data[{i}].index = {item.index}, expected {i}" + + +def test_combine_cached_embedding_response_corrects_index_fields(): + """ + Regression test for #20456 — even when API results are in order, the + index fields should be corrected to match the final merged positions. + + Cached: [cache_hit, None, cache_hit] + API returns: [idx=0] (single uncached item) + Expected: data[1].index should be 1, not 0. + """ + caching_handler = LLMCachingHandler( + original_function=lambda: None, request_kwargs={}, start_time=datetime.now() + ) + + start_time = datetime.now() + end_time = start_time + timedelta(seconds=1) + + cached_response = EmbeddingResponse( + data=[ + Embedding(embedding=[0.1, 0.2, 0.3], index=0, object="embedding"), + None, + Embedding(embedding=[0.7, 0.8, 0.9], index=2, object="embedding"), + ] + ) + caching_handler_response = CachingHandlerResponse( + final_embedding_cached_response=cached_response + ) + + api_response = EmbeddingResponse( + data=[Embedding(embedding=[0.4, 0.5, 0.6], index=0, object="embedding")] + ) + + result = caching_handler._combine_cached_embedding_response_with_api_result( + _caching_handler_response=caching_handler_response, + embedding_response=api_response, + start_time=start_time, + end_time=end_time, + ) + + assert len(result.data) == 3 + # Index must match final position + assert result.data[0].index == 0 + assert result.data[1].index == 1 # was api index=0, corrected to 1 + assert result.data[2].index == 2 + + @pytest.mark.asyncio async def test_embedding_cache_model_field_consistency(): """ diff --git a/tests/test_litellm/test_embedding_cache_index_alignment.py b/tests/test_litellm/test_embedding_cache_index_alignment.py new file mode 100644 index 00000000000..da8895b2b83 --- /dev/null +++ b/tests/test_litellm/test_embedding_cache_index_alignment.py @@ -0,0 +1,163 @@ +""" +Regression tests for #20456 — Async Batch Embedding + Redis Cache +results in Index Misalignment/Duplication. + +The root cause is threefold: +1. ``async_add_cache_pipeline`` stored embeddings using positional access + (``result.data[idx]``) instead of the ``index`` field, so when a + provider like vLLM returns results out of order, the wrong embedding + gets stored under the wrong cache key. +2. ``_combine_cached_embedding_response_with_api_result`` filled ``None`` + (uncached) slots with a sequential counter into the unsorted API + response, placing the wrong embeddings at the wrong positions. +3. After combining, the ``index`` field on each Embedding was not + corrected to match its final position in the merged list, leading to + duplicate or missing indices in ``[data.index for data in result.data]``. +""" + +import os +import sys + +sys.path.insert( + 0, os.path.abspath("../..") +) # Adds the parent directory to the system path + +from datetime import datetime, timedelta + +from litellm.caching.caching_handler import CachingHandlerResponse, LLMCachingHandler +from litellm.types.utils import Embedding, EmbeddingResponse + + +def _make_handler(): + return LLMCachingHandler( + original_function=lambda: None, request_kwargs={}, start_time=datetime.now() + ) + + +class TestCombineOutOfOrderAPIResponse: + """Tests for ``_combine_cached_embedding_response_with_api_result`` + when the provider returns results out of order.""" + + def test_single_uncached_in_middle(self): + """[hit, MISS, hit] with API returning a single item.""" + h = _make_handler() + cached = EmbeddingResponse( + data=[ + Embedding(embedding=[1.0], index=0, object="embedding"), + None, + Embedding(embedding=[3.0], index=2, object="embedding"), + ] + ) + api = EmbeddingResponse( + data=[Embedding(embedding=[2.0], index=0, object="embedding")] + ) + resp = CachingHandlerResponse(final_embedding_cached_response=cached) + result = h._combine_cached_embedding_response_with_api_result( + resp, api, datetime.now(), datetime.now() + timedelta(seconds=1) + ) + assert [d.embedding for d in result.data] == [[1.0], [2.0], [3.0]] + assert [d.index for d in result.data] == [0, 1, 2] + + def test_multiple_uncached_reversed(self): + """[MISS, hit, MISS, MISS] with API returning 3 items reversed.""" + h = _make_handler() + cached = EmbeddingResponse( + data=[ + None, + Embedding(embedding=[20.0], index=1, object="embedding"), + None, + None, + ] + ) + # API returns items in REVERSE order + api = EmbeddingResponse( + data=[ + Embedding(embedding=[40.0], index=2, object="embedding"), + Embedding(embedding=[30.0], index=1, object="embedding"), + Embedding(embedding=[10.0], index=0, object="embedding"), + ] + ) + resp = CachingHandlerResponse(final_embedding_cached_response=cached) + result = h._combine_cached_embedding_response_with_api_result( + resp, api, datetime.now(), datetime.now() + timedelta(seconds=1) + ) + assert [d.embedding for d in result.data] == [[10.0], [20.0], [30.0], [40.0]] + assert [d.index for d in result.data] == [0, 1, 2, 3] + + def test_all_uncached(self): + """All cache misses — pure API response reordering.""" + h = _make_handler() + cached = EmbeddingResponse(data=[None, None, None]) + # API returns shuffled + api = EmbeddingResponse( + data=[ + Embedding(embedding=[3.0], index=2, object="embedding"), + Embedding(embedding=[1.0], index=0, object="embedding"), + Embedding(embedding=[2.0], index=1, object="embedding"), + ] + ) + resp = CachingHandlerResponse(final_embedding_cached_response=cached) + result = h._combine_cached_embedding_response_with_api_result( + resp, api, datetime.now(), datetime.now() + timedelta(seconds=1) + ) + assert [d.embedding for d in result.data] == [[1.0], [2.0], [3.0]] + assert [d.index for d in result.data] == [0, 1, 2] + + def test_large_batch_with_scattered_cache_hits(self): + """Simulate a 10-item batch with 4 cache hits and 6 misses.""" + h = _make_handler() + data = [None] * 10 + # Cache hits at positions 1, 4, 7, 9 + for pos in [1, 4, 7, 9]: + data[pos] = Embedding( + embedding=[float(pos)], index=pos, object="embedding" + ) + cached = EmbeddingResponse(data=data) + + # 6 uncached positions: 0, 2, 3, 5, 6, 8 + # API returns them shuffled + miss_positions = [0, 2, 3, 5, 6, 8] + api_data = [ + Embedding(embedding=[float(p)], index=i, object="embedding") + for i, p in enumerate(reversed(miss_positions)) + ] + api = EmbeddingResponse(data=api_data) + + resp = CachingHandlerResponse(final_embedding_cached_response=cached) + result = h._combine_cached_embedding_response_with_api_result( + resp, api, datetime.now(), datetime.now() + timedelta(seconds=1) + ) + + assert len(result.data) == 10 + # All indices must be 0..9 + assert [d.index for d in result.data] == list(range(10)) + # Cache hits must be preserved + for pos in [1, 4, 7, 9]: + assert result.data[pos].embedding == [float(pos)] + + +class TestCachePipelineSorting: + """Tests for ``async_add_cache_pipeline`` sorting result.data by index + before storing.""" + + def test_sort_preserves_input_alignment(self): + """After sorting, result.data[i] should correspond to input[i].""" + # Simulate out-of-order API result + result = EmbeddingResponse( + data=[ + Embedding(embedding=[3.0], index=2, object="embedding"), + Embedding(embedding=[1.0], index=0, object="embedding"), + Embedding(embedding=[2.0], index=1, object="embedding"), + ] + ) + + # The fix sorts in-place + if result.data and hasattr(result.data[0], "index"): + result.data = sorted( + result.data, key=lambda e: getattr(e, "index", 0) + ) + + # After sorting, positional access matches input order + assert result.data[0].embedding == [1.0] + assert result.data[1].embedding == [2.0] + assert result.data[2].embedding == [3.0]