diff --git a/daser/connector/worker.py b/daser/connector/worker.py index 78143e1..a3d086a 100644 --- a/daser/connector/worker.py +++ b/daser/connector/worker.py @@ -1885,8 +1885,8 @@ def _submit_finished_save(self, save: _DeferredFinishedSave) -> Any | None: save: Deferred store plan built during ``wait_for_save``. Returns: - Future that completes after all store batches are committed, or - ``None`` when no store batch could be staged. + ``None`` once request KV has been copied into staging and background + store/commit futures have been tracked. Async/thread-safety: Called by vLLM's worker thread from ``get_finished`` while vLLM is @@ -1917,7 +1917,7 @@ def _submit_finished_save(self, save: _DeferredFinishedSave) -> Any | None: self._commit_after_store_futures(batch_futures, sorted(save.commit_keys)), ) self._track_save_future(commit_future, 0, None) - return commit_future + return None async def _commit_after_store_futures( self, diff --git a/daser/replacement/__init__.py b/daser/replacement/__init__.py index 2ed8753..363aabc 100644 --- a/daser/replacement/__init__.py +++ b/daser/replacement/__init__.py @@ -2,5 +2,10 @@ from daser.replacement.base import ReplacementPolicy from daser.replacement.lru import LRUReplacementPolicy +from daser.replacement.prefix_aware_lru import PrefixAwareLRUReplacementPolicy -__all__ = ["LRUReplacementPolicy", "ReplacementPolicy"] +__all__ = [ + "LRUReplacementPolicy", + "PrefixAwareLRUReplacementPolicy", + "ReplacementPolicy", +] diff --git a/daser/replacement/prefix_aware_lru.py b/daser/replacement/prefix_aware_lru.py new file mode 100644 index 0000000..f207261 --- /dev/null +++ b/daser/replacement/prefix_aware_lru.py @@ -0,0 +1,166 @@ +# SPDX-License-Identifier: Apache-2.0 + +"""Prefix-aware LRU replacement policy.""" + +# Standard +from collections.abc import Callable, Hashable, Iterable +from dataclasses import dataclass +from heapq import heappop, heappush +from typing import Generic, TypeVar + +# First Party +from daser.replacement.base import ReplacementPolicy + +K = TypeVar("K") + + +@dataclass(frozen=True, order=True) +class _Order: + epoch: int + prefix_rank: int + sequence: int + + +class PrefixAwareLRUReplacementPolicy(ReplacementPolicy[K], Generic[K]): + """LRU policy that ages request suffix keys before prefix keys. + + Async/thread-safety: + Not thread-safe by itself. Owners should call it from one event loop or + guard it with their own lock. + """ + + def __init__(self) -> None: + self._clock = 0 + self._sequence = 0 + self._orders: dict[K, _Order] = {} + self._heap: list[tuple[int, int, int, K]] = [] + self._group_epochs: dict[Hashable, int] = {} + self._key_groups: dict[K, Hashable] = {} + self._group_counts: dict[Hashable, int] = {} + + def insert(self, key: K) -> None: + """Insert or refresh a standalone key. + + Args: + key: Cache key to track. + """ + self._set_order(key, self._next_epoch(), 0, None) + + def insert_prefix(self, key: K, group: Hashable, prefix_index: int) -> None: + """Insert or refresh one key using prefix-aware group ordering. + + Args: + key: Cache key to track. + group: Stable request/chunk identity shared by sibling prefix keys. + prefix_index: Zero-based slot index inside the request prefix. + """ + epoch = self._group_epochs.get(group) + if epoch is None: + epoch = self._next_epoch() + self._group_epochs[group] = epoch + self._set_order(key, epoch, -int(prefix_index), group) + + def access(self, key: K) -> None: + """Mark an existing key as recently used. + + Args: + key: Cache key that was read or written. + """ + if key in self._orders: + self._set_order(key, self._next_epoch(), 0, None) + + def access_prefix(self, keys: Iterable[tuple[K, int]]) -> None: + """Refresh sibling prefix keys with suffix-before-prefix aging. + + Args: + keys: ``(key, prefix_index)`` pairs in any order. + """ + present = [ + (key, int(prefix_index)) + for key, prefix_index in keys + if key in self._orders + ] + if not present: + return + epoch = self._next_epoch() + for key, prefix_index in present: + self._set_order(key, epoch, -prefix_index, None) + + def remove(self, key: K) -> None: + """Remove a key from replacement tracking. + + Args: + key: Cache key to forget. + """ + if key not in self._orders: + return + self._orders.pop(key, None) + self._drop_key_group(key) + + def evict(self) -> K | None: + """Return and remove the least-recently-used key. + + Returns: + Victim key, or None when empty. + """ + while self._heap: + epoch, prefix_rank, sequence, key = heappop(self._heap) + order = self._orders.get(key) + if order != _Order(epoch, prefix_rank, sequence): + continue + self.remove(key) + return key + return None + + def evict_matching(self, predicate: Callable[[K], bool]) -> K | None: + """Return and remove the oldest key accepted by ``predicate``. + + Args: + predicate: Function returning True for eligible victim keys. + + Returns: + Victim key, or None when no tracked key matches. + """ + victims = [ + (order, key) for key, order in self._orders.items() if predicate(key) + ] + victim = min(victims) if victims else None + if victim is None: + return None + _order, key = victim + self.remove(key) + return key + + def _next_epoch(self) -> int: + """Return the next recency epoch.""" + self._clock += 1 + return self._clock + + def _set_order( + self, + key: K, + epoch: int, + prefix_rank: int, + group: Hashable | None, + ) -> None: + """Install a key order and push a lazy heap record.""" + self._drop_key_group(key) + self._sequence += 1 + order = _Order(epoch, prefix_rank, self._sequence) + self._orders[key] = order + if group is not None: + self._key_groups[key] = group + self._group_counts[group] = self._group_counts.get(group, 0) + 1 + heappush(self._heap, (order.epoch, order.prefix_rank, order.sequence, key)) + + def _drop_key_group(self, key: K) -> None: + """Remove a key from its prefix group accounting.""" + group = self._key_groups.pop(key, None) + if group is None: + return + count = self._group_counts.get(group, 0) - 1 + if count > 0: + self._group_counts[group] = count + return + self._group_counts.pop(group, None) + self._group_epochs.pop(group, None) diff --git a/daser/transfer/iouring/l1_cache.py b/daser/transfer/iouring/l1_cache.py index e7254b0..12b5507 100644 --- a/daser/transfer/iouring/l1_cache.py +++ b/daser/transfer/iouring/l1_cache.py @@ -17,7 +17,7 @@ from dataclasses import dataclass # First Party -from daser.replacement import LRUReplacementPolicy +from daser.replacement import PrefixAwareLRUReplacementPolicy from daser.transfer.iouring.pinned_pool import PinnedMemoryPool, PinnedMemorySlice @@ -59,7 +59,8 @@ def __init__( self._starts: list[int] = [] self._by_start: dict[int, tuple[int, int]] = {} self._used = 0 - self._policy = LRUReplacementPolicy[tuple[int, int]]() + self._policy = PrefixAwareLRUReplacementPolicy[tuple[int, int]]() + self._headroom_keys: set[tuple[int, int]] = set() self._pool_waiters: list[object] = [] self._is_pinned = pinned_predicate @@ -171,8 +172,14 @@ def record_hits(self, hits: list[L1RangeHit]) -> None: Args: hits: slices returned by ``resolve_subranges``. """ - for hit in hits: - self._policy.access(hit.key) + self._policy.access_prefix( + (hit.key, hit.key[0] + hit.source_offset) for hit in hits + ) + for hit in sorted( + hits, + key=lambda item: item.key[0] + item.source_offset, + reverse=True, + ): self._entries.move_to_end(hit.key) def touch(self, key: tuple[int, int]) -> None: @@ -190,6 +197,16 @@ def put(self, key: tuple[int, int], data: PinnedMemorySlice) -> None: self.drop_overlapping(key[0], key[1]) self._insert_entry(key, data) + def put_headroom(self, key: tuple[int, int], data: PinnedMemorySlice) -> None: + """Insert promoted bytes into the reserved L1 headroom. + + Args: + key: ``(file_offset, nbytes)`` range key. + data: pinned slice holding the range's bytes. + """ + self.drop_overlapping(key[0], key[1]) + self._insert_entry(key, data, headroom=True) + def reserve( self, key: tuple[int, int], @@ -197,6 +214,8 @@ def reserve( *, drop_overlaps: bool = True, preserve_overlaps: bool = False, + target_used_bytes: int | None = None, + prefer_headroom_victims: bool = False, ) -> PinnedMemorySlice | None: """Try to reserve pinned space for a store or promoted load. @@ -206,6 +225,12 @@ def reserve( drop_overlaps: drop resident ranges overlapping ``key`` first. preserve_overlaps: keep the non-overlapping remainder of dropped ranges when ``drop_overlaps`` is set. + target_used_bytes: optional resident-byte target after insertion. + When set, old residents are evicted before allocation until + ``bytes_used + nbytes`` fits the target or no victim is + available. The full L1 capacity remains the hard limit. + prefer_headroom_victims: evict promoted headroom entries before + falling back to the global replacement order. Returns: A pinned slice, or None when the pool is exhausted and no further @@ -221,16 +246,16 @@ def reserve( ) if drop_overlaps: self.drop_overlapping(key[0], key[1], preserve_remainder=preserve_overlaps) + target = self._l1_bytes + if target_used_bytes is not None: + target = min(self._l1_bytes, max(nbytes, target_used_bytes)) + while self._used + nbytes > target: + if not self._evict_next(prefer_headroom=prefer_headroom_victims): + return None data = self._pool.allocate(nbytes) while data is None: - victim = self._policy.evict() - if victim is None: + if not self._evict_next(prefer_headroom=prefer_headroom_victims): return None - removed = self._entries.pop(victim, None) - self._remove_index(victim) - if removed is not None: - self._used -= len(removed) - self.release(victim, removed) data = self._pool.allocate(nbytes) return data @@ -274,6 +299,28 @@ def release(self, key: tuple[int, int], data: PinnedMemorySlice) -> None: data.close() self.notify_pool_waiters() + def trim_to( + self, + target_used_bytes: int, + *, + prefer_headroom_victims: bool = False, + ) -> bool: + """Evict residents until usage fits ``target_used_bytes``. + + Args: + target_used_bytes: Desired resident-byte ceiling. + prefer_headroom_victims: Evict promoted headroom entries before + falling back to the global replacement order. + + Returns: + True when usage fits the target, False when no victim was available. + """ + target = min(self._l1_bytes, max(0, target_used_bytes)) + while self._used > target: + if not self._evict_next(prefer_headroom=prefer_headroom_victims): + return False + return True + def register_pool_waiter(self, waiter: object) -> None: """Register a future to wake when pool space or metadata changes.""" self._pool_waiters.append(waiter) @@ -311,6 +358,7 @@ def drop_overlapping( removed = self._entries.pop(victim, None) self._remove_index(victim) self._policy.remove(victim) + self._headroom_keys.discard(victim) preserved = ( self._preserve_non_overlapping(victim, removed, file_offset, end) if preserve_remainder and removed is not None @@ -322,7 +370,13 @@ def drop_overlapping( for preserved_key, payload in preserved: self._put_preserved_fragment(preserved_key, payload) - def _insert_entry(self, key: tuple[int, int], data: PinnedMemorySlice) -> None: + def _insert_entry( + self, + key: tuple[int, int], + data: PinnedMemorySlice, + *, + headroom: bool = False, + ) -> None: """Insert one non-overlapping entry and enforce capacity.""" if len(data) > self._l1_bytes: return @@ -330,17 +384,34 @@ def _insert_entry(self, key: tuple[int, int], data: PinnedMemorySlice) -> None: self._insert_index(key) self._entries.move_to_end(key) self._policy.insert(key) + if headroom: + self._headroom_keys.add(key) + else: + self._headroom_keys.discard(key) self._used += len(data) self.notify_pool_waiters() while self._used > self._l1_bytes: - victim = self._policy.evict() - if victim is None: + if not self._evict_next(): break - removed = self._entries.pop(victim, None) - self._remove_index(victim) - if removed is not None: - self._used -= len(removed) - self.release(victim, removed) + + def _evict_next(self, *, prefer_headroom: bool = False) -> bool: + """Evict one policy-selected resident entry.""" + victim = ( + self._policy.evict_matching(lambda key: key in self._headroom_keys) + if prefer_headroom + else None + ) + if victim is None: + victim = self._policy.evict() + if victim is None: + return False + removed = self._entries.pop(victim, None) + self._remove_index(victim) + self._headroom_keys.discard(victim) + if removed is not None: + self._used -= len(removed) + self.release(victim, removed) + return True def _preserve_non_overlapping( self, diff --git a/daser/transfer/iouring/layer.py b/daser/transfer/iouring/layer.py index 0d79fd4..26a1dbd 100644 --- a/daser/transfer/iouring/layer.py +++ b/daser/transfer/iouring/layer.py @@ -2,6 +2,7 @@ # Standard import asyncio +from dataclasses import dataclass from typing import Any # First Party @@ -16,6 +17,20 @@ logger = init_logger(__name__) _DIRECT_IO_ALIGNMENT = 4096 +_PROMOTION_HEADROOM_RATIO = 20 +_GROUPED_L2_WRITE_BATCH_BYTES = 64 * 1024 * 1024 + + +@dataclass +class _PendingL2WriteBatch: + """Mutable L2 batch assembled while grouped store inserts L1 entries.""" + + entries: list[tuple[tuple[int, int], int, PinnedMemorySlice]] + previous: list[asyncio.Task[None]] + ready: asyncio.Event + task: asyncio.Task[None] + end_offset: int + nbytes: int class TieredIOUringTransferLayer(TransferLayer): @@ -59,6 +74,15 @@ def __init__( if not skip_l2: self._l2 = L2IoEngine(path, l2_bytes, io_workers) self._l1_bytes = l1_bytes + headroom = max( + _DIRECT_IO_ALIGNMENT, + (l1_bytes + _PROMOTION_HEADROOM_RATIO - 1) // _PROMOTION_HEADROOM_RATIO, + ) + headroom = min(l1_bytes, self._align_direct_io(headroom)) + self._save_high_water_bytes = min( + l1_bytes, + max(0, l1_bytes - headroom), + ) self._l2_bytes = l2_bytes self._pending_l2: dict[tuple[int, int], asyncio.Task[None]] = {} self._pending_l2_buffers: dict[tuple[int, int], PinnedMemorySlice] = {} @@ -251,25 +275,7 @@ async def store_bytes(self, src: Any, file_offset: int, nbytes: int) -> int: self._l1.put(key, data) return nbytes - data = await self._reserve_l1_buffer(key, nbytes) - try: - self._copy_src_to_pinned(src, data, nbytes) - except BaseException: - data.close() - raise - async with self._lock: - self._raise_l2_error_locked() - previous = self._find_pending_l2_locked(file_offset, nbytes) - self._l1.put(key, data) - task = self._schedule_l2_write_locked( - key, - file_offset, - data, - previous, - ) - self._pending_l2[key] = task - self._pending_l2_buffers[key] = data - return nbytes + return await self._store_bytes_tiered(src, file_offset, nbytes) async def store_bytes_grouped( self, @@ -295,13 +301,65 @@ async def store_bytes_grouped( if self._l2 is None: return await self._store_bytes_grouped_l1_only(src, spans) + return await self._store_bytes_grouped_tiered(src, spans) + + async def _store_bytes_grouped_tiered( + self, + src: Any, + spans: list[dict[str, Any]], + ) -> int: + """Store grouped spans through L1 and schedule L2 persistence.""" total = 0 - for span in spans: - source_offset = int(span.get("source_offset", 0)) - nbytes = int(span["nbytes"]) - file_offset = int(span["file_offset"]) - source = self._slice_src(src, source_offset, nbytes) - total += await self.store_bytes(source, file_offset, nbytes) + max_span_bytes = 0 + current_batch: _PendingL2WriteBatch | None = None + try: + for span in spans: + source_offset = int(span.get("source_offset", 0)) + nbytes = int(span["nbytes"]) + file_offset = int(span["file_offset"]) + self._check_range(file_offset, nbytes) + current_batch = await self._seal_grouped_l2_batch_before_reserve( + current_batch, + file_offset, + nbytes, + ) + key = (file_offset, nbytes) + source = self._slice_src(src, source_offset, nbytes) + data = await self._reserve_l1_buffer(key, nbytes) + try: + self._copy_src_to_pinned(source, data, nbytes) + except BaseException: + data.close() + raise + async with self._lock: + self._raise_l2_error_locked() + previous = self._find_pending_l2_locked(file_offset, nbytes) + self._l1.put(key, data) + current_batch = self._append_grouped_l2_write_locked( + current_batch, + key, + file_offset, + data, + previous, + ) + self._pending_l2[key] = current_batch.task + self._pending_l2_buffers[key] = data + total += nbytes + max_span_bytes = max(max_span_bytes, nbytes) + except BaseException: + if current_batch is not None: + async with self._lock: + current_batch.ready.set() + raise + if total: + async with self._lock: + self._raise_l2_error_locked() + if current_batch is not None: + current_batch.ready.set() + self._l1.trim_to( + max(max_span_bytes, self._save_high_water_bytes), + prefer_headroom_victims=True, + ) return total async def drain(self) -> None: @@ -334,6 +392,38 @@ def close(self) -> None: self._l2.close() self._l1.close() + async def _store_bytes_tiered( + self, + src: Any, + file_offset: int, + nbytes: int, + ) -> int: + """Store one span through L1 and schedule L2 persistence.""" + key = (file_offset, nbytes) + data = await self._reserve_l1_buffer(key, nbytes) + try: + self._copy_src_to_pinned(src, data, nbytes) + except BaseException: + data.close() + raise + async with self._lock: + self._raise_l2_error_locked() + previous = self._find_pending_l2_locked(file_offset, nbytes) + self._l1.put(key, data) + task = self._schedule_l2_write_locked( + key, + file_offset, + data, + previous, + ) + self._pending_l2[key] = task + self._pending_l2_buffers[key] = data + self._l1.trim_to( + max(nbytes, self._save_high_water_bytes), + prefer_headroom_victims=True, + ) + return nbytes + def _is_pinned_by_l2(self, key: tuple[int, int], data: PinnedMemorySlice) -> bool: """Return whether ``data`` is still owned by an in-flight L2 write.""" return self._pending_l2_buffers.get(key) is data @@ -375,6 +465,14 @@ def _check_range(self, file_offset: int, nbytes: int) -> None: f"{_DIRECT_IO_ALIGNMENT} bytes: offset={file_offset} nbytes={nbytes}" ) + def _align_direct_io(self, value: int) -> int: + """Round ``value`` up to the direct-IO alignment.""" + return ( + (value + _DIRECT_IO_ALIGNMENT - 1) + // _DIRECT_IO_ALIGNMENT + * _DIRECT_IO_ALIGNMENT + ) + def _find_pending_l2_locked( self, file_offset: int, @@ -412,6 +510,15 @@ def _write_l2( if written != len(data): raise IOError(f"short io_uring write: {written} != {len(data)}") + def _write_l2_batch( + self, + entries: list[tuple[tuple[int, int], int, PinnedMemorySlice]], + uring: NativeIOUring, + ) -> None: + """Blocking io_uring L2 writes for one grouped-store batch.""" + for _key, file_offset, data in entries: + self._write_l2(file_offset, data, uring) + def _schedule_l2_write_locked( self, key: tuple[int, int], @@ -438,6 +545,63 @@ def _schedule_l2_write_locked( ) return asyncio.create_task(self._track_l2_write(key, future)) + def _start_grouped_l2_write_batch_locked(self) -> _PendingL2WriteBatch: + """Start a grouped L2 write task that waits until its batch is sealed.""" + ready = asyncio.Event() + entries: list[tuple[tuple[int, int], int, PinnedMemorySlice]] = [] + previous: list[asyncio.Task[None]] = [] + task = asyncio.create_task( + self._write_l2_batch_when_ready(entries, previous, ready) + ) + return _PendingL2WriteBatch( + entries=entries, + previous=previous, + ready=ready, + task=task, + end_offset=0, + nbytes=0, + ) + + def _append_grouped_l2_write_locked( + self, + batch: _PendingL2WriteBatch | None, + key: tuple[int, int], + file_offset: int, + data: PinnedMemorySlice, + previous: list[asyncio.Task[None]], + ) -> _PendingL2WriteBatch: + """Append one slot write to the current grouped L2 batch.""" + if batch is None: + batch = self._start_grouped_l2_write_batch_locked() + batch.entries.append((key, file_offset, data)) + for task in previous: + if task is not batch.task and task not in batch.previous: + batch.previous.append(task) + batch.end_offset = file_offset + len(data) + batch.nbytes += len(data) + return batch + + async def _seal_grouped_l2_batch_before_reserve( + self, + batch: _PendingL2WriteBatch | None, + file_offset: int, + nbytes: int, + ) -> _PendingL2WriteBatch | None: + """Seal a grouped L2 batch before the next reserve can block on it.""" + if batch is None: + return None + should_seal = file_offset != batch.end_offset or batch.nbytes + nbytes > max( + nbytes, _GROUPED_L2_WRITE_BATCH_BYTES + ) + if not should_seal: + async with self._lock: + should_seal = self._l1.bytes_used + nbytes > self._l1_bytes + if not should_seal: + return batch + async with self._lock: + batch.ready.set() + return None + async def _track_l2_write( self, key: tuple[int, int], @@ -464,6 +628,46 @@ async def _track_l2_write( ): pending_buffer.close() + async def _write_l2_batch_when_ready( + self, + entries: list[tuple[tuple[int, int], int, PinnedMemorySlice]], + previous: list[asyncio.Task[None]], + ready: asyncio.Event, + ) -> None: + """Persist a sealed grouped-store L2 batch and publish per-slot completion.""" + current = asyncio.current_task() + try: + await ready.wait() + if previous: + await asyncio.gather(*previous) + loop = asyncio.get_event_loop() + if self._l2 is None: + raise RuntimeError("L2 writes are disabled when skip_l2 is true") + batch_entries = list(entries) + await loop.run_in_executor( + self._l2.executor, + self._write_l2_batch, + batch_entries, + self._next_uring(), + ) + async with self._lock: + self._stats.l2_writes += len(batch_entries) + except BaseException as exc: + async with self._lock: + self._l2_errors.append(exc) + raise + finally: + async with self._lock: + for key, _file_offset, _data in entries: + if self._pending_l2.get(key) is current: + self._pending_l2.pop(key, None) + pending_buffer = self._pending_l2_buffers.pop(key, None) + if ( + pending_buffer is not None + and self._l1.get(key) is not pending_buffer + ): + pending_buffer.close() + async def _write_l2_async( self, key: tuple[int, int], @@ -577,7 +781,7 @@ async def _load_l2_miss_batch( nbytes = int(span["nbytes"]) key = (int(span["file_offset"]), nbytes) self._stats.l2_reads += 1 - self._l1.put(key, pinned) + self._l1.put_headroom(key, pinned) promoted.add(id(pinned)) except BaseException: if future_to_read: @@ -690,12 +894,15 @@ async def _reserve_l1_buffer( self, key: tuple[int, int], nbytes: int, + *, + target_used_bytes: int | None = None, ) -> PinnedMemorySlice: """Reserve pinned L1 space, waiting for pending L2 victims if needed. Args: key: L1 byte-range key being inserted. nbytes: Number of logical bytes needed. + target_used_bytes: optional resident-byte target after insertion. Returns: A pinned slice leased from the preallocated pool. @@ -708,7 +915,12 @@ async def _reserve_l1_buffer( while True: async with self._lock: self._raise_l2_error_locked() - data = self._l1.reserve(key, nbytes) + data = self._l1.reserve( + key, + nbytes, + target_used_bytes=target_used_bytes, + prefer_headroom_victims=True, + ) if data is not None: return data # Pool exhausted with no evictable victim: an in-flight L2 diff --git a/tests/connector/test_daser_connector.py b/tests/connector/test_daser_connector.py index d347c37..dbf9e12 100644 --- a/tests/connector/test_daser_connector.py +++ b/tests/connector/test_daser_connector.py @@ -992,8 +992,8 @@ def test_wait_for_save_groups_prefix_slot_stores_by_base_request() -> None: assert connector.committed_after == [(1, ["stored-0", "stored-1"])] -def test_get_finished_holds_blocks_until_deferred_store_completes() -> None: - """Worker should not release finished request blocks before store is done.""" +def test_get_finished_releases_after_deferred_store_is_staged() -> None: + """Worker releases request blocks once save work owns staging buffers.""" connector = _FinishedSaveProbe() connector.seed_finished_save( "req", @@ -1023,13 +1023,14 @@ def submit_pending(coro): finished_sending, finished_recving = connector.get_finished({"req"}) assert finished_recving is None - assert finished_sending is None + assert finished_sending == {"req"} assert connector.staged_batches - assert "req" in connector.pending_finished_save_ids() + assert connector.pending_finished_save_ids() == set() + assert len(connector.tracked) == 2 -def test_get_finished_reports_completed_deferred_store_on_later_step() -> None: - """Completed saves should be reported even after the original finish step.""" +def test_get_finished_does_not_report_deferred_store_twice() -> None: + """Background save completion should not emit a second request release.""" connector = _FinishedSaveProbe() pending_future = None @@ -1066,11 +1067,11 @@ def submit_pending(coro): ) connector.set_submit_store_coroutine(submit_pending) - assert connector.get_finished({"req"}) == (None, None) + assert connector.get_finished({"req"}) == ({"req"}, None) assert pending_future is not None pending_future.complete = True - assert connector.get_finished(set()) == ({"req"}, None) + assert connector.get_finished(set()) == (None, None) def test_get_finished_releases_request_when_no_store_batch_can_be_staged() -> None: @@ -3933,7 +3934,7 @@ def test_build_load_read_plan_deduplicates_identical_source_reads(): def test_build_staging_store_batches_deduplicates_identical_chunk_writes(): - """Multiple store specs for one allocation should produce one write span.""" + """Multiple store specs for one allocation should produce one span set.""" reqs_to_store = { "r0": ReqStoreSpec("k0", 10, 2, [4, 5], 320, 8), "r1": ReqStoreSpec("k0", 10, 2, [8, 9], 320, 8), @@ -3948,7 +3949,9 @@ def test_build_staging_store_batches_deduplicates_identical_chunk_writes(): assert len(batches) == 1 block_ids, spans = batches[0] assert block_ids == [4, 5] - assert spans == [StoreWriteSpan(0, 64, 320, "k0", 10, 2)] + assert spans == [ + StoreWriteSpan(0, 64, 320, "k0", 10, 2), + ] def test_build_load_copy_runs_merges_same_transform_ranges(): diff --git a/tests/server/test_ipc_server.py b/tests/server/test_ipc_server.py index e783ecf..e4da3c6 100644 --- a/tests/server/test_ipc_server.py +++ b/tests/server/test_ipc_server.py @@ -432,6 +432,78 @@ async def test_transfer_store_and_load_with_bytes_payload(tmp_path) -> None: await server.stop() +@pytest.mark.asyncio +async def test_transfer_store_preserves_span_order_when_backend_disables_coalesce( + tmp_path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """IPC preserves backend order when a transfer disables store coalescing.""" + seen_spans: list[list[dict[str, Any]]] = [] + + class OrderedStoreTransfer(TransferLayer): + coalesce_store_spans = False + + def __init__(self, **_kwargs: Any) -> None: + pass + + async def store_bytes(self, src: Any, file_offset: int, nbytes: int) -> int: + return nbytes + + async def load_bytes(self, dst: Any, file_offset: int, nbytes: int) -> int: + return nbytes + + async def store_bytes_grouped( + self, + src: Any, + spans: list[dict[str, Any]], + ) -> int: + seen_spans.append([dict(span) for span in spans]) + return sum(int(span["nbytes"]) for span in spans) + + def close(self) -> None: + pass + + monkeypatch.setattr( + "daser.server.ipc.server.TieredIOUringTransferLayer", + OrderedStoreTransfer, + ) + + core = make_core() + socket_path = str(tmp_path / "test.sock") + server = IPCServer(socket_path, core, make_runtime_config(tmp_path)) + await server.start() + try: + store = await _send_recv( + socket_path, + { + "op": "transfer_store", + "payload": {"data": b"a" * SLOT_SIZE * 2}, + "spans": [ + { + "source_offset": SLOT_SIZE, + "nbytes": SLOT_SIZE, + "file_offset": SLOT_SIZE, + }, + {"source_offset": 0, "nbytes": SLOT_SIZE, "file_offset": 0}, + ], + }, + ) + finally: + await server.stop() + + assert store == {"ok": True, "bytes": SLOT_SIZE * 2, "chunk_keys": []} + assert seen_spans == [ + [ + { + "source_offset": SLOT_SIZE, + "nbytes": SLOT_SIZE, + "file_offset": SLOT_SIZE, + }, + {"source_offset": 0, "nbytes": SLOT_SIZE, "file_offset": 0}, + ] + ] + + @pytest.mark.asyncio async def test_transfer_store_skips_stale_chunk_span(tmp_path) -> None: """IPC store ignores delayed spans whose chunk allocation was evicted.""" diff --git a/tests/transfer/test_replacement.py b/tests/transfer/test_replacement.py index bb572f8..f99f1f2 100644 --- a/tests/transfer/test_replacement.py +++ b/tests/transfer/test_replacement.py @@ -2,6 +2,7 @@ # First Party from daser.replacement.lru import LRUReplacementPolicy +from daser.replacement.prefix_aware_lru import PrefixAwareLRUReplacementPolicy def test_lru_policy_evicts_least_recently_used_key() -> None: @@ -27,3 +28,30 @@ def test_lru_policy_remove_disables_future_eviction() -> None: assert policy.evict() == "b" assert policy.evict() is None + + +def test_prefix_aware_lru_evicts_request_suffix_before_prefix() -> None: + """Prefix-aware LRU keeps earlier request slots newer than suffix slots.""" + policy = PrefixAwareLRUReplacementPolicy[str]() + + for index, key in enumerate(("a", "b", "c")): + policy.insert_prefix(key, ("req", 0, 3), index) + + assert policy.evict() == "c" + assert policy.evict() == "b" + assert policy.evict() == "a" + + +def test_prefix_aware_lru_prefers_older_request_before_newer_suffix() -> None: + """Request recency remains the primary LRU order.""" + policy = PrefixAwareLRUReplacementPolicy[str]() + + for index, key in enumerate(("a0", "b0")): + policy.insert_prefix(key, ("req0", 0, 2), index) + for index, key in enumerate(("a1", "b1")): + policy.insert_prefix(key, ("req1", 2, 2), index) + + assert policy.evict() == "b0" + assert policy.evict() == "a0" + assert policy.evict() == "b1" + assert policy.evict() == "a1" diff --git a/tests/transfer/test_tiered_iouring_transfer.py b/tests/transfer/test_tiered_iouring_transfer.py index adecced..89ac3d1 100644 --- a/tests/transfer/test_tiered_iouring_transfer.py +++ b/tests/transfer/test_tiered_iouring_transfer.py @@ -14,6 +14,7 @@ # First Party import daser.transfer.iouring.native as native_iouring from daser.transfer.iouring.native import NativeIOUring +from daser.transfer.iouring.pinned_pool import PinnedMemorySlice ALIGNMENT = 4096 @@ -94,6 +95,29 @@ def _read_l2_into( return super()._read_l2_into(file_offset, dst, uring) +class L2WriteBatchProbe(TieredIOUringTransferLayer): + """Test transfer layer that records grouped L2 write batches.""" + + def __init__(self, path: str, l1_bytes: int, l2_bytes: int) -> None: + super().__init__( + path=path, + l1_bytes=l1_bytes, + l2_bytes=l2_bytes, + ) + self.l2_write_batches: list[list[tuple[int, int]]] = [] + + def _write_l2_batch( + self, + entries: list[tuple[tuple[int, int], int, PinnedMemorySlice]], + uring: NativeIOUring, + ) -> None: + """Record grouped L2 writes before delegating to production IO.""" + self.l2_write_batches.append( + [(file_offset, len(data)) for _key, file_offset, data in entries] + ) + super()._write_l2_batch(entries, uring) + + class DelayedL2ReadProbe(TieredIOUringTransferLayer): """Test transfer layer that pauses selected L2 reads.""" @@ -421,8 +445,8 @@ def test_iouring_grouped_load_batches_l1_hits(tmp_path) -> None: """Grouped L1 loads batch host-to-destination copies.""" layer = GroupedCopyProbe( path=str(tmp_path / "daser.store"), - l1_bytes=ALIGNMENT * 2, - l2_bytes=ALIGNMENT * 3, + l1_bytes=ALIGNMENT * 3, + l2_bytes=ALIGNMENT * 4, ) try: @@ -453,6 +477,201 @@ def test_iouring_grouped_load_batches_l1_hits(tmp_path) -> None: layer.close() +def test_iouring_store_admission_leaves_promotion_headroom(tmp_path) -> None: + """Tiered stores keep a small promotion headroom below full L1 capacity.""" + + async def scenario() -> None: + layer = TieredIOUringTransferLayer( + path=str(tmp_path / "daser.store"), + l1_bytes=ALIGNMENT * 20, + l2_bytes=ALIGNMENT * 24, + ) + try: + for idx in range(20): + await layer.store_bytes( + _block(bytes([idx])), + file_offset=idx * ALIGNMENT, + nbytes=ALIGNMENT, + ) + await layer.drain() + assert layer.l1_bytes_used == ALIGNMENT * 19 + finally: + layer.close() + + _run(scenario()) + + +def test_iouring_store_spans_are_coalesced_by_ipc() -> None: + """io_uring opts into IPC store coalescing for cold-path performance.""" + + assert TieredIOUringTransferLayer.coalesce_store_spans is True + + +def test_iouring_grouped_store_keeps_coalesced_span_resident(tmp_path) -> None: + """A contiguous physical store remains one L1 range when it fits capacity.""" + + async def scenario() -> None: + layer = L2ReadProbe( + path=str(tmp_path / "daser.store"), + l1_bytes=ALIGNMENT * 5, + l2_bytes=ALIGNMENT * 8, + ) + try: + src = b"".join(_block(byte) for byte in (b"a", b"b", b"c", b"d", b"e")) + await layer.store_bytes_grouped( + bytearray(src), + [ + { + "source_offset": 0, + "file_offset": 0, + "nbytes": ALIGNMENT * 5, + "chunk_key": "req", + "start_slot": 0, + "num_slots": 5, + } + ], + ) + await layer.drain() + + assert layer.l1_bytes_used == ALIGNMENT * 5 + dst = bytearray(ALIGNMENT * 5) + await layer.load_bytes(dst, 0, ALIGNMENT * 5) + assert bytes(dst) == bytes(src) + assert layer.l2_read_ranges == [] + finally: + layer.close() + + _run(scenario()) + + +def test_iouring_grouped_store_batches_adjacent_l2_writes(tmp_path) -> None: + """Grouped stores keep L1 slots separate while batching adjacent L2 writes.""" + + async def scenario() -> None: + layer = L2WriteBatchProbe( + path=str(tmp_path / "daser.store"), + l1_bytes=ALIGNMENT * 8, + l2_bytes=ALIGNMENT * 8, + ) + try: + src = b"".join(_block(byte) for byte in (b"a", b"b", b"c", b"d")) + stored = await layer.store_bytes_grouped( + bytearray(src), + [ + { + "source_offset": idx * ALIGNMENT, + "file_offset": idx * ALIGNMENT, + "nbytes": ALIGNMENT, + "chunk_key": "req", + "start_slot": 0, + "num_slots": 4, + } + for idx in range(4) + ], + ) + await layer.drain() + + assert stored == ALIGNMENT * 4 + assert layer.l2_write_batches == [ + [(idx * ALIGNMENT, ALIGNMENT) for idx in range(4)] + ] + + dst = bytearray(ALIGNMENT * 4) + await layer.load_bytes_grouped( + dst, + [ + { + "target_offset": idx * ALIGNMENT, + "file_offset": idx * ALIGNMENT, + "nbytes": ALIGNMENT, + } + for idx in range(4) + ], + ) + assert bytes(dst) == bytes(src) + finally: + layer.close() + + _run(scenario()) + + +def test_iouring_l2_promotion_uses_headroom_then_evicts(tmp_path) -> None: + """L2 misses promote into headroom and evict old residents once full.""" + + async def scenario() -> None: + path = str(tmp_path / "daser.store") + writer = TieredIOUringTransferLayer( + path=path, + l1_bytes=ALIGNMENT * 5, + l2_bytes=ALIGNMENT * 8, + ) + try: + for idx, byte in enumerate((b"a", b"b", b"c", b"d", b"e", b"f")): + await writer.store_bytes( + _block(byte), + file_offset=idx * ALIGNMENT, + nbytes=ALIGNMENT, + ) + await writer.drain() + finally: + writer.close() + + layer = L2ReadProbe( + path=path, + l1_bytes=ALIGNMENT * 5, + l2_bytes=ALIGNMENT * 8, + ) + try: + for idx, byte in enumerate((b"a", b"b", b"c", b"d")): + await layer.store_bytes( + _block(byte), + file_offset=idx * ALIGNMENT, + nbytes=ALIGNMENT, + ) + await layer.drain() + assert layer.l1_bytes_used == ALIGNMENT * 4 + + dst = bytearray(ALIGNMENT) + await layer.load_bytes(dst, file_offset=ALIGNMENT * 4, nbytes=ALIGNMENT) + assert bytes(dst) == bytes(_block(b"e")) + assert layer.l1_bytes_used == ALIGNMENT * 5 + assert layer.l2_read_ranges == [(ALIGNMENT * 4, ALIGNMENT)] + + await layer.load_bytes(dst, file_offset=ALIGNMENT * 5, nbytes=ALIGNMENT) + assert bytes(dst) == bytes(_block(b"f")) + assert layer.l1_bytes_used == ALIGNMENT * 5 + assert layer.l2_read_ranges == [ + (ALIGNMENT * 4, ALIGNMENT), + (ALIGNMENT * 5, ALIGNMENT), + ] + + await layer.load_bytes(dst, file_offset=ALIGNMENT, nbytes=ALIGNMENT) + assert bytes(dst) == bytes(_block(b"b")) + assert layer.l2_read_ranges == [ + (ALIGNMENT * 4, ALIGNMENT), + (ALIGNMENT * 5, ALIGNMENT), + ] + + await layer.load_bytes(dst, file_offset=0, nbytes=ALIGNMENT) + assert bytes(dst) == bytes(_block(b"a")) + assert layer.l2_read_ranges == [ + (ALIGNMENT * 4, ALIGNMENT), + (ALIGNMENT * 5, ALIGNMENT), + ] + + await layer.load_bytes(dst, file_offset=ALIGNMENT * 4, nbytes=ALIGNMENT) + assert bytes(dst) == bytes(_block(b"e")) + assert layer.l2_read_ranges == [ + (ALIGNMENT * 4, ALIGNMENT), + (ALIGNMENT * 5, ALIGNMENT), + (ALIGNMENT * 4, ALIGNMENT), + ] + finally: + layer.close() + + _run(scenario()) + + def test_iouring_grouped_load_supports_sliceable_cuda_wrapper(tmp_path) -> None: """Grouped L1 loads can target CUDA wrapper objects from IPC."""