test: clean up comments and rename file-31 test

This commit is contained in:
Andreas Frings 2026-06-02 07:29:54 +02:00
parent d9a92227a5
commit 105f37df10

View file

@ -1,15 +1,17 @@
"""
Regression tests for the partial-cache-hit embedding index merge bug.
Exercises LLMCachingHandler._combine_cached_embedding_response_with_api_result
directly — no provider, no Redis, no event loop needed.
Calls _combine_cached_embedding_response_with_api_result directly with
synthetic data without any provider, Redis, or event loop.
"""
import os
import sys
from datetime import datetime
sys.path.insert(0, os.path.abspath("../../.."))
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
from litellm.caching.caching_handler import (
CachingHandlerResponse,
@ -19,29 +21,21 @@ from litellm.types.utils import Embedding, EmbeddingResponse
def _emb(index: int, marker: int) -> Embedding:
"""Build an Embedding with a recognisable marker so we can tell items apart."""
return Embedding(embedding=[float(marker)], index=index, object="embedding")
def _emb_dict(index: int, marker: int) -> dict:
"""Dict variant — many providers (e.g. hosted_vllm) return dicts, not Embedding objects."""
"""Dict variant; many providers (e.g. hosted_vllm) return dicts, not Embedding objects."""
return {"embedding": [float(marker)], "index": index, "object": "embedding"}
def _run_merge(
batch_size: int, cache_positions: set[int], provider_as_dict: bool = False
) -> list[int]:
"""
Simulate what the handler builds at the cache-lookup step, then call the
merge function with a synthetic provider response. Returns the resulting
indices.
"""
cached_data = []
for pos in range(batch_size):
if pos in cache_positions:
cached_data.append(
_emb(index=pos, marker=100 + pos)
) # cache hits use position-based index
cached_data.append(_emb(index=pos, marker=100 + pos))
else:
cached_data.append(None)
cached_response = EmbeddingResponse(model="x", data=cached_data, object="list")
@ -66,19 +60,16 @@ def _run_merge(
def test_partial_cache_hit_two_in_the_middle():
"""File-31-style: 2 cache hits in the middle of a small batch."""
actual = _run_merge(batch_size=8, cache_positions={1, 4})
assert actual == [0, 1, 2, 3, 4, 5, 6, 7]
def test_partial_cache_hit_one_at_start():
"""One cache hit at the start shifts every following item by 1."""
actual = _run_merge(batch_size=5, cache_positions={0})
assert actual == [0, 1, 2, 3, 4]
def test_partial_cache_hit_no_duplicates():
"""All indices in the merged response must be unique."""
for hits in ({1}, {1, 4}, {0, 2, 5}, set(range(15))):
actual = _run_merge(batch_size=16, cache_positions=hits)
assert len(set(actual)) == len(
@ -87,21 +78,20 @@ def test_partial_cache_hit_no_duplicates():
def test_all_cache_hits_no_provider_call_needed():
"""Sanity: every position is a cache hit, indices are correct."""
actual = _run_merge(batch_size=4, cache_positions={0, 1, 2, 3})
assert actual == [0, 1, 2, 3]
def test_no_cache_hits_provider_only():
"""Sanity: nothing cached, indices are 0..N-1 from the provider directly."""
actual = _run_merge(batch_size=4, cache_positions=set())
assert actual == [0, 1, 2, 3]
def test_real_world_file_31_pattern():
def test_partial_cache_hit_large_batch():
"""
Exactly mirrors the observed production failure:
128-element batch, cache hits at positions 16 and 17 produce drift +2 from pos 18.
Regression test for the index drift bug: with hits at positions 16 and 17
in a 128-element batch, all provider items from position 18 onward received
an index 2 lower than their final position.
"""
actual = _run_merge(batch_size=128, cache_positions={16, 17})
assert actual == list(range(128))