From 45c225b540271518030a610996da287f59423696 Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Thu, 3 Sep 2026 13:32:17 +0800 Subject: [PATCH 1/8] fix: ctrl c for multi node --- lightllm/server/api_start.py | 12 +- lightllm/utils/auto_shm_cleanup.py | 6 +- lightllm/utils/start_utils.py | 266 +++++++++++++-------- unit_tests/utils/test_start_utils.py | 332 +++++++++++++++++++++++++-- 4 files changed, 503 insertions(+), 113 deletions(-) diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 7c7ac9fe48..1676313371 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -424,7 +424,9 @@ def _hypercorn_config_args(args: StartArgs): def normal_or_p_d_start(args: StartArgs): - process_manager = _launch_subprocesses(args) + # Install this before _launch_subprocesses creates multiprocessing children. + process_manager.setup_signal_handlers() + _launch_subprocesses(args) # 启动 Hypercorn command = [ @@ -445,6 +447,7 @@ def normal_or_p_d_start(args: StartArgs): # 启动子进程 http_server_process = subprocess.Popen(command) + process_manager.setup_signal_handlers(http_server_process) if "s3://" in args.model_dir: from lightllm.utils.petrel_helper import s3_model_clear @@ -455,11 +458,11 @@ def normal_or_p_d_start(args: StartArgs): from lightllm.server.health_monitor.manager import start_health_check_process process_manager.start_submodule_processes(start_funcs=[start_health_check_process], start_args=[(args,)]) - process_manager.setup_signal_handlers(http_server_process) process_manager.supervise_processes(http_server_process) def pd_master_start(args: StartArgs): + process_manager.setup_signal_handlers() _set_envs_and_config(args) set_unique_server_name(args) if args.run_mode != "pd_master": @@ -507,13 +510,13 @@ def pd_master_start(args: StartArgs): ] http_server_process = subprocess.Popen(command) + process_manager.setup_signal_handlers(http_server_process) if args.health_monitor: from lightllm.server.health_monitor.manager import start_health_check_process process_manager.start_submodule_processes(start_funcs=[start_health_check_process], start_args=[(args,)]) - process_manager.setup_signal_handlers(http_server_process) process_manager.supervise_processes(http_server_process) @@ -521,6 +524,7 @@ def visual_only_start(args): from lightllm.server.core.objs.start_args_type import StartArgs args: StartArgs = args + process_manager.setup_signal_handlers() _set_envs_and_config(args) if args.afs_image_embed_dir is not None: os.makedirs(args.afs_image_embed_dir, mode=0o777, exist_ok=True) @@ -558,11 +562,11 @@ def visual_only_start(args): (args,), ], ) - process_manager.setup_signal_handlers() process_manager.supervise_processes() def config_server_start(args): + process_manager.setup_signal_handlers() set_unique_server_name(args) if args.run_mode != "config_server": return diff --git a/lightllm/utils/auto_shm_cleanup.py b/lightllm/utils/auto_shm_cleanup.py index 2417fef085..95d118d838 100644 --- a/lightllm/utils/auto_shm_cleanup.py +++ b/lightllm/utils/auto_shm_cleanup.py @@ -40,7 +40,7 @@ def _init_libc(self): self.libc = None def _register_handlers_for_cleanup(self): - atexit.register(self._cleanup) + atexit.register(self.cleanup) self.register_signal_handlers() def register_signal_handlers(self): @@ -51,13 +51,13 @@ def register_signal_handlers(self): self.signal_handlers_registered = True def _signal_cleanup_handler(self, signum, frame): - self._cleanup() + self.cleanup() parent = psutil.Process(os.getpid()) # 递归拿到所有子进程并终止 for ch in parent.children(recursive=True): ch.kill() - def _cleanup(self): + def cleanup(self): """清理:System V 执行 IPC_RMID,POSIX 执行 unlink。""" removed_sysv = 0 IPC_RMID = 0 diff --git a/lightllm/utils/start_utils.py b/lightllm/utils/start_utils.py index 65b6f5f41b..d6b6c0bbb2 100644 --- a/lightllm/utils/start_utils.py +++ b/lightllm/utils/start_utils.py @@ -11,42 +11,67 @@ logger = init_logger(__name__) +# Waiting for an unrelated/re-parented zombie can otherwise block forever. +PROCESS_SHUTDOWN_WAIT_TIMEOUT_SECONDS = 5 + + class SubmoduleManager: def __init__(self): self.processes = [] self.process_names = {} + self.http_server_process = None + self._handling_signal = False 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) - process = mp.Process( - target=start_func, - args=start_arg + (pipe_writer,), - ) - process.start() - pipe_readers.append(pipe_reader) - processes.append(process) - - # Wait for all processes to initialize - for index, pipe_reader in enumerate(pipe_readers): - init_state = pipe_reader.recv() - if init_state != "init ok": - logger.error(f"init func {start_funcs[index].__name__} : {str(init_state)}") - for proc in processes: - proc.kill() - sys.exit(1) - else: + try: + for start_func, start_arg in zip(start_funcs, start_args): + pipe_reader, pipe_writer = mp.Pipe(duplex=False) + process = mp.Process( + target=start_func, + args=start_arg + (pipe_writer,), + ) + pipe_readers.append(pipe_reader) + processes.append(process) + try: + process.start() + # Register before waiting for initialization so Ctrl-C can clean + # processes which are still starting up. + managed_process = psutil.Process(process.pid) + self.processes.append(managed_process) + self.process_names[managed_process] = managed_process.name() + managed_processes.append(managed_process) + finally: + pipe_writer.close() + + # Wait for all processes to initialize. + for index, pipe_reader in enumerate(pipe_readers): + init_state = pipe_reader.recv() + if init_state != "init ok": + logger.error(f"init func {start_funcs[index].__name__} : {str(init_state)}") + raise SystemExit(1) 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 + assert all(process.is_alive() for process in processes) + return managed_processes + except BaseException: + # recv() may be interrupted by Ctrl-C or raise EOFError when a child + # dies. All successfully-started children have already been managed. + try: + self.terminate_all_processes() + except Exception: + logger.exception("Failed to clean up submodules after initialization failure") + raise + finally: + for pipe_reader in pipe_readers: + try: + pipe_reader.close() + except (OSError, EOFError): + pass def register_process_tree(self, root_process): """Add persistent LightLLM descendants to supervision. @@ -71,70 +96,108 @@ def register_process_tree(self, root_process): self.processes.append(process) self.process_names[process] = process_name - def terminate_all_processes(self): - from lightllm.utils.envs_utils import get_env_start_args + def terminate_all_processes(self, wait_timeout=PROCESS_SHUTDOWN_WAIT_TIMEOUT_SECONDS): + """Kill all managed local process trees without indefinitely waiting.""" + processes_by_pid = {} + for process in self.processes: + for tree_process in _get_process_tree(process): + processes_by_pid.setdefault(tree_process.pid, tree_process) + + processes_to_wait_for = list(processes_by_pid.values()) + _kill_processes(processes_to_wait_for) - def kill_recursive(proc): + if processes_to_wait_for and wait_timeout > 0: try: - parent = psutil.Process(proc.pid) - children = parent.children(recursive=True) - for child in children: - logger.info(f"Killing child process {child.pid}") - child.kill() - logger.info(f"Killing parent process {proc.pid}") - parent.kill() - except psutil.NoSuchProcess: - logger.warning(f"Process {proc.pid} does not exist.") - - for proc in self.processes: - if proc.is_running(): - kill_recursive(proc) - proc.wait() - - # recover the gpu compute mode - is_enable_mps = get_env_start_args().enable_mps - if is_enable_mps: - from lightllm.utils.device_utils import stop_mps - - stop_mps() + _gone, alive = psutil.wait_procs(processes_to_wait_for, timeout=wait_timeout) + if alive: + logger.warning( + "Timed out waiting for processes to exit: %s", + ", ".join(str(process.pid) for process in alive), + ) + except (psutil.NoSuchProcess, psutil.ZombieProcess, psutil.AccessDenied): + # A process may disappear between kill() and wait_procs(). + pass + except Exception: + logger.exception("Failed while waiting for submodule processes to exit") + + # Recover GPU compute mode, but failure here must not prevent launcher exit. + try: + from lightllm.utils.envs_utils import get_env_start_args + + is_enable_mps = get_env_start_args().enable_mps + if is_enable_mps: + from lightllm.utils.device_utils import stop_mps + + stop_mps() + except Exception: + logger.exception("Failed to restore GPU compute mode during shutdown") logger.info("All processes terminated gracefully.") 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...") - if http_server_process is not None: - kill_recursive(http_server_process) + from lightllm.utils.auto_shm_cleanup import get_auto_cleanup - self.terminate_all_processes() - logger.info("All processes have been forcefully terminated.") - sys.exit(0) - - if sig == signal.SIGTERM: - logger.info("Received SIGTERM, shutting down gracefully...") - else: - logger.info("Received SIGHUP (terminal closed), shutting down gracefully...") + # Initialize shared-memory signal handling before the launcher takes over. + shm_cleanup = get_auto_cleanup() + # The installed closure deliberately reads this field at signal time: + # handlers are installed before submodules are launched, while Hypercorn + # is only available later in the startup sequence. + if http_server_process is not None: + self.http_server_process = http_server_process - if http_server_process is not None and http_server_process.poll() is None: - http_server_process.send_signal(signal.SIGTERM) - try: - http_server_process.wait(timeout=60) + def signal_handler(sig, _frame): + repeated_signal = self._handling_signal + self._handling_signal = True + try: + if repeated_signal: + logger.warning("Received a second shutdown signal; forcing immediate cleanup") + elif sig == signal.SIGINT: + logger.info("Received SIGINT (Ctrl+C), forcing immediate exit...") + elif sig == signal.SIGTERM: + logger.info("Received SIGTERM, shutting down gracefully...") + else: + logger.info("Received SIGHUP (terminal closed), shutting down gracefully...") + + if ( + not repeated_signal + and sig != signal.SIGINT + and self.http_server_process is not None + and self.http_server_process.poll() is None + ): + self.http_server_process.send_signal(signal.SIGTERM) + self.http_server_process.wait(timeout=60) logger.info("HTTP server exited gracefully") - except subprocess.TimeoutExpired: - logger.warning("HTTP server did not exit in time, killing it...") - kill_recursive(http_server_process) - - self.terminate_all_processes() - logger.info("All processes have been terminated gracefully.") - sys.exit(0) + except subprocess.TimeoutExpired: + logger.warning("HTTP server did not exit in time, killing it...") + except Exception: + logger.exception("Shutdown cleanup failed") + finally: + try: + for shutdown_signal in (signal.SIGTERM, signal.SIGINT, signal.SIGHUP): + signal.signal(shutdown_signal, signal.SIG_IGN) + self._cleanup_processes( + self.http_server_process, + wait_timeout=0 if repeated_signal else PROCESS_SHUTDOWN_WAIT_TIMEOUT_SECONDS, + ) + logger.info("All processes have been terminated.") + except Exception: + logger.exception("Failed to clean up processes during shutdown") + finally: + try: + shm_cleanup.cleanup() + except Exception: + logger.exception("Failed to clean up shared memory during shutdown") + finally: + # Do not let multiprocessing's atexit handler re-join a stuck + # child after Ctrl-C. Cleanup above is intentionally best effort. + os._exit(1 if repeated_signal else 0) signal.signal(signal.SIGTERM, signal_handler) signal.signal(signal.SIGINT, signal_handler) signal.signal(signal.SIGHUP, signal_handler) logger.info(f"start process pid {os.getpid()}") - if http_server_process is not None: - logger.info(f"http server pid {http_server_process.pid}") + if self.http_server_process is not None: + logger.info(f"http server pid {self.http_server_process.pid}") def supervise_processes(self, http_server_process=None): """Watch the HTTP server, when present, and all registered submodules. @@ -151,7 +214,7 @@ def supervise_processes(self, http_server_process=None): if http_return_code is not None: message = f"HTTP server exited unexpectedly with return code {http_return_code}" logger.error(message) - self._cleanup_after_process_failure(http_server_process) + self._cleanup_processes(http_server_process) raise RuntimeError(message) dead_processes = [ @@ -170,21 +233,21 @@ def supervise_processes(self, http_server_process=None): dead_process_descriptions = ", ".join(dead_process_descriptions) message = f"Critical LightLLM submodule exited unexpectedly: {dead_process_descriptions}" logger.error(message) - self._cleanup_after_process_failure(http_server_process) + self._cleanup_processes(http_server_process) raise RuntimeError(message) time.sleep(supervisor_interval_seconds) - def _cleanup_after_process_failure(self, http_server_process): - """Best-effort cleanup before the launcher exits with a failure.""" - if http_server_process is not None and http_server_process.poll() is None: - try: + def _cleanup_processes(self, http_server_process, wait_timeout=PROCESS_SHUTDOWN_WAIT_TIMEOUT_SECONDS): + """Best-effort cleanup before the launcher exits.""" + try: + if http_server_process is not None and http_server_process.poll() is None: kill_recursive(http_server_process) - except Exception: - logger.exception("Failed to terminate the HTTP server process tree") + except Exception: + logger.exception("Failed to terminate the HTTP server process tree") try: - self.terminate_all_processes() + self.terminate_all_processes(wait_timeout=wait_timeout) except Exception: logger.exception("Failed to terminate all LightLLM submodule processes") @@ -218,17 +281,34 @@ def start_submodule_processes(start_funcs=[], start_args=[]): return -def kill_recursive(proc): +def _get_process_tree(root_process): try: - parent = psutil.Process(proc.pid) - children = parent.children(recursive=True) - for child in children: - logger.info(f"Killing child process {child.pid}") - child.kill() - logger.info(f"Killing parent process {proc.pid}") - parent.kill() - except psutil.NoSuchProcess: - logger.warning(f"Process {proc.pid} does not exist.") + root = psutil.Process(root_process.pid) + descendants = root.children(recursive=True) + except (psutil.NoSuchProcess, psutil.ZombieProcess, psutil.AccessDenied): + return [] + + processes = [] + for process in list(reversed(descendants)) + [root]: + try: + if is_process_active(process.pid): + processes.append(process) + except (psutil.NoSuchProcess, psutil.ZombieProcess, psutil.AccessDenied): + continue + return processes + + +def _kill_processes(processes): + for process in processes: + try: + logger.info(f"Killing process {process.pid}") + process.kill() + except (psutil.NoSuchProcess, psutil.ZombieProcess, psutil.AccessDenied): + continue + + +def kill_recursive(proc): + _kill_processes(_get_process_tree(proc)) process_manager = SubmoduleManager() diff --git a/unit_tests/utils/test_start_utils.py b/unit_tests/utils/test_start_utils.py index 1ac147c1a5..b0ec7a4f65 100644 --- a/unit_tests/utils/test_start_utils.py +++ b/unit_tests/utils/test_start_utils.py @@ -1,6 +1,15 @@ +from types import SimpleNamespace + import pytest -from lightllm.utils import start_utils +from lightllm.utils import auto_shm_cleanup, start_utils + + +@pytest.fixture(autouse=True) +def mock_auto_shm_cleanup(monkeypatch): + monkeypatch.setattr(auto_shm_cleanup.AutoShmCleanup, "_init_libc", lambda self: None) + monkeypatch.setattr(auto_shm_cleanup.atexit, "register", lambda *args, **kwargs: None) + monkeypatch.setattr(auto_shm_cleanup, "_auto_cleanup", None) class FakeHttpServerProcess: @@ -29,6 +38,7 @@ def __init__(self, pid, running=True, children=None, name="process", exitcode=No self._name = name self.exitcode = exitcode self.wait_timeout = wait_timeout + self.kill_calls = 0 def is_running(self): return self.running @@ -44,12 +54,22 @@ def wait(self, timeout=None): raise start_utils.psutil.TimeoutExpired(timeout, pid=self.pid, name=self._name) return self.exitcode + def kill(self): + self.kill_calls += 1 + def test_start_submodule_processes_returns_and_manages_psutil_processes(monkeypatch): class FakePipeReader: def recv(self): return "init ok" + def close(self): + pass + + class FakePipeWriter: + def close(self): + pass + class FakeMpProcess: next_pid = 1000 @@ -63,7 +83,7 @@ def start(self): def is_alive(self): return True - monkeypatch.setattr(start_utils.mp, "Pipe", lambda duplex: (FakePipeReader(), object())) + monkeypatch.setattr(start_utils.mp, "Pipe", lambda duplex: (FakePipeReader(), FakePipeWriter())) monkeypatch.setattr(start_utils.mp, "Process", FakeMpProcess) monkeypatch.setattr( start_utils.psutil, @@ -85,7 +105,46 @@ def is_alive(self): } -def test_register_process_tree_adds_recursive_descendants(): +def test_start_submodule_processes_cleans_up_after_recv_error(monkeypatch): + class FakePipeReader: + def recv(self): + raise EOFError + + def close(self): + pass + + class FakePipeWriter: + def close(self): + pass + + class FakeMpProcess: + pid = 1000 + + def __init__(self, target, args): + pass + + def start(self): + pass + + managed_process = FakeProcess(pid=1000, name="process-1000") + process_manager = start_utils.SubmoduleManager() + cleanup_process_pids = [] + monkeypatch.setattr(start_utils.mp, "Pipe", lambda duplex: (FakePipeReader(), FakePipeWriter())) + monkeypatch.setattr(start_utils.mp, "Process", FakeMpProcess) + monkeypatch.setattr(start_utils.psutil, "Process", lambda pid: managed_process) + monkeypatch.setattr( + process_manager, + "terminate_all_processes", + lambda **_: cleanup_process_pids.extend(process.pid for process in process_manager.processes), + ) + + with pytest.raises(EOFError): + process_manager.start_submodule_processes(start_funcs=[lambda pipe_writer: None], start_args=[()]) + + assert cleanup_process_pids == [1000] + + +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 +152,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 +164,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) @@ -145,7 +206,13 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): "signal", lambda sig, handler: registered_handlers.__setitem__(sig, handler), ) - monkeypatch.setattr(process_manager, "terminate_all_processes", lambda: terminate_calls.append(True)) + monkeypatch.setattr( + process_manager, + "terminate_all_processes", + lambda **_: terminate_calls.append(True), + ) + exit_codes = [] + monkeypatch.setattr(start_utils.os, "_exit", lambda code: exit_codes.append(code)) process_manager.setup_signal_handlers(http_server_process) @@ -154,13 +221,232 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): start_utils.signal.SIGINT, start_utils.signal.SIGHUP, } - with pytest.raises(SystemExit) as exc_info: - registered_handlers[start_utils.signal.SIGTERM](start_utils.signal.SIGTERM, None) + registered_handlers[start_utils.signal.SIGTERM](start_utils.signal.SIGTERM, None) - assert exc_info.value.code == 0 assert http_server_process.sent_signals == [start_utils.signal.SIGTERM] assert http_server_process.wait_timeouts == [60] assert terminate_calls == [True] + assert exit_codes == [0] + + +@pytest.mark.parametrize( + "signal_to_send,repeated_signal,expected_exit_code", + [(start_utils.signal.SIGINT, False, 0), (start_utils.signal.SIGTERM, True, 1)], +) +def test_setup_signal_handlers_uses_http_process_set_after_handler_install( + monkeypatch, signal_to_send, repeated_signal, expected_exit_code +): + process_manager = start_utils.SubmoduleManager() + http_server_process = FakeHttpServerProcess() + registered_handlers = {} + killed_processes = [] + exit_codes = [] + monkeypatch.setattr( + start_utils.signal, + "signal", + lambda sig, handler: registered_handlers.__setitem__(sig, handler), + ) + monkeypatch.setattr(start_utils, "kill_recursive", lambda process: killed_processes.append(process)) + terminate_calls = [] + monkeypatch.setattr(process_manager, "terminate_all_processes", lambda **kwargs: terminate_calls.append(kwargs)) + monkeypatch.setattr(start_utils.os, "_exit", lambda code: exit_codes.append(code)) + + process_manager.setup_signal_handlers() + process_manager.setup_signal_handlers(http_server_process) + process_manager._handling_signal = repeated_signal + registered_handlers[signal_to_send](signal_to_send, None) + + assert killed_processes == [http_server_process] + assert terminate_calls == [ + {"wait_timeout": 0 if repeated_signal else start_utils.PROCESS_SHUTDOWN_WAIT_TIMEOUT_SECONDS} + ] + assert exit_codes == [expected_exit_code] + if repeated_signal: + assert http_server_process.sent_signals == [] + assert http_server_process.wait_timeouts == [] + + +def test_launcher_handler_remains_registered_after_shared_memory_registration(monkeypatch): + handlers = {} + exit_codes = [] + manager = start_utils.SubmoduleManager() + monkeypatch.setattr( + start_utils.signal, + "signal", + lambda sig, handler: handlers.__setitem__(sig, handler), + ) + monkeypatch.setattr(manager, "terminate_all_processes", lambda **_: None) + monkeypatch.setattr(start_utils.os, "_exit", lambda code: exit_codes.append(code)) + + manager.setup_signal_handlers() + launcher_handler = handlers[start_utils.signal.SIGINT] + cleaner = auto_shm_cleanup.get_auto_cleanup() + monkeypatch.setattr(cleaner, "cleanup", lambda: None) + auto_shm_cleanup.register_posix_shm_for_cleanup("test") + + assert handlers[start_utils.signal.SIGINT] is launcher_handler + launcher_handler(start_utils.signal.SIGINT, None) + assert exit_codes == [0] + + +@pytest.mark.parametrize("failure", ["poll", "send", "wait", "timeout", "kill"]) +def test_http_shutdown_failure_still_cleans_processes_and_shared_memory(monkeypatch, failure): + class FailingHttp: + pid = 4321 + + def poll(self): + if failure == "poll": + raise PermissionError("poll denied") + return None + + def send_signal(self, _sig): + if failure == "send": + raise PermissionError("send denied") + + def wait(self, timeout=None): + if failure == "wait": + raise PermissionError("wait denied") + if failure == "timeout": + raise start_utils.subprocess.TimeoutExpired("http", timeout) + + handlers = {} + terminate_calls = [] + kill_calls = [] + exit_codes = [] + manager = start_utils.SubmoduleManager() + monkeypatch.setattr( + start_utils.signal, + "signal", + lambda sig, handler: handlers.__setitem__(sig, handler), + ) + monkeypatch.setattr(manager, "terminate_all_processes", lambda **kwargs: terminate_calls.append(kwargs)) + monkeypatch.setattr( + start_utils, + "kill_recursive", + lambda process: (_ for _ in ()).throw(PermissionError("kill denied")) + if failure == "kill" + else kill_calls.append(process), + ) + monkeypatch.setattr(start_utils.os, "_exit", lambda code: exit_codes.append(code)) + + manager.setup_signal_handlers(FailingHttp()) + cleaner = auto_shm_cleanup.get_auto_cleanup() + shm_cleanup_calls = [] + monkeypatch.setattr(cleaner, "cleanup", lambda: shm_cleanup_calls.append(True)) + handlers[start_utils.signal.SIGTERM](start_utils.signal.SIGTERM, None) + + assert terminate_calls == [{"wait_timeout": start_utils.PROCESS_SHUTDOWN_WAIT_TIMEOUT_SECONDS}] + assert shm_cleanup_calls == [True] + assert exit_codes == [0] + if failure == "timeout": + assert kill_calls == [manager.http_server_process] + + +def test_signal_handler_exits_even_when_cleanup_raises(monkeypatch): + process_manager = start_utils.SubmoduleManager() + registered_handlers = {} + exit_codes = [] + monkeypatch.setattr( + start_utils.signal, + "signal", + lambda sig, handler: registered_handlers.__setitem__(sig, handler), + ) + monkeypatch.setattr( + process_manager, + "terminate_all_processes", + lambda **_: (_ for _ in ()).throw(RuntimeError()), + ) + monkeypatch.setattr(start_utils.os, "_exit", lambda code: exit_codes.append(code)) + + process_manager.setup_signal_handlers() + cleaner = auto_shm_cleanup.get_auto_cleanup() + shm_cleanup_calls = [] + monkeypatch.setattr(cleaner, "cleanup", lambda: shm_cleanup_calls.append(True)) + registered_handlers[start_utils.signal.SIGINT](start_utils.signal.SIGINT, None) + + assert exit_codes == [0] + assert shm_cleanup_calls == [True] + + +def test_terminate_skips_zombie_even_when_is_running_is_true(monkeypatch): + zombie = FakeProcess(pid=1234, running=True) + process_manager = start_utils.SubmoduleManager() + process_manager.processes = [zombie] + wait_calls = [] + monkeypatch.setattr(start_utils, "is_process_active", lambda pid: False) + monkeypatch.setattr(start_utils.psutil, "Process", lambda pid: zombie) + monkeypatch.setattr(start_utils.psutil, "wait_procs", lambda *args, **kwargs: wait_calls.append((args, kwargs))) + + process_manager.terminate_all_processes() + + assert zombie.kill_calls == 0 + assert wait_calls == [] + + +@pytest.mark.parametrize("wait_timeout", [0, start_utils.PROCESS_SHUTDOWN_WAIT_TIMEOUT_SECONDS]) +def test_terminate_wait_is_bounded_and_target_pids_are_deduplicated(monkeypatch, wait_timeout): + child = FakeProcess(pid=1002) + root = FakeProcess(pid=1001, children=[child]) + duplicate_root = FakeProcess(pid=1001, children=[child]) + process_by_pid = {1001: root, 1002: child} + process_manager = start_utils.SubmoduleManager() + process_manager.processes = [root, duplicate_root, child] + wait_calls = [] + monkeypatch.setattr(start_utils, "is_process_active", lambda pid: True) + monkeypatch.setattr(start_utils.psutil, "Process", lambda pid: process_by_pid[pid]) + monkeypatch.setattr( + start_utils.psutil, + "wait_procs", + lambda processes, timeout: ( + [], + wait_calls.append((processes, timeout)) or list(processes), + ), + ) + + process_manager.terminate_all_processes(wait_timeout=wait_timeout) + + assert child.kill_calls == 1 + assert root.kill_calls == 1 + if wait_timeout: + assert [process.pid for process in wait_calls[0][0]] == [1002, 1001] + assert wait_calls[0][1] == wait_timeout + else: + assert wait_calls == [] + + +def test_normal_start_installs_handlers_before_launching_submodules(monkeypatch): + from lightllm.server import api_start + + calls = [] + + class FakeManager: + def setup_signal_handlers(self, http_server_process=None): + calls.append(("setup", http_server_process)) + + def supervise_processes(self, http_server_process): + calls.append(("supervise", http_server_process)) + + http_server_process = object() + monkeypatch.setattr(api_start, "process_manager", FakeManager()) + monkeypatch.setattr(api_start, "_launch_subprocesses", lambda args: calls.append(("launch", None))) + monkeypatch.setattr( + api_start.subprocess, "Popen", lambda command: calls.append(("popen", command)) or http_server_process + ) + monkeypatch.setattr(api_start, "get_shm_port_args", lambda: SimpleNamespace(port=8000)) + args = SimpleNamespace( + hypercorn_config=None, + httpserver_workers=1, + host="127.0.0.1", + model_dir="/model", + health_monitor=False, + ) + + api_start.normal_or_p_d_start(args) + + assert calls[0] == ("setup", None) + assert calls[1] == ("launch", None) + assert calls[2][0] == "popen" + assert calls[3] == ("setup", http_server_process) def test_supervisor_fails_when_http_server_exits(monkeypatch): @@ -168,7 +454,11 @@ def test_supervisor_fails_when_http_server_exits(monkeypatch): process_manager = start_utils.SubmoduleManager() terminate_calls = [] kill_calls = [] - monkeypatch.setattr(process_manager, "terminate_all_processes", lambda: terminate_calls.append(True)) + monkeypatch.setattr( + process_manager, + "terminate_all_processes", + lambda **_: terminate_calls.append(True), + ) monkeypatch.setattr(start_utils, "kill_recursive", lambda process: kill_calls.append(process)) with pytest.raises(RuntimeError, match="HTTP server exited unexpectedly with return code 0"): @@ -186,7 +476,11 @@ def test_supervisor_fails_and_cleans_up_when_submodule_exits(monkeypatch): process_manager.process_names = {dead_process: dead_process.name()} terminate_calls = [] kill_calls = [] - monkeypatch.setattr(process_manager, "terminate_all_processes", lambda: terminate_calls.append(True)) + monkeypatch.setattr( + process_manager, + "terminate_all_processes", + lambda **_: terminate_calls.append(True), + ) monkeypatch.setattr(start_utils, "is_process_active", lambda pid: True) monkeypatch.setattr(start_utils, "kill_recursive", lambda process: kill_calls.append(process)) @@ -208,7 +502,11 @@ def test_supervisor_treats_zombie_submodule_as_dead(monkeypatch): process_manager.process_names = {zombie_process: zombie_process.name()} terminate_calls = [] kill_calls = [] - monkeypatch.setattr(process_manager, "terminate_all_processes", lambda: terminate_calls.append(True)) + monkeypatch.setattr( + process_manager, + "terminate_all_processes", + lambda **_: terminate_calls.append(True), + ) monkeypatch.setattr(start_utils, "is_process_active", lambda pid: False) monkeypatch.setattr(start_utils, "kill_recursive", lambda process: kill_calls.append(process)) @@ -230,7 +528,11 @@ def test_supervisor_keeps_polling_while_all_processes_are_alive(monkeypatch): process_manager.process_names = {child_process: child_process.name()} terminate_calls = [] sleep_calls = [] - monkeypatch.setattr(process_manager, "terminate_all_processes", lambda: terminate_calls.append(True)) + monkeypatch.setattr( + process_manager, + "terminate_all_processes", + lambda **_: terminate_calls.append(True), + ) def stop_after_first_poll(interval): sleep_calls.append(interval) @@ -253,7 +555,11 @@ def test_supervisor_supports_submodule_only_processes(monkeypatch): process_manager.process_names = {dead_process: dead_process.name()} terminate_calls = [] kill_calls = [] - monkeypatch.setattr(process_manager, "terminate_all_processes", lambda: terminate_calls.append(True)) + monkeypatch.setattr( + process_manager, + "terminate_all_processes", + lambda **_: terminate_calls.append(True), + ) monkeypatch.setattr(start_utils, "is_process_active", lambda pid: False) monkeypatch.setattr(start_utils, "kill_recursive", lambda process: kill_calls.append(process)) From 5337a53be6a388cdfb1d93645efc608d666f201d Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 10 Sep 2026 02:25:44 +0000 Subject: [PATCH 2/8] refactor: simplify launcher signal handling fix --- lightllm/server/api_start.py | 11 +- lightllm/utils/auto_shm_cleanup.py | 6 +- lightllm/utils/start_utils.py | 272 ++++++++-------------- unit_tests/utils/test_start_utils.py | 331 ++------------------------- 4 files changed, 124 insertions(+), 496 deletions(-) diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 1676313371..252cfcd472 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_signal_handlers() get_shm_port_args(create=True) # 多机用于收发node ip, 这个地方修改了args env,所以需要重新设置一下。 send_and_receive_node_ip(args) @@ -424,9 +425,7 @@ def _hypercorn_config_args(args: StartArgs): def normal_or_p_d_start(args: StartArgs): - # Install this before _launch_subprocesses creates multiprocessing children. - process_manager.setup_signal_handlers() - _launch_subprocesses(args) + process_manager = _launch_subprocesses(args) # 启动 Hypercorn command = [ @@ -462,7 +461,6 @@ def normal_or_p_d_start(args: StartArgs): def pd_master_start(args: StartArgs): - process_manager.setup_signal_handlers() _set_envs_and_config(args) set_unique_server_name(args) if args.run_mode != "pd_master": @@ -483,6 +481,7 @@ def pd_master_start(args: StartArgs): validate_ports([args.port]) set_env_start_args(args) + process_manager.setup_signal_handlers() get_shm_port_args(create=True) logger.info(f"all start args:{args}") @@ -524,7 +523,6 @@ def visual_only_start(args): from lightllm.server.core.objs.start_args_type import StartArgs args: StartArgs = args - process_manager.setup_signal_handlers() _set_envs_and_config(args) if args.afs_image_embed_dir is not None: os.makedirs(args.afs_image_embed_dir, mode=0o777, exist_ok=True) @@ -549,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_signal_handlers() get_shm_port_args(create=True) logger.info(f"all start args:{args}") @@ -566,7 +565,6 @@ def visual_only_start(args): def config_server_start(args): - process_manager.setup_signal_handlers() set_unique_server_name(args) if args.run_mode != "config_server": return @@ -576,6 +574,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_signal_handlers() get_shm_port_args(create=True) logger.info(f"all start args:{args}") diff --git a/lightllm/utils/auto_shm_cleanup.py b/lightllm/utils/auto_shm_cleanup.py index 95d118d838..2417fef085 100644 --- a/lightllm/utils/auto_shm_cleanup.py +++ b/lightllm/utils/auto_shm_cleanup.py @@ -40,7 +40,7 @@ def _init_libc(self): self.libc = None def _register_handlers_for_cleanup(self): - atexit.register(self.cleanup) + atexit.register(self._cleanup) self.register_signal_handlers() def register_signal_handlers(self): @@ -51,13 +51,13 @@ def register_signal_handlers(self): self.signal_handlers_registered = True def _signal_cleanup_handler(self, signum, frame): - self.cleanup() + self._cleanup() parent = psutil.Process(os.getpid()) # 递归拿到所有子进程并终止 for ch in parent.children(recursive=True): ch.kill() - def cleanup(self): + def _cleanup(self): """清理:System V 执行 IPC_RMID,POSIX 执行 unlink。""" removed_sysv = 0 IPC_RMID = 0 diff --git a/lightllm/utils/start_utils.py b/lightllm/utils/start_utils.py index d6b6c0bbb2..717f1682f2 100644 --- a/lightllm/utils/start_utils.py +++ b/lightllm/utils/start_utils.py @@ -7,20 +7,15 @@ import psutil from lightllm.utils.log_utils import init_logger from lightllm.utils.process_check import is_process_active +from lightllm.utils.auto_shm_cleanup import get_auto_cleanup logger = init_logger(__name__) -# Waiting for an unrelated/re-parented zombie can otherwise block forever. -PROCESS_SHUTDOWN_WAIT_TIMEOUT_SECONDS = 5 - - class SubmoduleManager: def __init__(self): self.processes = [] self.process_names = {} - self.http_server_process = None - self._handling_signal = False def start_submodule_processes(self, start_funcs=[], start_args=[]): assert len(start_funcs) == len(start_args) @@ -28,50 +23,34 @@ def start_submodule_processes(self, start_funcs=[], start_args=[]): processes = [] managed_processes = [] - try: - for start_func, start_arg in zip(start_funcs, start_args): - pipe_reader, pipe_writer = mp.Pipe(duplex=False) - process = mp.Process( - target=start_func, - args=start_arg + (pipe_writer,), - ) - pipe_readers.append(pipe_reader) - processes.append(process) - try: - process.start() - # Register before waiting for initialization so Ctrl-C can clean - # processes which are still starting up. - managed_process = psutil.Process(process.pid) - self.processes.append(managed_process) - self.process_names[managed_process] = managed_process.name() - managed_processes.append(managed_process) - finally: - pipe_writer.close() - - # Wait for all processes to initialize. - for index, pipe_reader in enumerate(pipe_readers): - init_state = pipe_reader.recv() - if init_state != "init ok": - logger.error(f"init func {start_funcs[index].__name__} : {str(init_state)}") - raise SystemExit(1) + for start_func, start_arg in zip(start_funcs, start_args): + pipe_reader, pipe_writer = mp.Pipe(duplex=False) + process = mp.Process( + target=start_func, + args=start_arg + (pipe_writer,), + ) + 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): + init_state = pipe_reader.recv() + if init_state != "init ok": + logger.error(f"init func {start_funcs[index].__name__} : {str(init_state)}") + for proc in processes: + proc.kill() + sys.exit(1) + else: logger.info(f"init func {start_funcs[index].__name__} : {str(init_state)}") - assert all(process.is_alive() for process in processes) - return managed_processes - except BaseException: - # recv() may be interrupted by Ctrl-C or raise EOFError when a child - # dies. All successfully-started children have already been managed. - try: - self.terminate_all_processes() - except Exception: - logger.exception("Failed to clean up submodules after initialization failure") - raise - finally: - for pipe_reader in pipe_readers: - try: - pipe_reader.close() - except (OSError, EOFError): - pass + assert all([proc.is_alive() for proc in processes]) + return managed_processes def register_process_tree(self, root_process): """Add persistent LightLLM descendants to supervision. @@ -96,108 +75,74 @@ def register_process_tree(self, root_process): self.processes.append(process) self.process_names[process] = process_name - def terminate_all_processes(self, wait_timeout=PROCESS_SHUTDOWN_WAIT_TIMEOUT_SECONDS): - """Kill all managed local process trees without indefinitely waiting.""" - processes_by_pid = {} - for process in self.processes: - for tree_process in _get_process_tree(process): - processes_by_pid.setdefault(tree_process.pid, tree_process) - - processes_to_wait_for = list(processes_by_pid.values()) - _kill_processes(processes_to_wait_for) + def terminate_all_processes(self): + from lightllm.utils.envs_utils import get_env_start_args - if processes_to_wait_for and wait_timeout > 0: + def kill_recursive(proc): try: - _gone, alive = psutil.wait_procs(processes_to_wait_for, timeout=wait_timeout) - if alive: - logger.warning( - "Timed out waiting for processes to exit: %s", - ", ".join(str(process.pid) for process in alive), - ) - except (psutil.NoSuchProcess, psutil.ZombieProcess, psutil.AccessDenied): - # A process may disappear between kill() and wait_procs(). - pass - except Exception: - logger.exception("Failed while waiting for submodule processes to exit") - - # Recover GPU compute mode, but failure here must not prevent launcher exit. - try: - from lightllm.utils.envs_utils import get_env_start_args - - is_enable_mps = get_env_start_args().enable_mps - if is_enable_mps: - from lightllm.utils.device_utils import stop_mps - - stop_mps() - except Exception: - logger.exception("Failed to restore GPU compute mode during shutdown") + parent = psutil.Process(proc.pid) + children = parent.children(recursive=True) + for child in children: + logger.info(f"Killing child process {child.pid}") + child.kill() + logger.info(f"Killing parent process {proc.pid}") + parent.kill() + except psutil.NoSuchProcess: + logger.warning(f"Process {proc.pid} does not exist.") + + for proc in self.processes: + if proc.is_running(): + kill_recursive(proc) + proc.wait() + + # recover the gpu compute mode + is_enable_mps = get_env_start_args().enable_mps + if is_enable_mps: + from lightllm.utils.device_utils import stop_mps + + stop_mps() logger.info("All processes terminated gracefully.") def setup_signal_handlers(self, http_server_process=None): - from lightllm.utils.auto_shm_cleanup import get_auto_cleanup - - # Initialize shared-memory signal handling before the launcher takes over. - shm_cleanup = get_auto_cleanup() - # The installed closure deliberately reads this field at signal time: - # handlers are installed before submodules are launched, while Hypercorn - # is only available later in the startup sequence. - if http_server_process is not None: - self.http_server_process = http_server_process + # AutoShmCleanup 会安装自己的信号处理器;先初始化它,再由 launcher + # 接管信号,避免多机 rendezvous 阶段的 SIGINT 被清理器处理后继续阻塞。 + get_auto_cleanup() def signal_handler(sig, _frame): - repeated_signal = self._handling_signal - self._handling_signal = True - try: - if repeated_signal: - logger.warning("Received a second shutdown signal; forcing immediate cleanup") - elif sig == signal.SIGINT: - logger.info("Received SIGINT (Ctrl+C), forcing immediate exit...") - elif sig == signal.SIGTERM: - logger.info("Received SIGTERM, shutting down gracefully...") - else: - logger.info("Received SIGHUP (terminal closed), shutting down gracefully...") - - if ( - not repeated_signal - and sig != signal.SIGINT - and self.http_server_process is not None - and self.http_server_process.poll() is None - ): - self.http_server_process.send_signal(signal.SIGTERM) - self.http_server_process.wait(timeout=60) - logger.info("HTTP server exited gracefully") - except subprocess.TimeoutExpired: - logger.warning("HTTP server did not exit in time, killing it...") - except Exception: - logger.exception("Shutdown cleanup failed") - finally: + if sig == signal.SIGINT: + logger.info("Received SIGINT (Ctrl+C), forcing immediate exit...") + if http_server_process is not None: + kill_recursive(http_server_process) + + self.terminate_all_processes() + logger.info("All processes have been forcefully terminated.") + sys.exit(0) + + if sig == signal.SIGTERM: + logger.info("Received SIGTERM, shutting down gracefully...") + else: + logger.info("Received SIGHUP (terminal closed), shutting down gracefully...") + + if http_server_process is not None and http_server_process.poll() is None: + http_server_process.send_signal(signal.SIGTERM) try: - for shutdown_signal in (signal.SIGTERM, signal.SIGINT, signal.SIGHUP): - signal.signal(shutdown_signal, signal.SIG_IGN) - self._cleanup_processes( - self.http_server_process, - wait_timeout=0 if repeated_signal else PROCESS_SHUTDOWN_WAIT_TIMEOUT_SECONDS, - ) - logger.info("All processes have been terminated.") - except Exception: - logger.exception("Failed to clean up processes during shutdown") - finally: - try: - shm_cleanup.cleanup() - except Exception: - logger.exception("Failed to clean up shared memory during shutdown") - finally: - # Do not let multiprocessing's atexit handler re-join a stuck - # child after Ctrl-C. Cleanup above is intentionally best effort. - os._exit(1 if repeated_signal else 0) + http_server_process.wait(timeout=60) + logger.info("HTTP server exited gracefully") + except subprocess.TimeoutExpired: + logger.warning("HTTP server did not exit in time, killing it...") + kill_recursive(http_server_process) + + self.terminate_all_processes() + logger.info("All processes have been terminated gracefully.") + sys.exit(0) signal.signal(signal.SIGTERM, signal_handler) signal.signal(signal.SIGINT, signal_handler) signal.signal(signal.SIGHUP, signal_handler) logger.info(f"start process pid {os.getpid()}") - if self.http_server_process is not None: - logger.info(f"http server pid {self.http_server_process.pid}") + if http_server_process is not None: + logger.info(f"http server pid {http_server_process.pid}") def supervise_processes(self, http_server_process=None): """Watch the HTTP server, when present, and all registered submodules. @@ -214,7 +159,7 @@ def supervise_processes(self, http_server_process=None): if http_return_code is not None: message = f"HTTP server exited unexpectedly with return code {http_return_code}" logger.error(message) - self._cleanup_processes(http_server_process) + self._cleanup_after_process_failure(http_server_process) raise RuntimeError(message) dead_processes = [ @@ -233,21 +178,21 @@ def supervise_processes(self, http_server_process=None): dead_process_descriptions = ", ".join(dead_process_descriptions) message = f"Critical LightLLM submodule exited unexpectedly: {dead_process_descriptions}" logger.error(message) - self._cleanup_processes(http_server_process) + self._cleanup_after_process_failure(http_server_process) raise RuntimeError(message) time.sleep(supervisor_interval_seconds) - def _cleanup_processes(self, http_server_process, wait_timeout=PROCESS_SHUTDOWN_WAIT_TIMEOUT_SECONDS): - """Best-effort cleanup before the launcher exits.""" - try: - if http_server_process is not None and http_server_process.poll() is None: + def _cleanup_after_process_failure(self, http_server_process): + """Best-effort cleanup before the launcher exits with a failure.""" + if http_server_process is not None and http_server_process.poll() is None: + try: kill_recursive(http_server_process) - except Exception: - logger.exception("Failed to terminate the HTTP server process tree") + except Exception: + logger.exception("Failed to terminate the HTTP server process tree") try: - self.terminate_all_processes(wait_timeout=wait_timeout) + self.terminate_all_processes() except Exception: logger.exception("Failed to terminate all LightLLM submodule processes") @@ -281,34 +226,17 @@ def start_submodule_processes(start_funcs=[], start_args=[]): return -def _get_process_tree(root_process): - try: - root = psutil.Process(root_process.pid) - descendants = root.children(recursive=True) - except (psutil.NoSuchProcess, psutil.ZombieProcess, psutil.AccessDenied): - return [] - - processes = [] - for process in list(reversed(descendants)) + [root]: - try: - if is_process_active(process.pid): - processes.append(process) - except (psutil.NoSuchProcess, psutil.ZombieProcess, psutil.AccessDenied): - continue - return processes - - -def _kill_processes(processes): - for process in processes: - try: - logger.info(f"Killing process {process.pid}") - process.kill() - except (psutil.NoSuchProcess, psutil.ZombieProcess, psutil.AccessDenied): - continue - - def kill_recursive(proc): - _kill_processes(_get_process_tree(proc)) + try: + parent = psutil.Process(proc.pid) + children = parent.children(recursive=True) + for child in children: + logger.info(f"Killing child process {child.pid}") + child.kill() + logger.info(f"Killing parent process {proc.pid}") + parent.kill() + except psutil.NoSuchProcess: + logger.warning(f"Process {proc.pid} does not exist.") process_manager = SubmoduleManager() diff --git a/unit_tests/utils/test_start_utils.py b/unit_tests/utils/test_start_utils.py index b0ec7a4f65..6acd2d82bb 100644 --- a/unit_tests/utils/test_start_utils.py +++ b/unit_tests/utils/test_start_utils.py @@ -1,15 +1,6 @@ -from types import SimpleNamespace - import pytest -from lightllm.utils import auto_shm_cleanup, start_utils - - -@pytest.fixture(autouse=True) -def mock_auto_shm_cleanup(monkeypatch): - monkeypatch.setattr(auto_shm_cleanup.AutoShmCleanup, "_init_libc", lambda self: None) - monkeypatch.setattr(auto_shm_cleanup.atexit, "register", lambda *args, **kwargs: None) - monkeypatch.setattr(auto_shm_cleanup, "_auto_cleanup", None) +from lightllm.utils import start_utils class FakeHttpServerProcess: @@ -38,7 +29,6 @@ def __init__(self, pid, running=True, children=None, name="process", exitcode=No self._name = name self.exitcode = exitcode self.wait_timeout = wait_timeout - self.kill_calls = 0 def is_running(self): return self.running @@ -54,22 +44,14 @@ def wait(self, timeout=None): raise start_utils.psutil.TimeoutExpired(timeout, pid=self.pid, name=self._name) return self.exitcode - def kill(self): - self.kill_calls += 1 - 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" - def close(self): - pass - - class FakePipeWriter: - def close(self): - pass - class FakeMpProcess: next_pid = 1000 @@ -83,7 +65,7 @@ def start(self): def is_alive(self): return True - monkeypatch.setattr(start_utils.mp, "Pipe", lambda duplex: (FakePipeReader(), FakePipeWriter())) + monkeypatch.setattr(start_utils.mp, "Pipe", lambda duplex: (FakePipeReader(), object())) monkeypatch.setattr(start_utils.mp, "Process", FakeMpProcess) monkeypatch.setattr( start_utils.psutil, @@ -105,45 +87,6 @@ def is_alive(self): } -def test_start_submodule_processes_cleans_up_after_recv_error(monkeypatch): - class FakePipeReader: - def recv(self): - raise EOFError - - def close(self): - pass - - class FakePipeWriter: - def close(self): - pass - - class FakeMpProcess: - pid = 1000 - - def __init__(self, target, args): - pass - - def start(self): - pass - - managed_process = FakeProcess(pid=1000, name="process-1000") - process_manager = start_utils.SubmoduleManager() - cleanup_process_pids = [] - monkeypatch.setattr(start_utils.mp, "Pipe", lambda duplex: (FakePipeReader(), FakePipeWriter())) - monkeypatch.setattr(start_utils.mp, "Process", FakeMpProcess) - monkeypatch.setattr(start_utils.psutil, "Process", lambda pid: managed_process) - monkeypatch.setattr( - process_manager, - "terminate_all_processes", - lambda **_: cleanup_process_pids.extend(process.pid for process in process_manager.processes), - ) - - with pytest.raises(EOFError): - process_manager.start_submodule_processes(start_funcs=[lambda pipe_writer: None], start_args=[()]) - - assert cleanup_process_pids == [1000] - - def test_register_process_tree_adds_recursive_descendants(monkeypatch): descendants = [ FakeProcess(pid=1001, name="lightllm::model_infer"), @@ -201,18 +144,14 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): process_manager = start_utils.SubmoduleManager() registered_handlers = {} terminate_calls = [] + auto_cleanup_calls = [] + monkeypatch.setattr(start_utils, "get_auto_cleanup", lambda: auto_cleanup_calls.append(True)) monkeypatch.setattr( start_utils.signal, "signal", lambda sig, handler: registered_handlers.__setitem__(sig, handler), ) - monkeypatch.setattr( - process_manager, - "terminate_all_processes", - lambda **_: terminate_calls.append(True), - ) - exit_codes = [] - monkeypatch.setattr(start_utils.os, "_exit", lambda code: exit_codes.append(code)) + monkeypatch.setattr(process_manager, "terminate_all_processes", lambda: terminate_calls.append(True)) process_manager.setup_signal_handlers(http_server_process) @@ -221,232 +160,14 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): start_utils.signal.SIGINT, start_utils.signal.SIGHUP, } - registered_handlers[start_utils.signal.SIGTERM](start_utils.signal.SIGTERM, None) + assert auto_cleanup_calls == [True] + with pytest.raises(SystemExit) as exc_info: + registered_handlers[start_utils.signal.SIGTERM](start_utils.signal.SIGTERM, None) + assert exc_info.value.code == 0 assert http_server_process.sent_signals == [start_utils.signal.SIGTERM] assert http_server_process.wait_timeouts == [60] assert terminate_calls == [True] - assert exit_codes == [0] - - -@pytest.mark.parametrize( - "signal_to_send,repeated_signal,expected_exit_code", - [(start_utils.signal.SIGINT, False, 0), (start_utils.signal.SIGTERM, True, 1)], -) -def test_setup_signal_handlers_uses_http_process_set_after_handler_install( - monkeypatch, signal_to_send, repeated_signal, expected_exit_code -): - process_manager = start_utils.SubmoduleManager() - http_server_process = FakeHttpServerProcess() - registered_handlers = {} - killed_processes = [] - exit_codes = [] - monkeypatch.setattr( - start_utils.signal, - "signal", - lambda sig, handler: registered_handlers.__setitem__(sig, handler), - ) - monkeypatch.setattr(start_utils, "kill_recursive", lambda process: killed_processes.append(process)) - terminate_calls = [] - monkeypatch.setattr(process_manager, "terminate_all_processes", lambda **kwargs: terminate_calls.append(kwargs)) - monkeypatch.setattr(start_utils.os, "_exit", lambda code: exit_codes.append(code)) - - process_manager.setup_signal_handlers() - process_manager.setup_signal_handlers(http_server_process) - process_manager._handling_signal = repeated_signal - registered_handlers[signal_to_send](signal_to_send, None) - - assert killed_processes == [http_server_process] - assert terminate_calls == [ - {"wait_timeout": 0 if repeated_signal else start_utils.PROCESS_SHUTDOWN_WAIT_TIMEOUT_SECONDS} - ] - assert exit_codes == [expected_exit_code] - if repeated_signal: - assert http_server_process.sent_signals == [] - assert http_server_process.wait_timeouts == [] - - -def test_launcher_handler_remains_registered_after_shared_memory_registration(monkeypatch): - handlers = {} - exit_codes = [] - manager = start_utils.SubmoduleManager() - monkeypatch.setattr( - start_utils.signal, - "signal", - lambda sig, handler: handlers.__setitem__(sig, handler), - ) - monkeypatch.setattr(manager, "terminate_all_processes", lambda **_: None) - monkeypatch.setattr(start_utils.os, "_exit", lambda code: exit_codes.append(code)) - - manager.setup_signal_handlers() - launcher_handler = handlers[start_utils.signal.SIGINT] - cleaner = auto_shm_cleanup.get_auto_cleanup() - monkeypatch.setattr(cleaner, "cleanup", lambda: None) - auto_shm_cleanup.register_posix_shm_for_cleanup("test") - - assert handlers[start_utils.signal.SIGINT] is launcher_handler - launcher_handler(start_utils.signal.SIGINT, None) - assert exit_codes == [0] - - -@pytest.mark.parametrize("failure", ["poll", "send", "wait", "timeout", "kill"]) -def test_http_shutdown_failure_still_cleans_processes_and_shared_memory(monkeypatch, failure): - class FailingHttp: - pid = 4321 - - def poll(self): - if failure == "poll": - raise PermissionError("poll denied") - return None - - def send_signal(self, _sig): - if failure == "send": - raise PermissionError("send denied") - - def wait(self, timeout=None): - if failure == "wait": - raise PermissionError("wait denied") - if failure == "timeout": - raise start_utils.subprocess.TimeoutExpired("http", timeout) - - handlers = {} - terminate_calls = [] - kill_calls = [] - exit_codes = [] - manager = start_utils.SubmoduleManager() - monkeypatch.setattr( - start_utils.signal, - "signal", - lambda sig, handler: handlers.__setitem__(sig, handler), - ) - monkeypatch.setattr(manager, "terminate_all_processes", lambda **kwargs: terminate_calls.append(kwargs)) - monkeypatch.setattr( - start_utils, - "kill_recursive", - lambda process: (_ for _ in ()).throw(PermissionError("kill denied")) - if failure == "kill" - else kill_calls.append(process), - ) - monkeypatch.setattr(start_utils.os, "_exit", lambda code: exit_codes.append(code)) - - manager.setup_signal_handlers(FailingHttp()) - cleaner = auto_shm_cleanup.get_auto_cleanup() - shm_cleanup_calls = [] - monkeypatch.setattr(cleaner, "cleanup", lambda: shm_cleanup_calls.append(True)) - handlers[start_utils.signal.SIGTERM](start_utils.signal.SIGTERM, None) - - assert terminate_calls == [{"wait_timeout": start_utils.PROCESS_SHUTDOWN_WAIT_TIMEOUT_SECONDS}] - assert shm_cleanup_calls == [True] - assert exit_codes == [0] - if failure == "timeout": - assert kill_calls == [manager.http_server_process] - - -def test_signal_handler_exits_even_when_cleanup_raises(monkeypatch): - process_manager = start_utils.SubmoduleManager() - registered_handlers = {} - exit_codes = [] - monkeypatch.setattr( - start_utils.signal, - "signal", - lambda sig, handler: registered_handlers.__setitem__(sig, handler), - ) - monkeypatch.setattr( - process_manager, - "terminate_all_processes", - lambda **_: (_ for _ in ()).throw(RuntimeError()), - ) - monkeypatch.setattr(start_utils.os, "_exit", lambda code: exit_codes.append(code)) - - process_manager.setup_signal_handlers() - cleaner = auto_shm_cleanup.get_auto_cleanup() - shm_cleanup_calls = [] - monkeypatch.setattr(cleaner, "cleanup", lambda: shm_cleanup_calls.append(True)) - registered_handlers[start_utils.signal.SIGINT](start_utils.signal.SIGINT, None) - - assert exit_codes == [0] - assert shm_cleanup_calls == [True] - - -def test_terminate_skips_zombie_even_when_is_running_is_true(monkeypatch): - zombie = FakeProcess(pid=1234, running=True) - process_manager = start_utils.SubmoduleManager() - process_manager.processes = [zombie] - wait_calls = [] - monkeypatch.setattr(start_utils, "is_process_active", lambda pid: False) - monkeypatch.setattr(start_utils.psutil, "Process", lambda pid: zombie) - monkeypatch.setattr(start_utils.psutil, "wait_procs", lambda *args, **kwargs: wait_calls.append((args, kwargs))) - - process_manager.terminate_all_processes() - - assert zombie.kill_calls == 0 - assert wait_calls == [] - - -@pytest.mark.parametrize("wait_timeout", [0, start_utils.PROCESS_SHUTDOWN_WAIT_TIMEOUT_SECONDS]) -def test_terminate_wait_is_bounded_and_target_pids_are_deduplicated(monkeypatch, wait_timeout): - child = FakeProcess(pid=1002) - root = FakeProcess(pid=1001, children=[child]) - duplicate_root = FakeProcess(pid=1001, children=[child]) - process_by_pid = {1001: root, 1002: child} - process_manager = start_utils.SubmoduleManager() - process_manager.processes = [root, duplicate_root, child] - wait_calls = [] - monkeypatch.setattr(start_utils, "is_process_active", lambda pid: True) - monkeypatch.setattr(start_utils.psutil, "Process", lambda pid: process_by_pid[pid]) - monkeypatch.setattr( - start_utils.psutil, - "wait_procs", - lambda processes, timeout: ( - [], - wait_calls.append((processes, timeout)) or list(processes), - ), - ) - - process_manager.terminate_all_processes(wait_timeout=wait_timeout) - - assert child.kill_calls == 1 - assert root.kill_calls == 1 - if wait_timeout: - assert [process.pid for process in wait_calls[0][0]] == [1002, 1001] - assert wait_calls[0][1] == wait_timeout - else: - assert wait_calls == [] - - -def test_normal_start_installs_handlers_before_launching_submodules(monkeypatch): - from lightllm.server import api_start - - calls = [] - - class FakeManager: - def setup_signal_handlers(self, http_server_process=None): - calls.append(("setup", http_server_process)) - - def supervise_processes(self, http_server_process): - calls.append(("supervise", http_server_process)) - - http_server_process = object() - monkeypatch.setattr(api_start, "process_manager", FakeManager()) - monkeypatch.setattr(api_start, "_launch_subprocesses", lambda args: calls.append(("launch", None))) - monkeypatch.setattr( - api_start.subprocess, "Popen", lambda command: calls.append(("popen", command)) or http_server_process - ) - monkeypatch.setattr(api_start, "get_shm_port_args", lambda: SimpleNamespace(port=8000)) - args = SimpleNamespace( - hypercorn_config=None, - httpserver_workers=1, - host="127.0.0.1", - model_dir="/model", - health_monitor=False, - ) - - api_start.normal_or_p_d_start(args) - - assert calls[0] == ("setup", None) - assert calls[1] == ("launch", None) - assert calls[2][0] == "popen" - assert calls[3] == ("setup", http_server_process) def test_supervisor_fails_when_http_server_exits(monkeypatch): @@ -454,11 +175,7 @@ def test_supervisor_fails_when_http_server_exits(monkeypatch): process_manager = start_utils.SubmoduleManager() terminate_calls = [] kill_calls = [] - monkeypatch.setattr( - process_manager, - "terminate_all_processes", - lambda **_: terminate_calls.append(True), - ) + monkeypatch.setattr(process_manager, "terminate_all_processes", lambda: terminate_calls.append(True)) monkeypatch.setattr(start_utils, "kill_recursive", lambda process: kill_calls.append(process)) with pytest.raises(RuntimeError, match="HTTP server exited unexpectedly with return code 0"): @@ -476,11 +193,7 @@ def test_supervisor_fails_and_cleans_up_when_submodule_exits(monkeypatch): process_manager.process_names = {dead_process: dead_process.name()} terminate_calls = [] kill_calls = [] - monkeypatch.setattr( - process_manager, - "terminate_all_processes", - lambda **_: terminate_calls.append(True), - ) + monkeypatch.setattr(process_manager, "terminate_all_processes", lambda: terminate_calls.append(True)) monkeypatch.setattr(start_utils, "is_process_active", lambda pid: True) monkeypatch.setattr(start_utils, "kill_recursive", lambda process: kill_calls.append(process)) @@ -502,11 +215,7 @@ def test_supervisor_treats_zombie_submodule_as_dead(monkeypatch): process_manager.process_names = {zombie_process: zombie_process.name()} terminate_calls = [] kill_calls = [] - monkeypatch.setattr( - process_manager, - "terminate_all_processes", - lambda **_: terminate_calls.append(True), - ) + monkeypatch.setattr(process_manager, "terminate_all_processes", lambda: terminate_calls.append(True)) monkeypatch.setattr(start_utils, "is_process_active", lambda pid: False) monkeypatch.setattr(start_utils, "kill_recursive", lambda process: kill_calls.append(process)) @@ -528,11 +237,7 @@ def test_supervisor_keeps_polling_while_all_processes_are_alive(monkeypatch): process_manager.process_names = {child_process: child_process.name()} terminate_calls = [] sleep_calls = [] - monkeypatch.setattr( - process_manager, - "terminate_all_processes", - lambda **_: terminate_calls.append(True), - ) + monkeypatch.setattr(process_manager, "terminate_all_processes", lambda: terminate_calls.append(True)) def stop_after_first_poll(interval): sleep_calls.append(interval) @@ -555,11 +260,7 @@ def test_supervisor_supports_submodule_only_processes(monkeypatch): process_manager.process_names = {dead_process: dead_process.name()} terminate_calls = [] kill_calls = [] - monkeypatch.setattr( - process_manager, - "terminate_all_processes", - lambda **_: terminate_calls.append(True), - ) + monkeypatch.setattr(process_manager, "terminate_all_processes", lambda: terminate_calls.append(True)) monkeypatch.setattr(start_utils, "is_process_active", lambda pid: False) monkeypatch.setattr(start_utils, "kill_recursive", lambda process: kill_calls.append(process)) From e49e5b641c9578002863f590a63c5325b3be5bcd Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 10 Sep 2026 03:08:47 +0000 Subject: [PATCH 3/8] refactor: centralize shared memory cleanup in launcher --- lightllm/server/embed_cache/utils.py | 4 +- .../server/multi_level_kv_cache/shm_objs.py | 5 +- .../model_infer/vision_peak_vram_hold.py | 3 +- lightllm/utils/auto_shm_cleanup.py | 135 -------------- lightllm/utils/kv_cache_utils.py | 2 - lightllm/utils/service_shm_cleanup.py | 172 ++++++++++++++++++ lightllm/utils/shm_port_args.py | 1 - lightllm/utils/shm_utils.py | 15 +- lightllm/utils/start_utils.py | 12 +- unit_tests/utils/test_service_shm_cleanup.py | 73 ++++++++ unit_tests/utils/test_start_utils.py | 26 ++- 11 files changed, 287 insertions(+), 161 deletions(-) delete mode 100644 lightllm/utils/auto_shm_cleanup.py create mode 100644 lightllm/utils/service_shm_cleanup.py create mode 100644 unit_tests/utils/test_service_shm_cleanup.py diff --git a/lightllm/server/embed_cache/utils.py b/lightllm/server/embed_cache/utils.py index 367bcc91a9..283a5006f8 100644 --- a/lightllm/server/embed_cache/utils.py +++ b/lightllm/server/embed_cache/utils.py @@ -1,5 +1,7 @@ import multiprocessing.shared_memory as shm +from lightllm.utils.envs_utils import get_unique_server_name + def create_shm(name, data): try: @@ -24,4 +26,4 @@ def free_shm(name): def get_shm_name_data(uid): - return str(uid) + "-data" + return f"{get_unique_server_name()}_{uid}-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..db55966970 100644 --- a/lightllm/server/multi_level_kv_cache/shm_objs.py +++ b/lightllm/server/multi_level_kv_cache/shm_objs.py @@ -3,7 +3,6 @@ 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 logger = init_logger(__name__) @@ -290,11 +289,9 @@ 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): 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/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/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/service_shm_cleanup.py b/lightllm/utils/service_shm_cleanup.py new file mode 100644 index 0000000000..29da62ecff --- /dev/null +++ b/lightllm/utils/service_shm_cleanup.py @@ -0,0 +1,172 @@ +import atexit +import ctypes +import json +import os +from multiprocessing import shared_memory +from pathlib import Path + +import psutil +from filelock import FileLock + +from lightllm.utils.log_utils import init_logger + + +logger = init_logger(__name__) + +SHM_DIR = Path("/dev/shm") +OWNER_DIR = Path("/tmp/lightllm_service_owners") +OWNER_LOCK_PATH = Path("/tmp/lightllm_service_owners.lock") +SYSTEM_V_SHM_KEY_NAMES = ("cpu_kv_cache_shm_id", "multi_modal_cache_shm_id") + +_registered_cleanups = {} + + +def _get_system_v_shm_keys(): + try: + start_args = json.loads(os.environ["LIGHTLLM_START_ARGS"]) + except (KeyError, json.JSONDecodeError): + return [] + return [int(start_args[name]) for name in SYSTEM_V_SHM_KEY_NAMES if start_args.get(name) is not None] + + +def _unlink_posix_shm(name): + shm = shared_memory.SharedMemory(name=name, create=False) + try: + shm.unlink() + finally: + shm.close() + + +def _remove_system_v_shm(key): + 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 + + shmid = libc.shmget(int(key), 0, 0) + if shmid >= 0: + return libc.shmctl(shmid, 0, None) == 0 # IPC_RMID + return False + + +def cleanup_service_shm(service_name, system_v_shm_keys=()): + """回收一个 LightLLM 服务在本机创建的共享内存。""" + if not service_name: + return + + prefix = f"{service_name}_" + removed_posix = 0 + try: + entries = list(SHM_DIR.iterdir()) + except FileNotFoundError: + entries = [] + + for entry in entries: + if not entry.name.startswith(prefix): + continue + try: + _unlink_posix_shm(entry.name) + removed_posix += 1 + except FileNotFoundError: + pass + except Exception: + logger.exception(f"Failed to unlink POSIX shm {entry.name}") + + removed_system_v = 0 + for key in system_v_shm_keys: + try: + removed_system_v += int(_remove_system_v_shm(key)) + except Exception: + logger.exception(f"Failed to remove System V shm key {key}") + + if removed_posix or removed_system_v: + logger.info( + f"Cleaned service shm for {service_name}: " f"POSIX={removed_posix}, System V keys={removed_system_v}" + ) + + +def _owner_is_alive(owner): + try: + process = psutil.Process(int(owner["pid"])) + return ( + process.is_running() + and process.status() != psutil.STATUS_ZOMBIE + and abs(process.create_time() - float(owner["create_time"])) < 0.01 + ) + except psutil.AccessDenied: + # 无权限确认的进程按存活处理,避免清理其他用户正在运行的服务。 + return True + except (KeyError, TypeError, ValueError, psutil.NoSuchProcess, psutil.ZombieProcess): + return False + + +def _load_owner(path): + try: + with path.open("r", encoding="utf-8") as file: + return json.load(file) + except FileNotFoundError: + return None + except (OSError, json.JSONDecodeError): + logger.warning(f"Ignore invalid LightLLM service owner file: {path}") + return None + + +def _recover_stale_services(): + for owner_path in OWNER_DIR.glob("*.json"): + owner = _load_owner(owner_path) + if owner is None or _owner_is_alive(owner): + continue + cleanup_service_shm(owner.get("service_name"), owner.get("system_v_shm_keys", ())) + try: + owner_path.unlink() + except FileNotFoundError: + pass + + +def register_launcher_shm_cleanup(service_name): + """登记 launcher 资源,并在退出或下次启动时回收该服务的共享内存。""" + if not service_name: + raise RuntimeError("service_name must be initialized before registering shm cleanup") + + owner_pid = os.getpid() + cleanup_key = (owner_pid, service_name) + if cleanup_key in _registered_cleanups: + return _registered_cleanups[cleanup_key] + + OWNER_DIR.mkdir(parents=True, exist_ok=True) + system_v_shm_keys = _get_system_v_shm_keys() + owner = { + "service_name": service_name, + "pid": owner_pid, + "create_time": psutil.Process(owner_pid).create_time(), + "system_v_shm_keys": system_v_shm_keys, + } + owner_path = OWNER_DIR / f"{service_name}.json" + + with FileLock(str(OWNER_LOCK_PATH)): + _recover_stale_services() + tmp_path = owner_path.with_suffix(f".{owner_pid}.tmp") + with tmp_path.open("w", encoding="utf-8") as file: + json.dump(owner, file) + os.replace(tmp_path, owner_path) + + def cleanup(): + nonlocal cleaned + # multiprocessing 子进程会继承 atexit 回调,只有创建记录的 launcher 可以执行全量回收。 + if cleaned or os.getpid() != owner_pid: + return + cleaned = True + cleanup_service_shm(service_name, system_v_shm_keys) + with FileLock(str(OWNER_LOCK_PATH)): + current_owner = _load_owner(owner_path) + if current_owner is not None and current_owner.get("pid") == owner_pid: + try: + owner_path.unlink() + except FileNotFoundError: + pass + + cleaned = False + _registered_cleanups[cleanup_key] = cleanup + atexit.register(cleanup) + return cleanup diff --git a/lightllm/utils/shm_port_args.py b/lightllm/utils/shm_port_args.py index b486c6e2f2..b62fd7cb32 100644 --- a/lightllm/utils/shm_port_args.py +++ b/lightllm/utils/shm_port_args.py @@ -85,7 +85,6 @@ def __init__(self, create: bool = False): 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..decc20fd5c 100644 --- a/lightllm/utils/shm_utils.py +++ b/lightllm/utils/shm_utils.py @@ -1,12 +1,11 @@ from multiprocessing import shared_memory from filelock import FileLock 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 create_or_link_shm(name, expected_size, force_mode=None): """ Args: name: name of the shared memory @@ -27,15 +26,15 @@ def create_or_link_shm(name, expected_size, force_mode=None, auto_cleanup=False) 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 +45,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 +63,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 +71,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 717f1682f2..38ae36f4ff 100644 --- a/lightllm/utils/start_utils.py +++ b/lightllm/utils/start_utils.py @@ -7,7 +7,8 @@ import psutil from lightllm.utils.log_utils import init_logger from lightllm.utils.process_check import is_process_active -from lightllm.utils.auto_shm_cleanup import get_auto_cleanup +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__) @@ -16,6 +17,7 @@ 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) @@ -95,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: @@ -104,9 +109,8 @@ def kill_recursive(proc): logger.info("All processes terminated gracefully.") def setup_signal_handlers(self, http_server_process=None): - # AutoShmCleanup 会安装自己的信号处理器;先初始化它,再由 launcher - # 接管信号,避免多机 rendezvous 阶段的 SIGINT 被清理器处理后继续阻塞。 - get_auto_cleanup() + # 共享内存由 launcher 在所有子进程退出后统一回收,不再让资源对象注册信号处理器。 + self._cleanup_service_shm = register_launcher_shm_cleanup(get_unique_server_name()) def signal_handler(sig, _frame): if sig == signal.SIGINT: 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..9b2ab96b92 --- /dev/null +++ b/unit_tests/utils/test_service_shm_cleanup.py @@ -0,0 +1,73 @@ +import json +import os + +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, "_unlink_posix_shm", lambda name: (shm_dir / name).unlink()) + monkeypatch.setattr( + service_shm_cleanup, + "_remove_system_v_shm", + lambda key: removed_system_v_keys.append(key) or True, + ) + + service_shm_cleanup.cleanup_service_shm("service_0", [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_register_launcher_cleanup_recovers_dead_owner_and_records_current_service(monkeypatch, tmp_path): + owner_dir = tmp_path / "owners" + owner_dir.mkdir() + old_owner_path = owner_dir / "old_service_0.json" + old_owner_path.write_text( + json.dumps( + { + "service_name": "old_service_0", + "pid": 999999999, + "create_time": 1.0, + "system_v_shm_keys": [11], + } + ), + encoding="utf-8", + ) + + cleanup_calls = [] + atexit_callbacks = [] + monkeypatch.setattr(service_shm_cleanup, "OWNER_DIR", owner_dir) + monkeypatch.setattr(service_shm_cleanup, "OWNER_LOCK_PATH", tmp_path / "owners.lock") + monkeypatch.setattr(service_shm_cleanup, "cleanup_service_shm", lambda *args: cleanup_calls.append(args)) + monkeypatch.setattr(service_shm_cleanup.atexit, "register", atexit_callbacks.append) + monkeypatch.setenv( + "LIGHTLLM_START_ARGS", + json.dumps({"cpu_kv_cache_shm_id": 21, "multi_modal_cache_shm_id": 22}), + ) + service_shm_cleanup._registered_cleanups.clear() + + cleanup = service_shm_cleanup.register_launcher_shm_cleanup("current_service_0") + + assert cleanup_calls == [("old_service_0", [11])] + assert not old_owner_path.exists() + current_owner_path = owner_dir / "current_service_0.json" + current_owner = json.loads(current_owner_path.read_text(encoding="utf-8")) + assert current_owner["pid"] == os.getpid() + assert current_owner["system_v_shm_keys"] == [21, 22] + assert atexit_callbacks == [cleanup] + + cleanup() + cleanup() + assert cleanup_calls[-1] == ("current_service_0", [21, 22]) + assert cleanup_calls.count(("current_service_0", [21, 22])) == 1 + assert not current_owner_path.exists() diff --git a/unit_tests/utils/test_start_utils.py b/unit_tests/utils/test_start_utils.py index 6acd2d82bb..e3f6e14081 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 @@ -144,8 +146,13 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): process_manager = start_utils.SubmoduleManager() registered_handlers = {} terminate_calls = [] - auto_cleanup_calls = [] - monkeypatch.setattr(start_utils, "get_auto_cleanup", lambda: auto_cleanup_calls.append(True)) + 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", @@ -160,7 +167,7 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): start_utils.signal.SIGINT, start_utils.signal.SIGHUP, } - assert auto_cleanup_calls == [True] + assert cleanup_calls == [("register", "service_0")] with pytest.raises(SystemExit) as exc_info: registered_handlers[start_utils.signal.SIGTERM](start_utils.signal.SIGTERM, None) @@ -170,6 +177,19 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): assert terminate_calls == [True] +def test_terminate_all_processes_runs_launcher_shm_cleanup(monkeypatch): + from lightllm.utils import envs_utils + + process_manager = start_utils.SubmoduleManager() + cleanup_calls = [] + process_manager._cleanup_service_shm = lambda: cleanup_calls.append(True) + monkeypatch.setattr(envs_utils, "get_env_start_args", lambda: SimpleNamespace(enable_mps=False)) + + process_manager.terminate_all_processes() + + assert cleanup_calls == [True] + + def test_supervisor_fails_when_http_server_exits(monkeypatch): http_server_process = FakeHttpServerProcess(return_code=0) process_manager = start_utils.SubmoduleManager() From a728db0b899111cfb85e77db77ba1d3b5cb2cdcf Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 10 Sep 2026 05:45:18 +0000 Subject: [PATCH 4/8] refactor: centralize service shared memory naming --- .../common/kv_cache_mem_manager/allocator.py | 7 ++---- .../kv_cache_mem_manager/mem_manager.py | 13 +++++----- lightllm/server/api_http.py | 2 +- lightllm/server/core/objs/req.py | 13 ++++------ .../server/core/objs/shm_objs_io_buffer.py | 5 ++-- lightllm/server/core/objs/shm_req_manager.py | 11 ++++----- lightllm/server/core/objs/token_metadata.py | 4 +--- lightllm/server/embed_cache/utils.py | 7 ++++-- lightllm/server/httpserver/manager.py | 9 ++++--- .../multi_level_kv_cache/cpu_cache_client.py | 10 ++++---- .../server/multi_level_kv_cache/shm_objs.py | 2 ++ lightllm/server/req_id_generator.py | 11 ++++----- .../dynamic_prompt/linear_att_radix_cache.py | 7 ++---- .../router/dynamic_prompt/radix_cache.py | 24 +++++++------------ lightllm/server/router/manager.py | 5 ++-- .../model_infer/mode_backend/base_backend.py | 5 +--- .../dp_backend/dp_shared_kv_trans.py | 4 ++-- lightllm/utils/health_check.py | 6 ++--- lightllm/utils/rl/bucketed_weight_transfer.py | 9 +++++-- lightllm/utils/service_shm_cleanup.py | 2 ++ lightllm/utils/shm_port_args.py | 15 ++++-------- lightllm/utils/shm_utils.py | 18 +++++++++++++- .../router/dynamic_prompt/test_radix_cache.py | 20 ++++++++-------- unit_tests/utils/test_shm_utils.py | 22 +++++++++++++++++ 24 files changed, 122 insertions(+), 109 deletions(-) create mode 100644 unit_tests/utils/test_shm_utils.py 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/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 283a5006f8..fcd48d421f 100644 --- a/lightllm/server/embed_cache/utils.py +++ b/lightllm/server/embed_cache/utils.py @@ -1,9 +1,10 @@ import multiprocessing.shared_memory as shm -from lightllm.utils.envs_utils import get_unique_server_name +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) @@ -14,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 f"{get_unique_server_name()}_{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 db55966970..5165ca3af1 100644 --- a/lightllm/server/multi_level_kv_cache/shm_objs.py +++ b/lightllm/server/multi_level_kv_cache/shm_objs.py @@ -3,6 +3,7 @@ from multiprocessing import shared_memory from typing import List, Optional from lightllm.utils.log_utils import init_logger +from lightllm.utils.shm_utils import get_service_shm_name logger = init_logger(__name__) @@ -290,6 +291,7 @@ def key(self, value: int): 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) logger.info(f"create lock shm {name}") 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/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/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 index 29da62ecff..dee942d440 100644 --- a/lightllm/utils/service_shm_cleanup.py +++ b/lightllm/utils/service_shm_cleanup.py @@ -30,6 +30,8 @@ def _get_system_v_shm_keys(): def _unlink_posix_shm(name): + # name 来自 /dev/shm 的实际条目,可能属于上一次异常退出的其他 service。 + # 这里必须直接使用完整名称,不能再按当前 service 添加前缀。 shm = shared_memory.SharedMemory(name=name, create=False) try: shm.unlink() diff --git a/lightllm/utils/shm_port_args.py b/lightllm/utils/shm_port_args.py index b62fd7cb32..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,16 +71,11 @@ 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, diff --git a/lightllm/utils/shm_utils.py b/lightllm/utils/shm_utils.py index decc20fd5c..77d383b9d6 100644 --- a/lightllm/utils/shm_utils.py +++ b/lightllm/utils/shm_utils.py @@ -1,14 +1,29 @@ 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 logger = init_logger(__name__) +def get_service_shm_name(name): + """为内部共享内存统一添加当前服务(UUID + node rank)前缀。 + + 已带当前服务前缀的完整名称保持不变,便于底层包装函数安全复用;未运行在 + launcher 环境中的独立工具和单元测试没有 service name,此时保留原名称。 + """ + name = str(name) + service_name = get_unique_server_name() + if not service_name: + return name + 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 @@ -22,6 +37,7 @@ def create_or_link_shm(name, expected_size, force_mode=None): 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": 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..15f86d2b1b 100644 --- a/unit_tests/server/router/dynamic_prompt/test_radix_cache.py +++ b/unit_tests/server/router/dynamic_prompt/test_radix_cache.py @@ -4,7 +4,7 @@ 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 +25,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 +51,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 +81,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 +96,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 +125,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 +149,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 +173,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 +199,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 +235,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_shm_utils.py b/unit_tests/utils/test_shm_utils.py new file mode 100644 index 0000000000..340c0d1268 --- /dev/null +++ b/unit_tests/utils/test_shm_utils.py @@ -0,0 +1,22 @@ +from lightllm.utils import shm_utils + + +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)] From ad6f241d19c5ea13c7a830fb20de045f66dd5d2c Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 10 Sep 2026 06:26:20 +0000 Subject: [PATCH 5/8] fix: require service name for shared memory --- lightllm/utils/shm_utils.py | 8 +++++--- unit_tests/server/core/objs/test_shm_array.py | 9 +++++++++ unit_tests/server/core/objs/test_shm_req_manager.py | 4 ++++ .../server/router/dynamic_prompt/test_radix_cache.py | 9 +++++++++ unit_tests/utils/test_shm_utils.py | 9 +++++++++ 5 files changed, 36 insertions(+), 3 deletions(-) diff --git a/lightllm/utils/shm_utils.py b/lightllm/utils/shm_utils.py index 77d383b9d6..1929865375 100644 --- a/lightllm/utils/shm_utils.py +++ b/lightllm/utils/shm_utils.py @@ -9,13 +9,15 @@ def get_service_shm_name(name): """为内部共享内存统一添加当前服务(UUID + node rank)前缀。 - 已带当前服务前缀的完整名称保持不变,便于底层包装函数安全复用;未运行在 - launcher 环境中的独立工具和单元测试没有 service name,此时保留原名称。 + 已带当前服务前缀的完整名称保持不变,便于底层包装函数安全复用。service name + 未初始化时直接报错,避免创建无法区分服务、也无法被 launcher 定向回收的裸名称。 """ name = str(name) service_name = get_unique_server_name() if not service_name: - return 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}" 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 15f86d2b1b..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,6 +1,15 @@ 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(): diff --git a/unit_tests/utils/test_shm_utils.py b/unit_tests/utils/test_shm_utils.py index 340c0d1268..45186fdce4 100644 --- a/unit_tests/utils/test_shm_utils.py +++ b/unit_tests/utils/test_shm_utils.py @@ -1,6 +1,15 @@ +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") From 13b206b87026ec8232efce65ba468291b38d389b Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 10 Sep 2026 08:22:16 +0000 Subject: [PATCH 6/8] refactor: simplify launcher shared memory cleanup --- lightllm/utils/service_shm_cleanup.py | 268 ++++++++----------- lightllm/utils/start_utils.py | 4 +- unit_tests/utils/test_service_shm_cleanup.py | 117 +++++--- unit_tests/utils/test_start_utils.py | 1 + 4 files changed, 200 insertions(+), 190 deletions(-) diff --git a/lightllm/utils/service_shm_cleanup.py b/lightllm/utils/service_shm_cleanup.py index dee942d440..19c9cbda01 100644 --- a/lightllm/utils/service_shm_cleanup.py +++ b/lightllm/utils/service_shm_cleanup.py @@ -1,174 +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 -from multiprocessing import shared_memory +import subprocess from pathlib import Path -import psutil -from filelock import FileLock - from lightllm.utils.log_utils import init_logger logger = init_logger(__name__) SHM_DIR = Path("/dev/shm") -OWNER_DIR = Path("/tmp/lightllm_service_owners") -OWNER_LOCK_PATH = Path("/tmp/lightllm_service_owners.lock") -SYSTEM_V_SHM_KEY_NAMES = ("cpu_kv_cache_shm_id", "multi_modal_cache_shm_id") - -_registered_cleanups = {} - - -def _get_system_v_shm_keys(): - try: - start_args = json.loads(os.environ["LIGHTLLM_START_ARGS"]) - except (KeyError, json.JSONDecodeError): - return [] - return [int(start_args[name]) for name in SYSTEM_V_SHM_KEY_NAMES if start_args.get(name) is not None] - - -def _unlink_posix_shm(name): - # name 来自 /dev/shm 的实际条目,可能属于上一次异常退出的其他 service。 - # 这里必须直接使用完整名称,不能再按当前 service 添加前缀。 - shm = shared_memory.SharedMemory(name=name, create=False) - try: - shm.unlink() - finally: - shm.close() - - -def _remove_system_v_shm(key): - 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 - - shmid = libc.shmget(int(key), 0, 0) - if shmid >= 0: - return libc.shmctl(shmid, 0, None) == 0 # IPC_RMID - return False - - -def cleanup_service_shm(service_name, system_v_shm_keys=()): - """回收一个 LightLLM 服务在本机创建的共享内存。""" - if not service_name: - return - - prefix = f"{service_name}_" - removed_posix = 0 - try: - entries = list(SHM_DIR.iterdir()) - except FileNotFoundError: - entries = [] - - for entry in entries: - if not entry.name.startswith(prefix): - continue - try: - _unlink_posix_shm(entry.name) - removed_posix += 1 - except FileNotFoundError: - pass - except Exception: - logger.exception(f"Failed to unlink POSIX shm {entry.name}") - removed_system_v = 0 - for key in system_v_shm_keys: + +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: - removed_system_v += int(_remove_system_v_shm(key)) - except Exception: - logger.exception(f"Failed to remove System V shm key {key}") - - if removed_posix or removed_system_v: - logger.info( - f"Cleaned service shm for {service_name}: " f"POSIX={removed_posix}, System V keys={removed_system_v}" - ) - - -def _owner_is_alive(owner): - try: - process = psutil.Process(int(owner["pid"])) - return ( - process.is_running() - and process.status() != psutil.STATUS_ZOMBIE - and abs(process.create_time() - float(owner["create_time"])) < 0.01 - ) - except psutil.AccessDenied: - # 无权限确认的进程按存活处理,避免清理其他用户正在运行的服务。 - return True - except (KeyError, TypeError, ValueError, psutil.NoSuchProcess, psutil.ZombieProcess): - return False - - -def _load_owner(path): - try: - with path.open("r", encoding="utf-8") as file: - return json.load(file) - except FileNotFoundError: - return None - except (OSError, json.JSONDecodeError): - logger.warning(f"Ignore invalid LightLLM service owner file: {path}") - return None - - -def _recover_stale_services(): - for owner_path in OWNER_DIR.glob("*.json"): - owner = _load_owner(owner_path) - if owner is None or _owner_is_alive(owner): - continue - cleanup_service_shm(owner.get("service_name"), owner.get("system_v_shm_keys", ())) + 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: - owner_path.unlink() + entries = [entry for entry in SHM_DIR.iterdir() if entry.name.startswith(f"{service_name}_")] except FileNotFoundError: - pass + 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 资源,并在退出或下次启动时回收该服务的共享内存。""" - if not service_name: - raise RuntimeError("service_name must be initialized before registering shm cleanup") - - owner_pid = os.getpid() - cleanup_key = (owner_pid, service_name) - if cleanup_key in _registered_cleanups: - return _registered_cleanups[cleanup_key] - - OWNER_DIR.mkdir(parents=True, exist_ok=True) - system_v_shm_keys = _get_system_v_shm_keys() - owner = { - "service_name": service_name, - "pid": owner_pid, - "create_time": psutil.Process(owner_pid).create_time(), - "system_v_shm_keys": system_v_shm_keys, - } - owner_path = OWNER_DIR / f"{service_name}.json" - - with FileLock(str(OWNER_LOCK_PATH)): - _recover_stale_services() - tmp_path = owner_path.with_suffix(f".{owner_pid}.tmp") - with tmp_path.open("w", encoding="utf-8") as file: - json.dump(owner, file) - os.replace(tmp_path, owner_path) - - def cleanup(): - nonlocal cleaned - # multiprocessing 子进程会继承 atexit 回调,只有创建记录的 launcher 可以执行全量回收。 - if cleaned or os.getpid() != owner_pid: - return - cleaned = True - cleanup_service_shm(service_name, system_v_shm_keys) - with FileLock(str(OWNER_LOCK_PATH)): - current_owner = _load_owner(owner_path) - if current_owner is not None and current_owner.get("pid") == owner_pid: - try: - owner_path.unlink() - except FileNotFoundError: - pass - - cleaned = False - _registered_cleanups[cleanup_key] = cleanup - atexit.register(cleanup) - return cleanup + """创建 launcher 的清理对象,并返回 cleanup 函数。""" + return ServiceShmCleanup(service_name).register() diff --git a/lightllm/utils/start_utils.py b/lightllm/utils/start_utils.py index 38ae36f4ff..130b6adb46 100644 --- a/lightllm/utils/start_utils.py +++ b/lightllm/utils/start_utils.py @@ -110,7 +110,9 @@ def kill_recursive(proc): def setup_signal_handlers(self, http_server_process=None): # 共享内存由 launcher 在所有子进程退出后统一回收,不再让资源对象注册信号处理器。 - self._cleanup_service_shm = register_launcher_shm_cleanup(get_unique_server_name()) + # 启动 HTTP server 前后都会调用本函数,但每个 launcher 只需要一个清理对象和一个 atexit 回调。 + if self._cleanup_service_shm is None: + self._cleanup_service_shm = register_launcher_shm_cleanup(get_unique_server_name()) def signal_handler(sig, _frame): if sig == signal.SIGINT: diff --git a/unit_tests/utils/test_service_shm_cleanup.py b/unit_tests/utils/test_service_shm_cleanup.py index 9b2ab96b92..d9fecc1360 100644 --- a/unit_tests/utils/test_service_shm_cleanup.py +++ b/unit_tests/utils/test_service_shm_cleanup.py @@ -1,5 +1,4 @@ import json -import os from lightllm.utils import service_shm_cleanup @@ -13,14 +12,14 @@ def test_cleanup_service_shm_only_removes_matching_service(monkeypatch, tmp_path removed_system_v_keys = [] monkeypatch.setattr(service_shm_cleanup, "SHM_DIR", shm_dir) - monkeypatch.setattr(service_shm_cleanup, "_unlink_posix_shm", lambda name: (shm_dir / name).unlink()) monkeypatch.setattr( - service_shm_cleanup, - "_remove_system_v_shm", - lambda key: removed_system_v_keys.append(key) or True, + service_shm_cleanup.ServiceShmCleanup, + "cleanup_system_v_shm", + staticmethod(lambda keys: removed_system_v_keys.extend(keys) or len(keys)), ) - service_shm_cleanup.cleanup_service_shm("service_0", [101, 102]) + 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() @@ -28,46 +27,96 @@ def test_cleanup_service_shm_only_removes_matching_service(monkeypatch, tmp_path assert removed_system_v_keys == [101, 102] -def test_register_launcher_cleanup_recovers_dead_owner_and_records_current_service(monkeypatch, tmp_path): - owner_dir = tmp_path / "owners" - owner_dir.mkdir() - old_owner_path = owner_dir / "old_service_0.json" - old_owner_path.write_text( - json.dumps( - { - "service_name": "old_service_0", - "pid": 999999999, - "create_time": 1.0, - "system_v_shm_keys": [11], - } - ), - encoding="utf-8", +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, "OWNER_DIR", owner_dir) - monkeypatch.setattr(service_shm_cleanup, "OWNER_LOCK_PATH", tmp_path / "owners.lock") - monkeypatch.setattr(service_shm_cleanup, "cleanup_service_shm", lambda *args: cleanup_calls.append(args)) + 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({"cpu_kv_cache_shm_id": 21, "multi_modal_cache_shm_id": 22}), + json.dumps(start_args), ) - service_shm_cleanup._registered_cleanups.clear() cleanup = service_shm_cleanup.register_launcher_shm_cleanup("current_service_0") - assert cleanup_calls == [("old_service_0", [11])] - assert not old_owner_path.exists() - current_owner_path = owner_dir / "current_service_0.json" - current_owner = json.loads(current_owner_path.read_text(encoding="utf-8")) - assert current_owner["pid"] == os.getpid() - assert current_owner["system_v_shm_keys"] == [21, 22] + assert cleanup_calls == [] assert atexit_callbacks == [cleanup] cleanup() cleanup() - assert cleanup_calls[-1] == ("current_service_0", [21, 22]) - assert cleanup_calls.count(("current_service_0", [21, 22])) == 1 - assert not current_owner_path.exists() + 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_start_utils.py b/unit_tests/utils/test_start_utils.py index e3f6e14081..e07ebf99c2 100644 --- a/unit_tests/utils/test_start_utils.py +++ b/unit_tests/utils/test_start_utils.py @@ -160,6 +160,7 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): ) monkeypatch.setattr(process_manager, "terminate_all_processes", lambda: terminate_calls.append(True)) + process_manager.setup_signal_handlers(http_server_process) process_manager.setup_signal_handlers(http_server_process) assert set(registered_handlers) == { From 9f8ba9367c093f52ef41d52a515781efaa60a612 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 10 Sep 2026 08:35:07 +0000 Subject: [PATCH 7/8] refactor: separate exit controller setup from signal registration --- lightllm/server/api_start.py | 13 +++++++------ lightllm/utils/start_utils.py | 12 +++++++++--- unit_tests/utils/test_start_utils.py | 20 +++++++++++++++----- 3 files changed, 31 insertions(+), 14 deletions(-) diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 252cfcd472..e3f06dd0af 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -338,7 +338,7 @@ def _launch_subprocesses(args: StartArgs): validate_ports(ports_to_check) set_env_start_args(args) - process_manager.setup_signal_handlers() + process_manager.setup_exit_controller() get_shm_port_args(create=True) # 多机用于收发node ip, 这个地方修改了args env,所以需要重新设置一下。 send_and_receive_node_ip(args) @@ -446,7 +446,6 @@ def normal_or_p_d_start(args: StartArgs): # 启动子进程 http_server_process = subprocess.Popen(command) - process_manager.setup_signal_handlers(http_server_process) if "s3://" in args.model_dir: from lightllm.utils.petrel_helper import s3_model_clear @@ -457,6 +456,7 @@ def normal_or_p_d_start(args: StartArgs): from lightllm.server.health_monitor.manager import start_health_check_process process_manager.start_submodule_processes(start_funcs=[start_health_check_process], start_args=[(args,)]) + process_manager.setup_signal_handlers(http_server_process) process_manager.supervise_processes(http_server_process) @@ -481,7 +481,7 @@ def pd_master_start(args: StartArgs): validate_ports([args.port]) set_env_start_args(args) - process_manager.setup_signal_handlers() + process_manager.setup_exit_controller() get_shm_port_args(create=True) logger.info(f"all start args:{args}") @@ -509,13 +509,13 @@ def pd_master_start(args: StartArgs): ] http_server_process = subprocess.Popen(command) - process_manager.setup_signal_handlers(http_server_process) if args.health_monitor: from lightllm.server.health_monitor.manager import start_health_check_process process_manager.start_submodule_processes(start_funcs=[start_health_check_process], start_args=[(args,)]) + process_manager.setup_signal_handlers(http_server_process) process_manager.supervise_processes(http_server_process) @@ -547,7 +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_signal_handlers() + process_manager.setup_exit_controller() get_shm_port_args(create=True) logger.info(f"all start args:{args}") @@ -561,6 +561,7 @@ def visual_only_start(args): (args,), ], ) + process_manager.setup_signal_handlers() process_manager.supervise_processes() @@ -574,7 +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_signal_handlers() + process_manager.setup_exit_controller() get_shm_port_args(create=True) logger.info(f"all start args:{args}") diff --git a/lightllm/utils/start_utils.py b/lightllm/utils/start_utils.py index 130b6adb46..b306a00a51 100644 --- a/lightllm/utils/start_utils.py +++ b/lightllm/utils/start_utils.py @@ -108,12 +108,18 @@ def kill_recursive(proc): stop_mps() logger.info("All processes terminated gracefully.") - def setup_signal_handlers(self, http_server_process=None): - # 共享内存由 launcher 在所有子进程退出后统一回收,不再让资源对象注册信号处理器。 - # 启动 HTTP server 前后都会调用本函数,但每个 launcher 只需要一个清理对象和一个 atexit 回调。 + def setup_exit_controller(self): + """初始化 launcher 退出清理控制器,并注册 atexit 兜底回调。 + + 在 service name 和启动参数写入环境后、创建共享内存或启动子进程前调用。 + 重复调用只注册一次;退出时由 launcher 在子进程停止后统一回收共享内存。 + """ if self._cleanup_service_shm is None: self._cleanup_service_shm = register_launcher_shm_cleanup(get_unique_server_name()) + 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/utils/test_start_utils.py b/unit_tests/utils/test_start_utils.py index e07ebf99c2..608cf04b04 100644 --- a/unit_tests/utils/test_start_utils.py +++ b/unit_tests/utils/test_start_utils.py @@ -147,7 +147,6 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): 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", @@ -160,7 +159,6 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): ) monkeypatch.setattr(process_manager, "terminate_all_processes", lambda: terminate_calls.append(True)) - process_manager.setup_signal_handlers(http_server_process) process_manager.setup_signal_handlers(http_server_process) assert set(registered_handlers) == { @@ -168,7 +166,7 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): start_utils.signal.SIGINT, start_utils.signal.SIGHUP, } - assert cleanup_calls == [("register", "service_0")] + assert cleanup_calls == [] with pytest.raises(SystemExit) as exc_info: registered_handlers[start_utils.signal.SIGTERM](start_utils.signal.SIGTERM, None) @@ -178,14 +176,26 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): assert terminate_calls == [True] -def test_terminate_all_processes_runs_launcher_shm_cleanup(monkeypatch): +def test_setup_exit_controller_registers_once_and_cleans_up_on_termination(monkeypatch): from lightllm.utils import envs_utils process_manager = start_utils.SubmoduleManager() + registration_calls = [] cleanup_calls = [] - process_manager._cleanup_service_shm = lambda: cleanup_calls.append(True) + 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: cleanup_calls.append(True)), + ) monkeypatch.setattr(envs_utils, "get_env_start_args", lambda: SimpleNamespace(enable_mps=False)) + process_manager.setup_exit_controller() + process_manager.setup_exit_controller() + + assert registration_calls == ["service_0"] + assert cleanup_calls == [] + process_manager.terminate_all_processes() assert cleanup_calls == [True] From 8a2ef663bc800af2e52dd8858fb9b2e097e1b485 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 10 Sep 2026 09:05:35 +0000 Subject: [PATCH 8/8] fix: handle shutdown signals during launcher initialization --- lightllm/utils/start_utils.py | 19 ++- unit_tests/utils/test_atexit_exit_modes.py | 188 +++++++++++++++++++++ unit_tests/utils/test_start_utils.py | 49 +++++- 3 files changed, 244 insertions(+), 12 deletions(-) create mode 100644 unit_tests/utils/test_atexit_exit_modes.py diff --git a/lightllm/utils/start_utils.py b/lightllm/utils/start_utils.py index b306a00a51..5f816f2832 100644 --- a/lightllm/utils/start_utils.py +++ b/lightllm/utils/start_utils.py @@ -109,16 +109,27 @@ def kill_recursive(proc): logger.info("All processes terminated gracefully.") def setup_exit_controller(self): - """初始化 launcher 退出清理控制器,并注册 atexit 兜底回调。 + """初始化 launcher 退出清理控制器,注册启动阶段信号处理和 atexit 回调。 在 service name 和启动参数写入环境后、创建共享内存或启动子进程前调用。 重复调用只注册一次;退出时由 launcher 在子进程停止后统一回收共享内存。 + 启动完成后由 setup_signal_handlers 替换信号处理函数,纳入 HTTP server 的退出流程。 """ - if self._cleanup_service_shm is None: - self._cleanup_service_shm = register_launcher_shm_cleanup(get_unique_server_name()) + 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: 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_start_utils.py b/unit_tests/utils/test_start_utils.py index 608cf04b04..250c41833f 100644 --- a/unit_tests/utils/test_start_utils.py +++ b/unit_tests/utils/test_start_utils.py @@ -141,12 +141,14 @@ 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", @@ -159,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) == { @@ -166,7 +171,12 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): start_utils.signal.SIGINT, start_utils.signal.SIGHUP, } - assert cleanup_calls == [] + 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) @@ -176,29 +186,52 @@ def test_setup_signal_handlers_registers_and_handles_sigterm(monkeypatch): assert terminate_calls == [True] -def test_setup_exit_controller_registers_once_and_cleans_up_on_termination(monkeypatch): +@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 = [] - cleanup_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: cleanup_calls.append(True)), + 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 cleanup_calls == [] + assert set(registered_handlers) == { + start_utils.signal.SIGTERM, + start_utils.signal.SIGINT, + start_utils.signal.SIGHUP, + } + assert registered_handlers == initial_handlers + assert shutdown_events == [] - process_manager.terminate_all_processes() + with pytest.raises(SystemExit) as exc_info: + registered_handlers[shutdown_signal](shutdown_signal, None) - assert cleanup_calls == [True] + assert exc_info.value.code == 0 + assert shutdown_events == ["kill", "wait", "cleanup"] def test_supervisor_fails_when_http_server_exits(monkeypatch):