diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index c997b436e..75dca450f 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -105,7 +105,11 @@ jobs: [ str(runtime_python), "-c", - "import art_vllm_runtime, torch, vllm; print('runtime imports ok')", + "import torch, vllm; " + "from art_vllm_runtime.policy_spans import " + "_patch_lora_update_coordinator; " + "_patch_lora_update_coordinator(); " + "print('runtime compatibility ok')", ], check=True, ) diff --git a/tests/integration/megatron/runtime_isolation/test_runtime_project_isolation.py b/tests/integration/megatron/runtime_isolation/test_runtime_project_isolation.py index 4eddbb176..5a9c57dcb 100644 --- a/tests/integration/megatron/runtime_isolation/test_runtime_project_isolation.py +++ b/tests/integration/megatron/runtime_isolation/test_runtime_project_isolation.py @@ -248,6 +248,60 @@ def test_runtime_general_plugin_loads_full_patch_set() -> None: assert 'art = "art_vllm_runtime.patches:apply_vllm_runtime_patches"' in pyproject +def test_lora_coordinator_supports_both_vllm_serving_layouts( + artifact_dir: Path, +) -> None: + payload = _runtime_python( + """ +import json +from types import SimpleNamespace +import art_vllm_runtime.policy_spans as policy +from vllm.entrypoints.openai.engine.serving import OpenAIServing + +policy._patch_lora_update_coordinator() +legacy_patched = getattr( + OpenAIServing.__init__, "__art_lora_update_patched__", False +) + +class GenerateBaseServing: + def __init__(self, models, engine_client): + self.models = models + self.engine_client = engine_client + +real_import_module = policy.importlib.import_module +def import_module(name): + if name == "vllm.entrypoints.openai.engine.serving": + raise ModuleNotFoundError(name, name=name) + if name == "vllm.entrypoints.generate.base.serving": + return SimpleNamespace(GenerateBaseServing=GenerateBaseServing) + return real_import_module(name) + +policy.importlib.import_module = import_module +policy._patch_lora_update_coordinator() +models = SimpleNamespace() +engine_client = SimpleNamespace() +GenerateBaseServing(models, engine_client) +print(json.dumps({ + "legacy_patched": legacy_patched, + "new_patched": getattr( + GenerateBaseServing.__init__, "__art_lora_update_patched__", False + ), + "shared_coordinator": ( + models._art_lora_update_coordinator + is engine_client._art_lora_update_coordinator + ), +})) +""", + artifact_dir, + "vllm_serving_layouts", + ) + assert json.loads(payload.splitlines()[-1]) == { + "legacy_patched": True, + "new_patched": True, + "shared_coordinator": True, + } + + def test_runtime_patch_adds_gemma4_moe_topk_alias(artifact_dir: Path) -> None: payload = _runtime_python( "import json; " diff --git a/vllm_runtime/src/art_vllm_runtime/policy_spans.py b/vllm_runtime/src/art_vllm_runtime/policy_spans.py index da90c7192..5e186d721 100644 --- a/vllm_runtime/src/art_vllm_runtime/policy_spans.py +++ b/vllm_runtime/src/art_vllm_runtime/policy_spans.py @@ -526,9 +526,16 @@ async def tracked_result_generator(): def _patch_lora_update_coordinator() -> None: - from vllm.entrypoints.openai.engine.serving import OpenAIServing - - original_init = OpenAIServing.__init__ + try: + module = importlib.import_module("vllm.entrypoints.openai.engine.serving") + serving_base = module.OpenAIServing + except ModuleNotFoundError as exc: + if exc.name != "vllm.entrypoints.openai.engine.serving": + raise + module = importlib.import_module("vllm.entrypoints.generate.base.serving") + serving_base = module.GenerateBaseServing + + original_init = serving_base.__init__ if not getattr(original_init, "__art_lora_update_patched__", False): def __init__(self: Any, *args: Any, **kwargs: Any) -> None: @@ -536,7 +543,7 @@ def __init__(self: Any, *args: Any, **kwargs: Any) -> None: lora_update_coordinator(self.models, self.engine_client) __init__.__art_lora_update_patched__ = True # type: ignore[attr-defined] - OpenAIServing.__init__ = __init__ # type: ignore[method-assign] + serving_base.__init__ = __init__ def _patch_engine_request_admission() -> None: