Files
lmcache--lmcache/tests/v1/distributed/test_fault_inject_l2_adapter.py
2026-07-13 12:24:33 +08:00

257 lines
8.6 KiB
Python

# SPDX-License-Identifier: Apache-2.0
"""Unit tests for FaultInjectL2Adapter.
The adapter is a decorator that drops a deterministic key subset at load to
simulate partial L2 retrieve failures. These tests wrap a real MockL2Adapter
and assert the load-result bitmap has exactly the expected bits cleared while
the lookup bitmap is left intact, using only public interface methods.
"""
# Standard
import select
import time
# Third Party
import pytest
import torch
# First Party
from lmcache.v1.distributed.api import MemoryLayoutDesc, ObjectKey
from lmcache.v1.distributed.l2_adapters.fault_inject_l2_adapter import (
FaultInjectL2Adapter,
)
from lmcache.v1.distributed.l2_adapters.mock_l2_adapter import (
MockL2Adapter,
MockL2AdapterConfig,
)
from lmcache.v1.memory_management import (
MemoryFormat,
MemoryObjMetadata,
TensorMemoryObj,
)
from lmcache.v1.platform import consume_fd
_EMPTY_LAYOUT = MemoryLayoutDesc(shapes=[], dtypes=[])
N_KEYS = 8
def _object_key(chunk_id: int) -> ObjectKey:
"""Build an ObjectKey from a chunk id (test fixture)."""
return ObjectKey(
chunk_hash=ObjectKey.IntHash2Bytes(chunk_id),
model_name="test_model",
kv_rank=0,
)
def _memory_obj(size: int = 256, fill_value: float = 1.0) -> TensorMemoryObj:
"""Build a filled TensorMemoryObj for use as a store/load buffer."""
raw = torch.empty(size, dtype=torch.float32)
raw.fill_(fill_value)
meta = MemoryObjMetadata(
shape=torch.Size([size]),
dtype=torch.float32,
address=0,
phy_size=size * 4,
fmt=MemoryFormat.KV_2LTD,
ref_count=1,
)
return TensorMemoryObj(raw, meta, parent_allocator=None)
def _wait_fd(event_fd: int, timeout: float = 5.0) -> bool:
"""Wait up to ``timeout`` seconds for ``event_fd`` to signal, draining it.
Returns True if the fd became readable before the timeout, False otherwise.
"""
poll = select.poll()
poll.register(event_fd, select.POLLIN)
if poll.poll(timeout * 1000):
try:
consume_fd(event_fd)
except BlockingIOError:
pass
return True
return False
def _make_adapter(rate: float = 0.0, seed: int = 0, gap_indices=(), gap_tail_ratios=()):
"""Build a FaultInjectL2Adapter wrapping a fresh MockL2Adapter.
Returns the ``(wrapper, inner)`` pair.
"""
inner = MockL2Adapter(MockL2AdapterConfig(max_size_gb=0.01, mock_bandwidth_gb=10.0))
wrapper = FaultInjectL2Adapter(
inner,
rate=rate,
seed=seed,
gap_indices=tuple(gap_indices),
gap_tail_ratios=tuple(gap_tail_ratios),
)
return wrapper, inner
def _store_all(adapter, keys):
"""Store one memory object per key and drain the completion event."""
fd = adapter.get_store_event_fd()
adapter.submit_store_task(keys, [_memory_obj() for _ in keys])
assert _wait_fd(fd)
adapter.pop_completed_store_tasks()
def _lookup_bitmap(adapter, keys):
"""Run a lookup-and-lock for ``keys`` and return its result bitmap."""
fd = adapter.get_lookup_and_lock_event_fd()
tid = adapter.submit_lookup_and_lock_task(keys, _EMPTY_LAYOUT)
assert _wait_fd(fd)
# query_*_result is non-idempotent (returns non-None once); poll briefly.
for _ in range(50):
bm = adapter.query_lookup_and_lock_result(tid)
if bm is not None:
return bm
time.sleep(0.01)
raise AssertionError("lookup result never ready")
def _load_bitmap(adapter, keys):
"""Run a load for ``keys`` and return its result bitmap."""
fd = adapter.get_load_event_fd()
tid = adapter.submit_load_task(keys, [_memory_obj() for _ in keys])
assert _wait_fd(fd)
for _ in range(50):
bm = adapter.query_load_result(tid)
if bm is not None:
return bm
time.sleep(0.01)
raise AssertionError("load result never ready")
# =============================================================================
# Pass-through (rate=0): no faults.
# =============================================================================
def test_rate_zero_is_passthrough():
"""rate=0 passes every key through unchanged (no drops)."""
adapter, inner = _make_adapter(rate=0.0)
try:
keys = [_object_key(i) for i in range(N_KEYS)]
_store_all(adapter, keys)
lookup = _lookup_bitmap(adapter, keys)
load = _load_bitmap(adapter, keys)
assert lookup.popcount() == N_KEYS
assert load.popcount() == N_KEYS
finally:
adapter.close()
# =============================================================================
# gap_indices: exact, deterministic drops.
# =============================================================================
def test_load_gap_cleared_lookup_intact():
"""A gapped load clears exactly those load bits, leaving lookup intact."""
gap = {1, 4, 6}
adapter, inner = _make_adapter(gap_indices=gap)
try:
keys = [_object_key(i) for i in range(N_KEYS)]
_store_all(adapter, keys)
# Lookup is never faulted (lookup says present).
lookup = _lookup_bitmap(adapter, keys)
assert lookup.popcount() == N_KEYS
load = _load_bitmap(adapter, keys)
for i in range(N_KEYS):
assert load.test(i) == (i not in gap)
finally:
adapter.close()
# =============================================================================
# gap_tail_ratios: distance-from-tail / load-length; workload-agnostic.
# =============================================================================
def test_gap_tail_ratio_position_and_scaling():
"""A tail-ratio drops the chunk at round((1-ratio)*(n-1)); the absolute
index is computed from the load batch, so it scales with length -- no
content or position is known in advance."""
keys = [_object_key(i) for i in range(N_KEYS)]
adapter, inner = _make_adapter(gap_tail_ratios=(0.5,))
try:
_store_all(adapter, keys)
_lookup_bitmap(adapter, keys)
mid = round(0.5 * (N_KEYS - 1))
load = _load_bitmap(adapter, keys)
assert not load.test(mid)
assert all(load.test(i) for i in range(N_KEYS) if i != mid)
# Same ratio on a SHORTER load drops a proportionally-different index
# (relative to that batch's tail) -- self-scaling, workload-agnostic.
short = keys[:4]
load2 = _load_bitmap(adapter, short)
assert not load2.test(round(0.5 * (len(short) - 1)))
finally:
adapter.close()
def test_gap_tail_ratio_endpoints():
"""ratio=0.0 drops the last chunk; ratio=1.0 drops the first."""
keys = [_object_key(i) for i in range(N_KEYS)]
adapter, inner = _make_adapter(gap_tail_ratios=(0.0, 1.0))
try:
_store_all(adapter, keys)
_lookup_bitmap(adapter, keys)
load = _load_bitmap(adapter, keys)
assert not load.test(0) and not load.test(N_KEYS - 1)
assert all(load.test(i) for i in range(1, N_KEYS - 1))
finally:
adapter.close()
# =============================================================================
# rate-based drops: deterministic across instances with the same seed.
# =============================================================================
def test_rate_drop_is_deterministic_across_instances():
"""Same seed drops the same keys; a different seed (very likely) differs."""
keys = [_object_key(i) for i in range(64)]
def dropped_positions(seed):
adapter, inner = _make_adapter(rate=0.3, seed=seed)
try:
_store_all(adapter, keys)
_lookup_bitmap(adapter, keys)
load = _load_bitmap(adapter, keys)
return {i for i in range(len(keys)) if not load.test(i)}
finally:
adapter.close()
a = dropped_positions(seed=42)
b = dropped_positions(seed=42)
c = dropped_positions(seed=1234)
assert a, "rate=0.3 should drop at least one of 64 keys"
assert a == b, "same seed must drop the same keys"
# A different seed should (with overwhelming probability) differ.
assert a != c
def test_rate_drop_within_tolerance():
"""Rate-based bucketing drops roughly ``rate * N`` keys (wide tolerance)."""
keys = [_object_key(i) for i in range(200)]
adapter, inner = _make_adapter(rate=0.25, seed=7)
try:
_store_all(adapter, keys)
_lookup_bitmap(adapter, keys)
load = _load_bitmap(adapter, keys)
dropped = sum(1 for i in range(len(keys)) if not load.test(i))
# Deterministic hash bucketing -> roughly rate * N, allow a wide band.
assert 25 <= dropped <= 75
finally:
adapter.close()
if __name__ == "__main__":
pytest.main([__file__, "-v"])