diff --git a/benchmarks/bench_start_servers.py b/benchmarks/bench_start_servers.py index 9f5539c..7fbdff0 100644 --- a/benchmarks/bench_start_servers.py +++ b/benchmarks/bench_start_servers.py @@ -43,6 +43,7 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace: parser.add_argument("--block-size", type=int, default=BLOCK_TOKENS) parser.add_argument("--l1-size", type=parse_size_bytes, default="256gib") parser.add_argument("--l2-size", type=parse_size_bytes, default="300gib") + parser.add_argument("--daser-prefetch-max-requests", type=int, default=0) parser.add_argument( "--cache-reuse-mode", choices=("chunk", "prefix"), default="chunk" ) @@ -87,6 +88,7 @@ async def main_async(args: argparse.Namespace) -> None: skip_l2=args.skip_l2, tensor_parallel_size=args.tensor_parallel_size, trust_remote_code=args.trust_remote_code, + daser_prefetch_max_requests=args.daser_prefetch_max_requests, ) manifest = await manager.start() print(f"manifest={args.store_dir}/manifest.json") diff --git a/benchmarks/run_bench.py b/benchmarks/run_bench.py index 0828aea..98c84a8 100644 --- a/benchmarks/run_bench.py +++ b/benchmarks/run_bench.py @@ -81,6 +81,7 @@ class RunBenchArgs: bench_seed: vLLM bench random seed. bench_burstiness: vLLM bench burstiness factor. evict: Whether to enable L2 and eviction sizing. + daser_prefetch_max_requests: Maximum concurrent DaseR prefetches. prometheus_url: Optional Prometheus base URL for scrape diagnostics. Thread-safety: @@ -117,6 +118,7 @@ class RunBenchArgs: bench_seed: int = 42 bench_burstiness: float = 1.0 evict: bool = False + daser_prefetch_max_requests: int = 0 prometheus_url: str = "http://127.0.0.1:9090" @@ -188,6 +190,7 @@ def parse_args(argv: list[str] | None = None) -> RunBenchArgs: parser.add_argument("--bench-seed", type=int, default=42) parser.add_argument("--bench-burstiness", type=float, default=1.0) parser.add_argument("--evict", action="store_true") + parser.add_argument("--daser-prefetch-max-requests", type=int, default=0) parser.add_argument( "--prometheus-url", default="http://127.0.0.1:9090", @@ -228,6 +231,7 @@ def parse_args(argv: list[str] | None = None) -> RunBenchArgs: bench_seed=args.bench_seed, bench_burstiness=args.bench_burstiness, evict=args.evict, + daser_prefetch_max_requests=args.daser_prefetch_max_requests, prometheus_url=args.prometheus_url, ) try: @@ -538,6 +542,13 @@ def _start_command( "--l2-size", str(derived_l2), ] + if backend_run.backend == "daser" and args.daser_prefetch_max_requests: + command.extend( + [ + "--daser-prefetch-max-requests", + str(args.daser_prefetch_max_requests), + ] + ) if backend_run.backend == "daser": command.extend(["--cache-reuse-mode", backend_run.reuse_mode]) if args.trust_remote_code: diff --git a/benchmarks/utils/loadgen.py b/benchmarks/utils/loadgen.py index 5bafde9..2fa0d9b 100644 --- a/benchmarks/utils/loadgen.py +++ b/benchmarks/utils/loadgen.py @@ -25,6 +25,7 @@ from benchmarks.utils.servers import LMCACHE_HTTP_PORT, BenchmarkManifest _LMCACHE_QUIESCENCE_TIMEOUT_SECONDS = 600.0 +_DASER_DRAIN_TIMEOUT_SECONDS = 360.0 @dataclass @@ -211,6 +212,7 @@ async def run_daser_prefix( metrics=await collect_phase_metrics(manifest, before_warm), elapsed_ms=warm_elapsed_ms, ) + await _wait_daser_drained(manifest) return {"cold": cold_phase, "warm": warm_phase} @@ -264,6 +266,7 @@ async def run_lmcache( metrics=await collect_phase_metrics(manifest, before_warm), elapsed_ms=warm_elapsed_ms, ) + await _wait_daser_drained(manifest) return {"cold": cold_phase, "warm": warm_phase} @@ -515,7 +518,9 @@ async def _wait_daser_drained(manifest: BenchmarkManifest) -> None: daser = manifest.endpoints.get("daser") if daser is None: return - async with httpx.AsyncClient(timeout=httpx.Timeout(30.0)) as client: + async with httpx.AsyncClient( + timeout=httpx.Timeout(_DASER_DRAIN_TIMEOUT_SECONDS) + ) as client: response = await client.post(f"{daser.url}/drain") response.raise_for_status() diff --git a/benchmarks/utils/servers.py b/benchmarks/utils/servers.py index d30af4f..f6cf784 100644 --- a/benchmarks/utils/servers.py +++ b/benchmarks/utils/servers.py @@ -132,6 +132,7 @@ def __init__( skip_l2: bool = False, tensor_parallel_size: int = 1, trust_remote_code: bool = False, + daser_prefetch_max_requests: int = 0, ) -> None: """Initialize the service manager. @@ -156,6 +157,7 @@ def __init__( skip_l2: Disable L2 persistence/adapters for L1-only no-evict runs. tensor_parallel_size: vLLM tensor-parallel rank count. trust_remote_code: allow model/tokenizer repository Python code. + daser_prefetch_max_requests: maximum concurrent scheduler prefetches. """ if tensor_parallel_size <= 0: raise ValueError("tensor_parallel_size must be positive") @@ -179,6 +181,7 @@ def __init__( self.skip_l2 = skip_l2 self.tensor_parallel_size = tensor_parallel_size self.trust_remote_code = trust_remote_code + self.daser_prefetch_max_requests = daser_prefetch_max_requests self.log_dir = self.store_dir / "logs" self.pid_file = self.store_dir / "pids.json" self.socket_path = self.store_dir / "daser.sock" @@ -351,6 +354,7 @@ async def start_vllm_daser(self) -> None: "kv_connector_extra_config": { "socket_path": str(self.socket_path), "cache_reuse_mode": self.reuse_mode, + "prefetch_max_requests": self.daser_prefetch_max_requests, }, } await self._start_vllm("vllm_daser.log", kv_config) diff --git a/benchmarks/utils/vllm_bench.py b/benchmarks/utils/vllm_bench.py index c292ad9..956ad1a 100644 --- a/benchmarks/utils/vllm_bench.py +++ b/benchmarks/utils/vllm_bench.py @@ -24,6 +24,7 @@ slot_size_for_block_tokens, ) from benchmarks.utils.loadgen import ( + _DASER_DRAIN_TIMEOUT_SECONDS, _wait_lmcache_quiescent, backend_server_hit_rate, collect_phase_metrics, @@ -268,6 +269,8 @@ def run_load( backend_dir, run_command=run_command, ) + if backend_run.backend == "daser": + _drain_daser(manifest, print_kv=print_kv) result_path.write_text(json.dumps(result, indent=2), encoding="utf-8") return if backend_run.backend == "vllm": @@ -308,6 +311,8 @@ def run_load( warm_metrics, warm_hit_rate = _run_phase( args, manifest, warm_raw, run_command=run_command ) + if backend_run.backend == "daser": + _drain_daser(manifest, print_kv=print_kv) if backend_run.backend == "daser" and args.evict: _require_daser_evict_tier_activity(warm_metrics) cold_summary = _normalise_result(cold_raw) @@ -451,7 +456,7 @@ def _drain_daser( endpoint = manifest.endpoints.get("daser") if endpoint is None: return - response = httpx.post(f"{endpoint.url}/drain", timeout=30.0) + response = httpx.post(f"{endpoint.url}/drain", timeout=_DASER_DRAIN_TIMEOUT_SECONDS) response.raise_for_status() print_kv("daser_drain_status", "ok") diff --git a/daser/connector/daser_connector.py b/daser/connector/daser_connector.py index 44ef88d..350803d 100644 --- a/daser/connector/daser_connector.py +++ b/daser/connector/daser_connector.py @@ -96,6 +96,9 @@ def __init__( socket_path = str(extra.get("socket_path", "/tmp/daser.sock")) if role == KVConnectorRole.SCHEDULER: + prefetch_max_requests = int(extra.get("prefetch_max_requests", 0)) + if prefetch_max_requests < 0: + raise ValueError("prefetch_max_requests must be non-negative") self._request_lifecycle = RequestLifecycle( ipc_client=IPCClientSync(socket_path), block_tokens=16, @@ -103,6 +106,8 @@ def __init__( model_id="default", cache_reuse_mode=str(extra.get("cache_reuse_mode", "chunk")), runtime_config_ready=False, + socket_path=socket_path, + prefetch_max_requests=prefetch_max_requests, ) self._request_lifecycle.refresh_runtime_config() else: @@ -135,6 +140,22 @@ def __init__( logger.info("[CONNECTOR] role=%s socket=%s", role.name, socket_path) + def shutdown(self) -> None: + """Stop connector-side scheduler or worker resources. + + Async/thread-safety: + Called by vLLM during shutdown. Scheduler prefetch workers are + joined and worker transfer pipelines use their existing shutdown + path. + """ + lifecycle = getattr(self, "_request_lifecycle", None) + if lifecycle is not None: + lifecycle.shutdown() + return + runtime = getattr(self, "_worker_runtime", None) + if runtime is not None: + runtime.shutdown() + @property def prefer_cross_layer_blocks(self) -> bool: """Request vLLM cross-layer KV cache blocks for bulk chunk transfers. diff --git a/daser/connector/ipc_client.py b/daser/connector/ipc_client.py index 4fedc06..3ff3c2a 100644 --- a/daser/connector/ipc_client.py +++ b/daser/connector/ipc_client.py @@ -3,9 +3,10 @@ # Standard import asyncio import contextlib +from dataclasses import dataclass import socket import threading -from typing import Any +from typing import Any, Literal # First Party from daser.ipc_protocol import pack_frame, read_frame, recv_frame @@ -14,6 +15,15 @@ logger = init_logger(__name__) +@dataclass(frozen=True) +class PrefetchLookupResult: + """Validated scheduler lookup result with exact host-tier admission state.""" + + chunks: list[dict[str, Any]] + spans: list[dict[str, int]] + tier: Literal["l1", "mixed", "l2"] | None + + def _raise_on_error(result: dict[str, Any]) -> dict[str, Any]: """Raise when the server returned an error frame, else return the result. @@ -165,6 +175,72 @@ def lookup( resp = self.call(payload) return resp.get("chunks", []) + def lookup_with_prefetch( + self, + lease_id: str, + tokens: list[int], + model_id: str, + external_prefix_queries: int, + num_computed_tokens: int = 0, + ) -> PrefetchLookupResult: + """Look up chunks and atomically classify their exact transfer spans. + + Args: + lease_id: Base vLLM request ID used to retain an all-L1 result. + tokens: Prompt token IDs sent to the retrieval index. + model_id: Model identifier. + external_prefix_queries: vLLM external-prefix query token count. + num_computed_tokens: Tokens already resident in vLLM's local cache. + + Returns: + Validated chunks, physical transfer spans, and tier classification. + + Thread-safety: + Uses the lock-protected scheduler IPC connection. The server makes + all-L1 classification and lease acquisition atomic. + """ + response = self.call( + { + "op": "lookup_prefetch", + "lease_id": lease_id, + "tokens": tokens, + "model_id": model_id, + "external_prefix_queries": int(external_prefix_queries), + "num_computed_tokens": int(num_computed_tokens), + } + ) + chunks = response.get("chunks", []) + spans = response.get("spans", []) + tier = response.get("tier") + if not isinstance(chunks, list) or not all( + isinstance(chunk, dict) for chunk in chunks + ): + raise RuntimeError("[IPC] invalid lookup_prefetch chunks response") + if not isinstance(spans, list) or not all( + isinstance(span, dict) and "file_offset" in span and "nbytes" in span + for span in spans + ): + raise RuntimeError("[IPC] invalid lookup_prefetch spans response") + if tier not in (None, "l1", "mixed", "l2"): + raise RuntimeError("[IPC] invalid lookup_prefetch tier response") + if bool(spans) != (tier is not None): + raise RuntimeError("[IPC] incomplete lookup_prefetch response") + if not chunks and spans: + raise RuntimeError( + "[IPC] empty lookup_prefetch response has transfer state" + ) + return PrefetchLookupResult( + chunks=[dict(chunk) for chunk in chunks], + spans=[ + { + "file_offset": int(span["file_offset"]), + "nbytes": int(span["nbytes"]), + } + for span in spans + ], + tier=tier, + ) + def record_external_prefix_cache(self, queries: int, hits: int) -> None: """Record vLLM-equivalent external prefix cache counters. @@ -263,6 +339,48 @@ def transfer_drain(self) -> None: """ self.call({"op": "transfer_drain"}) + def transfer_prefetch( + self, + spans: list[dict[str, int]], + lease_id: str | None = None, + ) -> dict[str, int]: + """Synchronously promote storage spans into the host-memory tier. + + Args: + spans: Storage spans containing ``file_offset`` and ``nbytes``. + lease_id: Optional admitted request ID that retains promoted bytes. + + Returns: + Requested, L1-resident, and L2-read byte counts. + + Thread-safety: + Intended for a dedicated scheduler prefetch thread because the RPC + waits for all required L2 reads. + """ + payload: dict[str, Any] = {"op": "transfer_prefetch", "spans": spans} + if lease_id is not None: + payload["lease_id"] = lease_id + response = self.call(payload) + fields = ("requested_bytes", "l1_bytes", "l2_bytes") + if any(field not in response for field in fields): + raise RuntimeError("[IPC] invalid transfer_prefetch response") + return {field: int(response[field]) for field in fields} + + def release_transfer_lease(self, lease_id: str) -> None: + """Idempotently release server transfer bytes retained for a request. + + Args: + lease_id: Base vLLM request ID to clean up. + + Returns: + None. + + Thread-safety: + Uses the lock-protected scheduler IPC connection and may race with + worker load completion safely. + """ + self.call({"op": "release_transfer_lease", "lease_id": lease_id}) + def commit_stats(self) -> dict[str, int]: """Return server-side connector commit counters. @@ -445,13 +563,19 @@ async def init_transfer(self) -> None: await self.call({"op": "init_transfer"}) async def transfer_store_bytes( - self, data: bytes, spans: list[dict[str, int]] + self, + data: bytes, + spans: list[dict[str, int]], + tp_rank: int = 0, + tp_size: int = 1, ) -> list[str]: """Store bytes through the server-owned transfer layer. Args: data: source bytes. spans: byte spans containing source_offset, nbytes, and file_offset. + tp_rank: tensor-parallel rank that owns the stored shard. + tp_size: total tensor-parallel ranks required before publication. Async/thread-safety: Opens a short-lived async IPC connection for this request. @@ -461,6 +585,8 @@ async def transfer_store_bytes( "op": "transfer_store", "payload": {"data": data}, "spans": spans, + "tp_rank": tp_rank, + "tp_size": tp_size, } ) chunk_keys = resp.get("chunk_keys", []) @@ -468,22 +594,28 @@ async def transfer_store_bytes( raise RuntimeError("[IPC] invalid transfer_store chunk_keys response") return [str(key) for key in chunk_keys] - async def transfer_load_bytes(self, spans: list[dict[str, int]]) -> bytes: + async def transfer_load_bytes( + self, + spans: list[dict[str, int]], + lease_id: str | None = None, + ) -> bytes: """Load bytes through the server-owned transfer layer. Args: spans: byte spans containing target_offset, nbytes, and file_offset. + lease_id: Optional base request ID retaining host-tier bytes. Returns: Loaded bytes in target-offset order. """ - resp = await self.call( - { - "op": "transfer_load", - "payload": {"return_data": True}, - "spans": spans, - } - ) + request: dict[str, Any] = { + "op": "transfer_load", + "payload": {"return_data": True}, + "spans": spans, + } + if lease_id is not None: + request["lease_id"] = lease_id + resp = await self.call(request) data = resp.get("data", b"") if not isinstance(data, bytes): raise RuntimeError("[IPC] invalid transfer_load data response") @@ -499,6 +631,8 @@ async def transfer_store_cuda( allocation_offset: int, producer_pid: int, spans: list[dict[str, Any]], + tp_rank: int = 0, + tp_size: int = 1, ) -> list[str]: """Store from a CUDA IPC buffer through the server transfer layer. @@ -513,6 +647,8 @@ async def transfer_store_cuda( ``allocation_base_ptr``. producer_pid: process ID that exported the pointer. spans: byte spans containing source_offset, nbytes, and file_offset. + tp_rank: tensor-parallel rank that owns the stored shard. + tp_size: total tensor-parallel ranks required before publication. """ resp = await self.call( { @@ -527,6 +663,8 @@ async def transfer_store_cuda( "producer_pid": producer_pid, }, "spans": spans, + "tp_rank": tp_rank, + "tp_size": tp_size, } ) chunk_keys = resp.get("chunk_keys", []) @@ -544,6 +682,7 @@ async def transfer_load_cuda( allocation_offset: int, producer_pid: int, spans: list[dict[str, int]], + lease_id: str | None = None, ) -> dict[str, Any]: """Load into a CUDA IPC buffer through the server transfer layer. @@ -558,26 +697,28 @@ async def transfer_load_cuda( ``allocation_base_ptr``. producer_pid: process ID that exported the pointer. spans: byte spans containing target_offset, nbytes, and file_offset. + lease_id: Optional base request ID retaining host-tier bytes. Returns: Server response including transferred bytes and optional timing counters. """ - return await self.call( - { - "op": "transfer_load", - "payload": { - "cuda_ipc_handle": cuda_ipc_handle, - "nbytes": nbytes, - "device_id": device_id, - "device_ptr": device_ptr, - "allocation_base_ptr": allocation_base_ptr, - "allocation_offset": allocation_offset, - "producer_pid": producer_pid, - }, - "spans": spans, - } - ) + request: dict[str, Any] = { + "op": "transfer_load", + "payload": { + "cuda_ipc_handle": cuda_ipc_handle, + "nbytes": nbytes, + "device_id": device_id, + "device_ptr": device_ptr, + "allocation_base_ptr": allocation_base_ptr, + "allocation_offset": allocation_offset, + "producer_pid": producer_pid, + }, + "spans": spans, + } + if lease_id is not None: + request["lease_id"] = lease_id + return await self.call(request) async def register_load_staging_cuda( self, @@ -631,6 +772,7 @@ async def transfer_load_registered_cuda( producer_pid: int, nbytes: int, spans: list[dict[str, int]], + lease_id: str | None = None, ) -> dict[str, Any]: """Load into a previously registered fixed CUDA staging buffer. @@ -640,6 +782,7 @@ async def transfer_load_registered_cuda( producer_pid: Process ID that registered the staging buffer. nbytes: logical bytes to write for this transfer. spans: byte spans containing target_offset, nbytes, and file_offset. + lease_id: Optional base request ID retaining host-tier bytes. Returns: Server response including transferred bytes and timing counters. @@ -648,14 +791,15 @@ async def transfer_load_registered_cuda( Runs on the worker load event loop and avoids per-load CUDA IPC handle export/open payloads on the hot path. """ - return await self.call( - { - "op": "transfer_load", - "payload": { - "load_staging_buffer_index": int(buffer_index), - "producer_pid": int(producer_pid), - "nbytes": int(nbytes), - }, - "spans": spans, - } - ) + request: dict[str, Any] = { + "op": "transfer_load", + "payload": { + "load_staging_buffer_index": int(buffer_index), + "producer_pid": int(producer_pid), + "nbytes": int(nbytes), + }, + "spans": spans, + } + if lease_id is not None: + request["lease_id"] = lease_id + return await self.call(request) diff --git a/daser/connector/metadata.py b/daser/connector/metadata.py index 33eacf3..b597498 100644 --- a/daser/connector/metadata.py +++ b/daser/connector/metadata.py @@ -19,6 +19,8 @@ class ReqLoadSpec: target_token_start: token offset where this chunk starts in the current prompt. pos_offset: target-aware position offset returned by the server. + lease_id: Base request ID retaining host-tier bytes, or empty when the + load follows the ordinary non-prefetch path. """ chunk_key: str @@ -29,6 +31,7 @@ class ReqLoadSpec: token_count: int target_token_start: int = 0 pos_offset: int = 0 + lease_id: str = "" @dataclass diff --git a/daser/connector/scheduler/lifecycle.py b/daser/connector/scheduler/lifecycle.py index 7c43e0b..3140be0 100644 --- a/daser/connector/scheduler/lifecycle.py +++ b/daser/connector/scheduler/lifecycle.py @@ -2,6 +2,7 @@ from __future__ import annotations +from concurrent.futures import Future, ThreadPoolExecutor import logging import math from typing import TYPE_CHECKING, Any @@ -12,6 +13,7 @@ from vllm.v1.request import Request from daser.connector.helpers import PendingStore +from daser.connector.ipc_client import PrefetchLookupResult from daser.connector.metadata import DaserConnectorMeta, ReqLoadSpec, ReqStoreSpec from daser.connector.scheduler.planning import ( _base_req_id, @@ -30,6 +32,21 @@ logger = init_logger(__name__) +def _prefetch_external_spans( + socket_path: str, + lease_id: str, + spans: list[dict[str, int]], +) -> dict[str, int]: + """Prefetch storage spans over a dedicated synchronous IPC connection.""" + from daser.connector.ipc_client import IPCClientSync + + client = IPCClientSync(socket_path) + try: + return client.transfer_prefetch(spans, lease_id=lease_id) + finally: + client.close() + + class RequestLifecycle: """Own scheduler request state and synchronous IPC orchestration. @@ -47,6 +64,8 @@ def __init__( model_id: str, cache_reuse_mode: str, runtime_config_ready: bool, + socket_path: str = "", + prefetch_max_requests: int = 0, ) -> None: self._ipc_sync = ipc_client self._block_tokens = block_tokens @@ -54,6 +73,8 @@ def __init__( self._model_id = model_id self._cache_reuse_mode = cache_reuse_mode self._runtime_config_ready = runtime_config_ready + self._socket_path = socket_path + self._prefetch_max_requests = prefetch_max_requests self._cache_reuse_strategy = build_cache_reuse_strategy( cache_reuse_mode, block_tokens, @@ -63,6 +84,11 @@ def __init__( self._pending_alloc: dict[str, PendingStore] = {} self._pending_async_saves: set[str] = set() self._req_tokens: dict[str, list[int]] = {} + self._prefetch_futures: dict[str, Future[dict[str, int]]] = {} + self._prefetch_executor: ThreadPoolExecutor | None = None + self._prefetch_signatures: dict[str, tuple[tuple[int, int], ...]] = {} + self._prefetch_lookup_results: dict[str, PrefetchLookupResult] = {} + self._leased_request_ids: set[str] = set() def get_num_new_matched_tokens( self, @@ -83,14 +109,45 @@ def get_num_new_matched_tokens( if not getattr(self, "_runtime_config_ready", True): self._refresh_runtime_config() + prefetch_limit = int(getattr(self, "_prefetch_max_requests", 0)) + socket_path = str(getattr(self, "_socket_path", "")) + prefetch_enabled = prefetch_limit > 0 and bool(socket_path) + prefetch_futures = getattr(self, "_prefetch_futures", {}) + prefetch_future = prefetch_futures.get(request.request_id) + prefetch_lookup_result = getattr(self, "_prefetch_lookup_results", {}).get( + request.request_id + ) + if prefetch_future is not None: + if not prefetch_future.done(): + return None, True + try: + prefetch_future.result() + getattr(self, "_leased_request_ids", set()).add(request.request_id) + except Exception as exc: # noqa: BLE001 + logger.warning( + "[CONNECTOR] prefetch failed req=%s: %s", + request.request_id[:8], + exc, + ) + getattr(self, "_prefetch_lookup_results", {}).pop( + request.request_id, None + ) + self._release_transfer_lease(request.request_id, force=True) + prefetch_lookup_result = None + prefetch_futures.pop(request.request_id, None) + start = num_computed_tokens available = len(tokens) - start if available < self._block_tokens: + getattr(self, "_prefetch_lookup_results", {}).pop(request.request_id, None) + self._release_transfer_lease(request.request_id) self._record_external_prefix_cache_miss(available) return 0, False skip_load = bool(_get_kv_transfer_flag(request, "daser_skip_load")) if skip_load: + getattr(self, "_prefetch_lookup_results", {}).pop(request.request_id, None) + self._release_transfer_lease(request.request_id) logger.debug("[CONNECTOR] skip load req=%s", request.request_id[:8]) self._record_external_prefix_cache_miss(available) full_aligned = (len(tokens) // self._block_tokens) * self._block_tokens @@ -109,19 +166,38 @@ def get_num_new_matched_tokens( full_aligned = (len(tokens) // self._block_tokens) * self._block_tokens skip_save = bool(_get_kv_transfer_flag(request, "daser_skip_save")) - try: - chunks = self._lookup_with_external_prefix_metrics( - prefix, - self._model_id, - max(0, len(tokens) - num_computed_tokens), - num_computed_tokens, - ) - except Exception as exc: - logger.warning("[CONNECTOR] lookup failed: %s", exc) - self._runtime_config_ready = False - return 0, False + if prefetch_lookup_result is None: + try: + if prefetch_enabled: + prefetch_lookup_result = self._lookup_with_prefetch_metrics( + request.request_id, + prefix, + self._model_id, + max(0, len(tokens) - num_computed_tokens), + num_computed_tokens, + ) + chunks = prefetch_lookup_result.chunks + if prefetch_lookup_result.tier == "l1": + getattr(self, "_leased_request_ids", set()).add( + request.request_id + ) + else: + chunks = self._lookup_with_external_prefix_metrics( + prefix, + self._model_id, + max(0, len(tokens) - num_computed_tokens), + num_computed_tokens, + ) + except Exception as exc: + logger.warning("[CONNECTOR] lookup failed: %s", exc) + self._release_transfer_lease(request.request_id, force=True) + self._runtime_config_ready = False + return 0, False + else: + chunks = prefetch_lookup_result.chunks if not chunks: + getattr(self, "_prefetch_lookup_results", {}).pop(request.request_id, None) pending_store = ( None if skip_save @@ -146,21 +222,82 @@ def get_num_new_matched_tokens( extra_tokens = _contiguous_prefix_tokens(chunks, num_computed_tokens) if extra_tokens <= 0: + getattr(self, "_prefetch_lookup_results", {}).pop(request.request_id, None) + self._release_transfer_lease(request.request_id) return 0, False available = len(tokens) - num_computed_tokens if extra_tokens >= available: extra_tokens = available - 1 if extra_tokens <= 0: + getattr(self, "_prefetch_lookup_results", {}).pop( + request.request_id, None + ) + self._release_transfer_lease(request.request_id) return 0, False + if ( + prefetch_enabled + and prefetch_lookup_result is not None + and prefetch_lookup_result.tier in ("mixed", "l2") + ): + spans = prefetch_lookup_result.spans + signature = tuple( + sorted( + (int(span["file_offset"]), int(span["nbytes"])) for span in spans + ) + ) + signatures = getattr(self, "_prefetch_signatures", {}) + if signature and signatures.get(request.request_id) != signature: + active = sum(not future.done() for future in prefetch_futures.values()) + if active >= prefetch_limit: + getattr(self, "_prefetch_lookup_results", {})[ + request.request_id + ] = prefetch_lookup_result + return None, True + executor = getattr(self, "_prefetch_executor", None) + if executor is None: + executor = ThreadPoolExecutor( + max_workers=prefetch_limit, + thread_name_prefix="daser-prefetch", + ) + self._prefetch_executor = executor + signatures[request.request_id] = signature + self._prefetch_signatures = signatures + getattr(self, "_prefetch_lookup_results", {})[request.request_id] = ( + prefetch_lookup_result + ) + prefetch_futures[request.request_id] = executor.submit( + _prefetch_external_spans, + socket_path, + request.request_id, + spans, + ) + self._prefetch_futures = prefetch_futures + return None, True + if len(chunks) == 1: self._pending_loads[request.request_id] = dict( - chunks[0], num_computed_tokens=num_computed_tokens + chunks[0], + num_computed_tokens=num_computed_tokens, + lease_id=( + request.request_id + if request.request_id in getattr(self, "_leased_request_ids", set()) + else "" + ), ) else: self._pending_loads[request.request_id] = { - str(i): dict(chunk, num_computed_tokens=num_computed_tokens) + str(i): dict( + chunk, + num_computed_tokens=num_computed_tokens, + lease_id=( + request.request_id + if request.request_id + in getattr(self, "_leased_request_ids", set()) + else "" + ), + ) for i, chunk in enumerate(chunks) } @@ -170,8 +307,34 @@ def get_num_new_matched_tokens( len(chunks), extra_tokens, ) + getattr(self, "_prefetch_lookup_results", {}).pop(request.request_id, None) return extra_tokens, True + def shutdown(self) -> None: + """Stop scheduler-side prefetch workers and release their resources. + + Async/thread-safety: + Called by vLLM after scheduler traffic stops. It waits for active + IPC calls to finish so their sockets are closed cleanly. + """ + futures = getattr(self, "_prefetch_futures", {}) + cleanup_ids = ( + set(futures) + | set(getattr(self, "_prefetch_lookup_results", {})) + | set(getattr(self, "_leased_request_ids", set())) + ) + for future in futures.values(): + future.cancel() + executor = getattr(self, "_prefetch_executor", None) + if executor is not None: + executor.shutdown(wait=True, cancel_futures=True) + self._prefetch_executor = None + futures.clear() + for req_id in cleanup_ids: + self._release_transfer_lease(req_id, force=True) + getattr(self, "_prefetch_lookup_results", {}).clear() + getattr(self, "_prefetch_signatures", {}).clear() + def update_state_after_alloc( self, request: "Request", @@ -201,6 +364,7 @@ def update_state_after_alloc( slot_size=self._slot_size, ): del self._pending_loads[req_id] + self._release_transfer_lease(req_id) self._record_pending_store_blocks(req_id, block_ids) return logger.debug( @@ -234,6 +398,9 @@ def update_state_after_alloc( chunk.get("chunk_key", "")[:8], chunk["block_ids"], ) + if not chunks: + self._pending_loads.pop(req_id, None) + self._release_transfer_lease(req_id) self._record_pending_store_blocks(req_id, block_ids) def build_connector_meta( @@ -313,21 +480,21 @@ def build_connector_meta( pending_async_saves.add(_base_req_id(req_id)) if logger.isEnabledFor(logging.DEBUG): - for req_id, spec in meta.reqs_to_load.items(): + for req_id, load_spec in meta.reqs_to_load.items(): logger.debug( "[CONNECTOR] meta LOAD req=%s start_slot=%d blocks=%d tokens=%d", req_id[:8], - spec.start_slot, - len(spec.block_ids), - spec.token_count, + load_spec.start_slot, + len(load_spec.block_ids), + load_spec.token_count, ) - for req_id, spec in meta.reqs_to_store.items(): + for req_id, store_spec in meta.reqs_to_store.items(): logger.debug( "[CONNECTOR] meta STORE req=%s start_slot=%d blocks=%d tokens=%d", req_id[:8], - spec.start_slot, - len(spec.block_ids), - spec.token_count, + store_spec.start_slot, + len(store_spec.block_ids), + store_spec.token_count, ) return meta @@ -355,11 +522,17 @@ def _drop_preempted_pending_state( Async/thread-safety: Runs on the scheduler thread before metadata is handed to workers. """ - preempted_req_ids = getattr(scheduler_output, "preempted_req_ids", set()) + preempted_req_ids: set[str] = getattr( + scheduler_output, "preempted_req_ids", set() + ) pending_async_saves = self._pending_async_save_ids() for req_id in preempted_req_ids: base_req_id = str(req_id) pending_async_saves.discard(base_req_id) + self._release_transfer_lease(base_req_id, force=True) + getattr(self, "_prefetch_lookup_results", {}).pop(base_req_id, None) + getattr(self, "_prefetch_signatures", {}).pop(base_req_id, None) + self._cancel_prefetch_future(base_req_id) for pending_req_id in list(self._pending_loads): if _matches_request_or_store_id(pending_req_id, base_req_id): self._pending_loads.pop(pending_req_id, None) @@ -469,6 +642,66 @@ def _lookup_with_external_prefix_metrics( raise return self._ipc_sync.lookup(tokens, model_id) + def _lookup_with_prefetch_metrics( + self, + lease_id: str, + tokens: list[int], + model_id: str, + queries: int, + num_computed_tokens: int, + ) -> PrefetchLookupResult: + """Run the atomic lookup/tier-classification scheduler RPC. + + Args: + lease_id: Base vLLM request ID for an all-L1 lease. + tokens: Token prefix sent to the DaseR lookup. + model_id: Model identifier. + queries: vLLM external-prefix query token count. + num_computed_tokens: Tokens already resident in local vLLM KV. + + Returns: + Validated lookup chunks, exact physical spans, and tier state. + + Thread-safety: + Runs on the scheduler thread through its synchronous IPC client. + """ + return self._ipc_sync.lookup_with_prefetch( + lease_id, + tokens, + model_id, + external_prefix_queries=queries, + num_computed_tokens=num_computed_tokens, + ) + + def _release_transfer_lease(self, req_id: str, *, force: bool = False) -> None: + """Best-effort idempotent cleanup for scheduler-owned lease state.""" + leased_ids: set[str] = getattr(self, "_leased_request_ids", set()) + if not force and req_id not in leased_ids: + return + release = getattr(self._ipc_sync, "release_transfer_lease", None) + try: + if release is not None: + release(req_id) + except Exception as exc: # noqa: BLE001 + logger.warning( + "[CONNECTOR] release_transfer_lease failed req=%s: %s", + req_id[:8], + exc, + ) + finally: + leased_ids.discard(req_id) + + def _cancel_prefetch_future(self, req_id: str) -> bool: + """Cancel queued prefetch or release its lease after an active RPC exits.""" + future = getattr(self, "_prefetch_futures", {}).pop(req_id, None) + if future is None: + return False + if not future.cancel(): + future.add_done_callback( + lambda _future: self._release_transfer_lease(req_id, force=True) + ) + return True + def _record_external_prefix_cache_miss(self, queries: int) -> None: """Record a connector external-prefix miss when lookup is skipped. @@ -511,6 +744,8 @@ def _refresh_runtime_config(self) -> None: logger.info("[CONNECTOR] runtime config unavailable: %s", exc) return self._slot_size = int(config.get("slot_size", self._slot_size)) + self._tensor_parallel_size = int(config.get("tensor_parallel_size", 1)) + self._rank_stride_bytes = int(config.get("rank_stride_bytes", 0)) block_tokens = int(config.get("block_tokens", self._block_tokens)) self._model_id = str(config.get("model_id", self._model_id)) cache_reuse_mode = str(config.get("cache_reuse_mode", self._cache_reuse_mode)) @@ -616,6 +851,9 @@ def drop_pending_alloc(self, req_id: str) -> None: req_id: vLLM request ID. """ self._pending_alloc.pop(req_id, None) + getattr(self, "_prefetch_signatures", {}).pop(req_id, None) + getattr(self, "_prefetch_lookup_results", {}).pop(req_id, None) + self._release_transfer_lease(req_id) def _drop_pending_store(self, req_id: str) -> None: """Remove a pending store and release its server writer claim. @@ -651,6 +889,10 @@ def _discard_pending_request(self, req_id: str) -> None: if pending_req_id.startswith(f"{req_id}:store:"): self._drop_pending_store(pending_req_id) self._pending_alloc.pop(req_id, None) + getattr(self, "_prefetch_signatures", {}).pop(req_id, None) + getattr(self, "_prefetch_lookup_results", {}).pop(req_id, None) + had_prefetch = self._cancel_prefetch_future(req_id) + self._release_transfer_lease(req_id, force=had_prefetch) def _pending_async_save_ids(self) -> set[str]: """Return request IDs whose worker-side saves are still pending. @@ -775,6 +1017,8 @@ def update_connector_output(self, connector_output: Any) -> None: self._req_tokens.pop(req_id, None) self._discard_pending_request(req_id) for req_id in getattr(connector_output, "finished_recving", None) or (): + getattr(self, "_prefetch_lookup_results", {}).pop(req_id, None) + self._release_transfer_lease(req_id) if req_id not in self._pending_alloc: self._req_tokens.pop(req_id, None) for pending_req_id in list(self._pending_loads): diff --git a/daser/connector/scheduler/planning.py b/daser/connector/scheduler/planning.py index 986011f..58c09b4 100644 --- a/daser/connector/scheduler/planning.py +++ b/daser/connector/scheduler/planning.py @@ -245,6 +245,7 @@ def _load_spec_from_chunk(chunk: dict[str, Any]) -> ReqLoadSpec: token_count=int(chunk["token_count"]), target_token_start=int(chunk.get("target_token_start", 0)), pos_offset=int(chunk.get("pos_offset", 0)), + lease_id=str(chunk.get("lease_id", "")), ) @@ -276,6 +277,7 @@ def _merge_adjacent_load_specs( prev_slots = len(prev.block_ids) adjacent = ( prev.pos_offset == spec.pos_offset + and prev.lease_id == spec.lease_id and prev.start_slot + prev_slots == spec.start_slot and prev.file_offset + prev_slots * slot_size == spec.file_offset and prev.target_token_start + prev.token_count == spec.target_token_start @@ -292,5 +294,6 @@ def _merge_adjacent_load_specs( token_count=prev.token_count + spec.token_count, target_token_start=prev.target_token_start, pos_offset=prev.pos_offset, + lease_id=prev.lease_id, ) return merged diff --git a/daser/connector/worker/load.py b/daser/connector/worker/load.py index 2820286..fc0d6da 100644 --- a/daser/connector/worker/load.py +++ b/daser/connector/worker/load.py @@ -51,6 +51,14 @@ def block_ids(self) -> list[int]: """Return all vLLM blocks affected by this request.""" return [block_id for spec in self.specs.values() for block_id in spec.block_ids] + @property + def lease_id(self) -> str | None: + """Return the single request lease shared by all grouped load specs.""" + lease_ids = {spec.lease_id for spec in self.specs.values() if spec.lease_id} + if len(lease_ids) > 1: + raise ValueError(f"grouped load has conflicting lease IDs: {lease_ids}") + return next(iter(lease_ids), None) + @dataclass class _InflightLoadBatch: @@ -492,12 +500,17 @@ def _submit_request( request=request, buffer_index=buffer_index, batches=batches, - active=self._submit_batch(batches.popleft(), buffer_index), + active=self._submit_batch( + request.lease_id, + batches.popleft(), + buffer_index, + ), completed=[], ) def _submit_batch( self, + lease_id: str | None, batch: _LoadBatch, buffer_index: int, ) -> _InflightLoadBatch: @@ -512,6 +525,7 @@ def _submit_batch( producer_pid=os.getpid(), nbytes=total_bytes, spans=spans, + lease_id=lease_id, ) else: cp_staging = cupy.asarray(staging) @@ -528,6 +542,7 @@ def _submit_batch( allocation_offset=allocation_offset, producer_pid=os.getpid(), spans=spans, + lease_id=lease_id, ) submitted_at = time.perf_counter() return _InflightLoadBatch( @@ -552,6 +567,7 @@ def _consume_request( state.request.future.set_result(None) return state.buffer_index, True state.active = self._submit_batch( + state.request.lease_id, state.batches.popleft(), state.buffer_index, ) diff --git a/daser/connector/worker/runtime.py b/daser/connector/worker/runtime.py index f4d75cb..0089ba7 100644 --- a/daser/connector/worker/runtime.py +++ b/daser/connector/worker/runtime.py @@ -483,11 +483,8 @@ def wait_for_save(self) -> None: if self._meta is None: return reqs_to_store = dict(self._meta.reqs_to_store) - commit_keys = { - spec.chunk_key for spec in reqs_to_store.values() if spec.block_ids - } - if commit_keys: - self._store_pipeline.queue_finished(reqs_to_store, commit_keys) + if reqs_to_store: + self._store_pipeline.queue_finished(reqs_to_store) def get_finished( self, finished_req_ids: set[str] diff --git a/daser/connector/worker/store.py b/daser/connector/worker/store.py index 40e2ec7..6cbbee0 100644 --- a/daser/connector/worker/store.py +++ b/daser/connector/worker/store.py @@ -40,7 +40,6 @@ class _DeferredFinishedSave: """Hold request store work until vLLM reports it finished.""" - commit_keys: set[str] reqs_to_store: dict[str, ReqStoreSpec] finished: bool = False future: Any | None = None @@ -68,6 +67,8 @@ def __init__(self, socket_path: str) -> None: self._staging_bytes = 0 self._staging_pool: FixedCudaStagingPool | None = None self._pending_finished_saves: dict[str, _DeferredFinishedSave] = {} + self._store_capacity = 1 + self._store_semaphore: asyncio.Semaphore | None = None self._kv_caches: dict[str, torch.Tensor] = {} self._layer_names: list[str] = [] self._layer_idx_map: dict[str, int] = {} @@ -116,6 +117,8 @@ def configure( self._tp_size = tp_size self._staging_bytes = staging_bytes self._staging_pool = staging_pool + self._store_capacity = staging_pool.depth + self._store_semaphore = None def initialize_transfer(self) -> None: """Initialize the store IPC transfer client on its event loop. @@ -150,14 +153,11 @@ def configure_rank_geometry( def queue_finished( self, reqs_to_store: dict[str, ReqStoreSpec], - commit_keys: set[str], ) -> None: """Queue stores until vLLM reports their requests finished. Args: reqs_to_store: Store metadata for the current worker step. - commit_keys: Chunk keys eligible for commit after transfer. - Async/thread-safety: Called on the worker thread to accumulate immutable store intent. CUDA ordering is captured later, when vLLM reports completion. @@ -166,11 +166,9 @@ def queue_finished( base_id = base_req_id(req_id) save = self._pending_finished_saves.get(base_id) if save is None: - save = _DeferredFinishedSave(set(), {}) + save = _DeferredFinishedSave({}) self._pending_finished_saves[base_id] = save save.reqs_to_store[req_id] = spec - if spec.chunk_key in commit_keys: - save.commit_keys.add(spec.chunk_key) def collect_finished(self, finished_req_ids: set[str]) -> set[str]: """Submit newly finished stores and collect completed requests. @@ -199,16 +197,12 @@ def collect_finished(self, finished_req_ids: set[str]) -> set[str]: finally: del self._pending_finished_saves[req_id] - capacity = self._staging_pool.depth if self._staging_pool is not None else 1 - inflight = sum( - save.future is not None for save in self._pending_finished_saves.values() - ) + # Submit the entire finished FIFO here. The private event loop applies + # the staging-depth bound, so one completed save releases the next + # queued save without requiring another vLLM connector step. for save in self._pending_finished_saves.values(): - if inflight >= capacity: - break if save.finished and save.future is None: self._submit_save(save) - inflight += 1 return finished def shutdown(self) -> None: @@ -264,10 +258,33 @@ def _run_loop(self) -> None: self._loop.run_forever() def _submit_save(self, save: _DeferredFinishedSave) -> None: - """Capture producer ordering and submit one save to the store thread.""" + """Capture producer ordering and enqueue one save in FIFO order.""" sample = next(iter(self._kv_caches.values()), None) producer_event = record_cuda_event(sample) if sample is not None else None - save.future = self._submit(self._store_finished_save(save, producer_event)) + save.future = self._submit(self._run_bounded_save(save, producer_event)) + + async def _run_bounded_save( + self, + save: _DeferredFinishedSave, + producer_event: torch.cuda.Event | None, + ) -> None: + """Run one save while preserving FIFO staging-buffer admission. + + Args: + save: Finished request save to transfer. + producer_event: CUDA event ordering the KV snapshot. + + Async/thread-safety: + Runs on the private store event loop. The semaphore prevents more + saves from entering synchronous staging acquisition than there are + preallocated buffers; semaphore waiters are admitted in FIFO order. + """ + semaphore = self._store_semaphore + if semaphore is None: + semaphore = asyncio.Semaphore(max(1, self._store_capacity)) + self._store_semaphore = semaphore + async with semaphore: + await self._store_finished_save(save, producer_event) def _plan_finished_save( self, @@ -296,22 +313,12 @@ async def _store_finished_save( save: _DeferredFinishedSave, producer_event: torch.cuda.Event | None, ) -> None: - stored_keys: list[str] = [] for block_ids, spans in self._plan_finished_save(save): staged = self._stage_batch(block_ids, spans, producer_event) try: - stored_keys.extend(await self._write_cuda_buffer(staged)) + await self._write_cuda_buffer(staged) finally: staged.lease.release() - requested = save.commit_keys - keys_to_commit = list( - dict.fromkeys(key for key in stored_keys if key in requested) - ) - await self._client.commit_chunks( - keys_to_commit, - tp_rank=self._tp_rank, - tp_size=self._tp_size, - ) def _stage_batch( self, @@ -382,6 +389,8 @@ async def _write_cuda_buffer(self, staged: "StagedStoreBatch") -> list[str]: allocation_base_ptr=allocation_base, allocation_offset=allocation_offset, producer_pid=os.getpid(), + tp_rank=self._tp_rank, + tp_size=self._tp_size, spans=[ { "source_offset": span.source_offset, diff --git a/daser/server/chunk_lifecycle.py b/daser/server/chunk_lifecycle.py index 10f70d4..0ad6524 100644 --- a/daser/server/chunk_lifecycle.py +++ b/daser/server/chunk_lifecycle.py @@ -26,6 +26,9 @@ def __init__(self) -> None: self._commit_waiters: dict[str, set[asyncio.Future[None]]] = {} self._commit_shards: dict[str, tuple[int, set[int]]] = {} self._publishing: set[str] = set() + self._written_ranges: dict[tuple[str, int], list[tuple[int, int]]] = {} + self._write_expectations: dict[tuple[str, int], tuple[int, int]] = {} + self._complete_write_ranks: set[tuple[str, int]] = set() def is_committed(self, chunk_key: str) -> bool: """Return whether ``chunk_key`` is committed and visible to lookup.""" @@ -55,8 +58,72 @@ def mark_committed(self, chunk_key: str) -> None: self._write_owners.add(chunk_key) self._commit_shards.pop(chunk_key, None) self._publishing.discard(chunk_key) + self._clear_written_ranges(chunk_key) self._notify_commit_waiters(chunk_key) + def record_written_range( + self, + chunk_key: str, + tp_rank: int, + expected_start: int, + expected_end: int, + range_start: int, + range_end: int, + ) -> bool: + """Record an accepted transfer range and report complete coverage. + + Args: + chunk_key: Cache key whose bytes were transferred. + tp_rank: Tensor-parallel rank owning the range. + expected_start: First file byte required for this rank's chunk. + expected_end: Exclusive file byte after this rank's chunk. + range_start: First file byte accepted by the transfer. + range_end: Exclusive file byte accepted by the transfer. + + Returns: + True once, when this rank has covered the complete expected range. + + Raises: + ValueError: if the range is invalid or the expected range changes. + + Async/thread-safety: + Runs on the server event loop. Ranges are merged in memory and do + not perform blocking I/O. + """ + if expected_start < 0 or expected_end <= expected_start: + raise ValueError("invalid expected write range") + if range_start < expected_start or range_end > expected_end: + raise ValueError("write range exceeds the allocated chunk") + if range_end <= range_start: + raise ValueError("write range must be non-empty") + identity = (chunk_key, tp_rank) + expected = self._write_expectations.setdefault( + identity, (expected_start, expected_end) + ) + if expected != (expected_start, expected_end): + raise ValueError("inconsistent expected write range") + if identity in self._complete_write_ranks: + return False + + ranges = self._written_ranges.setdefault(identity, []) + ranges.append((range_start, range_end)) + ranges.sort() + merged: list[tuple[int, int]] = [] + for start, end in ranges: + if merged and start <= merged[-1][1]: + merged[-1] = (merged[-1][0], max(merged[-1][1], end)) + else: + merged.append((start, end)) + self._written_ranges[identity] = merged + if ( + len(merged) == 1 + and merged[0][0] <= expected_start + and merged[0][1] >= expected_end + ): + self._complete_write_ranks.add(identity) + return True + return False + def record_commit_shard(self, chunk_key: str, tp_rank: int, tp_size: int) -> bool: """Record one TP rank and return whether this call should publish. @@ -102,6 +169,7 @@ def mark_evicted(self, chunk_key: str) -> None: self._write_owners.discard(chunk_key) self._commit_shards.pop(chunk_key, None) self._publishing.discard(chunk_key) + self._clear_written_ranges(chunk_key) self._evicted.add(chunk_key) def discard_owner(self, chunk_key: str) -> None: @@ -109,6 +177,7 @@ def discard_owner(self, chunk_key: str) -> None: self._write_owners.discard(chunk_key) self._commit_shards.pop(chunk_key, None) self._publishing.discard(chunk_key) + self._clear_written_ranges(chunk_key) def discard(self, chunk_key: str) -> None: """Drop committed and write-owner state without recording eviction.""" @@ -116,6 +185,19 @@ def discard(self, chunk_key: str) -> None: self._write_owners.discard(chunk_key) self._commit_shards.pop(chunk_key, None) self._publishing.discard(chunk_key) + self._clear_written_ranges(chunk_key) + + def _clear_written_ranges(self, chunk_key: str) -> None: + """Remove transfer coverage state for a chunk.""" + identities = [ + identity + for identity in self._write_expectations + if identity[0] == chunk_key + ] + for identity in identities: + self._write_expectations.pop(identity, None) + self._written_ranges.pop(identity, None) + self._complete_write_ranks.discard(identity) async def wait_for_committed( self, diff --git a/daser/server/core.py b/daser/server/core.py index 3a4e6a0..ae9ebe9 100644 --- a/daser/server/core.py +++ b/daser/server/core.py @@ -474,6 +474,74 @@ async def commit_chunk( tp_size, ) + async def record_store_ranges( + self, + spans: list[dict[str, Any]], + tp_rank: int, + tp_size: int, + local_slot_size: int, + rank_stride_bytes: int, + ) -> list[str]: + """Record completed transfer ranges and publish fully written chunks. + + Args: + spans: Accepted store spans with chunk allocation metadata. + tp_rank: Tensor-parallel rank that completed the transfer. + tp_size: Total tensor-parallel ranks required for publication. + local_slot_size: Bytes occupied by one rank-local KV slot. + rank_stride_bytes: File-byte distance between TP rank lanes. + + Returns: + Chunk keys whose rank shard became completely written and whose + existing commit path was invoked. + + Raises: + ValueError: if transfer geometry or allocation metadata is invalid. + + Async/thread-safety: + Runs on the server event loop. It only updates control-plane state + and awaits the existing retrieval-index commit operation. + """ + if tp_size <= 0 or not 0 <= tp_rank < tp_size: + raise ValueError(f"invalid TP rank {tp_rank} for size {tp_size}") + if local_slot_size <= 0: + raise ValueError("local_slot_size must be positive") + if tp_size > 1 and rank_stride_bytes <= 0: + raise ValueError("rank_stride_bytes must be positive for TP stores") + + ready: list[str] = [] + ready_set: set[str] = set() + for span in spans: + chunk_key = str(span.get("chunk_key", "")) + if not chunk_key: + continue + meta = self._cm.store.get(chunk_key) + if meta is None: + continue + start_slot = int(span.get("start_slot", -1)) + num_slots = int(span.get("num_slots", 0)) + if start_slot != meta.start_slot or num_slots != meta.num_slots: + raise ValueError(f"store span allocation mismatch: {chunk_key}") + expected_start = tp_rank * rank_stride_bytes + start_slot * local_slot_size + expected_end = expected_start + num_slots * local_slot_size + range_start = int(span["file_offset"]) + range_end = range_start + int(span["nbytes"]) + complete = self._lifecycle.record_written_range( + chunk_key, + tp_rank, + expected_start, + expected_end, + range_start, + range_end, + ) + if complete and chunk_key not in ready_set: + ready.append(chunk_key) + ready_set.add(chunk_key) + + for chunk_key in ready: + await self.commit_chunk(chunk_key, tp_rank=tp_rank, tp_size=tp_size) + return ready + def is_chunk_committed(self, chunk_key: str) -> bool: """Return whether a chunk key has been committed. diff --git a/daser/server/ipc/server.py b/daser/server/ipc/server.py index 07a6f0a..2103182 100644 --- a/daser/server/ipc/server.py +++ b/daser/server/ipc/server.py @@ -55,6 +55,65 @@ def _external_prefix_hits( return max(0, min(hits, queries)) +def _prefetch_spans_from_chunks( + chunks: list[ChunkInfo], + *, + external_start: int, + external_tokens: int, + block_tokens: int, + slot_size: int, + tensor_parallel_size: int, + rank_stride_bytes: int, +) -> list[dict[str, int]]: + """Translate an admitted external KV window into physical TP-lane spans. + + Args: + chunks: Server lookup chunks covering the prompt prefix. + external_start: Token offset where external cache loading begins. + external_tokens: Number of externally admitted tokens. + block_tokens: Tokens stored in one logical cache slot. + slot_size: Aggregate bytes per logical slot across TP ranks. + tensor_parallel_size: Number of physical rank lanes. + rank_stride_bytes: Byte distance between adjacent rank lanes. + + Returns: + Sorted block-aligned physical storage ranges for host-tier admission. + """ + if block_tokens <= 0: + raise ValueError("block_tokens must be positive for prefetch lookup") + if external_tokens <= 0 or external_start % block_tokens != 0: + return [] + if slot_size <= 0 or tensor_parallel_size <= 0: + raise ValueError("invalid transfer geometry for prefetch lookup") + if slot_size % tensor_parallel_size: + raise ValueError("slot_size must divide evenly across tensor-parallel ranks") + local_slot_size = slot_size // tensor_parallel_size + external_end = external_start + external_tokens + spans: list[dict[str, int]] = [] + for chunk in sorted(chunks, key=lambda item: int(item.target_token_start)): + target_start = int(chunk.target_token_start) + target_end = target_start + int(chunk.token_count) + load_start = max(target_start, external_start) + load_end = min(target_end, external_end) + load_start = ((load_start + block_tokens - 1) // block_tokens) * block_tokens + load_end = (load_end // block_tokens) * block_tokens + if load_end <= load_start: + continue + start_slot = int(chunk.start_slot) + ( + (load_start - target_start) // block_tokens + ) + nbytes = ((load_end - load_start) // block_tokens) * local_slot_size + for rank in range(tensor_parallel_size): + spans.append( + { + "file_offset": rank * rank_stride_bytes + + start_slot * local_slot_size, + "nbytes": nbytes, + } + ) + return spans + + def _coalesce_transfer_spans(spans: list[dict[str, Any]]) -> list[dict[str, int]]: """Merge adjacent transfer spans without changing byte contents. @@ -131,6 +190,7 @@ def __init__( str, Callable[[dict[str, Any]], Awaitable[dict[str, Any]]] ] = { "lookup": self._op_lookup, + "lookup_prefetch": self._op_lookup_prefetch, "record_external_prefix_cache": self._op_record_external_prefix_cache, "get_runtime_config": self._op_get_runtime_config, "alloc_chunk": self._op_alloc_chunk, @@ -141,12 +201,14 @@ def __init__( "commit_stats": self._op_commit_stats, "live_allocations": self._op_live_allocations, "transfer_drain": self._op_transfer_drain, + "transfer_prefetch": self._transfer_prefetch, "init_transfer": self._op_init_transfer, "transfer_store": self._transfer_store, "transfer_load": self._transfer_load, "register_load_staging": self._register_load_staging, "evict_chunk": self._op_evict_chunk, "release_chunk_writer": self._op_release_chunk_writer, + "release_transfer_lease": self._op_release_transfer_lease, } async def start(self) -> None: @@ -312,6 +374,41 @@ async def _op_lookup(self, msg: dict[str, Any]) -> dict[str, Any]: ) return {"chunks": [chunk.to_dict() for chunk in chunks]} + async def _op_lookup_prefetch(self, msg: dict[str, Any]) -> dict[str, Any]: + """Lookup, classify exact external spans, and lease an all-L1 result.""" + lease_id = str(msg.get("lease_id", "")) + if not lease_id: + raise ValueError("lookup_prefetch requires lease_id") + chunks = await self._core.lookup(msg["tokens"], msg["model_id"]) + queries = int(msg.get("external_prefix_queries", 0)) + num_computed_tokens = int(msg.get("num_computed_tokens", 0)) + hits = _external_prefix_hits(chunks, num_computed_tokens, queries) + await self._core.record_external_prefix_cache(queries=queries, hits=hits) + if not chunks or hits <= 0: + return {"chunks": [chunk.to_dict() for chunk in chunks], "spans": []} + + block_tokens = int(self._runtime_config.get("block_tokens", 0)) + spans = _prefetch_spans_from_chunks( + chunks, + external_start=num_computed_tokens, + external_tokens=hits, + block_tokens=block_tokens, + slot_size=int(self._runtime_config.get("slot_size", 0)), + tensor_parallel_size=int( + self._runtime_config.get("tensor_parallel_size", 1) + ), + rank_stride_bytes=int(self._runtime_config.get("rank_stride_bytes", 0)), + ) + if not spans: + return {"chunks": [chunk.to_dict() for chunk in chunks], "spans": []} + transfer = self._ensure_transfer() + tier = await transfer.classify_and_acquire_lease(lease_id, spans) + return { + "chunks": [chunk.to_dict() for chunk in chunks], + "spans": spans, + "tier": tier, + } + async def _op_record_external_prefix_cache( self, msg: dict[str, Any] ) -> dict[str, Any]: @@ -384,6 +481,44 @@ async def _op_transfer_drain(self, msg: dict[str, Any]) -> dict[str, Any]: await transfer.drain() return {"ok": True} + async def _transfer_prefetch(self, msg: dict[str, Any]) -> dict[str, Any]: + """Promote storage spans into the server-owned host-memory tier. + + Args: + msg: IPC request containing ``spans`` with file offsets and sizes. + + Returns: + Requested, L1-resident, and L2-read byte counts. + + Async/thread-safety: + Runs on the IPC event loop and delegates to the transfer layer's + asynchronous prefetch capability. + """ + transfer = self._ensure_transfer() + lease_id = str(msg.get("lease_id", "")) or None + spans = list(msg.get("spans", [])) + if lease_id is None: + result = await transfer.prefetch_bytes_grouped(spans) + else: + result = await transfer.prefetch_bytes_grouped(spans, lease_id=lease_id) + self._metrics.counter( + "daser_prefetch_operations_total", + "Host-tier prefetch operations by result.", + ).inc(labels={"status": "ok"}) + bytes_counter = self._metrics.counter( + "daser_prefetch_bytes_total", + "Host-tier prefetch bytes by tier.", + ) + bytes_counter.inc(result.requested_bytes, labels={"tier": "requested"}) + bytes_counter.inc(result.l1_bytes, labels={"tier": "l1"}) + bytes_counter.inc(result.l2_bytes, labels={"tier": "l2"}) + return { + "ok": True, + "requested_bytes": result.requested_bytes, + "l1_bytes": result.l1_bytes, + "l2_bytes": result.l2_bytes, + } + async def _op_init_transfer(self, msg: dict[str, Any]) -> dict[str, Any]: """Handle an ``init_transfer`` request.""" self._ensure_transfer() @@ -403,6 +538,16 @@ async def _op_release_chunk_writer(self, msg: dict[str, Any]) -> dict[str, Any]: ) return {"released": released} + async def _op_release_transfer_lease(self, msg: dict[str, Any]) -> dict[str, Any]: + """Idempotently release remaining host-tier bytes for one request.""" + lease_id = str(msg.get("lease_id", "")) + if not lease_id: + raise ValueError("release_transfer_lease requires lease_id") + transfer = self._transfer + if transfer is not None: + await transfer.release_lease(lease_id) + return {"ok": True} + async def _transfer_store(self, msg: dict[str, Any]) -> dict[str, Any]: """Store one or more spans through the server-owned transfer layer. @@ -422,6 +567,9 @@ async def _transfer_store(self, msg: dict[str, Any]) -> dict[str, Any]: transfer = self._ensure_transfer() total = 0 stored_chunk_keys: list[str] = [] + accepted_spans: list[dict[str, Any]] = [] + tp_rank = int(msg.get("tp_rank", 0)) + tp_size = int(msg.get("tp_size", 1)) buffer = self._payload_buffer(payload) try: live_spans: list[dict[str, Any]] = [] @@ -444,6 +592,15 @@ async def _transfer_store(self, msg: dict[str, Any]) -> dict[str, Any]: ) continue stored_chunk_keys.append(chunk_key) + accepted_spans.append( + { + "chunk_key": chunk_key, + "file_offset": file_offset, + "nbytes": nbytes, + "start_slot": int(span.get("start_slot", -1)), + "num_slots": int(span.get("num_slots", 0)), + } + ) live_spans.append(span) store_spans = ( @@ -455,6 +612,28 @@ async def _transfer_store(self, msg: dict[str, Any]) -> dict[str, Any]: finally: if isinstance(buffer, _UncachedCudaArray): buffer.close() + if accepted_spans: + configured_tp_size = int( + self._runtime_config.get("tensor_parallel_size", tp_size) + ) + if configured_tp_size != tp_size: + raise ValueError( + f"transfer TP size {tp_size} != configured {configured_tp_size}" + ) + slot_size = int(self._runtime_config.get("slot_size", 0)) + local_slot_size = int( + self._runtime_config.get( + "local_slot_size", slot_size // max(1, tp_size) + ) + ) + rank_stride_bytes = int(self._runtime_config.get("rank_stride_bytes", 0)) + await self._core.record_store_ranges( + accepted_spans, + tp_rank=tp_rank, + tp_size=tp_size, + local_slot_size=local_slot_size, + rank_stride_bytes=rank_stride_bytes, + ) self._record_transfer_metrics( op="store", backend=backend, @@ -478,6 +657,7 @@ async def _transfer_load(self, msg: dict[str, Any]) -> dict[str, Any]: """ payload = msg.get("payload", {}) spans = list(msg.get("spans", [])) + lease_id = str(msg.get("lease_id", "")) or None started_total = time.perf_counter() backend = str(self._runtime_config.get("transfer_mode", "gds")) transfer = self._ensure_transfer() @@ -497,17 +677,28 @@ async def _transfer_load(self, msg: dict[str, Any]) -> dict[str, Any]: close_one_shot_buffer = isinstance(buffer, _UncachedCudaArray) and ( "load_staging_buffer_index" not in payload ) + leased_load_started = False try: before = asdict(transfer.stats) started = time.perf_counter() load_start = time.perf_counter() - total = await transfer.load_bytes_grouped(buffer, spans) + if lease_id is None: + total = await transfer.load_bytes_grouped(buffer, spans) + else: + total = await transfer.load_leased_bytes_grouped( + buffer, + spans, + lease_id, + ) + leased_load_started = True load_ms = (time.perf_counter() - load_start) * 1000 synchronize = getattr(buffer, "synchronize", None) if synchronize is not None: sync_start = time.perf_counter() synchronize() sync_ms = (time.perf_counter() - sync_start) * 1000 + if lease_id is not None: + await transfer.release_lease_ranges(lease_id, spans) elapsed_ms = (time.perf_counter() - started) * 1000 after = asdict(transfer.stats) stats_delta = { @@ -544,11 +735,18 @@ async def _transfer_load(self, msg: dict[str, Any]) -> dict[str, Any]: elapsed_s=time.perf_counter() - started_total, ) return response + except BaseException: + if lease_id is not None: + if leased_load_started: + await transfer.release_lease_ranges(lease_id, []) + await transfer.release_lease(lease_id) + raise finally: if close_one_shot_buffer: close = getattr(buffer, "close", None) close_start = time.perf_counter() - close() + if close is not None: + close() close_ms = (time.perf_counter() - close_start) * 1000 logger.info( "[IPC] transfer_load close timing: bytes=%d close_ms=%.3f", diff --git a/daser/transfer/base.py b/daser/transfer/base.py index c72bb55..888be10 100644 --- a/daser/transfer/base.py +++ b/daser/transfer/base.py @@ -3,7 +3,9 @@ # Standard from abc import ABC, abstractmethod from dataclasses import dataclass -from typing import Any +from typing import Any, Literal + +TransferTier = Literal["l1", "mixed", "l2"] @dataclass @@ -15,12 +17,27 @@ class TransferStats: l1_misses: number of reads not found in the memory tier. l2_reads: number of reads issued to the SSD tier. l2_writes: number of writes issued to the SSD tier. + prefetch_requests: number of host-tier prefetch operations. + prefetch_l1_bytes: requested bytes already resident in L1. + prefetch_l2_bytes: requested bytes read from L2. """ l1_hits: int = 0 l1_misses: int = 0 l2_reads: int = 0 l2_writes: int = 0 + prefetch_requests: int = 0 + prefetch_l1_bytes: int = 0 + prefetch_l2_bytes: int = 0 + + +@dataclass(frozen=True) +class PrefetchResult: + """Byte attribution for one host-tier prefetch operation.""" + + requested_bytes: int + l1_bytes: int + l2_bytes: int class TransferLayer(ABC): @@ -142,6 +159,111 @@ async def load_bytes_grouped(self, dst: Any, spans: list[dict[str, int]]) -> int ) return total + async def prefetch_bytes_grouped( + self, + spans: list[dict[str, int]], + lease_id: str | None = None, + ) -> PrefetchResult: + """Promote storage spans into an optional host-memory tier. + + Args: + spans: Storage spans containing ``file_offset`` and ``nbytes``. + lease_id: Optional request identifier whose promoted bytes must + remain stable until a later load consumes them. + + Returns: + Byte attribution for the request. + + Raises: + NotImplementedError: If this backend has no host-memory tier. + + Async/thread-safety: + Implementations run on the server event loop and must await all + blocking I/O through their normal async backend path. + """ + raise NotImplementedError("transfer backend does not support host prefetch") + + async def classify_and_acquire_lease( + self, + lease_id: str, + spans: list[dict[str, int]], + ) -> TransferTier: + """Classify host-tier residency and atomically lease an all-L1 window. + + Args: + lease_id: Request identifier used by later load and release calls. + spans: Storage spans containing ``file_offset`` and ``nbytes``. + + Returns: + ``l1`` when every byte is resident, ``mixed`` for partial L1 + residency, or ``l2`` when no requested byte is resident. + + Raises: + NotImplementedError: If this backend has no host-memory tier. + + Async/thread-safety: + Classification and all-L1 lease acquisition must be atomic with + respect to cache eviction and overlapping stores. + """ + raise NotImplementedError("transfer backend does not support host leases") + + async def load_leased_bytes_grouped( + self, + dst: Any, + spans: list[dict[str, int]], + lease_id: str, + ) -> int: + """Load spans from bytes retained by a request lease. + + Args: + dst: Writable destination buffer. + spans: Span dicts with target offset, file offset, and byte count. + lease_id: Request identifier returned through lookup admission. + + Returns: + Total number of bytes loaded. + + Async/thread-safety: + Backends with host leases must keep the lease alive after this call; + the IPC boundary releases ranges only after destination sync. + """ + return await self.load_bytes_grouped(dst, spans) + + async def release_lease_ranges( + self, + lease_id: str, + spans: list[dict[str, int]], + ) -> None: + """Release successfully staged physical ranges from a request lease. + + Args: + lease_id: Request identifier owning the retained bytes. + spans: Physical storage ranges safe to release after staging sync. + + Returns: + None. + + Async/thread-safety: + Implementations must make repeated and overlapping releases + idempotent. + """ + return None + + async def release_lease(self, lease_id: str) -> None: + """Release every remaining range owned by a request lease. + + Args: + lease_id: Request identifier to clean up. + + Returns: + None. + + Async/thread-safety: + Implementations must make cleanup idempotent so cancellation can + race with load completion safely. + """ + return None + async def drain(self) -> None: # noqa: B027 """Wait for any background transfer work to complete. diff --git a/daser/transfer/iouring/l1_cache.py b/daser/transfer/iouring/l1_cache.py index e7254b0..e7cd27f 100644 --- a/daser/transfer/iouring/l1_cache.py +++ b/daser/transfer/iouring/l1_cache.py @@ -110,6 +110,30 @@ def find( return key, data, file_offset - key[0] return None + def has_overlap(self, file_offset: int, nbytes: int) -> bool: + """Return whether any resident range overlaps a requested span. + + Args: + file_offset: Start of the requested byte range. + nbytes: Number of requested bytes. + + Returns: + ``True`` when at least one resident L1 range overlaps the span. + """ + if nbytes <= 0: + return False + end = file_offset + nbytes + index = max(0, bisect.bisect_right(self._starts, file_offset) - 1) + while index < len(self._starts): + start = self._starts[index] + if start >= end: + break + key = self._by_start.get(start) + if key is not None and key[0] + key[1] > file_offset: + return True + index += 1 + return False + def resolve_subranges( self, target_offset: int, diff --git a/daser/transfer/iouring/layer.py b/daser/transfer/iouring/layer.py index 0d79fd4..f9aed9f 100644 --- a/daser/transfer/iouring/layer.py +++ b/daser/transfer/iouring/layer.py @@ -2,11 +2,17 @@ # Standard import asyncio +from dataclasses import dataclass from typing import Any # First Party from daser.logging import init_logger -from daser.transfer.base import TransferLayer, TransferStats +from daser.transfer.base import ( + PrefetchResult, + TransferLayer, + TransferStats, + TransferTier, +) from daser.transfer.iouring import copy_ops from daser.transfer.iouring.l1_cache import L1Cache, L1RangeHit from daser.transfer.iouring.l2_engine import L2IoEngine @@ -18,6 +24,108 @@ _DIRECT_IO_ALIGNMENT = 4096 +@dataclass(frozen=True) +class _LeasedSliceRange: + """Map one physical byte range to bytes retained in a pinned slice.""" + + file_offset: int + nbytes: int + data: PinnedMemorySlice + source_offset: int + + +@dataclass +class _RequestLease: + """Retain the unconsumed physical ranges for one inference request.""" + + remaining: list[tuple[int, int]] + hits: list[_LeasedSliceRange] + slice_ids: set[int] + released: asyncio.Future[None] + active_loads: int = 0 + release_all: bool = False + + +def _normalize_ranges(spans: list[dict[str, int]]) -> list[tuple[int, int]]: + """Return sorted, merged positive physical ranges from transfer spans.""" + ranges = sorted( + (int(span["file_offset"]), int(span["nbytes"])) + for span in spans + if int(span["nbytes"]) > 0 + ) + merged: list[tuple[int, int]] = [] + for start, size in ranges: + end = start + size + if merged and start <= merged[-1][0] + merged[-1][1]: + previous_start, previous_size = merged[-1] + merged[-1] = ( + previous_start, + max(previous_start + previous_size, end) - previous_start, + ) + else: + merged.append((start, size)) + return merged + + +def _subtract_ranges( + ranges: list[tuple[int, int]], + removals: list[tuple[int, int]], +) -> list[tuple[int, int]]: + """Subtract physical ranges, preserving any non-overlapping fragments.""" + result = list(ranges) + for remove_start, remove_size in removals: + remove_end = remove_start + remove_size + next_result: list[tuple[int, int]] = [] + for start, size in result: + end = start + size + if end <= remove_start or remove_end <= start: + next_result.append((start, size)) + continue + if start < remove_start: + next_result.append((start, remove_start - start)) + if remove_end < end: + next_result.append((remove_end, end - remove_end)) + result = next_result + return result + + +def _subtract_leased_hits( + hits: list[_LeasedSliceRange], + removals: list[tuple[int, int]], +) -> list[_LeasedSliceRange]: + """Subtract consumed ranges while preserving pinned-slice source offsets.""" + result = list(hits) + for remove_start, remove_size in removals: + remove_end = remove_start + remove_size + next_result: list[_LeasedSliceRange] = [] + for hit in result: + start = hit.file_offset + end = start + hit.nbytes + if end <= remove_start or remove_end <= start: + next_result.append(hit) + continue + if start < remove_start: + next_result.append( + _LeasedSliceRange( + file_offset=start, + nbytes=remove_start - start, + data=hit.data, + source_offset=hit.source_offset, + ) + ) + if remove_end < end: + next_result.append( + _LeasedSliceRange( + file_offset=remove_end, + nbytes=end - remove_end, + data=hit.data, + source_offset=hit.source_offset + remove_end - start, + ) + ) + result = next_result + return result + + class TieredIOUringTransferLayer(TransferLayer): """Async L1 pinned-memory + L2 SSD transfer layer. @@ -62,10 +170,16 @@ def __init__( self._l2_bytes = l2_bytes self._pending_l2: dict[tuple[int, int], asyncio.Task[None]] = {} self._pending_l2_buffers: dict[tuple[int, int], PinnedMemorySlice] = {} + self._pending_l1_promotions: dict[int, asyncio.Future[None]] = {} + self._pending_l1_promotion_epochs: dict[int, int] = {} + self._cache_epoch = 0 + self._cache_mutations: list[tuple[int, int, int]] = [] + self._request_leases: dict[str, _RequestLease] = {} + self._leased_slices: dict[int, tuple[PinnedMemorySlice, int]] = {} self._l1 = L1Cache( l1_bytes, alignment=_DIRECT_IO_ALIGNMENT, - pinned_predicate=self._is_pinned_by_l2, + pinned_predicate=self._is_slice_pinned, ) self._l2_errors: list[BaseException] = [] self._lock = asyncio.Lock() @@ -230,46 +344,71 @@ async def store_bytes(self, src: Any, file_offset: int, nbytes: int) -> int: self._check_range(file_offset, nbytes) key = (file_offset, nbytes) if self._l2 is None: - async with self._lock: - hit = self._l1.find(file_offset) - if hit is not None: - hit_key, cached, target_offset = hit - if target_offset + nbytes <= len(cached): - self._copy_src_to_pinned_at(src, cached, target_offset, nbytes) - self._l1.touch(hit_key) + while True: + async with self._lock: + waiters = self._overlapping_lease_waiters_locked( + file_offset, nbytes + ) + if not waiters: + hit = self._l1.find(file_offset) + if hit is not None: + hit_key, cached, target_offset = hit + if target_offset + nbytes <= len(cached): + self._copy_src_to_pinned_at( + src, cached, target_offset, nbytes + ) + self._l1.touch(hit_key) + self._record_cache_mutation_locked(file_offset, nbytes) + return nbytes + data = self._l1.reserve_or_raise( + key, + nbytes, + preserve_overlaps=True, + ) + try: + self._copy_src_to_pinned_at(src, data, 0, nbytes) + except BaseException: + data.close() + raise + self._record_cache_mutation_locked(file_offset, nbytes) + self._l1.put(key, data) return nbytes - data = self._l1.reserve_or_raise( - key, - nbytes, - preserve_overlaps=True, - ) - try: - self._copy_src_to_pinned_at(src, data, 0, nbytes) - except BaseException: - data.close() - raise - self._l1.put(key, data) - return nbytes + await asyncio.gather(*waiters) + await self._wait_for_overlapping_leases(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 - return nbytes + try: + while True: + async with self._lock: + self._raise_l2_error_locked() + waiters = self._overlapping_lease_waiters_locked( + file_offset, nbytes + ) + if not waiters: + previous = self._find_pending_l2_locked(file_offset, nbytes) + self._record_cache_mutation_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 + await asyncio.gather(*waiters) + except BaseException: + async with self._lock: + if not self._l1.contains_slice(data): + data.close() + self._l1.notify_pool_waiters() + raise async def store_bytes_grouped( self, @@ -304,6 +443,193 @@ async def store_bytes_grouped( total += await self.store_bytes(source, file_offset, nbytes) return total + async def prefetch_bytes_grouped( + self, + spans: list[dict[str, int]], + lease_id: str | None = None, + ) -> PrefetchResult: + """Promote L2-missing portions of spans into the pinned L1 tier. + + Args: + spans: Aligned storage spans containing ``file_offset`` and + ``nbytes``. + lease_id: Optional request identifier that blocks overlapping + stores and retains promoted slices for the later GPU load. + + Returns: + Requested bytes split between existing L1 data and L2 reads. + + Raises: + NotImplementedError: If the L2 tier is disabled. + + Async/thread-safety: + Metadata is protected by the transfer lock; io_uring reads are + awaited through the existing executor-backed read path. + """ + if self._l2 is None: + raise NotImplementedError("prefetch requires the io_uring L2 tier") + + requested_ranges = _normalize_ranges(spans) + requested_bytes = sum(size for _start, size in requested_ranges) + if lease_id is not None and requested_bytes > self._l1_bytes: + raise MemoryError( + f"request lease needs {requested_bytes} bytes but L1 capacity is " + f"{self._l1_bytes}" + ) + l1_bytes = 0 + misses: list[dict[str, int]] = [] + pending: list[asyncio.Task[None]] = [] + try: + async with self._lock: + self._raise_l2_error_locked() + if lease_id is not None: + self._replace_request_lease_locked(lease_id, requested_ranges, []) + for file_offset, nbytes in requested_ranges: + self._check_range(file_offset, nbytes) + l1_hits, span_misses = self._l1.resolve_subranges( + target_offset=0, + file_offset=file_offset, + nbytes=nbytes, + ) + if l1_hits: + self._l1.record_hits(l1_hits) + l1_bytes += sum(hit.nbytes for hit in l1_hits) + if lease_id is not None: + self._attach_l1_hits_to_lease_locked(lease_id, l1_hits) + for miss in span_misses: + pending.extend( + self._find_pending_l2_locked( + int(miss["file_offset"]), + int(miss["nbytes"]), + ) + ) + misses.append( + { + "target_offset": 0, + "file_offset": int(miss["file_offset"]), + "nbytes": int(miss["nbytes"]), + } + ) + + if pending: + await asyncio.gather(*set(pending)) + if misses: + await self._load_l2_misses_grouped(None, misses, lease_id=lease_id) + l2_bytes = sum(int(miss["nbytes"]) for miss in misses) + async with self._lock: + if lease_id is not None: + self._require_complete_lease_locked(lease_id) + self._stats.prefetch_requests += 1 + self._stats.prefetch_l1_bytes += l1_bytes + self._stats.prefetch_l2_bytes += l2_bytes + return PrefetchResult(requested_bytes, l1_bytes, l2_bytes) + except BaseException: + if lease_id is not None: + await self.release_lease(lease_id) + raise + + async def classify_and_acquire_lease( + self, + lease_id: str, + spans: list[dict[str, int]], + ) -> TransferTier: + """Classify exact spans and atomically retain an all-L1 request window.""" + requested_ranges = _normalize_ranges(spans) + if not requested_ranges: + return "l2" + async with self._lock: + self._raise_l2_error_locked() + existing = self._request_leases.get(lease_id) + if existing is not None: + if existing.remaining == requested_ranges: + return "l1" + self._release_request_lease_locked(lease_id) + + hits: list[L1RangeHit] = [] + l1_bytes = 0 + requested_bytes = 0 + has_miss = False + for file_offset, nbytes in requested_ranges: + self._check_range(file_offset, nbytes) + requested_bytes += nbytes + span_hits, misses = self._l1.resolve_subranges( + target_offset=0, + file_offset=file_offset, + nbytes=nbytes, + ) + hits.extend(span_hits) + l1_bytes += sum(hit.nbytes for hit in span_hits) + has_miss = has_miss or bool(misses) + if not has_miss: + self._l1.record_hits(hits) + self._replace_request_lease_locked(lease_id, requested_ranges, hits) + return "l1" + return "mixed" if l1_bytes else "l2" + + async def load_leased_bytes_grouped( + self, + dst: Any, + spans: list[dict[str, int]], + lease_id: str, + ) -> int: + """Copy request-leased L1 bytes without re-resolving cache metadata.""" + total = 0 + chunks: list[tuple[int, PinnedMemorySlice, int, int]] = [] + async with self._lock: + self._raise_l2_error_locked() + lease = self._request_leases.get(lease_id) + if lease is None: + raise KeyError(f"unknown transfer lease: {lease_id}") + lease.active_loads += 1 + try: + for span in spans: + target_offset = int(span.get("target_offset", 0)) + file_offset = int(span["file_offset"]) + nbytes = int(span["nbytes"]) + self._check_range(file_offset, nbytes) + total += nbytes + chunks.extend( + self._resolve_leased_range_locked( + lease, + target_offset, + file_offset, + nbytes, + ) + ) + if chunks: + self._stats.l1_hits += len(chunks) + self._copy_grouped_to_dst(dst, chunks) + except BaseException: + lease.active_loads -= 1 + if lease.release_all and lease.active_loads == 0: + self._release_request_lease_locked(lease_id, force=True) + raise + return total + + async def release_lease_ranges( + self, + lease_id: str, + spans: list[dict[str, int]], + ) -> None: + """Release staged physical ranges after the IPC destination is synced.""" + removals = _normalize_ranges(spans) + async with self._lock: + lease = self._request_leases.get(lease_id) + if lease is None: + return + if lease.active_loads > 0: + lease.active_loads -= 1 + if lease.release_all: + if lease.active_loads == 0: + self._release_request_lease_locked(lease_id, force=True) + return + self._consume_lease_ranges_locked(lease_id, removals) + + async def release_lease(self, lease_id: str) -> None: + """Idempotently release every remaining range for one request.""" + async with self._lock: + self._release_request_lease_locked(lease_id) + async def drain(self) -> None: """Wait until all pending L2 writes have completed. @@ -323,6 +649,11 @@ async def drain(self) -> None: def close(self) -> None: """Close the L2 file handle after pending writes are drained/cancelled.""" + for lease in self._request_leases.values(): + if not lease.released.done(): + lease.released.set_result(None) + self._request_leases.clear() + self._leased_slices.clear() for task in self._pending_l2.values(): task.cancel() self._pending_l2.clear() @@ -334,9 +665,11 @@ def close(self) -> None: self._l2.close() self._l1.close() - 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 + def _is_slice_pinned(self, key: tuple[int, int], data: PinnedMemorySlice) -> bool: + """Return whether an L2 writer or request lease still owns ``data``.""" + return ( + self._pending_l2_buffers.get(key) is data or id(data) in self._leased_slices + ) @property def l1_bytes_used(self) -> int: @@ -458,11 +791,8 @@ async def _track_l2_write( 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() + if pending_buffer is not None: + self._close_unowned_slice_locked(pending_buffer) async def _write_l2_async( self, @@ -497,11 +827,8 @@ async def _write_l2_async( 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() + if pending_buffer is not None: + self._close_unowned_slice_locked(pending_buffer) def _raise_l2_error_locked(self) -> None: """Raise and clear the first asynchronous L2 write failure.""" @@ -512,37 +839,45 @@ def _raise_l2_error_locked(self) -> None: async def _load_l2_misses_grouped( self, - dst: Any, + dst: Any | None, misses: list[dict[str, int]], + lease_id: str | None = None, ) -> None: """Read grouped L2 misses concurrently, then promote in request order.""" start = 0 while start < len(misses): batch = self._next_l2_miss_batch(misses, start) - await self._load_l2_miss_batch(dst, batch) + await self._load_l2_miss_batch(dst, batch, lease_id=lease_id) start += len(batch) async def _load_l2_miss_batch( self, - dst: Any, + dst: Any | None, misses: list[dict[str, int]], + lease_id: str | None = None, ) -> None: """Read one bounded L2 miss batch and promote it to L1.""" loop = asyncio.get_event_loop() if self._l2 is None: raise RuntimeError("L2 reads are disabled when skip_l2 is true") - reads: list[tuple[dict[str, int], PinnedMemorySlice]] = [] + reads: list[tuple[dict[str, int], PinnedMemorySlice, int, int]] = [] future_to_read: dict[ asyncio.Future[int], - tuple[dict[str, int], PinnedMemorySlice], + tuple[dict[str, int], PinnedMemorySlice, int, int], ] = {} - promoted: set[int] = set() try: for span in misses: nbytes = int(span["nbytes"]) key = (int(span["file_offset"]), nbytes) - pinned = await self._reserve_l1_buffer(key, nbytes) - reads.append((span, pinned)) + ( + pinned, + promotion_id, + epoch, + pending_writes, + ) = await self._reserve_l1_promotion(key, nbytes) + reads.append((span, pinned, promotion_id, epoch)) + if pending_writes: + await asyncio.gather(*set(pending_writes)) future = loop.run_in_executor( self._l2.executor, self._read_l2_into, @@ -550,7 +885,7 @@ async def _load_l2_miss_batch( pinned, self._next_uring(), ) - future_to_read[future] = (span, pinned) + future_to_read[future] = (span, pinned, promotion_id, epoch) pending: set[asyncio.Future[int]] = set(future_to_read) while pending: @@ -560,35 +895,56 @@ async def _load_l2_miss_batch( ) for future in done: future.result() - span, pinned = future_to_read[future] - self._copy_grouped_to_dst( - dst, - [ - ( - int(span["target_offset"]), - pinned, - 0, - int(span["nbytes"]), - ) - ], - ) + span, pinned, promotion_id, epoch = future_to_read[future] + if dst is not None: + self._copy_grouped_to_dst( + dst, + [ + ( + int(span["target_offset"]), + pinned, + 0, + int(span["nbytes"]), + ) + ], + ) async with self._lock: self._raise_l2_error_locked() nbytes = int(span["nbytes"]) key = (int(span["file_offset"]), nbytes) self._stats.l2_reads += 1 - self._l1.put(key, pinned) - promoted.add(id(pinned)) - except BaseException: + stale = self._promotion_is_stale_locked( + int(span["file_offset"]), nbytes, epoch + ) + if stale: + pinned.close() + self._l1.notify_pool_waiters() + else: + self._l1.put(key, pinned) + if lease_id is not None: + self._attach_l1_hits_to_lease_locked( + lease_id, + [ + L1RangeHit( + target_offset=0, + key=key, + data=pinned, + source_offset=0, + nbytes=nbytes, + ) + ], + ) + self._finish_l1_promotion_locked(promotion_id) + finally: if future_to_read: await asyncio.gather(*future_to_read, return_exceptions=True) - live_buffers = set() async with self._lock: - live_buffers = self._l1.resident_slice_ids() - for _span, pinned in reads: - if id(pinned) not in live_buffers and id(pinned) not in promoted: - pinned.close() - raise + live_buffers = self._l1.resident_slice_ids() | set(self._leased_slices) + for _span, pinned, promotion_id, _epoch in reads: + self._finish_l1_promotion_locked(promotion_id) + if id(pinned) not in live_buffers: + pinned.close() + self._l1.notify_pool_waiters() def _next_l2_miss_batch( self, @@ -652,39 +1008,52 @@ async def _store_bytes_grouped_l1_only( spans: list[dict[str, Any]], ) -> int: """Store grouped spans in L1 without scheduling L2 persistence.""" - total = 0 - async with self._lock: - self._raise_l2_error_locked() - 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) - key = (file_offset, nbytes) - hit = self._l1.find(file_offset) - source = self._slice_src(src, source_offset, nbytes) - if hit is not None: - hit_key, cached, target_offset = hit - if target_offset + nbytes <= len(cached): - self._copy_src_to_pinned_at( - source, cached, target_offset, nbytes + while True: + async with self._lock: + self._raise_l2_error_locked() + waiters = { + waiter + for span in spans + for waiter in self._overlapping_lease_waiters_locked( + int(span["file_offset"]), + int(span["nbytes"]), + ) + } + if not waiters: + total = 0 + 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) + key = (file_offset, nbytes) + hit = self._l1.find(file_offset) + source = self._slice_src(src, source_offset, nbytes) + if hit is not None: + hit_key, cached, target_offset = hit + if target_offset + nbytes <= len(cached): + self._copy_src_to_pinned_at( + source, cached, target_offset, nbytes + ) + self._l1.touch(hit_key) + self._record_cache_mutation_locked(file_offset, nbytes) + total += nbytes + continue + data = self._l1.reserve_or_raise( + key, + nbytes, + preserve_overlaps=True, ) - self._l1.touch(hit_key) + try: + self._copy_src_to_pinned_at(source, data, 0, nbytes) + except BaseException: + data.close() + raise + self._record_cache_mutation_locked(file_offset, nbytes) + self._l1.put(key, data) total += nbytes - continue - data = self._l1.reserve_or_raise( - key, - nbytes, - preserve_overlaps=True, - ) - try: - self._copy_src_to_pinned_at(source, data, 0, nbytes) - except BaseException: - data.close() - raise - self._l1.put(key, data) - total += nbytes - return total + return total + await asyncio.gather(*waiters) async def _reserve_l1_buffer( self, @@ -714,9 +1083,320 @@ async def _reserve_l1_buffer( # Pool exhausted with no evictable victim: an in-flight L2 # write still pins a slice. Wait for one to finish, then retry. wait_for: asyncio.Future[None] | None = next( - iter(self._pending_l2.values()), None + (task for task in self._pending_l2.values() if not task.done()), + None, + ) + if wait_for is None: + wait_for = next( + ( + future + for future in self._pending_l1_promotions.values() + if not future.done() + ), + None, + ) + if wait_for is None: + wait_for = asyncio.get_event_loop().create_future() + self._l1.register_pool_waiter(wait_for) + await wait_for + + async def _reserve_l1_promotion( + self, + key: tuple[int, int], + nbytes: int, + ) -> tuple[ + PinnedMemorySlice, + int, + int, + list[asyncio.Task[None]], + ]: + """Reserve and register a pinned slice for one L2 promotion. + + Args: + key: L2 range being promoted. + nbytes: Number of bytes to reserve. + + Returns: + The pinned slice, reservation identifier, cache mutation epoch, and + L2 writes that must finish before the read starts. + + Async/thread-safety: + The reservation and its waiter are registered while the transfer + lock is held. Other stores therefore cannot mistake an in-flight + promotion for free pool capacity. + """ + while True: + async with self._lock: + self._raise_l2_error_locked() + data = self._l1.reserve(key, nbytes) + if data is not None: + promotion_id = id(data) + self._pending_l1_promotions[promotion_id] = ( + asyncio.get_event_loop().create_future() + ) + self._pending_l1_promotion_epochs[promotion_id] = self._cache_epoch + pending_writes = self._find_pending_l2_locked(key[0], key[1]) + return ( + data, + promotion_id, + self._cache_epoch, + pending_writes, + ) + wait_for: asyncio.Future[None] | None = next( + (task for task in self._pending_l2.values() if not task.done()), + None, ) + if wait_for is None: + wait_for = next( + ( + future + for future in self._pending_l1_promotions.values() + if not future.done() + ), + None, + ) if wait_for is None: wait_for = asyncio.get_event_loop().create_future() self._l1.register_pool_waiter(wait_for) await wait_for + + def _record_cache_mutation_locked(self, file_offset: int, nbytes: int) -> None: + """Record a newly published store for concurrent promotion checks.""" + self._cache_epoch += 1 + self._cache_mutations.append( + (self._cache_epoch, file_offset, file_offset + nbytes) + ) + + def _promotion_is_stale_locked( + self, + file_offset: int, + nbytes: int, + epoch: int, + ) -> bool: + """Return whether a newer overlapping store invalidated a promotion.""" + end = file_offset + nbytes + if self._l1.has_overlap(file_offset, nbytes): + return True + return any( + mutation_epoch > epoch + and mutation_start < end + and file_offset < mutation_end + for mutation_epoch, mutation_start, mutation_end in self._cache_mutations + ) + + def _finish_l1_promotion_locked(self, promotion_id: int) -> None: + """Release one promotion reservation and wake blocked pool users.""" + waiter = self._pending_l1_promotions.pop(promotion_id, None) + self._pending_l1_promotion_epochs.pop(promotion_id, None) + if waiter is not None and not waiter.done(): + waiter.set_result(None) + if self._pending_l1_promotion_epochs: + oldest_epoch = min(self._pending_l1_promotion_epochs.values()) + self._cache_mutations = [ + mutation + for mutation in self._cache_mutations + if mutation[0] > oldest_epoch + ] + else: + self._cache_mutations.clear() + + def _replace_request_lease_locked( + self, + lease_id: str, + ranges: list[tuple[int, int]], + hits: list[L1RangeHit], + ) -> None: + """Replace one request lease while the transfer metadata lock is held.""" + if not lease_id: + raise ValueError("lease_id must not be empty") + existing = self._request_leases.get(lease_id) + if existing is not None and existing.active_loads: + raise RuntimeError(f"cannot replace active transfer lease: {lease_id}") + self._release_request_lease_locked(lease_id) + self._request_leases[lease_id] = _RequestLease( + remaining=list(ranges), + hits=[], + slice_ids=set(), + released=asyncio.get_running_loop().create_future(), + ) + self._attach_l1_hits_to_lease_locked(lease_id, hits) + + def _attach_l1_hits_to_lease_locked( + self, + lease_id: str, + hits: list[L1RangeHit], + ) -> None: + """Retain resident slice references for one admitted request.""" + lease = self._request_leases.get(lease_id) + if lease is None: + raise RuntimeError( + f"transfer lease was released during prefetch: {lease_id}" + ) + new_ids: set[int] = set() + for hit in hits: + lease.hits.append( + _LeasedSliceRange( + file_offset=hit.key[0] + hit.source_offset, + nbytes=hit.nbytes, + data=hit.data, + source_offset=hit.source_offset, + ) + ) + data_id = id(hit.data) + if data_id not in lease.slice_ids: + new_ids.add(data_id) + lease.slice_ids.add(data_id) + lease.hits.sort(key=lambda item: item.file_offset) + for data_id in new_ids: + data = next(hit.data for hit in hits if id(hit.data) == data_id) + existing = self._leased_slices.get(data_id) + if existing is None: + self._leased_slices[data_id] = (data, 1) + else: + self._leased_slices[data_id] = (existing[0], existing[1] + 1) + + def _require_complete_lease_locked(self, lease_id: str) -> None: + """Raise unless retained slice ranges cover the whole request lease.""" + lease = self._request_leases.get(lease_id) + if lease is None: + raise RuntimeError( + f"transfer lease was released during prefetch: {lease_id}" + ) + uncovered = list(lease.remaining) + for hit in lease.hits: + uncovered = _subtract_ranges( + uncovered, + [(hit.file_offset, hit.nbytes)], + ) + if uncovered: + raise RuntimeError(f"prefetch did not retain complete lease: {lease_id}") + + def _resolve_leased_range_locked( + self, + lease: _RequestLease, + target_offset: int, + file_offset: int, + nbytes: int, + ) -> list[tuple[int, PinnedMemorySlice, int, int]]: + """Resolve one load span exclusively through retained lease references.""" + chunks: list[tuple[int, PinnedMemorySlice, int, int]] = [] + cursor = file_offset + end = file_offset + nbytes + while cursor < end: + covering = next( + ( + hit + for hit in lease.hits + if hit.file_offset <= cursor < hit.file_offset + hit.nbytes + ), + None, + ) + if covering is None: + raise KeyError( + f"request lease does not cover load range [{file_offset}, {end})" + ) + covered = min(end, covering.file_offset + covering.nbytes) - cursor + chunks.append( + ( + target_offset + cursor - file_offset, + covering.data, + covering.source_offset + cursor - covering.file_offset, + covered, + ) + ) + cursor += covered + return chunks + + def _consume_lease_ranges_locked( + self, + lease_id: str, + removals: list[tuple[int, int]], + ) -> None: + """Remove staged ranges and release slice references no longer needed.""" + lease = self._request_leases.get(lease_id) + if lease is None: + return + old_slice_ids = set(lease.slice_ids) + lease.remaining = _subtract_ranges(lease.remaining, removals) + lease.hits = _subtract_leased_hits(lease.hits, removals) + lease.slice_ids = {id(hit.data) for hit in lease.hits} + for data_id in old_slice_ids - lease.slice_ids: + self._release_slice_reference_locked(data_id) + if lease.remaining: + return + self._request_leases.pop(lease_id, None) + if not lease.released.done(): + lease.released.set_result(None) + self._l1.notify_pool_waiters() + + def _release_request_lease_locked( + self, + lease_id: str, + *, + force: bool = False, + ) -> None: + """Drop one complete lease and wake overlapping store waiters.""" + lease = self._request_leases.get(lease_id) + if lease is None: + return + if lease.active_loads and not force: + lease.release_all = True + return + self._request_leases.pop(lease_id, None) + for data_id in lease.slice_ids: + self._release_slice_reference_locked(data_id) + if not lease.released.done(): + lease.released.set_result(None) + self._l1.notify_pool_waiters() + + def _release_slice_reference_locked(self, data_id: int) -> None: + """Release one request's reference to a retained pinned slice.""" + retained = self._leased_slices.get(data_id) + if retained is None: + return + data, count = retained + if count > 1: + self._leased_slices[data_id] = (data, count - 1) + return + self._leased_slices.pop(data_id, None) + self._close_unowned_slice_locked(data) + + def _close_unowned_slice_locked(self, data: PinnedMemorySlice) -> None: + """Close a detached slice after its final writer/lease reference leaves.""" + if self._l1.contains_slice(data): + return + if any(buffer is data for buffer in self._pending_l2_buffers.values()): + return + if id(data) in self._leased_slices: + return + data.close() + self._l1.notify_pool_waiters() + + def _overlapping_lease_waiters_locked( + self, + file_offset: int, + nbytes: int, + ) -> list[asyncio.Future[None]]: + """Return active request-lease waiters overlapping one store range.""" + end = file_offset + nbytes + return [ + lease.released + for lease in self._request_leases.values() + if any( + start < end and file_offset < start + size + for start, size in lease.remaining + ) + ] + + async def _wait_for_overlapping_leases( + self, + file_offset: int, + nbytes: int, + ) -> None: + """Wait until no request lease protects an overlapping physical range.""" + while True: + async with self._lock: + waiters = self._overlapping_lease_waiters_locked(file_offset, nbytes) + if not waiters: + return + await asyncio.gather(*waiters) diff --git a/docs/design/architecture.md b/docs/design/architecture.md index a1f10ac..92f7b09 100644 --- a/docs/design/architecture.md +++ b/docs/design/architecture.md @@ -224,18 +224,22 @@ IPC server 不提供文档管理 op。文档生命周期只属于 HTTP server ### 两阶段提交 `alloc_chunk` 只预留 slot 和 metadata,不把 chunk 插入 `RetrievalIndex`。 -vLLM worker 完成 KV 写入后调用 `commit_chunk`,chunk 才对 lookup 可见。 -这避免了部分写入的数据被其他请求读到。 +worker transfer 完成后由 DaseR server 检查完整的 L1 range coverage,并调用 +现有的 `commit_chunk` 路径,chunk 才对 lookup 可见。这避免了部分写入的 +数据被其他请求读到,同时避免 worker 发起重复 commit。 ### Worker 侧 staging,server 侧 transfer -store 路径不再每层单独发一次写 IO。Worker 在 `wait_for_save` 中按容量上限 -构造一个或多个 slot-major staging view,把待保存 blocks 的每层 KV 拷入 -staging,导出 CUDA IPC handle,并通过 IPC 请求 server 执行 transfer。server -写完对应 batch 后,worker 再统一调用 `commit_chunk`。staging 由 worker 侧 -小型 `CudaStagingPool` 复用,初始化时预分配一个 bounded buffer;单批和未完成 -后台 batch 的字节上限会根据 vLLM 分配 KV cache 后的当前可用显存推导,避免 -固定挤占显存。 +store 路径不再每层单独发一次写 IO。Worker 在 `wait_for_save` 中按请求进入 +顺序将 store 意图加入 FIFO,等 vLLM 报告 request finished 后再提交。后台 +`daser-store-io` loop 按 staging pool depth 限制 snapshot/transfer 并发,前一 +个任务释放 buffer 后自动推进下一个任务;不依赖后续 vLLM connector step。 +它构造一个或多个 slot-major staging view,把待保存 blocks 的每层 KV 拷入 +staging,导出 CUDA IPC handle,并通过 IPC 请求 server 执行 transfer。DaseR +server 在完整 L1 coverage 后负责 `commit_chunk`,而不是由 worker 发起第二个 +commit 路径。staging 由 worker 侧小型 `CudaStagingPool` 复用,初始化时预分配 +bounded buffer;单批和未完成后台 batch 的字节上限会根据 vLLM 分配 KV cache +后的当前可用显存推导,避免固定挤占显存。 load 路径在 `start_load_kv` 中把本 step 的命中 chunk 拆成 bounded staging batch,导出 CUDA IPC handle,请求 server 读回 spans,再按层批量拷回 vLLM diff --git a/docs/design/flows.md b/docs/design/flows.md index 343a046..0312510 100644 --- a/docs/design/flows.md +++ b/docs/design/flows.md @@ -107,6 +107,7 @@ sequenceDiagram participant WR as WorkerRuntime participant BG as StorePipeline / daser-store-io participant IPC as IPC server + participant C as ServerCore participant TL as TransferLayer W->>WR: bind_connector_metadata(reqs_to_store) @@ -126,12 +127,16 @@ sequenceDiagram TL-->>IPC: bytes accepted IPC-->>BG: stored chunk keys end - BG->>IPC: commit_chunks(...) after all batch futures complete + IPC->>C: record accepted ranges after transfer completion + C->>C: commit complete chunk coverage through existing TP quorum ``` -`wait_for_save` 在当前 worker step 内完成 KV -> staging snapshot,然后把 -server transfer 交给后台 `daser-store-io` loop。未完成 batch 的 staging lease 由 -future 持有,完成后归还 `CudaStagingPool`;`shutdown` 会阻塞等待所有 pending +`wait_for_save` 只把当前 worker step 的 KV store 意图加入待完成 FIFO,不会 +立即读取 live KV cache。`get_finished` 收到 vLLM 的 `finished_req_ids` 后, +按请求进入 FIFO 的顺序提交全部已完成请求。后台 `daser-store-io` loop 使用 +staging pool depth 的异步信号量限制实际 snapshot/transfer 并发;前一个 store +释放 staging lease 后,队列中的下一个 store 自动进入,不依赖后续 vLLM +connector step 或 benchmark dummy poll。`shutdown` 会阻塞等待所有 pending store future 完成。 ### 阶段三:Commit 发布 @@ -142,9 +147,10 @@ sequenceDiagram participant C as ServerCore participant RI as RetrievalIndex - IPC->>C: commit_chunk(chunk_key) + IPC->>C: record_store_ranges(chunk spans) + C->>C: commit complete coverage through TP quorum C->>RI: insert(ChunkMeta) - C-->>IPC: ok + C-->>IPC: transfer response ``` commit 完成后 chunk 才能被 lookup 命中。 diff --git a/docs/optimizations/2_e2e_lmcache_parity.md b/docs/optimizations/2_e2e_lmcache_parity.md index d9a60d7..6ece5da 100644 --- a/docs/optimizations/2_e2e_lmcache_parity.md +++ b/docs/optimizations/2_e2e_lmcache_parity.md @@ -53,9 +53,13 @@ Cold profiling showed the cold pass was dominated by save-side work: | `wait_for_save` writes | 3.15 s | | commit RPCs | 0.09 s | -`wait_for_save` now submits a background write-and-commit task and returns -without waiting for NVMe completion. Chunks are still inserted into the DaseR -index only after their writes finish, so readers cannot observe partial data. +`wait_for_save` now submits a background transfer task and returns without +waiting for NVMe completion. The DaseR server records accepted transfer ranges +and inserts a chunk into the index only after its full KV coverage is present +in the transfer destination, so readers cannot observe partial data. For +io_uring, this destination is L1 and L2 persistence remains asynchronous; GDS +uses the completed direct transfer. The worker no longer owns the final +`commit_chunks` call. The staging tensor is independent from vLLM's KV cache, so vLLM can safely reuse KV blocks while the background task drains. diff --git a/tests/connector/test_daser_connector.py b/tests/connector/test_daser_connector.py index 6ed81ed..a165575 100644 --- a/tests/connector/test_daser_connector.py +++ b/tests/connector/test_daser_connector.py @@ -1,6 +1,8 @@ # SPDX-License-Identifier: Apache-2.0 # Standard +import threading +import time from types import SimpleNamespace # Third Party @@ -18,6 +20,7 @@ rolling_prefix_key, rolling_prefix_keys, ) +from daser.connector.ipc_client import PrefetchLookupResult from daser.connector.metadata import ( DaserConnectorMeta, ReqLoadSpec, @@ -98,6 +101,103 @@ def test_rolling_prefix_keys_match_single_step_helper() -> None: ) +def test_scheduler_pipelines_l1_hit_past_bounded_l2_prefetch(monkeypatch) -> None: + """All-L1 admission bypasses two active L2 prefetches, but L2 stays bounded.""" + + class LookupIPC: + def __init__(self) -> None: + self.lookup_calls: list[str] = [] + self.released: list[str] = [] + + def lookup_with_prefetch( + self, + lease_id, + tokens, + model_id, + **kwargs, + ): + del tokens, model_id, kwargs + self.lookup_calls.append(lease_id) + index = len(self.lookup_calls) + tier = "l1" if lease_id == "l1-ready" else "l2" + return PrefetchLookupResult( + chunks=[ + { + "chunk_key": f"cached-{index}", + "start_slot": index * 2, + "num_slots": 2, + "file_offset": index * 64, + "token_count": 8, + "target_token_start": 0, + "pos_offset": 0, + } + ], + spans=[{"file_offset": index * 64, "nbytes": 64}], + tier=tier, + ) + + def release_transfer_lease(self, lease_id: str) -> None: + self.released.append(lease_id) + + prefetch_calls: list[tuple[str, str, list[dict[str, int]]]] = [] + release_prefetch = threading.Event() + + def fake_prefetch(socket_path, lease_id, spans): + prefetch_calls.append((socket_path, lease_id, spans)) + release_prefetch.wait(timeout=5.0) + return {"requested_bytes": 64, "l1_bytes": 0, "l2_bytes": 64} + + monkeypatch.setattr( + "daser.connector.scheduler.lifecycle._prefetch_external_spans", + fake_prefetch, + ) + ipc = LookupIPC() + lifecycle = RequestLifecycle( + ipc_client=ipc, + socket_path="/unused/daser.sock", + block_tokens=4, + slot_size=32, + model_id="model", + cache_reuse_mode="chunk", + runtime_config_ready=True, + prefetch_max_requests=2, + ) + + def request(req_id: str) -> SimpleNamespace: + return SimpleNamespace( + request_id=req_id, + prompt_token_ids=list(range(12)), + kv_transfer_params={"daser_skip_save": True}, + ) + + try: + assert lifecycle.get_num_new_matched_tokens(request("l2-a"), 0) == ( + None, + True, + ) + assert lifecycle.get_num_new_matched_tokens(request("l2-b"), 0) == ( + None, + True, + ) + deadline = time.monotonic() + 1.0 + while len(prefetch_calls) < 2 and time.monotonic() < deadline: + time.sleep(0.005) + assert len(prefetch_calls) == 2 + + assert lifecycle.get_num_new_matched_tokens(request("l1-ready"), 0) == ( + 8, + True, + ) + assert lifecycle.get_num_new_matched_tokens(request("l2-deferred"), 0) == ( + None, + True, + ) + assert len(prefetch_calls) == 2 + finally: + release_prefetch.set() + lifecycle.shutdown() + + class _SchedulerProbe(RequestLifecycle): """Scheduler-side probe that can emulate deferred runtime config.""" @@ -2891,6 +2991,8 @@ async def transfer_store_cuda(self, **kwargs): pipeline = StorePipeline.__new__(StorePipeline) pipeline._client = Client() # noqa: SLF001 + pipeline._tp_rank = 0 # noqa: SLF001 + pipeline._tp_size = 1 # noqa: SLF001 buffer = SimpleNamespace(device=torch.device("cuda:1"), nbytes=32) staged = StagedStoreBatch(buffer=buffer, spans=[], lease=object()) cupy_buffer = object() diff --git a/tests/connector/test_ipc_client.py b/tests/connector/test_ipc_client.py index a542eef..134058a 100644 --- a/tests/connector/test_ipc_client.py +++ b/tests/connector/test_ipc_client.py @@ -73,6 +73,34 @@ async def test_sync_client_lookup(tmp_path): await server.stop() +def test_sync_client_validates_prefetch_lookup_contract( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Scheduler lookup decodes exact spans and the host-tier state once.""" + + def fake_call(self, payload: dict) -> dict: + del self + assert payload["lease_id"] == "request-1" + return { + "chunks": [{"chunk_key": "cached"}], + "spans": [{"file_offset": 4096, "nbytes": 8192}], + "tier": "mixed", + } + + monkeypatch.setattr(IPCClientSync, "call", fake_call) + client = IPCClientSync("unused.sock") + + result = client.lookup_with_prefetch( + "request-1", + [1, 2, 3, 4], + "m", + external_prefix_queries=4, + ) + + assert result.tier == "mixed" + assert result.spans == [{"file_offset": 4096, "nbytes": 8192}] + + @pytest.mark.asyncio async def test_sync_client_get_runtime_config(tmp_path): server = make_server(tmp_path) @@ -362,12 +390,14 @@ async def fake_call(self, payload: dict) -> dict: allocation_offset=2272, producer_pid=43, spans=[{"target_offset": 0, "nbytes": 2048, "file_offset": 0}], + lease_id="request-cuda", ) assert recorded[0]["payload"]["allocation_base_ptr"] == 122880 assert recorded[0]["payload"]["allocation_offset"] == 576 assert recorded[1]["payload"]["allocation_base_ptr"] == 221184 assert recorded[1]["payload"]["allocation_offset"] == 2272 + assert recorded[1]["lease_id"] == "request-cuda" @pytest.mark.asyncio @@ -401,6 +431,7 @@ async def fake_call(self, payload: dict) -> dict: producer_pid=43, nbytes=16, spans=[{"target_offset": 0, "nbytes": 16, "file_offset": 0}], + lease_id="request-registered", ) assert response == {"ok": True, "bytes": 16} @@ -426,6 +457,7 @@ async def fake_call(self, payload: dict) -> dict: "nbytes": 16, }, "spans": [{"target_offset": 0, "nbytes": 16, "file_offset": 0}], + "lease_id": "request-registered", }, ] diff --git a/tests/connector/test_worker_pipelines.py b/tests/connector/test_worker_pipelines.py index 686cbe7..2ab7b02 100644 --- a/tests/connector/test_worker_pipelines.py +++ b/tests/connector/test_worker_pipelines.py @@ -35,7 +35,7 @@ def _store_spec(key: str, blocks: list[int]) -> ReqStoreSpec: return ReqStoreSpec(key, 0, len(blocks), blocks, 0, len(blocks)) -def test_store_pipeline_defers_and_dispatches_up_to_pool_depth() -> None: +def test_store_pipeline_dispatches_finished_saves_in_fifo_order() -> None: pipeline = StorePipeline.__new__(StorePipeline) pipeline._pending_finished_saves = {} # noqa: SLF001 pipeline._staging_pool = SimpleNamespace(depth=2) # noqa: SLF001 @@ -43,19 +43,18 @@ def test_store_pipeline_defers_and_dispatches_up_to_pool_depth() -> None: def submit(save: Any) -> None: future = _ManualFuture() - save.future = future submitted.append(future) + save.future = future pipeline._submit_save = submit # type: ignore[method-assign] # noqa: SLF001 pipeline.queue_finished( {req: _store_spec(req, [index]) for index, req in enumerate(("a", "b", "c"))}, - {"a", "b", "c"}, ) assert pipeline.collect_finished(set()) == set() assert submitted == [] assert pipeline.collect_finished({"a", "b", "c"}) == set() - assert len(submitted) == 2 + assert len(submitted) == 3 submitted[0].complete = True assert pipeline.collect_finished(set()) == {"a"} @@ -65,10 +64,41 @@ def submit(save: Any) -> None: assert pipeline.collect_finished(set()) == {"b", "c"} +@pytest.mark.asyncio +async def test_store_dispatcher_bounds_and_orders_background_saves() -> None: + """Finished saves run FIFO while respecting the staging depth.""" + pipeline = StorePipeline.__new__(StorePipeline) + pipeline._store_capacity = 1 # noqa: SLF001 + pipeline._store_semaphore = None # noqa: SLF001 + active = 0 + max_active = 0 + order: list[str] = [] + + async def save(save: Any, event: Any) -> None: + del event + nonlocal active, max_active + active += 1 + max_active = max(max_active, active) + order.append(save.req_id) + await asyncio.sleep(0) + active -= 1 + + pipeline._store_finished_save = save # type: ignore[method-assign] # noqa: SLF001 + saves = [SimpleNamespace(req_id=req_id) for req_id in ("a", "b", "c")] + await asyncio.gather( + *(pipeline._run_bounded_save(save, None) for save in saves) # noqa: SLF001 + ) + + assert order == ["a", "b", "c"] + assert max_active == 1 + + def test_store_pipeline_streams_request_larger_than_pool_depth() -> None: pipeline = StorePipeline.__new__(StorePipeline) pipeline._pending_finished_saves = {} # noqa: SLF001 pipeline._staging_pool = SimpleNamespace(depth=1) # noqa: SLF001 + pipeline._store_capacity = 1 # noqa: SLF001 + pipeline._store_semaphore = None # noqa: SLF001 pipeline._kv_caches = {"layer": torch.empty(1)} # noqa: SLF001 pipeline._local_slot_size = 16 # noqa: SLF001 pipeline._rank_stride_bytes = 0 # noqa: SLF001 @@ -84,11 +114,6 @@ class Lease: def release(self) -> None: released.append(1) - class Client: - async def commit_chunks(self, keys: list[str], **kwargs: int) -> None: - assert keys == ["large"] - assert kwargs == {"tp_rank": 0, "tp_size": 1} - def stage(block_ids: list[int], spans: list[Any], event: Any) -> StagedStoreBatch: del event return StagedStoreBatch(torch.empty(1), spans, Lease()) @@ -106,13 +131,11 @@ def submit(coro: Any) -> Future[None]: future.set_exception(exc) return future - pipeline._client = Client() # type: ignore[assignment] # noqa: SLF001 pipeline._stage_batch = stage # type: ignore[method-assign] # noqa: SLF001 pipeline._write_cuda_buffer = write # type: ignore[method-assign] # noqa: SLF001 pipeline._submit = submit # type: ignore[method-assign] # noqa: SLF001 - pipeline._submit_save = StorePipeline._submit_save.__get__(pipeline) # noqa: SLF001 pipeline._kv_caches = {} # noqa: SLF001 - pipeline.queue_finished({"large": _store_spec("large", [0, 1, 2, 3, 4])}, {"large"}) + pipeline.queue_finished({"large": _store_spec("large", [0, 1, 2, 3, 4])}) pipeline._kv_caches = {"layer": torch.empty(1)} # noqa: SLF001 assert pipeline.collect_finished({"large"}) == set() @@ -124,11 +147,13 @@ def submit(coro: Any) -> Future[None]: class _LoadClient: def __init__(self, fail_offset: int | None = None) -> None: self.calls: list[int] = [] + self.lease_ids: list[str | None] = [] self.fail_offset = fail_offset async def transfer_load_registered_cuda(self, **kwargs: Any) -> dict[str, Any]: offset = int(kwargs["spans"][0]["file_offset"]) self.calls.append(offset) + self.lease_ids.append(kwargs.get("lease_id")) if offset == self.fail_offset: raise RuntimeError("load failed") return { @@ -191,11 +216,17 @@ def test_load_pipeline_handles_empty_and_multibatch_requests( pipeline.start( { "empty": _load_spec("empty", []), - "large:load:0": _load_spec("large", [0, 1, 2, 3, 4]), + "large:load:0": ReqLoadSpec( + **{ + **vars(_load_spec("large", [0, 1, 2, 3, 4])), + "lease_id": "large", + } + ), } ) assert _wait_finished(pipeline, {"empty", "large"}) == {"empty", "large"} assert client.calls == [0, 32, 64] + assert client.lease_ids == ["large", "large", "large"] finally: pipeline.shutdown() diff --git a/tests/server/test_core.py b/tests/server/test_core.py index 0d0f52c..53c0db9 100644 --- a/tests/server/test_core.py +++ b/tests/server/test_core.py @@ -227,6 +227,98 @@ async def test_tensor_parallel_commit_waits_for_all_distinct_ranks() -> None: assert len(await core.lookup(tokens, "m")) == 1 +@pytest.mark.asyncio +async def test_store_range_commit_waits_for_complete_l1_coverage() -> None: + """Server publication waits until split transfer ranges cover a chunk.""" + core = make_core() + tokens = [1, 2, 3, 4] + key = first_rolling_key(tokens) + alloc = await core.alloc_chunk(key, token_count=len(tokens), model_id="m") + + assert ( + await core.record_store_ranges( + [ + { + "chunk_key": key, + "file_offset": alloc.file_offset, + "nbytes": SLOT_SIZE // 2, + "start_slot": alloc.start_slot, + "num_slots": alloc.num_slots, + } + ], + tp_rank=0, + tp_size=1, + local_slot_size=SLOT_SIZE, + rank_stride_bytes=0, + ) + == [] + ) + assert await core.lookup(tokens, "m") == [] + + assert await core.record_store_ranges( + [ + { + "chunk_key": key, + "file_offset": alloc.file_offset + SLOT_SIZE // 2, + "nbytes": SLOT_SIZE // 2, + "start_slot": alloc.start_slot, + "num_slots": alloc.num_slots, + } + ], + tp_rank=0, + tp_size=1, + local_slot_size=SLOT_SIZE, + rank_stride_bytes=0, + ) == [key] + assert len(await core.lookup(tokens, "m")) == 1 + + +@pytest.mark.asyncio +async def test_store_range_commit_preserves_tp_quorum() -> None: + """Complete L1 coverage on one TP rank is not public until quorum.""" + core = make_core() + tokens = [1, 2, 3, 4] + key = first_rolling_key(tokens) + alloc = await core.alloc_chunk(key, token_count=len(tokens), model_id="m") + local_slot_size = SLOT_SIZE // 2 + rank_stride = 64 * local_slot_size + + rank_span = { + "chunk_key": key, + "start_slot": alloc.start_slot, + "num_slots": alloc.num_slots, + } + assert await core.record_store_ranges( + [ + { + **rank_span, + "file_offset": alloc.start_slot * local_slot_size, + "nbytes": local_slot_size, + } + ], + tp_rank=0, + tp_size=2, + local_slot_size=local_slot_size, + rank_stride_bytes=rank_stride, + ) == [key] + assert await core.lookup(tokens, "m") == [] + + assert await core.record_store_ranges( + [ + { + **rank_span, + "file_offset": rank_stride + alloc.start_slot * local_slot_size, + "nbytes": local_slot_size, + } + ], + tp_rank=1, + tp_size=2, + local_slot_size=local_slot_size, + rank_stride_bytes=rank_stride, + ) == [key] + assert len(await core.lookup(tokens, "m")) == 1 + + @pytest.mark.asyncio async def test_restored_orphan_committed_chunk_can_be_reused(tmp_path) -> None: tokens = [1, 2, 3, 4] diff --git a/tests/server/test_ipc_server.py b/tests/server/test_ipc_server.py index 1994f79..e732301 100644 --- a/tests/server/test_ipc_server.py +++ b/tests/server/test_ipc_server.py @@ -431,6 +431,13 @@ async def test_transfer_store_and_load_with_bytes_payload(tmp_path) -> None: ], }, ) + prefetch = await _send_recv( + str(tmp_path / "test.sock"), + { + "op": "transfer_prefetch", + "spans": [{"nbytes": SLOT_SIZE * 2, "file_offset": 0}], + }, + ) load = await _send_recv( str(tmp_path / "test.sock"), { @@ -448,6 +455,12 @@ async def test_transfer_store_and_load_with_bytes_payload(tmp_path) -> None: ) assert store == {"ok": True, "bytes": SLOT_SIZE * 2, "chunk_keys": []} + assert prefetch == { + "ok": True, + "requested_bytes": SLOT_SIZE * 2, + "l1_bytes": SLOT_SIZE * 2, + "l2_bytes": 0, + } assert load == { "ok": True, "bytes": SLOT_SIZE * 2, @@ -457,6 +470,129 @@ async def test_transfer_store_and_load_with_bytes_payload(tmp_path) -> None: await server.stop() +@pytest.mark.asyncio +async def test_lookup_prefetch_classifies_and_loads_request_lease(tmp_path) -> None: + """Prefetch-aware lookup returns exact spans and a consumable all-L1 lease.""" + core = make_core() + socket_path = str(tmp_path / "test.sock") + server = IPCServer(socket_path, core, make_runtime_config(tmp_path)) + await server.start() + tokens = [1, 2, 3, 4, 5] + key = first_rolling_key(tokens) + try: + alloc = await _send_recv( + socket_path, + { + "op": "alloc_chunk", + "chunk_key": key, + "token_count": BLOCK_TOKENS, + "model_id": "m", + }, + ) + await _send_recv( + socket_path, + { + "op": "transfer_store", + "payload": {"data": b"k" * SLOT_SIZE}, + "spans": [ + { + "source_offset": 0, + "nbytes": SLOT_SIZE, + "file_offset": int(alloc["file_offset"]), + } + ], + }, + ) + await _send_recv( + socket_path, + {"op": "commit_chunk", "chunk_key": key}, + ) + + lookup = await _send_recv( + socket_path, + { + "op": "lookup_prefetch", + "lease_id": "request-1", + "tokens": tokens, + "model_id": "m", + "external_prefix_queries": len(tokens), + "num_computed_tokens": 0, + }, + ) + assert lookup["tier"] == "l1" + assert lookup["spans"] == [ + {"file_offset": int(alloc["file_offset"]), "nbytes": SLOT_SIZE} + ] + + load = await _send_recv( + socket_path, + { + "op": "transfer_load", + "lease_id": "request-1", + "payload": {"return_data": True}, + "spans": [ + { + "target_offset": 0, + "file_offset": int(alloc["file_offset"]), + "nbytes": SLOT_SIZE, + } + ], + }, + ) + assert load["data"] == b"k" * SLOT_SIZE + assert await _send_recv( + socket_path, + {"op": "release_transfer_lease", "lease_id": "request-1"}, + ) == {"ok": True} + finally: + await server.stop() + + +@pytest.mark.asyncio +async def test_transfer_store_commits_chunk_after_l1_copy(tmp_path) -> None: + """The server commits a fully transferred chunk without worker RPC.""" + core = make_core() + server = IPCServer(str(tmp_path / "test.sock"), core, make_runtime_config(tmp_path)) + await server.start() + tokens = [1, 2, 3, 4] + key = first_rolling_key(tokens) + try: + alloc = await _send_recv( + str(tmp_path / "test.sock"), + { + "op": "alloc_chunk", + "chunk_key": key, + "token_count": len(tokens), + "model_id": "m", + }, + ) + store = await _send_recv( + str(tmp_path / "test.sock"), + { + "op": "transfer_store", + "payload": {"data": b"a" * SLOT_SIZE}, + "spans": [ + { + "source_offset": 0, + "nbytes": SLOT_SIZE, + "file_offset": alloc["file_offset"], + "chunk_key": key, + "start_slot": alloc["start_slot"], + "num_slots": alloc["num_slots"], + } + ], + }, + ) + assert store["chunk_keys"] == [key] + lookup = await _send_recv( + str(tmp_path / "test.sock"), + {"op": "lookup", "tokens": tokens, "model_id": "m"}, + ) + assert [chunk["chunk_key"] for chunk in lookup["chunks"]] == [key] + finally: + await server.stop() + + @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_tiered_iouring_transfer.py b/tests/transfer/test_tiered_iouring_transfer.py index adecced..5b225cc 100644 --- a/tests/transfer/test_tiered_iouring_transfer.py +++ b/tests/transfer/test_tiered_iouring_transfer.py @@ -137,6 +137,73 @@ def _copy_grouped_to_dst( super()._copy_grouped_to_dst(dst, chunks) +class DelayedL2ReadCompletionProbe(TieredIOUringTransferLayer): + """Test transfer layer that pauses after selected L2 reads complete.""" + + def __init__( + self, + path: str, + l1_bytes: int, + l2_bytes: int, + delayed_offsets: set[int], + ) -> None: + super().__init__( + path=path, + l1_bytes=l1_bytes, + l2_bytes=l2_bytes, + ) + self.delayed_offsets = delayed_offsets + self.release_read = threading.Event() + self.read_started = threading.Event() + + def _read_l2_into( + self, + file_offset: int, + dst: object, + uring: NativeIOUring, + ) -> int: + """Pause promotion after the worker has captured the old L2 bytes.""" + result = super()._read_l2_into(file_offset, dst, uring) + if file_offset in self.delayed_offsets: + self.delayed_offsets.remove(file_offset) + self.read_started.set() + self.release_read.wait(timeout=5.0) + return result + + +class FailingL2ReadCompletionProbe(TieredIOUringTransferLayer): + """Test transfer layer that fails after claiming a promotion buffer.""" + + def __init__( + self, + path: str, + l1_bytes: int, + l2_bytes: int, + delayed_offsets: set[int], + ) -> None: + super().__init__( + path=path, + l1_bytes=l1_bytes, + l2_bytes=l2_bytes, + ) + self.delayed_offsets = delayed_offsets + self.release_read = threading.Event() + self.read_started = threading.Event() + + def _read_l2_into( + self, + file_offset: int, + dst: object, + uring: NativeIOUring, + ) -> int: + """Block then fail so pool waiters exercise promotion cleanup.""" + if file_offset in self.delayed_offsets: + self.delayed_offsets.remove(file_offset) + self.read_started.set() + self.release_read.wait(timeout=5.0) + raise IOError("synthetic L2 read failure") + + def _run(coro: object) -> object: """Run a coroutine on the current test event loop.""" return asyncio.get_event_loop().run_until_complete(coro) @@ -543,6 +610,49 @@ def test_iouring_promotes_l2_miss_to_l1(tmp_path) -> None: layer.close() +def test_iouring_prefetch_reads_only_l1_missing_ranges(tmp_path) -> None: + """Prefetch promotes only the missing part and reuses it on load.""" + + async def scenario() -> None: + path = str(tmp_path / "daser.store") + writer = TieredIOUringTransferLayer( + path=path, + l1_bytes=ALIGNMENT * 2, + l2_bytes=ALIGNMENT * 3, + ) + try: + await writer.store_bytes(_block(b"a"), 0, ALIGNMENT) + await writer.store_bytes(_block(b"b"), ALIGNMENT, ALIGNMENT) + await writer.drain() + finally: + writer.close() + + layer = TieredIOUringTransferLayer( + path=path, + l1_bytes=ALIGNMENT * 2, + l2_bytes=ALIGNMENT * 3, + ) + try: + await layer.load_bytes(bytearray(ALIGNMENT), 0, ALIGNMENT) + reads_before = layer.stats.l2_reads + result = await layer.prefetch_bytes_grouped( + [{"file_offset": 0, "nbytes": ALIGNMENT * 2}] + ) + assert result.requested_bytes == ALIGNMENT * 2 + assert result.l1_bytes == ALIGNMENT + assert result.l2_bytes == ALIGNMENT + assert layer.stats.l2_reads == reads_before + 1 + + dst = bytearray(ALIGNMENT * 2) + await layer.load_bytes(dst, 0, ALIGNMENT * 2) + assert bytes(dst) == bytes(_block(b"a") + _block(b"b")) + assert layer.stats.l2_reads == reads_before + 1 + finally: + layer.close() + + _run(scenario()) + + def test_iouring_grouped_l2_misses_are_bounded_by_l1_capacity(tmp_path) -> None: """Grouped L2 misses make progress when the request is larger than L1.""" @@ -804,6 +914,163 @@ async def scenario() -> None: _run(scenario()) +def test_iouring_prefetch_does_not_overwrite_concurrent_store(tmp_path) -> None: + """An older delayed promotion must not replace a newer L1 store.""" + + async def scenario() -> None: + path = str(tmp_path / "daser.store") + writer = TieredIOUringTransferLayer( + path=path, + l1_bytes=ALIGNMENT, + l2_bytes=ALIGNMENT * 2, + ) + try: + await writer.store_bytes(_block(b"o"), 0, ALIGNMENT) + await writer.drain() + finally: + writer.close() + + layer = DelayedL2ReadCompletionProbe( + path=path, + l1_bytes=ALIGNMENT * 2, + l2_bytes=ALIGNMENT * 2, + delayed_offsets={0}, + ) + try: + prefetch = asyncio.create_task( + layer.prefetch_bytes_grouped([{"file_offset": 0, "nbytes": ALIGNMENT}]) + ) + assert await asyncio.to_thread(layer.read_started.wait, timeout=2.0) + + await layer.store_bytes(_block(b"n"), 0, ALIGNMENT) + layer.release_read.set() + await asyncio.wait_for(prefetch, timeout=2.0) + await layer.drain() + + dst = bytearray(ALIGNMENT) + assert await layer.load_bytes(dst, 0, ALIGNMENT) == ALIGNMENT + assert bytes(dst) == bytes(_block(b"n")) + finally: + layer.release_read.set() + layer.close() + + _run(scenario()) + + +def test_iouring_failed_prefetch_releases_pool_waiters(tmp_path) -> None: + """A failed promotion must wake stores waiting for its pinned buffer.""" + + async def scenario() -> None: + path = str(tmp_path / "daser.store") + writer = TieredIOUringTransferLayer( + path=path, + l1_bytes=ALIGNMENT, + l2_bytes=ALIGNMENT * 2, + ) + try: + await writer.store_bytes(_block(b"o"), 0, ALIGNMENT) + await writer.drain() + finally: + writer.close() + + layer = FailingL2ReadCompletionProbe( + path=path, + l1_bytes=ALIGNMENT, + l2_bytes=ALIGNMENT * 2, + delayed_offsets={0}, + ) + try: + prefetch = asyncio.create_task( + layer.prefetch_bytes_grouped( + [{"file_offset": 0, "nbytes": ALIGNMENT}], + lease_id="failed-prefetch", + ) + ) + assert await asyncio.to_thread(layer.read_started.wait, timeout=2.0) + store = asyncio.create_task( + layer.store_bytes(_block(b"n"), ALIGNMENT, ALIGNMENT) + ) + await asyncio.sleep(0.05) + assert not store.done() + + layer.release_read.set() + with pytest.raises(IOError, match="synthetic"): + await asyncio.wait_for(prefetch, timeout=2.0) + assert await asyncio.wait_for(store, timeout=2.0) == ALIGNMENT + await layer.drain() + finally: + layer.release_read.set() + layer.close() + + _run(scenario()) + + +def test_iouring_classifies_and_leases_exact_l1_window(tmp_path) -> None: + """Tier classification is exact and a leased L1 range blocks overwrite.""" + + async def scenario() -> None: + path = str(tmp_path / "daser.store") + writer = TieredIOUringTransferLayer( + path=path, + l1_bytes=ALIGNMENT * 2, + l2_bytes=ALIGNMENT * 2, + ) + try: + await writer.store_bytes(_block(b"a"), 0, ALIGNMENT) + await writer.store_bytes(_block(b"b"), ALIGNMENT, ALIGNMENT) + await writer.drain() + finally: + writer.close() + + layer = TieredIOUringTransferLayer( + path=path, + l1_bytes=ALIGNMENT * 2, + l2_bytes=ALIGNMENT * 2, + ) + spans = [ + {"file_offset": 0, "nbytes": ALIGNMENT}, + {"file_offset": ALIGNMENT, "nbytes": ALIGNMENT}, + ] + load_spans = [ + {"target_offset": 0, **spans[0]}, + {"target_offset": ALIGNMENT, **spans[1]}, + ] + try: + assert await layer.classify_and_acquire_lease("all-l2", spans) == "l2" + await layer.load_bytes(bytearray(ALIGNMENT), 0, ALIGNMENT) + assert await layer.classify_and_acquire_lease("mixed", spans) == "mixed" + await layer.prefetch_bytes_grouped(spans, lease_id="leased") + assert await layer.classify_and_acquire_lease("all-l1", spans) == "l1" + await layer.release_lease("all-l1") + + overwrite = asyncio.create_task( + layer.store_bytes(_block(b"n"), 0, ALIGNMENT) + ) + await asyncio.sleep(0.05) + assert not overwrite.done() + + dst = bytearray(ALIGNMENT * 2) + assert ( + await layer.load_leased_bytes_grouped(dst, load_spans, "leased") + == ALIGNMENT * 2 + ) + assert bytes(dst) == bytes(_block(b"a") + _block(b"b")) + await layer.release_lease("leased") + await asyncio.sleep(0.05) + assert not overwrite.done() + await layer.release_lease_ranges("leased", load_spans) + assert await asyncio.wait_for(overwrite, timeout=1.0) == ALIGNMENT + finally: + await layer.release_lease("all-l2") + await layer.release_lease("mixed") + await layer.release_lease("leased") + await layer.release_lease("all-l1") + await layer.drain() + layer.close() + + _run(scenario()) + + def test_iouring_rejects_l2_overflow(tmp_path) -> None: """Writes beyond the configured L2 capacity are rejected.""" layer = TieredIOUringTransferLayer( diff --git a/tests/unit/test_benchmark_unified_utils.py b/tests/unit/test_benchmark_unified_utils.py index 564c955..2d7a189 100644 --- a/tests/unit/test_benchmark_unified_utils.py +++ b/tests/unit/test_benchmark_unified_utils.py @@ -47,6 +47,7 @@ PhaseResult, RequestResult, _metric_hit_ratios, + _wait_daser_drained, _wait_lmcache_quiescent, lmcache_metrics_url, run_daser_chunk, @@ -1254,7 +1255,7 @@ async def fake_run_vllm_phase_requests(*_args, **_kwargs): timeout=1.0, ) - assert calls == ["phase", "drain", "phase"] + assert calls == ["phase", "drain", "phase", "drain"] async def test_lmcache_waits_for_quiescence_without_extra_sleep(monkeypatch) -> None: @@ -2552,6 +2553,87 @@ def test_daser_drain_failure_aborts_benchmark(monkeypatch) -> None: ) +@pytest.mark.asyncio +async def test_daser_drain_uses_only_the_public_barrier(monkeypatch) -> None: + """The async drain calls DaseR without issuing benchmark-side requests.""" + calls: list[tuple[str, dict[str, Any] | None]] = [] + + class Response: + def raise_for_status(self) -> None: + return None + + class Client: + def __init__(self, **kwargs: Any) -> None: + del kwargs + + async def __aenter__(self) -> "Client": + return self + + async def __aexit__(self, *args: Any) -> None: + return None + + async def post(self, url: str, json: dict[str, Any] | None = None) -> Response: + calls.append((url, json)) + return Response() + + import benchmarks.utils.loadgen as loadgen + + monkeypatch.setattr(loadgen.httpx, "AsyncClient", Client) + manifest = SimpleNamespace( + endpoints={ + "vllm": SimpleNamespace(url="http://vllm"), + "daser": SimpleNamespace(url="http://daser"), + } + ) + + await _wait_daser_drained(manifest) + + drain_calls = [call for call in calls if call[0] == "http://daser/drain"] + vllm_calls = [call for call in calls if call[0].startswith("http://vllm")] + assert not vllm_calls + assert len(drain_calls) == 1 + + +@pytest.mark.asyncio +async def test_daser_drain_propagates_barrier_failure(monkeypatch) -> None: + """A failed DaseR drain remains a benchmark failure.""" + + class Response: + def __init__(self, url: str) -> None: + self.url = url + + def raise_for_status(self) -> None: + if self.url == "http://daser/drain": + raise RuntimeError("drain failed") + + class Client: + def __init__(self, **kwargs: Any) -> None: + del kwargs + + async def __aenter__(self) -> "Client": + return self + + async def __aexit__(self, *args: Any) -> None: + return None + + async def post(self, url: str, json: dict[str, Any] | None = None) -> Response: + del json + return Response(url) + + import benchmarks.utils.loadgen as loadgen + + monkeypatch.setattr(loadgen.httpx, "AsyncClient", Client) + manifest = SimpleNamespace( + endpoints={ + "vllm": SimpleNamespace(url="http://vllm"), + "daser": SimpleNamespace(url="http://daser"), + } + ) + + with pytest.raises(RuntimeError, match="drain failed"): + await _wait_daser_drained(manifest) + + def test_run_bench_vllm_bench_entrypoint_runs_openai_rows( tmp_path: Path, monkeypatch, @@ -2809,6 +2891,10 @@ def fake_run_command(command: list[str]) -> None: "benchmarks.run_bench._probe_daser_metrics", lambda *_args, **_kwargs: None, ) + monkeypatch.setattr( + "benchmarks.utils.vllm_bench._drain_daser", + lambda _manifest, **_kwargs: None, + ) async def fake_collect_phase_metrics( manifest: BenchmarkManifest,