mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(e2e): make concurrent replay consumption race-free
Greptile flagged that lazy per-slug pool initialization could double-build under concurrent replay calls, splitting consumption across a discarded pool. Pools are now built once at ReplaySource construction and per-key consumption is a single atomic deque pop, with a barrier-synchronized regression test that fails 10/10 under the lazy-init mutant
This commit is contained in:
parent
125587d286
commit
975e79dcef
2 changed files with 60 additions and 11 deletions
|
|
@ -391,18 +391,23 @@ def _miss_message(test_key: str, slug: str, canonical: CanonicalRequest, bundle:
|
|||
@dataclass(slots=True)
|
||||
class ReplaySource:
|
||||
"""One shared pool per test over a loaded bundle, so every client built in
|
||||
the session consumes the same recorded interactions. Calls match by
|
||||
canonical content key: order-independent across distinct keys (concurrent
|
||||
tests interleave calls nondeterministically), FIFO within one key (a poll
|
||||
loop replays its recorded responses in recorded order)."""
|
||||
the session consumes the same recorded interactions. Every pool is built
|
||||
once at construction and per-key consumption is a single atomic deque pop,
|
||||
so concurrent replay calls never race. Calls match by canonical content
|
||||
key: order-independent across distinct keys (concurrent tests interleave
|
||||
calls nondeterministically), FIFO within one key (a poll loop replays its
|
||||
recorded responses in recorded order)."""
|
||||
|
||||
bundle: LoadedBundle
|
||||
_pools: dict[str, dict[str, deque[Interaction]]] = field(default_factory=dict)
|
||||
_pools: dict[str, dict[str, deque[Interaction]]] = field(init=False)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self._pools = {
|
||||
slug: _build_pool(recorded) for slug, recorded in self.bundle.interactions.items()
|
||||
}
|
||||
|
||||
def _pool(self, slug: str) -> dict[str, deque[Interaction]]:
|
||||
if slug not in self._pools:
|
||||
self._pools[slug] = _build_pool(self.bundle.interactions.get(slug, ()))
|
||||
return self._pools[slug]
|
||||
return self._pools.get(slug, {})
|
||||
|
||||
def next_interaction(self, request: RecordedRequest) -> Interaction:
|
||||
test_key: Final = current_test_key()
|
||||
|
|
@ -412,12 +417,13 @@ class ReplaySource:
|
|||
queue: Final = pool.get(canonical.key)
|
||||
if queue is None:
|
||||
raise ReplayMiss(_miss_message(test_key, slug, canonical, self.bundle))
|
||||
if not queue:
|
||||
try:
|
||||
return queue.popleft()
|
||||
except IndexError:
|
||||
raise ReplayMiss(
|
||||
f"replay exhausted for {test_key}: every recorded interaction for key "
|
||||
f"{canonical.key} is already consumed; re-record with E2E_FIXTURE_MODE=record"
|
||||
)
|
||||
return queue.popleft()
|
||||
) from None
|
||||
|
||||
def leftover_error(self, test_key: str) -> str | None:
|
||||
"""Non-None when the test consumed fewer interactions than were recorded,
|
||||
|
|
|
|||
|
|
@ -14,6 +14,9 @@ pinned here too, including the stale message that names the bundle's age.
|
|||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import sys
|
||||
import threading
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from pathlib import Path
|
||||
|
|
@ -458,6 +461,46 @@ class TestReplayTransport:
|
|||
assert first == Success(status_code=200, data=Payload(value="first"))
|
||||
assert second == Success(status_code=200, data=Payload(value="second"))
|
||||
|
||||
def test_concurrent_replays_of_one_key_serve_each_recording_exactly_once(self, tmp_path: Path) -> None:
|
||||
"""A burst of parallel identical calls consumes one shared pool: no
|
||||
response duplicated, none forgotten, nothing left over at teardown.
|
||||
The tiny switch interval forces thread preemption inside pool setup
|
||||
and consumption, so a non-atomic pool build or pop fails this test."""
|
||||
root = tmp_path / "bundle"
|
||||
recorder = make_recorder(root)
|
||||
for ordinal in range(32):
|
||||
recorder.record(
|
||||
test_key=current_test_key(),
|
||||
request=recorded_request(
|
||||
"get", "/v1/models", headers=AuthHeaders(authorization="Bearer sk-x"), params=Query(q="all")
|
||||
),
|
||||
response=RecordedResult(kind="success", status_code=200, data={"value": f"v{ordinal:02d}"}),
|
||||
)
|
||||
source = replay_source(root)
|
||||
replay: Transport = ReplayTransport(source=source, master_key="sk-1234")
|
||||
barrier = threading.Barrier(8)
|
||||
|
||||
def consume_one() -> str:
|
||||
result = replay.get(
|
||||
"/v1/models", headers=replay.master, params=Query(q="all"), response_type=Payload
|
||||
)
|
||||
assert isinstance(result, Success)
|
||||
return result.data.value
|
||||
|
||||
def consume(_: int) -> tuple[str, ...]:
|
||||
barrier.wait()
|
||||
return tuple(consume_one() for _call in range(4))
|
||||
|
||||
previous_interval = sys.getswitchinterval()
|
||||
sys.setswitchinterval(1e-6)
|
||||
try:
|
||||
with ThreadPoolExecutor(max_workers=8) as executor:
|
||||
served = sorted(value for values in executor.map(consume, range(8)) for value in values)
|
||||
finally:
|
||||
sys.setswitchinterval(previous_interval)
|
||||
assert served == [f"v{ordinal:02d}" for ordinal in range(32)]
|
||||
assert source.leftover_error(current_test_key()) is None
|
||||
|
||||
|
||||
class TestRecordedKeySets:
|
||||
def test_two_separate_recordings_of_one_flow_produce_identical_key_sets(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue