diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 60b5fad4e8..72f2b8c891 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -71,7 +71,30 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: parser.add_argument( "--disable_pd_master_decode_capacity_limit", action="store_true", - help="Disable PD master admission control based on the total capacity of registered decode nodes.", + help=( + "Deprecated alias for --disable_pd_node_decode_admission. Set it consistently on PD Masters and " + "Decode nodes when intentionally allowing Decode nodes without authoritative admission." + ), + ) + parser.add_argument( + "--disable_pd_node_decode_admission", + action="store_true", + help="Disable the bounded request-slot admission queue on PD Decode nodes.", + ) + parser.add_argument( + "--pd_node_decode_admission_queue_size", + type=int, + default=None, + help=( + "Maximum number of Decode request slots allowed to wait on a PD Decode node. " + "Defaults to --running_max_req_size." + ), + ) + parser.add_argument( + "--pd_node_decode_admission_timeout", + type=float, + default=5.0, + help="Maximum seconds a request may wait in the PD Decode node admission queue. Default: 5.", ) parser.add_argument( "--pd_trans_mode", diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 37fe837ad1..8eac4f2dad 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -121,6 +121,11 @@ def _launch_subprocesses(args: StartArgs): args.mem_fraction > 0 and args.mem_fraction < 1 ), f"Invalid mem_fraction {args.mem_fraction}, The expected value is between 0 and 1." + if args.pd_node_decode_admission_queue_size is not None and args.pd_node_decode_admission_queue_size < 0: + raise ValueError("pd_node_decode_admission_queue_size must be non-negative") + if args.pd_node_decode_admission_timeout <= 0: + raise ValueError("pd_node_decode_admission_timeout must be positive") + if args.graph_max_len_in_batch == 0: args.graph_max_len_in_batch = args.max_req_total_len diff --git a/lightllm/server/core/objs/shm_req_manager.py b/lightllm/server/core/objs/shm_req_manager.py index c61376c8d0..4215e63d69 100644 --- a/lightllm/server/core/objs/shm_req_manager.py +++ b/lightllm/server/core/objs/shm_req_manager.py @@ -82,17 +82,34 @@ def init_alloc_state_shm(self): # alloc_req_index 和 release_req_index 是分配资源时使用的接口。 # 只有管理请求申请和释放的首节点才能调用这个接口。 def alloc_req_index(self): + indexes = self.alloc_req_indexes(1) + return None if indexes is None else indexes[0] + + def alloc_req_indexes(self, req_num: int): + """Atomically allocate all requested indexes or leave the pool unchanged.""" + if req_num < 1: + raise ValueError("req_num must be positive") + with self.manager_lock: - idx = self.linked_req_manager.alloc() - if idx is not None: + indexes = [] + for _ in range(req_num): + idx = self.linked_req_manager.alloc() + if idx is None: + for allocated_idx in reversed(indexes): + self.alloc_state_shm.arr[allocated_idx] = 0 + self.linked_req_manager.free(allocated_idx) + return None assert self.alloc_state_shm.arr[idx] == 0 self.alloc_state_shm.arr[idx] = 1 - return idx - return None + indexes.append(idx) + return indexes async def async_alloc_req_index(self): return self.alloc_req_index() + async def async_alloc_req_indexes(self, req_num: int): + return self.alloc_req_indexes(req_num) + def release_req_index(self, req_index_in_mem): assert req_index_in_mem < self.max_req_num with self.manager_lock: diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index a9aef608bd..76bb4715c9 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -23,6 +23,9 @@ class StartArgs: pd_master_port: int = field(default=1212) pd_master_mode: str = field(default="elastic") disable_pd_master_decode_capacity_limit: bool = field(default=False) + disable_pd_node_decode_admission: bool = field(default=False) + pd_node_decode_admission_queue_size: Optional[int] = field(default=None) + pd_node_decode_admission_timeout: float = field(default=5.0) pd_trans_mode: str = field(default="nccl", metadata={"choices": ["nccl", "nixl"]}) config_server_host: str = field(default=None) config_server_port: int = field(default=None) diff --git a/lightllm/server/httpserver/decode_admission.py b/lightllm/server/httpserver/decode_admission.py new file mode 100644 index 0000000000..18fc3f1add --- /dev/null +++ b/lightllm/server/httpserver/decode_admission.py @@ -0,0 +1,180 @@ +from __future__ import annotations + +import asyncio +import time +from collections import deque +from dataclasses import dataclass +from typing import Callable, Deque, Optional + +from lightllm.utils.error_utils import ServerBusyError + + +@dataclass(slots=True) +class _Waiter: + slots: int + enqueue_time: float + future: asyncio.Future + + +class DecodeAdmissionLease: + """A fixed number of request slots reserved on one Decode node.""" + + def __init__(self, controller: "DecodeAdmissionController", slots: int, waited_seconds: float) -> None: + self._controller = controller + self.slots = slots + self.waited_seconds = waited_seconds + self._released = False + + def release(self) -> None: + if self._released: + return + self._released = True + self._controller._release(self.slots) + + def split(self, slots: list[int]) -> list["DecodeAdmissionLease"]: + """Transfer this reservation into independently releasable child leases.""" + if self._released: + raise RuntimeError("Decode admission lease has already been released") + if not slots or any(slot_count < 1 for slot_count in slots) or sum(slots) != self.slots: + raise ValueError("child lease slots must be positive and sum to the parent lease") + + self._released = True + return [DecodeAdmissionLease(self._controller, slot_count, self.waited_seconds) for slot_count in slots] + + +class DecodeAdmissionLeaseHandle: + """Transfers a pre-acquired lease into an asynchronously started generator.""" + + def __init__(self, lease: DecodeAdmissionLease) -> None: + self._lease: Optional[DecodeAdmissionLease] = lease + + def take(self) -> DecodeAdmissionLease: + if self._lease is None: + raise RuntimeError("Decode admission lease has already been taken") + lease = self._lease + self._lease = None + return lease + + def release(self) -> None: + if self._lease is None: + return + self._lease.release() + self._lease = None + + +class DecodeAdmissionController: + """A bounded, cancellable FIFO for request slots owned by one Decode node.""" + + def __init__( + self, + capacity: int, + max_queued_slots: int, + timeout_seconds: float, + clock: Callable[[], float] = time.monotonic, + ) -> None: + if capacity < 1: + raise ValueError("decode admission capacity must be positive") + if max_queued_slots < 0: + raise ValueError("decode admission queue size must be non-negative") + if timeout_seconds <= 0: + raise ValueError("decode admission timeout must be positive") + + self.capacity = capacity + self.max_queued_slots = max_queued_slots + self.timeout_seconds = timeout_seconds + self._clock = clock + self._active_slots = 0 + self._queued_slots = 0 + self._waiters: Deque[_Waiter] = deque() + + @property + def active_slots(self) -> int: + return self._active_slots + + @property + def queued_slots(self) -> int: + return self._queued_slots + + @property + def queued_request_count(self) -> int: + return len(self._waiters) + + async def acquire(self, slots: int) -> DecodeAdmissionLease: + if slots < 1: + raise ValueError("decode admission slots must be positive") + if slots > self.capacity: + raise ServerBusyError(f"request needs {slots} Decode slots, but the node capacity is {self.capacity}") + + if not self._waiters and self._active_slots + slots <= self.capacity: + return self._activate(slots, waited_seconds=0.0) + + if self._queued_slots + slots > self.max_queued_slots: + raise ServerBusyError("Decode node admission queue is full") + + waiter = _Waiter( + slots=slots, + enqueue_time=self._clock(), + future=asyncio.get_running_loop().create_future(), + ) + self._waiters.append(waiter) + self._queued_slots += slots + self._drain() + + try: + return await asyncio.wait_for(asyncio.shield(waiter.future), timeout=self.timeout_seconds) + except asyncio.TimeoutError as exc: + lease = self._remove_waiter_or_take_lease(waiter) + if lease is not None: + lease.release() + raise ServerBusyError("Decode node admission queue wait timed out") from exc + except BaseException: + lease = self._remove_waiter_or_take_lease(waiter) + if lease is not None: + lease.release() + raise + + def _activate(self, slots: int, waited_seconds: float) -> DecodeAdmissionLease: + self._active_slots += slots + return DecodeAdmissionLease(self, slots, waited_seconds) + + def _release(self, slots: int) -> None: + self._active_slots -= slots + if self._active_slots < 0: + raise RuntimeError("Decode admission active slot count became negative") + self._drain() + + def _drain(self) -> None: + while self._waiters: + waiter = self._waiters[0] + if waiter.future.done(): + self._remove_waiter(waiter) + continue + if self._active_slots + waiter.slots > self.capacity: + break + + self._remove_waiter(waiter) + waiter.future.set_result( + self._activate( + waiter.slots, + waited_seconds=max(0.0, self._clock() - waiter.enqueue_time), + ) + ) + + def _remove_waiter(self, waiter: _Waiter) -> bool: + try: + self._waiters.remove(waiter) + except ValueError: + return False + self._queued_slots -= waiter.slots + return True + + def _remove_waiter_or_take_lease(self, waiter: _Waiter) -> Optional[DecodeAdmissionLease]: + if self._remove_waiter(waiter): + waiter.future.cancel() + self._drain() + return None + if waiter.future.done() and not waiter.future.cancelled(): + result = waiter.future.result() + if isinstance(result, DecodeAdmissionLease): + return result + return None diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 9e6f77e58d..b00feb0a00 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -34,6 +34,7 @@ from lightllm.server.metrics.manager import MetricClient from .rl_controller import HttpRlController from .manager_ext import HttpRlManagerHelper +from .decode_admission import DecodeAdmissionController, DecodeAdmissionLease, DecodeAdmissionLeaseHandle from lightllm.utils.statics_utils import MovingAverage from lightllm.utils.config_utils import get_vocab_size from lightllm.utils.envs_utils import get_unique_server_name @@ -117,6 +118,19 @@ def __init__( self.pd_mode: NodeRole = NodeRole(self.args.run_mode) assert self.pd_mode in [NodeRole.NORMAL, NodeRole.P, NodeRole.D] + self.decode_admission_controller: Optional[DecodeAdmissionController] = None + decode_admission_disabled = ( + args.disable_pd_node_decode_admission or args.disable_pd_master_decode_capacity_limit + ) + if self.pd_mode.is_D() and not self.is_multinode_tp_slave and not decode_admission_disabled: + max_queued_slots = args.pd_node_decode_admission_queue_size + if max_queued_slots is None: + max_queued_slots = args.running_max_req_size + self.decode_admission_controller = DecodeAdmissionController( + capacity=args.running_max_req_size, + max_queued_slots=max_queued_slots, + timeout_seconds=args.pd_node_decode_admission_timeout, + ) self.id_gen = ReqIDGenerator() self.first_time_costs = MovingAverage() self.per_token_costs = MovingAverage() @@ -322,6 +336,9 @@ async def generate( pd_upload_websocket: ClientConnection = None, # 用于等待 pd_master 下发的交换信息 pd_event: asyncio.Event = None, + # Decode Master 可在所有 choice 完成 Prefill 后原子预留 n 个槽, + # 再通过该 handle 把每个 choice 的 lease 转交给本地请求生命周期。 + decode_admission_lease_handle: Optional[DecodeAdmissionLeaseHandle] = None, ) -> AsyncGenerator[Tuple[int, str, dict, FinishStatus], None]: start_time = time.time() @@ -338,6 +355,11 @@ async def generate( ) running_request_registered = False + decode_admission_lease: Optional[DecodeAdmissionLease] = None + if decode_admission_lease_handle is not None: + decode_admission_lease = decode_admission_lease_handle.take() + unmanaged_req_indexes = [] + unmanaged_req_objs = [] if not self.pd_mode.is_P(): await self._register_running_request() running_request_registered = True @@ -426,21 +448,26 @@ async def generate( await self._register_running_request() running_request_registered = True - # 申请资源并存储 - 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) + # Decode admission happens immediately before the node reserves its + # local shm_req objects. All Masters therefore compete for the same + # authoritative pool, and an n-choice request reserves all n slots at once. + admission_controller = getattr(self, "decode_admission_controller", None) + if admission_controller is not None and decode_admission_lease is None: + decode_admission_lease = await admission_controller.acquire(sampling_params.n) + self.metric_client.histogram_observe( + "lightllm_request_queue_duration_bucket", + decode_admission_lease.waited_seconds, + ) + + alloced_req_indexes = await self._alloc_req_indexes( + sampling_params.n, + expect_available=decode_admission_lease is not None, + ) + unmanaged_req_indexes.extend(alloced_req_indexes) 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) + unmanaged_req_objs.append(req_obj) req_obj.init( group_request_id + i, prompt_ids, @@ -461,8 +488,17 @@ async def generate( f"{[(req_obj.request_id, req_obj.index_in_shm_mem) for req_obj in req_objs]}" ) - req_status = ReqStatus(group_request_id, multimodal_params, req_objs, start_time) + req_status = ReqStatus( + group_request_id, + multimodal_params, + req_objs, + start_time, + decode_admission_lease=decode_admission_lease, + ) self.req_id_to_out_inf[group_request_id] = req_status + decode_admission_lease = None + unmanaged_req_indexes.clear() + unmanaged_req_objs.clear() # RL:请求已登记到 req_id_to_out_inf 并即将转发下游,从 admission gate # 注销,避免 pause 统计里仍把它算作“等待准入”的 pending 请求。 if self.rl_controller is not None: @@ -517,6 +553,10 @@ async def generate( # 已经放入到 req_id_to_out_inf 中的请求对象,由统一的回收循环 # 进行回收。 if group_request_id not in self.req_id_to_out_inf: + for req_obj in reversed(unmanaged_req_objs): + await self.shm_req_manager.async_put_back_req_obj(req_obj) + for req_index in reversed(unmanaged_req_indexes): + await self.shm_req_manager.async_release_req_index(req_index) await self._release_multimodal_resources(multimodal_params) await self.abort(group_request_id) raise e @@ -527,8 +567,25 @@ async def generate( await self.rl_controller.unregister_generation_admission(group_request_id) if running_request_registered: await self._unregister_running_request() + if decode_admission_lease is not None: + decode_admission_lease.release() return + async def _alloc_req_indexes(self, req_num: int, expect_available: bool) -> List[int]: + """Wait for one atomic batch unless Decode admission already reserved it.""" + indexes = await self.shm_req_manager.async_alloc_req_indexes(req_num) + if indexes is not None: + return indexes + if expect_available: + raise RuntimeError("Decode admission reserved slots that are unavailable in shm_req") + + sleep_time = 0.1 + while indexes is None: + await asyncio.sleep(sleep_time) + sleep_time = min(1, sleep_time * 1.1) + indexes = await self.shm_req_manager.async_alloc_req_indexes(req_num) + return indexes + def _count_multimodal_tokens(self, multimodal_params: MultimodalParams) -> Tuple[int, int]: image_tokens = 0 audio_tokens = 0 @@ -900,6 +957,7 @@ async def recycle_resource_loop(self): await self.shm_req_manager.async_put_back_req_obj(req) await self.shm_req_manager.async_release_req_index(req.index_in_shm_mem) await self._release_multimodal_resources(req_status.group_req_objs.multimodal_params) + req_status.release_decode_admission() # 先保留这个关键得日志,用于方便定位重构中的问题。 if time.time() - pre_time_mark > 120: @@ -925,7 +983,10 @@ async def handle_loop(self): if self.is_multinode_tp_slave: asyncio.create_task(self.loop_for_request()) - if self.pd_mode.is_P_or_D(): + # Multi-node TP rank 0 is the single HTTP/PD ingress and forwards each + # admitted request to its slave ranks. Slaves must not register as + # independent P/D nodes or they could bypass rank 0 admission. + if self.pd_mode.is_P_or_D() and not self.is_multinode_tp_slave: from lightllm.server.httpserver.pd_loop import pd_handle_loop asyncio.create_task(pd_handle_loop(self)) @@ -1032,7 +1093,14 @@ async def _unregister_running_request(self): class ReqStatus: - def __init__(self, group_request_id, multimodal_params, req_objs: List[Req], start_time) -> None: + def __init__( + self, + group_request_id, + multimodal_params, + req_objs: List[Req], + start_time, + decode_admission_lease: Optional[DecodeAdmissionLease] = None, + ) -> None: self.lock = asyncio.Lock() self.event = asyncio.Event() self.group_req_objs = GroupReqObjs( @@ -1042,6 +1110,13 @@ def __init__(self, group_request_id, multimodal_params, req_objs: List[Req], sta time_mark=start_time, ) self.out_token_info_list = [] + self.decode_admission_lease = decode_admission_lease + + def release_decode_admission(self) -> None: + if self.decode_admission_lease is None: + return + self.decode_admission_lease.release() + self.decode_admission_lease = None def can_release(self): for req in self.group_req_objs.shm_req_objs: diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index e66540cb5e..bc21856c7a 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -11,21 +11,36 @@ import sys from typing import Dict, Optional, Union, List from websockets import ClientConnection -from lightllm.server.pd_io_struct import NodeRole, ObjType +from lightllm.server.pd_io_struct import PD_DECODE_ADMISSION_CAPABILITY_KEY, NodeRole, ObjType from lightllm.server.httpserver.async_queue import AsyncQueue from lightllm.utils.net_utils import get_hostname_ip from lightllm.utils.log_utils import init_logger from lightllm.utils.envs_utils import get_lightllm_websocket_max_message_size from lightllm.server.httpserver.manager import HttpServerManager +from lightllm.server.httpserver.decode_admission import DecodeAdmissionLeaseHandle from ..pd_io_struct import PD_Master_Obj from lightllm.server.core.objs import StartArgs from lightllm.server.core.objs import SamplingParams -from lightllm.utils.error_utils import PDPrefillNodeStopGenToken +from lightllm.utils.error_utils import PDPrefillNodeStopGenToken, ServerBusyError from lightllm.utils.shm_port_args import get_shm_port_args logger = init_logger(__name__) +def _build_pd_registration_info(manager: HttpServerManager) -> dict: + """构造保持旧顶层 schema 兼容的 P/D 节点注册信息。""" + args_dict = vars(manager.args).copy() + args_dict["host"] = manager.host_ip + if manager.pd_mode.is_D(): + args_dict[PD_DECODE_ADMISSION_CAPABILITY_KEY] = manager.decode_admission_controller is not None + return { + "node_id": manager.args.pd_node_id, + "client_ip_port": f"{manager.host_ip}:{get_shm_port_args().port}", + "mode": manager.pd_mode.value, + "start_args": args_dict, + } + + async def timer_log(manager: HttpServerManager): while True: await asyncio.sleep(30) @@ -34,6 +49,52 @@ async def timer_log(manager: HttpServerManager): return +async def _reserve_decode_slots( + manager: HttpServerManager, + reservation_id: int, + request_ids: tuple[int, ...], + reserved_lease_handles: Dict[int, DecodeAdmissionLeaseHandle], + websocket: ClientConnection, +) -> None: + installed_handles: Dict[int, DecodeAdmissionLeaseHandle] = {} + completed = False + lease = None + try: + controller = manager.decode_admission_controller + if controller is None: + raise ServerBusyError("Decode node admission is unavailable") + + lease = await controller.acquire(len(request_ids)) + manager.metric_client.histogram_observe("lightllm_request_queue_duration_bucket", lease.waited_seconds) + child_leases = lease.split([1] * len(request_ids)) + lease = None + installed_handles = { + request_id: DecodeAdmissionLeaseHandle(child_lease) + for request_id, child_lease in zip(request_ids, child_leases) + } + reserved_lease_handles.update(installed_handles) + await websocket.send(pickle.dumps((ObjType.PD_DECODE_SLOTS_RESERVED, reservation_id))) + completed = True + except asyncio.CancelledError: + raise + except ServerBusyError as error: + logger.warning(f"Decode reservation {reservation_id} rejected: {error.message}") + await websocket.send(pickle.dumps((ObjType.PD_UPLOAD_SERVER_BUSY, reservation_id, error.message))) + except BaseException as error: + logger.exception(f"Decode reservation {reservation_id} failed: {str(error)}") + await websocket.send( + pickle.dumps((ObjType.PD_UPLOAD_SERVER_BUSY, reservation_id, f"{type(error).__name__}: {str(error)}")) + ) + finally: + if not completed: + if lease is not None: + lease.release() + for request_id, handle in installed_handles.items(): + if reserved_lease_handles.get(request_id) is handle: + reserved_lease_handles.pop(request_id, None) + handle.release() + + 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") @@ -56,7 +117,7 @@ async def pd_handle_loop(manager: HttpServerManager): logger.info(f"get pd_master_objs {id_to_pd_master_obj}") if id_to_pd_master_obj is not None: - for node_id, pd_master_obj in id_to_handle_task.items(): + for node_id, pd_master_obj in list(id_to_handle_task.items()): if node_id not in id_to_pd_master_obj: id_to_handle_task[node_id].cancel() id_to_handle_task.pop(node_id, None) @@ -85,6 +146,9 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O forwarding_tokens_task = None heartbeat_task = None generation_tasks: Dict[int, asyncio.Task] = {} + reservation_tasks: Dict[int, asyncio.Task] = {} + request_id_to_reservation_task: Dict[int, asyncio.Task] = {} + reserved_lease_handles: Dict[int, DecodeAdmissionLeaseHandle] = {} try: uri = f"ws://{pd_master_obj.host_ip_port}/pd_register" async with websockets.connect( @@ -94,19 +158,11 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O # 下方应用层心跳已负责存活检测,禁用协议层 keepalive,避免繁忙连接被误断。 ping_interval=None, ) as websocket: - sock = websocket.transport.get_extra_info("socket") sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) - args_dict = vars(manager.args) - args_dict["host"] = manager.host_ip # 发送注册信息 - regist_json = { - "node_id": manager.args.pd_node_id, - "client_ip_port": f"{manager.host_ip}:{get_shm_port_args().port}", - "mode": manager.pd_mode.value, - "start_args": args_dict, - } + regist_json = _build_pd_registration_info(manager) await websocket.send(json.dumps(regist_json)) logger.info(f"Sent registration JSON: {regist_json}") @@ -120,9 +176,60 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O while True: recv_bytes = await websocket.recv() obj = pickle.loads(recv_bytes) - if obj[0] == ObjType.REQ: + if obj[0] == ObjType.PD_RESERVE_DECODE_SLOTS: + _, reservation_id, request_ids = obj + request_ids = tuple(request_ids) + if ( + not request_ids + or reservation_id in reservation_tasks + or len(set(request_ids)) != len(request_ids) + or any( + request_id in request_id_to_reservation_task + or request_id in reserved_lease_handles + or request_id in generation_tasks + for request_id in request_ids + ) + ): + await websocket.send( + pickle.dumps( + ( + ObjType.PD_UPLOAD_SERVER_BUSY, + reservation_id, + "invalid or duplicate Decode reservation", + ) + ) + ) + continue + + reservation_task = asyncio.create_task( + _reserve_decode_slots( + manager, reservation_id, request_ids, reserved_lease_handles, websocket + ) + ) + reservation_tasks[reservation_id] = reservation_task + for request_id in request_ids: + request_id_to_reservation_task[request_id] = reservation_task + + def remove_reservation_task( + task: asyncio.Task, + current_reservation_id: int = reservation_id, + current_request_ids: tuple[int, ...] = request_ids, + ): + if reservation_tasks.get(current_reservation_id) is task: + reservation_tasks.pop(current_reservation_id, None) + for request_id in current_request_ids: + if request_id_to_reservation_task.get(request_id) is task: + request_id_to_reservation_task.pop(request_id, None) + if not task.cancelled(): + error = task.exception() + if error is not None: + logger.error(f"Decode reservation task failed: {str(error)}") + + reservation_task.add_done_callback(remove_reservation_task) + elif obj[0] == ObjType.REQ: prompt, sampling_params, multimodal_params = obj[1] group_req_id = sampling_params.group_request_id + decode_admission_lease_handle = reserved_lease_handles.pop(group_req_id, None) pd_event = asyncio.Event() group_req_id_to_event[group_req_id] = pd_event generation_task = asyncio.create_task( @@ -134,6 +241,7 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O forwarding_queue=forwarding_queue, pd_upload_websocket=websocket, pd_event=pd_event, + decode_admission_lease_handle=decode_admission_lease_handle, ) ) generation_tasks[group_req_id] = generation_task @@ -146,6 +254,12 @@ def remove_generation_task(task: asyncio.Task, request_id: int = group_req_id): elif obj[0] == ObjType.ABORT: group_req_id = obj[1] logger.warning(f"recv cmd aborted req id {group_req_id}") + reservation_task = request_id_to_reservation_task.get(group_req_id) + if reservation_task is not None and not reservation_task.done(): + reservation_task.cancel() + reserved_handle = reserved_lease_handles.pop(group_req_id, None) + if reserved_handle is not None: + reserved_handle.release() generation_task = generation_tasks.get(group_req_id) if generation_task is not None and not generation_task.done(): generation_task.cancel() @@ -181,10 +295,14 @@ async def delayed_abort_task(group_req_id, retry_count): finally: child_tasks = [task for task in (forwarding_tokens_task, heartbeat_task) if task is not None] child_tasks.extend(generation_tasks.values()) + child_tasks.extend(reservation_tasks.values()) for task in child_tasks: task.cancel() if child_tasks: await asyncio.gather(*child_tasks, return_exceptions=True) + for handle in reserved_lease_handles.values(): + handle.release() + reserved_lease_handles.clear() await asyncio.sleep(10) await forwarding_queue.get_all_data() @@ -233,6 +351,7 @@ async def _pd_process_generate( forwarding_queue: AsyncQueue, pd_upload_websocket: ClientConnection, pd_event: asyncio.Event, + decode_admission_lease_handle: Optional[DecodeAdmissionLeaseHandle] = None, ): try: async for sub_req_id, request_output, metadata, finish_status in manager.generate( @@ -242,11 +361,19 @@ async def _pd_process_generate( request=None, pd_upload_websocket=pd_upload_websocket, pd_event=pd_event, + decode_admission_lease_handle=decode_admission_lease_handle, ): metadata["node_mode"] = manager.args.run_mode 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 ServerBusyError as e: + group_request_id = sampling_params.group_request_id + logger.warning(f"pd node rejected request {group_request_id}: {e.message}") + try: + await pd_upload_websocket.send(pickle.dumps((ObjType.PD_UPLOAD_SERVER_BUSY, group_request_id, e.message))) + except Exception: + logger.exception(f"report pd node request rejection failed, group_request_id: {group_request_id}") except asyncio.CancelledError: # PD master 主动 abort 或连接断开清理任务时会走取消路径,不需要反向重复上报。 pass @@ -261,10 +388,14 @@ async def _pd_process_generate( ) except Exception: logger.exception(f"report pd node generate error failed, group_request_id: {group_request_id}") + finally: + if decode_admission_lease_handle is not None: + decode_admission_lease_handle.release() # 转发token的task async def _up_tokens_to_pd_master(forwarding_queue: AsyncQueue, websocket: ClientConnection): + """批量向 PD Master 转发生成结果和最新负载。""" while True: handle_list = await forwarding_queue.wait_to_get_all_data() @@ -282,6 +413,7 @@ async def _send_heartbeat_to_pd_master(websocket: ClientConnection): # 获取节点负载信息 def _get_load_info() -> dict: + """汇总当前节点负载。""" from lightllm.server.api_http import g_objs diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index b422bf7703..ae2f18e93a 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -12,7 +12,13 @@ asyncio.set_event_loop_policy(uvloop.EventLoopPolicy()) from typing import Union, List, Tuple, Dict, Optional from lightllm.server.core.objs import FinishStatus -from ..pd_io_struct import PD_Client_Obj, PDUpKVStatus, ObjType, PDDecodeNodeInfo +from ..pd_io_struct import ( + PD_DECODE_ADMISSION_CAPABILITY_KEY, + PD_Client_Obj, + PDUpKVStatus, + ObjType, + PDDecodeNodeInfo, +) from lightllm.server.core.objs import SamplingParams, StartArgs from ..multimodal_params import MultimodalParams from ..tokenizer import get_tokenizer @@ -30,6 +36,60 @@ logger = init_logger(__name__) +class _DecodeReservationStatus: + def __init__(self) -> None: + self.event = asyncio.Event() + self.error_info: Optional[str] = None + + def raise_if_error(self) -> None: + if self.error_info is not None: + raise ServerBusyError(self.error_info) + + +class _DecodeAdmissionGroup: + """Wait for every choice's Prefill before reserving Decode slots once.""" + + def __init__( + self, + manager: "HttpServerManagerForPDMaster", + d_node: PD_Client_Obj, + reservation_id: int, + expected_count: int, + request: Request, + ) -> None: + self.manager = manager + self.d_node = d_node + self.reservation_id = reservation_id + self.expected_count = expected_count + self.request = request + self._lock = asyncio.Lock() + self._ready = asyncio.Event() + self._request_ids: List[int] = [] + self._error: Optional[BaseException] = None + + async def wait(self, group_request_id: int) -> None: + is_last = False + async with self._lock: + if group_request_id in self._request_ids: + raise RuntimeError(f"duplicate Decode reservation request id {group_request_id}") + self._request_ids.append(group_request_id) + is_last = len(self._request_ids) == self.expected_count + + if is_last: + try: + await self.manager.reserve_decode_slots( + self.d_node, self.reservation_id, tuple(self._request_ids), self.request + ) + except BaseException as error: + self._error = error + finally: + self._ready.set() + + await self._ready.wait() + if self._error is not None: + raise self._error + + class HttpServerManagerForPDMaster: def __init__( self, @@ -44,6 +104,7 @@ def __init__( self.pd_manager = PDManager(args) self.req_id_to_out_inf: Dict[int, ReqStatus] = {} + self.decode_reservation_statuses: Dict[int, _DecodeReservationStatus] = {} self.infos_queues = None # 这个需要延迟初始化,否则使用的loop不对 self.health_timeout = int(os.getenv("HEALTH_TIMEOUT", "200")) self.latest_success_infer_time = time.time() @@ -129,11 +190,6 @@ async def generate( multimodal_params: MultimodalParams, request: Request, ): - if not self.args.disable_pd_master_decode_capacity_limit: - decode_capacity = sum(node.start_args["running_max_req_size"] for node in self.pd_manager.decode_nodes) - if self.running_request_count >= decode_capacity: - raise ServerBusyError() - was_idle = self.running_request_count == 0 self.running_request_count += 1 if was_idle: @@ -168,14 +224,33 @@ async def _generate( origin_sampling_params = SamplingParams.from_buffer_copy(sampling_params) origin_group_request_id = self.id_gen.generate_id() - # Record one user request even when it is expanded into multiple independent - # n=1 requests below. The externally visible ids remain in the same request - # group so OpenAI streaming can derive choice_index from the sub request id. + # Record one user request even when it is expanded into independent n=1 + # requests. The Decode reservation coordinator still accounts for n as + # one atomic unit after all Prefill stages are ready. await self._log_req_header(request, origin_group_request_id) self.metric_client.counter_inc("lightllm_request_count") self.metric_client.histogram_observe("lightllm_request_max_new_tokens", origin_sampling_params.max_new_tokens) - choice_count = origin_sampling_params.n + p_node, d_node = await self.select_p_d_node(prompt, origin_sampling_params, multimodal_params) + if not p_node or not d_node: + logger.error(f"{origin_group_request_id}: No p_node or d_node found") + raise Exception(f"{origin_group_request_id}: No p_node or d_node found") + + choice_count = max(1, int(origin_sampling_params.n or 1)) + decode_reservation = None + args = getattr(self, "args", None) + decode_admission_disabled = args is not None and ( + args.disable_pd_node_decode_admission or args.disable_pd_master_decode_capacity_limit + ) + if choice_count > 1 and not decode_admission_disabled: + decode_reservation = _DecodeAdmissionGroup( + manager=self, + d_node=d_node, + reservation_id=origin_group_request_id, + expected_count=choice_count, + request=request, + ) + generators = [] for choice_index in range(choice_count): choice_sampling_params = SamplingParams.from_buffer_copy(origin_sampling_params) @@ -189,6 +264,9 @@ async def _generate( request, start_time, origin_group_request_id + choice_index, + p_node=p_node, + d_node=d_node, + decode_reservation=decode_reservation, ) ) @@ -205,25 +283,27 @@ async def _generate_one( request: Request, start_time: float, origin_request_id: int, + p_node: PD_Client_Obj, + d_node: PD_Client_Obj, + decode_reservation: Optional["_DecodeAdmissionGroup"] = None, ): # 先将请求根据max_new_tokens 参数进行分块操作,主要是 pd 分离场景中, # 只能使用保守调度,但是如果用户都设置一个很大的 max_new_tokens 值,会 # 导致极大显存预留,照成系统的吞吐能力下降,所以我们将请求分割成几段进行 # 推理,只要保证分块合理,实际分段推理是极少发生的情况,系统吞吐就不会受 # 到影响。 - max_new_tokens_list = self._split_max_new_tokens(max_new_tokens=origin_sampling_params.max_new_tokens) + if decode_reservation is None: + max_new_tokens_list = self._split_max_new_tokens(max_new_tokens=origin_sampling_params.max_new_tokens) + else: + # Choices have independent continuation histories, so the existing + # segmented continuation protocol cannot safely splice an n-way + # request. Each choice remains one n=1 request. + max_new_tokens_list = [origin_sampling_params.max_new_tokens] block_group_request_id = origin_request_id - p_node = None - d_node = None pending_prefill_load_chars = None try: - p_node, d_node = await self.select_p_d_node(prompt, origin_sampling_params, multimodal_params) - if not p_node or not d_node: - logger.error(f"{origin_request_id}: No p_node or d_node found") - raise Exception(f"{origin_request_id}: No p_node or d_node found") - history_gen_token_strs = [] origin_prompt_cache_len = None @@ -248,11 +328,11 @@ async def _generate_one( sampling_params, multimodal_params, request, + decode_reservation=decode_reservation, ) is_last_block = iter_index == len(max_new_tokens_list) - 1 prompt_tokens = sys.maxsize # 因为分段的原因 async for sub_req_id, request_output, metadata, finish_status in results_generator: - # pd 分离模式下,返回的 metadata 可能序号信息可能存在不准确性。 assert sub_req_id == block_group_request_id if finish_status.is_finished_length() and not is_last_block: finish_status = FinishStatus() # 转换为NoFinished @@ -359,6 +439,24 @@ async def raise_if_disconnected() -> None: except asyncio.TimeoutError: continue + async def reserve_decode_slots( + self, d_node: PD_Client_Obj, reservation_id: int, request_ids: Tuple[int, ...], request: Request + ) -> None: + status = _DecodeReservationStatus() + self.decode_reservation_statuses[reservation_id] = status + try: + await d_node.websocket.send_bytes( + pickle.dumps((ObjType.PD_RESERVE_DECODE_SLOTS, reservation_id, request_ids)) + ) + timeout = float(d_node.start_args.get("pd_node_decode_admission_timeout", 5)) + 5 + await self._wait_for_event_or_disconnect( + status.event, request, timeout=timeout, group_request_id=reservation_id, stage="decode admission" + ) + status.raise_if_error() + finally: + if self.decode_reservation_statuses.get(reservation_id) is status: + self.decode_reservation_statuses.pop(reservation_id, None) + async def _log_req_header(self, request: Request, group_request_id: int): x_request_id = request.headers.get("X-Request-Id", "") x_session_id = request.headers.get("X-Session-Id", "") @@ -378,6 +476,7 @@ async def fetch_pd_stream( sampling_params: SamplingParams, multimodal_params: MultimodalParams, request: Request, + decode_reservation: Optional["_DecodeAdmissionGroup"] = None, ): group_request_id = sampling_params.group_request_id sampling_params.pd_master_node_id.initialize(self.args.pd_node_id) @@ -408,6 +507,9 @@ async def fetch_pd_stream( prompt_ids = prefill_prompt_ids_event.prompt_ids logger.info(f"group_request_id: {group_request_id} get prefill prompt ids len {len(prompt_ids)}") + if decode_reservation is not None: + await decode_reservation.wait(group_request_id) + sampling_params.max_new_tokens = old_max_new_tokens await d_node.websocket.send_bytes( pickle.dumps((ObjType.REQ, (prompt_ids, sampling_params, MultimodalParams()))) @@ -515,6 +617,7 @@ async def _wait_to_token_package( sampling_params: SamplingParams, multimodal_params: MultimodalParams, request: Request, + decode_reservation: Optional["_DecodeAdmissionGroup"] = None, ): if sampling_params.disable_prompt_cache: assert False, "pd mode dont support set disable_prompt_cache to True" @@ -529,7 +632,13 @@ async def _wait_to_token_package( sub_req_id_to_mtp_verify_step_num: Dict[int, int] = {} async for sub_req_id, out_str, metadata, finish_status in self.fetch_pd_stream( - p_node, d_node, prompt, sampling_params, multimodal_params, request + p_node, + d_node, + prompt, + sampling_params, + multimodal_params, + request, + decode_reservation=decode_reservation, ): if await request.is_disconnected(): raise ClientDisconnected( @@ -677,6 +786,28 @@ async def handle_loop(self): ) else: await req_status.set_error(error_info) + elif obj[0] == ObjType.PD_UPLOAD_SERVER_BUSY: + _, group_req_id, error_info = obj + logger.warning( + f"received PD node server busy, group_req_id: {group_req_id}, reason: {error_info}" + ) + reservation_status = getattr(self, "decode_reservation_statuses", {}).get(group_req_id) + if reservation_status is not None: + reservation_status.error_info = error_info + reservation_status.event.set() + continue + req_status = self.req_id_to_out_inf.get(group_req_id) + if req_status is None: + logger.error(f"PD_UPLOAD_SERVER_BUSY fail find req status for group_req_id: {group_req_id}") + else: + await req_status.set_error(error_info, is_server_busy=True) + elif obj[0] == ObjType.PD_DECODE_SLOTS_RESERVED: + _, reservation_id = obj + reservation_status = getattr(self, "decode_reservation_statuses", {}).get(reservation_id) + if reservation_status is None: + logger.warning(f"received stale Decode reservation ack: {reservation_id}") + else: + reservation_status.event.set() else: logger.error(f"recevie error obj {obj}") except BaseException as e: @@ -703,6 +834,7 @@ def __init__(self, req_id, p_node, d_node) -> None: self.p_node: PD_Client_Obj = p_node self.d_node: PD_Client_Obj = d_node self.error_info: Optional[str] = None + self.is_server_busy = False async def wait_to_ready(self): try: @@ -710,9 +842,10 @@ async def wait_to_ready(self): except asyncio.TimeoutError: pass - async def set_error(self, error_info: str): + async def set_error(self, error_info: str, is_server_busy: bool = False): async with self.lock: self.error_info = error_info + self.is_server_busy = is_server_busy # 请求可能正在等待 Prefill prompt ids、Decode KV 资源或输出 token, # 设置全部事件,让请求自己的执行循环立即醒来并抛出异常。 self.event.set() @@ -721,6 +854,8 @@ async def set_error(self, error_info: str): def raise_if_error(self): if self.error_info is not None: + if self.is_server_busy: + raise ServerBusyError(self.error_info) logger.error( f"group_request_id: {self.req_id} detected PD node generate error, " f"raise exception to end the request flow early: {self.error_info}" @@ -813,7 +948,24 @@ async def check_pd_nodes_health(self): return True def register_pd(self, pd_info_json, websocket): - pd_client = PD_Client_Obj(**pd_info_json) + # Ignore the short-lived top-level capacity fields during a rolling + # upgrade from the former Master-side admission implementation. + pd_info = dict(pd_info_json) + pd_info.pop("capacity_share", None) + pd_info.pop("capacity_epoch", None) + start_args = pd_info.get("start_args") or {} + admission_disabled = ( + self.args.disable_pd_node_decode_admission or self.args.disable_pd_master_decode_capacity_limit + ) + if ( + pd_info.get("mode") == "decode" + and not admission_disabled + and start_args.get(PD_DECODE_ADMISSION_CAPABILITY_KEY) is not True + ): + raise ValueError( + "Decode node does not provide authoritative admission; upgrade Decode nodes before PD Masters" + ) + pd_client = PD_Client_Obj(**pd_info) 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( @@ -853,7 +1005,10 @@ def register_pd(self, pd_info_json, websocket): return def remove_pd(self, pd_info_json): - pd_client = PD_Client_Obj(**pd_info_json) + pd_info = dict(pd_info_json) + pd_info.pop("capacity_share", None) + pd_info.pop("capacity_epoch", None) + pd_client = PD_Client_Obj(**pd_info) 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] diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index 8a1e2bd42b..30ba92ca27 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -10,6 +10,8 @@ logger = init_logger(__name__) +PD_DECODE_ADMISSION_CAPABILITY_KEY = "__pd_decode_admission_v1" + # 节点的行为 class NodeRole(enum.Enum): @@ -42,6 +44,9 @@ class ObjType(enum.Enum): PD_REQ_DECODE_NODE_INFO = 5 # pd master 节点下发给 prefill 节点的请求对应的 decode 节点信息。 HEARTBEAT = 6 # P/D 节点向 pd master 上报的心跳。 PD_UPLOAD_GENERATE_ERROR = 7 # P/D 节点向 pd master 上报本地请求生成异常。 + PD_UPLOAD_SERVER_BUSY = 8 # P/D 节点向 pd master 上报本地准入拒绝。 + PD_RESERVE_DECODE_SLOTS = 9 # PD master 在 Decode 节点原子预留一组请求槽。 + PD_DECODE_SLOTS_RESERVED = 10 # Decode 节点确认一组请求槽已预留。 @dataclass @@ -92,7 +97,6 @@ class PDUpKVStatus: pd_kv_trans_params: bytes # pd kv 传输建立连接所使用的元数据对象 def __post_init__(self): - if not isinstance(self.group_request_id, int): error_info = "group_request_id only can be int" logger.error(error_info) diff --git a/test/test_pd_selector/test_pd_master_multi_choice.py b/test/test_pd_selector/test_pd_master_multi_choice.py index 7d0b32ded9..12b52e787d 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -1,6 +1,8 @@ from __future__ import annotations import asyncio +import pickle +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, call, patch import pytest @@ -8,20 +10,29 @@ from lightllm.server.core.objs import FinishStatus, SamplingParams from lightllm.server.httpserver.manager import HttpServerManager from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster +from lightllm.server.multimodal_params import MultimodalParams +from lightllm.server.pd_io_struct import ObjType def _manager() -> HttpServerManagerForPDMaster: manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) manager.pd_manager = MagicMock() + manager.args = SimpleNamespace( + disable_pd_node_decode_admission=False, + disable_pd_master_decode_capacity_limit=False, + pd_node_id=7, + ) manager.id_gen = MagicMock() manager.id_gen.generate_id.return_value = 800 manager.metric_client = MagicMock() manager._log_req_header = AsyncMock() manager.tokens = MagicMock(return_value=2) + manager.req_id_to_out_inf = {} + manager.decode_reservation_statuses = {} return manager -def test_pd_master_expands_n_into_concurrent_single_choice_requests(): +def test_pd_master_reserves_n_then_sends_independent_choices_to_one_decode_node(): async def run(): manager = _manager() manager.id_gen.generate_id.side_effect = [800, 808, 816, 824] @@ -33,38 +44,34 @@ async def run(): multimodal_params = MagicMock() multimodal_params.verify_and_preload = AsyncMock() request = MagicMock() - - started = set() - all_started = asyncio.Event() - captured_params = [] p_node = MagicMock(dispatched_prompt_chars=0, dispatched_req_num=0) d_node = MagicMock() manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) - manager._split_max_new_tokens = MagicMock(return_value=[4]) + manager._split_max_new_tokens = MagicMock() manager.remove_req = AsyncMock() + manager.reserve_decode_slots = AsyncMock() + captured_params = [] + captured_nodes = [] async def wait_to_token_package( selected_p_node, selected_d_node, start_time, prompt, - choice_sampling_params, + node_sampling_params, child_multimodal_params, child_request, + decode_reservation=None, ): - captured_params.append(choice_sampling_params) - started.add(choice_sampling_params.group_request_id) - if len(started) == 3: - all_started.set() - await all_started.wait() - for token_index in range(2): - finish_status = FinishStatus() if token_index == 0 else FinishStatus(FinishStatus.FINISHED_STOP) - yield ( - choice_sampling_params.group_request_id, - f"internal-{choice_sampling_params.group_request_id}-{token_index}", - {"prompt_tokens": 2}, - finish_status, - ) + captured_params.append(node_sampling_params) + captured_nodes.append((selected_p_node, selected_d_node)) + await decode_reservation.wait(node_sampling_params.group_request_id) + yield ( + node_sampling_params.group_request_id, + f"choice-{node_sampling_params.group_request_id}", + {"prompt_tokens": 2}, + FinishStatus(FinishStatus.FINISHED_STOP), + ) manager._wait_to_token_package = wait_to_token_package @@ -73,21 +80,21 @@ async def wait_to_token_package( async for result in manager._generate("prompt", sampling_params, multimodal_params, request): results.append(result) - assert [result[0] for result in results].count(800) == 2 - assert [result[0] for result in results].count(801) == 2 - assert [result[0] for result in results].count(802) == 2 - assert {result[1] for result in results} == { - "internal-808-0", - "internal-808-1", - "internal-816-0", - "internal-816-1", - "internal-824-0", - "internal-824-1", - } - assert all(params.n == 1 and params.best_of == 1 for params in captured_params) + assert {result[0] for result in results} == {800, 801, 802} + assert {result[1] for result in results} == {"choice-808", "choice-816", "choice-824"} + assert len(captured_params) == 3 + assert all(params.n == 1 for params in captured_params) + assert all(params.best_of == 1 for params in captured_params) assert {params.group_request_id for params in captured_params} == {808, 816, 824} - assert sampling_params.n == 3 - assert sampling_params.best_of == 3 + assert captured_nodes == [(p_node, d_node)] * 3 + manager.select_p_d_node.assert_awaited_once() + manager._split_max_new_tokens.assert_not_called() + manager.reserve_decode_slots.assert_awaited_once() + reserve_args = manager.reserve_decode_slots.await_args.args + assert reserve_args[0] is d_node + assert reserve_args[1] == 800 + assert set(reserve_args[2]) == {808, 816, 824} + assert reserve_args[3] is request multimodal_params.verify_and_preload.assert_awaited_once_with(request) manager._log_req_header.assert_awaited_once_with(request, 800) assert manager.metric_client.counter_inc.call_args_list == [ @@ -100,6 +107,57 @@ async def wait_to_token_package( asyncio.run(asyncio.wait_for(run(), timeout=2)) +def test_pd_master_waits_for_prefill_and_atomic_reservation_before_decode_request(): + async def run(): + manager = _manager() + sampling_params = SamplingParams() + sampling_params.group_request_id = 808 + sampling_params.max_new_tokens = 4 + p_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) + d_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) + request = SimpleNamespace(is_disconnected=AsyncMock(return_value=False)) + reservation_entered = asyncio.Event() + reservation_release = asyncio.Event() + + async def wait_for_reservation(group_request_id): + assert group_request_id == 808 + reservation_entered.set() + await reservation_release.wait() + + decode_reservation = SimpleNamespace(wait=wait_for_reservation) + generator = manager.fetch_pd_stream( + p_node, + d_node, + "prompt", + sampling_params, + MultimodalParams(), + request, + decode_reservation=decode_reservation, + ) + next_result = asyncio.create_task(generator.__anext__()) + try: + while 808 not in manager.req_id_to_out_inf: + await asyncio.sleep(0) + req_status = manager.req_id_to_out_inf[808] + req_status.prefill_prompt_ids_event.prompt_ids = [1, 2, 3] + req_status.prefill_prompt_ids_event.set() + + await asyncio.wait_for(reservation_entered.wait(), timeout=1) + d_node.websocket.send_bytes.assert_not_awaited() + + reservation_release.set() + while d_node.websocket.send_bytes.await_count == 0: + await asyncio.sleep(0) + decode_message = pickle.loads(d_node.websocket.send_bytes.await_args.args[0]) + assert decode_message[0] == ObjType.REQ + assert list(decode_message[1][0]) == [1, 2, 3] + finally: + next_result.cancel() + await asyncio.gather(next_result, return_exceptions=True) + + asyncio.run(asyncio.wait_for(run(), timeout=2)) + + def test_pd_master_n_one_uses_the_same_choice_merge_path(): async def run(): manager = _manager() @@ -111,6 +169,9 @@ async def run(): multimodal_params = MagicMock() multimodal_params.verify_and_preload = AsyncMock() request = MagicMock() + p_node = MagicMock() + d_node = MagicMock() + manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) merge_choice_generators = manager._merge_choice_generators manager._merge_choice_generators = MagicMock(side_effect=merge_choice_generators) @@ -121,6 +182,7 @@ async def generate_one( child_request, start_time, origin_request_id, + **_kwargs, ): yield ( origin_request_id, @@ -217,7 +279,6 @@ async def run(): manager.abort = AsyncMock() p_node = MagicMock(dispatched_prompt_chars=0, dispatched_req_num=0) d_node = MagicMock() - manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) async def failing_wait_to_token_package(*_args, **_kwargs): raise RuntimeError("generation failed") @@ -234,6 +295,8 @@ async def failing_wait_to_token_package(*_args, **_kwargs): MagicMock(), 0, 800, + p_node, + d_node, ): pass @@ -257,13 +320,14 @@ async def run(): dispatched_req_num=other_request_count, ) d_node = MagicMock() - manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) dispatched_nodes = [] dispatched_prompts = [] dispatched_loads = [] dispatched_req_counts = [] - async def wait_to_token_package(selected_p_node, _d_node, _start_time, block_prompt, sampling_params, *_args): + async def wait_to_token_package( + selected_p_node, _d_node, _start_time, block_prompt, sampling_params, *_args, **_kwargs + ): dispatched_nodes.append(selected_p_node) dispatched_prompts.append(block_prompt) dispatched_loads.append(selected_p_node.dispatched_prompt_chars) @@ -285,10 +349,11 @@ async def wait_to_token_package(selected_p_node, _d_node, _start_time, block_pro MagicMock(), 0, 800, + p_node, + d_node, ): results.append(result) - manager.select_p_d_node.assert_awaited_once() assert dispatched_nodes == [p_node, p_node] assert dispatched_prompts == ["prompt", "promptx"] assert dispatched_loads == [other_request_load + len("prompt"), other_request_load + len("promptx")] @@ -314,7 +379,6 @@ async def run(): dispatched_req_num=other_request_count, ) d_node = MagicMock() - manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) async def wait_to_token_package(*_args, **_kwargs): yield 808, "first", {"prompt_tokens": 1}, FinishStatus() @@ -328,6 +392,8 @@ async def wait_to_token_package(*_args, **_kwargs): MagicMock(), 0, 800, + p_node, + d_node, ) assert (await generator.__anext__())[1] == "first" diff --git a/unit_tests/server/core/objs/test_shm_req_manager.py b/unit_tests/server/core/objs/test_shm_req_manager.py index 56871b5466..59fba2c05f 100644 --- a/unit_tests/server/core/objs/test_shm_req_manager.py +++ b/unit_tests/server/core/objs/test_shm_req_manager.py @@ -1,6 +1,7 @@ import os import pytest import time +import numpy as np from unittest.mock import MagicMock from easydict import EasyDict @@ -81,6 +82,25 @@ def test_put_back_req_obj(shm_req_manager): shm_req_manager.release_req_index(index) +def test_alloc_req_indexes_allocates_all_slots_together(shm_req_manager): + indexes = shm_req_manager.alloc_req_indexes(3) + + assert indexes is not None + assert len(indexes) == 3 + assert all(shm_req_manager.alloc_state_shm.arr[index] == 1 for index in indexes) + + for index in indexes: + shm_req_manager.release_req_index(index) + + +def test_alloc_req_indexes_rolls_back_when_full_batch_is_unavailable(shm_req_manager): + before = shm_req_manager.alloc_state_shm.arr.copy() + free_slots = int(np.sum(before == 0)) + + assert shm_req_manager.alloc_req_indexes(free_slots + 1) is None + np.testing.assert_array_equal(shm_req_manager.alloc_state_shm.arr, before) + + def test_alloc_req_index_no_available(shm_req_manager): for _ in range(shm_req_manager.max_req_num): shm_req_manager.alloc_req_index() diff --git a/unit_tests/server/httpserver/test_decode_admission.py b/unit_tests/server/httpserver/test_decode_admission.py new file mode 100644 index 0000000000..6883e88192 --- /dev/null +++ b/unit_tests/server/httpserver/test_decode_admission.py @@ -0,0 +1,169 @@ +import asyncio + +import pytest + +from lightllm.server.httpserver.decode_admission import ( + DecodeAdmissionController, + DecodeAdmissionLeaseHandle, +) +from lightllm.utils.error_utils import ServerBusyError + + +def test_decode_admission_waits_and_grants_in_fifo_order(): + async def run(): + controller = DecodeAdmissionController(capacity=2, max_queued_slots=2, timeout_seconds=1) + active = await controller.acquire(2) + first_task = asyncio.create_task(controller.acquire(1)) + second_task = asyncio.create_task(controller.acquire(1)) + await asyncio.sleep(0) + + assert controller.active_slots == 2 + assert controller.queued_slots == 2 + assert not first_task.done() + assert not second_task.done() + + active.release() + first = await first_task + second = await second_task + assert controller.active_slots == 2 + + first.release() + second.release() + assert controller.active_slots == 0 + + asyncio.run(run()) + + +def test_decode_admission_keeps_strict_fifo_when_head_gang_does_not_fit(): + async def run(): + controller = DecodeAdmissionController(capacity=3, max_queued_slots=3, timeout_seconds=1) + first_active = await controller.acquire(1) + second_active = await controller.acquire(1) + third_active = await controller.acquire(1) + gang_task = asyncio.create_task(controller.acquire(2)) + await asyncio.sleep(0) + small_task = asyncio.create_task(controller.acquire(1)) + await asyncio.sleep(0) + + first_active.release() + await asyncio.sleep(0) + assert not gang_task.done() + assert not small_task.done() + assert controller.active_slots == 2 + + second_active.release() + gang = await gang_task + assert not small_task.done() + third_active.release() + small = await small_task + assert controller.active_slots == 3 + + gang.release() + small.release() + + asyncio.run(run()) + + +def test_decode_admission_cancellation_removes_waiter(): + async def run(): + controller = DecodeAdmissionController(capacity=1, max_queued_slots=1, timeout_seconds=1) + active = await controller.acquire(1) + waiting_task = asyncio.create_task(controller.acquire(1)) + await asyncio.sleep(0) + + waiting_task.cancel() + with pytest.raises(asyncio.CancelledError): + await waiting_task + + assert controller.queued_request_count == 0 + assert controller.queued_slots == 0 + active.release() + assert controller.active_slots == 0 + + asyncio.run(run()) + + +def test_decode_admission_timeout_removes_waiter(): + async def run(): + controller = DecodeAdmissionController(capacity=1, max_queued_slots=1, timeout_seconds=0.01) + active = await controller.acquire(1) + + with pytest.raises(ServerBusyError, match="wait timed out"): + await controller.acquire(1) + + assert controller.queued_request_count == 0 + active.release() + + asyncio.run(run()) + + +def test_decode_admission_rejects_full_queue_and_oversized_gang(): + async def run(): + controller = DecodeAdmissionController(capacity=2, max_queued_slots=1, timeout_seconds=1) + active = await controller.acquire(2) + waiting_task = asyncio.create_task(controller.acquire(1)) + await asyncio.sleep(0) + + with pytest.raises(ServerBusyError, match="queue is full"): + await controller.acquire(1) + with pytest.raises(ServerBusyError, match="needs 3 Decode slots"): + await controller.acquire(3) + + waiting_task.cancel() + await asyncio.gather(waiting_task, return_exceptions=True) + active.release() + + asyncio.run(run()) + + +def test_decode_admission_cancellation_after_concurrent_grant_releases_lease(): + async def run(): + controller = DecodeAdmissionController(capacity=1, max_queued_slots=1, timeout_seconds=1) + active = await controller.acquire(1) + waiting_task = asyncio.create_task(controller.acquire(1)) + await asyncio.sleep(0) + + active.release() + waiting_task.cancel() + await asyncio.gather(waiting_task, return_exceptions=True) + + assert controller.active_slots == 0 + assert controller.queued_request_count == 0 + + asyncio.run(run()) + + +def test_decode_admission_gang_lease_can_transfer_to_independent_requests(): + async def run(): + controller = DecodeAdmissionController(capacity=3, max_queued_slots=0, timeout_seconds=1) + gang = await controller.acquire(3) + first, second, third = gang.split([1, 1, 1]) + + first.release() + assert controller.active_slots == 2 + second.release() + assert controller.active_slots == 1 + third.release() + assert controller.active_slots == 0 + + with pytest.raises(RuntimeError, match="already been released"): + gang.split([1, 1, 1]) + + asyncio.run(run()) + + +def test_decode_admission_handle_transfers_or_releases_lease_once(): + async def run(): + controller = DecodeAdmissionController(capacity=2, max_queued_slots=0, timeout_seconds=1) + transferred_handle = DecodeAdmissionLeaseHandle(await controller.acquire(1)) + lease = transferred_handle.take() + transferred_handle.release() + assert controller.active_slots == 1 + lease.release() + + abandoned_handle = DecodeAdmissionLeaseHandle(await controller.acquire(1)) + abandoned_handle.release() + abandoned_handle.release() + assert controller.active_slots == 0 + + asyncio.run(run()) diff --git a/unit_tests/server/httpserver/test_pd_generate_error.py b/unit_tests/server/httpserver/test_pd_generate_error.py index 3c7a4242a0..08ea61d1a7 100644 --- a/unit_tests/server/httpserver/test_pd_generate_error.py +++ b/unit_tests/server/httpserver/test_pd_generate_error.py @@ -7,13 +7,14 @@ import pytest from lightllm.server.core.objs import FinishStatus, SamplingParams -from lightllm.server.httpserver.pd_loop import _pd_process_generate +from lightllm.server.httpserver.decode_admission import DecodeAdmissionController +from lightllm.server.httpserver.pd_loop import _pd_process_generate, _reserve_decode_slots from lightllm.server.httpserver_for_pd_master.manager import ( HttpServerManagerForPDMaster, ReqStatus, ) from lightllm.server.pd_io_struct import ObjType -from lightllm.utils.error_utils import PDPrefillNodeStopGenToken +from lightllm.utils.error_utils import PDPrefillNodeStopGenToken, ServerBusyError class _FailingManager: @@ -32,6 +33,14 @@ async def generate(self, **_kwargs): yield +class _BusyManager: + args = SimpleNamespace(run_mode="decode") + + async def generate(self, **_kwargs): + raise ServerBusyError("decode queue is full") + yield + + class _SuccessfulManager: args = SimpleNamespace(run_mode="prefill") @@ -107,6 +116,28 @@ async def run(): asyncio.run(run()) +def test_pd_node_reports_server_busy_separately(): + async def run(): + sampling_params = SamplingParams() + sampling_params.group_request_id = 123 + websocket = AsyncMock() + + await _pd_process_generate( + manager=_BusyManager(), + prompt="prompt", + sampling_params=sampling_params, + multimodal_params=MagicMock(), + forwarding_queue=MagicMock(), + pd_upload_websocket=websocket, + pd_event=asyncio.Event(), + ) + + obj = pickle.loads(websocket.send.await_args.args[0]) + assert obj == (ObjType.PD_UPLOAD_SERVER_BUSY, 123, "decode queue is full") + + asyncio.run(run()) + + def test_pd_node_success_forwards_token_without_error_report(): async def run(): sampling_params = SamplingParams() @@ -204,6 +235,100 @@ async def run(): asyncio.run(run()) +def test_pd_decode_node_reserves_n_slots_and_hands_out_independent_leases(): + async def run(): + manager = SimpleNamespace( + decode_admission_controller=DecodeAdmissionController(capacity=3, max_queued_slots=0, timeout_seconds=1), + metric_client=MagicMock(), + ) + websocket = AsyncMock() + handles = {} + + await _reserve_decode_slots(manager, 100, (11, 22, 33), handles, websocket) + + assert manager.decode_admission_controller.active_slots == 3 + assert set(handles) == {11, 22, 33} + assert pickle.loads(websocket.send.await_args.args[0]) == (ObjType.PD_DECODE_SLOTS_RESERVED, 100) + + first = handles.pop(11).take() + first.release() + handles.pop(22).release() + handles.pop(33).release() + assert manager.decode_admission_controller.active_slots == 0 + + asyncio.run(run()) + + +def test_pd_decode_node_cancels_queued_atomic_reservation_without_leaking(): + async def run(): + controller = DecodeAdmissionController(capacity=2, max_queued_slots=2, timeout_seconds=10) + active = await controller.acquire(2) + manager = SimpleNamespace(decode_admission_controller=controller, metric_client=MagicMock()) + websocket = AsyncMock() + handles = {} + task = asyncio.create_task(_reserve_decode_slots(manager, 100, (11, 22), handles, websocket)) + await asyncio.sleep(0) + + task.cancel() + with suppress(asyncio.CancelledError): + await task + + assert controller.queued_slots == 0 + assert handles == {} + websocket.send.assert_not_awaited() + active.release() + assert controller.active_slots == 0 + + asyncio.run(run()) + + +def test_pd_master_routes_decode_reservation_ack_and_busy_response(): + async def run(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = SimpleNamespace(config_server_host=None, config_server_port=None) + manager.pd_manager = MagicMock() + manager.timer_log = AsyncMock() + manager.infos_queues = None + manager.req_id_to_out_inf = {} + manager.decode_reservation_statuses = {} + d_node = SimpleNamespace( + start_args={"pd_node_decode_admission_timeout": 1}, + websocket=SimpleNamespace(send_bytes=AsyncMock()), + ) + request = SimpleNamespace(is_disconnected=AsyncMock(return_value=False)) + + handle_task = asyncio.create_task(manager.handle_loop()) + try: + while manager.infos_queues is None: + await asyncio.sleep(0) + + accepted = asyncio.create_task(manager.reserve_decode_slots(d_node, 100, (11, 22), request)) + while 100 not in manager.decode_reservation_statuses: + await asyncio.sleep(0) + await manager.put_to_handle_queue((ObjType.PD_DECODE_SLOTS_RESERVED, 100)) + await asyncio.wait_for(accepted, timeout=1) + + rejected = asyncio.create_task(manager.reserve_decode_slots(d_node, 200, (33, 44), request)) + while 200 not in manager.decode_reservation_statuses: + await asyncio.sleep(0) + await manager.put_to_handle_queue((ObjType.PD_UPLOAD_SERVER_BUSY, 200, "decode queue is full")) + with pytest.raises(ServerBusyError, match="decode queue is full"): + await asyncio.wait_for(rejected, timeout=1) + + sent = [pickle.loads(call.args[0]) for call in d_node.websocket.send_bytes.await_args_list] + assert sent == [ + (ObjType.PD_RESERVE_DECODE_SLOTS, 100, (11, 22)), + (ObjType.PD_RESERVE_DECODE_SLOTS, 200, (33, 44)), + ] + assert manager.decode_reservation_statuses == {} + finally: + handle_task.cancel() + with suppress(asyncio.CancelledError): + await handle_task + + asyncio.run(run()) + + def test_pd_master_generate_error_marks_request_and_wakes_all_waiters(): async def run(): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) @@ -211,6 +336,7 @@ async def run(): manager.pd_manager = MagicMock() manager.timer_log = AsyncMock() manager.infos_queues = None + manager.decode_reservation_statuses = {} p_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) d_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) @@ -244,6 +370,33 @@ async def run(): asyncio.run(run()) +def test_pd_master_server_busy_wakes_waiters_and_preserves_429(): + async def run(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = SimpleNamespace(config_server_host=None) + manager.pd_manager = MagicMock() + manager.timer_log = AsyncMock() + manager.infos_queues = None + + req_status = ReqStatus(123, MagicMock(), MagicMock()) + manager.req_id_to_out_inf = {123: req_status} + handle_task = asyncio.create_task(manager.handle_loop()) + try: + while manager.infos_queues is None: + await asyncio.sleep(0) + await manager.put_to_handle_queue((ObjType.PD_UPLOAD_SERVER_BUSY, 123, "decode queue is full")) + await asyncio.wait_for(req_status.event.wait(), timeout=1) + + with pytest.raises(ServerBusyError, match="decode queue is full"): + req_status.raise_if_error() + finally: + handle_task.cancel() + with suppress(asyncio.CancelledError): + await handle_task + + asyncio.run(run()) + + @pytest.mark.parametrize("event_name", ["prefill_prompt_ids_event", "up_status_event"]) def test_pd_master_generate_error_wakes_resource_wait(event_name): async def run(): diff --git a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py index a3dc268a93..6635cdfcd6 100644 --- a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -39,7 +39,7 @@ def gen_id(): def _collect(mgr, sampling_params, monkeypatch, split): mgr._split_max_new_tokens = lambda *a, **k: list(split) - async def fake_wait(p_node, d_node, start_time, prompt, sp, multimodal_params, request): + async def fake_wait(p_node, d_node, start_time, prompt, sp, multimodal_params, request, decode_reservation=None): sub_req_id = sp.group_request_id hit = sp.max_new_tokens * 10 yield sub_req_id, "x", {"prompt_tokens": 100, "prompt_cache_len": hit}, FinishStatus() diff --git a/unit_tests/server/httpserver/test_running_request_lifecycle.py b/unit_tests/server/httpserver/test_running_request_lifecycle.py index 879012c2ab..82e1d29fa4 100644 --- a/unit_tests/server/httpserver/test_running_request_lifecycle.py +++ b/unit_tests/server/httpserver/test_running_request_lifecycle.py @@ -6,6 +6,7 @@ import pytest from lightllm.server.core.objs import SamplingParams +from lightllm.server.httpserver.decode_admission import DecodeAdmissionController from lightllm.server.httpserver.manager import HttpServerManager from lightllm.server.pd_io_struct import NodeRole, ObjType from lightllm.utils.error_utils import PDPrefillNodeStopGenToken @@ -38,7 +39,9 @@ def _make_manager(mode: NodeRole): manager._register_running_request = AsyncMock() manager._unregister_running_request = AsyncMock() manager.metric_client = MagicMock() - manager.shm_req_manager = SimpleNamespace(async_alloc_req_index=AsyncMock(side_effect=RuntimeError("alloc failed"))) + manager.shm_req_manager = SimpleNamespace( + async_alloc_req_indexes=AsyncMock(side_effect=RuntimeError("alloc failed")) + ) return manager @@ -187,6 +190,68 @@ async def run(): asyncio.run(run()) +def test_decode_admission_lease_lives_until_request_resources_are_recycled(): + async def run(): + manager = _make_manager(NodeRole.D) + manager.decode_admission_controller = DecodeAdmissionController( + capacity=1, max_queued_slots=1, timeout_seconds=1 + ) + manager.args = SimpleNamespace(chunked_prefill_size=None) + manager.tokenizer = object() + manager.enable_multimodal = False + manager.transfer_to_next_module_or_node = AsyncMock() + manager._count_multimodal_tokens = MagicMock(return_value=(0, 0)) + + req_obj = MagicMock(request_id=123, index_in_shm_mem=7) + manager.shm_req_manager = SimpleNamespace( + async_alloc_req_indexes=AsyncMock(return_value=[7]), + async_get_req_obj_by_index=AsyncMock(return_value=req_obj), + async_put_back_req_obj=AsyncMock(), + async_release_req_index=AsyncMock(), + ) + + async def wait_to_token_package(*_args, **_kwargs): + yield 123, "token", {}, MagicMock() + + manager._wait_to_token_package = wait_to_token_package + generator = manager.generate( + prompt=[10, 11, 12], + sampling_params=_sampling_params(), + multimodal_params=_multimodal_params(), + request=None, + ) + + assert (await generator.__anext__())[1] == "token" + req_status = manager.req_id_to_out_inf[123] + assert manager.decode_admission_controller.active_slots == 1 + + await generator.aclose() + assert manager.decode_admission_controller.active_slots == 1 + + req_status.release_decode_admission() + assert manager.decode_admission_controller.active_slots == 0 + + asyncio.run(run()) + + +def test_multinode_tp_slave_does_not_register_as_independent_pd_node(): + async def run(): + manager = HttpServerManager.__new__(HttpServerManager) + manager.recycle_resource_loop = AsyncMock() + manager.loop_for_request = AsyncMock() + manager.is_multinode_tp_slave = True + manager.pd_mode = NodeRole.D + manager.zmq_recv_socket = SimpleNamespace(recv_pyobj=AsyncMock(side_effect=asyncio.CancelledError)) + + with patch("lightllm.server.httpserver.pd_loop.pd_handle_loop") as pd_handle_loop: + with pytest.raises(asyncio.CancelledError): + await manager.handle_loop() + + pd_handle_loop.assert_not_called() + + asyncio.run(run()) + + def test_running_request_helpers_are_atomic_and_refresh_timestamp_only_on_idle_to_running_transition(): async def run(): manager = HttpServerManager.__new__(HttpServerManager) diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index edb758a94b..c748e47f7d 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -6,7 +6,9 @@ from easydict import EasyDict from lightllm.server.core.objs.start_args_type import StartArgs +from lightllm.server.httpserver.pd_loop import _build_pd_registration_info from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster, PDManager +from lightllm.server.pd_io_struct import PD_DECODE_ADMISSION_CAPABILITY_KEY, NodeRole def test_auto_set_response_parsers_from_qwen35_model_config(tmp_path): @@ -184,6 +186,72 @@ def test_pd_manager_without_connected_nodes_is_healthy(): assert asyncio.run(manager.check_pd_nodes_health()) is True +def test_pd_master_rejects_decode_node_without_authoritative_admission(): + args = StartArgs() + manager = PDManager(args) + + with pytest.raises(ValueError, match="upgrade Decode nodes before PD Masters"): + manager.register_pd( + { + "node_id": 2, + "client_ip_port": "10.0.0.2:8000", + "mode": "decode", + "start_args": {"max_req_total_len": args.max_req_total_len}, + }, + websocket=object(), + ) + + assert manager.decode_nodes == [] + + +def test_pd_master_accepts_decode_node_with_authoritative_admission(): + args = StartArgs() + manager = PDManager(args) + + manager.register_pd( + { + "node_id": 2, + "client_ip_port": "10.0.0.2:8000", + "mode": "decode", + "start_args": {"max_req_total_len": args.max_req_total_len, PD_DECODE_ADMISSION_CAPABILITY_KEY: True}, + }, + websocket=object(), + ) + + assert [node.client_ip_port for node in manager.decode_nodes] == ["10.0.0.2:8000"] + + +def test_pd_master_allows_legacy_decode_node_only_when_admission_is_disabled(): + args = StartArgs(disable_pd_node_decode_admission=True) + manager = PDManager(args) + + manager.register_pd( + { + "node_id": 2, + "client_ip_port": "10.0.0.2:8000", + "mode": "decode", + "start_args": {"max_req_total_len": args.max_req_total_len}, + }, + websocket=object(), + ) + + assert [node.client_ip_port for node in manager.decode_nodes] == ["10.0.0.2:8000"] + + +def test_decode_registration_advertises_authoritative_admission(monkeypatch): + from lightllm.server.httpserver import pd_loop + + args = StartArgs(run_mode="decode", host="0.0.0.0") + manager = SimpleNamespace(args=args, host_ip="10.0.0.2", pd_mode=NodeRole.D, decode_admission_controller=object()) + monkeypatch.setattr(pd_loop, "get_shm_port_args", lambda: SimpleNamespace(port=8001)) + + registration = _build_pd_registration_info(manager) + + assert registration["start_args"][PD_DECODE_ADMISSION_CAPABILITY_KEY] is True + assert registration["start_args"]["host"] == "10.0.0.2" + assert args.host == "0.0.0.0" + + def test_prefill_registration_preserves_existing_inflight_prompt_chars(): args = StartArgs() manager = PDManager(args)