diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index 21632f9358e2..8129b84a8aac 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -90,6 +90,10 @@ # an overall deadline. KvCacheTransceiverV2 requires a configured transfer # timeout before it creates either sender or receiver sessions. _FALLBACK_TX_OVERALL_TIMEOUT_S = 60.0 +# Inactivity budget for RxSession.has_transferring_tasks(): a peer resolves only +# when its terminal result arrives, so a dead worker or a lost result would pin +# KV pages forever. Rearmed by every dispatch and result. +_PEER_DRAIN_TIMEOUT_S = float(os.environ.get("TRTLLM_KV_TRANSFER_PEER_DRAIN_TIMEOUT_S", "300")) @dataclass @@ -234,6 +238,9 @@ def __init__(self, params: DisaggregatedParams): self._event = threading.Event() self._exception: Optional[Exception] = None self.lock = threading.Lock() + # Peers with a queued-or-active write reading this task's source + # addresses. set add/discard is atomic under the GIL, so no lock. + self.pending_peers: set[int] = set() self._params = params self._unique_rid: Optional[int] = params.disagg_request_id self._perf_timer = PerfTimer() if perf_log_manager.enabled else None @@ -313,7 +320,8 @@ def __init__( self._peer_requests_timestamps: dict[int, float] = {} # unique_rid -> insert time self._peer_requests_lock = threading.Lock() self._messenger = ZMQMessenger(mode="ROUTER") - self._dealers = {} # used by listener thread only (single-threaded path) + self._dealers = {} # guarded by _dealers_lock; see _send_via_dealer + self._dealers_lock = threading.Lock() self._thread_local = threading.local() # per-thread DEALER cache for worker threads self._sessions = {} # unique_rid -> TxSession self._sessions_lock = threading.Lock() # Protects _sessions and _pre_cancelled_rids @@ -435,7 +443,15 @@ def _enqueue(self, write_meta: WriteMeta): # Route by (unique_rid, peer_rank) so that: # - Same peer's slices stay ordered on one thread (is_last_slice correctness) # - Different peers can run on different threads (better load balancing) + if self._shutdown: + # The workers may already have consumed their sentinel, so nothing + # would ever release the claim below and this side has no deadline. + raise RuntimeError(f"Sender is shutting down; refusing rid={write_meta.unique_rid}") thread_idx = hash((write_meta.unique_rid, write_meta.peer_rank)) % self._num_threads + # Claim the source-address lifetime before the item is visible to a + # worker; the worker's finally releases it. Unconditional so the two + # sides stay symmetric even for an already-cancelled session. + write_meta.task.pending_peers.add(write_meta.peer_rank) self._send_task_queues[thread_idx].put(write_meta) def _get_or_connect_thread_dealer(self, endpoint: Optional[str]) -> ZMQMessenger: @@ -477,6 +493,8 @@ def _process_task_queue(self, thread_idx: int): f"unique_rid={write_meta.unique_rid}: {e}" ) write_meta.task.fail(e) + finally: + write_meta.task.pending_peers.discard(write_meta.peer_rank) finally: # Clean up this thread's DEALER sockets. threading.local storage # is only accessible from the owning thread, so shutdown must @@ -573,7 +591,7 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): task.fail( RuntimeError(f"session {write_meta.unique_rid} {status.value}, transfer aborted") ) - self._get_or_connect_dealer(write_meta.peer_endpoint).send( + self._get_or_connect_thread_dealer(write_meta.peer_endpoint).send( _make_kv_result_msg( self._instance_rank, write_meta.unique_rid, @@ -603,7 +621,7 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): f"{write_meta.unique_rid} slice={write_meta.slice_id}: {e}" ) task.fail(RuntimeError(f"build_send_request failed: {e}")) - self._get_or_connect_dealer(write_meta.peer_endpoint).send( + self._get_or_connect_thread_dealer(write_meta.peer_endpoint).send( _make_kv_result_msg( self._instance_rank, write_meta.unique_rid, @@ -1128,20 +1146,32 @@ def _send_failed_result_to_receiver(self, info: RecvReqInfo): try: peer_ri = self._registrar.get_peer_rank_info(info.instance_name, info.instance_rank) slice_id = info.slice_id if info.slice_id is not None else 0 - self._get_or_connect_dealer(peer_ri.self_endpoint).send( + self._send_via_dealer( + peer_ri.self_endpoint, _make_kv_result_msg( self._instance_rank, info.unique_rid, slice_id, True, # is_last_slice AgentResult.FAILED, - ) + ), ) except Exception as e: logger.warning( f"_respond_with_kv: failed to abort receiver for rid={info.unique_rid}: {e}" ) + def _send_via_dealer(self, endpoint: Optional[str], message: list[bytes]) -> None: + """Serialize dealer creation and send; ZMQ sockets are not thread-safe. + + ``_dealers`` is reached from both the listener thread and the executor + thread (cancel paths), so an unguarded send can tear a multipart frame. + """ + with self._dealers_lock: + if self._shutdown: + return + self._get_or_connect_dealer(endpoint).send(message) + def _get_or_connect_dealer(self, endpoint: Optional[str]): if endpoint is None: raise ValueError("Sender: peer endpoint is None; peer may not have registered yet") @@ -1188,8 +1218,9 @@ def send_cancel_to_receivers(self, unique_rid: int) -> None: peer_ri = self._registrar.get_peer_rank_info( req_info.instance_name, req_info.instance_rank ) - self._get_or_connect_dealer(peer_ri.self_endpoint).send( - [MessageType.CANCEL_SESSION, str(unique_rid).encode("ascii")] + self._send_via_dealer( + peer_ri.self_endpoint, + [MessageType.CANCEL_SESSION, str(unique_rid).encode("ascii")], ) except Exception as e: logger.warning(f"send_cancel_to_receivers: failed for rid={unique_rid}: {e}") @@ -1219,12 +1250,13 @@ def shutdown(self): logger.warning( f"Failed to invalidate remote agent '{agent_name}' during shutdown: {e}" ) - for dealer in self._dealers.values(): + with self._dealers_lock: + dealers, self._dealers = list(self._dealers.values()), {} + for dealer in dealers: try: dealer.stop() except Exception as e: logger.warning(f"Failed to stop dealer during Sender shutdown: {e}") - self._dealers.clear() def __del__(self): try: @@ -1380,11 +1412,12 @@ def cancel(self) -> None: self._sender.send_cancel_to_receivers(self.disagg_request_id) def has_transferring_tasks(self) -> bool: - """True if any KV task is currently mid-write (TRANSFERRING). + """True while a queued or active write still reads this session's buffers. - cancel_request() must return False while this is True. + Per-task TaskStatus cannot express it: one peer failing marks the task + terminal while another peer still owns the registered source addresses. """ - return any(t.status == TaskStatus.TRANSFERRING for t in self.kv_tasks) + return any(t is not None and t.pending_peers for t in (*self.kv_tasks, self.aux_task)) def wait_complete(self, blocking: bool = True) -> Optional[WaitResult]: """Poll or block until KV (and optionally aux) transfer finishes. @@ -1546,6 +1579,12 @@ def __init__( self.status = TaskStatus.INIT self.expected_transfers = 0 self.last_slice_count = 0 + self.dispatched = False + # Peers that returned a terminal result, success or failure. Counting + # responders rather than dispatches is what keeps ADP broadcast honest: + # the broadcast reaches every DP group, but only ``expected_transfers`` + # of them own the request and reply. + self.responded_peer_ranks: set[int] = set() self._unique_rid = unique_rid self._kv_slice = kv_slice @@ -1594,7 +1633,8 @@ def __init__( self._registrar = peer_registrar self._agent = agent self._bounce = bounce - self._dealers = {} + self._dealers = {} # guarded by _dealers_lock; see _send_via_dealer + self._dealers_lock = threading.Lock() self._sender_ep_instance_map = {} # info_endpoint -> diagnostic message for peers that failed the # compatibility check. Requests targeting such a peer fail fast @@ -1619,12 +1659,13 @@ def shutdown(self): if getattr(self, "_shutdown", False): return self._shutdown = True - for dealer in self._dealers.values(): + with self._dealers_lock: + dealers, self._dealers = list(self._dealers.values()), {} + for dealer in dealers: try: dealer.stop() except Exception as e: logger.warning(f"Failed to stop dealer during Receiver shutdown: {e}") - self._dealers.clear() self._messenger.stop() def clear_session(self, unique_rid: int): @@ -1811,22 +1852,35 @@ 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 - ) # 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): + endpoint = peer_infos.sender_endpoints[rank] + # A partial dispatch failure below leaves fewer responders than + # expected_transfers; the drain deadline is what resolves that. + if not session.mark_peer_dispatched(task.slice_id, rank, endpoint): + exc = RuntimeError( + f"dispatch_task: RxSession {task._unique_rid} was cancelled " + f"before dispatch to peer_rank={rank}" + ) + task.fail(exc) + raise exc 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) + try: + self._request_sender_data(endpoint, receiver_req_bytes) + except Exception as exc: + task.fail(exc) + logger.error( + f"Receiver.dispatch_task failed for rid={task._unique_rid}, " + f"peer_rank={rank}: {exc}" + ) + raise return @staticmethod @@ -1842,6 +1896,17 @@ def _should_register_peer(self, params: DisaggregatedParams) -> bool: endpoint = self._extract_info_endpoint(params) return endpoint not in self._sender_ep_instance_map + def _send_via_dealer(self, endpoint: Optional[str], message: list[bytes]) -> None: + """Serialize dealer creation and send; ZMQ sockets are not thread-safe. + + ``_dealers`` is reached from both the listener thread and the executor + thread (cancel paths), so an unguarded send can tear a multipart frame. + """ + with self._dealers_lock: + if self._shutdown: + return + self._get_or_connect_dealer(endpoint).send(message) + def _get_or_connect_dealer(self, endpoint: Optional[str]): if endpoint is None: raise ValueError("Receiver: peer endpoint is None; peer may not have registered yet") @@ -1891,10 +1956,11 @@ def _get_sender_info(self, params: DisaggregatedParams) -> RankInfo: self._incompatible_peers[info_endpoint] = msg raise PeerIncompatibleError(msg) from e + rank_info = self._registrar.self_rank_info for endpoint in sender_info.sender_endpoints: - dealer = self._get_or_connect_dealer(endpoint) - rank_info = self._registrar.self_rank_info - dealer.send([MessageType.REGISTER_RANK_INFO, rank_info.to_bytes()]) + self._send_via_dealer( + endpoint, [MessageType.REGISTER_RANK_INFO, rank_info.to_bytes()] + ) self._sender_ep_instance_map[info_endpoint] = sender_info return sender_info @@ -1906,8 +1972,8 @@ def send_cancel_to_senders(self, unique_rid: int, sender_endpoints: set[str]) -> """Notify all senders involved in this session to cancel.""" for endpoint in sender_endpoints: try: - self._get_or_connect_dealer(endpoint).send( - [MessageType.CANCEL_SESSION, str(unique_rid).encode("ascii")] + self._send_via_dealer( + endpoint, [MessageType.CANCEL_SESSION, str(unique_rid).encode("ascii")] ) except Exception as e: logger.warning(f"send_cancel_to_senders: failed for rid={unique_rid}: {e}") @@ -1998,8 +2064,7 @@ def _process_aux_agent_result(self, _send_id: bytes, message: list[bytes]): def _request_sender_data(self, endpoint: str, receiver_info_bytes: bytes): # receiver_info serialized once and reused for every peer rank (block-table msgpack isn't free at fan-out). logger.debug("Sending data request to endpoint '%s'", endpoint) - messenger = self._get_or_connect_dealer(endpoint) - messenger.send([MessageType.REQUEST_DATA, receiver_info_bytes]) + self._send_via_dealer(endpoint, [MessageType.REQUEST_DATA, receiver_info_bytes]) def __del__(self): try: @@ -2044,8 +2109,15 @@ def __init__( self._kv_tasks: list[KVRecvTask] = [] self._aux_count = 0 self._aux_status: TaskStatus = TaskStatus.INIT + self._aux_responded_peer_ranks: set[int] = set() self._sender_endpoints: set[str] = set() + self._cancel_notified = False self.lock = threading.Lock() + # Monotonic timestamp of the last peer-lifetime event (dispatch or + # terminal result). Drives the drain deadline in + # has_transferring_tasks(); None means nothing was ever dispatched. + self._last_peer_progress: Optional[float] = None + self._drain_timeout_logged = False self._receiver.setup_session(self) @property @@ -2076,9 +2148,37 @@ def status(self) -> SessionStatus: return SessionStatus.TRANSFERRING return SessionStatus.INIT - def mark_transferring(self, slice_id: int): + def mark_peer_dispatched(self, slice_id: int, peer_rank: int, sender_endpoint: str) -> bool: + """Register a peer before contacting it; False if cancellation won the race.""" with self.lock: - self._kv_tasks[slice_id].status = TaskStatus.TRANSFERRING + if self._terminal_status in (SessionStatus.ERROR, SessionStatus.CANCELLED): + return False + task = self._kv_tasks[slice_id] + task.dispatched = True + task.status = TaskStatus.TRANSFERRING + # Cache the endpoint so cancel() can send CANCEL_SESSION to it. + self._sender_endpoints.add(sender_endpoint) + self._last_peer_progress = time.monotonic() + return True + + def _note_peer_progress_locked(self) -> None: + self._last_peer_progress = time.monotonic() + self._drain_timeout_logged = False + + def _has_unresolved_peers_locked(self) -> bool: + for task in self._kv_tasks: + if not task.dispatched: + continue + # Two distinct windows: a peer that has not replied yet, and (bounce + # path only) replies all in but the scatter into the KV pages still + # queued -- complete() runs in the scatter's on_done, not here. + if len(task.responded_peer_ranks) < task.expected_transfers: + return True + if task.status == TaskStatus.TRANSFERRING: + return True + if self._need_aux and len(self._aux_responded_peer_ranks) < task.expected_transfers: + return True + return False def receive(self, slice: KVSlice) -> None: if self.transfer_start_time is None: @@ -2114,6 +2214,16 @@ def process_kv_agent_result( f"Sender/receiver slice count mismatch." ) task = self._kv_tasks[sender_slice_id] + if peer_rank in task.responded_peer_ranks: + logger.warning( + f"RxSession {self.request_id} ignoring duplicate KV result for " + f"slice={sender_slice_id}, peer_rank={peer_rank}" + ) + return + # The result means the remote submit/wait released our destination. + # Resolve before processing so a raise cannot strand the peer. + task.responded_peer_ranks.add(peer_rank) + self._note_peer_progress_locked() if status == AgentResult.SUCCESS: from .bounce import scatter_write_result @@ -2203,10 +2313,18 @@ def on_done( f"Session {self.request_id} received unknown task status: {status.value}" ) - def process_aux_agent_result(self, _peer_rank: int, status: AgentResult): + def process_aux_agent_result(self, peer_rank: int, status: AgentResult): # Aux is session-level (not per-slice); expected_transfers is identical # across all kv_tasks, so any task provides the right count. with self.lock: + if peer_rank in self._aux_responded_peer_ranks: + logger.warning( + f"RxSession {self.request_id} ignoring duplicate aux result " + f"from peer_rank={peer_rank}" + ) + return + self._aux_responded_peer_ranks.add(peer_rank) + self._note_peer_progress_locked() if not self._kv_tasks: logger.warning( f"Aux result received before any KV tasks for request {self.request_id}" @@ -2278,8 +2396,9 @@ def cancel(self) -> None: The lock serializes with process_kv_agent_result() / process_aux_agent_result(). """ with self.lock: - if self._terminal_status == SessionStatus.CANCELLED: + if self._cancel_notified: return + self._cancel_notified = True self._terminal_status = SessionStatus.CANCELLED exc = RuntimeError(f"RxSession {self.disagg_request_id} cancelled") for task in self._kv_tasks: @@ -2293,15 +2412,31 @@ def cancel(self) -> None: # 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) + endpoints = set(self._sender_endpoints) # Send outside the lock to avoid holding it during I/O. - self._receiver.send_cancel_to_senders(self.disagg_request_id, self._sender_endpoints) + self._receiver.send_cancel_to_senders(self.disagg_request_id, endpoints) def has_transferring_tasks(self) -> bool: - """True if any KV task is currently mid-write (TRANSFERRING). + """True while a dispatched peer may still be writing into our buffers. - cancel_request() must return False while this is True. + Waits on responders bounded by expected_transfers, since under + attention-DP the broadcast reaches more peers than ever reply. """ - return any(t.status == TaskStatus.TRANSFERRING for t in self._kv_tasks) + with self.lock: + if not self._has_unresolved_peers_locked(): + return False + last = self._last_peer_progress + if last is not None and time.monotonic() - last > _PEER_DRAIN_TIMEOUT_S: + if not self._drain_timeout_logged: + self._drain_timeout_logged = True + logger.error( + f"RxSession {self.disagg_request_id}: no peer progress for " + f"{_PEER_DRAIN_TIMEOUT_S:g}s with peers still unresolved; " + "releasing destination buffers. Set " + "TRTLLM_KV_TRANSFER_PEER_DRAIN_TIMEOUT_S to tune." + ) + return False + return True def wait_complete(self, blocking: bool = False) -> Optional[WaitResult]: """Poll or block until transfer completes. diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index be3acc832558..a6754700278d 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -499,49 +499,48 @@ def _intersection(all_lists: List[List[int]], n_ranks: int) -> set: return {rid for rid, c in cnt.items() if c == n_ranks} def _consensus_outcome( - self, to_process, cancelled, failed, completed, allgather: Callable, need_sync: bool + self, to_process, cancelled, failed, completed, busy, allgather: Callable, need_sync: bool ): # CANCELLED/FAILED on any rank → global; COMPLETED only when ALL ranks agree. - # Batch the three id lists into one allgather to cut the per-step collective count. - if not need_sync: - all_c, all_f, all_done = [list(cancelled)], [list(failed)], [list(completed)] - else: - packed = list(allgather([list(cancelled), list(failed), list(completed)])) - all_c = [p[0] for p in packed] - all_f = [p[1] for p in packed] - all_done = [p[2] for p in packed] - n = len(all_c) - global_cancelled = self._union(all_c) - global_failed = self._union(all_f) - global_completed = self._intersection(all_done, n) - new_cancelled = [rid for rid in to_process if rid in global_cancelled] + # Batch the id lists into one allgather to cut the per-step collective count. + local_outcome = [list(cancelled), list(failed), list(completed), list(busy)] + packed = list(allgather(local_outcome)) if need_sync else [local_outcome] + n = len(packed) + global_cancelled = self._union([p[0] for p in packed]) + global_failed = self._union([p[1] for p in packed]) + global_completed = self._intersection([p[2] for p in packed], n) + # Busy on ANY rank means no rank may act: closing the session here would + # unregister buffers that rank is still writing through. + global_busy = self._union([p[3] for p in packed]) + ready = [rid for rid in to_process if rid not in global_busy] + new_cancelled = [rid for rid in ready if rid in global_cancelled] cancel_set = set(new_cancelled) - new_failed = [rid for rid in to_process if rid in global_failed and rid not in cancel_set] + new_failed = [rid for rid in ready if rid in global_failed and rid not in cancel_set] terminal = cancel_set | set(new_failed) - new_completed = [ - rid for rid in to_process if rid in global_completed and rid not in terminal - ] - return new_cancelled, new_failed, new_completed + new_completed = [rid for rid in ready if rid in global_completed and rid not in terminal] + new_busy = [rid for rid in to_process if rid in global_busy] + return new_cancelled, new_failed, new_completed, new_busy - def _gen_consensus_outcome(self, to_process, cancelled, failed, completed): + def _gen_consensus_outcome(self, to_process, cancelled, failed, completed, busy): return self._consensus_outcome( - to_process, cancelled, failed, completed, self._gen_allgather, self._gen_need_sync - ) + to_process, cancelled, failed, completed, busy, self._gen_allgather, self._gen_need_sync + )[:3] - def _ctx_consensus_outcome(self, to_process, cancelled, failed, completed): + def _ctx_consensus_outcome(self, to_process, cancelled, failed, completed, busy): # TP first, then PP. A local timeout remains nonterminal, so it is # represented by the absence of that request from completed. - c, f, d = self._consensus_outcome( + c, f, d, busy = self._consensus_outcome( to_process, cancelled, failed, completed, + busy, self._dist.tp_allgather, self._ctx_need_tp_sync, ) if self._ctx_need_pp_sync: pp_allgather: Callable = getattr(self._dist, "pp_allgather") - c, f, d = self._consensus_outcome(to_process, c, f, d, pp_allgather, True) + c, f, d, _ = self._consensus_outcome(to_process, c, f, d, busy, pp_allgather, True) return c, f, d def _sync_transfer_timing(self, reqs: list): @@ -601,10 +600,16 @@ def _collect_done(self, sessions: dict, reqs: dict): for rid, session in sessions.items(): if session.is_completed(): completed.append(rid) - elif session.has_failed(): + # A terminal status says nothing about whether a peer stopped writing. + elif session.has_failed() and not session.has_transferring_tasks(): failed.append(rid) return completed, failed + @staticmethod + def _busy_rids(sessions: dict, to_process: list) -> list: + """Rids a peer may still be writing into; no rank may close these yet.""" + return [rid for rid in to_process if sessions[rid].has_transferring_tasks()] + def _build_to_process( self, sessions: dict, consensus: list, wait_num: int, block_all: bool ) -> list: @@ -621,6 +626,10 @@ def _build_to_process( def _close_failed_sessions(self, sessions: dict, reqs: dict, failed: list): for rid in failed: reqs[rid].state = LlmRequestState.DISAGG_TRANS_ERROR + # Tell the peers to stop before unregistering: a drain deadline can + # expire before a backlogged sender has even started writing, and + # close() alone never notifies them. cancel() is idempotent. + sessions[rid].cancel() sessions[rid].close() del reqs[rid] del sessions[rid] @@ -674,6 +683,9 @@ def respond_and_send_async(self, req: LlmRequest): req.set_kv_cache_transfer_start(tensorrt_llm.bindings.global_steady_clock_now()) session = self._get_or_create_send_session(req) req.state = LlmRequestState.DISAGG_CONTEXT_TRANS_IN_PROGRESS + # See request_and_receive_async: _finalize_send assigns this too, but + # send() can raise once work is already queued for an earlier peer. + self._send_reqs[get_unique_rid(req)] = req session.send(self._create_kv_slice(req)) self._finalize_send(req, session) @@ -685,8 +697,11 @@ def request_and_receive_sync(self, req: LlmRequest): f"request_and_receive_sync: rid={rid} already has a recv session, skipping" ) return + self._ever_had_recv_session = True req.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS session = None + completed = False + raised = False try: session = self._transfer_worker.create_rx_session(req) self._recv_sessions[rid] = session @@ -702,16 +717,30 @@ def request_and_receive_sync(self, req: LlmRequest): self._apply_aux(session, req) self._assert_disagg_history_declared(req) req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE + completed = True else: req.state = LlmRequestState.DISAGG_TRANS_ERROR except Exception: req.state = LlmRequestState.DISAGG_TRANS_ERROR + raised = True raise finally: - if session is not None: - session.close() - self._recv_sessions.pop(rid, None) - self._recv_reqs.pop(rid, None) + retain = False + # The blocking wait can end while a peer is still writing; hand such + # a session to the async status path instead of closing it here. + if not completed and session is not None and session.has_transferring_tasks(): + session.cancel() + if session.has_transferring_tasks(): + retain = True + # The caller is about to see the exception; leave the error + # state visible and just let the session drain. + if not raised: + req.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS + if not retain: + if session is not None: + session.close() + self._recv_sessions.pop(rid, None) + self._recv_reqs.pop(rid, None) @nvtx_range("KvCacheTransceiverV2.request_and_receive_async") def request_and_receive_async(self, req: LlmRequest): @@ -726,10 +755,13 @@ def request_and_receive_async(self, req: LlmRequest): req.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS session = self._transfer_worker.create_rx_session(req) self._recv_sessions[rid] = session + # Register before dispatch: receive() can raise after earlier peers were + # already contacted, and the predicate then reports the session busy. + # Status polling must find the matching request for the whole drain. + self._recv_reqs[rid] = req kv_slice = self._create_kv_slice(req) req.py_kv_cache_xfer_bytes = self._slice_num_bytes(kv_slice) * self._kv_size_rank_factor session.receive(kv_slice) - self._recv_reqs[rid] = req def check_context_transfer_status( self, at_least_request_num: Optional[int], mark_complete: bool = False @@ -766,7 +798,8 @@ def check_context_transfer_status( session = self._send_sessions[rid] result = session.wait_complete(blocking=block_all) if session.status == SessionStatus.CANCELLED: - cancelled.append(rid) + if not session.has_transferring_tasks(): + cancelled.append(rid) elif result == WaitResult.COMPLETED: completed.append(rid) elif result is None: @@ -781,9 +814,12 @@ def check_context_transfer_status( logger.warning(f"TxSession rid={session.disagg_request_id} failed") failed.append(rid) - # All ranks must agree on per-rid outcome to avoid req.state divergence. cancelled, failed, completed = self._ctx_consensus_outcome( - to_process, cancelled, failed, completed + to_process, + cancelled, + failed, + completed, + self._busy_rids(self._send_sessions, to_process), ) for rid in cancelled: @@ -831,7 +867,8 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]): # cancel) or by a remote CANCEL_SESSION message (e.g. CTX # server timeout). Return the req objects so the caller can # distinguish the two cases and set the appropriate state. - cancelled.append(rid) + if not session.has_transferring_tasks(): + cancelled.append(rid) elif result == WaitResult.COMPLETED: req = self._recv_reqs[rid] if session.transfer_end_time is not None: @@ -843,9 +880,12 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]): failed.append(rid) # else: None — KV done but aux still in flight; re-poll next cycle - # All ranks must agree on per-rid outcome to avoid req.state divergence. cancelled, failed, completed = self._gen_consensus_outcome( - to_process, cancelled, failed, completed + to_process, + cancelled, + failed, + completed, + self._busy_rids(self._recv_sessions, to_process), ) cancelled_reqs = [] @@ -869,24 +909,60 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]): kv_cache_size=req.kv_cache_size, ) + # Post-transfer work is local and can raise (_apply_aux, + # _assert_disagg_history_declared). Letting it escape would skip the + # remaining rids and drop this rank out of the next collective, hanging + # the others. Finish the batch, then agree on the outcome. + local_postprocess_failed = [] + completed_reqs = {} for rid in completed: session = self._recv_sessions[rid] req = self._recv_reqs[rid] - # transfer_end already stamped at completion detection above. - req.set_kv_cache_size(getattr(req, "py_kv_cache_xfer_bytes", 0)) - if self._need_aux_transfer(req): - self._apply_aux(session, req) - self._assert_disagg_history_declared(req) - req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE - session.close() - del self._recv_reqs[rid] - del self._recv_sessions[rid] + completed_reqs[rid] = req + try: + # transfer_end already stamped at completion detection above. + # Only fall back to the dispatch-time estimate; the session's + # own byte count, when it has one, is the accurate value. + if session.kv_cache_size_bytes <= 0: + req.set_kv_cache_size(getattr(req, "py_kv_cache_xfer_bytes", 0)) + if self._need_aux_transfer(req): + self._apply_aux(session, req) + self._assert_disagg_history_declared(req) + except Exception as exc: + local_postprocess_failed.append(rid) + logger.error( + f"Disagg gen postprocess failed rank={self._dist.rank} rid={rid}: {exc}" + ) + finally: + # NIXL is done here, so unregistering is safe either way. + session.close() + del self._recv_reqs[rid] + del self._recv_sessions[rid] + + # A rid that failed postprocess on ANY rank must not be COMPLETE on the + # others, or req.state diverges. `completed` is already agreed, so this + # collective is symmetric. + if completed and self._gen_need_sync: + postprocess_failed = self._union(self._gen_allgather(list(local_postprocess_failed))) + else: + postprocess_failed = set(local_postprocess_failed) + for rid, req in completed_reqs.items(): + req.state = ( + LlmRequestState.DISAGG_TRANS_ERROR + if rid in postprocess_failed + else LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE + ) + transfer_failed = failed + failed = transfer_failed + [rid for rid in completed if rid in postprocess_failed] + completed = [rid for rid in completed if rid not in postprocess_failed] if failed: logger.warning( f"Disagg gen transfer FAILED rank={self._dist.rank} " f"rids={failed} gen_need_sync={self._gen_need_sync}" ) - self._close_failed_sessions(self._recv_sessions, self._recv_reqs, failed) + # Only transfer failures still own map entries; postprocess failures + # were already closed and removed above. + self._close_failed_sessions(self._recv_sessions, self._recv_reqs, transfer_failed) return completed, failed, cancelled_reqs @@ -966,29 +1042,14 @@ def cancel_request(self, req: LlmRequest) -> bool: # Not yet started (generation-first wait queue). self._wait_reqs.pop(rid, None) - has_transferring = False - - if rid in self._send_sessions: - self._send_sessions[rid].cancel() - if self._send_sessions[rid].has_transferring_tasks(): - has_transferring = True - else: - self._send_sessions[rid].close() - del self._send_reqs[rid] - del self._send_sessions[rid] - - if rid in self._recv_sessions: - self._recv_sessions[rid].cancel() - if self._recv_sessions[rid].has_transferring_tasks(): - has_transferring = True - else: - self._recv_sessions[rid].close() - del self._recv_reqs[rid] - del self._recv_sessions[rid] - - if has_transferring: - return False - return True + # Teardown is deferred to status polling: deleting the session here + # would drop this rank out of the cross-rank busy consensus. + matched = False + for sessions in (self._send_sessions, self._recv_sessions): + if rid in sessions: + matched = True + sessions[rid].cancel() + return not matched def get_disaggregated_params(self) -> Dict[str, Any]: # Keep this aligned with fields populated in respond_and_send_async(). diff --git a/tests/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index 79d77502baf4..8ddc582b1455 100644 --- a/tests/integration/test_lists/test-db/l0_h100.yml +++ b/tests/integration/test_lists/test-db/l0_h100.yml @@ -89,6 +89,7 @@ l0_h100: - unittest/disaggregated/test_kv_transfer.py - unittest/disaggregated/test_kv_transfer_mp.py - unittest/disaggregated/test_transceiver_bounded_polling.py + - unittest/disaggregated/test_rx_peer_drain.py - unittest/disaggregated/test_pool_matching.py - unittest/disaggregated/test_deepseek_v4_kv_transfer.py - unittest/disaggregated/test_minimax_m3_kv_transfer.py diff --git a/tests/unittest/disaggregated/test_rx_peer_drain.py b/tests/unittest/disaggregated/test_rx_peer_drain.py new file mode 100644 index 000000000000..1f0fee179779 --- /dev/null +++ b/tests/unittest/disaggregated/test_rx_peer_drain.py @@ -0,0 +1,292 @@ +# 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. +"""RxSession.has_transferring_tasks(): the predicate that pins receiver KV pages. + +It answers "may a remote peer still be writing into my destination buffers?". +cancel_request() must return False while it is True, so being wrong in one +direction corrupts KV and in the other direction hangs the request forever. + +The three properties covered here are exactly the three ways it can go wrong: + * ADP broadcast reaches more peers than will ever answer, + * a peer can disappear without answering at all, + * a duplicated answer must not count twice. +All three are load-bearing; there is no bespoke unwind path for a partial +dispatch, so the deadline has to carry that case too. +""" + +from __future__ import annotations + +import threading +import time + +import pytest + +from tensorrt_llm._torch.disaggregation.base.transfer import SessionStatus +from tensorrt_llm._torch.disaggregation.native import transfer as transfer_mod +from tensorrt_llm._torch.disaggregation.native.transfer import ( + AgentResult, + KVRecvTask, + RxSession, + TaskStatus, +) + + +class _FakeBounce: + """``bounced`` mimics the real transport: the scatter is queued, not inline.""" + + def __init__(self, bounced: bool = False) -> None: + self._bounced = bounced + + def record_failure(self, *_args, **_kwargs) -> None: + pass + + def is_bounced(self, *_args, **_kwargs) -> bool: + return self._bounced + + def record_result(self, *_args, **_kwargs) -> None: + pass # real impl defers on_done to the scatter worker + + def release_idle_reservation(self, *_args, **_kwargs) -> None: + pass + + def orphan_reservation(self, *_args, **_kwargs) -> None: + pass + + +class _FakeRegistrar: + class _RankInfo: + instance_name = "fake" + instance_rank = 0 + + self_rank_info = _RankInfo() + + +class _FakeReceiver: + def __init__(self, bounced: bool = False) -> None: + self._bounce = _FakeBounce(bounced) + self._registrar = _FakeRegistrar() + + +class _FakeParams: + disagg_request_id = 7 + ctx_request_id = None + + +class _FakeSessionArgs: + params = _FakeParams() + + +def _make_session( + num_peers_expected: int, *, need_aux: bool = False, bounced: bool = False +) -> RxSession: + """Build an RxSession without touching ZMQ, NIXL or the peer registrar.""" + session = object.__new__(RxSession) + session.lock = threading.Lock() + session.request_id = 7 + session._base_args = _FakeSessionArgs() + session._receiver = _FakeReceiver(bounced) + session._need_aux = need_aux + session._terminal_status = None + session._exception = None + session._kv_tasks = [] + session._aux_count = 0 + session._aux_status = TaskStatus.INIT + session._aux_responded_peer_ranks = set() + session._sender_endpoints = set() + session._cancel_notified = False + session._last_peer_progress = None + session._drain_timeout_logged = False + session.transfer_end_time = None + session.kv_cache_size_bytes = 0 + + task = object.__new__(KVRecvTask) + task._event = threading.Event() + task.slice_id = 0 + task.status = TaskStatus.INIT + task.expected_transfers = num_peers_expected + task.last_slice_count = 0 + task.dispatched = False + task.responded_peer_ranks = set() + task._exception = None + task._perf_timer = None + session._kv_tasks.append(task) + return session + + +def _dispatch(session: RxSession, peer_ranks) -> None: + for rank in peer_ranks: + assert session.mark_peer_dispatched(0, rank, f"tcp://peer{rank}") + + +def _kv_result(session: RxSession, peer_rank: int, status=AgentResult.SUCCESS) -> None: + session.process_kv_agent_result(peer_rank, 0, True, status) + + +# RxSession has a module-level attribute for the drain budget; patch it per test. +@pytest.fixture +def drain_timeout(monkeypatch): + def _set(seconds: float) -> None: + monkeypatch.setattr(transfer_mod, "_PEER_DRAIN_TIMEOUT_S", seconds) + + return _set + + +def test_untouched_session_is_not_transferring() -> None: + session = _make_session(2) + assert session.has_transferring_tasks() is False + + +def test_dispatched_peers_pin_the_session_until_they_answer() -> None: + session = _make_session(2) + _dispatch(session, [0, 1]) + assert session.has_transferring_tasks() is True + + _kv_result(session, 0) + assert session.has_transferring_tasks() is True + + _kv_result(session, 1) + assert session.has_transferring_tasks() is False + + +def test_adp_broadcast_ignores_peers_that_never_owned_the_request() -> None: + # Gen-first ADP broadcasts REQUEST_DATA to every DP group, but only one + # group holds the context request and replies. Waiting on dispatches rather + # than on expected responders would pin these pages for the four silent + # peers indefinitely. + session = _make_session(2) + _dispatch(session, [0, 1, 2, 3, 4, 5]) + + _kv_result(session, 0) + _kv_result(session, 1) + + assert session.has_transferring_tasks() is False + + +def test_failed_result_also_resolves_its_peer() -> None: + session = _make_session(2) + _dispatch(session, [0, 1]) + _kv_result(session, 0, AgentResult.FAILED) + _kv_result(session, 1, AgentResult.FAILED) + assert session.has_transferring_tasks() is False + + +def test_duplicate_result_does_not_resolve_a_second_peer() -> None: + session = _make_session(2) + _dispatch(session, [0, 1]) + _kv_result(session, 0) + _kv_result(session, 0) # duplicate; must be ignored + assert session.has_transferring_tasks() is True + + +def test_silent_peer_releases_the_session_after_the_drain_deadline(drain_timeout) -> None: + # A CTX worker killed mid-transfer never sends its terminal result. Without + # a deadline the request would stay uncancellable forever and, because the + # transceiver feeds this predicate into a cross-rank consensus, would stall + # every other request in the same polling batch. + drain_timeout(0.05) + session = _make_session(2) + _dispatch(session, [0, 1]) + _kv_result(session, 0) + assert session.has_transferring_tasks() is True + + time.sleep(0.1) + assert session.has_transferring_tasks() is False + + +def test_progress_rearms_the_drain_deadline(drain_timeout) -> None: + # The budget is an inactivity budget, not a total-duration budget: a healthy + # transfer that keeps producing results must never trip it. + drain_timeout(2.0) + session = _make_session(3) + _dispatch(session, [0, 1, 2]) + + for rank in (0, 1): + time.sleep(0.3) # well under the budget alone, over it cumulatively + assert session.has_transferring_tasks() is True + _kv_result(session, rank) + + _kv_result(session, 2) + assert session.has_transferring_tasks() is False + + +def test_partial_dispatch_failure_falls_back_to_the_drain_deadline(drain_timeout) -> None: + # Only 2 of 4 peers were reached before dispatch raised, so the responder + # count can never be satisfied. There is no bespoke unwind path: the + # inactivity deadline is what releases the buffers. + drain_timeout(0.05) + session = _make_session(4) + _dispatch(session, [0, 1]) + _kv_result(session, 0) + _kv_result(session, 1) + assert session.has_transferring_tasks() is True + + time.sleep(0.1) + assert session.has_transferring_tasks() is False + + +def test_aux_transfer_keeps_the_session_pinned_until_aux_answers() -> None: + session = _make_session(1, need_aux=True) + _dispatch(session, [0]) + _kv_result(session, 0) + assert session.has_transferring_tasks() is True + + session.process_aux_agent_result(0, AgentResult.SUCCESS) + assert session.has_transferring_tasks() is False + + +def test_bounced_transfer_stays_pinned_until_the_scatter_lands() -> None: + # On the bounce path a SUCCESS result only means the data reached the bounce + # arena; scatter_write_result queues the copy into the KV pages and + # task.complete() runs later, in the scatter worker's on_done. Releasing on + # responder count alone would free those pages under the pending scatter. + session = _make_session(1, bounced=True) + _dispatch(session, [0]) + _kv_result(session, 0) + + assert session._kv_tasks[0].responded_peer_ranks == {0} + assert session._kv_tasks[0].status == TaskStatus.TRANSFERRING + assert session.has_transferring_tasks() is True + + # What the scatter worker's on_done does once the copy has landed. + session._kv_tasks[0].complete() + assert session.has_transferring_tasks() is False + + +def test_cancel_is_idempotent_and_notifies_senders_once() -> None: + session = _make_session(2) + sent: list = [] + session._receiver.send_cancel_to_senders = lambda rid, eps: sent.append((rid, set(eps))) + _dispatch(session, [0, 1]) + + session.cancel() + session.cancel() + session.cancel() + + assert len(sent) == 1 + assert sent[0][1] == {"tcp://peer0", "tcp://peer1"} + assert session._terminal_status == SessionStatus.CANCELLED + + +def test_dispatch_is_refused_once_the_session_is_cancelled() -> None: + # dispatch_task must not contact a sender after cancellation won the race, + # or that peer writes into buffers we are about to release. + session = _make_session(2) + session._receiver.send_cancel_to_senders = lambda *_a: None + assert session.mark_peer_dispatched(0, 0, "tcp://peer0") is True + + session.cancel() + + assert session.mark_peer_dispatched(0, 1, "tcp://peer1") is False + assert session._kv_tasks[0].responded_peer_ranks == set() diff --git a/tests/unittest/disaggregated/test_transceiver_bounded_polling.py b/tests/unittest/disaggregated/test_transceiver_bounded_polling.py index 9185d9ef2283..421c970928dd 100644 --- a/tests/unittest/disaggregated/test_transceiver_bounded_polling.py +++ b/tests/unittest/disaggregated/test_transceiver_bounded_polling.py @@ -39,6 +39,15 @@ @dataclass class _FakeRequest: state: Optional[LlmRequestState] = None + kv_cache_transfer_time_ms: float = 0.0 + kv_cache_size: int = 0 + py_kv_cache_xfer_bytes: int = 0 + + def set_kv_cache_transfer_end(self, _t) -> None: + pass + + def set_kv_cache_size(self, size: int) -> None: + self.kv_cache_size = size class _FakeTransferWorker: @@ -58,14 +67,19 @@ def __init__( status: SessionStatus = SessionStatus.READY, is_completed: bool = False, has_failed: bool = False, + has_transferring_tasks: bool = False, ) -> None: self._rid = rid self._wait_result = wait_result self._status = status self._is_completed = is_completed self._has_failed = has_failed + self._has_transferring_tasks = has_transferring_tasks self.blocking_calls: list[bool] = [] self.closed = False + self.cancelled = False + self.transfer_end_time = None + self.kv_cache_size_bytes = 0 self.aux_slot: Optional[int] = 0 @property @@ -86,6 +100,12 @@ def is_completed(self) -> bool: def has_failed(self) -> bool: return self._has_failed + def has_transferring_tasks(self) -> bool: + return self._has_transferring_tasks + + def cancel(self) -> None: + self.cancelled = True + def close(self) -> None: self.closed = True self.aux_slot = None @@ -99,6 +119,7 @@ def __init__( on_wait: Optional[Callable[[Optional[float]], None]] = None, ) -> None: self.status = status + self.pending_peers: set[int] = set() self._wait_results = list(wait_result) if isinstance(wait_result, list) else [wait_result] self._on_wait = on_wait self.wait_calls: list[Optional[float]] = [] @@ -140,7 +161,7 @@ def _make_transceiver( transceiver._ctx_need_pp_sync = False transceiver._transfer_worker = _FakeTransferWorker() transceiver._ctx_consensus = lambda local_ids: list(local_ids) - transceiver._ctx_consensus_outcome = lambda _to_process, cancelled, failed, completed: ( + transceiver._ctx_consensus_outcome = lambda _to_process, cancelled, failed, completed, _busy: ( cancelled, failed, completed, @@ -359,27 +380,45 @@ def test_gen_transfer_status_enters_consensus_when_sync_required() -> None: def test_consensus_outcome_uses_single_batched_allgather() -> None: - # The cancelled/failed/completed id lists are exchanged with ONE allgather - # (packed as a list-of-lists) instead of three; verify a single call and that - # union (cancelled/failed) + intersection (completed) semantics are preserved. + # Every outcome list is exchanged with ONE allgather (packed as a + # list-of-lists); verify a single call and that union (cancelled/failed/busy) + # plus intersection (completed) semantics are preserved. transceiver = object.__new__(KvCacheTransceiverV2) calls: list = [] def fake_allgather(payload): calls.append(payload) - # rank0 = this rank's [cancelled, failed, completed]; rank1 = a peer rank. - return [payload, [[], [99], [7, 8]]] + # rank0 = this rank's [cancelled, failed, completed, busy]; rank1 = a peer. + return [payload, [[], [99], [7, 8], []]] to_process = [1, 2, 7, 8, 99] - new_cancelled, new_failed, new_completed = transceiver._consensus_outcome( - to_process, [1], [2], [7], fake_allgather, True + new_cancelled, new_failed, new_completed, new_busy = transceiver._consensus_outcome( + to_process, [1], [2], [7, 8], [8], fake_allgather, True ) - assert len(calls) == 1 # batched: a single allgather, not three - assert calls[0] == [[1], [2], [7]] + assert len(calls) == 1 # batched: a single allgather, not four + assert calls[0] == [[1], [2], [7, 8], [8]] assert new_cancelled == [1] # union of cancelled across ranks assert new_failed == [2, 99] # union of failed across ranks - assert new_completed == [7] # intersection only (8 is completed on the peer only) + assert new_completed == [7] # 8 is complete everywhere but still busy locally + assert new_busy == [8] + + +def test_consensus_outcome_defers_a_terminal_rid_that_is_busy_on_a_peer() -> None: + # A rank that has already drained must not close a session that another + # rank is still writing into. + transceiver = object.__new__(KvCacheTransceiverV2) + + def fake_allgather(payload): + # rank1 reports rid 5 as still busy. + return [payload, [[5], [], [], [5]]] + + cancelled, failed, completed, busy = transceiver._consensus_outcome( + [5], [5], [], [], [], fake_allgather, True + ) + + assert cancelled == [] + assert busy == [5] def test_ctx_tp_consensus_does_not_complete_when_peer_times_out() -> None: @@ -387,10 +426,10 @@ def test_ctx_tp_consensus_does_not_complete_when_peer_times_out() -> None: transceiver._ctx_need_tp_sync = True transceiver._ctx_need_pp_sync = False transceiver._dist = SimpleNamespace( - tp_allgather=lambda payload: [payload, [[], [], []]], + tp_allgather=lambda payload: [payload, [[], [], [], []]], ) - cancelled, failed, completed = transceiver._ctx_consensus_outcome([21], [], [], [21]) + cancelled, failed, completed = transceiver._ctx_consensus_outcome([21], [], [], [21], []) assert cancelled == [] assert failed == [] @@ -403,10 +442,10 @@ def test_ctx_pp_consensus_does_not_complete_when_peer_times_out() -> None: transceiver._ctx_need_pp_sync = True transceiver._dist = SimpleNamespace( tp_allgather=Mock(side_effect=AssertionError("TP allgather must be skipped")), - pp_allgather=lambda payload: [payload, [[], [], []]], + pp_allgather=lambda payload: [payload, [[], [], [], []]], ) - cancelled, failed, completed = transceiver._ctx_consensus_outcome([22], [], [], [22]) + cancelled, failed, completed = transceiver._ctx_consensus_outcome([22], [], [], [22], []) assert cancelled == [] assert failed == [] @@ -1002,3 +1041,111 @@ def test_prepare_context_requests_skips_consensus_when_nothing_waiting() -> None transceiver.prepare_context_requests([]) transceiver._ctx_consensus.assert_not_called() + + +def test_failed_session_is_cancelled_before_it_is_closed() -> None: + # close() never notifies the peers. A drain deadline can expire before a + # backlogged sender has even started writing, so the failed path has to + # cancel first or that sender writes into pages we already released. + session = _FakeSession(31, WaitResult.FAILED, has_failed=True) + req = _FakeRequest() + transceiver = object.__new__(KvCacheTransceiverV2) + + transceiver._close_failed_sessions({31: session}, {31: req}, [31]) + + assert session.cancelled is True + assert session.closed is True + assert req.state == LlmRequestState.DISAGG_TRANS_ERROR + + +def test_recv_request_is_registered_before_dispatch_can_raise() -> None: + # dispatch can raise after earlier peers were contacted, leaving a session + # the predicate calls busy. If the request were only registered after + # receive(), _close_failed_sessions would KeyError on it once the drain + # deadline expired, taking the executor loop down with it. + rid = 77 + req = SimpleNamespace( + request_id=rid, + py_disaggregated_params=None, + state=None, + py_kv_cache_xfer_bytes=0, + set_kv_cache_transfer_start=lambda _t: None, + ) + + class _RaisingSession(_FakeSession): + def receive(self, _slice) -> None: + raise RuntimeError("dispatch failed on peer 3 of 4") + + session = _RaisingSession(rid, WaitResult.FAILED) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._ever_had_recv_session = False + transceiver._recv_sessions = {} + transceiver._recv_reqs = {} + transceiver._transfer_worker = SimpleNamespace(create_rx_session=lambda _r: session) + transceiver._create_kv_slice = lambda _r: object() + transceiver._slice_num_bytes = lambda _s: 0 + transceiver._kv_size_rank_factor = 1 + + with pytest.raises(RuntimeError): + KvCacheTransceiverV2.request_and_receive_async(transceiver, req) + + # Both maps must agree, so the drain path can clean up without KeyError. + assert rid in transceiver._recv_sessions + assert rid in transceiver._recv_reqs + transceiver._close_failed_sessions(transceiver._recv_sessions, transceiver._recv_reqs, [rid]) + assert transceiver._recv_sessions == {} and transceiver._recv_reqs == {} + + +def _make_gen_postprocess_transceiver(rids, failing_rid, *, need_sync=False, peer_failed=()): + """check_gen_transfer_status with every rid already agreed COMPLETED.""" + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._ever_had_recv_session = True + transceiver._gen_need_sync = need_sync + transceiver._dist = SimpleNamespace(rank=0) + transceiver._recv_sessions = {r: _FakeSession(r, WaitResult.COMPLETED) for r in rids} + transceiver._recv_reqs = {r: _FakeRequest() for r in rids} + transceiver._gen_consensus = lambda ids: list(ids) + transceiver._build_to_process = lambda *_a, **_k: list(rids) + transceiver._busy_rids = lambda *_a, **_k: [] + transceiver._gen_consensus_outcome = lambda _tp, c, f, d, _b: (c, f, list(rids)) + transceiver._close_failed_sessions = Mock() + transceiver._need_aux_transfer = lambda _r: False + transceiver._sync_transfer_timing = lambda _r: None + + def _assert_history(req): + if transceiver._recv_reqs_snapshot.get(id(req)) == failing_rid: + raise RuntimeError("history not declared") + + transceiver._recv_reqs_snapshot = {id(r): rid for rid, r in transceiver._recv_reqs.items()} + transceiver._assert_disagg_history_declared = _assert_history + transceiver._gen_allgather = lambda payload: [payload, list(peer_failed)] + return transceiver + + +def test_gen_postprocess_failure_does_not_strand_later_requests() -> None: + # A raise on the first rid used to escape the loop, leaving the rest at + # IN_PROGRESS with their sessions still registered, and dropping this rank + # out of the next collective. + transceiver = _make_gen_postprocess_transceiver([41, 42], failing_rid=41) + + completed, failed, _ = transceiver.check_gen_transfer_status(at_least_request_num=0) + + assert completed == [42] + assert failed == [41] + # Both were closed and removed; the failed one must not be cleaned up twice. + assert transceiver._recv_sessions == {} and transceiver._recv_reqs == {} + transceiver._close_failed_sessions.assert_called_once() + assert transceiver._close_failed_sessions.call_args[0][2] == [] + + +def test_gen_postprocess_failure_on_a_peer_rank_is_not_completed_locally() -> None: + # Local postprocess succeeded for 51, but a peer rank failed it. Publishing + # TRANS_COMPLETE here while the peer publishes an error diverges req.state. + transceiver = _make_gen_postprocess_transceiver( + [51], failing_rid=None, need_sync=True, peer_failed=[51] + ) + + completed, failed, _ = transceiver.check_gen_transfer_status(at_least_request_num=0) + + assert completed == [] + assert failed == [51]