From 9bb4b212cf102bad48f7c4729b688d06295dd505 Mon Sep 17 00:00:00 2001 From: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> Date: Thu, 17 Sep 2026 13:42:25 -0700 Subject: [PATCH 1/3] test: reproduce unstable KV transfer outcomes Signed-off-by: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> --- .../disaggregated/test_task_handle.py | 141 ++++++++++++++++++ 1 file changed, 141 insertions(+) diff --git a/tests/unittest/disaggregated/test_task_handle.py b/tests/unittest/disaggregated/test_task_handle.py index 99380759618a..e50bb60de682 100644 --- a/tests/unittest/disaggregated/test_task_handle.py +++ b/tests/unittest/disaggregated/test_task_handle.py @@ -640,3 +640,144 @@ def test_a_send_session_reads_a_parked_cancel_as_the_peers_too(): outcome = TaskHandle(session, task, TOKENS).poll() assert isinstance(outcome, Cancelled) assert outcome.by_peer is True + + +# --------------------------------------------------------------------------- +# Logical decisions belong to the transition, not the first observer +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("poll_before_completion", [False, True]) +@pytest.mark.parametrize("ending", ["failure", "local_cancel", "peer_cancel"]) +def test_receive_outcome_does_not_depend_on_polling_before_late_completion( + ending: str, poll_before_completion: bool +) -> None: + session, handle = _receiving_from(1) + if ending == "failure": + session.fail_admission(RuntimeError("another publication failed")) + else: + session.cancel_local(by_peer=ending == "peer_cancel") + if poll_before_completion: + assert handle.poll() is not None + + # The writer finishes after the logical decision. Its report still drains the transfer. + _report(session, 0, AgentResult.SUCCESS) + + for observer in (handle, TaskHandle(session, session._kv_tasks[0], TOKENS)): + outcome = observer.poll() + if ending == "failure": + assert isinstance(outcome, Failed) + assert "another publication failed" in outcome.reason + else: + assert isinstance(outcome, Cancelled) + assert outcome.by_peer is (ending == "peer_cancel") + assert outcome.reports_pending is False + + +@pytest.mark.parametrize("by_peer", [False, True]) +def test_receive_cancel_keeps_its_outcome_when_a_writer_later_fails(by_peer: bool) -> None: + session, _ = _receiving_from(1) + session.cancel_local(by_peer=by_peer) + + _report(session, 0, AgentResult.FAILED) + + outcome = TaskHandle(session, session._kv_tasks[0], TOKENS).poll() + assert isinstance(outcome, Cancelled) + assert outcome.by_peer is by_peer + assert outcome.reports_pending is False + + +def test_receive_session_failure_precedes_later_cancel_without_an_observer() -> None: + session, handle = _receiving_from(1) + session.fail_admission(RuntimeError("first publication failure")) + + assert session.cancel_local() is True + + outcome = handle.poll() + assert isinstance(outcome, Failed) + assert "first publication failure" in outcome.reason + assert outcome.reports_pending is True + + +def _sending_pieces(count: int = 1) -> tuple[Sender, TxSession]: + sender = _wired_sender() + sender.dispatch_task = MagicMock() + sender._get_result_dealer = MagicMock() + sender._instance_rank = 0 + session = TxSession( + request_id=30, params=DisaggregatedParams(disagg_request_id=30), sender=sender + ) + for _ in range(count): + session.send(_sole_piece()) + return sender, session + + +@pytest.mark.parametrize("by_peer", [False, True]) +def test_queued_sender_abort_preserves_the_committed_cancellation(by_peer: bool) -> None: + sender, session = _sending_pieces() + task = session.kv_tasks[0] + session.cancel_local(by_peer=by_peer) + # Execute the real worker's pre-submission abort branch after cancellation won. + empty = SimpleNamespace(size=0) + write_meta = SimpleNamespace( + src_ptrs=empty, + dst_ptrs=empty, + sizes=empty, + unique_rid=30, + slice_id=0, + receiver_slice_id=0, + peer_rank=0, + peer_endpoint="tcp://receiver:1234", + task=task, + ) + + sender._deliver_kv_to_agent(write_meta) + + outcome = TaskHandle(session, task, TOKENS).poll() + assert isinstance(outcome, Cancelled) + assert outcome.by_peer is by_peer + sender._get_result_dealer.return_value.send.assert_called_once() + + +@pytest.mark.parametrize("poll_before_completion", [False, True]) +def test_sender_sibling_failure_is_stable_across_late_completion( + poll_before_completion: bool, +) -> None: + _, session = _sending_pieces(2) + failed, pending = session.kv_tasks + pending.status = TaskStatus.TRANSFERRING + handle = TaskHandle(session, pending, TOKENS) + + failed.fail(RuntimeError("first sibling failed")) + if poll_before_completion: + assert isinstance(handle.poll(), Failed) + pending.complete() + + for observer in (handle, TaskHandle(session, pending, TOKENS)): + outcome = observer.poll() + assert isinstance(outcome, Failed) + assert "first sibling failed" in outcome.reason + + +def test_terminal_failure_cause_is_stable_without_polling() -> None: + session, handle = _receiving_from(1) + task = session._kv_tasks[0] + task.fail(RuntimeError("original failure")) + task.fail(RuntimeError("cleanup failure")) + + outcome = handle.poll() + assert isinstance(outcome, Failed) + assert outcome.reason == "original failure" + + +@pytest.mark.parametrize("ending", ["failure", "local_cancel", "peer_cancel"]) +def test_delivered_outcome_precedes_later_session_terminal_events(ending: str) -> None: + session, handle = _receiving_from(1) + _report(session, 0, AgentResult.SUCCESS) + if ending == "failure": + session.fail_admission(RuntimeError("later failure")) + else: + session.cancel_local(by_peer=ending == "peer_cancel") + + assert isinstance(handle.poll(), Delivered) + assert handle.poll().token_end == TOKENS From 6f0b43a23dc86516c07beaf6f56ff75fb442168e Mon Sep 17 00:00:00 2001 From: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> Date: Thu, 17 Sep 2026 13:54:33 -0700 Subject: [PATCH 2/3] fix: commit KV transfer outcomes at state transitions Signed-off-by: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> --- .../_torch/disaggregation/base/backend.py | 6 +- .../_torch/disaggregation/native/handle.py | 79 +++-------- .../_torch/disaggregation/native/transfer.py | 123 ++++++++++++++---- .../unittest/disaggregated/test_peer_fetch.py | 63 +++++---- .../disaggregated/test_task_handle.py | 112 +++++++++++++--- .../test_transceiver_bounded_polling.py | 2 + .../test_transfer_ownership_regressions.py | 2 + 7 files changed, 249 insertions(+), 138 deletions(-) diff --git a/tensorrt_llm/_torch/disaggregation/base/backend.py b/tensorrt_llm/_torch/disaggregation/base/backend.py index a1d0dd5e0e9d..647e66c4db1a 100644 --- a/tensorrt_llm/_torch/disaggregation/base/backend.py +++ b/tensorrt_llm/_torch/disaggregation/base/backend.py @@ -197,8 +197,9 @@ class Cancelled: Outcome = Union[Delivered, Failed, Cancelled] """How one delivery ended. ``None`` rather than a member means it has not ended yet. -Latch which member it is, not the object: ``reports_pending`` turns from true to false over time, so -a stored outcome carries a stale one. +The logical member and its cause are committed at the task/session transition; +polling only observes that decision. ``reports_pending`` can still change, so a +stored outcome carries stale report progress, not a different logical result. ``reports_pending`` asks one question and only one: is a report this transfer was owed still to arrive. Receiving waits on the writers' reports, sending on word about its own writes. @@ -226,6 +227,7 @@ def poll(self) -> Optional[Outcome]: A failure is reported as soon as it is known, which may be before every report is in -- that second question is carried by the outcome itself. + The logical result is stable and does not depend on when the first poll occurs. """ ... diff --git a/tensorrt_llm/_torch/disaggregation/native/handle.py b/tensorrt_llm/_torch/disaggregation/native/handle.py index 0e18e4ddfbfd..1f757d3e1dd5 100644 --- a/tensorrt_llm/_torch/disaggregation/native/handle.py +++ b/tensorrt_llm/_torch/disaggregation/native/handle.py @@ -12,18 +12,10 @@ # 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. -"""How a native transfer's state reads as a contract outcome. +"""Observe a native task's committed logical result and current report progress. -Both directions land here: a send session and a receive session expose the same three things a -handle needs -- the session's own verdict, its exception, and per-task status. What differs is what -``reports_pending`` is owed by: on the receive side the writers' reports, on the send side word -about this side's own writes. Neither answers whether anyone is still touching the memory; that is a -separate question nothing here asks. - -TODO: Which ending a piece reports depends on when it is first polled. A piece whose session has -gone terminal while the piece itself is still writing latches the session's verdict, yet the piece -may finish afterwards -- so an early poll says failed and a late one says delivered. Nothing polls -these yet; it has to be settled before the surface is frozen. +Task/session transitions commit the result, not polling. Reports and physical +quiescence remain separate questions; an outcome never authorizes memory reuse. """ from __future__ import annotations @@ -32,56 +24,29 @@ from tensorrt_llm._torch.disaggregation.base import Cancelled, Delivered, Failed, Outcome -from .transfer import SessionStatus, TaskStatus +from .transfer import KVRecvTask, KVSendTask, RxSession, SessionStatus, TaskStatus, TxSession class TaskHandle: - """One piece of one request, seen through the session carrying it. - - Pieces of the same request share that session, so this piece's own state is read first and the - session's verdict only answers for a piece that has not ended on its own. - """ + """One piece's decision, including when the first observer arrives late.""" - def __init__(self, session, task, token_end: int): + def __init__( + self, session: RxSession | TxSession, task: KVRecvTask | KVSendTask, token_end: int + ) -> None: self._session = session self._task = task self._token_end = token_end - self._ended: Optional[Outcome] = None def poll(self) -> Optional[Outcome]: - if self._ended is not None: - # Which ending this was cannot change; only whether a writer still owes word can. - return self._rebuild(self._ended) - # This piece's own ending wins: bytes that landed landed, whatever became of its siblings, - # and "cancelled" means stopped short of delivering. - if self._task.status is TaskStatus.TRANSFERRED: - outcome = Delivered(token_end=self._token_end) - elif self._task.status is TaskStatus.ERROR: - if self._ended_by_cancel(): - outcome = Cancelled(by_peer=self._session.cancelled_by_peer, reports_pending=True) - else: - outcome = Failed(reason=self._why_failed(), reports_pending=True) - elif self._session.status is SessionStatus.CANCELLED: - outcome = Cancelled(by_peer=self._session.cancelled_by_peer, reports_pending=True) - elif self._session.status is SessionStatus.ERROR: - outcome = Failed(reason=self._why_failed(), reports_pending=True) - else: + result = self._task.logical_outcome + if result is None: return None - self._ended = outcome - return self._rebuild(outcome) - - def _rebuild(self, ended: Outcome) -> Outcome: - """The latched ending, with today's answer to whether a report is still owed. - - A session is shared by several pieces, so a sibling's failure moves the session's own - verdict after this piece has already ended; the ending latched here does not move with it. - """ - if isinstance(ended, Delivered): - return ended + if result.status is SessionStatus.TRANSFERRED: + return Delivered(token_end=self._token_end) owed = self._owed_a_report() - if isinstance(ended, Cancelled): - return Cancelled(by_peer=ended.by_peer, reports_pending=owed) - return Failed(reason=ended.reason, reports_pending=owed) + if result.status is SessionStatus.CANCELLED: + return Cancelled(by_peer=result.by_peer, reports_pending=owed) + return Failed(reason=result.reason, reports_pending=owed) def _owed_a_report(self) -> bool: """Whether anyone still owes word about this piece. @@ -98,20 +63,6 @@ def _owed_a_report(self) -> bool: return self._task.status is TaskStatus.TRANSFERRING return outstanding - def _ended_by_cancel(self) -> bool: - """Whether this piece's error is the cancellation itself. - - A cancel ends the pieces that had not started by failing them, which is the same task state - a real transfer error leaves; only the recorded cause tells the two apart, and it is the - very object the cancel installed rather than one that merely reads like it. - """ - cancelled = getattr(self._session, "_cancel_exception", None) - return cancelled is not None and self._task._exception is cancelled - - def _why_failed(self) -> str: - error = self._task._exception or self._session.exception - return str(error) if error is not None else "transfer failed without a recorded cause" - class NothingPublished: """Submission failed before any peer was told where to write, so there is nothing to wait on.""" diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index e262fd825ccc..06d4110d1e1c 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -202,6 +202,74 @@ class TaskStatus(Enum): ERROR = "ERROR" +@dataclass(frozen=True) +class _LogicalOutcome: + """A committed delivery decision, independent of physical/report progress.""" + + status: SessionStatus + reason: str = "" + by_peer: bool = False + + +class _LogicalOutcomes: + """Arbitrate task completion against session failure/cancellation at event time. + + A failed or cancelled session ends only pending tasks. Already committed + results, including their original cause, never change. No task, request or + backend handle is retained here: memory ownership remains a separate concern. + """ + + def __init__(self) -> None: + self._lock = threading.Lock() + self._terminal: Optional[_LogicalOutcome] = None + self._results: list[Optional[_LogicalOutcome]] = [] + + def add_task(self) -> int: + with self._lock: + index = len(self._results) + self._results.append(self._terminal) + return index + + def get(self, index: int) -> Optional[_LogicalOutcome]: + with self._lock: + return self._results[index] + + def complete(self, index: int) -> None: + with self._lock: + if self._results[index] is None: + self._results[index] = _LogicalOutcome(SessionStatus.TRANSFERRED) + + def fail(self, error: Exception) -> None: + self._end(_LogicalOutcome(SessionStatus.ERROR, reason=str(error))) + + def cancel(self, by_peer: bool) -> None: + self._end(_LogicalOutcome(SessionStatus.CANCELLED, by_peer=by_peer)) + + def _end(self, outcome: _LogicalOutcome) -> None: + with self._lock: + if self._terminal is not None: + return + self._terminal = outcome + for index, result in enumerate(self._results): + if result is None: + self._results[index] = self._terminal + + +class _LogicalTask: + def __init__(self) -> None: + self._logical_outcomes = _LogicalOutcomes() + self._logical_index = self._logical_outcomes.add_task() + + def bind_logical_outcomes(self, outcomes: _LogicalOutcomes) -> None: + """Join the session before this task is exposed to a worker or caller.""" + self._logical_outcomes = outcomes + self._logical_index = outcomes.add_task() + + @property + def logical_outcome(self) -> Optional[_LogicalOutcome]: + return self._logical_outcomes.get(self._logical_index) + + class _ReceiveOperationOwner: """Track destination access independently from a task's logical result.""" @@ -447,8 +515,9 @@ class _PhysicalOperation: status: Optional[object] = None -class SendTaskBase: +class SendTaskBase(_LogicalTask): def __init__(self, params: DisaggregatedParams): + super().__init__() self.status = TaskStatus.INIT self._event = threading.Event() self._exception: Optional[Exception] = None @@ -460,11 +529,13 @@ def __init__(self, params: DisaggregatedParams): self._physical_operations: dict[int, _PhysicalOperation] = {} def fail(self, exc: Exception) -> None: + self._logical_outcomes.fail(exc) self._exception = exc self.status = TaskStatus.ERROR self._event.set() def complete(self) -> None: + self._logical_outcomes.complete(self._logical_index) self.status = TaskStatus.TRANSFERRED self._event.set() @@ -1848,11 +1919,11 @@ def __init__( self._reported_aux_peer_ranks: set[int] = set() self._has_last_slice = False self.lock = threading.Lock() + self._logical_outcomes = _LogicalOutcomes() self._exception: Optional[Exception] = None self._closed = False self._terminal_status: Optional[SessionStatus] = None - self._cancel_exception: Optional[Exception] = None self.transfer_start_time = None self.transfer_end_time = None # Must be last: makes session visible to listener thread, @@ -1912,6 +1983,7 @@ def send(self, chunk: Chunk) -> None: prompt_len=self._base_args.prompt_len, ) task._unique_rid = self.disagg_request_id + task.bind_logical_outcomes(self._logical_outcomes) self.kv_tasks.append(task) self._has_last_slice |= chunk.is_last req_info_snapshot = dict(self._sender._get_req_info(task._unique_rid) or {}) @@ -1936,6 +2008,7 @@ def send_aux(self) -> AuxSendTask: params = self._base_args.params task = AuxSendTask(params, self.aux_slot) task._unique_rid = self.disagg_request_id + task.bind_logical_outcomes(self._logical_outcomes) self.aux_task = task self._report_unsubmitted_aux_failures(aux_failures) if terminal_error is not None: @@ -2035,19 +2108,16 @@ def cancel(self, by_peer: bool = False) -> bool: def cancel_local(self, by_peer: bool = False) -> bool: aux_failures: list[RecvReqInfo] = [] with self.lock: - # Only an earlier cancel refuses: a failed session may still have peers touching memory, - # and they are told to stop here. Which ending a piece reports is latched by its handle, - # so it does not depend on this. + # A later cancellation must still notify peers after logical failure. + # The logical arbiter preserves whichever outcome already committed. if self._terminal_status == SessionStatus.CANCELLED: return False + self._logical_outcomes.cancel(by_peer) self._terminal_status = SessionStatus.CANCELLED # Who asked is not recoverable later, and the two differ: the peer asking is a transfer # error, our own side asking is an ordinary end. self.cancelled_by_peer = by_peer exc = RuntimeError(f"TxSession {self.disagg_request_id} cancelled") - # Kept so a task failed by the line below is recognisable as cancelled rather than - # broken; its own status cannot say which, since both end it the same way. - self._cancel_exception = exc for task in self.kv_tasks: if task.status == TaskStatus.INIT: task.fail(exc) @@ -2159,6 +2229,7 @@ def wait_for_task(task: SendTaskBase) -> Optional[WaitResult]: self._exception = RuntimeError( "required auxiliary transfer was not dispatched" ) + self._logical_outcomes.fail(self._exception) self._terminal_status = SessionStatus.ERROR return WaitResult.FAILED result = wait_for_task(self.aux_task) @@ -2177,6 +2248,7 @@ def set_exception(self, reason: str = "") -> None: aux_failures: list[RecvReqInfo] = [] with self.lock: self._exception = RuntimeError(msg) + self._logical_outcomes.fail(self._exception) self._terminal_status = SessionStatus.ERROR for task in self.kv_tasks: if not task.is_done: @@ -2219,7 +2291,7 @@ def __del__(self): logger.warning(f"TxSession.__del__: exception during close: {e}") -class KVRecvTask: +class KVRecvTask(_LogicalTask): def __init__( self, unique_rid: Optional[int], @@ -2228,6 +2300,7 @@ def __init__( params: DisaggregatedParams, aux_slot: Optional[int], ): + super().__init__() self._event = threading.Event() self.slice_id = slice_id self.status = TaskStatus.INIT @@ -2245,6 +2318,7 @@ def __init__( self._ownership_state_lock: Optional[threading.Lock] = None def fail(self, exc: Exception) -> None: + self._logical_outcomes.fail(exc) if self._ownership_state_lock is None: self._exception = exc self.status = TaskStatus.ERROR @@ -2257,12 +2331,14 @@ def fail(self, exc: Exception) -> None: def complete(self) -> None: if self._ownership_state_lock is None: + self._logical_outcomes.complete(self._logical_index) self.status = TaskStatus.TRANSFERRED self._event.set() return with self._ownership_state_lock: if self.status == TaskStatus.ERROR: return + self._logical_outcomes.complete(self._logical_index) self.status = TaskStatus.TRANSFERRED self._event.set() @@ -2949,7 +3025,6 @@ def __init__( self._exception: Optional[Exception] = None self._closed = False self._terminal_status: Optional[SessionStatus] = None - self._cancel_exception: Optional[Exception] = None self.transfer_start_time = None self.transfer_end_time = None self.kv_cache_size_bytes: int = 0 @@ -2968,6 +3043,7 @@ def __init__( # publication wins the state transition. self._publication_lock = threading.Lock() self.lock = threading.Lock() + self._logical_outcomes = _LogicalOutcomes() try: self._receiver.setup_session(self) except Exception: @@ -3013,6 +3089,7 @@ def _poison_receiver_ownership(self, error: Exception) -> None: def _record_ownership_evidence_error(self, error: Exception) -> None: """Record a fatal ownership-evidence error and close receiver admission.""" + self._logical_outcomes.fail(error) self._exception = error if self._terminal_status is None: self._terminal_status = SessionStatus.ERROR @@ -3113,6 +3190,7 @@ def receive(self, chunk: Chunk) -> None: params, aux_slot=self.aux_slot, ) + task.bind_logical_outcomes(self._logical_outcomes) self._kv_tasks.append(task) self._receiver.dispatch_task(task) @@ -3129,6 +3207,7 @@ def prepare_receive(self, chunk: Chunk) -> Optional[KVRecvTask]: params, aux_slot=self.aux_slot, ) + task.bind_logical_outcomes(self._logical_outcomes) task.begin_publication() self._kv_tasks.append(task) return task @@ -3147,6 +3226,7 @@ def dispatch_prepared_receive(self, task: KVRecvTask) -> None: def fail_admission(self, error: Exception) -> None: """Fail logical admission without releasing possibly published destinations.""" with self.lock: + self._logical_outcomes.fail(error) self._exception = error if self._terminal_status is None: self._terminal_status = SessionStatus.ERROR @@ -3259,10 +3339,10 @@ def on_done( 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. + # Runs on the scatter worker thread for the bounced path. Do not acquire + # RxSession.lock here: the non-bounced path invokes this callback inline + # while already holding it. Task outcome/ownership locks are independent. + # complete() sets status before _event for wait_complete's status-first poll. if self._enforce_physical_ownership: task.finish_local_completion() if not success: @@ -3286,8 +3366,8 @@ def on_done( ) 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. + # landed. Plain attribute write (atomic under the GIL); 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( @@ -3377,12 +3457,14 @@ def process_aux_agent_result(self, peer_rank: int, status: AgentResult): self._exception = RuntimeError( f"Session {self.request_id} received too many aux transfers" ) + self._logical_outcomes.fail(self._exception) if self._terminal_status is None: self._terminal_status = SessionStatus.ERROR logger.error(str(self._exception)) elif status == AgentResult.FAILED: self._aux_status = TaskStatus.ERROR self._exception = RuntimeError(f"Session {self.request_id} aux transfer failed") + self._logical_outcomes.fail(self._exception) if self._terminal_status is None: self._terminal_status = SessionStatus.ERROR else: @@ -3438,19 +3520,16 @@ def resources_drained(self) -> bool: def cancel_local(self, by_peer: bool = False) -> bool: """Commit cancellation under the same lock used for publication.""" with self.lock: - # Only an earlier cancel refuses: a failed session may still have peers touching memory, - # and they are told to stop here. Which ending a piece reports is latched by its handle, - # so it does not depend on this. + # A later cancellation must still notify peers after logical failure. + # The logical arbiter preserves whichever outcome already committed. if self._terminal_status == SessionStatus.CANCELLED: return False + self._logical_outcomes.cancel(by_peer) self._terminal_status = SessionStatus.CANCELLED # Who asked is not recoverable later, and the two differ: the peer asking is a transfer # error, our own side asking is an ordinary end. self.cancelled_by_peer = by_peer exc = RuntimeError(f"RxSession {self.disagg_request_id} cancelled") - # Kept so a task failed by the loop below is recognisable as cancelled rather than - # broken; its own status cannot say which, since both end it the same way. - self._cancel_exception = exc for task in self._kv_tasks: rid_slice = (self.disagg_request_id, task.slice_id) if task.status == TaskStatus.INIT: diff --git a/tests/unittest/disaggregated/test_peer_fetch.py b/tests/unittest/disaggregated/test_peer_fetch.py index 701b0ea02712..a64f136384cb 100644 --- a/tests/unittest/disaggregated/test_peer_fetch.py +++ b/tests/unittest/disaggregated/test_peer_fetch.py @@ -2,9 +2,8 @@ # SPDX-License-Identifier: Apache-2.0 """How the native backend answers the contract's two questions. -The mapping is the whole point of the adapter: which member of ``Outcome`` a native session state -becomes, and whether a writer is still owed word about this piece. Both are exercised against stub -sessions, because the interesting states are the ones a real transfer reaches only under a race. +The adapter observes committed task outcomes and whether a writer still owes word about a piece. +Stub sessions isolate admission, but task transitions use the native logical outcome arbiter. The last section covers what the transceiver holds once admission returns, since an adapter that starts nothing still decides whether a request stays paired with a session the sweep can retire. @@ -29,7 +28,12 @@ ) from tensorrt_llm._torch.disaggregation.base.transfer import get_unique_rid from tensorrt_llm._torch.disaggregation.native.fetch import PeerFetch -from tensorrt_llm._torch.disaggregation.native.transfer import SessionStatus, TaskStatus +from tensorrt_llm._torch.disaggregation.native.transfer import ( + KVRecvTask, + SessionStatus, + TaskStatus, + _LogicalOutcomes, +) from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 from tensorrt_llm.bindings import LlmRequestState @@ -45,25 +49,12 @@ def _owes_a_report(outcome) -> bool: return outcome is None or outcome.reports_pending -class _StubTask: - """A receive task answers for its writers. - - Without that the handle falls back to the send-side probe, and the fan-in assertions would be - testing the wrong branch. Transferring means no writer has reported yet; a test about a settled - piece says so. - """ - - def __init__(self, status, reports_outstanding=True): - self.status = status - self._exception = None - self.reports_outstanding = reports_outstanding - - class _StubSession: """Only what the adapter reads, in states a real session reaches only under a race.""" def __init__(self): self._kv_tasks = [] + self._logical_outcomes = _LogicalOutcomes() self.status = SessionStatus.INIT self.cancelled_by_peer = False self.exception = None @@ -77,18 +68,27 @@ def __init__(self): def receive(self, chunk): if self.raise_on_receive is not None: if self.append_before_raising: - self._kv_tasks.append(_StubTask(TaskStatus.TRANSFERRING)) + self._append_task(chunk) raise self.raise_on_receive if not self.admits: # A closed or already terminal session returns without taking a task. return - self._kv_tasks.append(_StubTask(TaskStatus.TRANSFERRING)) + self._append_task(chunk) + + def _append_task(self, chunk): + task = KVRecvTask(42, chunk, len(self._kv_tasks), DisaggregatedParams(), aux_slot=None) + task.bind_logical_outcomes(self._logical_outcomes) + task.status = TaskStatus.TRANSFERRING + task.expected_transfers = 1 + self._kv_tasks.append(task) def fail_admission(self, error): + self._logical_outcomes.fail(error) self.status = SessionStatus.ERROR self.exception = error def cancel_local(self, by_peer=False): + self._logical_outcomes.cancel(by_peer) self.cancel_committed += 1 return self.cancel_committed == 1 @@ -160,7 +160,7 @@ def test_pieces_of_one_request_share_a_session(): def test_each_attempt_holds_its_own_piece(): worker, peer, first = _fetch_one() second = peer.fetch(_extent()) - worker.session._kv_tasks[0].status = TaskStatus.TRANSFERRED + worker.session._kv_tasks[0].complete() assert isinstance(first.poll(), Delivered) assert second.poll() is None @@ -187,8 +187,8 @@ def test_no_conclusion_while_the_piece_is_in_flight(): def test_delivery_reports_how_far_the_piece_reaches(): worker, _, attempt = _fetch_one(end=64) - worker.session._kv_tasks[0].status = TaskStatus.TRANSFERRED - worker.session._kv_tasks[0].reports_outstanding = False + worker.session._kv_tasks[0].complete() + worker.session._kv_tasks[0].note_writer_report(0, True) outcome = attempt.poll() assert isinstance(outcome, Delivered) assert outcome.token_end == 64 @@ -198,8 +198,7 @@ def test_delivery_reports_how_far_the_piece_reaches(): def test_failure_is_reported_before_every_writer_has_gone_quiet(): """The gate has to stay shut on a failure whose writers have not reported.""" worker, _, attempt = _fetch_one() - worker.session.status = SessionStatus.ERROR - worker.session.exception = RuntimeError("peer died") + worker.session.fail_admission(RuntimeError("peer died")) outcome = attempt.poll() assert isinstance(outcome, Failed) assert "peer died" in outcome.reason @@ -209,27 +208,25 @@ def test_failure_is_reported_before_every_writer_has_gone_quiet(): def test_a_settled_failure_opens_the_gate(): worker, _, attempt = _fetch_one() - worker.session.status = SessionStatus.ERROR - worker.session._kv_tasks[0].status = TaskStatus.ERROR - worker.session._kv_tasks[0].reports_outstanding = False + worker.session.fail_admission(RuntimeError("transfer failed")) + worker.session._kv_tasks[0].note_writer_report(0, False) assert _owes_a_report(attempt.poll()) is False def test_a_cancellation_says_who_asked(): worker, _, attempt = _fetch_one() - worker.session.status = SessionStatus.CANCELLED - worker.session.cancelled_by_peer = True + worker.session.cancel_local(by_peer=True) outcome = attempt.poll() assert isinstance(outcome, Cancelled) assert outcome.by_peer is True assert outcome.reports_pending is True -def test_failure_without_a_recorded_cause_still_says_something(): +def test_failure_preserves_its_recorded_cause(): worker, _, attempt = _fetch_one() - worker.session.status = SessionStatus.ERROR + worker.session.fail_admission(RuntimeError("transfer failed")) assert isinstance(attempt.poll(), Failed) - assert attempt.poll().reason + assert attempt.poll().reason == "transfer failed" # --------------------------------------------------------------------------- diff --git a/tests/unittest/disaggregated/test_task_handle.py b/tests/unittest/disaggregated/test_task_handle.py index e50bb60de682..626a5d619cf6 100644 --- a/tests/unittest/disaggregated/test_task_handle.py +++ b/tests/unittest/disaggregated/test_task_handle.py @@ -83,18 +83,23 @@ def _stub_receiver(): return receiver -def _receiving_from(writers: int, rid: int = 7) -> tuple[RxSession, TaskHandle]: +def _receiving_from( + writers: int, rid: int = 7, *, owns_transfers: bool = False +) -> tuple[RxSession, TaskHandle]: """One piece published to ``writers`` peers, and the handle the caller polls it through.""" + receiver = _stub_receiver() + receiver._enforce_physical_ownership = owns_transfers session = RxSession( request_id=rid, params=DisaggregatedParams(disagg_request_id=rid), - receiver=_stub_receiver(), + receiver=receiver, prompt_len=TOKENS, ) session.receive(_sole_piece()) task = session._kv_tasks[0] task.expected_transfers = writers - session.mark_transferring(task.slice_id) + cohort = set(range(writers)) if owns_transfers else None + session.mark_transferring(task.slice_id, cohort) return session, TaskHandle(session, task, TOKENS) @@ -255,20 +260,16 @@ def test_closing_leaves_the_tasks_where_they_were(): assert handle.poll() is None -def test_a_task_that_counts_nothing_falls_back_to_its_own_state(): - """With no count to read, the handle reads the state it does have.""" - session = SimpleNamespace( - status=SessionStatus.ERROR, - exception=RuntimeError("peer died"), - cancelled_by_peer=False, - ) - task = SimpleNamespace(status=TaskStatus.TRANSFERRING, _exception=None) - handle = TaskHandle(session, task, TOKENS) +def test_report_progress_remains_separate_from_a_committed_failure(): + session, handle = _receiving_from(1) + task = session._kv_tasks[0] + task.fail(RuntimeError("peer died")) assert handle.poll().reports_pending is True - task.status = TaskStatus.ERROR + _report(session, 0, AgentResult.SUCCESS) assert handle.poll().reports_pending is False + assert isinstance(handle.poll(), Failed) def test_a_send_task_owes_until_every_peer_write_is_done(): @@ -351,7 +352,7 @@ def test_a_siblings_failure_does_not_reopen_a_delivered_piece(): """Pieces share a session, so its verdict moves after one of them has already ended.""" task = KVRecvTask(9, _sole_piece(), 0, DisaggregatedParams(disagg_request_id=9), aux_slot=None) task.expected_transfers = 1 - task.status = TaskStatus.TRANSFERRED + task.complete() task.note_writer_report(0, True) session = SimpleNamespace( status=SessionStatus.TRANSFERRED, @@ -479,7 +480,7 @@ def test_a_piece_that_landed_is_delivered_even_if_the_request_was_cancelled(): ) task.expected_transfers = 1 task.note_writer_report(0, True) - task.status = TaskStatus.TRANSFERRED + task.complete() session = SimpleNamespace( status=SessionStatus.CANCELLED, exception=None, @@ -636,7 +637,9 @@ def test_a_send_session_reads_a_parked_cancel_as_the_peers_too(): ) assert session.status is SessionStatus.CANCELLED - task = SimpleNamespace(status=TaskStatus.TRANSFERRING, _exception=None) + sender.dispatch_task = MagicMock() + session.send(_sole_piece()) + task = session.kv_tasks[0] outcome = TaskHandle(session, task, TOKENS).poll() assert isinstance(outcome, Cancelled) assert outcome.by_peer is True @@ -713,9 +716,13 @@ def _sending_pieces(count: int = 1) -> tuple[Sender, TxSession]: @pytest.mark.parametrize("by_peer", [False, True]) -def test_queued_sender_abort_preserves_the_committed_cancellation(by_peer: bool) -> None: +@pytest.mark.parametrize("queued_status", [TaskStatus.INIT, TaskStatus.TRANSFERRING]) +def test_queued_sender_abort_preserves_the_committed_cancellation( + by_peer: bool, queued_status: TaskStatus +) -> None: sender, session = _sending_pieces() task = session.kv_tasks[0] + task.status = queued_status session.cancel_local(by_peer=by_peer) # Execute the real worker's pre-submission abort branch after cancellation won. empty = SimpleNamespace(size=0) @@ -781,3 +788,74 @@ def test_delivered_outcome_precedes_later_session_terminal_events(ending: str) - assert isinstance(handle.poll(), Delivered) assert handle.poll().token_end == TOKENS + + +def test_delivered_sender_piece_survives_a_direct_sibling_failure_without_polling() -> None: + _, session = _sending_pieces(2) + delivered, failed = session.kv_tasks + delivered.complete() + + failed.fail(RuntimeError("sibling failed later")) + + assert isinstance(TaskHandle(session, delivered, TOKENS).poll(), Delivered) + assert isinstance(TaskHandle(session, failed, TOKENS).poll(), Failed) + + +@pytest.mark.parametrize("owns_transfers", [False, True]) +@pytest.mark.parametrize("scatter_succeeded", [False, True]) +@pytest.mark.parametrize("cancel_first", [False, True]) +def test_scatter_callback_and_cancel_commit_in_event_order( + monkeypatch: pytest.MonkeyPatch, + owns_transfers: bool, + scatter_succeeded: bool, + cancel_first: bool, +) -> None: + from tensorrt_llm._torch.disaggregation.native import bounce + + session, handle = _receiving_from(1, owns_transfers=owns_transfers) + deferred = [] + monkeypatch.setattr(bounce, "scatter_write_result", lambda *args: deferred.append(args[-1])) + _report(session, 0, AgentResult.SUCCESS) + assert len(deferred) == 1 + assert handle.poll() is None + if owns_transfers: + assert session.resources_drained() is False + + ready, resume, finished = threading.Event(), threading.Event(), threading.Event() + + def finish_scatter() -> None: + ready.set() + if resume.wait(timeout=5): + deferred[0](scatter_succeeded) + finished.set() + + worker = threading.Thread(target=finish_scatter) + worker.start() + try: + assert ready.wait(timeout=5) + if cancel_first: + session.cancel_local(by_peer=True) + assert isinstance(handle.poll(), Cancelled) + if owns_transfers: + assert session.resources_drained() is False + resume.set() + assert finished.wait(timeout=5) + if not cancel_first: + session.cancel_local(by_peer=True) + finally: + resume.set() + worker.join(timeout=5) + assert not worker.is_alive() + + for observer in (handle, TaskHandle(session, session._kv_tasks[0], TOKENS)): + outcome = observer.poll() + if cancel_first: + assert isinstance(outcome, Cancelled) + assert outcome.by_peer is True + elif scatter_succeeded: + assert isinstance(outcome, Delivered) + else: + assert isinstance(outcome, Failed) + assert "bounce scatter failed" in outcome.reason + if owns_transfers: + assert session.resources_drained() is True diff --git a/tests/unittest/disaggregated/test_transceiver_bounded_polling.py b/tests/unittest/disaggregated/test_transceiver_bounded_polling.py index 0e4bb46daa84..e61a416baa68 100644 --- a/tests/unittest/disaggregated/test_transceiver_bounded_polling.py +++ b/tests/unittest/disaggregated/test_transceiver_bounded_polling.py @@ -36,6 +36,7 @@ TransferWorker, TransferWorkerConfig, TxSession, + _LogicalOutcomes, ) from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 from tensorrt_llm.bindings import LlmRequestState @@ -171,6 +172,7 @@ def _make_tx_session( deadline_monotonic_s: Optional[float] = None, ) -> TxSession: session = object.__new__(TxSession) + session._logical_outcomes = _LogicalOutcomes() session._timeout_s = timeout_s session._overall_timeout_s = None session._deadline_monotonic_s = deadline_monotonic_s diff --git a/tests/unittest/disaggregated/test_transfer_ownership_regressions.py b/tests/unittest/disaggregated/test_transfer_ownership_regressions.py index 583204fc260a..5c6119cd6a07 100644 --- a/tests/unittest/disaggregated/test_transfer_ownership_regressions.py +++ b/tests/unittest/disaggregated/test_transfer_ownership_regressions.py @@ -1601,12 +1601,14 @@ def test_terminal_sender_settles_known_unsubmitted_aux_peer_once() -> None: session._enforce_physical_ownership = True session._need_aux = True session._reported_aux_peer_ranks = set() + session._logical_outcomes = transfer_mod._LogicalOutcomes() session.kv_tasks = [] session.lock = threading.Lock() session._exception = None session._terminal_status = None session._closed = False session.aux_task = transfer_mod.AuxSendTask(params, slot=0) + session.aux_task.bind_logical_outcomes(session._logical_outcomes) assert session.aux_task.begin_physical_operation(active_info.instance_rank) session.set_exception("request failed before auxiliary submission") From f372733fe8906c0f4c1225d8ac5278b8627e9c9a Mon Sep 17 00:00:00 2001 From: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> Date: Thu, 17 Sep 2026 13:58:57 -0700 Subject: [PATCH 3/3] fix: retire retained KV operations after late backend completion Signed-off-by: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> --- .../_torch/disaggregation/native/transfer.py | 216 +++++++- .../test_transfer_late_settlement.py | 523 ++++++++++++++++++ .../test_transfer_ownership_regressions.py | 4 + 3 files changed, 733 insertions(+), 10 deletions(-) create mode 100644 tests/unittest/disaggregated/test_transfer_late_settlement.py diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index 06d4110d1e1c..aad4c8049c13 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -281,6 +281,7 @@ def __init__(self) -> None: self._writer_cohort: Optional[frozenset[int]] = None self._writer_results: dict[int, bool] = {} self._in_doubt_writers: set[int] = set() + self._settled_writers: set[int] = set() self._publication_failed = False self._local_completion_pending = False self._invalid_evidence = False @@ -323,7 +324,7 @@ def abort_publication(self, published_writers: set[int]) -> None: 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: + if not (self._writer_results.keys() | self._in_doubt_writers) <= published: self._invalid_evidence = True raise RuntimeError("terminal evidence came from an unpublished writer") self._expected_writers = len(published) @@ -355,10 +356,34 @@ def record_writer_in_doubt(self, peer_rank: int) -> bool: 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") - if peer_rank in self._in_doubt_writers: + if peer_rank in self._in_doubt_writers or peer_rank in self._settled_writers: return False self._in_doubt_writers.add(peer_rank) - self._invalid_evidence = True + return True + + def record_writer_settlement(self, peer_rank: int) -> bool: + """Accept physical DONE for a writer that previously reported IN_DOUBT. + + Ordinary FAILED is not this proof: it may describe a later, unsubmitted + chunk while the earlier ambiguous write is still touching the destination. + """ + with self._lock: + if self._expected_writers is None or ( + self._writer_cohort is not None and peer_rank not in self._writer_cohort + ): + self._invalid_evidence = True + raise RuntimeError(f"writer {peer_rank} settled outside the published cohort") + if peer_rank in self._settled_writers: + return False + if peer_rank not in self._in_doubt_writers: + self._invalid_evidence = True + raise RuntimeError(f"writer {peer_rank} settled without prior ambiguous evidence") + if self._writer_results.get(peer_rank) is True: + self._invalid_evidence = True + raise RuntimeError(f"writer {peer_rank} settled after contradictory success") + self._writer_results[peer_rank] = False + self._in_doubt_writers.remove(peer_rank) + self._settled_writers.add(peer_rank) return True def record_writer_result( @@ -378,6 +403,11 @@ def record_writer_result( 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") + if peer_rank in self._in_doubt_writers: + if succeeded: + self._invalid_evidence = True + raise RuntimeError(f"writer {peer_rank} reported success while in doubt") + return False, False previous = self._writer_results.get(peer_rank) if previous is not None: if previous != succeeded: @@ -408,6 +438,7 @@ def all_writers_reported(self) -> bool: return ( self._expected_writers is not None and len(self._writer_results) == self._expected_writers + and not self._in_doubt_writers ) @property @@ -420,6 +451,7 @@ def resources_drained(self) -> bool: writers_drained and not self._publication_pending and not self._local_completion_pending + and not self._in_doubt_writers and not self._invalid_evidence ) @@ -428,6 +460,7 @@ class AgentResult(Enum): SUCCESS = "SUCCESS" FAILED = "FAILED" IN_DOUBT = "IN_DOUBT" + FAILED_QUIESCED = "FAILED_QUIESCED" # KV_AGENT_RESULT prefix in one struct frame (was ascii frames serialized/parsed under the @@ -438,11 +471,13 @@ class AgentResult(Enum): AgentResult.SUCCESS: 0, AgentResult.FAILED: 1, AgentResult.IN_DOUBT: 2, + AgentResult.FAILED_QUIESCED: 3, } _AGENT_RESULT_BY_CODE = { 0: AgentResult.SUCCESS, 1: AgentResult.FAILED, 2: AgentResult.IN_DOUBT, + 3: AgentResult.FAILED_QUIESCED, } @@ -485,11 +520,11 @@ class _PhysicalOperationState(Enum): ADMITTED -> NOT_SUBMITTED ADMITTED -> SUBMITTING -> SUBMITTED -> BACKEND_DONE SUBMITTING -> IN_DOUBT - SUBMITTED -> IN_DOUBT + SUBMITTED -> IN_DOUBT -> BACKEND_DONE (same retained status reports DONE) Only NOT_SUBMITTED and BACKEND_DONE prove that the operation can no longer - access the source. IN_DOUBT deliberately has no retirement transition in - this bridge. Repeating a terminal transition is idempotent, and repeating + access the source. An IN_DOUBT operation without a retained backend status + cannot retire. Repeating a terminal transition is idempotent, and repeating IN_DOUBT preserves the retained backend evidence. """ @@ -515,6 +550,15 @@ class _PhysicalOperation: status: Optional[object] = None +@dataclass +class _PendingSettlement: + write_meta: WriteMeta + initial_report: Optional[list[bytes]] + send_slot_id: Optional[int] = None + backend_done: bool = False + report_error_logged: bool = False + + class SendTaskBase(_LogicalTask): def __init__(self, params: DisaggregatedParams): super().__init__() @@ -627,6 +671,39 @@ def retire_backend_done_physical_operation(self, peer_rank: int) -> None: operation.status = None operation.state = _PhysicalOperationState.BACKEND_DONE + def poll_in_doubt_physical_operation(self, peer_rank: int) -> bool: + """Retire once, only after a fresh DONE query on the retained status. + + Polling does not change the task's logical outcome. Keep strong local + roots across the query, and reject a result if its operation changed. + """ + with self._physical_lock: + operation = self._physical_operations.get(peer_rank) + if operation is None or operation.state is not _PhysicalOperationState.IN_DOUBT: + return False + request, status = operation.request, operation.status + if status is None: + return False + try: + completed = status.is_completed() + except Exception: + # A backend query failure is not evidence that its accessors stopped. + return False + if completed is not True: + return False + with self._physical_lock: + if ( + self._physical_operations.get(peer_rank) is not operation + or operation.state is not _PhysicalOperationState.IN_DOUBT + or operation.request is not request + or operation.status is not status + ): + return False + operation.request = None + operation.status = None + operation.state = _PhysicalOperationState.BACKEND_DONE + return True + def has_started_physical_operation(self, peer_rank: int) -> bool: with self._physical_lock: return peer_rank in self._physical_operations @@ -727,6 +804,11 @@ def __init__( self._send_task_queues: List[queue.Queue] = [ queue.Queue() for _ in range(self._num_threads) ] + # Each dictionary belongs to its worker, preserving that peer stream's + # result ordering and ZMQ socket affinity even after the first failure. + self._pending_settlements: list[dict[tuple[SendTaskBase, int], _PendingSettlement]] = [ + {} for _ in range(self._num_threads) + ] self._worker_threads: List[threading.Thread] = [ threading.Thread(target=self._process_task_queue, args=(i,), daemon=True) for i in range(self._num_threads) @@ -965,7 +1047,22 @@ def _process_task_queue(self, thread_idx: int): task_queue = self._send_task_queues[thread_idx] try: while True: - write_meta = task_queue.get() + if self._pending_settlements[thread_idx]: + self._poll_in_doubt_transfers(thread_idx) + if any( + pending.initial_report is not None + for pending in self._pending_settlements[thread_idx].values() + ): + # A later queued FAILED may describe only an unsubmitted + # chunk. It must never overtake the earlier IN_DOUBT. + time.sleep(0.01) + continue + try: + write_meta = task_queue.get(timeout=0.01) + except queue.Empty: + continue + else: + write_meta = task_queue.get() if write_meta is None: break if isinstance(write_meta, tuple): @@ -1005,6 +1102,69 @@ def _process_task_queue(self, thread_idx: int): ) dealers.clear() + def _retain_in_doubt_transfer( + self, + write_meta: WriteMeta, + initial_report: list[bytes], + send_slot_id: Optional[int] = None, + ) -> None: + thread_idx = hash((write_meta.unique_rid, write_meta.peer_rank)) % self._num_threads + key = (write_meta.task, write_meta.peer_rank) + if key in self._pending_settlements[thread_idx]: + return + pending = _PendingSettlement(write_meta, initial_report, send_slot_id) + self._pending_settlements[thread_idx][key] = pending + try: + self._get_result_dealer(write_meta.peer_endpoint).send(initial_report) + except Exception as error: + logger.warning(f"Failed to report ambiguous transfer; retaining evidence: {error}") + pending.report_error_logged = True + else: + pending.initial_report = None + + def _poll_in_doubt_transfers(self, thread_idx: int) -> None: + pending_transfers = self._pending_settlements[thread_idx] + for key, pending in list(pending_transfers.items()): + meta = pending.write_meta + try: + dealer = self._get_result_dealer(meta.peer_endpoint) + if pending.initial_report is not None: + dealer.send(pending.initial_report) + pending.initial_report = None + if not pending.backend_done: + if not meta.task.poll_in_doubt_physical_operation(meta.peer_rank): + continue + pending.backend_done = True + if pending.send_slot_id is not None: + self._bounce.release_send(pending.send_slot_id) + pending.send_slot_id = None + if meta.meta_type == WriteMetaType.AUX: + message = _make_aux_result_msg( + self._instance_rank, meta.unique_rid, AgentResult.FAILED_QUIESCED + ) + else: + message = _make_kv_result_msg( + self._instance_rank, + meta.unique_rid, + meta.receiver_slice_id, + True, + AgentResult.FAILED_QUIESCED, + ) + dealer.send(message) + except Exception as error: + # Sending may have escaped before raising. A duplicate settlement + # is idempotent; losing the only remaining report is not safe. + if not pending.report_error_logged: + logger.warning(f"Failed to report physical settlement; will retry: {error}") + pending.report_error_logged = True + continue + with meta.task.lock: + if meta.meta_type == WriteMetaType.AUX: + meta.task._transfer_count += 1 + else: + meta.task.transferred_count += 1 + del pending_transfers[key] + @staticmethod @nvtx_range("_make_agent_request") def _make_agent_request(write_meta: WriteMeta, device_id: int) -> "TransferRequest": @@ -1189,10 +1349,10 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): transfer_size=transfer_size, tail=tail, ) - self._get_result_dealer(write_meta.peer_endpoint).send(result_msg) - if agent_result == AgentResult.IN_DOUBT: + self._retain_in_doubt_transfer(write_meta, result_msg, send_slot_id) return + self._get_result_dealer(write_meta.peer_endpoint).send(result_msg) if timer: timer.record_task_end(write_meta.peer_rank) @@ -1287,6 +1447,12 @@ def _deliver_aux_to_agent(self, write_meta: WriteMeta): # claimant. Keep the claim at the send boundary so future cleanup-path # changes cannot publish contradictory evidence. if not owned or session._claim_aux_terminal_result(write_meta.peer_rank): + if agent_result == AgentResult.IN_DOUBT: + self._retain_in_doubt_transfer( + write_meta, + _make_aux_result_msg(self._instance_rank, write_meta.unique_rid, agent_result), + ) + return self._get_result_dealer(write_meta.peer_endpoint).send( _make_aux_result_msg( self._instance_rank, @@ -1780,7 +1946,9 @@ def _route_result_messages_to_receiver( *, defer_to_worker: bool, ) -> None: - if defer_to_worker: + if defer_to_worker or self._enforce_physical_ownership: + # Listener-side rejection must share the worker's ordered stream; + # it cannot overtake an ambiguous write's pending IN_DOUBT report. thread_idx = hash((info.unique_rid, info.instance_rank)) % self._num_threads for message in messages: self._send_task_queues[thread_idx].put((endpoint, message)) @@ -1847,6 +2015,8 @@ def shutdown(self): return if self._enforce_physical_ownership and self._sessions: raise RuntimeError("Sender refuses shutdown while transfer ownership is active") + if any(self._pending_settlements): + raise RuntimeError("Sender refuses shutdown while settlement reports are pending") self._shutdown_requested = True self._messenger.stop() @@ -2410,6 +2580,9 @@ def cancel_unpublished(self) -> bool: def record_writer_in_doubt(self, peer_rank: int) -> bool: return self._get_physical_owner().record_writer_in_doubt(peer_rank) + def record_writer_settlement(self, peer_rank: int) -> bool: + return self._get_physical_owner().record_writer_settlement(peer_rank) + def record_writer_result( self, peer_rank: int, @@ -3267,6 +3440,7 @@ def process_kv_agent_result( AgentResult.SUCCESS, AgentResult.FAILED, AgentResult.IN_DOUBT, + AgentResult.FAILED_QUIESCED, ): raise ValueError( f"Session {self.request_id} received unknown task status: {status.value}" @@ -3291,6 +3465,19 @@ def process_kv_agent_result( task.fail(error) self._record_ownership_evidence_error(error) return + if status == AgentResult.FAILED_QUIESCED: + if not self._enforce_physical_ownership: + raise RuntimeError("received physical settlement without ownership enabled") + try: + accepted = task.record_writer_settlement(peer_rank) + except Exception as error: + self._record_ownership_evidence_error(error) + raise + if accepted: + self._receiver._bounce.record_failure( + (self.disagg_request_id, task.slice_id), peer_rank + ) + return if self._enforce_physical_ownership: if status == AgentResult.FAILED or is_last_slice: try: @@ -3432,6 +3619,15 @@ def process_aux_agent_result(self, peer_rank: int, status: AgentResult): self._aux_status = TaskStatus.ERROR self._record_ownership_evidence_error(error) return + if status == AgentResult.FAILED_QUIESCED: + if self._aux_physical_owner is None: + raise RuntimeError("received auxiliary settlement without ownership enabled") + try: + self._aux_physical_owner.record_writer_settlement(peer_rank) + except Exception as error: + self._record_ownership_evidence_error(error) + raise + return if self._aux_physical_owner is not None: try: accepted, all_succeeded = self._aux_physical_owner.record_writer_result( diff --git a/tests/unittest/disaggregated/test_transfer_late_settlement.py b/tests/unittest/disaggregated/test_transfer_late_settlement.py new file mode 100644 index 000000000000..13ce118b2e81 --- /dev/null +++ b/tests/unittest/disaggregated/test_transfer_late_settlement.py @@ -0,0 +1,523 @@ +# 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. +"""Late physical evidence must not turn an unsuccessful transfer into delivery.""" + +import queue +import threading +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._torch.disaggregation.base import Cancelled, Chunk, Failed, TokenRange +from tensorrt_llm._torch.disaggregation.base.transfer import SessionStatus +from tensorrt_llm._torch.disaggregation.native.handle import TaskHandle +from tensorrt_llm.disaggregated_params import DisaggregatedParams, DisaggScheduleStyle + + +def _params() -> DisaggregatedParams: + return DisaggregatedParams( + disagg_request_id=401, schedule_style=DisaggScheduleStyle.GENERATION_FIRST + ) + + +def _chunk() -> Chunk: + return Chunk([], [], TokenRange(0, 16), True) + + +def _in_doubt_task(status: object | None) -> transfer_mod.SendTaskBase: + task = transfer_mod.SendTaskBase(_params()) + assert task.begin_physical_operation(7) + task.begin_backend_submission(7, object()) + if status is not None: + task.record_backend_submission(7, status) + task.mark_physical_operation_in_doubt(7) + return task + + +def _owner(cohort: set[int]) -> transfer_mod._ReceiveOperationOwner: + owner = transfer_mod._ReceiveOperationOwner() + owner.begin_publication() + owner.seal_writer_cohort(len(cohort), cohort) + owner.finish_publication() + return owner + + +def _sender() -> transfer_mod.Sender: + sender = object.__new__(transfer_mod.Sender) + sender._enforce_physical_ownership = True + sender._sessions_lock, sender._sessions = threading.Lock(), {} + sender._shutdown = sender._shutdown_requested = False + sender._ownership_poisoned, sender._ownership_poison_lock = None, threading.Lock() + sender._instance_rank = 7 + sender._device_id = 0 + sender._num_threads = 1 + sender._pending_settlements = [{}] + sender._bounce = Mock() + return sender + + +@pytest.mark.cpu_only +@pytest.mark.parametrize("query", [None, False, "DONE", RuntimeError("query failed")]) +def test_late_settlement_requires_positive_retained_status(query: object) -> None: + status = None if query is None else Mock() + if status is not None: + if isinstance(query, Exception): + status.is_completed.side_effect = query + else: + status.is_completed.return_value = query + task = _in_doubt_task(status) + operation = task._physical_operations[7] + request = operation.request + + assert not task.poll_in_doubt_physical_operation(7) + assert operation.state is transfer_mod._PhysicalOperationState.IN_DOUBT + assert operation.request is request + assert operation.status is status + assert not task.resources_drained + + +@pytest.mark.cpu_only +def test_late_settlement_rejects_replaced_status() -> None: + status = Mock() + task = _in_doubt_task(status) + replacement = Mock() + + def replace() -> bool: + task._physical_operations[7].status = replacement + return True + + status.is_completed.side_effect = replace + assert not task.poll_in_doubt_physical_operation(7) + assert not task.resources_drained + assert task._physical_operations[7].status is replacement + + +@pytest.mark.cpu_only +def test_concurrent_late_done_retires_once_without_changing_failure() -> None: + barrier = threading.Barrier(2, timeout=10) + status = Mock() + + def done() -> bool: + barrier.wait() + return True + + status.is_completed.side_effect = done + task = _in_doubt_task(status) + original_error = RuntimeError("original failure") + task.fail(original_error) + outcomes = [False, False] + + def poll(index: int) -> None: + outcomes[index] = task.poll_in_doubt_physical_operation(7) + + threads = [threading.Thread(target=poll, args=(i,)) for i in range(2)] + for thread in threads: + thread.start() + for thread in threads: + thread.join(timeout=10) + assert not thread.is_alive() + assert sum(outcomes) == 1 + assert task.resources_drained + assert task.status is transfer_mod.TaskStatus.ERROR + assert task._exception is original_error + assert task._physical_operations[7].request is None + assert task._physical_operations[7].status is None + assert not task.poll_in_doubt_physical_operation(7) + + +@pytest.mark.cpu_only +def test_late_done_does_not_retire_active_sibling() -> None: + task = _in_doubt_task(Mock(is_completed=Mock(return_value=True))) + assert task.begin_physical_operation(8) + assert task.poll_in_doubt_physical_operation(7) + assert not task.resources_drained + task.retire_unsubmitted_physical_operation(8) + assert task.resources_drained + + +@pytest.mark.cpu_only +@pytest.mark.parametrize("count_only", [False, True]) +def test_only_explicit_settlement_closes_ambiguous_writer(count_only: bool) -> None: + owner = transfer_mod._ReceiveOperationOwner() + owner.begin_publication() + # Gen-first ADP counts the immutable, no-retry selected group's writers. + owner.seal_writer_cohort(2, None if count_only else {7, 8}) + owner.finish_publication() + assert owner.record_writer_in_doubt(7) + assert owner.record_writer_result(7, False, wait_for_local_completion=False) == (False, False) + assert not owner.all_writers_reported + assert not owner.resources_drained + assert owner.record_writer_settlement(7) + assert not owner.resources_drained + assert not owner.record_writer_settlement(7) + assert not owner.record_writer_in_doubt(7) + owner.record_writer_result(8, False, wait_for_local_completion=False) + assert owner.resources_drained + + +@pytest.mark.cpu_only +@pytest.mark.parametrize("foreign", [False, True]) +def test_reordered_or_foreign_settlement_stays_invalid(foreign: bool) -> None: + owner = _owner({7}) + with pytest.raises(RuntimeError): + owner.record_writer_settlement(8 if foreign else 7) + owner.record_writer_in_doubt(7) + assert owner.record_writer_settlement(7) + assert not owner.resources_drained + + +@pytest.mark.cpu_only +def test_settlement_does_not_clear_contradictory_evidence() -> None: + owner = _owner({7}) + owner.record_writer_in_doubt(7) + with pytest.raises(RuntimeError, match="success while in doubt"): + owner.record_writer_result(7, True, wait_for_local_completion=True) + assert owner.record_writer_settlement(7) + assert not owner.resources_drained + + +@pytest.mark.cpu_only +def test_abort_publication_cannot_exclude_ambiguous_writer() -> None: + owner = _owner({7, 8}) + owner.record_writer_in_doubt(7) + with pytest.raises(RuntimeError, match="unpublished writer"): + owner.abort_publication({8}) + assert owner.record_writer_settlement(7) + owner.record_writer_result(8, False, wait_for_local_completion=False) + assert not owner.resources_drained + + +@pytest.mark.cpu_only +def test_settlement_does_not_clear_local_completion_or_publication() -> None: + owner = _owner({7}) + owner._local_completion_pending = True + owner._publication_pending = True + owner.record_writer_in_doubt(7) + owner.record_writer_settlement(7) + assert not owner.resources_drained + owner.finish_local_completion() + assert not owner.resources_drained + owner.finish_publication() + assert owner.resources_drained + + +@pytest.mark.cpu_only +@pytest.mark.parametrize("auxiliary", [False, True]) +@pytest.mark.parametrize("poll_early", [False, True]) +def test_sender_late_done_reports_physical_failure_only( + monkeypatch: pytest.MonkeyPatch, auxiliary: bool, poll_early: bool +) -> None: + sender = _sender() + task = ( + transfer_mod.AuxSendTask(_params(), slot=0) + if auxiliary + else transfer_mod.KVSendTask(_chunk(), _params(), slice_id=0) + ) + session = SimpleNamespace( + kv_tasks=[task], + aux_task=task, + status=SessionStatus.READY, + exception=None, + lock=threading.Lock(), + _claim_aux_terminal_result=Mock(return_value=True), + set_exception=lambda reason: task.fail(RuntimeError(reason)), + ) + sender._sessions[401] = session + if not auxiliary: + task.expected_transfers = 1 + status = Mock(wait=Mock(return_value=False), is_completed=Mock(return_value=False)) + status.last_status_str.return_value = "ERROR" + sender._agent = SimpleNamespace(name="nixl", submit_transfer_requests=lambda request: status) + request = SimpleNamespace(op="WRITE", remote_name="gen") + monkeypatch.setattr(transfer_mod.Sender, "_make_agent_request", Mock(return_value=request)) + dealer = Mock() + sender._get_result_dealer = Mock(return_value=dealer) + meta = transfer_mod.WriteMeta( + task=task, + expected_transfers=1, + peer_name="gen", + peer_rank=7, + peer_endpoint="receiver", + unique_rid=401, + src_ptrs=np.array([0x1000]), + dst_ptrs=np.array([0x2000]), + sizes=np.array([0x100]), + slice_id=0, + receiver_slice_id=3, + is_last_slice=False, + meta_type=transfer_mod.WriteMetaType.AUX if auxiliary else transfer_mod.WriteMetaType.KV, + ) + assert task.begin_physical_operation(7) + if auxiliary: + sender._deliver_aux_to_agent(meta) + else: + sender._deliver_kv_to_agent(meta) + original_error = task._exception + handle = None if auxiliary else TaskHandle(session, task, token_end=16) + if poll_early and handle is not None: + assert isinstance(handle.poll(), Failed) + assert not task.resources_drained + sender._poll_in_doubt_transfers(0) + assert dealer.send.call_count == 1 + status.is_completed.return_value = True + sender._poll_in_doubt_transfers(0) + sender._poll_in_doubt_transfers(0) + + assert task.resources_drained + assert task.status is transfer_mod.TaskStatus.ERROR + assert task._exception is original_error + if handle is not None: + assert isinstance(handle.poll(), Failed) + assert sender._ownership_poisoned is not None + assert not sender._pending_settlements[0] + assert dealer.send.call_count == 2 + messages = [call.args[0] for call in dealer.send.call_args_list] + if auxiliary: + assert [message[-1].decode() for message in messages] == ["IN_DOUBT", "FAILED_QUIESCED"] + assert task._transfer_count == 1 + session._claim_aux_terminal_result.assert_called_once_with(7) + else: + decoded = [transfer_mod._KV_RESULT_PREFIX.unpack(message[1]) for message in messages] + assert [transfer_mod._AGENT_RESULT_BY_CODE[item[4]] for item in decoded] == [ + transfer_mod.AgentResult.IN_DOUBT, + transfer_mod.AgentResult.FAILED_QUIESCED, + ] + assert decoded[1][:3] == (7, 401, 3) + assert task.transferred_count == 1 + + +@pytest.mark.cpu_only +@pytest.mark.parametrize("auxiliary", [False, True]) +def test_receiver_late_settlement_retains_logical_failure_and_quarantine(auxiliary: bool) -> None: + receiver = object.__new__(transfer_mod.Receiver) + receiver._enforce_physical_ownership = True + receiver._sessions_lock, receiver._sessions = threading.Lock(), {} + receiver._pre_cancelled_rids = {} + receiver._shutdown = False + receiver._ownership_admission_lock = threading.Lock() + receiver._ownership_poisoned = None + receiver._bounce = Mock() + receiver._bounce.is_bounced.return_value = False + session = transfer_mod.RxSession(request_id=401, params=_params(), receiver=receiver) + task = session.prepare_receive(_chunk()) + assert task is not None + task.expected_transfers = 1 + assert session.try_begin_transfer(task.slice_id, set(), writer_cohort={7}) + result = transfer_mod.AgentResult + if auxiliary: + session.process_kv_agent_result(7, 0, True, result.FAILED) + else: + session.process_aux_agent_result(7, result.FAILED) + + def report(status: transfer_mod.AgentResult) -> None: + if auxiliary: + session.process_aux_agent_result(7, status) + else: + session.process_kv_agent_result(7, 0, True, status) + + report(result.IN_DOUBT) + error = session.exception + assert not session.resources_drained() + report(result.FAILED) + assert not session.resources_drained() + report(result.FAILED_QUIESCED) + report(result.FAILED_QUIESCED) + assert session.resources_drained() + assert session.status is SessionStatus.ERROR + assert session.exception is error + assert receiver._ownership_poisoned is not None + receiver._bounce.record_failure.assert_called_once_with((401, 0), 7) + assert session.close() + + +@pytest.mark.cpu_only +def test_settlement_retries_in_order_without_releasing_twice() -> None: + sender = _sender() + task = transfer_mod.KVSendTask(_chunk(), _params(), slice_id=0) + status = Mock(is_completed=Mock(return_value=True)) + assert task.begin_physical_operation(7) + task.begin_backend_submission(7, object()) + task.record_backend_submission(7, status) + task.mark_physical_operation_in_doubt(7) + task.fail(RuntimeError("original failure")) + meta = transfer_mod.WriteMeta( + task=task, + expected_transfers=1, + peer_name="gen", + peer_rank=7, + peer_endpoint="receiver", + unique_rid=401, + src_ptrs=np.array([]), + dst_ptrs=np.array([]), + sizes=np.array([]), + receiver_slice_id=3, + ) + initial = transfer_mod._make_kv_result_msg(7, 401, 3, False, transfer_mod.AgentResult.IN_DOUBT) + dealer = Mock() + dealer.send.side_effect = [RuntimeError("initial send"), None, RuntimeError("late send"), None] + sender._get_result_dealer = Mock(return_value=dealer) + sender._retain_in_doubt_transfer(meta, initial, send_slot_id=12) + status.is_completed.assert_not_called() + sender._poll_in_doubt_transfers(0) + assert task.resources_drained + assert sender._pending_settlements[0] + sender._bounce.release_send.assert_called_once_with(12) + with pytest.raises(RuntimeError, match="settlement reports are pending"): + sender.shutdown() + sender._poll_in_doubt_transfers(0) + assert not sender._pending_settlements[0] + sender._bounce.release_send.assert_called_once_with(12) + status.is_completed.assert_called_once() + assert task.transferred_count == 1 + codes = [ + transfer_mod._KV_RESULT_PREFIX.unpack(call.args[0][1])[4] + for call in dealer.send.call_args_list + ] + assert codes == [2, 2, 3, 3] + sender._shutdown = True + + +@pytest.mark.cpu_only +@pytest.mark.parametrize("direction", ["send", "receive"]) +@pytest.mark.parametrize("terminal_kind", ["failed", "local_cancel", "peer_cancel"]) +@pytest.mark.parametrize("poll_early", [False, True]) +def test_committed_session_outcome_survives_late_settlement( + direction: str, terminal_kind: str, poll_early: bool +) -> None: + """Integration seam with the event-time logical-outcome commit change.""" + status = Mock(is_completed=Mock(return_value=True)) + if direction == "send": + sender = Mock() + sender._enforce_physical_ownership = True + sender._get_req_info.return_value = {} + session = transfer_mod.TxSession(request_id=401, params=_params(), sender=sender) + session.send(_chunk()) + task = session.kv_tasks[0] + task.expected_transfers = 1 + task.status = transfer_mod.TaskStatus.TRANSFERRING + assert task.begin_physical_operation(7) + task.begin_backend_submission(7, object()) + task.record_backend_submission(7, status) + else: + receiver = Mock() + receiver._enforce_physical_ownership = True + receiver._get_ownership_admission_lock.return_value = threading.Lock() + receiver._bounce.is_bounced.return_value = False + session = transfer_mod.RxSession(request_id=401, params=_params(), receiver=receiver) + task = session.prepare_receive(_chunk()) + assert task is not None + task.expected_transfers = 1 + assert session.try_begin_transfer(task.slice_id, set(), writer_cohort={7}) + cancelled = terminal_kind != "failed" + if cancelled: + assert session.cancel_local(by_peer=terminal_kind == "peer_cancel") + if direction == "send": + task.mark_physical_operation_in_doubt(7) + task.fail(RuntimeError("backend outcome ambiguous")) + else: + session.process_kv_agent_result(7, 0, True, transfer_mod.AgentResult.IN_DOUBT) + handle = TaskHandle(session, task, token_end=16) + expected_type = Cancelled if cancelled else Failed + if poll_early: + early_outcome = handle.poll() + assert isinstance(early_outcome, expected_type) + if cancelled: + assert early_outcome.by_peer is (terminal_kind == "peer_cancel") + assert not task.resources_drained + if direction == "send": + assert task.poll_in_doubt_physical_operation(7) + else: + session.process_kv_agent_result(7, 0, True, transfer_mod.AgentResult.FAILED_QUIESCED) + session.process_aux_agent_result(7, transfer_mod.AgentResult.FAILED) + assert task.resources_drained + for outcome in (handle.poll(), TaskHandle(session, task, token_end=16).poll()): + assert isinstance(outcome, expected_type) + if cancelled: + assert outcome.by_peer is (terminal_kind == "peer_cancel") + assert session.close() + + +@pytest.mark.cpu_only +def test_listener_failure_uses_worker_stream_in_ownership_bridge() -> None: + sender = _sender() + sender._send_task_queues = [queue.Queue()] + sender._get_or_connect_dealer = Mock() + info = SimpleNamespace(unique_rid=401, instance_rank=7) + message = [transfer_mod.MessageType.KV_AGENT_RESULT, b"failed"] + sender._route_result_messages_to_receiver(info, "receiver", [message], defer_to_worker=False) + sender._get_or_connect_dealer.assert_not_called() + assert sender._send_task_queues[0].get_nowait() == ("receiver", message) + sender._shutdown = True + + +@pytest.mark.cpu_only +@pytest.mark.parametrize("initial_send_fails", [False, True]) +def test_worker_polls_retained_status_when_queue_is_idle( + monkeypatch: pytest.MonkeyPatch, initial_send_fails: bool +) -> None: + sender = _sender() + sender._thread_local = threading.local() + task = transfer_mod.KVSendTask(_chunk(), _params(), slice_id=0) + status = Mock(is_completed=Mock(side_effect=[False, True])) + assert task.begin_physical_operation(7) + task.begin_backend_submission(7, object()) + task.record_backend_submission(7, status) + task.mark_physical_operation_in_doubt(7) + task.fail(RuntimeError("original failure")) + meta = transfer_mod.WriteMeta( + task=task, + expected_transfers=1, + peer_name="gen", + peer_rank=7, + peer_endpoint="receiver", + unique_rid=401, + src_ptrs=np.array([]), + dst_ptrs=np.array([]), + sizes=np.array([]), + ) + work_queue = Mock(get=Mock(side_effect=[queue.Empty, None])) + dealer = Mock() + if initial_send_fails: + attempts = 0 + + def send(_message: list[bytes]) -> None: + nonlocal attempts + attempts += 1 + if attempts <= 2: + # Even later safely rejected work must not overtake IN_DOUBT. + work_queue.get.assert_not_called() + raise RuntimeError("initial send failed") + + dealer.send.side_effect = send + sender._get_result_dealer = Mock(return_value=dealer) + sender._retain_in_doubt_transfer( + meta, + transfer_mod._make_kv_result_msg(7, 401, 0, False, transfer_mod.AgentResult.IN_DOUBT), + ) + sender._send_task_queues = [work_queue] + monkeypatch.setattr(transfer_mod.time, "sleep", Mock()) + monkeypatch.setattr(transfer_mod.torch.cuda, "set_device", Mock()) + monkeypatch.setattr(transfer_mod.cudart, "cudaSetDevice", Mock(return_value=0)) + monkeypatch.setattr(transfer_mod, "CUASSERT", Mock()) + sender._process_task_queue(0) + assert task.resources_drained + assert not sender._pending_settlements[0] + assert dealer.send.call_count == (4 if initial_send_fails else 2) + assert status.is_completed.call_count == 2 + assert task.status is transfer_mod.TaskStatus.ERROR + sender._shutdown = True diff --git a/tests/unittest/disaggregated/test_transfer_ownership_regressions.py b/tests/unittest/disaggregated/test_transfer_ownership_regressions.py index 5c6119cd6a07..9e907a80c457 100644 --- a/tests/unittest/disaggregated/test_transfer_ownership_regressions.py +++ b/tests/unittest/disaggregated/test_transfer_ownership_regressions.py @@ -1341,6 +1341,9 @@ def _make_owned_sender() -> transfer_mod.Sender: sender._ownership_poisoned, sender._ownership_poison_lock = None, threading.Lock() sender._loaded_remote_agents_lock, sender._loaded_remote_agents = threading.Lock(), set() sender._instance_rank = 0 + sender._num_threads = 1 + sender._pending_settlements = [{}] + sender._send_task_queues = [queue.Queue()] return sender @@ -1401,6 +1404,7 @@ def test_pre_cancelled_sender_settles_saved_generation_first_request(monkeypatch def test_sender_failed_result_routes_messages_directly_in_order(monkeypatch) -> None: rid = 98 sender = object.__new__(transfer_mod.Sender) + sender._enforce_physical_ownership = False sender._instance_rank = 5 sender._registrar = SimpleNamespace( get_peer_rank_info=Mock(return_value=SimpleNamespace(self_endpoint="receiver"))