From df2636a9d8139dd1d53ec1f1e666abe66fb76a3c Mon Sep 17 00:00:00 2001 From: David Zheng <153074367+dzhengAP@users.noreply.github.com> Date: Fri, 8 May 2026 21:32:04 -0700 Subject: [PATCH] [Bugfix] Fix LOGITPROC_SOURCE_ENTRYPOINT test to use spawn-compatible dist-info registration for XPU/ROCm (#42040) Signed-off-by: dqzhengAP Signed-off-by: David Zheng <153074367+dzhengAP@users.noreply.github.com> Signed-off-by: Andreas Karatzas Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> Co-authored-by: Andreas Karatzas Co-authored-by: Kunshang Ji --- tests/utils.py | 67 +++++++++++++++--- .../logits_processors/test_custom_offline.py | 23 ++----- .../logits_processors/test_custom_online.py | 19 ++---- tests/v1/logits_processors/utils.py | 68 ++++++++++++++++--- 4 files changed, 129 insertions(+), 48 deletions(-) diff --git a/tests/utils.py b/tests/utils.py index cff601374b0..68a9031a2c4 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -20,9 +20,9 @@ import time import warnings from collections.abc import Callable, Iterable, Sequence from contextlib import ExitStack, contextmanager -from multiprocessing import Process +from multiprocessing import Process, get_context from pathlib import Path -from typing import Any, Literal +from typing import Any, Literal, cast from unittest.mock import patch import anthropic @@ -127,6 +127,21 @@ ROCM_ENGINE_KWARGS: dict = ( ) +def requires_spawn_multiprocessing() -> bool: + """Whether this platform requires spawn instead of fork for test processes.""" + return current_platform.is_rocm() or current_platform.is_xpu() + + +def _run_in_new_process_group( + child_process_fxn: Callable[[dict[str, str] | None, str, list[str]], None], + env_dict: dict[str, str] | None, + model: str, + vllm_serve_args: list[str], +) -> None: + os.setsid() + child_process_fxn(env_dict, model, vllm_serve_args) + + class RemoteVLLMServer: """Base class for launching vLLM server subprocesses for testing. @@ -738,8 +753,11 @@ class RemoteOpenAIServerCustom(RemoteOpenAIServer): def _start_server( self, model: str, vllm_serve_args: list[str], env_dict: dict[str, str] | None ) -> None: - self.proc: Process = Process( - target=self.child_process_fxn, args=(env_dict, model, vllm_serve_args) + method = "spawn" if requires_spawn_multiprocessing() else "fork" + ctx = get_context(method) + self.proc: Process = cast(Any, ctx).Process( + target=_run_in_new_process_group, + args=(self.child_process_fxn, env_dict, model, vllm_serve_args), ) # type: ignore[assignment] self.proc.start() @@ -769,12 +787,40 @@ class RemoteOpenAIServerCustom(RemoteOpenAIServer): def _poll(self) -> int | None: return self.proc.exitcode - def __exit__(self, exc_type, exc_value, traceback): - self.proc.terminate() - self.proc.join(8) + def _terminate_process_tree(self) -> None: + pid = self.proc.pid + if pid is None: + return + + pgid: int | None + try: + pgid = os.getpgid(pid) + # _run_in_new_process_group should make the child the group + # leader. Avoid signaling pytest's process group if startup failed + # before os.setsid() ran. + if pgid != pid: + pgid = None + except (ProcessLookupError, OSError): + pgid = None + + with contextlib.suppress(ProcessLookupError, OSError): + self.proc.terminate() + print(f"[RemoteOpenAIServerCustom] Sent SIGTERM to process {pid}") + + self.proc.join(15) if self.proc.is_alive(): - # force kill if needed - self.proc.kill() + print( + f"[RemoteOpenAIServerCustom] Server {pid} did not respond " + "to SIGTERM, sending SIGKILL to process group" + ) + if pgid is not None: + with contextlib.suppress(ProcessLookupError, OSError): + os.killpg(pgid, signal.SIGKILL) + else: + self.proc.kill() + self.proc.join(10) + + self._kill_process_group_survivors(pgid) def _test_completion( @@ -1633,8 +1679,7 @@ def create_new_process_for_each_test( A decorator to run test functions in separate processes. """ if method is None: - use_spawn = current_platform.is_rocm() or current_platform.is_xpu() - method = "spawn" if use_spawn else "fork" + method = "spawn" if requires_spawn_multiprocessing() else "fork" assert method in ["spawn", "fork"], "Method must be either 'spawn' or 'fork'" diff --git a/tests/v1/logits_processors/test_custom_offline.py b/tests/v1/logits_processors/test_custom_offline.py index 29ec72186b8..325ca48b597 100644 --- a/tests/v1/logits_processors/test_custom_offline.py +++ b/tests/v1/logits_processors/test_custom_offline.py @@ -16,8 +16,8 @@ from tests.v1.logits_processors.utils import ( DummyLogitsProcessor, WrappedPerReqLogitsProcessor, prompts, + setup_fake_entrypoint, ) -from tests.v1.logits_processors.utils import entry_points as fake_entry_points from vllm import LLM, SamplingParams from vllm.v1.sample.logits_processor import ( STR_POOLING_REJECTS_LOGITSPROCS, @@ -145,13 +145,9 @@ def test_custom_logitsprocs(monkeypatch, logitproc_source: CustomLogitprocSource if logitproc_source == CustomLogitprocSource.LOGITPROC_SOURCE_ENTRYPOINT: # Scenario: vLLM loads a logitproc from a preconfigured entrypoint - # To that end, mock a dummy logitproc entrypoint - import importlib.metadata - - importlib.metadata.entry_points = fake_entry_points # type: ignore - - # fork is required for workers to see entrypoint patch - monkeypatch.setenv("VLLM_WORKER_MULTIPROC_METHOD", "fork") + # To that end, register a real dist-info package so spawned + # workers can discover the entrypoint via PYTHONPATH + setup_fake_entrypoint(monkeypatch) _run_test({}, logitproc_loaded=True) return @@ -266,14 +262,9 @@ def test_rejects_custom_logitsprocs( # Scenario: vLLM loads a model and ignores a logitproc that is # available at a preconfigured entrypoint - # Patch in dummy logitproc entrypoint - import importlib.metadata - - importlib.metadata.entry_points = fake_entry_points # type: ignore - - # fork is required for entrypoint patch to be visible to workers, - # although they should ignore the entrypoint patch anyway - monkeypatch.setenv("VLLM_WORKER_MULTIPROC_METHOD", "fork") + # Register real dist-info package so spawned workers can + # discover the entrypoint via PYTHONPATH (spawn-compatible) + setup_fake_entrypoint(monkeypatch) llm = LLM(**llm_kwargs) # Require that no custom logitsprocs have been loaded diff --git a/tests/v1/logits_processors/test_custom_online.py b/tests/v1/logits_processors/test_custom_online.py index 05ac7034973..3b7add3b80f 100644 --- a/tests/v1/logits_processors/test_custom_online.py +++ b/tests/v1/logits_processors/test_custom_online.py @@ -18,8 +18,8 @@ from tests.v1.logits_processors.utils import ( MODEL_NAME, TEMP_GREEDY, prompts, + setup_fake_entrypoint, ) -from tests.v1.logits_processors.utils import entry_points as fake_entry_points def _server_with_logitproc_entrypoint( @@ -27,16 +27,9 @@ def _server_with_logitproc_entrypoint( model: str, vllm_serve_args: list[str], ) -> None: - """Start vLLM server, inject dummy logitproc entrypoint""" - - # Patch `entry_points` to inject logitproc entrypoint - import importlib.metadata - - importlib.metadata.entry_points = fake_entry_points # type: ignore + """Start vLLM server with dummy logitproc entrypoint.""" from vllm.entrypoints.cli import main - # fork is required for workers to see entrypoint patch - os.environ["VLLM_WORKER_MULTIPROC_METHOD"] = "fork" if env_dict is not None: os.environ.update(env_dict) @@ -50,7 +43,7 @@ def _server_with_logitproc_fqcn( model: str, vllm_serve_args: list[str], ) -> None: - """Start vLLM server, inject module with dummy logitproc""" + """Start vLLM server with dummy logitproc specified by FQCN.""" from vllm.entrypoints.cli import main if env_dict is not None: @@ -80,8 +73,8 @@ def default_server_args(): def server(default_server_args, request, monkeypatch): """Consider two server configurations: (1) --logits-processors cli arg specifies dummy logits processor via fully- - qualified class name (FQCN); patch in a dummy logits processor module - (2) No --logits-processors cli arg; patch in a dummy logits processor + qualified class name (FQCN) + (2) No --logits-processors cli arg; inject a dummy logits processor entrypoint """ @@ -94,6 +87,7 @@ def server(default_server_args, request, monkeypatch): _server_fxn = _server_with_logitproc_fqcn else: # Launch server, inject dummy logitproc entrypoint + setup_fake_entrypoint(monkeypatch) args = default_server_args _server_fxn = _server_with_logitproc_entrypoint @@ -119,7 +113,6 @@ api_keyword_args = { } -@pytest.mark.asyncio @pytest.mark.parametrize( "model_name", [MODEL_NAME], diff --git a/tests/v1/logits_processors/utils.py b/tests/v1/logits_processors/utils.py index e54da72e5e2..fc8ce50c05f 100644 --- a/tests/v1/logits_processors/utils.py +++ b/tests/v1/logits_processors/utils.py @@ -1,12 +1,15 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -import types +import os +import tempfile from enum import Enum, auto +from pathlib import Path from typing import Any import torch +from tests.utils import requires_spawn_multiprocessing from vllm.config import VllmConfig from vllm.logger import init_logger from vllm.sampling_params import SamplingParams @@ -102,11 +105,6 @@ class DummyLogitsProcessor(LogitsProcessor): return logits -"""Dummy module with dummy logitproc class""" -dummy_module = types.ModuleType(DUMMY_LOGITPROC_MODULE) -dummy_module.DummyLogitsProcessor = DummyLogitsProcessor # type: ignore - - class EntryPoint: """Dummy entrypoint class for logitsprocs testing""" @@ -187,5 +185,59 @@ class WrappedPerReqLogitsProcessor(AdapterLogitsProcessor): return DummyPerReqLogitsProcessor(target_token) -"""Fake version of importlib.metadata.entry_points""" -entry_points = lambda group: EntryPoints(group) +def register_fake_entrypoint(monkeypatch) -> str: + """Register the dummy logitsproc entrypoint in a way that is visible + to spawned subprocesses by creating a real dist-info directory on disk. + + Unlike monkey-patching importlib.metadata.entry_points (which only works + with fork), this approach writes a real dist-info package that + importlib.metadata can discover in any subprocess via PYTHONPATH. + + Returns the temp directory path. + """ + tmpdir = Path(tempfile.mkdtemp(prefix="dummy-logitproc-")) + dist_info = tmpdir / "dummy_logitproc-0.1.dist-info" + dist_info.mkdir() + + # Write METADATA file (required by importlib.metadata) + (dist_info / "METADATA").write_text( + "Metadata-Version: 2.1\nName: dummy-logitproc\nVersion: 0.1\n", + encoding="utf-8", + ) + + # Write entry_points.txt + (dist_info / "entry_points.txt").write_text( + f"[{LOGITSPROCS_GROUP}]\n" + f"{DUMMY_LOGITPROC_ENTRYPOINT} = {DUMMY_LOGITPROC_FQCN}\n", + encoding="utf-8", + ) + + # Add to PYTHONPATH so spawned subprocesses can discover it + existing = os.environ.get("PYTHONPATH", "") + monkeypatch.setenv( + "PYTHONPATH", str(tmpdir) + (os.pathsep + existing if existing else "") + ) + + # Also update sys.path for the current process so the driver can + # discover the entrypoint. + monkeypatch.syspath_prepend(str(tmpdir)) + + return str(tmpdir) + + +def fake_entry_points(group: str) -> EntryPoints: + """Fake version of importlib.metadata.entry_points.""" + return EntryPoints(group) + + +def setup_fake_entrypoint(monkeypatch) -> None: + """Expose the dummy logitproc entrypoint for the current platform.""" + if requires_spawn_multiprocessing(): + register_fake_entrypoint(monkeypatch) + monkeypatch.setenv("VLLM_WORKER_MULTIPROC_METHOD", "spawn") + return + + import importlib.metadata + + monkeypatch.setattr(importlib.metadata, "entry_points", fake_entry_points) + monkeypatch.setenv("VLLM_WORKER_MULTIPROC_METHOD", "fork")