diff --git a/tensorrt_llm/_torch/disaggregation/native/bounce/core.py b/tensorrt_llm/_torch/disaggregation/native/bounce/core.py index 0eb621d5d189..c193de0ce996 100644 --- a/tensorrt_llm/_torch/disaggregation/native/bounce/core.py +++ b/tensorrt_llm/_torch/disaggregation/native/bounce/core.py @@ -75,6 +75,8 @@ class TransferContext: _writer_ok: Dict[int, bool] = field(default_factory=dict) # per successful writer: where it wrote, plus the fragments to scatter back _scatter_descs: List[tuple] = field(default_factory=list) + _writer_cohort: Optional[frozenset[int]] = None + _publication_failed: bool = False _orphaned: bool = False scatter_state: ScatterState = ScatterState.IDLE state: TransferState = TransferState.INIT @@ -93,10 +95,16 @@ def _writers_final(self) -> bool: ) def _all_writers_reported(self) -> bool: + if self._writer_cohort is not None: + return self._writer_cohort.issubset(self._writer_ok) return len(self._writer_ok) >= self.num_writers def _all_writers_succeeded(self) -> bool: - return self._all_writers_reported() and all(self._writer_ok.values()) + if not self._all_writers_reported(): + return False + if self._writer_cohort is not None: + return all(self._writer_ok[rank] for rank in self._writer_cohort) + return all(self._writer_ok.values()) # Mutations: call only while holding the transport's reservation lock. def record_writer_result( @@ -115,6 +123,8 @@ def record_writer_result( # notification, or a stray failure that would flip a good transfer to failed; drop it. if self._writers_final() or peer_rank in self._writer_ok: return + if self._writer_cohort is not None and peer_rank not in self._writer_cohort: + raise RuntimeError(f"writer {peer_rank} is outside the published bounce cohort") self._writer_ok[peer_rank] = succeeded if succeeded and dst_ptrs is not None and int(dst_ptrs.size) > 0: self._scatter_descs.append( @@ -130,6 +140,25 @@ def mark_orphaned(self) -> None: return self._orphaned = True + def abort_publication(self, published_writers: set[int]) -> None: + """Close a failed fan-out around only the writers that were published. + + Published writers still have to report terminal evidence. Successful + partial data is intentionally not scattered because the request is + already incomplete. + """ + if self._writers_final(): + return + published = frozenset(published_writers) + if len(published) > self.num_writers or not self._writer_ok.keys() <= published: + raise RuntimeError( + "published bounce cohort is inconsistent with terminal evidence: " + f"reported={sorted(self._writer_ok)} published={sorted(published)} " + f"reserved_count={self.num_writers}" + ) + self._writer_cohort = published + self._publication_failed = True + def begin_scatter(self) -> None: self.state = TransferState.SCATTERING self.scatter_state = ScatterState.QUEUED @@ -146,6 +175,7 @@ def ready_to_scatter(self) -> bool: return ( self.state is TransferState.ACTIVE and not self._orphaned + and not self._publication_failed and self._all_writers_succeeded() and bool(self._scatter_descs) ) @@ -157,6 +187,8 @@ def ready_to_settle(self) -> bool: return True # in doubt: settle now and quarantine if not self._all_writers_reported(): return False # a writer has not reported yet + if self._publication_failed: + return True # every published writer drained; incomplete data is discarded if self._all_writers_succeeded() and self._scatter_descs: return self.scatter_state in (ScatterState.DONE, ScatterState.FAILED) return True # nothing to scatter, or a failure among them: either way drained @@ -170,7 +202,11 @@ def settle(self) -> Optional[Settlement]: if self._orphaned: self.state = TransferState.QUARANTINED return Settlement(self.slot_id, Disposition.QUARANTINE, False, self.on_done) - success = self._all_writers_succeeded() and self.scatter_state is not ScatterState.FAILED + success = ( + not self._publication_failed + and self._all_writers_succeeded() + and self.scatter_state is not ScatterState.FAILED + ) self.state = TransferState.COMPLETED if success else TransferState.FAILED return Settlement(self.slot_id, Disposition.RELEASE, success, self.on_done) @@ -220,6 +256,10 @@ def release_idle_reservation(self, rid_slice) -> None: def orphan_reservation(self, rid_slice) -> None: """Give up on an in-flight reservation (cancel/timeout/lost result); quarantine, don't leak.""" + @abstractmethod + def abort_publication(self, rid_slice, published_writers: set[int]) -> None: + """Limit a failed fan-out to writers whose REQUEST_DATA was queued.""" + @abstractmethod def record_result( self, rid_slice, peer_rank, dst_ptrs=None, sizes=None, src_base=None, on_done=None diff --git a/tensorrt_llm/_torch/disaggregation/native/bounce/impl.py b/tensorrt_llm/_torch/disaggregation/native/bounce/impl.py index 1718b60eb522..22de32f9384a 100644 --- a/tensorrt_llm/_torch/disaggregation/native/bounce/impl.py +++ b/tensorrt_llm/_torch/disaggregation/native/bounce/impl.py @@ -379,6 +379,13 @@ def orphan_reservation(self, rid_slice: RidSlice) -> None: or leaking it. Idempotent; a no-op once the transfer has settled.""" self._apply(rid_slice, lambda ctx: ctx.mark_orphaned()) + def abort_publication(self, rid_slice: RidSlice, published_writers: set[int]) -> None: + """Retain a failed fan-out until every successfully published writer drains.""" + self._apply( + rid_slice, + lambda ctx: ctx.abort_publication(published_writers), + ) + def _apply(self, rid_slice: RidSlice, mutate: Callable[[TransferContext], None]) -> None: """Mutate the state under the lock, then do what it asks (scatter or settle) with the lock released, never holding it across a CUDA sync, a queue put, or a callback. No-op if the @@ -527,6 +534,9 @@ def release_idle_reservation(self, rid_slice) -> None: def orphan_reservation(self, rid_slice) -> None: pass + def abort_publication(self, rid_slice, published_writers: set[int]) -> None: + pass + def record_result( self, rid_slice, peer_rank, dst_ptrs=None, sizes=None, src_base=None, on_done=None ): diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index 4d9d9495d35d..253f54ca931c 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -23,7 +23,7 @@ import weakref from dataclasses import dataclass from enum import Enum -from typing import TYPE_CHECKING, List, Optional, Union +from typing import TYPE_CHECKING, Callable, List, Optional, Union import msgpack import numpy as np @@ -226,6 +226,131 @@ class TaskStatus(Enum): ERROR = "ERROR" +class _ReceiveOperationOwner: + """Track destination access independently from a task's logical result.""" + + def __init__(self) -> None: + self._lock = threading.Lock() + self._publication_pending = False + self._cancelled_unpublished = False + self._expected_writers: Optional[int] = None + self._writer_cohort: Optional[frozenset[int]] = None + self._writer_results: dict[int, bool] = {} + self._publication_failed = False + self._local_completion_pending = False + self._invalid_evidence = False + + def begin_publication(self) -> None: + with self._lock: + if self._publication_pending or self._expected_writers is not None: + raise RuntimeError("destination publication was already started") + self._publication_pending = True + + def seal_writer_cohort( + self, + expected_writers: int, + writer_cohort: Optional[set[int]] = None, + ) -> None: + if expected_writers < 0: + raise ValueError(f"expected_writers must be non-negative, got {expected_writers}") + cohort = None if writer_cohort is None else frozenset(writer_cohort) + if cohort is not None and len(cohort) != expected_writers: + raise ValueError( + f"writer cohort has {len(cohort)} member(s), expected {expected_writers}" + ) + with self._lock: + if self._expected_writers is not None: + if self._expected_writers != expected_writers or self._writer_cohort != cohort: + raise RuntimeError("writer cohort was already sealed differently") + return + self._expected_writers = expected_writers + self._writer_cohort = cohort + + def finish_publication(self) -> None: + """Record that every authorized REQUEST_DATA message was sent.""" + with self._lock: + self._publication_pending = False + + def abort_publication(self, published_writers: set[int]) -> None: + """Close a failed fan-out around the writers whose sends succeeded.""" + with self._lock: + published = frozenset(published_writers) + if self._writer_cohort is not None and not published.issubset(self._writer_cohort): + self._invalid_evidence = True + raise RuntimeError("publication recorded a writer outside the sealed cohort") + if not self._writer_results.keys() <= published: + self._invalid_evidence = True + raise RuntimeError("terminal evidence came from an unpublished writer") + self._expected_writers = len(published) + self._writer_cohort = frozenset(published) + self._publication_failed = True + self._publication_pending = False + + def cancel_unpublished(self) -> bool: + """Close a publication that did not authorize a remote writer.""" + with self._lock: + if self._expected_writers is not None: + return self._cancelled_unpublished + self._cancelled_unpublished = True + self._expected_writers = 0 + self._writer_cohort = frozenset() + self._publication_pending = False + return True + + def record_writer_result( + self, + peer_rank: int, + succeeded: bool, + *, + wait_for_local_completion: bool, + ) -> tuple[bool, bool]: + """Record one writer and return ``(accepted, all_succeeded)``.""" + with self._lock: + if self._expected_writers is None: + self._invalid_evidence = True + raise RuntimeError( + f"writer {peer_rank} reported terminal evidence before publication" + ) + if self._writer_cohort is not None and peer_rank not in self._writer_cohort: + self._invalid_evidence = True + raise RuntimeError(f"writer {peer_rank} is outside the sealed cohort") + previous = self._writer_results.get(peer_rank) + if previous is not None: + if previous != succeeded: + self._invalid_evidence = True + raise RuntimeError( + f"writer {peer_rank} reported contradictory terminal evidence" + ) + return False, False + if len(self._writer_results) >= self._expected_writers: + return False, False + self._writer_results[peer_rank] = succeeded + all_reported = len(self._writer_results) == self._expected_writers + all_succeeded = ( + all_reported and not self._publication_failed and all(self._writer_results.values()) + ) + if all_succeeded and wait_for_local_completion: + self._local_completion_pending = True + return True, all_succeeded + + def finish_local_completion(self) -> None: + with self._lock: + self._local_completion_pending = False + + @property + def resources_drained(self) -> bool: + with self._lock: + writers_drained = self._expected_writers is not None and ( + len(self._writer_results) == self._expected_writers + ) + return ( + writers_drained + and not self._publication_pending + and not self._local_completion_pending + and not self._invalid_evidence + ) + + class AgentResult(Enum): SUCCESS = "SUCCESS" FAILED = "FAILED" @@ -1612,15 +1737,30 @@ def __init__( self._exception: Optional[Exception] = None self._aux_slot = aux_slot self._perf_timer = PerfTimer() if perf_log_manager.enabled else None + self._physical_owner: Optional[_ReceiveOperationOwner] = None + self._ownership_state_lock: Optional[threading.Lock] = None def fail(self, exc: Exception) -> None: - self._exception = exc - self.status = TaskStatus.ERROR - self._event.set() + if self._ownership_state_lock is None: + self._exception = exc + self.status = TaskStatus.ERROR + self._event.set() + return + with self._ownership_state_lock: + self._exception = exc + self.status = TaskStatus.ERROR + self._event.set() def complete(self) -> None: - self.status = TaskStatus.TRANSFERRED - self._event.set() + if self._ownership_state_lock is None: + self.status = TaskStatus.TRANSFERRED + self._event.set() + return + with self._ownership_state_lock: + if self.status == TaskStatus.ERROR: + return + self.status = TaskStatus.TRANSFERRED + self._event.set() def wait(self, timeout: Optional[float] = None) -> bool: """Block until terminal state. Returns True if done, False on timeout.""" @@ -1630,6 +1770,49 @@ def wait(self, timeout: Optional[float] = None) -> bool: def is_done(self) -> bool: return self._event.is_set() + def begin_publication(self) -> None: + if self._physical_owner is None: + self._physical_owner = _ReceiveOperationOwner() + self._ownership_state_lock = threading.Lock() + self._physical_owner.begin_publication() + + def _get_physical_owner(self) -> _ReceiveOperationOwner: + if self._physical_owner is None: + raise RuntimeError("receive operation owner was not initialized") + return self._physical_owner + + def seal_writer_cohort(self, writer_cohort: Optional[set[int]] = None) -> None: + self._get_physical_owner().seal_writer_cohort(self.expected_transfers, writer_cohort) + + def finish_publication(self) -> None: + self._get_physical_owner().finish_publication() + + def abort_publication(self, published_writers: set[int]) -> None: + self._get_physical_owner().abort_publication(published_writers) + + def cancel_unpublished(self) -> bool: + return self._get_physical_owner().cancel_unpublished() + + def record_writer_result( + self, + peer_rank: int, + succeeded: bool, + *, + wait_for_local_completion: bool, + ) -> tuple[bool, bool]: + return self._get_physical_owner().record_writer_result( + peer_rank, + succeeded, + wait_for_local_completion=wait_for_local_completion, + ) + + def finish_local_completion(self) -> None: + self._get_physical_owner().finish_local_completion() + + @property + def resources_drained(self) -> bool: + return self._get_physical_owner().resources_drained + def print_perf_info(self, peer_rank: int, instance_name: str, instance_rank: int): if self._perf_timer is None: return @@ -1649,10 +1832,14 @@ def __init__( peer_registrar: PeerRegistrar, agent: BaseTransferAgent, bounce=None, + enforce_physical_ownership: bool = False, ): self._registrar = peer_registrar self._agent = agent self._bounce = bounce + # Internal component gate only. Production wiring is added after the + # complete Phase 1 ownership path is available and qualified. + self._enforce_physical_ownership = enforce_physical_ownership self._dealers = {} self._sender_ep_instance_map = {} # info_endpoint -> diagnostic message for peers that failed the @@ -1693,23 +1880,37 @@ def clear_session(self, unique_rid: int): def setup_session(self, rx_session: RxSessionBase): pre_cancel = False with self._sessions_lock: - self._sessions[rx_session.disagg_request_id] = weakref.ref(rx_session) + if getattr(rx_session, "_enforce_physical_ownership", False): + self._sessions[rx_session.disagg_request_id] = rx_session + else: + self._sessions[rx_session.disagg_request_id] = weakref.ref(rx_session) if rx_session.disagg_request_id in self._pre_cancelled_rids: pre_cancel = True self._pre_cancelled_rids.discard(rx_session.disagg_request_id) if pre_cancel: - rx_session.cancel() + if getattr(rx_session, "_enforce_physical_ownership", False): + rx_session.cancel_local() + else: + rx_session.cancel() def _get_session(self, unique_rid: Optional[int]) -> Optional["RxSession"]: with self._sessions_lock: - session_ref = self._sessions.get(unique_rid) - if session_ref is None: - return None - session = session_ref() - if session is None: - logger.warning(f"RxSession {unique_rid} has been garbage collected") + session_entry = self._sessions.get(unique_rid) + return self._resolve_session_entry(unique_rid, session_entry) + + @staticmethod + def _resolve_session_entry( + unique_rid: Optional[int], + session_entry: Optional[Union["RxSession", weakref.ReferenceType["RxSession"]]], + ) -> Optional["RxSession"]: + if session_entry is None: return None - return session + if isinstance(session_entry, weakref.ReferenceType): + session = session_entry() + if session is None: + logger.warning(f"RxSession {unique_rid} has been garbage collected") + return session + return session_entry def _build_recv_req_info(self, task: KVRecvTask) -> RecvReqInfo: self_ri = self._registrar.self_rank_info @@ -1814,6 +2015,8 @@ def dispatch_task(self, task: KVRecvTask) -> None: "dispatch_task: context peer incompatible, failing request " f"unique_rid={task._unique_rid}: {e}" ) + if self._enforce_physical_ownership: + task.cancel_unpublished() task.fail(e) return @@ -1884,22 +2087,64 @@ def dispatch_task(self, task: KVRecvTask) -> None: f"dispatch_task: RxSession {task._unique_rid} not found; " "session may have been closed before dispatch" ) - session.mark_transferring(task.slice_id) - # Cache sender endpoints so cancel() can send CANCEL_SESSION to them. - session._sender_endpoints.update( - peer_infos.sender_endpoints[rank] for rank in peer_overlap.ranks - ) + sender_endpoints = {peer_infos.sender_endpoints[rank] for rank in peer_overlap.ranks} # Fan-in: each sender gets its own sub-region base (writers must not overwrite); else serialize once. fanin_bounce = bounced and task.expected_transfers > 1 key = (receiver_req.unique_rid, receiver_req.slice_id) receiver_req_bytes = None if fanin_bounce else receiver_req.to_bytes() - for i, rank in enumerate(peer_overlap.ranks): - if task._perf_timer is not None: - task._perf_timer.record_task_start(rank) + if not self._enforce_physical_ownership: + session.mark_transferring(task.slice_id) + # Cache sender endpoints so cancel() can send CANCEL_SESSION to them. + session._sender_endpoints.update(sender_endpoints) + for i, rank in enumerate(peer_overlap.ranks): + if task._perf_timer is not None: + task._perf_timer.record_task_start(rank) + if fanin_bounce: + receiver_req.bounce_dst_base = self._bounce.writer_base(key, i) + receiver_req_bytes = receiver_req.to_bytes() + self._request_sender_data(peer_infos.sender_endpoints[rank], receiver_req_bytes) + return + + peer_ranks = list(peer_overlap.ranks) + serialized_requests: list[tuple[int, bytes]] = [] + for i, rank in enumerate(peer_ranks): if fanin_bounce: receiver_req.bounce_dst_base = self._bounce.writer_base(key, i) - receiver_req_bytes = receiver_req.to_bytes() - self._request_sender_data(peer_infos.sender_endpoints[rank], receiver_req_bytes) + payload = receiver_req.to_bytes() + else: + assert receiver_req_bytes is not None + payload = receiver_req_bytes + serialized_requests.append((rank, payload)) + + published_writers: set[int] = set() + + def publish_requests() -> None: + for rank, payload in serialized_requests: + if task._perf_timer is not None: + task._perf_timer.record_task_start(rank) + published_writers.add(rank) + try: + self._request_sender_data(peer_infos.sender_endpoints[rank], payload) + except Exception: + published_writers.discard(rank) + raise + + # Gen-first ADP publishes the destination to every eligible DP group, + # but the qualified immutable-request, no-retry/no-reroute profile + # admits exactly one group. Track that group's distinct terminal + # writers by count instead of sealing the full broadcast union. + writer_cohort = ( + set(peer_ranks) if sender_dp_rank is not None or peer_infos.dp_size == 1 else None + ) + if not session.try_begin_transfer( + task.slice_id, + sender_endpoints, + writer_cohort, + publish=publish_requests, + published_writers=published_writers, + ): + self._bounce.release_idle_reservation(key) + task.cancel_unpublished() return @staticmethod @@ -2017,15 +2262,15 @@ def _handle_cancel_session(self, message: list[bytes]): unique_rid = int(message[1]) session = None with self._sessions_lock: - session_ref = self._sessions.get(unique_rid) - if session_ref is None: + session_entry = self._sessions.get(unique_rid) + session = self._resolve_session_entry(unique_rid, session_entry) + if session is None: self._pre_cancelled_rids.add(unique_rid) - else: - session = session_ref() - if session is None: - self._pre_cancelled_rids.add(unique_rid) if session is not None: - session.cancel() + if self._enforce_physical_ownership: + session.cancel_local() + else: + session.cancel() def _process_kv_agent_result(self, _send_id: bytes, message: list[bytes]): if message[0] != MessageType.KV_AGENT_RESULT: @@ -2104,6 +2349,7 @@ def __init__( ) self._timeout_s = timeout_s self._need_aux = params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST + self._enforce_physical_ownership = getattr(receiver, "_enforce_physical_ownership", False) self._receiver: Receiver # narrow base class type for Pylance self.request_id = request_id self._aux_buffer = aux_buffer @@ -2118,6 +2364,11 @@ def __init__( self._aux_count = 0 self._aux_status: TaskStatus = TaskStatus.INIT self._sender_endpoints: set[str] = set() + # Serialize REQUEST_DATA publication with cancellation notification + # without holding the session state lock across a potentially blocking + # network send. The ordering is publication -> cancellation whenever + # publication wins the state transition. + self._publication_lock = threading.Lock() self.lock = threading.Lock() self._receiver.setup_session(self) @@ -2149,13 +2400,66 @@ def status(self) -> SessionStatus: return SessionStatus.TRANSFERRING return SessionStatus.INIT - def mark_transferring(self, slice_id: int): - with self.lock: - self._kv_tasks[slice_id].status = TaskStatus.TRANSFERRING + def try_begin_transfer( + self, + slice_id: int, + sender_endpoints: set[str], + writer_cohort: Optional[set[int]] = None, + *, + publish: Optional[Callable[[], None]] = None, + published_writers: Optional[set[int]] = None, + ) -> bool: + """Seal and publish a writer cohort unless cancellation won first.""" + if publish is not None and published_writers is None: + raise ValueError("publication must track successfully queued writers") + with self._publication_lock: + with self.lock: + if self._closed or self._terminal_status is not None: + return False + task = self._kv_tasks[slice_id] + task.seal_writer_cohort(writer_cohort) + task.status = TaskStatus.TRANSFERRING + self._sender_endpoints.update(sender_endpoints) + if publish is not None: + assert published_writers is not None + try: + publish() + if writer_cohort is not None and published_writers != writer_cohort: + raise RuntimeError("publication did not queue the complete writer cohort") + except Exception: + # ZMQ delivers multipart messages atomically. A successful + # send transfers responsibility to that writer; a failed + # send does not. Retain only the successfully queued prefix + # and wait for each of those writers to settle. + self._receiver._bounce.abort_publication( + (self.disagg_request_id, task.slice_id), + published_writers, + ) + task.abort_publication(published_writers) + raise + task.finish_publication() + return True + + def mark_transferring(self, slice_id: int, writer_cohort: Optional[set[int]] = None) -> None: + if not self._enforce_physical_ownership: + with self.lock: + self._kv_tasks[slice_id].status = TaskStatus.TRANSFERRING + return + if writer_cohort is None: + raise ValueError("ownership-mode publication requires an explicit writer cohort") + if not self.try_begin_transfer(slice_id, set(), writer_cohort): + raise RuntimeError( + f"RxSession {self.disagg_request_id} became terminal before publication" + ) def receive(self, slice: KVSlice) -> None: if self.transfer_start_time is None: self.transfer_start_time = tensorrt_llm.bindings.global_steady_clock_now() + if self._enforce_physical_ownership: + task = self.prepare_receive(slice) + if task is not None: + self.dispatch_prepared_receive(task) + return params = self._base_args.params slice_id = len(self._kv_tasks) task = KVRecvTask( @@ -2168,6 +2472,37 @@ def receive(self, slice: KVSlice) -> None: self._kv_tasks.append(task) self._receiver.dispatch_task(task) + def prepare_receive(self, slice: KVSlice) -> Optional[KVRecvTask]: + """Create an unpublished task under the cancellation linearization lock.""" + with self.lock: + if self._closed or self._terminal_status is not None: + return None + params = self._base_args.params + task = KVRecvTask( + self.disagg_request_id, + slice, + len(self._kv_tasks), + params, + aux_slot=self.aux_slot, + ) + task.begin_publication() + self._kv_tasks.append(task) + return task + + def dispatch_prepared_receive(self, task: KVRecvTask) -> None: + try: + self._receiver.dispatch_task(task) + except Exception as error: + if task.cancel_unpublished(): + self._receiver._bounce.release_idle_reservation( + (self.disagg_request_id, task.slice_id) + ) + with self.lock: + task.fail(error) + if self._terminal_status is None: + self._terminal_status = SessionStatus.ERROR + raise + def process_kv_agent_result( self, peer_rank: int, @@ -2180,7 +2515,6 @@ def process_kv_agent_result( transfer_size: int = 0, ): with self.lock: - self.kv_cache_size_bytes += transfer_size assert receiver_slice_id < len(self._kv_tasks), ( f"Receiver got receiver_slice_id={receiver_slice_id} but only has " f"{len(self._kv_tasks)} receive task(s) for request {self.request_id}. " @@ -2188,65 +2522,84 @@ def process_kv_agent_result( f"address a task this receiver posted." ) task = self._kv_tasks[receiver_slice_id] + if status not in (AgentResult.SUCCESS, AgentResult.FAILED): + raise ValueError( + f"Session {self.request_id} received unknown task status: {status.value}" + ) + if self._enforce_physical_ownership: + if status == AgentResult.FAILED or is_last_slice: + accepted, all_succeeded = task.record_writer_result( + peer_rank, + status == AgentResult.SUCCESS, + wait_for_local_completion=(status == AgentResult.SUCCESS), + ) + if not accepted: + return + else: + all_succeeded = False + else: + all_succeeded = False + if status == AgentResult.SUCCESS and is_last_slice: + task.last_slice_count += 1 + all_succeeded = task.last_slice_count == task.expected_transfers + self.kv_cache_size_bytes += transfer_size if status == AgentResult.SUCCESS: from .bounce import scatter_write_result on_done = None - if is_last_slice: - task.last_slice_count += 1 - if task.last_slice_count == task.expected_transfers: - # Completing message: defer task.complete()+perf until the scatter has actually - # landed. scatter_write_result fires this inline for the non-bounced path, or on - # the scatter worker (after cudaStreamSynchronize) for the bounced path, so the - # gen consumer never observes completion before the KV is scattered into place. - request_id = self.request_id - ri = self._receiver._registrar.self_rank_info - instance_name, instance_rank = ri.instance_name, ri.instance_rank - - def on_done( - success, - task=task, - peer_rank=peer_rank, - receiver_slice_id=receiver_slice_id, - request_id=request_id, - instance_name=instance_name, - instance_rank=instance_rank, - ): - # Runs on the scatter worker thread for the bounced path. Touches only this - # task's own status/_event/_perf_timer (no RxSession.lock, no shared session - # state), so it is lock-free. complete() sets status before _event, keeping - # wait_complete's status-first poll correct. - if not success: - task.fail( - RuntimeError( - f"KV bounce scatter failed for request {request_id} " - f"slice={receiver_slice_id}" - ) - ) - return - if task.status == TaskStatus.ERROR: - return # a concurrent FAILED writer already failed it; don't un-fail - try: - if task._perf_timer is not None: - task._perf_timer.record_task_end(peer_rank) - task.print_perf_info(peer_rank, instance_name, instance_rank) - except Exception as e: # perf is best-effort; never block completion - logger.warning( - f"KV transfer perf logging failed for request {request_id} " - f"slice={receiver_slice_id}: {e}" - ) - task.complete() - # Transfer end for perf/time-sync: only meaningful once every slice has - # landed. Plain attribute write (atomic under the GIL); on_done must stay - # lock-free, and consumers only read it after wait_complete succeeds. - if all(t.status == TaskStatus.TRANSFERRED for t in self._kv_tasks): - self.transfer_end_time = ( - tensorrt_llm.bindings.global_steady_clock_now() + if is_last_slice and all_succeeded: + # Completing message: defer task.complete()+perf until the scatter has actually + # landed. scatter_write_result fires this inline for the non-bounced path, or on + # the scatter worker (after cudaStreamSynchronize) for the bounced path, so the + # gen consumer never observes completion before the KV is scattered into place. + request_id = self.request_id + ri = self._receiver._registrar.self_rank_info + instance_name, instance_rank = ri.instance_name, ri.instance_rank + + def on_done( + success, + task=task, + peer_rank=peer_rank, + receiver_slice_id=receiver_slice_id, + request_id=request_id, + instance_name=instance_name, + instance_rank=instance_rank, + ): + # Runs on the scatter worker thread for the bounced path. Touches only this + # task's own status/_event/_perf_timer (no RxSession.lock, no shared session + # state), so it is lock-free. complete() sets status before _event, keeping + # wait_complete's status-first poll correct. + if self._enforce_physical_ownership: + task.finish_local_completion() + if not success: + task.fail( + RuntimeError( + f"KV bounce scatter failed for request {request_id} " + f"slice={receiver_slice_id}" ) - logger.debug( - f"KV transfer complete for request {request_id} " - f"slice={receiver_slice_id}" ) + return + if task.status == TaskStatus.ERROR: + return # a concurrent FAILED writer already failed it; don't un-fail + try: + if task._perf_timer is not None: + task._perf_timer.record_task_end(peer_rank) + task.print_perf_info(peer_rank, instance_name, instance_rank) + except Exception as e: # perf is best-effort; never block completion + logger.warning( + f"KV transfer perf logging failed for request {request_id} " + f"slice={receiver_slice_id}: {e}" + ) + task.complete() + # Transfer end for perf/time-sync: only meaningful once every slice has + # landed. Plain attribute write (atomic under the GIL); on_done must stay + # lock-free, and consumers only read it after wait_complete succeeds. + if all(t.status == TaskStatus.TRANSFERRED for t in self._kv_tasks): + self.transfer_end_time = tensorrt_llm.bindings.global_steady_clock_now() + logger.debug( + f"KV transfer complete for request {request_id} " + f"slice={receiver_slice_id}" + ) scatter_write_result( self._receiver._bounce, @@ -2259,7 +2612,8 @@ def on_done( ) elif status == AgentResult.FAILED: detail = ( - f"KV transfer failed for request {self.request_id} slice={receiver_slice_id} " + f"KV transfer failed for request {self.request_id} " + f"slice={receiver_slice_id} " f"peer_rank={peer_rank} is_last_slice={is_last_slice} " f"(reported by remote agent; see sender-side log for nixl_status)" ) @@ -2272,10 +2626,6 @@ def on_done( task.fail(RuntimeError(detail)) if self._terminal_status is None: # Don't overwrite CANCELLED with ERROR self._terminal_status = SessionStatus.ERROR - else: - raise ValueError( - f"Session {self.request_id} received unknown task status: {status.value}" - ) def process_aux_agent_result(self, _peer_rank: int, status: AgentResult): # Aux is session-level (not per-slice); expected_transfers is identical @@ -2337,45 +2687,70 @@ def is_completed(self) -> bool: """Non-blocking check: has the transfer completed successfully?""" status = self.status if self._need_aux: - return status == SessionStatus.FULLY_TRANSFERRED - return status in (SessionStatus.KV_TRANSFERRED, SessionStatus.FULLY_TRANSFERRED) + logically_completed = status == SessionStatus.FULLY_TRANSFERRED + else: + logically_completed = status in ( + SessionStatus.KV_TRANSFERRED, + SessionStatus.FULLY_TRANSFERRED, + ) + return logically_completed and ( + not self._enforce_physical_ownership or self.resources_drained() + ) def has_failed(self) -> bool: """Non-blocking check: has the transfer failed or been cancelled?""" return self.status in (SessionStatus.ERROR, SessionStatus.CANCELLED) - def cancel(self) -> None: - """Cancel the session and notify the remote sender. + def resources_drained(self) -> bool: + if not self._enforce_physical_ownership: + return not any(task.status == TaskStatus.TRANSFERRING for task in self._kv_tasks) + return all(task.resources_drained for task in self._kv_tasks) - Safe to call multiple times. TRANSFERRING tasks keep running (mid-write). - Only INIT tasks have their events signalled immediately. - The lock serializes with process_kv_agent_result() / process_aux_agent_result(). - """ + def cancel_local(self) -> bool: + """Commit cancellation under the same lock used for publication.""" with self.lock: if self._terminal_status == SessionStatus.CANCELLED: - return + return False self._terminal_status = SessionStatus.CANCELLED exc = RuntimeError(f"RxSession {self.disagg_request_id} cancelled") for task in self._kv_tasks: rid_slice = (self.disagg_request_id, task.slice_id) if task.status == TaskStatus.INIT: - # INIT = reserved but no write in flight, so freeing its bounce reservation here - # is safe. + if self._enforce_physical_ownership: + task.cancel_unpublished() self._receiver._bounce.release_idle_reservation(rid_slice) task.fail(exc) - elif task.status == TaskStatus.TRANSFERRING: - # A write may still be mid-flight, so quarantine the region rather than freeing - # it; this keeps a cancelled transfer from leaking. No-op when bounce is off. - self._receiver._bounce.orphan_reservation(rid_slice) - # Send outside the lock to avoid holding it during I/O. - self._receiver.send_cancel_to_senders(self.disagg_request_id, self._sender_endpoints) + else: + has_active_access = ( + not task.resources_drained + if self._enforce_physical_ownership + else task.status == TaskStatus.TRANSFERRING + ) + if has_active_access: + self._receiver._bounce.orphan_reservation(rid_slice) + return True + + def capture_cancel_targets(self) -> set[str]: + with self.lock: + return set(self._sender_endpoints) + + def notify_cancel(self, sender_endpoints: Optional[set[str]] = None) -> None: + if sender_endpoints is None: + sender_endpoints = self.capture_cancel_targets() + with self._publication_lock: + self._receiver.send_cancel_to_senders(self.disagg_request_id, sender_endpoints) + + def cancel(self) -> None: + """Cancel locally, then notify every writer without holding the lock.""" + if self.cancel_local(): + self.notify_cancel() def has_transferring_tasks(self) -> bool: - """True if any KV task is currently mid-write (TRANSFERRING). + """True while a destination accessor may remain active. cancel_request() must return False while this is True. """ - return any(t.status == TaskStatus.TRANSFERRING for t in self._kv_tasks) + return not self.resources_drained() def wait_complete(self, blocking: bool = False) -> Optional[WaitResult]: """Poll or block until transfer completes. @@ -2385,6 +2760,8 @@ def wait_complete(self, blocking: bool = False) -> Optional[WaitResult]: With blocking=True: waits up to _timeout_s for each task. Returns WaitResult.COMPLETED on full success, WaitResult.FAILED on error/timeout. """ + if self.has_failed() and self._enforce_physical_ownership: + return WaitResult.FAILED if self.resources_drained() else None if not blocking: # Use task.status instead of task.wait(timeout=0): task.complete() # sets status before event, so a GIL switch between the two steps @@ -2415,10 +2792,18 @@ def wait_complete(self, blocking: bool = False) -> Optional[WaitResult]: time.sleep(0.001) return WaitResult.COMPLETED - def close(self): - if getattr(self, "_closed", False): - return - self._closed = True + def close(self) -> bool: + if self._enforce_physical_ownership: + with self.lock: + if self._closed: + return True + if not self.resources_drained(): + return False + self._closed = True + else: + if getattr(self, "_closed", False): + return True + self._closed = True if self._aux_buffer is not None and self.aux_slot is not None: self._aux_buffer.free_slot(self.aux_slot) self.aux_slot = None @@ -2429,16 +2814,24 @@ def close(self): for task in self._kv_tasks: self._receiver._bounce.orphan_reservation((self.disagg_request_id, task.slice_id)) self._receiver.clear_session(self.disagg_request_id) + return True def __enter__(self): return self def __exit__(self, _exc_type, _exc, _tb): - self.close() + if not self.close(): + logger.warning( + f"RxSession {self.disagg_request_id} refused close while resources remain active" + ) def __del__(self): try: - self.close() + if not self.close(): + logger.warning( + f"RxSession {self.disagg_request_id} refused destructor close while " + "resources remain active" + ) except Exception as e: logger.warning(f"RxSession.__del__: exception during close: {e}") diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index c8bea047f730..62107d31e447 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -86,6 +86,7 @@ def __init__( kv_cache_manager: KVCacheManager, cache_transceiver_config: CacheTransceiverConfig, ): + self._shutdown_complete = False self._dist: Distributed = dist self._kv_cache_manager = kv_cache_manager self._mapping = mapping @@ -257,18 +258,24 @@ def summarize( ) def shutdown(self) -> None: - if getattr(self, "_shutdown", False): + if getattr(self, "_shutdown_complete", False): return - self._shutdown = True - for session in list(self._send_sessions.values()): - session.close() - for session in list(self._recv_sessions.values()): - session.close() + # This flag records completed teardown, not an attempted shutdown. If + # an active owner refuses closure, leave it false so shutdown can be + # retried after the physical operation drains. + # Close receive owners before touching send-side or worker state. The + # separate drain snapshot would not be sufficient: late evidence can make an + # owner unretirable before close() acquires its lock. + for rid, session in list(self._recv_sessions.items()): + self._close_session_or_raise(session, rid, "shutdown") + for rid, session in list(self._send_sessions.items()): + self._close_session_or_raise(session, rid, "shutdown") self._send_sessions.clear() self._send_reqs.clear() self._recv_sessions.clear() self._recv_reqs.clear() self._transfer_worker.shutdown() + self._shutdown_complete = True def __enter__(self): return self @@ -552,8 +559,27 @@ def _consensus_outcome( return new_cancelled, new_failed, new_completed def _gen_consensus_outcome(self, to_process, cancelled, failed, completed): - return self._consensus_outcome( - to_process, cancelled, failed, completed, self._gen_allgather, self._gen_need_sync + # A failure/cancellation may be global, but reuse is safe only after + # every participating rank has drained its local physical accessor. + locally_retirable = [] + for rid in to_process: + session = self._recv_sessions[rid] + if not self._ownership_blocks_retirement(session): + locally_retirable.append(rid) + new_cancelled, new_failed, new_completed, globally_retirable = self._consensus_outcome( + to_process, + cancelled, + failed, + completed, + self._gen_allgather, + self._gen_need_sync, + locally_retirable, + ) + retirable = set(globally_retirable) + return ( + [rid for rid in new_cancelled if rid in retirable], + [rid for rid in new_failed if rid in retirable], + new_completed, ) def _ctx_consensus_outcome(self, to_process, cancelled, failed, completed, locally_quiesced): @@ -631,6 +657,8 @@ def _collect_done(self, sessions: dict, reqs: dict): if session.is_completed(): completed.append(rid) elif session.has_failed(): + if self._ownership_blocks_retirement(session): + continue failed.append(rid) return completed, failed @@ -651,15 +679,17 @@ def _close_failed_sessions( self, sessions: dict, reqs: dict, failed: list, mark_retired: bool = False ): for rid in failed: + session = sessions.get(rid) + if session is None: + continue + self._close_session_or_raise(session, rid, "failed") req = reqs.pop(rid, None) if req is not None: if not mark_retired: req.state = LlmRequestState.DISAGG_TRANS_ERROR else: req.py_kv_send_session_retired = True - session = sessions.pop(rid, None) - if session is not None: - session.close() + sessions.pop(rid, None) def _retire_send_session(self, rid: int, req: Optional[LlmRequest] = None) -> None: """Close a send session and prevent later chunks from recreating it.""" @@ -672,6 +702,26 @@ def _retire_send_session(self, rid: int, req: Optional[LlmRequest] = None) -> No if req is not None: req.py_kv_send_session_retired = True + @staticmethod + def _ownership_blocks_retirement(session: object) -> bool: + """Whether an ownership-enabled session still holds physical resources.""" + if not getattr(session, "_enforce_physical_ownership", False): + return False + resources_drained = getattr(session, "resources_drained", None) + return resources_drained is None or not resources_drained() + + def _close_session_or_raise(self, session: object, rid: int, outcome: str) -> None: + """Close one terminal session or fail-stop instead of diverging locally.""" + if self._ownership_blocks_retirement(session): + raise RuntimeError( + f"refusing to retire {outcome} KV transfer rid={rid}: " + "physical resources remain active" + ) + if session.close() is False: + raise RuntimeError( + f"refusing to retire {outcome} KV transfer rid={rid}: session close refused" + ) + def _apply_aux(self, session, req: LlmRequest): """Unpack aux tokens from session into request's context_phase_params.""" session.unpack_aux(req) @@ -815,6 +865,7 @@ def respond_and_send_async(self, req: LlmRequest) -> None: @nvtx_range("KvCacheTransceiverV2.request_and_receive_sync") def request_and_receive_sync(self, req: LlmRequest) -> None: rid = get_unique_rid(req) + self._ever_had_recv_session = True if rid in self._recv_sessions: logger.warning( f"request_and_receive_sync: rid={rid} already has a recv session, skipping" @@ -843,10 +894,15 @@ def request_and_receive_sync(self, req: LlmRequest) -> None: req.state = LlmRequestState.DISAGG_TRANS_ERROR raise finally: - if session is not None: - session.close() - self._recv_sessions.pop(rid, None) - self._recv_reqs.pop(rid, None) + close_succeeded = session is None or session.close() is not False + if close_succeeded: + self._recv_sessions.pop(rid, None) + self._recv_reqs.pop(rid, None) + else: + logger.error( + f"request_and_receive_sync: retaining rid={rid} because receive " + "resources remain active" + ) @nvtx_range("KvCacheTransceiverV2.request_and_receive_async") def request_and_receive_async(self, req: LlmRequest) -> None: @@ -976,6 +1032,8 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]) -> GenT session = self._recv_sessions[rid] result = session.wait_complete(blocking=block_all) if session.status == SessionStatus.CANCELLED: + if self._ownership_blocks_retirement(session): + continue # Session cancelled — either by local cancel_request() (user # cancel) or by a remote CANCEL_SESSION message (e.g. CTX # server timeout). Return the req objects so the caller can @@ -989,6 +1047,12 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]) -> GenT req.set_kv_cache_size(session.kv_cache_size_bytes) completed.append(rid) elif result == WaitResult.FAILED: + # RxSession.wait_complete() already withholds FAILED while an + # ownership-enabled accessor remains active. Keep the caller + # boundary fail-closed as well so alternate/test session + # implementations cannot authorize request retirement early. + if self._ownership_blocks_retirement(session): + continue failed.append(rid) # else: None — KV done but aux still in flight; re-poll next cycle @@ -999,8 +1063,9 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]) -> GenT cancelled_reqs = [] for rid in cancelled: + session = self._recv_sessions[rid] + self._close_session_or_raise(session, rid, "cancelled") cancelled_reqs.append(self._recv_reqs[rid]) - self._recv_sessions[rid].close() del self._recv_reqs[rid] del self._recv_sessions[rid] @@ -1026,8 +1091,8 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]) -> GenT if self._need_aux_transfer(req): self._apply_aux(session, req) self._assert_disagg_history_declared(req) + self._close_session_or_raise(session, rid, "completed") req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE - session.close() del self._recv_reqs[rid] del self._recv_sessions[rid] if failed: @@ -1139,8 +1204,9 @@ def cancel_request(self, req: LlmRequest) -> bool: self._recv_sessions[rid].cancel() if self._recv_sessions[rid].has_transferring_tasks(): has_transferring = True + elif self._recv_sessions[rid].close() is False: + has_transferring = True else: - self._recv_sessions[rid].close() del self._recv_reqs[rid] del self._recv_sessions[rid] diff --git a/tests/unittest/disaggregated/test_bounce.py b/tests/unittest/disaggregated/test_bounce.py index fbc1194f12c9..24f0122da025 100644 --- a/tests/unittest/disaggregated/test_bounce.py +++ b/tests/unittest/disaggregated/test_bounce.py @@ -978,6 +978,29 @@ def test_fanin_failed_then_success_releases(self): assert ret.success is False assert c.state is bcore.TransferState.FAILED + def test_partial_publication_waits_for_published_writers_without_scatter(self): + c = self._ctx(2) + c.record_writer_result(7, succeeded=True, src_base=0, **self._dst()) + + c.abort_publication(published_writers={7}) + + assert not c.ready_to_scatter() + assert c.ready_to_settle() + ret = c.settle() + assert ret.disposition is bcore.Disposition.RELEASE + assert ret.success is False + assert c.state is bcore.TransferState.FAILED + + def test_failed_publication_with_no_published_writer_settles_immediately(self): + c = self._ctx(2) + + c.abort_publication(published_writers=set()) + + assert c.ready_to_settle() + ret = c.settle() + assert ret.disposition is bcore.Disposition.RELEASE + assert ret.success is False + def test_orphan_quarantines(self): c = self._ctx(2) c.record_writer_result(7, succeeded=True, src_base=0, **self._dst()) diff --git a/tests/unittest/disaggregated/test_cache_reuse_adapter.py b/tests/unittest/disaggregated/test_cache_reuse_adapter.py index f8b1f773925f..059ae07efd14 100644 --- a/tests/unittest/disaggregated/test_cache_reuse_adapter.py +++ b/tests/unittest/disaggregated/test_cache_reuse_adapter.py @@ -859,6 +859,7 @@ def _tc(): tc._send_reqs = {} tc._recv_reqs = {} tc._transfer_worker = MagicMock() + tc._shutdown_complete = False return tc def test_enter_returns_self(self): @@ -871,7 +872,7 @@ def test_exit_calls_shutdown(self): with tc: pass tc._transfer_worker.shutdown.assert_called_once() - assert tc._shutdown is True + assert tc._shutdown_complete is True def test_exit_calls_shutdown_on_exception(self): tc = self._tc() @@ -880,10 +881,10 @@ def test_exit_calls_shutdown_on_exception(self): raise RuntimeError("boom") # __exit__ still ran shutdown despite the in-block exception. tc._transfer_worker.shutdown.assert_called_once() - assert tc._shutdown is True + assert tc._shutdown_complete is True def test_shutdown_is_idempotent(self): tc = self._tc() tc.shutdown() - tc.shutdown() # second call short-circuits on the _shutdown guard. + tc.shutdown() # second call short-circuits after completed teardown. tc._transfer_worker.shutdown.assert_called_once() diff --git a/tests/unittest/disaggregated/test_chunked_transfer.py b/tests/unittest/disaggregated/test_chunked_transfer.py index 7cd2bad8b94b..89002470f95f 100644 --- a/tests/unittest/disaggregated/test_chunked_transfer.py +++ b/tests/unittest/disaggregated/test_chunked_transfer.py @@ -69,6 +69,7 @@ def _stub_sender(): def _stub_receiver(): """Create a stub receiver with no-op methods needed by RxSession.""" receiver = MagicMock() + receiver._enforce_physical_ownership = False receiver.setup_session = MagicMock() receiver.dispatch_task = MagicMock() return receiver @@ -903,10 +904,10 @@ def test_failed_pipelined_send_retires_without_mutating_request_state(): py_kv_send_session_retired=False, ) session = MagicMock() + session._enforce_physical_ownership = False + transceiver = object.__new__(KvCacheTransceiverV2) - KvCacheTransceiverV2._close_failed_sessions( - MagicMock(), {42: session}, {42: request}, [42], mark_retired=True - ) + transceiver._close_failed_sessions({42: session}, {42: request}, [42], mark_retired=True) assert request.state == LlmRequestState.CONTEXT_INIT assert request.py_kv_send_session_retired diff --git a/tests/unittest/disaggregated/test_transfer_ownership_regressions.py b/tests/unittest/disaggregated/test_transfer_ownership_regressions.py new file mode 100644 index 000000000000..055c9c5198a7 --- /dev/null +++ b/tests/unittest/disaggregated/test_transfer_ownership_regressions.py @@ -0,0 +1,1162 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Regressions for receive-side KV transfer ownership boundaries.""" + +from __future__ import annotations + +import queue +import threading +from collections.abc import Callable +from types import SimpleNamespace +from unittest.mock import Mock + +import numpy as np +import pytest + +import tensorrt_llm._torch.disaggregation.native.transfer as transfer_mod +from tensorrt_llm import DisaggregatedParams +from tensorrt_llm._torch.disaggregation.base.transfer import KVSlice, SessionStatus, WaitResult +from tensorrt_llm._torch.disaggregation.native.bounce.core import TransferContext +from tensorrt_llm._torch.disaggregation.native.transfer import ( + AgentResult, + KVRecvTask, + MessageType, + Receiver, + RxSession, +) +from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 +from tensorrt_llm.disaggregated_params import DisaggScheduleStyle + + +class _BounceProbe: + """No-bounce probe that records physical-owner cleanup decisions.""" + + def __init__(self) -> None: + self.failed_writers: list[tuple[tuple[int, int], int]] = [] + self.orphaned: list[tuple[int, int]] = [] + + def record_failure(self, rid_slice: tuple[int, int], peer_rank: int) -> None: + self.failed_writers.append((rid_slice, peer_rank)) + + def reserve(self, _receiver_req, _num_writers: int, *, extra_bytes: int = 0) -> bool: + del extra_bytes + return False + + def release_idle_reservation(self, _rid_slice: tuple[int, int]) -> None: + return + + def orphan_reservation(self, rid_slice: tuple[int, int]) -> None: + self.orphaned.append(rid_slice) + + def abort_publication( + self, + _rid_slice: tuple[int, int], + _published_writers: set[int], + ) -> None: + return + + def is_bounced(self, _rid_slice: tuple[int, int]) -> bool: + return False + + +class _LateReservationBounce(_BounceProbe): + """Track a reservation created after cancellation already ran cleanup.""" + + def __init__(self) -> None: + super().__init__() + self.active_reservations: set[tuple[int, int]] = set() + + def release_idle_reservation(self, rid_slice: tuple[int, int]) -> None: + self.active_reservations.discard(rid_slice) + + +class _PartialFanInBounce(_BounceProbe): + """CPU model of one bounced fan-in reservation.""" + + def __init__(self) -> None: + super().__init__() + self.context: TransferContext | None = None + self.release_count = 0 + self.scatter_count = 0 + + def reserve(self, receiver_req, num_writers: int, *, extra_bytes: int = 0) -> bool: + del extra_bytes + self.context = TransferContext( + rid_slice=(receiver_req.unique_rid, receiver_req.slice_id), + slot_id=0, + base_addr=0x1000, + per_writer_bytes=0x100, + num_writers=num_writers, + ) + return True + + def writer_base(self, rid_slice: tuple[int, int], writer_index: int) -> int | None: + if self.context is None or self.context.rid_slice != rid_slice: + return None + return self.context.writer_base(writer_index) + + def is_bounced(self, rid_slice: tuple[int, int]) -> bool: + return self.context is not None and self.context.rid_slice == rid_slice + + def _advance(self) -> None: + if self.context is None: + return + if self.context.ready_to_scatter(): + self.scatter_count += 1 + self.context.begin_scatter() + self.context.finish_scatter(True) + if not self.context.ready_to_settle(): + return + settlement = self.context.settle() + self.context = None + assert settlement is not None + self.release_count += 1 + if settlement.on_done is not None: + settlement.on_done(settlement.success) + + def abort_publication( + self, + rid_slice: tuple[int, int], + published_writers: set[int], + ) -> None: + assert self.context is not None and self.context.rid_slice == rid_slice + self.context.abort_publication(published_writers) + self._advance() + + def record_result( + self, + rid_slice: tuple[int, int], + peer_rank: int, + dst_ptrs=None, + sizes=None, + src_base=None, + on_done=None, + ) -> None: + assert self.context is not None and self.context.rid_slice == rid_slice + if on_done is not None: + self.context.on_done = on_done + self.context.record_writer_result( + peer_rank, + succeeded=True, + src_base=src_base, + dst_ptrs=dst_ptrs, + sizes=sizes, + ) + self._advance() + + +def _start_checked_thread( + target: Callable[[], None], + results: queue.Queue[Exception | None], +) -> threading.Thread: + """Start a daemon worker and surface expected assertion/runtime failures.""" + + def run() -> None: + try: + target() + except (AssertionError, RuntimeError, ValueError) as error: + results.put(error) + else: + results.put(None) + + thread = threading.Thread(target=run, daemon=True) + thread.start() + return thread + + +def _raise_thread_errors(results: queue.Queue[Exception | None], expected: int) -> None: + completed = [] + for _ in range(expected): + try: + completed.append(results.get_nowait()) + except queue.Empty as error: + raise AssertionError("worker terminated with an unexpected exception") from error + if not results.empty(): + raise AssertionError("worker reported more than one result") + for error in completed: + if error is not None: + raise error + + +class _TrackingLock: + """Record when one selected thread attempts to enter a critical section.""" + + def __init__(self, race_outcomes: queue.Queue[str]) -> None: + self._lock = threading.Lock() + self._race_outcomes = race_outcomes + self.tracked_thread_id: int | None = None + + def acquire(self, blocking: bool = True, timeout: float = -1) -> bool: + if threading.get_ident() == self.tracked_thread_id: + if self._lock.acquire(blocking=False): + return True + self._race_outcomes.put("blocked") + if not blocking: + return False + if timeout == -1: + return self._lock.acquire(blocking) + return self._lock.acquire(blocking, timeout) + + def release(self) -> None: + self._lock.release() + + def __enter__(self) -> "_TrackingLock": + self.acquire() + return self + + def __exit__(self, _exc_type, _exc_value, _traceback) -> None: + self.release() + + +class _ReceiverProbe: + """Minimal receiver that publishes one destination to two writers.""" + + def __init__(self) -> None: + self._bounce = _BounceProbe() + self._enforce_physical_ownership = True + self._session: RxSession | None = None + self.clear_count = 0 + self.cancel_count = 0 + + def setup_session(self, session: RxSession) -> None: + self._session = session + + def dispatch_task(self, task: KVRecvTask) -> None: + assert self._session is not None + task.expected_transfers = 2 + self._session.mark_transferring(task.slice_id, writer_cohort={0, 1}) + + def send_cancel_to_senders(self, _unique_rid: int, _sender_endpoints: set[str]) -> None: + self.cancel_count += 1 + + def clear_session(self, _unique_rid: int) -> None: + self.clear_count += 1 + + +class _OneSlotAllocator: + """Model the caller that releases a request allocation on a True result.""" + + def __init__(self, owner: int) -> None: + self.owner: int | None = owner + self.release_count = 0 + + @property + def is_reusable(self) -> bool: + return self.owner is None + + def apply_reuse_decision(self, safe_to_reuse: bool) -> None: + if safe_to_reuse and self.owner is not None: + self.owner = None + self.release_count += 1 + + +def _make_rx_session(receiver: object, rid: int) -> RxSession: + return RxSession( + request_id=rid, + params=DisaggregatedParams(disagg_request_id=rid), + receiver=receiver, + ) + + +@pytest.mark.cpu_only +def test_failed_writer_cannot_authorize_reuse_while_sibling_is_active() -> None: + sibling_started = threading.Event() + release_sibling = threading.Event() + thread_results: queue.Queue[Exception | None] = queue.Queue() + receiver = _ReceiverProbe() + session = _make_rx_session(receiver, rid=41) + session.receive(KVSlice(is_last_slice=True)) + request = SimpleNamespace( + request_id=41, + py_disaggregated_params=DisaggregatedParams(disagg_request_id=41), + ) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._wait_reqs = {} + transceiver._send_sessions = {} + transceiver._send_reqs = {} + transceiver._recv_sessions = {41: session} + transceiver._recv_reqs = {41: request} + allocator = _OneSlotAllocator(owner=41) + + def finish_sibling_writer() -> None: + sibling_started.set() + release_sibling.wait() + session.process_kv_agent_result( + peer_rank=1, + receiver_slice_id=0, + is_last_slice=True, + status=AgentResult.SUCCESS, + ) + + sibling_thread = _start_checked_thread(finish_sibling_writer, thread_results) + try: + assert sibling_started.wait(timeout=10) + + # Writer 0 is terminal, but writer 1 has not reported a terminal result + # and may still write to the same receive-side KV allocation. + session.process_kv_agent_result( + peer_rank=0, + receiver_slice_id=0, + is_last_slice=True, + status=AgentResult.FAILED, + ) + + safe_to_reuse = transceiver.cancel_request(request) + allocator.apply_reuse_decision(safe_to_reuse) + + assert safe_to_reuse is False, ( + "cancel_request() authorized receive-side KV reuse before every " + "published writer reached a terminal physical state" + ) + assert not allocator.is_reusable + assert 41 in transceiver._recv_sessions + assert receiver.clear_count == 0 + finally: + release_sibling.set() + sibling_thread.join(timeout=10) + assert not sibling_thread.is_alive() + _raise_thread_errors(thread_results, expected=1) + + # The allocation becomes reusable only after the remaining writer reports + # a terminal physical result. Repeated cancellation must not clean it up + # more than once. + safe_to_reuse = transceiver.cancel_request(request) + allocator.apply_reuse_decision(safe_to_reuse) + assert safe_to_reuse is True + assert allocator.is_reusable + assert allocator.release_count == 1 + assert 41 not in transceiver._recv_sessions + assert receiver.clear_count == 1 + repeated_decision = transceiver.cancel_request(request) + assert repeated_decision is True + allocator.apply_reuse_decision(repeated_decision) + assert allocator.release_count == 1 + assert receiver.clear_count == 1 + + +@pytest.mark.cpu_only +def test_pre_cancelled_rx_session_never_publishes_destination( + monkeypatch: pytest.MonkeyPatch, +) -> None: + rid = 73 + receiver = object.__new__(Receiver) + receiver._sessions_lock = threading.Lock() + receiver._sessions = {} + receiver._pre_cancelled_rids = {rid} + receiver._bounce = _BounceProbe() + receiver._enforce_physical_ownership = True + receiver._shutdown = True + receiver.dispatch_task = Mock() + receiver.send_cancel_to_senders = Mock() + + monkeypatch.setattr( + transfer_mod.tensorrt_llm.bindings, + "global_steady_clock_now", + lambda: 0, + ) + + session = _make_rx_session(receiver, rid) + assert session.status == SessionStatus.CANCELLED + + session.receive(KVSlice(is_last_slice=True)) + + assert (len(session._kv_tasks), receiver.dispatch_task.call_count) == (0, 0), ( + "a pre-cancelled receive session created and published a destination task" + ) + + +@pytest.mark.cpu_only +def test_remote_cancel_resolves_strong_owned_session() -> None: + rid = 77 + receiver = object.__new__(Receiver) + receiver._sessions_lock = threading.Lock() + receiver._sessions = {} + receiver._pre_cancelled_rids = set() + receiver._bounce = _BounceProbe() + receiver._enforce_physical_ownership = True + receiver._shutdown = True + receiver.send_cancel_to_senders = Mock() + session = _make_rx_session(receiver, rid) + + assert receiver._sessions[rid] is session + receiver._handle_cancel_session([MessageType.CANCEL_SESSION, str(rid).encode("ascii")]) + + assert session.status == SessionStatus.CANCELLED + receiver.send_cancel_to_senders.assert_not_called() + + +@pytest.mark.cpu_only +def test_remote_cancelled_session_is_retained_until_writers_drain() -> None: + rid = 78 + request = SimpleNamespace(request_id=rid) + session = SimpleNamespace( + _enforce_physical_ownership=True, + status=SessionStatus.CANCELLED, + is_completed=Mock(return_value=False), + has_failed=Mock(return_value=True), + wait_complete=Mock(return_value=None), + resources_drained=Mock(return_value=False), + close=Mock(return_value=False), + ) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._ever_had_recv_session = True + transceiver._gen_need_sync = False + transceiver._mapping = SimpleNamespace( + pp_size=1, + enable_attention_dp=False, + world_size=1, + ) + transceiver._dist = SimpleNamespace(rank=0) + transceiver._gen_allgather = Mock() + transceiver._recv_sessions = {rid: session} + transceiver._recv_reqs = {rid: request} + + completed, failed, cancelled = transceiver.check_gen_transfer_status(None) + + assert (completed, failed, cancelled) == ([], [], []) + assert transceiver._recv_sessions[rid] is session + assert transceiver._recv_reqs[rid] is request + session.close.assert_not_called() + + +@pytest.mark.cpu_only +def test_failed_receive_session_is_retained_until_writers_drain() -> None: + rid = 87 + request = SimpleNamespace(request_id=rid) + session = SimpleNamespace( + _enforce_physical_ownership=True, + status=SessionStatus.ERROR, + is_completed=Mock(return_value=False), + has_failed=Mock(return_value=True), + wait_complete=Mock(return_value=WaitResult.FAILED), + resources_drained=Mock(return_value=False), + close=Mock(return_value=True), + ) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._ever_had_recv_session = True + transceiver._gen_need_sync = False + transceiver._mapping = SimpleNamespace( + pp_size=1, + enable_attention_dp=False, + world_size=1, + ) + transceiver._dist = SimpleNamespace(rank=0) + transceiver._gen_allgather = Mock() + transceiver._recv_sessions = {rid: session} + transceiver._recv_reqs = {rid: request} + + completed, failed, cancelled = transceiver.check_gen_transfer_status(None) + + assert (completed, failed, cancelled) == ([], [], []) + assert transceiver._recv_sessions[rid] is session + assert transceiver._recv_reqs[rid] is request + session.close.assert_not_called() + + session.resources_drained.return_value = True + completed, failed, cancelled = transceiver.check_gen_transfer_status(None) + + assert (completed, failed, cancelled) == ([], [rid], []) + assert rid not in transceiver._recv_sessions + assert rid not in transceiver._recv_reqs + session.close.assert_called_once_with() + + +@pytest.mark.cpu_only +def test_failed_receive_consensus_waits_for_every_rank_to_drain() -> None: + rid = 88 + request = SimpleNamespace(request_id=rid) + session = SimpleNamespace( + _enforce_physical_ownership=True, + status=SessionStatus.ERROR, + is_completed=Mock(return_value=False), + has_failed=Mock(return_value=True), + wait_complete=Mock(return_value=WaitResult.FAILED), + resources_drained=Mock(return_value=False), + close=Mock(return_value=True), + ) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._ever_had_recv_session = True + transceiver._gen_need_sync = True + transceiver._mapping = SimpleNamespace( + pp_size=1, + enable_attention_dp=False, + world_size=2, + ) + transceiver._dist = SimpleNamespace(rank=0) + transceiver._recv_sessions = {rid: session} + transceiver._recv_reqs = {rid: request} + transceiver._gen_allgather = Mock( + side_effect=[ + [[], []], + [ + [[], [rid], [], [rid]], + [[], [], [], []], + ], + ] + ) + + completed, failed, cancelled = transceiver.check_gen_transfer_status(None) + + assert (completed, failed, cancelled) == ([], [], []) + assert transceiver._recv_sessions[rid] is session + session.close.assert_not_called() + + session.resources_drained.return_value = True + transceiver._gen_allgather = Mock( + side_effect=[ + [[rid], [rid]], + [ + [[], [rid], [], [rid]], + [[], [rid], [], [rid]], + ], + ] + ) + + completed, failed, cancelled = transceiver.check_gen_transfer_status(None) + + assert (completed, failed, cancelled) == ([], [rid], []) + assert rid not in transceiver._recv_sessions + assert rid not in transceiver._recv_reqs + session.close.assert_called_once_with() + + +@pytest.mark.cpu_only +def test_non_terminal_writer_result_does_not_authorize_reuse() -> None: + receiver = _ReceiverProbe() + session = _make_rx_session(receiver, rid=80) + session.receive(KVSlice(is_last_slice=True)) + + session.process_kv_agent_result( + peer_rank=0, + receiver_slice_id=0, + is_last_slice=False, + status=AgentResult.SUCCESS, + ) + + assert session.status == SessionStatus.TRANSFERRING + assert not session.resources_drained() + + +@pytest.mark.cpu_only +def test_gen_first_no_retry_adp_count_seal_waits_for_one_writer_group() -> None: + rid = 94 + receiver = object.__new__(Receiver) + receiver._sessions_lock = threading.Lock() + receiver._sessions = {} + receiver._pre_cancelled_rids = set() + receiver._bounce = _BounceProbe() + receiver._enforce_physical_ownership = True + receiver._shutdown = True + receiver_req = SimpleNamespace( + unique_rid=rid, + slice_id=0, + bounce_dst_base=None, + to_bytes=Mock(return_value=b"receiver-request"), + ) + receiver._build_recv_req_info = Mock(return_value=receiver_req) + overlaps = { + 0: SimpleNamespace(ranks=[0, 1]), + 1: SimpleNamespace(ranks=[2, 3]), + } + receiver._registrar = SimpleNamespace( + get_peer_overlap=Mock(side_effect=lambda _peer, dp_rank: overlaps[dp_rank]), + self_extractor=SimpleNamespace(page_table=None), + self_rank_info=SimpleNamespace( + cp_size=1, + instance_name="gen", + instance_rank=0, + ), + ) + receiver._get_sender_info = Mock( + return_value=SimpleNamespace( + sender_endpoints={rank: f"tcp://sender-{rank}" for rank in range(4)}, + page_table=None, + dp_size=2, + cp_size=1, + ) + ) + receiver._request_sender_data = Mock() + session = RxSession( + request_id=rid, + params=DisaggregatedParams( + disagg_request_id=rid, + ctx_dp_rank=None, + schedule_style=DisaggScheduleStyle.GENERATION_FIRST, + ), + receiver=receiver, + ) + task = session.prepare_receive(KVSlice(is_last_slice=True)) + assert task is not None + + session.dispatch_prepared_receive(task) + + assert task.expected_transfers == 2 + assert {call.args[0] for call in receiver._request_sender_data.call_args_list} == { + f"tcp://sender-{rank}" for rank in range(4) + } + assert receiver._request_sender_data.call_count == 4 + session.process_kv_agent_result( + peer_rank=2, + receiver_slice_id=0, + is_last_slice=True, + status=AgentResult.FAILED, + ) + assert task.status == transfer_mod.TaskStatus.ERROR + assert not task.resources_drained + + session.process_kv_agent_result( + peer_rank=3, + receiver_slice_id=0, + is_last_slice=True, + status=AgentResult.SUCCESS, + ) + assert task.resources_drained + assert task.status == transfer_mod.TaskStatus.ERROR + assert session.close() is True + + +@pytest.mark.cpu_only +@pytest.mark.parametrize("writer_settles_before_failure", [False, True]) +def test_partial_bounced_publication_waits_for_queued_writer_success( + monkeypatch: pytest.MonkeyPatch, + writer_settles_before_failure: bool, +) -> None: + rid = 86 + queued_endpoints: list[str] = [] + bounce = _PartialFanInBounce() + receiver = object.__new__(Receiver) + receiver._sessions_lock = threading.Lock() + receiver._sessions = {} + receiver._pre_cancelled_rids = set() + receiver._bounce = bounce + receiver._enforce_physical_ownership = True + receiver._shutdown = True + receiver._dealers = {} + receiver_req = SimpleNamespace( + unique_rid=rid, + slice_id=0, + mamba_state_index=None, + bounce_dst_base=None, + to_bytes=Mock(side_effect=[b"writer-0", b"writer-1"]), + ) + receiver._build_recv_req_info = Mock(return_value=receiver_req) + overlap = SimpleNamespace( + ranks=[0, 1], + duplicate_head_factor=1, + overlap_pp_size=1, + ) + receiver._registrar = SimpleNamespace( + get_peer_overlap=Mock(return_value=overlap), + self_extractor=SimpleNamespace(page_table=None), + self_rank_info=SimpleNamespace(instance_name="gen", instance_rank=0, cp_size=1), + ) + receiver._get_sender_info = Mock( + return_value=SimpleNamespace( + sender_endpoints={0: "tcp://sender-0", 1: "tcp://sender-1"}, + page_table=None, + cp_size=1, + dp_size=1, + ) + ) + + def request_sender_data(endpoint: str, _payload: bytes) -> None: + if endpoint == "tcp://sender-1": + raise RuntimeError("writer 1 publication failed after writer 0") + queued_endpoints.append(endpoint) + if writer_settles_before_failure: + session.process_kv_agent_result( + peer_rank=0, + receiver_slice_id=0, + is_last_slice=True, + status=AgentResult.SUCCESS, + dst_ptrs=np.array([0x2000], dtype=np.int64), + sizes=np.array([0x100], dtype=np.int64), + src_base=0x1000, + ) + + receiver._request_sender_data = request_sender_data + monkeypatch.setattr( + transfer_mod.tensorrt_llm.bindings, + "global_steady_clock_now", + lambda: 0, + ) + session = RxSession( + request_id=rid, + params=DisaggregatedParams(disagg_request_id=rid, ctx_dp_rank=0), + receiver=receiver, + ) + + with pytest.raises(RuntimeError, match="writer 1 publication failed"): + session.receive(KVSlice(is_last_slice=True)) + + assert queued_endpoints == ["tcp://sender-0"] + if not writer_settles_before_failure: + assert bounce.context is not None + assert bounce.release_count == 0 + assert not session.resources_drained() + + # Only writer 0's REQUEST_DATA was successfully queued. The destination + # remains owned until that writer reports terminal physical evidence. + session.process_kv_agent_result( + peer_rank=0, + receiver_slice_id=0, + is_last_slice=True, + status=AgentResult.SUCCESS, + dst_ptrs=np.array([0x2000], dtype=np.int64), + sizes=np.array([0x100], dtype=np.int64), + src_base=0x1000, + ) + + assert session.status == SessionStatus.ERROR + assert session.resources_drained() + assert bounce.context is None + assert bounce.scatter_count == 0 + assert bounce.release_count == 1 + assert session.close() is True + assert rid not in receiver._sessions + + +@pytest.mark.cpu_only +def test_aborted_publication_cannot_complete_during_failure_unwind() -> None: + published_writers = {0} + + class _FailureWindowReceiver(_ReceiverProbe): + def dispatch_task(self, task: KVRecvTask) -> None: + assert self._session is not None + task.expected_transfers = 1 + original_abort = task.abort_publication + + def abort_then_deliver_success(writers: set[int]) -> None: + original_abort(writers) + self._session.process_kv_agent_result( + peer_rank=0, + receiver_slice_id=0, + is_last_slice=True, + status=AgentResult.SUCCESS, + ) + assert task.status == transfer_mod.TaskStatus.TRANSFERRING + + task.abort_publication = abort_then_deliver_success + + def fail_after_publication() -> None: + raise RuntimeError("publication failed after writer 0 was queued") + + self._session.try_begin_transfer( + task.slice_id, + {"tcp://sender-0"}, + writer_cohort={0}, + publish=fail_after_publication, + published_writers=published_writers, + ) + + receiver = _FailureWindowReceiver() + session = _make_rx_session(receiver, rid=91) + + with pytest.raises(RuntimeError, match="publication failed"): + session.receive(KVSlice(is_last_slice=True)) + + assert session.status == SessionStatus.ERROR + assert session.resources_drained() + assert session.close() is True + + +@pytest.mark.cpu_only +def test_out_of_cohort_writer_cannot_authorize_reuse() -> None: + receiver = _ReceiverProbe() + session = _make_rx_session(receiver, rid=82) + session.receive(KVSlice(is_last_slice=True)) + + with pytest.raises(RuntimeError, match="outside the sealed cohort"): + session.process_kv_agent_result( + peer_rank=2, + receiver_slice_id=0, + is_last_slice=True, + status=AgentResult.SUCCESS, + ) + + assert not session.resources_drained() + # Invalid evidence intentionally fails closed. Avoid a noisy best-effort + # destructor close for this hand-built, resource-free unit-test fixture. + session._closed = True + receiver._session = None + + +@pytest.mark.cpu_only +def test_cancel_after_publication_cannot_overtake_request_data( + monkeypatch: pytest.MonkeyPatch, +) -> None: + rid = 79 + request_data_started = threading.Event() + finish_request_data = threading.Event() + race_outcomes: queue.Queue[str] = queue.Queue() + protocol_order: list[str] = [] + initial_cancel_outcome: str | None = None + thread_results: queue.Queue[Exception | None] = queue.Queue() + + receiver = object.__new__(Receiver) + receiver._sessions_lock = threading.Lock() + receiver._sessions = {} + receiver._pre_cancelled_rids = set() + receiver._bounce = _BounceProbe() + receiver._enforce_physical_ownership = True + receiver._shutdown = True + receiver._dealers = {} + receiver._build_recv_req_info = Mock( + return_value=SimpleNamespace( + unique_rid=rid, + slice_id=0, + mamba_state_index=None, + bounce_dst_base=None, + to_bytes=Mock(return_value=b"receiver-request"), + ) + ) + overlap = SimpleNamespace(ranks=[0]) + receiver._registrar = SimpleNamespace( + get_peer_overlap=Mock(return_value=overlap), + self_extractor=SimpleNamespace(page_table=None), + self_rank_info=SimpleNamespace(cp_size=1), + ) + receiver._get_sender_info = Mock( + return_value=SimpleNamespace( + sender_endpoints={0: "tcp://sender-0"}, + page_table=None, + tp_size=1, + pp_size=1, + cp_size=1, + dp_size=1, + attention=None, + ) + ) + + def request_sender_data(_endpoint: str, _receiver_info_bytes: bytes) -> None: + protocol_order.append("request_data_started") + request_data_started.set() + assert finish_request_data.wait(timeout=10) + protocol_order.append("request_data_sent") + + def send_cancel_to_senders(_unique_rid: int, _sender_endpoints: set[str]) -> None: + protocol_order.append("cancel_sent") + race_outcomes.put("cancel_sent") + + receiver._request_sender_data = request_sender_data + receiver.send_cancel_to_senders = send_cancel_to_senders + + monkeypatch.setattr( + transfer_mod.tensorrt_llm.bindings, + "global_steady_clock_now", + lambda: 0, + ) + + session = RxSession( + request_id=rid, + params=DisaggregatedParams(disagg_request_id=rid, ctx_dp_rank=0), + receiver=receiver, + ) + publication_lock = _TrackingLock(race_outcomes) + session._publication_lock = publication_lock + + def receive() -> None: + session.receive(KVSlice(is_last_slice=True)) + + def cancel() -> None: + publication_lock.tracked_thread_id = threading.get_ident() + session.cancel() + + def finish_writer() -> None: + session.process_kv_agent_result( + peer_rank=0, + receiver_slice_id=0, + is_last_slice=True, + status=AgentResult.FAILED, + ) + + receive_thread = _start_checked_thread(receive, thread_results) + cancel_thread = None + result_thread = None + try: + assert request_data_started.wait(timeout=10) + cancel_thread = _start_checked_thread(cancel, thread_results) + initial_cancel_outcome = race_outcomes.get(timeout=10) + result_thread = _start_checked_thread(finish_writer, thread_results) + result_thread.join(timeout=10) + assert not result_thread.is_alive() + assert not session.resources_drained(), ( + "terminal writer evidence authorized reuse before REQUEST_DATA publication finished" + ) + finally: + finish_request_data.set() + receive_thread.join(timeout=10) + if cancel_thread is not None: + cancel_thread.join(timeout=10) + if result_thread is not None: + result_thread.join(timeout=10) + + assert not receive_thread.is_alive() + assert cancel_thread is not None and not cancel_thread.is_alive() + _raise_thread_errors(thread_results, expected=3) + assert initial_cancel_outcome == "blocked" + assert session.resources_drained() + assert protocol_order == ["request_data_started", "request_data_sent", "cancel_sent"], ( + "cancellation overtook an already-authorized REQUEST_DATA publication" + ) + + +@pytest.mark.cpu_only +def test_cancel_before_dispatch_releases_late_idle_reservation() -> None: + rid = 81 + receiver = _ReceiverProbe() + bounce = _LateReservationBounce() + receiver._bounce = bounce + session = _make_rx_session(receiver, rid) + task = session.prepare_receive(KVSlice(is_last_slice=True)) + assert task is not None + + assert session.cancel_local() + + bounce.active_reservations.add((rid, task.slice_id)) + with pytest.raises(RuntimeError, match="became terminal before publication"): + session.dispatch_prepared_receive(task) + + assert bounce.active_reservations == set(), ( + "cancel-before-dispatch left a reservation allocated after cancellation cleanup" + ) + assert session.status == SessionStatus.CANCELLED + assert task.resources_drained + + +@pytest.mark.cpu_only +def test_cancel_request_retains_session_when_close_refuses() -> None: + rid = 83 + request = SimpleNamespace( + request_id=rid, + py_disaggregated_params=DisaggregatedParams(disagg_request_id=rid), + ) + session = SimpleNamespace( + cancel=Mock(), + has_transferring_tasks=Mock(return_value=False), + close=Mock(return_value=False), + ) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._wait_reqs = {} + transceiver._send_sessions = {} + transceiver._send_reqs = {} + transceiver._recv_sessions = {rid: session} + transceiver._recv_reqs = {rid: request} + + assert transceiver.cancel_request(request) is False + assert transceiver._recv_sessions[rid] is session + assert transceiver._recv_reqs[rid] is request + session.close.assert_called_once_with() + + +@pytest.mark.cpu_only +def test_failed_session_is_not_reported_retired_when_close_refuses() -> None: + rid = 89 + initial_state = object() + request = SimpleNamespace(state=initial_state) + session = SimpleNamespace( + _enforce_physical_ownership=True, + resources_drained=Mock(return_value=True), + close=Mock(return_value=False), + ) + sessions = {rid: session} + requests = {rid: request} + failed = [rid] + transceiver = object.__new__(KvCacheTransceiverV2) + + with pytest.raises(RuntimeError, match="session close refused"): + transceiver._close_failed_sessions(sessions, requests, failed) + + assert failed == [rid] + assert request.state is initial_state + assert sessions == {rid: session} + assert requests == {rid: request} + + session.close.return_value = True + failed = [rid] + transceiver._close_failed_sessions(sessions, requests, failed) + + assert failed == [rid] + assert request.state is not initial_state + assert sessions == {} + assert requests == {} + + +@pytest.mark.cpu_only +def test_completed_session_is_not_reported_retired_when_close_refuses() -> None: + rid = 90 + initial_state = object() + request = SimpleNamespace( + state=initial_state, + py_kv_cache_xfer_bytes=0, + set_kv_cache_size=Mock(), + ) + session = SimpleNamespace( + status=SessionStatus.KV_TRANSFERRED, + transfer_end_time=None, + kv_cache_size_bytes=0, + is_completed=Mock(return_value=True), + has_failed=Mock(return_value=False), + wait_complete=Mock(return_value=WaitResult.COMPLETED), + close=Mock(return_value=False), + ) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._ever_had_recv_session = True + transceiver._gen_need_sync = False + transceiver._mapping = SimpleNamespace(pp_size=1, enable_attention_dp=False, world_size=1) + transceiver._recv_sessions = {rid: session} + transceiver._recv_reqs = {rid: request} + transceiver._gen_allgather = Mock() + transceiver._gen_consensus = Mock(side_effect=lambda rids: rids) + transceiver._need_aux_transfer = Mock(return_value=False) + transceiver._assert_disagg_history_declared = Mock() + + with pytest.raises(RuntimeError, match="session close refused"): + transceiver.check_gen_transfer_status(None) + + assert request.state is initial_state + assert transceiver._recv_sessions == {rid: session} + assert transceiver._recv_reqs == {rid: request} + + session.close.return_value = True + completed, failed, cancelled = transceiver.check_gen_transfer_status(None) + + assert (completed, failed, cancelled) == ([rid], [], []) + assert request.state is not initial_state + assert transceiver._recv_sessions == {} + assert transceiver._recv_reqs == {} + + +@pytest.mark.cpu_only +def test_cancelled_session_close_refusal_fails_stop_after_consensus() -> None: + rid = 92 + request = SimpleNamespace(state=object()) + session = SimpleNamespace( + status=SessionStatus.CANCELLED, + is_completed=Mock(return_value=False), + has_failed=Mock(return_value=True), + wait_complete=Mock(return_value=WaitResult.FAILED), + close=Mock(return_value=False), + ) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._ever_had_recv_session = True + transceiver._gen_need_sync = False + transceiver._recv_sessions = {rid: session} + transceiver._recv_reqs = {rid: request} + transceiver._gen_allgather = Mock() + transceiver._gen_consensus = Mock(side_effect=lambda rids: rids) + + with pytest.raises(RuntimeError, match="session close refused"): + transceiver.check_gen_transfer_status(None) + + assert transceiver._recv_sessions == {rid: session} + assert transceiver._recv_reqs == {rid: request} + + session.close.return_value = True + completed, failed, cancelled = transceiver.check_gen_transfer_status(None) + + assert completed == [] + assert failed == [] + assert cancelled == [request] + assert transceiver._recv_sessions == {} + assert transceiver._recv_reqs == {} + + +@pytest.mark.cpu_only +def test_collect_done_waits_for_physical_drain() -> None: + rid = 84 + drained = False + session = SimpleNamespace( + _enforce_physical_ownership=True, + is_completed=Mock(return_value=False), + has_failed=Mock(return_value=True), + resources_drained=lambda: drained, + ) + transceiver = object.__new__(KvCacheTransceiverV2) + + assert transceiver._collect_done({rid: session}, {rid: object()}) == ([], []) + + drained = True + assert transceiver._collect_done({rid: session}, {rid: object()}) == ([], [rid]) + + +@pytest.mark.cpu_only +def test_shutdown_refuses_to_drop_active_receive_owner() -> None: + rid = 85 + session = SimpleNamespace( + _enforce_physical_ownership=True, + resources_drained=Mock(return_value=False), + close=Mock(return_value=True), + ) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._shutdown_complete = False + transceiver._wait_reqs = {} + transceiver._send_sessions = {} + transceiver._send_reqs = {} + transceiver._recv_sessions = {rid: session} + transceiver._recv_reqs = {rid: object()} + transceiver._transfer_worker = SimpleNamespace(shutdown=Mock()) + + with pytest.raises(RuntimeError, match="physical resources remain active"): + transceiver.shutdown() + + assert not transceiver._shutdown_complete + assert transceiver._recv_sessions[rid] is session + session.close.assert_not_called() + transceiver._transfer_worker.shutdown.assert_not_called() + + session.resources_drained.return_value = True + transceiver.shutdown() + + assert transceiver._shutdown_complete + assert transceiver._recv_sessions == {} + session.close.assert_called_once_with() + transceiver._transfer_worker.shutdown.assert_called_once_with() + + +@pytest.mark.cpu_only +def test_shutdown_fails_stop_when_receive_close_refuses_after_preflight() -> None: + rid = 93 + send_session = SimpleNamespace(close=Mock(return_value=None)) + recv_session = SimpleNamespace( + _enforce_physical_ownership=True, + resources_drained=Mock(return_value=True), + close=Mock(return_value=False), + ) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._shutdown_complete = False + transceiver._send_sessions = {rid + 1: send_session} + transceiver._send_reqs = {rid + 1: object()} + transceiver._recv_sessions = {rid: recv_session} + transceiver._recv_reqs = {rid: object()} + transceiver._transfer_worker = SimpleNamespace(shutdown=Mock()) + + with pytest.raises(RuntimeError, match="session close refused"): + transceiver.shutdown() + + assert not transceiver._shutdown_complete + assert transceiver._recv_sessions == {rid: recv_session} + assert transceiver._send_sessions == {rid + 1: send_session} + send_session.close.assert_not_called() + transceiver._transfer_worker.shutdown.assert_not_called() + + recv_session.close.return_value = True + transceiver.shutdown() + + assert transceiver._shutdown_complete + assert transceiver._recv_sessions == {} + assert transceiver._send_sessions == {} + send_session.close.assert_called_once_with() + transceiver._transfer_worker.shutdown.assert_called_once_with()