diff --git a/components/src/dynamo/common/rl/admin.py b/components/src/dynamo/common/rl/admin.py index e99533fb0358..cc7e81bef744 100644 --- a/components/src/dynamo/common/rl/admin.py +++ b/components/src/dynamo/common/rl/admin.py @@ -93,9 +93,17 @@ def __init__( runtime: Any, *, logger_: logging.Logger | None = None, + world_size: int | None = None, ) -> None: + if world_size is not None and ( + isinstance(world_size, bool) + or not isinstance(world_size, int) + or world_size < 1 + ): + raise ValueError("world_size must be a positive integer") self._runtime = runtime self._logger = logger_ or logger + self._world_size = world_size self.routes: dict[str, RLRouteHandler] = {} def add_route(self, name: str, handler: RLRouteHandler) -> None: @@ -117,6 +125,9 @@ def describe(self) -> dict[str, Any]: if system_url: response["system_url"] = system_url + if self._world_size is not None: + response["world_size"] = self._world_size + return response async def dispatch( diff --git a/components/src/dynamo/common/tests/test_rl_admin.py b/components/src/dynamo/common/tests/test_rl_admin.py index 176005cd2de5..65c0f3bf116d 100644 --- a/components/src/dynamo/common/tests/test_rl_admin.py +++ b/components/src/dynamo/common/tests/test_rl_admin.py @@ -53,6 +53,23 @@ async def ping(body: dict) -> dict: assert routes_with_kwargs == routes +def test_route_registry_describes_world_size() -> None: + registry = RLRouteRegistry(_Runtime("http://worker:8081"), world_size=4) + + assert registry.describe() == { + "status": "ok", + "routes": [], + "system_url": "http://worker:8081", + "world_size": 4, + } + + +@pytest.mark.parametrize("world_size", [True, 0, -1, 1.5, "4"]) +def test_route_registry_rejects_invalid_world_size(world_size) -> None: + with pytest.raises(ValueError, match="world_size"): + RLRouteRegistry(_Runtime(), world_size=world_size) + + def test_route_registry_rejects_request_plane_admin_execution() -> None: registry = RLRouteRegistry(_Runtime()) diff --git a/components/src/dynamo/vllm/engine_generate.py b/components/src/dynamo/vllm/engine_generate.py index 2390a5930eeb..332042378bc8 100644 --- a/components/src/dynamo/vllm/engine_generate.py +++ b/components/src/dynamo/vllm/engine_generate.py @@ -1,14 +1,26 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 - """Runtime capability metadata for vLLM's native Generate API.""" +from __future__ import annotations + import json +from dataclasses import dataclass +from functools import lru_cache +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + from vllm.multimodal.inputs import MultiModalKwargsItems, PlaceholderRange + from vllm.sampling_params import SamplingParams -from dynamo.llm import ModelInput, ModelRuntimeConfig, ModelType, WorkerType +from dynamo.common.utils.guided_json import reject_nonprogressing_guided_json_ref_cycles +from dynamo.llm import HttpError, ModelInput, ModelRuntimeConfig, ModelType, WorkerType + +from .kv_hints import _apply_kv_hint VLLM_GENERATE_CAPABILITY = "vllm_inference_v1_generate" VLLM_ENABLE_TOWER_CONNECTOR_LORA_RUNTIME_KEY = "vllm_enable_tower_connector_lora" +DYNAMO_CACHE_SALT_PREFIX = "dynamo-cache-salt:" def publish_engine_generate_capability( @@ -19,15 +31,12 @@ def publish_engine_generate_capability( tower_connector_lora_enabled: bool, ) -> bool: """Publish native Generate support and its MM-routing-relevant config.""" - if model_input != ModelInput.Tokens: + if model_input != ModelInput.Tokens or worker_type not in ( + WorkerType.Aggregated, + WorkerType.Decode, + ): return False - if worker_type == WorkerType.Prefill: - supported = model_type == ModelType.Prefill - else: - supported = worker_type in (WorkerType.Decode, WorkerType.Aggregated) and ( - model_type.supports_chat() or model_type == ModelType.Completions - ) - if not supported: + if not (model_type.supports_chat() or model_type == ModelType.Completions): return False runtime_config.set_engine_specific( @@ -39,3 +48,261 @@ def publish_engine_generate_capability( json.dumps(tower_connector_lora_enabled), ) return True + + +@dataclass(frozen=True) +class EngineGenerateInput: + prompt: Any + sampling_params: SamplingParams + priority: int + + +@lru_cache(maxsize=1) +def _native_generate_api() -> tuple[Any, Any]: + try: + from vllm.entrypoints.scale_out.token_in_token_out.mm_serde import ( + decode_mm_kwargs_item, + ) + from vllm.entrypoints.scale_out.token_in_token_out.protocol import ( + GenerateRequest, + ) + except ModuleNotFoundError as exc: + expected_module = "vllm.entrypoints.scale_out.token_in_token_out" + if exc.name is None or not expected_module.startswith(exc.name): + raise + from vllm.entrypoints.serve.disagg.mm_serde import decode_mm_kwargs_item + from vllm.entrypoints.serve.disagg.protocol import GenerateRequest + + return decode_mm_kwargs_item, GenerateRequest + + +def _image_features( + request: dict[str, Any], + features: dict[str, Any], +) -> tuple[ + dict[str, list[str]], + MultiModalKwargsItems, + dict[str, list[PlaceholderRange]], +]: + import torch + from vllm.multimodal.inputs import ( + MultiModalKwargsItem, + MultiModalKwargsItems, + PlaceholderRange, + ) + + mm_hashes = features.get("mm_hashes") + mm_placeholders = features.get("mm_placeholders") + kwargs_data = features.get("kwargs_data") + if not isinstance(mm_hashes, dict) or not isinstance(mm_placeholders, dict): + raise TypeError("TITO features require mm_hashes and mm_placeholders objects") + if not isinstance(kwargs_data, dict): + raise TypeError("TITO features kwargs_data must be an object") + + modalities = set(mm_hashes) | set(mm_placeholders) + modalities.update(kwargs_data) + if modalities != {"image"}: + raise ValueError("TITO preprocessed features currently support image only") + + hashes = mm_hashes.get("image") + ranges = mm_placeholders.get("image") + if not isinstance(hashes, list) or not isinstance(ranges, list): + raise TypeError("TITO image hashes and placeholders must be lists") + if len(hashes) != len(ranges): + raise ValueError("TITO image hash and placeholder counts must match") + if not all(isinstance(value, str) and value for value in hashes): + raise ValueError("TITO image hashes must be non-empty strings") + + routing_hashes = (request.get("extra_args") or {}).get("dynamo_mm_routing_hashes") + if routing_hashes is not None: + if ( + not isinstance(routing_hashes, list) + or len(routing_hashes) != len(hashes) + or not all(isinstance(value, str) and value for value in routing_hashes) + ): + raise ValueError("TITO image routing hash count or value is invalid") + hashes = routing_hashes + + prompt_length = len(request["token_ids"]) + restored_ranges: list[PlaceholderRange] = [] + for item in ranges: + if not isinstance(item, dict): + raise TypeError("TITO image placeholders must be objects") + offset = item.get("offset") + length = item.get("length") + if ( + isinstance(offset, bool) + or not isinstance(offset, int) + or offset < 0 + or isinstance(length, bool) + or not isinstance(length, int) + or length < 1 + or offset + length > prompt_length + ): + raise ValueError("TITO image placeholder range is invalid") + is_embed_raw = item.get("is_embed") + if is_embed_raw is not None: + if not isinstance(is_embed_raw, (list, tuple)): + raise TypeError("TITO image placeholder is_embed must be a sequence") + if len(is_embed_raw) != length: + raise ValueError( + "TITO image placeholder is_embed must match placeholder length" + ) + if not all(isinstance(value, bool) for value in is_embed_raw): + raise ValueError( + "TITO image placeholder is_embed values must be booleans" + ) + is_embed = ( + None + if is_embed_raw is None + else torch.as_tensor(is_embed_raw, dtype=torch.bool) + ) + restored_ranges.append( + PlaceholderRange(offset=offset, length=length, is_embed=is_embed) + ) + + restored_kwargs: list[MultiModalKwargsItem] + image_data = kwargs_data.get("image") + if not isinstance(image_data, list) or len(image_data) != len(hashes): + raise ValueError("TITO image tensor and hash counts must match") + decode_mm_kwargs_item, _ = _native_generate_api() + restored_kwargs = [decode_mm_kwargs_item(value) for value in image_data] + + return ( + {"image": hashes}, + MultiModalKwargsItems({"image": restored_kwargs}), + {"image": restored_ranges}, + ) + + +def adapt_engine_generate_request( + request: dict[str, Any], + *, + enable_multimodal: bool, + decode_capable: bool, + vllm_config: Any, + default_sampling_params: dict[str, Any], + allow_multimodal_features: bool = True, +) -> EngineGenerateInput | None: + """Adapt one Rust-frontend TITO envelope at the Python engine boundary.""" + import msgspec + from vllm.inputs import TokensPrompt, mm_input + from vllm.sampling_params import RequestOutputKind, SamplingParams + + extra_args = request.get("extra_args") + if not isinstance(extra_args, dict) or "vllm_tito" not in extra_args: + return None + envelope = extra_args["vllm_tito"] + if not isinstance(envelope, dict): + raise TypeError("extra_args.vllm_tito must be an object") + if not decode_capable: + raise ValueError("TITO requests require an aggregated or decode vLLM worker") + if envelope.get("content_parts"): + raise ValueError("TITO raw multimodal content_parts are not supported") + + raw_sampling_params = envelope.get("sampling_params") + if not isinstance(raw_sampling_params, dict): + raise TypeError("extra_args.vllm_tito.sampling_params must be an object") + features = envelope.get("features") + if features is not None and not allow_multimodal_features: + raise ValueError("TITO multimodal features require an aggregated vLLM worker") + if isinstance(features, dict): + kwargs_data = features.get("kwargs_data") + if not isinstance(kwargs_data, dict): + raise TypeError("TITO features kwargs_data must be an object") + token_ids = list(request.get("token_ids") or []) + raw_prompt_start = raw_sampling_params.get("routed_experts_prompt_start", 0) + if ( + isinstance(raw_prompt_start, bool) + or not isinstance(raw_prompt_start, int) + or raw_prompt_start < 0 + ): + raise ValueError( + "sampling_params.routed_experts_prompt_start must be a non-negative integer" + ) + if raw_prompt_start >= len(token_ids): + raise ValueError( + "sampling_params.routed_experts_prompt_start must be smaller than " + "the prompt token count" + ) + reconstructed = {**envelope, "token_ids": token_ids} + _, generate_request_type = _native_generate_api() + native_request = generate_request_type.model_validate(reconstructed) + sampling_params = native_request.sampling_params + if isinstance(sampling_params, dict): + sampling_params = msgspec.convert(sampling_params, type=SamplingParams) + if not isinstance(sampling_params, SamplingParams): + raise TypeError("vLLM GenerateRequest returned invalid sampling_params") + + structured_outputs = sampling_params.structured_outputs + if structured_outputs is not None and structured_outputs.json is not None: + try: + reject_nonprogressing_guided_json_ref_cycles(structured_outputs.json) + except HttpError as exc: + raise ValueError(str(exc)) from exc + + if native_request.kv_transfer_params is not None: + sampling_params.extra_args = { + **(sampling_params.extra_args or {}), + "kv_transfer_params": native_request.kv_transfer_params, + } + _apply_kv_hint(sampling_params, request.get("kv_hint")) + if not sampling_params.stop: + sampling_params.detokenize = False + max_num_seqs = vllm_config.scheduler_config.max_num_seqs + if sampling_params.n > max_num_seqs: + raise ValueError( + "sampling_params.n must be at most the server's max_num_seqs " + f"({max_num_seqs}), got {sampling_params.n}." + ) + if not native_request.is_sampling_param_provided("max_tokens"): + # Older supported vLLM builds do not expose this helper. Keep the import + # lazy so workers that do not adapt native Generate requests still start. + from vllm.entrypoints.serve.utils.api_utils import get_max_tokens + + model_config = vllm_config.model_config + generation_config = getattr(model_config, "generation_config", "vllm") + override_generation_config = ( + getattr(model_config, "override_generation_config", {}) or {} + ) + override_max_tokens = ( + default_sampling_params.get("max_tokens") + if generation_config not in ("auto", "vllm") + else override_generation_config.get("max_new_tokens") + ) + sampling_params.max_tokens = get_max_tokens( + max_model_len=model_config.max_model_len, + max_tokens=None, + input_length=len(token_ids), + default_sampling_params=default_sampling_params, + override_max_tokens=override_max_tokens, + ) + sampling_params.output_kind = RequestOutputKind.DELTA + + cache_salt = envelope.get("cache_salt") + engine_cache_salt = ( + f"{DYNAMO_CACHE_SALT_PREFIX}{cache_salt}" if cache_salt else None + ) + if features is None: + prompt = TokensPrompt(prompt_token_ids=token_ids) + if engine_cache_salt is not None: + prompt["cache_salt"] = engine_cache_salt + else: + if not enable_multimodal: + raise ValueError("TITO multimodal features require --enable-multimodal") + if not isinstance(features, dict): + raise ValueError("extra_args.vllm_tito.features must be an object") + mm_hashes, mm_kwargs, mm_placeholders = _image_features(request, features) + prompt = mm_input( + prompt_token_ids=token_ids, + mm_kwargs=mm_kwargs, + mm_hashes=mm_hashes, + mm_placeholders=mm_placeholders, + cache_salt=engine_cache_salt, + ) + + return EngineGenerateInput( + prompt=prompt, + sampling_params=sampling_params, + priority=native_request.priority, + ) diff --git a/components/src/dynamo/vllm/handlers.py b/components/src/dynamo/vllm/handlers.py index 36305c283051..c83aa199e643 100644 --- a/components/src/dynamo/vllm/handlers.py +++ b/components/src/dynamo/vllm/handlers.py @@ -91,13 +91,17 @@ KvConnectorProtocol, make_kv_connector_protocol, ) -from dynamo.vllm.kv_hints import publish_kv_hint_capabilities +from dynamo.vllm.kv_hints import _apply_kv_hint, publish_kv_hint_capabilities from .args import Config from .cache_info import get_configured_kv_event_block_size from .capacity import publish_vllm_token_budget from .constants import DisaggregationMode, EmbeddingTransferMode from .dp_topology import get_dp_range_for_worker +from .engine_generate import ( + adapt_engine_generate_request, + publish_engine_generate_capability, +) from .engine_monitor import VllmEngineMonitor from .lora_state import LoRAState from .multimodal_utils.custom_encoder import ( @@ -524,6 +528,14 @@ def _attach_routed_experts_engine_data( engine_data["routed_experts"] = routed_experts +def _attach_sampling_mask_engine_data( + tok: Dict[str, Any], sampling_mask: list[list[int]] +) -> None: + engine_data = tok.setdefault("engine_data", {}) + if isinstance(engine_data, dict): + engine_data["sampling_mask"] = sampling_mask + + def _iter_nvext_sources(request: Dict[str, Any]) -> Iterator[Dict[str, Any]]: """Yield each nvext dict on the request, in priority order: @@ -573,6 +585,12 @@ def _apply_nvext_cache_salt(request: Dict[str, Any], prompt: Any) -> None: if cache_salt: prompt["cache_salt"] = f"{_DYNAMO_CACHE_SALT_PREFIX}{cache_salt}" return + extra_args = request.get("extra_args") + vllm_tito = extra_args.get("vllm_tito") if isinstance(extra_args, dict) else None + if isinstance(vllm_tito, dict): + cache_salt = vllm_tito.get("cache_salt") + if cache_salt: + prompt["cache_salt"] = f"{_DYNAMO_CACHE_SALT_PREFIX}{cache_salt}" def _prompt_token_ids_for_engine_data( @@ -894,27 +912,6 @@ def build_sampling_params( return sampling_params -def _apply_kv_hint(sampling_params: SamplingParams, kv_hint: Any) -> None: - """Attach the complete Dynamo KV hint message to vLLM's private input.""" - if not isinstance(kv_hint, Mapping): - return - - extra_args = ( - dict(sampling_params.extra_args) - if isinstance(sampling_params.extra_args, dict) - else {} - ) - existing_kv_transfer_params = extra_args.get(_KV_TRANSFER_PARAMS_EXTRA_ARGS_KEY) - kv_transfer_params = ( - dict(existing_kv_transfer_params) - if isinstance(existing_kv_transfer_params, dict) - else {} - ) - kv_transfer_params[_KV_HINT_EXTRA_ARGS_KEY] = dict(kv_hint) - extra_args[_KV_TRANSFER_PARAMS_EXTRA_ARGS_KEY] = kv_transfer_params - sampling_params.extra_args = extra_args - - def _update_kv_transfer_params( sampling_params: SamplingParams, kv_transfer_params: Mapping[str, Any], @@ -1098,6 +1095,17 @@ def apply_data_parallel_runtime_config( runtime_config.data_parallel_size = dp_range[1] +def resolve_rl_weight_world_size(parallel_config: Any) -> int: + if ( + parallel_config.data_parallel_size != 1 + or parallel_config.distributed_executor_backend == "external_launcher" + ): + raise ValueError( + "Dynamo vLLM RL currently does not support data parallelism and external launcher" + ) + return parallel_config.world_size + + RequestT = TypeVar("RequestT") ResponseT = TypeVar("ResponseT") @@ -1266,7 +1274,16 @@ def __init__( self.shutdown_event = shutdown_event # Request-plane RL method map served by rl_dispatch on # dyn://..rl when --enable-rl / DYN_ENABLE_RL is set. - self.rl_route_registry = RLRouteRegistry(self.runtime, logger_=logger) + rl_world_size = None + if config.enable_rl: + rl_world_size = resolve_rl_weight_world_size( + engine.vllm_config.parallel_config + ) + self.rl_route_registry = RLRouteRegistry( + self.runtime, + logger_=logger, + world_size=rl_world_size, + ) # Load the custom encoder last. If a later init step raised, executor # GC would eventually reap the idle actor thread — but only once the @@ -1911,6 +1928,15 @@ async def pause_generation(self, body: dict) -> dict: "message": f"Invalid mode '{mode}'; expected keep|wait|abort", } async with self._pause_lock: + active_loras = sorted(self._lora_state.active_requests) + if mode == "keep" and active_loras: + return { + "status": "error", + "message": ( + "Cannot pause generation in keep mode with active LoRA requests: " + + ", ".join(active_loras) + ), + } try: try: await self.engine_client.pause_generation( @@ -2534,6 +2560,19 @@ async def _register_lora_discovery(self, lora_name: str, lora_id: int) -> None: runtime_config.tool_call_parser = self.config.dyn_tool_call_parser runtime_config.reasoning_parser = self.config.dyn_reasoning_parser + if lora_worker_type in (WorkerType.Aggregated, WorkerType.Decode): + lora_config = self.engine_client.vllm_config.lora_config + publish_engine_generate_capability( + runtime_config, + ModelInput.Tokens, + lora_model_type, + lora_worker_type, + bool( + lora_config + and getattr(lora_config, "enable_tower_connector_lora", False) + ), + ) + lora_needs: list[list[WorkerType]] = [lora_needs_set] if lora_needs_set else [] await register_model( @@ -2574,20 +2613,24 @@ async def _generate_with_lora_admission_lock( lora_request: LoRARequest | None, create_generator: Callable[[LoRARequest | None], AsyncIterator[Any]], ) -> AsyncIterator[Any]: - """Yield results after atomically admitting a lazy LoRA request. + """Yield results after atomically admitting a LoRA request. vLLM admits an ``AsyncLLM.generate`` request on its first iteration. Holding the adapter lifecycle lock through that iteration prevents an - unload from deleting bookkeeping before lazy activation completes. + unload from removing lazy or preloaded adapter state before admission. """ - if lora_request is None or self._preload_lora_into_engine(): - self._track_lora_request_activation(lora_request) + if lora_request is None: async for result in create_generator(lora_request): yield result return lock = self._get_lora_lock(lora_request.lora_name) - async with lock: + async with lock, self._pause_lock: + if self._paused: + raise RuntimeError( + f"Cannot admit LoRA request '{lora_request.lora_name}' while " + "generation is paused" + ) # The adapter may have been unloaded or reloaded at a different path # while this request waited. Look it up again while holding the lock. admitted_lora_request = self._resolve_lora_request(lora_request.lora_name) @@ -2600,16 +2643,20 @@ async def _generate_with_lora_admission_lock( raise ValueError( f"unknown model or LoRA adapter: '{lora_request.lora_name}'" ) - generator = create_generator(admitted_lora_request) self._track_lora_request_activation(admitted_lora_request) - try: - first_result = await anext(generator) - except StopAsyncIteration: - return + self._lora_state.begin_request(admitted_lora_request.lora_name) - yield first_result - async for result in generator: - yield result + engine_generator = create_generator(admitted_lora_request) + try: + async for result in engine_generator: + yield result + finally: + try: + close = getattr(engine_generator, "aclose", None) + if close is not None: + await close() + finally: + self._lora_state.end_request(admitted_lora_request.lora_name) def _preload_lora_into_engine(self) -> bool: """Whether lifecycle registration should eagerly activate the adapter. @@ -2718,7 +2765,20 @@ async def load_lora(self, request=None): if is_hot_swap and old_info is not None and old_engine_loaded: try: - await self.engine_client.remove_lora(old_info.id) + async with self._pause_lock: + if getattr( + self, "_paused", False + ) and self._lora_state.active_requests.get( + lora_name, 0 + ): + raise RuntimeError( + f"Cannot hot-swap LoRA '{lora_name}' while generation " + "is paused with active requests; resume generation or " + "abort the requests first" + ) + await self._lora_state.wait_until_idle(lora_name) + async with self._pause_lock: + await self.engine_client.remove_lora(old_info.id) self._engine_loaded_loras.discard(lora_name) except Exception as e: if capacity_reserved: @@ -2745,13 +2805,14 @@ async def load_lora(self, request=None): ) if preload_into_engine: try: - await self.engine_client.add_lora( - LoRARequest( - lora_name=lora_name, - lora_int_id=lora_id, - lora_path=lora_path, + async with self._pause_lock: + await self.engine_client.add_lora( + LoRARequest( + lora_name=lora_name, + lora_int_id=lora_id, + lora_path=lora_path, + ) ) - ) self._engine_loaded_loras.add(lora_name) except Exception as e: if ( @@ -2760,13 +2821,14 @@ async def load_lora(self, request=None): and old_engine_loaded ): try: - await self.engine_client.add_lora( - LoRARequest( - lora_name=lora_name, - lora_int_id=old_info.id, - lora_path=old_info.path, + async with self._pause_lock: + await self.engine_client.add_lora( + LoRARequest( + lora_name=lora_name, + lora_int_id=old_info.id, + lora_path=old_info.path, + ) ) - ) self._engine_loaded_loras.add(lora_name) except Exception as rollback_error: self._lora_state.loaded_loras.pop(lora_name, None) @@ -2797,7 +2859,8 @@ async def load_lora(self, request=None): if is_hot_swap: try: - await self.engine_client.reset_prefix_cache() + async with self._pause_lock: + await self.engine_client.reset_prefix_cache() except Exception as e: # The new adapter is already active in the engine, but # the prefix cache still holds entries computed under @@ -2810,16 +2873,20 @@ async def load_lora(self, request=None): if old_info is not None: try: if preload_into_engine: - await self.engine_client.remove_lora(lora_id) + async with self._pause_lock: + await self.engine_client.remove_lora( + lora_id + ) self._engine_loaded_loras.discard(lora_name) if old_engine_loaded: - await self.engine_client.add_lora( - LoRARequest( - lora_name=lora_name, - lora_int_id=old_info.id, - lora_path=old_info.path, + async with self._pause_lock: + await self.engine_client.add_lora( + LoRARequest( + lora_name=lora_name, + lora_int_id=old_info.id, + lora_path=old_info.path, + ) ) - ) self._engine_loaded_loras.add(lora_name) self._lora_state.loaded_loras[lora_name] = old_info rolled_back = ( @@ -2870,7 +2937,8 @@ async def load_lora(self, request=None): logger.debug( f"Rolling back: removing LoRA '{lora_name}' from engine" ) - await self.engine_client.remove_lora(lora_id) + async with self._pause_lock: + await self.engine_client.remove_lora(lora_id) self._engine_loaded_loras.discard(lora_name) self._lora_state.loaded_loras.pop(lora_name, None) logger.debug( @@ -2948,6 +3016,23 @@ async def unload_lora(self, request=None): logger.debug(f"Unloading LoRA adapter: {lora_name}") lora_id = lora.id + if lora_name in self._engine_loaded_loras: + async with self._pause_lock: + if getattr( + self, "_paused", False + ) and self._lora_state.active_requests.get(lora_name, 0): + yield { + "status": "error", + "message": ( + f"Cannot unload LoRA '{lora_name}' while generation " + "is paused with active requests; resume generation or " + "abort the requests first" + ), + "lora_name": lora_name, + } + return + await self._lora_state.wait_until_idle(lora_name) + # Stop advertising the adapter before mutating engine or # tracking state. Otherwise requests can still route here # after _resolve_lora_request has forgotten the adapter and @@ -2982,7 +3067,8 @@ async def unload_lora(self, request=None): # reached vLLM. if lora_name in self._engine_loaded_loras: try: - await self.engine_client.remove_lora(lora_id) + async with self._pause_lock: + await self.engine_client.remove_lora(lora_id) except Exception as e: if not self._is_lora_not_loaded_error(e): raise @@ -3310,6 +3396,7 @@ async def generate_tokens( total_output_tokens_by_index: dict[int, int] = {} raw_routed_experts_by_output: dict[int, Any] = {} + raw_sampling_mask_by_output: dict[int, Any] = {} # vLLM surfaces prompt_logprobs once (at end-of-prefill) and clears # them on subsequent chunks, so the generation-finish chunk often # carries None. Capture the first non-None payload and attach it to @@ -3373,6 +3460,9 @@ async def generate_tokens( raw_routed_experts = getattr(output, "routed_experts", None) if raw_routed_experts is not None: raw_routed_experts_by_output[output_idx] = raw_routed_experts + raw_sampling_mask = getattr(output, "sampling_mask", None) + if raw_sampling_mask is not None: + raw_sampling_mask_by_output[output_idx] = raw_sampling_mask # vLLM DELTA outputs already align token_ids/logprobs to this chunk. tokenizer = getattr(self.engine_client, "tokenizer", None) @@ -3396,6 +3486,11 @@ async def generate_tokens( _attach_prompt_logprobs_engine_data( out, prompt_logprobs_payload ) + kv_transfer_params = getattr(res, "kv_transfer_params", None) + if kv_transfer_params is not None: + engine_data = out.setdefault("engine_data", {}) + if isinstance(engine_data, dict): + engine_data["kv_transfer_params"] = kv_transfer_params # Emit the EFFECTIVE trim offset: clamp the requested # routed_experts_prompt_start to the prompt length. vLLM # clamps the returned routing rows the same way, so an @@ -3414,6 +3509,25 @@ async def generate_tokens( ) if routed_experts is not None: _attach_routed_experts_engine_data(out, routed_experts) + sampling_mask = raw_sampling_mask_by_output.get(output_idx) + if sampling_mask is not None: + rows = getattr(sampling_mask, "token_ids", None) + if rows is None: + raise TypeError( + "vLLM sampling mask is missing token_ids" + ) + normalized_rows = [list(row) for row in rows] + output_token_count = total_output_tokens_by_index.get( + output_idx, 0 + ) + if len(normalized_rows) != output_token_count: + raise ValueError( + "vLLM sampling mask row count " + f"{len(normalized_rows)} does not match " + f"completion token count {output_token_count}" + ) + _attach_sampling_mask_engine_data(out, normalized_rows) + # Log completion with LoRA info (debug level to avoid log spam) self._log_with_lora_context( "Completed token generation for request {request_id}{lora_info}: " @@ -3674,7 +3788,6 @@ async def _generate_token_mode(self, request, context, request_id): ) is_decode_only = False mode = DisaggregationMode.AGGREGATED - has_external_encoder_result = request.get("encoder_result") is not None if has_external_encoder_result and mode != DisaggregationMode.AGGREGATED: yield { @@ -3688,6 +3801,17 @@ async def _generate_token_mode(self, request, context, request_id): return has_mm_data = request.get("multi_modal_data") is not None assembled_prompt: EmbedsPrompt | TokensPrompt | None = None + engine_generate_input = None + extra_args = request.get("extra_args") + if isinstance(extra_args, dict) and "vllm_tito" in extra_args: + engine_generate_input = adapt_engine_generate_request( + request, + enable_multimodal=self._multimodal_request_processor.enable_multimodal, + decode_capable=mode != DisaggregationMode.PREFILL, + allow_multimodal_features=mode == DisaggregationMode.AGGREGATED, + vllm_config=self.engine_client.vllm_config, + default_sampling_params=self.default_sampling_params, + ) if has_external_encoder_result: assembled_prompt = await self._assemble_external_encoder_prompt( @@ -3696,6 +3820,10 @@ async def _generate_token_mode(self, request, context, request_id): multi_modal_data = None mm_processor_kwargs = None pre_rendered = None + elif engine_generate_input is not None: + multi_modal_data = None + mm_processor_kwargs = None + pre_rendered = engine_generate_input.prompt elif ( mode == DisaggregationMode.AGGREGATED and self._custom_encoder is not None @@ -3764,11 +3892,15 @@ async def _generate_token_mode(self, request, context, request_id): _apply_nvext_cache_salt(request, prompt) # Build sampling params from request - sampling_params = build_sampling_params( - request, - self.default_sampling_params, - self.model_max_len, - enable_rl=self.config.enable_rl, + sampling_params = ( + engine_generate_input.sampling_params + if engine_generate_input is not None + else build_sampling_params( + request, + self.default_sampling_params, + self.model_max_len, + enable_rl=self.config.enable_rl, + ) ) if kv_params is not None: @@ -3793,7 +3925,11 @@ async def _generate_token_mode(self, request, context, request_id): ) routing = request.get("routing") or {} dp_rank = self._to_local_dp_rank(routing.get("dp_rank")) - priority = -int(routing.get("priority", 0)) + priority = ( + engine_generate_input.priority + if engine_generate_input is not None + else -int(routing.get("priority", 0)) + ) trace_headers = context.trace_headers() reasoning_ended, reasoning_parser_kwargs = _request_reasoning_metadata(request) diff --git a/components/src/dynamo/vllm/kv_hints.py b/components/src/dynamo/vllm/kv_hints.py index 48e901a13f77..2f4fab66b9f8 100644 --- a/components/src/dynamo/vllm/kv_hints.py +++ b/components/src/dynamo/vllm/kv_hints.py @@ -21,6 +21,9 @@ ) from dynamo.llm import ModelRuntimeConfig, WorkerType +_KV_TRANSFER_PARAMS_EXTRA_ARGS_KEY = "kv_transfer_params" +_KV_HINT_EXTRA_ARGS_KEY = "kv_hint" + @dataclass(frozen=True) class KvTransferHintSource: @@ -28,6 +31,27 @@ class KvTransferHintSource: worker_type: str +def _apply_kv_hint(sampling_params: Any, kv_hint: Any) -> None: + """Attach the complete Dynamo KV hint message to vLLM's private input.""" + if not isinstance(kv_hint, Mapping): + return + + extra_args = ( + dict(sampling_params.extra_args) + if isinstance(sampling_params.extra_args, dict) + else {} + ) + existing_kv_transfer_params = extra_args.get(_KV_TRANSFER_PARAMS_EXTRA_ARGS_KEY) + kv_transfer_params = ( + dict(existing_kv_transfer_params) + if isinstance(existing_kv_transfer_params, dict) + else {} + ) + kv_transfer_params[_KV_HINT_EXTRA_ARGS_KEY] = dict(kv_hint) + extra_args[_KV_TRANSFER_PARAMS_EXTRA_ARGS_KEY] = kv_transfer_params + sampling_params.extra_args = extra_args + + def _secondary_tiers(engine_args: AsyncEngineArgs) -> list[Mapping[str, Any]]: """Return mapping-shaped secondary-tier configs from KVTransferConfig.""" kv_config = getattr(engine_args, "kv_transfer_config", None) diff --git a/components/src/dynamo/vllm/lora_state.py b/components/src/dynamo/vllm/lora_state.py index dbe3f2f36ebf..cac80dfe4f2d 100644 --- a/components/src/dynamo/vllm/lora_state.py +++ b/components/src/dynamo/vllm/lora_state.py @@ -23,6 +23,8 @@ def __init__(self): str, asyncio.Lock ] = weakref.WeakValueDictionary() self.lora_load_locks_guard = threading.Lock() + self.active_requests: dict[str, int] = {} + self.request_drained: dict[str, asyncio.Event] = {} def resolve_request( self, @@ -76,6 +78,28 @@ def get_lock(self, lora_name: str) -> asyncio.Lock: self.lora_load_locks[lora_name] = lock return lock + def begin_request(self, lora_name: str) -> None: + """Track a request; every call must be paired with ``end_request``.""" + count = self.active_requests.get(lora_name, 0) + if count == 0: + self.request_drained[lora_name] = asyncio.Event() + self.active_requests[lora_name] = count + 1 + + def end_request(self, lora_name: str) -> None: + """Release one request previously tracked by ``begin_request``.""" + count = self.active_requests[lora_name] + if count > 1: + self.active_requests[lora_name] = count - 1 + return + del self.active_requests[lora_name] + self.request_drained.pop(lora_name).set() + + async def wait_until_idle(self, lora_name: str) -> None: + """Wait until all tracked requests for an adapter have ended.""" + drained = self.request_drained.get(lora_name) + if drained is not None: + await drained.wait() + def list_lora_ids(self) -> dict[str, int]: """Return map of loaded LoRA names to integer IDs. diff --git a/components/src/dynamo/vllm/omni/base_handler.py b/components/src/dynamo/vllm/omni/base_handler.py index 93e79c1e2bf8..d0cf517e617f 100644 --- a/components/src/dynamo/vllm/omni/base_handler.py +++ b/components/src/dynamo/vllm/omni/base_handler.py @@ -76,6 +76,8 @@ def __init__( self.shutdown_event = shutdown_event self._lora_state = LoRAState() + self._paused = False + self._pause_lock = asyncio.Lock() # Properties loaded_loras, _lora_load_locks, _lora_load_locks_guard are now # available through LoRAState. No direct assignment needed since properties # are backed by _lora_state after initialization. diff --git a/components/src/dynamo/vllm/tests/omni/test_omni_base_handler.py b/components/src/dynamo/vllm/tests/omni/test_omni_base_handler.py index 85a64496f1d9..4f477add0ec5 100644 --- a/components/src/dynamo/vllm/tests/omni/test_omni_base_handler.py +++ b/components/src/dynamo/vllm/tests/omni/test_omni_base_handler.py @@ -85,6 +85,18 @@ def _parallel_config_without(*excluded_fields): return dataclasses.make_dataclass("LegacyDiffusionParallelConfig", fields) +def test_init_initializes_pause_state(): + config = _make_config() + with ( + patch.object(BaseOmniHandler, "_build_omni_kwargs", return_value={}), + patch("dynamo.vllm.omni.base_handler.AsyncOmni", return_value=MagicMock()), + ): + handler = BaseOmniHandler(None, config, {}) + + assert handler._paused is False + assert not handler._pause_lock.locked() + + class TestDiffusionParallelConfigCoverage: def test_all_diffusion_parallel_config_fields_covered(self): """Every DiffusionParallelConfig field must be in OmniParallelKwargs, engine_args, or _SKIP_FIELDS. diff --git a/components/src/dynamo/vllm/tests/omni/test_omni_handler.py b/components/src/dynamo/vllm/tests/omni/test_omni_handler.py index a27ce447902d..6ba903f1108b 100644 --- a/components/src/dynamo/vllm/tests/omni/test_omni_handler.py +++ b/components/src/dynamo/vllm/tests/omni/test_omni_handler.py @@ -76,6 +76,8 @@ def _make_handler(stage_types=("diffusion",)): # BaseOmniHandler.__init__ is mocked out in tests; recreate LoRA state attrs # expected by BaseWorkerHandler helpers called by OmniHandler. handler._lora_state = LoRAState() + handler._paused = False + handler._pause_lock = asyncio.Lock() handler.loaded_loras = handler._lora_state.loaded_loras handler._lora_load_locks = handler._lora_state.lora_load_locks handler._lora_load_locks_guard = handler._lora_state.lora_load_locks_guard diff --git a/components/src/dynamo/vllm/tests/test_runtime_metadata.py b/components/src/dynamo/vllm/tests/test_runtime_metadata.py index e9cc0e72c197..dce48f68e99a 100644 --- a/components/src/dynamo/vllm/tests/test_runtime_metadata.py +++ b/components/src/dynamo/vllm/tests/test_runtime_metadata.py @@ -1,7 +1,10 @@ # SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +import builtins +import importlib import json +import sys from types import SimpleNamespace from unittest.mock import Mock, call @@ -24,6 +27,24 @@ ] +def test_engine_generate_metadata_imports_without_vllm(monkeypatch): + module_name = "dynamo.vllm.engine_generate" + loaded_module = sys.modules.pop(module_name) + original_import = builtins.__import__ + + def reject_vllm_import(name, *args, **kwargs): + if name == "vllm" or name.startswith("vllm."): + raise ModuleNotFoundError("vLLM is not installed", name=name) + return original_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", reject_vllm_import) + try: + metadata_module = importlib.import_module(module_name) + assert metadata_module.VLLM_GENERATE_CAPABILITY + finally: + sys.modules[module_name] = loaded_module + + def test_spec_decode_runtime_data_uses_vllm_speculative_config(): config = SimpleNamespace( engine_args=SimpleNamespace( @@ -78,7 +99,7 @@ def test_vllm_token_budget_matches_rejection_policy(): "expected", ), [ - (ModelInput.Tokens, ModelType.Prefill, WorkerType.Prefill, False, True), + (ModelInput.Tokens, ModelType.Prefill, WorkerType.Prefill, False, False), (ModelInput.Tokens, ModelType.Chat, WorkerType.Decode, True, True), ( ModelInput.Tokens, diff --git a/components/src/dynamo/vllm/tests/test_vllm_engine_generate.py b/components/src/dynamo/vllm/tests/test_vllm_engine_generate.py new file mode 100644 index 000000000000..0c383fd090a5 --- /dev/null +++ b/components/src/dynamo/vllm/tests/test_vllm_engine_generate.py @@ -0,0 +1,432 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import importlib.util +from types import SimpleNamespace + +import pytest + +pytestmark = [ + pytest.mark.unit, + pytest.mark.pre_merge, + pytest.mark.vllm, + pytest.mark.core, + pytest.mark.gpu_0, + pytest.mark.skipif( + importlib.util.find_spec("vllm") is None, + reason="vllm not installed in this container", + ), +] + + +_VALID_MM_KWARGS_BASE64 = ( + "gaxwaXhlbF92YWx1ZXOCpGRhdGGTpXVpbnQ4kQPHAwMBAgOlZmllbGSS" + "p2JhdGNoZWSBq2tlZXBfb25fY3B1wg==" +) + + +def _request(*, sampling_params=None, features=None, token_ids=None, **envelope): + payload = { + "request_id": "request-1", + "sampling_params": sampling_params or {}, + **envelope, + } + if features is not None: + payload["features"] = features + return { + "model": "test-model", + "token_ids": token_ids or [11, 22, 33], + "extra_args": {"vllm_tito": payload}, + } + + +def _vllm_config(*, max_num_seqs=8, max_model_len=128): + return SimpleNamespace( + scheduler_config=SimpleNamespace(max_num_seqs=max_num_seqs), + model_config=SimpleNamespace( + max_model_len=max_model_len, + generation_config="vllm", + override_generation_config={}, + ), + ) + + +def test_tito_adapter_uses_outer_tokens_and_rl_sampling_defaults(): + from vllm.sampling_params import RequestOutputKind + + from dynamo.vllm.engine_generate import adapt_engine_generate_request + + request = _request( + token_ids=[41, 42], + sampling_params={ + "temperature": 0.7, + "top_k": 8, + "max_tokens": 5, + "routed_experts_prompt_start": 1, + }, + ) + request["extra_args"]["vllm_tito"]["token_ids"] = [999] + + adapted = adapt_engine_generate_request( + request, + enable_multimodal=False, + decode_capable=True, + vllm_config=_vllm_config(), + default_sampling_params={}, + ) + + assert adapted is not None + assert adapted.prompt["prompt_token_ids"] == [41, 42] + assert adapted.sampling_params.temperature == pytest.approx(0.7) + assert adapted.sampling_params.top_k == 8 + assert adapted.sampling_params.max_tokens == 5 + assert adapted.sampling_params.routed_experts_prompt_start == 1 + assert adapted.sampling_params.detokenize is False + assert adapted.sampling_params.output_kind is RequestOutputKind.DELTA + + +def test_tito_adapter_preserves_kv_transfer_params_in_sampling_extra_args(): + from dynamo.vllm.engine_generate import adapt_engine_generate_request + + request = _request( + sampling_params={"max_tokens": 5, "extra_args": {"existing": "value"}}, + kv_transfer_params={"connector_data": {"block_ids": [1, 2]}}, + ) + request["kv_hint"] = {"source": "worker-a"} + adapted = adapt_engine_generate_request( + request, + enable_multimodal=False, + decode_capable=True, + vllm_config=_vllm_config(), + default_sampling_params={}, + ) + + assert adapted is not None + assert adapted.sampling_params.extra_args == { + "existing": "value", + "kv_transfer_params": { + "connector_data": {"block_ids": [1, 2]}, + "kv_hint": {"source": "worker-a"}, + }, + } + + +@pytest.mark.parametrize("prompt_start", [True, 1.0, -1]) +def test_tito_adapter_rejects_invalid_routed_experts_prompt_start(prompt_start): + from dynamo.vllm.engine_generate import adapt_engine_generate_request + + with pytest.raises(ValueError, match="routed_experts_prompt_start"): + adapt_engine_generate_request( + _request( + sampling_params={ + "max_tokens": 1, + "routed_experts_prompt_start": prompt_start, + } + ), + enable_multimodal=False, + decode_capable=True, + vllm_config=_vllm_config(), + default_sampling_params={}, + ) + + +@pytest.mark.parametrize("prompt_start", [3, 99]) +def test_tito_adapter_rejects_out_of_range_routed_experts_start(prompt_start): + from dynamo.vllm.engine_generate import adapt_engine_generate_request + + with pytest.raises(ValueError, match="smaller than the prompt token count"): + adapt_engine_generate_request( + _request( + sampling_params={ + "max_tokens": 1, + "routed_experts_prompt_start": prompt_start, + } + ), + enable_multimodal=False, + decode_capable=True, + vllm_config=_vllm_config(), + default_sampling_params={}, + ) + + +def test_tito_adapter_rejects_nonprogressing_guided_json_cycle(): + from dynamo.vllm.engine_generate import adapt_engine_generate_request + + schema = { + "$defs": {"A": {"allOf": [{"$ref": "#/$defs/A"}]}}, + "$ref": "#/$defs/A", + } + with pytest.raises(ValueError, match=r"non-progressing local \$ref cycle"): + adapt_engine_generate_request( + _request( + sampling_params={ + "max_tokens": 1, + "structured_outputs": {"json": schema}, + } + ), + enable_multimodal=False, + decode_capable=True, + vllm_config=_vllm_config(), + default_sampling_params={}, + ) + + +def test_tito_adapter_preserves_stop_strings_and_stop_token_ids(): + from dynamo.vllm.engine_generate import adapt_engine_generate_request + + adapted = adapt_engine_generate_request( + _request(sampling_params={"max_tokens": 1, "stop": ["END"]}), + enable_multimodal=False, + decode_capable=True, + vllm_config=_vllm_config(), + default_sampling_params={}, + ) + assert adapted is not None + assert adapted.sampling_params.stop == ["END"] + assert adapted.sampling_params.detokenize is True + + adapted = adapt_engine_generate_request( + _request(sampling_params={"max_tokens": 1, "stop_token_ids": [42]}), + enable_multimodal=False, + decode_capable=True, + vllm_config=_vllm_config(), + default_sampling_params={}, + ) + assert adapted is not None + assert adapted.sampling_params.stop_token_ids == [42] + assert adapted.sampling_params.detokenize is False + + +def test_tito_adapter_builds_preprocessed_image_input_without_reprocessing(): + from dynamo.vllm.engine_generate import adapt_engine_generate_request + + features = { + "mm_hashes": {"image": ["renderer-hash"]}, + "mm_placeholders": { + "image": [ + { + "offset": 1, + "length": 2, + "is_embed": [False, True], + } + ] + }, + "kwargs_data": {"image": [_VALID_MM_KWARGS_BASE64]}, + } + request = _request(features=features) + request["extra_args"]["dynamo_mm_routing_hashes"] = ["a" * 16 + "0" * 48] + + adapted = adapt_engine_generate_request( + request, + enable_multimodal=True, + decode_capable=True, + vllm_config=_vllm_config(), + default_sampling_params={}, + ) + + assert adapted is not None + assert adapted.prompt["type"] == "multimodal" + assert adapted.prompt["prompt_token_ids"] == [11, 22, 33] + assert adapted.prompt["mm_hashes"] == {"image": ["a" * 16 + "0" * 48]} + assert len(adapted.prompt["mm_kwargs"]["image"]) == 1 + assert adapted.prompt["mm_kwargs"]["image"][0] is not None + placeholder = adapted.prompt["mm_placeholders"]["image"][0] + assert (placeholder.offset, placeholder.length) == (1, 2) + assert placeholder.is_embed.tolist() == [False, True] + + +def test_tito_adapter_rejects_preprocessed_features_on_disaggregated_decode(): + from dynamo.vllm.engine_generate import adapt_engine_generate_request + + with pytest.raises(ValueError, match="aggregated vLLM worker"): + adapt_engine_generate_request( + _request( + features={ + "mm_hashes": {"image": ["renderer-hash"]}, + "mm_placeholders": {"image": [{"offset": 0, "length": 1}]}, + "kwargs_data": {"image": [_VALID_MM_KWARGS_BASE64]}, + } + ), + enable_multimodal=True, + decode_capable=True, + allow_multimodal_features=False, + vllm_config=_vllm_config(), + default_sampling_params={}, + ) + + +@pytest.mark.parametrize( + ("enable_multimodal", "decode_capable", "features", "match"), + [ + ( + False, + True, + { + "mm_hashes": {"image": ["x"]}, + "mm_placeholders": {"image": [{"offset": 0, "length": 1}]}, + "kwargs_data": {"image": [_VALID_MM_KWARGS_BASE64]}, + }, + "multimodal", + ), + (True, False, None, "aggregated or decode"), + ( + True, + True, + { + "mm_hashes": {"audio": ["x"]}, + "mm_placeholders": {"audio": [{"offset": 0, "length": 1}]}, + "kwargs_data": {"audio": [_VALID_MM_KWARGS_BASE64]}, + }, + "image", + ), + ], +) +def test_tito_adapter_rejects_unsupported_execution_paths( + enable_multimodal, decode_capable, features, match +): + from dynamo.vllm.engine_generate import adapt_engine_generate_request + + with pytest.raises(ValueError, match=match): + adapt_engine_generate_request( + _request(features=features), + enable_multimodal=enable_multimodal, + decode_capable=decode_capable, + vllm_config=_vllm_config(), + default_sampling_params={}, + ) + + +@pytest.mark.parametrize( + "features", + [ + { + "mm_hashes": {}, + "mm_placeholders": {"image": []}, + "kwargs_data": {"image": []}, + }, + { + "mm_hashes": {"image": []}, + "mm_placeholders": {}, + "kwargs_data": {"image": []}, + }, + ], +) +def test_tito_adapter_rejects_asymmetric_image_feature_objects(features): + from dynamo.vllm.engine_generate import adapt_engine_generate_request + + with pytest.raises(TypeError, match="hashes and placeholders must be lists"): + adapt_engine_generate_request( + _request(features=features, sampling_params={"max_tokens": 1}), + enable_multimodal=True, + decode_capable=True, + vllm_config=_vllm_config(), + default_sampling_params={}, + ) + + +def test_tito_adapter_rejects_routing_hash_count_mismatch(): + from dynamo.vllm.engine_generate import adapt_engine_generate_request + + request = _request( + features={ + "mm_hashes": {"image": ["one"]}, + "mm_placeholders": {"image": [{"offset": 0, "length": 1}]}, + "kwargs_data": {"image": [_VALID_MM_KWARGS_BASE64]}, + } + ) + request["extra_args"]["dynamo_mm_routing_hashes"] = ["one", "two"] + + with pytest.raises(ValueError, match="routing hash"): + adapt_engine_generate_request( + request, + enable_multimodal=True, + decode_capable=True, + vllm_config=_vllm_config(), + default_sampling_params={}, + ) + + +@pytest.mark.parametrize("kwargs_data", [None, []]) +def test_tito_adapter_rejects_non_object_kwargs_data(kwargs_data): + from dynamo.vllm.engine_generate import adapt_engine_generate_request + + request = _request( + features={ + "mm_hashes": {"image": ["one"]}, + "mm_placeholders": {"image": [{"offset": 0, "length": 1}]}, + "kwargs_data": kwargs_data, + } + ) + + with pytest.raises(TypeError, match="kwargs_data must be an object"): + adapt_engine_generate_request( + request, + enable_multimodal=True, + decode_capable=True, + vllm_config=_vllm_config(), + default_sampling_params={}, + ) + + +@pytest.mark.parametrize( + ("is_embed", "error", "match"), + [ + ({"unexpected": True}, TypeError, "sequence"), + ([False, True], ValueError, "placeholder length"), + ([1], ValueError, "booleans"), + ], +) +def test_tito_adapter_rejects_invalid_placeholder_mask_before_tensor_conversion( + is_embed, error, match +): + from dynamo.vllm.engine_generate import adapt_engine_generate_request + + request = _request( + token_ids=[11], + features={ + "mm_hashes": {"image": ["one"]}, + "mm_placeholders": { + "image": [{"offset": 0, "length": 1, "is_embed": is_embed}] + }, + "kwargs_data": {"image": [_VALID_MM_KWARGS_BASE64]}, + }, + ) + + with pytest.raises(error, match=match): + adapt_engine_generate_request( + request, + enable_multimodal=True, + decode_capable=True, + vllm_config=_vllm_config(), + default_sampling_params={}, + ) + + +def test_tito_adapter_rejects_sampling_choices_above_scheduler_capacity(): + from dynamo.vllm.engine_generate import adapt_engine_generate_request + + with pytest.raises(ValueError, match="max_num_seqs"): + adapt_engine_generate_request( + _request(sampling_params={"n": 5, "max_tokens": 4}), + enable_multimodal=False, + decode_capable=True, + vllm_config=_vllm_config(max_num_seqs=4), + default_sampling_params={}, + ) + + +def test_tito_adapter_resolves_omitted_max_tokens_from_server_limits(): + from dynamo.vllm.engine_generate import adapt_engine_generate_request + + adapted = adapt_engine_generate_request( + _request(token_ids=[1, 2, 3], sampling_params={}), + enable_multimodal=False, + decode_capable=True, + vllm_config=_vllm_config(max_model_len=20), + default_sampling_params={}, + ) + + assert adapted is not None + assert adapted.sampling_params.max_tokens == 17 diff --git a/components/src/dynamo/vllm/tests/test_vllm_legacy_lora.py b/components/src/dynamo/vllm/tests/test_vllm_legacy_lora.py index 8b99606e4576..4c5dd83d5f8a 100644 --- a/components/src/dynamo/vllm/tests/test_vllm_legacy_lora.py +++ b/components/src/dynamo/vllm/tests/test_vllm_legacy_lora.py @@ -54,6 +54,7 @@ def _make_prefill_handler(): vllm_config=SimpleNamespace( additional_config={DYNAMO_KV_EVENT_BLOCK_SIZE_KEY: 1056}, cache_config=SimpleNamespace(block_size=16), + lora_config=None, ), ) handler.generate_endpoint = object() @@ -67,6 +68,8 @@ def _make_prefill_handler(): handler._served_model_aliases = ("llama2-7b-alias",) handler._lora_state = LoRAState() handler._engine_loaded_loras = set() + handler._pause_lock = asyncio.Lock() + handler._paused = False return handler @@ -191,6 +194,58 @@ async def test_decode_load_still_eagerly_adds_to_engine(monkeypatch): handler.engine_client.add_lora.assert_awaited_once() +@pytest.mark.asyncio +async def test_hot_swap_rejects_paused_adapter_with_active_request(monkeypatch): + handler = _make_prefill_handler() + handler.config.disaggregation_mode = DisaggregationMode.AGGREGATED + handler._paused = True + handler._lora_state.loaded_loras = {"adapterA": LoRAInfo(id=123, path="/cache/old")} + handler._engine_loaded_loras = {"adapterA"} + handler._lora_state.begin_request("adapterA") + manager = SimpleNamespace( + download_lora=AsyncMock( + return_value={"status": "success", "local_path": "/cache/new"} + ) + ) + monkeypatch.setenv("DYN_LORA_HOTSWAP_ENABLED", "true") + monkeypatch.setattr(handlers_mod, "get_lora_manager", lambda: manager) + + results = [ + result + async for result in handler.load_lora( + {"lora_name": "adapterA", "source": {"uri": "file:///adapter"}} + ) + ] + + assert results[-1]["status"] == "error" + assert "paused" in results[-1]["message"] + handler.engine_client.remove_lora.assert_not_awaited() + handler._lora_state.end_request("adapterA") + + +@pytest.mark.asyncio +async def test_unload_rejects_paused_adapter_with_active_request(monkeypatch): + handler = _make_prefill_handler() + handler._paused = True + handler._lora_state.loaded_loras = { + "adapterA": LoRAInfo(id=123, path="/cache/adapter") + } + handler._engine_loaded_loras = {"adapterA"} + handler._lora_state.begin_request("adapterA") + unregister = AsyncMock() + monkeypatch.setattr(handlers_mod, "unregister_model", unregister) + + results = [ + result async for result in handler.unload_lora({"lora_name": "adapterA"}) + ] + + assert results[-1]["status"] == "error" + assert "paused" in results[-1]["message"] + unregister.assert_not_awaited() + handler.engine_client.remove_lora.assert_not_awaited() + handler._lora_state.end_request("adapterA") + + @pytest.mark.asyncio async def test_prefill_publish_failure_rolls_back_metadata_only(monkeypatch): handler = _make_prefill_handler() @@ -290,7 +345,10 @@ async def test_legacy_prefill_unload_removes_request_activated_adapter(monkeypat @pytest.mark.asyncio -async def test_legacy_prefill_request_admission_serializes_with_unload(monkeypatch): +@pytest.mark.timeout(5) +async def test_legacy_lora_request_admission_serializes_with_unload( + monkeypatch, +): handler = _make_prefill_handler() handler._lora_state.loaded_loras = { "adapterA": LoRAInfo(id=123, path="/cache/adapter") @@ -331,13 +389,248 @@ async def _unload(): allow_admission.set() await admission_task - results = await unload_task + await asyncio.sleep(0) + assert not unload_task.done() await admission.aclose() + results = await unload_task assert results[-1]["status"] == "success" handler.engine_client.remove_lora.assert_awaited_once_with(123) +@pytest.mark.asyncio +@pytest.mark.timeout(5) +async def test_legacy_lora_request_closes_engine_generator_before_drain(): + handler = _make_prefill_handler() + handler._lora_state.loaded_loras = { + "adapterA": LoRAInfo(id=123, path="/cache/adapter") + } + cleanup_started = asyncio.Event() + allow_cleanup = asyncio.Event() + engine_generator_closed = asyncio.Event() + + async def _generate(_lora_request): + try: + yield SimpleNamespace() + finally: + cleanup_started.set() + await allow_cleanup.wait() + engine_generator_closed.set() + + admission = handler._generate_with_lora_admission_lock( + handler._resolve_lora_request("adapterA"), + _generate, + ) + + await anext(admission) + assert handler._lora_state.active_requests == {"adapterA": 1} + + close_task = asyncio.create_task(admission.aclose()) + await cleanup_started.wait() + + assert handler._lora_state.active_requests == {"adapterA": 1} + assert not close_task.done() + + allow_cleanup.set() + await close_task + assert engine_generator_closed.is_set() + assert handler._lora_state.active_requests == {} + + +@pytest.mark.asyncio +async def test_legacy_lora_request_drain_preserves_concurrent_generation(): + handler = _make_prefill_handler() + handler.config.disaggregation_mode = DisaggregationMode.AGGREGATED + handler._lora_state.loaded_loras = { + "adapterA": LoRAInfo(id=123, path="/cache/adapter") + } + started = 0 + both_started = asyncio.Event() + release = asyncio.Event() + + async def _blocked_generate(_lora_request): + nonlocal started + started += 1 + if started == 2: + both_started.set() + await release.wait() + yield SimpleNamespace() + + admissions = [ + handler._generate_with_lora_admission_lock( + handler._resolve_lora_request("adapterA"), + _blocked_generate, + ) + for _ in range(2) + ] + tasks = [asyncio.create_task(anext(admission)) for admission in admissions] + + await asyncio.wait_for(both_started.wait(), timeout=1) + assert handler._lora_state.active_requests == {"adapterA": 2} + + release.set() + await asyncio.gather(*tasks) + await asyncio.gather(*(admission.aclose() for admission in admissions)) + + assert handler._lora_state.active_requests == {} + assert handler._lora_state.request_drained == {} + + +@pytest.mark.asyncio +async def test_lora_unload_drain_allows_other_adapter_admission(monkeypatch): + handler = _make_prefill_handler() + handler._lora_state.loaded_loras = { + "adapterA": LoRAInfo(id=123, path="/cache/adapter-a"), + "adapterB": LoRAInfo(id=456, path="/cache/adapter-b"), + } + handler._engine_loaded_loras = {"adapterA"} + handler._lora_state.begin_request("adapterA") + monkeypatch.setattr(handlers_mod, "unregister_model", AsyncMock()) + + async def _generate(_lora_request): + yield SimpleNamespace() + + async def _run_unload(): + return [ + result async for result in handler.unload_lora({"lora_name": "adapterA"}) + ] + + unload_task = asyncio.create_task(_run_unload()) + await asyncio.sleep(0) + + admission = handler._generate_with_lora_admission_lock( + handler._resolve_lora_request("adapterB"), + _generate, + ) + await asyncio.wait_for(anext(admission), timeout=1) + await admission.aclose() + assert handler._lora_state.active_requests == {"adapterA": 1} + + handler._lora_state.end_request("adapterA") + results = await unload_task + + assert results[-1]["status"] == "success" + + +@pytest.mark.asyncio +async def test_lora_download_allows_other_adapter_admission(monkeypatch): + handler = _make_prefill_handler() + handler._lora_state.loaded_loras = { + "adapterB": LoRAInfo(id=456, path="/cache/adapter-b") + } + download_started = asyncio.Event() + allow_download = asyncio.Event() + + async def _blocked_download(_uri): + download_started.set() + await allow_download.wait() + return {"status": "success", "local_path": "/cache/adapter-a"} + + async def _generate(_lora_request): + yield SimpleNamespace() + + manager = SimpleNamespace(download_lora=_blocked_download) + monkeypatch.setattr(handlers_mod, "get_lora_manager", lambda: manager) + monkeypatch.setattr(handlers_mod, "register_model", AsyncMock()) + + async def _run_load(): + return [ + result + async for result in handler.load_lora( + {"lora_name": "adapterA", "source": {"uri": "file:///adapter-a"}} + ) + ] + + load_task = asyncio.create_task(_run_load()) + await download_started.wait() + + admission = handler._generate_with_lora_admission_lock( + handler._resolve_lora_request("adapterB"), + _generate, + ) + await asyncio.wait_for(anext(admission), timeout=1) + await admission.aclose() + + allow_download.set() + results = await load_task + assert results[-1]["status"] == "success" + + +@pytest.mark.asyncio +@pytest.mark.timeout(5) +async def test_lora_admission_waits_for_pause_and_rejects_when_paused(): + handler = _make_prefill_handler() + handler._lora_state.loaded_loras = { + "adapterA": LoRAInfo(id=123, path="/cache/adapter") + } + pause_started = asyncio.Event() + allow_pause = asyncio.Event() + generation_started = asyncio.Event() + + async def _pause_generation(**_kwargs): + pause_started.set() + await allow_pause.wait() + + async def _generate(_lora_request): + generation_started.set() + yield SimpleNamespace() + + handler.engine_client.pause_generation = AsyncMock(side_effect=_pause_generation) + pause_task = asyncio.create_task(handler.pause_generation({"mode": "keep"})) + await pause_started.wait() + + admission = handler._generate_with_lora_admission_lock( + handler._resolve_lora_request("adapterA"), + _generate, + ) + admission_task = asyncio.create_task(anext(admission)) + await asyncio.sleep(0) + + assert not generation_started.is_set() + assert not admission_task.done() + + allow_pause.set() + pause_result = await pause_task + assert pause_result["status"] == "ok" + with pytest.raises(RuntimeError, match="generation is paused"): + await admission_task + + assert handler._lora_state.active_requests == {} + assert not generation_started.is_set() + + +@pytest.mark.asyncio +@pytest.mark.timeout(5) +async def test_legacy_unload_cancellation_does_not_unregister_active_lora( + monkeypatch, +): + handler = _make_prefill_handler() + handler._lora_state.loaded_loras = { + "adapterA": LoRAInfo(id=123, path="/cache/adapter") + } + handler._engine_loaded_loras = {"adapterA"} + handler._lora_state.begin_request("adapterA") + unregister = AsyncMock() + monkeypatch.setattr(handlers_mod, "unregister_model", unregister) + + async def _run_unload(): + return [ + result async for result in handler.unload_lora({"lora_name": "adapterA"}) + ] + + unload_task = asyncio.create_task(_run_unload()) + await asyncio.sleep(0) + + unregister.assert_not_awaited() + unload_task.cancel() + with pytest.raises(asyncio.CancelledError): + await unload_task + assert handler._lora_state.loaded_loras["adapterA"].id == 123 + assert "adapterA" in handler._engine_loaded_loras + + handler._lora_state.end_request("adapterA") + + @pytest.mark.asyncio async def test_legacy_prefill_request_rejects_adapter_unloaded_before_admission( monkeypatch, caplog diff --git a/components/src/dynamo/vllm/tests/test_vllm_worker_handler.py b/components/src/dynamo/vllm/tests/test_vllm_worker_handler.py index fad0a4e58683..eae8f28bf3cc 100644 --- a/components/src/dynamo/vllm/tests/test_vllm_worker_handler.py +++ b/components/src/dynamo/vllm/tests/test_vllm_worker_handler.py @@ -40,6 +40,47 @@ ] +def test_rl_weight_world_size_accepts_tensor_parallel_topology(): + parallel_config = SimpleNamespace( + data_parallel_size=1, + distributed_executor_backend="mp", + world_size=4, + ) + + assert mod.resolve_rl_weight_world_size(parallel_config) == 4 + + +@pytest.mark.parametrize( + "parallel_config", + [ + SimpleNamespace( + data_parallel_size=2, + distributed_executor_backend="mp", + world_size=4, + ), + SimpleNamespace( + data_parallel_size=1, + distributed_executor_backend="external_launcher", + world_size=4, + ), + ], +) +def test_rl_weight_world_size_rejects_unsupported_topologies(parallel_config): + with pytest.raises(ValueError, match="data parallelism and external launcher"): + mod.resolve_rl_weight_world_size(parallel_config) + + +def test_native_generate_cache_salt_applies_to_prefill_prompt(): + request = { + "extra_args": {"vllm_tito": {"cache_salt": "policy-7"}}, + } + prompt = {"prompt_token_ids": [1, 2, 3]} + + mod._apply_nvext_cache_salt(request, prompt) + + assert prompt["cache_salt"] == "dynamo-cache-salt:policy-7" + + # ── Helpers ────────────────────────────────────────────────────────── @@ -62,6 +103,7 @@ def _make_config( # so set to LOCAL mode. config.embedding_transfer_mode = EmbeddingTransferMode.LOCAL config.enable_multimodal = enable_multimodal + config.enable_rl = False config.multimodal_embedding_cache_capacity_gb = ( multimodal_embedding_cache_capacity_gb ) @@ -144,6 +186,68 @@ def _make_engine_response(request_id: str = "req-1", finished: bool = True): return resp +@pytest.mark.parametrize( + ("disaggregation_mode", "expected_worker_type"), + [ + (None, mod.WorkerType.Aggregated), + ("PREFILL", mod.WorkerType.Prefill), + ("DECODE", mod.WorkerType.Decode), + ], +) +def test_lora_discovery_publishes_engine_generate_capability( + disaggregation_mode, expected_worker_type +): + config = _make_config(disaggregation_mode=disaggregation_mode) + handler = _make_handler(config) + handler.config = config + handler.generate_endpoint = MagicMock() + handler.dp_range = (0, 1) + handler.model_max_len = 4096 + handler.config.route_to_encoder = False + handler.config.engine_args.max_loras = 2 + handler.engine_client = SimpleNamespace( + vllm_config=SimpleNamespace( + lora_config=SimpleNamespace(enable_tower_connector_lora=True) + ) + ) + runtime_config = MagicMock() + + with ( + patch.object(mod, "ModelRuntimeConfig", return_value=runtime_config), + patch.object(mod, "apply_data_parallel_runtime_config"), + patch.object(mod, "publish_kv_hint_capabilities"), + patch.object(mod, "publish_vllm_token_budget"), + patch.object(mod, "state_agent_settings", return_value=None), + patch.object(mod, "get_configured_kv_event_block_size", return_value=16), + patch.object( + mod, + "publish_engine_generate_capability", + create=True, + return_value=True, + ) as publish_generate, + patch.object(mod, "register_model", new=AsyncMock()) as register_model, + ): + asyncio.run(handler._register_lora_discovery("adapter-v1", 42)) + + if expected_worker_type == mod.WorkerType.Prefill: + publish_generate.assert_not_called() + else: + publish_generate.assert_called_once() + ( + runtime_arg, + input_arg, + model_type_arg, + worker_arg, + tower_lora_arg, + ) = publish_generate.call_args.args + assert runtime_arg is runtime_config + assert input_arg == mod.ModelInput.Tokens + assert model_type_arg.supports_chat() + assert worker_arg == expected_worker_type + assert tower_lora_arg is True + assert register_model.await_args.kwargs["runtime_config"] is runtime_config + + @pytest.mark.asyncio async def test_clear_kv_blocks_resets_vllm_external_cache(): handler = _make_handler() @@ -392,6 +496,126 @@ async def fake_generate(*args, **kwargs): np.testing.assert_array_equal(decoded, routed_experts.reshape(-1)) @pytest.mark.asyncio + async def test_generate_tokens_emits_sampling_mask_only_on_final_chunk(self): + from vllm.sampling_params import SamplingParams + + handler = _make_handler() + handler._extract_logprobs = MagicMock(return_value=(None, None)) + + async def fake_generate(*args, **kwargs): + yield SimpleNamespace( + outputs=[ + SimpleNamespace( + index=0, + token_ids=[11], + routed_experts=None, + sampling_mask=None, + finish_reason=None, + stop_reason=None, + ) + ], + prompt_token_ids=[1, 2], + prompt_logprobs=None, + ) + yield SimpleNamespace( + outputs=[ + SimpleNamespace( + index=0, + token_ids=[12], + routed_experts=None, + sampling_mask=SimpleNamespace(token_ids=[[11, 21], [12, 22]]), + finish_reason="stop", + stop_reason=None, + ) + ], + prompt_token_ids=[1, 2], + prompt_logprobs=None, + ) + + handler.engine_client = MagicMock() + handler.engine_client.generate = fake_generate + + chunks = [] + async for chunk in handler.generate_tokens( + PatchedTokensPrompt(prompt_token_ids=[1]), + SamplingParams(max_tokens=2), + "req-mask", + ): + chunks.append(chunk) + + assert "engine_data" not in chunks[0] + assert chunks[1]["engine_data"]["sampling_mask"] == [[11, 21], [12, 22]] + + @pytest.mark.asyncio + async def test_generate_tokens_emits_final_kv_transfer_params(self): + from vllm.sampling_params import SamplingParams + + handler = _make_handler() + handler._extract_logprobs = MagicMock(return_value=(None, None)) + + async def fake_generate(*args, **kwargs): + yield SimpleNamespace( + outputs=[ + SimpleNamespace( + index=0, + token_ids=[11], + finish_reason="stop", + stop_reason=None, + ) + ], + prompt_token_ids=[1, 2], + prompt_logprobs=None, + kv_transfer_params={"connector": "nixl"}, + ) + + handler.engine_client = MagicMock() + handler.engine_client.generate = fake_generate + + chunks = [ + chunk + async for chunk in handler.generate_tokens( + PatchedTokensPrompt(prompt_token_ids=[1]), + SamplingParams(max_tokens=1), + "req-kv", + ) + ] + + assert chunks[-1]["engine_data"]["kv_transfer_params"] == {"connector": "nixl"} + + @pytest.mark.asyncio + async def test_generate_tokens_rejects_sampling_mask_length_mismatch(self): + from vllm.sampling_params import SamplingParams + + handler = _make_handler() + handler._extract_logprobs = MagicMock(return_value=(None, None)) + + async def fake_generate(*args, **kwargs): + yield SimpleNamespace( + outputs=[ + SimpleNamespace( + index=0, + token_ids=[11, 12], + routed_experts=None, + sampling_mask=SimpleNamespace(token_ids=[[11]]), + finish_reason="stop", + stop_reason=None, + ) + ], + prompt_token_ids=[1, 2], + prompt_logprobs=None, + ) + + handler.engine_client = MagicMock() + handler.engine_client.generate = fake_generate + + with pytest.raises(ValueError, match="sampling mask"): + async for _ in handler.generate_tokens( + PatchedTokensPrompt(prompt_token_ids=[1]), + SamplingParams(max_tokens=2), + "req-mask", + ): + pass + async def test_generate_tokens_routed_experts_start_echoes_prompt_start(self): """routed_experts.start echoes SamplingParams.routed_experts_prompt_start (the offset vLLM trimmed) so the RL consumer can align the completion.""" @@ -1934,6 +2158,23 @@ async def test_admin_rejects_non_dict_body(self): assert resp["status"] == "error", (fn.__name__, body, resp) assert "JSON object" in resp["message"] + @pytest.mark.asyncio + async def test_keep_pause_rejects_active_lora_requests(self): + handler = _make_handler() + handler._pause_lock = asyncio.Lock() + handler._paused = False + handler._lora_state = mod.LoRAState() + handler.engine_client = MagicMock() + handler.engine_client.pause_generation = AsyncMock() + handler._lora_state.begin_request("adapterA") + + resp = await handler.pause_generation({"mode": "keep"}) + + assert resp["status"] == "error" + assert "active LoRA requests" in resp["message"] + handler.engine_client.pause_generation.assert_not_awaited() + handler._lora_state.end_request("adapterA") + @pytest.mark.asyncio async def test_distributed_update_can_match_async_rl_semantics(self): handler = _make_handler() diff --git a/lib/llm/src/block_manager/offload/filter.rs b/lib/llm/src/block_manager/offload/filter.rs index 36c8efe02e44..248fff59fcf1 100644 --- a/lib/llm/src/block_manager/offload/filter.rs +++ b/lib/llm/src/block_manager/offload/filter.rs @@ -7,7 +7,7 @@ use std::sync::{Arc, Mutex, MutexGuard}; use tokio::runtime::Handle; use tokio::sync::Notify; -use tokio::time::Duration; +use tokio::time::{Duration, Instant}; use tokio_util::sync::CancellationToken; use crate::tokens::SequenceHash; @@ -49,7 +49,8 @@ impl FrequencyFilter { CriticalTaskExecutionHandle::new_with_runtime( move |cancel_token| async move { - let mut interval = tokio::time::interval(flush_interval); + let mut interval = + tokio::time::interval_at(Instant::now() + flush_interval, flush_interval); loop { tokio::select! { // Observe cancellation and exit the loop. @@ -198,7 +199,15 @@ mod tests { tokio::time::sleep(Duration::from_millis(300)).await; - // The count should have decayed from 4 to 2. + // The count should have decayed from 4 to 3. + { + let frequency_map = filter.frequency_map.lock().unwrap(); + assert_eq!(*frequency_map.get(&hash(0)).unwrap(), 3); + } + + tokio::time::sleep(Duration::from_millis(250)).await; + + // The count should have decayed from 3 to 2. { let frequency_map = filter.frequency_map.lock().unwrap(); assert_eq!(*frequency_map.get(&hash(0)).unwrap(), 2); @@ -206,7 +215,7 @@ mod tests { tokio::time::sleep(Duration::from_millis(250)).await; - // The count should have decayed from 2 to 1, and should be pruned. + // The count should have decayed from 2 to 1. { let frequency_map = filter.frequency_map.lock().unwrap(); assert_eq!(*frequency_map.get(&hash(0)).unwrap(), 1); @@ -214,7 +223,7 @@ mod tests { tokio::time::sleep(Duration::from_millis(250)).await; - // The count should have decayed from 1 to 0, and should be pruned. + // The count should have decayed from 1 to 0 and been pruned. { let frequency_map = filter.frequency_map.lock().unwrap(); assert!(frequency_map.get(&hash(0)).is_none()); diff --git a/lib/llm/src/http/service/generate.rs b/lib/llm/src/http/service/generate.rs index cd340a6fc91e..d79e293f350b 100644 --- a/lib/llm/src/http/service/generate.rs +++ b/lib/llm/src/http/service/generate.rs @@ -27,7 +27,7 @@ use serde::Serialize; use tracing::Instrument; use super::disconnect::create_connection_monitor; -use super::error::SanitizedError; +use super::error::{SanitizedError, find_canonical_error_in_chain}; use super::metrics::{ CancellationLabels, ErrorType, HttpQueueGuard, InflightGuard, ResponseMetricCollector, }; @@ -217,6 +217,23 @@ fn generate_internal_error_response() -> Response { ) } +fn generate_invalid_request_response( + error: &(dyn std::error::Error + 'static), +) -> Option { + let error = find_canonical_error_in_chain(error)?; + if error.class() != dynamo_runtime::error::ErrorClass::InvalidRequest { + return None; + } + Some(generate_error_response( + StatusCode::BAD_REQUEST, + "invalid_request_error", + error + .public_message() + .unwrap_or("Invalid request") + .to_string(), + )) +} + /// Borrowed worker envelope for vLLM-specific request fields. /// /// `token_ids` are intentionally absent: `PreprocessedRequest.token_ids` is @@ -1056,10 +1073,13 @@ async fn generate_dispatch( || super::metrics::request_was_cancelled(error.as_ref()); let was_rejected = super::metrics::request_was_rejected(error.as_ref()); let was_unavailable = super::metrics::request_was_unavailable(error.as_ref()); + let invalid_request = generate_invalid_request_response(error.as_ref()); inflight_guard.mark_error(if deadline_exceeded || was_cancelled { ErrorType::Cancelled } else if was_rejected || was_unavailable { ErrorType::Unavailable + } else if invalid_request.is_some() { + ErrorType::Validation } else { ErrorType::Internal }); @@ -1088,6 +1108,10 @@ async fn generate_dispatch( ); return generate_unavailable_response(); } + if let Some(response) = invalid_request { + tracing::debug!(%request_id, %error, "invalid generate request"); + return response; + } tracing::error!(%request_id, error = %format!("{error:#}"), "engine generate call failed"); return generate_internal_error_response(); } @@ -1141,6 +1165,11 @@ async fn generate_dispatch( tracing::warn!(%request_id, %error, "generate stream failed: no worker available"); return generate_unavailable_response(); } + if let Some(response) = generate_invalid_request_response(error.as_ref()) { + inflight_guard.mark_error(ErrorType::Validation); + tracing::debug!(%request_id, %error, "invalid generate request"); + return response; + } inflight_guard.mark_error(ErrorType::Internal); tracing::error!(%request_id, %error, "failed to fold generate stream"); generate_internal_error_response() @@ -1262,6 +1291,8 @@ pub(crate) mod tests { struct WorkerUnavailableStreamEngine; + struct InvalidArgumentStreamEngine; + fn worker_unavailable_error() -> dynamo_runtime::error::DynamoError { dynamo_runtime::error::DynamoError::builder() .error_type(dynamo_runtime::error::ErrorType::WorkerUnavailable) @@ -1269,6 +1300,16 @@ pub(crate) mod tests { .build() } + fn invalid_argument_error() -> dynamo_runtime::error::DynamoError { + let message = "TITO requests currently require an aggregated vLLM worker"; + dynamo_runtime::error::DynamoError::builder() + .error_type(dynamo_runtime::error::ErrorType::Backend( + dynamo_runtime::error::BackendError::InvalidArgument, + )) + .message(message) + .build() + } + struct MigrationMetricBackend { calls: AtomicU32, } @@ -1347,6 +1388,19 @@ pub(crate) mod tests { } } + #[async_trait::async_trait] + impl AsyncEngine, ManyOut>, Error> + for InvalidArgumentStreamEngine + { + async fn generate( + &self, + request: SingleIn, + ) -> Result>, Error> { + let stream = futures::stream::iter([Annotated::from_err(invalid_argument_error())]); + Ok(ResponseStream::new(Box::pin(stream), request.context())) + } + } + #[async_trait::async_trait] impl AsyncEngine, ManyOut>, Error> for TerminalEngine @@ -2768,6 +2822,34 @@ pub(crate) mod tests { assert_worker_unavailable_returns_503(Arc::new(WorkerUnavailableStreamEngine)).await; } + #[tokio::test] + async fn backend_invalid_argument_stream_returns_400() { + let (response, state) = + dispatch_engine(Arc::new(InvalidArgumentStreamEngine), "req-invalid").await; + + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + let body = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .expect("read error response"); + let body: serde_json::Value = serde_json::from_slice(&body).expect("parse error response"); + assert_eq!(body["error"]["message"], "Invalid request"); + assert!(!body.to_string().contains("aggregated vLLM worker")); + assert_eq!(body["error"]["type"], "invalid_request_error"); + assert_eq!(body["error"]["code"], 400); + + let metric_model = state.manager().metric_model_for("test-model"); + assert_eq!( + state.metrics_clone().get_request_counter( + metric_model, + &Endpoint::Generate, + &RequestType::Unary, + &Status::Error, + &ErrorType::Validation, + ), + 1 + ); + } + #[tokio::test] async fn immediate_engine_cancellation_returns_499() { let engine: crate::types::openai::generate::GenerateStreamingEngine = diff --git a/lib/llm/src/model_card.rs b/lib/llm/src/model_card.rs index 9b22c5a9efd1..950fb635ba51 100644 --- a/lib/llm/src/model_card.rs +++ b/lib/llm/src/model_card.rs @@ -1259,6 +1259,13 @@ impl ModelDeploymentCard { bytes_to_hash.extend_from_slice(b"\0vllm_enable_tower_connector_lora\0true"); } + if self.runtime_config.runtime_flag_enabled( + crate::local_model::runtime_config::VLLM_INFERENCE_V1_GENERATE_CAPABILITY, + ) { + bytes_to_hash + .extend_from_slice(b"\0vllm_inference_v1_generate\0true"); + } + // The Qwen video contract is resolved per cohort, not per card. // Nemotron contracts still partition WorkerSets by checksum. append_runtime_contract_checksum( @@ -3246,6 +3253,26 @@ mod ownership_tests { assert_ne!(missing.mdcsum(), enabled.mdcsum()); } + #[test] + fn vllm_generate_capability_isolates_worker_sets() { + use crate::local_model::runtime_config::VLLM_INFERENCE_V1_GENERATE_CAPABILITY; + + let missing = ModelDeploymentCard::with_name_only("model"); + let mut disabled = ModelDeploymentCard::with_name_only("model"); + disabled.runtime_config.runtime_data.insert( + VLLM_INFERENCE_V1_GENERATE_CAPABILITY.to_string(), + false.into(), + ); + let mut enabled = ModelDeploymentCard::with_name_only("model"); + enabled.runtime_config.runtime_data.insert( + VLLM_INFERENCE_V1_GENERATE_CAPABILITY.to_string(), + true.into(), + ); + + assert_eq!(missing.mdcsum(), disabled.mdcsum()); + assert_ne!(missing.mdcsum(), enabled.mdcsum()); + } + #[test] fn qwen_video_processor_contract_stays_out_of_the_checksum() { use crate::local_model::runtime_config::VLLM_QWEN_VIDEO_PROCESSOR_CONTRACT_RUNTIME_KEY; diff --git a/lib/llm/src/protocols/openai/generate.rs b/lib/llm/src/protocols/openai/generate.rs index d6c7f47a3447..fcb1ed67912d 100644 --- a/lib/llm/src/protocols/openai/generate.rs +++ b/lib/llm/src/protocols/openai/generate.rs @@ -13,6 +13,7 @@ use std::collections::HashMap; use anyhow::Result; +use base64::Engine as _; use dynamo_runtime::error::{BackendError, DynamoError, ErrorType as DynamoErrorType}; use futures::{Stream, StreamExt, pin_mut}; use serde::de::DeserializeOwned; @@ -367,6 +368,57 @@ where .map_err(|error| format!("sampling_params.{name}: {error}")) } +#[derive(Serialize, Deserialize, Debug, Clone, PartialEq, Eq)] +#[serde(untagged)] +pub enum GenerateRoutedExperts { + Legacy(String), + Tensor(GenerateRoutedExpertsTensor), +} + +#[derive(Serialize, Deserialize, Debug, Clone, PartialEq, Eq)] +pub struct GenerateRoutedExpertsTensor { + pub data: String, + pub shape: Vec, + pub start: usize, + pub dtype: String, +} + +impl GenerateRoutedExperts { + fn validate(&self) -> Result<()> { + let Self::Tensor(tensor) = self else { + return Ok(()); + }; + anyhow::ensure!( + tensor.shape.len() == 3 && tensor.shape.iter().all(|dimension| *dimension > 0), + "structured routed_experts shape must contain three positive dimensions" + ); + let width = match tensor.dtype.as_str() { + "uint8" | "int8" => 1, + "uint16" | "int16" => 2, + "uint32" | "int32" => 4, + "uint64" | "int64" => 8, + other => anyhow::bail!("unsupported routed_experts dtype {other:?}"), + }; + let elements = tensor.shape.iter().try_fold(1usize, |total, dimension| { + total + .checked_mul(*dimension) + .ok_or_else(|| anyhow::anyhow!("routed_experts shape overflows")) + })?; + let expected_bytes = elements + .checked_mul(width) + .ok_or_else(|| anyhow::anyhow!("routed_experts byte count overflows"))?; + let decoded = base64::engine::general_purpose::STANDARD + .decode(&tensor.data) + .map_err(|error| anyhow::anyhow!("invalid routed_experts base64 data: {error}"))?; + anyhow::ensure!( + decoded.len() == expected_bytes, + "routed_experts decoded byte count {} does not match shape and dtype ({expected_bytes})", + decoded.len() + ); + Ok(()) + } +} + /// A single choice in a `GenerateResponse`. #[derive(Serialize, Deserialize, Debug, Clone)] pub struct GenerateResponseChoice { @@ -378,7 +430,7 @@ pub struct GenerateResponseChoice { pub finish_reason: Option, - pub routed_experts: Option, + pub routed_experts: Option, pub sampling_mask: Option>>, } @@ -408,7 +460,7 @@ struct GenerateChoiceAcc { token_ids: Vec, logprobs: Option>, finish_reason: Option, - routed_experts: Option, + routed_experts: Option, sampling_mask: Option>>, } @@ -645,11 +697,14 @@ impl GenerateAggregator { }); if let Some(engine_data) = output.engine_data.as_ref() { if let Some(routed_experts) = engine_data.get("routed_experts") { - choice.routed_experts = Some( + let routed_experts: GenerateRoutedExperts = serde_json::from_value(routed_experts.clone()).map_err(|error| { anyhow::anyhow!("invalid generate routed_experts payload: {error}") - })?, - ); + })?; + routed_experts.validate().map_err(|error| { + anyhow::anyhow!("invalid generate routed_experts payload: {error}") + })?; + choice.routed_experts = Some(routed_experts); } if let Some(sampling_mask) = engine_data.get("sampling_mask") { choice.sampling_mask = Some( @@ -1195,8 +1250,10 @@ mod tests { assert_eq!(response.choices[0].token_ids, Some(vec![100, 101])); assert_eq!(response.choices[0].finish_reason.as_deref(), Some("length")); assert_eq!( - response.choices[0].routed_experts.as_deref(), - Some("encoded-experts") + response.choices[0].routed_experts.as_ref(), + Some(&GenerateRoutedExperts::Legacy( + "encoded-experts".to_string() + )) ); assert_eq!( response.choices[0].sampling_mask, @@ -1264,16 +1321,73 @@ mod tests { assert_eq!(response.choices[0].index, 0); assert_eq!( - response.choices[0].routed_experts.as_deref(), - Some("experts-0") + response.choices[0].routed_experts.as_ref(), + Some(&GenerateRoutedExperts::Legacy("experts-0".to_string())) ); assert_eq!(response.choices[1].index, 1); assert_eq!( - response.choices[1].routed_experts.as_deref(), - Some("experts-1") + response.choices[1].routed_experts.as_ref(), + Some(&GenerateRoutedExperts::Legacy("experts-1".to_string())) + ); + } + + #[tokio::test] + async fn generate_response_accepts_structured_routed_experts() { + let payload = json!({ + "data": "AQI=", + "shape": [2, 1, 1], + "start": 3, + "dtype": "uint8" + }); + let stream = futures::stream::iter([Annotated::from_data(LLMEngineOutput { + token_ids: vec![100], + index: Some(0), + finish_reason: Some(crate::protocols::common::FinishReason::Stop), + engine_data: Some(json!({"routed_experts": payload})), + ..Default::default() + })]); + + let response = + GenerateResponse::from_annotated_stream(stream, "req-routed-structured".to_string()) + .await + .expect("aggregate structured routed experts"); + + assert_eq!( + response.choices[0].routed_experts, + Some(GenerateRoutedExperts::Tensor(GenerateRoutedExpertsTensor { + data: "AQI=".to_string(), + shape: vec![2, 1, 1], + start: 3, + dtype: "uint8".to_string(), + })) ); } + #[tokio::test] + async fn generate_response_rejects_structured_routed_expert_byte_mismatch() { + let stream = futures::stream::iter([Annotated::from_data(LLMEngineOutput { + token_ids: vec![100], + index: Some(0), + finish_reason: Some(crate::protocols::common::FinishReason::Stop), + engine_data: Some(json!({ + "routed_experts": { + "data": "AQI=", + "shape": [3, 1, 1], + "start": 0, + "dtype": "uint8" + } + })), + ..Default::default() + })]); + + let error = + GenerateResponse::from_annotated_stream(stream, "req-routed-structured".to_string()) + .await + .expect_err("routed expert byte mismatch must fail"); + + assert!(error.to_string().contains("decoded byte count")); + } + #[tokio::test] async fn generate_response_rejects_malformed_routed_experts() { let stream = futures::stream::iter([Annotated::from_data(LLMEngineOutput {