From e4e812222b6e98716087ee8b0950c4b5b08676e2 Mon Sep 17 00:00:00 2001 From: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> Date: Mon, 28 Sep 2026 13:20:56 -0700 Subject: [PATCH 1/3] [None][fix] Bound remote MPI task failure and launcher shutdown Signed-off-by: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> --- tensorrt_llm/llmapi/mgmn_leader_node.py | 15 +- tensorrt_llm/llmapi/mpi_session.py | 274 ++++++++++-------- tensorrt_llm/llmapi/trtllm-llmapi-launch | 160 +++++++--- .../integration/test_lists/test-db/l0_cpu.yml | 3 + .../executor/test_proxy_fast_death.py | 26 -- .../llmapi/_run_mpi_lifecycle_task.py | 164 +++++++++++ .../llmapi/test_mpi_launcher_shutdown.py | 238 +++++++++++++++ tests/unittest/llmapi/test_mpi_lifecycle.py | 243 ++++++++++++++++ .../llmapi/test_mpi_server_lifecycle.py | 264 +++++++++++++++++ tests/unittest/llmapi/test_mpi_session.py | 18 +- 10 files changed, 1209 insertions(+), 196 deletions(-) create mode 100644 tests/unittest/llmapi/_run_mpi_lifecycle_task.py create mode 100644 tests/unittest/llmapi/test_mpi_launcher_shutdown.py create mode 100644 tests/unittest/llmapi/test_mpi_lifecycle.py create mode 100644 tests/unittest/llmapi/test_mpi_server_lifecycle.py diff --git a/tensorrt_llm/llmapi/mgmn_leader_node.py b/tensorrt_llm/llmapi/mgmn_leader_node.py index 0f0d62e72c98..9d6f87d5a5c5 100644 --- a/tensorrt_llm/llmapi/mgmn_leader_node.py +++ b/tensorrt_llm/llmapi/mgmn_leader_node.py @@ -1,3 +1,5 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 ''' This script is used to start the MPICommSession in the rank0 and wait for the MPI Proxy process to connect and get the MPI task to run. @@ -41,14 +43,21 @@ def stop_server_main(): is_server=False, socket_type=zmq.PAIR) + # Only queue the stop message once a peer is connected, and bound both + # sending and context termination if that peer disappears during shutdown. + queue.socket.setsockopt(zmq.IMMEDIATE, 1) + queue.socket.setsockopt(zmq.SNDTIMEO, 5000) + queue.socket.setsockopt(zmq.LINGER, 1000) try: logger_debug( f"RemoteMpiCommSessionClient [rank{global_mpi_rank()}] send shutdown signal to server\n", "green") queue.put(None) # ask RemoteMpiCommSessionServer to shutdown - except zmq.error.ZMQError as e: - logger_debug(f"Error during RemoteMpiCommSessionClient shutdown: {e}\n", - "red") + except zmq.error.ZMQError as error: + raise click.ClickException( + f"Failed to send MPI Comm server shutdown: {error}") from error + finally: + queue.close() @click.command() diff --git a/tensorrt_llm/llmapi/mpi_session.py b/tensorrt_llm/llmapi/mpi_session.py index c1fb41c42346..ecfc672548cf 100644 --- a/tensorrt_llm/llmapi/mpi_session.py +++ b/tensorrt_llm/llmapi/mpi_session.py @@ -10,8 +10,9 @@ import threading import time import traceback +from collections import deque from collections.abc import Callable -from concurrent.futures import Future, ThreadPoolExecutor, as_completed +from concurrent.futures import Future, ThreadPoolExecutor from concurrent.futures import wait as futures_wait from typing import Any, Dict, List, NamedTuple, Optional, Tuple, TypeVar @@ -342,6 +343,20 @@ def _process_start_time(pid: int) -> Optional[bytes]: _DEFAULT_IDENTITY_TIMEOUT = 300.0 +def _mgmn_shutdown_grace_seconds() -> float: + """Bound draining a failed batch and shutting down the server-owned world.""" + raw = os.environ.get("TLLM_MGMN_SHUTDOWN_GRACE_SECONDS", "60") + try: + value = float(raw) + if math.isfinite(value) and value > 0: + return value + except ValueError: + pass + logger.warning( + f"Ignoring invalid TLLM_MGMN_SHUTDOWN_GRACE_SECONDS={raw!r}; using 60s") + return 60.0 + + def _identity_barrier_timeout() -> float: """Deadline for the ``wait_shutdown`` worker-identity barrier, in seconds. @@ -872,7 +887,6 @@ def __init__(self, socket_type=zmq.PAIR, use_hmac_encryption=True) self.comm = comm - self.results = [] # the results may arrive in any order if self.comm is not None: self.session = MpiCommSession(n_workers=self.comm.Get_size(), @@ -891,7 +905,7 @@ def task_wrapper(task: Callable[..., T], *args, **kwargs) -> T: f"MpiCommSession rank{mpi_rank()} start task [{task}] with args: {args} and kwargs: {kwargs}\n", "green") - # wait for all ranks to start the task + # Pin one task to each rank before any worker can accept another task. mpi_barrier() try: @@ -901,130 +915,158 @@ def task_wrapper(task: Callable[..., T], *args, **kwargs) -> T: f"MpiCommSession rank{mpi_rank()} task [{task}] failed with exception: {e}\n", "red") traceback.print_exc() - raise e + raise finally: logger_debug( f"MpiCommSession rank{mpi_rank()} task [{task}] finished\n", "green") - mpi_barrier() + # Let exceptions reach their futures even when a peer is stuck. + # The server gates the next batch on completion of all futures. - def serve(self): - logger_debug(f"RemoteMpiCommSessionServer listening on {self.addr}\n", - "yellow") - pending_futures = [] - while True: - # Wait for any pending futures from previous tasks to complete - # This ensures all ranks are ready before accepting the next task - if pending_futures: - logger_debug( - f"RemoteMpiCommSessionServer waiting for {len(pending_futures)} pending futures to complete\n", - "grey") - n_failed = 0 - first_exc = None - # Use as_completed so that failures are logged as soon as - # they occur rather than blocking behind a stuck future. - for future in as_completed(pending_futures): - try: - future.result() # Wait for completion - except Exception as e: - n_failed += 1 - if first_exc is None: - first_exc = e - print_colored( - f"RemoteMpiCommSessionServer: MPI worker future " - f"failed: {type(e).__name__}: {e}\n", "red") - if n_failed == len(pending_futures): - # All workers failed — no point waiting further. - break - if n_failed: - logger.error( - f"RemoteMpiCommSessionServer: {n_failed}/" - f"{len(pending_futures)} MPI worker(s) failed. " - f"First error: {first_exc}") - pending_futures.clear() - logger_debug( - "RemoteMpiCommSessionServer all pending futures completed\n", - "grey") + def _shutdown_session(self, grace: float): + """Close the world owned by this server, or abort it within the deadline. - message: Optional[RemoteTask] = self.queue.get() - if message is None: - logger_debug( - f"RemoteMpiCommSessionServer [rank{global_mpi_rank()}] received shutdown signal\n", - "green") - self.session.shutdown_abort() - break - else: - logger_debug( - f"RemoteMpiCommSessionServer [rank{global_mpi_rank()}] received task [{message.task}] from {self.addr}\n", - "green") - futures = self.session.submit( - RemoteMpiCommSessionServer.task_wrapper, message.task, - *message.args, **message.kwargs) - self.num_results = self.session.n_workers - assert len(futures) == self.num_results == mpi_world_size() - # Store futures to wait for them before the next task - pending_futures = list(futures) - for future in futures: - if message.sync: - future.add_done_callback(self.mpi_future_callback) - else: - # Fire-and-forget tasks have no result channel, but a - # crashed worker must still reach the client (the - # client-side session has no futures to watch); see - # RemoteWorkerDeath. - future.add_done_callback(self.mpi_async_error_callback) - - def mpi_async_error_callback(self, future): - """Forward a worker exception to the client for async tasks. - - Runs on the executor's callback thread, like the existing sync-path - mpi_future_callback (same pre-existing cross-thread ZMQ-put pattern). - Best-effort: the socket may already be closed at shutdown. + MpiCommSession.shutdown deliberately leaves the shared COMM_WORLD pool + alive for reuse. Only the server's final shutdown owns closing it. """ - if future.cancelled(): - return - exc = future.exception() - if exc is None: + if grace <= 0: + self.session.abort() return - print_colored( - f"RemoteMpiCommSessionServer: async MPI worker failed, forwarding " - f"to client: {type(exc).__name__}: {exc}\n", "red") - try: - self.queue.put(RemoteWorkerDeath.from_exception(exc)) - except Exception as e: - logger_debug(f"Failed to forward worker death to client: {e}\n", - "red") + executor = None + if (isinstance(self.session, MpiCommSession) + and self.session.mpi_pool is not None + and self.session.mpi_pool is MPINodeState._global_mpi_pool): + executor = MPINodeState._global_comm_executor + finished = threading.Event() + errors = [] + + def shutdown(): + try: + self.session.shutdown() + if executor is not None: + executor.__exit__(None, None, None) + MPINodeState._global_comm_executor = None + MPINodeState._global_mpi_pool = None + except BaseException as error: + errors.append(error) + finally: + finished.set() + + thread = threading.Thread(target=shutdown, + name="RemoteMpiSessionShutdown", + daemon=True) + thread.start() + if not finished.wait(grace): + logger.error(f"Remote MPI shutdown exceeded {grace}s; aborting") + self.session.abort() + elif errors: + logger.error(f"Remote MPI shutdown failed: {errors[0]!r}; aborting") + self.session.abort() + else: + thread.join() - def mpi_future_callback(self, future): - logger_debug(f"rank{global_mpi_rank()} got future: {future}\n", "red") - if future.exception() is not None: - logger_debug( - f"mpi_future got exception: {future.exception()}, quitting\n", - "red") - self.queue.put(future.exception()) - return + def serve(self): + """Handle completions and control messages without waiting on a rank. - result = future.result() - self.results.append(result) - logger_debug( - f"RemoteMpiCommSessionServer working status: {len(self.results)}/{self.num_results}\n", - "grey") - if len(self.results) == self.num_results: - logger_debug( - "RemoteMpiCommSessionServer received all results, sending to client\n", - "green") + All socket operations run here, never in executor callbacks. Normal + tasks are serialized by batch; a stop request can bypass that gate. + """ + grace = _mgmn_shutdown_grace_seconds() + # A disconnected client must not prevent world teardown. Allow a short + # delivery window, including for the first worker-death notification. + self.queue.socket.setsockopt(zmq.SNDTIMEO, 1000) + self.queue.socket.setsockopt(zmq.LINGER, 1000) + queued_tasks = deque() + pending = [] + results = [] + sync = False + first_error = None + drain_deadline = None + stop_deadline = None + try: + while True: + for future in pending[:]: + if not future.done(): + continue + pending.remove(future) + try: + results.append(future.result()) + except Exception as error: + logger.error(f"Remote MPI worker failed: {error!r}") + if first_error is None: + first_error = error + # Preserve the sync API's one exception response, + # but never mix partial results into the next batch. + response = (error if sync else + RemoteWorkerDeath.from_exception(error)) + try: + self.queue.put(response) + except Exception as send_error: + logger.error( + f"Failed to send MPI worker error: {send_error!r}" + ) + raise + if not sync: + raise RuntimeError( + "Remote MPI asynchronous task failed" + ) from error + drain_deadline = time.monotonic() + grace + + if drain_deadline is not None and pending: + if time.monotonic() >= drain_deadline: + raise RuntimeError( + "MPI workers did not drain after a task failure" + ) from first_error + + if not pending: + if sync and first_error is None: + self.queue.put(results) + # Results and error state belong to a single submitted batch. + results = [] + sync = False + first_error = None + drain_deadline = None + if queued_tasks: + message = queued_tasks.popleft() + pending = list( + self.session.submit(self.task_wrapper, message.task, + *message.args, + **message.kwargs)) + assert len(pending) == self.session.n_workers + sync = message.sync + + if stop_deadline is not None: + if not pending and not queued_tasks: + break + if time.monotonic() >= stop_deadline: + raise RuntimeError( + "Remote MPI shutdown deadline expired") + # Preserve tasks submitted before stop, but bound the whole + # drain so a wedged rank cannot strand server teardown. + time.sleep(0.1) + continue + + # Keep receiving stop requests even while workers are stuck. + # Other requests wait until every future in the batch is done. + if self.queue.poll(0.1): + message = self.queue.get() + if message is None: + stop_deadline = time.monotonic() + grace + else: + queued_tasks.append(message) + finally: try: - self.queue.put_noblock(self.results, retry=2) - except zmq.ZMQError as e: - # The client could be shutdown first. - if e.errno == zmq.EAGAIN: - pass - else: - raise e - - logger_debug("RemoteMpiCommSessionServer sent results to client\n", - "green") - self.results.clear() + deadlines = [ + deadline for deadline in (drain_deadline, stop_deadline) + if deadline is not None + ] + if deadlines: + grace = max(0.0, + min(grace, + min(deadlines) - time.monotonic())) + self._shutdown_session(grace) + finally: + self.queue.close() def find_free_port() -> int: diff --git a/tensorrt_llm/llmapi/trtllm-llmapi-launch b/tensorrt_llm/llmapi/trtllm-llmapi-launch index e9323d93f3ba..e17b7c180c43 100755 --- a/tensorrt_llm/llmapi/trtllm-llmapi-launch +++ b/tensorrt_llm/llmapi/trtllm-llmapi-launch @@ -10,11 +10,16 @@ mpi_rank=${SLURM_PROCID:-${OMPI_COMM_WORLD_RANK:-${PMI_RANK:-${PMI_ID:-0}}}} flashinfer_workspace_lock_held=0 flashinfer_temporary_workspace= +owned_child_pids=() log_stderr() { echo -e "\033[33m$@\033[0m" >&2; } log_stderr "mpi_rank: $mpi_rank" -pid=$(ps -o pid= -p $$ | tr -d ' ') +stop_timeout=${TLLM_LLMAPI_LAUNCH_STOP_TIMEOUT:-120} +if [[ ! "$stop_timeout" =~ ^[1-9][0-9]*$ ]] || [ "${#stop_timeout}" -gt 6 ]; then + log_stderr "TLLM_LLMAPI_LAUNCH_STOP_TIMEOUT must be a positive integer of at most six digits" + exit 2 +fi # Tell TRTLLM to spawn a additional process for the Proxy export TLLM_SPAWN_PROXY_PROCESS=1 @@ -31,9 +36,38 @@ function mpi_world_size { fi } -function cleanup_flashinfer_temporary_workspace { +function cleanup_owned_children { + local child_pid + local cleanup_deadline=$((SECONDS + 2)) + local children_alive + + # Each background child leads a process group created by job control. This + # also reaches its descendants, without targeting the caller's process group. + for child_pid in "${owned_child_pids[@]}"; do + kill -TERM -- "-$child_pid" 2>/dev/null || true + done + while [ "$SECONDS" -lt "$cleanup_deadline" ]; do + children_alive=0 + for child_pid in "${owned_child_pids[@]}"; do + if kill -0 -- "-$child_pid" 2>/dev/null; then + children_alive=1 + fi + done + [ "$children_alive" -eq 0 ] && break + sleep 0.1 + done + for child_pid in "${owned_child_pids[@]}"; do + kill -KILL -- "-$child_pid" 2>/dev/null || true + wait "$child_pid" 2>/dev/null || true + done +} + +function cleanup_launcher { local exit_status=$? + trap - INT TERM + cleanup_owned_children + if [ -n "$flashinfer_temporary_workspace" ]; then if ! rm -rf -- "$flashinfer_temporary_workspace"; then log_stderr "Failed to remove temporary FlashInfer JIT workspace: $flashinfer_temporary_workspace" @@ -44,7 +78,9 @@ function cleanup_flashinfer_temporary_workspace { exit "$exit_status" } -trap cleanup_flashinfer_temporary_workspace EXIT +trap cleanup_launcher EXIT +trap 'exit 130' INT +trap 'exit 143' TERM function use_unique_flashinfer_workspace { local temporary_root=${TMPDIR:-/tmp} @@ -204,8 +240,15 @@ if [ -z "$mpi_rank" ] || [ "$mpi_rank" -eq 0 ]; then I_MPI_ HYDRA_ KMP_ MPICH_ MV2_ CRAY_ ) - ( - # Remove MPI-related variables only in the subshell context + function exec_child_command { + if [ "$flashinfer_workspace_lock_held" -eq 1 ]; then + exec 200>&- + fi + exec "$@" + } + + function exec_without_mpi_environment { + local var prefix for var in $(compgen -e); do for prefix in "${mpi_blacklist[@]}"; do if [[ "$var" == "$prefix"* ]]; then @@ -214,53 +257,76 @@ if [ -z "$mpi_rank" ] || [ "$mpi_rank" -eq 0 ]; then fi done done + exec_child_command "$@" + } - # Turn off "exit on error" so the following lines always run - set +e - - # Execute the task with cleaned environment - run_child_command "${task_with_command[@]}" - task_exit_code=$? - log_stderr "Rank${mpi_rank} Task exit code: $task_exit_code" - - # Stop the MPI Comm server - run_child_command python3 -m tensorrt_llm.llmapi.mgmn_leader_node --action stop - mpi_exit_code=$? - log_stderr "Rank${mpi_rank} MPI Comm server exit code: $mpi_exit_code" - - # Propagate task exit status - if [ $task_exit_code -ne 0 ]; then - exit $task_exit_code - else - exit $mpi_exit_code + # Job control gives each child its own process group. Keep the actual child + # PIDs so cleanup never depends on a process-name search or a shell wrapper. + set -m + set +e + exec_child_command python3 -m tensorrt_llm.llmapi.mgmn_leader_node & + server_pid=$! + owned_child_pids+=("$server_pid") + exec_without_mpi_environment "${task_with_command[@]}" & + task_pid=$! + owned_child_pids+=("$task_pid") + log_stderr "Rank${mpi_rank} task PID: $task_pid; MPI Comm server PID: $server_pid" + + while kill -0 "$task_pid" 2>/dev/null; do + if ! kill -0 "$server_pid" 2>/dev/null; then + wait "$server_pid" + server_exit_code=$? + log_stderr "MPI Comm server exited before the task (status $server_exit_code); terminating task" + # Even a successful server exit is a failure while its task is alive. + [ "$server_exit_code" -eq 0 ] && server_exit_code=1 + exit "$server_exit_code" fi - ) & + sleep 0.1 + done + wait "$task_pid" + task_exit_code=$? + log_stderr "Rank${mpi_rank} Task exit code: $task_exit_code" + + # The deadline covers the stop helper's imports, connection, send and queue + # cleanup as well as the MPI server. Neither child can block this supervisor. + shutdown_deadline=$((SECONDS + stop_timeout)) + stop_pid= + stop_exit_code=0 + if kill -0 "$server_pid" 2>/dev/null; then + exec_without_mpi_environment python3 -m tensorrt_llm.llmapi.mgmn_leader_node --action stop & + stop_pid=$! + owned_child_pids+=("$stop_pid") + fi - # Turn off "exit on error" so the following lines always run - set +e + shutdown_timed_out=0 + while kill -0 "$server_pid" 2>/dev/null || \ + { [ -n "$stop_pid" ] && kill -0 "$stop_pid" 2>/dev/null; }; do + if [ "$SECONDS" -ge "$shutdown_deadline" ]; then + log_stderr "MPI Comm shutdown exceeded ${stop_timeout}s; terminating owned processes" + shutdown_timed_out=1 + break + fi + sleep 0.1 + done - # Capture subshell PID - subshell_pid=$! - log_stderr "Rank${mpi_rank} Subshell PID: $subshell_pid" - - log_stderr "Rank${mpi_rank} run mgmn leader node with mpi_world_size: $(mpi_world_size) ..." - log_stderr "Rank0 host: $HOSTNAME" - run_child_command python3 -m tensorrt_llm.llmapi.mgmn_leader_node - mgmn_leader_node_exit_code=$? - log_stderr "Rank${mpi_rank} MGMN leader node exit code: $mgmn_leader_node_exit_code" - - # Wait for subshell - wait $subshell_pid - # This is subshell's exit code - subshell_exit_code=$? - log_stderr "Rank${mpi_rank} Subshell exit code: $subshell_exit_code" - - # Propagate subshell exit status - if [ $subshell_exit_code -ne 0 ]; then - exit $subshell_exit_code - else - exit $mgmn_leader_node_exit_code + if [ "$shutdown_timed_out" -ne 0 ]; then + [ "$task_exit_code" -ne 0 ] && exit "$task_exit_code" + exit 124 fi + if [ -n "$stop_pid" ]; then + wait "$stop_pid" + stop_exit_code=$? + log_stderr "Rank${mpi_rank} stop helper exit code: $stop_exit_code" + fi + wait "$server_pid" + server_exit_code=$? + log_stderr "Rank${mpi_rank} MPI Comm server exit code: $server_exit_code" + + # Preserve the original task failure; otherwise surface teardown failures. + [ "$task_exit_code" -ne 0 ] && exit "$task_exit_code" + [ "$stop_exit_code" -ne 0 ] && exit "$stop_exit_code" + exit "$server_exit_code" + else # Turn off "exit on error" so the following lines always run set +e diff --git a/tests/integration/test_lists/test-db/l0_cpu.yml b/tests/integration/test_lists/test-db/l0_cpu.yml index 4b1f1e4efab7..d92e1dfd410d 100644 --- a/tests/integration/test_lists/test-db/l0_cpu.yml +++ b/tests/integration/test_lists/test-db/l0_cpu.yml @@ -169,6 +169,9 @@ l0_cpu: - unittest/llmapi/test_llm_telemetry.py - unittest/llmapi/test_llm_utils.py - unittest/llmapi/test_mpi_session.py ISOLATION + - unittest/llmapi/test_mpi_server_lifecycle.py + - unittest/llmapi/test_mpi_launcher_shutdown.py + - unittest/llmapi/test_mpi_lifecycle.py ISOLATION - unittest/llmapi/test_reasoning_parser.py - unittest/llmapi/test_request_priority.py - unittest/llmapi/test_rl_control_auth.py diff --git a/tests/unittest/executor/test_proxy_fast_death.py b/tests/unittest/executor/test_proxy_fast_death.py index e2ab75d66c37..8db8961c4057 100644 --- a/tests/unittest/executor/test_proxy_fast_death.py +++ b/tests/unittest/executor/test_proxy_fast_death.py @@ -457,32 +457,6 @@ def test_remote_worker_death_roundtrip(): assert "rank 3 exploded" in str(exc) -def test_server_async_callback_forwards_only_failures(): - from concurrent.futures import Future - - from tensorrt_llm.llmapi.mpi_session import RemoteMpiCommSessionServer, RemoteWorkerDeath - - server = object.__new__(RemoteMpiCommSessionServer) - sent = [] - server.queue = type("Q", (), {"put": lambda self, m: sent.append(m)})() - - ok = Future() - ok.set_result(42) - server.mpi_async_error_callback(ok) - assert sent == [] - - cancelled = Future() - cancelled.cancel() - server.mpi_async_error_callback(cancelled) - assert sent == [] - - failed = Future() - failed.set_exception(RuntimeError("worker segfault")) - server.mpi_async_error_callback(failed) - assert len(sent) == 1 and isinstance(sent[0], RemoteWorkerDeath) - assert sent[0].message == "worker segfault" - - class _FakeZmqQueue: """poll()/get() stub fed with a fixed message sequence.""" diff --git a/tests/unittest/llmapi/_run_mpi_lifecycle_task.py b/tests/unittest/llmapi/_run_mpi_lifecycle_task.py new file mode 100644 index 000000000000..1004b3ff5206 --- /dev/null +++ b/tests/unittest/llmapi/_run_mpi_lifecycle_task.py @@ -0,0 +1,164 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Importable MPI tasks and the non-MPI engine used by lifecycle regressions.""" + +import argparse +import json +import os +import time +from pathlib import Path + +import psutil + + +def _mark(directory: Path, name: str, value: object = True) -> None: + temporary = directory / f".{name}.{os.getpid()}.tmp" + temporary.write_text(json.dumps(value)) + temporary.replace(directory / f"{name}.json") + + +def _record_identity(directory: Path, name: str) -> None: + process = psutil.Process() + _mark(directory, f"identity-{name}", {"pid": process.pid, "created": process.create_time()}) + + +def _wait_for_markers(directory: Path, names: list[str], timeout: float = 20) -> None: + deadline = time.monotonic() + timeout + while not all((directory / f"{name}.json").exists() for name in names): + if time.monotonic() >= deadline: + raise TimeoutError(f"Missing lifecycle readiness markers: {names}") + time.sleep(0.01) + + +def worker_task(directory_name: str, scenario: str, batch: int, size: int) -> tuple[int, int]: + """Record per-rank execution, then finish, fail, or wait for world teardown.""" + from mpi4py import MPI + + directory = Path(directory_name) + rank = MPI.COMM_WORLD.Get_rank() + assert MPI.COMM_WORLD.Get_size() == size + _record_identity(directory, f"rank-{rank}") + + if batch: + assert all( + (directory / f"finished-{batch - 1}-{peer}.json").exists() for peer in range(size) + ), "The next task started before every rank finished the previous task" + event_file = directory / f"events-{rank}.jsonl" + with event_file.open("a") as events: + events.write(json.dumps({"batch": batch, "scenario": scenario}) + "\n") + _mark(directory, f"started-{batch}-{rank}") + + if scenario == "async_drain" and batch == 0: + _wait_for_markers(directory, ["engine-exiting"]) + + if scenario in ("mixed_failure", "all_hang"): + _wait_for_markers(directory, [f"started-{batch}-{peer}" for peer in range(size)]) + if scenario == "mixed_failure" and rank % 2 == 0: + raise RuntimeError("injected MPI lifecycle failure") + time.sleep(3600) + raise AssertionError("The owner did not terminate the stuck worker world") + + # Different completion times exercise ordering between queued batches. + time.sleep(0.01 * (rank + 1)) + _mark(directory, f"finished-{batch}-{rank}") + if scenario == "sync_failure" and rank == 0: + raise RuntimeError("injected recoverable sync failure") + return rank, batch + + +def main() -> int: + """Run one engine scenario; readiness is signalled through persistent files.""" + parser = argparse.ArgumentParser() + parser.add_argument("scenario") + parser.add_argument("directory", type=Path) + parser.add_argument("--ranks", type=int, required=True) + args = parser.parse_args() + directory = args.directory + _record_identity(directory, "engine") + _mark(directory, "engine-started") + if args.scenario == "no_submission": + return 3 + + # Import through the module name so MPI receives an importable callable, + # rather than a function belonging to the engine's __main__ module. + from _run_mpi_lifecycle_task import worker_task as remote_task + + from tensorrt_llm.executor.utils import ( + get_spawn_proxy_process_ipc_addr_env, + get_spawn_proxy_process_ipc_hmac_key_env, + ) + from tensorrt_llm.llmapi.mpi_session import RemoteMpiCommSessionClient + + address = get_spawn_proxy_process_ipc_addr_env() + key = get_spawn_proxy_process_ipc_hmac_key_env() + client = RemoteMpiCommSessionClient(address, hmac_key=key) + client.SYNC_IDLE_INTERVAL = 0.01 + + def run_sync(scenario: str, batch: int) -> object: + return client.submit_sync(remote_task, str(directory), scenario, batch, args.ranks) + + def assert_results(response: object, batch: int) -> None: + assert isinstance(response, list), response + assert sorted(response) == [(rank, batch) for rank in range(args.ranks)], response + + if args.scenario in ("mixed_failure", "all_hang"): + client.submit(remote_task, str(directory), args.scenario, 0, args.ranks) + _wait_for_markers(directory, [f"started-0-{rank}" for rank in range(args.ranks)]) + _mark(directory, "workers-started") + if args.scenario == "all_hang": + _mark(directory, "engine-exiting", {"status": 1, "monotonic": time.monotonic()}) + return 1 + deadline = time.monotonic() + 30 + while time.monotonic() < deadline: + error = client.check_worker_error() + if error is not None: + assert "injected MPI lifecycle failure" in str(error), error + _mark( + directory, + "error-observed", + {"error": str(error), "monotonic": time.monotonic()}, + ) + # The owner must end this run even while the engine remains + # alive; this is not the engine-exit shutdown path. + time.sleep(3600) + raise AssertionError("The engine survived fatal worker-world teardown") + time.sleep(0.01) + raise TimeoutError("No worker failure reached the live engine") + + if args.scenario == "all_return": + assert_results(run_sync("return", 0), 0) + elif args.scenario == "sync_recovery": + response = run_sync("sync_failure", 0) + assert isinstance(response, Exception), response + assert "injected recoverable sync failure" in str(response), response + _mark(directory, "sync-error-observed") + assert_results(run_sync("return", 1), 1) + elif args.scenario == "reuse": + for batch in range(3): + client.submit(remote_task, str(directory), "return", batch, args.ranks) + assert_results(run_sync("return", 3), 3) + for batch in range(4, 7): + assert_results(run_sync("return", batch), batch) + client.shutdown() + reused = RemoteMpiCommSessionClient(address, hmac_key=key) + assert reused is client + assert_results(run_sync("return", 7), 7) + elif args.scenario == "async_drain": + for batch in range(3): + client.submit(remote_task, str(directory), "async_drain", batch, args.ranks) + _wait_for_markers(directory, [f"started-0-{rank}" for rank in range(args.ranks)]) + # The first batch cannot finish until every request is queued and the + # engine exits. There is deliberately no synchronous flush request. + _mark(directory, "engine-exiting", {"status": 0, "monotonic": time.monotonic()}) + return 0 + else: + raise ValueError(f"Unknown lifecycle scenario: {args.scenario}") + + client.shutdown() + _mark(directory, "engine-completed") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tests/unittest/llmapi/test_mpi_launcher_shutdown.py b/tests/unittest/llmapi/test_mpi_launcher_shutdown.py new file mode 100644 index 000000000000..92a863fa2f6a --- /dev/null +++ b/tests/unittest/llmapi/test_mpi_launcher_shutdown.py @@ -0,0 +1,238 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Exercise launcher supervision without importing TensorRT-LLM or initializing MPI.""" + +import json +import os +import signal +import subprocess +import sys +from pathlib import Path + +import pytest + +pytestmark = pytest.mark.cpu_only + +_LAUNCHER = Path(__file__).parents[3] / "tensorrt_llm" / "llmapi" / "trtllm-llmapi-launch" +_MPI_PREFIXES = ( + "OMPI_", + "PMIX_", + "PMI_", + "SLURM_", + "MPI_", + "UCX_", + "I_MPI_", + "HYDRA_", + "KMP_", + "MPICH_", + "MV2_", + "CRAY_", +) + +# FIFOs establish readiness explicitly: the engine cannot finish before the +# server starts, and an early-exiting server waits for the engine's child. +_STUB = r""" +import json +import os +import signal +import subprocess +import sys +from pathlib import Path + +root = Path(os.environ["LAUNCHER_TEST_ROOT"]) +mode = os.environ["LAUNCHER_TEST_MODE"] +role = "task" if sys.argv[1] == "task" else "server" +if "--action" in sys.argv: + role = "stop" +if sys.argv[1] == "-S": + os.execv(sys.executable, [sys.executable, *sys.argv[1:]]) + +try: + os.fstat(200) + lock_fd_open = True +except OSError: + lock_fd_open = False +(root / f"{role}.json").write_text(json.dumps({ + "pid": os.getpid(), + "pgid": os.getpgrp(), + "pmi_rank": os.environ.get("PMI_RANK"), + "lock_fd_open": lock_fd_open, + "workspace": os.environ.get("FLASHINFER_WORKSPACE_BASE"), +})) + +def send(name): + with (root / name).open("w") as stream: + stream.write("ready\n") + +def receive(name): + with (root / name).open() as stream: + assert stream.readline() == "ready\n" + +def hang(): + signal.signal(signal.SIGTERM, signal.SIG_IGN) + while True: + signal.pause() + +if role == "server": + send("server_ready") + receive("task_ready") + if mode == "server_exits": + sys.exit(17) + receive("stop_requested") + if mode == "server_hangs": + hang() +elif role == "stop": + if mode == "stop_hangs": + hang() + send("stop_requested") +else: + receive("server_ready") + if mode == "server_exits": + child = subprocess.Popen([ + sys.executable, "-c", + "import signal; print('ready', flush=True); signal.pause()", + ], stdout=subprocess.PIPE, text=True) + assert child.stdout.readline() == "ready\n" + (root / "task_child.json").write_text(json.dumps({ + "pid": child.pid, + "pgid": os.getpgid(child.pid), + })) + def terminate(signum, frame): + raise SystemExit(128 + signum) + signal.signal(signal.SIGTERM, terminate) + try: + send("task_ready") + while True: + signal.pause() + finally: + child.terminate() + child.wait(timeout=5) + else: + send("task_ready") + sys.exit(int(os.environ.get("LAUNCHER_TEST_TASK_STATUS", "0"))) +""" + + +def _launcher_env(tmp_path: Path, mode: str) -> dict[str, str]: + stub_bin = tmp_path / "bin" + stub_bin.mkdir() + python_stub = stub_bin / "python3" + python_stub.write_text(f"#!{sys.executable}\n{_STUB}") + python_stub.chmod(0o755) + for name in ("server_ready", "task_ready", "stop_requested"): + os.mkfifo(tmp_path / name) + env = { + name: value + for name, value in os.environ.items() + if not name.startswith(_MPI_PREFIXES) + and not name.startswith(("FLASHINFER_", "TRTLLM_FLASHINFER_", "TLLM_SPAWN_PROXY_PROCESS")) + } + env.update( + { + "HOME": str(tmp_path / "home"), + "PATH": f"{stub_bin}{os.pathsep}{env['PATH']}", + "PMI_RANK": "0", + "PMI_SIZE": "1", + "TLLM_LLMAPI_LAUNCH_STOP_TIMEOUT": "1", + "TLLM_SPAWN_PROXY_PROCESS_IPC_ADDR": f"ipc://{tmp_path / 'ipc'}", + "LAUNCHER_TEST_ROOT": str(tmp_path), + "LAUNCHER_TEST_MODE": mode, + } + ) + return env + + +def _process_records(tmp_path: Path) -> dict[str, dict]: + return {path.stem: json.loads(path.read_text()) for path in tmp_path.glob("*.json")} + + +def _run_launcher(tmp_path: Path, env: dict[str, str]) -> subprocess.CompletedProcess[str]: + command = ["bash", str(_LAUNCHER), str(tmp_path / "bin" / "python3"), "task"] + with subprocess.Popen( # nosec B603 + command, + env=env, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + start_new_session=True, + ) as process: + try: + stdout, stderr = process.communicate(timeout=12) + _assert_exited(tmp_path) + return subprocess.CompletedProcess(command, process.returncode, stdout, stderr) + finally: + # Clean up this test's groups even if a broken launcher times out. + process_groups = {process.pid} + process_groups.update(record["pgid"] for record in _process_records(tmp_path).values()) + for group in process_groups: + try: + os.killpg(group, signal.SIGKILL) + except ProcessLookupError: + pass + process.wait(timeout=5) + + +def _assert_exited(tmp_path: Path) -> None: + for role, record in _process_records(tmp_path).items(): + try: + os.kill(record["pid"], 0) + except ProcessLookupError: + continue + pytest.fail(f"Launcher left {role} process {record['pid']} alive") + + +@pytest.mark.parametrize("task_status", [0, 7]) +def test_launcher_preserves_task_status_and_child_environment( + tmp_path: Path, task_status: int +) -> None: + env = _launcher_env(tmp_path, "clean") + env["LAUNCHER_TEST_TASK_STATUS"] = str(task_status) + result = _run_launcher(tmp_path, env) + assert result.returncode == task_status, result.stderr + records = _process_records(tmp_path) + assert set(records) == {"task", "server", "stop"} + assert records["server"]["pmi_rank"] == "0" + for role in ("task", "stop"): + assert records[role]["pmi_rank"] is None + for record in records.values(): + assert not record["lock_fd_open"] + assert record["workspace"].endswith("/rank-0") + assert record["pid"] == record["pgid"] + _assert_exited(tmp_path) + + +@pytest.mark.parametrize("mode", ["server_hangs", "stop_hangs"]) +@pytest.mark.parametrize("task_status", [0, 7]) +def test_launcher_deadline_covers_stop_helper_and_server( + tmp_path: Path, + mode: str, + task_status: int, +) -> None: + env = _launcher_env(tmp_path, mode) + env["LAUNCHER_TEST_TASK_STATUS"] = str(task_status) + result = _run_launcher(tmp_path, env) + assert result.returncode == (task_status or 124), result.stderr + assert "MPI Comm shutdown exceeded 1s" in result.stderr + assert "stop" in _process_records(tmp_path) + _assert_exited(tmp_path) + + +def test_launcher_terminates_engine_and_child_when_server_exits(tmp_path: Path) -> None: + env = _launcher_env(tmp_path, "server_exits") + result = _run_launcher(tmp_path, env) + assert result.returncode == 17, result.stderr + assert "MPI Comm server exited before the task" in result.stderr + assert "task_child" in _process_records(tmp_path) + assert "stop" not in _process_records(tmp_path) + _assert_exited(tmp_path) + + +@pytest.mark.parametrize("value", ["0", "-1", "invalid", "1.5", "9999999"]) +def test_launcher_rejects_invalid_shutdown_deadline(tmp_path: Path, value: str) -> None: + env = _launcher_env(tmp_path, "clean") + env["TLLM_LLMAPI_LAUNCH_STOP_TIMEOUT"] = value + result = _run_launcher(tmp_path, env) + assert result.returncode == 2, result.stderr + assert "TLLM_LLMAPI_LAUNCH_STOP_TIMEOUT must be a positive integer" in result.stderr + assert not _process_records(tmp_path) diff --git a/tests/unittest/llmapi/test_mpi_lifecycle.py b/tests/unittest/llmapi/test_mpi_lifecycle.py new file mode 100644 index 000000000000..2a6e29ac7296 --- /dev/null +++ b/tests/unittest/llmapi/test_mpi_lifecycle.py @@ -0,0 +1,243 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Single-node, model-free regressions for the remote MPI worker lifecycle.""" + +import importlib.util +import json +import os +import shutil +import signal +import subprocess +import sys +import time +from pathlib import Path + +import psutil +import pytest + +from tensorrt_llm.bindings.BuildInfo import ENABLE_MULTI_DEVICE + +_MPI_RUNTIME_AVAILABLE = ( + sys.platform == "linux" + and ENABLE_MULTI_DEVICE + and shutil.which("mpirun") is not None + and importlib.util.find_spec("mpi4py") is not None +) + + +def _live_processes(identities: dict[int, float]) -> list[psutil.Process]: + result = [] + for pid, created in identities.items(): + try: + process = psutil.Process(pid) + if process.create_time() == created and process.status() != psutil.STATUS_ZOMBIE: + result.append(process) + except psutil.NoSuchProcess: + pass + return result + + +def _observe_processes(root: psutil.Process, directory: Path, identities: dict[int, float]) -> None: + try: + for child in root.children(recursive=True): + try: + identities[child.pid] = child.create_time() + except psutil.NoSuchProcess: + pass + except psutil.NoSuchProcess: + pass + for path in directory.glob("identity-*.json"): + identity = json.loads(path.read_text()) + identities[identity["pid"]] = identity["created"] + + +def _cleanup_processes(identities: dict[int, float]) -> None: + processes = _live_processes(identities) + for process in processes: + try: + process.terminate() + except psutil.NoSuchProcess: + pass + _, alive = psutil.wait_procs(processes, timeout=2) + for process in alive: + try: + process.kill() + except psutil.NoSuchProcess: + pass + psutil.wait_procs(alive, timeout=2) + + +@pytest.mark.cpu_only +@pytest.mark.skipif( + not _MPI_RUNTIME_AVAILABLE, reason="Linux and a multi-device MPI runtime required" +) +@pytest.mark.parametrize("ranks", [2, 4]) +@pytest.mark.parametrize( + "scenario", + [ + "mixed_failure", + "all_hang", + "all_return", + "no_submission", + "reuse", + "sync_recovery", + "async_drain", + ], +) +def test_remote_mpi_worker_lifecycle(scenario: str, ranks: int, tmp_path: Path) -> None: + """Verify bounded failure and clean reuse without leaving owned processes alive.""" + launcher = shutil.which("trtllm-llmapi-launch") + assert launcher is not None, "The matching trtllm-llmapi-launch must be installed" + test_directory = Path(__file__).resolve().parent + # This is a separate local MPI job, including when pytest itself runs in + # a Slurm step. Inherited rank identities must not override its own ranks. + env = { + key: value + for key, value in os.environ.items() + if not key.startswith(("SLURM_", "OMPI_COMM_", "PMIX_", "PMI_")) + and key + not in ( + "TLLM_SPAWN_PROXY_PROCESS", + "TLLM_SPAWN_PROXY_PROCESS_IPC_ADDR", + "TLLM_SPAWN_PROXY_PROCESS_IPC_HMAC_KEY", + "tllm_mpi_size", + ) + } + env.update( + OMPI_ALLOW_RUN_AS_ROOT="1", + OMPI_ALLOW_RUN_AS_ROOT_CONFIRM="1", + PRTE_ALLOW_RUN_AS_ROOT="1", + PRTE_ALLOW_RUN_AS_ROOT_CONFIRM="1", + TLLM_MGMN_SHUTDOWN_GRACE_SECONDS="5", + TLLM_LLMAPI_LAUNCH_STOP_TIMEOUT="90", + TLLM_LLMAPI_ZMQ_DEBUG="1", + TLLM_LOG_LEVEL="info", + PYTHONUNBUFFERED="1", + ) + env["PYTHONPATH"] = os.pathsep.join( + [str(test_directory)] + ([env["PYTHONPATH"]] if env.get("PYTHONPATH") else []) + ) + command = [ + "mpirun", + "--allow-run-as-root", + "--host", + "localhost", + "--oversubscribe", + "--bind-to", + "none", + ] + for variable in ( + "OMPI_ALLOW_RUN_AS_ROOT", + "OMPI_ALLOW_RUN_AS_ROOT_CONFIRM", + "PRTE_ALLOW_RUN_AS_ROOT", + "PRTE_ALLOW_RUN_AS_ROOT_CONFIRM", + "TLLM_MGMN_SHUTDOWN_GRACE_SECONDS", + "TLLM_LLMAPI_LAUNCH_STOP_TIMEOUT", + "TLLM_LLMAPI_ZMQ_DEBUG", + "PYTHONPATH", + "PYTHONUNBUFFERED", + ): + command.extend(["-x", variable]) + command.extend( + [ + "-np", + str(ranks), + launcher, + sys.executable, + "-m", + "_run_mpi_lifecycle_task", + scenario, + str(tmp_path), + "--ranks", + str(ranks), + ] + ) + + identities: dict[int, float] = {} + log_path = tmp_path / "launcher.log" + started = time.monotonic() + timed_out = False + survivors: list[int] = [] + with log_path.open("w") as output: + process = subprocess.Popen( + command, + env=env, + cwd=test_directory, + stdout=output, + stderr=subprocess.STDOUT, + start_new_session=True, + ) + root = psutil.Process(process.pid) + try: + while process.poll() is None: + _observe_processes(root, tmp_path, identities) + if time.monotonic() - started >= 180: + timed_out = True + try: + os.killpg(process.pid, signal.SIGKILL) + except ProcessLookupError: + pass + break + time.sleep(0.05) + process.wait(timeout=10) + exited = time.monotonic() + _observe_processes(root, tmp_path, identities) + cleanup_deadline = time.monotonic() + 3 + while _live_processes(identities) and time.monotonic() < cleanup_deadline: + time.sleep(0.05) + survivors = [child.pid for child in _live_processes(identities)] + finally: + if process.poll() is None: + try: + os.killpg(process.pid, signal.SIGKILL) + except ProcessLookupError: + pass + process.wait(timeout=10) + _cleanup_processes(identities) + + evidence = { + "command": command, + "returncode": process.returncode, + "timed_out": timed_out, + "elapsed_seconds": time.monotonic() - started, + "survivors_before_cleanup": survivors, + "process_identities": identities, + } + (tmp_path / "result.json").write_text(json.dumps(evidence, indent=2)) + log = log_path.read_text() + print(f"MPI lifecycle evidence: {tmp_path}\n{json.dumps(evidence, indent=2)}\n{log}") + assert not timed_out, f"The outer safety timeout ended the MPI run:\n{log}" + assert not survivors, f"Owned processes survived MPI launcher exit: {survivors}\n{log}" + assert not _live_processes(identities), "Failed to clean up the test's own processes" + assert (tmp_path / "engine-started.json").exists(), log + assert "ZMQ thread safety violation" not in log, log + + if scenario in ("mixed_failure", "all_hang"): + assert process.returncode != 0, log + assert (tmp_path / "workers-started.json").exists(), log + expected_marker = "error-observed" if scenario == "mixed_failure" else "engine-exiting" + marker = tmp_path / f"{expected_marker}.json" + assert marker.exists(), log + triggered = json.loads(marker.read_text())["monotonic"] + # The stop helper imports TensorRT-LLM afresh after engine exit. The + # mixed-failure path needs no new interpreter and has a tighter bound. + teardown_budget = 20 if scenario == "mixed_failure" else 120 + assert exited - triggered < teardown_budget, f"Teardown exceeded its bounded budget:\n{log}" + elif scenario == "no_submission": + assert process.returncode == 3, log + elif scenario == "async_drain": + assert process.returncode == 0, log + assert (tmp_path / "engine-exiting.json").exists(), log + else: + assert process.returncode == 0, log + assert (tmp_path / "engine-completed.json").exists(), log + + batches = {"all_return": 1, "reuse": 8, "sync_recovery": 2, "async_drain": 3}.get(scenario, 1) + if scenario != "no_submission": + for rank in range(ranks): + events = [ + json.loads(line) + for line in (tmp_path / f"events-{rank}.jsonl").read_text().splitlines() + ] + assert [event["batch"] for event in events] == list(range(batches)), events diff --git a/tests/unittest/llmapi/test_mpi_server_lifecycle.py b/tests/unittest/llmapi/test_mpi_server_lifecycle.py new file mode 100644 index 000000000000..7f29520e5609 --- /dev/null +++ b/tests/unittest/llmapi/test_mpi_server_lifecycle.py @@ -0,0 +1,264 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Control-loop and owner-teardown regressions; real MPI coverage is separate.""" + +import threading +from collections import deque +from concurrent.futures import Future +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +import zmq + +from tensorrt_llm.llmapi import mpi_session as mpi + +pytestmark = pytest.mark.cpu_only + + +def _future(result=None, error=None): + future = Future() + if error is None: + future.set_result(result) + else: + future.set_exception(error) + return future + + +def _task(): + return 42 + + +class _Queue: + def __init__(self, messages, on_poll=None): + self.messages = deque(messages) + self.sent = [] + self.socket = Mock() + self.closed = False + self.on_poll = on_poll + self.thread_ids = set() + self.polls = 0 + + def poll(self, timeout): + self.thread_ids.add(threading.get_ident()) + self.polls += 1 + assert self.polls < 100, "server failed to make progress" + if self.on_poll: + self.on_poll(self) + return bool(self.messages) + + def get(self): + self.thread_ids.add(threading.get_ident()) + return self.messages.popleft() + + def put(self, message): + self.thread_ids.add(threading.get_ident()) + self.sent.append(message) + + def close(self): + self.closed = True + + +def _server(queue, batches): + server = object.__new__(mpi.RemoteMpiCommSessionServer) + server.queue = queue + server.session = SimpleNamespace( + n_workers=2, submit=Mock(side_effect=batches), shutdown=Mock(), abort=Mock() + ) + server._shutdown_session = Mock() + return server + + +def test_failed_task_publishes_without_a_final_collective(monkeypatch): + barrier = Mock() + monkeypatch.setattr(mpi, "mpi_barrier", barrier) + monkeypatch.setattr(mpi, "mpi_rank", lambda: 0) + monkeypatch.setattr(mpi, "mpi_world_size", lambda: 2) + + def fail(): + raise ValueError("injected failure") + + with pytest.raises(ValueError, match="injected failure"): + mpi.RemoteMpiCommSessionServer.task_wrapper(fail) + barrier.assert_called_once() + + +def test_stop_is_processed_while_all_futures_are_pending(monkeypatch): + monkeypatch.setattr(mpi, "_mgmn_shutdown_grace_seconds", lambda: 0.01) + queue = _Queue([mpi.RemoteTask(_task, (), {}), None]) + server = _server(queue, [[Future(), Future()]]) + with pytest.raises(RuntimeError, match="shutdown deadline expired"): + server.serve() + server.session.submit.assert_called_once() + server._shutdown_session.assert_called_once_with(0) + assert queue.closed + + +def test_async_failure_triggers_shutdown_without_waiting_for_peer(): + queue = _Queue([mpi.RemoteTask(_task, (), {})]) + server = _server(queue, [[Future(), _future(error=ValueError("rank failed"))]]) + with pytest.raises(RuntimeError, match="asynchronous task failed"): + server.serve() + assert queue.sent == [mpi.RemoteWorkerDeath("ValueError", "rank failed")] + assert queue.thread_ids == {threading.get_ident()} + server._shutdown_session.assert_called_once() + assert queue.closed + + +def test_failed_error_delivery_still_shuts_down(): + queue = _Queue([mpi.RemoteTask(_task, (), {})]) + queue.put = Mock(side_effect=zmq.Again()) + server = _server(queue, [[Future(), _future(error=ValueError("failed"))]]) + with pytest.raises(zmq.Again): + server.serve() + server._shutdown_session.assert_called_once() + assert queue.closed + queue.socket.setsockopt.assert_any_call(zmq.SNDTIMEO, 1000) + queue.socket.setsockopt.assert_any_call(zmq.LINGER, 1000) + + +def test_sync_error_does_not_leak_responses_into_next_batch(): + def stop_after_responses(queue): + if len(queue.sent) == 2: + queue.messages.append(None) + + queue = _Queue([mpi.RemoteTask(_task, (), {}, True)] * 2, on_poll=stop_after_responses) + error = ValueError("first batch failed") + server = _server(queue, [[_future(7), _future(error=error)], [_future(8), _future(9)]]) + server.serve() + assert queue.sent == [error, [8, 9]] + assert queue.thread_ids == {threading.get_ident()} + assert server.session.submit.call_count == 2 + + +def test_sync_multiple_failures_send_one_response(): + def stop_after_response(queue): + if queue.sent: + queue.messages.append(None) + + queue = _Queue([mpi.RemoteTask(_task, (), {}, True)], on_poll=stop_after_response) + first_error = ValueError("first") + server = _server(queue, [[_future(error=first_error), _future(error=ValueError("second"))]]) + server.serve() + assert queue.sent == [first_error] + + +def test_next_batch_waits_for_every_previous_future(): + pending = Future() + queue = _Queue([mpi.RemoteTask(_task, (0,), {}), mpi.RemoteTask(_task, (1,), {}, True)]) + server = _server(queue, [[_future(0), pending], [_future(1), _future(1)]]) + + def advance(queue): + if queue.polls == 3: + assert server.session.submit.call_count == 1 + pending.set_result(0) + if queue.sent: + queue.messages.append(None) + + queue.on_poll = advance + server.serve() + assert [call.args[2] for call in server.session.submit.call_args_list] == [0, 1] + assert queue.sent == [[1, 1]] + + +def test_sync_error_bounds_peer_drain(monkeypatch): + clock = iter([0, 0, 61, 61]) + monkeypatch.setattr(mpi.time, "monotonic", lambda: next(clock)) + queue = _Queue([mpi.RemoteTask(_task, (), {}, True)]) + error = ValueError("rank failed") + server = _server(queue, [[_future(error=error), Future()]]) + with pytest.raises(RuntimeError, match="did not drain"): + server.serve() + assert queue.sent == [error] + server._shutdown_session.assert_called_once() + + +def test_stop_drains_preceding_async_batches(): + pending = Future() + + def release_first_batch_at_stop(queue): + if queue.polls == 4: + pending.set_result(0) + + queue = _Queue( + [mpi.RemoteTask(_task, (i,), {}) for i in range(3)] + [None], + on_poll=release_first_batch_at_stop, + ) + server = _server(queue, [[_future(0), pending], [_future(1)] * 2, [_future(2)] * 2]) + server.serve() + assert [call.args[2] for call in server.session.submit.call_args_list] == [0, 1, 2] + server._shutdown_session.assert_called_once() + assert queue.closed + + +@pytest.mark.parametrize("failure", ["session", "executor", None]) +def test_final_owner_shutdown_closes_shared_executor(monkeypatch, failure): + class TestSession(mpi.MpiCommSession): + def __del__(self): + pass + + session = object.__new__(TestSession) + session.mpi_pool = object() + session.shutdown = Mock(side_effect=RuntimeError("shutdown") if failure == "session" else None) + session.abort = Mock() + executor = Mock() + executor.__exit__ = Mock( + side_effect=RuntimeError("executor") if failure == "executor" else None + ) + monkeypatch.setattr(mpi.MPINodeState, "_global_mpi_pool", session.mpi_pool) + monkeypatch.setattr(mpi.MPINodeState, "_global_comm_executor", executor) + server = object.__new__(mpi.RemoteMpiCommSessionServer) + server.session = session + server._shutdown_session(1) + session.shutdown.assert_called_once_with() + if failure: + session.abort.assert_called_once_with() + else: + executor.__exit__.assert_called_once_with(None, None, None) + session.abort.assert_not_called() + assert mpi.MPINodeState._global_comm_executor is None + assert mpi.MPINodeState._global_mpi_pool is None + + +def test_unrelated_global_executor_is_not_closed(monkeypatch): + executor = Mock() + executor.__exit__ = Mock() + monkeypatch.setattr(mpi.MPINodeState, "_global_comm_executor", executor) + server = object.__new__(mpi.RemoteMpiCommSessionServer) + server.session = SimpleNamespace(shutdown=Mock(), abort=Mock()) + server._shutdown_session(1) + executor.__exit__.assert_not_called() + server.session.abort.assert_not_called() + + +def test_shutdown_timeout_aborts_the_world(): + entered = threading.Event() + release = threading.Event() + finished = threading.Event() + + def shutdown(): + entered.set() + release.wait(5) + finished.set() + + server = object.__new__(mpi.RemoteMpiCommSessionServer) + server.session = SimpleNamespace(shutdown=shutdown, abort=Mock()) + try: + server._shutdown_session(0.02) + assert entered.is_set() + server.session.abort.assert_called_once_with() + finally: + release.set() + assert finished.wait(1) + + +@pytest.mark.parametrize("raw", ["", "0", "-1", "nan", "inf", "bad"]) +def test_invalid_shutdown_grace_uses_default(monkeypatch, raw): + monkeypatch.setenv("TLLM_MGMN_SHUTDOWN_GRACE_SECONDS", raw) + assert mpi._mgmn_shutdown_grace_seconds() == 60 + + +def test_shutdown_grace_accepts_positive_finite_value(monkeypatch): + monkeypatch.setenv("TLLM_MGMN_SHUTDOWN_GRACE_SECONDS", "0.5") + assert mpi._mgmn_shutdown_grace_seconds() == 0.5 diff --git a/tests/unittest/llmapi/test_mpi_session.py b/tests/unittest/llmapi/test_mpi_session.py index 3f94a2e97c93..32ee0c9dabf9 100644 --- a/tests/unittest/llmapi/test_mpi_session.py +++ b/tests/unittest/llmapi/test_mpi_session.py @@ -290,10 +290,19 @@ def _launcher_env(tmp_path: Path, home: str) -> dict: stub_bin = tmp_path / "bin" stub_bin.mkdir(exist_ok=True) python_stub = stub_bin / "python3" - python_stub.write_text("#!/bin/sh\n" - "if [ \"$1\" = \"-c\" ]; then\n" - " echo ipc:///tmp/trtllm-pmi-workspace-test\n" - "fi\n") + stop_fifo = tmp_path / "launcher-stop" + os.mkfifo(stop_fifo) + python_stub.write_text( + "#!/bin/sh\n" + "if [ \"$1\" = \"-c\" ]; then\n" + " echo ipc:///tmp/trtllm-pmi-workspace-test\n" + "elif [ \"$1\" = \"-m\" ]; then\n" + " if [ \"$4\" = \"stop\" ]; then\n" + " echo stop > \"$LAUNCHER_TEST_STOP_FIFO\"\n" + " else\n" + " read message < \"$LAUNCHER_TEST_STOP_FIFO\"\n" + " fi\n" + "fi\n") python_stub.chmod(0o755) openssl_stub = stub_bin / "openssl" openssl_stub.write_text("#!/bin/sh\nprintf '%064d\\n' 0\n") @@ -303,6 +312,7 @@ def _launcher_env(tmp_path: Path, home: str) -> dict: for name in _LAUNCHER_ENV_SCRUB: env.pop(name, None) env["PMI_RANK"] = "0" + env["LAUNCHER_TEST_STOP_FIFO"] = str(stop_fifo) env["HOME"] = home env["PATH"] = f"{stub_bin}{os.pathsep}{env['PATH']}" return env From 21b5622cd6db2c08cfe0bab349de9d0f12bf109a Mon Sep 17 00:00:00 2001 From: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> Date: Mon, 28 Sep 2026 13:22:45 -0700 Subject: [PATCH 2/3] [None][test] Cover nonroot failure with peers blocked in MPI collective Signed-off-by: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> --- tests/unittest/llmapi/_run_mpi_lifecycle_task.py | 9 +++++++-- tests/unittest/llmapi/test_mpi_lifecycle.py | 7 ++++--- tests/unittest/llmapi/test_mpi_server_lifecycle.py | 10 ++++++++++ 3 files changed, 21 insertions(+), 5 deletions(-) diff --git a/tests/unittest/llmapi/_run_mpi_lifecycle_task.py b/tests/unittest/llmapi/_run_mpi_lifecycle_task.py index 1004b3ff5206..06e4b0037804 100644 --- a/tests/unittest/llmapi/_run_mpi_lifecycle_task.py +++ b/tests/unittest/llmapi/_run_mpi_lifecycle_task.py @@ -52,10 +52,15 @@ def worker_task(directory_name: str, scenario: str, batch: int, size: int) -> tu if scenario == "async_drain" and batch == 0: _wait_for_markers(directory, ["engine-exiting"]) - if scenario in ("mixed_failure", "all_hang"): + if scenario in ("mixed_failure", "mixed_collective", "all_hang"): _wait_for_markers(directory, [f"started-{batch}-{peer}" for peer in range(size)]) if scenario == "mixed_failure" and rank % 2 == 0: raise RuntimeError("injected MPI lifecycle failure") + if scenario == "mixed_collective": + if rank == size - 1: + raise RuntimeError("injected MPI lifecycle failure") + MPI.COMM_WORLD.Barrier() + raise AssertionError("The collective completed without the failed rank") time.sleep(3600) raise AssertionError("The owner did not terminate the stuck worker world") @@ -102,7 +107,7 @@ def assert_results(response: object, batch: int) -> None: assert isinstance(response, list), response assert sorted(response) == [(rank, batch) for rank in range(args.ranks)], response - if args.scenario in ("mixed_failure", "all_hang"): + if args.scenario in ("mixed_failure", "mixed_collective", "all_hang"): client.submit(remote_task, str(directory), args.scenario, 0, args.ranks) _wait_for_markers(directory, [f"started-0-{rank}" for rank in range(args.ranks)]) _mark(directory, "workers-started") diff --git a/tests/unittest/llmapi/test_mpi_lifecycle.py b/tests/unittest/llmapi/test_mpi_lifecycle.py index 2a6e29ac7296..9b5ad3e40b23 100644 --- a/tests/unittest/llmapi/test_mpi_lifecycle.py +++ b/tests/unittest/llmapi/test_mpi_lifecycle.py @@ -77,6 +77,7 @@ def _cleanup_processes(identities: dict[int, float]) -> None: "scenario", [ "mixed_failure", + "mixed_collective", "all_hang", "all_return", "no_submission", @@ -213,16 +214,16 @@ def test_remote_mpi_worker_lifecycle(scenario: str, ranks: int, tmp_path: Path) assert (tmp_path / "engine-started.json").exists(), log assert "ZMQ thread safety violation" not in log, log - if scenario in ("mixed_failure", "all_hang"): + if scenario in ("mixed_failure", "mixed_collective", "all_hang"): assert process.returncode != 0, log assert (tmp_path / "workers-started.json").exists(), log - expected_marker = "error-observed" if scenario == "mixed_failure" else "engine-exiting" + expected_marker = "engine-exiting" if scenario == "all_hang" else "error-observed" marker = tmp_path / f"{expected_marker}.json" assert marker.exists(), log triggered = json.loads(marker.read_text())["monotonic"] # The stop helper imports TensorRT-LLM afresh after engine exit. The # mixed-failure path needs no new interpreter and has a tighter bound. - teardown_budget = 20 if scenario == "mixed_failure" else 120 + teardown_budget = 120 if scenario == "all_hang" else 20 assert exited - triggered < teardown_budget, f"Teardown exceeded its bounded budget:\n{log}" elif scenario == "no_submission": assert process.returncode == 3, log diff --git a/tests/unittest/llmapi/test_mpi_server_lifecycle.py b/tests/unittest/llmapi/test_mpi_server_lifecycle.py index 7f29520e5609..6bf901ffae5f 100644 --- a/tests/unittest/llmapi/test_mpi_server_lifecycle.py +++ b/tests/unittest/llmapi/test_mpi_server_lifecycle.py @@ -192,6 +192,16 @@ def release_first_batch_at_stop(queue): assert queue.closed +def test_stop_passes_remaining_deadline_to_final_shutdown(monkeypatch): + monkeypatch.setattr(mpi, "_mgmn_shutdown_grace_seconds", lambda: 10) + clock = iter([100, 103]) + monkeypatch.setattr(mpi.time, "monotonic", lambda: next(clock)) + queue = _Queue([mpi.RemoteTask(_task, (), {}), None]) + server = _server(queue, [[_future(1), _future(2)]]) + server.serve() + server._shutdown_session.assert_called_once_with(7) + + @pytest.mark.parametrize("failure", ["session", "executor", None]) def test_final_owner_shutdown_closes_shared_executor(monkeypatch, failure): class TestSession(mpi.MpiCommSession): From 8d8c2752d19d3b895aec8e099632979524ef02c7 Mon Sep 17 00:00:00 2001 From: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> Date: Mon, 28 Sep 2026 14:48:16 -0700 Subject: [PATCH 3/3] [None][fix] Preserve launcher stdin and follower signal behavior Signed-off-by: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> --- tensorrt_llm/llmapi/trtllm-llmapi-launch | 7 ++- .../llmapi/test_mpi_launcher_shutdown.py | 58 ++++++++++++++++++- tests/unittest/llmapi/test_mpi_lifecycle.py | 4 ++ 3 files changed, 65 insertions(+), 4 deletions(-) diff --git a/tensorrt_llm/llmapi/trtllm-llmapi-launch b/tensorrt_llm/llmapi/trtllm-llmapi-launch index e17b7c180c43..40960d31aa83 100755 --- a/tensorrt_llm/llmapi/trtllm-llmapi-launch +++ b/tensorrt_llm/llmapi/trtllm-llmapi-launch @@ -79,8 +79,6 @@ function cleanup_launcher { } trap cleanup_launcher EXIT -trap 'exit 130' INT -trap 'exit 143' TERM function use_unique_flashinfer_workspace { local temporary_root=${TMPDIR:-/tmp} @@ -215,6 +213,9 @@ unset TLLM_SPAWN_PROXY_PROCESS_IPC_HMAC_KEY # The launcher owns the workspace lock. Long-lived child commands close its # descriptor so detached descendants cannot keep a stale slot locked. if [ -z "$mpi_rank" ] || [ "$mpi_rank" -eq 0 ]; then + trap 'exit 130' INT + trap 'exit 143' TERM + if ! hmac_key=$(openssl rand -hex 32); then log_stderr "Failed to generate TLLM_SPAWN_PROXY_PROCESS_IPC_HMAC_KEY" exit 1 @@ -267,7 +268,7 @@ if [ -z "$mpi_rank" ] || [ "$mpi_rank" -eq 0 ]; then exec_child_command python3 -m tensorrt_llm.llmapi.mgmn_leader_node & server_pid=$! owned_child_pids+=("$server_pid") - exec_without_mpi_environment "${task_with_command[@]}" & + exec_without_mpi_environment "${task_with_command[@]}" < /dev/null & task_pid=$! owned_child_pids+=("$task_pid") log_stderr "Rank${mpi_rank} task PID: $task_pid; MPI Comm server PID: $server_pid" diff --git a/tests/unittest/llmapi/test_mpi_launcher_shutdown.py b/tests/unittest/llmapi/test_mpi_launcher_shutdown.py index 92a863fa2f6a..a03ce414263b 100644 --- a/tests/unittest/llmapi/test_mpi_launcher_shutdown.py +++ b/tests/unittest/llmapi/test_mpi_launcher_shutdown.py @@ -5,9 +5,11 @@ import json import os +import pty import signal import subprocess import sys +import time from pathlib import Path import pytest @@ -75,6 +77,10 @@ def hang(): signal.pause() if role == "server": + if mode == "follower": + (root / "worker-ready").touch() + while True: + signal.pause() send("server_ready") receive("task_ready") if mode == "server_exits": @@ -88,6 +94,8 @@ def hang(): send("stop_requested") else: receive("server_ready") + if mode == "read_stdin": + assert sys.stdin.read() == "" if mode == "server_exits": child = subprocess.Popen([ sys.executable, "-c", @@ -147,11 +155,23 @@ def _process_records(tmp_path: Path) -> dict[str, dict]: return {path.stem: json.loads(path.read_text()) for path in tmp_path.glob("*.json")} -def _run_launcher(tmp_path: Path, env: dict[str, str]) -> subprocess.CompletedProcess[str]: +def _run_launcher( + tmp_path: Path, env: dict[str, str], *, terminal_fd: int | None = None +) -> subprocess.CompletedProcess[str]: command = ["bash", str(_LAUNCHER), str(tmp_path / "bin" / "python3"), "task"] + if terminal_fd is not None: + command = [ + sys.executable, + "-c", + "import fcntl, os, sys, termios; " + "fcntl.ioctl(0, termios.TIOCSCTTY, 0); " + "os.execvp(sys.argv[1], sys.argv[1:])", + *command, + ] with subprocess.Popen( # nosec B603 command, env=env, + stdin=terminal_fd, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, @@ -228,6 +248,42 @@ def test_launcher_terminates_engine_and_child_when_server_exits(tmp_path: Path) _assert_exited(tmp_path) +def test_launcher_task_sees_eof_with_terminal_stdin(tmp_path: Path) -> None: + env = _launcher_env(tmp_path, "read_stdin") + master_fd, slave_fd = pty.openpty() + with os.fdopen(master_fd, "rb"), os.fdopen(slave_fd, "rb") as terminal: + result = _run_launcher(tmp_path, env, terminal_fd=terminal.fileno()) + assert result.returncode == 0, result.stderr + _assert_exited(tmp_path) + + +def test_follower_launcher_responds_to_sigterm(tmp_path: Path) -> None: + env = _launcher_env(tmp_path, "follower") + env["PMI_RANK"] = "1" + with subprocess.Popen( # nosec B603 + ["bash", str(_LAUNCHER), "/usr/bin/true"], + env=env, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + start_new_session=True, + ) as process: + try: + deadline = time.monotonic() + 5 + while not (tmp_path / "worker-ready").exists(): + assert process.poll() is None, "Follower launcher exited before worker startup" + assert time.monotonic() < deadline, "Follower worker did not start" + time.sleep(0.01) + process.send_signal(signal.SIGTERM) + assert process.wait(timeout=2) == -signal.SIGTERM + finally: + # The foreground worker belongs to this test's isolated session. + try: + os.killpg(process.pid, signal.SIGKILL) + except ProcessLookupError: + pass + process.wait(timeout=5) + + @pytest.mark.parametrize("value", ["0", "-1", "invalid", "1.5", "9999999"]) def test_launcher_rejects_invalid_shutdown_deadline(tmp_path: Path, value: str) -> None: env = _launcher_env(tmp_path, "clean") diff --git a/tests/unittest/llmapi/test_mpi_lifecycle.py b/tests/unittest/llmapi/test_mpi_lifecycle.py index 9b5ad3e40b23..8369cc98b9e1 100644 --- a/tests/unittest/llmapi/test_mpi_lifecycle.py +++ b/tests/unittest/llmapi/test_mpi_lifecycle.py @@ -242,3 +242,7 @@ def test_remote_mpi_worker_lifecycle(scenario: str, ranks: int, tmp_path: Path) for line in (tmp_path / f"events-{rank}.jsonl").read_text().splitlines() ] assert [event["batch"] for event in events] == list(range(batches)), events + if scenario in ("all_return", "reuse", "sync_recovery", "async_drain"): + assert (tmp_path / f"finished-{batches - 1}-{rank}.json").exists(), ( + f"Rank {rank} started but did not finish the final batch:\n{log}" + )