diff --git a/scripts/disagg_transfer_diagnostics.py b/scripts/disagg_transfer_diagnostics.py new file mode 100644 index 000000000000..3ece4ac128ee --- /dev/null +++ b/scripts/disagg_transfer_diagnostics.py @@ -0,0 +1,1377 @@ +#!/usr/bin/env python3 +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Summarize opt-in disaggregated KV-transfer diagnostic events. + +The runtime writes compact JSON objects after ``[DISAGG_TRANSFER_DIAG]``. This +tool tolerates unrelated and malformed log lines. Requests are correlated across +processes only with a shared run UUID; otherwise analysis stays process-local. +Durations require matching run/process UUIDs, host, and PID. Legacy records +without a process UUID remain readable but cannot establish safe timing. + +Startup ``diagnostic_capabilities`` records describe which event groups each +executor/transceiver can emit. Unsupported boundaries are not missing events; +logs without this metadata retain identity-safe measured pairs but have +unassessed expectations. +""" + +from __future__ import annotations + +import argparse +import fileinput +import json +import math +import statistics +import sys +from bisect import bisect_left +from collections import Counter, defaultdict +from dataclasses import dataclass +from pathlib import Path +from typing import Iterable, Sequence +from uuid import UUID + +DIAGNOSTICS_LOG_PREFIX = "[DISAGG_TRANSFER_DIAG] " +_SUPPORTED_SCHEMA_VERSION = 1 +_GEN_KV_ADMISSION_EVENT = "gen_kv_admission_result" +_CAPABILITY_FIELDS = ("executor_events", "scheduler_kv_admission_events", "python_transfer_events") +_TIMELINE_FIELDS = ( + "event", + "side", + "wall_ns", + "monotonic_ns", + "host", + "pid", + "run_uuid", + "process_uuid", + "run_uuid_status", + "instance", + "rank", + "tp_rank", + "pp_rank", + "cp_rank", + "dp_rank", + "local_request_id", + "outcome", + "policy", + "slice_id", + "peer_rank", + "peer_instance", + "receiver_slice_id", + "is_last_slice", + "worker_queue_index", + "transfer_entries", + "ownership_enabled", + "source_kv_reuse_block_count", + "state", + "reason", + "timeout_ms", + "timeout_expected", + "timeout_owner", + "timer_start_monotonic_ns", + "elapsed_ms", + "cancellation_requested", + "session_status", + "session_found", + "resources_drained", + "transfer_bytes", + "expected_receivers", + "expected_writers", + "writer_cohort_known", + "prompt_tokens", + "tokens_per_block", + "cache_present", + "capacity_tokens", + "history_tokens", + "capacity_block_equivalent", + "request_blocks", + "active_transfer_blocks", + "admitted_transfer_blocks", + "transfer_block_budget", + "legacy_budget_outcome", + "legacy_active_transfer_blocks", + "legacy_admitted_transfer_blocks", + "legacy_limited_by_budget", + "source_kv_request_owned", + "source_kv_reuse_pinned", + "init_requests", + "transfers_in_progress", + "transfers_complete", + "kv_admitted_this_iteration", + "decode_requests", + "kv_pool_max_blocks", + "kv_pool_free_blocks", + "kv_pool_used_blocks", + "index_free_slots", + "dropped_events", + "capability_schema_version", + "transceiver_runtime", + *_CAPABILITY_FIELDS, +) + +Event = dict[str, object] +ClockDomain = tuple[str, int, str, str | None] +Participant = tuple[str | None, str | None, int | None, int | None, str | None, str | None] +CapabilityScope = tuple[str, int, str, str | None, int | None] + +_STRING_FIELDS = frozenset( + { + "event", + "host", + "instance", + "legacy_budget_outcome", + "outcome", + "peer_instance", + "policy", + "reason", + "session_status", + "side", + "state", + "timeout_owner", + } +) +_IDENTIFIER_FIELDS = frozenset({"request_id", "local_request_id"}) +_NONNEGATIVE_INTEGER_FIELDS = frozenset( + { + "active_transfer_blocks", + "admitted_transfer_blocks", + "capacity_block_equivalent", + "capacity_tokens", + "cp_rank", + "decode_requests", + "dp_rank", + "dropped_events", + "expected_receivers", + "expected_writers", + "history_tokens", + "index_free_slots", + "init_requests", + "kv_admitted_this_iteration", + "kv_pool_free_blocks", + "kv_pool_max_blocks", + "kv_pool_used_blocks", + "legacy_active_transfer_blocks", + "legacy_admitted_transfer_blocks", + "monotonic_ns", + "peer_rank", + "pid", + "pp_rank", + "prompt_tokens", + "rank", + "receiver_slice_id", + "request_blocks", + "slice_id", + "source_kv_reuse_block_count", + "timer_start_monotonic_ns", + "timeout_ms", + "tokens_per_block", + "tp_rank", + "transfer_block_budget", + "transfer_bytes", + "transfer_entries", + "transfers_complete", + "transfers_in_progress", + "wall_ns", + "worker_queue_index", + } +) +_NONNEGATIVE_NUMBER_FIELDS = frozenset({"elapsed_ms"}) +_BOOLEAN_FIELDS = frozenset( + { + "cache_present", + "cancellation_requested", + "is_last_slice", + "legacy_limited_by_budget", + "ownership_enabled", + "resources_drained", + "session_found", + "source_kv_request_owned", + "source_kv_reuse_pinned", + "timeout_expected", + "writer_cohort_known", + } +) +_MAX_TIMESTAMP_NS = (1 << 63) - 1 +_MAX_IDENTIFIER = (1 << 64) - 1 + + +@dataclass(frozen=True) +class _ParsedEvent: + record: Event + line_number: int + + +@dataclass(frozen=True) +class ParseResult: + """Events and accounting produced while scanning mixed runtime logs.""" + + events: list[_ParsedEvent] + total_lines: int + ignored_lines: int + malformed_diagnostic_lines: int + + +def _is_diagnostic_scalar(value: object) -> bool: + return ( + value is None + or isinstance(value, (str, int, bool)) + or (isinstance(value, float) and math.isfinite(value)) + ) + + +def _has_valid_field_types(record: Event) -> bool: + """Reject records that cannot conform to the runtime's scalar schema.""" + if not all(_is_diagnostic_scalar(value) for value in record.values()): + return False + schema_version = record.get("schema_version") + if ( + not isinstance(schema_version, int) + or isinstance(schema_version, bool) + or schema_version != _SUPPORTED_SCHEMA_VERSION + ): + return False + for field in _IDENTIFIER_FIELDS: + value = record.get(field) + if value is not None and ( + not isinstance(value, int) + or isinstance(value, bool) + or not 0 <= value <= _MAX_IDENTIFIER + ): + return False + for field in _NONNEGATIVE_INTEGER_FIELDS: + value = record.get(field) + if value is not None and ( + not isinstance(value, int) + or isinstance(value, bool) + or not 0 <= value <= _MAX_TIMESTAMP_NS + ): + return False + for field in _NONNEGATIVE_NUMBER_FIELDS: + value = record.get(field) + if value is None: + continue + if not isinstance(value, (int, float)) or isinstance(value, bool) or value < 0: + return False + if isinstance(value, int): + if value > _MAX_TIMESTAMP_NS: + return False + elif not math.isfinite(value): + return False + for field in _STRING_FIELDS: + value = record.get(field) + if value is not None and not isinstance(value, str): + return False + for field in _BOOLEAN_FIELDS: + value = record.get(field) + if value is not None and not isinstance(value, bool): + return False + return True + + +@dataclass(frozen=True) +class _Phase: + name: str + start_event: str + end_event: str + correlation_fields: tuple[str, ...] = () + start_outcomes: frozenset[str] = frozenset() + end_outcomes: frozenset[str] = frozenset() + start_fields: tuple[tuple[str, object], ...] = () + end_fields: tuple[tuple[str, object], ...] = () + single_pair: bool = False + signed_offset: bool = False + report_unmatched: bool = True + + +_PHASES = ( + _Phase( + "gen_gate1_admission_wait", + "gen_ingress", + _GEN_KV_ADMISSION_EVENT, + end_outcomes=frozenset({"admitted"}), + single_pair=True, + ), + _Phase( + "gen_transfer_window_admission_wait", + _GEN_KV_ADMISSION_EVENT, + "gen_transfer_window_result", + start_outcomes=frozenset({"admitted"}), + end_outcomes=frozenset({"admitted"}), + single_pair=True, + report_unmatched=False, + ), + _Phase("gen_receive_lifetime", "gen_receive_start", "gen_transfer_settled"), + _Phase( + "gen_transfer_to_service", + "gen_transfer_settled", + "gen_decode_ready", + start_outcomes=frozenset({"completed"}), + ), + _Phase( + "ctx_receiver_readiness_offset", + "ctx_send_ready", + "ctx_all_receivers_ready", + signed_offset=True, + ), + _Phase( + "ctx_worker_queue_wait", + "ctx_transfer_queued", + "ctx_worker_dequeued", + correlation_fields=("slice_id", "peer_rank"), + ), + _Phase( + "ctx_worker_preparation", + "ctx_worker_dequeued", + "ctx_backend_submit_start", + correlation_fields=("slice_id", "peer_rank"), + ), + _Phase( + "ctx_backend_submission", + "ctx_backend_submit_start", + "ctx_backend_submitted", + correlation_fields=("slice_id", "peer_rank"), + ), + _Phase( + "ctx_backend_service", + "ctx_backend_submitted", + "ctx_backend_complete", + correlation_fields=("slice_id", "peer_rank"), + ), + _Phase("ctx_transfer_lifetime", "ctx_send_ready", "ctx_transfer_settled"), + _Phase( + "ctx_source_kv_request_ownership", + "ctx_send_ready", + "ctx_source_kv_released", + ), + _Phase( + "gen_writer_first_response", + "gen_request_data_sent", + "gen_writer_result_received", + correlation_fields=("slice_id", "peer_rank"), + single_pair=True, + ), + _Phase( + "gen_destination_drain", + "gen_writer_result_received", + "gen_destination_complete", + correlation_fields=("slice_id", "peer_rank"), + start_fields=(("outcome", "success"), ("is_last_slice", True)), + report_unmatched=False, + ), + _Phase( + "ctx_send_to_timeout_start", + "ctx_send_ready", + "transfer_timeout_started", + correlation_fields=("side",), + start_fields=(("timeout_expected", True),), + end_fields=(("side", "ctx"),), + ), + _Phase( + "gen_receive_to_timeout_start", + "gen_receive_start", + "transfer_timeout_started", + correlation_fields=("side",), + start_fields=(("timeout_expected", True),), + end_fields=(("side", "gen"),), + ), + _Phase( + "transfer_timeout_window", + "transfer_timeout_started", + "transfer_timeout_observed", + correlation_fields=("side", "timeout_owner"), + report_unmatched=False, + ), +) + +_EVENT_CAPABILITIES = { + event: "python_transfer_events" + for phase in _PHASES + for event in (phase.start_event, phase.end_event) +} +_EVENT_CAPABILITIES.update( + dict.fromkeys( + ( + "gen_ingress", + "gen_transfer_window_result", + "gen_decode_ready", + "ctx_send_ready", + "ctx_source_kv_released", + "ctx_source_unpinned", + "transfer_timeout_started", + "transfer_timeout_observed", + ), + "executor_events", + ) +) +_EVENT_CAPABILITIES.update( + dict.fromkeys( + (_GEN_KV_ADMISSION_EVENT, "gen_kv_pool_snapshot"), "scheduler_kv_admission_events" + ) +) + + +@dataclass(frozen=True) +class _Capabilities: + groups: dict[str, bool | None] + runtime: str | None = None + issues: tuple[str, ...] = () + + def supports(self, event: str) -> bool | None: + return self.groups[_EVENT_CAPABILITIES[event]] + + def to_json(self) -> dict[str, object]: + known = sum(value is not None for value in self.groups.values()) + return { + "status": "known" + if known == len(_CAPABILITY_FIELDS) + else "partial" + if known + else "unknown", + "transceiver_runtime": self.runtime, + **self.groups, + "issues": list(self.issues), + } + + +def _capability_scope(record: Event) -> CapabilityScope | None: + domain = _clock_domain(record) + return (*domain, _participant(record)[3]) if domain is not None else None + + +def _known_capability_version(record: Event) -> bool: + return ( + type(record.get("capability_schema_version")) is int + and record["capability_schema_version"] == 1 + ) + + +def _capability_index(events: list[_ParsedEvent]) -> dict[CapabilityScope, _Capabilities]: + declarations: dict[CapabilityScope, list[Event]] = defaultdict(list) + for event in events: + if event.record["event"] == "diagnostic_capabilities": + scope = _capability_scope(event.record) + if scope is not None: + declarations[scope].append(event.record) + result = {} + for scope, records in declarations.items(): + issues = set() + groups: dict[str, bool | None] = {} + if not all(_known_capability_version(record) for record in records): + issues.add("invalid_capability_schema_version") + for field in _CAPABILITY_FIELDS: + field_records = [record for record in records if field in record] + values = [record[field] for record in field_records] + if any(type(value) is not bool for value in values) or not all( + _known_capability_version(record) for record in field_records + ): + issues.add(f"invalid_{field}") + groups[field] = None + elif len(set(values)) > 1: + issues.add(f"conflicting_{field}") + groups[field] = None + else: + groups[field] = values[0] is True if values else None + runtimes = { + record["transceiver_runtime"] for record in records if "transceiver_runtime" in record + } + runtime = next(iter(runtimes)) if len(runtimes) == 1 else None + if runtimes and ( + len(runtimes) != 1 + or runtime not in ("CPP", "PYTHON") + or any( + not _known_capability_version(record) + for record in records + if "transceiver_runtime" in record + ) + ): + issues.add("invalid_or_conflicting_transceiver_runtime") + groups["python_transfer_events"] = None + runtime = None + elif runtime is not None and groups["python_transfer_events"] not in ( + None, + runtime == "PYTHON", + ): + issues.add("runtime_capability_mismatch") + groups["python_transfer_events"] = None + result[scope] = _Capabilities( + groups, runtime if isinstance(runtime, str) else None, tuple(sorted(issues)) + ) + return result + + +def _participant_capabilities( + events: list[_ParsedEvent], index: dict[CapabilityScope, _Capabilities] +) -> _Capabilities: + scope = _capability_scope(events[0].record) + profile = index.get(scope) if scope is not None else None + if profile is None: + return _Capabilities( + dict.fromkeys(_CAPABILITY_FIELDS), issues=("missing_capability_metadata",) + ) + # Contradictory records must not turn an observed Python path into a false + # C++ exemption. Keep actual measurements, but mark the declaration unknown. + groups = profile.groups.copy() + issues = set(profile.issues) + for event in events: + field = _EVENT_CAPABILITIES.get(str(event.record["event"])) + if field is not None and groups[field] is False: + groups[field] = None + issues.add(f"observed_unsupported_{field}") + return _Capabilities(groups, profile.runtime, tuple(sorted(issues))) + + +def parse_lines(lines: Iterable[str]) -> ParseResult: + """Extract valid diagnostic JSON objects from mixed log lines.""" + events: list[_ParsedEvent] = [] + total_lines = 0 + ignored_lines = 0 + malformed_lines = 0 + + for line_number, line in enumerate(lines, start=1): + total_lines += 1 + marker = line.find(DIAGNOSTICS_LOG_PREFIX) + if marker < 0: + ignored_lines += 1 + continue + + payload = line[marker + len(DIAGNOSTICS_LOG_PREFIX) :].strip() + try: + record = json.loads(payload) + except (json.JSONDecodeError, RecursionError, ValueError): + malformed_lines += 1 + continue + if ( + not isinstance(record, dict) + or not isinstance(record.get("event"), str) + or not _has_valid_field_types(record) + ): + malformed_lines += 1 + continue + events.append(_ParsedEvent(record=record, line_number=line_number)) + + return ParseResult( + events=events, + total_lines=total_lines, + ignored_lines=ignored_lines, + malformed_diagnostic_lines=malformed_lines, + ) + + +def _uuid_value(value: object) -> str | None: + if not isinstance(value, str): + return None + try: + return str(UUID(value)) + except ValueError: + return None + + +def _identity_issues(record: Event) -> list[str]: + issues = [] + for field in ("run_uuid", "process_uuid"): + if _uuid_value(record.get(field)) is None: + invalid = record.get(field) is not None or ( + field == "run_uuid" and record.get("run_uuid_status") == "invalid" + ) + issues.append(f"{'invalid' if invalid else 'missing'}_{field}") + return issues + + +def _clock_domain(record: Event) -> ClockDomain | None: + host = record.get("host") + pid = record.get("pid") + if not isinstance(host, str) or not host or not isinstance(pid, int) or isinstance(pid, bool): + return None + process_uuid = _uuid_value(record.get("process_uuid")) + if process_uuid is None: + return None + return host, pid, process_uuid, _uuid_value(record.get("run_uuid")) + + +def _monotonic_ns(record: Event) -> int | None: + value = ( + record.get("timer_start_monotonic_ns", record.get("monotonic_ns")) + if record.get("event") == "transfer_timeout_started" + else record.get("monotonic_ns") + ) + if isinstance(value, int) and not isinstance(value, bool) and 0 <= value <= _MAX_TIMESTAMP_NS: + return value + return None + + +def _matches_phase_boundary( + event: _ParsedEvent, + name: str, + outcomes: frozenset[str], + required_fields: tuple[tuple[str, object], ...], +) -> bool: + record = event.record + return ( + record.get("event") == name + and (not outcomes or record.get("outcome") in outcomes) + and all(record.get(field) == value for field, value in required_fields) + ) + + +def _correlation(record: Event, fields: tuple[str, ...]) -> tuple[object, ...]: + return tuple(record.get(field) for field in fields) + + +def _domain_json(domain: ClockDomain) -> dict[str, object]: + return { + "host": domain[0], + "pid": domain[1], + "process_uuid": domain[2], + "run_uuid": domain[3], + } + + +def _participant(record: Event) -> Participant: + side = record.get("side") + host = record.get("host") + pid = record.get("pid") + rank = record.get("rank") + return ( + side if isinstance(side, str) else None, + host if isinstance(host, str) else None, + pid if isinstance(pid, int) and not isinstance(pid, bool) else None, + rank if isinstance(rank, int) and not isinstance(rank, bool) else None, + _uuid_value(record.get("process_uuid")), + _uuid_value(record.get("run_uuid")), + ) + + +def _participant_json(participant: Participant) -> dict[str, object]: + side, host, pid, rank, process_uuid, run_uuid = participant + return { + "side": side, + "host": host, + "pid": pid, + "rank": rank, + "process_uuid": process_uuid, + "run_uuid": run_uuid, + } + + +def _timeline(events: list[_ParsedEvent]) -> list[dict[str, object]]: + timeline = [] + for event in sorted(events, key=_timeline_sort_key): + entry = {field: event.record[field] for field in _TIMELINE_FIELDS if field in event.record} + entry["line_number"] = event.line_number + timeline.append(entry) + return timeline + + +def _timeline_sort_key(event: _ParsedEvent) -> tuple[int, int, int]: + wall_ns = event.record.get("wall_ns") + if isinstance(wall_ns, int) and not isinstance(wall_ns, bool) and wall_ns >= 0: + return 0, wall_ns, event.line_number + return 1, 0, event.line_number + + +def _broadcast_writer_expectations( + starts: list[_ParsedEvent], ends: list[_ParsedEvent] +) -> tuple[set[tuple[ClockDomain, Participant, object]], list[dict[str, object]]]: + """Do not require every candidate in a GEN-first ADP broadcast to respond.""" + cohorts: dict[tuple[ClockDomain, Participant, object], list[Event]] = defaultdict(list) + responders: dict[tuple[ClockDomain, Participant, object], set[object]] = defaultdict(set) + for boundaries, is_start in ((starts, True), (ends, False)): + for event in boundaries: + record = event.record + domain = _clock_domain(record) + if ( + domain is None + or _monotonic_ns(record) is None + or record.get("slice_id") is None + or record.get("peer_rank") is None + ): + continue + key = domain, _participant(record), record["slice_id"] + if is_start: + cohorts[key].append(record) + elif record.get("session_found") is not False: + responders[key].add(record["peer_rank"]) + + relaxed = set() + gaps = [] + for (domain, participant, slice_id), records in cohorts.items(): + flags = {record.get("writer_cohort_known") for record in records} + if flags == {True}: + continue + key = domain, participant, slice_id + relaxed.add(key) + observed = len(responders[key]) + counts = {record.get("expected_writers") for record in records} + expected = next(iter(counts)) if len(counts) == 1 else None + issue = None + if flags != {False}: + # Older or partial logs cannot prove that every published peer was + # selected. Still retain all observed per-peer timing below. + issue = "missing_or_conflicting_writer_cohort" + elif expected is None or not isinstance(expected, int) or expected <= 0: + issue = "missing_invalid_or_conflicting_expected_writers" + if issue is None and observed == expected: + continue + gap: dict[str, object] = { + "phase": "gen_writer_first_response", + "clock_domain": _domain_json(domain), + "correlation": {"slice_id": slice_id}, + "observed_writers": observed, + } + if participant[3] is not None: + gap["rank"] = participant[3] + if issue is not None: + gap.update(reason="unknown_writer_cohort", detail=issue) + else: + gap.update( + reason="missing_writer_responses" + if observed < expected + else "unexpected_writer_responses", + expected_writers=expected, + count=abs(expected - observed), + ) + gaps.append(gap) + return relaxed, gaps + + +def _derive_phase( + events: list[_ParsedEvent], phase: _Phase, profiles: dict[Participant, _Capabilities] +) -> tuple[list[dict[str, object]], list[dict[str, object]]]: + starts = [ + event + for event in events + if _matches_phase_boundary( + event, + phase.start_event, + phase.start_outcomes, + phase.start_fields, + ) + ] + ends = [ + event + for event in events + if _matches_phase_boundary( + event, + phase.end_event, + phase.end_outcomes, + phase.end_fields, + ) + ] + if not starts and not ends: + return [], [] + if not phase.report_unmatched and (not starts or not ends): + return [], [] + + Boundary = tuple[ClockDomain, Participant, tuple[object, ...], int] + + def timed(boundaries: list[_ParsedEvent]) -> tuple[list[Boundary], int, int, int]: + result = [] + invalid_clock_count = 0 + missing_correlation_count = 0 + unverified_identity_count = 0 + for event in boundaries: + if _uuid_value(event.record.get("process_uuid")) is None: + unverified_identity_count += 1 + continue + domain = _clock_domain(event.record) + timestamp = _monotonic_ns(event.record) + if domain is None or timestamp is None: + invalid_clock_count += 1 + continue + correlation = _correlation(event.record, phase.correlation_fields) + if any(value is None for value in correlation): + missing_correlation_count += 1 + continue + result.append((domain, _participant(event.record), correlation, timestamp)) + return result, invalid_clock_count, missing_correlation_count, unverified_identity_count + + timed_starts, invalid_clock_starts, missing_correlation_starts, unverified_starts = timed( + starts + ) + timed_ends, invalid_clock_ends, missing_correlation_ends, unverified_ends = timed(ends) + grouped_starts: dict[tuple[ClockDomain, Participant, tuple[object, ...]], list[int]] = ( + defaultdict(list) + ) + grouped_ends: dict[tuple[ClockDomain, Participant, tuple[object, ...]], list[int]] = ( + defaultdict(list) + ) + for domain, participant, correlation, timestamp in timed_starts: + grouped_starts[(domain, participant, correlation)].append(timestamp) + for domain, participant, correlation, timestamp in timed_ends: + grouped_ends[(domain, participant, correlation)].append(timestamp) + + durations: list[dict[str, object]] = [] + relaxed_writer_cohorts, unmeasured = ( + _broadcast_writer_expectations(starts, ends) + if phase.name == "gen_writer_first_response" + else (set(), []) + ) + all_keys = sorted(grouped_starts.keys() | grouped_ends.keys(), key=lambda key: repr(key)) + for domain, participant, correlation in all_keys: + domain_starts = sorted(grouped_starts.get((domain, participant, correlation), [])) + domain_ends = sorted(grouped_ends.get((domain, participant, correlation), [])) + used_end_indexes: set[int] = set() + unmatched_starts = 0 + + for start_ns in domain_starts[:1] if phase.single_pair else domain_starts: + candidates = [ + index + for index, end_ns in enumerate(domain_ends) + if index not in used_end_indexes and (phase.signed_offset or end_ns >= start_ns) + ] + if not candidates: + unmatched_starts += 1 + continue + key = ( + (lambda index: abs(domain_ends[index] - start_ns)) + if phase.signed_offset + else (lambda index: domain_ends[index]) + ) + end_index = min(candidates, key=key) + used_end_indexes.add(end_index) + duration_ns = domain_ends[end_index] - start_ns + result: dict[str, object] = { + "phase": phase.name, + "clock_domain": _domain_json(domain), + } + if participant[3] is not None: + result["rank"] = participant[3] + if phase.signed_offset: + result.update( + { + "signed_offset_ns": duration_ns, + "signed_offset_ms": duration_ns / 1_000_000, + "readiness_wait_ms": max(duration_ns, 0) / 1_000_000, + "readiness_lead_ms": max(-duration_ns, 0) / 1_000_000, + } + ) + else: + result.update( + { + "duration_ns": duration_ns, + "duration_ms": duration_ns / 1_000_000, + } + ) + if phase.correlation_fields: + result["correlation"] = dict(zip(phase.correlation_fields, correlation)) + durations.append(result) + + if phase.single_pair: + missing_end_count = int(bool(domain_starts) and not used_end_indexes) + missing_start_count = int(not domain_starts and bool(domain_ends)) + else: + missing_end_count = unmatched_starts + missing_start_count = len(domain_ends) - len(used_end_indexes) + if ( + relaxed_writer_cohorts + and not domain_ends + and (domain, participant, correlation[0]) in relaxed_writer_cohorts + ): + missing_end_count = 0 + if phase.report_unmatched and (missing_end_count or missing_start_count): + for reason, count in ( + ("missing_end", missing_end_count), + ("missing_start", missing_start_count), + ): + if count == 0: + continue + gap: dict[str, object] = { + "phase": phase.name, + "reason": _phase_gap_reason(profiles[participant], phase, reason), + "count": count, + "clock_domain": _domain_json(domain), + } + if participant[3] is not None: + gap["rank"] = participant[3] + if gap["reason"] != reason: + gap["missing_boundary"] = reason.removeprefix("missing_") + if phase.correlation_fields: + gap["correlation"] = dict(zip(phase.correlation_fields, correlation)) + unmeasured.append(gap) + + for reason, start_count, end_count in ( + ("unverified_process_identity", unverified_starts, unverified_ends), + ("invalid_clock_metadata", invalid_clock_starts, invalid_clock_ends), + ( + "missing_correlation_fields", + missing_correlation_starts, + missing_correlation_ends, + ), + ): + if start_count or end_count: + unmeasured.append( + { + "phase": phase.name, + "reason": reason, + "start_count": start_count, + "end_count": end_count, + } + ) + if durations: + return durations, unmeasured + start_domains = {start[0] for start in timed_starts} + end_domains = {end[0] for end in timed_ends} + if not phase.single_pair and unmeasured: + if starts and ends and start_domains and start_domains.isdisjoint(end_domains): + unmeasured.append({"phase": phase.name, "reason": "clock_domain_mismatch"}) + return [], unmeasured + if not starts or not ends: + # Participant-specific gaps above already describe capability support; + # do not add a second, unconditional missing-boundary summary. + return [], unmeasured + elif len(timed_starts) != len(starts) or len(timed_ends) != len(ends): + return [], unmeasured + elif start_domains.isdisjoint(end_domains): + reason = "clock_domain_mismatch" + else: + reason = "no_ordered_matching_boundaries" + summary = {"phase": phase.name, "reason": reason} + if summary not in unmeasured: + unmeasured.append(summary) + return [], unmeasured + + +def _has_event( + events: list[_ParsedEvent], + name: str, + *, + outcomes: frozenset[str] = frozenset(), + excluded_policies: frozenset[str] = frozenset(), + required_fields: tuple[tuple[str, object], ...] = (), +) -> bool: + return any( + event.record.get("event") == name + and (not outcomes or event.record.get("outcome") in outcomes) + and event.record.get("policy") not in excluded_policies + and all(event.record.get(field) == value for field, value in required_fields) + for event in events + ) + + +def _phase_gap_reason(profile: _Capabilities, phase: _Phase, reason: str) -> str: + support = (profile.supports(phase.start_event), profile.supports(phase.end_event)) + if False in support: + return "unsupported_capability" + if None in support: + return "unknown_capability" + return reason + + +def _boundary_expectations( + events: list[_ParsedEvent], profile: _Capabilities +) -> dict[str, list[str]]: + expected: set[str] = set() + observed = {event.record.get("event") for event in events} + + if _has_event(events, "gen_ingress"): + expected.add(_GEN_KV_ADMISSION_EVENT) + if _has_event( + events, + "gen_transfer_window_result", + outcomes=frozenset({"admitted"}), + excluded_policies=frozenset({"not_applicable"}), + ): + expected.add("gen_receive_start") + if _has_event(events, "gen_receive_start"): + expected.update(("gen_request_data_sent", "gen_transfer_settled")) + if _has_event(events, "gen_request_data_sent"): + expected.add("gen_writer_result_received") + writer_results = [ + event for event in events if event.record.get("event") == "gen_writer_result_received" + ] + if ( + writer_results + and all(event.record.get("outcome") == "success" for event in writer_results) + and any(event.record.get("is_last_slice") is True for event in writer_results) + ): + expected.add("gen_destination_complete") + if _has_event(events, "gen_transfer_settled", outcomes=frozenset({"completed"})): + expected.add("gen_decode_ready") + if _has_event(events, "transfer_timeout_observed"): + expected.add("transfer_timeout_started") + + if _has_event(events, "ctx_send_ready"): + expected.update( + ( + "ctx_all_receivers_ready", + "ctx_source_kv_released", + "ctx_transfer_settled", + ) + ) + if _has_event(events, "ctx_transfer_queued"): + expected.add("ctx_worker_dequeued") + if _has_event(events, "ctx_worker_dequeued"): + expected.add("ctx_backend_submit_start") + if _has_event(events, "ctx_backend_submit_start"): + expected.add("ctx_backend_submitted") + if _has_event(events, "ctx_backend_submitted"): + expected.add("ctx_backend_complete") + + result: dict[str, list[str]] = { + "missing_boundaries": [], + "unsupported_boundaries": [], + "unassessed_boundaries": [], + } + for event in sorted(expected - observed): + support = profile.supports(event) + key = ( + "missing_boundaries" + if support is True + else "unsupported_boundaries" + if support is False + else "unassessed_boundaries" + ) + result[key].append(event) + return result + + +def _request_sort_key(request_id: object) -> tuple[str, str]: + return type(request_id).__name__, str(request_id) + + +def _request_group_key(event: _ParsedEvent) -> tuple[str, str, str, str]: + record = event.record + process_uuid = _uuid_value(record.get("process_uuid")) + run_uuid = _uuid_value(record.get("run_uuid")) + if process_uuid is None: + # Even records from one input file can span restarts. Keep unidentified + # records separate instead of manufacturing a request from reused IDs. + scope, identity = "unverified", str(event.line_number) + elif run_uuid is not None: + scope, identity = "run", run_uuid + else: + scope, identity = "process", repr((process_uuid, record.get("host"), record.get("pid"))) + return scope, identity, *_request_sort_key(record["request_id"]) + + +def _participant_summaries( + grouped: dict[Participant, list[_ParsedEvent]], profiles: dict[Participant, _Capabilities] +) -> list[dict[str, object]]: + summaries = [] + for participant, participant_events in sorted(grouped.items(), key=lambda item: repr(item[0])): + event_counts = Counter(str(event.record["event"]) for event in participant_events) + summary = _participant_json(participant) + summary.update( + { + "event_count": len(participant_events), + "event_counts": dict(sorted(event_counts.items())), + "capabilities": profiles[participant].to_json(), + **_boundary_expectations(participant_events, profiles[participant]), + } + ) + summaries.append(summary) + return summaries + + +def _summarize_request( + request_id: object, + events: list[_ParsedEvent], + capability_index: dict[CapabilityScope, _Capabilities], +) -> dict[str, object]: + event_counts = Counter(str(event.record["event"]) for event in events) + sides = sorted( + {side for event in events if isinstance((side := event.record.get("side")), str)} + ) + domains = sorted( + {domain for event in events if (domain := _clock_domain(event.record)) is not None} + ) + durations: list[dict[str, object]] = [] + unmeasured: list[dict[str, object]] = [] + grouped: dict[Participant, list[_ParsedEvent]] = defaultdict(list) + for event in events: + grouped[_participant(event.record)].append(event) + profiles = { + participant: _participant_capabilities(participant_events, capability_index) + for participant, participant_events in grouped.items() + } + participants = _participant_summaries(grouped, profiles) + for phase in _PHASES: + phase_durations, phase_unmeasured = _derive_phase(events, phase, profiles) + durations.extend(phase_durations) + unmeasured.extend(phase_unmeasured) + + return { + "request_id": request_id, + "run_uuid": _uuid_value(events[0].record.get("run_uuid")), + "correlation_scope": _request_group_key(events[0])[0], + "identity_issues": sorted( + {issue for event in events for issue in _identity_issues(event.record)} + ), + "event_count": len(events), + "event_counts": dict(sorted(event_counts.items())), + "sides": sides, + "clock_domains": [_domain_json(domain) for domain in domains], + **{ + field: sorted( + {boundary for participant in participants for boundary in participant[field]} + ) + for field in ("missing_boundaries", "unsupported_boundaries", "unassessed_boundaries") + }, + "participants": participants, + "durations": durations, + "unmeasured_phases": unmeasured, + "timeline": _timeline(events), + } + + +def _percentile(values: list[float], fraction: float) -> float: + if len(values) == 1: + return values[0] + position = fraction * (len(values) - 1) + lower = int(position) + upper = min(lower + 1, len(values) - 1) + weight = position - lower + return values[lower] * (1 - weight) + values[upper] * weight + + +def _phase_summary(requests: list[dict[str, object]]) -> dict[str, dict[str, object]]: + samples: dict[str, list[float]] = defaultdict(list) + for request in requests: + durations = request["durations"] + assert isinstance(durations, list) + for duration in durations: + assert isinstance(duration, dict) + phase = duration.get("phase") + value = duration.get("duration_ms") + if value is None: + value = duration.get("signed_offset_ms") + if isinstance(phase, str) and isinstance(value, (int, float)): + samples[phase].append(float(value)) + + summary: dict[str, dict[str, object]] = {} + for phase, unsorted_values in sorted(samples.items()): + values = sorted(unsorted_values) + summary[phase] = { + "count": len(values), + "min_ms": values[0], + "mean_ms": statistics.fmean(values), + "p50_ms": _percentile(values, 0.50), + "p95_ms": _percentile(values, 0.95), + "max_ms": values[-1], + } + return summary + + +def _snapshot_cadence(events: list[_ParsedEvent]) -> list[dict[str, object]]: + samples: dict[tuple[ClockDomain, int | None], list[int]] = defaultdict(list) + for event in events: + if event.record.get("event") != "gen_kv_pool_snapshot": + continue + domain = _clock_domain(event.record) + timestamp = _monotonic_ns(event.record) + rank = event.record.get("rank") + rank = rank if isinstance(rank, int) and not isinstance(rank, bool) else None + if domain is not None and timestamp is not None: + samples[(domain, rank)].append(timestamp) + + result = [] + for (domain, rank), timestamps in sorted(samples.items(), key=lambda item: repr(item[0])): + ordered = sorted(timestamps) + intervals_ms = [ + (current - previous) / 1_000_000 for previous, current in zip(ordered, ordered[1:]) + ] + item: dict[str, object] = { + **_domain_json(domain), + "rank": rank, + "snapshot_count": len(ordered), + "interval_count": len(intervals_ms), + } + if intervals_ms: + item.update( + { + "min_ms": min(intervals_ms), + "mean_ms": statistics.fmean(intervals_ms), + "max_ms": max(intervals_ms), + } + ) + result.append(item) + return result + + +def _transfer_settled_to_next_decision(events: list[_ParsedEvent]) -> list[dict[str, object]]: + decisions: dict[tuple[ClockDomain, int | None], list[int]] = defaultdict(list) + settled_events: dict[tuple[ClockDomain, int | None], list[int]] = defaultdict(list) + for event in events: + domain = _clock_domain(event.record) + timestamp = _monotonic_ns(event.record) + rank = event.record.get("rank") + rank = rank if isinstance(rank, int) and not isinstance(rank, bool) else None + if domain is None or timestamp is None: + continue + key = domain, rank + if event.record.get("event") == "gen_kv_pool_snapshot": + decisions[key].append(timestamp) + elif ( + event.record.get("event") == "gen_transfer_settled" + and event.record.get("resources_drained") is True + ): + settled_events[key].append(timestamp) + + result = [] + for (domain, rank), settled_times in sorted( + settled_events.items(), key=lambda item: repr(item[0]) + ): + decision_times = sorted(decisions.get((domain, rank), [])) + delays_ms = [] + unmatched = 0 + for settled_time in settled_times: + index = bisect_left(decision_times, settled_time) + if index == len(decision_times): + unmatched += 1 + else: + delays_ms.append((decision_times[index] - settled_time) / 1_000_000) + item: dict[str, object] = { + **_domain_json(domain), + "rank": rank, + "settled_count": len(settled_times), + "matched_count": len(delays_ms), + "unmatched_count": unmatched, + } + if delays_ms: + item.update( + { + "min_ms": min(delays_ms), + "mean_ms": statistics.fmean(delays_ms), + "max_ms": max(delays_ms), + } + ) + result.append(item) + return result + + +def analyze(parse_result: ParseResult) -> dict[str, object]: + """Group parsed events and produce request- and phase-level summaries.""" + grouped: dict[tuple[str, str, str, str], tuple[object, list[_ParsedEvent]]] = {} + ungrouped_events: list[_ParsedEvent] = [] + ungrouped_event_counts: Counter[str] = Counter() + event_counts: Counter[str] = Counter() + schema_versions: Counter[str] = Counter() + + for event in parse_result.events: + record = event.record + event_name = str(record["event"]) + event_counts[event_name] += 1 + schema_versions[str(record.get("schema_version", "missing"))] += 1 + request_id = record.get("request_id") + if request_id is None: + ungrouped_events.append(event) + ungrouped_event_counts[event_name] += 1 + continue + key = _request_group_key(event) + if key not in grouped: + grouped[key] = (request_id, []) + grouped[key][1].append(event) + + capabilities = _capability_index(parse_result.events) + requests = [ + _summarize_request(request_id, events, capabilities) + for _, (request_id, events) in sorted(grouped.items()) + ] + return { + "identity_semantics": { + "requests": "Join across processes only by a shared run_uuid and request_id.", + "process_local": ( + "Without a valid run_uuid, correlate only within one process_uuid, host and PID." + ), + "unverified": ( + "Without a valid process_uuid, retain individual records " + "without joining or deriving timing." + ), + "configuration": ( + "Set TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS_RUN_ID to the same fresh UUID " + "on CTX and GEN for each launch." + ), + }, + "identity_issues": dict( + sorted( + Counter( + issue + for event in parse_result.events + for issue in _identity_issues(event.record) + ).items() + ) + ), + "capability_semantics": { + "scope": ( + "Startup declarations are matched by run/process UUIDs, host, pid and rank, " + "then applied per participant." + ), + "missing_boundaries": "Expected events supported by this participant but not observed.", + "unsupported_boundaries": "Events not instrumented by this participant's implementation.", + "unassessed_boundaries": ( + "Expected events whose capability metadata is absent or ambiguous; " + "not a healthy result." + ), + "durations": ( + "Observed matching pairs remain measurable without capability metadata; " + "unsupported or unknown gaps are labeled separately." + ), + }, + "clock_semantics": { + "durations": ( + "Derived only from monotonic_ns within matching run/process UUIDs, host and PID." + ), + "timeline": ( + "wall_ns is preserved for manual cross-process correlation; cross-host ordering " + "is clock-sync-sensitive and no wall-clock deltas are derived." + ), + "gen_transfer_settled_to_next_scheduler_decision": ( + "A same-domain measure from successful transceiver session retirement " + "to the next admission opportunity." + ), + }, + "summary": { + "total_lines": parse_result.total_lines, + "ignored_lines": parse_result.ignored_lines, + "malformed_diagnostic_lines": parse_result.malformed_diagnostic_lines, + "parsed_events": len(parse_result.events), + "events_without_request_id": sum(ungrouped_event_counts.values()), + "request_count": len(requests), + }, + "event_counts": dict(sorted(event_counts.items())), + "schema_versions": dict(sorted(schema_versions.items())), + "ungrouped_event_counts": dict(sorted(ungrouped_event_counts.items())), + "aggregate_timeline": _timeline(ungrouped_events), + "scheduler_decision_cadence": _snapshot_cadence(ungrouped_events), + "gen_transfer_settled_to_next_scheduler_decision": ( + _transfer_settled_to_next_decision(parse_result.events) + ), + "phase_durations": _phase_summary(requests), + "requests": requests, + } + + +def analyze_lines(lines: Iterable[str]) -> dict[str, object]: + """Parse and analyze mixed runtime log lines.""" + return analyze(parse_lines(lines)) + + +def _parse_args(argv: Sequence[str] | None) -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "logs", + nargs="*", + help="Log files to analyze. Read standard input when no file is provided.", + ) + parser.add_argument("-o", "--output", type=Path, help="Write JSON to this file.") + parser.add_argument("--indent", type=int, default=2, help="JSON indentation (default: 2).") + return parser.parse_args(argv) + + +def main(argv: Sequence[str] | None = None) -> int: + """Run the command-line analyzer.""" + args = _parse_args(argv) + with fileinput.input(files=args.logs or ("-",), encoding="utf-8", errors="replace") as lines: + result = analyze_lines(lines) + + if args.output is None: + json.dump(result, sys.stdout, indent=args.indent, sort_keys=True) + sys.stdout.write("\n") + else: + with args.output.open("w", encoding="utf-8") as output: + json.dump(result, output, indent=args.indent, sort_keys=True) + output.write("\n") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tensorrt_llm/_torch/disaggregation/diagnostics.py b/tensorrt_llm/_torch/disaggregation/diagnostics.py new file mode 100644 index 000000000000..51ae83e3c9de --- /dev/null +++ b/tensorrt_llm/_torch/disaggregation/diagnostics.py @@ -0,0 +1,357 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Opt-in request-edge diagnostics for disaggregated KV transfer. + +Set ``TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS=1`` before process startup for a +targeted diagnostic run. Standard and performance runs leave it unset, making +each instrumented edge a module-attribute check with no clock read, request +inspection, serialization, or log emission. + +For cross-process correlation, set +``TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS_RUN_ID`` to the same UUID for all CTX and +GEN workers in one diagnostic launch. Use a new UUID for each launch. An unset +or invalid value limits analysis to individual process lifetimes. +""" + +from __future__ import annotations + +import atexit +import functools +import json +import os +import queue +import socket +import threading +import time +import uuid +from contextlib import contextmanager +from typing import TYPE_CHECKING, Iterator, Optional, TypeAlias + +if TYPE_CHECKING: + from tensorrt_llm._torch.disaggregation.native.rank_info import RankInfo + +_DIAGNOSTICS_ENV = "TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS" +_DIAGNOSTICS_RUN_ID_ENV = "TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS_RUN_ID" +DIAGNOSTICS_LOG_PREFIX = "[DISAGG_TRANSFER_DIAG] " +DIAGNOSTICS_SCHEMA_VERSION = 1 +_DIAGNOSTIC_QUEUE_CAPACITY = 32_768 +_DIAGNOSTIC_SHUTDOWN_TIMEOUT_S = 1.0 + +# Read once so the disabled path is a single caller-side branch. Tests may +# monkeypatch this module attribute without reloading importers. +DISAGG_TRANSFER_DIAGNOSTICS_ENABLED = os.getenv(_DIAGNOSTICS_ENV) == "1" + +DiagnosticValue: TypeAlias = str | int | float | bool | None +DiagnosticRecord: TypeAlias = dict[str, DiagnosticValue] +DiagnosticTimestamp: TypeAlias = tuple[int, int] + + +@contextmanager +def suppress_diagnostic_errors() -> Iterator[None]: + """Contain telemetry-only preparation without hiding runtime failures. + + Callers must keep core scheduling, transfer, and cleanup work outside this + scope. The scope exists because Python evaluates ``emit_event`` arguments + before that function can apply its own best-effort exception handling. + """ + try: + yield + except Exception: + # Diagnostics are best-effort and must never affect request progress. + return + + +def capture_timestamp() -> DiagnosticTimestamp: + """Capture one local event boundary before publishing shared state.""" + return time.monotonic_ns(), time.time_ns() + + +@functools.lru_cache(maxsize=1) +def _host_identity() -> str: + """Return the execution-node identity, computed only when diagnostics run.""" + hostname = os.getenv("SLURMD_NODENAME") + if not hostname: + try: + hostname = socket.gethostname() + except OSError: + hostname = os.getenv("HOSTNAME", "unknown") + return "_".join(hostname.split()) if hostname else "unknown" + + +class _AsyncDiagnosticSink: + """Serialize and write records away from request-progress threads.""" + + def __init__(self, pid: int) -> None: + self.pid = pid + self._identity: DiagnosticRecord = { + "process_uuid": str(uuid.uuid4()), + "run_uuid": None, + "run_uuid_status": "unset", + } + run_id = os.getenv(_DIAGNOSTICS_RUN_ID_ENV) + if run_id is not None: + try: + self._identity["run_uuid"] = str(uuid.UUID(run_id)) + except ValueError: + self._identity["run_uuid_status"] = "invalid" + else: + self._identity["run_uuid_status"] = "shared" + self._queue: queue.Queue[DiagnosticRecord] = queue.Queue(maxsize=_DIAGNOSTIC_QUEUE_CAPACITY) + self._stop_requested = threading.Event() + self._submit_lock = threading.Lock() + self._drop_lock = threading.Lock() + self._dropped = 0 + self._thread = threading.Thread( + target=self._run, + name="disagg-transfer-diagnostics", + daemon=True, + ) + self._thread.start() + + def submit(self, record: DiagnosticRecord) -> None: + with self._submit_lock: + if self._stop_requested.is_set(): + return + try: + self._queue.put_nowait(record) + except queue.Full: + # Never block request progress on diagnostics. A later writer + # pass reports the loss explicitly when the sink catches up. + self._record_dropped(1) + + def _record_dropped(self, count: int) -> None: + with self._drop_lock: + self._dropped += count + + def _take_dropped(self) -> int: + with self._drop_lock: + dropped = self._dropped + self._dropped = 0 + return dropped + + @staticmethod + def _write(record: DiagnosticRecord) -> None: + """Write one complete record to the process's OS-level stdout. + + File descriptor 1 is intentional: launcher and CI process-log capture + should receive diagnostics independently of Python logging or + ``sys.stdout`` replacement. Consequently, in-process redirection of + ``sys.stdout`` alone does not capture these records. Partial writes + are completed here, although records larger than the platform's + ``PIPE_BUF`` may still interleave with other processes sharing a pipe. + An unrecoverable descriptor error drops the record silently because + the failed output channel cannot reliably carry its own loss report. + """ + line = ( + f"{DIAGNOSTICS_LOG_PREFIX}{json.dumps(record, separators=(',', ':'), sort_keys=True)}\n" + ) + encoded_line = line.encode("utf-8") + offset = 0 + while offset < len(encoded_line): + written = os.write(1, encoded_line[offset:]) + if written <= 0: + raise OSError("diagnostic stdout write made no progress") + offset += written + + def _write_identified_record(self, record: DiagnosticRecord) -> None: + # Attach cached identities on the writer thread, overriding caller + # details so no event can impersonate a different process lifetime. + record.update(self._identity) + record["host"] = _host_identity() + record["pid"] = self.pid + self._write(record) + + def _write_drop_record(self) -> None: + dropped = self._take_dropped() + if dropped == 0: + return + try: + self._write_identified_record( + { + "schema_version": DIAGNOSTICS_SCHEMA_VERSION, + "event": "diagnostics_events_dropped", + "side": "runtime", + "request_id": None, + "local_request_id": None, + "monotonic_ns": time.monotonic_ns(), + "wall_ns": time.time_ns(), + "dropped_events": dropped, + } + ) + except Exception: + # Preserve overflow accounting across transient output failures. + self._record_dropped(dropped) + raise + + def _run(self) -> None: + while True: + try: + record = self._queue.get(timeout=0.1) + except queue.Empty: + if self._stop_requested.is_set(): + try: + self._write_drop_record() + except Exception: + pass + return + continue + + try: + self._write_drop_record() + except Exception: + pass + try: + self._write_identified_record(record) + except Exception: + # Diagnostics are best-effort and must never affect request + # progress. Account for transient failures so a later healthy + # write can make the missing record visible. + self._record_dropped(1) + finally: + self._queue.task_done() + + def flush(self) -> None: + """Wait for already accepted records; used only by focused tests.""" + self._queue.join() + + def close(self) -> None: + """Stop accepting records and make a bounded best effort to drain. + + If stdout remains backpressured past the deadline, queued records are + converted into one pending drop count so the writer can report the + loss if the output channel recovers before process exit. The record + currently blocked in an OS write cannot be recovered or counted. + """ + with self._submit_lock: + self._stop_requested.set() + self._thread.join(timeout=_DIAGNOSTIC_SHUTDOWN_TIMEOUT_S) + if not self._thread.is_alive(): + return + + abandoned = 0 + while True: + try: + self._queue.get_nowait() + except queue.Empty: + break + else: + abandoned += 1 + self._queue.task_done() + if abandoned: + self._record_dropped(abandoned) + + +_sink_lock = threading.Lock() +_sink: Optional[_AsyncDiagnosticSink] = None + + +def _reset_sink_after_fork() -> None: + """Discard inherited thread and lock state in a forked child.""" + global _sink, _sink_lock + _sink = None + _sink_lock = threading.Lock() + + +def _get_sink(pid: int) -> _AsyncDiagnosticSink: + """Return a per-process sink, replacing inherited pre-fork state.""" + global _sink + sink = _sink + if sink is not None and sink.pid == pid: + return sink + with _sink_lock: + sink = _sink + if sink is None or sink.pid != pid: + sink = _AsyncDiagnosticSink(pid) + _sink = sink + return sink + + +def _flush_diagnostic_sink_for_tests() -> None: + sink = _sink + if sink is not None: + sink.flush() + + +def _reset_diagnostic_sink_for_tests() -> None: + global _sink + with _sink_lock: + sink = _sink + _sink = None + if sink is not None: + sink.close() + + +def _shutdown_diagnostic_sink() -> None: + sink = _sink + if sink is not None and sink.pid == os.getpid(): + sink.close() + + +atexit.register(_shutdown_diagnostic_sink) +if hasattr(os, "register_at_fork"): + os.register_at_fork(after_in_child=_reset_sink_after_fork) + + +def emit_event( + event: str, + *, + side: str, + request_id: Optional[int], + local_request_id: Optional[int] = None, + rank_info: Optional["RankInfo"] = None, + rank: Optional[int] = None, + instance: Optional[str] = None, + slice_id: Optional[int] = None, + peer_rank: Optional[int] = None, + timestamp: Optional[DiagnosticTimestamp] = None, + **details: DiagnosticValue, +) -> None: + """Emit one compact JSON request-edge event when explicitly enabled. + + Callers must also guard this function with + ``DISAGG_TRANSFER_DIAGNOSTICS_ENABLED`` so argument construction and + request-state inspection are absent from the disabled path. This internal + check keeps accidental unguarded calls inexpensive and harmless. + """ + if not DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + return + + try: + pid = os.getpid() + if timestamp is None: + timestamp = capture_timestamp() + monotonic_ns, wall_ns = timestamp + record: DiagnosticRecord = { + "schema_version": DIAGNOSTICS_SCHEMA_VERSION, + "event": event, + "side": side, + "request_id": request_id, + "local_request_id": local_request_id, + "pid": pid, + "monotonic_ns": monotonic_ns, + "wall_ns": wall_ns, + } + if rank_info is not None: + record.update( + { + "instance": rank_info.instance_name, + "rank": rank_info.instance_rank, + "tp_rank": rank_info.tp_rank, + "pp_rank": rank_info.pp_rank, + "cp_rank": rank_info.cp_rank, + "dp_rank": rank_info.dp_rank, + } + ) + else: + record["instance"] = instance + record["rank"] = rank + if slice_id is not None: + record["slice_id"] = slice_id + if peer_rank is not None: + record["peer_rank"] = peer_rank + record.update(details) + _get_sink(pid).submit(record) + except Exception: + # Diagnostics are best-effort and must never affect request progress. + return diff --git a/tensorrt_llm/_torch/disaggregation/kv_cache_transceiver.py b/tensorrt_llm/_torch/disaggregation/kv_cache_transceiver.py index 9ee46e9f9ef2..1627bc84312a 100644 --- a/tensorrt_llm/_torch/disaggregation/kv_cache_transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/kv_cache_transceiver.py @@ -19,6 +19,7 @@ MambaHybridCacheManagerV2, MixedMambaHybridCacheManager) from ..pyexecutor.llm_request import LlmRequest from ..pyexecutor.resource_manager import KVCacheManager +from . import diagnostics as disagg_diagnostics CacheTransceiverCpp = tensorrt_llm.bindings.internal.batch_manager.CacheTransceiver AttentionTypeCpp = tensorrt_llm.bindings.internal.batch_manager.AttentionType @@ -198,6 +199,7 @@ def create_kv_cache_transceiver( "MambaHybridCacheManagerV2 requires transceiver_runtime='PYTHON' " "with backend='NIXL'; it cannot use the C++ transceiver.") + transceiver: KvCacheTransceiver if use_python_transceiver: if isinstance(mamba_cache_manager, CppMambaHybridCacheManager): raise ValueError( @@ -216,13 +218,27 @@ def create_kv_cache_transceiver( KvCacheTransceiverV2 logger.info("Using KvCacheTransceiverV2") # MixedMambaHybridCacheManager contains both the KV and Mamba pools. - return KvCacheTransceiverV2(mapping, dist, kv_cache_manager, - cache_transceiver_config) - - # Default: use C++ transceiver (transceiver_runtime is None or "CPP") - return BindKvCacheTransceiver(mapping, dist, kv_cache_manager, - attention_type, cache_transceiver_config, - mamba_cache_manager) + transceiver = KvCacheTransceiverV2(mapping, dist, kv_cache_manager, + cache_transceiver_config) + else: + # Default: use C++ transceiver (transceiver_runtime is None or "CPP") + transceiver = BindKvCacheTransceiver(mapping, dist, kv_cache_manager, + attention_type, + cache_transceiver_config, + mamba_cache_manager) + + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + "diagnostic_capabilities", + side="runtime", + request_id=None, + rank=mapping.rank, + capability_schema_version=1, + transceiver_runtime="PYTHON" + if use_python_transceiver else "CPP", + python_transfer_events=use_python_transceiver) + return transceiver class CtxTransferStatus(NamedTuple): diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index b0fa94029892..5b92c52c9773 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -37,6 +37,7 @@ import tensorrt_llm.bindings from tensorrt_llm import logger +from tensorrt_llm._torch.disaggregation import diagnostics as disagg_diagnostics from tensorrt_llm._torch.disaggregation.base.agent import ( BaseTransferAgent, MemoryDescs, @@ -647,6 +648,9 @@ def __init__( self._enforce_physical_ownership = enforce_physical_ownership self._peer_requests: dict = {} self._peer_requests_timestamps: dict[int, float] = {} # unique_rid -> insert time + self._peer_requests_ready_timestamps: dict[ + int, Optional[disagg_diagnostics.DiagnosticTimestamp] + ] = {} self._peer_requests_lock = threading.Lock() self._messenger = ZMQMessenger(mode="ROUTER") self._dealers = {} # used by listener thread only (single-threaded path) @@ -697,6 +701,24 @@ def _is_req_ready(self, unique_rid: int, expected_count: int) -> bool: return False return len(requests) == expected_count + def _get_req_ready_state( + self, unique_rid: int, expected_count: int + ) -> tuple[bool, Optional[disagg_diagnostics.DiagnosticTimestamp]]: + """Return readiness and preserve the timestamp of its first observation.""" + with self._peer_requests_lock: + requests = self._peer_requests.get(unique_rid) + ready = bool(requests) and len(requests) == expected_count + if not ready or not disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + return ready, None + if unique_rid not in self._peer_requests_ready_timestamps: + ready_timestamp = None + with disagg_diagnostics.suppress_diagnostic_errors(): + ready_timestamp = disagg_diagnostics.capture_timestamp() + self._peer_requests_ready_timestamps[unique_rid] = ready_timestamp + else: + ready_timestamp = self._peer_requests_ready_timestamps[unique_rid] + return True, ready_timestamp + def _get_req_info(self, unique_rid: Optional[int]) -> Optional[dict]: with self._peer_requests_lock: return self._peer_requests.get(unique_rid) @@ -712,6 +734,7 @@ def _remove_req_info(self, unique_rid: int): with self._peer_requests_lock: self._peer_requests.pop(unique_rid, None) self._peer_requests_timestamps.pop(unique_rid, None) + self._peer_requests_ready_timestamps.pop(unique_rid, None) def sweep_stale_req_infos(self): """Evict RecvReqInfo entries that have no matching TxSession and exceed the TTL. @@ -734,6 +757,7 @@ def sweep_stale_req_infos(self): if rid not in self._sessions and rid in self._peer_requests: self._peer_requests.pop(rid, None) self._peer_requests_timestamps.pop(rid, None) + self._peer_requests_ready_timestamps.pop(rid, None) logger.debug(f"Swept stale RecvReqInfo for rid={rid}") def setup_session(self, tx_session: "TxSession"): @@ -780,9 +804,28 @@ def setup_session(self, tx_session: "TxSession"): req_info.instance_name, req_info.instance_rank ) expected_count = len(self._registrar.get_peer_overlap(peer_ri, peer_ri.dp_rank).ranks) - if self._is_req_ready(unique_rid, expected_count): + ready, ready_timestamp = self._get_req_ready_state(unique_rid, expected_count) + if ready: + became_ready = False with tx_session.lock: - tx_session.receiver_ready = True + if not tx_session.receiver_ready: + tx_session.receiver_ready = True + became_ready = True + if ( + became_ready + and ready_timestamp is not None + and disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED + ): + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + "ctx_all_receivers_ready", + side="ctx", + request_id=unique_rid, + local_request_id=tx_session.request_id, + rank_info=self._registrar.self_rank_info, + expected_receivers=expected_count, + timestamp=ready_timestamp, + ) return def _get_session(self, unique_rid: Optional[int]) -> Optional["TxSession"]: @@ -818,10 +861,25 @@ def _begin_task_operation(self, task: SendTaskBase, peer_rank: int) -> Optional[ return True def _enqueue(self, write_meta: WriteMeta): + thread_idx = hash((write_meta.unique_rid, write_meta.peer_rank)) % self._num_threads + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + if write_meta.meta_type == WriteMetaType.KV and write_meta.src_ptrs.size > 0: + disagg_diagnostics.emit_event( + "ctx_transfer_queued", + side="ctx", + request_id=write_meta.unique_rid, + rank_info=self._registrar.self_rank_info, + slice_id=write_meta.slice_id, + peer_rank=write_meta.peer_rank, + receiver_slice_id=write_meta.receiver_slice_id, + is_last_slice=write_meta.is_last_slice, + transfer_bytes=int(write_meta.sizes.sum()), + worker_queue_index=thread_idx, + ) # 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) - thread_idx = hash((write_meta.unique_rid, write_meta.peer_rank)) % self._num_threads self._send_task_queues[thread_idx].put(write_meta) def _get_or_connect_thread_dealer(self, endpoint: Optional[str]) -> ZMQMessenger: @@ -851,9 +909,13 @@ def _submit_transfer( task: SendTaskBase, peer_rank: int, request: TransferRequest, + on_submitted: Optional[Callable[[], None]] = None, ) -> tuple[bool, Optional[str]]: if not self._enforce_physical_ownership: status = self._agent.submit_transfer_requests(request) + if on_submitted is not None: + with disagg_diagnostics.suppress_diagnostic_errors(): + on_submitted() completed = status.wait() detail = None if completed else getattr(status, "last_status_str", lambda: None)() return completed, detail @@ -875,6 +937,9 @@ def _submit_transfer( self._ownership_poisoned = error task.mark_physical_operation_in_doubt(peer_rank) return False, str(error) + if on_submitted is not None: + with disagg_diagnostics.suppress_diagnostic_errors(): + on_submitted() try: if not status.wait(): # A non-success query is a logical transfer failure, but the @@ -913,6 +978,24 @@ def _process_task_queue(self, thread_idx: int): except Exception as e: logger.warning(f"failed to send transfer rejection: {e}") continue + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + if ( + write_meta.meta_type == WriteMetaType.KV + and write_meta.src_ptrs.size > 0 + ): + disagg_diagnostics.emit_event( + "ctx_worker_dequeued", + side="ctx", + request_id=write_meta.unique_rid, + rank_info=self._registrar.self_rank_info, + slice_id=write_meta.slice_id, + peer_rank=write_meta.peer_rank, + receiver_slice_id=write_meta.receiver_slice_id, + is_last_slice=write_meta.is_last_slice, + transfer_bytes=int(write_meta.sizes.sum()), + worker_queue_index=thread_idx, + ) try: if write_meta.meta_type == WriteMetaType.AUX: logger.debug( @@ -1044,6 +1127,29 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): agent_result = AgentResult.SUCCESS send_slot_id = None + backend_submitted = False + submission_callback: Optional[Callable[[], None]] = None + diagnostic_transfer_bytes = None + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + diagnostic_transfer_bytes = int(write_meta.sizes.sum()) + + def emit_submitted() -> None: + nonlocal backend_submitted + backend_submitted = True + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + "ctx_backend_submitted", + side="ctx", + request_id=write_meta.unique_rid, + rank_info=self._registrar.self_rank_info, + slice_id=write_meta.slice_id, + peer_rank=write_meta.peer_rank, + receiver_slice_id=write_meta.receiver_slice_id, + transfer_bytes=diagnostic_transfer_bytes, + ) + + submission_callback = emit_submitted if write_meta.src_ptrs.size > 0: try: request, send_slot_id = build_send_request( @@ -1075,8 +1181,25 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): timer.record_transfer_start(write_meta.peer_rank) transfer_finished = False try: + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + "ctx_backend_submit_start", + side="ctx", + request_id=write_meta.unique_rid, + rank_info=self._registrar.self_rank_info, + slice_id=write_meta.slice_id, + peer_rank=write_meta.peer_rank, + receiver_slice_id=write_meta.receiver_slice_id, + transfer_bytes=diagnostic_transfer_bytes, + transfer_entries=int(write_meta.sizes.size), + ) + transfer_finished, last_status = self._submit_transfer( - task, write_meta.peer_rank, request + task, + write_meta.peer_rank, + request, + on_submitted=submission_callback, ) if transfer_finished: del request @@ -1110,6 +1233,24 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): task.retire_unsubmitted_physical_operation(write_meta.peer_rank) if timer: timer.record_transfer_end(write_meta.peer_rank) + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + if ( + write_meta.src_ptrs.size > 0 + and backend_submitted + and agent_result != AgentResult.IN_DOUBT + ): + disagg_diagnostics.emit_event( + "ctx_backend_complete", + side="ctx", + request_id=write_meta.unique_rid, + rank_info=self._registrar.self_rank_info, + slice_id=write_meta.slice_id, + peer_rank=write_meta.peer_rank, + receiver_slice_id=write_meta.receiver_slice_id, + transfer_bytes=diagnostic_transfer_bytes, + outcome=("completed" if agent_result == AgentResult.SUCCESS else "failed"), + ) # Report every chunk so failures reach the receiver immediately. tail = ( @@ -1664,6 +1805,17 @@ def _respond_with_kv(self, _send_id: bytes, message: list[bytes]): # _sessions_lock prevents a race between session lookup and req_info save. # session.lock atomically saves peer info and snapshots tasks against send(). info: RecvReqInfo = RecvReqInfo.from_bytes(message[1]) + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + "ctx_request_data_received", + side="ctx", + request_id=info.unique_rid, + rank_info=self._registrar.self_rank_info, + slice_id=info.slice_id, + peer_rank=info.instance_rank, + peer_instance=info.instance_name, + ) with self._sessions_lock: session = self._get_session(info.unique_rid) if session is None: @@ -1676,7 +1828,7 @@ def _respond_with_kv(self, _send_id: bytes, message: list[bytes]): terminal = True include_aux = bool(session._claim_unsubmitted_aux_failures_locked((info,))) else: - self._save_peer_req_info(info) + became_ready, ready_timestamp = self._save_peer_req_info(info) tasks = list(session.kv_tasks) terminal = session.has_failed() include_aux = terminal and bool( @@ -1688,6 +1840,25 @@ def _respond_with_kv(self, _send_id: bytes, message: list[bytes]): if terminal: self._send_failed_result_to_receiver(info, include_aux=include_aux) return + if ( + became_ready + and ready_timestamp is not None + and disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED + ): + with disagg_diagnostics.suppress_diagnostic_errors(): + peer_ri = self._registrar.get_peer_rank_info(info.instance_name, info.instance_rank) + expected_transfers = len( + self._registrar.get_peer_overlap(peer_ri, peer_ri.dp_rank).ranks + ) + disagg_diagnostics.emit_event( + "ctx_all_receivers_ready", + side="ctx", + request_id=info.unique_rid, + local_request_id=session.request_id, + rank_info=self._registrar.self_rank_info, + expected_receivers=expected_transfers, + timestamp=ready_timestamp, + ) for task in tasks: self._dispatch_task_to_peer(task, info) @@ -1818,15 +1989,21 @@ def _get_or_connect_dealer(self, endpoint: Optional[str]): self._dealers[endpoint] = ZMQMessenger(mode="DEALER", endpoint=endpoint) return self._dealers[endpoint] - def _save_peer_req_info(self, peer_transfer_req_info: RecvReqInfo): + def _save_peer_req_info( + self, peer_transfer_req_info: RecvReqInfo + ) -> tuple[bool, Optional[disagg_diagnostics.DiagnosticTimestamp]]: req_info = peer_transfer_req_info self._add_req_info(req_info.unique_rid, req_info.instance_rank, req_info) peer_ri = self._registrar.get_peer_rank_info(req_info.instance_name, req_info.instance_rank) expected_transfers = len(self._registrar.get_peer_overlap(peer_ri, peer_ri.dp_rank).ranks) - if self._is_req_ready(req_info.unique_rid, expected_transfers): + became_ready = False + ready, ready_timestamp = self._get_req_ready_state(req_info.unique_rid, expected_transfers) + if ready: session = self._get_session(req_info.unique_rid) if session is not None and not session.receiver_ready: session.receiver_ready = True + became_ready = True + return became_ready, ready_timestamp def has_all_peer_req_infos(self, unique_rid: int) -> bool: req_info = self._get_first_req_info(unique_rid) @@ -2735,7 +2912,27 @@ def dispatch_task(self, task: KVRecvTask) -> None: if fanin_bounce: receiver_req.bounce_dst_base = self._bounce.writer_base(key, i) receiver_req_bytes = receiver_req.to_bytes() + send_timestamp = None + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + send_timestamp = disagg_diagnostics.capture_timestamp() self._request_sender_data(peer_infos.sender_endpoints[rank], receiver_req_bytes) + if send_timestamp is not None: + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + "gen_request_data_sent", + side="gen", + request_id=receiver_req.unique_rid, + rank_info=self._registrar.self_rank_info, + slice_id=receiver_req.slice_id, + peer_rank=rank, + expected_writers=task.expected_transfers, + writer_cohort_known=( + sender_dp_rank is not None or peer_infos.dp_size == 1 + ), + ownership_enabled=False, + timestamp=send_timestamp, + ) return peer_ranks = list(peer_overlap.ranks) @@ -2750,17 +2947,26 @@ def dispatch_task(self, task: KVRecvTask) -> None: serialized_requests.append((rank, payload)) published_writers: set[int] = set() + published_writer_timestamps: Optional[ + list[tuple[int, disagg_diagnostics.DiagnosticTimestamp]] + ] = [] if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED else None def publish_requests() -> None: for rank, payload in serialized_requests: if task._perf_timer is not None: task._perf_timer.record_task_start(rank) + send_timestamp = None + if published_writer_timestamps is not None: + with disagg_diagnostics.suppress_diagnostic_errors(): + send_timestamp = disagg_diagnostics.capture_timestamp() published_writers.add(rank) try: self._request_sender_data(peer_infos.sender_endpoints[rank], payload) except Exception: published_writers.discard(rank) raise + if published_writer_timestamps is not None and send_timestamp is not None: + published_writer_timestamps.append((rank, send_timestamp)) # Gen-first ADP publishes the destination to every eligible DP group, # but the qualified immutable-request, no-retry/no-reroute profile @@ -2782,6 +2988,23 @@ def publish_requests() -> None: task.cancel_unpublished() finally: release_admission() + if published_writer_timestamps: + with disagg_diagnostics.suppress_diagnostic_errors(): + for rank, send_timestamp in published_writer_timestamps: + disagg_diagnostics.emit_event( + "gen_request_data_sent", + side="gen", + request_id=receiver_req.unique_rid, + rank_info=self._registrar.self_rank_info, + slice_id=receiver_req.slice_id, + peer_rank=rank, + expected_writers=task.expected_transfers, + writer_cohort_known=( + sender_dp_rank is not None or peer_infos.dp_size == 1 + ), + ownership_enabled=True, + timestamp=send_timestamp, + ) return @staticmethod @@ -2926,16 +3149,36 @@ def _process_kv_agent_result(self, _send_id: bytes, message: list[bytes]): dst_ptrs, sizes, src_base = decode_result_tail(message) session = self._get_session(unique_rid) + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + diagnostic_result = _AGENT_RESULT_BY_CODE.get(status_code) + disagg_diagnostics.emit_event( + "gen_writer_result_received", + side="gen", + request_id=unique_rid, + rank_info=self._registrar.self_rank_info, + slice_id=receiver_slice_id, + peer_rank=peer_rank, + outcome=( + diagnostic_result.value.lower() + if diagnostic_result is not None + else f"unknown:{status_code}" + ), + is_last_slice=is_last_slice, + transfer_bytes=transfer_size, + session_found=session is not None, + ) if session is None: logger.warning( f"_process_kv_agent_result: session {unique_rid} not found (already closed?), dropping status" ) return + agent_result = _AGENT_RESULT_BY_CODE[status_code] session.process_kv_agent_result( peer_rank, receiver_slice_id, is_last_slice, - _AGENT_RESULT_BY_CODE[status_code], + agent_result, dst_ptrs=dst_ptrs, sizes=sizes, src_base=src_base, @@ -3290,6 +3533,10 @@ def process_kv_agent_result( # the scatter worker (after cudaStreamSynchronize) for the bounced path, so the # gen consumer never observes completion before the KV is scattered into place. request_id = self.request_id + disagg_request_id = None + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_request_id = self.disagg_request_id ri = self._receiver._registrar.self_rank_info instance_name, instance_rank = ri.instance_name, ri.instance_rank @@ -3299,6 +3546,8 @@ def on_done( peer_rank=peer_rank, receiver_slice_id=receiver_slice_id, request_id=request_id, + disagg_request_id=disagg_request_id, + rank_info=ri, instance_name=instance_name, instance_rank=instance_rank, ): @@ -3309,12 +3558,32 @@ def on_done( if self._enforce_physical_ownership: task.finish_local_completion() if not success: + destination_timestamp = None + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + destination_timestamp = disagg_diagnostics.capture_timestamp() task.fail( RuntimeError( f"KV bounce scatter failed for request {request_id} " f"slice={receiver_slice_id}" ) ) + if ( + destination_timestamp is not None + and disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED + ): + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + "gen_destination_complete", + side="gen", + request_id=disagg_request_id, + local_request_id=request_id, + rank_info=rank_info, + slice_id=receiver_slice_id, + peer_rank=peer_rank, + outcome="failed", + timestamp=destination_timestamp, + ) return if task.status == TaskStatus.ERROR: return # a concurrent FAILED writer already failed it; don't un-fail @@ -3327,12 +3596,32 @@ def on_done( f"KV transfer perf logging failed for request {request_id} " f"slice={receiver_slice_id}: {e}" ) + destination_timestamp = None + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + destination_timestamp = disagg_diagnostics.capture_timestamp() task.complete() # Transfer end for perf/time-sync: only meaningful once every slice has # landed. Plain attribute write (atomic under the GIL); on_done must stay # lock-free, and consumers only read it after wait_complete succeeds. if all(t.status == TaskStatus.TRANSFERRED for t in self._kv_tasks): self.transfer_end_time = tensorrt_llm.bindings.global_steady_clock_now() + if ( + destination_timestamp is not None + and disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED + ): + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + "gen_destination_complete", + side="gen", + request_id=disagg_request_id, + local_request_id=request_id, + rank_info=rank_info, + slice_id=receiver_slice_id, + peer_rank=peer_rank, + outcome="completed", + timestamp=destination_timestamp, + ) logger.debug( f"KV transfer complete for request {request_id} " f"slice={receiver_slice_id}" diff --git a/tensorrt_llm/_torch/disaggregation/orchestration/coordinator.py b/tensorrt_llm/_torch/disaggregation/orchestration/coordinator.py index 05ffb8e895ff..8c0eefbe15f4 100644 --- a/tensorrt_llm/_torch/disaggregation/orchestration/coordinator.py +++ b/tensorrt_llm/_torch/disaggregation/orchestration/coordinator.py @@ -13,6 +13,7 @@ from dataclasses import dataclass, fields from typing import TYPE_CHECKING, Callable, List, Set, Tuple +from tensorrt_llm._torch.disaggregation import diagnostics as disagg_diagnostics from tensorrt_llm._torch.disaggregation.base.transfer import get_unique_rid from tensorrt_llm._torch.disaggregation.kv_cache_transceiver import ( is_disagg_inflight_cancel_enabled, @@ -228,16 +229,33 @@ def flag_if_timed_out(req: LlmRequest, kind: str) -> None: return elapsed_ms = (time.monotonic() - req.py_kv_transfer_start_time) * 1000 if elapsed_ms > timeout_ms and not req.py_kv_transfer_timed_out: - verb = ( - "Requesting cancellation for" - if self.inflight_cancel_active() - else "Observed timeout on" - ) + cancel_enabled = self.inflight_cancel_active() + verb = "Requesting cancellation for" if cancel_enabled else "Observed timeout on" logger.warning( f"{verb} {kind} request {req.py_request_id} due to KV cache " f"transfer timeout: elapsed {elapsed_ms:.0f}ms > " f"kv_transfer_timeout_ms={timeout_ms}ms" ) + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + "transfer_timeout_observed", + side="ctx" if kind == "context" else "gen", + request_id=get_unique_rid(req), + local_request_id=req.py_request_id, + rank=self._dist.rank, + elapsed_ms=elapsed_ms, + timeout_ms=timeout_ms, + timeout_owner="pyexecutor", + timer_start_monotonic_ns=int( + req.py_kv_transfer_start_time * 1_000_000_000 + ), + state=req.state.name, + cancellation_requested=cancel_enabled, + tp_rank=self._dist.tp_rank, + pp_rank=self._dist.pp_rank, + cp_rank=self._dist.cp_rank, + ) req.py_kv_transfer_timed_out = True # Context requests start their clock on the last chunk, which is also @@ -321,6 +339,23 @@ def send_completed_context(self, requests: List[LlmRequest]) -> None: # sends the final slice and (for the Python transceiver) moves # the request toward completion. self._transfers.start_transfer(req) + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + "ctx_send_ready", + side="ctx", + request_id=get_unique_rid(req), + local_request_id=req.py_request_id, + rank=self._dist.rank, + prompt_tokens=req.prompt_len, + state=req.state.name, + source_kv_request_owned=True, + source_kv_reuse_pinned=self._transfers.should_store_blocks, + timeout_expected=self._transceiver.kv_transfer_timeout_ms is not None, + tp_rank=self._dist.tp_rank, + pp_rank=self._dist.pp_rank, + cp_rank=self._dist.cp_rank, + ) self._transceiver.respond_and_send_async(req) # Bridge validation can reject before a transfer session exists. # Release the claim right away: there is no physical accessor @@ -331,9 +366,43 @@ def send_completed_context(self, requests: List[LlmRequest]) -> None: and not self._transceiver.has_inflight_transfer(req) ): self.release_transfer(req) + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + "ctx_transfer_settled", + side="ctx", + request_id=get_unique_rid(req), + local_request_id=req.py_request_id, + rank=self._dist.rank, + instance=getattr(self._transceiver, "_instance_name", None), + outcome="failed", + session_status=None, + resources_drained=True, + tp_rank=self._dist.tp_rank, + pp_rank=self._dist.pp_rank, + cp_rank=self._dist.cp_rank, + dp_rank=getattr(self._dist, "dp_rank", None), + ) continue if self._transceiver.kv_transfer_timeout_ms is not None: - req.py_kv_transfer_start_time = time.monotonic() + timeout_start = time.monotonic() + req.py_kv_transfer_start_time = timeout_start + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + "transfer_timeout_started", + side="ctx", + request_id=get_unique_rid(req), + local_request_id=req.py_request_id, + rank=self._dist.rank, + timeout_ms=self._transceiver.kv_transfer_timeout_ms, + timeout_owner="pyexecutor", + timer_start_monotonic_ns=int(timeout_start * 1_000_000_000), + state=req.state.name, + tp_rank=self._dist.tp_rank, + pp_rank=self._dist.pp_rank, + cp_rank=self._dist.cp_rank, + ) elif ( self._transceiver.pipeline_transfer_enabled and req.state != LlmRequestState.GENERATION_COMPLETE @@ -541,6 +610,26 @@ def _cancel_timed_out_gen_transfers(self) -> None: continue elapsed_ms = (current_time - request.py_kv_transfer_start_time) * 1000 if elapsed_ms > timeout_ms and not request.py_kv_transfer_timed_out: + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + "transfer_timeout_observed", + side="gen", + request_id=get_unique_rid(request), + local_request_id=request.py_request_id, + rank=self._dist.rank, + elapsed_ms=elapsed_ms, + timeout_ms=timeout_ms, + timeout_owner="pyexecutor", + timer_start_monotonic_ns=int( + request.py_kv_transfer_start_time * 1_000_000_000 + ), + state=request.state.name, + cancellation_requested=True, + tp_rank=self._dist.tp_rank, + pp_rank=self._dist.pp_rank, + cp_rank=self._dist.cp_rank, + ) logger.warning( f"Requesting cancellation for generation request " f"{request.py_request_id} due to KV cache transfer timeout" diff --git a/tensorrt_llm/_torch/disaggregation/orchestration/transfer_manager.py b/tensorrt_llm/_torch/disaggregation/orchestration/transfer_manager.py index b487311acfc9..c24ed41bea2c 100644 --- a/tensorrt_llm/_torch/disaggregation/orchestration/transfer_manager.py +++ b/tensorrt_llm/_torch/disaggregation/orchestration/transfer_manager.py @@ -3,6 +3,8 @@ from typing import Dict, Optional +from tensorrt_llm._torch.disaggregation import diagnostics as disagg_diagnostics +from tensorrt_llm._torch.disaggregation.base.transfer import get_unique_rid from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, LlmRequestState from tensorrt_llm._torch.pyexecutor.resource_manager import ResourceManager, ResourceManagerType from tensorrt_llm.logger import logger @@ -23,7 +25,7 @@ class AsyncTransferManager: """ class RequestTransferMetadata: - def __init__(self, block_id: Optional[int]): + def __init__(self, block_id: Optional[list[int]]): self.block_id = block_id self.counter = 0 @@ -116,6 +118,22 @@ def end_transfer(self, request: LlmRequest) -> bool: if request.state != LlmRequestState.DISAGG_TRANS_ERROR: request.state = LlmRequestState.DISAGG_CONTEXT_COMPLETE + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + if self.should_store_blocks and request.is_context_only_request: + assert transfer_metadata.block_id is not None + mapping = getattr(self.kv_cache_manager, "mapping", None) + disagg_diagnostics.emit_event( + "ctx_source_unpinned", + side="ctx", + request_id=get_unique_rid(request), + local_request_id=request.py_request_id, + rank=getattr(mapping, "rank", None), + source_kv_reuse_pinned=False, + state=request.state.name, + source_kv_reuse_block_count=len(transfer_metadata.block_id), + ) + return True return False diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index fbde326401c0..39650591cc54 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -25,6 +25,7 @@ import tensorrt_llm.bindings from tensorrt_llm import logger +from tensorrt_llm._torch.disaggregation import diagnostics as disagg_diagnostics from tensorrt_llm._torch.disaggregation.base.agent import use_pure_python_transfer_agent from tensorrt_llm._torch.disaggregation.base.transfer import ( KVSlice, @@ -834,6 +835,32 @@ def _close_session_or_raise(self, session: object, rid: int, outcome: str) -> No f"refusing to retire {outcome} KV transfer rid={rid}: session close refused" ) + def _emit_transfer_settled( + self, + side: str, + rid: Optional[int], + req: Optional[LlmRequest], + session: Optional[object], + outcome: str, + ) -> None: + """Report a transfer only after its session has been physically retired.""" + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + f"{side}_transfer_settled", + side=side, + request_id=rid, + local_request_id=(req.py_request_id if req is not None else None), + rank=self._mapping.rank, + instance=self._instance_name, + outcome=outcome, + session_status=(session.status.value if session is not None else None), + resources_drained=(session is None or not session.has_transferring_tasks()), + tp_rank=self._mapping.tp_rank, + pp_rank=self._mapping.pp_rank, + cp_rank=self._mapping.cp_rank, + dp_rank=self._dp_rank, + ) + def _apply_aux(self, session, req: LlmRequest): """Unpack aux tokens from session into request's context_phase_params.""" session.unpack_aux(req) @@ -1011,31 +1038,64 @@ def request_and_receive_sync(self, req: LlmRequest) -> None: return req.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS session = None + receive_started = False + settlement_outcome = None try: session = self._transfer_worker.create_rx_session(req) self._recv_sessions[rid] = session self._recv_reqs[rid] = req kv_slice = self._create_kv_slice(req) + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + diagnostic_transfer_bytes = ( + self._slice_num_bytes(kv_slice) * self._kv_size_rank_factor + ) + disagg_diagnostics.emit_event( + "gen_receive_start", + side="gen", + request_id=rid, + local_request_id=req.py_request_id, + rank=self._mapping.rank, + instance=self._instance_name, + slice_id=0, + transfer_bytes=diagnostic_transfer_bytes, + timeout_expected=False, + tp_rank=self._mapping.tp_rank, + pp_rank=self._mapping.pp_rank, + cp_rank=self._mapping.cp_rank, + dp_rank=self._dp_rank, + ) + receive_started = True session.receive(kv_slice) result = session.wait_complete(blocking=True) if result == WaitResult.COMPLETED: # KV-transfer timing setters deferred to #15871 (clock-source consistency); size only. - req.set_kv_cache_size(self._slice_num_bytes(kv_slice) * self._kv_size_rank_factor) + transfer_bytes = self._slice_num_bytes(kv_slice) * self._kv_size_rank_factor + req.set_kv_cache_size(transfer_bytes) if self._need_aux_transfer(req): self._apply_aux(session, req) self._assert_disagg_history_declared(req) req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE + settlement_outcome = "completed" else: req.state = LlmRequestState.DISAGG_TRANS_ERROR + settlement_outcome = "failed" except Exception: req.state = LlmRequestState.DISAGG_TRANS_ERROR + settlement_outcome = "failed" raise finally: close_succeeded = session is None or session.close() is not False if close_succeeded: self._recv_sessions.pop(rid, None) self._recv_reqs.pop(rid, None) + if ( + disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED + and receive_started + and settlement_outcome is not None + ): + self._emit_transfer_settled("gen", rid, req, session, settlement_outcome) else: logger.error( f"request_and_receive_sync: retaining rid={rid} because receive " @@ -1074,6 +1134,23 @@ def request_and_receive_async(self, req: LlmRequest) -> None: try: kv_slice = self._create_kv_slice(req) req.py_kv_cache_xfer_bytes = self._slice_num_bytes(kv_slice) * self._kv_size_rank_factor + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + "gen_receive_start", + side="gen", + request_id=rid, + local_request_id=req.py_request_id, + rank=self._mapping.rank, + instance=self._instance_name, + slice_id=0, + transfer_bytes=req.py_kv_cache_xfer_bytes, + timeout_expected=self.kv_transfer_timeout_ms is not None, + tp_rank=self._mapping.tp_rank, + pp_rank=self._mapping.pp_rank, + cp_rank=self._mapping.cp_rank, + dp_rank=self._dp_rank, + ) session.receive(kv_slice) except Exception as error: if bridge_enabled: @@ -1149,16 +1226,43 @@ def check_context_transfer_status( quiesced = set(quiesced_ids) cancelled = [rid for rid in cancelled if rid in quiesced] failed = [rid for rid in failed if rid in quiesced] + diagnostics_enabled = disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED for rid in cancelled: + diagnostic_record = None + if diagnostics_enabled: + with disagg_diagnostics.suppress_diagnostic_errors(): + diagnostic_record = ( + self._send_reqs.get(rid), + self._send_sessions[rid], + ) self._retire_send_session(rid, outcome="cancelled") + if diagnostic_record is not None: + req, session = diagnostic_record + self._emit_transfer_settled("ctx", rid, req, session, "cancelled") for rid in completed: req = self._send_reqs[rid] + diagnostic_session = None + if diagnostics_enabled: + with disagg_diagnostics.suppress_diagnostic_errors(): + diagnostic_session = self._send_sessions[rid] self._retire_send_session(rid, outcome="completed") if mark_complete: req.state = LlmRequestState.DISAGG_CONTEXT_COMPLETE + if diagnostic_session is not None: + self._emit_transfer_settled("ctx", rid, req, diagnostic_session, "completed") + + failed_records = [] + if diagnostics_enabled: + with disagg_diagnostics.suppress_diagnostic_errors(): + for rid in failed: + session = self._send_sessions.get(rid) + if session is not None: + failed_records.append((rid, self._send_reqs.get(rid), session)) self._close_failed_sessions(self._send_sessions, self._send_reqs, failed, mark_retired=True) + for rid, req, session in failed_records: + self._emit_transfer_settled("ctx", rid, req, session, "failed") # Sweep orphaned RecvReqInfo entries from ADP broadcast on non-assigned # DP ranks (entries that will never have a TxSession created for them). @@ -1222,14 +1326,18 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]) -> GenT cancelled, failed, completed = self._gen_consensus_outcome( to_process, cancelled, failed, completed ) + diagnostics_enabled = disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED cancelled_reqs = [] for rid in cancelled: session = self._recv_sessions[rid] + req = self._recv_reqs[rid] self._close_session_or_raise(session, rid, "cancelled") - cancelled_reqs.append(self._recv_reqs[rid]) + cancelled_reqs.append(req) del self._recv_reqs[rid] del self._recv_sessions[rid] + if diagnostics_enabled: + self._emit_transfer_settled("gen", rid, req, session, "cancelled") # Log gen-side transfer summary after consensus. if completed and os.getenv("TRTLLM_KVCACHE_TIME_OUTPUT_PATH"): @@ -1257,12 +1365,23 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]) -> GenT req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE del self._recv_reqs[rid] del self._recv_sessions[rid] + if diagnostics_enabled: + self._emit_transfer_settled("gen", rid, req, session, "completed") if failed: logger.warning( f"Disagg gen transfer FAILED rank={self._dist.rank} " f"rids={failed} gen_need_sync={self._gen_need_sync}" ) + failed_records = [] + if diagnostics_enabled: + with disagg_diagnostics.suppress_diagnostic_errors(): + for rid in failed: + session = self._recv_sessions.get(rid) + if session is not None: + failed_records.append((rid, self._recv_reqs.get(rid), session)) self._close_failed_sessions(self._recv_sessions, self._recv_reqs, failed) + for rid, req, session in failed_records: + self._emit_transfer_settled("gen", rid, req, session, "failed") return GenTransferStatus(completed, failed, cancelled_reqs) @@ -1350,16 +1469,39 @@ def cancel_request(self, req: LlmRequest) -> bool: """ rid = get_unique_rid(req) + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + for side, sessions in ( + ("ctx", self._send_sessions), + ("gen", self._recv_sessions), + ): + if rid not in sessions: + continue + disagg_diagnostics.emit_event( + "transfer_cancel_requested", + side=side, + request_id=rid, + local_request_id=req.py_request_id, + rank=self._mapping.rank, + instance=self._instance_name, + session_status=sessions[rid].status.value, + tp_rank=self._mapping.tp_rank, + pp_rank=self._mapping.pp_rank, + cp_rank=self._mapping.cp_rank, + dp_rank=self._dp_rank, + ) + # 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(): + session = self._send_sessions[rid] + session.cancel() + if session.has_transferring_tasks(): has_transferring = True - elif self._send_sessions[rid].close() is False: + elif session.close() is False: has_transferring = True else: self._retire_send_session( @@ -1368,16 +1510,21 @@ def cancel_request(self, req: LlmRequest) -> bool: outcome="cancelled", session_already_closed=True, ) + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + self._emit_transfer_settled("ctx", rid, req, session, "cancelled") if rid in self._recv_sessions: - self._recv_sessions[rid].cancel() - if self._recv_sessions[rid].has_transferring_tasks(): + session = self._recv_sessions[rid] + session.cancel() + if session.has_transferring_tasks(): has_transferring = True - elif self._recv_sessions[rid].close() is False: + elif session.close() is False: has_transferring = True else: del self._recv_reqs[rid] del self._recv_sessions[rid] + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + self._emit_transfer_settled("gen", rid, req, session, "cancelled") if has_transferring: return False diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index dc2f60fceb6d..394d08efeec4 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -52,6 +52,7 @@ from tensorrt_llm.tools.profiler.host_profile_tools.host_profiler import \ host_profiler_context +from ..disaggregation import diagnostics as disagg_diagnostics from ..disaggregation.base.transfer import get_unique_rid from ..disaggregation.kv_cache_transceiver import KvCacheTransceiver from ..disaggregation.orchestration.admission import \ @@ -107,7 +108,8 @@ ResourceManagerType, request_context) from .sampler import (AsyncWorkerMixin, Sampler, SamplerEvent, SampleState, SampleStateTensors) -from .scheduler import (RequestScheduler, ScheduledRequests, +from .scheduler import (KVCacheV2Scheduler, MultimodalScheduler, + RequestScheduler, ScheduledRequests, SerializableSchedulerOutput, WaitingQueue, create_waiting_queue) from .scheduler.adp_router import ADPRouter @@ -948,6 +950,7 @@ def on_detected(): if kv_cache_transceiver is not None: self.hang_detector.register_status_provider( kv_cache_transceiver.get_status_dump) + self._emit_disagg_diagnostic_capabilities() cache_transceiver_config = getattr(self.llm_args, "cache_transceiver_config", None) max_tokens_in_buffer = getattr(cache_transceiver_config, @@ -1038,6 +1041,25 @@ def on_detected(): if start_worker: self.start_worker() + def _emit_disagg_diagnostic_capabilities(self) -> None: + """Declare this executor's event groups once, before serving requests.""" + if (not disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED + or self.kv_cache_transceiver is None): + return + with disagg_diagnostics.suppress_diagnostic_errors(): + scheduler = self.scheduler + while isinstance(scheduler, MultimodalScheduler): + scheduler = scheduler.scheduler + disagg_diagnostics.emit_event( + "diagnostic_capabilities", + side="runtime", + request_id=None, + rank=self.global_rank, + capability_schema_version=1, + executor_events=True, + scheduler_kv_admission_events=isinstance( + scheduler, KVCacheV2Scheduler)) + def _maybe_init_kv_connector_manager(self): if self.kv_connector_manager is not None: if self.kv_cache_transceiver is not None: @@ -3645,11 +3667,59 @@ def _apply_disagg_transfer_admission( # Real synchronous gen_only transfers still honor the budget to bound # the number of blocking transfers started in one executor iteration. if self._is_disagg_gen_only_no_context_benchmark(): + if (disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED + and fitting_disagg_gen_init_requests): + with disagg_diagnostics.suppress_diagnostic_errors(): + self._emit_disagg_transfer_window_results( + fitting_disagg_gen_init_requests, + fitting_disagg_gen_init_requests, + policy="not_applicable", + ) return fitting_disagg_gen_init_requests, False controller = self._get_disagg_transfer_admission_controller() if not (self._disagg_transfer_window_is_active() and fitting_disagg_gen_init_requests): + if (disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED + and fitting_disagg_gen_init_requests): + with disagg_diagnostics.suppress_diagnostic_errors(): + if not controller.enabled(): + policy = "disabled" + elif self._is_disagg_transfer_window_bypass_eligible(): + policy = "bypassed" + else: + policy = "inactive" + legacy_result = None + if policy == "bypassed": + # This opt-in counterfactual scans active requests and + # candidates on every bypass admission decision. Keep + # diagnostics disabled outside targeted debugging runs. + try: + legacy_result = controller.select( + self.active_requests, + fitting_disagg_gen_init_requests) + except Exception: + # Preserve the actual bypass result when the + # counterfactual diagnostic cannot be computed. + pass + self._emit_disagg_transfer_window_results( + fitting_disagg_gen_init_requests, + fitting_disagg_gen_init_requests, + policy=policy, + controller=controller, + legacy_admitted_requests=( + legacy_result.admitted_requests + if legacy_result is not None else None), + legacy_active_transfer_blocks=( + legacy_result.active_transfer_blocks + if legacy_result is not None else None), + legacy_admitted_transfer_blocks=( + legacy_result.admitted_transfer_blocks + if legacy_result is not None else None), + legacy_limited_by_budget=( + legacy_result.limited_by_budget + if legacy_result is not None else None), + ) return fitting_disagg_gen_init_requests, False admission_result = controller.select(self.active_requests, @@ -3663,16 +3733,95 @@ def _apply_disagg_transfer_admission( f"{admission_result.admitted_transfer_blocks}, " f"budget={controller.max_transfer_blocks}") + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + self._emit_disagg_transfer_window_results( + fitting_disagg_gen_init_requests, + admission_result.admitted_requests, + policy="enforced", + controller=controller, + active_transfer_blocks=admission_result. + active_transfer_blocks, + admitted_transfer_blocks=admission_result. + admitted_transfer_blocks, + ) + self._revert_deferred_disagg_gen_init_alloc( fitting_disagg_gen_init_requests, - admission_result.admitted_requests) + admission_result.admitted_requests, + reason="transfer_window", + ) return (admission_result.admitted_requests, admission_result.is_blocked_by_active_transfers()) + def _emit_disagg_transfer_window_results( + self, + candidates: List[LlmRequest], + admitted_requests: List[LlmRequest], + *, + policy: str, + controller: Optional[DisaggTransferAdmissionController] = None, + active_transfer_blocks: Optional[int] = None, + admitted_transfer_blocks: Optional[int] = None, + legacy_admitted_requests: Optional[List[LlmRequest]] = None, + legacy_active_transfer_blocks: Optional[int] = None, + legacy_admitted_transfer_blocks: Optional[int] = None, + legacy_limited_by_budget: Optional[bool] = None) -> None: + """Emit per-request transfer-window decisions for offline analysis.""" + if not disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + return + + with disagg_diagnostics.suppress_diagnostic_errors(): + admitted_ids = { + request.py_request_id + for request in admitted_requests + } + legacy_admitted_ids = ({ + request.py_request_id + for request in legacy_admitted_requests + } if legacy_admitted_requests is not None else None) + tokens_per_block = getattr(controller, "tokens_per_block", 0) + for request in candidates: + prompt_tokens = getattr(request, "total_input_len_cp", None) + if prompt_tokens is None: + prompt_tokens = request.prompt_len + request_blocks = ((prompt_tokens + tokens_per_block - 1) // + tokens_per_block + if tokens_per_block > 0 else None) + disagg_diagnostics.emit_event( + "gen_transfer_window_result", + side="gen", + request_id=get_unique_rid(request), + local_request_id=request.py_request_id, + rank=self.global_rank, + outcome=("admitted" if request.py_request_id in admitted_ids + else "deferred"), + policy=policy, + prompt_tokens=prompt_tokens, + request_blocks=request_blocks, + active_transfer_blocks=active_transfer_blocks, + admitted_transfer_blocks=admitted_transfer_blocks, + transfer_block_budget=getattr(controller, + "max_transfer_blocks", None), + legacy_budget_outcome=( + "admitted" if request.py_request_id + in legacy_admitted_ids else "deferred") + if legacy_admitted_ids is not None else None, + legacy_active_transfer_blocks=legacy_active_transfer_blocks, + legacy_admitted_transfer_blocks= + legacy_admitted_transfer_blocks, + legacy_limited_by_budget=legacy_limited_by_budget, + tp_rank=self.dist.tp_rank, + pp_rank=self.dist.pp_rank, + cp_rank=self.dist.cp_rank, + ) + def _revert_deferred_disagg_gen_init_alloc( - self, candidates: List[LlmRequest], - admitted_requests: List[LlmRequest]) -> None: + self, + candidates: List[LlmRequest], + admitted_requests: List[LlmRequest], + reason: str = "pp_reconciliation") -> None: """Revert Scheduler V2 allocations absent from an admitted request set. Scheduler V2 allocates KV while evaluating generation-init requests. @@ -3693,6 +3842,20 @@ def _revert_deferred_disagg_gen_init_alloc( ] if deferred_requests: self._revert_ctx_alloc(deferred_requests) + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + for request in deferred_requests: + disagg_diagnostics.emit_event( + "gen_kv_rollback", + side="gen", + request_id=get_unique_rid(request), + local_request_id=request.py_request_id, + rank=self.global_rank, + reason=reason, + tp_rank=self.dist.tp_rank, + pp_rank=self.dist.pp_rank, + cp_rank=self.dist.cp_rank, + ) @staticmethod def _dist_size(dist, name: str) -> int: @@ -6039,6 +6202,26 @@ def _respond_if_invalid(request: LlmRequest) -> bool: ] self.active_requests.extend(validated_requests) + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + for request in validated_requests: + if not request.is_disagg_generation_init_state: + continue + prompt_tokens = getattr(request, "total_input_len_cp", None) + if prompt_tokens is None: + prompt_tokens = request.prompt_len + disagg_diagnostics.emit_event( + "gen_ingress", + side="gen", + request_id=get_unique_rid(request), + local_request_id=request.py_request_id, + rank=self.global_rank, + prompt_tokens=prompt_tokens, + state=request.state.name, + tp_rank=self.dist.tp_rank, + pp_rank=self.dist.pp_rank, + cp_rank=self.dist.cp_rank, + ) return validated_requests def _add_kv_cache_events(self): @@ -7239,6 +7422,22 @@ def _prepare_disagg_gen_transmission_complete(self, scheduled_batch): req.add_new_token(first_gen_tokens[beam], beam) self._maybe_prepend_logprobs_and_logits(req, beam_width) + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + prompt_tokens = getattr(req, "total_input_len_cp", None) + if prompt_tokens is None: + prompt_tokens = req.prompt_len + disagg_diagnostics.emit_event( + "gen_decode_ready", + side="gen", + request_id=get_unique_rid(req), + local_request_id=req.py_request_id, + rank=self.global_rank, + prompt_tokens=prompt_tokens, + tp_rank=self.dist.tp_rank, + pp_rank=self.dist.pp_rank, + cp_rank=self.dist.cp_rank, + ) def _update_sampler_state_for_disagg_gen_request(self, req, beam_width, first_gen_tokens) -> bool: @@ -7421,7 +7620,26 @@ def _recv_disagg_gen_cache(self, new_gen_reqs): if self.kv_cache_transceiver.kv_transfer_timeout_ms is not None: for req in new_gen_reqs: if req.state == LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS: - req.py_kv_transfer_start_time = time.monotonic() + timeout_start = time.monotonic() + req.py_kv_transfer_start_time = timeout_start + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + disagg_diagnostics.emit_event( + "transfer_timeout_started", + side="gen", + request_id=get_unique_rid(req), + local_request_id=req.py_request_id, + rank=self.global_rank, + timeout_ms=self.kv_cache_transceiver. + kv_transfer_timeout_ms, + timeout_owner="pyexecutor", + timer_start_monotonic_ns=int(timeout_start * + 1_000_000_000), + state=req.state.name, + tp_rank=self.dist.tp_rank, + pp_rank=self.dist.pp_rank, + cp_rank=self.dist.cp_rank, + ) self.disagg.reap_gen_receives(0) @@ -7952,6 +8170,22 @@ def _terminate_request(self, request: LlmRequest) -> None: def _free_request_resources(self, request: LlmRequest) -> None: """Release execution resources without removing response routing.""" self.resource_manager.free_resources(request) + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + if request.is_context_only_request: + disagg_diagnostics.emit_event( + "ctx_source_kv_released", + side="ctx", + request_id=get_unique_rid(request), + local_request_id=request.py_request_id, + rank=self.global_rank, + prompt_tokens=request.prompt_len, + state=request.state.name, + source_kv_request_owned=False, + tp_rank=self.dist.tp_rank, + pp_rank=self.dist.pp_rank, + cp_rank=self.dist.cp_rank, + ) self._prefetched_request_ids.discard(request.py_request_id) self.disagg.forget_request(request.py_request_id) diff --git a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py index 1996d219425a..6d28c5e052bd 100644 --- a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py @@ -20,6 +20,8 @@ from tensorrt_llm.llmapi.llm_args import CapacitySchedulerPolicy, ContextChunkingPolicy from tensorrt_llm.logger import logger +from ...disaggregation import diagnostics as disagg_diagnostics +from ...disaggregation.base.transfer import get_unique_rid from ..llm_request import LlmRequest, LlmRequestState, get_draft_token_length from .scheduler import ( RequestList, @@ -271,6 +273,10 @@ def schedule_request( has_chunking, ) = self._schedule_loop(active_requests, inflight_request_ids) + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + self._emit_disagg_kv_pool_snapshot(active_requests, disagg_candidates) + # Sort by LoRA task ID scheduled_encoder.sort(key=_get_lora_task_id) self._sort_requests(scheduled_ctx, scheduled_gen, has_chunking) @@ -285,6 +291,54 @@ def schedule_request( num_fitting_requests=(len(scheduled_encoder) + len(scheduled_ctx) + len(scheduled_gen)), ) + def _emit_disagg_kv_pool_snapshot( + self, + active_requests: RequestList, + disagg_candidates: RequestList, + ) -> None: + """Emit one post-scheduler physical KV-pool snapshot when relevant.""" + if not disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + return + + with disagg_diagnostics.suppress_diagnostic_errors(): + init_requests = sum( + request.is_disagg_generation_init_state for request in active_requests + ) + transfers_in_progress = sum( + request.is_disagg_generation_transmission_in_progress for request in active_requests + ) + transfers_complete = sum( + request.is_disagg_generation_transmission_complete for request in active_requests + ) + if not (init_requests or transfers_in_progress or transfers_complete): + return + + stats = self.kv_cache_manager.get_kv_cache_stats() + mapping = self.kv_cache_manager.mapping + index_mapper = getattr(self.kv_cache_manager, "index_mapper", None) + num_free_slots = getattr(index_mapper, "num_free_slots", None) + disagg_diagnostics.emit_event( + "gen_kv_pool_snapshot", + side="gen", + request_id=None, + rank=mapping.rank, + init_requests=init_requests, + transfers_in_progress=transfers_in_progress, + transfers_complete=transfers_complete, + kv_admitted_this_iteration=len(disagg_candidates), + decode_requests=sum( + request.state == LlmRequestState.GENERATION_IN_PROGRESS + for request in active_requests + ), + kv_pool_max_blocks=stats.max_num_blocks, + kv_pool_free_blocks=stats.free_num_blocks, + kv_pool_used_blocks=stats.used_num_blocks, + index_free_slots=(num_free_slots() if callable(num_free_slots) else None), + tp_rank=mapping.tp_rank, + pp_rank=mapping.pp_rank, + cp_rank=mapping.cp_rank, + ) + # ---- Main scheduling loop ---- def _schedule_loop(self, active_requests, inflight_request_ids): @@ -700,7 +754,36 @@ def _try_schedule_disagg_gen_init( # Cache-transceiver mode disables the separate one-model draft manager, # so disagg generation init has no paired-reuse path. Supporting one # would also require draft KV transfer and history_length=prompt_len. - if not self.kv_cache_manager.prepare_disagg_gen_init(req): + prepared = self.kv_cache_manager.prepare_disagg_gen_init(req) + if disagg_diagnostics.DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + with disagg_diagnostics.suppress_diagnostic_errors(): + kv_cache = self.kv_cache_manager.kv_cache_map.get(req.py_request_id) + capacity_tokens = getattr(kv_cache, "capacity", None) + history_tokens = getattr(kv_cache, "history_length", None) + capacity_block_equivalent = ( + (capacity_tokens + self.tokens_per_block - 1) // self.tokens_per_block + if capacity_tokens is not None and self.tokens_per_block > 0 + else None + ) + prompt_tokens = getattr(req, "total_input_len_cp", None) + if prompt_tokens is None: + prompt_tokens = req.prompt_len + disagg_diagnostics.emit_event( + "gen_kv_admission_result", + side="gen", + request_id=get_unique_rid(req), + local_request_id=req.py_request_id, + rank=self.kv_cache_manager.mapping.rank, + outcome="admitted" if prepared else "deferred", + reason=None if prepared else "kv_or_index_capacity", + prompt_tokens=prompt_tokens, + tokens_per_block=self.tokens_per_block, + cache_present=kv_cache is not None, + capacity_tokens=capacity_tokens, + history_tokens=history_tokens, + capacity_block_equivalent=capacity_block_equivalent, + ) + if not prepared: logger.debug("prepare_disagg_gen_init failed for request %s", req.py_request_id) return ScheduleAction.SKIP, 0 return ScheduleAction.SCHEDULED, 0 diff --git a/tests/unittest/_torch/disaggregation/test_disagg_coordinator_diagnostics.py b/tests/unittest/_torch/disaggregation/test_disagg_coordinator_diagnostics.py new file mode 100644 index 000000000000..add41264fe77 --- /dev/null +++ b/tests/unittest/_torch/disaggregation/test_disagg_coordinator_diagnostics.py @@ -0,0 +1,214 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Telemetry through the progress/error paths moved into the coordinator. + +Use the real coordinator and transfer manager with the existing contract fakes. +FakeDist checks collective ordering and payloads, not real GPU communication. +""" + +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +from coordinator_harness import CoordinatorHarness, TransferRequest +from fake_dist import FakeDistGroup + +from tensorrt_llm._torch.disaggregation import diagnostics as disagg_diagnostics +from tensorrt_llm._torch.disaggregation.orchestration import coordinator as coordinator_module +from tensorrt_llm.bindings import LlmRequestState + +pytestmark = pytest.mark.cpu_only + + +@pytest.fixture(params=["disabled", "enabled", "raises"]) +def diagnostics(request, monkeypatch) -> tuple[bool, Mock]: + enabled = request.param != "disabled" + emit = Mock( + side_effect=RuntimeError("diagnostic sink failed") if request.param == "raises" else None + ) + monkeypatch.setattr(disagg_diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", enabled) + monkeypatch.setattr(disagg_diagnostics, "emit_event", emit) + monkeypatch.setattr(coordinator_module, "is_disagg_inflight_cancel_enabled", lambda: False) + monkeypatch.delenv("TRTLLM_DISAGG_BENCHMARK_GEN_ONLY", raising=False) + monkeypatch.delenv("TRTLLM_DISABLE_KV_CACHE_TRANSFER_OVERLAP", raising=False) + return enabled, emit + + +def _harness(group: FakeDistGroup, rank: int = 0, **kwargs) -> CoordinatorHarness: + dist = group.rank(rank) + dist.pp_rank = dist.cp_rank = 0 + h = CoordinatorHarness(dist=dist, **kwargs) + h.kv_cache_manager.mapping = SimpleNamespace(rank=rank) + # Match the real KV manager's block-ID collection, rather than the harness's + # scalar placeholder, so the unpin diagnostic must actually reach emit_event. + h.kv_cache_manager.store_blocks_for_reuse.side_effect = lambda req, _: [req.py_request_id] + return h + + +def _request(h: CoordinatorHarness) -> TransferRequest: + req = TransferRequest( + 7, prompt_len=128, py_disaggregated_params=SimpleNamespace(disagg_request_id=7007) + ) + h.active.append(req) + return req + + +def _assert_trace(diagnostics: tuple[bool, Mock], expected: list[tuple[str, int]]) -> None: + enabled, emit = diagnostics + if not enabled: + emit.assert_not_called() + return + calls = emit.call_args_list + assert [(call.args[0], call.kwargs["rank"]) for call in calls] == expected + for call in calls: + assert call.kwargs["request_id"] == 7007 + assert call.kwargs["local_request_id"] == 7 + assert call.kwargs["side"] == "ctx" + + +@pytest.mark.parametrize("synchronous", [False, True]) +def test_idle_progress_unpins_and_terminates_once(diagnostics, monkeypatch, synchronous) -> None: + """Both single-rank idle paths reach the unpin edge, even if emitting fails.""" + if synchronous: + monkeypatch.setenv("TRTLLM_DISABLE_KV_CACHE_TRANSFER_OVERLAP", "1") + h = _harness(FakeDistGroup(world_size=1, tp_size=1)) + req = _request(h) + h.send(req) + h.transceiver.finish_send(req) + + h.coordinator.poll_progress_when_idle() + h.coordinator.poll_progress_when_idle() + + assert not h.in_transfer(req) + assert req.state == LlmRequestState.DISAGG_CONTEXT_COMPLETE + assert h.active == [] + assert h.effects.terminated == [req] + assert h.effects.failed == [] + h.kv_cache_manager.store_blocks_for_reuse.assert_called_once_with(req, True) + h.kv_cache_manager.unpin_blocks_by_id.assert_called_once_with([7]) + assert h.dist.calls == [] + _assert_trace(diagnostics, [("ctx_send_ready", 0), ("ctx_source_unpinned", 0)]) + enabled, emit = diagnostics + if enabled: + unpinned = emit.call_args.kwargs + assert unpinned["source_kv_reuse_pinned"] is False + assert unpinned["source_kv_reuse_block_count"] == 1 + assert unpinned["state"] == "DISAGG_CONTEXT_COMPLETE" + + +def test_idle_error_cleanup_waits_for_last_owner(diagnostics) -> None: + """No unpin event or error cleanup is allowed while the connector owns KV.""" + h = _harness(FakeDistGroup(world_size=1, tp_size=1)) + req = _request(h) + h.send(req) + h.transfers.start_transfer(req) # A second claim held by the KV connector. + h.transceiver.finish_send(req, outcome="error") + + h.coordinator.poll_progress_when_idle() + + assert h.in_transfer(req) + assert req.state == LlmRequestState.DISAGG_TRANS_ERROR + assert h.effects.failed == [] + h.kv_cache_manager.unpin_blocks_by_id.assert_not_called() + _assert_trace(diagnostics, [("ctx_send_ready", 0)]) + + h.coordinator.release_transfer(req) + h.coordinator.check_transfer_errors("context requests") + + assert not h.in_transfer(req) + assert req.state == LlmRequestState.DISAGG_TRANS_ERROR + assert h.effects.failed == [("Error in kv cache transfer for context requests", [req], False)] + assert h.effects.terminated == [] + h.kv_cache_manager.unpin_blocks_by_id.assert_called_once_with([7]) + _assert_trace(diagnostics, [("ctx_send_ready", 0), ("ctx_source_unpinned", 0)]) + enabled, emit = diagnostics + if enabled: + assert emit.call_args.kwargs["state"] == "DISAGG_TRANS_ERROR" + + +def test_rank_skew_preserves_error_vote_and_unpin_edges(diagnostics) -> None: + """Diagnostic failures must not bypass the cross-rank last-owner barrier.""" + group = FakeDistGroup(world_size=2, tp_size=2) + ranks = [_harness(group, rank, enable_attention_dp=True) for rank in range(2)] + requests = [_request(h) for h in ranks] + for h, req in zip(ranks, requests): + h.send(req) + requests[0].state = LlmRequestState.DISAGG_TRANS_ERROR + ranks[1].transceiver.finish_send(requests[1], outcome="error") + + group.run(lambda rank: ranks[rank].coordinator.poll_progress_when_idle()) + group.run(lambda rank: ranks[rank].coordinator.handle_errors_synced()) + + assert ranks[0].in_transfer(requests[0]) + assert not ranks[1].in_transfer(requests[1]) + assert all(h.effects.failed == [] for h in ranks) + ranks[0].kv_cache_manager.unpin_blocks_by_id.assert_not_called() + assert ranks[0].dist.calls == [("tp_allgather", {"error_ids": [7], "blocked_ids": [7]})] + assert ranks[1].dist.calls == [("tp_allgather", {"error_ids": [7], "blocked_ids": []})] + _assert_trace( + diagnostics, [("ctx_send_ready", 0), ("ctx_send_ready", 1), ("ctx_source_unpinned", 1)] + ) + + ranks[0].transceiver.finish_send(requests[0], outcome="error") + group.run(lambda rank: ranks[rank].coordinator.poll_progress_when_idle()) + group.run(lambda rank: ranks[rank].coordinator.handle_errors_synced()) + + for h, req in zip(ranks, requests): + assert not h.in_transfer(req) + assert req.state == LlmRequestState.DISAGG_TRANS_ERROR + assert h.effects.failed == [("Disagg KV cache transfer error", [req], False)] + assert h.effects.terminated == [] + h.kv_cache_manager.unpin_blocks_by_id.assert_called_once_with([7]) + assert len(h.dist.calls) == 2 + assert h.dist.calls[1] == ("tp_allgather", {"error_ids": [7], "blocked_ids": []}) + _assert_trace( + diagnostics, + [ + ("ctx_send_ready", 0), + ("ctx_send_ready", 1), + ("ctx_source_unpinned", 1), + ("ctx_source_unpinned", 0), + ], + ) + + +def test_timeout_observation_and_idle_cancellation_preserve_trace(diagnostics, clock) -> None: + """The recorded timer start survives timeout observation and idle cleanup.""" + h = _harness(FakeDistGroup(world_size=1, tp_size=1), kv_transfer_timeout_ms=1000) + req = _request(h) + h.send(req) + start = req.py_kv_transfer_start_time + assert start == clock["t"] + clock["t"] += 2.0 + + h.coordinator.check_transfer_timeouts() + h.coordinator.check_transfer_timeouts() # Do not report the same timeout twice. + h.coordinator.poll_progress_when_idle() + h.coordinator.poll_progress_when_idle() + + assert req.py_kv_transfer_timed_out + assert req.py_kv_transfer_start_time is None + assert not h.in_transfer(req) + assert req.state == LlmRequestState.DISAGG_CONTEXT_COMPLETE + assert h.transceiver.call_log.count("cancel_request:7") == 1 + h.kv_cache_manager.unpin_blocks_by_id.assert_called_once_with([7]) + assert h.effects.terminated == [req] + assert h.dist.calls == [] + _assert_trace( + diagnostics, + [ + ("ctx_send_ready", 0), + ("transfer_timeout_started", 0), + ("transfer_timeout_observed", 0), + ("ctx_source_unpinned", 0), + ], + ) + enabled, emit = diagnostics + if enabled: + started = emit.call_args_list[1].kwargs + observed = emit.call_args_list[2].kwargs + assert started["timer_start_monotonic_ns"] == int(start * 1_000_000_000) + assert observed["timer_start_monotonic_ns"] == started["timer_start_monotonic_ns"] + assert started["timeout_owner"] == observed["timeout_owner"] == "pyexecutor" + assert started["timeout_ms"] == observed["timeout_ms"] == 1000 + assert observed["elapsed_ms"] == 2000.0 diff --git a/tests/unittest/disaggregated/test_disagg_transfer_diagnostics.py b/tests/unittest/disaggregated/test_disagg_transfer_diagnostics.py new file mode 100644 index 000000000000..9ec80eb2ca96 --- /dev/null +++ b/tests/unittest/disaggregated/test_disagg_transfer_diagnostics.py @@ -0,0 +1,639 @@ +# 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. + +import json +import os +import threading +import uuid +from collections.abc import Iterator +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +from tensorrt_llm._torch.disaggregation import diagnostics + +pytestmark = pytest.mark.cpu_only + + +@pytest.fixture(autouse=True) +def _isolate_diagnostic_sink(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: + diagnostics._reset_diagnostic_sink_for_tests() + monkeypatch.delenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS_RUN_ID", raising=False) + yield + diagnostics._reset_diagnostic_sink_for_tests() + + +def test_suppress_diagnostic_errors_does_not_interrupt_request_progress() -> None: + progress = [] + + with diagnostics.suppress_diagnostic_errors(): + progress.append("diagnostic_started") + raise RuntimeError("diagnostic preparation failed") + + progress.append("request_progressed") + assert progress == ["diagnostic_started", "request_progressed"] + + +def test_disabled_emit_event_does_no_diagnostic_work(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", False) + + unexpected_call = MagicMock(side_effect=AssertionError("disabled diagnostics did work")) + monkeypatch.setattr(diagnostics, "_host_identity", unexpected_call) + monkeypatch.setattr(diagnostics, "_get_sink", unexpected_call) + monkeypatch.setattr( + diagnostics, + "os", + SimpleNamespace(getpid=unexpected_call, write=unexpected_call, getenv=unexpected_call), + ) + monkeypatch.setattr( + diagnostics, + "time", + SimpleNamespace(monotonic_ns=unexpected_call, time_ns=unexpected_call), + ) + monkeypatch.setattr(diagnostics, "json", SimpleNamespace(dumps=unexpected_call)) + monkeypatch.setattr( + diagnostics, "uuid", SimpleNamespace(uuid4=unexpected_call, UUID=unexpected_call) + ) + + diagnostics.emit_event("ctx_send_ready", side="ctx", request_id=17) + + unexpected_call.assert_not_called() + + +@pytest.mark.skipif(not hasattr(os, "fork"), reason="requires POSIX fork support") +def test_forked_child_replaces_inherited_diagnostic_sink(monkeypatch: pytest.MonkeyPatch) -> None: + run_uuid = str(uuid.uuid4()) + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS_RUN_ID", run_uuid) + diagnostics._reset_diagnostic_sink_for_tests() + parent_pid = os.getpid() + parent_sink = diagnostics._get_sink(parent_pid) + read_fd, write_fd = os.pipe() + + child_pid = os.fork() + if child_pid == 0: + os.close(read_fd) + try: + reset_in_child = diagnostics._sink is None + child_sink = diagnostics._get_sink(os.getpid()) + result = ( + reset_in_child + and child_sink is not parent_sink + and child_sink.pid == os.getpid() + and child_sink._thread.is_alive() + and child_sink._identity["process_uuid"] != parent_sink._identity["process_uuid"] + and child_sink._identity["run_uuid"] + == parent_sink._identity["run_uuid"] + == run_uuid + ) + diagnostics._reset_diagnostic_sink_for_tests() + os.write(write_fd, b"ok" if result else b"failed") + except BaseException as error: + os.write(write_fd, f"error: {error!r}".encode()) + finally: + os.close(write_fd) + os._exit(0) + + os.close(write_fd) + try: + child_result = os.read(read_fd, 4096) + _, status = os.waitpid(child_pid, 0) + finally: + os.close(read_fd) + diagnostics._reset_diagnostic_sink_for_tests() + + assert os.waitstatus_to_exitcode(status) == 0 + assert child_result == b"ok" + + +@pytest.mark.parametrize( + ("configured_run_id", "expected_run_id", "expected_status"), + [ + (None, None, "unset"), + ("", None, "invalid"), + ("not-a-run-uuid", None, "invalid"), + ( + "9C74CDA0AA094C668F24B8B38F40A958", + "9c74cda0-aa09-4c66-8f24-b8b38f40a958", + "shared", + ), + ], +) +def test_sink_records_validated_run_identity( + monkeypatch: pytest.MonkeyPatch, + configured_run_id: str | None, + expected_run_id: str | None, + expected_status: str, +) -> None: + if configured_run_id is not None: + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS_RUN_ID", configured_run_id) + records = [] + monkeypatch.setattr( + diagnostics._AsyncDiagnosticSink, + "_write", + staticmethod(lambda record: records.append(record.copy())), + ) + sink = diagnostics._get_sink(os.getpid()) + sink.submit({"event": "diagnostic_capabilities", "request_id": None}) + sink.flush() + + assert len(records) == 1 + record = records[0] + assert record["run_uuid"] == expected_run_id + assert record["run_uuid_status"] == expected_status + assert str(uuid.UUID(record["process_uuid"])) == record["process_uuid"] + + +def test_event_identity_is_cached_and_cannot_be_overridden(monkeypatch: pytest.MonkeyPatch) -> None: + run_uuid = str(uuid.uuid4()) + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS_RUN_ID", run_uuid) + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "_host_identity", lambda: "node-a") + records = [] + monkeypatch.setattr( + diagnostics._AsyncDiagnosticSink, + "_write", + staticmethod(lambda record: records.append(record.copy())), + ) + sink = diagnostics._get_sink(os.getpid()) + unexpected_call = MagicMock(side_effect=AssertionError("event regenerated identity")) + monkeypatch.setattr( + diagnostics, "uuid", SimpleNamespace(uuid4=unexpected_call, UUID=unexpected_call) + ) + monkeypatch.setattr( + diagnostics, "os", SimpleNamespace(getpid=os.getpid, getenv=unexpected_call) + ) + + for event in ("diagnostic_capabilities", "ctx_send_ready"): + diagnostics.emit_event( + event, + side="ctx", + request_id=17, + run_uuid="forged-run", + process_uuid="forged-process", + run_uuid_status="invalid", + host="forged-host", + pid=-1, + ) + sink.flush() + + unexpected_call.assert_not_called() + assert len(records) == 2 + for record in records: + assert record["run_uuid"] == run_uuid + assert record["run_uuid_status"] == "shared" + assert record["process_uuid"] == sink._identity["process_uuid"] + assert record["host"] == "node-a" + assert record["pid"] == os.getpid() + + +@pytest.mark.parametrize("restart", ["same_pid", "changed_pid"]) +def test_recreated_sink_has_fresh_process_identity( + monkeypatch: pytest.MonkeyPatch, restart: str +) -> None: + run_uuid = str(uuid.uuid4()) + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS_RUN_ID", run_uuid) + pid = os.getpid() + original_sink = diagnostics._get_sink(pid) + assert diagnostics._get_sink(pid) is original_sink + if restart == "same_pid": + diagnostics._reset_diagnostic_sink_for_tests() + replacement_sink = diagnostics._get_sink(pid) + else: + # Model PID replacement without inheriting a live parent thread. + original_sink.close() + replacement_sink = diagnostics._get_sink(pid + 1) + + assert replacement_sink is not original_sink + assert replacement_sink._identity["process_uuid"] != original_sink._identity["process_uuid"] + assert replacement_sink._identity["run_uuid"] == original_sink._identity["run_uuid"] == run_uuid + + +@pytest.mark.parametrize("failing_dependency", ["uuid", "environment"]) +def test_identity_initialization_failure_does_not_affect_request_progress( + monkeypatch: pytest.MonkeyPatch, failing_dependency: str +) -> None: + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + failure = MagicMock(side_effect=OSError("diagnostic identity unavailable")) + if failing_dependency == "uuid": + monkeypatch.setattr(diagnostics, "uuid", SimpleNamespace(uuid4=failure)) + else: + monkeypatch.setattr(diagnostics, "os", SimpleNamespace(getpid=os.getpid, getenv=failure)) + + diagnostics.emit_event("gen_decode_ready", side="gen", request_id=42) + + failure.assert_called_once() + assert diagnostics._sink is None + + +def test_enabled_emit_event_records_request_and_rank_context( + monkeypatch: pytest.MonkeyPatch, +) -> None: + diagnostics._reset_diagnostic_sink_for_tests() + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "_host_identity", lambda: "node-a") + request_thread_id = threading.get_ident() + writes = [] + + def write(fd: int, data: bytes) -> int: + writes.append((threading.get_ident(), fd, data)) + return len(data) + + monkeypatch.setattr( + diagnostics, + "os", + SimpleNamespace( + getpid=lambda: 321, + write=write, + getenv=lambda _key: None, + ), + ) + monkeypatch.setattr( + diagnostics, + "time", + SimpleNamespace(monotonic_ns=lambda: 111, time_ns=lambda: 222), + ) + rank_info = SimpleNamespace( + instance_name="ctx_0", + instance_rank=4, + tp_rank=1, + pp_rank=2, + cp_rank=3, + dp_rank=0, + ) + + diagnostics.emit_event( + "ctx_backend_submitted", + side="ctx", + request_id=99, + local_request_id=7, + rank_info=rank_info, + slice_id=5, + peer_rank=8, + transfer_bytes=4096, + source_kv_request_owned=True, + source_kv_reuse_pinned=False, + timestamp=(1_234, 5_678), + ) + diagnostics._flush_diagnostic_sink_for_tests() + + assert len(writes) == 1 + writer_thread_id, fd, encoded_message = writes[0] + assert writer_thread_id != request_thread_id + assert fd == 1 + message = encoded_message.decode("utf-8").removesuffix("\n") + assert message.startswith(diagnostics.DIAGNOSTICS_LOG_PREFIX) + payload = message.removeprefix(diagnostics.DIAGNOSTICS_LOG_PREFIX) + process_uuid = json.loads(payload)["process_uuid"] + assert str(uuid.UUID(process_uuid)) == process_uuid + assert json.loads(payload) == { + "schema_version": diagnostics.DIAGNOSTICS_SCHEMA_VERSION, + "event": "ctx_backend_submitted", + "side": "ctx", + "request_id": 99, + "local_request_id": 7, + "host": "node-a", + "pid": 321, + "process_uuid": process_uuid, + "run_uuid": None, + "run_uuid_status": "unset", + "monotonic_ns": 1_234, + "wall_ns": 5_678, + "instance": "ctx_0", + "rank": 4, + "tp_rank": 1, + "pp_rank": 2, + "cp_rank": 3, + "dp_rank": 0, + "slice_id": 5, + "peer_rank": 8, + "transfer_bytes": 4096, + "source_kv_request_owned": True, + "source_kv_reuse_pinned": False, + } + assert payload == json.dumps(json.loads(payload), separators=(",", ":"), sort_keys=True) + diagnostics._reset_diagnostic_sink_for_tests() + + +def test_enabled_emit_event_accepts_executor_rank_context(monkeypatch: pytest.MonkeyPatch) -> None: + diagnostics._reset_diagnostic_sink_for_tests() + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "_host_identity", lambda: "node-b") + writes = [] + + def write(fd: int, data: bytes) -> int: + writes.append((fd, data)) + return len(data) + + monkeypatch.setattr( + diagnostics, + "os", + SimpleNamespace( + getpid=lambda: 654, + write=write, + getenv=lambda _key: None, + ), + ) + monkeypatch.setattr( + diagnostics, + "time", + SimpleNamespace(monotonic_ns=lambda: 9_000, time_ns=lambda: 10_000), + ) + diagnostics.emit_event( + "gen_kv_admission_result", + side="gen", + request_id=101, + rank=6, + instance="gen_0", + outcome="deferred", + ) + diagnostics._flush_diagnostic_sink_for_tests() + + assert len(writes) == 1 + _, encoded_message = writes[0] + message = encoded_message.decode("utf-8").removesuffix("\n") + payload = json.loads(message.removeprefix(diagnostics.DIAGNOSTICS_LOG_PREFIX)) + assert payload["rank"] == 6 + assert payload["instance"] == "gen_0" + assert payload["outcome"] == "deferred" + assert "slice_id" not in payload + assert "peer_rank" not in payload + diagnostics._reset_diagnostic_sink_for_tests() + + +def test_sink_write_completes_partial_stdout_writes( + monkeypatch: pytest.MonkeyPatch, +) -> None: + record = {"event": "partial_write", "request_id": 42} + expected = ( + f"{diagnostics.DIAGNOSTICS_LOG_PREFIX}" + f"{json.dumps(record, separators=(',', ':'), sort_keys=True)}\n" + ).encode("utf-8") + accepted_chunks = [] + calls = [] + + def write(fd: int, data: bytes) -> int: + calls.append((fd, data)) + written = min(7, len(data)) + accepted_chunks.append(data[:written]) + return written + + monkeypatch.setattr(diagnostics, "os", SimpleNamespace(write=write)) + + diagnostics._AsyncDiagnosticSink._write(record) + + assert len(calls) > 2 + assert all(fd == 1 for fd, _ in calls) + assert [data for _, data in calls] == [ + expected[offset:] for offset in range(0, len(expected), 7) + ] + assert b"".join(accepted_chunks) == expected + + +def test_sink_write_rejects_zero_progress(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr( + diagnostics, + "os", + SimpleNamespace(write=lambda _fd, _data: 0), + ) + + with pytest.raises(OSError, match="made no progress"): + diagnostics._AsyncDiagnosticSink._write({"event": "no_progress"}) + + +def test_async_sink_failure_does_not_stop_later_diagnostic_writes( + monkeypatch: pytest.MonkeyPatch, +) -> None: + diagnostics._reset_diagnostic_sink_for_tests() + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "_host_identity", lambda: "node-a") + writes = [] + + def write(fd: int, data: bytes) -> int: + writes.append((fd, data)) + if len(writes) == 1: + raise OSError("diagnostic sink unavailable") + return len(data) + + monkeypatch.setattr( + diagnostics, + "os", + SimpleNamespace( + getpid=lambda: 123, + write=write, + getenv=lambda _key: None, + ), + ) + + diagnostics.emit_event("gen_decode_ready", side="gen", request_id=42) + diagnostics._flush_diagnostic_sink_for_tests() + sink = diagnostics._sink + assert sink is not None + assert sink._thread.is_alive() + + diagnostics.emit_event("gen_decode_ready", side="gen", request_id=43) + diagnostics._flush_diagnostic_sink_for_tests() + + assert len(writes) == 3 + dropped = json.loads( + writes[1][1].decode("utf-8").removeprefix(diagnostics.DIAGNOSTICS_LOG_PREFIX) + ) + assert dropped["event"] == "diagnostics_events_dropped" + assert dropped["dropped_events"] == 1 + assert b'"request_id":43' in writes[2][1] + diagnostics._reset_diagnostic_sink_for_tests() + + +def test_drop_record_failure_preserves_loss_count_and_current_record( + monkeypatch: pytest.MonkeyPatch, +) -> None: + records = [] + failed_once = False + + def write(record) -> None: + nonlocal failed_once + records.append(record.copy()) + if record.get("event") == "diagnostics_events_dropped" and not failed_once: + failed_once = True + raise OSError("transient diagnostic sink failure") + + monkeypatch.setattr(diagnostics._AsyncDiagnosticSink, "_write", staticmethod(write)) + sink = diagnostics._AsyncDiagnosticSink(os.getpid()) + try: + sink._record_dropped(3) + sink.submit({"event": "first"}) + sink.submit({"event": "second"}) + sink.flush() + finally: + sink.close() + + assert [record["event"] for record in records] == [ + "diagnostics_events_dropped", + "first", + "diagnostics_events_dropped", + "second", + ] + assert records[0]["dropped_events"] == records[2]["dropped_events"] == 3 + + +def test_bounded_close_accounts_for_abandoned_queued_records( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(diagnostics, "_DIAGNOSTIC_SHUTDOWN_TIMEOUT_S", 0.01) + writer_started = threading.Event() + release_writer = threading.Event() + records = [] + + def write(record) -> None: + if record.get("event") == "first": + writer_started.set() + assert release_writer.wait(timeout=1.0) + records.append(record.copy()) + + monkeypatch.setattr(diagnostics._AsyncDiagnosticSink, "_write", staticmethod(write)) + sink = diagnostics._AsyncDiagnosticSink(os.getpid()) + try: + sink.submit({"event": "first"}) + assert writer_started.wait(timeout=1.0) + sink.submit({"event": "second"}) + sink.submit({"event": "third"}) + + sink.close() + assert sink._thread.is_alive() + release_writer.set() + sink._thread.join(timeout=1.0) + finally: + release_writer.set() + sink.close() + + assert not sink._thread.is_alive() + assert [record["event"] for record in records] == [ + "first", + "diagnostics_events_dropped", + ] + assert records[-1]["dropped_events"] == 2 + + +def test_sink_creation_failure_does_not_affect_request_progress( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + get_sink = MagicMock(side_effect=RuntimeError("diagnostic sink unavailable")) + monkeypatch.setattr(diagnostics, "_get_sink", get_sink) + + diagnostics.emit_event("gen_decode_ready", side="gen", request_id=42) + + get_sink.assert_called_once() + + +def test_full_diagnostic_queue_reports_dropped_events( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(diagnostics, "_DIAGNOSTIC_QUEUE_CAPACITY", 1) + writer_started = threading.Event() + release_writer = threading.Event() + records = [] + + def write(record) -> None: + records.append(record) + if record.get("event") == "first": + writer_started.set() + assert release_writer.wait(timeout=1.0) + + monkeypatch.setattr(diagnostics._AsyncDiagnosticSink, "_write", staticmethod(write)) + sink = diagnostics._AsyncDiagnosticSink(os.getpid()) + try: + sink.submit({"event": "first"}) + assert writer_started.wait(timeout=1.0) + sink.submit({"event": "second"}) + sink.submit({"event": "dropped"}) + release_writer.set() + sink.flush() + finally: + release_writer.set() + sink.close() + + drop_record = next( + record for record in records if record["event"] == "diagnostics_events_dropped" + ) + assert drop_record["dropped_events"] == 1 + for record in records: + assert record["process_uuid"] == sink._identity["process_uuid"] + assert record["run_uuid"] is None + assert record["run_uuid_status"] == "unset" + + +def test_scheduler_kv_admission_guard_avoids_telemetry_state_inspection( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from tensorrt_llm._torch.pyexecutor.scheduler.scheduler_v2 import ( + KVCacheV2Scheduler, + ScheduleAction, + ) + + class _OpaqueRequest: + @property + def py_request_id(self) -> int: + raise AssertionError("disabled diagnostics inspected the request") + + @property + def prompt_len(self) -> int: + raise AssertionError("disabled diagnostics inspected the request") + + class _KVCacheManager: + def prepare_disagg_gen_init(self, _request: _OpaqueRequest) -> bool: + return True + + @property + def kv_cache_map(self) -> dict[int, object]: + raise AssertionError("disabled diagnostics inspected the KV cache map") + + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", False) + scheduler = object.__new__(KVCacheV2Scheduler) + scheduler.kv_cache_manager = _KVCacheManager() + scheduler.tokens_per_block = 32 + + action, tokens = scheduler._try_schedule_disagg_gen_init(_OpaqueRequest(), None) + + assert action is ScheduleAction.SCHEDULED + assert tokens == 0 + + +def test_scheduler_kv_admission_continues_when_diagnostic_inspection_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from tensorrt_llm._torch.pyexecutor.scheduler.scheduler_v2 import ( + KVCacheV2Scheduler, + ScheduleAction, + ) + + class _KVCacheManager: + def prepare_disagg_gen_init(self, _request) -> bool: + return True + + @property + def kv_cache_map(self): + raise RuntimeError("diagnostic KV inspection failed") + + request = SimpleNamespace(py_request_id=17) + scheduler = object.__new__(KVCacheV2Scheduler) + scheduler.kv_cache_manager = _KVCacheManager() + scheduler.tokens_per_block = 32 + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + + action, tokens = scheduler._try_schedule_disagg_gen_init(request, None) + + assert action is ScheduleAction.SCHEDULED + assert tokens == 0 diff --git a/tests/unittest/disaggregated/test_disagg_transfer_diagnostics_wiring.py b/tests/unittest/disaggregated/test_disagg_transfer_diagnostics_wiring.py new file mode 100644 index 000000000000..6ce71f85aa67 --- /dev/null +++ b/tests/unittest/disaggregated/test_disagg_transfer_diagnostics_wiring.py @@ -0,0 +1,1571 @@ +# 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. +"""CPU-only wiring tests for disaggregated-transfer diagnostic events.""" + +from __future__ import annotations + +import queue +import threading +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest + +import tensorrt_llm._torch.disaggregation.native.transfer as transfer_module +import tensorrt_llm._torch.disaggregation.transceiver as python_transceiver_module +from tensorrt_llm import DisaggregatedParams +from tensorrt_llm._torch.disaggregation import diagnostics, kv_cache_transceiver +from tensorrt_llm._torch.disaggregation.base.transfer import KVSlice, SessionStatus, WaitResult +from tensorrt_llm._torch.disaggregation.native.transfer import ( + AgentResult, + Receiver, + RxSession, + Sender, + TaskStatus, +) +from tensorrt_llm._torch.disaggregation.orchestration import coordinator as coordinator_module +from tensorrt_llm._torch.disaggregation.orchestration.admission import ( + DisaggTransferAdmissionController, +) +from tensorrt_llm._torch.disaggregation.orchestration.coordinator import DisaggTransferCoordinator +from tensorrt_llm._torch.disaggregation.orchestration.transfer_manager import AsyncTransferManager +from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 +from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequestState +from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor +from tensorrt_llm._torch.pyexecutor.resource_manager import ResourceManagerType +from tensorrt_llm._torch.pyexecutor.scheduler import MultimodalScheduler, SimpleScheduler +from tensorrt_llm._torch.pyexecutor.scheduler.scheduler_v2 import KVCacheV2Scheduler, ScheduleAction +from tensorrt_llm.llmapi.llm_args import CacheTransceiverConfig + +pytestmark = pytest.mark.cpu_only + + +def _disagg_request(local_id: int, canonical_id: int, prompt_len: int = 65) -> SimpleNamespace: + return SimpleNamespace( + py_request_id=local_id, + request_id=local_id, + py_disaggregated_params=SimpleNamespace(disagg_request_id=canonical_id), + prompt_len=prompt_len, + ) + + +def _diagnostic_transceiver() -> KvCacheTransceiverV2: + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._mapping = SimpleNamespace(rank=4, tp_rank=1, pp_rank=0, cp_rank=0) + transceiver._instance_name = "diagnostic-test" + transceiver._dp_rank = 2 + return transceiver + + +@pytest.mark.parametrize("diagnostic_mode", ["disabled", "enabled", "failing"]) +@pytest.mark.parametrize( + ("configured_runtime", "pipelined", "expected_runtime"), + [ + ("PYTHON", False, "PYTHON"), + ("CPP", False, "CPP"), + (None, False, "CPP"), + ("auto", False, "CPP"), + ("auto", True, "PYTHON"), + ], +) +def test_factory_reports_effective_diagnostic_capabilities( + monkeypatch: pytest.MonkeyPatch, + diagnostic_mode: str, + configured_runtime: str | None, + pipelined: bool, + expected_runtime: str, +) -> None: + config = CacheTransceiverConfig( + backend="NIXL", + transceiver_runtime=configured_runtime, + enable_pipelined_transfer=pipelined, + ) + transceiver = object() + constructors = { + "PYTHON": Mock(return_value=transceiver), + "CPP": Mock(return_value=transceiver), + } + monkeypatch.setattr(python_transceiver_module, "KvCacheTransceiverV2", constructors["PYTHON"]) + monkeypatch.setattr(kv_cache_transceiver, "BindKvCacheTransceiver", constructors["CPP"]) + monkeypatch.setattr(kv_cache_transceiver, "is_disagg_inflight_cancel_enabled", lambda: False) + monkeypatch.setattr( + diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", diagnostic_mode != "disabled" + ) + + def record_capabilities(*_args, **_kwargs) -> None: + constructors[expected_runtime].assert_called_once() + if diagnostic_mode == "failing": + raise RuntimeError("diagnostic emission failed") + + emit_event = Mock(side_effect=record_capabilities) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + result = kv_cache_transceiver.create_kv_cache_transceiver( + SimpleNamespace(rank=4), Mock(), Mock(), Mock(), config + ) + + assert result is transceiver + constructors[expected_runtime].assert_called_once() + constructors["CPP" if expected_runtime == "PYTHON" else "PYTHON"].assert_not_called() + if diagnostic_mode == "disabled": + emit_event.assert_not_called() + else: + emit_event.assert_called_once_with( + "diagnostic_capabilities", + side="runtime", + request_id=None, + rank=4, + capability_schema_version=1, + transceiver_runtime=expected_runtime, + python_transfer_events=expected_runtime == "PYTHON", + ) + + +@pytest.mark.parametrize("runtime", ["PYTHON", "CPP"]) +def test_factory_does_not_report_capabilities_when_construction_fails( + monkeypatch: pytest.MonkeyPatch, runtime: str +) -> None: + constructor = Mock(side_effect=RuntimeError("transceiver construction failed")) + monkeypatch.setattr(python_transceiver_module, "KvCacheTransceiverV2", constructor) + monkeypatch.setattr(kv_cache_transceiver, "BindKvCacheTransceiver", constructor) + monkeypatch.setattr(kv_cache_transceiver, "is_disagg_inflight_cancel_enabled", lambda: False) + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + emit_event = Mock() + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + with pytest.raises(RuntimeError, match="transceiver construction failed"): + kv_cache_transceiver.create_kv_cache_transceiver( + SimpleNamespace(rank=4), + Mock(), + Mock(), + Mock(), + CacheTransceiverConfig(backend="NIXL", transceiver_runtime=runtime), + ) + + constructor.assert_called_once() + emit_event.assert_not_called() + + +@pytest.mark.parametrize("diagnostic_mode", ["disabled", "enabled", "failing"]) +@pytest.mark.parametrize("scheduler_v2", [False, True]) +@pytest.mark.parametrize("wrapped", [False, True]) +def test_executor_reports_scheduler_capabilities_independently_of_transceiver( + monkeypatch: pytest.MonkeyPatch, + diagnostic_mode: str, + scheduler_v2: bool, + wrapped: bool, +) -> None: + scheduler = object.__new__(KVCacheV2Scheduler if scheduler_v2 else SimpleScheduler) + if wrapped: + wrapper = object.__new__(MultimodalScheduler) + wrapper.scheduler = scheduler + scheduler = wrapper + executor = object.__new__(PyExecutor) + executor.scheduler = scheduler + # Deliberately do not correlate scheduler V2 with the Python runtime. + executor.kv_cache_transceiver = SimpleNamespace(consumes_transfer_buffer=scheduler_v2) + executor.global_rank = 4 + monkeypatch.setattr( + diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", diagnostic_mode != "disabled" + ) + emit_event = Mock( + side_effect=RuntimeError("diagnostic emission failed") + if diagnostic_mode == "failing" + else None + ) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + executor._emit_disagg_diagnostic_capabilities() + + if diagnostic_mode == "disabled": + emit_event.assert_not_called() + else: + emit_event.assert_called_once_with( + "diagnostic_capabilities", + side="runtime", + request_id=None, + rank=4, + capability_schema_version=1, + executor_events=True, + scheduler_kv_admission_events=scheduler_v2, + ) + + +def test_capability_reporting_skips_disabled_transceiver( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + emit_event = Mock() + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + executor = object.__new__(PyExecutor) + executor.kv_cache_transceiver = None + + assert ( + kv_cache_transceiver.create_kv_cache_transceiver(Mock(), Mock(), Mock(), Mock(), None) + is None + ) + executor._emit_disagg_diagnostic_capabilities() + + emit_event.assert_not_called() + + +def test_sender_reports_worker_dequeue_before_transfer_preparation( + monkeypatch: pytest.MonkeyPatch, +) -> None: + task_queue = queue.Queue() + write_meta = SimpleNamespace( + meta_type=transfer_module.WriteMetaType.KV, + src_ptrs=SimpleNamespace(size=1), + sizes=SimpleNamespace(sum=lambda: 4096), + unique_rid=1010, + slice_id=2, + peer_rank=3, + receiver_slice_id=4, + is_last_slice=True, + ) + task_queue.put(write_meta) + task_queue.put(None) + sender = object.__new__(Sender) + sender._device_id = 0 + sender._send_task_queues = [task_queue] + sender._registrar = SimpleNamespace(self_rank_info=SimpleNamespace()) + sender._thread_local = threading.local() + operations = [] + + def emit_event(event: str, **kwargs) -> None: + operations.append((event, kwargs)) + + sender._deliver_kv_to_agent = Mock( + side_effect=lambda meta: operations.append(("deliver", meta)) + ) + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + monkeypatch.setattr(transfer_module.torch.cuda, "set_device", Mock()) + monkeypatch.setattr(transfer_module, "CUASSERT", Mock()) + monkeypatch.setattr( + transfer_module, + "cudart", + SimpleNamespace(cudaSetDevice=Mock(return_value=0)), + ) + + sender._process_task_queue(0) + + assert [operation[0] for operation in operations] == [ + "ctx_worker_dequeued", + "deliver", + ] + event = operations[0][1] + assert event["request_id"] == 1010 + assert event["slice_id"] == 2 + assert event["peer_rank"] == 3 + assert event["worker_queue_index"] == 0 + assert event["transfer_bytes"] == 4096 + + +def test_gen_ingress_uses_full_cp_prompt_length(monkeypatch: pytest.MonkeyPatch) -> None: + request = _disagg_request(10, 1010, prompt_len=1) + request.total_input_len_cp = 257 + request.is_disagg_generation_init_state = True + request.state = LlmRequestState.DISAGG_GENERATION_INIT + executor = object.__new__(PyExecutor) + executor.waiting_queue = [] + executor.active_requests = [] + executor._fetch_new_requests = Mock(return_value=[request]) + executor._validate_request = Mock() + executor._mm_encoder_item_scheduling_enabled = False + executor.global_rank = 4 + executor.dist = SimpleNamespace(tp_rank=1, pp_rank=0, cp_rank=3) + + emit_event = Mock() + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + assert executor._fetch_and_activate_new_requests() == [request] + assert executor.active_requests == [request] + event = emit_event.call_args + assert event.args == ("gen_ingress",) + assert event.kwargs["prompt_tokens"] == 257 + assert event.kwargs["cp_rank"] == 3 + + +def test_scheduler_emits_admitted_and_deferred_kv_admission_results( + monkeypatch: pytest.MonkeyPatch, +) -> None: + admitted = _disagg_request(11, 1011) + admitted.total_input_len_cp = 257 + deferred = _disagg_request(12, 1012) + results = iter((True, False)) + manager = SimpleNamespace( + prepare_disagg_gen_init=lambda _request: next(results), + kv_cache_map={ + admitted.py_request_id: SimpleNamespace(capacity=65, history_length=64), + }, + mapping=SimpleNamespace(rank=3), + ) + scheduler = object.__new__(KVCacheV2Scheduler) + scheduler.kv_cache_manager = manager + scheduler.tokens_per_block = 32 + + emit_event = Mock() + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + admitted_result = scheduler._try_schedule_disagg_gen_init(admitted, None) + deferred_result = scheduler._try_schedule_disagg_gen_init(deferred, None) + + assert admitted_result == (ScheduleAction.SCHEDULED, 0) + assert deferred_result == (ScheduleAction.SKIP, 0) + assert emit_event.call_count == 2 + + admitted_event = emit_event.call_args_list[0] + assert admitted_event.args == ("gen_kv_admission_result",) + assert admitted_event.kwargs == { + "side": "gen", + "request_id": 1011, + "local_request_id": 11, + "rank": 3, + "outcome": "admitted", + "reason": None, + "prompt_tokens": 257, + "tokens_per_block": 32, + "cache_present": True, + "capacity_tokens": 65, + "history_tokens": 64, + "capacity_block_equivalent": 3, + } + + deferred_event = emit_event.call_args_list[1] + assert deferred_event.args == ("gen_kv_admission_result",) + assert deferred_event.kwargs["request_id"] == 1012 + assert deferred_event.kwargs["local_request_id"] == 12 + assert deferred_event.kwargs["outcome"] == "deferred" + assert deferred_event.kwargs["reason"] == "kv_or_index_capacity" + assert deferred_event.kwargs["prompt_tokens"] == 65 + assert deferred_event.kwargs["cache_present"] is False + assert deferred_event.kwargs["capacity_tokens"] is None + assert deferred_event.kwargs["capacity_block_equivalent"] is None + + +def test_scheduler_kv_pool_snapshot_reports_pressure_and_skips_irrelevant_batches( + monkeypatch: pytest.MonkeyPatch, +) -> None: + get_stats = Mock( + return_value=SimpleNamespace( + max_num_blocks=100, + free_num_blocks=40, + used_num_blocks=60, + ), + ) + manager = SimpleNamespace( + get_kv_cache_stats=get_stats, + mapping=SimpleNamespace(rank=3, tp_rank=1, pp_rank=0, cp_rank=0), + index_mapper=SimpleNamespace(num_free_slots=lambda: 7), + ) + scheduler = object.__new__(KVCacheV2Scheduler) + scheduler.kv_cache_manager = manager + emit_event = Mock() + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + decode = SimpleNamespace( + is_disagg_generation_init_state=False, + is_disagg_generation_transmission_in_progress=False, + is_disagg_generation_transmission_complete=False, + state=LlmRequestState.GENERATION_IN_PROGRESS, + ) + scheduler._emit_disagg_kv_pool_snapshot([decode], []) + + get_stats.assert_not_called() + emit_event.assert_not_called() + + pending = SimpleNamespace( + is_disagg_generation_init_state=True, + is_disagg_generation_transmission_in_progress=False, + is_disagg_generation_transmission_complete=False, + state=LlmRequestState.DISAGG_GENERATION_INIT, + ) + transferring = SimpleNamespace( + is_disagg_generation_init_state=False, + is_disagg_generation_transmission_in_progress=True, + is_disagg_generation_transmission_complete=False, + state=LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS, + ) + transferred = SimpleNamespace( + is_disagg_generation_init_state=False, + is_disagg_generation_transmission_in_progress=False, + is_disagg_generation_transmission_complete=True, + state=LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE, + ) + scheduler._emit_disagg_kv_pool_snapshot( + [pending, transferring, transferred, decode], + [pending], + ) + + get_stats.assert_called_once_with() + event = emit_event.call_args + assert event.args == ("gen_kv_pool_snapshot",) + assert event.kwargs == { + "side": "gen", + "request_id": None, + "rank": 3, + "init_requests": 1, + "transfers_in_progress": 1, + "transfers_complete": 1, + "kv_admitted_this_iteration": 1, + "decode_requests": 1, + "kv_pool_max_blocks": 100, + "kv_pool_free_blocks": 40, + "kv_pool_used_blocks": 60, + "index_free_slots": 7, + "tp_rank": 1, + "pp_rank": 0, + "cp_rank": 0, + } + + +def test_gen_timeout_start_and_observation_share_request_identity( + monkeypatch: pytest.MonkeyPatch, +) -> None: + request = _disagg_request(21, 2021) + request.state = LlmRequestState.DISAGG_GENERATION_INIT + request.py_kv_transfer_start_time = None + request.py_kv_transfer_timed_out = False + request.is_disagg_generation_transmission_in_progress = False + + def start_receive(req) -> None: + req.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS + req.is_disagg_generation_transmission_in_progress = True + + transceiver = Mock() + transceiver.kv_transfer_timeout_ms = 100 + transceiver.request_and_receive_async.side_effect = start_receive + + executor = object.__new__(PyExecutor) + executor.kv_cache_transceiver = transceiver + executor.global_rank = 4 + executor.dist = SimpleNamespace(tp_rank=1, pp_rank=0, cp_rank=0) + executor._is_disagg_gen_only_no_context_benchmark = Mock(return_value=False) + executor._uses_async_disagg_gen_transfer = Mock(return_value=True) + executor._disagg_coordinator = SimpleNamespace(reap_gen_receives=Mock()) + + coordinator = DisaggTransferCoordinator( + transceiver=transceiver, + transfer_manager=SimpleNamespace(requests_in_transfer=lambda: {}), + kv_cache_manager=None, + dist=SimpleNamespace(rank=4, tp_rank=1, pp_rank=0, cp_rank=0), + effects=None, + registry=SimpleNamespace(active_requests=lambda: [request]), + enable_attention_dp=False, + force_terminate_ctx_for_partial_reuse=False, + delegates=None, + ) + + emit_event = Mock() + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + monkeypatch.setattr( + "tensorrt_llm._torch.pyexecutor.py_executor.time.monotonic", + lambda: 10.0, + ) + monkeypatch.setattr( + coordinator_module, + "is_disagg_inflight_cancel_enabled", + lambda: False, + ) + + executor._recv_disagg_gen_cache([request]) + # py_executor and coordinator import the same time module, so advance the + # shared clock only after the receive path records its start timestamp. + monkeypatch.setattr(coordinator_module.time, "monotonic", lambda: 10.2) + coordinator.check_transfer_timeouts() + + assert request.py_kv_transfer_start_time == 10.0 + assert request.py_kv_transfer_timed_out + assert [entry.args[0] for entry in emit_event.call_args_list] == [ + "transfer_timeout_started", + "transfer_timeout_observed", + ] + started = emit_event.call_args_list[0].kwargs + observed = emit_event.call_args_list[1].kwargs + assert started["side"] == observed["side"] == "gen" + assert started["request_id"] == observed["request_id"] == 2021 + assert started["local_request_id"] == observed["local_request_id"] == 21 + assert started["timeout_owner"] == observed["timeout_owner"] == "pyexecutor" + assert started["timer_start_monotonic_ns"] == 10_000_000_000 + assert observed["timer_start_monotonic_ns"] == 10_000_000_000 + assert observed["elapsed_ms"] == pytest.approx(200.0) + assert observed["cancellation_requested"] is False + + +def test_gen_decode_ready_uses_full_cp_prompt_length(monkeypatch: pytest.MonkeyPatch) -> None: + request = _disagg_request(22, 2022, prompt_len=1) + request.total_input_len_cp = 257 + request.is_disagg_generation_transmission_complete = True + request.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE + request.context_phase_params = SimpleNamespace(first_gen_tokens=[7], draft_tokens=None) + request.py_beam_width = 1 + request.add_new_token = Mock() + seq_slot_manager = SimpleNamespace(prepare_resources=Mock()) + executor = object.__new__(PyExecutor) + executor.resource_manager = SimpleNamespace( + resource_managers={ResourceManagerType.SEQ_SLOT_MANAGER: seq_slot_manager}, + ) + executor._setup_sampler_step = Mock() + executor.model_engine = SimpleNamespace(enable_spec_decode=False) + executor.kv_cache_transceiver = None + executor._update_sampler_state_for_disagg_gen_request = Mock(return_value=True) + executor._maybe_prepend_logprobs_and_logits = Mock() + executor.global_rank = 4 + executor.dist = SimpleNamespace(tp_rank=1, pp_rank=0, cp_rank=3) + + emit_event = Mock() + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + executor._prepare_disagg_gen_transmission_complete( + SimpleNamespace(generation_requests=[request]), + ) + + event = emit_event.call_args + assert event.args == ("gen_decode_ready",) + assert event.kwargs["prompt_tokens"] == 257 + assert event.kwargs["cp_rank"] == 3 + request.add_new_token.assert_called_once_with(7, 0) + + +def test_ctx_send_ready_and_timeout_start_follow_transfer_handoff( + monkeypatch: pytest.MonkeyPatch, +) -> None: + operations = [] + request = SimpleNamespace( + is_context_only_request=True, + is_finished_due_to_cancellation=False, + is_child=False, + py_request_id=31, + request_id=31, + py_disaggregated_params=SimpleNamespace(disagg_request_id=3031), + is_context_finished=True, + is_finished_due_to_length=False, + prompt_len=128, + state=SimpleNamespace(name="CONTEXT_IN_PROGRESS"), + py_kv_transfer_start_time=None, + ) + transfer_manager = SimpleNamespace( + start_transfer=lambda req: operations.append(("start_transfer", req)), + should_store_blocks=True, + ) + transceiver = SimpleNamespace( + has_retired_send_session=lambda _req: False, + respond_and_send_async=lambda req: operations.append(("respond", req)), + kv_transfer_timeout_ms=100, + pipeline_transfer_enabled=False, + ) + coordinator = DisaggTransferCoordinator( + transceiver=transceiver, + transfer_manager=transfer_manager, + kv_cache_manager=None, + dist=SimpleNamespace(rank=4, tp_rank=1, pp_rank=0, cp_rank=0), + effects=None, + registry=SimpleNamespace(canceled_request_ids=lambda: []), + enable_attention_dp=False, + force_terminate_ctx_for_partial_reuse=False, + delegates=None, + ) + + def emit_event(event: str, **kwargs) -> None: + operations.append((event, kwargs)) + + def monotonic() -> float: + operations.append(("monotonic", None)) + return 12.5 + + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + monkeypatch.setattr(coordinator_module.time, "monotonic", monotonic) + + coordinator.send_completed_context([request]) + + assert [operation[0] for operation in operations] == [ + "start_transfer", + "ctx_send_ready", + "respond", + "monotonic", + "transfer_timeout_started", + ] + send_ready = operations[1][1] + timeout_started = operations[4][1] + assert send_ready["request_id"] == timeout_started["request_id"] == 3031 + assert send_ready["source_kv_request_owned"] is True + assert send_ready["source_kv_reuse_pinned"] is True + assert send_ready["timeout_expected"] is True + assert timeout_started["timeout_ms"] == 100 + assert timeout_started["timer_start_monotonic_ns"] == 12_500_000_000 + assert request.py_kv_transfer_start_time == 12.5 + + +def test_ctx_send_continues_when_diagnostic_preparation_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + operations = [] + request = SimpleNamespace( + is_context_only_request=True, + is_finished_due_to_cancellation=False, + is_child=False, + py_request_id=31, + request_id=31, + py_disaggregated_params=SimpleNamespace(disagg_request_id=3031), + is_context_finished=True, + is_finished_due_to_length=False, + prompt_len=128, + state=SimpleNamespace(name="CONTEXT_IN_PROGRESS"), + py_kv_transfer_start_time=None, + ) + transfer_manager = SimpleNamespace( + start_transfer=lambda req: operations.append(("start_transfer", req)), + should_store_blocks=True, + ) + transceiver = SimpleNamespace( + has_retired_send_session=lambda _req: False, + respond_and_send_async=lambda req: operations.append(("respond", req)), + kv_transfer_timeout_ms=100, + pipeline_transfer_enabled=False, + ) + coordinator = DisaggTransferCoordinator( + transceiver=transceiver, + transfer_manager=transfer_manager, + kv_cache_manager=None, + dist=SimpleNamespace(rank=4, tp_rank=1, pp_rank=0, cp_rank=0), + effects=None, + registry=SimpleNamespace(canceled_request_ids=lambda: []), + enable_attention_dp=False, + force_terminate_ctx_for_partial_reuse=False, + delegates=None, + ) + + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + failing_emit = Mock(side_effect=RuntimeError("diagnostics failed")) + monkeypatch.setattr(diagnostics, "emit_event", failing_emit) + monkeypatch.setattr(coordinator_module.time, "monotonic", lambda: 12.5) + + coordinator.send_completed_context([request]) + + assert failing_emit.call_count > 0 + assert [operation[0] for operation in operations] == ["start_transfer", "respond"] + assert request.py_kv_transfer_start_time == 12.5 + + +def test_bridge_validation_rejection_emits_failed_ctx_settlement_after_release( + monkeypatch: pytest.MonkeyPatch, +) -> None: + operations = [] + request = SimpleNamespace( + is_context_only_request=True, + is_finished_due_to_cancellation=False, + is_child=False, + py_request_id=32, + request_id=32, + py_disaggregated_params=SimpleNamespace(disagg_request_id=3032), + is_context_finished=True, + is_finished_due_to_length=False, + prompt_len=128, + state=LlmRequestState.CONTEXT_INIT, + py_kv_transfer_start_time=None, + ) + transfer_manager = SimpleNamespace( + start_transfer=lambda req: operations.append(("start_transfer", req)), + should_store_blocks=True, + ) + + def reject(req) -> None: + operations.append(("respond", req)) + req.state = LlmRequestState.DISAGG_TRANS_ERROR + + transceiver = SimpleNamespace( + _fp4_mla_bridge_enabled=True, + _instance_name="ctx-test", + has_retired_send_session=lambda _req: False, + respond_and_send_async=reject, + has_inflight_transfer=lambda _req: False, + kv_transfer_timeout_ms=100, + pipeline_transfer_enabled=False, + ) + coordinator = DisaggTransferCoordinator( + transceiver=transceiver, + transfer_manager=transfer_manager, + kv_cache_manager=None, + dist=SimpleNamespace(rank=4, tp_rank=1, pp_rank=0, cp_rank=0, dp_rank=2), + effects=None, + registry=SimpleNamespace(canceled_request_ids=lambda: []), + enable_attention_dp=False, + force_terminate_ctx_for_partial_reuse=False, + delegates=None, + ) + coordinator.release_transfer = lambda req: operations.append(("release", req)) + + def emit_event(event: str, **kwargs) -> None: + operations.append((event, kwargs)) + + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + coordinator.send_completed_context([request]) + + assert [operation[0] for operation in operations] == [ + "start_transfer", + "ctx_send_ready", + "respond", + "release", + "ctx_transfer_settled", + ] + settlement = operations[4][1] + assert settlement["request_id"] == 3032 + assert settlement["local_request_id"] == 32 + assert settlement["outcome"] == "failed" + assert settlement["session_status"] is None + assert settlement["resources_drained"] is True + assert request.py_kv_transfer_start_time is None + + +@pytest.mark.parametrize( + ("wait_result", "session_status", "outcome", "request_state"), + [ + ( + WaitResult.COMPLETED, + SessionStatus.FULLY_TRANSFERRED, + "completed", + LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE, + ), + ( + WaitResult.FAILED, + SessionStatus.ERROR, + "failed", + LlmRequestState.DISAGG_TRANS_ERROR, + ), + ], +) +def test_sync_receive_emits_start_and_terminal_settlement( + monkeypatch: pytest.MonkeyPatch, + wait_result: WaitResult, + session_status: SessionStatus, + outcome: str, + request_state: LlmRequestState, +) -> None: + operations = [] + session = SimpleNamespace( + status=session_status, + receive=Mock(side_effect=lambda _slice: operations.append("receive")), + wait_complete=Mock(side_effect=lambda blocking: operations.append("wait") or wait_result), + has_transferring_tasks=Mock(return_value=False), + close=Mock(side_effect=lambda: operations.append("close") or True), + ) + request = _disagg_request(41, 4041) + request.state = LlmRequestState.DISAGG_GENERATION_INIT + request.set_kv_cache_size = Mock() + transceiver = _diagnostic_transceiver() + transceiver._validate_bridge_req = Mock(return_value=True) + transceiver._recv_sessions = {} + transceiver._recv_reqs = {} + transceiver._transfer_worker = SimpleNamespace(create_rx_session=Mock(return_value=session)) + transceiver._create_kv_slice = Mock(return_value=KVSlice(is_last_slice=True)) + transceiver._slice_num_bytes = Mock(return_value=64) + transceiver._kv_size_rank_factor = 2 + transceiver._need_aux_transfer = Mock(return_value=False) + transceiver._assert_disagg_history_declared = Mock() + + def emit_event(event: str, **kwargs) -> None: + if event == "gen_transfer_settled": + assert transceiver._recv_sessions == {} + assert transceiver._recv_reqs == {} + operations.append(event) + emitted.append((event, kwargs)) + + emitted = [] + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + transceiver.request_and_receive_sync(request) + + assert operations == [ + "gen_receive_start", + "receive", + "wait", + "close", + "gen_transfer_settled", + ] + assert [event for event, _kwargs in emitted] == [ + "gen_receive_start", + "gen_transfer_settled", + ] + receive_start = emitted[0][1] + settlement = emitted[1][1] + assert receive_start["request_id"] == settlement["request_id"] == 4041 + assert receive_start["local_request_id"] == settlement["local_request_id"] == 41 + assert receive_start["transfer_bytes"] == 128 + assert receive_start["timeout_expected"] is False + assert settlement["outcome"] == outcome + assert settlement["session_status"] == session_status.value + assert settlement["resources_drained"] is True + assert request.state == request_state + assert transceiver._recv_sessions == {} + assert transceiver._recv_reqs == {} + transceiver._validate_bridge_req.assert_called_once_with(request, synchronous=True) + session.wait_complete.assert_called_once_with(blocking=True) + if wait_result == WaitResult.COMPLETED: + request.set_kv_cache_size.assert_called_once_with(128) + transceiver._assert_disagg_history_declared.assert_called_once_with(request) + else: + request.set_kv_cache_size.assert_not_called() + transceiver._assert_disagg_history_declared.assert_not_called() + + +def test_sync_receive_does_not_report_settlement_while_close_is_refused( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = SimpleNamespace( + status=SessionStatus.ERROR, + receive=Mock(side_effect=RuntimeError("receive failed")), + has_transferring_tasks=Mock(return_value=True), + close=Mock(return_value=False), + ) + request = _disagg_request(44, 4044) + request.state = LlmRequestState.DISAGG_GENERATION_INIT + transceiver = _diagnostic_transceiver() + transceiver._validate_bridge_req = Mock(return_value=True) + transceiver._recv_sessions = {} + transceiver._recv_reqs = {} + transceiver._transfer_worker = SimpleNamespace(create_rx_session=Mock(return_value=session)) + transceiver._create_kv_slice = Mock(return_value=KVSlice(is_last_slice=True)) + transceiver._slice_num_bytes = Mock(return_value=64) + transceiver._kv_size_rank_factor = 2 + + emit_event = Mock() + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + with pytest.raises(RuntimeError, match="receive failed"): + transceiver.request_and_receive_sync(request) + + assert [call.args[0] for call in emit_event.call_args_list] == ["gen_receive_start"] + assert transceiver._recv_sessions == {4044: session} + assert transceiver._recv_reqs == {4044: request} + + +def test_sync_receive_disabled_diagnostics_preserves_receive_failure_order( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = SimpleNamespace( + receive=Mock(side_effect=RuntimeError("receive failed")), + close=Mock(return_value=True), + ) + request = _disagg_request(45, 4045) + request.state = LlmRequestState.DISAGG_GENERATION_INIT + transceiver = _diagnostic_transceiver() + transceiver._validate_bridge_req = Mock(return_value=True) + transceiver._recv_sessions = {} + transceiver._recv_reqs = {} + transceiver._transfer_worker = SimpleNamespace(create_rx_session=Mock(return_value=session)) + kv_slice = KVSlice(is_last_slice=True) + transceiver._create_kv_slice = Mock(return_value=kv_slice) + transceiver._slice_num_bytes = Mock(side_effect=AssertionError("must not size before receive")) + + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", False) + + with pytest.raises(RuntimeError, match="receive failed"): + transceiver.request_and_receive_sync(request) + + session.receive.assert_called_once_with(kv_slice) + transceiver._slice_num_bytes.assert_not_called() + assert request.state == LlmRequestState.DISAGG_TRANS_ERROR + assert transceiver._recv_sessions == {} + assert transceiver._recv_reqs == {} + + +@pytest.mark.parametrize("side", ("ctx", "gen")) +def test_fast_cancel_emits_one_terminal_settlement( + monkeypatch: pytest.MonkeyPatch, + side: str, +) -> None: + request = _disagg_request(42, 4042) + request.py_kv_send_session_retired = False + session = SimpleNamespace( + status=SessionStatus.READY, + has_transferring_tasks=Mock(return_value=False), + close=Mock(return_value=True), + ) + session.cancel = Mock(side_effect=lambda: setattr(session, "status", SessionStatus.CANCELLED)) + transceiver = _diagnostic_transceiver() + transceiver._wait_reqs = {} + transceiver._send_sessions = {4042: session} if side == "ctx" else {} + transceiver._send_reqs = {4042: request} if side == "ctx" else {} + transceiver._recv_sessions = {4042: session} if side == "gen" else {} + transceiver._recv_reqs = {4042: request} if side == "gen" else {} + + emitted = [] + + def emit_event(event: str, **kwargs) -> None: + if event.endswith("_transfer_settled"): + assert 4042 not in transceiver._send_sessions + assert 4042 not in transceiver._recv_sessions + emitted.append((event, kwargs)) + + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + assert transceiver.cancel_request(request) is True + assert transceiver.cancel_request(request) is True + + assert [event for event, _kwargs in emitted] == [ + "transfer_cancel_requested", + f"{side}_transfer_settled", + ] + settlement = emitted[1][1] + assert settlement["outcome"] == "cancelled" + assert settlement["session_status"] == SessionStatus.CANCELLED.value + assert settlement["resources_drained"] is True + session.close.assert_called_once_with() + sessions = transceiver._send_sessions if side == "ctx" else transceiver._recv_sessions + requests = transceiver._send_reqs if side == "ctx" else transceiver._recv_reqs + assert sessions == {} + assert requests == {} + if side == "ctx": + assert request.py_kv_send_session_retired is True + + +@pytest.mark.parametrize("side", ("ctx", "gen")) +def test_active_cancel_does_not_emit_terminal_settlement( + monkeypatch: pytest.MonkeyPatch, + side: str, +) -> None: + request = _disagg_request(43, 4043) + session = SimpleNamespace( + status=SessionStatus.TRANSFERRING, + cancel=Mock(), + has_transferring_tasks=Mock(return_value=True), + close=Mock(), + ) + transceiver = _diagnostic_transceiver() + transceiver._wait_reqs = {} + transceiver._send_sessions = {4043: session} if side == "ctx" else {} + transceiver._send_reqs = {4043: request} if side == "ctx" else {} + transceiver._recv_sessions = {4043: session} if side == "gen" else {} + transceiver._recv_reqs = {4043: request} if side == "gen" else {} + + emit_event = Mock() + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + assert transceiver.cancel_request(request) is False + + assert [call.args[0] for call in emit_event.call_args_list] == ["transfer_cancel_requested"] + session.close.assert_not_called() + + +def test_bypassed_transfer_window_reports_legacy_budget_counterfactual( + monkeypatch: pytest.MonkeyPatch, +) -> None: + active = _disagg_request(20, 2020, prompt_len=32) + active.is_disagg_generation_transmission_in_progress = True + candidates = [ + _disagg_request(21, 2021, prompt_len=64), + _disagg_request(22, 2022, prompt_len=32), + ] + controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=64, + tokens_per_block=32, + ) + executor = object.__new__(PyExecutor) + executor.active_requests = [active] + executor.global_rank = 4 + executor.dist = SimpleNamespace(tp_rank=1, pp_rank=0, cp_rank=0) + executor._is_disagg_gen_only_no_context_benchmark = Mock(return_value=False) + executor._get_disagg_transfer_admission_controller = Mock(return_value=controller) + executor._disagg_transfer_window_is_active = Mock(return_value=False) + executor._is_disagg_transfer_window_bypass_eligible = Mock(return_value=True) + + emit_event = Mock() + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + admitted, wait_for_progress = executor._apply_disagg_transfer_admission(candidates) + + assert admitted == candidates + assert wait_for_progress is False + assert emit_event.call_count == 2 + for event, request in zip(emit_event.call_args_list, candidates): + assert event.args == ("gen_transfer_window_result",) + assert event.kwargs["request_id"] == request.py_disaggregated_params.disagg_request_id + assert event.kwargs["outcome"] == "admitted" + assert event.kwargs["policy"] == "bypassed" + assert event.kwargs["legacy_budget_outcome"] == "deferred" + assert event.kwargs["legacy_active_transfer_blocks"] == 1 + assert event.kwargs["legacy_admitted_transfer_blocks"] == 0 + assert event.kwargs["legacy_limited_by_budget"] is True + assert event.kwargs["transfer_block_budget"] == 2 + + +def test_transfer_window_result_uses_full_cp_prompt_length( + monkeypatch: pytest.MonkeyPatch, +) -> None: + request = _disagg_request(23, 2023, prompt_len=1) + request.total_input_len_cp = 257 + controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=512, + tokens_per_block=32, + ) + executor = object.__new__(PyExecutor) + executor.global_rank = 4 + executor.dist = SimpleNamespace(tp_rank=1, pp_rank=0, cp_rank=3) + + emit_event = Mock() + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + executor._emit_disagg_transfer_window_results( + [request], + [request], + policy="enforced", + controller=controller, + active_transfer_blocks=0, + admitted_transfer_blocks=9, + ) + + event = emit_event.call_args + assert event.args == ("gen_transfer_window_result",) + assert event.kwargs["prompt_tokens"] == 257 + assert event.kwargs["request_blocks"] == 9 + assert event.kwargs["cp_rank"] == 3 + + +def test_disabled_diagnostics_do_not_evaluate_bypass_counterfactual( + monkeypatch: pytest.MonkeyPatch, +) -> None: + candidate = _disagg_request(21, 2021) + controller = Mock() + controller.enabled.return_value = True + controller.select.side_effect = AssertionError( + "disabled diagnostics evaluated the legacy transfer window" + ) + executor = object.__new__(PyExecutor) + executor.active_requests = [] + executor._is_disagg_gen_only_no_context_benchmark = Mock(return_value=False) + executor._get_disagg_transfer_admission_controller = Mock(return_value=controller) + executor._disagg_transfer_window_is_active = Mock(return_value=False) + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", False) + + admitted, wait_for_progress = executor._apply_disagg_transfer_admission([candidate]) + + assert admitted == [candidate] + assert wait_for_progress is False + controller.select.assert_not_called() + + +def test_diagnostic_preparation_failure_does_not_change_bypassed_admission( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class _BrokenLegacyResult: + @property + def admitted_requests(self) -> list: + raise RuntimeError("diagnostic counterfactual inspection failed") + + candidate = _disagg_request(21, 2021) + controller = Mock() + controller.enabled.return_value = True + controller.select.return_value = _BrokenLegacyResult() + executor = object.__new__(PyExecutor) + executor.active_requests = [] + executor._is_disagg_gen_only_no_context_benchmark = Mock(return_value=False) + executor._get_disagg_transfer_admission_controller = Mock(return_value=controller) + executor._disagg_transfer_window_is_active = Mock(return_value=False) + executor._is_disagg_transfer_window_bypass_eligible = Mock(return_value=True) + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + + admitted, wait_for_progress = executor._apply_disagg_transfer_admission([candidate]) + + assert admitted == [candidate] + assert wait_for_progress is False + executor._is_disagg_transfer_window_bypass_eligible.assert_called_once_with() + controller.select.assert_called_once_with([], [candidate]) + + +def test_transfer_window_helper_contains_diagnostic_preparation_failure( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class _OpaqueCandidate: + @property + def py_request_id(self) -> int: + raise RuntimeError("diagnostic request inspection failed") + + candidate = _OpaqueCandidate() + executor = object.__new__(PyExecutor) + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + + executor._emit_disagg_transfer_window_results( + [candidate], + [candidate], + policy="bypassed", + ) + + +def test_admission_rollback_continues_when_diagnostic_emission_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + candidates = [ + _disagg_request(31, 3031, prompt_len=32), + _disagg_request(32, 3032, prompt_len=32), + ] + controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=32, + tokens_per_block=32, + ) + executor = object.__new__(PyExecutor) + executor.active_requests = [] + executor._is_disagg_gen_only_no_context_benchmark = Mock(return_value=False) + executor._get_disagg_transfer_admission_controller = Mock(return_value=controller) + executor._disagg_transfer_window_is_active = Mock(return_value=True) + executor._emit_disagg_transfer_window_results = Mock( + side_effect=RuntimeError("diagnostics failed") + ) + executor._revert_deferred_disagg_gen_init_alloc = Mock() + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + + admitted, wait_for_progress = executor._apply_disagg_transfer_admission(candidates) + + assert admitted == [candidates[0]] + assert wait_for_progress is False + executor._revert_deferred_disagg_gen_init_alloc.assert_called_once_with( + candidates, + [candidates[0]], + reason="transfer_window", + ) + + +def test_pp_reconciliation_emits_rollback_after_releasing_kv( + monkeypatch: pytest.MonkeyPatch, +) -> None: + admitted = _disagg_request(33, 3033) + deferred = _disagg_request(34, 3034) + operations = [] + executor = object.__new__(PyExecutor) + executor._is_kv_manager_v2 = True + executor._revert_ctx_alloc = lambda requests: operations.append(("revert", requests)) + executor.global_rank = 4 + executor.dist = SimpleNamespace(tp_rank=1, pp_rank=2, cp_rank=3) + + def emit_event(event: str, **kwargs) -> None: + operations.append((event, kwargs)) + + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + executor._revert_deferred_disagg_gen_init_alloc( + [admitted, deferred], + [admitted], + ) + + assert operations[0] == ("revert", [deferred]) + assert operations[1][0] == "gen_kv_rollback" + event = operations[1][1] + assert event["request_id"] == 3034 + assert event["reason"] == "pp_reconciliation" + assert event["pp_rank"] == 2 + assert event["cp_rank"] == 3 + + +def test_ctx_settlement_continues_when_diagnostic_emission_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + request_id = 41 + session = SimpleNamespace( + wait_complete=Mock(return_value=WaitResult.COMPLETED), + status=SessionStatus.READY, + has_transferring_tasks=Mock(return_value=False), + ) + request = SimpleNamespace(py_request_id=4) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._ever_had_send_session = True + transceiver._ctx_need_tp_sync = False + transceiver._ctx_need_pp_sync = False + transceiver._send_sessions = {request_id: session} + transceiver._send_reqs = {request_id: request} + transceiver._collect_done = Mock(return_value=([request_id], [])) + transceiver._ctx_consensus = Mock(side_effect=lambda request_ids: request_ids) + transceiver._build_to_process = Mock(return_value=[request_id]) + transceiver._ctx_consensus_outcome = Mock(return_value=([], [], [request_id], [request_id])) + + def retire_send_session(rid: int, **_kwargs) -> None: + transceiver._send_sessions.pop(rid, None) + transceiver._send_reqs.pop(rid, None) + + transceiver._retire_send_session = Mock(side_effect=retire_send_session) + transceiver._close_failed_sessions = Mock() + transceiver._transfer_worker = SimpleNamespace(sweep_stale_req_infos=Mock()) + transceiver._mapping = SimpleNamespace(rank=1, tp_rank=0, pp_rank=0, cp_rank=0) + transceiver._instance_name = "ctx" + transceiver._dp_rank = 0 + + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + emitted_events = [] + + def fail_after_ctx_release(_event: str, **_kwargs) -> None: + emitted_events.append( + ( + _event, + tuple(transceiver._send_sessions), + tuple(transceiver._send_reqs), + ) + ) + raise RuntimeError("diagnostics failed") + + monkeypatch.setattr(diagnostics, "emit_event", fail_after_ctx_release) + + status = transceiver.check_context_transfer_status(at_least_request_num=0) + + assert status.completed_request_ids == [request_id] + assert status.error_request_ids == [] + assert emitted_events == [("ctx_transfer_settled", (), ())] + transceiver._retire_send_session.assert_called_once_with( + request_id, + outcome="completed", + ) + transceiver._transfer_worker.sweep_stale_req_infos.assert_called_once_with() + + +def test_ctx_settlement_continues_when_diagnostic_session_snapshot_is_stale( + monkeypatch: pytest.MonkeyPatch, +) -> None: + request_id = 43 + session = SimpleNamespace( + wait_complete=Mock(return_value=WaitResult.COMPLETED), + status=SessionStatus.READY, + has_transferring_tasks=Mock(return_value=False), + ) + request = SimpleNamespace(py_request_id=6) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._ever_had_send_session = True + transceiver._ctx_need_tp_sync = False + transceiver._ctx_need_pp_sync = False + transceiver._send_sessions = {request_id: session} + transceiver._send_reqs = {request_id: request} + transceiver._collect_done = Mock(return_value=([request_id], [])) + transceiver._ctx_consensus = Mock(side_effect=lambda request_ids: request_ids) + transceiver._build_to_process = Mock(return_value=[request_id]) + + def retire_before_diagnostic_snapshot(*_args) -> tuple[list, list, list, list]: + transceiver._send_sessions.pop(request_id) + return [], [], [request_id], [request_id] + + transceiver._ctx_consensus_outcome = Mock(side_effect=retire_before_diagnostic_snapshot) + transceiver._retire_send_session = Mock( + side_effect=lambda rid, **_kwargs: transceiver._send_reqs.pop(rid, None) + ) + transceiver._close_failed_sessions = Mock() + transceiver._transfer_worker = SimpleNamespace(sweep_stale_req_infos=Mock()) + transceiver._mapping = SimpleNamespace(rank=1, tp_rank=0, pp_rank=0, cp_rank=0) + transceiver._instance_name = "ctx" + transceiver._dp_rank = 0 + emit_event = Mock() + + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + status = transceiver.check_context_transfer_status(at_least_request_num=0) + + assert status.completed_request_ids == [request_id] + assert status.error_request_ids == [] + transceiver._retire_send_session.assert_called_once_with( + request_id, + outcome="completed", + ) + emit_event.assert_not_called() + + +def test_gen_settlement_continues_when_diagnostic_emission_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + request_id = 42 + session = SimpleNamespace( + wait_complete=Mock(return_value=WaitResult.COMPLETED), + status=SessionStatus.READY, + transfer_end_time=None, + kv_cache_size_bytes=0, + has_transferring_tasks=Mock(return_value=False), + ) + request = SimpleNamespace( + py_request_id=5, + set_kv_cache_size=Mock(), + ) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._ever_had_recv_session = True + transceiver._gen_need_sync = False + transceiver._recv_sessions = {request_id: session} + transceiver._recv_reqs = {request_id: request} + transceiver._collect_done = Mock(return_value=([request_id], [])) + transceiver._gen_consensus = Mock(side_effect=lambda request_ids: request_ids) + transceiver._build_to_process = Mock(return_value=[request_id]) + transceiver._gen_consensus_outcome = Mock(return_value=([], [], [request_id])) + transceiver._need_aux_transfer = Mock(return_value=False) + transceiver._assert_disagg_history_declared = Mock() + transceiver._close_session_or_raise = Mock() + transceiver._close_failed_sessions = Mock() + transceiver._mapping = SimpleNamespace(rank=1, tp_rank=0, pp_rank=0, cp_rank=0) + transceiver._instance_name = "gen" + transceiver._dp_rank = 0 + transceiver._dist = SimpleNamespace(rank=1) + + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + emitted_events = [] + + def fail_after_gen_release(_event: str, **_kwargs) -> None: + emitted_events.append( + ( + _event, + tuple(transceiver._recv_sessions), + tuple(transceiver._recv_reqs), + ) + ) + raise RuntimeError("diagnostics failed") + + monkeypatch.setattr(diagnostics, "emit_event", fail_after_gen_release) + + status = transceiver.check_gen_transfer_status(at_least_request_num=0) + + assert status.completed_request_ids == [request_id] + assert status.error_request_ids == [] + assert status.cancelled_requests == [] + assert emitted_events == [("gen_transfer_settled", (), ())] + request.set_kv_cache_size.assert_called_once_with(0) + transceiver._close_session_or_raise.assert_called_once_with( + session, + request_id, + "completed", + ) + assert request.state == LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE + assert transceiver._recv_sessions == {} + assert transceiver._recv_reqs == {} + + +@pytest.mark.parametrize("enforce_physical_ownership", [False, True]) +def test_backend_wait_continues_when_submission_diagnostic_callback_fails( + enforce_physical_ownership: bool, +) -> None: + status = SimpleNamespace(wait=Mock(return_value=True)) + sender = object.__new__(Sender) + sender._enforce_physical_ownership = enforce_physical_ownership + sender._agent = SimpleNamespace(submit_transfer_requests=Mock(return_value=status)) + sender._ownership_poison_lock = threading.Lock() + sender._ownership_poisoned = None + task = Mock() + request = Mock() + callback = Mock(side_effect=RuntimeError("diagnostics failed")) + + result = sender._submit_transfer( + task, + 7, + request, + on_submitted=callback, + ) + + assert result == (True, None) + callback.assert_called_once_with() + status.wait.assert_called_once_with() + if enforce_physical_ownership: + task.begin_backend_submission.assert_called_once_with(7, request) + task.record_backend_submission.assert_called_once_with(7, status) + task.retire_backend_done_physical_operation.assert_called_once_with(7) + else: + task.begin_backend_submission.assert_not_called() + task.record_backend_submission.assert_not_called() + task.retire_backend_done_physical_operation.assert_not_called() + + +def test_request_cleanup_continues_when_diagnostic_emission_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + request = SimpleNamespace( + is_context_only_request=True, + py_request_id=51, + request_id=51, + py_disaggregated_params=SimpleNamespace(disagg_request_id=5051), + prompt_len=128, + state=SimpleNamespace(name="DISAGG_CONTEXT_COMPLETE"), + ) + executor = object.__new__(PyExecutor) + executor.resource_manager = SimpleNamespace(free_resources=Mock()) + executor._prefetched_request_ids = {request.py_request_id} + executor._disagg_coordinator = SimpleNamespace(forget_request=Mock()) + executor.global_rank = 4 + executor.dist = SimpleNamespace(tp_rank=1, pp_rank=0, cp_rank=0) + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + failing_emit = Mock(side_effect=RuntimeError("diagnostics failed")) + monkeypatch.setattr(diagnostics, "emit_event", failing_emit) + + executor._free_request_resources(request) + + failing_emit.assert_called_once() + executor.resource_manager.free_resources.assert_called_once_with(request) + assert request.py_request_id not in executor._prefetched_request_ids + executor._disagg_coordinator.forget_request.assert_called_once_with(request.py_request_id) + + +def test_source_unpin_continues_when_diagnostic_inspection_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class _KVCacheManager: + def __init__(self) -> None: + self.unpin_blocks_by_id = Mock() + self.mapping_accessed = False + + @property + def mapping(self) -> SimpleNamespace: + self.mapping_accessed = True + raise RuntimeError("diagnostic mapping inspection failed") + + request = SimpleNamespace( + is_context_only_request=True, + py_request_id=61, + state=LlmRequestState.DISAGG_CONTEXT_TRANS_IN_PROGRESS, + ) + block_ids = [9] + metadata = AsyncTransferManager.RequestTransferMetadata(block_id=block_ids) + metadata.start_transfer() + manager = object.__new__(AsyncTransferManager) + manager.should_store_blocks = True + manager.kv_cache_manager = _KVCacheManager() + manager._requests_in_transfer = {request.py_request_id: request} + manager._request_transfer_metadata = {request.py_request_id: metadata} + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + + should_terminate = manager.end_transfer(request) + + assert should_terminate is True + assert manager.kv_cache_manager.mapping_accessed + manager.kv_cache_manager.unpin_blocks_by_id.assert_called_once_with(block_ids) + assert request.state == LlmRequestState.DISAGG_CONTEXT_COMPLETE + assert manager._requests_in_transfer == {} + assert manager._request_transfer_metadata == {} + + +def test_source_unpin_summarizes_real_block_id_shape( + monkeypatch: pytest.MonkeyPatch, +) -> None: + block_ids = [9, 10, 11] + request = SimpleNamespace( + is_context_only_request=True, + py_request_id=62, + py_disaggregated_params=SimpleNamespace(disagg_request_id=6062), + state=LlmRequestState.DISAGG_CONTEXT_TRANS_IN_PROGRESS, + ) + metadata = AsyncTransferManager.RequestTransferMetadata(block_id=block_ids) + metadata.start_transfer() + manager = object.__new__(AsyncTransferManager) + manager.should_store_blocks = True + manager.kv_cache_manager = SimpleNamespace( + mapping=SimpleNamespace(rank=4), + unpin_blocks_by_id=Mock(), + ) + manager._requests_in_transfer = {request.py_request_id: request} + manager._request_transfer_metadata = {request.py_request_id: metadata} + emit_event = Mock() + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + + assert manager.end_transfer(request) is True + + manager.kv_cache_manager.unpin_blocks_by_id.assert_called_once_with(block_ids) + event = emit_event.call_args + assert event.args == ("ctx_source_unpinned",) + assert event.kwargs["source_kv_reuse_block_count"] == 3 + + +def test_native_writer_result_precedes_local_destination_completion( + monkeypatch: pytest.MonkeyPatch, +) -> None: + request_id = 3031 + writer_rank = 7 + receiver = object.__new__(Receiver) + receiver._enforce_physical_ownership = True + receiver._sessions = {} + receiver._sessions_lock = threading.Lock() + receiver._pre_cancelled_rids = set() + receiver._shutdown = False + receiver._bounce = Mock() + receiver._bounce.is_bounced.return_value = False + receiver._registrar = SimpleNamespace( + self_rank_info=SimpleNamespace(instance_name="gen", instance_rank=2), + ) + + session = RxSession( + request_id=31, + params=DisaggregatedParams(disagg_request_id=request_id), + receiver=receiver, + ) + + def dispatch(task) -> None: + task.expected_transfers = 1 + session.mark_transferring(task.slice_id, writer_cohort={writer_rank}) + + receiver.dispatch_task = dispatch + emit_event = Mock() + monkeypatch.setattr(diagnostics, "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(diagnostics, "emit_event", emit_event) + monkeypatch.setattr( + transfer_module.tensorrt_llm.bindings, + "global_steady_clock_now", + lambda: 0, + ) + + session.receive(KVSlice(is_last_slice=True)) + destination_timestamp_captured = threading.Event() + destination_timestamp = (123_000, 456_000) + + def capture_timestamp() -> tuple[int, int]: + destination_timestamp_captured.set() + return destination_timestamp + + task = session._kv_tasks[0] + original_complete = task.complete + + def complete_after_timestamp_capture() -> None: + assert destination_timestamp_captured.is_set() + original_complete() + + monkeypatch.setattr(diagnostics, "capture_timestamp", capture_timestamp) + monkeypatch.setattr(task, "complete", complete_after_timestamp_capture) + message = transfer_module._make_kv_result_msg( + writer_rank, + request_id, + 0, + True, + AgentResult.SUCCESS, + transfer_size=4096, + ) + receiver._process_kv_agent_result(b"sender", message) + + assert session._kv_tasks[0].status is TaskStatus.TRANSFERRED + assert session.resources_drained() + assert [entry.args[0] for entry in emit_event.call_args_list] == [ + "gen_writer_result_received", + "gen_destination_complete", + ] + writer_event = emit_event.call_args_list[0].kwargs + destination_event = emit_event.call_args_list[1].kwargs + assert writer_event["request_id"] == destination_event["request_id"] == request_id + assert writer_event["slice_id"] == destination_event["slice_id"] == 0 + assert writer_event["peer_rank"] == destination_event["peer_rank"] == writer_rank + assert writer_event["outcome"] == "success" + assert writer_event["transfer_bytes"] == 4096 + assert destination_event["outcome"] == "completed" + assert destination_event["timestamp"] == destination_timestamp + assert session.close() diff --git a/tests/unittest/disaggregated/test_transfer_ownership_regressions.py b/tests/unittest/disaggregated/test_transfer_ownership_regressions.py index a92c83ce7de1..b92c86ff32a7 100644 --- a/tests/unittest/disaggregated/test_transfer_ownership_regressions.py +++ b/tests/unittest/disaggregated/test_transfer_ownership_regressions.py @@ -567,7 +567,9 @@ def test_non_terminal_writer_result_does_not_authorize_reuse() -> None: @pytest.mark.cpu_only -def test_gen_first_no_retry_adp_count_seal_waits_for_one_writer_group() -> None: +def test_gen_first_no_retry_adp_count_seal_waits_for_one_writer_group( + monkeypatch: pytest.MonkeyPatch, +) -> None: rid = 94 receiver = object.__new__(Receiver) receiver._sessions_lock = threading.Lock() @@ -604,7 +606,43 @@ def test_gen_first_no_retry_adp_count_seal_waits_for_one_writer_group() -> None: cp_size=1, ) ) - receiver._request_sender_data = Mock() + publication_order: list[tuple[str, int]] = [] + published_ranks: list[int] = [] + captured_timestamps: list[tuple[int, int]] = [] + publication_timestamps: dict[int, tuple[int, int]] = {} + fast_response_timestamps: dict[int, tuple[int, int]] = {} + + def request_sender_data(endpoint: str, _payload: bytes) -> None: + rank = int(endpoint.rsplit("-", 1)[1]) + assert publication_order[-1][0] == "capture" + publication_timestamps[rank] = captured_timestamps[-1] + published_ranks.append(rank) + publication_order.append(("send", rank)) + # Model a listener response that races ahead of the deferred event + # emission after this send returns. + fast_response_timestamps[rank] = ( + captured_timestamps[-1][0] + 50, + captured_timestamps[-1][1] + 50, + ) + publication_order.append(("result", rank)) + + def capture_timestamp() -> tuple[int, int]: + index = len(captured_timestamps) + timestamp = (index + 100, index + 200) + captured_timestamps.append(timestamp) + publication_order.append(("capture", index)) + return timestamp + + def emit_event(event: str, **kwargs) -> None: + if event == "gen_request_data_sent": + rank = kwargs["peer_rank"] + assert kwargs["writer_cohort_known"] is False + assert kwargs["expected_writers"] == 2 + assert kwargs["timestamp"] == publication_timestamps[rank] + assert kwargs["timestamp"][0] < fast_response_timestamps[rank][0] + publication_order.append(("emit", rank)) + + receiver._request_sender_data = request_sender_data session = RxSession( request_id=rid, params=DisaggregatedParams( @@ -617,13 +655,32 @@ def test_gen_first_no_retry_adp_count_seal_waits_for_one_writer_group() -> None: task = session.prepare_receive(KVSlice(is_last_slice=True)) assert task is not None + monkeypatch.setattr( + transfer_mod.disagg_diagnostics, + "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", + True, + ) + monkeypatch.setattr( + transfer_mod.disagg_diagnostics, + "capture_timestamp", + capture_timestamp, + ) + monkeypatch.setattr(transfer_mod.disagg_diagnostics, "emit_event", emit_event) session.dispatch_prepared_receive(task) assert task.expected_transfers == 2 - assert {call.args[0] for call in receiver._request_sender_data.call_args_list} == { + assert {f"tcp://sender-{rank}" for rank in published_ranks} == { f"tcp://sender-{rank}" for rank in range(4) } - assert receiver._request_sender_data.call_count == 4 + assert len(published_ranks) == 4 + assert publication_order == [ + *( + entry + for index, rank in enumerate(published_ranks) + for entry in (("capture", index), ("send", rank), ("result", rank)) + ), + *(("emit", rank) for rank in published_ranks), + ] session.process_kv_agent_result( peer_rank=2, receiver_slice_id=0, @@ -713,6 +770,13 @@ def request_sender_data(endpoint: str, _payload: bytes) -> None: ) receiver._request_sender_data = request_sender_data + emit_event = Mock() + monkeypatch.setattr( + transfer_mod.disagg_diagnostics, + "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", + True, + ) + monkeypatch.setattr(transfer_mod.disagg_diagnostics, "emit_event", emit_event) monkeypatch.setattr( transfer_mod.tensorrt_llm.bindings, "global_steady_clock_now", @@ -732,6 +796,14 @@ def request_sender_data(endpoint: str, _payload: bytes) -> None: session.receive(KVSlice(is_last_slice=True)) assert queued_endpoints == ["tcp://sender-0"] + request_data_events = [ + call for call in emit_event.call_args_list if call.args == ("gen_request_data_sent",) + ] + assert len(request_data_events) == 1 + assert request_data_events[0].kwargs["request_id"] == rid + assert request_data_events[0].kwargs["peer_rank"] == 0 + assert request_data_events[0].kwargs["ownership_enabled"] is True + assert request_data_events[0].kwargs["writer_cohort_known"] is True if not writer_settles_before_failure: assert bounce.context is not None assert bounce.release_count == 0 @@ -1321,6 +1393,11 @@ def _make_owned_sender() -> transfer_mod.Sender: sender = object.__new__(transfer_mod.Sender) sender._enforce_physical_ownership = True sender._sessions_lock, sender._sessions = threading.Lock(), {} + sender._peer_requests_lock = threading.Lock() + sender._peer_requests = {} + sender._peer_requests_timestamps = {} + sender._peer_requests_ready_timestamps = {} + sender._pre_cancelled_rids = set() sender._shutdown = sender._shutdown_requested = False sender._ownership_poisoned, sender._ownership_poison_lock = None, threading.Lock() sender._loaded_remote_agents_lock, sender._loaded_remote_agents = threading.Lock(), set() @@ -1328,6 +1405,135 @@ def _make_owned_sender() -> transfer_mod.Sender: return sender +@pytest.mark.cpu_only +def test_gen_first_ready_timestamp_survives_until_sender_session_setup( + monkeypatch: pytest.MonkeyPatch, +) -> None: + rid = 107 + ready_timestamp = (123, 456) + sender = _make_owned_sender() + sender._registrar = SimpleNamespace( + self_rank_info=SimpleNamespace(instance_name="ctx", instance_rank=0), + get_peer_rank_info=Mock(return_value=SimpleNamespace(dp_rank=0)), + get_peer_overlap=Mock(return_value=SimpleNamespace(ranks=[2])), + ) + info = transfer_mod.RecvReqInfo( + sender_req_id=18, + instance_name="gen", + instance_rank=2, + block_ids_per_layer_groups=[], + unique_rid=rid, + slice_id=0, + ) + capture_timestamp = Mock(return_value=ready_timestamp) + emit_event = Mock() + monkeypatch.setattr( + transfer_mod.disagg_diagnostics, + "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", + True, + ) + monkeypatch.setattr( + transfer_mod.disagg_diagnostics, + "capture_timestamp", + capture_timestamp, + ) + monkeypatch.setattr(transfer_mod.disagg_diagnostics, "emit_event", emit_event) + + sender._respond_with_kv(b"receiver", [MessageType.REQUEST_DATA, info.to_bytes()]) + + assert sender._sessions == {} + assert sender._peer_requests_ready_timestamps[rid] == ready_timestamp + tx_session = SimpleNamespace( + disagg_request_id=rid, + request_id=19, + lock=threading.Lock(), + receiver_ready=False, + ) + + sender.setup_session(tx_session) + + assert tx_session.receiver_ready + capture_timestamp.assert_called_once_with() + ready_events = [ + call for call in emit_event.call_args_list if call.args == ("ctx_all_receivers_ready",) + ] + assert len(ready_events) == 1 + assert ready_events[0].kwargs["timestamp"] == ready_timestamp + + sender.clear_session(rid) + + assert rid not in sender._peer_requests + assert rid not in sender._peer_requests_ready_timestamps + + +@pytest.mark.cpu_only +@pytest.mark.parametrize("diagnostics_enabled", [False, True]) +def test_kv_result_size_is_diagnostics_independent_without_perf_timer( + monkeypatch: pytest.MonkeyPatch, + diagnostics_enabled: bool, +) -> None: + rid = 108 + peer_rank = 2 + sender = _make_owned_sender() + sender._enforce_physical_ownership = False + sender._device_id = 0 + sender._bounce = Mock() + sender._agent = SimpleNamespace( + submit_transfer_requests=Mock(return_value=SimpleNamespace(wait=Mock(return_value=True))) + ) + sender._registrar = SimpleNamespace( + self_rank_info=SimpleNamespace(instance_name="ctx", instance_rank=0) + ) + dealer = Mock() + sender._get_result_dealer = Mock(return_value=dealer) + task = transfer_mod.KVSendTask( + KVSlice(is_last_slice=True), + DisaggregatedParams(disagg_request_id=rid), + slice_id=0, + ) + task._perf_timer = None + session = SimpleNamespace( + kv_tasks=[task], + lock=threading.Lock(), + status=SessionStatus.READY, + set_exception=Mock(), + transfer_end_time=None, + ) + sender._sessions = {rid: session} + write_meta = transfer_mod.WriteMeta( + task=task, + expected_transfers=1, + peer_name="gen2", + peer_rank=peer_rank, + peer_endpoint="receiver", + unique_rid=rid, + src_ptrs=np.array([0x1000, 0x2000], dtype=np.int64), + dst_ptrs=np.array([0x3000, 0x4000], dtype=np.int64), + sizes=np.array([0x100, 0x80], dtype=np.int64), + dst_device_id=0, + slice_id=0, + receiver_slice_id=0, + is_last_slice=True, + ) + monkeypatch.setattr( + transfer_mod.disagg_diagnostics, + "DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", + diagnostics_enabled, + ) + monkeypatch.setattr(transfer_mod.disagg_diagnostics, "emit_event", Mock()) + monkeypatch.setattr(transfer_mod.Sender, "_make_agent_request", Mock(return_value=Mock())) + monkeypatch.setattr( + transfer_mod.tensorrt_llm.bindings, + "global_steady_clock_now", + lambda: 0, + ) + + sender._deliver_kv_to_agent(write_meta) + + result = transfer_mod._KV_RESULT_PREFIX.unpack(dealer.send.call_args.args[0][1]) + assert result[5] == 0 + + @pytest.mark.cpu_only def test_pre_cancelled_sender_settles_saved_generation_first_request(monkeypatch) -> None: rid = 97 @@ -1622,7 +1828,7 @@ def test_late_request_data_to_terminal_sender_settles_aux( unique_rid=rid, ) sender = _make_owned_sender() - sender._save_peer_req_info = Mock() + sender._save_peer_req_info = Mock(return_value=(False, None)) sender._send_failed_result_to_receiver = Mock() session = SimpleNamespace( lock=threading.Lock(), diff --git a/tests/unittest/tools/test_disagg_transfer_diagnostics.py b/tests/unittest/tools/test_disagg_transfer_diagnostics.py new file mode 100644 index 000000000000..7086d964d283 --- /dev/null +++ b/tests/unittest/tools/test_disagg_transfer_diagnostics.py @@ -0,0 +1,1059 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Tests for the disaggregated-transfer diagnostic log analyzer.""" + +import json +from pathlib import Path +from uuid import NAMESPACE_DNS, uuid5 + +import pytest + +__extra_import_path__ = ["~/scripts"] +from disagg_transfer_diagnostics import DIAGNOSTICS_LOG_PREFIX, analyze_lines, main # noqa: E402 + +pytestmark = pytest.mark.cpu_only + + +def _domain(host: str = "node-a", pid: int = 17) -> dict[str, object]: + return { + "host": host, + "pid": pid, + "run_uuid": "11111111-1111-4111-8111-111111111111", + "process_uuid": str(uuid5(NAMESPACE_DNS, f"{host}:{pid}")), + } + + +def _event( + name: str, + request_id: int | None, + monotonic_ns: int, + *, + host: str = "node-a", + pid: int = 17, + side: str = "gen", + **details: object, +) -> str: + record: dict[str, object] = { + "schema_version": 1, + "event": name, + "request_id": request_id, + "side": side, + **_domain(host, pid), + "monotonic_ns": monotonic_ns, + "wall_ns": 100_000_000_000 + monotonic_ns, + } + if name == "gen_request_data_sent": + record["writer_cohort_known"] = True + record.update(details) + return f"worker-prefix {DIAGNOSTICS_LOG_PREFIX}{json.dumps(record)}\n" + + +def _request(result: dict[str, object], request_id: int) -> dict[str, object]: + requests = result["requests"] + assert isinstance(requests, list) + return next(request for request in requests if request["request_id"] == request_id) + + +def _capabilities(*, pid: int = 17, rank: int | None = None) -> str: + return _event( + "diagnostic_capabilities", + None, + 0, + pid=pid, + rank=rank, + side="runtime", + capability_schema_version=1, + transceiver_runtime="PYTHON", + python_transfer_events=True, + executor_events=True, + scheduler_kv_admission_events=True, + ) + + +def test_noisy_logs_are_counted_and_grouped_by_canonical_request() -> None: + lines = [ + "ordinary runtime log\n", + f"{DIAGNOSTICS_LOG_PREFIX}not-json\n", + f"{DIAGNOSTICS_LOG_PREFIX}[]\n", + _event("gen_ingress", 22, 1_000_000), + _event("gen_kv_admission_result", 22, 3_000_000, outcome="admitted"), + _event( + "gen_transfer_window_result", + 22, + 6_000_000, + outcome="admitted", + policy="bypassed", + legacy_budget_outcome="deferred", + legacy_active_transfer_blocks=64, + legacy_admitted_transfer_blocks=0, + legacy_limited_by_budget=True, + ), + _event("transfer_timeout_observed", None, 7_000_000), + ] + + result = analyze_lines(lines) + + assert result["summary"] == { + "total_lines": 7, + "ignored_lines": 1, + "malformed_diagnostic_lines": 2, + "parsed_events": 4, + "events_without_request_id": 1, + "request_count": 1, + } + assert result["event_counts"] == { + "gen_kv_admission_result": 1, + "gen_ingress": 1, + "gen_transfer_window_result": 1, + "transfer_timeout_observed": 1, + } + + request = _request(result, 22) + assert request["missing_boundaries"] == [] + assert request["unassessed_boundaries"] == ["gen_receive_start"] + assert [event["event"] for event in request["timeline"]] == [ + "gen_ingress", + "gen_kv_admission_result", + "gen_transfer_window_result", + ] + transfer_window = request["timeline"][-1] + assert transfer_window["legacy_budget_outcome"] == "deferred" + assert transfer_window["legacy_limited_by_budget"] is True + assert "clock-sync-sensitive" in result["clock_semantics"]["timeline"] + durations = {duration["phase"]: duration["duration_ms"] for duration in request["durations"]} + assert durations == { + "gen_gate1_admission_wait": 2.0, + "gen_transfer_window_admission_wait": 3.0, + } + + +def test_semantically_malformed_correlation_fields_are_quarantined() -> None: + lines = [ + _event("ctx_transfer_queued", 1, 1, side="ctx", slice_id={}), + _event("ctx_transfer_queued", 2, 2, side="ctx", peer_rank=[]), + _event("ctx_transfer_queued", 3, 3, side="ctx", slice_id=True), + _event("ctx_transfer_queued", 4, 4, side="ctx", slice_id=0, peer_rank=1), + ] + + result = analyze_lines(lines) + + assert result["summary"] == { + "total_lines": 4, + "ignored_lines": 0, + "malformed_diagnostic_lines": 3, + "parsed_events": 1, + "events_without_request_id": 0, + "request_count": 1, + } + assert result["requests"][0]["request_id"] == 4 + + +@pytest.mark.parametrize( + "payload", + ( + '{"event":"gen_ingress","request_id":NaN}', + '{"event":"gen_ingress","request_id":Infinity}', + '{"event":"gen_ingress","request_id":18446744073709551616}', + '{"event":"gen_ingress","request_id":' + "1" * 5_000 + "}", + ), +) +def test_non_finite_and_oversized_json_numbers_are_quarantined(payload: str) -> None: + result = analyze_lines([f"{DIAGNOSTICS_LOG_PREFIX}{payload}\n"]) + + assert result["summary"] == { + "total_lines": 1, + "ignored_lines": 0, + "malformed_diagnostic_lines": 1, + "parsed_events": 0, + "events_without_request_id": 0, + "request_count": 0, + } + + +def test_excessively_nested_json_is_quarantined() -> None: + nested = "[" * 2_000 + "0" + "]" * 2_000 + payload = f'{{"event":"gen_ingress","detail":{nested}}}' + + result = analyze_lines([f"{DIAGNOSTICS_LOG_PREFIX}{payload}\n"]) + + assert result["summary"]["malformed_diagnostic_lines"] == 1 + assert result["summary"]["parsed_events"] == 0 + + +@pytest.mark.parametrize("field", ("is_last_slice", "session_found")) +def test_numeric_boolean_field_is_quarantined(field: str) -> None: + line = _event( + "gen_writer_result_received", + 7, + 1, + outcome="success", + **{field: 1}, + ) + + result = analyze_lines([line]) + + assert result["summary"]["malformed_diagnostic_lines"] == 1 + assert result["summary"]["parsed_events"] == 0 + + +@pytest.mark.parametrize("schema_version", (None, "1", True, 0, 2)) +def test_unsupported_schema_version_is_quarantined(schema_version: object) -> None: + line = _event("gen_ingress", 5, 1) + record = json.loads(line.split(DIAGNOSTICS_LOG_PREFIX, 1)[1]) + if schema_version is None: + record.pop("schema_version") + else: + record["schema_version"] = schema_version + + result = analyze_lines([f"{DIAGNOSTICS_LOG_PREFIX}{json.dumps(record)}\n"]) + + assert result["summary"]["malformed_diagnostic_lines"] == 1 + assert result["summary"]["parsed_events"] == 0 + + +@pytest.mark.parametrize( + ("field", "value"), + ( + ("pid", -5), + ("slice_id", -1), + ("rank", "0"), + ("source_kv_reuse_block_count", True), + ("outcome", 7), + ("elapsed_ms", -0.5), + ("elapsed_ms", int("9" * 400)), + ("monotonic_ns", 1 << 63), + ), +) +def test_invalid_known_field_is_quarantined(field: str, value: object) -> None: + line = _event("ctx_transfer_settled", 5, 1, side="ctx", outcome="completed") + record = json.loads(line.split(DIAGNOSTICS_LOG_PREFIX, 1)[1]) + record[field] = value + + result = analyze_lines([f"{DIAGNOSTICS_LOG_PREFIX}{json.dumps(record)}\n"]) + + assert result["summary"]["malformed_diagnostic_lines"] == 1 + assert result["summary"]["parsed_events"] == 0 + + +def test_missing_fanout_correlation_is_unmeasured() -> None: + result = analyze_lines( + [ + _event("ctx_transfer_queued", 6, 1, side="ctx", slice_id=0), + _event("ctx_worker_dequeued", 6, 2, side="ctx", slice_id=0), + ] + ) + + request = _request(result, 6) + assert all(duration["phase"] != "ctx_worker_queue_wait" for duration in request["durations"]) + assert { + "phase": "ctx_worker_queue_wait", + "reason": "missing_correlation_fields", + "start_count": 1, + "end_count": 1, + } in request["unmeasured_phases"] + + +def test_single_pair_invalid_metadata_is_reported_once() -> None: + result = analyze_lines( + [ + _event("gen_request_data_sent", 8, 1, host="", slice_id=0, peer_rank=1), + _event( + "gen_writer_result_received", + 8, + 2, + host="", + outcome="success", + slice_id=0, + peer_rank=1, + ), + ] + ) + + invalid = [ + phase + for phase in _request(result, 8)["unmeasured_phases"] + if phase["phase"] == "gen_writer_first_response" + and phase["reason"] == "invalid_clock_metadata" + ] + assert invalid == [ + { + "phase": "gen_writer_first_response", + "reason": "invalid_clock_metadata", + "start_count": 1, + "end_count": 1, + } + ] + + +def test_aggregate_kv_pool_snapshots_keep_headroom_fields() -> None: + result = analyze_lines( + [ + _event( + "gen_kv_pool_snapshot", + None, + 12_000_000, + init_requests=4, + transfers_in_progress=2, + transfers_complete=1, + kv_admitted_this_iteration=1, + decode_requests=8, + kv_pool_max_blocks=100, + kv_pool_free_blocks=25, + kv_pool_used_blocks=75, + index_free_slots=6, + ), + _event("gen_kv_pool_snapshot", None, 22_000_000), + _event( + "diagnostics_events_dropped", + None, + 32_000_000, + side="runtime", + dropped_events=3, + ), + ] + ) + + assert result["summary"]["request_count"] == 0 + snapshot = result["aggregate_timeline"][0] + assert snapshot["event"] == "gen_kv_pool_snapshot" + assert snapshot["init_requests"] == 4 + assert snapshot["kv_pool_free_blocks"] == 25 + assert snapshot["index_free_slots"] == 6 + assert result["aggregate_timeline"][-1]["dropped_events"] == 3 + assert result["scheduler_decision_cadence"] == [ + { + **_domain(), + "rank": None, + "snapshot_count": 2, + "interval_count": 1, + "min_ms": 10.0, + "mean_ms": 10.0, + "max_ms": 10.0, + } + ] + + +def test_gen_transfer_settled_uses_the_next_same_rank_scheduler_decision() -> None: + result = analyze_lines( + [ + _event( + "gen_transfer_settled", + 70, + 10_000_000, + rank=0, + outcome="completed", + resources_drained=True, + ), + _event("gen_kv_pool_snapshot", None, 16_000_000, rank=0), + _event( + "gen_transfer_settled", + 71, + 20_000_000, + rank=0, + outcome="completed", + resources_drained=True, + ), + _event( + "gen_transfer_settled", + 72, + 21_000_000, + rank=0, + outcome="failed", + resources_drained=False, + ), + _event("gen_kv_pool_snapshot", None, 25_000_000, rank=1), + ] + ) + + assert result["gen_transfer_settled_to_next_scheduler_decision"] == [ + { + **_domain(), + "rank": 0, + "settled_count": 2, + "matched_count": 1, + "unmatched_count": 1, + "min_ms": 6.0, + "mean_ms": 6.0, + "max_ms": 6.0, + } + ] + + +def test_boundary_completeness_is_reported_per_participant() -> None: + result = analyze_lines( + [ + _capabilities(rank=0), + _capabilities(pid=18, rank=1), + _event("gen_ingress", 31, 1_000_000, rank=0), + _event("gen_kv_admission_result", 31, 2_000_000, rank=0, outcome="admitted"), + _event("gen_ingress", 31, 3_000_000, pid=18, rank=1), + ] + ) + + request = _request(result, 31) + # A peer's event cannot satisfy rank 1's expected admission boundary. + assert request["missing_boundaries"] == ["gen_kv_admission_result"] + rank_one = next( + participant for participant in request["participants"] if participant["rank"] == 1 + ) + assert rank_one["missing_boundaries"] == ["gen_kv_admission_result"] + + +def test_kv_admission_does_not_infer_transfer_window_ownership() -> None: + result = analyze_lines( + [ + _event("gen_ingress", 33, 1_000_000, rank=0, pp_rank=0), + _event( + "gen_kv_admission_result", + 33, + 2_000_000, + rank=0, + pp_rank=0, + outcome="admitted", + ), + _event("gen_ingress", 33, 1_000_000, pid=18, rank=1, pp_rank=1), + _event( + "gen_kv_admission_result", + 33, + 2_000_000, + pid=18, + rank=1, + pp_rank=1, + outcome="admitted", + ), + ] + ) + + participants = _request(result, 33)["participants"] + rank_zero = next(participant for participant in participants if participant["rank"] == 0) + rank_one = next(participant for participant in participants if participant["rank"] == 1) + assert rank_zero["missing_boundaries"] == [] + assert rank_one["missing_boundaries"] == [] + assert all( + phase["phase"] != "gen_transfer_window_admission_wait" + for phase in _request(result, 33)["unmeasured_phases"] + ) + + +def test_gate2_retries_form_one_admission_span_without_false_missing_pairs() -> None: + result = analyze_lines( + [ + _event("gen_ingress", 32, 1_000_000), + _event("gen_kv_admission_result", 32, 2_000_000, outcome="admitted"), + _event("gen_transfer_window_result", 32, 3_000_000, outcome="deferred"), + _event("gen_kv_admission_result", 32, 5_000_000, outcome="admitted"), + _event("gen_transfer_window_result", 32, 8_000_000, outcome="admitted"), + ] + ) + + request = _request(result, 32) + durations = {duration["phase"]: duration["duration_ms"] for duration in request["durations"]} + assert durations["gen_gate1_admission_wait"] == 1.0 + assert durations["gen_transfer_window_admission_wait"] == 6.0 + assert not any( + phase["phase"] in {"gen_gate1_admission_wait", "gen_transfer_window_admission_wait"} + for phase in request["unmeasured_phases"] + ) + + +def test_cross_domain_monotonic_timestamps_are_not_subtracted() -> None: + result = analyze_lines( + [ + _event("ctx_send_ready", 9, 900_000_000, host="ctx-node", pid=11, side="ctx"), + _event( + "ctx_all_receivers_ready", + 9, + 100, + host="another-node", + pid=22, + side="ctx", + ), + ] + ) + + request = _request(result, 9) + assert all( + duration["phase"] != "ctx_receiver_readiness_offset" for duration in request["durations"] + ) + assert { + "phase": "ctx_receiver_readiness_offset", + "reason": "clock_domain_mismatch", + } in request["unmeasured_phases"] + assert request["missing_boundaries"] == [] + assert request["unassessed_boundaries"] == [ + "ctx_all_receivers_ready", + "ctx_source_kv_released", + "ctx_transfer_settled", + ] + assert [event["event"] for event in request["timeline"]] == [ + "ctx_all_receivers_ready", + "ctx_send_ready", + ] + + +def test_receiver_ready_before_final_send_is_reported_as_readiness_lead() -> None: + result = analyze_lines( + [ + _event("ctx_all_receivers_ready", 10, 3_000_000, side="ctx"), + _event("ctx_send_ready", 10, 8_000_000, side="ctx"), + ] + ) + + readiness = next( + duration + for duration in _request(result, 10)["durations"] + if duration["phase"] == "ctx_receiver_readiness_offset" + ) + assert readiness["signed_offset_ms"] == -5.0 + assert readiness["readiness_lead_ms"] == 5.0 + assert readiness["readiness_wait_ms"] == 0.0 + + +def test_fanout_boundaries_are_correlated_by_slice_and_peer() -> None: + result = analyze_lines( + [ + _event("gen_request_data_sent", 41, 1_000_000, slice_id=0, peer_rank=0), + _event("gen_request_data_sent", 41, 2_000_000, slice_id=0, peer_rank=1), + _event("gen_writer_result_received", 41, 7_000_000, slice_id=0, peer_rank=1), + _event("gen_writer_result_received", 41, 4_000_000, slice_id=0, peer_rank=0), + ] + ) + + durations = [ + duration + for duration in _request(result, 41)["durations"] + if duration["phase"] == "gen_writer_first_response" + ] + by_peer = { + duration["correlation"]["peer_rank"]: duration["duration_ms"] for duration in durations + } + assert by_peer == {0: 3.0, 1: 5.0} + + +def test_writer_timeline_preserves_late_session_evidence() -> None: + result = analyze_lines( + [ + _event( + "gen_writer_result_received", + 43, + 4_000_000, + slice_id=0, + peer_rank=2, + session_found=False, + ), + ] + ) + + timeline = _request(result, 43)["timeline"] + assert timeline[0]["session_found"] is False + + +def test_writer_first_response_reports_an_unmatched_fanout_peer() -> None: + result = analyze_lines( + [ + _capabilities(), + _event("gen_request_data_sent", 42, 1_000_000, slice_id=0, peer_rank=0), + _event("gen_request_data_sent", 42, 2_000_000, slice_id=0, peer_rank=1), + _event("gen_writer_result_received", 42, 4_000_000, slice_id=0, peer_rank=0), + ] + ) + + request = _request(result, 42) + assert any( + duration["phase"] == "gen_writer_first_response" + and duration["correlation"]["peer_rank"] == 0 + for duration in request["durations"] + ) + assert { + "phase": "gen_writer_first_response", + "reason": "missing_end", + "count": 1, + "clock_domain": _domain(), + "correlation": {"slice_id": 0, "peer_rank": 1}, + } in request["unmeasured_phases"] + + +def _broadcast_send(peer_rank: int, *, expected_writers: int | None = 1, **details: object) -> str: + return _event( + "gen_request_data_sent", + 58, + 1_000_000, + slice_id=0, + peer_rank=peer_rank, + writer_cohort_known=False, + expected_writers=expected_writers, + **details, + ) + + +def _broadcast_response(peer_rank: int, **details: object) -> str: + return _event( + "gen_writer_result_received", + 58, + 3_000_000, + peer_rank=peer_rank, + **{"slice_id": 0, "outcome": "success", "is_last_slice": True, **details}, + ) + + +def _writer_gaps(result: dict[str, object]) -> list[dict[str, object]]: + return [ + gap + for gap in _request(result, 58)["unmeasured_phases"] + if gap["phase"] == "gen_writer_first_response" + ] + + +def test_complete_adp_broadcast_does_not_require_unselected_peers_to_respond() -> None: + result = analyze_lines( + [ + _capabilities(), + _broadcast_send(0), + _broadcast_send(1), + _broadcast_response(1), + _event("gen_destination_complete", 58, 4_000_000, slice_id=0, peer_rank=1), + ] + ) + + assert _writer_gaps(result) == [] + duration = next( + sample + for sample in _request(result, 58)["durations"] + if sample["phase"] == "gen_writer_first_response" + ) + assert duration["duration_ms"] == 2.0 + assert duration["correlation"] == {"slice_id": 0, "peer_rank": 1} + assert _request(result, 58)["timeline"][0]["writer_cohort_known"] is False + + +def test_incomplete_adp_broadcast_reports_missing_writer_count_not_candidate_peers() -> None: + result = analyze_lines( + [ + _broadcast_send(0, expected_writers=2), + _broadcast_send(1, expected_writers=2), + _broadcast_send(2, expected_writers=2), + _broadcast_response(1, is_last_slice=False), + _broadcast_response(1), + _broadcast_response(2, session_found=False), + ] + ) + + assert _writer_gaps(result) == [ + { + "phase": "gen_writer_first_response", + "reason": "missing_writer_responses", + "clock_domain": _domain(), + "correlation": {"slice_id": 0}, + "observed_writers": 1, + "expected_writers": 2, + "count": 1, + } + ] + + +@pytest.mark.parametrize("counts", ((None, None), (0, 0), (1, 2), (None, 1))) +def test_ambiguous_broadcast_writer_count_is_unknown(counts: tuple[int | None, int | None]) -> None: + result = analyze_lines( + [ + _broadcast_send(0, expected_writers=counts[0]), + _broadcast_send(1, expected_writers=counts[1]), + _broadcast_response(1), + ] + ) + + gaps = _writer_gaps(result) + assert len(gaps) == 1 + assert gaps[0]["reason"] == "unknown_writer_cohort" + assert gaps[0]["detail"] == "missing_invalid_or_conflicting_expected_writers" + + +@pytest.mark.parametrize("cohort_flag", (None, True)) +def test_missing_or_conflicting_writer_cohort_does_not_assume_exact_peers( + cohort_flag: bool | None, +) -> None: + result = analyze_lines( + [ + _broadcast_send(0), + _event( + "gen_request_data_sent", + 58, + 1_000_000, + slice_id=0, + peer_rank=1, + writer_cohort_known=cohort_flag, + expected_writers=1, + ), + _broadcast_response(1), + ] + ) + + assert [gap["reason"] for gap in _writer_gaps(result)] == ["unknown_writer_cohort"] + + +@pytest.mark.parametrize( + "other_identity", + ( + {"host": "node-b"}, + {"pid": 18}, + {"rank": 1}, + {"slice_id": 1}, + {"process_uuid": "22222222-2222-4222-8222-222222222222"}, + ), +) +def test_other_participant_or_slice_cannot_complete_broadcast_cohort( + other_identity: dict[str, object], +) -> None: + result = analyze_lines( + [ + _broadcast_send(0), + _broadcast_response(0, **other_identity), + ] + ) + + missing = [gap for gap in _writer_gaps(result) if gap["reason"] == "missing_writer_responses"] + assert len(missing) == 1 + assert missing[0]["observed_writers"] == 0 + assert missing[0]["expected_writers"] == 1 + + +def test_duplicate_broadcast_publications_and_chunk_results_count_distinct_writers() -> None: + result = analyze_lines( + [ + _broadcast_send(0), + _broadcast_send(0), + _broadcast_send(1), + _broadcast_response(1, is_last_slice=False), + _broadcast_response(1), + _broadcast_response(1), + ] + ) + + assert _writer_gaps(result) == [] + durations = [ + sample + for sample in _request(result, 58)["durations"] + if sample["phase"] == "gen_writer_first_response" + ] + assert len(durations) == 1 + + +def test_broadcast_does_not_hide_response_that_precedes_publication() -> None: + result = analyze_lines( + [ + _capabilities(), + _broadcast_send(0, expected_writers=2), + _broadcast_send(1, expected_writers=2), + _event("gen_writer_result_received", 58, 0, slice_id=0, peer_rank=0), + _broadcast_response(1), + ] + ) + + gaps = _writer_gaps(result) + assert len(gaps) == 1 + assert gaps[0]["reason"] == "missing_end" + assert gaps[0]["correlation"] == {"slice_id": 0, "peer_rank": 0} + + +def test_pipelined_chunks_use_first_response_and_final_successful_destination() -> None: + result = analyze_lines( + [ + _event("gen_request_data_sent", 55, 1_000_000, slice_id=0, peer_rank=3), + _event("gen_request_data_sent", 55, 1_500_000, slice_id=0, peer_rank=4), + _event( + "gen_writer_result_received", + 55, + 2_000_000, + slice_id=0, + peer_rank=3, + outcome="success", + is_last_slice=False, + ), + _event( + "gen_writer_result_received", + 55, + 5_000_000, + slice_id=0, + peer_rank=4, + outcome="success", + is_last_slice=True, + ), + _event( + "gen_writer_result_received", + 55, + 6_000_000, + slice_id=0, + peer_rank=3, + outcome="success", + is_last_slice=True, + ), + _event( + "gen_destination_complete", + 55, + 9_000_000, + slice_id=0, + peer_rank=3, + outcome="completed", + ), + ] + ) + + request = _request(result, 55) + writer_durations = { + duration["correlation"]["peer_rank"]: duration["duration_ms"] + for duration in request["durations"] + if duration["phase"] == "gen_writer_first_response" + } + destination_duration = next( + duration + for duration in request["durations"] + if duration["phase"] == "gen_destination_drain" + ) + assert writer_durations == {3: 1.0, 4: 3.5} + assert destination_duration["duration_ms"] == 3.0 + assert not any( + phase["phase"] in {"gen_writer_first_response", "gen_destination_drain"} + for phase in request["unmeasured_phases"] + ) + + +def test_failed_writer_result_does_not_expect_destination_completion() -> None: + result = analyze_lines( + [ + _event("gen_request_data_sent", 56, 1_000_000, slice_id=0, peer_rank=3), + _event( + "gen_writer_result_received", + 56, + 4_000_000, + slice_id=0, + peer_rank=3, + outcome="failed", + is_last_slice=True, + ), + ] + ) + + request = _request(result, 56) + assert "gen_destination_complete" not in request["missing_boundaries"] + assert all(duration["phase"] != "gen_destination_drain" for duration in request["durations"]) + + +def test_mixed_writer_results_do_not_expect_destination_completion() -> None: + result = analyze_lines( + [ + _event("gen_request_data_sent", 57, 1_000_000, slice_id=0, peer_rank=3), + _event("gen_request_data_sent", 57, 1_500_000, slice_id=0, peer_rank=4), + _event( + "gen_writer_result_received", + 57, + 4_000_000, + slice_id=0, + peer_rank=3, + outcome="success", + is_last_slice=True, + ), + _event( + "gen_writer_result_received", + 57, + 5_000_000, + slice_id=0, + peer_rank=4, + outcome="failed", + is_last_slice=True, + ), + ] + ) + + request = _request(result, 57) + assert "gen_destination_complete" not in request["missing_boundaries"] + assert all(duration["phase"] != "gen_destination_drain" for duration in request["durations"]) + + +def test_timeout_phases_use_the_endpoint_local_monotonic_clock() -> None: + result = analyze_lines( + [ + _event( + "ctx_send_ready", + 77, + 1_000_000, + side="ctx", + timeout_expected=True, + ), + _event( + "transfer_timeout_started", + 77, + 2_000_000, + side="ctx", + timer_start_monotonic_ns=1_500_000, + timeout_ms=60, + timeout_owner="pyexecutor", + state="DISAGG_CONTEXT_TRANS_IN_PROGRESS", + ), + _event( + "transfer_timeout_observed", + 77, + 62_000_000, + side="ctx", + timeout_owner="pyexecutor", + ), + ] + ) + + durations = { + duration["phase"]: duration["duration_ms"] for duration in _request(result, 77)["durations"] + } + assert durations["ctx_send_to_timeout_start"] == 0.5 + assert durations["transfer_timeout_window"] == 60.5 + timeout_start = next( + event + for event in _request(result, 77)["timeline"] + if event["event"] == "transfer_timeout_started" + ) + assert timeout_start["timer_start_monotonic_ns"] == 1_500_000 + assert timeout_start["timeout_ms"] == 60 + assert timeout_start["timeout_owner"] == "pyexecutor" + assert timeout_start["state"] == "DISAGG_CONTEXT_TRANS_IN_PROGRESS" + + +def test_healthy_or_inapplicable_timeout_phases_are_not_reported_missing() -> None: + result = analyze_lines( + [ + _event( + "ctx_send_ready", + 78, + 1_000_000, + side="ctx", + timeout_expected=True, + ), + _event( + "transfer_timeout_started", + 78, + 2_000_000, + side="ctx", + timeout_owner="pyexecutor", + ), + _event( + "gen_receive_start", + 79, + 3_000_000, + timeout_expected=False, + ), + _event( + "ctx_send_ready", + 80, + 4_000_000, + side="ctx", + timeout_expected=False, + ), + ] + ) + + healthy = _request(result, 78) + assert all( + phase["phase"] != "transfer_timeout_window" for phase in healthy["unmeasured_phases"] + ) + sync_gen = _request(result, 79) + assert all( + phase["phase"] != "gen_receive_to_timeout_start" for phase in sync_gen["unmeasured_phases"] + ) + timeout_disabled_ctx = _request(result, 80) + assert all( + phase["phase"] != "ctx_send_to_timeout_start" + for phase in timeout_disabled_ctx["unmeasured_phases"] + ) + + +def test_ctx_worker_queue_wait_is_correlated_per_slice_and_peer() -> None: + result = analyze_lines( + [ + _capabilities(), + _event( + "ctx_transfer_queued", + 88, + 2_000_000, + side="ctx", + slice_id=1, + peer_rank=0, + receiver_slice_id=2, + is_last_slice=True, + ), + _event( + "ctx_transfer_queued", + 88, + 3_000_000, + side="ctx", + slice_id=1, + peer_rank=1, + receiver_slice_id=2, + is_last_slice=True, + ), + _event( + "ctx_worker_dequeued", + 88, + 7_000_000, + side="ctx", + slice_id=1, + peer_rank=0, + ), + _event( + "ctx_backend_submit_start", + 88, + 11_000_000, + side="ctx", + slice_id=1, + peer_rank=0, + ), + _event( + "ctx_backend_submitted", + 88, + 12_000_000, + side="ctx", + slice_id=1, + peer_rank=0, + ), + _event( + "ctx_backend_complete", + 88, + 20_000_000, + side="ctx", + slice_id=1, + peer_rank=0, + ), + ] + ) + + request = _request(result, 88) + queue_wait = next( + duration + for duration in request["durations"] + if duration["phase"] == "ctx_worker_queue_wait" + ) + assert queue_wait["duration_ms"] == 5.0 + assert queue_wait["correlation"] == {"slice_id": 1, "peer_rank": 0} + preparation = next( + duration + for duration in request["durations"] + if duration["phase"] == "ctx_worker_preparation" + ) + assert preparation["duration_ms"] == 4.0 + assert preparation["correlation"] == {"slice_id": 1, "peer_rank": 0} + assert { + "phase": "ctx_worker_queue_wait", + "reason": "missing_end", + "count": 1, + "clock_domain": _domain(), + "correlation": {"slice_id": 1, "peer_rank": 1}, + } in request["unmeasured_phases"] + queued = request["timeline"][0] + assert queued["receiver_slice_id"] == 2 + assert queued["is_last_slice"] is True + + +def test_cli_reads_files_and_emits_json(tmp_path: Path, capfd: pytest.CaptureFixture[str]) -> None: + log = tmp_path / "worker.log" + log.write_text(_event("gen_ingress", 5, 100), encoding="utf-8") + + assert main([str(log), "--indent", "0"]) == 0 + + result = json.loads(capfd.readouterr().out) + assert result["summary"]["request_count"] == 1 + assert result["requests"][0]["request_id"] == 5 diff --git a/tests/unittest/tools/test_disagg_transfer_diagnostics_capabilities.py b/tests/unittest/tools/test_disagg_transfer_diagnostics_capabilities.py new file mode 100644 index 000000000000..dc8d2ac0165d --- /dev/null +++ b/tests/unittest/tools/test_disagg_transfer_diagnostics_capabilities.py @@ -0,0 +1,369 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Runtime-capability regressions for the transfer diagnostic analyzer.""" + +import json +from uuid import NAMESPACE_DNS, uuid5 + +import pytest + +__extra_import_path__ = ["~/scripts"] +from disagg_transfer_diagnostics import DIAGNOSTICS_LOG_PREFIX, analyze_lines # noqa: E402 + +pytestmark = pytest.mark.cpu_only + + +def _event( + name: str, + timestamp: int = 1_000_000, + *, + request_id: int | None = 7, + side: str = "ctx", + host: str = "node-a", + pid: int = 17, + rank: int = 0, + **details: object, +) -> str: + record = { + "schema_version": 1, + "event": name, + "request_id": request_id, + "side": side, + "host": host, + "pid": pid, + "rank": rank, + "run_uuid": "11111111-1111-4111-8111-111111111111", + "process_uuid": str(uuid5(NAMESPACE_DNS, f"{host}:{pid}")), + "monotonic_ns": timestamp, + "wall_ns": 100_000_000_000 + timestamp, + **details, + } + return f"{DIAGNOSTICS_LOG_PREFIX}{json.dumps(record)}\n" + + +def _capability_record(**details: object) -> str: + return _event( + "diagnostic_capabilities", + 0, + request_id=None, + side="runtime", + **{"capability_schema_version": 1, **details}, + ) + + +def _profile(runtime: str = "CPP", *, scheduler_v2: bool = True, **identity: object) -> list[str]: + return [ + _capability_record( + transceiver_runtime=runtime, + python_transfer_events=runtime == "PYTHON", + **identity, + ), + _capability_record( + executor_events=True, + scheduler_kv_admission_events=scheduler_v2, + **identity, + ), + ] + + +def _request(lines: list[str]) -> dict[str, object]: + result = analyze_lines(lines) + assert len(result["requests"]) == 1 + return result["requests"][0] + + +def _phase_reasons(request: dict[str, object], phase: str) -> set[str]: + return {gap["reason"] for gap in request["unmeasured_phases"] if gap["phase"] == phase} + + +def test_healthy_cpp_trace_reports_unsupported_python_boundaries_not_missing() -> None: + request = _request( + _profile() + + [ + _event("ctx_send_ready"), + _event("ctx_source_kv_released", 3_000_000), + _event("ctx_source_unpinned", 4_000_000), + _event("gen_ingress", side="gen"), + _event("gen_kv_admission_result", 2_000_000, side="gen", outcome="admitted"), + _event("gen_transfer_window_result", 3_000_000, side="gen", outcome="admitted"), + _event("gen_decode_ready", 4_000_000, side="gen"), + ] + ) + + assert request["missing_boundaries"] == [] + assert request["unassessed_boundaries"] == [] + assert request["unsupported_boundaries"] == [ + "ctx_all_receivers_ready", + "ctx_transfer_settled", + "gen_receive_start", + ] + for participant in request["participants"]: + assert participant["capabilities"] == { + "status": "known", + "transceiver_runtime": "CPP", + "python_transfer_events": False, + "executor_events": True, + "scheduler_kv_admission_events": True, + "issues": [], + } + assert participant["missing_boundaries"] == [] + assert _phase_reasons(request, "ctx_transfer_lifetime") == {"unsupported_capability"} + assert _phase_reasons(request, "gen_transfer_to_service") == {"unsupported_capability"} + assert any( + duration["phase"] == "ctx_source_kv_request_ownership" and duration["duration_ms"] == 2.0 + for duration in request["durations"] + ) + + +def test_cpp_runtime_still_reports_missing_shared_executor_boundary() -> None: + request = _request(_profile() + [_event("ctx_send_ready")]) + + assert request["missing_boundaries"] == ["ctx_source_kv_released"] + assert request["participants"][0]["missing_boundaries"] == ["ctx_source_kv_released"] + assert _phase_reasons(request, "ctx_source_kv_request_ownership") == {"missing_end"} + assert _phase_reasons(request, "ctx_receiver_readiness_offset") == {"unsupported_capability"} + + +def test_incomplete_python_trace_keeps_missing_transfer_boundaries() -> None: + request = _request( + _profile("PYTHON") + [_event("ctx_send_ready"), _event("ctx_source_kv_released", 3_000_000)] + ) + + assert request["missing_boundaries"] == ["ctx_all_receivers_ready", "ctx_transfer_settled"] + assert request["unsupported_boundaries"] == [] + assert request["unassessed_boundaries"] == [] + assert _phase_reasons(request, "ctx_transfer_lifetime") == {"missing_end"} + + +def test_scheduler_v1_support_is_independent_of_python_transfer_support() -> None: + request = _request( + _profile("PYTHON", scheduler_v2=False) + + [ + _event("gen_ingress", side="gen"), + _event("gen_transfer_window_result", 2_000_000, side="gen", outcome="admitted"), + ] + ) + + assert request["missing_boundaries"] == ["gen_receive_start"] + assert request["unsupported_boundaries"] == ["gen_kv_admission_result"] + assert request["unassessed_boundaries"] == [] + assert _phase_reasons(request, "gen_gate1_admission_wait") == {"unsupported_capability"} + assert request["participants"][0]["capabilities"]["scheduler_kv_admission_events"] is False + + +@pytest.mark.parametrize("other_identity", ({"host": "node-b"}, {"pid": 18}, {"rank": 1})) +def test_capabilities_do_not_leak_across_participants(other_identity: dict[str, object]) -> None: + request = _request( + _profile("CPP") + + _profile("PYTHON", **other_identity) + + [ + _event("gen_transfer_window_result", side="gen", outcome="admitted"), + _event("gen_transfer_window_result", side="gen", outcome="admitted", **other_identity), + ] + ) + + participants = { + participant["capabilities"]["transceiver_runtime"]: participant + for participant in request["participants"] + } + assert len(participants) == 2 + assert participants["CPP"]["missing_boundaries"] == [] + assert participants["CPP"]["unsupported_boundaries"] == ["gen_receive_start"] + assert participants["PYTHON"]["missing_boundaries"] == ["gen_receive_start"] + assert participants["PYTHON"]["unsupported_boundaries"] == [] + assert request["missing_boundaries"] == ["gen_receive_start"] + + +def test_absent_metadata_is_unknown_not_implicitly_cpp() -> None: + request = _request([_event("ctx_send_ready")]) + + capabilities = request["participants"][0]["capabilities"] + assert capabilities["status"] == "unknown" + assert capabilities["transceiver_runtime"] is None + assert capabilities["python_transfer_events"] is None + assert capabilities["executor_events"] is None + assert request["missing_boundaries"] == [] + assert request["unsupported_boundaries"] == [] + assert request["unassessed_boundaries"] == [ + "ctx_all_receivers_ready", + "ctx_source_kv_released", + "ctx_transfer_settled", + ] + assert _phase_reasons(request, "ctx_transfer_lifetime") == {"unknown_capability"} + + +@pytest.mark.parametrize( + "invalid_fields", + ( + {"capability_schema_version": 2}, + {"capability_schema_version": True}, + {"python_transfer_events": 0}, + {"python_transfer_events": "false"}, + {"transceiver_runtime": "AUTO"}, + ), +) +def test_invalid_factory_metadata_does_not_silence_missing_python_edges( + invalid_fields: dict[str, object], +) -> None: + declaration = { + "transceiver_runtime": "CPP", + "python_transfer_events": False, + **invalid_fields, + } + request = _request([_capability_record(**declaration), _event("ctx_send_ready")]) + + capabilities = request["participants"][0]["capabilities"] + assert capabilities["status"] != "known" + assert capabilities["python_transfer_events"] is None + assert capabilities["issues"] + assert request["missing_boundaries"] == [] + assert request["unsupported_boundaries"] == [] + assert "ctx_transfer_settled" in request["unassessed_boundaries"] + assert _phase_reasons(request, "ctx_transfer_lifetime") == {"unknown_capability"} + + +def test_conflicting_runtime_declarations_are_not_resolved_by_input_order() -> None: + profiles = _profile("CPP") + _profile("PYTHON") + for declarations in (profiles, list(reversed(profiles))): + request = _request(declarations + [_event("ctx_send_ready")]) + + capabilities = request["participants"][0]["capabilities"] + assert capabilities["status"] == "partial" + assert capabilities["transceiver_runtime"] is None + assert capabilities["python_transfer_events"] is None + assert capabilities["executor_events"] is True + assert capabilities["issues"] + assert request["missing_boundaries"] == ["ctx_source_kv_released"] + assert request["unsupported_boundaries"] == [] + assert request["unassessed_boundaries"] == [ + "ctx_all_receivers_ready", + "ctx_transfer_settled", + ] + + +def test_executor_only_metadata_preserves_shared_checks_without_assuming_runtime() -> None: + request = _request( + [ + _capability_record(executor_events=True, scheduler_kv_admission_events=True), + _event("ctx_send_ready"), + ] + ) + + capabilities = request["participants"][0]["capabilities"] + assert capabilities["status"] == "partial" + assert capabilities["executor_events"] is True + assert capabilities["python_transfer_events"] is None + assert request["missing_boundaries"] == ["ctx_source_kv_released"] + assert request["unassessed_boundaries"] == [ + "ctx_all_receivers_ready", + "ctx_transfer_settled", + ] + + +def test_factory_only_metadata_does_not_imply_executor_instrumentation() -> None: + request = _request( + [ + _capability_record(transceiver_runtime="CPP", python_transfer_events=False), + _event("ctx_send_ready"), + ] + ) + + capabilities = request["participants"][0]["capabilities"] + assert capabilities["status"] == "partial" + assert capabilities["python_transfer_events"] is False + assert capabilities["executor_events"] is None + assert request["missing_boundaries"] == [] + assert request["unsupported_boundaries"] == ["ctx_all_receivers_ready", "ctx_transfer_settled"] + assert request["unassessed_boundaries"] == ["ctx_source_kv_released"] + assert _phase_reasons(request, "ctx_source_kv_request_ownership") == {"unknown_capability"} + + +def test_observed_pairs_still_measure_durations_without_capability_metadata() -> None: + request = _request( + [ + _event("ctx_send_ready"), + _event("ctx_all_receivers_ready", 2_000_000), + _event("ctx_source_kv_released", 3_000_000), + _event("ctx_transfer_settled", 5_000_000, outcome="completed"), + ] + ) + + assert request["participants"][0]["capabilities"]["status"] == "unknown" + assert request["missing_boundaries"] == [] + assert request["unassessed_boundaries"] == [] + assert request["unmeasured_phases"] == [] + durations = {duration["phase"]: duration for duration in request["durations"]} + assert durations["ctx_transfer_lifetime"]["duration_ms"] == 4.0 + assert durations["ctx_source_kv_request_ownership"]["duration_ms"] == 2.0 + assert durations["ctx_receiver_readiness_offset"]["signed_offset_ms"] == 1.0 + + +def test_observed_python_event_overrides_unsupported_claim_without_hiding_gaps() -> None: + request = _request( + _profile("CPP") + + [ + _event("ctx_send_ready"), + _event("ctx_all_receivers_ready", 2_000_000), + _event("ctx_source_kv_released", 3_000_000), + ] + ) + + capabilities = request["participants"][0]["capabilities"] + assert capabilities["status"] == "partial" + assert capabilities["python_transfer_events"] is None + assert "observed_unsupported_python_transfer_events" in capabilities["issues"] + assert capabilities["executor_events"] is True + assert request["missing_boundaries"] == [] + assert request["unsupported_boundaries"] == [] + assert request["unassessed_boundaries"] == ["ctx_transfer_settled"] + assert _phase_reasons(request, "ctx_transfer_lifetime") == {"unknown_capability"} + durations = {duration["phase"]: duration for duration in request["durations"]} + assert durations["ctx_receiver_readiness_offset"]["signed_offset_ms"] == 1.0 + assert durations["ctx_source_kv_request_ownership"]["duration_ms"] == 2.0 + + +def test_boundaries_on_different_ranks_in_one_process_are_not_paired() -> None: + request = _request( + _profile("CPP", rank=0) + + _profile("CPP", rank=1) + + [ + _event("ctx_send_ready", rank=0), + _event("ctx_source_kv_released", 3_000_000, rank=1), + ] + ) + + assert not any( + duration["phase"] == "ctx_source_kv_request_ownership" for duration in request["durations"] + ) + gaps = { + (gap["rank"], gap["reason"]) + for gap in request["unmeasured_phases"] + if gap["phase"] == "ctx_source_kv_request_ownership" + } + assert gaps == {(0, "missing_end"), (1, "missing_start")} + assert request["missing_boundaries"] == ["ctx_source_kv_released"] + + +@pytest.mark.parametrize("runtime,python_events", (("CPP", True), ("PYTHON", False))) +def test_runtime_flag_mismatch_leaves_python_support_unknown( + runtime: str, python_events: bool +) -> None: + request = _request( + [ + _capability_record(transceiver_runtime=runtime, python_transfer_events=python_events), + _capability_record(executor_events=True, scheduler_kv_admission_events=True), + _event("ctx_send_ready"), + ] + ) + + capabilities = request["participants"][0]["capabilities"] + assert capabilities["status"] == "partial" + assert capabilities["python_transfer_events"] is None + assert "runtime_capability_mismatch" in capabilities["issues"] + assert capabilities["executor_events"] is True + assert request["missing_boundaries"] == ["ctx_source_kv_released"] + assert request["unsupported_boundaries"] == [] + assert request["unassessed_boundaries"] == [ + "ctx_all_receivers_ready", + "ctx_transfer_settled", + ] diff --git a/tests/unittest/tools/test_disagg_transfer_diagnostics_identity.py b/tests/unittest/tools/test_disagg_transfer_diagnostics_identity.py new file mode 100644 index 000000000000..02da9c6bd01b --- /dev/null +++ b/tests/unittest/tools/test_disagg_transfer_diagnostics_identity.py @@ -0,0 +1,273 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Run and process identity regressions for the transfer diagnostic analyzer.""" + +import json + +import pytest + +__extra_import_path__ = ["~/scripts"] +from disagg_transfer_diagnostics import DIAGNOSTICS_LOG_PREFIX, analyze_lines # noqa: E402 + +pytestmark = pytest.mark.cpu_only + +_RUN_A = "205d1c66-cc2a-4bb9-a9c4-1559c6f95fd8" +_RUN_B = "e199df92-446e-4a55-847a-18078be332d8" +_PROCESS_A = "6d905e36-6967-4db1-b44b-e8ab47cb3053" +_PROCESS_B = "4f88878a-b7bf-4c30-ad73-b14615738452" + + +def _event( + name: str, + timestamp: int = 1_000_000, + *, + request_id: int | str | None = 7, + side: str = "ctx", + **details: object, +) -> str: + record = { + "schema_version": 1, + "event": name, + "request_id": request_id, + "side": side, + "host": "node-a", + "pid": 17, + "rank": 0, + "run_uuid": _RUN_A, + "process_uuid": _PROCESS_A, + "monotonic_ns": timestamp, + "wall_ns": 100_000_000_000 + timestamp, + **details, + } + return f"{DIAGNOSTICS_LOG_PREFIX}{json.dumps(record)}\n" + + +def _ownership_pair(**identity: object) -> list[str]: + return [ + _event("ctx_send_ready", **identity), + _event("ctx_source_kv_released", 3_000_000, **identity), + ] + + +def _ownership_durations(request: dict[str, object]) -> list[float]: + return [ + duration["duration_ms"] + for duration in request["durations"] + if duration["phase"] == "ctx_source_kv_request_ownership" + ] + + +def test_different_runs_cannot_pair_reused_request_ids_and_pids() -> None: + result = analyze_lines( + [ + _event("ctx_send_ready"), + _event("ctx_source_kv_released", 61_000_000, run_uuid=_RUN_B), + ] + ) + + assert len(result["requests"]) == 2 + assert {request["run_uuid"] for request in result["requests"]} == {_RUN_A, _RUN_B} + assert all(request["correlation_scope"] == "run" for request in result["requests"]) + assert all(request["durations"] == [] for request in result["requests"]) + assert "ctx_source_kv_request_ownership" not in result["phase_durations"] + + +def test_shared_run_correlates_ctx_and_gen_without_cross_process_timing() -> None: + result = analyze_lines( + _ownership_pair() + + [ + _event( + "gen_transfer_settled", + side="gen", + process_uuid=_PROCESS_B, + pid=18, + outcome="completed", + ), + _event("gen_decode_ready", 4_000_000, side="gen", process_uuid=_PROCESS_B, pid=18), + ] + ) + + assert len(result["requests"]) == 1 + request = result["requests"][0] + assert request["run_uuid"] == _RUN_A + assert request["correlation_scope"] == "run" + assert {participant["side"] for participant in request["participants"]} == {"ctx", "gen"} + assert {participant["process_uuid"] for participant in request["participants"]} == { + _PROCESS_A, + _PROCESS_B, + } + assert all(participant["run_uuid"] == _RUN_A for participant in request["participants"]) + assert _ownership_durations(request) == [2.0] + assert any( + duration["phase"] == "gen_transfer_to_service" and duration["duration_ms"] == 3.0 + for duration in request["durations"] + ) + assert {duration["clock_domain"]["process_uuid"] for duration in request["durations"]} == { + _PROCESS_A, + _PROCESS_B, + } + + +def test_restart_with_reused_pid_cannot_pair_local_boundaries() -> None: + result = analyze_lines( + [ + _event("ctx_send_ready"), + _event("ctx_source_kv_released", 61_000_000, process_uuid=_PROCESS_B), + ] + ) + + assert len(result["requests"]) == 1 + request = result["requests"][0] + assert request["durations"] == [] + assert len(request["participants"]) == 2 + assert {domain["process_uuid"] for domain in request["clock_domains"]} == { + _PROCESS_A, + _PROCESS_B, + } + assert all(domain["run_uuid"] == _RUN_A for domain in request["clock_domains"]) + + +@pytest.mark.parametrize("new_identity", ({"run_uuid": _RUN_B}, {"process_uuid": _PROCESS_B})) +def test_capability_metadata_does_not_leak_across_runs_or_restarts( + new_identity: dict[str, str], +) -> None: + result = analyze_lines( + [ + _event( + "diagnostic_capabilities", + 0, + request_id=None, + side="runtime", + capability_schema_version=1, + transceiver_runtime="CPP", + executor_events=True, + scheduler_kv_admission_events=True, + python_transfer_events=False, + ), + _event("ctx_send_ready", **new_identity), + ] + ) + + request = result["requests"][0] + assert request["participants"][0]["capabilities"]["status"] == "unknown" + assert request["unsupported_boundaries"] == [] + assert "ctx_transfer_settled" in request["unassessed_boundaries"] + + +@pytest.mark.parametrize("new_identity", ({"run_uuid": _RUN_B}, {"process_uuid": _PROCESS_B})) +def test_scheduler_intervals_do_not_bridge_runs_or_restarts( + new_identity: dict[str, str], +) -> None: + result = analyze_lines( + [ + _event("gen_kv_pool_snapshot", request_id=None, side="gen"), + _event("gen_transfer_settled", 2_000_000, side="gen", resources_drained=True), + _event("gen_kv_pool_snapshot", 3_000_000, request_id=None, side="gen", **new_identity), + ] + ) + + cadence = result["scheduler_decision_cadence"] + assert len(cadence) == 2 + assert all(item["snapshot_count"] == 1 and item["interval_count"] == 0 for item in cadence) + settled = result["gen_transfer_settled_to_next_scheduler_decision"] + assert len(settled) == 1 + assert settled[0]["settled_count"] == 1 + assert settled[0]["matched_count"] == 0 + assert settled[0]["unmatched_count"] == 1 + + +@pytest.mark.parametrize("run_uuid", (None, "", "not-a-run-uuid")) +def test_missing_or_invalid_shared_run_allows_only_process_local_analysis( + run_uuid: str | None, +) -> None: + result = analyze_lines( + _ownership_pair(run_uuid=run_uuid) + + _ownership_pair(run_uuid=run_uuid, process_uuid=_PROCESS_B) + ) + + assert len(result["requests"]) == 2 + for request in result["requests"]: + assert request["run_uuid"] is None + assert request["correlation_scope"] == "process" + assert len(request["participants"]) == 1 + assert _ownership_durations(request) == [2.0] + + +@pytest.mark.parametrize("process_uuid", (None, "", "not-a-process-uuid")) +def test_invalid_process_identity_keeps_events_readable_without_joining( + process_uuid: str | None, +) -> None: + result = analyze_lines(_ownership_pair(process_uuid=process_uuid)) + + assert result["summary"]["parsed_events"] == 2 + assert len(result["requests"]) == 2 + for request in result["requests"]: + assert request["correlation_scope"] == "unverified" + assert len(request["timeline"]) == 1 + assert request["durations"] == [] + + +def test_legacy_records_without_identity_do_not_derive_timing() -> None: + lines = [] + for line in _ownership_pair(): + record = json.loads(line.removeprefix(DIAGNOSTICS_LOG_PREFIX)) + record.pop("run_uuid") + record.pop("process_uuid") + lines.append(f"{DIAGNOSTICS_LOG_PREFIX}{json.dumps(record)}\n") + + result = analyze_lines(lines) + + assert len(result["requests"]) == 2 + assert all(request["correlation_scope"] == "unverified" for request in result["requests"]) + assert all(request["durations"] == [] for request in result["requests"]) + + +def test_unverified_process_identity_does_not_derive_scheduler_aggregates() -> None: + result = analyze_lines( + [ + _event("gen_kv_pool_snapshot", request_id=None, side="gen", process_uuid=None), + _event( + "gen_transfer_settled", + 2_000_000, + side="gen", + process_uuid=None, + resources_drained=True, + ), + _event( + "gen_kv_pool_snapshot", + 3_000_000, + request_id=None, + side="gen", + process_uuid=None, + ), + ] + ) + + assert result["scheduler_decision_cadence"] == [] + assert result["gen_transfer_settled_to_next_scheduler_decision"] == [] + assert len(result["aggregate_timeline"]) == 2 + + +def test_uuid_spelling_is_normalized_for_matching() -> None: + result = analyze_lines( + [ + _event("ctx_send_ready", run_uuid=_RUN_A.upper(), process_uuid=_PROCESS_A.upper()), + _event("ctx_source_kv_released", 3_000_000), + ] + ) + + assert len(result["requests"]) == 1 + request = result["requests"][0] + assert request["run_uuid"] == _RUN_A + assert request["participants"][0]["process_uuid"] == _PROCESS_A + assert _ownership_durations(request) == [2.0] + + +def test_invalid_string_request_id_is_not_coerced_into_valid_run_group() -> None: + result = analyze_lines(_ownership_pair(request_id=7) + _ownership_pair(request_id="7")) + + assert result["summary"]["malformed_diagnostic_lines"] == 2 + assert len(result["requests"]) == 1 + assert result["requests"][0]["request_id"] == 7 + assert result["requests"][0]["event_count"] == 2 + assert _ownership_durations(result["requests"][0]) == [2.0]