mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
d9c69ae9e5
commit
c18f1b9482
5 changed files with 302 additions and 1 deletions
|
|
@ -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"]):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
163
tests/test_litellm/test_embedding_cache_index_alignment.py
Normal file
163
tests/test_litellm/test_embedding_cache_index_alignment.py
Normal 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]
|
||||
Loading…
Add table
Reference in a new issue