Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 12 additions & 1 deletion lightllm/server/api_http.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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():
Expand Down
10 changes: 9 additions & 1 deletion lightllm/server/api_openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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="".
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 将命令写入到推理进程中
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 _
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 _
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import queue
import threading
import psutil
import torch.multiprocessing as mp
Expand Down Expand Up @@ -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()
4 changes: 4 additions & 0 deletions lightllm/utils/error_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"""

Expand Down
Original file line number Diff line number Diff line change
@@ -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
29 changes: 29 additions & 0 deletions unit_tests/server/test_generation_error_response.py
Original file line number Diff line number Diff line change
@@ -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"]
Loading