feat(cli): add extensible serve backend plugins (#34753)

This commit is contained in:
Mick
2026-08-14 13:57:59 +08:00
committed by GitHub
parent 463981922c
commit 46d84f4b48
7 changed files with 796 additions and 85 deletions
+1
View File
@@ -1011,6 +1011,7 @@
"pages": [
"docs/developer_guide/overview",
"docs/developer_guide/contribution_guide",
"docs/developer_guide/serve_backend_plugins",
{
"group": "Development",
"pages": [
+1
View File
@@ -4,6 +4,7 @@ description: Contributing to SGLang — development setup, benchmarking, and eva
---
- [Contribution Guide](./contribution_guide)
- [Add an out-of-tree serve backend](/docs/developer_guide/serve_backend_plugins)
- [Development Guide (Docker)](./development_guide_using_docker)
- [JIT Kernels](./development_jit_kernel_guide)
- [Quantization Contribution Guide](./quantization_contribution_guide)
@@ -0,0 +1,157 @@
---
title: "Add an out-of-tree serve backend"
description: "Connect an ecosystem runtime to sglang serve through the versioned serve backend plugin API."
---
A serve backend plugin connects an out-of-tree runtime to the SGLang-owned CLI:
```bash
sglang serve MODEL_PATH --model-type BACKEND_NAME [BACKEND_OPTIONS]
```
The extension retains its own argument parser, process topology, API endpoints, hardware requirements, and release cycle. The core CLI owns backend selection and common child-process cleanup.
## Prerequisites
- Install your extension and a compatible SGLang version in the same Python 3.10+ environment.
- Declare the SGLang version range tested by your extension in its package dependencies.
- Keep the backend factory and detector independent of GPU initialization and model loading.
The plugin API itself is platform-independent. Your backend defines its supported operating systems, accelerators, parallelism options, and authentication requirements.
## Keep one owner for the executable
Only the `sglang` distribution should publish a console script named `sglang`. Your extension registers package metadata under `sglang.serve_backends`; it must not publish another `sglang` script.
This prevents installation order from replacing the command and ensures uninstalling an extension does not remove the core executable. You can retain a project-specific executable as a compatibility alias:
```bash
my-runtime serve MODEL_PATH
sglang serve MODEL_PATH --model-type my_runtime
```
Both commands should call the same backend implementation.
## Register a backend factory
Add a zero-argument factory to your extension's `pyproject.toml`:
```toml
[project]
name = "my-sglang-runtime"
version = "0.1.0"
dependencies = ["sglang"]
[project.entry-points."sglang.serve_backends"]
my_runtime = "my_sglang_runtime.sglang_backend:create_backend"
```
The entry point name, `my_runtime`, becomes an accepted `--model-type` value. Choose a distinctive name. `auto` and SGLang's in-tree backend names are reserved.
## Implement the backend
Create `my_sglang_runtime/sglang_backend.py`:
```python
import argparse
from sglang.cli.serve_backends import (
ServeBackend,
ServeBackendDetection,
ServeRequest,
)
def detect(request: ServeRequest) -> ServeBackendDetection:
if request.model_path is None:
return ServeBackendDetection.UNKNOWN
if supports_model(request.model_path):
return ServeBackendDetection.MATCH
return ServeBackendDetection.NO_MATCH
def run(request: ServeRequest) -> None:
parser = argparse.ArgumentParser(prog="sglang serve")
parser.add_argument("--model-path", required=True)
parser.add_argument("--pipeline-parallel", type=int, default=1)
args, remaining = parser.parse_known_args(request.argv)
launch_runtime(args, remaining)
def create_backend() -> ServeBackend:
return ServeBackend(api_version=1, run=run, detect=detect)
```
Replace `supports_model()` and `launch_runtime()` with your extension's lightweight metadata check and blocking server launcher. A real `run()` call should block for the server lifetime. It must also honor `-h` and `--help` without launching a server; `argparse` does this automatically.
Among serve backend entry points, explicit selection imports only the selected provider. Automatic selection loads installed backend factories and invokes their detectors, so importing this module and calling `detect()` must not initialize accelerators, import model weights, or start workers.
## Handle forwarded arguments
SGLang removes `--model-type` and normalizes a positional Hugging Face model ID or local model directory before dispatch. For example:
```bash
sglang serve org/model --model-type my_runtime --pipeline-parallel 2
```
Your backend receives:
```python
("--model-path", "org/model", "--pipeline-parallel", "2")
```
The selected backend owns all remaining argument parsing and validation. Target backend-specific help with:
```bash
sglang serve --model-type my_runtime --help
```
## Support config-only runtimes
SGLang requires a model path by default. If your runtime resolves its model and parallelism settings from a configuration file, disable that validation:
```python
def create_backend() -> ServeBackend:
return ServeBackend(
api_version=1,
run=run,
detect=detect,
requires_model_path=False,
)
```
You can then accept commands such as:
```bash
sglang serve --model-type my_runtime --config pipeline.yaml
```
Config-only requests generally require explicit `--model-type` unless your detector can identify the backend from the remaining arguments.
## Understand automatic routing
The default `--model-type auto` follows these rules:
1. Backends without a detector remain explicit-only.
2. One `MATCH` selects that backend.
3. Multiple matches fail and require an explicit `--model-type`.
4. `UNKNOWN` and detector failures do not claim the request.
5. No matches preserve the existing LLM fallback.
The registry does not resolve overlap by package installation order or a hidden priority. A backend can opt out of automatic routing by omitting its detector:
```python
ServeBackend(api_version=1, run=run, requires_model_path=False)
```
## Maintain compatibility
Declare the API version implemented by your extension as a literal. Do not copy SGLang's current version constant at runtime; a fixed value lets a future SGLang release detect an older plugin contract. SGLang rejects incompatible, duplicate, and reserved backend registrations with an actionable error.
The public extension contract consists of:
- `ServeRequest`: normalized backend arguments and the optional model path
- `ServeBackend`: the runner, optional detector, and model-path requirement
- `ServeBackendDetection`: `MATCH`, `NO_MATCH`, or `UNKNOWN`
Test explicit selection, backend-specific help, automatic detection, ambiguous models, and installation or removal alongside the core `sglang` package.
+147 -83
View File
@@ -4,7 +4,14 @@ import argparse
import logging
import os
from sglang.cli.utils import get_is_diffusion_model, get_model_path
from sglang.cli.serve_backends import (
SERVE_BACKEND_API_VERSION,
ServeBackend,
ServeBackendDetection,
ServeBackendRegistry,
ServeRequest,
)
from sglang.cli.utils import get_is_diffusion_model, get_model_path, try_get_model_path
from sglang.srt.utils import kill_process_tree
from sglang.srt.utils.common import suppress_noisy_warnings
@@ -14,36 +21,35 @@ logger = logging.getLogger(__name__)
def _extract_model_type_override(extra_argv):
"""Extract and remove --model-type override from argv."""
model_type = "auto"
"""Extract and remove the backend selector from argv.
Validation is deferred to :class:`ServeBackendRegistry` because installed
out-of-tree entry points dynamically extend the accepted values.
"""
backend_name = "auto"
filtered_argv = []
i = 0
while i < len(extra_argv):
arg = extra_argv[i]
if arg == "--model-type":
if i + 1 >= len(extra_argv):
raise Exception(
"Error: --model-type requires a value. "
"Valid values are: auto, llm, diffusion."
)
model_type = extra_argv[i + 1]
raise ValueError("Error: --model-type requires a value.")
backend_name = extra_argv[i + 1]
i += 2
continue
if arg.startswith("--model-type="):
model_type = arg.split("=", 1)[1]
backend_name = arg.split("=", 1)[1]
i += 1
continue
filtered_argv.append(arg)
i += 1
if model_type not in ("auto", "llm", "diffusion"):
raise Exception(
f"Error: invalid --model-type '{model_type}'. "
"Valid values are: auto, llm, diffusion."
)
return model_type, filtered_argv
if not backend_name:
raise ValueError("Error: --model-type requires a non-empty value.")
return backend_name, filtered_argv
def _normalize_positional_model_path(extra_argv):
@@ -53,91 +59,149 @@ def _normalize_positional_model_path(extra_argv):
return extra_argv, False
def serve(args, extra_argv):
if any(h in extra_argv for h in ("-h", "--help")):
# Since the server type is determined by the model, and we don't have a model path,
# we can't show the exact help. Instead, we show a general help message and then
# the help for both possible server types.
print(
"Usage: sglang serve <model-name-or-path> [additional-arguments]\n"
" or: sglang serve --model-path <model-name-or-path> [additional-arguments]\n\n"
"This command can launch either a standard language model server or a diffusion model server.\n"
"The server type is determined by the --model-path.\n"
"Optional override: --model-type {auto,llm,diffusion} "
"(default: auto, fallback to LLM on detection failure)."
def _print_llm_help(_request: ServeRequest) -> None:
from sglang.srt.server_args import prepare_server_args
try:
prepare_server_args(["--help"])
except SystemExit:
pass # argparse --help calls sys.exit
def _print_diffusion_help(_request: ServeRequest) -> None:
try:
from sglang.multimodal_gen.runtime.entrypoints.cli.serve import (
add_multimodal_gen_serve_args,
)
print("\n--- Help for Standard Language Model Server ---")
from sglang.srt.server_args import prepare_server_args
parser = argparse.ArgumentParser(
prog="sglang serve",
description="SGLang Diffusion Model Serving",
)
add_multimodal_gen_serve_args(parser)
parser.print_help()
except ImportError:
print(
"Diffusion model support is not available. "
'Install with: pip install "sglang[diffusion]"'
)
try:
prepare_server_args(["--help"])
except SystemExit:
pass # argparse --help calls sys.exit
print("\n--- Help for Diffusion Model Server ---")
try:
from sglang.multimodal_gen.runtime.entrypoints.cli.serve import (
add_multimodal_gen_serve_args,
)
def _run_llm(request: ServeRequest) -> None:
if any(arg in request.argv for arg in ("-h", "--help")):
_print_llm_help(request)
return
parser = argparse.ArgumentParser(
prog="sglang serve",
description="SGLang Diffusion Model Serving",
)
add_multimodal_gen_serve_args(parser)
parser.print_help()
except ImportError:
print(
"Diffusion model support is not available. "
'Install with: pip install "sglang[diffusion]"'
)
from sglang.launch_server import run_server
from sglang.srt.server_args import prepare_server_args
server_args = prepare_server_args(list(request.argv))
run_server(server_args)
def _detect_diffusion(request: ServeRequest) -> ServeBackendDetection:
if request.model_path is None:
return ServeBackendDetection.UNKNOWN
if get_is_diffusion_model(request.model_path):
return ServeBackendDetection.MATCH
return ServeBackendDetection.NO_MATCH
def _run_diffusion(request: ServeRequest) -> None:
if any(arg in request.argv for arg in ("-h", "--help")):
_print_diffusion_help(request)
return
from sglang.multimodal_gen.runtime.entrypoints.cli.serve import (
add_multimodal_gen_serve_args,
execute_serve_cmd,
)
parser = argparse.ArgumentParser(description="SGLang Diffusion Model Serving")
add_multimodal_gen_serve_args(parser)
parsed_args, remaining_argv = parser.parse_known_args(list(request.argv))
if request.model_path_is_positional:
parsed_args._sglang_explicit_arg_names = {"model_path"}
execute_serve_cmd(parsed_args, remaining_argv)
def _create_backend_registry() -> ServeBackendRegistry:
return ServeBackendRegistry(
{
"llm": ServeBackend(
api_version=SERVE_BACKEND_API_VERSION,
run=_run_llm,
),
"diffusion": ServeBackend(
api_version=SERVE_BACKEND_API_VERSION,
run=_run_diffusion,
detect=_detect_diffusion,
),
}
)
def _print_general_help(registry: ServeBackendRegistry) -> None:
available = ",".join(("auto", *registry.available_names))
print(
"Usage: sglang serve <model-name-or-path> [additional-arguments]\n"
" or: sglang serve --model-path <model-name-or-path> "
"[additional-arguments]\n\n"
"The serving backend is detected from the model by default. Installed "
"out-of-tree projects can add more backends.\n"
f"Optional override: --model-type {{{available}}} "
"(default: auto, fallback to LLM when no backend matches).\n"
"Use `sglang serve --model-type BACKEND --help` for backend-specific "
"help."
)
print("\n--- Help for Standard Language Model Server ---")
_print_llm_help(ServeRequest(argv=("--help",), model_path=None))
print("\n--- Help for Diffusion Model Server ---")
_print_diffusion_help(ServeRequest(argv=("--help",), model_path=None))
def serve(args, extra_argv):
del args # The top-level parser currently has no serve-owned fields.
backend_name, dispatch_argv = _extract_model_type_override(extra_argv)
dispatch_argv, positional_model_path = _normalize_positional_model_path(
dispatch_argv
)
request = ServeRequest(
argv=tuple(dispatch_argv),
model_path=try_get_model_path(dispatch_argv),
model_path_is_positional=positional_model_path,
)
registry = _create_backend_registry()
if any(h in request.argv for h in ("-h", "--help")):
if backend_name == "auto":
_print_general_help(registry)
else:
registry.get(backend_name).backend.run(request)
return
from sglang.srt.plugins import load_plugins
load_plugins()
model_type, dispatch_argv = _extract_model_type_override(extra_argv)
dispatch_argv, positional_model_path = _normalize_positional_model_path(
dispatch_argv
)
model_path = get_model_path(dispatch_argv)
try:
if model_type == "auto":
is_diffusion_model = get_is_diffusion_model(model_path)
if is_diffusion_model:
logger.info("Diffusion model detected")
if backend_name == "auto":
registered = registry.auto_detect(request)
logger.info("Selected serve backend %r", registered.name)
else:
is_diffusion_model = model_type == "diffusion"
registered = registry.get(backend_name)
logger.info(
"Dispatch override enabled: --model-type=%s " "(skip auto detection)",
model_type,
backend_name,
)
if is_diffusion_model:
# Logic for Diffusion Models
from sglang.multimodal_gen.runtime.entrypoints.cli.serve import (
add_multimodal_gen_serve_args,
execute_serve_cmd,
)
if registered.backend.requires_model_path and request.model_path is None:
get_model_path(request.argv) # Raise the existing actionable CLI error.
parser = argparse.ArgumentParser(
description="SGLang Diffusion Model Serving"
)
add_multimodal_gen_serve_args(parser)
parsed_args, remaining_argv = parser.parse_known_args(dispatch_argv)
if positional_model_path:
parsed_args._sglang_explicit_arg_names = {"model_path"}
execute_serve_cmd(parsed_args, remaining_argv)
else:
# Logic for Standard Language Models
from sglang.launch_server import run_server
from sglang.srt.server_args import prepare_server_args
server_args = prepare_server_args(dispatch_argv)
run_server(server_args)
registered.backend.run(request)
finally:
kill_process_tree(os.getpid(), include_parent=False)
+231
View File
@@ -0,0 +1,231 @@
# SPDX-License-Identifier: Apache-2.0
"""Extension API and discovery for ``sglang serve`` backends.
Out-of-tree projects register a zero-argument factory in the
``sglang.serve_backends`` entry point group. The factory returns a
:class:`ServeBackend`; its entry point name becomes a valid ``--model-type``.
The module is intentionally lightweight. Importing it must not initialize an
inference runtime or import an out-of-tree backend implementation.
"""
from __future__ import annotations
import logging
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from enum import Enum
from importlib.metadata import EntryPoint, entry_points
logger = logging.getLogger(__name__)
SERVE_BACKENDS_GROUP = "sglang.serve_backends"
SERVE_BACKEND_API_VERSION = 1
RESERVED_SERVE_BACKEND_NAMES = frozenset({"auto"})
class ServeBackendDetection(str, Enum):
"""Result returned by a serve backend's optional detector."""
MATCH = "match"
NO_MATCH = "no_match"
UNKNOWN = "unknown"
@dataclass(frozen=True)
class ServeRequest:
"""Normalized command line forwarded from ``sglang serve`` to a backend."""
argv: tuple[str, ...]
model_path: str | None
model_path_is_positional: bool = False
ServeBackendRunner = Callable[[ServeRequest], None]
ServeBackendDetector = Callable[[ServeRequest], ServeBackendDetection]
@dataclass(frozen=True)
class ServeBackend:
"""Implementation contract for built-in and out-of-tree serve backends.
``run`` must parse the backend-owned arguments in ``request.argv``. It must
also honor ``-h`` and ``--help`` without launching a server. For a real
launch, it should block for the server lifetime; SGLang applies its common
child-process cleanup after ``run`` returns or raises.
``detect`` is optional. Backends without a detector remain available by
explicit ``--model-type`` but do not participate in automatic routing.
Detector implementations should be lightweight and return ``UNKNOWN`` on
inconclusive I/O or metadata errors.
"""
api_version: int
run: ServeBackendRunner
detect: ServeBackendDetector | None = None
requires_model_path: bool = True
@dataclass(frozen=True)
class RegisteredServeBackend:
"""A loaded backend together with its discovery metadata."""
name: str
backend: ServeBackend
distribution: str | None = None
class ServeBackendRegistry:
"""Registry of built-in and installed out-of-tree serve backends."""
def __init__(self, builtins: Mapping[str, ServeBackend]) -> None:
invalid_builtin_names = set(builtins) & RESERVED_SERVE_BACKEND_NAMES
if invalid_builtin_names:
names = ", ".join(sorted(invalid_builtin_names))
raise ValueError(f"Reserved serve backend names cannot be used: {names}")
self._builtins = dict(builtins)
self._entry_points = self._discover_entry_points()
self._loaded: dict[str, RegisteredServeBackend] = {
name: RegisteredServeBackend(name=name, backend=backend)
for name, backend in self._builtins.items()
}
reserved = (set(self._builtins) | RESERVED_SERVE_BACKEND_NAMES) & set(
self._entry_points
)
if reserved:
names = ", ".join(sorted(reserved))
raise RuntimeError(
"Out-of-tree serve backends cannot replace reserved or built-in "
f"backends: {names}"
)
@staticmethod
def _discover_entry_points() -> dict[str, list[EntryPoint]]:
discovered: dict[str, list[EntryPoint]] = {}
for entry_point in entry_points(group=SERVE_BACKENDS_GROUP):
discovered.setdefault(entry_point.name, []).append(entry_point)
return discovered
@property
def available_names(self) -> tuple[str, ...]:
"""Return backend names without importing out-of-tree packages."""
external_names = sorted(set(self._entry_points) - set(self._builtins))
return (*self._builtins, *external_names)
def get(self, name: str) -> RegisteredServeBackend:
"""Return one backend, importing only the explicitly requested plugin."""
if name in self._loaded:
return self._loaded[name]
candidates = self._entry_points.get(name, [])
if not candidates:
available = ", ".join(("auto", *self.available_names))
raise ValueError(
f"Unknown serve backend {name!r}. Available values: {available}."
)
if len(candidates) > 1:
providers = ", ".join(
sorted(
self._entry_point_provider(candidate) for candidate in candidates
)
)
raise RuntimeError(
f"Multiple distributions register serve backend {name!r}: "
f"{providers}. Uninstall one provider or choose another backend name."
)
entry_point = candidates[0]
try:
factory = entry_point.load()
if not callable(factory):
raise TypeError("the entry point must resolve to a callable factory")
backend = factory()
except Exception as exc:
raise RuntimeError(
f"Failed to load serve backend {name!r} from "
f"{self._entry_point_provider(entry_point)}: {exc}"
) from exc
if not isinstance(backend, ServeBackend):
raise TypeError(
f"Serve backend {name!r} factory returned {type(backend).__name__}; "
"expected sglang.cli.serve_backends.ServeBackend."
)
if backend.api_version != SERVE_BACKEND_API_VERSION:
raise RuntimeError(
f"Serve backend {name!r} uses API version {backend.api_version}; "
f"this SGLang release requires version {SERVE_BACKEND_API_VERSION}."
)
registered = RegisteredServeBackend(
name=name,
backend=backend,
distribution=self._entry_point_distribution(entry_point),
)
self._loaded[name] = registered
return registered
def auto_detect(self, request: ServeRequest) -> RegisteredServeBackend:
"""Resolve a unique detector match, falling back to the ``llm`` backend."""
matches: list[RegisteredServeBackend] = []
for name in self.available_names:
if name == "llm":
# LLM preserves the historical fallback role instead of matching
# every Hugging Face repository.
continue
try:
registered = self.get(name)
detector = registered.backend.detect
if detector is None:
continue
result = detector(request)
if not isinstance(result, ServeBackendDetection):
raise TypeError(
"detector must return ServeBackendDetection, got "
f"{type(result).__name__}"
)
if result is ServeBackendDetection.MATCH:
matches.append(registered)
except Exception as exc:
# An unrelated optional extension must not make the default LLM
# path unusable. Explicit selection remains strict via get().
logger.warning(
"Skipping automatic detection for serve backend %r: %s",
name,
exc,
)
if len(matches) > 1:
names = ", ".join(match.name for match in matches)
raise RuntimeError(
f"Multiple serve backends matched this request: {names}. "
"Select one explicitly with --model-type BACKEND."
)
if matches:
return matches[0]
return self.get("llm")
@classmethod
def _entry_point_distribution(cls, entry_point: EntryPoint) -> str | None:
distribution = getattr(entry_point, "dist", None)
return getattr(distribution, "name", None)
@classmethod
def _entry_point_provider(cls, entry_point: EntryPoint) -> str:
return cls._entry_point_distribution(entry_point) or entry_point.value
__all__ = [
"SERVE_BACKEND_API_VERSION",
"SERVE_BACKENDS_GROUP",
"ServeBackend",
"ServeBackendDetection",
"ServeBackendRegistry",
"ServeRequest",
]
+10 -2
View File
@@ -96,8 +96,9 @@ def get_is_diffusion_model(model_path: str) -> bool:
return False
def get_model_path(extra_argv):
# Find the model_path argument
def try_get_model_path(extra_argv) -> str | None:
"""Return a model path from command-line arguments when one is present."""
model_path = None
for i, arg in enumerate(extra_argv):
if arg in ("--model-path", "--model"):
@@ -108,6 +109,13 @@ def get_model_path(extra_argv):
model_path = arg.split("=", 1)[1]
break
return model_path
def get_model_path(extra_argv):
# Find the model_path argument
model_path = try_get_model_path(extra_argv)
if model_path is None:
# Fallback for --help or other cases where model-path is not provided
if any(h in extra_argv for h in ["-h", "--help"]):
@@ -0,0 +1,249 @@
# SPDX-License-Identifier: Apache-2.0
import unittest
from unittest.mock import MagicMock, patch
from sglang.cli.serve import serve
from sglang.cli.serve_backends import (
SERVE_BACKEND_API_VERSION,
RegisteredServeBackend,
ServeBackend,
ServeBackendDetection,
ServeBackendRegistry,
ServeRequest,
)
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
def _make_entry_point(name, factory, distribution=None):
entry_point = MagicMock()
entry_point.name = name
entry_point.value = f"fake_{name}:create_backend"
entry_point.load.return_value = factory
if distribution is None:
entry_point.dist = None
else:
entry_point.dist = MagicMock()
entry_point.dist.name = distribution
return entry_point
def _backend(*, detector=None, requires_model_path=True, api_version=None):
if api_version is None:
api_version = SERVE_BACKEND_API_VERSION
return ServeBackend(
api_version=api_version,
run=MagicMock(),
detect=detector,
requires_model_path=requires_model_path,
)
class TestServeBackendRegistry(unittest.TestCase):
@patch("sglang.cli.serve_backends.entry_points")
def test_listing_does_not_import_out_of_tree_backends(self, mock_entry_points):
factory = MagicMock(return_value=_backend())
entry_point = _make_entry_point("omni", factory)
mock_entry_points.return_value = [entry_point]
registry = ServeBackendRegistry({"llm": _backend()})
self.assertEqual(registry.available_names, ("llm", "omni"))
entry_point.load.assert_not_called()
factory.assert_not_called()
@patch("sglang.cli.serve_backends.entry_points")
def test_explicit_selection_only_loads_selected_backend(self, mock_entry_points):
omni_entry_point = _make_entry_point(
"omni", MagicMock(return_value=_backend()), "sglang-omni"
)
other_entry_point = _make_entry_point(
"other", MagicMock(return_value=_backend()), "other-project"
)
mock_entry_points.return_value = [omni_entry_point, other_entry_point]
registry = ServeBackendRegistry({"llm": _backend()})
registered = registry.get("omni")
self.assertEqual(registered.name, "omni")
self.assertEqual(registered.distribution, "sglang-omni")
omni_entry_point.load.assert_called_once_with()
other_entry_point.load.assert_not_called()
@patch("sglang.cli.serve_backends.entry_points")
def test_auto_detection_selects_one_unique_match(self, mock_entry_points):
matching = _backend(
detector=MagicMock(return_value=ServeBackendDetection.MATCH)
)
non_matching = _backend(
detector=MagicMock(return_value=ServeBackendDetection.NO_MATCH)
)
mock_entry_points.return_value = [
_make_entry_point("omni", lambda: matching),
_make_entry_point("speech", lambda: non_matching),
]
registry = ServeBackendRegistry({"llm": _backend()})
request = ServeRequest(argv=("--model-path", "model"), model_path="model")
self.assertEqual(registry.auto_detect(request).name, "omni")
@patch("sglang.cli.serve_backends.entry_points")
def test_auto_detection_rejects_ambiguous_matches(self, mock_entry_points):
mock_entry_points.return_value = [
_make_entry_point(
"omni",
lambda: _backend(
detector=MagicMock(return_value=ServeBackendDetection.MATCH)
),
),
_make_entry_point(
"speech",
lambda: _backend(
detector=MagicMock(return_value=ServeBackendDetection.MATCH)
),
),
]
registry = ServeBackendRegistry({"llm": _backend()})
request = ServeRequest(argv=("--model-path", "model"), model_path="model")
with self.assertRaisesRegex(RuntimeError, "Multiple serve backends matched"):
registry.auto_detect(request)
@patch("sglang.cli.serve_backends.entry_points")
def test_broken_optional_detector_does_not_block_llm_fallback(
self, mock_entry_points
):
broken_entry_point = _make_entry_point("broken", MagicMock())
broken_entry_point.load.side_effect = ImportError("optional dependency missing")
mock_entry_points.return_value = [broken_entry_point]
llm = _backend()
registry = ServeBackendRegistry({"llm": llm})
request = ServeRequest(argv=("--model-path", "model"), model_path="model")
with self.assertLogs("sglang.cli.serve_backends", level="WARNING"):
selected = registry.auto_detect(request)
self.assertIs(selected.backend, llm)
with self.assertRaisesRegex(RuntimeError, "Failed to load serve backend"):
registry.get("broken")
@patch("sglang.cli.serve_backends.entry_points")
def test_duplicate_provider_names_are_rejected_when_selected(
self, mock_entry_points
):
mock_entry_points.return_value = [
_make_entry_point("omni", lambda: _backend(), "provider-a"),
_make_entry_point("omni", lambda: _backend(), "provider-b"),
]
registry = ServeBackendRegistry({"llm": _backend()})
with self.assertRaisesRegex(RuntimeError, "provider-a, provider-b"):
registry.get("omni")
@patch("sglang.cli.serve_backends.entry_points")
def test_out_of_tree_backend_cannot_replace_builtin(self, mock_entry_points):
mock_entry_points.return_value = [
_make_entry_point("llm", lambda: _backend(), "bad-provider")
]
with self.assertRaisesRegex(RuntimeError, "cannot replace"):
ServeBackendRegistry({"llm": _backend()})
@patch("sglang.cli.serve_backends.entry_points")
def test_backend_api_version_is_validated(self, mock_entry_points):
mock_entry_points.return_value = [
_make_entry_point(
"future",
lambda: _backend(api_version=SERVE_BACKEND_API_VERSION + 1),
)
]
registry = ServeBackendRegistry({"llm": _backend()})
with self.assertRaisesRegex(RuntimeError, "uses API version"):
registry.get("future")
class TestServeBackendDispatch(unittest.TestCase):
@patch("sglang.cli.serve.kill_process_tree")
@patch("sglang.srt.plugins.load_plugins")
@patch("sglang.cli.serve._create_backend_registry")
def test_explicit_backend_receives_normalized_model_path(
self, mock_registry_factory, mock_load_plugins, mock_kill
):
backend = _backend()
registry = MagicMock()
registry.get.return_value = RegisteredServeBackend("omni", backend)
mock_registry_factory.return_value = registry
serve(
None,
["Example/Model", "--model-type", "omni", "--pipeline-parallel", "2"],
)
request = backend.run.call_args.args[0]
self.assertEqual(
request.argv,
("--model-path", "Example/Model", "--pipeline-parallel", "2"),
)
self.assertEqual(request.model_path, "Example/Model")
self.assertTrue(request.model_path_is_positional)
mock_load_plugins.assert_called_once_with()
mock_kill.assert_called_once()
@patch("sglang.cli.serve.kill_process_tree")
@patch("sglang.srt.plugins.load_plugins")
@patch("sglang.cli.serve._create_backend_registry")
def test_explicit_backend_can_support_config_only_requests(
self, mock_registry_factory, _mock_load_plugins, _mock_kill
):
backend = _backend(requires_model_path=False)
registry = MagicMock()
registry.get.return_value = RegisteredServeBackend("pipeline", backend)
mock_registry_factory.return_value = registry
serve(None, ["--model-type=pipeline", "--config", "pipeline.yaml"])
request = backend.run.call_args.args[0]
self.assertIsNone(request.model_path)
self.assertEqual(request.argv, ("--config", "pipeline.yaml"))
@patch("sglang.cli.serve.kill_process_tree")
@patch("sglang.srt.plugins.load_plugins")
@patch("sglang.cli.serve._create_backend_registry")
def test_auto_detection_uses_registry(
self, mock_registry_factory, _mock_load_plugins, _mock_kill
):
backend = _backend()
registry = MagicMock()
registry.auto_detect.return_value = RegisteredServeBackend("omni", backend)
mock_registry_factory.return_value = registry
serve(None, ["Example/Model"])
registry.auto_detect.assert_called_once()
backend.run.assert_called_once()
@patch("sglang.cli.serve.kill_process_tree")
@patch("sglang.srt.plugins.load_plugins")
@patch("sglang.cli.serve._create_backend_registry")
def test_targeted_help_is_forwarded_without_startup_or_model_path(
self, mock_registry_factory, mock_load_plugins, mock_kill
):
backend = _backend()
registry = MagicMock()
registry.get.return_value = RegisteredServeBackend("omni", backend)
mock_registry_factory.return_value = registry
serve(None, ["--model-type", "omni", "--help"])
request = backend.run.call_args.args[0]
self.assertEqual(request.argv, ("--help",))
self.assertIsNone(request.model_path)
mock_load_plugins.assert_not_called()
mock_kill.assert_not_called()
if __name__ == "__main__":
unittest.main()