From f13e07a6786e7d229d5715a8b3ea572bf5e272f1 Mon Sep 17 00:00:00 2001 From: Limark Dcunha Date: Mon, 28 Sep 2026 22:53:14 -0400 Subject: [PATCH 1/3] Add SGLang prefill/decode disaggregation support Signed-off-by: Limark Dcunha --- .../advanced-guides/asyncio-best-practices.md | 19 +- .../serve/core/configs/engine_adapter.py | 67 +++ .../serve/core/configs/llm_config.py | 63 +-- .../serve/engines/common/__init__.py | 0 .../engines/common/kv_transfer/__init__.py | 0 .../serve/engines/common/kv_transfer/base.py | 176 ++++++++ .../{vllm => common}/kv_transfer/factory.py | 24 +- .../engines/sglang/kv_transfer/__init__.py | 0 .../sglang/kv_transfer/pd_connector.py | 223 ++++++++++ .../serve/engines/sglang/sglang_engine.py | 211 +++++++++- .../serve/engines/vllm/kv_transfer/base.py | 215 ++-------- .../serve/engines/vllm/kv_transfer/lmcache.py | 4 +- .../serve/engines/vllm/kv_transfer/moriio.py | 4 +- .../vllm/kv_transfer/multi_connector.py | 9 +- .../serve/engines/vllm/kv_transfer/nixl.py | 4 +- .../prefill_decode/builder.py | 118 +++++- .../prefill_decode/pd_server.py | 171 ++++++-- .../configs/test_llm_config_num_devices.py | 26 ++ .../vllm/kv_transfer_backends/test_factory.py | 6 +- .../test_moriio_connector.py | 6 +- .../test_multi_connector.py | 6 +- python/ray/serve/_private/constants.py | 10 +- python/ray/serve/tests/BUILD.bazel | 3 +- .../serve/test_llm_serve_sglang_pd.py | 380 ++++++++++++++++++ release/release_tests.yaml | 22 + 25 files changed, 1492 insertions(+), 275 deletions(-) create mode 100644 python/ray/llm/_internal/serve/core/configs/engine_adapter.py create mode 100644 python/ray/llm/_internal/serve/engines/common/__init__.py create mode 100644 python/ray/llm/_internal/serve/engines/common/kv_transfer/__init__.py create mode 100644 python/ray/llm/_internal/serve/engines/common/kv_transfer/base.py rename python/ray/llm/_internal/serve/engines/{vllm => common}/kv_transfer/factory.py (85%) create mode 100644 python/ray/llm/_internal/serve/engines/sglang/kv_transfer/__init__.py create mode 100644 python/ray/llm/_internal/serve/engines/sglang/kv_transfer/pd_connector.py create mode 100644 python/ray/llm/tests/serve/cpu/configs/test_llm_config_num_devices.py create mode 100644 release/llm_tests/serve/test_llm_serve_sglang_pd.py diff --git a/doc/source/serve/advanced-guides/asyncio-best-practices.md b/doc/source/serve/advanced-guides/asyncio-best-practices.md index 48f10006eb2a..d05a60a545d0 100644 --- a/doc/source/serve/advanced-guides/asyncio-best-practices.md +++ b/doc/source/serve/advanced-guides/asyncio-best-practices.md @@ -80,8 +80,8 @@ For a synchronous deployment: How this method executes depends on configuration: -- With `RAY_SERVE_RUN_SYNC_IN_THREADPOOL=0` (current default), `__call__` runs directly on the user event loop and blocks it for 1 second. -- With `RAY_SERVE_RUN_SYNC_IN_THREADPOOL=1`, Serve offloads `__call__` to a threadpool so the event loop stays responsive. +- With `RAY_SERVE_RUN_SYNC_IN_THREADPOOL=1` (the default), Serve offloads `__call__` to a threadpool so the event loop stays responsive. +- With `RAY_SERVE_RUN_SYNC_IN_THREADPOOL=0`, `__call__` runs directly on the user event loop and blocks it for 1 second. ### FastAPI ingress (`@serve.ingress`) @@ -274,16 +274,7 @@ Ray Serve exposes several environment variables that control how user code inter ### `RAY_SERVE_RUN_SYNC_IN_THREADPOOL` -By default (`RAY_SERVE_RUN_SYNC_IN_THREADPOOL=0`), which means synchronous methods in a deployment run directly on the user event loop. To help you migrate to a safer model, Serve emits a warning like: - -> `RAY_SERVE_RUN_SYNC_IN_THREADPOOL_WARNING`: Calling sync method '...' directly on the asyncio loop. In a future version, sync methods will be run in a threadpool by default... - -This warning means: - -- You have a `def` method that is currently running on the event loop. -- In a future version, that method runs in a threadpool instead. - -You can opt in to the future behavior now by setting: +By default (`RAY_SERVE_RUN_SYNC_IN_THREADPOOL=1`), synchronous methods in a deployment run in a threadpool: ```bash export RAY_SERVE_RUN_SYNC_IN_THREADPOOL=1 @@ -294,11 +285,13 @@ When this flag is `1`: - Serve runs synchronous methods in a threadpool. - The event loop is free to keep serving other requests while sync methods run. -Before enabling this in production, make sure: +Make sure: - Your handler code and any shared state are thread-safe. - Your model objects can safely be used from multiple threads, or you protect them with locks. +Set this flag to `0` to retain the legacy behavior of running synchronous methods directly on the user event loop. + ### `RAY_SERVE_RUN_USER_CODE_IN_SEPARATE_THREAD` By default, Serve runs user code in a separate event loop from the replica's main/control loop: diff --git a/python/ray/llm/_internal/serve/core/configs/engine_adapter.py b/python/ray/llm/_internal/serve/core/configs/engine_adapter.py new file mode 100644 index 000000000000..a8d400185778 --- /dev/null +++ b/python/ray/llm/_internal/serve/core/configs/engine_adapter.py @@ -0,0 +1,67 @@ +"""Per-engine adapters. + +Each adapter declares how an engine names its KV connector and which +engine-config class builds its ``EngineConfig``. This replaces ``if vLLM / +elif SGLang`` branching in ``llm_config.py`` so adding an engine is additive: +add an adapter and register it here. +""" + +import abc +from typing import TYPE_CHECKING, Any, Dict, Optional, Type + +if TYPE_CHECKING: + pass + + +class EngineAdapter(abc.ABC): + @abc.abstractmethod + def connector_name(self, engine_kwargs: Dict[str, Any]) -> Optional[str]: + """The KV-connector registry name implied by engine_kwargs, or None.""" + ... + + @abc.abstractmethod + def engine_config_cls(self) -> Type: + """The engine-config class whose ``from_llm_config`` builds the config.""" + ... + + +class VLLMAdapter(EngineAdapter): + def connector_name(self, engine_kwargs: Dict[str, Any]) -> Optional[str]: + cfg = engine_kwargs.get("kv_transfer_config") + if not cfg: + return None + kv_connector = cfg.get("kv_connector") + if not kv_connector: + # Fail fast: a kv_transfer_config with no kv_connector is a + # misconfiguration, not a "no connector" case. + raise ValueError("Connector type is not specified.") + return kv_connector + + def engine_config_cls(self) -> Type: + from ray.llm._internal.serve.engines.vllm.vllm_models import VLLMEngineConfig + + return VLLMEngineConfig + + +class SGLangAdapter(EngineAdapter): + def connector_name(self, engine_kwargs: Dict[str, Any]) -> Optional[str]: + return ( + "SGLang" if engine_kwargs.get("disaggregation_transfer_backend") else None + ) + + def engine_config_cls(self) -> Type: + from ray.llm._internal.serve.engines.sglang.sglang_engine import ( + SGLangEngineConfig, + ) + + return SGLangEngineConfig + + +_ADAPTERS = {"vLLM": VLLMAdapter, "SGLang": SGLangAdapter} + + +def get_engine_adapter(llm_engine: str) -> EngineAdapter: + try: + return _ADAPTERS[llm_engine]() + except KeyError: + raise ValueError(f"Unsupported engine: {llm_engine}") diff --git a/python/ray/llm/_internal/serve/core/configs/llm_config.py b/python/ray/llm/_internal/serve/core/configs/llm_config.py index 55dde6d37b3c..daabbcde752d 100644 --- a/python/ray/llm/_internal/serve/core/configs/llm_config.py +++ b/python/ray/llm/_internal/serve/core/configs/llm_config.py @@ -45,14 +45,14 @@ TPUConfig, infer_hardware_kind_from_bundles, ) -from ray.llm._internal.serve.engines.vllm.kv_transfer.factory import ( +from ray.llm._internal.serve.engines.common.kv_transfer.factory import ( KVConnectorBackendFactory, ) from ray.llm._internal.serve.observability.logging import get_logger from ray.serve._private.config import DeploymentConfig, handle_num_replicas_auto if TYPE_CHECKING: - from ray.llm._internal.serve.engines.vllm.kv_transfer.base import ( + from ray.llm._internal.serve.engines.common.kv_transfer.base import ( BaseConnectorBackend, ) @@ -86,6 +86,7 @@ class LLMEngine(str, Enum): """Enum that represents an LLMEngine.""" vLLM = "vLLM" + SGLang = "SGLang" class LoraConfig(BaseModelExtended): @@ -146,7 +147,7 @@ class ModelLoadingConfig(BaseModelExtended): ) -EngineConfigType = Union[None, "VLLMEngineConfig"] # noqa: F821 +EngineConfigType = Union[None, "VLLMEngineConfig", "SGLangEngineConfig"] # noqa: F821 class LLMConfig(BaseModelExtended): @@ -590,19 +591,23 @@ def get_engine_config(self) -> EngineConfigType: if self._engine_config: return self._engine_config - if self.llm_engine == LLMEngine.vLLM: - from ray.llm._internal.serve.engines.vllm.vllm_models import ( - VLLMEngineConfig, - ) - - self._engine_config = VLLMEngineConfig.from_llm_config(self) - else: - # Note (genesu): This should never happen because we validate the engine - # in the config. - raise ValueError(f"Unsupported engine: {self.llm_engine}") + from ray.llm._internal.serve.core.configs.engine_adapter import ( + get_engine_adapter, + ) + adapter = get_engine_adapter(self.llm_engine) + self._engine_config = adapter.engine_config_cls().from_llm_config(self) return self._engine_config + @property + def num_devices(self) -> int: + """Total devices per replica (tensor-parallel × pipeline-parallel). + + Neutral accessor used by connector backends for port spacing, so the + connector base never reaches into an engine-specific config type. + """ + return self.get_engine_config().num_devices + def update_engine_kwargs(self, **kwargs: Any) -> None: """Update the engine_kwargs and the engine_config engine_kwargs. @@ -612,27 +617,35 @@ def update_engine_kwargs(self, **kwargs: Any) -> None: self.engine_kwargs.update(kwargs) # engine_config may be created before engine starts, this makes sure # the engine_config is updated with the latest engine_kwargs. - if self._engine_config: + # Not every engine config mirrors engine_kwargs + # (the minimal SGLangEngineConfig does not — SGLangServer reads llm_config.engine_kwargs directly), so + # only sync configs that carry it. + if self._engine_config is not None and hasattr( + self._engine_config, "engine_kwargs" + ): self._engine_config.engine_kwargs.update(kwargs) def setup_engine_backend(self): self._setup_kv_connector_backend() def _setup_kv_connector_backend(self): - """Private method to setup kv connector depending on the local deployment state""" - # 1. validate that the backend is one of the backends supported (Nixl or LMCache) - kv_transfer_config = self.engine_kwargs.get("kv_transfer_config") - if not kv_transfer_config: - return + """Create and set up the KV connector backend for this engine, if any. - kv_connector = kv_transfer_config.get("kv_connector") - if not kv_connector: - raise ValueError("Connector type is not specified.") + The connector name is resolved through the per-engine adapter so this + path carries no engine-specific knowledge (vLLM reads ``kv_transfer_config.kv_connector``; + SGLang reads ``disaggregation_transfer_backend``). + """ - # 2. Setup the backend using factory - kv_connector_backend = KVConnectorBackendFactory.create_backend( - kv_connector, self + from ray.llm._internal.serve.core.configs.engine_adapter import ( + get_engine_adapter, ) + + adapter = get_engine_adapter(self.llm_engine) + name = adapter.connector_name(self.engine_kwargs) + if not name: + return + + kv_connector_backend = KVConnectorBackendFactory.create_backend(name, self) kv_connector_backend.setup() # 3. Stash the instance so the P/D orchestrator can reach the connector's # coordination protocol (request shaping, peer binding, handoff diff --git a/python/ray/llm/_internal/serve/engines/common/__init__.py b/python/ray/llm/_internal/serve/engines/common/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/python/ray/llm/_internal/serve/engines/common/kv_transfer/__init__.py b/python/ray/llm/_internal/serve/engines/common/kv_transfer/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/python/ray/llm/_internal/serve/engines/common/kv_transfer/base.py b/python/ray/llm/_internal/serve/engines/common/kv_transfer/base.py new file mode 100644 index 000000000000..10410e96fc3d --- /dev/null +++ b/python/ray/llm/_internal/serve/engines/common/kv_transfer/base.py @@ -0,0 +1,176 @@ +import abc +import random +import string +from typing import TYPE_CHECKING, Any, Dict, Optional, Union + +from ray import serve + +if TYPE_CHECKING: + from ray.llm._internal.serve.core.configs.llm_config import LLMConfig + from ray.llm._internal.serve.core.configs.openai_api_models import ( + ChatCompletionRequest, + CompletionRequest, + ) + + # The two OpenAI request models the P/D orchestrator shapes. Defined under + # TYPE_CHECKING (and used as a string annotation) to avoid an import cycle + # between this module and the config/openai-models modules. + RequestType = Union[ChatCompletionRequest, CompletionRequest] + + +def clamp_request_to_single_token(request: "RequestType") -> None: + """Clamp a prefill request to a single, non-streaming token (in place).""" + request.max_tokens = 1 + if hasattr(request, "max_completion_tokens"): + request.max_completion_tokens = 1 + request.stream = False + if hasattr(request, "stream_options"): + request.stream_options = None + + +class BaseConnectorBackend(abc.ABC): + # ---- P/D coordination protocol ---- + # + # These class attributes and methods let the P/D orchestrator + # (``PDOrchestratorMixin``) delegate request shaping, peer addressing, and + # handoff discipline to the connector. They are connector-agnostic: a + # connector picks a quadrant of (``requires_peer_binding``, + # ``concurrent_handoff``) and implements ``prepare_prefill_request`` / + # ``prepare_decode_request`` accordingly. + # + # ``requires_peer_binding``: + # * False -> the orchestrator dispatches prefill via the standard handle + # path; the peer (if any) is resolved post-hoc from the prefill response. + # * True -> the orchestrator selects the prefill replica first + # (``choose_replica``) and passes its ``replica_metadata`` to the backend + # as ``peer`` (pre-dispatch addressing). + # + # ``concurrent_handoff``: + # * False -> prefill runs to its first chunk before local decode starts + # (sequential handoff). + # * True -> prefill dispatch and local decode run concurrently. + # + # The two flags are independent; the known combos: + # * (False, False) — e.g. NixlConnector / LMCacheConnectorV1: decode learns + # everything it needs (remote engine id / block ids) from the prefill + # response, so it must wait for that response (sequential). + # * (True, True) — e.g. MoRIIO WRITE mode / SGLang: the peer address is + # bound up front and prefill *pushes* KV to decode, so decode needs + # nothing from prefill's response and can start immediately (concurrent). + # Concurrent handoff is only possible because peer binding happens up + # front. + # * (True, False) — e.g. MoRIIO READ mode: the peer is bound up front, + # but decode *pulls* KV using block ids returned in prefill's response, + # so it still waits for prefill to finish (sequential). + requires_peer_binding: bool = False + concurrent_handoff: bool = False + + def __init__(self, llm_config: "LLMConfig"): + """Base class for connector backends. + + Args: + llm_config: The llm configuration for this engine + """ + self.llm_config = llm_config + + def _get_unique_suffix(self, len: int = 6) -> str: + """Generates unique alphanumeric suffix. + + Args: + len: Length of the suffix to generate. + Returns: + A unique alphanumeric suffix string of specified length. + """ + return "".join(random.choices(string.ascii_letters + string.digits, k=len)) + + def _compute_port_offset(self) -> int: + """Compute a deterministic port offset for this replica. + + Uses data_parallel_rank if DP case, otherwise falls back to + the replica rank assigned by Ray Serve (TP/PP case). + + For TP/PP cases, multiply by num_devices (tp × pp) to reserve + sufficient port space, since each worker needs a unique port. + Each TP worker adds its tp_rank (0, 1, ..., tp_size-1) to the + base port at bind time, and PP stages also need separate ports. + + ``num_devices`` is read from the neutral ``LLMConfig`` accessor, so this + base never reaches into an engine-specific config type. + + Returns: + Non-negative integer offset to add to a base port. + """ + # Prefer explicit DP rank when available + dp_rank = self.llm_config.engine_kwargs.get("data_parallel_rank") + if isinstance(dp_rank, int) and dp_rank >= 0: + # vLLM already accounts for TP spacing in DP offset calculation + # (data_parallel_rank × tp_size), don't multiply here + return dp_rank + + # NOTE (jeffreywang): A missing replica context must fail loudly, not + # silently return a 0 offset that collides colocated replicas on the + # same side-channel port. get_replica_context() raises RayServeException + # outside a replica. + rc = serve.get_replica_context() + num_devices = self.llm_config.num_devices + return rc.rank.rank * num_devices + + @abc.abstractmethod + def prepare_prefill_request( + self, *, request: "RequestType", peer: Optional[Dict[str, Any]] + ) -> "RequestType": + """Shape the request sent to the remote prefill engine. + + Args: + request: The incoming chat/completion request. + peer: The selected prefill replica's ``replica_metadata`` dict when + the connector opted into pre-dispatch peer binding + (``requires_peer_binding=True``), else None. + + Returns: + A new request object to dispatch to the prefill engine. + """ + ... + + @abc.abstractmethod + def prepare_decode_request( + self, + *, + request: "RequestType", + peer: Optional[Dict[str, Any]], + prefill_response: Optional[Any], + ) -> "RequestType": + """Shape the request run on the local decode engine. + + Args: + request: The incoming chat/completion request. + peer: The selected prefill replica's ``replica_metadata`` dict when + the connector opted into pre-dispatch peer binding, else None. + prefill_response: The captured prefill response chunk whose + ``kv_transfer_params`` may be forwarded, or None when no chunk is + captured before decode starts (concurrent-handoff mode). + + Returns: + A new request object to run on the local decode engine. + """ + ... + + def setup(self) -> None: + """Setup the connector backend. + + This method is called to setup the connector backend. + """ + pass + + def replica_metadata(self) -> Dict[str, Any]: + """Static per-replica coordination data published to the orchestrator. + + Surfaced via the replica-metadata hook on ``ReplicaSelection`` so that a + connector opting into ``requires_peer_binding`` can address the selected + prefill peer. The default backend publishes nothing; connectors that need + to advertise an address (e.g. MoRIIO's zmq endpoint) override this. + + Returns: + A JSON-serializable dict of per-replica metadata (empty by default). + """ + return {} diff --git a/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/factory.py b/python/ray/llm/_internal/serve/engines/common/kv_transfer/factory.py similarity index 85% rename from python/ray/llm/_internal/serve/engines/vllm/kv_transfer/factory.py rename to python/ray/llm/_internal/serve/engines/common/kv_transfer/factory.py index 54b7c12be713..2eb29a7cd629 100644 --- a/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/factory.py +++ b/python/ray/llm/_internal/serve/engines/common/kv_transfer/factory.py @@ -7,9 +7,8 @@ from typing import TYPE_CHECKING, Type, Union -from ray.llm._internal.serve.engines.vllm.kv_transfer.base import ( +from ray.llm._internal.serve.engines.common.kv_transfer.base import ( BaseConnectorBackend, - DefaultConnectorBackend, ) from ray.llm._internal.serve.observability.logging import get_logger from ray.llm._internal.serve.utils.registry import get_registry @@ -20,6 +19,21 @@ logger = get_logger(__name__) + +def _default_connector_backend() -> Type["BaseConnectorBackend"]: + """The factory's fallback backend for unregistered connector names. + + The default policy is expressed with vLLM's ``kv_transfer_params``, so the + concrete fallback lives on the vLLM side. Imported lazily so this neutral + factory module carries no module-level vLLM dependency. + """ + from ray.llm._internal.serve.engines.vllm.kv_transfer.base import ( + DefaultConnectorBackend, + ) + + return DefaultConnectorBackend + + # Get the registry instance for KV connector backends _kv_backend_registry = get_registry("kv_connector_backend") @@ -75,11 +89,12 @@ def get_backend_class(cls, name: str) -> Type["BaseConnectorBackend"]: try: return _kv_backend_registry.get(name) except ValueError: + default_backend = _default_connector_backend() logger.warning( f"Unsupported connector backend: {name}. " - f"Using default: {DefaultConnectorBackend.__name__}." + f"Using default: {default_backend.__name__}." ) - return DefaultConnectorBackend + return default_backend except Exception as e: raise ImportError( f"Failed to load connector backend '{name}': {type(e).__name__}: {e}" @@ -122,6 +137,7 @@ def unregister_backend(cls, name: str) -> None: "NixlConnector": "ray.llm._internal.serve.engines.vllm.kv_transfer.nixl:NixlConnectorBackend", "MultiConnector": "ray.llm._internal.serve.engines.vllm.kv_transfer.multi_connector:MultiConnectorBackend", "MoRIIOConnector": "ray.llm._internal.serve.engines.vllm.kv_transfer.moriio:MoRIIOConnectorBackend", + "SGLang": "ray.llm._internal.serve.engines.sglang.kv_transfer.pd_connector:SGLangConnectorBackend", } diff --git a/python/ray/llm/_internal/serve/engines/sglang/kv_transfer/__init__.py b/python/ray/llm/_internal/serve/engines/sglang/kv_transfer/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/python/ray/llm/_internal/serve/engines/sglang/kv_transfer/pd_connector.py b/python/ray/llm/_internal/serve/engines/sglang/kv_transfer/pd_connector.py new file mode 100644 index 000000000000..97051f1ff405 --- /dev/null +++ b/python/ray/llm/_internal/serve/engines/sglang/kv_transfer/pd_connector.py @@ -0,0 +1,223 @@ +"""SGLang P/D connector backend for Ray Serve LLM. + +SGLang runs prefill and decode concurrently; prefill PUSHES the KV cache to +decode through a bootstrap server that lives on the prefill worker. The decode +side must know the prefill node's bootstrap (host, port) up front. B +Both protocol flags are on: + * ``concurrent_handoff = True`` (prefill pushes; decode needs nothing from + prefill's response) + * ``requires_peer_binding = True`` (decode binds to the selected prefill + replica's bootstrap address before dispatch). + +``setup()`` picks a free bootstrap port (SGLang's default 8998 collides when +replicas share a node) and sets the SGLang server ``host`` to the routable node +IP (SGLang binds the bootstrap server to ``server_args.host``, default 127.0.0.1 +-> a remote decode would get Connection refused). + +It runs on BOTH sides -- every replica picks its own port and GPU block -- but +only prefill's address is ever consumed: the decode orchestrator selects a +prefill replica, reads its ``replica_metadata``, and stamps that address onto +both requests via ``peer``. Decode never publishes an address of its own. + +``bootstrap_room`` is derived deterministically from the incoming request id, so +the two stateless ``prepare_*`` calls agree without per-request backend state. +""" + +import hashlib +import secrets +from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple + +import ray +from ray import serve +from ray.llm._internal.serve.engines.common.kv_transfer.base import ( + BaseConnectorBackend, + clamp_request_to_single_token, +) + +if TYPE_CHECKING: + from ray.llm._internal.serve.engines.common.kv_transfer.base import RequestType + +# SGLang's default disaggregation bootstrap port. Colocated replicas collide on +# it, so setup() adds _compute_port_offset() on top of the base. +DEFAULT_BOOTSTRAP_PORT_BASE = 8998 + +# experimental_configs key for overriding the bootstrap port base. The builder +# shifts decode's base off prefill's default (see builder.py) so a colocated P+D +# pair on one node doesn't collide; per-replica offset is applied on top. +BOOTSTRAP_PORT_BASE_KEY = "SGLANG_BOOTSTRAP_PORT_BASE" + +# Width of the bootstrap_room id, used for BOTH the rid-hash mask and the +# random fallback below, so the two stateless prepare_* calls agree. +# the modulo only needs a non-negative int, and one spare +# bit keeps the value clear of any signed-64 boundary as it crosses JSON and +# SGLang's own typing. Any width <= 63 is correct; this is the conservative one. +_ROOM_BITS = 62 + +# Attribute the minted room is cached under when the client sent no ``rid``. +# Set on the incoming request so both prepare_* calls read the same value. +_ROOM_ATTR = "_ray_sglang_bootstrap_room" + + +class SGLangConnectorBackend(BaseConnectorBackend): + """SGLang P/D connector: concurrent handoff, prefill-address-first.""" + + concurrent_handoff: bool = True + requires_peer_binding: bool = True + + # Set by setup(); published via replica_metadata(). + _bootstrap_host: Optional[str] = None + _bootstrap_port: Optional[int] = None + + @staticmethod + def _check_request_model_has_bootstrap_fields() -> None: + """Fail early if the resolved OpenAI request model lacks bootstrap fields. + + Ray's ``ChatCompletionRequest`` resolves to SGLang's model only in a + SGLang-only environment (the import chain in ``openai_api_models`` tries + vLLM first). If vLLM is also installed, it resolves to vLLM's model, + which has no ``bootstrap_room`` — assigning it in ``prepare_*`` then + raises deep in request handling. Surface it at startup instead. + """ + from ray.llm._internal.serve.core.configs.openai_api_models import ( + ChatCompletionRequest, + ) + + if "bootstrap_room" not in ChatCompletionRequest.model_fields: + raise RuntimeError( + "SGLang P/D requires SGLang's OpenAI request models, but the " + "resolved ChatCompletionRequest has no 'bootstrap_room' field. " + "This happens when vLLM is installed alongside SGLang (Ray's " + "import chain then picks vLLM's request model). SGLang P/D needs " + "a SGLang-only environment." + ) + + def setup(self) -> None: + """Pick a free bootstrap port + set host to the node IP, before engine start.""" + self._check_request_model_has_bootstrap_fields() + offset = self._compute_port_offset() + engine_kwargs = self.llm_config.engine_kwargs + + # SGLang binds the bootstrap server to server_args.host (default + # 127.0.0.1). The remote decode dials the node IP we advertise, so bind + # the routable IP or it gets Connection refused. + host = ray.util.get_node_ip_address() + engine_kwargs["host"] = host + + # A user-pinned explicit port wins (advanced/escape hatch). Otherwise + # compute base + per-replica offset so colocated replicas never share a + # port. The base is overridable via experimental_configs (the builder + # shifts decode's base off prefill's default); the offset is derived from + # the replica rank (or DP rank), matching the MoRIIO connector. + port = engine_kwargs.get("disaggregation_bootstrap_port") + if port is None: + base = int( + self.llm_config.experimental_configs.get( + BOOTSTRAP_PORT_BASE_KEY, DEFAULT_BOOTSTRAP_PORT_BASE + ) + ) + port = base + offset + engine_kwargs["disaggregation_bootstrap_port"] = port + + # base_gpu_id needs the same per-replica shift as the bootstrap port. + # engine_kwargs is shared across a side's replicas, so without a shift + # every replica inherits the same base_gpu_id and drives the same + # physical devices. The configured value stays the side's starting + # device; each replica moves up by its own device block. + # + # Deliberately NOT reusing the port offset computed above. In the DP + # case that offset is the bare data_parallel_rank (base.py:105), which + # spaces ports but not devices: with tp_size=2, DP ranks 0 and 1 give + # base_gpu_id 0 and 1, so replica 0 (GPUs 0-1) and replica 1 (GPUs 1-2) + # both claim GPU 1. _compute_gpu_offset() multiplies by num_devices + # instead, so the blocks tile without overlap. + engine_kwargs["base_gpu_id"] = ( + int(engine_kwargs.get("base_gpu_id", 0)) + self._compute_gpu_offset() + ) + + self._bootstrap_host = host + self._bootstrap_port = port + + def _compute_gpu_offset(self) -> int: + """This replica's device-block offset within its side of the P/D pair. + + Each replica occupies ``num_devices`` (tp x pp) consecutive GPUs, so the + n-th replica *on this node* starts at ``n * num_devices``. + + Keyed on ``local_rank``, not the cluster-wide ``rank``: base_gpu_id + indexes node-local CUDA devices, so a global rank would walk off the end + of a multi-node deployment's per-node device range (the P/D release + compute is two GPU workers). A missing replica context raises rather + than silently returning 0, which would double-book GPU 0. + """ + rc = serve.get_replica_context() + return rc.rank.local_rank * self.llm_config.num_devices + + def replica_metadata(self) -> Dict[str, Any]: + """Publish this (prefill) replica's bootstrap address for the decode peer.""" + return { + "bootstrap_host": self._bootstrap_host, + "bootstrap_port": self._bootstrap_port, + } + + def _peer_address(self, peer: Optional[Dict[str, Any]]) -> Tuple[str, int]: + host = (peer or {}).get("bootstrap_host") + port = (peer or {}).get("bootstrap_port") + if not host or not port: + raise ValueError( + "SGLang peer is missing bootstrap_host/bootstrap_port: the " + "selected prefill replica did not publish its bootstrap address " + "(is the prefill deployment using llm_engine='SGLang' with a " + "disaggregation_transfer_backend?)." + ) + return host, port + + def _bootstrap_room(self, request: Any) -> int: + """Per-request room id, shared by both prepare_* calls on this request. + + SGLang's ``rid`` is client-optional and defaults to ``None`` - most + OpenAI-compatible clients never set it, so hashing it would collide + every concurrent request onto one room and let the bootstrap server + mix KV caches. When ``rid`` is absent we mint a fresh random room and + cache it on the request, so the prefill and decode calls agree without + depending on object identity or per-backend request state. + """ + rid = getattr(request, "rid", None) + if rid is not None: + digest = hashlib.sha256(str(rid).encode()).hexdigest() + return int(digest, 16) & ((1 << _ROOM_BITS) - 1) + + room = getattr(request, _ROOM_ATTR, None) + if room is None: + room = secrets.randbits(_ROOM_BITS) + object.__setattr__(request, _ROOM_ATTR, room) + return room + + def _stamp(self, request: Any, peer: Optional[Dict[str, Any]]) -> Any: + host, port = self._peer_address(peer) + out = request.model_copy(deep=True) + out.bootstrap_host = host + out.bootstrap_port = port + out.bootstrap_room = self._bootstrap_room(request) + return out + + def prepare_prefill_request( + self, *, request: "RequestType", peer: Optional[Dict[str, Any]] + ) -> "RequestType": + prefill_request = self._stamp(request, peer) + + # Prefill only produces the KV cache, so it is clamped; decode is NOT + # clamped -- it generates the real output. + clamp_request_to_single_token(prefill_request) + return prefill_request + + def prepare_decode_request( + self, + *, + request: "RequestType", + peer: Optional[Dict[str, Any]], + prefill_response: Optional[Any], + ) -> "RequestType": + # Concurrent handoff: the orchestrator passes prefill_response=None + # (pd_server.py, concurrent branch) -- decode needs only the SAME + # prefill bootstrap host/port/room to rendezvous. + return self._stamp(request, peer) diff --git a/python/ray/llm/_internal/serve/engines/sglang/sglang_engine.py b/python/ray/llm/_internal/serve/engines/sglang/sglang_engine.py index 2a7156ee2557..e5ec871126cd 100644 --- a/python/ray/llm/_internal/serve/engines/sglang/sglang_engine.py +++ b/python/ray/llm/_internal/serve/engines/sglang/sglang_engine.py @@ -15,6 +15,7 @@ import time import uuid from typing import ( + TYPE_CHECKING, Any, AsyncGenerator, List, @@ -25,8 +26,12 @@ from pydantic import BaseModel +from ray.llm._internal.common.utils.cloud_utils import CloudMirrorConfig from ray.llm._internal.serve.constants import ENABLE_WORKER_PROCESS_SETUP_HOOK -from ray.llm._internal.serve.core.configs.llm_config import LLMConfig +from ray.llm._internal.serve.core.configs.llm_config import ( + DiskMultiplexConfig, + LLMConfig, +) from ray.llm._internal.serve.core.configs.openai_api_models import ( ChatCompletionRequest, ChatCompletionResponse, @@ -41,11 +46,83 @@ TokenizeRequest, TokenizeResponse, ) +from ray.llm._internal.serve.core.engine.protocol import LLMEngine from ray.llm._internal.serve.core.protocol import RawRequestInfo + +if TYPE_CHECKING: + from ray.llm._internal.serve.core.configs.openai_api_models import ( + ErrorResponse, + TranscriptionRequest, + TranscriptionResponse, + ) from ray.llm._internal.serve.core.server.llm_server import ( _add_openai_models_retrieve_route, ) + +class SGLangEngineConfig(BaseModel): + """Minimal engine config for SGLang, exposing the fields telemetry needs. + + Unlike VLLMEngineConfig, this does not drive engine construction — + SGLangServer.__init__ passes engine_kwargs to sglang.Engine directly. + What it does carry: the model source fields. Node initialization reads + them to download the weights, then writes hf_model_id back as the on-disk + path. Usage telemetry reads them too. + + It stays minimal so get_engine_config() works without vLLM installed. + """ + + tensor_parallel_degree: int + num_devices: int + model_id: str + # The HuggingFace id / local path SGLang loads from. None for a pure + # CloudMirrorConfig source (mirror fetched under model_id); actual_hf_model_id + # then falls back to model_id. Overwritten with the on-disk path after + # download (see SGLangServer._download_and_resolve_model). + hf_model_id: Optional[str] = None + # Cloud bucket the weights are mirrored from, when model_source is a remote + # URI or a CloudMirrorConfig. None for local / HF-hub sources. + mirror_config: Optional["CloudMirrorConfig"] = None + + @property + def actual_hf_model_id(self) -> str: + return self.hf_model_id or self.model_id + + @classmethod + def from_llm_config(cls, llm_config: "LLMConfig") -> "SGLangEngineConfig": + from ray.llm._internal.common.utils.cloud_utils import ( + CloudMirrorConfig, + is_remote_path, + ) + + tp_size = llm_config.engine_kwargs.get("tp_size", 1) + pp_size = llm_config.engine_kwargs.get("pp_size", 1) + + # Mirror the vLLM mapping: resolve model_source into (hf_model_id, + # mirror_config). A remote URI or CloudMirrorConfig is a download + # address, not a HF id — the weights are fetched under model_id. + hf_model_id, mirror_config = None, None + model_source = llm_config.model_loading_config.model_source + if model_source is None: + hf_model_id = llm_config.model_id + elif isinstance(model_source, str): + if is_remote_path(model_source): + hf_model_id = llm_config.model_id + mirror_config = CloudMirrorConfig(bucket_uri=model_source) + else: + hf_model_id = model_source + else: + mirror_config = model_source + + return cls( + tensor_parallel_degree=tp_size, + num_devices=tp_size * pp_size, + model_id=llm_config.model_id, + hf_model_id=hf_model_id, + mirror_config=mirror_config, + ) + + logger = logging.getLogger(__name__) @@ -88,7 +165,7 @@ class SGLangWakeupConfig(BaseModel): _SLEEP_TAGS: frozenset[str] = frozenset({"kv_cache", "weights", "cuda_graph"}) -class SGLangServer: +class SGLangServer(LLMEngine): def __init__(self, llm_config: LLMConfig): self._llm_config = llm_config @@ -105,6 +182,12 @@ def __init__(self, llm_config: LLMConfig): "dependencies." ) from e + # Must run BEFORE engine_kwargs is snapshotted below: for SGLang PD the + # connector's setup() mutates engine_kwargs in place (host, + # disaggregation_bootstrap_port), and a snapshot taken first would not + # carry those into sglang.Engine. + llm_config.setup_engine_backend() + # Route SGLang engine metrics through Ray's metric agent (Ray's # Prometheus endpoint / dashboard). RayEngine runs schedulers as Ray # actors and auto-wires the Ray-backed stat_loggers when enable_metrics @@ -141,13 +224,77 @@ def noop_signal_handler(sig, action): # Returns default handler to satisfy signal.signal() return signature return signal.SIG_DFL + # Inject model_path from model_loading_config unless the user set it + # explicitly in engine_kwargs (explicit always wins). + engine_init_kwargs = dict(self.engine_kwargs) + + if "model_path" not in engine_init_kwargs: + engine_init_kwargs["model_path"] = self._download_and_resolve_model( + llm_config + ) + + if not engine_init_kwargs.get("model_path"): + raise ValueError( + "SGLang engine requires 'model_path' but it could not be determined. " + "Set it via model_loading_config.model_source or " + "engine_kwargs['model_path'] directly." + ) + try: # Override signal.signal with our no-op function signal.signal = noop_signal_handler - self.engine = Engine(**self.engine_kwargs) + self.engine = Engine(**engine_init_kwargs) finally: signal.signal = original_signal_func + @staticmethod + def _download_and_resolve_model(llm_config: LLMConfig) -> str: + """Download the model (if mirrored) and return the path SGLang loads from. + + SGLang builds its in-process engine here in ``__init__`` and, unlike the + vLLM path, does not go through ``initialize_node``. So the download must + happen on this replica's node before the engine starts: + + Mirrored weights end up loading from the on-disk copy Ray fetched, never + a HuggingFace id pointing at weights meant to come from the mirror. + """ + from ray.llm._internal.common.utils.download_utils import ( + STREAMING_LOAD_FORMATS, + NodeModelDownloadable, + download_model_files, + get_model_location_on_disk, + ) + + engine_config = llm_config.get_engine_config() + + # STREAMING_LOAD_FORMATS pull weights lazily at load time; don't + # pre-download (matches the vLLM callback ctx decision). + if llm_config.engine_kwargs.get("load_format") in STREAMING_LOAD_FORMATS: + download_model = NodeModelDownloadable.NONE + else: + download_model = NodeModelDownloadable.MODEL_AND_TOKENIZER + + local_path = download_model_files( + model_id=engine_config.actual_hf_model_id, + mirror_config=engine_config.mirror_config, + download_model=download_model, + download_extra_files=True, + ) + + # download_model_files returns the local path for a mirror, else the id. + # For a local path or plain HF id (no mirror), still prefer an existing + # on-disk snapshot (mirror or HF cache). + if not (local_path and local_path != engine_config.actual_hf_model_id): + local_path = get_model_location_on_disk(engine_config.actual_hf_model_id) + + # Write the resolved path back onto the cached engine config so later + # readers of actual_hf_model_id see the on-disk location, not the + # pre-download id (mirrors the vLLM path in vllm_engine.py). + if local_path and local_path != engine_config.actual_hf_model_id: + engine_config.hf_model_id = local_path + + return local_path + @staticmethod def _build_sampling_params(request: Any) -> dict[str, Any]: sampling_params: dict[str, Any] = {} @@ -278,6 +425,23 @@ async def check_health(self) -> None: # this integration does not run. Keep the protocol hook as a no-op. return + def routing_stats(self) -> dict: + # Required by the EngineProtocol ABC. SGLang has no KV-events routing + # integration (unlike VLLMEngine), so there is nothing to surface -- + # LLMServer.record_routing_stats forwards this empty dict to Serve. + return {} + + async def build_asgi_app(self) -> Any: + """Engine-protocol entry point; forwards to the Serve deployment hook. + + ``SGLangServer`` fills two roles. Non-PD deploys it directly as + ``server_cls``, where Ray Serve calls ``__serve_build_asgi_app__`` on + the replica. P/D instead uses it as ``LLMServer``'s engine, and + ``LLMServer.__serve_build_asgi_app__`` calls ``engine.build_asgi_app()`` + -- the engine-protocol name. Both roles build the same app. + """ + return await self.__serve_build_asgi_app__() + async def __serve_build_asgi_app__(self) -> Any: """Return SGLang's native OpenAI ASGI app for Ray Serve direct streaming. @@ -336,6 +500,20 @@ def _build_generate_kwargs( sampling_params = self._build_sampling_params(request) if sampling_params: generate_kwargs["sampling_params"] = sampling_params + + # bootstrap_* are stamped by the SGLang PD connector (pd_connector.py): + # on the decode request before the local engine call, and on the prefill + # request before it is forwarded to the prefill server. + bootstrap_room = getattr(request, "bootstrap_room", None) + if bootstrap_room is not None: + generate_kwargs["bootstrap_room"] = bootstrap_room + bootstrap_host = getattr(request, "bootstrap_host", None) + if bootstrap_host is not None: + generate_kwargs["bootstrap_host"] = bootstrap_host + bootstrap_port = getattr(request, "bootstrap_port", None) + if bootstrap_port is not None: + generate_kwargs["bootstrap_port"] = bootstrap_port + return generate_kwargs async def _generate_raw( @@ -578,6 +756,33 @@ async def completions( yield resp + async def resolve_lora(self, lora_model: DiskMultiplexConfig) -> None: + """Not supported: SGLang multi-LoRA is not wired up in Ray Serve LLM. + + ``LLMServer`` only calls this for a multiplexed request, which requires + ``lora_config``; SGLang deployments do not set it, so this is + unreachable in practice. Declared because ``LLMEngine`` requires it. + """ + raise NotImplementedError( + "LoRA multiplexing is not supported by the SGLang engine." + ) + + async def transcriptions( + self, + request: "TranscriptionRequest", + raw_request_info: Optional[RawRequestInfo] = None, + ) -> AsyncGenerator[Union[str, "TranscriptionResponse", "ErrorResponse"], None]: + """Not supported: SGLang has no audio transcription endpoint here. + + Declared because ``LLMEngine`` requires it; the ``yield`` keeps this an + async generator so callers fail on iteration, matching the protocol's + return type. + """ + raise NotImplementedError( + "Transcriptions are not supported by the SGLang engine." + ) + yield # pragma: no cover - unreachable, marks this an async generator + async def embeddings( self, request: EmbeddingRequest, diff --git a/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/base.py b/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/base.py index bfd0294d53af..4a7ca47099af 100644 --- a/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/base.py +++ b/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/base.py @@ -1,28 +1,39 @@ -import abc -import random -import string -from typing import TYPE_CHECKING, Any, Dict, Optional, Union +"""vLLM connector base + the default (vLLM-shaped) P/D protocol policy. -from ray import serve +The engine-neutral connector base lives in +``ray.llm._internal.serve.engines.common.kv_transfer.base`` and knows nothing +about any specific engine. This module adds the vLLM-specific layer on top: + + * ``VLLMConnectorBackend`` — adds the ``kv_transfer_config`` property that vLLM + connectors need (the neutral base does not carry it). + * ``DefaultPDProtocolMixin`` / ``DefaultConnectorBackend`` — the standard + (no-peer-binding, sequential) P/D policy, expressed with vLLM's + ``kv_transfer_params``. vLLM connectors (nixl, lmcache) inherit it; a + non-vLLM engine implements its own request shaping instead. + +``BaseConnectorBackend`` and ``clamp_request_to_single_token`` are re-exported +from the neutral base so existing imports of them from this path keep working. +""" + +from typing import TYPE_CHECKING, Any, Dict, Optional + +from ray.llm._internal.serve.engines.common.kv_transfer.base import ( # noqa: F401 + BaseConnectorBackend, + clamp_request_to_single_token, +) if TYPE_CHECKING: - from ray.llm._internal.serve.core.configs.llm_config import LLMConfig - from ray.llm._internal.serve.core.configs.openai_api_models import ( - ChatCompletionRequest, - CompletionRequest, + from ray.llm._internal.serve.engines.common.kv_transfer.base import ( # noqa: F401 + RequestType, ) - # The two OpenAI request models the P/D orchestrator shapes. Defined under - # TYPE_CHECKING (and used as a string annotation) to avoid an import cycle - # between this module and the config/openai-models modules. - RequestType = Union[ChatCompletionRequest, CompletionRequest] - def base_prefill_kv_transfer_params() -> Dict[str, Any]: """The ``kv_transfer_params`` common to a prefill (producer) request. Tells the prefill engine to produce KV for a remote decode. Connectors layer - their own keys (e.g. a transfer id, DP/TP routing) on top of these. + their own keys (e.g. a transfer id, DP/TP routing) on top of these. This is a + vLLM concept (``kv_transfer_params``), hence it lives on the vLLM side. """ return { "do_remote_decode": True, @@ -32,60 +43,13 @@ def base_prefill_kv_transfer_params() -> Dict[str, Any]: } -def clamp_request_to_single_token(request: "RequestType") -> None: - """Clamp a prefill request to a single, non-streaming token (in place).""" - request.max_tokens = 1 - if hasattr(request, "max_completion_tokens"): - request.max_completion_tokens = 1 - request.stream = False - if hasattr(request, "stream_options"): - request.stream_options = None - +class VLLMConnectorBackend(BaseConnectorBackend): + """Connector base for vLLM connectors that need ``kv_transfer_config``. -class BaseConnectorBackend(abc.ABC): - # ---- P/D coordination protocol ---- - # - # These class attributes and methods let the P/D orchestrator - # (``PDOrchestratorMixin``) delegate request shaping, peer addressing, and - # handoff discipline to the connector. They are connector-agnostic: a - # connector picks a quadrant of (``requires_peer_binding``, - # ``concurrent_handoff``) and implements ``prepare_prefill_request`` / - # ``prepare_decode_request`` accordingly. - # - # ``requires_peer_binding``: - # * False -> the orchestrator dispatches prefill via the standard handle - # path; the peer (if any) is resolved post-hoc from the prefill response. - # * True -> the orchestrator selects the prefill replica first - # (``choose_replica``) and passes its ``replica_metadata`` to the backend - # as ``peer`` (pre-dispatch addressing). - # - # ``concurrent_handoff``: - # * False -> prefill runs to its first chunk before local decode starts - # (sequential handoff). - # * True -> prefill dispatch and local decode run concurrently. - # - # The two flags are independent; the known combos: - # * (False, False) — e.g. NixlConnector / LMCacheConnectorV1: decode learns - # everything it needs (remote engine id / block ids) from the prefill - # response, so it must wait for that response (sequential). - # * (True, True) — e.g. MoRIIO WRITE mode: the peer address is bound - # into the request id before dispatch and prefill *pushes* KV to decode, - # so decode needs nothing from prefill's response and can start - # immediately (concurrent). Concurrent handoff is only possible because - # peer binding happens up front. - # * (True, False) — e.g. MoRIIO READ mode: the peer is bound up front, - # but decode *pulls* KV using block ids returned in prefill's response, - # so it still waits for prefill to finish (sequential). - requires_peer_binding: bool = False - concurrent_handoff: bool = False - - def __init__(self, llm_config: "LLMConfig"): - """Base class for connector backends. - - Args: - llm_config: The llm configuration for this engine - """ - self.llm_config = llm_config + The neutral ``BaseConnectorBackend`` is engine-agnostic and does not carry + ``kv_transfer_config`` (a vLLM concept). vLLM connectors inherit this + subclass instead so they keep the property. + """ @property def kv_transfer_config(self) -> Dict[str, Any]: @@ -96,118 +60,17 @@ def kv_transfer_config(self) -> Dict[str, Any]: ), "In Connector backend, kv_transfer_config is not set" return kv_transfer_config - def _get_unique_suffix(self, len: int = 6) -> str: - """Generates unique alphanumeric suffix. - - Args: - len: Length of the suffix to generate. - Returns: - A unique alphanumeric suffix string of specified length. - """ - return "".join(random.choices(string.ascii_letters + string.digits, k=len)) - - def _compute_port_offset(self) -> int: - """Compute a deterministic port offset for this replica. - - Uses data_parallel_rank if DP case, otherwise falls back to - the replica rank assigned by Ray Serve (TP/PP case). - - For TP/PP cases, multiply by num_devices (tp × pp) to reserve - sufficient port space, since each worker needs a unique port. - Each TP worker adds its tp_rank (0, 1, ..., tp_size-1) to the - base port at bind time, and PP stages also need separate ports. - - Returns: - Non-negative integer offset to add to a base port. - """ - # Prefer explicit DP rank when available - dp_rank = self.llm_config.engine_kwargs.get("data_parallel_rank") - if isinstance(dp_rank, int) and dp_rank >= 0: - # vLLM already accounts for TP spacing in DP offset calculation - # (data_parallel_rank × tp_size), don't multiply here - return dp_rank - - # NOTE (jeffreywang): A missing replica context must fail loudly, not - # silently return a 0 offset that collides colocated replicas on the - # same NIXL side-channel port. get_replica_context() raises RayServeException - # outside a replica. - rc = serve.get_replica_context() - engine_config = self.llm_config.get_engine_config() - num_devices = engine_config.num_devices - return rc.rank.rank * num_devices - - @abc.abstractmethod - def prepare_prefill_request( - self, *, request: "RequestType", peer: Optional[Dict[str, Any]] - ) -> "RequestType": - """Shape the request sent to the remote prefill engine. - - Args: - request: The incoming chat/completion request. - peer: The selected prefill replica's ``replica_metadata`` dict when - the connector opted into pre-dispatch peer binding - (``requires_peer_binding=True``), else None. - - Returns: - A new request object to dispatch to the prefill engine. - """ - ... - - @abc.abstractmethod - def prepare_decode_request( - self, - *, - request: "RequestType", - peer: Optional[Dict[str, Any]], - prefill_response: Optional[Any], - ) -> "RequestType": - """Shape the request run on the local decode engine. - - Args: - request: The incoming chat/completion request. - peer: The selected prefill replica's ``replica_metadata`` dict when - the connector opted into pre-dispatch peer binding, else None. - prefill_response: The captured prefill response chunk whose - ``kv_transfer_params`` may be forwarded, or None when no chunk is - captured before decode starts (concurrent-handoff mode). - - Returns: - A new request object to run on the local decode engine. - """ - ... - - def setup(self) -> None: - """Setup the connector backend. - - This method is called to setup the connector backend. - """ - pass - - def replica_metadata(self) -> Dict[str, Any]: - """Static per-replica coordination data published to the orchestrator. - - Surfaced via the replica-metadata hook on ``ReplicaSelection`` so that a - connector opting into ``requires_peer_binding`` can address the selected - prefill peer. The default backend publishes nothing; connectors that need - to advertise an address (e.g. MoRIIO's zmq endpoint) override this. - - Returns: - A JSON-serializable dict of per-replica metadata (empty by default). - """ - return {} - class DefaultPDProtocolMixin: """The default P/D protocol policy: no peer binding, sequential handoff. - Implements ``prepare_prefill_request`` / ``prepare_decode_request`` for - connectors that follow the standard policy: the prefill engine is told to - produce KV for a remote decode (clamped to a single non-streaming token), - and the decode engine forwards the ``kv_transfer_params`` that the prefill - engine returned on its first response chunk. + Prefill is told to produce KV for a remote decode and clamped to a single + non-streaming token; decode forwards the ``kv_transfer_params`` that prefill + returned on its first response chunk. - Mix this in *before* ``BaseConnectorBackend`` in a backend's bases so its - concrete methods satisfy the abstract methods. + Mix in *before* ``VLLMConnectorBackend`` -- the reverse order leaves the + ABC's abstract ``prepare_*`` ahead in the Method Resolution Order and the class won't + instantiate. """ def prepare_prefill_request( @@ -252,7 +115,7 @@ def prepare_decode_request( return decode_request -class DefaultConnectorBackend(DefaultPDProtocolMixin, BaseConnectorBackend): +class DefaultConnectorBackend(DefaultPDProtocolMixin, VLLMConnectorBackend): """Concrete connector backend using the default P/D protocol policy. Used as the factory fallback for connectors that are not registered with a diff --git a/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/lmcache.py b/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/lmcache.py index b451e24b0d64..ef30a0579f9f 100644 --- a/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/lmcache.py +++ b/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/lmcache.py @@ -1,6 +1,6 @@ from ray.llm._internal.serve.engines.vllm.kv_transfer.base import ( - BaseConnectorBackend, DefaultPDProtocolMixin, + VLLMConnectorBackend, ) from ray.llm._internal.serve.observability.logging import get_logger @@ -16,7 +16,7 @@ def _check_lmcache_installed(): ) -class LMCacheConnectorV1Backend(DefaultPDProtocolMixin, BaseConnectorBackend): +class LMCacheConnectorV1Backend(DefaultPDProtocolMixin, VLLMConnectorBackend): KV_CONNECTOR_EXTRA_CONFIG_FIELD_NAME = "kv_connector_extra_config" LMCACHE_RPC_PORT_FIELD_NAME = "lmcache_rpc_port" diff --git a/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/moriio.py b/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/moriio.py index 5c9fd62f1c2b..e1cc41856e32 100644 --- a/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/moriio.py +++ b/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/moriio.py @@ -35,7 +35,7 @@ import ray from ray.llm._internal.serve.engines.vllm.kv_transfer.base import ( - BaseConnectorBackend, + VLLMConnectorBackend, base_prefill_kv_transfer_params, clamp_request_to_single_token, ) @@ -113,7 +113,7 @@ def _read_mode_enabled(extra_config: Dict[str, Any]) -> bool: ) -class MoRIIOConnectorBackend(BaseConnectorBackend): +class MoRIIOConnectorBackend(VLLMConnectorBackend): """Set up MoRIIO ports/extra_config and implement the PD connector protocol.""" # The advertised zmq address ("host:IP,handshake:PORT,notify:PORT"), diff --git a/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/multi_connector.py b/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/multi_connector.py index a60985a49ea7..c114fbbc0f81 100644 --- a/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/multi_connector.py +++ b/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/multi_connector.py @@ -1,18 +1,19 @@ import copy from typing import TYPE_CHECKING, List +from ray.llm._internal.serve.engines.common.kv_transfer.factory import ( + KVConnectorBackendFactory, +) from ray.llm._internal.serve.engines.vllm.kv_transfer.base import ( BaseConnectorBackend, -) -from ray.llm._internal.serve.engines.vllm.kv_transfer.factory import ( - KVConnectorBackendFactory, + VLLMConnectorBackend, ) if TYPE_CHECKING: from ray.llm._internal.serve.core.configs.llm_config import LLMConfig -class MultiConnectorBackend(BaseConnectorBackend): +class MultiConnectorBackend(VLLMConnectorBackend): """Wraps multiple sub-connectors. The P/D protocol (``prepare_prefill_request`` / ``prepare_decode_request`` and diff --git a/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/nixl.py b/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/nixl.py index d45448e4d5b5..653064e14869 100644 --- a/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/nixl.py +++ b/python/ray/llm/_internal/serve/engines/vllm/kv_transfer/nixl.py @@ -2,12 +2,12 @@ import ray from ray.llm._internal.serve.engines.vllm.kv_transfer.base import ( - BaseConnectorBackend, DefaultPDProtocolMixin, + VLLMConnectorBackend, ) -class NixlConnectorBackend(DefaultPDProtocolMixin, BaseConnectorBackend): +class NixlConnectorBackend(DefaultPDProtocolMixin, VLLMConnectorBackend): def _set_side_channel_port(self): from vllm import envs as vllm_envs diff --git a/python/ray/llm/_internal/serve/serving_patterns/prefill_decode/builder.py b/python/ray/llm/_internal/serve/serving_patterns/prefill_decode/builder.py index af7a8ffa8554..84673ad95473 100644 --- a/python/ray/llm/_internal/serve/serving_patterns/prefill_decode/builder.py +++ b/python/ray/llm/_internal/serve/serving_patterns/prefill_decode/builder.py @@ -79,6 +79,24 @@ def _validate_ingress_cls_config( return IngressClsConfig.model_validate(value) return value + @model_validator(mode="after") + def _unalias_configs(self): + """Give the two sides independent configs before anything mutates them. + + Passing one ``LLMConfig`` instance for both fields is natural (the sides + are identical apart from their role) and the field validator hands the + same object back for both. The role-specific validators below then mutate + ``decode_config`` in place -- disaggregation_mode, the shifted bootstrap + port bases -- so a shared instance would leave BOTH sides in decode mode + on the same ports, breaking the KV handoff with no error. + + Runs first: every mutating validator is defined after this one, and + ``mode="after"`` validators run in definition order. + """ + if self.prefill_config is self.decode_config: + self.decode_config = self.decode_config.model_copy(deep=True) + return self + @model_validator(mode="after") def _validate_model_ids(self): """Validate that prefill and decode configs use the same model ID.""" @@ -87,15 +105,109 @@ def _validate_model_ids(self): return self @model_validator(mode="after") - def _validate_kv_transfer_config(self): - """Validate that kv_transfer_config is set for both prefill and decode configs.""" + def _validate_same_engine(self): + """Prefill and decode must use the same ``llm_engine``. + + The decode orchestrator drives both sides through one connector protocol; + a mixed pair (e.g. prefill vLLM + decode SGLang) passes the per-side + transfer checks but has no compatible P/D wiring and fails at runtime. + Reject it up front. + """ + if self.prefill_config.llm_engine != self.decode_config.llm_engine: + raise ValueError( + "P/D prefill and decode must use the same llm_engine " + f"(got prefill={self.prefill_config.llm_engine!r}, " + f"decode={self.decode_config.llm_engine!r})." + ) + return self + + @model_validator(mode="after") + def _validate_transfer_config(self): + """Each engine needs its own PD transfer config. + + vLLM requires ``kv_transfer_config``; SGLang requires + ``disaggregation_transfer_backend``. + """ for config in [self.prefill_config, self.decode_config]: - if config.engine_kwargs.get("kv_transfer_config") is None: + if config.llm_engine == "SGLang": + if not config.engine_kwargs.get("disaggregation_transfer_backend"): + raise ValueError( + "disaggregation_transfer_backend is required for SGLang " + "P/D disaggregation" + ) + elif config.engine_kwargs.get("kv_transfer_config") is None: raise ValueError( "kv_transfer_config is required for P/D disaggregation" ) return self + @model_validator(mode="after") + def _reject_sglang_data_parallel(self): + """SGLang P/D with data_parallel_size>1 is not supported yet. + + DP P/D uses DPPD{Prefill,Decode}Server, whose gang scheduling comes from + DPServer.get_deployment_options / __init__ — both read engine-config + fields (accelerator, placement_bundles) the minimal SGLangEngineConfig + does not carry. Rather than silently drop gang scheduling, fail fast. + Tracked as a follow-up (TODO: link a ticket here). + """ + for label, config in ( + ("prefill_config", self.prefill_config), + ("decode_config", self.decode_config), + ): + if config.llm_engine != "SGLang": + continue + # SGLang's own engine kwarg is dp_size, not vLLM's + # data_parallel_size; check both so neither name bypasses the guard. + for key in ("data_parallel_size", "dp_size"): + dp_size = config.engine_kwargs.get(key, 1) + if isinstance(dp_size, int) and dp_size > 1: + raise NotImplementedError( + f"SGLang P/D disaggregation does not support " + f"{key}>1 yet (got {dp_size} on {label}). " + f"Use {key}=1." + ) + return self + + @model_validator(mode="after") + def _set_sglang_disaggregation_mode(self): + """Set disaggregation_mode from the role, so users never set it by hand. + + Assigned, not defaulted: the field is owned by the position in the graph + (prefill_config vs decode_config), so a stale or copied-across value must + be corrected. Reusing one config dict for both sides is the easy mistake, + and leaving decode in "prefill" mode breaks the KV handoff silently. + """ + if self.prefill_config.llm_engine == "SGLang": + self.prefill_config.engine_kwargs["disaggregation_mode"] = "prefill" + if self.decode_config.llm_engine == "SGLang": + self.decode_config.engine_kwargs["disaggregation_mode"] = "decode" + return self + + @model_validator(mode="after") + def _default_decode_sglang_bootstrap_port_base(self): + """Shift decode's SGLang bootstrap port base off prefill's default so a + colocated P+D pair doesn't collide (mirrors the NIXL/MoRIIO shifts). + + The decode engine runs a (mostly unused) bootstrap server too; pinning a + distinct port avoids a same-node bind clash on 8998. + """ + if self.decode_config.llm_engine != "SGLang": + return self + from ray.llm._internal.serve.engines.sglang.kv_transfer.pd_connector import ( + BOOTSTRAP_PORT_BASE_KEY, + DEFAULT_BOOTSTRAP_PORT_BASE, + ) + + # Shift the decode BASE (not the final port): the connector adds a + # per-replica offset on top, so colocated decode replicas still get + # distinct ports. The +1000 stride is well above any realistic + # tp_size*pp_size offset. Mirrors _default_decode_moriio_port_base. + self.decode_config.experimental_configs.setdefault( + BOOTSTRAP_PORT_BASE_KEY, DEFAULT_BOOTSTRAP_PORT_BASE + 1000 + ) + return self + @model_validator(mode="after") def _default_decode_nixl_port_base(self): """Shift decode's NIXL base off prefill's default (20000) so colocated replicas don't collide.""" diff --git a/python/ray/llm/_internal/serve/serving_patterns/prefill_decode/pd_server.py b/python/ray/llm/_internal/serve/serving_patterns/prefill_decode/pd_server.py index 45c25563bc71..330326e77088 100644 --- a/python/ray/llm/_internal/serve/serving_patterns/prefill_decode/pd_server.py +++ b/python/ray/llm/_internal/serve/serving_patterns/prefill_decode/pd_server.py @@ -24,6 +24,7 @@ ChatCompletionResponse, CompletionRequest, CompletionResponse, + ErrorInfo, ErrorResponse, ) from ray.llm._internal.serve.core.ingress.utils import ( @@ -34,7 +35,7 @@ ) from ray.llm._internal.serve.core.protocol import RawRequestInfo from ray.llm._internal.serve.core.server.llm_server import LLMServer -from ray.llm._internal.serve.engines.vllm.kv_transfer.base import BaseConnectorBackend +from ray.llm._internal.serve.engines.common.kv_transfer.base import BaseConnectorBackend from ray.llm._internal.serve.serving_patterns.data_parallel.dp_server import DPServer from ray.llm._internal.serve.utils.broadcast import broadcast from ray.serve._private.http_util import session_id_from_headers @@ -43,6 +44,7 @@ logger = logging.getLogger(__name__) + RequestType = Union[ChatCompletionRequest, CompletionRequest] _PREWARM_PROMPT = " x" @@ -78,6 +80,7 @@ async def _pd_http_response(gen) -> Response: Returns a JSON response when the first chunk is an error or a complete (non-streaming) response, otherwise an SSE stream. Uses the same response helpers as ``OpenAiIngress`` so the wire format matches the standard path. + """ first, gen = await _peek_at_generator(gen) if isinstance(first, list): @@ -335,7 +338,7 @@ async def _concurrent_decode( raw_request_info: Optional[RawRequestInfo], *, cancel_on_failure: bool = True, - ): + ) -> AsyncGenerator[Any, None]: """Run local decode while a remote prefill drains concurrently. While prefill is in flight, each decode chunk is raced against the @@ -356,47 +359,101 @@ async def _concurrent_decode( completion accounting fires when the response completes, so the stream must be drained to exhaustion, never abandoned. Prefill is clamped to a single token, so draining is bounded either way. + + Yields: + Any: Decode response chunks, or a single ``ErrorResponse`` if the + remote prefill failed, was cancelled, or returned an error. """ prefill_task = asyncio.create_task(_drain_prefill(prefill_resp)) completed = False local_gen = None next_fut = None + + def _prefill_error() -> Optional[ErrorResponse]: + """An error to surface to the client, if prefill finished badly. + + Covers all three "prefill is done and it's bad" cases -- returned + an ``ErrorResponse``, raised, or got cancelled -- so a dead + prefill always aborts decode instead of leaving it to hang on KV + that will never arrive. Never raises itself: ``.exception()`` and + ``.cancelled()`` just report state, and the real exception object + is still awaited (and logged) by the ``finally`` block below. + """ + if not prefill_task.done(): + return None + if prefill_task.cancelled(): + return ErrorResponse( + error=ErrorInfo( + message="Remote prefill was cancelled", + type="internal_error", + code=500, + ) + ) + exc = prefill_task.exception() + if exc is not None: + return ErrorResponse( + error=ErrorInfo( + message=f"Remote prefill failed: {exc}", + type="internal_error", + code=500, + ) + ) + result = prefill_task.result() + return result if isinstance(result, ErrorResponse) else None + try: local_gen = await getattr(super(), method)(decode_request, raw_request_info) gen = local_gen.__aiter__() - while True: - # Surface a failed prefill as soon as it is observed. - if prefill_task.done() and isinstance( - prefill_task.result(), ErrorResponse - ): - err = prefill_task.result() - logger.error("Remote prefill returned error: %s", err) - yield err - return + + # Phase 1: race against prefill while it's still in flight. + while not prefill_task.done(): if next_fut is None: next_fut = asyncio.ensure_future(gen.__anext__()) - # Race the next decode chunk against the in-flight prefill; - # once prefill has completed (successfully), just stream. - awaitables = {next_fut} - if not prefill_task.done(): - awaitables.add(prefill_task) - done, _ = await asyncio.wait( - awaitables, return_when=asyncio.FIRST_COMPLETED + + await asyncio.wait( + {next_fut, prefill_task}, return_when=asyncio.FIRST_COMPLETED ) - if next_fut in done: + + if next_fut.done(): try: chunk = next_fut.result() except StopAsyncIteration: + completed = True break - next_fut = None + finally: + if next_fut.done(): + next_fut = None yield chunk - # else: prefill finished first; loop back to inspect it. - completed = True + + err = _prefill_error() + if err is not None: + logger.error("Remote prefill returned error: %s", err) + yield err + return + + # Phase 2: prefill is done, nothing left to race -- drain decode + # directly through the same iterator Phase 1 was pulling from. + if not completed: + if next_fut is not None: + try: + chunk = next_fut.result() if next_fut.done() else await next_fut + except StopAsyncIteration: + completed = True + else: + yield chunk + finally: + next_fut = None + if not completed: + async for chunk in gen: + yield chunk + completed = True + finally: if next_fut is not None and not next_fut.done(): next_fut.cancel() with contextlib.suppress(BaseException): await next_fut + if not completed: # Abort the local decode request if we bailed early. if local_gen is not None: @@ -582,12 +639,78 @@ async def _drain_prefill(prefill_resp) -> Optional[ErrorResponse]: return None +# --------------------------------------------------------------------------- +# Engine class selection (per llm_engine) +# --------------------------------------------------------------------------- + + +def _resolve_pd_engine_class(llm_config: "LLMConfig"): + """Engine class for a PD server, chosen by ``llm_engine``. + + SGLang must use ``SGLangServer``; everything else falls through to the base + LLMServer default (VLLMEngine). The generic PD servers don't set + ``_default_engine_cls``, so without this they would always pick VLLMEngine. + """ + if llm_config.llm_engine == "SGLang": + from ray.llm._internal.serve.engines.sglang.sglang_engine import SGLangServer + + return SGLangServer + from ray.llm._internal.serve.engines.vllm.vllm_engine import VLLMEngine + + return VLLMEngine + + +class _PDEngineSelectionMixin: + """Make a PD server pick its engine class from ``llm_config.llm_engine``. + + Preserves the base ``RAYLLM_VLLM_ENGINE_CLS`` escape hatch (used by tests + that patch the engine class) and otherwise dispatches on the engine. + """ + + @classmethod + def _resolve_engine_class(cls, llm_config: "LLMConfig"): + return _resolve_pd_engine_class(llm_config) + + def _get_default_engine_class(self): + import os + + from ray._common.utils import import_attr + from ray.llm._internal.serve.constants import RAYLLM_VLLM_ENGINE_CLS_ENV + + engine_cls_path = os.environ.get(RAYLLM_VLLM_ENGINE_CLS_ENV) + if engine_cls_path: + return import_attr(engine_cls_path) + return self._resolve_engine_class(self._llm_config) + + @classmethod + def get_deployment_options(cls, llm_config: "LLMConfig"): + """Deployment options for the PD server, per engine. + + The base ``LLMServer.get_deployment_options`` reads + ``engine_config.accelerator`` / ``placement_strategy`` — fields the + minimal ``SGLangEngineConfig`` does not carry. SGLang builds its bundles + from ``SGLangServer.get_deployment_options`` instead. + + The SGLang branch is unconditional, which would drop DPServer's gang + scheduling for data_parallel_size>1 -- but the builder rejects that + combination up front (``_reject_sglang_data_parallel``), so it can't + reach here. + """ + if llm_config.llm_engine == "SGLang": + from ray.llm._internal.serve.engines.sglang.sglang_engine import ( + SGLangServer, + ) + + return SGLangServer.get_deployment_options(llm_config) + return super().get_deployment_options(llm_config) + + # --------------------------------------------------------------------------- # PDPrefillServer # --------------------------------------------------------------------------- -class PDPrefillServer(LLMServer): +class PDPrefillServer(_PDEngineSelectionMixin, LLMServer): """Prefill-side LLM server for P/D disaggregation. This is a standard LLMServer with an additional ``prewarm_prefill`` @@ -634,7 +757,7 @@ async def prewarm_prefill( # --------------------------------------------------------------------------- -class PDDecodeServer(PDOrchestratorMixin, LLMServer): +class PDDecodeServer(_PDEngineSelectionMixin, PDOrchestratorMixin, LLMServer): """Decode-side LLM server that orchestrates remote prefill. This deployment owns a real engine (decode config) and holds a handle diff --git a/python/ray/llm/tests/serve/cpu/configs/test_llm_config_num_devices.py b/python/ray/llm/tests/serve/cpu/configs/test_llm_config_num_devices.py new file mode 100644 index 000000000000..20878e56adce --- /dev/null +++ b/python/ray/llm/tests/serve/cpu/configs/test_llm_config_num_devices.py @@ -0,0 +1,26 @@ +import sys + +import pytest + +from ray.serve.llm import LLMConfig + + +def test_num_devices_vllm_default(): + cfg = LLMConfig( + model_loading_config=dict(model_source="facebook/opt-125m"), + llm_engine="vLLM", + engine_kwargs=dict(tensor_parallel_size=2, pipeline_parallel_size=2), + ) + assert cfg.num_devices == 4 + + +def test_num_devices_defaults_to_one(): + cfg = LLMConfig( + model_loading_config=dict(model_source="facebook/opt-125m"), + llm_engine="vLLM", + ) + assert cfg.num_devices == 1 + + +if __name__ == "__main__": + sys.exit(pytest.main(["-v", "-x", __file__])) diff --git a/python/ray/llm/tests/serve/cpu/deployments/llm/vllm/kv_transfer_backends/test_factory.py b/python/ray/llm/tests/serve/cpu/deployments/llm/vllm/kv_transfer_backends/test_factory.py index 7d5eef97635f..69c847a411e2 100644 --- a/python/ray/llm/tests/serve/cpu/deployments/llm/vllm/kv_transfer_backends/test_factory.py +++ b/python/ray/llm/tests/serve/cpu/deployments/llm/vllm/kv_transfer_backends/test_factory.py @@ -5,14 +5,14 @@ import pytest from ray import serve +from ray.llm._internal.serve.engines.common.kv_transfer.factory import ( + KVConnectorBackendFactory, +) from ray.llm._internal.serve.engines.vllm.kv_transfer.base import ( BaseConnectorBackend, DefaultConnectorBackend, DefaultPDProtocolMixin, ) -from ray.llm._internal.serve.engines.vllm.kv_transfer.factory import ( - KVConnectorBackendFactory, -) from ray.serve.llm import LLMConfig diff --git a/python/ray/llm/tests/serve/cpu/deployments/llm/vllm/kv_transfer_backends/test_moriio_connector.py b/python/ray/llm/tests/serve/cpu/deployments/llm/vllm/kv_transfer_backends/test_moriio_connector.py index 3aed8b567641..6ff893eea376 100644 --- a/python/ray/llm/tests/serve/cpu/deployments/llm/vllm/kv_transfer_backends/test_moriio_connector.py +++ b/python/ray/llm/tests/serve/cpu/deployments/llm/vllm/kv_transfer_backends/test_moriio_connector.py @@ -5,12 +5,12 @@ import pytest +from ray.llm._internal.serve.engines.common.kv_transfer.factory import ( + KVConnectorBackendFactory, +) from ray.llm._internal.serve.engines.vllm.kv_transfer.base import ( BaseConnectorBackend, ) -from ray.llm._internal.serve.engines.vllm.kv_transfer.factory import ( - KVConnectorBackendFactory, -) from ray.llm._internal.serve.engines.vllm.kv_transfer.moriio import ( _DECODE_ZMQ_RE, _PREFILL_ZMQ_RE, diff --git a/python/ray/llm/tests/serve/cpu/deployments/llm/vllm/kv_transfer_backends/test_multi_connector.py b/python/ray/llm/tests/serve/cpu/deployments/llm/vllm/kv_transfer_backends/test_multi_connector.py index 5b5522629f44..57f92bb33ff4 100644 --- a/python/ray/llm/tests/serve/cpu/deployments/llm/vllm/kv_transfer_backends/test_multi_connector.py +++ b/python/ray/llm/tests/serve/cpu/deployments/llm/vllm/kv_transfer_backends/test_multi_connector.py @@ -3,12 +3,12 @@ import pytest +from ray.llm._internal.serve.engines.common.kv_transfer.factory import ( + KVConnectorBackendFactory, +) from ray.llm._internal.serve.engines.vllm.kv_transfer.base import ( BaseConnectorBackend, ) -from ray.llm._internal.serve.engines.vllm.kv_transfer.factory import ( - KVConnectorBackendFactory, -) from ray.llm._internal.serve.engines.vllm.kv_transfer.multi_connector import ( MultiConnectorBackend, ) diff --git a/python/ray/serve/_private/constants.py b/python/ray/serve/_private/constants.py index 2deb89fc30f9..7d78a97c97cb 100644 --- a/python/ray/serve/_private/constants.py +++ b/python/ray/serve/_private/constants.py @@ -647,14 +647,12 @@ ) # Run sync methods defined in the replica in a thread pool by default. -RAY_SERVE_RUN_SYNC_IN_THREADPOOL = get_env_bool("RAY_SERVE_RUN_SYNC_IN_THREADPOOL", "0") +RAY_SERVE_RUN_SYNC_IN_THREADPOOL = get_env_bool("RAY_SERVE_RUN_SYNC_IN_THREADPOOL", "1") RAY_SERVE_RUN_SYNC_IN_THREADPOOL_WARNING = ( - "Calling sync method '{method_name}' directly on the " - "asyncio loop. In a future version, sync methods will be run in a " - "threadpool by default. Ensure your sync methods are thread safe " - "or keep the existing behavior by making them `async def`. Opt " - "into the new behavior by setting " + "Calling sync method '{method_name}' directly on the asyncio loop because " + "RAY_SERVE_RUN_SYNC_IN_THREADPOOL=0. This can block other requests. Make " + "the method `async def` or restore threadpool dispatch by setting " "RAY_SERVE_RUN_SYNC_IN_THREADPOOL=1." ) diff --git a/python/ray/serve/tests/BUILD.bazel b/python/ray/serve/tests/BUILD.bazel index b118ed36e596..eba3aae51d86 100644 --- a/python/ray/serve/tests/BUILD.bazel +++ b/python/ray/serve/tests/BUILD.bazel @@ -812,8 +812,7 @@ py_test_module_list( ], ) -# Test currently off-by-default behavior to run replica sync methods in a threadpool. -# TODO(edoakes): remove this once the FF is flipped on by default. +# Test the default behavior that runs replica sync methods in a threadpool. py_test_module_list( size = "medium", env = {"RAY_SERVE_RUN_SYNC_IN_THREADPOOL": "1"}, diff --git a/release/llm_tests/serve/test_llm_serve_sglang_pd.py b/release/llm_tests/serve/test_llm_serve_sglang_pd.py new file mode 100644 index 000000000000..df779713b3d7 --- /dev/null +++ b/release/llm_tests/serve/test_llm_serve_sglang_pd.py @@ -0,0 +1,380 @@ +"""Release tests for SGLang Prefill-Decode disaggregation on Ray Serve. + +SGLang PD now travels the generic decode-as-orchestrator graph +(``build_pd_openai_app`` -> ``PDDecodeServer`` -> ``PDPrefillServer``). All the +SGLang-specific coordination (bootstrap host/port/room) lives in a single +connector, ``SGLangConnectorBackend``, selected automatically from +``disaggregation_transfer_backend``. + +A two-GPU node is required for the end-to-end tests (prefill on GPU 0, decode on +GPU 1). They use the real NIXL KV transport — SGLang's "fake" transport has no +bootstrap-server class, so a prefill-mode engine cannot start under it. +""" + +import sys +import concurrent.futures +from types import SimpleNamespace +from typing import Optional + +import pytest +from openai import OpenAI + +from ray import serve +from ray._common.test_utils import wait_for_condition +from ray.llm._internal.serve.engines.sglang.kv_transfer.pd_connector import ( + BOOTSTRAP_PORT_BASE_KEY, + DEFAULT_BOOTSTRAP_PORT_BASE, + SGLangConnectorBackend, +) +from ray.llm._internal.serve.serving_patterns.prefill_decode.builder import ( + build_pd_openai_app, +) +from ray.serve._private.constants import SERVE_DEFAULT_APP_NAME +from ray.serve.llm import LLMConfig +from ray.serve.schema import ApplicationStatus + +MODEL_ID = "Qwen/Qwen2.5-0.5B-Instruct" +RAY_MODEL_ID = "qwen-0.5b-sglang-pd" + + +def _app_is_running(): + try: + return ( + serve.status().applications[SERVE_DEFAULT_APP_NAME].status + == ApplicationStatus.RUNNING + ) + except (KeyError, AttributeError): + return False + + +def _sglang_config(base_gpu_id: int) -> dict: + return LLMConfig( + model_loading_config={ + "model_id": RAY_MODEL_ID, + "model_source": MODEL_ID, + }, + deployment_config={ + "autoscaling_config": { + "min_replicas": 1, + "max_replicas": 1, + } + }, + engine_kwargs={ + # disaggregation_mode is set automatically by the builder. + "disaggregation_transfer_backend": "nixl", + "tp_size": 1, + "mem_fraction_static": 0.4, + "base_gpu_id": base_gpu_id, + }, + llm_engine="SGLang", + ).model_dump() + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + + +@pytest.fixture(scope="module") +def sglang_pd_client(): + """Start a SGLang PD deployment using the real NIXL KV transport. + + Requires a node with at least 2 GPUs (prefill on GPU 0, decode on GPU 1). + NIXL is pre-installed in the llm-cu130 BYOD image. + """ + + app = build_pd_openai_app( + { + "prefill_config": _sglang_config(base_gpu_id=0), + "decode_config": _sglang_config(base_gpu_id=1), + } + ) + serve.run(app, blocking=False) + wait_for_condition(_app_is_running, timeout=300) + + client = OpenAI(base_url="http://localhost:8000/v1", api_key="fake-key") + yield client + + serve.shutdown() + + +# --------------------------------------------------------------------------- +# Tests — real NIXL transport (two-GPU node required) +# --------------------------------------------------------------------------- + + +def test_sglang_pd_chat(sglang_pd_client): + """Verify chat completions work end-to-end over NIXL KV transfer.""" + + resp = sglang_pd_client.chat.completions.create( + model=RAY_MODEL_ID, + messages=[{"role": "user", "content": "What is the capital of France?"}], + max_tokens=64, + temperature=0.0, + ) + assert resp.choices[0].message.content.strip() + + +def test_sglang_pd_completions(sglang_pd_client): + """Verify completions work end-to-end over NIXL KV transfer.""" + + resp = sglang_pd_client.completions.create( + model=RAY_MODEL_ID, + prompt="The capital of France is", + max_tokens=64, + temperature=0.0, + ) + assert resp.choices[0].text.strip() + + +def test_sglang_pd_streaming_chat(sglang_pd_client): + """Verify streaming chat completions produce incremental chunks.""" + + stream = sglang_pd_client.chat.completions.create( + model=RAY_MODEL_ID, + messages=[{"role": "user", "content": "Count to 5"}], + max_tokens=64, + temperature=0.0, + stream=True, + ) + + chunks = list(stream) + assert len(chunks) > 1, "Expected multiple streaming chunks" + + collected_text = "" + finish_reason = None + for chunk in chunks: + delta = chunk.choices[0].delta + if delta.content is not None: + collected_text += delta.content + if chunk.choices[0].finish_reason is not None: + finish_reason = chunk.choices[0].finish_reason + + assert collected_text.strip(), "Streaming produced no text" + assert finish_reason is not None, "Final chunk must have a finish_reason" + + +def test_sglang_pd_streaming_completions(sglang_pd_client): + """Verify streaming completions produce incremental chunks.""" + + stream = sglang_pd_client.completions.create( + model=RAY_MODEL_ID, + prompt="The capital of France is", + max_tokens=32, + temperature=0.0, + stream=True, + ) + + chunks = list(stream) + assert len(chunks) > 1, "Expected multiple streaming chunks" + + collected_text = "" + finish_reason = None + for chunk in chunks: + if chunk.choices[0].text is not None: + collected_text += chunk.choices[0].text + if chunk.choices[0].finish_reason is not None: + finish_reason = chunk.choices[0].finish_reason + + assert collected_text.strip(), "Streaming produced no text" + assert finish_reason is not None, "Final chunk must have a finish_reason" + + +def test_sglang_pd_concurrent_requests(sglang_pd_client): + """Verify multiple concurrent requests each complete successfully. + + Each request gets its own unique bootstrap_room — if rooms collide, + SGLang's bootstrap server would mix up KV caches between requests. + """ + + def send_request(i): + return sglang_pd_client.chat.completions.create( + model=RAY_MODEL_ID, + messages=[{"role": "user", "content": f"Say the number {i}"}], + max_tokens=10, + temperature=0.0, + ) + + with concurrent.futures.ThreadPoolExecutor(max_workers=4) as executor: + futures = [executor.submit(send_request, i) for i in range(4)] + results = [f.result() for f in futures] + + for resp in results: + assert resp.choices[0].message.content.strip() + + +# --------------------------------------------------------------------------- +# Unit tests — no GPU required (connector-level) +# --------------------------------------------------------------------------- + + +def _connector() -> SGLangConnectorBackend: + cfg = LLMConfig( + model_loading_config={"model_id": RAY_MODEL_ID, "model_source": MODEL_ID}, + llm_engine="SGLang", + engine_kwargs={"disaggregation_transfer_backend": "nixl"}, + ) + return SGLangConnectorBackend(cfg) + + +def _req(rid: Optional[str] = None) -> SimpleNamespace: + """A stand-in request. ``rid`` is None by default, matching real clients.""" + return SimpleNamespace(rid=rid, model_copy=lambda deep: SimpleNamespace()) + + +def test_sglang_pd_flags_on(): + backend = _connector() + assert backend.requires_peer_binding is True + assert backend.concurrent_handoff is True + + +def test_sglang_pd_bootstrap_field_injection(): + """Both prefill and decode requests carry the PREFILL bootstrap address. + + The bootstrap server runs on the prefill worker; ``peer`` is always the + prefill replica's metadata, so both sides use that address. Both + ``prepare_*`` calls run on the SAME request object (as in pd_server) and + must agree on ``bootstrap_room``. + """ + backend = _connector() + peer = {"bootstrap_host": "10.0.0.5", "bootstrap_port": 9201} + request = _req() + + prefill_req = backend.prepare_prefill_request(request=request, peer=peer) + assert prefill_req.bootstrap_host == "10.0.0.5" + assert prefill_req.bootstrap_port == 9201 + + decode_req = backend.prepare_decode_request( + request=request, peer=peer, prefill_response=None + ) + assert decode_req.bootstrap_host == "10.0.0.5" + assert decode_req.bootstrap_port == 9201 + assert prefill_req.bootstrap_room == decode_req.bootstrap_room + + +def test_sglang_pd_bootstrap_room_uniqueness(): + """Distinct requests get distinct rooms even when no client sets ``rid``. + + ``rid`` is optional and defaults to None, so this is the common path: if + rooms were derived from it, every concurrent request would share one room + and the bootstrap server could mix KV caches. + """ + backend = _connector() + peer = {"bootstrap_host": "10.0.0.5", "bootstrap_port": 9201} + rooms = { + backend.prepare_prefill_request(request=_req(), peer=peer).bootstrap_room + for _ in range(1000) + } + assert len(rooms) == 1000, "bootstrap_room values are not unique" + + +def test_sglang_pd_bootstrap_room_honors_client_rid(): + """A client-supplied ``rid`` still drives the room, and stays stable.""" + backend = _connector() + peer = {"bootstrap_host": "10.0.0.5", "bootstrap_port": 9201} + + room_a = backend.prepare_prefill_request( + request=_req("r1"), peer=peer + ).bootstrap_room + room_b = backend.prepare_prefill_request( + request=_req("r1"), peer=peer + ).bootstrap_room + room_c = backend.prepare_prefill_request( + request=_req("r2"), peer=peer + ).bootstrap_room + + assert room_a == room_b + assert room_a != room_c + + +def test_sglang_pd_replica_metadata_publishes_address(): + backend = _connector() + backend._bootstrap_host = "10.0.0.7" + backend._bootstrap_port = 8211 + assert backend.replica_metadata() == { + "bootstrap_host": "10.0.0.7", + "bootstrap_port": 8211, + } + + +def test_sglang_pd_missing_peer_address_raises(): + backend = _connector() + with pytest.raises(ValueError): + backend.prepare_prefill_request(request=_req("r1"), peer={}) + + +def test_sglang_pd_builder_sets_disaggregation_mode(): + """The builder sets disaggregation_mode and accepts SGLang configs without + kv_transfer_config.""" + from ray.llm._internal.serve.serving_patterns.prefill_decode.builder import ( + PDServingArgs, + ) + + args = PDServingArgs.model_validate( + { + "prefill_config": _sglang_config(base_gpu_id=0), + "decode_config": _sglang_config(base_gpu_id=1), + } + ) + assert args.prefill_config.engine_kwargs["disaggregation_mode"] == "prefill" + assert args.decode_config.engine_kwargs["disaggregation_mode"] == "decode" + # Decode's bootstrap BASE is shifted (per-replica offset added later in the + # connector's setup()), so colocated decode replicas get distinct ports. + assert ( + args.decode_config.experimental_configs[BOOTSTRAP_PORT_BASE_KEY] + == DEFAULT_BOOTSTRAP_PORT_BASE + 1000 + ) + + +def test_sglang_pd_rejects_data_parallel(): + """SGLang P/D with data_parallel_size>1 fails fast (unsupported).""" + import copy + + from ray.llm._internal.serve.serving_patterns.prefill_decode.builder import ( + PDServingArgs, + ) + + prefill = copy.deepcopy(_sglang_config(base_gpu_id=0)) + decode = copy.deepcopy(_sglang_config(base_gpu_id=1)) + decode["engine_kwargs"]["data_parallel_size"] = 2 + + with pytest.raises(NotImplementedError, match="data_parallel_size"): + PDServingArgs.model_validate( + {"prefill_config": prefill, "decode_config": decode} + ) + + +def test_sglang_pd_rejects_mixed_engines(): + """Prefill and decode must use the same llm_engine.""" + import copy + + from ray.llm._internal.serve.serving_patterns.prefill_decode.builder import ( + PDServingArgs, + ) + + sglang = _sglang_config(base_gpu_id=1) + vllm = copy.deepcopy(sglang) + vllm["llm_engine"] = "vLLM" + vllm["engine_kwargs"] = {"kv_transfer_config": {"kv_connector": "NixlConnector"}} + + with pytest.raises(ValueError, match="same llm_engine"): + PDServingArgs.model_validate({"prefill_config": vllm, "decode_config": sglang}) + + +def test_sglang_connector_requires_bootstrap_fields(): + """The connector fails early if the request model lacks bootstrap fields.""" + from ray.llm._internal.serve.core.configs.openai_api_models import ( + ChatCompletionRequest, + ) + + if "bootstrap_room" in ChatCompletionRequest.model_fields: + # SGLang-only env: the guard passes silently. + SGLangConnectorBackend._check_request_model_has_bootstrap_fields() + else: + # vLLM installed: the guard must raise. + with pytest.raises(RuntimeError, match="bootstrap_room"): + SGLangConnectorBackend._check_request_model_has_bootstrap_fields() + + +if __name__ == "__main__": + sys.exit(pytest.main(["-xvs", __file__])) diff --git a/release/release_tests.yaml b/release/release_tests.yaml index baedcd5dd6c1..51a82709ebd5 100644 --- a/release/release_tests.yaml +++ b/release/release_tests.yaml @@ -5203,6 +5203,28 @@ timeout: 3600 script: pytest -vs test_llm_serve_sglang.py +- name: llm_serve_sglang_pd_nixl + frequency: manual + python: "3.12" + group: llm-serve + team: llm + working_dir: llm_tests/serve + + cluster: + byod: + type: llm-cu130 + runtime_env: + - RAY_EXPERIMENTAL_NOSET_CUDA_VISIBLE_DEVICES=1 + - UCX_TLS=all + - UCX_NET_DEVICES=all + post_build_script: byod_llm_sglang_test.sh + anyscale_sdk_2026: true + cluster_compute: llm_2x_4xl4.yaml + + run: + timeout: 3600 + script: pytest -vs test_llm_serve_sglang_pd.py + - name: llm_serve_kv_router frequency: manual python: "3.12" From 5b774726628242fe944db0dccc63858cf9e3bfdf Mon Sep 17 00:00:00 2001 From: Limark Dcunha Date: Mon, 28 Sep 2026 22:58:35 -0400 Subject: [PATCH 2/3] Reverting unnecessary changes Signed-off-by: Limark Dcunha --- .../advanced-guides/asyncio-best-practices.md | 19 +++++++++++++------ python/ray/serve/_private/constants.py | 10 ++++++---- python/ray/serve/tests/BUILD.bazel | 3 ++- 3 files changed, 21 insertions(+), 11 deletions(-) diff --git a/doc/source/serve/advanced-guides/asyncio-best-practices.md b/doc/source/serve/advanced-guides/asyncio-best-practices.md index d05a60a545d0..48f10006eb2a 100644 --- a/doc/source/serve/advanced-guides/asyncio-best-practices.md +++ b/doc/source/serve/advanced-guides/asyncio-best-practices.md @@ -80,8 +80,8 @@ For a synchronous deployment: How this method executes depends on configuration: -- With `RAY_SERVE_RUN_SYNC_IN_THREADPOOL=1` (the default), Serve offloads `__call__` to a threadpool so the event loop stays responsive. -- With `RAY_SERVE_RUN_SYNC_IN_THREADPOOL=0`, `__call__` runs directly on the user event loop and blocks it for 1 second. +- With `RAY_SERVE_RUN_SYNC_IN_THREADPOOL=0` (current default), `__call__` runs directly on the user event loop and blocks it for 1 second. +- With `RAY_SERVE_RUN_SYNC_IN_THREADPOOL=1`, Serve offloads `__call__` to a threadpool so the event loop stays responsive. ### FastAPI ingress (`@serve.ingress`) @@ -274,7 +274,16 @@ Ray Serve exposes several environment variables that control how user code inter ### `RAY_SERVE_RUN_SYNC_IN_THREADPOOL` -By default (`RAY_SERVE_RUN_SYNC_IN_THREADPOOL=1`), synchronous methods in a deployment run in a threadpool: +By default (`RAY_SERVE_RUN_SYNC_IN_THREADPOOL=0`), which means synchronous methods in a deployment run directly on the user event loop. To help you migrate to a safer model, Serve emits a warning like: + +> `RAY_SERVE_RUN_SYNC_IN_THREADPOOL_WARNING`: Calling sync method '...' directly on the asyncio loop. In a future version, sync methods will be run in a threadpool by default... + +This warning means: + +- You have a `def` method that is currently running on the event loop. +- In a future version, that method runs in a threadpool instead. + +You can opt in to the future behavior now by setting: ```bash export RAY_SERVE_RUN_SYNC_IN_THREADPOOL=1 @@ -285,13 +294,11 @@ When this flag is `1`: - Serve runs synchronous methods in a threadpool. - The event loop is free to keep serving other requests while sync methods run. -Make sure: +Before enabling this in production, make sure: - Your handler code and any shared state are thread-safe. - Your model objects can safely be used from multiple threads, or you protect them with locks. -Set this flag to `0` to retain the legacy behavior of running synchronous methods directly on the user event loop. - ### `RAY_SERVE_RUN_USER_CODE_IN_SEPARATE_THREAD` By default, Serve runs user code in a separate event loop from the replica's main/control loop: diff --git a/python/ray/serve/_private/constants.py b/python/ray/serve/_private/constants.py index 7d78a97c97cb..2deb89fc30f9 100644 --- a/python/ray/serve/_private/constants.py +++ b/python/ray/serve/_private/constants.py @@ -647,12 +647,14 @@ ) # Run sync methods defined in the replica in a thread pool by default. -RAY_SERVE_RUN_SYNC_IN_THREADPOOL = get_env_bool("RAY_SERVE_RUN_SYNC_IN_THREADPOOL", "1") +RAY_SERVE_RUN_SYNC_IN_THREADPOOL = get_env_bool("RAY_SERVE_RUN_SYNC_IN_THREADPOOL", "0") RAY_SERVE_RUN_SYNC_IN_THREADPOOL_WARNING = ( - "Calling sync method '{method_name}' directly on the asyncio loop because " - "RAY_SERVE_RUN_SYNC_IN_THREADPOOL=0. This can block other requests. Make " - "the method `async def` or restore threadpool dispatch by setting " + "Calling sync method '{method_name}' directly on the " + "asyncio loop. In a future version, sync methods will be run in a " + "threadpool by default. Ensure your sync methods are thread safe " + "or keep the existing behavior by making them `async def`. Opt " + "into the new behavior by setting " "RAY_SERVE_RUN_SYNC_IN_THREADPOOL=1." ) diff --git a/python/ray/serve/tests/BUILD.bazel b/python/ray/serve/tests/BUILD.bazel index eba3aae51d86..b118ed36e596 100644 --- a/python/ray/serve/tests/BUILD.bazel +++ b/python/ray/serve/tests/BUILD.bazel @@ -812,7 +812,8 @@ py_test_module_list( ], ) -# Test the default behavior that runs replica sync methods in a threadpool. +# Test currently off-by-default behavior to run replica sync methods in a threadpool. +# TODO(edoakes): remove this once the FF is flipped on by default. py_test_module_list( size = "medium", env = {"RAY_SERVE_RUN_SYNC_IN_THREADPOOL": "1"}, From 2e8c4ac405a28a59e40412b855850e867c99b319 Mon Sep 17 00:00:00 2001 From: Limark Dcunha Date: Tue, 29 Sep 2026 10:53:38 -0400 Subject: [PATCH 3/3] Fix release test schema and config test validation errors Signed-off-by: Limark Dcunha --- .../serve/cpu/configs/test_llm_config_num_devices.py | 8 ++++++-- release/release_tests.yaml | 6 +++--- 2 files changed, 9 insertions(+), 5 deletions(-) diff --git a/python/ray/llm/tests/serve/cpu/configs/test_llm_config_num_devices.py b/python/ray/llm/tests/serve/cpu/configs/test_llm_config_num_devices.py index 20878e56adce..bef233f08be7 100644 --- a/python/ray/llm/tests/serve/cpu/configs/test_llm_config_num_devices.py +++ b/python/ray/llm/tests/serve/cpu/configs/test_llm_config_num_devices.py @@ -7,7 +7,9 @@ def test_num_devices_vllm_default(): cfg = LLMConfig( - model_loading_config=dict(model_source="facebook/opt-125m"), + model_loading_config=dict( + model_id="opt-125m", model_source="facebook/opt-125m" + ), llm_engine="vLLM", engine_kwargs=dict(tensor_parallel_size=2, pipeline_parallel_size=2), ) @@ -16,7 +18,9 @@ def test_num_devices_vllm_default(): def test_num_devices_defaults_to_one(): cfg = LLMConfig( - model_loading_config=dict(model_source="facebook/opt-125m"), + model_loading_config=dict( + model_id="opt-125m", model_source="facebook/opt-125m" + ), llm_engine="vLLM", ) assert cfg.num_devices == 1 diff --git a/release/release_tests.yaml b/release/release_tests.yaml index 51a82709ebd5..b9f85c8e9e0e 100644 --- a/release/release_tests.yaml +++ b/release/release_tests.yaml @@ -5214,9 +5214,9 @@ byod: type: llm-cu130 runtime_env: - - RAY_EXPERIMENTAL_NOSET_CUDA_VISIBLE_DEVICES=1 - - UCX_TLS=all - - UCX_NET_DEVICES=all + RAY_EXPERIMENTAL_NOSET_CUDA_VISIBLE_DEVICES: "1" + UCX_TLS: all + UCX_NET_DEVICES: all post_build_script: byod_llm_sglang_test.sh anyscale_sdk_2026: true cluster_compute: llm_2x_4xl4.yaml