mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
test: clean up comments and rename file-31 test
This commit is contained in:
parent
d9a92227a5
commit
105f37df10
1 changed files with 11 additions and 21 deletions
|
|
@ -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))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue