diff --git a/lightllm/server/api_http_pd.py b/lightllm/server/api_http_pd.py index 1d8b2112f..c23e5e7a7 100644 --- a/lightllm/server/api_http_pd.py +++ b/lightllm/server/api_http_pd.py @@ -57,7 +57,7 @@ async def register_and_keep_alive(websocket: WebSocket): logger.exception(str(e)) finally: logger.error(f"client {regist_json} removed") - await g_objs.httpserver_manager.remove_pd(regist_json) + await g_objs.httpserver_manager.remove_pd(regist_json, websocket) return diff --git a/lightllm/server/httpserver/async_queue.py b/lightllm/server/httpserver/async_queue.py index a9f0c9068..866d75726 100644 --- a/lightllm/server/httpserver/async_queue.py +++ b/lightllm/server/httpserver/async_queue.py @@ -7,9 +7,9 @@ def __init__(self): self.event = asyncio.Event() self.lock = asyncio.Lock() - async def wait_to_ready(self): + async def wait_to_ready(self, timeout=3): try: - await asyncio.wait_for(self.event.wait(), timeout=3) + await asyncio.wait_for(self.event.wait(), timeout=timeout) except asyncio.TimeoutError: pass @@ -26,7 +26,7 @@ async def put(self, obj): self.event.set() return - async def wait_to_get_all_data(self): - await self.wait_to_ready() + async def wait_to_get_all_data(self, timeout=3): + await self.wait_to_ready(timeout=timeout) handle_list = await self.get_all_data() return handle_list diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 195ef17e4..7e04ba81c 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -311,6 +311,40 @@ def alloc_req_id(self, sampling_params): assert False, "dead code path" return group_request_id + async def _alloc_req_objs(self, group_request_id, prompt_ids, sampling_params) -> List[Req]: + """Allocate and initialize request slots without leaking partial allocations.""" + alloced_req_indexes = [] + req_objs = [] + try: + while len(alloced_req_indexes) < sampling_params.n: + alloc_req_index = await self.shm_req_manager.async_alloc_req_index() + sleep_time = 0.1 + while alloc_req_index is None: + await asyncio.sleep(sleep_time) + sleep_time *= 1.1 + sleep_time = min(1, sleep_time) + + alloc_req_index = await self.shm_req_manager.async_alloc_req_index() + alloced_req_indexes.append(alloc_req_index) + + for i, req_index in enumerate(alloced_req_indexes): + req_obj = await self.shm_req_manager.async_get_req_obj_by_index(req_index) + req_objs.append(req_obj) + req_obj.init( + group_request_id + i, + prompt_ids, + sampling_params, + self.tokenizer, + chunked_prefill_size=self.args.chunked_prefill_size, + ) + return req_objs + except BaseException: + for req_obj in req_objs: + await self.shm_req_manager.async_put_back_req_obj(req_obj) + for req_index in alloced_req_indexes: + await self.shm_req_manager.async_release_req_index(req_index) + raise + async def generate( self, prompt: Union[str, List[int]], @@ -414,28 +448,7 @@ async def generate( raise PDPrefillNodeStopGenToken(group_request_id=group_request_id) # 申请资源并存储 - alloced_req_indexes = [] - while len(alloced_req_indexes) < sampling_params.n: - alloc_req_index = await self.shm_req_manager.async_alloc_req_index() - sleep_time = 0.1 - while alloc_req_index is None: - await asyncio.sleep(sleep_time) - sleep_time *= 1.1 - sleep_time = min(1, sleep_time) - - alloc_req_index = await self.shm_req_manager.async_alloc_req_index() - alloced_req_indexes.append(alloc_req_index) - req_objs: List[Req] = [] - for i, req_index in enumerate(alloced_req_indexes): - req_obj = await self.shm_req_manager.async_get_req_obj_by_index(req_index) - req_obj.init( - group_request_id + i, - prompt_ids, - sampling_params, - self.tokenizer, - chunked_prefill_size=self.args.chunked_prefill_size, - ) - req_objs.append(req_obj) + req_objs = await self._alloc_req_objs(group_request_id, prompt_ids, sampling_params) self._log_stage_timing( group_request_id, start_time, diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index a69364c95..ccd45181e 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -34,6 +34,48 @@ async def timer_log(manager: HttpServerManager): return +async def _abort_pd_request( + manager: HttpServerManager, + group_req_id: int, + generation_tasks: Dict[int, asyncio.Task], +) -> bool: + generation_task = generation_tasks.get(group_req_id) + if generation_task is not None and not generation_task.done() and not generation_task.cancelling(): + # The request may still be waiting for a shared-memory slot and therefore + # not yet be visible to HttpServerManager.abort(). Cancelling its owning + # task makes the abort effective at every point in the admission path. + generation_task.cancel() + + return (await manager.abort(group_req_id)) or generation_task is not None + + +async def _recv_or_raise_on_background_failure(websocket, background_tasks): + recv_task = asyncio.create_task(websocket.recv()) + try: + done, _ = await asyncio.wait( + [recv_task, *background_tasks], + return_when=asyncio.FIRST_COMPLETED, + ) + for task in background_tasks: + if task not in done: + continue + + recv_task.cancel() + await asyncio.gather(recv_task, return_exceptions=True) + if task.cancelled(): + raise asyncio.CancelledError + error = task.exception() + if error is not None: + raise error + raise RuntimeError("PD connection background task exited unexpectedly") + + return recv_task.result() + finally: + if not recv_task.done(): + recv_task.cancel() + await asyncio.gather(recv_task, return_exceptions=True) + + async def pd_handle_loop(manager: HttpServerManager): if manager.args.host in ["127.0.0.1", "localhost"]: logger.error("pd mode must specify host ip, not use 127.0.0.1 or localhost") @@ -84,6 +126,8 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O while True: forwarding_tokens_task = None heartbeat_task = None + connection_failure = None + generation_tasks: Dict[int, asyncio.Task] = {} try: uri = f"ws://{pd_master_obj.host_ip_port}/pd_register" async with websockets.connect( @@ -113,18 +157,22 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O # 转发任务 forwarding_tokens_task = asyncio.create_task(_up_tokens_to_pd_master(forwarding_queue, websocket)) heartbeat_task = asyncio.create_task(_send_heartbeat_to_pd_master(websocket)) + connection_failure = asyncio.get_running_loop().create_future() group_req_id_to_event: Dict[int, asyncio.Event] = weakref.WeakValueDictionary() # 接收 pd master 发来的请求,并推理后,将生成的token转发回pd master。 while True: - recv_bytes = await websocket.recv() + recv_bytes = await _recv_or_raise_on_background_failure( + websocket, + (forwarding_tokens_task, heartbeat_task, connection_failure), + ) obj = pickle.loads(recv_bytes) if obj[0] == ObjType.REQ: prompt, sampling_params, multimodal_params = obj[1] group_req_id = sampling_params.group_request_id pd_event = asyncio.Event() group_req_id_to_event[group_req_id] = pd_event - asyncio.create_task( + generation_task = asyncio.create_task( _pd_process_generate( manager=manager, prompt=prompt, @@ -135,18 +183,22 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O pd_event=pd_event, ) ) + generation_tasks[group_req_id] = generation_task + + def remove_generation_task(done_task, request_id=group_req_id): + if generation_tasks.get(request_id) is done_task: + generation_tasks.pop(request_id, None) + if not done_task.cancelled(): + error = done_task.exception() + if error is not None and not connection_failure.done(): + connection_failure.set_exception(error) + + generation_task.add_done_callback(remove_generation_task) elif obj[0] == ObjType.ABORT: group_req_id = obj[1] logger.warning(f"recv cmd aborted req id {group_req_id}") - if not (await manager.abort(group_req_id)): - - async def delayed_abort_task(group_req_id, retry_count): - for _ in range(retry_count): - await asyncio.sleep(5.0) - if await manager.abort(group_req_id): - break - - asyncio.create_task(delayed_abort_task(group_req_id=group_req_id, retry_count=4)) + group_req_id_to_event.pop(group_req_id, None) + await _abort_pd_request(manager, group_req_id, generation_tasks) elif obj[0] == ObjType.PD_REQ_DECODE_NODE_INFO: _, group_req_id, decode_node_info = obj @@ -169,7 +221,14 @@ async def delayed_abort_task(group_req_id, retry_count): logger.exception(str(e)) finally: child_tasks = [task for task in (forwarding_tokens_task, heartbeat_task) if task is not None] + child_tasks.extend(generation_tasks.values()) + if connection_failure is not None: + child_tasks.append(connection_failure) for task in child_tasks: + if task.done(): + continue + if isinstance(task, asyncio.Task) and task.cancelling(): + continue task.cancel() if child_tasks: await asyncio.gather(*child_tasks, return_exceptions=True) @@ -235,8 +294,12 @@ async def _pd_process_generate( await forwarding_queue.put((sub_req_id, request_output, metadata, finish_status)) except PDPrefillNodeStopGenToken as e: logger.info(f"pd prefill node stop gen token for group_request_id {e.group_request_id}") - except BaseException as e: - logger.error(str(e)) + except asyncio.CancelledError: + logger.info(f"pd request task cancelled for group_request_id {sampling_params.group_request_id}") + raise + except BaseException: + logger.exception(f"PD request generation failed for group_request_id {sampling_params.group_request_id}") + raise # 转发token的task diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 8d307abdc..550706bf5 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -29,6 +29,8 @@ logger = init_logger(__name__) +_PREFILL_TOKEN_WAIT_TIMEOUT = 5 + class HttpServerManagerForPDMaster: def __init__( @@ -74,11 +76,26 @@ def is_healthy(self): return False async def register_pd(self, pd_info_json, websocket): - self.pd_manager.register_pd(pd_info_json, websocket) + replaced_node = self.pd_manager.register_pd(pd_info_json, websocket) + if replaced_node is not None: + await self._fail_requests_for_node(replaced_node, "PD node connection was replaced") + return + + async def remove_pd(self, pd_info_json, websocket): + removed_node = self.pd_manager.remove_pd(pd_info_json, websocket) + if removed_node is not None: + await self._fail_requests_for_node(removed_node, "PD node disconnected") return - async def remove_pd(self, pd_info_json): - self.pd_manager.remove_pd(pd_info_json) + async def _fail_requests_for_node(self, node: PD_Client_Obj, reason: str): + error_message = f"{reason}: {node.mode} node {node.client_ip_port}" + affected_requests = [ + req_status + for req_status in list(self.req_id_to_out_inf.values()) + if req_status.p_node is node or req_status.d_node is node + ] + for req_status in affected_requests: + await req_status.fail(RuntimeError(error_message)) return async def update_req_status(self, upkv_status: PDUpKVStatus): @@ -395,6 +412,7 @@ async def fetch_pd_stream( group_request_id=group_request_id, stage="prefill", ) + req_status.raise_if_failed() except ServerBusyError: logger.warning(f"group_request_id: {group_request_id} wait prefill prompt ids time out") raise @@ -415,6 +433,7 @@ async def fetch_pd_stream( group_request_id=group_request_id, stage="decode", ) + req_status.raise_if_failed() except ServerBusyError: logger.warning(f"group_request_id: {group_request_id} wait decode stage time out err, server is busy now.") raise @@ -427,31 +446,61 @@ async def fetch_pd_stream( pickle.dumps((ObjType.PD_REQ_DECODE_NODE_INFO, group_request_id, decode_node_info)) ) - first_token_gen = False + first_token_emitted = False + needs_prefill_first_token = decode_node_info.ready_kv_len != len(prompt_ids) - 1 + buffered_decode_tokens = [] + prefill_token_deadline = None + prompt_cache_len_from_prefill = None while True: - await req_status.wait_to_ready() + wait_timeout = 5 + if prefill_token_deadline is not None: + wait_timeout = min(wait_timeout, max(0, prefill_token_deadline - time.monotonic())) + new_tokens = await req_status.out_tokens.wait_to_get_all_data(timeout=wait_timeout) + req_status.raise_if_failed() + assert group_request_id in self.req_id_to_out_inf, f"error state req_id {group_request_id}" if await request.is_disconnected(): raise ClientDisconnected( group_request_id=group_request_id, reason="fetch_pd_stream decode period check network disconnected", ) - if await req_status.can_read(self.req_id_to_out_inf): - token_list = await req_status.pop_all_tokens() - for sub_req_id, request_output, metadata, finish_status in token_list: - output_index = metadata.get("count_output_tokens") - # 因为 pd 的 prefill 和 decode 节点都有可能上报首token,所以需要做一下过滤。 - if output_index == 1: - if first_token_gen is False: - first_token_gen = True - node_run_mode = metadata.pop("node_mode", None) - if node_run_mode == "prefill": - if old_max_new_tokens != 1 and finish_status.is_finished_length(): - finish_status = FinishStatus(FinishStatus.NO_FINISH) - yield sub_req_id, request_output, metadata, finish_status - else: - continue - else: - yield sub_req_id, request_output, metadata, finish_status + + if not new_tokens and not buffered_decode_tokens: + continue + + # P 首 token 携带缓存统计,需先于已缓冲的 D 输出返回。 + if needs_prefill_first_token: + prefill_token = next((token for token in new_tokens if token[2].get("node_mode") == "prefill"), None) + if prefill_token is None: + buffered_decode_tokens.extend(new_tokens) + if prefill_token_deadline is None: + prefill_token_deadline = time.monotonic() + _PREFILL_TOKEN_WAIT_TIMEOUT + if time.monotonic() < prefill_token_deadline: + continue + logger.warning(f"{group_request_id}: prefill token missing; releasing decode output") + new_tokens = buffered_decode_tokens + buffered_decode_tokens = [] + else: + new_tokens.remove(prefill_token) + new_tokens = [prefill_token, *buffered_decode_tokens, *new_tokens] + buffered_decode_tokens.clear() + needs_prefill_first_token = False + prefill_token_deadline = None + + for sub_req_id, request_output, metadata, finish_status in new_tokens: + output_index = metadata.get("count_output_tokens") + node_run_mode = metadata.pop("node_mode", None) + if output_index == 1: + # D 首 token 可能是 KV 传输失败产生的唯一结束标记,不能按重复 token 丢弃。 + if first_token_emitted and not (node_run_mode == "decode" and finish_status.is_finished()): + continue + first_token_emitted = True + if node_run_mode == "prefill": + prompt_cache_len_from_prefill = metadata.get("prompt_cache_len", 0) + if old_max_new_tokens != 1 and finish_status.is_finished_length(): + finish_status = FinishStatus(FinishStatus.NO_FINISH) + if prompt_cache_len_from_prefill is not None: + metadata["prompt_cache_len"] = prompt_cache_len_from_prefill + yield sub_req_id, request_output, metadata, finish_status return @@ -595,18 +644,15 @@ async def handle_loop(self): group_req_id = convert_sub_id_to_group_id(sub_req_id) try: req_status: ReqStatus = self.req_id_to_out_inf[group_req_id] - async with req_status.lock: - req_status.out_token_info_list.append((sub_req_id, text, metadata, finish_status)) - req_status.event.set() + await req_status.out_tokens.put((sub_req_id, text, metadata, finish_status)) except: pass elif obj[0] == ObjType.PD_UPLOAD_PREFILL_PROMPT_IDS: _, group_req_id, prompt_ids = obj try: req_status: ReqStatus = self.req_id_to_out_inf[group_req_id] - async with req_status.lock: - req_status.prefill_prompt_ids_event.prompt_ids = prompt_ids - req_status.prefill_prompt_ids_event.set() + req_status.prefill_prompt_ids_event.prompt_ids = prompt_ids + req_status.prefill_prompt_ids_event.set() except: logger.error( f"PD_UPLOAD_PREFILL_PROMPT_IDS fail find req status for group_req_id: {group_req_id}" @@ -629,34 +675,26 @@ def _split_max_new_tokens(self, max_new_tokens: int) -> List[int]: class ReqStatus: def __init__(self, req_id, p_node, d_node) -> None: self.req_id = req_id - self.lock = asyncio.Lock() - self.event = asyncio.Event() + self.out_tokens = AsyncQueue() self.up_status_event = asyncio.Event() self.prefill_prompt_ids_event = asyncio.Event() - self.out_token_info_list: List[Tuple[int, str, dict, FinishStatus]] = [] self.p_node: PD_Client_Obj = p_node self.d_node: PD_Client_Obj = d_node + self.error: Optional[BaseException] = None + + async def fail(self, error: BaseException): + async with self.out_tokens.lock: + if self.error is None: + self.error = error + self.out_tokens.event.set() + self.up_status_event.set() + self.prefill_prompt_ids_event.set() + return - async def wait_to_ready(self): - try: - await asyncio.wait_for(self.event.wait(), timeout=5) - except asyncio.TimeoutError: - pass - - async def can_read(self, req_id_to_out_inf): - async with self.lock: - self.event.clear() - assert self.req_id in req_id_to_out_inf, f"error state req_id {self.req_id}" - if len(self.out_token_info_list) == 0: - return False - else: - return True - - async def pop_all_tokens(self): - async with self.lock: - ans = self.out_token_info_list.copy() - self.out_token_info_list.clear() - return ans + def raise_if_failed(self): + if self.error is not None: + raise self.error + return class PDManager: @@ -720,6 +758,7 @@ async def check_pd_nodes_health(self): def register_pd(self, pd_info_json, websocket): pd_client = PD_Client_Obj(**pd_info_json) + replaced_node = self.url_to_pd_nodes.get(pd_client.client_ip_port) client_max_req_total_len = pd_client.start_args["max_req_total_len"] if client_max_req_total_len != self.args.max_req_total_len: logger.error( @@ -756,10 +795,16 @@ def register_pd(self, pd_info_json, websocket): self.selector.update_nodes(self.prefill_nodes, self.decode_nodes) logger.info(f"mode: {pd_client.mode} url: {pd_client.client_ip_port} registed") - return + if replaced_node is not None and replaced_node.websocket is not websocket: + return replaced_node + return None - def remove_pd(self, pd_info_json): + def remove_pd(self, pd_info_json, websocket): pd_client = PD_Client_Obj(**pd_info_json) + registered_node = self.url_to_pd_nodes.get(pd_client.client_ip_port) + if registered_node is None or registered_node.websocket is not websocket: + logger.info(f"ignore stale disconnect for {pd_client.mode} node {pd_client.client_ip_port}") + return None self.url_to_pd_nodes.pop(pd_client.client_ip_port, None) self.prefill_nodes = [e for e in self.prefill_nodes if e.client_ip_port != pd_client.client_ip_port] @@ -768,7 +813,7 @@ def remove_pd(self, pd_info_json): self.selector.update_nodes(self.prefill_nodes, self.decode_nodes) logger.info(f"mode: {pd_client.mode} url: {pd_client.client_ip_port} removed") - return + return registered_node def update_node_load_info(self, load_info: Optional[dict]): """更新节点负载信息 diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py index cd9b015b1..7373671ab 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py @@ -21,7 +21,7 @@ from ..kv_transporter import create_kv_transporter from lightllm.utils.error_utils import log_exception from lightllm.utils.envs_utils import get_unique_server_name -from lightllm.utils.process_check import start_parent_check_thread +from lightllm.utils.process_check import install_fatal_thread_excepthook, start_parent_check_thread logger = init_logger(__name__) @@ -47,6 +47,7 @@ def _init_env( task_out_queue: mp.Queue, up_status_in_queue: Optional[mp.SimpleQueue], ): + install_fatal_thread_excepthook() start_parent_check_thread() import lightllm.utils.rpyc_fix_utils as _ @@ -369,7 +370,6 @@ def request_page_loop(self): self.waiting_dict.pop(key, None) logger.error(f"send write ready task to prefill node failed: {trans_task.to_str()}") logger.exception(str(e)) - self.transporter.remove_remote_agent(peer_name=trans_task.prefill_agent_name) trans_task.error_info = f"send write ready task to prefill node failed: {str(e)}" self.failed_queue.put(trans_task) continue diff --git a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py index bd5e11f05..fa0512649 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py @@ -1,9 +1,9 @@ import pickle import copy import os +import threading import time -from dataclasses import dataclass -from typing import Dict +from typing import Dict, Optional, Tuple from torch import Tensor from lightllm.server.pd_io_struct import PDChunckedTransTask, PDAgentMetadata from lightllm.utils.log_utils import init_logger @@ -15,6 +15,7 @@ from nixl._api import nixl_agent as NixlWrapper from nixl._api import nixlBind from nixl._api import nixl_agent_config + from nixl._api import nixl_thread_sync_t logger.info("Nixl is available") except ImportError: @@ -26,20 +27,28 @@ class NixlKVTransporter: def __init__(self, node_id: int, tp_idx: int, kv_move_buffer: Tensor): self.node_id = node_id self.tp_idx = tp_idx + # A transporter is shared by the notification, transfer, status, and + # failure-handling threads. NIXL's internal synchronization protects + # individual calls; this lock also makes LightLLM's compound peer + # lifecycle operations atomic. + self._nixl_lock = threading.RLock() self.capture_telemetry = os.getenv("LIGHTLLM_NIXL_CAPTURE_TELEMETRY", "0").lower() in ( "1", "true", "yes", "on", ) - conf = None + conf = nixl_agent_config(sync_mode=nixl_thread_sync_t.NIXL_THREAD_SYNC_STRICT) if self.capture_telemetry: - conf = nixl_agent_config() conf.capture_telemetry = True logger.info("NIXL telemetry enabled") self.nixl_agent = NixlWrapper(self.agent_name, conf) self._register_kv_move_buffer(kv_move_buffer=kv_move_buffer) self.remote_agents: Dict[str, PDAgentMetadata] = {} + self._peer_generations: Dict[str, int] = {} + self._next_peer_generation = 0 + self._broken_remote_agents: set[str] = set() + self._active_xfers: Dict[int, Tuple[object, str, int]] = {} return @property @@ -48,14 +57,17 @@ def agent_name(self) -> str: @property def agent_metadata(self): - return self.nixl_agent.get_agent_metadata() + with self._nixl_lock: + return self.nixl_agent.get_agent_metadata() @property def local_page_mem_desc(self): - return self.nixl_agent.get_serialized_descs(self.page_reg_desc) + with self._nixl_lock: + return self.nixl_agent.get_serialized_descs(self.page_reg_desc) def get_new_notifs(self) -> Dict[str, list[bytes]]: - return self.nixl_agent.get_new_notifs() + with self._nixl_lock: + return self.nixl_agent.get_new_notifs() def _register_kv_move_buffer(self, kv_move_buffer: Tensor): self.num_pages, self.page_size, self.num_layers, self.kv_head_num, self.head_dims = kv_move_buffer.shape @@ -73,11 +85,54 @@ def _create_paged_xfer_handles(self, reg_desc: "nixlBind.nixlRegDList", page_num return self.nixl_agent.prep_xfer_dlist(agent_name, descs, "VRAM") def connect_add_remote_agent(self, remote_agent: PDAgentMetadata): + with self._nixl_lock: + self._ensure_remote_agent_locked(remote_agent) + return + + @staticmethod + def _same_remote_agent(left: PDAgentMetadata, right: PDAgentMetadata) -> bool: + return ( + left.agent_name == right.agent_name + and left.agent_metadata == right.agent_metadata + and left.num_pages == right.num_pages + and left.page_reg_desc == right.page_reg_desc + ) + + def _ensure_remote_agent_locked(self, remote_agent: PDAgentMetadata) -> Tuple[PDAgentMetadata, int]: + peer_name = remote_agent.agent_name + current_agent = self.remote_agents.get(peer_name) + current_generation = self._peer_generations.get(peer_name) + + if current_agent is not None and not self._same_remote_agent(current_agent, remote_agent): + self._mark_remote_agent_broken_locked(peer_name, current_generation) + current_agent = self.remote_agents.get(peer_name) + + if peer_name in self._broken_remote_agents: + if current_agent is not None: + self._remove_remote_agent_locked(peer_name, expected_generation=current_generation) + else: + try: + self.nixl_agent.remove_remote_agent(peer_name) + self._broken_remote_agents.discard(peer_name) + except BaseException as e: + raise RuntimeError(f"NIXL remote agent {peer_name} is broken and could not be removed") from e + current_agent = self.remote_agents.get(peer_name) + + if current_agent is None: + self._connect_add_remote_agent_locked(remote_agent) + current_agent = self.remote_agents[peer_name] + current_generation = self._peer_generations[peer_name] + + if peer_name in self._broken_remote_agents: + raise RuntimeError(f"NIXL remote agent {peer_name} is unavailable while active transfers drain") + + return current_agent, current_generation + + def _connect_add_remote_agent_locked(self, remote_agent: PDAgentMetadata): if remote_agent.agent_name in self.remote_agents: return start_time = time.time() - peer_name = self.nixl_agent.add_remote_agent(remote_agent.agent_metadata) if isinstance(peer_name, bytes): peer_name = peer_name.decode() @@ -86,96 +141,135 @@ def connect_add_remote_agent(self, remote_agent: PDAgentMetadata): peer_name == remote_agent.agent_name ), f"Peer name {peer_name} does not match remote name {remote_agent.agent_name}" - page_mem_desc = self.nixl_agent.deserialize_descs(remote_agent.page_reg_desc) - kv_page_xfer_handles = self._create_paged_xfer_handles( - page_mem_desc, remote_agent.num_pages, agent_name=peer_name - ) - remote_agent.page_xfer_handles = kv_page_xfer_handles + self._next_peer_generation += 1 + generation = self._next_peer_generation + self.remote_agents[peer_name] = remote_agent + self._peer_generations[peer_name] = generation + try: + page_mem_desc = self.nixl_agent.deserialize_descs(remote_agent.page_reg_desc) + remote_agent.page_xfer_handles = self._create_paged_xfer_handles( + page_mem_desc, remote_agent.num_pages, agent_name=peer_name + ) + except BaseException: + self._broken_remote_agents.add(peer_name) + self._remove_remote_agent_locked(peer_name, expected_generation=generation) + raise logger.info( - f"Added remote agent {peer_name} with mem desc {page_mem_desc} cost time: {time.time() - start_time} s" + f"Added remote agent {peer_name} generation {generation} " + f"with mem desc {page_mem_desc} cost time: {time.time() - start_time} s" ) - - self.remote_agents[remote_agent.agent_name] = remote_agent + self._broken_remote_agents.discard(peer_name) return def remove_remote_agent(self, peer_name: str): - if peer_name in self.remote_agents: + with self._nixl_lock: + generation = self._peer_generations.get(peer_name) + if generation is None: + logger.warning(f"try to remove remote agent, but peer name {peer_name} agent did not exist") + return + self._mark_remote_agent_broken_locked(peer_name, generation) + return + + def _has_active_xfers_locked(self, peer_name: str, generation: int) -> bool: + return any( + active_peer_name == peer_name and active_generation == generation + for _, active_peer_name, active_generation in self._active_xfers.values() + ) + + def _mark_remote_agent_broken_locked(self, peer_name: str, generation: Optional[int]): + if generation is None or self._peer_generations.get(peer_name) != generation: + return + self._broken_remote_agents.add(peer_name) + self._remove_remote_agent_locked(peer_name, expected_generation=generation) + + def _remove_remote_agent_locked(self, peer_name: str, expected_generation: Optional[int] = None) -> bool: + generation = self._peer_generations.get(peer_name) + if generation is None: + return False + if expected_generation is not None and generation != expected_generation: + return False + if self._has_active_xfers_locked(peer_name, generation): + self._broken_remote_agents.add(peer_name) + logger.warning( + f"defer removing remote agent {peer_name} generation {generation} until active transfers drain" + ) + return False + + remote_agent = self.remote_agents[peer_name] + try: + self.nixl_agent.remove_remote_agent(remote_agent.agent_name) + except BaseException as e: + self._broken_remote_agents.add(peer_name) + logger.error(f"remove remote agent {peer_name} generation {generation} failed") + logger.exception(str(e)) + return False + + self.remote_agents.pop(peer_name, None) + self._peer_generations.pop(peer_name, None) + self._broken_remote_agents.discard(peer_name) + if remote_agent.page_xfer_handles is not None: try: - remote_agent: PDAgentMetadata = self.remote_agents.pop(peer_name, None) - assert remote_agent.agent_name == peer_name - self.nixl_agent.remove_remote_agent(remote_agent.agent_name) - if remote_agent.page_xfer_handles is not None: - self.nixl_agent.release_dlist_handle(remote_agent.page_xfer_handles) + self.nixl_agent.release_dlist_handle(remote_agent.page_xfer_handles) except BaseException as e: - logger.error(f"remove remote agent {peer_name} failed") + logger.error(f"release remote agent {peer_name} descriptor handle failed") logger.exception(str(e)) - else: - logger.warning(f"try to remove remote agent, but peer name {peer_name} agent did not exist") + finally: + remote_agent.page_xfer_handles = None + return True + + def _send_notif_locked(self, peer_name: str, generation: int, notif_msg: bytes): + try: + self.nixl_agent.send_notif(remote_agent_name=peer_name, notif_msg=notif_msg) + except BaseException: + self._mark_remote_agent_broken_locked(peer_name, generation) + raise def send_write_done_task_to_decode_node(self, trans_task: PDChunckedTransTask): decode_agent_name = trans_task.decode_agent_name - if decode_agent_name not in self.remote_agents: - logger.warning(f"decode_agent_name {decode_agent_name} not exist") - _remote_agent = trans_task.create_decode_agent_obj() - self.connect_add_remote_agent(_remote_agent) - - new_trans_task: PDChunckedTransTask = copy.copy(trans_task) - new_trans_task.write_stage = "done" - new_trans_task.mem_indexes = None - new_trans_task.xfer_handle = None - new_trans_task.decode_agent_metadata = None - new_trans_task.decode_page_reg_desc = None - new_trans_task.prefill_agent_name = self.agent_name - new_trans_task.prefill_agent_metadata = self.agent_metadata - new_trans_task.prefill_num_pages = self.num_pages - new_trans_task.prefill_page_reg_desc = self.local_page_mem_desc - self.nixl_agent.send_notif( - remote_agent_name=decode_agent_name, - notif_msg=pickle.dumps(new_trans_task), - ) + with self._nixl_lock: + _, generation = self._ensure_remote_agent_locked(trans_task.create_decode_agent_obj()) + new_trans_task: PDChunckedTransTask = copy.copy(trans_task) + new_trans_task.write_stage = "done" + new_trans_task.mem_indexes = None + new_trans_task.xfer_handle = None + new_trans_task.decode_agent_metadata = None + new_trans_task.decode_page_reg_desc = None + new_trans_task.prefill_agent_name = self.agent_name + new_trans_task.prefill_agent_metadata = self.agent_metadata + new_trans_task.prefill_num_pages = self.num_pages + new_trans_task.prefill_page_reg_desc = self.local_page_mem_desc + self._send_notif_locked(decode_agent_name, generation, pickle.dumps(new_trans_task)) return def send_write_request_task_to_decode_node(self, trans_task: PDChunckedTransTask): decode_agent_name = trans_task.decode_agent_name - if decode_agent_name not in self.remote_agents: - logger.warning(f"decode_agent_name {decode_agent_name} not exist") - _remote_agent = trans_task.create_decode_agent_obj() - self.connect_add_remote_agent(_remote_agent) - - new_trans_task: PDChunckedTransTask = copy.copy(trans_task) - new_trans_task.write_stage = "request" - new_trans_task.mem_indexes = None - new_trans_task.xfer_handle = None - new_trans_task.prefill_agent_name = self.agent_name - new_trans_task.prefill_agent_metadata = self.agent_metadata - new_trans_task.prefill_num_pages = self.num_pages - new_trans_task.prefill_page_reg_desc = self.local_page_mem_desc - self.nixl_agent.send_notif( - remote_agent_name=decode_agent_name, - notif_msg=pickle.dumps(new_trans_task), - ) + with self._nixl_lock: + _, generation = self._ensure_remote_agent_locked(trans_task.create_decode_agent_obj()) + new_trans_task: PDChunckedTransTask = copy.copy(trans_task) + new_trans_task.write_stage = "request" + new_trans_task.mem_indexes = None + new_trans_task.xfer_handle = None + new_trans_task.prefill_agent_name = self.agent_name + new_trans_task.prefill_agent_metadata = self.agent_metadata + new_trans_task.prefill_num_pages = self.num_pages + new_trans_task.prefill_page_reg_desc = self.local_page_mem_desc + self._send_notif_locked(decode_agent_name, generation, pickle.dumps(new_trans_task)) return def send_write_ready_task_to_prefill_node(self, trans_task: PDChunckedTransTask): prefill_agent_name = trans_task.prefill_agent_name - if prefill_agent_name not in self.remote_agents: - logger.warning(f"prefill_agent_name {prefill_agent_name} not exist") - _remote_agent = trans_task.create_prefill_agent_obj() - self.connect_add_remote_agent(_remote_agent) - - new_trans_task: PDChunckedTransTask = copy.copy(trans_task) - new_trans_task.write_stage = "ready" - new_trans_task.mem_indexes = None - new_trans_task.xfer_handle = None - new_trans_task.decode_agent_name = self.agent_name - new_trans_task.decode_agent_metadata = self.agent_metadata - new_trans_task.decode_num_pages = self.num_pages - new_trans_task.decode_page_reg_desc = self.local_page_mem_desc - self.nixl_agent.send_notif( - remote_agent_name=prefill_agent_name, - notif_msg=pickle.dumps(new_trans_task), - ) + with self._nixl_lock: + _, generation = self._ensure_remote_agent_locked(trans_task.create_prefill_agent_obj()) + new_trans_task: PDChunckedTransTask = copy.copy(trans_task) + new_trans_task.write_stage = "ready" + new_trans_task.mem_indexes = None + new_trans_task.xfer_handle = None + new_trans_task.decode_agent_name = self.agent_name + new_trans_task.decode_agent_metadata = self.agent_metadata + new_trans_task.decode_num_pages = self.num_pages + new_trans_task.decode_page_reg_desc = self.local_page_mem_desc + self._send_notif_locked(prefill_agent_name, generation, pickle.dumps(new_trans_task)) return def send_error_info_to_prefill_node(self, trans_task: PDChunckedTransTask): @@ -186,53 +280,41 @@ def send_error_info_to_prefill_node(self, trans_task: PDChunckedTransTask): try: prefill_agent_name = trans_task.prefill_agent_name - if prefill_agent_name not in self.remote_agents: - logger.warning(f"prefill_agent_name {prefill_agent_name} not exist") - _remote_agent = trans_task.create_prefill_agent_obj() - self.connect_add_remote_agent(_remote_agent) - assert trans_task.error_info is not None - new_trans_task: PDChunckedTransTask = copy.copy(trans_task) - new_trans_task.write_stage = "error" - new_trans_task.mem_indexes = None - new_trans_task.xfer_handle = None - new_trans_task.decode_agent_name = self.agent_name - new_trans_task.decode_agent_metadata = self.agent_metadata - new_trans_task.decode_num_pages = self.num_pages - new_trans_task.decode_page_reg_desc = self.local_page_mem_desc - self.nixl_agent.send_notif( - remote_agent_name=prefill_agent_name, - notif_msg=pickle.dumps(new_trans_task), - ) + with self._nixl_lock: + _, generation = self._ensure_remote_agent_locked(trans_task.create_prefill_agent_obj()) + assert trans_task.error_info is not None + new_trans_task: PDChunckedTransTask = copy.copy(trans_task) + new_trans_task.write_stage = "error" + new_trans_task.mem_indexes = None + new_trans_task.xfer_handle = None + new_trans_task.decode_agent_name = self.agent_name + new_trans_task.decode_agent_metadata = self.agent_metadata + new_trans_task.decode_num_pages = self.num_pages + new_trans_task.decode_page_reg_desc = self.local_page_mem_desc + self._send_notif_locked(prefill_agent_name, generation, pickle.dumps(new_trans_task)) except BaseException as e: logger.error(f"send error info to prefill node failed: {trans_task.to_str()}") logger.exception(str(e)) - self.remove_remote_agent(peer_name=prefill_agent_name) return def send_error_info_to_decode_node(self, trans_task: PDChunckedTransTask): try: decode_agent_name = trans_task.decode_agent_name - if decode_agent_name not in self.remote_agents: - logger.warning(f"decode_agent_name {decode_agent_name} not exist") - _remote_agent = trans_task.create_decode_agent_obj() - self.connect_add_remote_agent(_remote_agent) - assert trans_task.error_info is not None - new_trans_task: PDChunckedTransTask = copy.copy(trans_task) - new_trans_task.write_stage = "error" - new_trans_task.mem_indexes = None - new_trans_task.xfer_handle = None - new_trans_task.prefill_agent_name = self.agent_name - new_trans_task.prefill_agent_metadata = self.agent_metadata - new_trans_task.prefill_num_pages = self.num_pages - new_trans_task.prefill_page_reg_desc = self.local_page_mem_desc - self.nixl_agent.send_notif( - remote_agent_name=decode_agent_name, - notif_msg=pickle.dumps(new_trans_task), - ) + with self._nixl_lock: + _, generation = self._ensure_remote_agent_locked(trans_task.create_decode_agent_obj()) + assert trans_task.error_info is not None + new_trans_task: PDChunckedTransTask = copy.copy(trans_task) + new_trans_task.write_stage = "error" + new_trans_task.mem_indexes = None + new_trans_task.xfer_handle = None + new_trans_task.prefill_agent_name = self.agent_name + new_trans_task.prefill_agent_metadata = self.agent_metadata + new_trans_task.prefill_num_pages = self.num_pages + new_trans_task.prefill_page_reg_desc = self.local_page_mem_desc + self._send_notif_locked(decode_agent_name, generation, pickle.dumps(new_trans_task)) except BaseException as e: logger.error(f"send error info to decode node failed: {trans_task.to_str()}") logger.exception(str(e)) - self.remove_remote_agent(peer_name=decode_agent_name) return def write_blocks_paged( @@ -243,46 +325,85 @@ def write_blocks_paged( prefill node call this function to write kv blocks into decode node pages """ decode_agent_name = trans_task.decode_agent_name - if decode_agent_name not in self.remote_agents: - logger.warning(f"decode_agent_name {decode_agent_name} not exist") - _remote_agent = trans_task.create_decode_agent_obj() - self.connect_add_remote_agent(_remote_agent) - - assert trans_task.src_page_index is not None and trans_task.dst_page_index is not None - remote_agent: PDAgentMetadata = self.remote_agents[decode_agent_name] - src_handle = self.page_local_xfer_handles - dst_handle = remote_agent.page_xfer_handles - handle = self.nixl_agent.make_prepped_xfer( - "WRITE", - src_handle, - [trans_task.src_page_index], - dst_handle, - [trans_task.dst_page_index], - b"", - ) - if not handle: - raise RuntimeError(f"make_prepped_xfer failed for task: {trans_task.to_str()}") - - self.nixl_agent.transfer(handle) - - return handle + with self._nixl_lock: + remote_agent, generation = self._ensure_remote_agent_locked(trans_task.create_decode_agent_obj()) + assert trans_task.src_page_index is not None and trans_task.dst_page_index is not None + handle = None + try: + handle = self.nixl_agent.make_prepped_xfer( + "WRITE", + self.page_local_xfer_handles, + [trans_task.src_page_index], + remote_agent.page_xfer_handles, + [trans_task.dst_page_index], + b"", + ) + if not handle: + raise RuntimeError(f"make_prepped_xfer failed for task: {trans_task.to_str()}") + self.nixl_agent.transfer(handle) + self._active_xfers[id(handle)] = (handle, decode_agent_name, generation) + return handle + except BaseException: + if handle: + try: + self.nixl_agent.release_xfer_handle(handle=handle) + except BaseException as release_error: + logger.error(f"release failed transfer handle for remote agent {decode_agent_name}") + logger.exception(str(release_error)) + self._mark_remote_agent_broken_locked(decode_agent_name, generation) + raise def check_task_status(self, trans_task: PDChunckedTransTask) -> str: assert trans_task.xfer_handle is not None handle = trans_task.xfer_handle - xfer_state = self.nixl_agent.check_xfer_state(handle) - if xfer_state == "ERR": - logger.warning(f"Transfer failed with trans task {trans_task.to_str()} for handle {handle}") - return xfer_state + with self._nixl_lock: + active_xfer = self._active_xfers.get(id(handle)) + try: + xfer_state = self.nixl_agent.check_xfer_state(handle) + except BaseException: + if active_xfer is not None: + _, peer_name, generation = active_xfer + self._mark_remote_agent_broken_locked(peer_name, generation) + raise + if xfer_state == "ERR": + logger.warning(f"Transfer failed with trans task {trans_task.to_str()} for handle {handle}") + if active_xfer is not None: + _, peer_name, generation = active_xfer + self._mark_remote_agent_broken_locked(peer_name, generation) + return xfer_state def release_xfer_handle(self, handle): - self.nixl_agent.release_xfer_handle(handle=handle) + with self._nixl_lock: + active_xfer = self._active_xfers.get(id(handle)) + self.nixl_agent.release_xfer_handle(handle=handle) + self._active_xfers.pop(id(handle), None) + if active_xfer is not None: + _, peer_name, generation = active_xfer + if peer_name in self._broken_remote_agents: + self._remove_remote_agent_locked(peer_name, expected_generation=generation) return + def get_xfer_telemetry(self, handle): + with self._nixl_lock: + return self.nixl_agent.get_xfer_telemetry(handle) + + def query_xfer_backend(self, handle): + with self._nixl_lock: + return self.nixl_agent.query_xfer_backend(handle) + def shutdown(self): - self.nixl_agent.deregister_memory(self.page_reg_desc) - self.nixl_agent.release_dlist_handle(self.page_local_xfer_handles) - agent_names = list(self.remote_agents.keys()) - for agent_name in agent_names: - self.remove_remote_agent(agent_name) + with self._nixl_lock: + for handle, _, _ in list(self._active_xfers.values()): + try: + self.nixl_agent.release_xfer_handle(handle=handle) + except BaseException as e: + logger.error("release active transfer handle during NIXL shutdown failed") + logger.exception(str(e)) + self._active_xfers.clear() + for agent_name in list(self.remote_agents.keys()): + generation = self._peer_generations.get(agent_name) + self._broken_remote_agents.add(agent_name) + self._remove_remote_agent_locked(agent_name, expected_generation=generation) + self.nixl_agent.deregister_memory(self.page_reg_desc) + self.nixl_agent.release_dlist_handle(self.page_local_xfer_handles) return diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py index 40bd2e42a..f483aa032 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py @@ -15,7 +15,7 @@ from ..kv_transporter import create_kv_transporter from lightllm.utils.error_utils import log_exception from lightllm.utils.envs_utils import get_unique_server_name -from lightllm.utils.process_check import start_parent_check_thread +from lightllm.utils.process_check import install_fatal_thread_excepthook, start_parent_check_thread logger = init_logger(__name__) @@ -41,6 +41,7 @@ def _init_env( task_in_queue: mp.Queue, task_out_queue: mp.Queue, ): + install_fatal_thread_excepthook() start_parent_check_thread() import lightllm.utils.rpyc_fix_utils as _ @@ -225,7 +226,6 @@ def ready_transfer_loop(self): logger.error(f"send WRITE request to decode failed: {trans_task.to_str()}") logger.exception(str(e)) trans_task.error_info = f"send WRITE request to decode failed: {str(e)}" - self.transporter.remove_remote_agent(peer_name=trans_task.decode_agent_name) self.failed_queue.put(trans_task) continue return @@ -316,7 +316,6 @@ def write_peer_kv_loop(self): except BaseException as e: logger.error(f"write_blocks_paged failed: {trans_task.to_str()}") logger.exception(str(e)) - self.transporter.remove_remote_agent(peer_name=trans_task.decode_agent_name) trans_task.error_info = f"write_blocks_paged failed: {str(e)}" self.failed_queue.put(trans_task) continue @@ -331,48 +330,76 @@ def update_task_status_loop( time.sleep(0.001) continue - with self.waiting_dict_lock: - tasks = list(self.waiting_dict.values()) - for trans_task in tasks: - if trans_task.xfer_handle is None: - continue + self._update_task_status_once() + time.sleep(0.001) - # 传输任务状态检查 + def _update_task_status_once(self): + with self.waiting_dict_lock: + tasks = list(self.waiting_dict.values()) + for trans_task in tasks: + if trans_task.xfer_handle is None: + continue + + # A single native NIXL error must fail and release this task without + # silently terminating the status thread for every later transfer. + try: ret = self.transporter.check_task_status(trans_task=trans_task) - if ret == "DONE": - trans_task = self.waiting_dict.pop(trans_task.get_key(), None) - if self.transporter.capture_telemetry: - telem = self.transporter.nixl_agent.get_xfer_telemetry(trans_task.xfer_handle) + except BaseException as e: + failed_task = self.waiting_dict.pop(trans_task.get_key(), None) + if failed_task is not None: + logger.exception(f"check transfer status failed: {failed_task.to_str()}") + failed_task.error_info = f"check transfer status failed: {str(e)}" + self.failed_queue.put(failed_task) + continue + + if ret == "DONE": + completed_task = self.waiting_dict.pop(trans_task.get_key(), None) + if completed_task is None: + continue + if self.transporter.capture_telemetry: + try: + telem = self.transporter.get_xfer_telemetry(completed_task.xfer_handle) total_us = telem.xferDuration post_us = telem.postDuration backend_us = telem.xferDuration - telem.postDuration - nixl_backend = self.transporter.nixl_agent.query_xfer_backend(trans_task.xfer_handle) + nixl_backend = self.transporter.query_xfer_backend(completed_task.xfer_handle) logger.info( - f"write trans task request_id={trans_task.request_id} " - f"kv=[{trans_task.start_kv_index},{trans_task.end_kv_index}) " - f"src_page={trans_task.src_page_index} dst_page={trans_task.dst_page_index} " + f"write trans task request_id={completed_task.request_id} " + f"kv=[{completed_task.start_kv_index},{completed_task.end_kv_index}) " + f"src_page={completed_task.src_page_index} " + f"dst_page={completed_task.dst_page_index} " f"xfer time: {total_us:.3f} us, " f"post time: {post_us:.3f} us, backend time: {backend_us:.3f} us, " f"nixl_backend: {nixl_backend}, total_bytes: {telem.totalBytes}" ) - self.transporter.send_write_done_task_to_decode_node(trans_task) - logger.info( - f"send WRITE done nixl notify " - f"request_id={trans_task.request_id} " - f"kv=[{trans_task.start_kv_index},{trans_task.end_kv_index}) " - f"src_page={trans_task.src_page_index} dst_page={trans_task.dst_page_index}" - ) - self.success_queue.put(trans_task) - elif ret == "ERR": - trans_task = self.waiting_dict.pop(trans_task.get_key(), None) - trans_task.error_info = "xfer error" - self.failed_queue.put(trans_task) - elif trans_task.time_out(): - trans_task = self.waiting_dict.pop(trans_task.get_key(), None) - trans_task.error_info = "time out in update_task_status_loop" - self.failed_queue.put(trans_task) + except BaseException: + logger.exception(f"get transfer telemetry failed: {completed_task.to_str()}") + + try: + self.transporter.send_write_done_task_to_decode_node(completed_task) + except BaseException as e: + logger.exception(f"send WRITE done nixl notify failed: {completed_task.to_str()}") + completed_task.error_info = f"send WRITE done nixl notify failed: {str(e)}" + self.failed_queue.put(completed_task) + continue - time.sleep(0.001) + logger.info( + f"send WRITE done nixl notify " + f"request_id={completed_task.request_id} " + f"kv=[{completed_task.start_kv_index},{completed_task.end_kv_index}) " + f"src_page={completed_task.src_page_index} dst_page={completed_task.dst_page_index}" + ) + self.success_queue.put(completed_task) + elif ret == "ERR": + failed_task = self.waiting_dict.pop(trans_task.get_key(), None) + if failed_task is not None: + failed_task.error_info = "xfer error" + self.failed_queue.put(failed_task) + elif trans_task.time_out(): + failed_task = self.waiting_dict.pop(trans_task.get_key(), None) + if failed_task is not None: + failed_task.error_info = "time out in update_task_status_loop" + self.failed_queue.put(failed_task) @log_exception def success_loop(self): diff --git a/lightllm/utils/process_check.py b/lightllm/utils/process_check.py index 00cc258bf..db508c8a6 100644 --- a/lightllm/utils/process_check.py +++ b/lightllm/utils/process_check.py @@ -8,6 +8,21 @@ logger = init_logger(__name__) +def install_fatal_thread_excepthook(): + """Terminate the process when an unexpected daemon-thread exception escapes.""" + + def exit_process(args): + try: + logger.error( + f"fatal background thread failure in {getattr(args.thread, 'name', 'unknown')}", + exc_info=(args.exc_type, args.exc_value, args.exc_traceback), + ) + finally: + os._exit(1) + + threading.excepthook = exit_process + + def is_process_active(pid): try: process = psutil.Process(pid) diff --git a/unit_tests/server/httpserver/test_pd_connection_tasks.py b/unit_tests/server/httpserver/test_pd_connection_tasks.py new file mode 100644 index 000000000..eaf2a8e08 --- /dev/null +++ b/unit_tests/server/httpserver/test_pd_connection_tasks.py @@ -0,0 +1,61 @@ +import asyncio +from types import SimpleNamespace + +import pytest + +from lightllm.server.httpserver.async_queue import AsyncQueue +from lightllm.server.httpserver.pd_loop import ( + _pd_process_generate, + _recv_or_raise_on_background_failure, +) + + +def test_background_failure_interrupts_blocked_pd_receive(): + async def run(): + recv_started = asyncio.Event() + recv_cancelled = asyncio.Event() + + class BlockingWebsocket: + async def recv(self): + recv_started.set() + try: + await asyncio.Future() + except asyncio.CancelledError: + recv_cancelled.set() + raise + + websocket = BlockingWebsocket() + failure = asyncio.get_running_loop().create_future() + receive_task = asyncio.create_task(_recv_or_raise_on_background_failure(websocket, (failure,))) + await recv_started.wait() + + failure.set_exception(RuntimeError("generation failed")) + with pytest.raises(RuntimeError, match="generation failed"): + await receive_task + assert recv_cancelled.is_set() + + asyncio.run(run()) + + +def test_pd_generation_failure_is_not_swallowed(): + async def run(): + class FailingManager: + args = SimpleNamespace(run_mode="prefill") + + async def generate(self, **_kwargs): + yield 1, "token", {}, None + raise RuntimeError("generation failed") + + sampling_params = SimpleNamespace(group_request_id=123) + with pytest.raises(RuntimeError, match="generation failed"): + await _pd_process_generate( + manager=FailingManager(), + prompt="prompt", + sampling_params=sampling_params, + multimodal_params={}, + forwarding_queue=AsyncQueue(), + pd_upload_websocket=object(), + pd_event=asyncio.Event(), + ) + + asyncio.run(run()) diff --git a/unit_tests/server/httpserver/test_pd_pending_abort.py b/unit_tests/server/httpserver/test_pd_pending_abort.py new file mode 100644 index 000000000..0ece4d4d8 --- /dev/null +++ b/unit_tests/server/httpserver/test_pd_pending_abort.py @@ -0,0 +1,70 @@ +import asyncio +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from lightllm.server.httpserver.manager import HttpServerManager +from lightllm.server.httpserver.pd_loop import _abort_pd_request + + +def test_pd_abort_cancels_request_before_http_manager_registration(): + async def run(): + manager = MagicMock() + manager.abort = AsyncMock(return_value=False) + request_started = asyncio.Event() + + async def wait_for_request_slot(): + request_started.set() + await asyncio.Future() + + generation_task = asyncio.create_task(wait_for_request_slot()) + generation_tasks = {123: generation_task} + await request_started.wait() + + assert await _abort_pd_request(manager, 123, generation_tasks) + assert generation_task.cancelling() == 1 + assert await _abort_pd_request(manager, 123, generation_tasks) + assert generation_task.cancelling() == 1 + with pytest.raises(asyncio.CancelledError): + await generation_task + + assert manager.abort.await_count == 2 + manager.abort.assert_awaited_with(123) + assert generation_task.cancelled() + + asyncio.run(run()) + + +def test_cancelled_slot_allocation_releases_partially_allocated_indexes(): + async def run(): + manager = HttpServerManager.__new__(HttpServerManager) + manager.shm_req_manager = MagicMock() + manager.shm_req_manager.async_release_req_index = AsyncMock() + manager.shm_req_manager.async_put_back_req_obj = AsyncMock() + manager.tokenizer = MagicMock() + manager.args = MagicMock(chunked_prefill_size=16) + + waiting_for_second_slot = asyncio.Event() + allocation_count = 0 + + async def alloc_req_index(): + nonlocal allocation_count + allocation_count += 1 + if allocation_count == 1: + return 7 + waiting_for_second_slot.set() + return None + + manager.shm_req_manager.async_alloc_req_index = alloc_req_index + sampling_params = MagicMock(n=2) + allocation_task = asyncio.create_task(manager._alloc_req_objs(123, [1, 2], sampling_params)) + await waiting_for_second_slot.wait() + allocation_task.cancel() + + with pytest.raises(asyncio.CancelledError): + await allocation_task + + manager.shm_req_manager.async_release_req_index.assert_awaited_once_with(7) + manager.shm_req_manager.async_put_back_req_obj.assert_not_awaited() + + asyncio.run(run()) diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index edb758a94..afef60108 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -6,7 +6,20 @@ from easydict import EasyDict from lightllm.server.core.objs.start_args_type import StartArgs -from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster, PDManager +from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster, PDManager, ReqStatus + + +def _prefill_pd_info(args, node_id=1, client_ip_port="10.0.0.1:8000"): + return { + "node_id": node_id, + "client_ip_port": client_ip_port, + "mode": "prefill", + "start_args": { + "max_req_total_len": args.max_req_total_len, + "max_image_pixels": args.max_image_pixels, + "disable_image_resize": args.disable_image_resize, + }, + } def test_auto_set_response_parsers_from_qwen35_model_config(tmp_path): @@ -243,6 +256,61 @@ def pd_info(node_id, client_ip_port): assert [node.dispatched_req_num for node in manager.prefill_nodes] == [1, 0] +def test_pd_reconnection_fails_old_requests_and_ignores_stale_disconnect(): + async def run(): + args = StartArgs() + pd_manager = PDManager(args) + http_manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + http_manager.pd_manager = pd_manager + http_manager.req_id_to_out_inf = {} + pd_info = _prefill_pd_info(args) + old_websocket = object() + new_websocket = object() + + pd_manager.register_pd(pd_info, old_websocket) + old_node = pd_manager.prefill_nodes[0] + req_status = ReqStatus(8, old_node, object()) + http_manager.req_id_to_out_inf[8] = req_status + + await http_manager.register_pd(pd_info, new_websocket) + + with pytest.raises(RuntimeError, match="connection was replaced"): + req_status.raise_if_failed() + assert req_status.prefill_prompt_ids_event.is_set() + assert req_status.up_status_event.is_set() + assert req_status.out_tokens.event.is_set() + + await http_manager.remove_pd(pd_info, old_websocket) + assert pd_manager.url_to_pd_nodes[pd_info["client_ip_port"]].websocket is new_websocket + + asyncio.run(run()) + + +def test_pd_disconnect_wakes_requests_owned_by_the_disconnected_node(): + async def run(): + args = StartArgs() + pd_manager = PDManager(args) + http_manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + http_manager.pd_manager = pd_manager + http_manager.req_id_to_out_inf = {} + pd_info = _prefill_pd_info(args) + websocket = object() + + pd_manager.register_pd(pd_info, websocket) + node = pd_manager.prefill_nodes[0] + req_status = ReqStatus(9, node, object()) + http_manager.req_id_to_out_inf[9] = req_status + token_waiter = asyncio.create_task(req_status.out_tokens.wait_to_get_all_data(timeout=10)) + + await http_manager.remove_pd(pd_info, websocket) + assert await asyncio.wait_for(token_waiter, timeout=1) == [] + with pytest.raises(RuntimeError, match="PD node disconnected"): + req_status.raise_if_failed() + assert pd_info["client_ip_port"] not in pd_manager.url_to_pd_nodes + + asyncio.run(run()) + + def test_pd_master_inference_health_matches_normal_node_semantics(monkeypatch): monkeypatch.setattr("lightllm.server.httpserver_for_pd_master.manager.time.time", lambda: 1000) manager = SimpleNamespace( diff --git a/unit_tests/test_nixl_kv_transporter_thread_safety.py b/unit_tests/test_nixl_kv_transporter_thread_safety.py new file mode 100644 index 000000000..97276e60c --- /dev/null +++ b/unit_tests/test_nixl_kv_transporter_thread_safety.py @@ -0,0 +1,321 @@ +import threading +import time +import queue +from types import SimpleNamespace + +import pytest + +from lightllm.server.pd_io_struct import PDChunckedTransTask, PDAgentMetadata +from lightllm.server.router.model_infer.mode_backend.pd import nixl_kv_transporter +from lightllm.server.router.model_infer.mode_backend.pd.prefill_node_impl.prefill_trans_process import ( + _PrefillTransModule, +) + + +class _FakeBuffer: + shape = (1, 1, 1, 1, 1) + + @staticmethod + def element_size(): + return 1 + + +class _FakeConfig: + def __init__(self, sync_mode=None): + self.sync_mode = sync_mode + self.capture_telemetry = False + + +class _FakeSyncMode: + NIXL_THREAD_SYNC_STRICT = object() + + +class _FakeNixlAgent: + instances = [] + + def __init__(self, name, config): + self.name = name + self.config = config + self.add_calls = [] + self.remove_calls = [] + self.released_dlists = [] + self.released_xfers = [] + self.send_calls = [] + self.fail_next_send = False + self._native_call_lock = threading.Lock() + self._native_calls = 0 + self.max_native_calls = 0 + self.__class__.instances.append(self) + + def _enter_native_call(self): + with self._native_call_lock: + self._native_calls += 1 + self.max_native_calls = max(self.max_native_calls, self._native_calls) + + def _leave_native_call(self): + with self._native_call_lock: + self._native_calls -= 1 + + def register_memory(self, _buffer): + return [(1000, 1, 0, "")] + + def get_xfer_descs(self, pages_data, _mem_type): + return pages_data + + def prep_xfer_dlist(self, agent_name, descs, _mem_type): + return (agent_name, tuple(descs)) + + def get_agent_metadata(self): + return b"local-metadata" + + def get_serialized_descs(self, _reg_desc): + return b"local-page-desc" + + def get_new_notifs(self): + return {} + + def add_remote_agent(self, metadata): + self._enter_native_call() + try: + time.sleep(0.01) + self.add_calls.append(metadata) + return metadata.decode() + finally: + self._leave_native_call() + + def deserialize_descs(self, _page_reg_desc): + return [(2000, 1, 0, "")] + + def remove_remote_agent(self, peer_name): + self._enter_native_call() + try: + self.remove_calls.append(peer_name) + finally: + self._leave_native_call() + + def release_dlist_handle(self, handle): + self.released_dlists.append(handle) + + def send_notif(self, remote_agent_name, notif_msg): + self._enter_native_call() + try: + self.send_calls.append((remote_agent_name, notif_msg)) + if self.fail_next_send: + self.fail_next_send = False + raise RuntimeError("remote disconnected") + finally: + self._leave_native_call() + + def make_prepped_xfer(self, *_args): + return object() + + def transfer(self, _handle): + return "PROC" + + def check_xfer_state(self, _handle): + return "DONE" + + def release_xfer_handle(self, handle): + self.released_xfers.append(handle) + + def deregister_memory(self, _reg_desc): + return None + + +@pytest.fixture +def transporter(monkeypatch): + _FakeNixlAgent.instances.clear() + monkeypatch.setattr(nixl_kv_transporter, "NixlWrapper", _FakeNixlAgent) + monkeypatch.setattr(nixl_kv_transporter, "nixl_agent_config", _FakeConfig, raising=False) + monkeypatch.setattr(nixl_kv_transporter, "nixl_thread_sync_t", _FakeSyncMode, raising=False) + return nixl_kv_transporter.NixlKVTransporter(node_id=1, tp_idx=0, kv_move_buffer=_FakeBuffer()) + + +def _remote_agent(name="peer"): + return PDAgentMetadata( + agent_name=name, + agent_metadata=name.encode(), + num_pages=1, + page_reg_desc=b"remote-page-desc", + ) + + +def _task(remote_agent): + return PDChunckedTransTask( + request_id=1, + start_kv_index=0, + end_kv_index=1, + time_out_secs=60, + pd_master_node_id=1, + prefill_dp_index=0, + decode_dp_index=0, + src_device_id=0, + dst_device_id=0, + mem_indexes=[0], + prefill_agent_name=remote_agent.agent_name, + prefill_agent_metadata=remote_agent.agent_metadata, + prefill_num_pages=remote_agent.num_pages, + prefill_page_reg_desc=remote_agent.page_reg_desc, + decode_agent_name="decode", + decode_agent_metadata=b"decode", + decode_num_pages=1, + decode_page_reg_desc=b"decode-page-desc", + first_gen_token_id=None, + first_gen_token_logprob=None, + src_page_index=0, + dst_page_index=0, + ) + + +def _run_concurrently(callables): + barrier = threading.Barrier(len(callables) + 1) + errors = [] + + def run(callable_): + barrier.wait() + try: + callable_() + except BaseException as exc: + errors.append(exc) + + threads = [threading.Thread(target=run, args=(callable_,)) for callable_ in callables] + for thread in threads: + thread.start() + barrier.wait() + for thread in threads: + thread.join(timeout=2) + assert not thread.is_alive() + return errors + + +def test_uses_strict_nixl_sync_mode(transporter): + assert transporter.nixl_agent.config.sync_mode is _FakeSyncMode.NIXL_THREAD_SYNC_STRICT + + +def test_concurrent_peer_admission_calls_native_add_once(transporter): + remote_agent = _remote_agent() + + errors = _run_concurrently( + [ + lambda: transporter.connect_add_remote_agent(remote_agent), + lambda: transporter.connect_add_remote_agent(remote_agent), + ] + ) + + assert errors == [] + assert transporter.nixl_agent.add_calls == [b"peer"] + assert transporter.nixl_agent.max_native_calls == 1 + + +def test_failed_send_is_removed_before_one_serialized_reconnect(transporter): + remote_agent = _remote_agent() + task = _task(remote_agent) + transporter.connect_add_remote_agent(remote_agent) + transporter.nixl_agent.fail_next_send = True + + errors = _run_concurrently( + [ + lambda: transporter.send_write_ready_task_to_prefill_node(task), + lambda: transporter.send_write_ready_task_to_prefill_node(task), + ] + ) + + assert len(errors) == 1 + assert str(errors[0]) == "remote disconnected" + assert transporter.nixl_agent.add_calls == [b"peer", b"peer"] + assert transporter.nixl_agent.remove_calls == ["peer"] + assert transporter._peer_generations == {"peer": 2} + assert transporter.nixl_agent.max_native_calls == 1 + + +def test_active_transfer_defers_removal_and_stale_generation_cannot_remove_reconnect(transporter): + remote_agent = _remote_agent("decode") + task = _task(_remote_agent()) + handle = transporter.write_blocks_paged(task) + first_generation = transporter._peer_generations["decode"] + + with transporter._nixl_lock: + transporter._mark_remote_agent_broken_locked("decode", first_generation) + + assert "decode" in transporter.remote_agents + assert "decode" in transporter._broken_remote_agents + with pytest.raises(RuntimeError, match="active transfers drain"): + transporter.connect_add_remote_agent(remote_agent) + + transporter.release_xfer_handle(handle) + transporter.connect_add_remote_agent(remote_agent) + second_generation = transporter._peer_generations["decode"] + assert second_generation > first_generation + + with transporter._nixl_lock: + transporter._mark_remote_agent_broken_locked("decode", first_generation) + + assert transporter._peer_generations == {"decode": second_generation} + assert "decode" in transporter.remote_agents + + +def _status_module(transporter, task): + module = _PrefillTransModule.__new__(_PrefillTransModule) + module.transporter = transporter + module.waiting_dict_lock = threading.Lock() + module.waiting_dict = {task.get_key(): task} + module.success_queue = queue.Queue() + module.failed_queue = queue.Queue() + return module + + +def _status_task(): + return SimpleNamespace( + xfer_handle=object(), + error_info=None, + request_id=1, + start_kv_index=0, + end_kv_index=1, + src_page_index=0, + dst_page_index=0, + get_key=lambda: "task-1", + to_str=lambda: "task-1", + time_out=lambda: False, + ) + + +def test_status_query_failure_routes_task_to_failure_cleanup(): + class FailingTransporter: + capture_telemetry = False + + @staticmethod + def check_task_status(trans_task): + raise RuntimeError("native status failure") + + task = _status_task() + module = _status_module(FailingTransporter(), task) + + module._update_task_status_once() + + assert module.waiting_dict == {} + assert module.failed_queue.get_nowait() is task + assert task.error_info == "check transfer status failed: native status failure" + assert module.success_queue.empty() + + +def test_done_notification_failure_routes_task_to_failure_cleanup(): + class FailingTransporter: + capture_telemetry = False + + @staticmethod + def check_task_status(trans_task): + return "DONE" + + @staticmethod + def send_write_done_task_to_decode_node(trans_task): + raise RuntimeError("notify failure") + + task = _status_task() + module = _status_module(FailingTransporter(), task) + + module._update_task_status_once() + + assert module.waiting_dict == {} + assert module.failed_queue.get_nowait() is task + assert task.error_info == "send WRITE done nixl notify failed: notify failure" + assert module.success_queue.empty()