fix(caching): sort embedding results by index to prevent cache misalignment (#20456)

When providers like vLLM return embedding results out of order (i.e.
result.data[0].index != 0), three things go wrong:

1. async_add_cache_pipeline stores embeddings using positional access
   (result.data[idx]) instead of the index field, so the wrong embedding
   gets cached under the wrong key.

2. _combine_cached_embedding_response_with_api_result fills None
   (uncached) slots with a sequential counter into the unsorted API
   response, placing wrong embeddings at wrong positions.

3. After combining, the index field on API result items retains the
   provider-relative value (0, 1, 2, ...) instead of the final merged
   position, causing duplicate/missing indices visible to the caller.

Fixes:
- Sort result.data by index in async_add_cache_pipeline before storing
- Sort embedding_response.data by index in _combine_cached_... before
  sequential filling
- Correct index fields on all items after merging
- Fix pre-existing unused-import lint error (F401) in
  policy_resolve_endpoints.py
- Add 9 regression tests covering out-of-order responses, index
  correction, large batches with scattered cache hits
This commit is contained in:
skylarkoo7 2026-02-11 13:42:31 +05:30
parent d9c69ae9e5
commit c18f1b9482
5 changed files with 302 additions and 1 deletions

View file

@ -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"]):

View file

@ -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"

View file

@ -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

View file

@ -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():
"""

View file

@ -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]