diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 2b0fb414a..3f6be17fb 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -160,6 +160,9 @@ class CommonKVManager(BaseKVManager): self.is_hybrid_mla_backend = getattr(args, "is_hybrid_mla_backend", False) self.disaggregation_mode = disaggregation_mode self.server_args = server_args + self.enable_deferred_decode_kv_release = ( + envs.SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE.get() + ) # for p/d multi node infer self.bootstrap_host = server_args.host self.bootstrap_port = server_args.disaggregation_bootstrap_port @@ -251,6 +254,10 @@ class CommonKVManager(BaseKVManager): self.session_pool_lock = threading.Lock() self.addr_to_rooms_tracker: Dict[str, Set[int]] = defaultdict(set) self.prefill_response_tracker: Dict[int, Set[int]] = defaultdict(set) + # Deferred KV release: room -> prefill ranks that acked their transfer + # drained. Entry exists only while the room is held, so a stale/late + # ack for a reused bootstrap_room is dropped. + self._deferred_abort_ack_tracker: Dict[int, Set[int]] = {} # Heartbeat interval should be at least 2 seconds self.heartbeat_interval = max( envs.SGLANG_DISAGGREGATION_HEARTBEAT_INTERVAL.get(), 2.0 @@ -335,6 +342,28 @@ class CommonKVManager(BaseKVManager): with self.failure_lock: self.failure_records[bootstrap_room] = failure_reason + def register_deferred_abort_room(self, bootstrap_room: int) -> None: + """Arm drain-ack accounting for a held room; a fresh set wipes stale acks + from a prior request that reused this bootstrap_room.""" + self._deferred_abort_ack_tracker[bootstrap_room] = set() + + def note_abort_ack(self, bootstrap_room: int, prefill_rank: int) -> None: + """Record a prefill rank's drain ack (decode receiver thread). Only counts + while the room is held; grabs the set by reference to avoid racing clear.""" + acks = self._deferred_abort_ack_tracker.get(bootstrap_room) + if acks is not None: + acks.add(prefill_rank) + + def is_abort_release_safe(self, bootstrap_room: int, required_acks: int) -> bool: + """True once every prefill rank that could still write these pages has acked.""" + return ( + len(self._deferred_abort_ack_tracker.get(bootstrap_room, ())) + >= required_acks + ) + + def clear_deferred_abort_state(self, bootstrap_room: int) -> None: + self._deferred_abort_ack_tracker.pop(bootstrap_room, None) + def get_kv_replica_factor(self) -> int: if self._kv_replica_factor is None: logger.warning_once( @@ -1237,6 +1266,10 @@ class CommonKVSender(BaseKVSender): self.kv_mgr.req_to_decode_prefix_len.pop(self.bootstrap_room, None) if hasattr(self.kv_mgr, "transfer_infos"): self.kv_mgr.transfer_infos.pop(self.bootstrap_room, None) + if hasattr(self.kv_mgr, "_deferred_ack_targets"): + # Drop a held ack target if the room concluded without draining + # (e.g. aborted before any chunk enqueued); else it leaks on prefill. + self.kv_mgr._deferred_ack_targets.pop(self.bootstrap_room, None) def abort(self): self.kv_mgr.record_failure( diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 8bb4f8b2a..af1f9186d 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -1808,6 +1808,15 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin): self.spec_algorithm = scheduler.spec_algorithm self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get() self.staging_handler = None + self.enable_deferred_kv_release = ( + envs.SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE.get() + ) + self.deferred_kv_release_timeout = ( + envs.SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE_TIMEOUT.get() + ) + # Aborted-mid-transfer requests whose KV pages/slot are held until drained + # or timed out. Entries: (decode_req, deadline, metadata_idx, required_acks). + self._deferred_releases: List[Tuple[DecodeRequest, float, int, int]] = [] def add(self, decode_req: DecodeRequest) -> None: self.queue.append(decode_req) @@ -2027,6 +2036,9 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin): transferred_reqs = [] indices_to_remove = set() + # Queue-removed but held for deferred release; excluded from the metadata + # teardown below. + deferred_indices = set() for i, (decode_req, poll) in enumerate(zip(self.queue, polls)): if rids_to_check is not None and decode_req.req.rid not in rids_to_check: continue @@ -2064,11 +2076,23 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin): ) if self.scheduler.enable_hisparse: self.scheduler.hisparse_coordinator.request_finished(decode_req.req) - # release pre-allocated kv cache, but don't insert into the tree since it's failed - release_kv_cache(decode_req.req, self.tree_cache, is_insert=False) - decode_req.kv_receiver.clear() - decode_req.kv_receiver = None - indices_to_remove.add(i) + if ( + self.enable_deferred_kv_release + and decode_req.kv_receiver.abort_notified + ): + # Decode-initiated abort: a prefill write may still target + # these pages, so hold them until the drain ack or timeout. + # (A prefill-initiated failure has already stopped writing -> + # immediate release below.) + self._defer_release(decode_req) + deferred_indices.add(i) + indices_to_remove.add(i) + else: + # release pre-allocated kv cache, but don't insert into the tree since it's failed + release_kv_cache(decode_req.req, self.tree_cache, is_insert=False) + decode_req.kv_receiver.clear() + decode_req.kv_receiver = None + indices_to_remove.add(i) if self.scheduler.metrics_reporter.enable_metrics: self.scheduler.metrics_collector.increment_transfer_failed_reqs() continue @@ -2106,6 +2130,9 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin): raise ValueError(f"Unexpected poll case: {poll}") for i in indices_to_remove: + if i in deferred_indices: + # Held for deferred release; metadata buffer freed at resolve time. + continue if self.enable_staging and self.staging_handler.is_staging_room( self.queue[i].req.bootstrap_room ): @@ -2125,9 +2152,67 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin): return transferred_reqs + def _defer_release(self, decode_req: DecodeRequest) -> None: + deadline = time.monotonic() + self.deferred_kv_release_timeout + # Require an ack from every notified prefill rank (dummy-proof). Snapshot + # now -- the receiver may be cleared by resolve time. + required_acks = len(decode_req.kv_receiver.bootstrap_infos) + self._deferred_releases.append( + (decode_req, deadline, decode_req.metadata_buffer_index, required_acks) + ) + + def _do_release(self, decode_req: DecodeRequest, idx: int) -> None: + room = decode_req.req.bootstrap_room + if self.enable_staging and self.staging_handler.is_staging_room(room): + self.staging_handler.unregister_decode_req(room) + # release pre-allocated kv cache, but don't insert into the tree since it's failed + release_kv_cache(decode_req.req, self.tree_cache, is_insert=False) + self.metadata_buffers.bootstrap_room[idx] = 0 + self.req_to_metadata_buffer_idx_allocator.free(idx) + decode_req.kv_receiver.kv_mgr.clear_deferred_abort_state(room) + decode_req.kv_receiver.clear() + decode_req.kv_receiver = None + + def has_pending_deferred_releases(self) -> bool: + return bool(self._deferred_releases) + + def resolve_deferred_releases(self) -> None: + """Release held requests once every prefill rank acks the drain, or the + hold times out.""" + if not self._deferred_releases: + return + now = time.monotonic() + still_held = [] + to_release = [] + for decode_req, deadline, idx, required_acks in self._deferred_releases: + room = decode_req.req.bootstrap_room + kv_mgr = decode_req.kv_receiver.kv_mgr + drained = kv_mgr.is_abort_release_safe(room, required_acks) + if not drained and now < deadline: + still_held.append((decode_req, deadline, idx, required_acks)) + else: + to_release.append((decode_req, idx, room, drained)) + # Commit the survivors before releasing so a _do_release exception can't + # leave a released entry in the list (double-free / None receiver on retry). + self._deferred_releases = still_held + for decode_req, idx, room, drained in to_release: + if not drained: + logger.warning( + f"Deferred KV release for room {room} timed out after " + f"{self.deferred_kv_release_timeout}s without a full drain " + f"ack from prefill; releasing anyway." + ) + try: + self._do_release(decode_req, idx) + except Exception: + # Isolate a failed release so the rest still run; entry already dropped. + logger.exception(f"Deferred KV release failed for room {room}") + def release_memory_occupation(self): """Clean up in-flight transfers before releasing GPU memory.""" self.queue.clear() + # Pool is being torn down; drop held entries without per-request release. + self._deferred_releases.clear() def resume_memory_occupation(self): """Queues are already cleared on release; new transfers can be accepted.""" @@ -2346,6 +2431,10 @@ class SchedulerDisaggregationDecodeMixin: if get_disagg().disaggregation_decode_enable_offload_kvcache: self.decode_offload_manager.check_offload_progress() + # Resolve held releases every iteration (before the retraction/polling + # gates below) so their timeouts fire under memory pressure. + self.disagg_decode_transfer_queue.resolve_deferred_releases() + # try to resume retracted requests if there are enough space for another `num_reserved_decode_tokens` decode steps resumed_reqs = self.disagg_decode_prealloc_queue.resume_retracted_reqs() self.waiting_queue.extend(resumed_reqs) diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 66ac63b34..3fbc37316 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -8,7 +8,7 @@ import struct import threading import time from collections import defaultdict -from typing import List, Optional, Set, Tuple, Union +from typing import Dict, List, Optional, Set, Tuple, Union import numpy as np import numpy.typing as npt @@ -214,6 +214,10 @@ class MooncakeKVManager(CommonKVManager): # Per-room count of chunks not yet transferred; teardown waits for # zero so a deferred chunk is not dropped by an early conclude. self._staging_outstanding = defaultdict(int) + # Deferred KV release: aborted room -> (decode_ip, decode_port), ack + # held until the transfer drains. Written by the bootstrap thread, + # popped by the single transfer worker that owns the room. + self._deferred_ack_targets: Dict[int, Tuple[str, int]] = {} self.session_lock = threading.Lock() # Determine the number of threads to use for kv sender cpu_count = os.cpu_count() @@ -795,13 +799,7 @@ class MooncakeKVManager(CommonKVManager): ) for (src_ptr, dst_ptr, item_len) in layers_params ] - for future in concurrent.futures.as_completed(futures): - status = future.result() - if status != 0: - for f in futures: - f.cancel() - return status - return 0 + return self._await_transfer_futures(futures) else: # Combining all layers' params in one batch transfer is more efficient # compared to using multiple threads @@ -853,6 +851,25 @@ class MooncakeKVManager(CommonKVManager): "must enable it and use the same page size and model spec." ) + def _await_transfer_futures(self, futures) -> int: + """Await a chunk's per-layer RDMA writes; return the first non-zero status. + cancel() is a no-op for a running future, so with deferred release on we + still drain the running ones before returning (no write may outlive this + call, which the drain-ack relies on). Off: original early-return.""" + ret = 0 + for future in concurrent.futures.as_completed(futures): + try: + status = future.result() + except concurrent.futures.CancelledError: + continue + if status != 0 and ret == 0: + ret = status + for f in futures: + f.cancel() + if not self.enable_deferred_decode_kv_release: + return ret + return ret + def send_kvcache( self, mooncake_session_id: str, @@ -978,13 +995,7 @@ class MooncakeKVManager(CommonKVManager): executor.submit(process_layer, src_ptr, dst_ptr, token_item_len) for src_ptr, dst_ptr, token_item_len in layers_params ] - for future in concurrent.futures.as_completed(futures): - status = future.result() - if status != 0: - for pending in futures: - pending.cancel() - return status - return 0 + return self._await_transfer_futures(futures) transfer_blocks = [] for src_ptr, dst_ptr, token_item_len in layers_params: @@ -1112,14 +1123,7 @@ class MooncakeKVManager(CommonKVManager): executor.submit(process_layer_tp_aware, src_v_ptrs[i], dst_v_ptrs[i]) ) - for future in concurrent.futures.as_completed(futures): - status = future.result() - if status != 0: - for f in futures: - f.cancel() - return status - - return 0 + return self._await_transfer_futures(futures) def send_aux( self, @@ -1612,6 +1616,39 @@ class MooncakeKVManager(CommonKVManager): is_ipv6=na.is_ipv6, ) + def _prefill_unique_rank(self) -> int: + """Stable per-sender id, matching what the transfer worker syncs on Success.""" + return ( + self.attn_tp_rank * (self.pp_size * self.attn_cp_size) + + self.pp_rank * self.attn_cp_size + + self.attn_cp_rank + ) + + def _send_abort_ack(self, decode_ip: str, decode_port: int, room: int) -> None: + """Best-effort ack that this rank's transfer for an aborted room drained.""" + try: + na = NetworkAddress(decode_ip, decode_port) + self._send_multipart_locked( + na.to_tcp(), + [ + b"ABORT_ACK", + str(room).encode("ascii"), + str(self._prefill_unique_rank()).encode("ascii"), + ], + is_ipv6=na.is_ipv6, + ) + except Exception as e: + logger.debug(f"Failed to send drained ABORT_ACK for room {room}: {e}") + + def _maybe_ack_drained_abort(self, room: int) -> None: + """Send the deferred ack once an aborted room's chunks have drained + (outstanding == 0). pop() makes it fire at most once.""" + if self._staging_outstanding.get(room, 0) > 0: + return + target = self._deferred_ack_targets.pop(room, None) + if target is not None: + self._send_abort_ack(target[0], target[1], room) + def transfer_worker( self, queue: FastQueue, @@ -1651,6 +1688,9 @@ class MooncakeKVManager(CommonKVManager): thread_finish_flag=True, ) self._staging_outstanding.pop(kv_chunk.room, None) + if self.enable_deferred_decode_kv_release: + # Skipped => nothing written for this aborted room; ack. + self._maybe_ack_drained_abort(kv_chunk.room) continue # Count each chunk once; the flag survives re-enqueue on defer. @@ -1925,6 +1965,10 @@ class MooncakeKVManager(CommonKVManager): continue self._staging_outstanding[kv_chunk.room] -= 1 + if self.enable_deferred_decode_kv_release: + # In-flight write finished; if aborted and nothing outstanding, + # the pages are idle -> release the held ack. + self._maybe_ack_drained_abort(kv_chunk.room) # Tear down only when no chunk is still outstanding and the room # has concluded: already cleared, Success, or a Failed *last* # chunk. A non-last Failed chunk keeps the room (more chunks may @@ -1984,11 +2028,36 @@ class MooncakeKVManager(CommonKVManager): room_to_be_aborted = int(waiting_req_bytes[1].decode("ascii")) decode_ip = waiting_req_bytes[2].decode("ascii") decode_port = int(waiting_req_bytes[3].decode("ascii")) - # No need to abort the room if it has already succeeded - if ( + room_active = ( room_to_be_aborted in self.request_status and self.check_status(room_to_be_aborted) != KVPoll.Success - ): + ) + if self.enable_deferred_decode_kv_release: + # Mark Failed FIRST (stops add_transfer_request enqueuing + # new chunks), THEN register the ack target: registering + # first would let the worker drain+ack while the room is + # not yet Failed, so a newly enqueued chunk could still + # write to the freed pages. The worker (not this thread) + # acks once its in-flight write drains; if nothing is in + # flight, decode falls back to the release timeout. + if room_active: + self.update_status(room_to_be_aborted, KVPoll.Failed) + self._deferred_ack_targets[room_to_be_aborted] = ( + decode_ip, + decode_port, + ) + logger.debug( + f"Received abort notification for room {room_to_be_aborted}, " + f"marked as Failed; ACK deferred until transfer drains" + ) + else: + # Already completed/unknown: no in-flight write, ack now. + self._send_abort_ack( + decode_ip, decode_port, room_to_be_aborted + ) + continue + # No need to abort the room if it has already succeeded + if room_active: self.update_status(room_to_be_aborted, KVPoll.Failed) logger.debug( f"Received abort notification for room {room_to_be_aborted}, " @@ -2102,9 +2171,14 @@ class MooncakeKVManager(CommonKVManager): # Prefill acknowledges abort notification if msg[0] == b"ABORT_ACK": - # TODO(shangming): use this info to implement the deferred release mechanism if needed ack_aborted_room = int(msg[1].decode("ascii")) logger.debug(f"Received ABORT_ACK for room {ack_aborted_room}") + # Deferred release: the 3-frame ack carries the prefill rank + # and means its transfer drained; aggregate for is_abort_release_safe. + if self.enable_deferred_decode_kv_release and len(msg) >= 3: + self.note_abort_ack( + ack_aborted_room, int(msg[2].decode("ascii")) + ) continue bootstrap_room, status, prefill_rank = msg diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index fa802f44d..bc0e6004c 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -619,6 +619,11 @@ class Envs: SGLANG_DISAGGREGATION_FORCE_QUERY_PREFILL_DP_RANK = EnvBool(False) SGLANG_DISAGGREGATION_SAMPLING_MASK_MAX_TOKENS = EnvInt(0) SGLANG_DISAGGREGATION_BOOTSTRAP_ENTRY_CLEANUP_INTERVAL = EnvInt(120) + # Deferred decode-side KV release: on abort, hold an in-flight request's KV + # pages/slot until the prefill acks the transfer drained, or the timeout + # below fires. Off by default (no behavior/perf impact when disabled). + SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE = EnvBool(False) + SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE_TIMEOUT = EnvFloat(30.0) # =================================================================== # Distributed and model-parallel runtime diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 7dc5be833..ecf79dc20 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -4054,7 +4054,14 @@ class Scheduler( # memory leak check (skipped for hisparse — pool counters intentionally # diverge during host-backup, see _get_swa_token_info clamp). - if not self.enable_hisparse: + # Also skipped while deferred KV releases are pending: they hold pages out + # of the allocator by design, so the pool is transiently below `total` and + # would trip the idle leak invariant. Resumes once the holds resolve. + deferred_pending = ( + self.disaggregation_mode == DisaggregationMode.DECODE + and self.disagg_decode_transfer_queue.has_pending_deferred_releases() + ) + if not self.enable_hisparse and not deferred_pending: has_leak, messages = self.invariant_checker._check_all_pools( self.pool_stats_observer.get_pool_stats(), ) @@ -4542,7 +4549,21 @@ class Scheduler( for decode_req in self.disagg_decode_transfer_queue.queue: if recv_req.abort_all or decode_req.req.rid.startswith(recv_req.rid): logger.debug(f"Abort transfer queue request. {decode_req.req.rid=}") - decode_req.kv_receiver.abort() + receiver = decode_req.kv_receiver + receiver.abort() + # Arm drain-ack accounting once the ABORT is sent, so acks + # arriving before this req is deferred (e.g. during the next + # forward step) are captured. A fresh set also drops stale acks + # from a prior request that reused this bootstrap_room. A + # redundant abort only re-wipes -- holds longer, never releases + # early -- so no transition guard is needed. + if ( + receiver.kv_mgr.enable_deferred_decode_kv_release + and receiver.abort_notified + ): + receiver.kv_mgr.register_deferred_abort_room( + decode_req.req.bootstrap_room + ) # Abort requests whose KV is already backed up for retraction. if self.disagg_decode_prealloc_queue.retracted_queue: diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index 42ead03e0..7fa52ae77 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -1468,6 +1468,9 @@ class SchedulerPPMixin: def process_decode_transfer_queue( self: Scheduler, release_rids: Optional[List[str]] ): + # Resolve held deferred releases every call, independent of release_rids, + # so ack/timeout-driven releases still fire when no rids are being polled. + self.disagg_decode_transfer_queue.resolve_deferred_releases() if release_rids is not None: released_reqs = self.disagg_decode_transfer_queue.pop_transferred( release_rids diff --git a/test/registered/unit/disaggregation/test_decode_queue_cleanup.py b/test/registered/unit/disaggregation/test_decode_queue_cleanup.py index 6cddeb03b..a0d07f686 100644 --- a/test/registered/unit/disaggregation/test_decode_queue_cleanup.py +++ b/test/registered/unit/disaggregation/test_decode_queue_cleanup.py @@ -211,6 +211,7 @@ class TestDecodeQueueCleanup(CustomTestCase): queue = DecodeTransferQueue.__new__(DecodeTransferQueue) queue.queue = [decode_req] queue.enable_staging = False + queue.enable_deferred_kv_release = False queue.gloo_group = MagicMock() queue.req_to_metadata_buffer_idx_allocator = MagicMock() queue.tp_rank = 0 diff --git a/test/registered/unit/disaggregation/test_deferred_decode_kv_release.py b/test/registered/unit/disaggregation/test_deferred_decode_kv_release.py new file mode 100644 index 000000000..3f33c3d07 --- /dev/null +++ b/test/registered/unit/disaggregation/test_deferred_decode_kv_release.py @@ -0,0 +1,247 @@ +"""Unit tests for the deferred decode-side KV release mechanism. + +When a decode request is aborted while its prefill->decode KV transfer may still +be in flight, the decode side holds its KV pages / req-slot instead of freeing +them immediately (which could let the still-in-flight write land on pages already +reused by another request). The pages are released once every prefill rank acks +that its transfer drained (CommonKVManager.is_abort_release_safe), or a timeout +fires. See DecodeTransferQueue.resolve_deferred_releases. +""" + +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +from sglang.srt.disaggregation import decode as decode_mod +from sglang.srt.disaggregation.common.conn import CommonKVManager +from sglang.srt.disaggregation.decode import DecodeTransferQueue +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +def _make_manager(): + """A bare CommonKVManager carrying only the deferred-ack state the helpers + touch (avoids the heavy real __init__).""" + mgr = CommonKVManager.__new__(CommonKVManager) + mgr._deferred_abort_ack_tracker = {} + return mgr + + +class TestAbortAckAggregation(CustomTestCase): + def test_release_safe_only_after_all_required_ranks_ack(self): + mgr = _make_manager() + room = 100 + mgr.register_deferred_abort_room(room) + self.assertFalse(mgr.is_abort_release_safe(room, required_acks=2)) + + mgr.note_abort_ack(room, 0) + self.assertFalse(mgr.is_abort_release_safe(room, required_acks=2)) + + mgr.note_abort_ack(room, 1) + self.assertTrue(mgr.is_abort_release_safe(room, required_acks=2)) + + def test_duplicate_rank_ack_does_not_over_count(self): + mgr = _make_manager() + room = 101 + mgr.register_deferred_abort_room(room) + mgr.note_abort_ack(room, 0) + mgr.note_abort_ack(room, 0) # same rank twice + # Two acks arrived but from one rank: not safe for a 2-rank prefill. + self.assertFalse(mgr.is_abort_release_safe(room, required_acks=2)) + + def test_single_rank_fast_path(self): + mgr = _make_manager() + room = 102 + mgr.register_deferred_abort_room(room) + mgr.note_abort_ack(room, 0) + self.assertTrue(mgr.is_abort_release_safe(room, required_acks=1)) + + def test_clear_deferred_abort_state(self): + mgr = _make_manager() + room = 103 + mgr.register_deferred_abort_room(room) + mgr.note_abort_ack(room, 0) + mgr.clear_deferred_abort_state(room) + self.assertNotIn(room, mgr._deferred_abort_ack_tracker) + self.assertFalse(mgr.is_abort_release_safe(room, required_acks=1)) + + def test_ack_before_register_is_dropped(self): + # An ack for a room that isn't actively held must not be recorded (it + # would otherwise pollute a later request reusing the same room). + mgr = _make_manager() + room = 104 + mgr.note_abort_ack(room, 0) # no register yet + self.assertNotIn(room, mgr._deferred_abort_ack_tracker) + self.assertFalse(mgr.is_abort_release_safe(room, required_acks=1)) + + def test_late_ack_after_release_does_not_pollute_reused_room(self): + # Regression for bootstrap_room reuse: req A (room R) releases, then a + # late ack from A arrives, then req B reuses room R. B must start from a + # clean slate and not inherit A's ack (which would release B early while + # its transfer is still in flight -> KV corruption). + mgr = _make_manager() + room = 105 + + # Req A: held, one of two ranks acks, then released (e.g. timed out). + mgr.register_deferred_abort_room(room) + mgr.note_abort_ack(room, 0) + mgr.clear_deferred_abort_state(room) + + # Late ack from A's other rank arrives after release -> dropped. + mgr.note_abort_ack(room, 1) + self.assertNotIn(room, mgr._deferred_abort_ack_tracker) + + # Req B reuses room R. + mgr.register_deferred_abort_room(room) + # Only B's rank-0 has acked so far; a 2-rank prefill is NOT safe yet. + mgr.note_abort_ack(room, 0) + self.assertFalse(mgr.is_abort_release_safe(room, required_acks=2)) + mgr.note_abort_ack(room, 1) + self.assertTrue(mgr.is_abort_release_safe(room, required_acks=2)) + + def test_register_resets_stale_acks(self): + mgr = _make_manager() + room = 106 + mgr.register_deferred_abort_room(room) + mgr.note_abort_ack(room, 0) + mgr.note_abort_ack(room, 1) + self.assertTrue(mgr.is_abort_release_safe(room, required_acks=2)) + # Re-registering (a later reuse) wipes the prior acks. + mgr.register_deferred_abort_room(room) + self.assertFalse(mgr.is_abort_release_safe(room, required_acks=2)) + + +class _FakeIdxAllocator: + def __init__(self): + self.freed = [] + + def free(self, idx): + self.freed.append(idx) + + +def _make_queue(timeout=30.0): + q = DecodeTransferQueue.__new__(DecodeTransferQueue) + q._deferred_releases = [] + q.deferred_kv_release_timeout = timeout + q.enable_staging = False + q.staging_handler = None + q.tree_cache = object() + q.metadata_buffers = SimpleNamespace(bootstrap_room={}) + q.req_to_metadata_buffer_idx_allocator = _FakeIdxAllocator() + return q + + +def _make_decode_req(room, idx, mgr, n_prefill_ranks=1): + receiver = SimpleNamespace( + kv_mgr=mgr, + # One entry per prefill rank the decode notified of the abort; its length + # is the required drain-ack count (see DecodeTransferQueue._defer_release). + bootstrap_infos=[{"rank": r} for r in range(n_prefill_ranks)], + clear=lambda: None, + ) + return SimpleNamespace( + req=SimpleNamespace(bootstrap_room=room), + kv_receiver=receiver, + metadata_buffer_index=idx, + ) + + +class TestResolveDeferredReleases(CustomTestCase): + def test_noop_when_nothing_deferred(self): + q = _make_queue() + with patch.object(decode_mod, "release_kv_cache") as rel: + q.resolve_deferred_releases() + rel.assert_not_called() + + def test_holds_until_drained_then_releases(self): + mgr = _make_manager() + room, idx = 200, 7 + q = _make_queue() + dreq = _make_decode_req(room, idx, mgr, n_prefill_ranks=2) + # In production the room is armed in abort_request when the ABORT is + # sent, before the scheduler defers here. + mgr.register_deferred_abort_room(room) + q._defer_release(dreq) + + with patch.object(decode_mod, "release_kv_cache") as rel: + # Not yet acked -> held, not released. + q.resolve_deferred_releases() + rel.assert_not_called() + self.assertEqual(len(q._deferred_releases), 1) + + # One of two ranks acked -> still held. + mgr.note_abort_ack(room, 0) + q.resolve_deferred_releases() + rel.assert_not_called() + self.assertEqual(len(q._deferred_releases), 1) + + # Both ranks acked -> released exactly once. + mgr.note_abort_ack(room, 1) + q.resolve_deferred_releases() + rel.assert_called_once_with(dreq.req, q.tree_cache, is_insert=False) + + # Held state fully cleaned up. + self.assertEqual(q._deferred_releases, []) + self.assertEqual(q.req_to_metadata_buffer_idx_allocator.freed, [idx]) + self.assertEqual(q.metadata_buffers.bootstrap_room[idx], 0) + self.assertNotIn(room, mgr._deferred_abort_ack_tracker) + self.assertIsNone(dreq.kv_receiver) + + def test_releases_on_timeout_without_ack(self): + mgr = _make_manager() + room, idx = 300, 3 + q = _make_queue(timeout=30.0) + dreq = _make_decode_req(room, idx, mgr, n_prefill_ranks=1) + # Force an already-expired deadline (no ack will ever arrive). + q._deferred_releases.append((dreq, float("-inf"), idx, 1)) + + with patch.object(decode_mod, "release_kv_cache") as rel: + q.resolve_deferred_releases() + rel.assert_called_once_with(dreq.req, q.tree_cache, is_insert=False) + + self.assertEqual(q._deferred_releases, []) + self.assertEqual(q.req_to_metadata_buffer_idx_allocator.freed, [idx]) + self.assertIsNone(dreq.kv_receiver) + + def test_failed_release_is_isolated_and_not_retried(self): + # A raising _do_release must drop the entry (no double-free on retry) and + # not brick resolve for the remaining entries or subsequent calls. + mgr = _make_manager() + q = _make_queue() + good = _make_decode_req(700, 1, mgr) + bad = _make_decode_req(701, 2, mgr) + # Both already past deadline -> both selected for release. + q._deferred_releases.append((bad, float("-inf"), 2, 1)) + q._deferred_releases.append((good, float("-inf"), 1, 1)) + + calls = [] + + def fake_release(req, tree_cache, is_insert): + calls.append(req) + if req is bad.req: + raise RuntimeError("boom") + + with patch.object(decode_mod, "release_kv_cache", side_effect=fake_release): + q.resolve_deferred_releases() # must not raise + # The good one still released despite the bad one throwing. + self.assertIn(good.req, calls) + # Nothing left held, and a second call is a clean no-op (no retry). + self.assertEqual(q._deferred_releases, []) + q.resolve_deferred_releases() + + def test_defer_release_records_deadline_and_idx(self): + mgr = _make_manager() + q = _make_queue(timeout=12.5) + dreq = _make_decode_req(room=400, idx=9, mgr=mgr) + q._defer_release(dreq) + self.assertEqual(len(q._deferred_releases), 1) + held_req, deadline, held_idx, required = q._deferred_releases[0] + self.assertIs(held_req, dreq) + self.assertEqual(held_idx, 9) + self.assertIsInstance(deadline, float) + + +if __name__ == "__main__": + unittest.main()