diff --git a/lightllm/common/kv_cache_mem_manager/allocator.py b/lightllm/common/kv_cache_mem_manager/allocator.py index 850c158778..1331943b1a 100644 --- a/lightllm/common/kv_cache_mem_manager/allocator.py +++ b/lightllm/common/kv_cache_mem_manager/allocator.py @@ -1,7 +1,6 @@ import torch from lightllm.server.router.dynamic_prompt.shared_arr import SharedInt from lightllm.utils.dist_utils import get_current_rank_in_node -from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.log_utils import init_logger from typing import Union, List @@ -25,10 +24,8 @@ def __init__(self, size: int) -> None: self.can_use_mem_size = self.size rank_in_node = get_current_rank_in_node() - # 用共享内存进行共享,router 模块读取进行精确的调度估计, nccl port 作为一个单机中单实列的标记。防止冲突。 - self.shared_can_use_token_num = SharedInt( - f"{get_unique_server_name()}_mem_manger_can_use_token_num_{rank_in_node}" - ) + # 用共享内存进行共享,router 模块读取进行精确的调度估计;基础层会统一添加服务前缀以防止实例冲突。 + self.shared_can_use_token_num = SharedInt(f"mem_manger_can_use_token_num_{rank_in_node}") self.shared_can_use_token_num.set_value(self.can_use_mem_size) return diff --git a/lightllm/common/kv_cache_mem_manager/mem_manager.py b/lightllm/common/kv_cache_mem_manager/mem_manager.py index d217e05c78..9171566599 100755 --- a/lightllm/common/kv_cache_mem_manager/mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/mem_manager.py @@ -14,7 +14,8 @@ get_current_rank_in_node, get_node_world_size, ) -from lightllm.utils.envs_utils import get_unique_server_name, get_env_start_args +from lightllm.utils.envs_utils import get_env_start_args +from lightllm.utils.shm_utils import get_service_shm_name from lightllm.utils.config_utils import get_num_key_value_heads from lightllm.common.kv_trans_kernel.nixl_kv_trans import page_io from lightllm.utils.device_utils import kv_trans_use_p2p @@ -237,10 +238,10 @@ def write_to_shm(self, req_manager): # 避免过多无用的数据复制和传输开销。 self.req_to_token_indexs: torch.Tensor = req_manager.req_to_token_indexs - lock = FileLock(f"/tmp/{get_unique_server_name()}_mem_manager_lock") + lock = FileLock(f"/tmp/{get_service_shm_name('mem_manager_lock')}") with lock: node_world_size = get_node_world_size() - shm_name = f"{get_unique_server_name()}_mem_manager_{get_current_rank_in_node()}" + shm_name = f"mem_manager_{get_current_rank_in_node()}" obj_bytes_array = [ForkingPickler.dumps(self).tobytes() for _ in range(node_world_size * 2)] obj_size = len(obj_bytes_array[0]) shm = create_or_link_shm( @@ -256,8 +257,8 @@ def write_to_shm(self, req_manager): @staticmethod def loads_from_shm(rank_in_node: int) -> "MemoryManager": - shm_name = f"{get_unique_server_name()}_mem_manager_{rank_in_node}" - lock = FileLock(f"/tmp/{get_unique_server_name()}_mem_manager_lock") + shm_name = f"mem_manager_{rank_in_node}" + lock = FileLock(f"/tmp/{get_service_shm_name('mem_manager_lock')}") logger.info(f"get memmanager from shm {shm_name}") with lock: shm = create_or_link_shm(name=shm_name, expected_size=-1, force_mode="link") @@ -285,7 +286,7 @@ def __init__(self) -> None: # 兼容多机 dp size=1 纯 tp 模式的情况 self.is_multinode_tp = args.dp == 1 and args.nnodes > 1 self.shared_tp_infos = [ - SharedInt(f"{get_unique_server_name()}_mem_manger_can_use_token_num_{rank_in_node}") + SharedInt(f"mem_manger_can_use_token_num_{rank_in_node}") for rank_in_node in range(0, self.node_world_size, self.dp_world_size) ] diff --git a/lightllm/server/api_http.py b/lightllm/server/api_http.py index a2a6036f28..9f3f7dd19d 100755 --- a/lightllm/server/api_http.py +++ b/lightllm/server/api_http.py @@ -111,7 +111,7 @@ def set_args(self, args: StartArgs): self.metric_client = MetricClient(get_shm_port_args().metric_port) self.httpserver_manager = HttpServerManager(args=args) dp_size_in_node = max(1, args.dp // args.nnodes) # 兼容多机纯tp的运行模式,这时候 1 // 2 == 0, 需要兼容 - self.shared_token_load = TokenLoad(f"{get_unique_server_name()}_shared_token_load", dp_size_in_node) + self.shared_token_load = TokenLoad("shared_token_load", dp_size_in_node) g_objs = G_Objs() diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 7c7ac9fe48..e3f06dd0af 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -338,6 +338,7 @@ def _launch_subprocesses(args: StartArgs): validate_ports(ports_to_check) set_env_start_args(args) + process_manager.setup_exit_controller() get_shm_port_args(create=True) # 多机用于收发node ip, 这个地方修改了args env,所以需要重新设置一下。 send_and_receive_node_ip(args) @@ -480,6 +481,7 @@ def pd_master_start(args: StartArgs): validate_ports([args.port]) set_env_start_args(args) + process_manager.setup_exit_controller() get_shm_port_args(create=True) logger.info(f"all start args:{args}") @@ -545,6 +547,7 @@ def visual_only_start(args): ports_to_check.append(args.visual_rpyc_port) validate_ports(ports_to_check) set_env_start_args(args) + process_manager.setup_exit_controller() get_shm_port_args(create=True) logger.info(f"all start args:{args}") @@ -572,6 +575,7 @@ def config_server_start(args): ports_to_check.append(args.config_server_visual_redis_port) validate_ports(ports_to_check) set_env_start_args(args) + process_manager.setup_exit_controller() get_shm_port_args(create=True) logger.info(f"all start args:{args}") diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index 9729a8205c..49dcc5cfb1 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -9,7 +9,6 @@ from .shm_array import ShmArray from .token_chunck_hash_list import TokenHashList, CpuCachePageList, TokenPageLenList from lightllm.server.req_id_generator import convert_sub_id_to_group_id -from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.config_utils import is_linear_att_mixed_model from lightllm.utils.kv_cache_utils import compute_token_list_hash @@ -280,22 +279,19 @@ def _fill_linear_att_token_hash(self): return def create_prompt_ids_shm_array(self): - service_uni_name = get_unique_server_name() - name = f"{service_uni_name}_shm_prompts_{self.index_in_shm_mem}" + name = f"shm_prompts_{self.index_in_shm_mem}" self.shm_prompt_ids = ShmArray(name, (self.alloc_shm_numpy_len,), dtype=np.int64) self.shm_prompt_ids.create_shm() return def link_prompt_ids_shm_array(self): - service_uni_name = get_unique_server_name() - name = f"{service_uni_name}_shm_prompts_{self.index_in_shm_mem}" + name = f"shm_prompts_{self.index_in_shm_mem}" self.shm_prompt_ids = ShmArray(name, (self.alloc_shm_numpy_len,), dtype=np.int64) self.shm_prompt_ids.link_shm() return def create_logprobs_shm_array(self): - service_uni_name = get_unique_server_name() - name = f"{service_uni_name}_shm_logprobs_{self.index_in_shm_mem}" + name = f"shm_logprobs_{self.index_in_shm_mem}" self.shm_logprobs = ShmArray( name, (self.alloc_shm_numpy_len,), @@ -308,8 +304,7 @@ def create_logprobs_shm_array(self): return def link_logprobs_shm_array(self): - service_uni_name = get_unique_server_name() - name = f"{service_uni_name}_shm_logprobs_{self.index_in_shm_mem}" + name = f"shm_logprobs_{self.index_in_shm_mem}" self.shm_logprobs = ShmArray( name, (self.alloc_shm_numpy_len,), diff --git a/lightllm/server/core/objs/shm_objs_io_buffer.py b/lightllm/server/core/objs/shm_objs_io_buffer.py index 05b6087601..d1988762d5 100644 --- a/lightllm/server/core/objs/shm_objs_io_buffer.py +++ b/lightllm/server/core/objs/shm_objs_io_buffer.py @@ -2,7 +2,6 @@ import pickle from lightllm.server.core.objs.atomic_lock import AtomicShmLock from lightllm.utils.envs_utils import get_env_start_args -from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.log_utils import init_logger from lightllm.utils.shm_utils import create_or_link_shm @@ -14,8 +13,8 @@ class ShmObjsIOBuffer: def __init__(self, tail_str=""): self.args = get_env_start_args() - self.name = f"{get_unique_server_name()}_ShmReqsBufferParams_{tail_str}" - self.lock = AtomicShmLock(lock_name=f"{get_unique_server_name()}_ShmReqsBufferParams_atomlock_{tail_str}") + self.name = f"ShmReqsBufferParams_{tail_str}" + self.lock = AtomicShmLock(lock_name=f"ShmReqsBufferParams_atomlock_{tail_str}") self._create_or_link_shm() self.node_world_size = self.args.tp // self.args.nnodes diff --git a/lightllm/server/core/objs/shm_req_manager.py b/lightllm/server/core/objs/shm_req_manager.py index c61376c8d0..d7783dcf50 100644 --- a/lightllm/server/core/objs/shm_req_manager.py +++ b/lightllm/server/core/objs/shm_req_manager.py @@ -1,6 +1,5 @@ import ctypes import numpy as np -from lightllm.utils.envs_utils import get_unique_server_name from multiprocessing import shared_memory from lightllm.utils.log_utils import init_logger from .req import Req, ChunkedPrefillReq @@ -44,7 +43,7 @@ def init_reqs_shm(self): self._init_reqs_shm() def _init_reqs_shm(self): - shm_name = f"{get_unique_server_name()}_req_shm_total" + shm_name = "req_shm_total" self.reqs_shm = create_or_link_shm(shm_name, self.req_shm_byte_size) return @@ -56,7 +55,7 @@ def init_to_req_objs(self): return def init_to_req_locks(self): - array_lock_name = f"{get_unique_server_name()}_array_reqs_lock" + array_lock_name = "array_reqs_lock" self.reqs_lock = AtomicShmArrayLock(array_lock_name, self.max_req_num) return @@ -64,13 +63,13 @@ def get_req_lock_by_index(self, req_index_in_mem: int) -> AtomicLockItem: return self.reqs_lock.get_lock_context(req_index_in_mem) def init_manager_lock(self): - lock_name = f"{get_unique_server_name()}_shm_reqs_manager_lock" + lock_name = "shm_reqs_manager_lock" self.manager_lock = AtomicShmLock(lock_name) return def init_alloc_state_shm(self): - shm_name = f"{get_unique_server_name()}_req_alloc_states" - req_link_list_name = f"{get_unique_server_name()}_req_linked_states" + shm_name = "req_alloc_states" + req_link_list_name = "req_linked_states" self.linked_req_manager = ReqLinkedListManager(req_link_list_name, self.max_req_num) self.alloc_state_shm = ShmArray(shm_name, (self.max_req_num,), np.int32) self.alloc_state_shm.create_shm() diff --git a/lightllm/server/core/objs/token_metadata.py b/lightllm/server/core/objs/token_metadata.py index 11427eb76a..adb3f836be 100644 --- a/lightllm/server/core/objs/token_metadata.py +++ b/lightllm/server/core/objs/token_metadata.py @@ -4,7 +4,6 @@ import numpy as np -from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.log_utils import init_logger from lightllm.utils.shm_utils import create_or_link_shm @@ -175,5 +174,4 @@ def _build_routed_experts_response(self, packed_routed: Optional[np.ndarray]) -> } def _shm_name(self) -> str: - service_uni_name = get_unique_server_name() - return f"{service_uni_name}_shm_final_token_metadata_{self.req.index_in_shm_mem}" + return f"shm_final_token_metadata_{self.req.index_in_shm_mem}" diff --git a/lightllm/server/embed_cache/utils.py b/lightllm/server/embed_cache/utils.py index 367bcc91a9..fcd48d421f 100644 --- a/lightllm/server/embed_cache/utils.py +++ b/lightllm/server/embed_cache/utils.py @@ -1,7 +1,10 @@ import multiprocessing.shared_memory as shm +from lightllm.utils.shm_utils import get_service_shm_name + def create_shm(name, data): + name = get_service_shm_name(name) try: data_size = len(data) shared_memory = shm.SharedMemory(name=name, create=True, size=data_size) @@ -12,16 +15,18 @@ def create_shm(name, data): def read_shm(name): + name = get_service_shm_name(name) shared_memory = shm.SharedMemory(name=name) data = shared_memory.buf.tobytes() return data def free_shm(name): + name = get_service_shm_name(name) shared_memory = shm.SharedMemory(name=name) shared_memory.close() shared_memory.unlink() def get_shm_name_data(uid): - return str(uid) + "-data" + return f"{uid}-data" diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index cea3cb6fc9..2be5aae107 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -36,7 +36,6 @@ from .manager_ext import HttpRlManagerHelper 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 from lightllm.utils.shm_port_args import get_shm_port_args from lightllm.utils.error_utils import ( ClientDisconnected, @@ -62,7 +61,7 @@ def __init__( self.multinode_req_manager = None self.nnodes = args.nnodes - self._shm_lock_pool = AtomicShmArrayLock(f"{get_unique_server_name()}_lightllm_resource_lock", 2) + self._shm_lock_pool = AtomicShmArrayLock("lightllm_resource_lock", 2) self._resource_lock = AsyncLock(self._shm_lock_pool.get_lock_context(0)) self._run_reqs_count_lock = AsyncLock(self._shm_lock_pool.get_lock_context(1)) self.node_rank = args.node_rank @@ -129,19 +128,19 @@ def __init__( self.vocab_size = max(get_vocab_size(args.model_dir), self.tokenizer.vocab_size) # Timemark of the latest successful inference, used by passive /health checks. - self.latest_success_infer_time_mark = SharedInt(f"{get_unique_server_name()}_latest_success_infer_time_mark") + self.latest_success_infer_time_mark = SharedInt("latest_success_infer_time_mark") self.latest_success_infer_time_mark.set_value(int(time.time())) self.rl_controller: Optional[HttpRlController] = HttpRlController(self) if args.enable_rl else None - self.run_reqs_count_mark = SharedInt(f"{get_unique_server_name()}_run_reqs_count_mark") + self.run_reqs_count_mark = SharedInt("run_reqs_count_mark") self.run_reqs_count_mark.set_value(0) # 用于记录真实的--max_total_token_num 参数,当这个参数在启动参数中没有设置的时候,其是在推理进程中被分析出来的, # 这个时候如果 --max_req_total_len > --max_total_token_num 时,如果httpserver放过一些非法的输入进入后续的模块可能 # 会触发整个系统崩溃,所以httpserver需要知道真实的 max_total_token_num的数据,用于提前拦截非法请求等参数。 # router 进程会在启动后向这个共享内存写入正确的max_total_token_num 参数,用于后续的请求控制。 - self.shm_max_total_token_num = SharedInt(f"{get_unique_server_name()}_shm_max_total_token_num") + self.shm_max_total_token_num = SharedInt("shm_max_total_token_num") return def _log_stage_timing(self, group_request_id: int, start_time: float, stage: str, **kwargs): diff --git a/lightllm/server/multi_level_kv_cache/cpu_cache_client.py b/lightllm/server/multi_level_kv_cache/cpu_cache_client.py index 33da63ab56..fb9e74c47b 100644 --- a/lightllm/server/multi_level_kv_cache/cpu_cache_client.py +++ b/lightllm/server/multi_level_kv_cache/cpu_cache_client.py @@ -1,5 +1,5 @@ import ctypes -from lightllm.utils.envs_utils import get_env_start_args, get_unique_server_name, get_disk_cache_prompt_limit_length +from lightllm.utils.envs_utils import get_env_start_args, get_disk_cache_prompt_limit_length from typing import List, Optional, Tuple from lightllm.utils.log_utils import init_logger from lightllm.common.cpu_cache import CpuCacheCreator, CpuCacheTensorSpec @@ -20,7 +20,7 @@ def __init__(self, only_create_meta_data: bool, init_shm_data: bool): # to do here need calcu from from settings. self.kv_cache_tensor_meta = calcu_cpu_cache_meta() self.page_num: int = self.kv_cache_tensor_meta.page_num - self.lock = AtomicShmLock(lock_name=f"{get_unique_server_name()}_cpu_kv_cache_client_lock") + self.lock = AtomicShmLock(lock_name="cpu_kv_cache_client_lock") self._create_cpu_status_list(init_shm_data) if not only_create_meta_data: @@ -268,18 +268,18 @@ def recycle_pages(self, page_list: List[int]): def _create_cpu_status_list(self, init_shm_data: bool): self.page_items = ShmLinkedList( - name=f"{get_unique_server_name()}_cpu_kv_cache_page_items", + name="cpu_kv_cache_page_items", item_class=_CpuPageStatus, capacity=self.page_num, init_shm_data=init_shm_data, ) self.page_hash_dict = ShmDict( - name=f"{get_unique_server_name()}_cpu_kv_cache_hash", + name="cpu_kv_cache_hash", capacity=self.page_num * 2, init_shm_data=init_shm_data, ) self.offload_page_indexes = IntList( - name=f"{get_unique_server_name()}_cpu_kv_cache_offload_page_indexes", + name="cpu_kv_cache_offload_page_indexes", capacity=self.page_num * 2, init_shm_data=init_shm_data, ) diff --git a/lightllm/server/multi_level_kv_cache/shm_objs.py b/lightllm/server/multi_level_kv_cache/shm_objs.py index 50f3abfc7b..5165ca3af1 100644 --- a/lightllm/server/multi_level_kv_cache/shm_objs.py +++ b/lightllm/server/multi_level_kv_cache/shm_objs.py @@ -3,7 +3,7 @@ from multiprocessing import shared_memory from typing import List, Optional from lightllm.utils.log_utils import init_logger -from lightllm.utils.auto_shm_cleanup import register_posix_shm_for_cleanup +from lightllm.utils.shm_utils import get_service_shm_name logger = init_logger(__name__) @@ -290,11 +290,10 @@ def key(self, value: int): self.key_high = (value >> 64) & 0xFFFFFFFFFFFFFFFF -def _create_shm(name: str, byte_size: int, auto_cleanup: bool = False): +def _create_shm(name: str, byte_size: int): + name = get_service_shm_name(name) try: shm = shared_memory.SharedMemory(name=name, create=True, size=byte_size) - if auto_cleanup: - register_posix_shm_for_cleanup(name) logger.info(f"create lock shm {name}") except: shm = shared_memory.SharedMemory(name=name, create=False, size=byte_size) diff --git a/lightllm/server/req_id_generator.py b/lightllm/server/req_id_generator.py index 8b8d4a5dc5..720c5f8d53 100644 --- a/lightllm/server/req_id_generator.py +++ b/lightllm/server/req_id_generator.py @@ -19,17 +19,17 @@ class ReqIDGenerator: def __init__(self): from lightllm.server.core.objs.atomic_lock import AtomicShmLock from lightllm.server.core.objs.shm_array import ShmArray - from lightllm.utils.envs_utils import get_unique_server_name, get_env_start_args + from lightllm.utils.envs_utils import get_env_start_args self.args = get_env_start_args() self.use_config_server = ( self.args.config_server_host and self.args.config_server_port and self.args.run_mode == "pd_master" ) - self.current_id = ShmArray(f"{get_unique_server_name()}_req_id_gen", (2,), dtype=np.int64) + self.current_id = ShmArray("req_id_gen", (2,), dtype=np.int64) self.current_id.create_shm() self.current_id.arr[0] = 0 self.current_id.arr[1] = 0 - self.lock = AtomicShmLock(f"{get_unique_server_name()}_req_id_gen_lock") + self.lock = AtomicShmLock("req_id_gen_lock") self._wait_all_workers_ready() logger.info("ReqIDGenerator init finished") @@ -37,12 +37,9 @@ def _wait_all_workers_ready(self): if self.args.httpserver_workers == 1: return - from lightllm.utils.envs_utils import get_unique_server_name from lightllm.server.core.objs.shm_array import ShmArray - _sync_shm = ShmArray( - f"{get_unique_server_name()}_httpworker_start_sync", (self.args.httpserver_workers,), dtype=np.int64 - ) + _sync_shm = ShmArray("httpworker_start_sync", (self.args.httpserver_workers,), dtype=np.int64) _sync_shm.create_shm() # 等待所有 httpserver 的 worker 启动完成,防止重新初始化对应的请求id 对应的shm try_count = 0 diff --git a/lightllm/server/router/dynamic_prompt/linear_att_radix_cache.py b/lightllm/server/router/dynamic_prompt/linear_att_radix_cache.py index 6a8e0a3917..f2e0b42bcf 100644 --- a/lightllm/server/router/dynamic_prompt/linear_att_radix_cache.py +++ b/lightllm/server/router/dynamic_prompt/linear_att_radix_cache.py @@ -112,7 +112,6 @@ def is_leaf(self): class LinearAttPagedRadixCache: def __init__( self, - unique_name: str, total_token_num: int, rank_in_node: int, hash_page_size: int, @@ -147,11 +146,9 @@ def __init__( key=lambda x: x.get_compare_key_for_buffer_idx() ) - self.refed_tokens_num = SharedArray(f"{unique_name}_refed_tokens_num_{rank_in_node}", (1,), dtype=np.int64) + self.refed_tokens_num = SharedArray(f"refed_tokens_num_{rank_in_node}", (1,), dtype=np.int64) self.refed_tokens_num.arr[0] = 0 - self.tree_total_tokens_num = SharedArray( - f"{unique_name}_tree_total_tokens_num_{rank_in_node}", (1,), dtype=np.int64 - ) + self.tree_total_tokens_num = SharedArray(f"tree_total_tokens_num_{rank_in_node}", (1,), dtype=np.int64) self.tree_total_tokens_num.arr[0] = 0 self.linear_att_small_page_buffers: LinearAttCacheManager = linear_att_small_page_buffers diff --git a/lightllm/server/router/dynamic_prompt/radix_cache.py b/lightllm/server/router/dynamic_prompt/radix_cache.py index c103a61473..69176950b5 100644 --- a/lightllm/server/router/dynamic_prompt/radix_cache.py +++ b/lightllm/server/router/dynamic_prompt/radix_cache.py @@ -99,11 +99,7 @@ def match(t1: torch.Tensor, t2: torch.Tensor) -> int: class RadixCache: - """ - unique_name 主要用于解决单机,多实列部署时的shm冲突 - """ - - def __init__(self, unique_name, total_token_num, rank_in_node, mem_manager=None): + def __init__(self, total_token_num, rank_in_node, mem_manager=None): from lightllm.common.kv_cache_mem_manager import MemoryManager self.total_token_num = total_token_num @@ -119,11 +115,9 @@ def __init__(self, unique_name, total_token_num, rank_in_node, mem_manager=None) self.evict_tree_set: Set[TreeNode] = SortedSet(key=lambda x: x.get_compare_key()) # 自定义比较器 self.evict_tree_set.add(self.root_node) - self.refed_tokens_num = SharedArray(f"{unique_name}_refed_tokens_num_{rank_in_node}", (1,), dtype=np.int64) + self.refed_tokens_num = SharedArray(f"refed_tokens_num_{rank_in_node}", (1,), dtype=np.int64) self.refed_tokens_num.arr[0] = 0 - self.tree_total_tokens_num = SharedArray( - f"{unique_name}_tree_total_tokens_num_{rank_in_node}", (1,), dtype=np.int64 - ) + self.tree_total_tokens_num = SharedArray(f"tree_total_tokens_num_{rank_in_node}", (1,), dtype=np.int64) self.tree_total_tokens_num.arr[0] = 0 def insert(self, key, value=None) -> Tuple[int, Optional[TreeNode]]: @@ -515,11 +509,9 @@ class _RadixCacheReadOnlyClient: router 端只读用的客户端,用于从共享内存中读取树结构中的信息,用于进行prompt cache 的调度估计。 """ - def __init__(self, unique_name, total_token_num, rank_in_node): - self.refed_tokens_num = SharedArray(f"{unique_name}_refed_tokens_num_{rank_in_node}", (1,), dtype=np.int64) - self.tree_total_tokens_num = SharedArray( - f"{unique_name}_tree_total_tokens_num_{rank_in_node}", (1,), dtype=np.int64 - ) + def __init__(self, total_token_num, rank_in_node): + self.refed_tokens_num = SharedArray(f"refed_tokens_num_{rank_in_node}", (1,), dtype=np.int64) + self.tree_total_tokens_num = SharedArray(f"tree_total_tokens_num_{rank_in_node}", (1,), dtype=np.int64) def get_refed_tokens_num(self): return self.refed_tokens_num.arr[0] @@ -532,9 +524,9 @@ def get_unrefed_tokens_num(self): class RadixCacheReadOnlyClient: - def __init__(self, unique_name, total_token_num, node_world_size, dp_world_size): + def __init__(self, total_token_num, node_world_size, dp_world_size): self.dp_rank_clients: List[_RadixCacheReadOnlyClient] = [ - _RadixCacheReadOnlyClient(unique_name, total_token_num, rank_in_node) + _RadixCacheReadOnlyClient(total_token_num, rank_in_node) for rank_in_node in range(0, node_world_size, dp_world_size) ] diff --git a/lightllm/server/router/manager.py b/lightllm/server/router/manager.py index b1375d754c..6586d62583 100644 --- a/lightllm/server/router/manager.py +++ b/lightllm/server/router/manager.py @@ -64,7 +64,7 @@ def __init__(self, args: StartArgs): self.load_way = args.load_way self.max_total_token_num = args.max_total_token_num # 存储在共享内存中的真实token容量数据 - self.shm_max_total_token_num = SharedInt(f"{get_unique_server_name()}_shm_max_total_token_num") + self.shm_max_total_token_num = SharedInt("shm_max_total_token_num") self.shm_req_manager = ShmReqManager() # 用共享内存进行共享,router 模块读取进行精确的调度估计 self.read_only_statics_mem_manager = ReadOnlyStaticsMemoryManager() @@ -72,7 +72,7 @@ def __init__(self, args: StartArgs): self.radix_cache_client = None # 共享变量,用于存储router端调度分析得到的机器负载信息 - self.shared_token_load = TokenLoad(f"{get_unique_server_name()}_shared_token_load", self.dp_size_in_node) + self.shared_token_load = TokenLoad("shared_token_load", self.dp_size_in_node) for dp_index in range(self.dp_size_in_node): self.shared_token_load.set_estimated_peak_token_count(0, dp_index) self.shared_token_load.set_current_load(0.0, dp_index) @@ -197,7 +197,6 @@ async def wait_to_model_ready(self): if not self.args.disable_dynamic_prompt_cache: self.radix_cache_client = RadixCacheReadOnlyClient( - get_unique_server_name(), self.max_total_token_num, node_world_size=self.node_world_size, dp_world_size=self.dp_world_size, diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 560c3b6de4..9ac4b5be51 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -21,7 +21,6 @@ from lightllm.server.router.dynamic_prompt.radix_cache import RadixCache from lightllm.common.basemodel.batch_objs import ModelOutput, ModelInput from lightllm.utils.dist_utils import init_distributed_env -from lightllm.utils.envs_utils import get_unique_server_name from lightllm.server.core.objs import ShmReqManager, StartArgs from lightllm.server.core.objs.io_objs import AbortedReqCmd, StopStrMatchedReqCmd from lightllm.server.router.model_infer.infer_batch import g_infer_context @@ -122,7 +121,7 @@ def init_model(self, kvargs): ) dist_group_manager.create_groups(group_size=group_size) # set the default group - self.shared_token_load = TokenLoad(f"{get_unique_server_name()}_shared_token_load", self.dp_size_in_node) + self.shared_token_load = TokenLoad("shared_token_load", self.dp_size_in_node) if self.args.enable_multimodal: g_infer_context.init_cpu_embed_cache_client() @@ -166,7 +165,6 @@ def init_model(self, kvargs): else: if self.is_linear_att_mixed_model: self.radix_cache = LinearAttPagedRadixCache( - unique_name=get_unique_server_name(), total_token_num=self.model.mem_manager.size, rank_in_node=self.rank_in_node, hash_page_size=self.args.linear_att_hash_page_size, @@ -176,7 +174,6 @@ def init_model(self, kvargs): ) else: self.radix_cache = RadixCache( - unique_name=get_unique_server_name(), total_token_num=self.model.mem_manager.size, rank_in_node=self.rank_in_node, mem_manager=self.model.mem_manager, diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py index 2fa2c9cb9a..2b2e435b62 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py @@ -5,7 +5,7 @@ import torch from typing import List from lightllm.common.kv_cache_mem_manager import MemoryManager -from lightllm.utils.envs_utils import get_unique_server_name, get_env_start_args +from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.dist_utils import get_dp_rank_in_node from lightllm.server.core.objs.shm_array import ShmArray from ...infer_batch import InferReq @@ -26,7 +26,7 @@ def __init__(self, max_req_num: int, dp_size_in_node: int, backend): # 0 代表 kv_len, 1 代表 radix_cache_len self.shared_req_infos = ShmArray( - name=f"{get_unique_server_name()}_dp_shared_req_infos", + name="dp_shared_req_infos", shape=(self.max_req_num, dp_size_in_node, 2), dtype=np.int64, ) diff --git a/lightllm/server/visualserver/model_infer/vision_peak_vram_hold.py b/lightllm/server/visualserver/model_infer/vision_peak_vram_hold.py index a1af03939e..c47eeecfc0 100644 --- a/lightllm/server/visualserver/model_infer/vision_peak_vram_hold.py +++ b/lightllm/server/visualserver/model_infer/vision_peak_vram_hold.py @@ -12,7 +12,6 @@ from PIL import Image from lightllm.server.embed_cache.utils import create_shm, free_shm, get_shm_name_data from lightllm.server.multimodal_params import ImageItem -from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.log_utils import init_logger logger = init_logger(__name__) @@ -120,7 +119,7 @@ def _build_worst_case_image_items( color = (255, 255, 255) if batch_id % 2 else (0, 0, 0) image_bytes = self._gen_rgb_jpeg_bytes(width, height, color=color) item = ImageItem(type="base64", data="") - item.uuid = f"{get_unique_server_name()}_vision_peak_hold_dp{dp_rank_id}_{batch_id}" + item.uuid = f"vision_peak_hold_dp{dp_rank_id}_{batch_id}" item.image_w = width item.image_h = height # InternVL encode() reads image_patch_max_num from extra_params (normally set by diff --git a/lightllm/utils/auto_shm_cleanup.py b/lightllm/utils/auto_shm_cleanup.py deleted file mode 100644 index 2417fef085..0000000000 --- a/lightllm/utils/auto_shm_cleanup.py +++ /dev/null @@ -1,135 +0,0 @@ -import os -import ctypes -import atexit -import signal -import threading -import psutil -from multiprocessing import shared_memory -from typing import Set, Optional -from lightllm.utils.log_utils import init_logger - -logger = init_logger(__name__) - - -class AutoShmCleanup: - """ - 自动清理 System V 和 POSIX 共享内存 - shared_memory.SharedMemory虽然有自动请理功能,但如果自动清理时仍有进程占用会清理失败,这里可做最后兜底清理 - """ - - def __init__(self): - self.libc = None - self._init_libc() - # System V - self.registered_shm_keys = [] - self.registered_shm_ids = [] - # POSIX - self.registered_posix_shm_names = [] - self.signal_handlers_registered = False - self._register_handlers_for_cleanup() - - def _init_libc(self): - try: - self.libc = ctypes.CDLL("/usr/lib/x86_64-linux-gnu/libc.so.6") - self.libc.shmget.argtypes = (ctypes.c_long, ctypes.c_size_t, ctypes.c_int) - self.libc.shmget.restype = ctypes.c_int - self.libc.shmctl.argtypes = (ctypes.c_int, ctypes.c_int, ctypes.c_void_p) - self.libc.shmctl.restype = ctypes.c_int - except Exception as e: - logger.debug(f"libc init failed: {e}") - self.libc = None - - def _register_handlers_for_cleanup(self): - atexit.register(self._cleanup) - self.register_signal_handlers() - - def register_signal_handlers(self): - if self.signal_handlers_registered or not threading.current_thread() is threading.main_thread(): - return - for sig in (signal.SIGTERM, signal.SIGINT, signal.SIGHUP): - signal.signal(sig, self._signal_cleanup_handler) - self.signal_handlers_registered = True - - def _signal_cleanup_handler(self, signum, frame): - self._cleanup() - parent = psutil.Process(os.getpid()) - # 递归拿到所有子进程并终止 - for ch in parent.children(recursive=True): - ch.kill() - - def _cleanup(self): - """清理:System V 执行 IPC_RMID,POSIX 执行 unlink。""" - removed_sysv = 0 - IPC_RMID = 0 - for shmid in self.registered_shm_ids: - try: - if self.libc.shmctl(shmid, IPC_RMID, None) == 0: - removed_sysv += 1 - except Exception as e: - logger.warning(f"cleanup: shmid {shmid} clean failed, reason: {e}") - pass - for key in self.registered_shm_keys: - shmid = self.libc.shmget(key, 0, 0) - try: - if shmid >= 0 and self.libc.shmctl(shmid, IPC_RMID, None) == 0: - removed_sysv += 1 - except Exception as e: - logger.warning(f"cleanup: shmid {shmid} clean failed, reason: {e}") - pass - if removed_sysv: - logger.info(f"cleanup: removed {removed_sysv} System V shm segments") - - removed_posix = 0 - for name in self.registered_posix_shm_names: - try: - shm = shared_memory.SharedMemory(name=name, create=False) - try: - shm.unlink() - removed_posix += 1 - except FileNotFoundError: - pass - except Exception as e: - logger.warning(f"cleanup: posix shm {name} clean failed, reason: {e}") - pass - finally: - shm.close() - except FileNotFoundError: - pass - except Exception as e: - logger.warning(f"cleanup: posix {name} clean failed, reason: {e}") - pass - if removed_posix: - logger.info(f"cleanup: unlinked {removed_posix} POSIX shm segments") - - def register_sysv_shm(self, key: int, shmid: Optional[int] = None): - """注册 System V 共享内存。""" - self.registered_shm_keys.append(key) - if shmid is not None: - self.registered_shm_ids.append(shmid) - return - - def register_posix_shm(self, name: str): - """注册 POSIX 共享内存。""" - self.registered_posix_shm_names.append(name) - return - - -# 全局自动清理器实例 -_auto_cleanup = None - - -def get_auto_cleanup() -> AutoShmCleanup: - """获取全局自动清理器实例""" - global _auto_cleanup - if _auto_cleanup is None: - _auto_cleanup = AutoShmCleanup() - _auto_cleanup.register_signal_handlers() - return _auto_cleanup - - -def register_sysv_shm_for_cleanup(key: int, shmid: Optional[int] = None): - get_auto_cleanup().register_sysv_shm(key, shmid) - - -def register_posix_shm_for_cleanup(name: str): - get_auto_cleanup().register_posix_shm(name) diff --git a/lightllm/utils/health_check.py b/lightllm/utils/health_check.py index d2a776b862..5090eff8e6 100644 --- a/lightllm/utils/health_check.py +++ b/lightllm/utils/health_check.py @@ -5,7 +5,6 @@ from lightllm.server.router.dynamic_prompt.shared_arr import SharedInt from lightllm.utils.log_utils import init_logger -from lightllm.utils.envs_utils import get_unique_server_name if TYPE_CHECKING: from lightllm.server.core.objs.shm_req_manager import ShmReqManager @@ -18,9 +17,8 @@ class HealthObj: grace_timeout: int = int(os.getenv("HEALTH_TIMEOUT", "200")) def __post_init__(self): - uid = get_unique_server_name() - self.latest_success_infer_time_mark = SharedInt(f"{uid}_latest_success_infer_time_mark") - self.run_reqs_count_mark = SharedInt(f"{uid}_run_reqs_count_mark") + self.latest_success_infer_time_mark = SharedInt("latest_success_infer_time_mark") + self.run_reqs_count_mark = SharedInt("run_reqs_count_mark") def check(self, shm_req_manager: "ShmReqManager") -> bool: """On-the-fly health check: recent success is ok; otherwise require no in-flight shm requests.""" diff --git a/lightllm/utils/kv_cache_utils.py b/lightllm/utils/kv_cache_utils.py index e81caafe7a..44902e4ca4 100644 --- a/lightllm/utils/kv_cache_utils.py +++ b/lightllm/utils/kv_cache_utils.py @@ -28,7 +28,6 @@ from typing import List, Tuple, Optional from tqdm import tqdm -from lightllm.utils.auto_shm_cleanup import register_sysv_shm_for_cleanup from lightllm.utils.dist_utils import get_current_device_id from lightllm.common.linear_att_cache_manager.config_objs import LinearAttCacheConfig @@ -212,7 +211,6 @@ def create_shm_kv_cache_ptr(key: int, size: int) -> int: else: raise Exception(f"Error creating regular shared memory (errno={err})") - register_sysv_shm_for_cleanup(key, shmid) logger.info(f"Shared memory ID: {shmid}") # 附加共享内存 diff --git a/lightllm/utils/rl/bucketed_weight_transfer.py b/lightllm/utils/rl/bucketed_weight_transfer.py index 4849497f1a..0ac3ed71e0 100644 --- a/lightllm/utils/rl/bucketed_weight_transfer.py +++ b/lightllm/utils/rl/bucketed_weight_transfer.py @@ -50,7 +50,12 @@ class TensorMetadata(TypedDict): def create_shared_memory(size: int, name: str): - """Create shared memory for weight transfer. If already exists, attach to it.""" + """Create shared memory for weight transfer. If already exists, attach to it. + + ``name`` is part of the sender/receiver transfer protocol and may be created + by an external RL process, so it is already a complete name rather than a + LightLLM service-local logical name. + """ try: shm = shared_memory.SharedMemory(name=name, create=True, size=size) except FileExistsError: @@ -60,7 +65,7 @@ def create_shared_memory(size: int, name: str): def rebuild_shared_memory(name: str, size: int, dtype=torch.uint8): - """Rebuild tensor from shared memory.""" + """Rebuild tensor from an external sender's complete shared-memory name.""" shm = shared_memory.SharedMemory(name=name) tensor = torch.frombuffer(shm.buf[:size], dtype=dtype) diff --git a/lightllm/utils/service_shm_cleanup.py b/lightllm/utils/service_shm_cleanup.py new file mode 100644 index 0000000000..19c9cbda01 --- /dev/null +++ b/lightllm/utils/service_shm_cleanup.py @@ -0,0 +1,132 @@ +"""由 launcher 统一管理 LightLLM 服务创建的共享内存。 + +设计背景 +-------- +LightLLM 的 router、model、HTTP server 等子进程通过共享内存交换状态。业务层使用逻辑名称, +共享内存基础层会统一添加 ``{service_name}_`` 前缀,因此同一台机器上的多个服务不会重名。 +这些资源由 launcher 统一回收,子进程不单独注册信号处理函数。 + +支持的退出场景 +-------------- +1. SIGINT/SIGTERM/SIGHUP:launcher 先停止子进程,再主动调用本模块返回的 cleanup 函数。 +2. Python 正常退出或未捕获异常:``atexit`` 作为兜底执行同一个 cleanup 函数。 +3. 多实例并存:``/dev/shm`` 名称带有 service name,清理时只处理当前服务。 + +SIGKILL、OOM Killer 等场景无法执行 Python 清理回调,本模块不记录 owner 文件,也不在下次 +启动时补偿清理这些异常残留。该功能定位为 launcher 有机会退出时执行的简单兜底清理。 + +所有启动模式都会创建 ShmPortArgs 等 POSIX 共享内存,因此统一按 service name 清理。System V +共享内存只可能由 normal、prefill、decode 推理节点创建,并继续按照 CPU KV Cache 和多模态缓存 +功能开关选择有效 key;pd_master、visual_only、config_server 不执行 System V SHM 清理。 + +本模块只负责服务内部共享内存。由外部 RL 进程创建并通过协议传入完整名称的共享内存,不属于 +当前 launcher 的服务前缀命名空间,因此不在这里按名称扫描回收。 +""" + +import atexit +import ctypes +import json +import os +import subprocess +from pathlib import Path + +from lightllm.utils.log_utils import init_logger + + +logger = init_logger(__name__) + +SHM_DIR = Path("/dev/shm") + + +class ServiceShmCleanup: + """管理一个 launcher 在当前节点上拥有的共享内存。""" + + def __init__(self, service_name): + if not service_name: + raise RuntimeError("service_name must be initialized before registering shm cleanup") + + self.service_name = service_name + + try: + self.start_args = json.loads(os.environ["LIGHTLLM_START_ARGS"]) + except (KeyError, json.JSONDecodeError): + self.start_args = {} + if not isinstance(self.start_args, dict): + self.start_args = {} + + @staticmethod + def cleanup_posix_shm(service_name): + """删除名称严格属于目标 service 的 POSIX 共享内存。""" + try: + entries = [entry for entry in SHM_DIR.iterdir() if entry.name.startswith(f"{service_name}_")] + except FileNotFoundError: + return 0 + + if not entries: + return 0 + + try: + # 这是服务退出后的兜底路径:先按严格前缀筛选,再启动一次 rm 批量删除。 + # 使用参数列表和 "--",避免 shell 管道、通配符展开及名称转义问题。 + subprocess.run(["rm", "-f", "--", *(str(entry) for entry in entries)], check=True) + except (OSError, subprocess.CalledProcessError): + logger.exception(f"Failed to remove POSIX shm for service {service_name}") + return 0 + return len(entries) + + @staticmethod + def cleanup_system_v_shm(keys): + """删除启动参数中记录的 System V 共享内存。""" + libc = ctypes.CDLL("/usr/lib/x86_64-linux-gnu/libc.so.6", use_errno=True) + libc.shmget.argtypes = (ctypes.c_long, ctypes.c_size_t, ctypes.c_int) + libc.shmget.restype = ctypes.c_int + libc.shmctl.argtypes = (ctypes.c_int, ctypes.c_int, ctypes.c_void_p) + libc.shmctl.restype = ctypes.c_int + + removed = 0 + for key in keys: + try: + # shmget(key, size, shmflg):这里只查找已有段,不创建新段;size=0 不申请空间, + # shmflg=0 表示不附加 IPC_CREAT 等标志。成功返回非负的内核共享内存 ID, + # 失败返回 -1。 + shmid = libc.shmget(int(key), 0, 0) + if shmid < 0: + continue + + # shmctl(shmid, cmd, buf):cmd=0 是 IPC_RMID,buf 在该命令下不使用,所以传 None。 + # IPC_RMID 将共享内存标记为删除;最后一个已 attach 的进程 detach 后才真正释放。 + # shmctl 成功返回 0,失败返回 -1。 + removed += int(libc.shmctl(shmid, 0, None) == 0) + except Exception: + logger.exception(f"Failed to remove System V shm key {key}") + return removed + + def cleanup_service_resources(self): + """回收当前服务实际创建的 POSIX 和 System V 共享内存。""" + removed_posix = self.cleanup_posix_shm(self.service_name) + system_v_shm_keys = [] + if self.start_args.get("run_mode") in ["normal", "prefill", "decode"]: + if self.start_args.get("enable_cpu_cache") and self.start_args.get("cpu_kv_cache_shm_id") is not None: + system_v_shm_keys.append(int(self.start_args["cpu_kv_cache_shm_id"])) + if self.start_args.get("enable_multimodal") and self.start_args.get("multi_modal_cache_shm_id") is not None: + system_v_shm_keys.append(int(self.start_args["multi_modal_cache_shm_id"])) + removed_system_v = self.cleanup_system_v_shm(system_v_shm_keys) if system_v_shm_keys else 0 + if removed_posix or removed_system_v: + logger.info( + f"Cleaned service shm for {self.service_name}: " + f"POSIX={removed_posix}, System V keys={removed_system_v}" + ) + + def register(self): + """安装正常退出时的兜底清理回调。""" + atexit.register(self.cleanup) + return self.cleanup + + def cleanup(self): + """执行幂等的兜底回收,允许主动退出流程和 atexit 重复调用。""" + self.cleanup_service_resources() + + +def register_launcher_shm_cleanup(service_name): + """创建 launcher 的清理对象,并返回 cleanup 函数。""" + return ServiceShmCleanup(service_name).register() diff --git a/lightllm/utils/shm_port_args.py b/lightllm/utils/shm_port_args.py index b486c6e2f2..b0fe9c5523 100644 --- a/lightllm/utils/shm_port_args.py +++ b/lightllm/utils/shm_port_args.py @@ -24,7 +24,7 @@ -------- 使用 ShmPortArgs 之前,必须先初始化以下环境信息(否则会直接抛错): 1. `set_unique_server_name(args)` - → 写入 `LIGHTLLM_UNIQUE_SERVICE_NAME_ID`,供 `get_unique_server_name()` 使用,用于拼 shm 名。 + → 写入 `LIGHTLLM_UNIQUE_SERVICE_NAME_ID`,共享内存基础层会统一添加该服务前缀。 2. `set_env_start_args(args)` → 写入 `LIGHTLLM_START_ARGS`,供 `get_env_start_args()` 使用,用于读取用户已设置端口、 以及 visual_dp / audio_dp 等分配参数。 @@ -56,9 +56,9 @@ from filelock import FileLock -from lightllm.utils.envs_utils import get_env_start_args, get_unique_server_name +from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.log_utils import init_logger -from lightllm.utils.shm_utils import create_or_link_shm +from lightllm.utils.shm_utils import create_or_link_shm, get_service_shm_name logger = init_logger(__name__) @@ -71,21 +71,15 @@ class ShmPortArgs: _instance: "ShmPortArgs | None" = None def __init__(self, create: bool = False): - uni = get_unique_server_name() - if not uni: - raise RuntimeError( - "LIGHTLLM_UNIQUE_SERVICE_NAME_ID is unset; " "call set_unique_server_name(args) before ShmPortArgs" - ) if "LIGHTLLM_START_ARGS" not in os.environ: raise RuntimeError("LIGHTLLM_START_ARGS is unset; call set_env_start_args(args) before ShmPortArgs") - self._shm_name = f"{uni}_shm_port_args" - self._lock = FileLock(f"/tmp/{self._shm_name}.lock") + self._shm_name = "shm_port_args" + self._lock = FileLock(f"/tmp/{get_service_shm_name(self._shm_name)}.lock") self.shm = create_or_link_shm( self._shm_name, self._SHM_SIZE, force_mode="create" if create else "link", - auto_cleanup=create, ) if create: self._save({}) diff --git a/lightllm/utils/shm_utils.py b/lightllm/utils/shm_utils.py index 0a25d82143..1929865375 100644 --- a/lightllm/utils/shm_utils.py +++ b/lightllm/utils/shm_utils.py @@ -1,15 +1,31 @@ from multiprocessing import shared_memory from filelock import FileLock +from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.log_utils import init_logger -from lightllm.utils.auto_shm_cleanup import register_posix_shm_for_cleanup logger = init_logger(__name__) -def create_or_link_shm(name, expected_size, force_mode=None, auto_cleanup=False): +def get_service_shm_name(name): + """为内部共享内存统一添加当前服务(UUID + node rank)前缀。 + + 已带当前服务前缀的完整名称保持不变,便于底层包装函数安全复用。service name + 未初始化时直接报错,避免创建无法区分服务、也无法被 launcher 定向回收的裸名称。 + """ + name = str(name) + service_name = get_unique_server_name() + if not service_name: + raise RuntimeError( + "LIGHTLLM_UNIQUE_SERVICE_NAME_ID is unset; " "call set_unique_server_name(args) before using shared memory" + ) + prefix = f"{service_name}_" + return name if name.startswith(prefix) else f"{prefix}{name}" + + +def create_or_link_shm(name, expected_size, force_mode=None): """ Args: - name: name of the shared memory + name: logical name of the shared memory; the current service prefix is added here expected_size: expected size of the shared memory, if expected_size == -1, no check for size linked. force_mode: force mode - 'create': force create new shared memory, if exists, delete and create @@ -23,19 +39,20 @@ def create_or_link_shm(name, expected_size, force_mode=None, auto_cleanup=False) FileNotFoundError: when force_mode='link' but shared memory not exists ValueError: when force_mode='link' but size mismatch """ + name = get_service_shm_name(name) lock_name = f"/tmp/{name}.lock" if force_mode == "create": with FileLock(lock_name): - return _force_create_shm(name, expected_size, auto_cleanup) + return _force_create_shm(name, expected_size) elif force_mode == "link": return _force_link_shm(name, expected_size) else: with FileLock(lock_name): - return _smart_create_or_link_shm(name, expected_size, auto_cleanup) + return _smart_create_or_link_shm(name, expected_size) -def _force_create_shm(name, expected_size, auto_cleanup): +def _force_create_shm(name, expected_size): """强制创建新的共享内存""" try: existing_shm = shared_memory.SharedMemory(name=name) @@ -46,8 +63,6 @@ def _force_create_shm(name, expected_size, auto_cleanup): # 创建新的共享内存 shm = shared_memory.SharedMemory(name=name, create=True, size=expected_size) - if auto_cleanup: - register_posix_shm_for_cleanup(name) return shm @@ -66,7 +81,7 @@ def _force_link_shm(name, expected_size): raise e -def _smart_create_or_link_shm(name, expected_size, auto_cleanup): +def _smart_create_or_link_shm(name, expected_size): """优先连接,不存在则创建""" try: shm = _force_link_shm(name=name, expected_size=expected_size) @@ -74,4 +89,4 @@ def _smart_create_or_link_shm(name, expected_size, auto_cleanup): except: pass - return _force_create_shm(name=name, expected_size=expected_size, auto_cleanup=auto_cleanup) + return _force_create_shm(name=name, expected_size=expected_size) diff --git a/lightllm/utils/start_utils.py b/lightllm/utils/start_utils.py index 65b6f5f41b..5f816f2832 100644 --- a/lightllm/utils/start_utils.py +++ b/lightllm/utils/start_utils.py @@ -7,6 +7,8 @@ import psutil from lightllm.utils.log_utils import init_logger from lightllm.utils.process_check import is_process_active +from lightllm.utils.envs_utils import get_unique_server_name +from lightllm.utils.service_shm_cleanup import register_launcher_shm_cleanup logger = init_logger(__name__) @@ -15,11 +17,13 @@ class SubmoduleManager: def __init__(self): self.processes = [] self.process_names = {} + self._cleanup_service_shm = None def start_submodule_processes(self, start_funcs=[], start_args=[]): assert len(start_funcs) == len(start_args) pipe_readers = [] processes = [] + managed_processes = [] for start_func, start_arg in zip(start_funcs, start_args): pipe_reader, pipe_writer = mp.Pipe(duplex=False) @@ -30,6 +34,11 @@ def start_submodule_processes(self, start_funcs=[], start_args=[]): process.start() pipe_readers.append(pipe_reader) processes.append(process) + # 初始化完成前也可能收到退出信号,因此子进程启动后立即纳入管理。 + managed_process = psutil.Process(process.pid) + managed_processes.append(managed_process) + self.processes.append(managed_process) + self.process_names[managed_process] = managed_process.name() # Wait for all processes to initialize for index, pipe_reader in enumerate(pipe_readers): @@ -43,10 +52,7 @@ def start_submodule_processes(self, start_funcs=[], start_args=[]): logger.info(f"init func {start_funcs[index].__name__} : {str(init_state)}") assert all([proc.is_alive() for proc in processes]) - processes = [psutil.Process(proc.pid) for proc in processes] - self.processes.extend(processes) - self.process_names.update((process, process.name()) for process in processes) - return processes + return managed_processes def register_process_tree(self, root_process): """Add persistent LightLLM descendants to supervision. @@ -91,6 +97,9 @@ def kill_recursive(proc): kill_recursive(proc) proc.wait() + if self._cleanup_service_shm is not None: + self._cleanup_service_shm() + # recover the gpu compute mode is_enable_mps = get_env_start_args().enable_mps if is_enable_mps: @@ -99,7 +108,29 @@ def kill_recursive(proc): stop_mps() logger.info("All processes terminated gracefully.") + def setup_exit_controller(self): + """初始化 launcher 退出清理控制器,注册启动阶段信号处理和 atexit 回调。 + + 在 service name 和启动参数写入环境后、创建共享内存或启动子进程前调用。 + 重复调用只注册一次;退出时由 launcher 在子进程停止后统一回收共享内存。 + 启动完成后由 setup_signal_handlers 替换信号处理函数,纳入 HTTP server 的退出流程。 + """ + if self._cleanup_service_shm is not None: + return + self._cleanup_service_shm = register_launcher_shm_cleanup(get_unique_server_name()) + + def signal_handler(sig, _frame): + logger.info(f"Received {signal.Signals(sig).name} during startup, shutting down...") + self.terminate_all_processes() + sys.exit(0) + + signal.signal(signal.SIGTERM, signal_handler) + signal.signal(signal.SIGINT, signal_handler) + signal.signal(signal.SIGHUP, signal_handler) + def setup_signal_handlers(self, http_server_process=None): + """在子进程启动完成后安装退出信号处理函数,覆盖启动阶段的处理函数。""" + def signal_handler(sig, _frame): if sig == signal.SIGINT: logger.info("Received SIGINT (Ctrl+C), forcing immediate exit...") diff --git a/unit_tests/server/core/objs/test_shm_array.py b/unit_tests/server/core/objs/test_shm_array.py index 8e55f4a6e1..af02c7de0b 100644 --- a/unit_tests/server/core/objs/test_shm_array.py +++ b/unit_tests/server/core/objs/test_shm_array.py @@ -4,6 +4,15 @@ import numpy as np from multiprocessing import shared_memory from lightllm.server.core.objs.shm_array import ShmArray # Replace 'your_module' with the actual module name +from lightllm.utils import shm_utils + + +@pytest.fixture(scope="module", autouse=True) +def service_name(): + monkeypatch = pytest.MonkeyPatch() + monkeypatch.setattr(shm_utils, "get_unique_server_name", lambda: "test_shm_array_service_0") + yield + monkeypatch.undo() @pytest.fixture(scope="module") 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..3d605e6845 100644 --- a/unit_tests/server/core/objs/test_shm_req_manager.py +++ b/unit_tests/server/core/objs/test_shm_req_manager.py @@ -5,11 +5,14 @@ from easydict import EasyDict from lightllm.utils.envs_utils import set_env_start_args, get_env_start_args +from lightllm.utils import shm_utils from lightllm.server.core.objs.shm_req_manager import ShmReqManager @pytest.fixture(scope="module", autouse=True) def setup_env(): + monkeypatch = pytest.MonkeyPatch() + monkeypatch.setattr(shm_utils, "get_unique_server_name", lambda: "test_shm_req_manager_service_0") original = os.environ.get("LIGHTLLM_START_ARGS") set_env_start_args( EasyDict( @@ -33,6 +36,7 @@ def setup_env(): os.environ.pop("LIGHTLLM_START_ARGS", None) if hasattr(get_env_start_args, "cache_clear"): get_env_start_args.cache_clear() + monkeypatch.undo() @pytest.fixture(scope="module") diff --git a/unit_tests/server/router/dynamic_prompt/test_radix_cache.py b/unit_tests/server/router/dynamic_prompt/test_radix_cache.py index dfeda0b6f7..bcd2adc155 100644 --- a/unit_tests/server/router/dynamic_prompt/test_radix_cache.py +++ b/unit_tests/server/router/dynamic_prompt/test_radix_cache.py @@ -1,10 +1,19 @@ import pytest import torch from lightllm.server.router.dynamic_prompt.radix_cache import RadixCache +from lightllm.utils import shm_utils + + +@pytest.fixture(scope="module", autouse=True) +def service_name(): + monkeypatch = pytest.MonkeyPatch() + monkeypatch.setattr(shm_utils, "get_unique_server_name", lambda: "test_radix_cache_service_0") + yield + monkeypatch.undo() def test_case1(): - tree = RadixCache("unique_name", 100, 0) + tree = RadixCache(100, 0) ans, _ = tree.insert(torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9], dtype=torch.int64, device="cpu")) assert ans == 0 tree.print_self() @@ -25,7 +34,7 @@ def test_case1(): def test_case2(): - tree = RadixCache("unique_name", 100, 1) + tree = RadixCache(100, 1) ans, _ = tree.insert(torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9], dtype=torch.int64, device="cpu")) ans, _ = tree.insert(torch.tensor([0, 1, 2, 3, 4, 7, 8, 9], dtype=torch.int64, device="cpu")) tree.print_self() @@ -51,7 +60,7 @@ def test_case2(): def test_case3(): - tree = RadixCache("unique_name", 100, 2) + tree = RadixCache(100, 2) ans, _ = tree.insert(torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9], dtype=torch.int64, device="cpu")) ans, _ = tree.insert(torch.tensor([0, 1, 2, 3, 4, 7, 8, 9], dtype=torch.int64, device="cpu")) tree.print_self() @@ -81,7 +90,7 @@ def test_case3(): def test_case4(): - tree = RadixCache("unique_name", 100, 2) + tree = RadixCache(100, 2) ans, _ = tree.insert(torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9], dtype=torch.int64, device="cpu")) ans, _ = tree.insert(torch.tensor([0, 1, 2, 3, 4, 7, 8, 9], dtype=torch.int64, device="cpu")) tree.print_self() @@ -96,7 +105,7 @@ def test_case5(): 测试场景:一个简单的父子节点链 (A -> B),在 ref_counter 都为 0 时,应该成功合并。 """ print("\nTest Case 5: Merging simple parent-child nodes when ref_counter is 0\n") - tree = RadixCache("unique_name", 100, 0) + tree = RadixCache(100, 0) _, node_a = tree.insert(torch.tensor([1, 2, 3], dtype=torch.int64)) _, node_b = tree.insert(torch.tensor([1, 2, 3, 4, 5], dtype=torch.int64)) @@ -125,7 +134,7 @@ def test_case6(): 测试场景:一个长的节点链 (A -> B -> C),在 ref_counter 都为 0 时,应该级联合并成一个节点。 """ print("\nTest Case 6: Merging long nodes when ref_counter is 0\n") - tree = RadixCache("unique_name", 100, 0) + tree = RadixCache(100, 0) _, node_a = tree.insert(torch.tensor([1], dtype=torch.int64)) _, node_b = tree.insert(torch.tensor([1, 2], dtype=torch.int64)) _, node_c = tree.insert(torch.tensor([1, 2, 3, 4], dtype=torch.int64)) @@ -149,7 +158,7 @@ def test_case7(): 测试场景:由于父节点或子节点的 ref_counter > 0,合并不应该发生。 """ print("\nTest Case 7: Merging when parent or child ref_counter > 0\n") - tree = RadixCache("unique_name", 100, 0) + tree = RadixCache(100, 0) _, node_a = tree.insert(torch.tensor([1, 2, 3], dtype=torch.int64)) _, node_b = tree.insert(torch.tensor([1, 2, 3, 4, 5], dtype=torch.int64)) @@ -173,7 +182,7 @@ def test_case8(): 测试场景:由于父节点有多个子节点,合并不应该发生。 """ print("\nTest Case 8: Merging when parent has multiple children\n") - tree = RadixCache("unique_name", 100, 0) + tree = RadixCache(100, 0) _, node_a = tree.insert(torch.tensor([1, 2], dtype=torch.int64)) _, node_b = tree.insert(torch.tensor([1, 2, 3], dtype=torch.int64)) @@ -199,7 +208,7 @@ def test_case9(): 测试场景:在一个复杂的树中,只有满足条件的分支被合并。 """ print("\nTest Case 9: Merging in a complex tree with mixed conditions\n") - tree = RadixCache("unique_name", 100, 0) + tree = RadixCache(100, 0) # 分支1: 可合并的链 A -> B _, node_a = tree.insert(torch.tensor([1, 2], dtype=torch.int64)) @@ -235,7 +244,7 @@ def test_case10(): 测试场景:测试 flush_cache 函数 """ print("\nTest Case 10: Testing flush_cache function\n") - tree = RadixCache("unique_name", 100, 0) + tree = RadixCache(100, 0) tree.insert(torch.tensor([1, 2, 3], dtype=torch.int64)) tree.insert(torch.tensor([1, 2, 3, 4, 5], dtype=torch.int64)) tree_node, size, values = tree.match_prefix( diff --git a/unit_tests/utils/test_atexit_exit_modes.py b/unit_tests/utils/test_atexit_exit_modes.py new file mode 100644 index 0000000000..0883f0efb7 --- /dev/null +++ b/unit_tests/utils/test_atexit_exit_modes.py @@ -0,0 +1,188 @@ +"""用独立进程验证 atexit;所有信号只发送给本测试创建的进程。""" + +import selectors +import signal +import subprocess +import sys +import textwrap + +import pytest + + +pytestmark = pytest.mark.skipif(sys.platform != "linux", reason="Exit codes and signals below target Linux") + +PROBE_SCRIPT = textwrap.dedent( + """ + import atexit + import multiprocessing as mp + import os + import resource + import signal + import sys + + + def record(marker_path, event): + with open(marker_path, "a", encoding="utf-8") as output: + output.write(event + "\\n") + + + def cleanup(marker_path, interrupt=False): + record(marker_path, "started") + if interrupt: + print("READY", flush=True) + while True: + signal.pause() + record(marker_path, "finished") + + + def worker(marker_path): + atexit.register(cleanup, marker_path) + record(marker_path, "worker_registered") + + + def main(): + # SIGABRT/SIGSEGV 实验不生成 core dump。 + resource.setrlimit(resource.RLIMIT_CORE, (0, 0)) + mode, marker_path, signal_number = sys.argv[1:] + signal_number = int(signal_number) + signal.signal(signal.SIGINT, signal.default_int_handler) + + if mode.startswith("multiprocessing_"): + context = mp.get_context(mode.removeprefix("multiprocessing_")) + process = context.Process(target=worker, args=(marker_path,)) + process.start() + process.join(timeout=10) + if process.is_alive(): + process.kill() + process.join(timeout=10) + raise RuntimeError("multiprocessing worker failed to exit") + assert process.exitcode == 0, process.exitcode + return + + atexit.register(cleanup, marker_path, mode == "interrupt_cleanup") + + if mode in ("normal_return", "interrupt_cleanup"): + return + if mode == "sys_exit_0": + sys.exit(0) + if mode == "sys_exit_7": + sys.exit(7) + if mode == "unhandled_exception": + raise RuntimeError("intentional test exception") + if mode == "keyboard_interrupt": + raise KeyboardInterrupt + if mode == "os_exit_0": + os._exit(0) + if mode == "os_exit_7": + os._exit(7) + if mode == "abort": + os.abort() + if mode == "exec_replace": + os.execv(sys.executable, [sys.executable, "-I", "-c", "pass"]) + + if mode == "signal_sys_exit": + signal.signal(signal_number, lambda signum, frame: sys.exit(0)) + elif mode == "signal_os_exit": + signal.signal(signal_number, lambda signum, frame: os._exit(7)) + elif mode == "signal_os_default" or (mode == "signal_default" and signal_number != signal.SIGINT): + if signal_number != signal.SIGKILL: + signal.signal(signal_number, signal.SIG_DFL) + elif mode != "signal_default": + raise ValueError(mode) + + # 父进程收到 READY 后才发信号,确保回调和信号处理方式已注册。 + print("READY", flush=True) + while True: + signal.pause() + + + if __name__ == "__main__": + main() + """ +) + + +@pytest.fixture +def probe_script(tmp_path): + script_path = tmp_path / "atexit_probe.py" + script_path.write_text(PROBE_SCRIPT, encoding="utf-8") + return script_path + + +def run_probe(probe_script, tmp_path, mode, signum=0): + marker_path = tmp_path / "callback_events.txt" + with subprocess.Popen( + [sys.executable, "-I", str(probe_script), mode, str(marker_path), str(int(signum))], + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + start_new_session=True, + ) as process: + try: + if signum: + with selectors.DefaultSelector() as selector: + selector.register(process.stdout, selectors.EVENT_READ) + assert selector.select(timeout=10), "Probe did not become ready" + assert process.stdout.readline().strip() == "READY", "Probe exited before registering its callback" + process.send_signal(signum) + stdout, stderr = process.communicate(timeout=30) + finally: + if process.poll() is None: + process.kill() + process.communicate(timeout=10) + + events = marker_path.read_text(encoding="utf-8").splitlines() if marker_path.exists() else [] + print(f"{mode} signal={int(signum)} returncode={process.returncode} events={events}") + return process.returncode, events, stdout + stderr + + +@pytest.mark.parametrize( + "mode,signum,expected_returncode,callback_completed", + [ + pytest.param("normal_return", 0, 0, True, id="normal-return"), + pytest.param("sys_exit_0", 0, 0, True, id="sys-exit-0"), + pytest.param("sys_exit_7", 0, 7, True, id="sys-exit-nonzero"), + pytest.param("unhandled_exception", 0, 1, True, id="unhandled-exception"), + pytest.param("keyboard_interrupt", 0, -signal.SIGINT, True, id="unhandled-keyboard-interrupt"), + pytest.param("signal_default", signal.SIGINT, -signal.SIGINT, True, id="sigint-python-default"), + pytest.param("signal_os_default", signal.SIGINT, -signal.SIGINT, False, id="sigint-os-default"), + pytest.param("signal_default", signal.SIGTERM, -signal.SIGTERM, False, id="sigterm-default"), + pytest.param("signal_default", signal.SIGHUP, -signal.SIGHUP, False, id="sighup-default"), + pytest.param("signal_default", signal.SIGQUIT, -signal.SIGQUIT, False, id="sigquit-default"), + pytest.param("signal_default", signal.SIGSEGV, -signal.SIGSEGV, False, id="sigsegv-default"), + pytest.param("signal_default", signal.SIGKILL, -signal.SIGKILL, False, id="sigkill"), + pytest.param("signal_sys_exit", signal.SIGINT, 0, True, id="sigint-handler-sys-exit"), + pytest.param("signal_sys_exit", signal.SIGTERM, 0, True, id="sigterm-handler-sys-exit"), + pytest.param("signal_sys_exit", signal.SIGHUP, 0, True, id="sighup-handler-sys-exit"), + pytest.param("signal_os_exit", signal.SIGTERM, 7, False, id="sigterm-handler-os-exit"), + pytest.param("os_exit_0", 0, 0, False, id="os-exit-0"), + pytest.param("os_exit_7", 0, 7, False, id="os-exit-nonzero"), + pytest.param("abort", 0, -signal.SIGABRT, False, id="abort"), + pytest.param("exec_replace", 0, 0, False, id="exec-replaces-interpreter"), + ], +) +def test_atexit_on_process_exit(probe_script, tmp_path, mode, signum, expected_returncode, callback_completed): + returncode, events, output = run_probe(probe_script, tmp_path, mode, signum) + + assert returncode == expected_returncode, output + assert events == (["started", "finished"] if callback_completed else []), output + + +@pytest.mark.skipif( + sys.version_info[:2] != (3, 10), reason="This experiment records Python 3.10 multiprocessing behavior" +) +@pytest.mark.parametrize("start_method,callback_completed", [("spawn", True), ("fork", False)]) +def test_atexit_in_multiprocessing_worker(probe_script, tmp_path, start_method, callback_completed): + returncode, events, output = run_probe(probe_script, tmp_path, f"multiprocessing_{start_method}") + + assert returncode == 0, output + # 本机 Python 3.10:spawn_main() 使用 sys.exit(),fork 的 _launch() 使用 os._exit()。 + expected_events = ["worker_registered"] + (["started", "finished"] if callback_completed else []) + assert events == expected_events, output + + +@pytest.mark.parametrize("signum", [signal.SIGINT, signal.SIGKILL]) +def test_signal_can_interrupt_atexit_callback(probe_script, tmp_path, signum): + _returncode, events, output = run_probe(probe_script, tmp_path, "interrupt_cleanup", signum) + + assert events == ["started"], output diff --git a/unit_tests/utils/test_service_shm_cleanup.py b/unit_tests/utils/test_service_shm_cleanup.py new file mode 100644 index 0000000000..d9fecc1360 --- /dev/null +++ b/unit_tests/utils/test_service_shm_cleanup.py @@ -0,0 +1,122 @@ +import json + +from lightllm.utils import service_shm_cleanup + + +def test_cleanup_service_shm_only_removes_matching_service(monkeypatch, tmp_path): + shm_dir = tmp_path / "shm" + shm_dir.mkdir() + matching_names = ["service_0_req_pool", "service_0_token_load"] + for name in [*matching_names, "service_1_req_pool", "other_service_0_value"]: + (shm_dir / name).touch() + + removed_system_v_keys = [] + monkeypatch.setattr(service_shm_cleanup, "SHM_DIR", shm_dir) + monkeypatch.setattr( + service_shm_cleanup.ServiceShmCleanup, + "cleanup_system_v_shm", + staticmethod(lambda keys: removed_system_v_keys.extend(keys) or len(keys)), + ) + + service_shm_cleanup.ServiceShmCleanup.cleanup_posix_shm("service_0") + service_shm_cleanup.ServiceShmCleanup.cleanup_system_v_shm([101, 102]) + + assert all(not (shm_dir / name).exists() for name in matching_names) + assert (shm_dir / "service_1_req_pool").exists() + assert (shm_dir / "other_service_0_value").exists() + assert removed_system_v_keys == [101, 102] + + +def test_system_v_shm_keys_follow_feature_switches(monkeypatch): + start_args = { + "run_mode": "prefill", + "enable_cpu_cache": True, + "enable_multimodal": False, + "cpu_kv_cache_shm_id": 21, + "multi_modal_cache_shm_id": 22, + } + removed_system_v_keys = [] + monkeypatch.setenv("LIGHTLLM_START_ARGS", json.dumps(start_args)) + monkeypatch.setattr( + service_shm_cleanup.ServiceShmCleanup, + "cleanup_posix_shm", + staticmethod(lambda service_name: 0), + ) + monkeypatch.setattr( + service_shm_cleanup.ServiceShmCleanup, + "cleanup_system_v_shm", + staticmethod(lambda keys: removed_system_v_keys.extend(keys) or len(keys)), + ) + + service_shm_cleanup.ServiceShmCleanup("current_service_0").cleanup_service_resources() + + assert removed_system_v_keys == [21] + + +def test_non_inference_mode_skips_system_v_shm_cleanup(monkeypatch): + start_args = { + "run_mode": "visual_only", + "enable_cpu_cache": True, + "enable_multimodal": True, + "cpu_kv_cache_shm_id": 21, + "multi_modal_cache_shm_id": 22, + } + system_v_cleanup_calls = [] + monkeypatch.setenv("LIGHTLLM_START_ARGS", json.dumps(start_args)) + monkeypatch.setattr( + service_shm_cleanup.ServiceShmCleanup, + "cleanup_posix_shm", + staticmethod(lambda service_name: 0), + ) + monkeypatch.setattr( + service_shm_cleanup.ServiceShmCleanup, + "cleanup_system_v_shm", + staticmethod(lambda keys: system_v_cleanup_calls.append(keys) or 0), + ) + + service_shm_cleanup.ServiceShmCleanup("current_service_0").cleanup_service_resources() + + assert system_v_cleanup_calls == [] + + +def test_register_launcher_cleanup_uses_current_start_args(monkeypatch): + cleanup_calls = [] + atexit_callbacks = [] + monkeypatch.setattr( + service_shm_cleanup.ServiceShmCleanup, + "cleanup_posix_shm", + staticmethod(lambda service_name: cleanup_calls.append(("posix", service_name)) or 0), + ) + monkeypatch.setattr( + service_shm_cleanup.ServiceShmCleanup, + "cleanup_system_v_shm", + staticmethod(lambda keys: cleanup_calls.append(("system_v", keys)) or 0), + ) + monkeypatch.setattr(service_shm_cleanup.atexit, "register", atexit_callbacks.append) + start_args = { + "run_mode": "normal", + "model_dir": "/models/test", + "tp": 2, + "enable_cpu_cache": True, + "enable_multimodal": True, + "cpu_kv_cache_shm_id": 21, + "multi_modal_cache_shm_id": 22, + } + monkeypatch.setenv( + "LIGHTLLM_START_ARGS", + json.dumps(start_args), + ) + + cleanup = service_shm_cleanup.register_launcher_shm_cleanup("current_service_0") + + assert cleanup_calls == [] + assert atexit_callbacks == [cleanup] + + cleanup() + cleanup() + assert cleanup_calls == [ + ("posix", "current_service_0"), + ("system_v", [21, 22]), + ("posix", "current_service_0"), + ("system_v", [21, 22]), + ] diff --git a/unit_tests/utils/test_shm_utils.py b/unit_tests/utils/test_shm_utils.py new file mode 100644 index 0000000000..45186fdce4 --- /dev/null +++ b/unit_tests/utils/test_shm_utils.py @@ -0,0 +1,31 @@ +import pytest + +from lightllm.utils import shm_utils + + +def test_get_service_shm_name_requires_service_name(monkeypatch): + monkeypatch.setattr(shm_utils, "get_unique_server_name", lambda: None) + + with pytest.raises(RuntimeError, match="LIGHTLLM_UNIQUE_SERVICE_NAME_ID is unset"): + shm_utils.get_service_shm_name("req_pool") + + +def test_get_service_shm_name_adds_prefix_once(monkeypatch): + monkeypatch.setattr(shm_utils, "get_unique_server_name", lambda: "service_uuid_0") + + assert shm_utils.get_service_shm_name("req_pool") == "service_uuid_0_req_pool" + assert shm_utils.get_service_shm_name("service_uuid_0_req_pool") == "service_uuid_0_req_pool" + + +def test_create_or_link_shm_passes_scoped_name_to_shared_memory_layer(monkeypatch): + created_names = [] + monkeypatch.setattr(shm_utils, "get_unique_server_name", lambda: "service_uuid_1") + monkeypatch.setattr( + shm_utils, + "_force_create_shm", + lambda name, expected_size: created_names.append((name, expected_size)) or object(), + ) + + shm_utils.create_or_link_shm("token_load", 128, force_mode="create") + + assert created_names == [("service_uuid_1_token_load", 128)] diff --git a/unit_tests/utils/test_start_utils.py b/unit_tests/utils/test_start_utils.py index 1ac147c1a5..250c41833f 100644 --- a/unit_tests/utils/test_start_utils.py +++ b/unit_tests/utils/test_start_utils.py @@ -1,3 +1,5 @@ +from types import SimpleNamespace + import pytest from lightllm.utils import start_utils @@ -48,6 +50,8 @@ def wait(self, timeout=None): def test_start_submodule_processes_returns_and_manages_psutil_processes(monkeypatch): class FakePipeReader: def recv(self): + # 子进程应在等待初始化结果之前就进入 manager,保证此时 Ctrl-C 可以清理它们。 + assert len(process_manager.processes) == 2 return "init ok" class FakeMpProcess: @@ -85,7 +89,7 @@ def is_alive(self): } -def test_register_process_tree_adds_recursive_descendants(): +def test_register_process_tree_adds_recursive_descendants(monkeypatch): descendants = [ FakeProcess(pid=1001, name="lightllm::model_infer"), FakeProcess(pid=1002, name="lightllm::pd_manager"), @@ -93,6 +97,7 @@ def test_register_process_tree_adds_recursive_descendants(): ] router_process = FakeProcess(pid=1000, children=descendants) process_manager = start_utils.SubmoduleManager() + monkeypatch.setattr(start_utils, "is_process_active", lambda pid: True) process_manager.register_process_tree(router_process) @@ -104,12 +109,13 @@ def test_register_process_tree_adds_recursive_descendants(): } -def test_register_process_tree_filters_short_lived_helper_processes(): +def test_register_process_tree_filters_short_lived_helper_processes(monkeypatch): model_process = FakeProcess(pid=1001, name="lightllm::model_infer") compile_worker = FakeProcess(pid=1002, name="python") pd_process = FakeProcess(pid=1003, name="lightllm::decode_trans") router_process = FakeProcess(pid=1000, children=[model_process, compile_worker, pd_process]) process_manager = start_utils.SubmoduleManager() + monkeypatch.setattr(start_utils, "is_process_active", lambda pid: True) process_manager.register_process_tree(router_process) @@ -135,11 +141,19 @@ def name(self): assert process_manager.process_names == {} -def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): +@pytest.mark.parametrize("initialize_exit_controller", [False, True]) +def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch, initialize_exit_controller): http_server_process = FakeHttpServerProcess() process_manager = start_utils.SubmoduleManager() registered_handlers = {} terminate_calls = [] + cleanup_calls = [] + monkeypatch.setattr(start_utils, "get_unique_server_name", lambda: "service_0") + monkeypatch.setattr( + start_utils, + "register_launcher_shm_cleanup", + lambda service_name: cleanup_calls.append(("register", service_name)) or (lambda: None), + ) monkeypatch.setattr( start_utils.signal, "signal", @@ -147,6 +161,9 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): ) monkeypatch.setattr(process_manager, "terminate_all_processes", lambda: terminate_calls.append(True)) + if initialize_exit_controller: + process_manager.setup_exit_controller() + startup_handlers = registered_handlers.copy() process_manager.setup_signal_handlers(http_server_process) assert set(registered_handlers) == { @@ -154,6 +171,12 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): start_utils.signal.SIGINT, start_utils.signal.SIGHUP, } + assert cleanup_calls == ([("register", "service_0")] if initialize_exit_controller else []) + if initialize_exit_controller: + assert all(registered_handlers[sig] is not handler for sig, handler in startup_handlers.items()) + runtime_handlers = registered_handlers.copy() + process_manager.setup_exit_controller() + assert registered_handlers == runtime_handlers with pytest.raises(SystemExit) as exc_info: registered_handlers[start_utils.signal.SIGTERM](start_utils.signal.SIGTERM, None) @@ -163,6 +186,54 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): assert terminate_calls == [True] +@pytest.mark.parametrize( + "shutdown_signal", [start_utils.signal.SIGTERM, start_utils.signal.SIGINT, start_utils.signal.SIGHUP] +) +def test_setup_exit_controller_registers_once_and_cleans_up_on_signal(monkeypatch, shutdown_signal): + from lightllm.utils import envs_utils + + process_manager = start_utils.SubmoduleManager() + registration_calls = [] + registered_handlers = {} + shutdown_events = [] + managed_process = FakeProcess(pid=1234) + managed_process.kill = lambda: shutdown_events.append("kill") + managed_process.wait = lambda: shutdown_events.append("wait") + process_manager.processes = [managed_process] + monkeypatch.setattr(start_utils.psutil, "Process", lambda pid: managed_process) + monkeypatch.setattr(start_utils, "get_unique_server_name", lambda: "service_0") + monkeypatch.setattr( + start_utils, + "register_launcher_shm_cleanup", + lambda service_name: registration_calls.append(service_name) or (lambda: shutdown_events.append("cleanup")), + ) + monkeypatch.setattr( + start_utils.signal, + "signal", + lambda sig, handler: registered_handlers.__setitem__(sig, handler), + ) + monkeypatch.setattr(envs_utils, "get_env_start_args", lambda: SimpleNamespace(enable_mps=False)) + + process_manager.setup_exit_controller() + initial_handlers = registered_handlers.copy() + process_manager.setup_exit_controller() + + assert registration_calls == ["service_0"] + assert set(registered_handlers) == { + start_utils.signal.SIGTERM, + start_utils.signal.SIGINT, + start_utils.signal.SIGHUP, + } + assert registered_handlers == initial_handlers + assert shutdown_events == [] + + with pytest.raises(SystemExit) as exc_info: + registered_handlers[shutdown_signal](shutdown_signal, None) + + assert exc_info.value.code == 0 + assert shutdown_events == ["kill", "wait", "cleanup"] + + def test_supervisor_fails_when_http_server_exits(monkeypatch): http_server_process = FakeHttpServerProcess(return_code=0) process_manager = start_utils.SubmoduleManager()