From f6506cb1727ed9524a634db273bcb6f00a5d663b Mon Sep 17 00:00:00 2001 From: baishihao Date: Wed, 30 Sep 2026 12:02:30 +0800 Subject: [PATCH] fix(pd): wait for KV transport readiness and track transfer progress --- lightllm/server/api_http.py | 13 ++- lightllm/server/api_openai.py | 10 ++- .../mode_backend/pd/base_kv_move_manager.py | 6 ++ .../decode_kv_move_manager.py | 7 +- .../decode_node_impl/decode_trans_process.py | 41 +++++++-- .../prefill_kv_move_manager.py | 7 +- .../prefill_trans_process.py | 5 ++ .../mode_backend/pd/trans_process_obj.py | 19 ++++ lightllm/utils/error_utils.py | 4 + .../mode_backend/test_pd_transfer_progress.py | 90 +++++++++++++++++++ .../server/test_generation_error_response.py | 29 ++++++ 11 files changed, 220 insertions(+), 11 deletions(-) create mode 100644 unit_tests/server/router/model_infer/mode_backend/test_pd_transfer_progress.py create mode 100644 unit_tests/server/test_generation_error_response.py diff --git a/lightllm/server/api_http.py b/lightllm/server/api_http.py index 9f3f7dd19d..3184d39420 100755 --- a/lightllm/server/api_http.py +++ b/lightllm/server/api_http.py @@ -46,7 +46,13 @@ from .api_lightllm import lightllm_get_score from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.log_utils import init_logger -from lightllm.utils.error_utils import ClientDisconnected, InvalidRequestError, SERVER_BUSY_MESSAGE, ServerBusyError +from lightllm.utils.error_utils import ( + ClientDisconnected, + GenerationError, + InvalidRequestError, + SERVER_BUSY_MESSAGE, + ServerBusyError, +) from lightllm.server.metrics.manager import MetricClient from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.shm_port_args import get_shm_port_args @@ -181,6 +187,11 @@ async def invalid_request_exception_handler(request: Request, exc: InvalidReques return create_error_response(HTTPStatus.BAD_REQUEST, str(exc)) +@app.exception_handler(GenerationError) +async def generation_exception_handler(request: Request, exc: GenerationError) -> JSONResponse: + return create_error_response(HTTPStatus.INTERNAL_SERVER_ERROR, str(exc)) + + @app.get("/liveness") @app.post("/liveness") def liveness(): diff --git a/lightllm/server/api_openai.py b/lightllm/server/api_openai.py index f258155b4c..b408c9bb6e 100644 --- a/lightllm/server/api_openai.py +++ b/lightllm/server/api_openai.py @@ -30,7 +30,13 @@ from .httpserver_for_pd_master.manager import HttpServerManagerForPDMaster from .api_lightllm import lightllm_get_score from lightllm.utils.envs_utils import get_env_start_args, get_lightllm_websocket_max_message_size -from lightllm.utils.error_utils import ClientDisconnected, InvalidRequestError, SERVER_BUSY_MESSAGE, ServerBusyError +from lightllm.utils.error_utils import ( + ClientDisconnected, + GenerationError, + InvalidRequestError, + SERVER_BUSY_MESSAGE, + ServerBusyError, +) from lightllm.utils.log_utils import init_logger from lightllm.server.metrics.manager import MetricClient @@ -523,6 +529,8 @@ async def stream_results() -> AsyncGenerator[bytes, None]: delta = request_output current_finish_reason = finish_status.get_finish_reason() + if current_finish_reason == "error" and completion_tokens == 1: + raise GenerationError("Generation failed before producing output") # Emit the initial role-only chunk once per choice, as required by the # OpenAI SSE spec: role appears only in the first delta with content="". diff --git a/lightllm/server/router/model_infer/mode_backend/pd/base_kv_move_manager.py b/lightllm/server/router/model_infer/mode_backend/pd/base_kv_move_manager.py index 125edede25..1e170975a3 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/base_kv_move_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/base_kv_move_manager.py @@ -44,6 +44,12 @@ def __init__( start_func=start_trans_process_func, up_status_in_queue=up_status_in_queue, ) + + for trans_process in self.kv_trans_processes: + if not trans_process.wait_until_ready(): + raise RuntimeError(f"KV trans module for device {trans_process.device_id} failed to initialize") + + for trans_process in self.kv_trans_processes: threading.Thread(target=self.task_ret_handle_loop, args=(trans_process,), daemon=True).start() # 通过 io buffer 将命令写入到推理进程中 diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_kv_move_manager.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_kv_move_manager.py index 121a528a43..fcfbaf1936 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_kv_move_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_kv_move_manager.py @@ -21,8 +21,11 @@ def start_decode_kv_move_manager_process(args, info_queue: mp.Queue): event = mp.Event() proc = mp.Process(target=_init_env, args=(args, info_queue, event)) proc.start() - event.wait() - assert proc.is_alive() + while not event.wait(timeout=1): + if not proc.is_alive(): + raise RuntimeError("decode kv move manager process failed during initialization") + if not proc.is_alive(): + raise RuntimeError("decode kv move manager process exited during initialization") logger.info("decode kv move manager process started") return 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 b406405e8a..ff1d3de186 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 @@ -47,6 +47,7 @@ def _init_env( task_out_queue: mp.Queue, up_status_in_queue: Optional[mp.SimpleQueue], ): + module_ready = False install_fatal_thread_excepthook() start_parent_check_thread() import lightllm.utils.rpyc_fix_utils as _ @@ -100,11 +101,15 @@ def _init_env( up_status_in_queue=up_status_in_queue, ) assert manager is not None + task_out_queue.put("module_ready") + module_ready = True while True: time.sleep(100) except Exception as e: + if not module_ready: + task_out_queue.put("init_failed") logger.exception(str(e)) logger.error(f"Fatal error happened in kv trans process: {e}") pass @@ -141,6 +146,7 @@ def __init__( self.recv_task_group_queue = queue.Queue() self.waiting_dict_lock = threading.Lock() self.waiting_dict: Dict[str, PDChunckedTransTask] = {} + self.request_last_progress_time: Dict[int, float] = {} self.request_page_task_queue = queue.Queue() self.ready_page_task_queue = queue.Queue() self.success_queue = queue.Queue() @@ -239,6 +245,19 @@ def dispatch_task_loop(self): self.up_status_in_queue.put(up_status) + def _pop_waiting_task_for_notify(self, notify_task: PDChunckedTransTask): + with self.waiting_dict_lock: + local_trans_task = self.waiting_dict.pop(notify_task.get_key(), None) + if local_trans_task is None: + return None + + # Decode creates every page task before prefill starts producing pages. + # A matched notify is forward progress for the request, so future pages + # use an idle timeout instead of their original creation time. + self.request_last_progress_time[local_trans_task.request_id] = time.time() + + return local_trans_task + @log_exception def accept_peer_task_loop( self, @@ -287,8 +306,7 @@ def accept_peer_task_loop( # 到了请求页面的阶段 remote_trans_task = notify_obj if remote_trans_task.write_stage == "request": - with self.waiting_dict_lock: - local_trans_task = self.waiting_dict.pop(remote_trans_task.get_key(), None) + local_trans_task = self._pop_waiting_task_for_notify(remote_trans_task) if local_trans_task is not None: local_trans_task.prefill_agent_name = remote_trans_task.prefill_agent_name local_trans_task.prefill_agent_metadata = remote_trans_task.prefill_agent_metadata @@ -316,8 +334,7 @@ def accept_peer_task_loop( # prefill 写完数据到了 done 阶段 if remote_trans_task.write_stage == "done": - with self.waiting_dict_lock: - local_trans_task = self.waiting_dict.pop(remote_trans_task.get_key(), None) + local_trans_task = self._pop_waiting_task_for_notify(remote_trans_task) if local_trans_task is not None: local_trans_task.first_gen_token_id = remote_trans_task.first_gen_token_id local_trans_task.first_gen_token_logprob = remote_trans_task.first_gen_token_logprob @@ -345,9 +362,23 @@ def accept_peer_task_loop( def _check_tasks_time_out(self): with self.waiting_dict_lock: timeout_tasks = [] + pending_request_ids = set() + now = time.time() for key, trans_task in list(self.waiting_dict.items()): - if trans_task.time_out(): + if trans_task.start_trans_time is None: + request_last_progress = self.request_last_progress_time.get(trans_task.request_id) + is_timeout = ( + request_last_progress is not None and now - request_last_progress > trans_task.time_out_secs + ) + else: + is_timeout = trans_task.time_out() + if is_timeout: timeout_tasks.append(self.waiting_dict.pop(key)) + elif trans_task.start_trans_time is None: + pending_request_ids.add(trans_task.request_id) + for request_id in list(self.request_last_progress_time): + if request_id not in pending_request_ids: + self.request_last_progress_time.pop(request_id) for trans_task in timeout_tasks: trans_task.error_info = "time out in accept_peer_task_loop" diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py index b23a5c4141..e09ba79c95 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py @@ -20,8 +20,11 @@ def start_prefill_kv_move_manager_process(args, info_queue: mp.Queue): event = mp.Event() proc = mp.Process(target=_init_env, args=(args, info_queue, event)) proc.start() - event.wait() - assert proc.is_alive() + while not event.wait(timeout=1): + if not proc.is_alive(): + raise RuntimeError("prefill kv move manager process failed during initialization") + if not proc.is_alive(): + raise RuntimeError("prefill kv move manager process exited during initialization") logger.info("prefill kv move manager process started") 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 534966ebe9..de9894074c 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 @@ -41,6 +41,7 @@ def _init_env( task_in_queue: mp.Queue, task_out_queue: mp.Queue, ): + module_ready = False install_fatal_thread_excepthook() start_parent_check_thread() import lightllm.utils.rpyc_fix_utils as _ @@ -74,11 +75,15 @@ def _init_env( mem_managers=mem_managers, ) assert manager is not None + task_out_queue.put("module_ready") + module_ready = True while True: time.sleep(100) except Exception as e: + if not module_ready: + task_out_queue.put("init_failed") logger.exception(str(e)) logger.error(f"Fatal error happened in kv trans process: {e}") pass diff --git a/lightllm/server/router/model_infer/mode_backend/pd/trans_process_obj.py b/lightllm/server/router/model_infer/mode_backend/pd/trans_process_obj.py index 073ecf23d2..84173a44ba 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/trans_process_obj.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/trans_process_obj.py @@ -1,3 +1,4 @@ +import queue import threading import psutil import torch.multiprocessing as mp @@ -58,5 +59,23 @@ def is_trans_process_health(self): except: return False + def wait_until_ready(self): + for _ in range(600): + try: + status = self.task_out_queue.get(timeout=1) + except queue.Empty: + if not self.process.is_alive(): + logger.error(f"KV trans process for device {self.device_id} exited during initialization") + return False + continue + + if status != "module_ready": + logger.error(f"KV trans module for device {self.device_id} failed to initialize: {status}") + return False + return True + + logger.error(f"KV trans module for device {self.device_id} initialization timed out") + return False + def killself(self): self.process.kill() diff --git a/lightllm/utils/error_utils.py b/lightllm/utils/error_utils.py index acf76ad54d..25d7744e50 100644 --- a/lightllm/utils/error_utils.py +++ b/lightllm/utils/error_utils.py @@ -10,6 +10,10 @@ class InvalidRequestError(ValueError): """Request validation failed before generation started.""" +class GenerationError(Exception): + """Generation stopped because of an internal server failure.""" + + class ServerBusyError(Exception): """Custom exception for server busy/overload situations""" diff --git a/unit_tests/server/router/model_infer/mode_backend/test_pd_transfer_progress.py b/unit_tests/server/router/model_infer/mode_backend/test_pd_transfer_progress.py new file mode 100644 index 0000000000..8e4b8ff613 --- /dev/null +++ b/unit_tests/server/router/model_infer/mode_backend/test_pd_transfer_progress.py @@ -0,0 +1,90 @@ +import queue +import threading +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest + +from lightllm.server.router.model_infer.mode_backend.pd.decode_node_impl import decode_trans_process +from lightllm.server.router.model_infer.mode_backend.pd.trans_process_obj import KVTransProcess + + +@pytest.mark.parametrize("status, expected", [("module_ready", True), ("init_failed", False)]) +def test_transfer_process_requires_module_ready(status, expected): + worker = KVTransProcess(process=Mock(), task_out_queue=queue.Queue(), device_id=0) + worker.task_out_queue.put(status) + assert worker.wait_until_ready() is expected + + +def test_transfer_process_waits_through_slow_initialization(): + output = Mock() + output.get.side_effect = [queue.Empty(), queue.Empty(), "module_ready"] + worker = KVTransProcess(process=Mock(), task_out_queue=output, device_id=0) + worker.process.is_alive.return_value = True + assert worker.wait_until_ready() + assert output.get.call_count == 3 + + +def test_transfer_process_detects_exit_before_ready(): + worker = KVTransProcess(process=Mock(), task_out_queue=Mock(), device_id=0) + worker.task_out_queue.get.side_effect = queue.Empty() + worker.process.is_alive.return_value = False + assert not worker.wait_until_ready() + assert worker.task_out_queue.get.call_count == 1 + + +def test_transfer_process_initialization_has_bounded_wait(): + worker = KVTransProcess(process=Mock(), task_out_queue=Mock(), device_id=0) + worker.task_out_queue.get.side_effect = queue.Empty() + worker.process.is_alive.return_value = True + assert not worker.wait_until_ready() + assert worker.task_out_queue.get.call_count == 600 + + +def make_task(request_id, page, started=None): + task = SimpleNamespace(request_id=request_id, start_trans_time=started, time_out_secs=10) + task.get_key = lambda: f"{request_id}:{page}" + task.time_out = Mock(return_value=False) + return task + + +def test_waiting_pages_timeout_from_request_progress_not_creation(monkeypatch): + now = [100.0] + monkeypatch.setattr(decode_trans_process.time, "time", lambda: now[0]) + module = decode_trans_process._DecodeTransModule.__new__(decode_trans_process._DecodeTransModule) + module.waiting_dict_lock = threading.Lock() + first, second, unrelated = make_task(1, 0), make_task(1, 1), make_task(2, 0) + module.waiting_dict = {task.get_key(): task for task in (first, second, unrelated)} + module.request_last_progress_time = {} + module.failed_queue = queue.Queue() + + # No transfer progress yet: the PD master owns the prefill-stage deadline. + module._check_tasks_time_out() + assert len(module.waiting_dict) == 3 + assert module._pop_waiting_task_for_notify(first) is first + assert module.request_last_progress_time == {1: 100.0} + + now[0] = 109.0 + module._check_tasks_time_out() + assert second.get_key() in module.waiting_dict + now[0] = 111.0 + module._check_tasks_time_out() + assert module.failed_queue.get_nowait() is second + assert unrelated.get_key() in module.waiting_dict + assert not module.request_last_progress_time + assert module._pop_waiting_task_for_notify(first) is None + + +def test_started_transfer_keeps_per_page_timeout(monkeypatch): + monkeypatch.setattr(decode_trans_process.time, "time", lambda: 100.0) + module = decode_trans_process._DecodeTransModule.__new__(decode_trans_process._DecodeTransModule) + task = make_task(1, 0, started=1.0) + task.time_out.return_value = True + module.waiting_dict = {task.get_key(): task} + module.waiting_dict_lock = threading.Lock() + module.request_last_progress_time = {1: 99.0} + module.failed_queue = queue.Queue() + module._check_tasks_time_out() + task.time_out.assert_called_once_with() + assert module.failed_queue.get_nowait() is task + assert not module.waiting_dict and not module.request_last_progress_time diff --git a/unit_tests/server/test_generation_error_response.py b/unit_tests/server/test_generation_error_response.py new file mode 100644 index 0000000000..a08a5c6c30 --- /dev/null +++ b/unit_tests/server/test_generation_error_response.py @@ -0,0 +1,29 @@ +from unittest.mock import Mock + +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from lightllm.server import api_stream_obj +from lightllm.server.core.objs import StartArgs +from lightllm.server.api_http import g_objs +from lightllm.server.api_http import generation_exception_handler +from lightllm.utils.error_utils import GenerationError + + +def test_generation_failure_before_first_chunk_returns_http_500(monkeypatch): + monkeypatch.setattr(g_objs, "metric_client", Mock()) + monkeypatch.setattr(api_stream_obj, "get_env_start_args", lambda: StartArgs()) + app = FastAPI() + app.add_exception_handler(GenerationError, generation_exception_handler) + + @app.get("/generate") + async def generate(): + async def failed_stream(): + raise GenerationError("Generation failed before producing output") + yield + + return api_stream_obj.CustomStreamingResponse(failed_stream(), media_type="text/event-stream") + + response = TestClient(app).get("/generate") + assert response.status_code == 500 + assert "Generation failed" in response.json()["error"]["message"]