From 56aca6c98f3068dab611250a9588336f3720e870 Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Thu, 11 Jun 2026 20:48:00 +0800 Subject: [PATCH 01/16] Fault Tolerance Framework (#229) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * Fault Tolerant EP: Implement fault-report * Milestone 1 of Internal Process-level Fault Tolerance * Milestone 1 of Internal Process-level Fault Tolerance (#61) * feat(fault-tolerance): add class skeletons for fault tolerance Signed-off-by: fangyuchu * config: add configuration options for fault tolerance Signed-off-by: fangyuchu * 增加generate_identity和generate_identitys函数 Generate a unique identity for ZMQ ROUTER node * add service startup configuradtion fault report addr * add init WorkerGuard * add engine_core_cmd_addr、fault_report_addr、client_cmd_addr、engine_core_identitys in EngineZmqAddresses init engine_core_cmd_addr、fault_report_addr、client_cmd_addr in launch_core_engines func add _report_engine_dead func in CoreEngineProcManager * init ClientGuard init EngineZmqAddresses engine_core_identitys * init EngineCoreGuard * change generate_identitys to generate_identity_group * code typesetting is optimized * code typesetting is optimized * changed code format ensure every line < 88 chars * changed code format ensure every line < 88 chars fix error Value of type "dict[Any, Any] | None" is not indexable [index] * fix bug Error: vllm/v1/engine/utils.py:122:89: E501 Line too long (117 > 88) Error: vllm/v1/engine/utils.py:1059:9: F402 Import `uuid` from line 6 shadowed by loop variable * fix Error: vllm/v1/engine/utils.py:1045: error: Need type annotation for "uuids" (hint: "uuids: set[] = ...") [var-annotated] * fix error: Value of type "dict[Any, Any] | None" is not indexable [index] * fix error: Value of type "dict[Any, Any] | None" is not indexable [index] Signed-off-by: a798347923 <2645302020@qq.com> * add _send_msg in EngineCoreGuard Signed-off-by: a798347923 <2645302020@qq.com> * add import torch.cuda * add _recv_cmd function docstring that clearly explains the meaning of the return value. * changed recv_fault_msg to recv_msg add ClientGuard __init__ func parameter types * add engine monitor Signed-off-by: TianZhuo <2770730562@qq.com> * Delete requirements/test.txt~ Signed-off-by: a798347923 <39047817+a798347923@users.noreply.github.com> * Delete vllm/v1/engine/core_client.py~ Signed-off-by: a798347923 <39047817+a798347923@users.noreply.github.com> * simply _send_msg and _recv_cmd in EngineCoreGuard * simply recv_msg in ClientGuard * engine: add fault tolerance features for EngineCore. Signed-off-by: fangyuchu * engine: add timeout mechanism in retry. Signed-off-by: fangyuchu * add engine monitor * Delete vllm/v1/engine/exceptions.py~ Signed-off-by: 205150940 <112750056+205150940@users.noreply.github.com> * updata actor_index * updata enginedead flag * handle fault and report exception Signed-off-by: w00689259 * fix engine_actor * fix engine_actor fault_info * handle fault and report exception Signed-off-by: w00689259 * delete num_identity * changed try expect * fix debug error * fix one bug. Signed-off-by: fangyuchu * add fault_report_addr in FaultToleranceConfig * add handle fault&get_fault_info api Signed-off-by: w00689259 * remove fault_report_address in CoreEngineActorManager __init__ Signed-off-by: a798347923 <2645302020@qq.com> * ruff format Signed-off-by: a798347923 <2645302020@qq.com> * add handle fault&get_fault_info api Signed-off-by: w00689259 * fix one bug. Signed-off-by: fangyuchu * add fault_report_port in FaultToleranceConfig Signed-off-by: a798347923 <2645302020@qq.com> * add zmq_addr concatenate with fault_report_addr and fault_report_port Signed-off-by: a798347923 <2645302020@qq.com> * fault reporter bug fix Signed-off-by: w00689259 * fault reporter bug fix Signed-off-by: w00689259 * fault reporter bug fix Signed-off-by: w00689259 * fault reporter bug fix Signed-off-by: w00689259 * fault reporter bug fix Signed-off-by: w00689259 * fault reporter bug fix Signed-off-by: w00689259 * fix some bug * fault reporter bug fix Signed-off-by: w00689259 * fault reporter bug fix Signed-off-by: w00689259 * remove fault_report_addr in FaultToleranceConfig Signed-off-by: a798347923 <2645302020@qq.com> * refactor: relocate method serialization functions to serial_util.py Signed-off-by: fangyuchu * fix actor bug * fix actor bug * add engine_core_cmd_addr in FaultToleranceConfig Signed-off-by: a798347923 <2645302020@qq.com> * add and use _stop_worker_execution in EngineCoreGuard Signed-off-by: a798347923 <2645302020@qq.com> * add and use run in WorkerGuard Signed-off-by: a798347923 <2645302020@qq.com> * fix actor bug * fix bug * fix sentinel * fix bug vllm/v1/engine/core.py:847: error: Missing positional argument "tp_size" in call to "EngineCoreGuard" Signed-off-by: a798347923 <2645302020@qq.com> * fix bug error: Missing positional arguments "length", "byteorder" in call to "to_bytes" of "int" Signed-off-by: a798347923 <2645302020@qq.com> * fix bug in fault tolerance mode Signed-off-by: w00689259 * fix bug in fault tolerance mode Signed-off-by: w00689259 * change fault_report_port to internal_fault_report_port add external_fault_notify_port Signed-off-by: a798347923 <2645302020@qq.com> * change fault_report_port to internal_fault_report_port add external_fault_notify_port Signed-off-by: a798347923 <2645302020@qq.com> * add _recv_cmd func use deserialize_method_call and run_method in run func Signed-off-by: a798347923 <2645302020@qq.com> * Update core.py fix bug error: Need type annotation for "kwargs" (hint: "kwargs: dict[, ] = ...") Signed-off-by: a798347923 <39047817+a798347923@users.noreply.github.com> * add self.ctx.term() in shutdown() Signed-off-by: a798347923 <2645302020@qq.com> * changed import deserialize_method_call,serialize_method_call Signed-off-by: a798347923 <2645302020@qq.com> * changed init worker_guard in init_device Signed-off-by: a798347923 <2645302020@qq.com> * Update core.py add import serialize_method_call Signed-off-by: a798347923 <39047817+a798347923@users.noreply.github.com> * Update gpu_worker.py changed init WorkerGuard in init_device Signed-off-by: a798347923 <39047817+a798347923@users.noreply.github.com> * Update gpu_worker.py FIX BUG self.worker_guard: WorkerGuard|None = None Signed-off-by: a798347923 <39047817+a798347923@users.noreply.github.com> * Update gpu_worker.py fix bug error: Argument 1 to "deserialize_method_call" has incompatible type "str | None"; expected "str" [arg-type] Signed-off-by: a798347923 <39047817+a798347923@users.noreply.github.com> * Update gpu_worker.py ruff format Signed-off-by: a798347923 <39047817+a798347923@users.noreply.github.com> * Update core.py ruff-format Signed-off-by: a798347923 <39047817+a798347923@users.noreply.github.com> * actively send exception information Signed-off-by: w00689259 * actively send exception information Signed-off-by: w00689259 * actively send exception information Signed-off-by: w00689259 * change engine_core_cmd_addr(str) to engine_core_cmd_addrs(list[str]) in EngineZmqAddresses Signed-off-by: a798347923 <2645302020@qq.com> * change engine_core_cmd_addr(str) to engine_core_cmd_addrs(list[str]) in EngineZmqAddresses Signed-off-by: a798347923 <2645302020@qq.com> * Update utils.py delete engine_core_cmd_addr in EngineZmqAddresses Signed-off-by: a798347923 <39047817+a798347923@users.noreply.github.com> * Remove redundant configuration: fault-pub-port Signed-off-by: fangyuchu * Send pause instructions after receiving fault info in ClientGuard Signed-off-by: fangyuchu * change engine_core_guard_identities from dict[int, bytes] to list[bytes] Signed-off-by: a798347923 <2645302020@qq.com> * fix bug "only the worker guard of engine core 0 can receive messages sent from engine core guard Signed-off-by: a798347923 <2645302020@qq.com> * change local_rank to rank_in_group in WorkerGuard Signed-off-by: a798347923 <2645302020@qq.com> * changed del self.client_cmd_registry[int(unhealthy_engine.engine_id)] Signed-off-by: a798347923 <2645302020@qq.com> * add gloo communication timeout * fix some bug * add stateless_process_group gloo_comm_timeout * reconstruct fault receiver&fault handler Signed-off-by: w00689259 * fix some bug * reconstruct fault receiver&fault handler Signed-off-by: w00689259 * reconstruct fault receiver&fault handler Signed-off-by: w00689259 * fix return format Signed-off-by: w00689259 * fix return format Signed-off-by: w00689259 * fix return format Signed-off-by: w00689259 * add abort request * fix some bug * fix some bug * fix some bug * add dt for client guard Signed-off-by: w00689259 * add dt for client guard Signed-off-by: w00689259 * add dt for client guard Signed-off-by: w00689259 * Implementation of two types of pause: a soft one by using flag signals and a hard one by aborting nccl communicators. Signed-off-by: fangyuchu * Refine certain log forms and fix a minor bug in pause function. Signed-off-by: fangyuchu * Refactor and abstract the recv_msg logic in CG,ECG,WG. Signed-off-by: fangyuchu * Add and check method uuid when sending commands and receiving results. Signed-off-by: fangyuchu * Abstract the logic of sending instructions and waiting responses from FaultHandler Signed-off-by: fangyuchu * Add options in EngineCoreGuard to recv execution results from WorkerGuard Signed-off-by: fangyuchu * Support worker reinitialization after hard pause; add task queue in FaultHandler to ensure sequential task execution Signed-off-by: fangyuchu * resolve conflicts Signed-off-by: w00689259 * resolve conflicts Signed-off-by: w00689259 * resolve conflicts Signed-off-by: w00689259 * resolve conflicts Signed-off-by: w00689259 * resolve conflicts Signed-off-by: w00689259 * resolve conflicts Signed-off-by: w00689259 * add engine core ut Signed-off-by: w00689259 * add engine core ut Signed-off-by: w00689259 * Ensure WorkerGuard command execution returns result; fix missing set_device when TP>1 Signed-off-by: fangyuchu * rename& format logger Signed-off-by: w00689259 * rename& format logger Signed-off-by: w00689259 * feat(nccl): enable non-blocking NCCL communicators to support ncclCommAbort Signed-off-by: fangyuchu * reinit dp_group * fix bug * fix bug * fix bug * fix bug (#54) * Move requests to waiting queue instead of abandoing them directly. Signed-off-by: fangyuchu * add annotation Signed-off-by: w00689259 * fix typos Signed-off-by: fangyuchu --------- Signed-off-by: fangyuchu Signed-off-by: a798347923 <2645302020@qq.com> Signed-off-by: TianZhuo <2770730562@qq.com> Signed-off-by: a798347923 <39047817+a798347923@users.noreply.github.com> Signed-off-by: 205150940 <112750056+205150940@users.noreply.github.com> Signed-off-by: w00689259 Signed-off-by: zWaNg3 <37772915+zWaNg3@users.noreply.github.com> Co-authored-by: zWaNg3 <37772915+zWaNg3@users.noreply.github.com> Co-authored-by: a798347923 <2645302020@qq.com> Co-authored-by: TianZhuo <2770730562@qq.com> Co-authored-by: 205150940 <112750056+205150940@users.noreply.github.com> Co-authored-by: a798347923 <39047817+a798347923@users.noreply.github.com> Co-authored-by: w00689259 * Fix DT and zmq socket closing issues, updated names per feedback and reinitialize dp_group with new port Signed-off-by: fangyuchu * Improve documentation and logging in API server Signed-off-by: fangyuchu * Fix hanging issue in DT; fix hang when aborting communicators from Python side; use queue.Queue for engine_exception_q Signed-off-by: fangyuchu * Refactor fault tolerance modules by renaming classes to Sentinel and converting engine_registry to a dict Signed-off-by: fangyuchu * reject requests when engine is in fault status Signed-off-by: fangyuchu * clear batch_queue for async scheduling Signed-off-by: fangyuchu * Fix incorrect initialization of worker_cmd_socket in multi-node setups Signed-off-by: fangyuchu * Switch from field to Field Signed-off-by: fangyuchu * Unify start_engine_core_monitor in MPClient and CoreEngineProcManager to reduce duplication Signed-off-by: fangyuchu * refactor(Sentinel): Abstract and refactor class to standardize fault … (#84) * refactor(Sentinel): Abstract and refactor class to standardize fault tolerance logic Signed-off-by: w00689259 * refactor(Sentinel): Abstract and refactor class to standardize fault tolerance logic Signed-off-by: w00689259 * refactor(Sentinel): Abstract and refactor class to standardize fault tolerance logic Signed-off-by: w00689259 * refactor(Sentinel): Abstract and refactor class to standardize fault tolerance logic Signed-off-by: w00689259 * refactor(Sentinel): Abstract and refactor class to standardize fault tolerance logic Signed-off-by: w00689259 * refactor(Sentinel): Abstract and refactor class to standardize fault tolerance logic Signed-off-by: w00689259 * refactor(Sentinel): Abstract and refactor class to standardize fault tolerance logic Signed-off-by: w00689259 * refactor(Sentinel): Abstract and refactor class to standardize fault tolerance logic Signed-off-by: w00689259 * refactor(Sentinel): Abstract and refactor class to standardize fault tolerance logic Signed-off-by: w00689259 * refactor(Sentinel): Abstract and refactor class to standardize fault tolerance logic Signed-off-by: w00689259 * refactor(Sentinel): Abstract and refactor class to standardize fault tolerance logic Signed-off-by: w00689259 * refactor(Sentinel): Abstract and refactor class to standardize fault tolerance logic Signed-off-by: w00689259 --------- Signed-off-by: w00689259 Co-authored-by: w00689259 Signed-off-by: zWaNg3 <389750525@qq.com> * fix bug in tests Signed-off-by: w00689259 Signed-off-by: zWaNg3 <389750525@qq.com> * fix bug in tests Signed-off-by: w00689259 Signed-off-by: zWaNg3 <389750525@qq.com> * refactor: improve naming and add comments for readability Signed-off-by: fangyuchu * Pass fault_tolerance_config through process group creation for future extensibility Signed-off-by: fangyuchu * Switch to native preempt_request implementation Signed-off-by: fangyuchu * Rename base_sentinel.py Signed-off-by: fangyuchu * refactor(api_server): Split fault_tolerance interfaces into standalone files Signed-off-by: fangyuchu * Use zmq poll for socket receive in Sentinel DT to avoid hanging Signed-off-by: fangyuchu * Add shutdown-on-fault-tolerance-failure config option Signed-off-by: fangyuchu * ClientSentinel: add extra check to prevent repeated pause commands on error Signed-off-by: fangyuchu * feat(pause): apply pause with target index Signed-off-by: zWaNg3 <389750525@qq.com> * Add middleware for fault tolerance Signed-off-by: fangyuchu * fix engine_actor monitoring function bug Signed-off-by: TianZhuo <2770730562@qq.com> * fix engine_actor monitoring bug Signed-off-by: TianZhuo <2770730562@qq.com> * logger output format Signed-off-by: TianZhuo <2770730562@qq.com> * refactor(client_sentinel): support ClientSentinel-Client communication; refactor internal socket logic Signed-off-by: zWaNg3 <389750525@qq.com> * feat: add FaultToleranceRequest and FaultToleranceResult Signed-off-by: fangyuchu * feat: add EngineStatusType enum and support paused state Signed-off-by: fangyuchu * Unify the logic of engine monitor for engine process manager and engine actor manager Signed-off-by: fangyuchu * Fix the hanging issue of ClientSentinel in the shutdown Signed-off-by: fangyuchu * refactor(client_sentinel): rename process_ft_requests_loop function and run function Signed-off-by: zWaNg3 <389750525@qq.com> * Use VllmConfig as the input of Sentinel Modules Signed-off-by: fangyuchu * Remove redundant @dataclass from FaultToleranceConfig Signed-off-by: fangyuchu * Move hardcoded vllm_fault topic string into FaultToleranceConfig Signed-off-by: fangyuchu * Update corresponding tests to new ClientSentinel design. Signed-off-by: fangyuchu * Update engine core sentinel tests. Signed-off-by: fangyuchu * Fix incorrect device settings in the pause of worker sentinel. Signed-off-by: fangyuchu * Code cleanup and readability improvements Signed-off-by: fangyuchu * Simplify FaultInfo and improve the readability Signed-off-by: fangyuchu * Move sentinels into one file Signed-off-by: fangyuchu * Remove recv_router_dealer_message Signed-off-by: fangyuchu * Simplify the code in BaseSentinel Signed-off-by: fangyuchu * Simplify the code in EngineCoreSentinel Signed-off-by: fangyuchu * Introduce fault_tolerance utils and address dataclass Signed-off-by: fangyuchu * refactor: split different sentinels into separate files Signed-off-by: fangyuchu * refactor: split worker sentinel into v1/worker/sentinel for better plugin support and hardware adaptation Signed-off-by: fangyuchu * remove ThreadSafeDict Signed-off-by: fangyuchu * Simplify the communication between client, client sentinel and engine core sentinel (#137) * refactor(client_sentinel): use core_client input_socket to broadcast ft_requst Signed-off-by: zWaNg3 <389750525@qq.com> * refactor(client_sentinel): use core_client input_socket to broadcast ft_requst Signed-off-by: zWaNg3 <389750525@qq.com> * refactor(client_sentinel): use core_client input_socket to broadcast ft_requst Signed-off-by: zWaNg3 <389750525@qq.com> * add _send_utility_result in ClientSentinel Signed-off-by: fangyuchu * Processes fault-tolerant requests and forwards them to output. Signed-off-by: yzchang-plus <1078477584@qq.com> * replace uncertain code with TODO Signed-off-by: yzchang-plus <1078477584@qq.com> * add monitoring logic in client sentinel and implement thread-safe pause in monitoring. Signed-off-by: fangyuchu * refactor(client_sentinel): send ft request using input_address Signed-off-by: zWaNg3 <389750525@qq.com> * refactor(client_sentinel): return ft result to client Signed-off-by: zWaNg3 <389750525@qq.com> * Use call_utility_async for interactions between client, client_sentinel and engine core sentinel Signed-off-by: fangyuchu * Rename engine recovery timeout config Signed-off-by: fangyuchu * Remove upstream and downstream concept from the base sentinel Signed-off-by: fangyuchu * Support passing stateless dp port to retry Signed-off-by: fangyuchu * Improve the shutdown of client sentinel Signed-off-by: fangyuchu * Add try except for handle_fault in engine core Signed-off-by: fangyuchu --------- Signed-off-by: zWaNg3 <389750525@qq.com> Signed-off-by: fangyuchu Signed-off-by: yzchang-plus <1078477584@qq.com> Co-authored-by: zWaNg3 <389750525@qq.com> Co-authored-by: yzchang-plus <1078477584@qq.com> --------- Signed-off-by: fangyuchu Signed-off-by: a798347923 <2645302020@qq.com> Signed-off-by: TianZhuo <2770730562@qq.com> Signed-off-by: a798347923 <39047817+a798347923@users.noreply.github.com> Signed-off-by: 205150940 <112750056+205150940@users.noreply.github.com> Signed-off-by: w00689259 Signed-off-by: zWaNg3 <37772915+zWaNg3@users.noreply.github.com> Signed-off-by: zWaNg3 <389750525@qq.com> Signed-off-by: yzchang-plus <1078477584@qq.com> Co-authored-by: zWaNg3 <37772915+zWaNg3@users.noreply.github.com> Co-authored-by: a798347923 <2645302020@qq.com> Co-authored-by: TianZhuo <2770730562@qq.com> Co-authored-by: 205150940 <112750056+205150940@users.noreply.github.com> Co-authored-by: a798347923 <39047817+a798347923@users.noreply.github.com> Co-authored-by: w00689259 Co-authored-by: zWaNg3 <389750525@qq.com> Co-authored-by: yzchang-plus <1078477584@qq.com> Signed-off-by: fangyuchu * refactor(dt tests of sentinels): add dt tests for sentinels Signed-off-by: zWaNg3 <389750525@qq.com> Signed-off-by: fangyuchu * Remove torch.cuda API call (#148) * Remove torch.cuda API call Signed-off-by: fangyuchu * Remove unwanted shutdown Signed-off-by: fangyuchu --------- Signed-off-by: fangyuchu * Fault Tolerant EP: Implement fault-report Signed-off-by: fangyuchu * merge engine monitor codes Signed-off-by: fangyuchu * Move FT router attachment point and simplify FaultInfo initialization logic Signed-off-by: fangyuchu * Revise DT for Fault Report Signed-off-by: fangyuchu * Fix incorrect count of engine core index Signed-off-by: fangyuchu * Update engine process monitoring codes Signed-off-by: fangyuchu * [Bugfix] revise engine monitor logic on account of dead processes Signed-off-by: fangyuchu * Improve the format of the fault report json Signed-off-by: fangyuchu * Fix incorrect shutdown of engine manager Signed-off-by: fangyuchu * Avoid error logging in normal shutdown Signed-off-by: fangyuchu * handle zmq error Signed-off-by: fangyuchu --------- Signed-off-by: fangyuchu Signed-off-by: a798347923 <2645302020@qq.com> Signed-off-by: TianZhuo <2770730562@qq.com> Signed-off-by: a798347923 <39047817+a798347923@users.noreply.github.com> Signed-off-by: 205150940 <112750056+205150940@users.noreply.github.com> Signed-off-by: w00689259 Signed-off-by: zWaNg3 <37772915+zWaNg3@users.noreply.github.com> Signed-off-by: zWaNg3 <389750525@qq.com> Signed-off-by: yzchang-plus <1078477584@qq.com> Co-authored-by: zWaNg3 <37772915+zWaNg3@users.noreply.github.com> Co-authored-by: a798347923 <2645302020@qq.com> Co-authored-by: TianZhuo <2770730562@qq.com> Co-authored-by: 205150940 <112750056+205150940@users.noreply.github.com> Co-authored-by: a798347923 <39047817+a798347923@users.noreply.github.com> Co-authored-by: w00689259 Co-authored-by: zWaNg3 <389750525@qq.com> Co-authored-by: yzchang-plus <1078477584@qq.com> Signed-off-by: fangyuchu * Support elastic ep Signed-off-by: fangyuchu * Clean notify_engine_down logic Signed-off-by: fangyuchu * refactor: Refactor Sentinel class and move fault_tolerance_config from VllmConfig to ParallelConfig Signed-off-by: zWaNg3 <389750525@qq.com> * chore: remove Chinese comments Signed-off-by: Jade Zheng * refactor: Refactor Sentinel class and move fault_tolerance_config from VllmConfig to ParallelConfig Signed-off-by: zWaNg3 <389750525@qq.com> * set engine_recovery_timeout_sec as attr of engine core sentinel Signed-off-by: fangyuchu * Improve the comments Signed-off-by: fangyuchu * Fix DT Signed-off-by: fangyuchu * refactor: optimize sentinel logging,and refactor ctx initialization Signed-off-by: zWaNg3 <389750525@qq.com> * fix: fix fault_tolerance_addresses overwrite by rank0 in external/hybrid mode Signed-off-by: zWaNg3 <389750525@qq.com> * Rename func in client sentinel. Signed-off-by: fangyuchu * enable fault tolerance if ft config is provided Signed-off-by: fangyuchu * Add comments; change engine_id in FaultInfo to int type Signed-off-by: zWaNg3 <389750525@qq.com> * Add schema when reporting engine status and DT for engine status enum Signed-off-by: fangyuchu * Remove monitoring thread from EngineCoreSentinel and use weakref for host in sentinels Signed-off-by: fangyuchu * Add fault tolerance instruction workflow and pause operation Signed-off-by: fangyuchu * Set gloo timeout seconds in parallel config and pass directly to create process group Signed-off-by: fangyuchu * Create worker sentinel through collective_rpc Signed-off-by: fangyuchu * fix: Fix DT test cases and add gloo_timeout_seconds validation Signed-off-by: zWaNg3 <389750525@qq.com> * Simplify pause logic in model runner and check mask for ft backend in AsyncGPUModelRunnerOutput Signed-off-by: fangyuchu * fix dt Signed-off-by: fangyuchu * Clean inputs for sentinels Signed-off-by: fangyuchu * Remove duplicated EngineLoopPausedError msg. Signed-off-by: fangyuchu * add new engine status: hang to indicate the responsiveness of workers Signed-off-by: fangyuchu * feat(fault-tolerance): add retry mechanism to the fault tolerance framework Signed-off-by: fangyuchu * add hung state for engine and use host in sentinels in the recovery Signed-off-by: fangyuchu * Fix bugs in retry Signed-off-by: fangyuchu * test simplification Signed-off-by: fangyuchu * test simplification 2 Signed-off-by: fangyuchu * test simplification 3 Signed-off-by: fangyuchu * remove status collection from client side Signed-off-by: fangyuchu * test simplification 4 Signed-off-by: fangyuchu * test simplification 4 Signed-off-by: fangyuchu * set vllm_config context in retry and add handle_command in worker sentinel Signed-off-by: fangyuchu * Abort in-flight requests on FT fault Signed-off-by: fangyuchu * Add fault detection and buffer cleanup for deepep_ll and nixl_ep Signed-off-by: fangyuchu * Fix buffer cleanup for deepep_ll and nixl_ep all2all managers - Add synchronize after low_latency_clean_mask_buffer in DeepEPLL - Simplify NixlEP buffer cleanup: zero entire local buffer instead of using parameterized clean_buffer with cached dimensions - Add synchronize barriers after buffer/mask cleanup operations Signed-off-by: fangyuchu * Simplify fault tolerance code Signed-off-by: fangyuchu * simplify all2all ft implementations Signed-off-by: fangyuchu * fix precommit and fix mask bugs Signed-off-by: fangyuchu * unify get_status method and solve race condition in reinit dp group Signed-off-by: fangyuchu * remove unneccessary drain executor response mqs Signed-off-by: fangyuchu --------- Signed-off-by: fangyuchu Signed-off-by: a798347923 <2645302020@qq.com> Signed-off-by: TianZhuo <2770730562@qq.com> Signed-off-by: a798347923 <39047817+a798347923@users.noreply.github.com> Signed-off-by: 205150940 <112750056+205150940@users.noreply.github.com> Signed-off-by: w00689259 Signed-off-by: zWaNg3 <37772915+zWaNg3@users.noreply.github.com> Signed-off-by: zWaNg3 <389750525@qq.com> Signed-off-by: yzchang-plus <1078477584@qq.com> Signed-off-by: Jade Zheng Co-authored-by: zWaNg3 <37772915+zWaNg3@users.noreply.github.com> Co-authored-by: a798347923 <2645302020@qq.com> Co-authored-by: TianZhuo <2770730562@qq.com> Co-authored-by: 205150940 <112750056+205150940@users.noreply.github.com> Co-authored-by: a798347923 <39047817+a798347923@users.noreply.github.com> Co-authored-by: w00689259 Co-authored-by: zWaNg3 <389750525@qq.com> Co-authored-by: yzchang-plus <1078477584@qq.com> Co-authored-by: Jade Zheng Signed-off-by: fangyuchu --- vllm/config/__init__.py | 3 + vllm/config/fault_tolerance.py | 18 ++ vllm/config/parallel.py | 12 ++ .../device_communicators/all2all.py | 25 ++- .../base_device_communicator.py | 4 + vllm/engine/arg_utils.py | 38 ++++ vllm/engine/protocol.py | 11 ++ vllm/entrypoints/openai/api_server.py | 7 + .../serve/fault_tolerance/__init__.py | 0 .../serve/fault_tolerance/api_router.py | 77 ++++++++ vllm/v1/engine/__init__.py | 6 + vllm/v1/engine/async_llm.py | 10 + vllm/v1/engine/core.py | 25 +++ vllm/v1/engine/core_client.py | 31 +++ vllm/v1/engine/utils.py | 3 - vllm/v1/fault_tolerance/__init__.py | 8 + .../fault_tolerance/engine_core_sentinel.py | 177 ++++++++++++++++++ vllm/v1/fault_tolerance/utils.py | 17 ++ vllm/v1/worker/gpu_worker.py | 9 +- vllm/v1/worker/sentinel/__init__.py | 0 .../v1/worker/sentinel/gpu_worker_sentinel.py | 76 ++++++++ 21 files changed, 552 insertions(+), 5 deletions(-) create mode 100644 vllm/config/fault_tolerance.py create mode 100644 vllm/entrypoints/serve/fault_tolerance/__init__.py create mode 100644 vllm/entrypoints/serve/fault_tolerance/api_router.py create mode 100644 vllm/v1/fault_tolerance/__init__.py create mode 100644 vllm/v1/fault_tolerance/engine_core_sentinel.py create mode 100644 vllm/v1/fault_tolerance/utils.py create mode 100644 vllm/v1/worker/sentinel/__init__.py create mode 100644 vllm/v1/worker/sentinel/gpu_worker_sentinel.py diff --git a/vllm/config/__init__.py b/vllm/config/__init__.py index 6070a3f82382..24a8b5a31713 100644 --- a/vllm/config/__init__.py +++ b/vllm/config/__init__.py @@ -13,6 +13,7 @@ from vllm.config.diffusion import DiffusionConfig from vllm.config.ec_manager_config import EncoderCacheManagerConfig from vllm.config.ec_transfer import ECTransferConfig +from vllm.config.fault_tolerance import FaultToleranceConfig from vllm.config.kernel import KernelConfig from vllm.config.kv_events import KVEventsConfig from vllm.config.kv_transfer import KVTransferConfig @@ -124,6 +125,8 @@ "StructuredOutputsConfig", # From vllm.config.profiler "ProfilerConfig", + # From vllm.config.fault_tolerance + "FaultToleranceConfig", # From vllm.config.utils "ConfigType", "SupportsMetricsInfo", diff --git a/vllm/config/fault_tolerance.py b/vllm/config/fault_tolerance.py new file mode 100644 index 000000000000..d4b41c9c1a3c --- /dev/null +++ b/vllm/config/fault_tolerance.py @@ -0,0 +1,18 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + + +from vllm.config.utils import config + + +@config +class FaultToleranceConfig: + """Configuration for fault tolerance.""" + + engine_recovery_timeout_sec: int = 120 + """Timeout (in seconds) to wait for error handling instructions + before raising an exception. If the EngineCore encounters an + error, it waits up to this many seconds for instructions on how + to handle the error. If no instructions are received within this + time, the original error is raised. + """ diff --git a/vllm/config/parallel.py b/vllm/config/parallel.py index ce038a99ea9a..85adb182c06a 100644 --- a/vllm/config/parallel.py +++ b/vllm/config/parallel.py @@ -13,6 +13,7 @@ from typing_extensions import Self import vllm.envs as envs +from vllm.config.fault_tolerance import FaultToleranceConfig from vllm.config.utils import config from vllm.logger import init_logger from vllm.platforms import current_platform @@ -22,6 +23,7 @@ from ray.runtime_env import RuntimeEnv from ray.util.placement_group import PlacementGroup + from vllm.config.fault_tolerance import FaultToleranceConfig from vllm.v1.executor import Executor else: RuntimeEnv = Any @@ -393,6 +395,16 @@ class is dynamically inherited by the worker class. This is used to inject should only be set by API server scale-out. """ + enable_fault_tolerance: bool = False + """Enable fault tolerance for detailed error recovery, + such as scaling down fault DPEngineCore. + """ + + fault_tolerance_config: FaultToleranceConfig = Field( + default_factory=FaultToleranceConfig + ) + """The configurations for fault tolerance.""" + @field_validator("disable_nccl_for_dp_synchronization", mode="wrap") @classmethod def _skip_none_validation(cls, value: Any, handler: Callable) -> Any: diff --git a/vllm/distributed/device_communicators/all2all.py b/vllm/distributed/device_communicators/all2all.py index 679764f6a82e..f96ccb6ccc2d 100644 --- a/vllm/distributed/device_communicators/all2all.py +++ b/vllm/distributed/device_communicators/all2all.py @@ -8,6 +8,7 @@ import torch.distributed as dist import vllm.envs as envs +from vllm.config import get_current_vllm_config from vllm.distributed import get_dp_group, get_ep_group, get_pcp_group from vllm.distributed.utils import StatelessProcessGroup from vllm.forward_context import get_forward_context @@ -278,7 +279,9 @@ class DeepEPLLAll2AllManager(DeepEPAll2AllManagerBase): def __init__(self, cpu_group, tcp_store_group=None): super().__init__(cpu_group, tcp_store_group) - self.support_fault_tolerance = False # TODO: set to True when FT is supported. + self.support_fault_tolerance = ( + get_current_vllm_config().parallel_config.enable_fault_tolerance + ) def _make_all2all_kwargs( self, @@ -360,6 +363,16 @@ def query_fault(self) -> torch.Tensor: has_fault = (current != DeepEPLLAll2AllManager._last_mask).any() return has_fault + def clean_buffers(self) -> None: + buf = DeepEPLLAll2AllManager._buffer + if buf is None: + return + buf.get_local_buffer_tensor(dtype=torch.int8, use_rdma_buffer=True).zero_() + torch.accelerator.synchronize() + buf.low_latency_clean_mask_buffer() + torch.accelerator.synchronize() + DeepEPLLAll2AllManager._last_mask = None + @dataclass class _NixlEPBufferState: @@ -565,6 +578,16 @@ def query_fault(self) -> torch.Tensor: has_fault = (current != last).any() return has_fault + def clean_buffers(self) -> None: + if NixlEPAll2AllManager._buffer is None: + return + state = NixlEPAll2AllManager._buffer + state.buffer.get_local_buffer_tensor(dtype=torch.int8).zero_() + torch.accelerator.synchronize() + state.buffer.clean_mask_buffer() + torch.accelerator.synchronize() + NixlEPAll2AllManager._last_active_mask = None + class FlashInferNVLinkTwoSidedManager(All2AllManagerBase): """ diff --git a/vllm/distributed/device_communicators/base_device_communicator.py b/vllm/distributed/device_communicators/base_device_communicator.py index 70f1fb5d62c5..f526ba6314f1 100644 --- a/vllm/distributed/device_communicators/base_device_communicator.py +++ b/vllm/distributed/device_communicators/base_device_communicator.py @@ -111,6 +111,10 @@ def query_fault(self) -> torch.Tensor: """Returns has_fault scalar.""" raise NotImplementedError + def clean_buffers(self) -> None: + """Clean RDMA buffers and mask state during FT retry.""" + raise NotImplementedError + def set_num_sms(self, num_sms: int): pass diff --git a/vllm/engine/arg_utils.py b/vllm/engine/arg_utils.py index a14de27190ec..82294132eb64 100644 --- a/vllm/engine/arg_utils.py +++ b/vllm/engine/arg_utils.py @@ -42,6 +42,7 @@ DiffusionConfig, ECTransferConfig, EPLBConfig, + FaultToleranceConfig, KernelConfig, KVEventsConfig, KVTransferConfig, @@ -722,6 +723,12 @@ class EngineArgs: optimization_level: OptimizationLevel = VllmConfig.optimization_level performance_mode: PerformanceMode = VllmConfig.performance_mode + # fault tolerance fields (`None` means not explicitly provided). + fault_tolerance_config: FaultToleranceConfig | None = get_field( + ParallelConfig, "fault_tolerance_config" + ) + enable_fault_tolerance: bool = ParallelConfig.enable_fault_tolerance + kv_offloading_size: float | None = CacheConfig.kv_offloading_size kv_offloading_backend: KVOffloadingBackend = CacheConfig.kv_offloading_backend tokens_only: bool = False @@ -754,6 +761,10 @@ def __post_init__(self): self.weight_transfer_config = WeightTransferConfig( **self.weight_transfer_config ) + if isinstance(self.fault_tolerance_config, dict): + self.fault_tolerance_config = FaultToleranceConfig( + **self.fault_tolerance_config + ) if isinstance(self.ir_op_priority, dict): self.ir_op_priority = IrOpPriorityConfig(**self.ir_op_priority) @@ -1151,6 +1162,16 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: parallel_group.add_argument( "--worker-extension-cls", **parallel_kwargs["worker_extension_cls"] ) + parallel_group.add_argument( + "--enable-fault-tolerance", **parallel_kwargs["enable_fault_tolerance"] + ) + parallel_group.add_argument( + "--fault-tolerance-config", + **{ + **parallel_kwargs["fault_tolerance_config"], + "default": None, + }, + ) # KV cache arguments cache_kwargs = get_kwargs(CacheConfig) @@ -1617,6 +1638,13 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: def from_cli_args(cls, args: argparse.Namespace): # Get the list of attributes of this dataclass. attrs = [attr.name for attr in dataclasses.fields(cls)] + + # If --fault-tolerance-config is provided, enable fault tolerance by default. + if args.fault_tolerance_config is not None: + args.enable_fault_tolerance = True + if args.enable_fault_tolerance and args.fault_tolerance_config is None: + args.fault_tolerance_config = FaultToleranceConfig() + # Set the attributes from the parsed arguments. engine_args = cls( **{attr: getattr(args, attr) for attr in attrs if hasattr(args, attr)} @@ -2022,6 +2050,12 @@ def create_engine_config( data_parallel_external_lb = ( self.data_parallel_external_lb or self.data_parallel_rank is not None ) + if self.enable_fault_tolerance and not data_parallel_external_lb: + raise ValueError( + "Fault tolerance requires external load balancer mode " + "(--data-parallel-external-lb or --data-parallel-rank). " + "Internal LB mode is not supported." + ) if ( self.data_parallel_size > 1 and data_parallel_external_lb @@ -2179,6 +2213,10 @@ def create_engine_config( _api_process_count=self._api_process_count, _api_process_rank=self._api_process_rank, assigned_physical_gpu_ids=self._resolve_device_ids(), + enable_fault_tolerance=self.enable_fault_tolerance, + fault_tolerance_config=( + self.fault_tolerance_config or FaultToleranceConfig() + ), numa_bind=self.numa_bind, numa_bind_nodes=self.numa_bind_nodes, numa_bind_cpus=self.numa_bind_cpus, diff --git a/vllm/engine/protocol.py b/vllm/engine/protocol.py index c54123bea9e5..ef3be178ac8d 100644 --- a/vllm/engine/protocol.py +++ b/vllm/engine/protocol.py @@ -20,6 +20,7 @@ from vllm.tasks import SupportedTask from vllm.v1.engine import EngineCoreRequest from vllm.v1.engine.input_processor import InputProcessor +from vllm.v1.fault_tolerance.utils import FaultToleranceRequest, FaultToleranceResult if TYPE_CHECKING: from vllm.v1.engine import PauseMode @@ -234,6 +235,16 @@ async def collective_rpc( """Perform a collective RPC call to the given path.""" raise NotImplementedError + async def handle_fault( + self, fault_tolerance_request: FaultToleranceRequest + ) -> FaultToleranceResult: + """send fault tolerance instruction to the engine""" + raise NotImplementedError + + async def get_status(self): + """Get fault tolerance status of all engines.""" + raise NotImplementedError + async def get_supported_tasks(self) -> tuple[SupportedTask, ...]: """Get supported tasks""" raise NotImplementedError diff --git a/vllm/entrypoints/openai/api_server.py b/vllm/entrypoints/openai/api_server.py index 59c7ee84caef..9103dd7fae95 100644 --- a/vllm/entrypoints/openai/api_server.py +++ b/vllm/entrypoints/openai/api_server.py @@ -269,6 +269,13 @@ def build_app( register_pooling_api_routers(app, supported_tasks, model_config) + if args.enable_fault_tolerance: + from vllm.entrypoints.serve.fault_tolerance.api_router import ( + register_fault_tolerance_api_router, + ) + + register_fault_tolerance_api_router(app) + # Endpoint plugins are attached last so their routes are registered after all core # routers. This runs even for the CPU only render server. A plugin eligible for # the `render` task still gets its routes registered. It receives diff --git a/vllm/entrypoints/serve/fault_tolerance/__init__.py b/vllm/entrypoints/serve/fault_tolerance/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/vllm/entrypoints/serve/fault_tolerance/api_router.py b/vllm/entrypoints/serve/fault_tolerance/api_router.py new file mode 100644 index 000000000000..cae5742365eb --- /dev/null +++ b/vllm/entrypoints/serve/fault_tolerance/api_router.py @@ -0,0 +1,77 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import json +import uuid +from http import HTTPStatus + +from fastapi import APIRouter, Depends, FastAPI, HTTPException, Request +from fastapi.responses import JSONResponse + +from vllm.engine.protocol import EngineClient +from vllm.entrypoints.openai.engine.protocol import ErrorResponse +from vllm.entrypoints.serve.utils.api_utils import validate_json_request +from vllm.logger import init_logger +from vllm.v1.fault_tolerance.utils import FaultToleranceRequest + +logger = init_logger(__name__) + +router = APIRouter() + +_ALLOWED_INSTRUCTIONS = {"retry"} + + +def _validate_payload(body: dict) -> tuple[str, dict]: + instruction = body.get("instruction") + params = body.get("params") + if not instruction or not isinstance(params, dict): + raise HTTPException(400, "'instruction' and 'params' are required.") + if instruction not in _ALLOWED_INSTRUCTIONS: + raise HTTPException(400, f"Invalid instruction: '{instruction}'.") + if "timeout" not in params or not isinstance(params["timeout"], (int, float)): + raise HTTPException(400, "Missing or invalid 'timeout' parameter.") + return instruction, params + + +@router.post( + "/fault_tolerance/apply", + dependencies=[Depends(validate_json_request)], + responses={ + HTTPStatus.OK.value: {"model": dict}, + HTTPStatus.BAD_REQUEST.value: {"model": ErrorResponse}, + HTTPStatus.REQUEST_TIMEOUT.value: {"model": ErrorResponse}, + HTTPStatus.INTERNAL_SERVER_ERROR.value: {"model": ErrorResponse}, + }, +) +async def process_fault_tolerance_instruction(raw_request: Request): + try: + body = await raw_request.json() + except json.JSONDecodeError as e: + raise HTTPException(400, "Invalid JSON format") from e + + instruction, params = _validate_payload(body) + ft_request = FaultToleranceRequest( + instruction=instruction, + params=params, + request_id=str(uuid.uuid4()), + ) + + client: EngineClient = raw_request.app.state.engine_client + try: + ft_result = await client.handle_fault(ft_request) + except Exception as e: + logger.error("Failed to handle fault: %s", e) + raise HTTPException(500, "Failed to handle fault.") from e + + if ft_result.success: + return JSONResponse({"message": "Instruction executed successfully."}) + raise HTTPException(500, f"Instruction failed: {ft_result.reason}") + + +@router.get("/fault_tolerance/status") +async def get_status(raw_request: Request): + client: EngineClient = raw_request.app.state.engine_client + return JSONResponse(content=await client.get_status()) + + +def register_fault_tolerance_api_router(app: FastAPI): + app.include_router(router) diff --git a/vllm/v1/engine/__init__.py b/vllm/v1/engine/__init__.py index 919402a16ab0..3a916af832f4 100644 --- a/vllm/v1/engine/__init__.py +++ b/vllm/v1/engine/__init__.py @@ -282,3 +282,9 @@ class ReconfigureRankType(enum.IntEnum): KEEP_CURRENT_RANK = -1 SHUTDOWN_CURRENT_RANK = -2 + + +class EngineStatusType(enum.IntEnum): + HEALTHY = 0 + DEAD = 1 + UNHEALTHY = 2 diff --git a/vllm/v1/engine/async_llm.py b/vllm/v1/engine/async_llm.py index 93e02abf7479..f1e7132339c3 100644 --- a/vllm/v1/engine/async_llm.py +++ b/vllm/v1/engine/async_llm.py @@ -44,6 +44,7 @@ from vllm.v1.engine.output_processor import OutputProcessor, RequestOutputCollector from vllm.v1.engine.parallel_sampling import ParentRequest from vllm.v1.executor import Executor +from vllm.v1.fault_tolerance.utils import FaultToleranceRequest, FaultToleranceResult from vllm.v1.metrics.loggers import ( StatLoggerFactory, StatLoggerManager, @@ -1041,6 +1042,15 @@ async def scale_elastic_ep( finally: set_scaling_elastic_ep(False) + async def handle_fault( + self, fault_tolerance_request: FaultToleranceRequest + ) -> FaultToleranceResult: + """send fault tolerance instruction to the engine""" + return await self.engine_core.handle_fault(fault_tolerance_request) + + async def get_status(self): + return await self.engine_core.get_status() + @property def is_running(self) -> bool: # Is None before the loop is started. diff --git a/vllm/v1/engine/core.py b/vllm/v1/engine/core.py index 476f53d46115..8a62c8b2b1be 100644 --- a/vllm/v1/engine/core.py +++ b/vllm/v1/engine/core.py @@ -78,6 +78,11 @@ get_physical_gpu_ids_for_local_dp_rank, ) from vllm.v1.executor import Executor +from vllm.v1.fault_tolerance.engine_core_sentinel import ( + FT_UTILITY_METHOD, + EngineCoreSentinel, + fault_tolerant_wrapper, +) from vllm.v1.kv_cache_interface import KVCacheConfig, get_kv_cache_spec_kind from vllm.v1.metrics.stats import SchedulerIterationDetails, SchedulerStats from vllm.v1.outputs import ModelRunnerOutput @@ -1070,6 +1075,16 @@ def __init__( internal_dp_balancing, ) + # Initialize fault tolerance settings. + self.enable_fault_tolerance = ( + vllm_config.parallel_config.enable_fault_tolerance + ) + if self.enable_fault_tolerance: + self.ft_sentinel = EngineCoreSentinel( + engine=self, + parallel_config=vllm_config.parallel_config, + ) + # Background Threads and Queues for IO. These enable us to # overlap ZMQ socket IO with GPU since they release the GIL, # and to overlap some serialization/deserialization with the @@ -1355,6 +1370,7 @@ def is_running(self) -> bool: """Returns true if shutdown has not been requested.""" return self.shutdown_state == EngineShutdownState.RUNNING + @fault_tolerant_wrapper def run_busy_loop(self): """Core busy loop of the EngineCore.""" while self._handle_shutdown(): @@ -1672,6 +1688,14 @@ def process_input_sockets( except Exception: self._handle_request_preproc_error(req) continue + elif request_type == EngineCoreRequestType.UTILITY: + request = generic_decoder.decode(data_frames) + client_idx, call_id, method, args = request + if method == FT_UTILITY_METHOD: + self.ft_sentinel.handle_command( + client_idx, call_id, args[0] + ) + continue else: request = generic_decoder.decode(data_frames) @@ -2021,6 +2045,7 @@ def _should_throttle_prefills(self) -> bool: and self.step_counter % self.prefill_schedule_interval != 0 ) + @fault_tolerant_wrapper def run_busy_loop(self): """Core busy loop of the EngineCore for data parallel case.""" diff --git a/vllm/v1/engine/core_client.py b/vllm/v1/engine/core_client.py index bcb441e7564a..2f0c0910f071 100644 --- a/vllm/v1/engine/core_client.py +++ b/vllm/v1/engine/core_client.py @@ -56,6 +56,11 @@ launch_core_engines, ) from vllm.v1.executor import Executor +from vllm.v1.fault_tolerance.engine_core_sentinel import FT_UTILITY_METHOD +from vllm.v1.fault_tolerance.utils import ( + FaultToleranceRequest, + FaultToleranceResult, +) from vllm.v1.pool.late_interaction import get_late_interaction_engine_index from vllm.v1.serial_utils import MsgpackDecoder, MsgpackEncoder, bytestr @@ -272,6 +277,14 @@ async def collective_rpc_async( ) -> list[_R]: raise NotImplementedError + async def handle_fault( + self, fault_tolerance_request: FaultToleranceRequest + ) -> FaultToleranceResult: + raise NotImplementedError + + async def get_status(self): + raise NotImplementedError + class InprocClient(EngineCoreClient): """ @@ -1196,6 +1209,24 @@ async def collective_rpc_async( "collective_rpc", method, timeout, args, kwargs ) + async def handle_fault( + self, ft_request: FaultToleranceRequest + ) -> FaultToleranceResult: + res = await self.call_utility_async(FT_UTILITY_METHOD, ft_request) + result = FaultToleranceResult(**res) + return result + + async def get_status(self): + ft_request = FaultToleranceRequest(instruction="status", params={}) + res = await self.call_utility_async(FT_UTILITY_METHOD, ft_request) + return { + "schema_version": 1, + "total_engines": len(self.engine_ranks_managed), + "engines": [ + {"id": res["engine_id"], "status": res["status"]}, + ], + } + class DPAsyncMPClient(AsyncMPClient): """Asyncio-compatible client for multi-proc, multi-engine (data parallel) diff --git a/vllm/v1/engine/utils.py b/vllm/v1/engine/utils.py index 093f065475ab..db1896b09468 100644 --- a/vllm/v1/engine/utils.py +++ b/vllm/v1/engine/utils.py @@ -234,9 +234,6 @@ def monitor_engine_liveness(self) -> None: if exitcode != 0 and not self.manager_stopped.is_set(): self.failed_proc_name = proc.name if died_sentinels: - # Any engine exit currently triggers a shutdown. Future - # work (e.g., Elastic and fault-tolerant EP) will add finer-grained - # handling for different exit scenarios. break self.shutdown() diff --git a/vllm/v1/fault_tolerance/__init__.py b/vllm/v1/fault_tolerance/__init__.py new file mode 100644 index 000000000000..ee54b4eb2e9c --- /dev/null +++ b/vllm/v1/fault_tolerance/__init__.py @@ -0,0 +1,8 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from .engine_core_sentinel import EngineCoreSentinel, fault_tolerant_wrapper + +__all__ = [ + "EngineCoreSentinel", + "fault_tolerant_wrapper", +] diff --git a/vllm/v1/fault_tolerance/engine_core_sentinel.py b/vllm/v1/fault_tolerance/engine_core_sentinel.py new file mode 100644 index 000000000000..5b6fdd568b74 --- /dev/null +++ b/vllm/v1/fault_tolerance/engine_core_sentinel.py @@ -0,0 +1,177 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""EngineCoreSentinel and fault_tolerant_wrapper for the engine core.""" + +import threading +from collections.abc import Callable +from typing import TYPE_CHECKING + +from vllm.config import set_current_vllm_config +from vllm.distributed import stateless_destroy_torch_distributed_process_group +from vllm.distributed.utils import stateless_init_torch_distributed_process_group +from vllm.logger import init_logger +from vllm.utils.network_utils import get_open_port +from vllm.v1.engine import EngineCoreOutputs, EngineStatusType, UtilityOutput +from vllm.v1.fault_tolerance.utils import FaultToleranceRequest +from vllm.v1.request import RequestStatus +from vllm.v1.serial_utils import UtilityResult, run_method + +if TYPE_CHECKING: + from vllm.v1.engine.core import EngineCoreProc + +logger = init_logger(__name__) + +FT_UTILITY_METHOD = "handle_fault_tolerance" + + +class EngineCoreSentinel: + """Manages fault tolerance state for a single engine core.""" + + def __init__(self, engine: "EngineCoreProc", parallel_config): + self.engine = engine + self.engine_index = engine.engine_index + self.parallel_config = parallel_config + ft_config = parallel_config.fault_tolerance_config + self.engine_recovery_timeout_sec = ft_config.engine_recovery_timeout_sec + + self.resumed = threading.Event() + self.resumed.set() + self.status_type = EngineStatusType.HEALTHY + self._dp_reinit_epoch = 0 + + # ------------------------------------------------------------------ + # Command dispatch (called from process_input_sockets thread) + # ------------------------------------------------------------------ + + def handle_command(self, client_idx: int, call_id: int, ft_args: dict): + """Dispatch an FT command by instruction name and enqueue result.""" + ft_request = FaultToleranceRequest(**ft_args) + try: + result = run_method(self, ft_request.instruction, (ft_request,), {}) + except Exception as e: + logger.exception("[FT] Instruction '%s' failed", ft_request.instruction) + result = { + "request_id": ft_request.request_id, + "success": False, + "reason": str(e), + } + + uo = UtilityOutput(call_id) + uo.result = UtilityResult(result) + self.engine.output_queue.put_nowait( + (client_idx, EngineCoreOutputs(utility_output=uo)) + ) + + # ------------------------------------------------------------------ + # Fault handling (called by wrapper, runs in busy-loop thread) + # ------------------------------------------------------------------ + + def on_fault(self, exc: Exception): + """Called by the wrapper when the busy loop raises an exception.""" + self.resumed.clear() + logger.warning( + "[FT] Busy loop raised %s. Waiting for recovery.", type(exc).__name__ + ) + + engine = self.engine + aborted = engine.scheduler.finish_requests(None, RequestStatus.FINISHED_ABORTED) + engine._send_abort_outputs(aborted) + if engine.batch_queue is not None: + engine.batch_queue.clear() + + self.status_type = EngineStatusType.UNHEALTHY + logger.info( + "[FT] Engine %d status -> UNHEALTHY:", self.engine_index, exc_info=exc + ) + + # ------------------------------------------------------------------ + # Instruction handlers (method name == instruction string) + # ------------------------------------------------------------------ + + def status(self, ft_request: FaultToleranceRequest) -> dict: + return { + "request_id": ft_request.request_id, + "success": True, + "engine_id": self.engine_index, + "status": self.status_type.name.lower(), + } + + def retry(self, ft_request: FaultToleranceRequest) -> dict: + engine = self.engine + executor = engine.model_executor + + with set_current_vllm_config(engine.vllm_config): + ft_request.params.update(self._reinit_dp_group()) + if hasattr(engine, "step_counter"): + engine.step_counter = 0 + + executor.collective_rpc("handle_ft_command", args=(ft_request,)) + + self.status_type = EngineStatusType.HEALTHY + logger.info("[FT] Engine %d status -> HEALTHY", self.engine_index) + self.resumed.set() + return {"request_id": ft_request.request_id, "success": True} + + # ------------------------------------------------------------------ + # Recovery helpers + # ------------------------------------------------------------------ + + def _reinit_dp_group(self) -> dict: + """Reinit DP process group if in DP mode. Returns worker params.""" + engine = self.engine + if not hasattr(engine, "dp_group") or not hasattr(engine, "dp_store"): + return {} + + parallel_config = engine.vllm_config.parallel_config + worker_key = f"ft_worker_dp_port_{self._dp_reinit_epoch}" + engine_key = f"ft_engine_dp_port_{self._dp_reinit_epoch}" + self._dp_reinit_epoch += 1 + + if parallel_config.data_parallel_rank == 0: + worker_port = get_open_port() + engine_port = get_open_port() + engine.dp_store.set(worker_key, str(worker_port).encode()) + engine.dp_store.set(engine_key, str(engine_port).encode()) + else: + worker_port = int(engine.dp_store.get(worker_key).decode()) + engine_port = int(engine.dp_store.get(engine_key).decode()) + + stateless_destroy_torch_distributed_process_group(engine.dp_group) + engine.dp_group, engine.dp_store = ( + stateless_init_torch_distributed_process_group( + parallel_config.data_parallel_master_ip, + engine_port, + parallel_config.data_parallel_rank, + parallel_config.data_parallel_size, + backend="gloo", + return_store=True, + ) + ) + return {"new_stateless_dp_group_port": worker_port} + + +def fault_tolerant_wrapper(busy_loop_func: Callable): + """Wrap the busy loop to catch faults and delegate recovery.""" + + def run_with_fault_tolerance(self: "EngineCoreProc"): + while True: + try: + busy_loop_func(self) + except SystemExit: + raise + except Exception as exc: + if not self.enable_fault_tolerance: + raise + self.ft_sentinel.on_fault(exc) + recovered = self.ft_sentinel.resumed.wait( + timeout=self.ft_sentinel.engine_recovery_timeout_sec + ) + if recovered: + continue + logger.error( + "[FT] No recovery within %ds timeout.", + self.ft_sentinel.engine_recovery_timeout_sec, + ) + raise + + return run_with_fault_tolerance diff --git a/vllm/v1/fault_tolerance/utils.py b/vllm/v1/fault_tolerance/utils.py new file mode 100644 index 000000000000..0c1b1689b01a --- /dev/null +++ b/vllm/v1/fault_tolerance/utils.py @@ -0,0 +1,17 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from typing import Any + +import msgspec + + +class FaultToleranceResult(msgspec.Struct): + request_id: str + success: bool + reason: str | None = None + + +class FaultToleranceRequest(msgspec.Struct): + instruction: str + params: dict[str, Any] + request_id: str = "" diff --git a/vllm/v1/worker/gpu_worker.py b/vllm/v1/worker/gpu_worker.py index ca63e0a117c7..6c031f793ed6 100644 --- a/vllm/v1/worker/gpu_worker.py +++ b/vllm/v1/worker/gpu_worker.py @@ -70,6 +70,7 @@ ModelRunnerOutput, ) from vllm.v1.utils import compute_iteration_details, report_usage_stats +from vllm.v1.worker.sentinel.gpu_worker_sentinel import WorkerSentinel from vllm.v1.worker.startup_plan import ( maybe_apply_startup_plan, maybe_save_startup_plan, @@ -146,7 +147,9 @@ def __init__( from vllm.distributed.elastic_ep.elastic_execute import ElasticEPScalingExecutor self.elastic_ep_executor = ElasticEPScalingExecutor(self) - + self.worker_sentinel: WorkerSentinel | None = None + if self.parallel_config.enable_fault_tolerance: + self.worker_sentinel = WorkerSentinel(worker=self, device=self.device) # Buffers saved before sleep self._sleep_saved_buffers: dict[str, torch.Tensor] = {} self._sleep_rebuild_draft_metadata_buffers = False @@ -414,6 +417,10 @@ def init_device(self): # If usage stat is enabled, collect relevant info. report_usage_stats(self.vllm_config) + def handle_ft_command(self, ft_request): + assert self.worker_sentinel is not None + return self.worker_sentinel.handle_command(ft_request) + # FIXME(youkaichao & ywang96): Use TorchDispatchMode instead of memory pool # to hijack tensor allocation. def load_model(self, *, load_dummy_weights: bool = False) -> None: diff --git a/vllm/v1/worker/sentinel/__init__.py b/vllm/v1/worker/sentinel/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/vllm/v1/worker/sentinel/gpu_worker_sentinel.py b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py new file mode 100644 index 000000000000..adcd06835c0f --- /dev/null +++ b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py @@ -0,0 +1,76 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from typing import TYPE_CHECKING + +import torch + +from vllm.config import set_current_vllm_config +from vllm.distributed import ( + get_dp_group, + get_ep_group, + stateless_init_torch_distributed_process_group, +) +from vllm.logger import init_logger +from vllm.v1.fault_tolerance.utils import FaultToleranceRequest +from vllm.v1.serial_utils import run_method + +if TYPE_CHECKING: + from vllm.v1.worker.gpu_worker import Worker + +logger = init_logger(__name__) + +_FT_BACKEND_SET = {"deepep_low_latency", "nixl_ep"} + + +class WorkerSentinel: + """Holds FT state for a single worker (mask tensors, DP config). + + Methods are called via collective_rpc from EngineCoreSentinel. + """ + + def __init__(self, worker: "Worker", device: torch.device): + self.worker = worker + self.device = device + self.dp_rank = worker.parallel_config.data_parallel_rank + self.dp_size = worker.parallel_config.data_parallel_size + self.data_parallel_master_ip = worker.parallel_config.data_parallel_master_ip + + self.use_ft_backend = ( + worker.parallel_config.all2all_backend in _FT_BACKEND_SET + and self.dp_size > 1 + ) + + def handle_command(self, ft_request: FaultToleranceRequest): + """Dispatch an FT command by instruction name.""" + with set_current_vllm_config(self.worker.vllm_config): + return run_method(self, ft_request.instruction, (ft_request,), {}) + + def retry(self, ft_request: FaultToleranceRequest): + torch.accelerator.synchronize() + params = ft_request.params + self._clean_worker_state() + if self.dp_size > 1: + port = params["new_stateless_dp_group_port"] + get_dp_group().cpu_group = stateless_init_torch_distributed_process_group( + self.data_parallel_master_ip, + port, + self.dp_rank, + self.dp_size, + backend="gloo", + ) + if self.use_ft_backend: + comm = get_ep_group().device_communicator + assert comm and comm.all2all_manager + mgr = comm.all2all_manager + mgr.clean_buffers() + + def _clean_worker_state(self): + self.worker.model_runner.execute_model_state = None + self.worker.model_runner.kv_connector_output = None + input_batch = self.worker.model_runner.input_batch + cached_req_ids = input_batch.req_id_to_index.keys() + for req_id in list(cached_req_ids): + input_batch.remove_request(req_id) + input_batch.condense() + input_batch.refresh_metadata() + input_batch.req_prompt_embeds.clear() From 7f7c375d63170b81dff422621cbdcb64f4cbfcc9 Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Fri, 3 Jul 2026 09:21:58 +0800 Subject: [PATCH 02/16] unify param name for nixl and deepep and destroy old cpu_group in retry (#230) * Clean comments * Check ft is enabled with ft all2all backend at initialization (#237) * adapt to nixl-1.13.0 (#238) Signed-off-by: fangyuchu --- .../device_communicators/all2all.py | 6 ++++-- vllm/engine/arg_utils.py | 6 ++++++ .../fault_tolerance/engine_core_sentinel.py | 16 -------------- vllm/v1/fault_tolerance/utils.py | 4 ++++ .../v1/worker/sentinel/gpu_worker_sentinel.py | 21 ++++++++----------- 5 files changed, 23 insertions(+), 30 deletions(-) diff --git a/vllm/distributed/device_communicators/all2all.py b/vllm/distributed/device_communicators/all2all.py index f96ccb6ccc2d..ee404f688a19 100644 --- a/vllm/distributed/device_communicators/all2all.py +++ b/vllm/distributed/device_communicators/all2all.py @@ -582,11 +582,13 @@ def clean_buffers(self) -> None: if NixlEPAll2AllManager._buffer is None: return state = NixlEPAll2AllManager._buffer - state.buffer.get_local_buffer_tensor(dtype=torch.int8).zero_() + state.buffer.get_local_buffer_tensor( + dtype=torch.int8, use_rdma_buffer=True + ).zero_() torch.accelerator.synchronize() state.buffer.clean_mask_buffer() torch.accelerator.synchronize() - NixlEPAll2AllManager._last_active_mask = None + NixlEPAll2AllManager._last_mask = None class FlashInferNVLinkTwoSidedManager(All2AllManagerBase): diff --git a/vllm/engine/arg_utils.py b/vllm/engine/arg_utils.py index 82294132eb64..78cd8d0c8a6a 100644 --- a/vllm/engine/arg_utils.py +++ b/vllm/engine/arg_utils.py @@ -117,6 +117,7 @@ from vllm.utils.network_utils import get_ip from vllm.utils.torch_utils import resolve_kv_cache_dtype_string from vllm.v1.attention.backends.registry import AttentionBackendEnum +from vllm.v1.fault_tolerance.utils import FT_BACKEND_SET from vllm.v1.sample.logits_processor import LogitsProcessor from vllm.version import __version__ as VLLM_VERSION @@ -2056,6 +2057,11 @@ def create_engine_config( "(--data-parallel-external-lb or --data-parallel-rank). " "Internal LB mode is not supported." ) + if self.enable_fault_tolerance and self.all2all_backend not in FT_BACKEND_SET: + raise ValueError( + "Fault tolerance requires an FT-capable all2all backend " + f"(deepep_low_latency or nixl_ep), but got '{self.all2all_backend}'." + ) if ( self.data_parallel_size > 1 and data_parallel_external_lb diff --git a/vllm/v1/fault_tolerance/engine_core_sentinel.py b/vllm/v1/fault_tolerance/engine_core_sentinel.py index 5b6fdd568b74..d59e70893278 100644 --- a/vllm/v1/fault_tolerance/engine_core_sentinel.py +++ b/vllm/v1/fault_tolerance/engine_core_sentinel.py @@ -39,10 +39,6 @@ def __init__(self, engine: "EngineCoreProc", parallel_config): self.status_type = EngineStatusType.HEALTHY self._dp_reinit_epoch = 0 - # ------------------------------------------------------------------ - # Command dispatch (called from process_input_sockets thread) - # ------------------------------------------------------------------ - def handle_command(self, client_idx: int, call_id: int, ft_args: dict): """Dispatch an FT command by instruction name and enqueue result.""" ft_request = FaultToleranceRequest(**ft_args) @@ -62,10 +58,6 @@ def handle_command(self, client_idx: int, call_id: int, ft_args: dict): (client_idx, EngineCoreOutputs(utility_output=uo)) ) - # ------------------------------------------------------------------ - # Fault handling (called by wrapper, runs in busy-loop thread) - # ------------------------------------------------------------------ - def on_fault(self, exc: Exception): """Called by the wrapper when the busy loop raises an exception.""" self.resumed.clear() @@ -84,10 +76,6 @@ def on_fault(self, exc: Exception): "[FT] Engine %d status -> UNHEALTHY:", self.engine_index, exc_info=exc ) - # ------------------------------------------------------------------ - # Instruction handlers (method name == instruction string) - # ------------------------------------------------------------------ - def status(self, ft_request: FaultToleranceRequest) -> dict: return { "request_id": ft_request.request_id, @@ -112,10 +100,6 @@ def retry(self, ft_request: FaultToleranceRequest) -> dict: self.resumed.set() return {"request_id": ft_request.request_id, "success": True} - # ------------------------------------------------------------------ - # Recovery helpers - # ------------------------------------------------------------------ - def _reinit_dp_group(self) -> dict: """Reinit DP process group if in DP mode. Returns worker params.""" engine = self.engine diff --git a/vllm/v1/fault_tolerance/utils.py b/vllm/v1/fault_tolerance/utils.py index 0c1b1689b01a..712567d0c2f3 100644 --- a/vllm/v1/fault_tolerance/utils.py +++ b/vllm/v1/fault_tolerance/utils.py @@ -4,6 +4,10 @@ import msgspec +# All2all backends that support fault-tolerant timeout + rank masking, +# required for FT under DP+EP MoE deployments. +FT_BACKEND_SET = frozenset({"deepep_low_latency", "nixl_ep"}) + class FaultToleranceResult(msgspec.Struct): request_id: str diff --git a/vllm/v1/worker/sentinel/gpu_worker_sentinel.py b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py index adcd06835c0f..904f8ff25909 100644 --- a/vllm/v1/worker/sentinel/gpu_worker_sentinel.py +++ b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py @@ -8,6 +8,7 @@ from vllm.distributed import ( get_dp_group, get_ep_group, + stateless_destroy_torch_distributed_process_group, stateless_init_torch_distributed_process_group, ) from vllm.logger import init_logger @@ -19,8 +20,6 @@ logger = init_logger(__name__) -_FT_BACKEND_SET = {"deepep_low_latency", "nixl_ep"} - class WorkerSentinel: """Holds FT state for a single worker (mask tensors, DP config). @@ -35,11 +34,6 @@ def __init__(self, worker: "Worker", device: torch.device): self.dp_size = worker.parallel_config.data_parallel_size self.data_parallel_master_ip = worker.parallel_config.data_parallel_master_ip - self.use_ft_backend = ( - worker.parallel_config.all2all_backend in _FT_BACKEND_SET - and self.dp_size > 1 - ) - def handle_command(self, ft_request: FaultToleranceRequest): """Dispatch an FT command by instruction name.""" with set_current_vllm_config(self.worker.vllm_config): @@ -50,6 +44,8 @@ def retry(self, ft_request: FaultToleranceRequest): params = ft_request.params self._clean_worker_state() if self.dp_size > 1: + old_cpu_group = get_dp_group().cpu_group + stateless_destroy_torch_distributed_process_group(old_cpu_group) port = params["new_stateless_dp_group_port"] get_dp_group().cpu_group = stateless_init_torch_distributed_process_group( self.data_parallel_master_ip, @@ -58,11 +54,12 @@ def retry(self, ft_request: FaultToleranceRequest): self.dp_size, backend="gloo", ) - if self.use_ft_backend: - comm = get_ep_group().device_communicator - assert comm and comm.all2all_manager - mgr = comm.all2all_manager - mgr.clean_buffers() + self._get_all2all_manager().clean_buffers() + + def _get_all2all_manager(self): + comm = get_ep_group().device_communicator + assert comm and comm.all2all_manager + return comm.all2all_manager def _clean_worker_state(self): self.worker.model_runner.execute_model_state = None From 579be781f9de6e7928747000360bec404249c614 Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Fri, 3 Jul 2026 11:43:36 +0800 Subject: [PATCH 03/16] Surface exception to status Signed-off-by: fangyuchu --- vllm/engine/arg_utils.py | 6 ------ vllm/v1/engine/core_client.py | 4 +--- vllm/v1/fault_tolerance/engine_core_sentinel.py | 11 +++++++---- vllm/v1/fault_tolerance/utils.py | 4 ---- vllm/v1/worker/sentinel/gpu_worker_sentinel.py | 9 +++++++++ 5 files changed, 17 insertions(+), 17 deletions(-) diff --git a/vllm/engine/arg_utils.py b/vllm/engine/arg_utils.py index 78cd8d0c8a6a..82294132eb64 100644 --- a/vllm/engine/arg_utils.py +++ b/vllm/engine/arg_utils.py @@ -117,7 +117,6 @@ from vllm.utils.network_utils import get_ip from vllm.utils.torch_utils import resolve_kv_cache_dtype_string from vllm.v1.attention.backends.registry import AttentionBackendEnum -from vllm.v1.fault_tolerance.utils import FT_BACKEND_SET from vllm.v1.sample.logits_processor import LogitsProcessor from vllm.version import __version__ as VLLM_VERSION @@ -2057,11 +2056,6 @@ def create_engine_config( "(--data-parallel-external-lb or --data-parallel-rank). " "Internal LB mode is not supported." ) - if self.enable_fault_tolerance and self.all2all_backend not in FT_BACKEND_SET: - raise ValueError( - "Fault tolerance requires an FT-capable all2all backend " - f"(deepep_low_latency or nixl_ep), but got '{self.all2all_backend}'." - ) if ( self.data_parallel_size > 1 and data_parallel_external_lb diff --git a/vllm/v1/engine/core_client.py b/vllm/v1/engine/core_client.py index 2f0c0910f071..0114bf10b48d 100644 --- a/vllm/v1/engine/core_client.py +++ b/vllm/v1/engine/core_client.py @@ -1222,9 +1222,7 @@ async def get_status(self): return { "schema_version": 1, "total_engines": len(self.engine_ranks_managed), - "engines": [ - {"id": res["engine_id"], "status": res["status"]}, - ], + "engines": [res], } diff --git a/vllm/v1/fault_tolerance/engine_core_sentinel.py b/vllm/v1/fault_tolerance/engine_core_sentinel.py index d59e70893278..7065681ba226 100644 --- a/vllm/v1/fault_tolerance/engine_core_sentinel.py +++ b/vllm/v1/fault_tolerance/engine_core_sentinel.py @@ -37,6 +37,7 @@ def __init__(self, engine: "EngineCoreProc", parallel_config): self.resumed = threading.Event() self.resumed.set() self.status_type = EngineStatusType.HEALTHY + self.fault_info: str | None = None self._dp_reinit_epoch = 0 def handle_command(self, client_idx: int, call_id: int, ft_args: dict): @@ -72,17 +73,19 @@ def on_fault(self, exc: Exception): engine.batch_queue.clear() self.status_type = EngineStatusType.UNHEALTHY + self.fault_info = f"{type(exc).__name__}: {exc}" logger.info( "[FT] Engine %d status -> UNHEALTHY:", self.engine_index, exc_info=exc ) def status(self, ft_request: FaultToleranceRequest) -> dict: - return { - "request_id": ft_request.request_id, - "success": True, - "engine_id": self.engine_index, + result = { + "id": self.engine_index, "status": self.status_type.name.lower(), } + if self.status_type == EngineStatusType.UNHEALTHY: + result["fault_info"] = self.fault_info + return result def retry(self, ft_request: FaultToleranceRequest) -> dict: engine = self.engine diff --git a/vllm/v1/fault_tolerance/utils.py b/vllm/v1/fault_tolerance/utils.py index 712567d0c2f3..0c1b1689b01a 100644 --- a/vllm/v1/fault_tolerance/utils.py +++ b/vllm/v1/fault_tolerance/utils.py @@ -4,10 +4,6 @@ import msgspec -# All2all backends that support fault-tolerant timeout + rank masking, -# required for FT under DP+EP MoE deployments. -FT_BACKEND_SET = frozenset({"deepep_low_latency", "nixl_ep"}) - class FaultToleranceResult(msgspec.Struct): request_id: str diff --git a/vllm/v1/worker/sentinel/gpu_worker_sentinel.py b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py index 904f8ff25909..3b0d2d9162e6 100644 --- a/vllm/v1/worker/sentinel/gpu_worker_sentinel.py +++ b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py @@ -20,6 +20,10 @@ logger = init_logger(__name__) +# All2all backends that support fault-tolerant timeout + rank masking, +# required for FT under DP+EP MoE deployments. +FT_BACKEND_SET = frozenset({"deepep_low_latency", "nixl_ep"}) + class WorkerSentinel: """Holds FT state for a single worker (mask tensors, DP config). @@ -33,6 +37,11 @@ def __init__(self, worker: "Worker", device: torch.device): self.dp_rank = worker.parallel_config.data_parallel_rank self.dp_size = worker.parallel_config.data_parallel_size self.data_parallel_master_ip = worker.parallel_config.data_parallel_master_ip + all2all_backend = worker.parallel_config.all2all_backend + assert all2all_backend in FT_BACKEND_SET, ( + f"Fault tolerance requires an FT-capable all2all backend " + f"(one of {sorted(FT_BACKEND_SET)}), but got '{all2all_backend}'." + ) def handle_command(self, ft_request: FaultToleranceRequest): """Dispatch an FT command by instruction name.""" From 397ae140a01a3c828eb484a3be090c78e1bf36cd Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Fri, 3 Jul 2026 14:23:14 +0800 Subject: [PATCH 04/16] change status report mode from pull to push Signed-off-by: fangyuchu --- vllm/v1/engine/__init__.py | 2 ++ vllm/v1/engine/core_client.py | 23 ++++++++++++--- .../fault_tolerance/engine_core_sentinel.py | 28 +++++++++++++------ 3 files changed, 41 insertions(+), 12 deletions(-) diff --git a/vllm/v1/engine/__init__.py b/vllm/v1/engine/__init__.py index 3a916af832f4..4ac27be5068f 100644 --- a/vllm/v1/engine/__init__.py +++ b/vllm/v1/engine/__init__.py @@ -31,6 +31,8 @@ EEP_NOTIFICATION_CALL_ID = -1 +FT_STATUS_CALL_ID = -2 + class EEPNotificationType(enum.Enum): NEW_CORE_ENGINES_INIT_READY = "NEW_CORE_ENGINES_INIT_READY" diff --git a/vllm/v1/engine/core_client.py b/vllm/v1/engine/core_client.py index 0114bf10b48d..c64bc0b5d4c1 100644 --- a/vllm/v1/engine/core_client.py +++ b/vllm/v1/engine/core_client.py @@ -35,6 +35,7 @@ ) from vllm.v1.engine import ( EEP_NOTIFICATION_CALL_ID, + FT_STATUS_CALL_ID, EEPNotificationType, EngineCoreOutputs, EngineCoreReadyResponse, @@ -984,6 +985,14 @@ def __init__( self.client_count = client_count self.client_index = client_index self.outputs_queue = asyncio.Queue[EngineCoreOutputs | Exception]() + + # locally-cached engine status + self._engine_status: dict[int, dict] = {} + if self.vllm_config.parallel_config.enable_fault_tolerance: + self._engine_status = { + rank: {"id": rank, "status": "healthy"} + for rank in self.engine_ranks_managed + } try: # If we are running in an asyncio event loop, start the queue task. # Otherwise, it will be started lazily. If it is not started here, @@ -1007,7 +1016,7 @@ def _ensure_output_queue_task(self): output_handler: ( Callable[[AsyncMPClient, EngineCoreOutputs], Awaitable[None]] | None ) = getattr(self.__class__, "process_engine_outputs", None) - _self_ref = weakref.ref(self) if output_handler else None + _self_ref = weakref.ref(self) output_socket = resources.output_socket assert output_socket is not None @@ -1038,6 +1047,14 @@ async def process_outputs_socket(): asyncio.create_task( notification_callback_handler(_self, notification_data) ) + elif outputs.utility_output.call_id == FT_STATUS_CALL_ID: + _self = _self_ref() + if not _self: + return + if outputs.utility_output.result is not None: + _self._engine_status[outputs.engine_index] = ( + outputs.utility_output.result.result + ) else: _process_utility_output( outputs.utility_output, utility_results @@ -1217,12 +1234,10 @@ async def handle_fault( return result async def get_status(self): - ft_request = FaultToleranceRequest(instruction="status", params={}) - res = await self.call_utility_async(FT_UTILITY_METHOD, ft_request) return { "schema_version": 1, "total_engines": len(self.engine_ranks_managed), - "engines": [res], + "engines": list(self._engine_status.values()), } diff --git a/vllm/v1/fault_tolerance/engine_core_sentinel.py b/vllm/v1/fault_tolerance/engine_core_sentinel.py index 7065681ba226..b953968e75ca 100644 --- a/vllm/v1/fault_tolerance/engine_core_sentinel.py +++ b/vllm/v1/fault_tolerance/engine_core_sentinel.py @@ -11,7 +11,12 @@ from vllm.distributed.utils import stateless_init_torch_distributed_process_group from vllm.logger import init_logger from vllm.utils.network_utils import get_open_port -from vllm.v1.engine import EngineCoreOutputs, EngineStatusType, UtilityOutput +from vllm.v1.engine import ( + FT_STATUS_CALL_ID, + EngineCoreOutputs, + EngineStatusType, + UtilityOutput, +) from vllm.v1.fault_tolerance.utils import FaultToleranceRequest from vllm.v1.request import RequestStatus from vllm.v1.serial_utils import UtilityResult, run_method @@ -77,15 +82,21 @@ def on_fault(self, exc: Exception): logger.info( "[FT] Engine %d status -> UNHEALTHY:", self.engine_index, exc_info=exc ) + self._push_status() - def status(self, ft_request: FaultToleranceRequest) -> dict: - result = { - "id": self.engine_index, - "status": self.status_type.name.lower(), - } + def _push_status(self): + """Push current health to the client so it can refresh its cache.""" + payload = {"id": self.engine_index, "status": self.status_type.name.lower()} if self.status_type == EngineStatusType.UNHEALTHY: - result["fault_info"] = self.fault_info - return result + payload["fault_info"] = self.fault_info + outputs = EngineCoreOutputs( + utility_output=UtilityOutput( + call_id=FT_STATUS_CALL_ID, + result=UtilityResult(payload), + ) + ) + outputs.engine_index = self.engine_index + self.engine.output_queue.put_nowait((0, outputs)) def retry(self, ft_request: FaultToleranceRequest) -> dict: engine = self.engine @@ -101,6 +112,7 @@ def retry(self, ft_request: FaultToleranceRequest) -> dict: self.status_type = EngineStatusType.HEALTHY logger.info("[FT] Engine %d status -> HEALTHY", self.engine_index) self.resumed.set() + self._push_status() return {"request_id": ft_request.request_id, "success": True} def _reinit_dp_group(self) -> dict: From c20d2f1b9e49ba986ac12ac38c29d95fafc8ce03 Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Fri, 3 Jul 2026 14:56:57 +0800 Subject: [PATCH 05/16] simplify fault tolerance config args Signed-off-by: fangyuchu --- vllm/engine/arg_utils.py | 20 ++++---------------- 1 file changed, 4 insertions(+), 16 deletions(-) diff --git a/vllm/engine/arg_utils.py b/vllm/engine/arg_utils.py index 82294132eb64..de2e213e48df 100644 --- a/vllm/engine/arg_utils.py +++ b/vllm/engine/arg_utils.py @@ -723,8 +723,7 @@ class EngineArgs: optimization_level: OptimizationLevel = VllmConfig.optimization_level performance_mode: PerformanceMode = VllmConfig.performance_mode - # fault tolerance fields (`None` means not explicitly provided). - fault_tolerance_config: FaultToleranceConfig | None = get_field( + fault_tolerance_config: FaultToleranceConfig = get_field( ParallelConfig, "fault_tolerance_config" ) enable_fault_tolerance: bool = ParallelConfig.enable_fault_tolerance @@ -762,6 +761,7 @@ def __post_init__(self): **self.weight_transfer_config ) if isinstance(self.fault_tolerance_config, dict): + self.enable_fault_tolerance = True self.fault_tolerance_config = FaultToleranceConfig( **self.fault_tolerance_config ) @@ -1166,11 +1166,7 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: "--enable-fault-tolerance", **parallel_kwargs["enable_fault_tolerance"] ) parallel_group.add_argument( - "--fault-tolerance-config", - **{ - **parallel_kwargs["fault_tolerance_config"], - "default": None, - }, + "--fault-tolerance-config", **parallel_kwargs["fault_tolerance_config"] ) # KV cache arguments @@ -1639,12 +1635,6 @@ def from_cli_args(cls, args: argparse.Namespace): # Get the list of attributes of this dataclass. attrs = [attr.name for attr in dataclasses.fields(cls)] - # If --fault-tolerance-config is provided, enable fault tolerance by default. - if args.fault_tolerance_config is not None: - args.enable_fault_tolerance = True - if args.enable_fault_tolerance and args.fault_tolerance_config is None: - args.fault_tolerance_config = FaultToleranceConfig() - # Set the attributes from the parsed arguments. engine_args = cls( **{attr: getattr(args, attr) for attr in attrs if hasattr(args, attr)} @@ -2214,9 +2204,7 @@ def create_engine_config( _api_process_rank=self._api_process_rank, assigned_physical_gpu_ids=self._resolve_device_ids(), enable_fault_tolerance=self.enable_fault_tolerance, - fault_tolerance_config=( - self.fault_tolerance_config or FaultToleranceConfig() - ), + fault_tolerance_config=self.fault_tolerance_config, numa_bind=self.numa_bind, numa_bind_nodes=self.numa_bind_nodes, numa_bind_cpus=self.numa_bind_cpus, From e5ee32302a693cfaeaa25ab4cd2d9f01ac521d9c Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Fri, 3 Jul 2026 15:49:01 +0800 Subject: [PATCH 06/16] use existing get_all2all_manager implementation Signed-off-by: fangyuchu --- vllm/v1/worker/sentinel/gpu_worker_sentinel.py | 9 ++------- 1 file changed, 2 insertions(+), 7 deletions(-) diff --git a/vllm/v1/worker/sentinel/gpu_worker_sentinel.py b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py index 3b0d2d9162e6..aca1c9277445 100644 --- a/vllm/v1/worker/sentinel/gpu_worker_sentinel.py +++ b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py @@ -7,11 +7,11 @@ from vllm.config import set_current_vllm_config from vllm.distributed import ( get_dp_group, - get_ep_group, stateless_destroy_torch_distributed_process_group, stateless_init_torch_distributed_process_group, ) from vllm.logger import init_logger +from vllm.model_executor.layers.fused_moe.all2all_utils import get_ep_all2all_manager from vllm.v1.fault_tolerance.utils import FaultToleranceRequest from vllm.v1.serial_utils import run_method @@ -63,12 +63,7 @@ def retry(self, ft_request: FaultToleranceRequest): self.dp_size, backend="gloo", ) - self._get_all2all_manager().clean_buffers() - - def _get_all2all_manager(self): - comm = get_ep_group().device_communicator - assert comm and comm.all2all_manager - return comm.all2all_manager + get_ep_all2all_manager().clean_buffers() def _clean_worker_state(self): self.worker.model_runner.execute_model_state = None From 3ed1aab5cc6e7e819239a90012ca11d95873dc60 Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Mon, 6 Jul 2026 20:55:24 +0800 Subject: [PATCH 07/16] support tp>1 (#241) Signed-off-by: fangyuchu --- vllm/v1/fault_tolerance/engine_core_sentinel.py | 11 ++++++----- vllm/v1/worker/sentinel/gpu_worker_sentinel.py | 3 ++- 2 files changed, 8 insertions(+), 6 deletions(-) diff --git a/vllm/v1/fault_tolerance/engine_core_sentinel.py b/vllm/v1/fault_tolerance/engine_core_sentinel.py index b953968e75ca..817e4280be79 100644 --- a/vllm/v1/fault_tolerance/engine_core_sentinel.py +++ b/vllm/v1/fault_tolerance/engine_core_sentinel.py @@ -2,6 +2,7 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project """EngineCoreSentinel and fault_tolerant_wrapper for the engine core.""" +import json import threading from collections.abc import Callable from typing import TYPE_CHECKING @@ -122,17 +123,17 @@ def _reinit_dp_group(self) -> dict: return {} parallel_config = engine.vllm_config.parallel_config - worker_key = f"ft_worker_dp_port_{self._dp_reinit_epoch}" + worker_key = f"ft_worker_dp_ports_{self._dp_reinit_epoch}" engine_key = f"ft_engine_dp_port_{self._dp_reinit_epoch}" self._dp_reinit_epoch += 1 if parallel_config.data_parallel_rank == 0: - worker_port = get_open_port() + worker_ports = [get_open_port() for _ in range(parallel_config.world_size)] engine_port = get_open_port() - engine.dp_store.set(worker_key, str(worker_port).encode()) + engine.dp_store.set(worker_key, json.dumps(worker_ports).encode()) engine.dp_store.set(engine_key, str(engine_port).encode()) else: - worker_port = int(engine.dp_store.get(worker_key).decode()) + worker_ports = json.loads(engine.dp_store.get(worker_key).decode()) engine_port = int(engine.dp_store.get(engine_key).decode()) stateless_destroy_torch_distributed_process_group(engine.dp_group) @@ -146,7 +147,7 @@ def _reinit_dp_group(self) -> dict: return_store=True, ) ) - return {"new_stateless_dp_group_port": worker_port} + return {"new_stateless_dp_group_ports": worker_ports} def fault_tolerant_wrapper(busy_loop_func: Callable): diff --git a/vllm/v1/worker/sentinel/gpu_worker_sentinel.py b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py index aca1c9277445..33d1145234be 100644 --- a/vllm/v1/worker/sentinel/gpu_worker_sentinel.py +++ b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py @@ -55,7 +55,8 @@ def retry(self, ft_request: FaultToleranceRequest): if self.dp_size > 1: old_cpu_group = get_dp_group().cpu_group stateless_destroy_torch_distributed_process_group(old_cpu_group) - port = params["new_stateless_dp_group_port"] + world_size = self.worker.parallel_config.world_size + port = params["new_stateless_dp_group_ports"][self.worker.rank % world_size] get_dp_group().cpu_group = stateless_init_torch_distributed_process_group( self.data_parallel_master_ip, port, From 89b31e0675e338e8e5f1e80835cfbe935e4ec40f Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Tue, 7 Jul 2026 12:11:27 +0800 Subject: [PATCH 08/16] Enhance the pass of fault tolerance results (#242) Signed-off-by: fangyuchu --- vllm/engine/arg_utils.py | 7 ++++++- .../serve/fault_tolerance/api_router.py | 12 +++++++----- vllm/v1/engine/core_client.py | 4 ++-- .../v1/fault_tolerance/engine_core_sentinel.py | 18 +++++++++--------- 4 files changed, 24 insertions(+), 17 deletions(-) diff --git a/vllm/engine/arg_utils.py b/vllm/engine/arg_utils.py index de2e213e48df..0244ff192cd6 100644 --- a/vllm/engine/arg_utils.py +++ b/vllm/engine/arg_utils.py @@ -761,7 +761,12 @@ def __post_init__(self): **self.weight_transfer_config ) if isinstance(self.fault_tolerance_config, dict): - self.enable_fault_tolerance = True + if not self.enable_fault_tolerance: + logger.warning( + "--fault-tolerance-config was passed. Fault tolerance is being " + "automatically enabled." + ) + self.enable_fault_tolerance = True self.fault_tolerance_config = FaultToleranceConfig( **self.fault_tolerance_config ) diff --git a/vllm/entrypoints/serve/fault_tolerance/api_router.py b/vllm/entrypoints/serve/fault_tolerance/api_router.py index cae5742365eb..c1d1f2fed34e 100644 --- a/vllm/entrypoints/serve/fault_tolerance/api_router.py +++ b/vllm/entrypoints/serve/fault_tolerance/api_router.py @@ -21,14 +21,16 @@ def _validate_payload(body: dict) -> tuple[str, dict]: + if not isinstance(body, dict): + raise HTTPException(400, "Request body must be a JSON object.") instruction = body.get("instruction") - params = body.get("params") - if not instruction or not isinstance(params, dict): - raise HTTPException(400, "'instruction' and 'params' are required.") + if not instruction: + raise HTTPException(400, "'instruction' is required.") if instruction not in _ALLOWED_INSTRUCTIONS: raise HTTPException(400, f"Invalid instruction: '{instruction}'.") - if "timeout" not in params or not isinstance(params["timeout"], (int, float)): - raise HTTPException(400, "Missing or invalid 'timeout' parameter.") + params = body.get("params", {}) + if not isinstance(params, dict): + raise HTTPException(400, "'params' must be an object.") return instruction, params diff --git a/vllm/v1/engine/core_client.py b/vllm/v1/engine/core_client.py index c64bc0b5d4c1..c43d52c0d830 100644 --- a/vllm/v1/engine/core_client.py +++ b/vllm/v1/engine/core_client.py @@ -16,6 +16,7 @@ from threading import Thread from typing import Any, TypeAlias, TypeVar +import msgspec import msgspec.msgpack import zmq import zmq.asyncio @@ -1230,8 +1231,7 @@ async def handle_fault( self, ft_request: FaultToleranceRequest ) -> FaultToleranceResult: res = await self.call_utility_async(FT_UTILITY_METHOD, ft_request) - result = FaultToleranceResult(**res) - return result + return msgspec.convert(res, FaultToleranceResult) async def get_status(self): return { diff --git a/vllm/v1/fault_tolerance/engine_core_sentinel.py b/vllm/v1/fault_tolerance/engine_core_sentinel.py index 817e4280be79..02bf1ab0b38e 100644 --- a/vllm/v1/fault_tolerance/engine_core_sentinel.py +++ b/vllm/v1/fault_tolerance/engine_core_sentinel.py @@ -7,6 +7,8 @@ from collections.abc import Callable from typing import TYPE_CHECKING +import msgspec + from vllm.config import set_current_vllm_config from vllm.distributed import stateless_destroy_torch_distributed_process_group from vllm.distributed.utils import stateless_init_torch_distributed_process_group @@ -18,7 +20,7 @@ EngineStatusType, UtilityOutput, ) -from vllm.v1.fault_tolerance.utils import FaultToleranceRequest +from vllm.v1.fault_tolerance.utils import FaultToleranceRequest, FaultToleranceResult from vllm.v1.request import RequestStatus from vllm.v1.serial_utils import UtilityResult, run_method @@ -53,14 +55,12 @@ def handle_command(self, client_idx: int, call_id: int, ft_args: dict): result = run_method(self, ft_request.instruction, (ft_request,), {}) except Exception as e: logger.exception("[FT] Instruction '%s' failed", ft_request.instruction) - result = { - "request_id": ft_request.request_id, - "success": False, - "reason": str(e), - } + result = FaultToleranceResult( + request_id=ft_request.request_id, success=False, reason=str(e) + ) uo = UtilityOutput(call_id) - uo.result = UtilityResult(result) + uo.result = UtilityResult(msgspec.structs.asdict(result)) self.engine.output_queue.put_nowait( (client_idx, EngineCoreOutputs(utility_output=uo)) ) @@ -99,7 +99,7 @@ def _push_status(self): outputs.engine_index = self.engine_index self.engine.output_queue.put_nowait((0, outputs)) - def retry(self, ft_request: FaultToleranceRequest) -> dict: + def retry(self, ft_request: FaultToleranceRequest) -> FaultToleranceResult: engine = self.engine executor = engine.model_executor @@ -114,7 +114,7 @@ def retry(self, ft_request: FaultToleranceRequest) -> dict: logger.info("[FT] Engine %d status -> HEALTHY", self.engine_index) self.resumed.set() self._push_status() - return {"request_id": ft_request.request_id, "success": True} + return FaultToleranceResult(request_id=ft_request.request_id, success=True) def _reinit_dp_group(self) -> dict: """Reinit DP process group if in DP mode. Returns worker params.""" From ced92b5064e08e206d6e15ffba1f9d2d3ee43b84 Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Tue, 14 Jul 2026 11:23:54 +0800 Subject: [PATCH 09/16] clean states for model runner v2 Signed-off-by: fangyuchu --- .../v1/worker/sentinel/gpu_worker_sentinel.py | 26 ++++++++++++------- 1 file changed, 16 insertions(+), 10 deletions(-) diff --git a/vllm/v1/worker/sentinel/gpu_worker_sentinel.py b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py index 33d1145234be..41c0542c4553 100644 --- a/vllm/v1/worker/sentinel/gpu_worker_sentinel.py +++ b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py @@ -1,6 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, cast import torch @@ -16,6 +16,7 @@ from vllm.v1.serial_utils import run_method if TYPE_CHECKING: + from vllm.v1.worker.gpu.model_runner import GPUModelRunner as GPUModelRunnerV2 from vllm.v1.worker.gpu_worker import Worker logger = init_logger(__name__) @@ -67,12 +68,17 @@ def retry(self, ft_request: FaultToleranceRequest): get_ep_all2all_manager().clean_buffers() def _clean_worker_state(self): - self.worker.model_runner.execute_model_state = None - self.worker.model_runner.kv_connector_output = None - input_batch = self.worker.model_runner.input_batch - cached_req_ids = input_batch.req_id_to_index.keys() - for req_id in list(cached_req_ids): - input_batch.remove_request(req_id) - input_batch.condense() - input_batch.refresh_metadata() - input_batch.req_prompt_embeds.clear() + model_runner = self.worker.model_runner + model_runner.execute_model_state = None + if self.worker.use_v2_model_runner: + runner = cast("GPUModelRunnerV2", model_runner) + for req_id in list(runner.req_states.req_id_to_index): + runner._remove_request(req_id) + else: + model_runner.kv_connector_output = None + input_batch = model_runner.input_batch + for req_id in list(input_batch.req_id_to_index): + input_batch.remove_request(req_id) + input_batch.condense() + input_batch.refresh_metadata() + input_batch.req_prompt_embeds.clear() From 6ee0981e6c24bbea6787797b50c8888a963bdd46 Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Tue, 14 Jul 2026 11:24:43 +0800 Subject: [PATCH 10/16] add support for fault detection throuth mask for model runner v2 Signed-off-by: fangyuchu --- vllm/v1/worker/gpu/async_utils.py | 15 +++++++++++++++ vllm/v1/worker/gpu/model_runner.py | 8 ++++++++ 2 files changed, 23 insertions(+) diff --git a/vllm/v1/worker/gpu/async_utils.py b/vllm/v1/worker/gpu/async_utils.py index e4659104f49e..4570d7267344 100644 --- a/vllm/v1/worker/gpu/async_utils.py +++ b/vllm/v1/worker/gpu/async_utils.py @@ -5,6 +5,7 @@ import numpy as np import torch +from vllm.model_executor.layers.fused_moe.all2all_utils import get_ep_all2all_manager from vllm.v1.outputs import AsyncModelRunnerOutput, LogprobsTensors, ModelRunnerOutput from vllm.v1.worker.gpu.sample.output import SamplerOutput @@ -17,6 +18,7 @@ def __init__( num_sampled_tokens: torch.Tensor, main_stream: torch.cuda.Stream, copy_stream: torch.cuda.Stream, + check_ep_fault: bool = False, ): # NOTE(woosuk): We must retain references to the GPU tensors, # as the copy operations are performed on a different CUDA stream than @@ -26,6 +28,7 @@ def __init__( self.num_sampled_tokens = num_sampled_tokens # Blocking (sleep) event to avoid busy-polling the CUDA driver lock. self.copy_event = torch.cuda.Event(blocking=True) + self._has_fault: torch.Tensor | None = None with stream(copy_stream, main_stream): copy_stream.wait_stream(main_stream) @@ -44,6 +47,9 @@ def __init__( k: v.to_cpu_nonblocking() if v is not None else None for k, v in self.model_runner_output.prompt_logprobs_dict.items() } + if check_ep_fault: + has_fault = get_ep_all2all_manager().query_fault() + self._has_fault = has_fault.to("cpu", non_blocking=True) self.copy_event.record(copy_stream) def get_output(self) -> ModelRunnerOutput: @@ -67,6 +73,15 @@ def get_output(self) -> ModelRunnerOutput: if self.logprobs_tensors is not None: self.model_runner_output.logprobs = self.logprobs_tensors.tolists() self.model_runner_output.prompt_logprobs_dict = self.prompt_logprobs_dict + + if self._has_fault is not None and self._has_fault.item(): + mask = get_ep_all2all_manager().query_active_mask() + raise RuntimeError( + "Fault detected in EP all2all communication: " + "one or more ranks timed out during dispatch/combine. " + f"Mask: {mask.cpu().tolist()}" + ) + return self.model_runner_output diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index 3f86b5595cb9..fdedbfb86d03 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -38,6 +38,7 @@ ) from vllm.forward_context import BatchDescriptor, set_forward_context from vllm.logger import init_logger +from vllm.model_executor.layers.fused_moe.all2all_utils import get_ep_all2all_manager from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( initialize_mamba_ssu_backend, ) @@ -175,6 +176,12 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device): self.dp_size = self.parallel_config.data_parallel_size self.dp_rank = self.parallel_config.data_parallel_rank + # Detect EP all2all peer faults to prevent emitting corrupted output. + # Only meaningful for MoE + DP with an FT-capable all2all backend. + self.check_ep_fault = False + if self.dp_size > 1 and self.model_config.is_moe: + self.check_ep_fault = get_ep_all2all_manager().support_fault_tolerance + # Decode context parallelism. self.dcp_size = self.parallel_config.decode_context_parallel_size self.use_dcp = self.dcp_size > 1 @@ -1488,6 +1495,7 @@ def sample_tokens( num_sampled_tokens=num_sampled, main_stream=self.main_stream, copy_stream=self.output_copy_stream, + check_ep_fault=self.check_ep_fault, ) mm_inputs: tuple[list[torch.Tensor], torch.Tensor] | None = None From cac0e6a2da0f267982ed378b5e8c251c49387d15 Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Wed, 15 Jul 2026 18:23:32 +0800 Subject: [PATCH 11/16] [FT] validate single API server, fix clean_buffers ordering and worker state cleanup Signed-off-by: fangyuchu --- tests/test_config.py | 10 +++++++++ vllm/config/fault_tolerance.py | 6 +++--- vllm/config/parallel.py | 7 +++++++ .../base_device_communicator.py | 20 ++++++++++++++++-- .../fault_tolerance/engine_core_sentinel.py | 2 +- vllm/v1/worker/gpu_worker.py | 2 +- .../v1/worker/sentinel/gpu_worker_sentinel.py | 21 ++++++++++++------- 7 files changed, 53 insertions(+), 15 deletions(-) diff --git a/tests/test_config.py b/tests/test_config.py index 71e078ef3a2a..1fc00a8f8a94 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -1547,6 +1547,16 @@ def test_needs_dp_coordination( assert vllm_config.needs_dp_coordinator == expected_needs_coordinator +def test_fault_tolerance_requires_single_api_server(): + """Fault tolerance assumes one AsyncMPClient manages all engines, so it + is incompatible with API server scale-out (_api_process_count > 1).""" + with pytest.raises(ValueError, match="single API server"): + ParallelConfig(enable_fault_tolerance=True, _api_process_count=2) + + # Single API server (the FT-supported topology) is accepted. + ParallelConfig(enable_fault_tolerance=True, _api_process_count=1) + + def test_renderer_num_workers_with_mm_cache(): """Disallow renderer_num_workers > 1 when mm processor cache is enabled, since neither cache type is thread-safe.""" diff --git a/vllm/config/fault_tolerance.py b/vllm/config/fault_tolerance.py index d4b41c9c1a3c..7ed095d204f6 100644 --- a/vllm/config/fault_tolerance.py +++ b/vllm/config/fault_tolerance.py @@ -12,7 +12,7 @@ class FaultToleranceConfig: engine_recovery_timeout_sec: int = 120 """Timeout (in seconds) to wait for error handling instructions before raising an exception. If the EngineCore encounters an - error, it waits up to this many seconds for instructions on how - to handle the error. If no instructions are received within this - time, the original error is raised. + error, it waits up to this many seconds for vLLM to receive + instructions on how to handle the error and then recover from the fault. + If vLLM does not recover during this time, the original error is raised. """ diff --git a/vllm/config/parallel.py b/vllm/config/parallel.py index 85adb182c06a..949eb298a170 100644 --- a/vllm/config/parallel.py +++ b/vllm/config/parallel.py @@ -457,6 +457,13 @@ def _validate_parallel_config(self) -> Self: f"but found: {self._api_process_rank}" ) + if self.enable_fault_tolerance and self._api_process_count > 1: + raise ValueError( + "Fault tolerance requires a single API server process " + f"(--api-server-count=1), but got {self._api_process_count}. " + "The FT system assumes one AsyncMPClient manages all engines." + ) + if self.all2all_backend in ["pplx", "naive"]: logger.warning( "The '%s' all2all backend has been removed. " diff --git a/vllm/distributed/device_communicators/base_device_communicator.py b/vllm/distributed/device_communicators/base_device_communicator.py index f526ba6314f1..dc2671433891 100644 --- a/vllm/distributed/device_communicators/base_device_communicator.py +++ b/vllm/distributed/device_communicators/base_device_communicator.py @@ -105,14 +105,30 @@ def dispatch( raise NotImplementedError def query_active_mask(self) -> torch.Tensor: + """Return the all2all liveness mask for the EP ranks. + + Returns: + An int32 device tensor where 0 marks a live rank and 1 marks a + masked (dead/unreachable) rank. + """ raise NotImplementedError def query_fault(self) -> torch.Tensor: - """Returns has_fault scalar.""" + """Return a scalar bool tensor, True if a new fault appeared. + + Compares the current mask against the baseline recorded at the last + recovery point. + """ raise NotImplementedError def clean_buffers(self) -> None: - """Clean RDMA buffers and mask state during FT retry.""" + """Reset this rank's RDMA buffers and all2all mask state (rank-local). + + Post-fault cleanup: a dispatch/combine that hit a dead peer or timed + out can leave partially-written or stale tokens in the RDMA receive + buffer, so it is zeroed to stop the next forward from reading that + contaminated data. + """ raise NotImplementedError def set_num_sms(self, num_sms: int): diff --git a/vllm/v1/fault_tolerance/engine_core_sentinel.py b/vllm/v1/fault_tolerance/engine_core_sentinel.py index 02bf1ab0b38e..b9cfd631d8f8 100644 --- a/vllm/v1/fault_tolerance/engine_core_sentinel.py +++ b/vllm/v1/fault_tolerance/engine_core_sentinel.py @@ -79,7 +79,7 @@ def on_fault(self, exc: Exception): engine.batch_queue.clear() self.status_type = EngineStatusType.UNHEALTHY - self.fault_info = f"{type(exc).__name__}: {exc}" + self.fault_info = f"{type(exc).__name__}" logger.info( "[FT] Engine %d status -> UNHEALTHY:", self.engine_index, exc_info=exc ) diff --git a/vllm/v1/worker/gpu_worker.py b/vllm/v1/worker/gpu_worker.py index 6c031f793ed6..7c3d03366881 100644 --- a/vllm/v1/worker/gpu_worker.py +++ b/vllm/v1/worker/gpu_worker.py @@ -149,7 +149,7 @@ def __init__( self.elastic_ep_executor = ElasticEPScalingExecutor(self) self.worker_sentinel: WorkerSentinel | None = None if self.parallel_config.enable_fault_tolerance: - self.worker_sentinel = WorkerSentinel(worker=self, device=self.device) + self.worker_sentinel = WorkerSentinel(worker=self) # Buffers saved before sleep self._sleep_saved_buffers: dict[str, torch.Tensor] = {} self._sleep_rebuild_draft_metadata_buffers = False diff --git a/vllm/v1/worker/sentinel/gpu_worker_sentinel.py b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py index 41c0542c4553..80050cf9d391 100644 --- a/vllm/v1/worker/sentinel/gpu_worker_sentinel.py +++ b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py @@ -32,17 +32,17 @@ class WorkerSentinel: Methods are called via collective_rpc from EngineCoreSentinel. """ - def __init__(self, worker: "Worker", device: torch.device): + def __init__(self, worker: "Worker"): self.worker = worker - self.device = device self.dp_rank = worker.parallel_config.data_parallel_rank self.dp_size = worker.parallel_config.data_parallel_size self.data_parallel_master_ip = worker.parallel_config.data_parallel_master_ip all2all_backend = worker.parallel_config.all2all_backend - assert all2all_backend in FT_BACKEND_SET, ( - f"Fault tolerance requires an FT-capable all2all backend " - f"(one of {sorted(FT_BACKEND_SET)}), but got '{all2all_backend}'." - ) + if all2all_backend not in FT_BACKEND_SET: + raise ValueError( + f"Fault tolerance requires an FT-capable all2all backend " + f"(one of {sorted(FT_BACKEND_SET)}), but got '{all2all_backend}'." + ) def handle_command(self, ft_request: FaultToleranceRequest): """Dispatch an FT command by instruction name.""" @@ -54,6 +54,7 @@ def retry(self, ft_request: FaultToleranceRequest): params = ft_request.params self._clean_worker_state() if self.dp_size > 1: + get_ep_all2all_manager().clean_buffers() old_cpu_group = get_dp_group().cpu_group stateless_destroy_torch_distributed_process_group(old_cpu_group) world_size = self.worker.parallel_config.world_size @@ -65,7 +66,6 @@ def retry(self, ft_request: FaultToleranceRequest): self.dp_size, backend="gloo", ) - get_ep_all2all_manager().clean_buffers() def _clean_worker_state(self): model_runner = self.worker.model_runner @@ -76,9 +76,14 @@ def _clean_worker_state(self): runner._remove_request(req_id) else: model_runner.kv_connector_output = None + input_batch = model_runner.input_batch - for req_id in list(input_batch.req_id_to_index): + cached_req_ids = list(input_batch.req_id_to_index) + for req_id in cached_req_ids: + model_runner.requests.pop(req_id, None) + model_runner.num_prompt_logprobs.pop(req_id, None) input_batch.remove_request(req_id) + input_batch.condense() input_batch.refresh_metadata() input_batch.req_prompt_embeds.clear() From 90cb72e01afa1b5f8440b258deae180772afe796 Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Thu, 16 Jul 2026 16:53:51 +0800 Subject: [PATCH 12/16] [FT] make apply endpoint async and gate recovery by engine status Signed-off-by: fangyuchu --- .../serve/fault_tolerance/api_router.py | 47 ++++++++++++++----- vllm/v1/engine/core_client.py | 8 +++- .../fault_tolerance/engine_core_sentinel.py | 38 +++++++++++---- 3 files changed, 70 insertions(+), 23 deletions(-) diff --git a/vllm/entrypoints/serve/fault_tolerance/api_router.py b/vllm/entrypoints/serve/fault_tolerance/api_router.py index c1d1f2fed34e..960833af2155 100644 --- a/vllm/entrypoints/serve/fault_tolerance/api_router.py +++ b/vllm/entrypoints/serve/fault_tolerance/api_router.py @@ -4,7 +4,7 @@ import uuid from http import HTTPStatus -from fastapi import APIRouter, Depends, FastAPI, HTTPException, Request +from fastapi import APIRouter, BackgroundTasks, Depends, FastAPI, HTTPException, Request from fastapi.responses import JSONResponse from vllm.engine.protocol import EngineClient @@ -38,13 +38,13 @@ def _validate_payload(body: dict) -> tuple[str, dict]: "/fault_tolerance/apply", dependencies=[Depends(validate_json_request)], responses={ - HTTPStatus.OK.value: {"model": dict}, + HTTPStatus.ACCEPTED.value: {"model": dict}, HTTPStatus.BAD_REQUEST.value: {"model": ErrorResponse}, - HTTPStatus.REQUEST_TIMEOUT.value: {"model": ErrorResponse}, - HTTPStatus.INTERNAL_SERVER_ERROR.value: {"model": ErrorResponse}, }, ) -async def process_fault_tolerance_instruction(raw_request: Request): +async def process_fault_tolerance_instruction( + raw_request: Request, background_tasks: BackgroundTasks +): try: body = await raw_request.json() except json.JSONDecodeError as e: @@ -58,15 +58,36 @@ async def process_fault_tolerance_instruction(raw_request: Request): ) client: EngineClient = raw_request.app.state.engine_client + # Recovery runs cross-rank collective ops that only complete once every rank + # has been dispatched. Run it in the background and return immediately so the + # orchestrator can dispatch to all ranks without blocking; completion is + # observed by polling GET /fault_tolerance/status. + background_tasks.add_task(_run_fault_recovery, client, ft_request) + return JSONResponse( + status_code=HTTPStatus.ACCEPTED.value, + content={ + "message": "Request accepted; poll /fault_tolerance/status for updates.", + "request_id": ft_request.request_id, + }, + background=background_tasks, + ) + + +async def _run_fault_recovery( + client: EngineClient, ft_request: FaultToleranceRequest +) -> None: + """Drive recovery to completion after the 202 response is sent.""" try: - ft_result = await client.handle_fault(ft_request) - except Exception as e: - logger.error("Failed to handle fault: %s", e) - raise HTTPException(500, "Failed to handle fault.") from e - - if ft_result.success: - return JSONResponse({"message": "Instruction executed successfully."}) - raise HTTPException(500, f"Instruction failed: {ft_result.reason}") + result = await client.handle_fault(ft_request) + except Exception: + logger.exception("[FT] Recovery dispatch failed.") + return + if not result.success: + logger.error( + "[FT] Recovery failed for request %s: %s", + ft_request.request_id, + result.reason, + ) @router.get("/fault_tolerance/status") diff --git a/vllm/v1/engine/core_client.py b/vllm/v1/engine/core_client.py index c43d52c0d830..f83e32096a90 100644 --- a/vllm/v1/engine/core_client.py +++ b/vllm/v1/engine/core_client.py @@ -1231,7 +1231,13 @@ async def handle_fault( self, ft_request: FaultToleranceRequest ) -> FaultToleranceResult: res = await self.call_utility_async(FT_UTILITY_METHOD, ft_request) - return msgspec.convert(res, FaultToleranceResult) + result = msgspec.convert(res, FaultToleranceResult) + if not result.success: + status = self._engine_status.get(self.engine_ranks_managed[0]) + if status is not None: + status["last_ft_request_id"] = result.request_id + status["ft_error"] = result.reason + return result async def get_status(self): return { diff --git a/vllm/v1/fault_tolerance/engine_core_sentinel.py b/vllm/v1/fault_tolerance/engine_core_sentinel.py index b9cfd631d8f8..1d82dd9f7435 100644 --- a/vllm/v1/fault_tolerance/engine_core_sentinel.py +++ b/vllm/v1/fault_tolerance/engine_core_sentinel.py @@ -49,15 +49,27 @@ def __init__(self, engine: "EngineCoreProc", parallel_config): self._dp_reinit_epoch = 0 def handle_command(self, client_idx: int, call_id: int, ft_args: dict): - """Dispatch an FT command by instruction name and enqueue result.""" + """Dispatch an FT command by instruction name.""" ft_request = FaultToleranceRequest(**ft_args) - try: - result = run_method(self, ft_request.instruction, (ft_request,), {}) - except Exception as e: - logger.exception("[FT] Instruction '%s' failed", ft_request.instruction) + if self.status_type != EngineStatusType.UNHEALTHY: + reason = ( + f"[FT] Rejecting {ft_request.instruction} on engine " + f"{self.engine_index}: status is {self.status_type.name}" + ) + logger.warning(reason) result = FaultToleranceResult( - request_id=ft_request.request_id, success=False, reason=str(e) + request_id=ft_request.request_id, + success=False, + reason=reason, ) + else: + try: + result = run_method(self, ft_request.instruction, (ft_request,), {}) + except Exception as e: + logger.exception("[FT] Instruction '%s' failed", ft_request.instruction) + result = FaultToleranceResult( + request_id=ft_request.request_id, success=False, reason=str(e) + ) uo = UtilityOutput(call_id) uo.result = UtilityResult(msgspec.structs.asdict(result)) @@ -77,11 +89,19 @@ def on_fault(self, exc: Exception): engine._send_abort_outputs(aborted) if engine.batch_queue is not None: engine.batch_queue.clear() - - self.status_type = EngineStatusType.UNHEALTHY + if ( + hasattr(engine.model_executor, "is_failed") + and engine.model_executor.is_failed + ): + self.status_type = EngineStatusType.DEAD + else: + self.status_type = EngineStatusType.UNHEALTHY self.fault_info = f"{type(exc).__name__}" logger.info( - "[FT] Engine %d status -> UNHEALTHY:", self.engine_index, exc_info=exc + "[FT] Engine %d status -> %s:", + self.engine_index, + self.status_type.name, + exc_info=exc, ) self._push_status() From df0befe6a0f510c51006eba2ba83c35bf1f4bbde Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Fri, 17 Jul 2026 17:59:10 +0800 Subject: [PATCH 13/16] add e2e test for fault tolerance Signed-off-by: fangyuchu --- tests/v1/fault_tolerance/__init__.py | 2 + .../test_fault_tolerance_e2e.py | 247 ++++++++++++++++++ 2 files changed, 249 insertions(+) create mode 100644 tests/v1/fault_tolerance/__init__.py create mode 100644 tests/v1/fault_tolerance/test_fault_tolerance_e2e.py diff --git a/tests/v1/fault_tolerance/__init__.py b/tests/v1/fault_tolerance/__init__.py new file mode 100644 index 000000000000..208f01a7cb5e --- /dev/null +++ b/tests/v1/fault_tolerance/__init__.py @@ -0,0 +1,2 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project diff --git a/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py b/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py new file mode 100644 index 000000000000..ce53721e4c70 --- /dev/null +++ b/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py @@ -0,0 +1,247 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""End-to-end tests for the elastic fault-tolerance framework. + +Requires nixl_ep FT hardware; gated behind ``has_nixl_ep()``. +""" + +import contextlib +import os +import threading +import time + +import psutil +import pytest +import requests + +from tests.utils import multi_gpu_test +from vllm.utils.import_utils import has_nixl_ep + +MODEL_NAME = os.getenv("MODEL_NAME", "ibm-research/PowerMoE-3b") +DP_SIZE = int(os.getenv("DP_SIZE", "2")) + +# Fault-detection timeout budget: +# - CPU: Gloo DP allreduce timeout (10s) detects the dead peer. +# - nixl_ep: kernel masks the dead rank after Buffer's default timeout_ms=30000 (30s). +# - Deadline (45s): slowest fallback (30s) + margin. +CPU_DISTRIBUTED_TIMEOUT_S = 10 +FAULT_DETECTION_DEADLINE_S = 45 + + +def _ft_server_args() -> list[str]: + return [ + "--enforce-eager", + "--dtype", + "bfloat16", + "--max-model-len", + "2048", + "--max-num-seqs", + "128", + "--enable-expert-parallel", + "--all2all-backend", + "nixl_ep", + "--enable-fault-tolerance", + "--cpu-distributed-timeout-seconds", + str(CPU_DISTRIBUTED_TIMEOUT_S), + "--fault-tolerance-config", + '{"engine_recovery_timeout_sec": 120}', + ] + + +def _server_for_rank(servers, rank: int): + """Locate the server for a DP rank.""" + for server, sargs in servers: + if "--data-parallel-rank" in sargs: + idx = sargs.index("--data-parallel-rank") + if int(sargs[idx + 1]) == rank: + return server + raise AssertionError(f"no server found for DP rank {rank}") + + +def _get_ft_status(server) -> dict: + resp = requests.get(server.url_for("fault_tolerance/status"), timeout=10) + resp.raise_for_status() + return resp.json() + + +def _kill_worker_process(server) -> None: + """SIGKILL only the worker proc, leaving EngineCore and API server alive.""" + workers = [ + p + for p in psutil.Process(server.proc.pid).children(recursive=True) + if "Worker" in " ".join(p.cmdline()) + ] + assert len(workers) == 1, f"expected 1 worker proc, found: {workers}" + workers[0].kill() + + +def _poll_status(server, deadline_s: int, predicate): + """Poll ``/fault_tolerance/status`` until ``predicate(engine)`` is true. + + Returns the first matching engine dict, or ``None`` on timeout. Request + errors are suppressed so a briefly-unreachable server doesn't abort the poll. + """ + start = time.time() + while time.time() - start < deadline_s: + with contextlib.suppress(Exception): + for engine in _get_ft_status(server)["engines"]: + if predicate(engine): + return engine + time.sleep(1.0) + return None + + +def _wait_for_engine_fault(server, client, deadline_s: int): + """Drive the server and wait until an engine reports ``dead``/``unhealthy``. + + Requests run in the background so the engine keeps stepping into the failed + component. Both statuses match, so an unexpected one fails the caller's + assertion instead of looking like a hang. Returns the faulted engine dict, + or ``None`` on timeout. + """ + stop = threading.Event() + + def _drive(): + while not stop.is_set(): + with contextlib.suppress(Exception): # errors once faulted -- expected + client.completions.create( + model=MODEL_NAME, + prompt="Hello, my name is", + max_tokens=5, + temperature=0.0, + ) + time.sleep(1.0) + + driver = threading.Thread(target=_drive, daemon=True) + driver.start() + try: + return _poll_status( + server, + deadline_s, + lambda engine: engine["status"] in ("dead", "unhealthy"), + ) + finally: + stop.set() + driver.join(timeout=5) + + +def _wait_for_ft_apply_outcome(server, request_id: str, deadline_s: int) -> str | None: + """Wait until ``/status`` records the outcome of the given FT apply request. + + ``/apply`` dispatches recovery in a background task; once the engine records + the request id, its ``ft_error`` field holds the rejection reason (``None`` + on success). Returns that error, or ``None`` if the request succeeded or was + never recorded before the deadline (indistinguishable here). + """ + engine = _poll_status( + server, + deadline_s, + lambda engine: engine.get("last_ft_request_id") == request_id, + ) + return engine.get("ft_error") if engine else None + + +@pytest.mark.skipif(not has_nixl_ep(), reason="Requires nixl_ep all2all backend") +@multi_gpu_test(num_gpus=2) +def test_worker_kill_survivor_unhealthy_and_dead_rejects_retry(): + """One worker kill surfaces two status transitions at once. + + SIGKILLing only rank 1's worker leaves both EngineCores alive, so the same + fault is seen two ways: + + - Survivor (rank 0): detects the dead peer via Gloo allreduce / nixl_ep + kernel timeout. Its own executor is fine, so ``on_fault`` marks it + UNHEALTHY with a ``fault_info``. + - Victim (rank 1): detects its own executor failure and marks itself DEAD. + + Recovery is gated on UNHEALTHY: the DEAD engine accepts ``retry`` at the + HTTP layer (202 = background dispatch) but rejects it in the engine, + recording the reason as ``ft_error``. + """ + from tests.v1.distributed.test_external_lb_dp import ExternalLBServerManager + + manager = ExternalLBServerManager( + MODEL_NAME, + DP_SIZE, + api_server_count=1, # FT requires a single API server per engine + base_server_args=_ft_server_args(), + tp_size=1, + ) + + with manager as servers: + assert len(servers) == DP_SIZE + survivor = _server_for_rank(servers, 0) + victim = _server_for_rank(servers, 1) + + # 1. Confirm both engines are healthy and serving. + for server in (survivor, victim): + client = server.get_client() + client.completions.create( + model=MODEL_NAME, + prompt="Hello, my name is", + max_tokens=5, + temperature=0.0, + ) + status = _get_ft_status(server) + assert status["schema_version"] == 1, status + assert status["total_engines"] == 1, status # one engine per server + assert all(e["status"] == "healthy" for e in status["engines"]), status + + # 2. Kill only the victim's worker; both EngineCores stay alive. + _kill_worker_process(victim) + + # 3. Drive both engines so each keeps stepping into the failed + # component; poll in parallel so the deadline applies independently. + results: dict = {} + + def _wait(server, client, key): + results[key] = _wait_for_engine_fault( + server, client, FAULT_DETECTION_DEADLINE_S + ) + + pollers = [ + threading.Thread( + target=_wait, args=(survivor, survivor.get_client(), "survivor") + ), + threading.Thread( + target=_wait, args=(victim, victim.get_client(), "victim") + ), + ] + for p in pollers: + p.start() + for p in pollers: + p.join() + survivor_faulted = results["survivor"] + victim_faulted = results["victim"] + + assert survivor_faulted is not None, ( + "survivor did not report the peer fault within " + f"{FAULT_DETECTION_DEADLINE_S}s -- it likely hung" + ) + # The survivor's own executor is fine, so it must be UNHEALTHY, not DEAD. + assert survivor_faulted["status"] == "unhealthy", survivor_faulted + assert survivor_faulted.get("fault_info"), survivor_faulted + + assert victim_faulted is not None, ( + "victim did not report its worker's death within " + f"{FAULT_DETECTION_DEADLINE_S}s" + ) + assert victim_faulted["status"] == "dead", victim_faulted + + # 4. retry is accepted at the HTTP layer (202 = background dispatch)... + resp = requests.post( + victim.url_for("fault_tolerance/apply"), + json={"instruction": "retry", "params": {}}, + timeout=10, + ) + assert resp.status_code == 202, resp.text + request_id = resp.json()["request_id"] + + # 5. ...but the DEAD engine must reject it: recovery requires UNHEALTHY. + ft_error = _wait_for_ft_apply_outcome( + victim, request_id, FAULT_DETECTION_DEADLINE_S + ) + assert ft_error is not None, ( + "rejection was never recorded in /fault_tolerance/status" + ) + assert "status is DEAD" in ft_error, ft_error From a44ee3db42fdc4084c504039ba6f83eeffdbced3 Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Sat, 18 Jul 2026 16:57:13 +0800 Subject: [PATCH 14/16] add e2e test for retry recovery Signed-off-by: fangyuchu --- .../test_fault_tolerance_e2e.py | 294 +++++++++++++----- 1 file changed, 211 insertions(+), 83 deletions(-) diff --git a/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py b/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py index ce53721e4c70..b0fe39f9bd11 100644 --- a/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py +++ b/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py @@ -28,6 +28,78 @@ FAULT_DETECTION_DEADLINE_S = 45 +# Patches ``gpu.dp_utils.sync_cudagraph_and_dp_padding`` to raise on ``rank`` at +# a chosen step. Gated on VLLM_FT_TEST_INJECT_FAULT. +_FAULT_INJECT_SITECUSTOMIZE = """\ +import builtins +import os +import sys + +_SPEC = os.environ.get("VLLM_FT_TEST_INJECT_FAULT") +_MODULE = "vllm.v1.worker.gpu.dp_utils" +_ATTR = "sync_cudagraph_and_dp_padding" + +if _SPEC: + _f = dict(kv.split("=", 1) for kv in _SPEC.split(",")) + _RANK, _STEP = int(_f["rank"]), int(_f["step"]) + _steps = [0] + + def _patch(m): + import inspect + _orig = getattr(m, _ATTR) + _sig = inspect.signature(_orig) + def _wrapped(*args, **kwargs): + result = _orig(*args, **kwargs) + bound = _sig.bind(*args, **kwargs) + bound.apply_defaults() + dp_rank = bound.arguments.get("dp_rank") + if dp_rank == _RANK: + _steps[0] += 1 + if _steps[0] == _STEP: + raise RuntimeError( + "FT test fault injection (rank=%d step=%d)" % (_RANK, _STEP) + ) + return result + + setattr(m, _ATTR, _wrapped) + + _real_import = builtins.__import__ + + def _hook(name, *a, **k): + module = _real_import(name, *a, **k) + m = sys.modules.get(_MODULE) + # During vLLM's circular import the module lands in sys.modules before + # its functions are defined; hasattr guards against patching too early. + if ( + m is not None + and hasattr(m, _ATTR) + and not getattr(m, "_ft_patched", False) + ): + m._ft_patched = True + _patch(m) + return module + + builtins.__import__ = _hook +""" + + +def _install_fault_injection(monkeypatch, tmp_path, rank: int, step: int) -> None: + """Arrange for the DP-sync fn to raise on ``rank`` at serving ``step``. + + Writes a ``sitecustomize.py`` and prepends its dir to PYTHONPATH so every + vLLM subprocess picks it up; the fault spec is read from the environment. + """ + site_dir = tmp_path / "ft_inject" + site_dir.mkdir() + (site_dir / "sitecustomize.py").write_text(_FAULT_INJECT_SITECUSTOMIZE) + existing = os.environ.get("PYTHONPATH", "") + monkeypatch.setenv( + "PYTHONPATH", + str(site_dir) + (os.pathsep + existing if existing else ""), + ) + monkeypatch.setenv("VLLM_FT_TEST_INJECT_FAULT", f"rank={rank},step={step}") + + def _ft_server_args() -> list[str]: return [ "--enforce-eager", @@ -48,6 +120,19 @@ def _ft_server_args() -> list[str]: ] +def _ft_manager(): + """Build the shared DP+EP fault-tolerant server topology (one engine/server).""" + from tests.v1.distributed.test_external_lb_dp import ExternalLBServerManager + + return ExternalLBServerManager( + MODEL_NAME, + DP_SIZE, + api_server_count=1, # FT requires a single API server per engine + base_server_args=_ft_server_args(), + tp_size=1, + ) + + def _server_for_rank(servers, rank: int): """Locate the server for a DP rank.""" for server, sargs in servers: @@ -58,12 +143,42 @@ def _server_for_rank(servers, rank: int): raise AssertionError(f"no server found for DP rank {rank}") +def _complete(client): + """Issue the one standard completion the tests use everywhere.""" + return client.completions.create( + model=MODEL_NAME, + prompt="Hello, my name is", + max_tokens=5, + temperature=0.0, + timeout=10.0, + ) + + def _get_ft_status(server) -> dict: resp = requests.get(server.url_for("fault_tolerance/status"), timeout=10) resp.raise_for_status() return resp.json() +def _assert_serving_and_healthy(server) -> dict: + """Serve one request and assert every engine reports healthy. Returns status.""" + _complete(server.get_client()) + status = _get_ft_status(server) + assert all(e["status"] == "healthy" for e in status["engines"]), status + return status + + +def _apply_ft(server, instruction: str, params: dict | None = None) -> dict: + """POST an FT instruction; assert it is accepted (202) and return the body.""" + resp = requests.post( + server.url_for("fault_tolerance/apply"), + json={"instruction": instruction, "params": params or {}}, + timeout=10, + ) + assert resp.status_code == 202, resp.text + return resp.json() + + def _kill_worker_process(server) -> None: """SIGKILL only the worker proc, leaving EngineCore and API server alive.""" workers = [ @@ -91,48 +206,41 @@ def _poll_status(server, deadline_s: int, predicate): return None -def _wait_for_engine_fault(server, client, deadline_s: int): - """Drive the server and wait until an engine reports ``dead``/``unhealthy``. +def _wait_for_status(server, statuses, deadline_s: int = FAULT_DETECTION_DEADLINE_S): + """Poll until an engine reports one of ``statuses`` (a set of status strings).""" + return _poll_status(server, deadline_s, lambda e: e["status"] in statuses) + + +@contextlib.contextmanager +def _driving(*servers): + """Pump completions at each server in the background for the block's duration. - Requests run in the background so the engine keeps stepping into the failed - component. Both statuses match, so an unexpected one fails the caller's - assertion instead of looking like a hang. Returns the faulted engine dict, - or ``None`` on timeout. + Keeps every engine stepping into its failed component so a fault surfaces, + and lets each ``retry`` reach its cross-rank collective. Errors are expected + once faulted and are ignored. """ stop = threading.Event() - def _drive(): + def _drive(server): + client = server.get_client() while not stop.is_set(): - with contextlib.suppress(Exception): # errors once faulted -- expected - client.completions.create( - model=MODEL_NAME, - prompt="Hello, my name is", - max_tokens=5, - temperature=0.0, - ) - time.sleep(1.0) - - driver = threading.Thread(target=_drive, daemon=True) - driver.start() + with contextlib.suppress(Exception): + _complete(client) + time.sleep(0.2) + + threads = [threading.Thread(target=_drive, args=(s,), daemon=True) for s in servers] + for t in threads: + t.start() try: - return _poll_status( - server, - deadline_s, - lambda engine: engine["status"] in ("dead", "unhealthy"), - ) + yield finally: stop.set() - driver.join(timeout=5) + for t in threads: + t.join(timeout=2) def _wait_for_ft_apply_outcome(server, request_id: str, deadline_s: int) -> str | None: - """Wait until ``/status`` records the outcome of the given FT apply request. - - ``/apply`` dispatches recovery in a background task; once the engine records - the request id, its ``ft_error`` field holds the rejection reason (``None`` - on success). Returns that error, or ``None`` if the request succeeded or was - never recorded before the deadline (indistinguishable here). - """ + """Wait until ``/status`` records the outcome of the given FT apply request.""" engine = _poll_status( server, deadline_s, @@ -141,6 +249,69 @@ def _wait_for_ft_apply_outcome(server, request_id: str, deadline_s: int) -> str return engine.get("ft_error") if engine else None +@pytest.mark.skipif(not has_nixl_ep(), reason="Requires nixl_ep all2all backend") +@multi_gpu_test(num_gpus=2) +def test_injected_fault_retry_recovers_all_ranks(monkeypatch, tmp_path): + """An exception injected into the inference path drives full retry recovery. + + Injecting an exception into ``sync_cudagraph_and_dp_padding`` at a chosen + step on rank 1. + + - Rank 1 raises inside the busy loop and goes UNHEALTHY. + - Rank 0 detects the now-absent peer via the communication timeout and also + goes UNHEALTHY. + + Both being UNHEALTHY is the precondition for ``retry``. The fault is patched + into the DP-sync fn from the test (via a generated ``sitecustomize``). + """ + fault_step = int(os.getenv("FT_FAULT_STEP", "50")) + _install_fault_injection(monkeypatch, tmp_path, rank=1, step=fault_step) + + with _ft_manager() as servers: + assert len(servers) == DP_SIZE + rank0 = _server_for_rank(servers, 0) + rank1 = _server_for_rank(servers, 1) + + # 1. Both engines healthy and serving. + for server in (rank0, rank1): + status = _assert_serving_and_healthy(server) + assert status["schema_version"] == 1, status + assert status["total_engines"] == 1, status # one engine per server + + # 2. Drive both ranks so rank 1 accumulates execute_model steps and trips + # the injected fault; rank 0 then times out on the DP allreduce. + with _driving(rank0, rank1): + faulted = { + rank: _wait_for_status(server, {"unhealthy"}) + for rank, server in ((0, rank0), (1, rank1)) + } + + for rank, engine in faulted.items(): + assert engine is not None, ( + f"rank {rank} did not report UNHEALTHY within " + f"{FAULT_DETECTION_DEADLINE_S}s -- it likely hung" + ) + # The rank that raised carries the fault info from its own exception. + assert faulted[1].get("fault_info"), faulted[1] + + # 3. retry both engines. + for server in (rank0, rank1): + _apply_ft(server, "retry") + + # 4. Recovery completes: both engines return to healthy. + for rank, server in ((0, rank0), (1, rank1)): + healthy = _wait_for_status(server, {"healthy"}) + assert healthy is not None, ( + f"rank {rank} did not recover to healthy within " + f"{FAULT_DETECTION_DEADLINE_S}s" + ) + + # 5. Post-recovery inference works on both ranks. + for server in (rank0, rank1): + completion = _complete(server.get_client()) + assert completion.choices[0].text is not None + + @pytest.mark.skipif(not has_nixl_ep(), reason="Requires nixl_ep all2all backend") @multi_gpu_test(num_gpus=2) def test_worker_kill_survivor_unhealthy_and_dead_rejects_retry(): @@ -158,61 +329,24 @@ def test_worker_kill_survivor_unhealthy_and_dead_rejects_retry(): HTTP layer (202 = background dispatch) but rejects it in the engine, recording the reason as ``ft_error``. """ - from tests.v1.distributed.test_external_lb_dp import ExternalLBServerManager - - manager = ExternalLBServerManager( - MODEL_NAME, - DP_SIZE, - api_server_count=1, # FT requires a single API server per engine - base_server_args=_ft_server_args(), - tp_size=1, - ) - - with manager as servers: + with _ft_manager() as servers: assert len(servers) == DP_SIZE survivor = _server_for_rank(servers, 0) victim = _server_for_rank(servers, 1) # 1. Confirm both engines are healthy and serving. for server in (survivor, victim): - client = server.get_client() - client.completions.create( - model=MODEL_NAME, - prompt="Hello, my name is", - max_tokens=5, - temperature=0.0, - ) - status = _get_ft_status(server) - assert status["schema_version"] == 1, status - assert status["total_engines"] == 1, status # one engine per server - assert all(e["status"] == "healthy" for e in status["engines"]), status + _assert_serving_and_healthy(server) # 2. Kill only the victim's worker; both EngineCores stay alive. _kill_worker_process(victim) - # 3. Drive both engines so each keeps stepping into the failed - # component; poll in parallel so the deadline applies independently. - results: dict = {} - - def _wait(server, client, key): - results[key] = _wait_for_engine_fault( - server, client, FAULT_DETECTION_DEADLINE_S - ) - - pollers = [ - threading.Thread( - target=_wait, args=(survivor, survivor.get_client(), "survivor") - ), - threading.Thread( - target=_wait, args=(victim, victim.get_client(), "victim") - ), - ] - for p in pollers: - p.start() - for p in pollers: - p.join() - survivor_faulted = results["survivor"] - victim_faulted = results["victim"] + # 3. Drive both engines so each keeps stepping into the failed component. + # Both fault ~together via timeout, so sequential polling under one + # driving context still sees each within the deadline. + with _driving(survivor, victim): + survivor_faulted = _wait_for_status(survivor, {"dead", "unhealthy"}) + victim_faulted = _wait_for_status(victim, {"dead", "unhealthy"}) assert survivor_faulted is not None, ( "survivor did not report the peer fault within " @@ -229,13 +363,7 @@ def _wait(server, client, key): assert victim_faulted["status"] == "dead", victim_faulted # 4. retry is accepted at the HTTP layer (202 = background dispatch)... - resp = requests.post( - victim.url_for("fault_tolerance/apply"), - json={"instruction": "retry", "params": {}}, - timeout=10, - ) - assert resp.status_code == 202, resp.text - request_id = resp.json()["request_id"] + request_id = _apply_ft(victim, "retry")["request_id"] # 5. ...but the DEAD engine must reject it: recovery requires UNHEALTHY. ft_error = _wait_for_ft_apply_outcome( From 4f6fe5839ff06cf0493c1c34cbceb6e4a147100f Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Tue, 21 Jul 2026 19:03:05 +0800 Subject: [PATCH 15/16] refactor e2e tests with concurrent polling and add Buildkite CI config (#250) Signed-off-by: fangyuchu --- .buildkite/test_areas/fault_tolerance.yaml | 26 ++++ .../test_fault_tolerance_e2e.py | 127 +++++++++--------- 2 files changed, 91 insertions(+), 62 deletions(-) create mode 100644 .buildkite/test_areas/fault_tolerance.yaml diff --git a/.buildkite/test_areas/fault_tolerance.yaml b/.buildkite/test_areas/fault_tolerance.yaml new file mode 100644 index 000000000000..e2f700a8bd86 --- /dev/null +++ b/.buildkite/test_areas/fault_tolerance.yaml @@ -0,0 +1,26 @@ +group: Fault Tolerance +depends_on: + - image-build +steps: +- label: Fault Tolerance E2E (2xH100) + key: fault-tolerance-e2e-2xh100 + timeout_in_minutes: 35 + device: h100 + num_devices: 2 + working_dir: "/vllm-workspace/tests" + source_file_dependencies: + - vllm/v1/fault_tolerance/ + - vllm/v1/worker/sentinel/ + - vllm/entrypoints/serve/fault_tolerance/ + - vllm/distributed/elastic_ep/ + - vllm/distributed/device_communicators/ + - vllm/v1/engine/ + - vllm/v1/worker/ + - tests/v1/fault_tolerance/ + - tests/v1/distributed/test_external_lb_dp.py + commands: + # Base image has no nixl; install it or has_nixl_ep() skips the tests. + - bash /vllm-workspace/.buildkite/scripts/install-kv-connectors.sh + # https://github.com/NVIDIA/nccl/issues/1838 + - export NCCL_CUMEM_HOST_ENABLE=0 + - pytest -v -s v1/fault_tolerance/test_fault_tolerance_e2e.py diff --git a/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py b/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py index b0fe39f9bd11..406cc6827949 100644 --- a/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py +++ b/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py @@ -9,16 +9,18 @@ import os import threading import time +from concurrent.futures import ThreadPoolExecutor +from typing import Any import psutil import pytest import requests -from tests.utils import multi_gpu_test +from tests.utils import RemoteOpenAIServer, multi_gpu_test from vllm.utils.import_utils import has_nixl_ep MODEL_NAME = os.getenv("MODEL_NAME", "ibm-research/PowerMoE-3b") -DP_SIZE = int(os.getenv("DP_SIZE", "2")) +DP_SIZE = 2 # Fault-detection timeout budget: # - CPU: Gloo DP allreduce timeout (10s) detects the dead peer. @@ -154,18 +156,25 @@ def _complete(client): ) +def _in_parallel(fn, servers) -> list: + """Run ``fn(server)`` for all servers concurrently; return results in order.""" + with ThreadPoolExecutor(max_workers=len(servers)) as ex: + return list(ex.map(fn, servers)) + + def _get_ft_status(server) -> dict: resp = requests.get(server.url_for("fault_tolerance/status"), timeout=10) resp.raise_for_status() return resp.json() -def _assert_serving_and_healthy(server) -> dict: - """Serve one request and assert every engine reports healthy. Returns status.""" - _complete(server.get_client()) - status = _get_ft_status(server) - assert all(e["status"] == "healthy" for e in status["engines"]), status - return status +def _assert_serving_and_healthy(servers) -> None: + """Wait until every engine is healthy, then serve one request per server.""" + healthy = _wait_for_engines( + list(servers), match_key="status", match_values={"healthy"} + ) + assert all(healthy), healthy + _in_parallel(lambda s: _complete(s.get_client()), servers) def _apply_ft(server, instruction: str, params: dict | None = None) -> dict: @@ -190,34 +199,40 @@ def _kill_worker_process(server) -> None: workers[0].kill() -def _poll_status(server, deadline_s: int, predicate): - """Poll ``/fault_tolerance/status`` until ``predicate(engine)`` is true. +def _wait_for_engines( + servers: list[RemoteOpenAIServer], + match_key: str, + match_values: set[str], + deadline_s: int = FAULT_DETECTION_DEADLINE_S, +) -> list[dict[str, Any] | None]: + """Poll ``/fault_tolerance/status`` until each server's engine status matches. - Returns the first matching engine dict, or ``None`` on timeout. Request - errors are suppressed so a briefly-unreachable server doesn't abort the poll. + A server matches when its engine-status dict has ``match_key`` equal to + one of ``match_values``. Returns one engine-status dict per server. Servers still + unmatched after ``deadline_s`` get None. """ + results: dict[int, dict[str, Any]] = {} + pending = dict(enumerate(servers)) start = time.time() - while time.time() - start < deadline_s: - with contextlib.suppress(Exception): - for engine in _get_ft_status(server)["engines"]: - if predicate(engine): - return engine - time.sleep(1.0) - return None - - -def _wait_for_status(server, statuses, deadline_s: int = FAULT_DETECTION_DEADLINE_S): - """Poll until an engine reports one of ``statuses`` (a set of status strings).""" - return _poll_status(server, deadline_s, lambda e: e["status"] in statuses) + while pending and time.time() - start < deadline_s: + for i, server in list(pending.items()): + with contextlib.suppress(Exception): + for engine_status in _get_ft_status(server)["engines"]: + if engine_status.get(match_key) in match_values: + results[i] = engine_status + del pending[i] + break + if pending: + time.sleep(1.0) + return [results.get(i) for i in range(len(servers))] @contextlib.contextmanager def _driving(*servers): """Pump completions at each server in the background for the block's duration. - Keeps every engine stepping into its failed component so a fault surfaces, - and lets each ``retry`` reach its cross-rank collective. Errors are expected - once faulted and are ignored. + Keeps every engine stepping into its failed component so a fault surfaces. + Errors are expected once faulted and are ignored. """ stop = threading.Event() @@ -240,13 +255,14 @@ def _drive(server): def _wait_for_ft_apply_outcome(server, request_id: str, deadline_s: int) -> str | None: - """Wait until ``/status`` records the outcome of the given FT apply request.""" - engine = _poll_status( - server, - deadline_s, - lambda engine: engine.get("last_ft_request_id") == request_id, - ) - return engine.get("ft_error") if engine else None + """Wait until ``/fault_tolerance/status`` records the FT apply outcome.""" + engine_status = _wait_for_engines( + [server], + match_key="last_ft_request_id", + match_values={request_id}, + deadline_s=deadline_s, + )[0] + return engine_status.get("ft_error") if engine_status else None @pytest.mark.skipif(not has_nixl_ep(), reason="Requires nixl_ep all2all backend") @@ -273,43 +289,30 @@ def test_injected_fault_retry_recovers_all_ranks(monkeypatch, tmp_path): rank1 = _server_for_rank(servers, 1) # 1. Both engines healthy and serving. - for server in (rank0, rank1): - status = _assert_serving_and_healthy(server) - assert status["schema_version"] == 1, status - assert status["total_engines"] == 1, status # one engine per server + _assert_serving_and_healthy((rank0, rank1)) # 2. Drive both ranks so rank 1 accumulates execute_model steps and trips # the injected fault; rank 0 then times out on the DP allreduce. with _driving(rank0, rank1): - faulted = { - rank: _wait_for_status(server, {"unhealthy"}) - for rank, server in ((0, rank0), (1, rank1)) - } + faulted = _wait_for_engines( + [rank0, rank1], match_key="status", match_values={"unhealthy"} + ) - for rank, engine in faulted.items(): - assert engine is not None, ( + for rank, engine_status in enumerate(faulted): + assert engine_status is not None, ( f"rank {rank} did not report UNHEALTHY within " f"{FAULT_DETECTION_DEADLINE_S}s -- it likely hung" ) # The rank that raised carries the fault info from its own exception. + assert faulted[1] is not None assert faulted[1].get("fault_info"), faulted[1] # 3. retry both engines. for server in (rank0, rank1): _apply_ft(server, "retry") - # 4. Recovery completes: both engines return to healthy. - for rank, server in ((0, rank0), (1, rank1)): - healthy = _wait_for_status(server, {"healthy"}) - assert healthy is not None, ( - f"rank {rank} did not recover to healthy within " - f"{FAULT_DETECTION_DEADLINE_S}s" - ) - - # 5. Post-recovery inference works on both ranks. - for server in (rank0, rank1): - completion = _complete(server.get_client()) - assert completion.choices[0].text is not None + # 4. Recovery completes: both engines return to healthy and serve again. + _assert_serving_and_healthy((rank0, rank1)) @pytest.mark.skipif(not has_nixl_ep(), reason="Requires nixl_ep all2all backend") @@ -335,18 +338,18 @@ def test_worker_kill_survivor_unhealthy_and_dead_rejects_retry(): victim = _server_for_rank(servers, 1) # 1. Confirm both engines are healthy and serving. - for server in (survivor, victim): - _assert_serving_and_healthy(server) + _assert_serving_and_healthy((survivor, victim)) # 2. Kill only the victim's worker; both EngineCores stay alive. _kill_worker_process(victim) # 3. Drive both engines so each keeps stepping into the failed component. - # Both fault ~together via timeout, so sequential polling under one - # driving context still sees each within the deadline. with _driving(survivor, victim): - survivor_faulted = _wait_for_status(survivor, {"dead", "unhealthy"}) - victim_faulted = _wait_for_status(victim, {"dead", "unhealthy"}) + survivor_faulted, victim_faulted = _wait_for_engines( + [survivor, victim], + match_key="status", + match_values={"dead", "unhealthy"}, + ) assert survivor_faulted is not None, ( "survivor did not report the peer fault within " From 5be79f80bb55669a94dcd7cad2f5f6e87d74e7b9 Mon Sep 17 00:00:00 2001 From: fangyuchu Date: Thu, 23 Jul 2026 21:18:03 +0800 Subject: [PATCH 16/16] set cpu timeout to default value of nixl-ep in test (#251) Signed-off-by: fangyuchu --- tests/v1/fault_tolerance/test_fault_tolerance_e2e.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py b/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py index 406cc6827949..f6d15343b56f 100644 --- a/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py +++ b/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py @@ -23,10 +23,10 @@ DP_SIZE = 2 # Fault-detection timeout budget: -# - CPU: Gloo DP allreduce timeout (10s) detects the dead peer. +# - CPU: Gloo DP allreduce timeout (30s) detects the dead peer. # - nixl_ep: kernel masks the dead rank after Buffer's default timeout_ms=30000 (30s). # - Deadline (45s): slowest fallback (30s) + margin. -CPU_DISTRIBUTED_TIMEOUT_S = 10 +CPU_DISTRIBUTED_TIMEOUT_S = 30 FAULT_DETECTION_DEADLINE_S = 45