From 105f37df10fd1edd6c1acfe97baeac09269cfa26 Mon Sep 17 00:00:00 2001 From: Andreas Frings Date: Tue, 2 Jun 2026 07:29:54 +0200 Subject: [PATCH] test: clean up comments and rename file-31 test --- .../caching/test_partial_cache_merge.py | 32 +++++++------------ 1 file changed, 11 insertions(+), 21 deletions(-) diff --git a/tests/test_litellm/caching/test_partial_cache_merge.py b/tests/test_litellm/caching/test_partial_cache_merge.py index 94f64209043..3a1d9a10366 100644 --- a/tests/test_litellm/caching/test_partial_cache_merge.py +++ b/tests/test_litellm/caching/test_partial_cache_merge.py @@ -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))