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:
mateo-berri 2026-08-19 14:55:53 -07:00
parent 125587d286
commit 975e79dcef
2 changed files with 60 additions and 11 deletions

View file

@ -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,

View file

@ -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(