From c356ed03dd8f2c6ac8e4246eebacbf35a8ad7622 Mon Sep 17 00:00:00 2001 From: Simo Lin Date: Wed, 7 Jan 2026 07:50:28 -0800 Subject: [PATCH] refactor(e2e): unify RouterInstance into Gateway class, split conftest.py into modular fixtures (#16671) --- sgl-model-gateway/e2e_test/backends.py | 224 ++-- sgl-model-gateway/e2e_test/conftest.py | 1015 ++--------------- .../e2e_test/fixtures/__init__.py | 48 +- sgl-model-gateway/e2e_test/fixtures/hooks.py | 358 ++++++ .../e2e_test/fixtures/markers.py | 57 + sgl-model-gateway/e2e_test/fixtures/pool.py | 194 ++++ sgl-model-gateway/e2e_test/fixtures/ports.py | 6 + .../e2e_test/fixtures/router_manager.py | 6 + .../e2e_test/fixtures/setup_backend.py | 394 +++++++ sgl-model-gateway/e2e_test/infra/gateway.py | 69 +- 10 files changed, 1281 insertions(+), 1090 deletions(-) create mode 100644 sgl-model-gateway/e2e_test/fixtures/hooks.py create mode 100644 sgl-model-gateway/e2e_test/fixtures/markers.py create mode 100644 sgl-model-gateway/e2e_test/fixtures/pool.py create mode 100644 sgl-model-gateway/e2e_test/fixtures/setup_backend.py diff --git a/sgl-model-gateway/e2e_test/backends.py b/sgl-model-gateway/e2e_test/backends.py index 926f0b5ad..e3f0319dd 100644 --- a/sgl-model-gateway/e2e_test/backends.py +++ b/sgl-model-gateway/e2e_test/backends.py @@ -1,157 +1,100 @@ -"""Cloud backend configurations for E2E tests. +"""Cloud runtime configurations for E2E tests. -This module handles cloud API backends (OpenAI, xAI) that don't need local GPU workers. -For local backends (gRPC, HTTP), use ModelPool from infra/ to launch workers, -then launch the router separately pointing to those workers. +This module handles cloud API runtimes (OpenAI, xAI) that don't need local GPU workers. +For local runtimes (gRPC, HTTP), use ModelPool from infra/ to launch workers, +then launch the gateway separately pointing to those workers. + +Cloud runtimes vs History backends: +- Cloud runtimes: Where models run (openai, xai) +- History backends: Gateway plugin for conversation storage (memory, oracle) +These are orthogonal - any cloud runtime can use any history backend. """ from __future__ import annotations import logging -import os -import subprocess -from dataclasses import dataclass from typing import Any -from infra import get_open_port, kill_process_tree, wait_for_health +from infra import Gateway logger = logging.getLogger(__name__) -@dataclass -class RouterInstance: - """A running router instance (for cloud backends).""" - - base_url: str - router_process: subprocess.Popen - backend: str - - def shutdown(self) -> None: - """Shutdown the router.""" - if self.router_process.poll() is None: - kill_process_tree(self.router_process.pid) - - -def launch_cloud_router( - backend: str, # "openai" or "xai" - *, - history_backend: str = "memory", - router_args: list[str] | None = None, - timeout: float = 60, - show_output: bool | None = None, -) -> RouterInstance: - """Launch router with cloud API backend (OpenAI/xAI). - - Args: - backend: "openai" or "xai" - history_backend: "memory" or "oracle" - router_args: Additional router arguments - timeout: Startup timeout in seconds - show_output: Show subprocess output - - Returns: - RouterInstance with running router - """ - if show_output is None: - show_output = os.environ.get("SHOW_ROUTER_LOGS", "0") == "1" - - router_port = get_open_port() - prometheus_port = get_open_port() - base_url = f"http://127.0.0.1:{router_port}" - - # Get API key and worker URL - if backend == "openai": - worker_url = "https://api.openai.com" - api_key = os.environ.get("OPENAI_API_KEY") - if not api_key: - raise ValueError("OPENAI_API_KEY environment variable required") - elif backend == "xai": - worker_url = "https://api.x.ai" - api_key = os.environ.get("XAI_API_KEY") - if not api_key: - raise ValueError("XAI_API_KEY environment variable required") - else: - raise ValueError(f"Unsupported cloud backend: {backend}") - - logger.info("Launching %s router on port %d", backend, router_port) - - cmd = [ - "python3", - "-m", - "sglang_router.launch_router", - "--host", - "127.0.0.1", - "--port", - str(router_port), - "--prometheus-port", - str(prometheus_port), - "--backend", - "openai", - "--worker-urls", - worker_url, - "--history-backend", - history_backend, - "--log-level", - "warn", - ] - - if router_args: - cmd.extend(router_args) - - env = os.environ.copy() - if backend == "openai": - env["OPENAI_API_KEY"] = api_key - else: - env["XAI_API_KEY"] = api_key - - router_proc = subprocess.Popen( - cmd, - env=env, - stdout=None if show_output else subprocess.PIPE, - stderr=None if show_output else subprocess.PIPE, - start_new_session=True, - ) - - try: - wait_for_health(base_url, timeout=timeout) - except TimeoutError: - kill_process_tree(router_proc.pid) - raise - - logger.info("%s router ready at %s", backend, base_url) - - return RouterInstance( - base_url=base_url, - router_process=router_proc, - backend=backend, - ) - - -# Cloud backend configurations -CLOUD_BACKENDS: dict[str, dict[str, Any]] = { +# Cloud runtime configurations (where models run) +CLOUD_RUNTIMES: dict[str, dict[str, Any]] = { "openai": { - "description": "OpenAI API backend", + "description": "OpenAI API", "model": "gpt-4o-mini", "api_key_env": "OPENAI_API_KEY", - "history_backend": "memory", }, "xai": { - "description": "xAI API backend", + "description": "xAI API", "model": "grok-2-latest", "api_key_env": "XAI_API_KEY", - "history_backend": "memory", }, +} + +# Keep CLOUD_BACKENDS as alias for backward compatibility during migration +# TODO: Remove after e2e_response_api migration +CLOUD_BACKENDS: dict[str, dict[str, Any]] = { + **CLOUD_RUNTIMES, + # Legacy entry for tests that parameterize on history backend "oracle_store": { "description": "OpenAI API with Oracle history backend", "model": "gpt-4o-mini", "api_key_env": "OPENAI_API_KEY", "history_backend": "oracle", + "_runtime": "openai", # Actual runtime to use }, } +def get_cloud_runtime_config(runtime: str) -> dict[str, Any]: + """Get configuration for a cloud runtime.""" + if runtime not in CLOUD_RUNTIMES: + raise KeyError( + f"Unknown cloud runtime: {runtime}. Available: {list(CLOUD_RUNTIMES.keys())}" + ) + return CLOUD_RUNTIMES[runtime] + + +def launch_cloud_gateway( + runtime: str, # "openai" or "xai" + *, + history_backend: str = "memory", + extra_args: list[str] | None = None, + timeout: float = 60, + show_output: bool | None = None, +) -> Gateway: + """Launch gateway with cloud API runtime. + + Args: + runtime: Cloud runtime ("openai" or "xai") + history_backend: History storage backend ("memory" or "oracle") + extra_args: Additional router arguments + timeout: Startup timeout in seconds + show_output: Show subprocess output + + Returns: + Gateway instance with running router + """ + if runtime not in CLOUD_RUNTIMES: + raise ValueError(f"Unknown cloud runtime: {runtime}") + + gateway = Gateway() + gateway.start( + cloud_backend=runtime, + history_backend=history_backend, + timeout=timeout, + show_output=show_output, + extra_args=extra_args, + ) + return gateway + + +# Backward compatibility aliases - TODO: Remove after migration def get_cloud_backend_config(backend: str) -> dict[str, Any]: - """Get configuration for a cloud backend.""" + """Deprecated: Use get_cloud_runtime_config instead.""" if backend not in CLOUD_BACKENDS: raise KeyError( f"Unknown cloud backend: {backend}. Available: {list(CLOUD_BACKENDS.keys())}" @@ -159,28 +102,23 @@ def get_cloud_backend_config(backend: str) -> dict[str, Any]: return CLOUD_BACKENDS[backend] -def launch_cloud_backend(backend: str, **kwargs: Any) -> RouterInstance: - """Launch a cloud backend router. - - Args: - backend: Backend name from CLOUD_BACKENDS - **kwargs: Override launcher kwargs - - Returns: - RouterInstance with running router - """ +def launch_cloud_backend(backend: str, **kwargs: Any) -> Gateway: + """Deprecated: Use launch_cloud_gateway instead.""" cfg = get_cloud_backend_config(backend) - # Determine actual backend type (openai or xai) - if backend == "oracle_store": - actual_backend = "openai" - else: - actual_backend = backend + # Handle legacy oracle_store entry + runtime = cfg.get("_runtime", backend) + history_backend = kwargs.pop( + "history_backend", cfg.get("history_backend", "memory") + ) - history_backend = kwargs.pop("history_backend", cfg["history_backend"]) - - return launch_cloud_router( - actual_backend, + return launch_cloud_gateway( + runtime, history_backend=history_backend, + extra_args=kwargs.pop("router_args", None), **kwargs, ) + + +# Keep old function name as alias +launch_cloud_router = launch_cloud_gateway diff --git a/sgl-model-gateway/e2e_test/conftest.py b/sgl-model-gateway/e2e_test/conftest.py index 98785d7f6..b144ba627 100644 --- a/sgl-model-gateway/e2e_test/conftest.py +++ b/sgl-model-gateway/e2e_test/conftest.py @@ -90,62 +90,14 @@ PD disaggregation mode: from __future__ import annotations import logging -import os -import subprocess import sys from importlib.util import find_spec from pathlib import Path -from typing import TYPE_CHECKING, Any - -import pytest - -if TYPE_CHECKING: - from infra import ModelPool - # --------------------------------------------------------------------------- -# Logging setup (clean output without pytest's "---- live log ----" dividers) +# Path setup (must happen before other imports) # --------------------------------------------------------------------------- -def _setup_logging() -> None: - """Configure clean logging to stdout with timestamps.""" - # Custom format: timestamp [logger] message - fmt = "%(asctime)s.%(msecs)03d [%(name)s] %(message)s" - datefmt = "%H:%M:%S" - # Create handler for stdout - handler = logging.StreamHandler(sys.stdout) - handler.setFormatter(logging.Formatter(fmt, datefmt)) - - # Configure our e2e_test and infra modules for INFO level - for logger_name in ("e2e_test", "infra"): - log = logging.getLogger(logger_name) - log.setLevel(logging.INFO) - log.addHandler(handler) - log.propagate = False # Don't double-log - - # Suppress noisy third-party loggers - for logger_name in ("openai", "httpx", "httpcore", "numexpr"): - logging.getLogger(logger_name).setLevel(logging.WARNING) - - -_setup_logging() - -logger = logging.getLogger(__name__) - - -# --------------------------------------------------------------------------- -# Test visibility hooks -# --------------------------------------------------------------------------- -def pytest_runtest_logstart(nodeid: str, location: tuple) -> None: - """Print clear test header at start of each test.""" - # Extract test name from nodeid (e.g., "test_mmlu.py::TestMMLU::test_mmlu_basic[grpc]") - test_name = nodeid.split("::")[-1] if "::" in nodeid else nodeid - print(f"\n{'=' * LOG_SEPARATOR_WIDTH}") - print(f"TEST: {test_name}") - print(f"{'=' * LOG_SEPARATOR_WIDTH}") - - -# Path setup for imports _ROOT = Path(__file__).resolve().parents[1] # sgl-model-gateway/ _E2E_TEST = Path(__file__).resolve().parent # e2e_test/ _SRC = _ROOT / "bindings" / "python" @@ -160,897 +112,78 @@ _wheel_installed = find_spec("sglang_router.sglang_router_rs") is not None if not _wheel_installed and str(_SRC) not in sys.path: sys.path.insert(0, str(_SRC)) -# Import constants after path setup -from infra import ( - DEFAULT_MODEL, - DEFAULT_ROUTER_TIMEOUT, - ENV_BACKENDS, - ENV_MODEL, - ENV_MODELS, - ENV_SHOW_ROUTER_LOGS, - ENV_SKIP_BACKEND_SETUP, - ENV_SKIP_MODEL_POOL, - ENV_STARTUP_TIMEOUT, - LOCAL_MODES, - LOG_SEPARATOR_WIDTH, - PARAM_MODEL, - PARAM_SETUP_BACKEND, - ConnectionMode, - WorkerIdentity, - WorkerType, + +# --------------------------------------------------------------------------- +# Logging setup (clean output without pytest's "---- live log ----" dividers) +# --------------------------------------------------------------------------- + + +def _setup_logging() -> None: + """Configure clean logging to stdout with timestamps.""" + fmt = "%(asctime)s.%(msecs)03d [%(name)s] %(message)s" + datefmt = "%H:%M:%S" + + handler = logging.StreamHandler(sys.stdout) + handler.setFormatter(logging.Formatter(fmt, datefmt)) + + for logger_name in ("e2e_test", "infra", "fixtures"): + log = logging.getLogger(logger_name) + log.setLevel(logging.INFO) + log.addHandler(handler) + log.propagate = False + + for logger_name in ("openai", "httpx", "httpcore", "numexpr"): + logging.getLogger(logger_name).setLevel(logging.WARNING) + + +_setup_logging() + +logger = logging.getLogger(__name__) + + +# --------------------------------------------------------------------------- +# Test visibility hooks +# --------------------------------------------------------------------------- + + +def pytest_runtest_logstart(nodeid: str, location: tuple) -> None: + """Print clear test header at start of each test.""" + from infra import LOG_SEPARATOR_WIDTH + + test_name = nodeid.split("::")[-1] if "::" in nodeid else nodeid + print(f"\n{'=' * LOG_SEPARATOR_WIDTH}") + print(f"TEST: {test_name}") + print(f"{'=' * LOG_SEPARATOR_WIDTH}") + + +# --------------------------------------------------------------------------- +# Import pytest hooks and fixtures from fixtures/ package +# --------------------------------------------------------------------------- + +# Import fixtures - pytest discovers these by name +# Import hooks - pytest discovers these by name +from fixtures import ( + backend_router, + model_base_url, + model_client, + model_pool, + pytest_collection_finish, + pytest_collection_modifyitems, + pytest_configure, + setup_backend, ) -# --------------------------------------------------------------------------- -# Test collection: scan for required workers -# --------------------------------------------------------------------------- - -# Track max worker counts: (model_id, mode, worker_type) -> max_count -# This unified approach handles regular, prefill, and decode workers the same way -_worker_counts: dict[tuple[str, ConnectionMode, WorkerType], int] = {} - -# Track first-seen order to preserve test collection order -_first_seen_order: list[tuple[str, ConnectionMode, WorkerType]] = [] - -# Track max GPU requirement for any single test (for validation) -_max_test_gpu_requirement: int = 0 -_max_test_name: str = "" - -_needs_default_model: bool = False # True if any e2e test lacks explicit model marker - - -def pytest_collection_modifyitems( - session: pytest.Session, - config: pytest.Config, - items: list[pytest.Item], -) -> None: - """Scan collected tests to determine required workers. - - This runs after test collection but before tests execute. - It extracts worker requirements from markers in test collection order, - tracking the max count needed for each (model, mode, worker_type) combination. - - Also tracks the max GPU requirement for any single test for validation. - """ - global _worker_counts, _first_seen_order, _needs_default_model - global _max_test_gpu_requirement, _max_test_name - - from infra import MODEL_SPECS - - def track_worker( - model_id: str, mode: ConnectionMode, worker_type: WorkerType, count: int - ) -> None: - """Track a worker requirement, updating max count if needed.""" - key = (model_id, mode, worker_type) - if key not in _worker_counts: - _first_seen_order.append(key) - _worker_counts[key] = count - else: - _worker_counts[key] = max(_worker_counts[key], count) - - def calculate_test_gpus( - model_id: str, prefill: int, decode: int, regular: int - ) -> int: - """Calculate GPU requirement for a single test.""" - if model_id not in MODEL_SPECS: - return 0 - tp = MODEL_SPECS[model_id].get("tp", 1) - return tp * (prefill + decode + regular) - - for item in items: - # Extract model from marker or use default - model_marker = item.get_closest_marker(PARAM_MODEL) - model_id = model_marker.args[0] if model_marker and model_marker.args else None - - # Check parametrize for model - if model_id is None: - for marker in item.iter_markers("parametrize"): - if marker.args and len(marker.args) >= 2: - param_name = marker.args[0] - if param_name == PARAM_MODEL or PARAM_MODEL in param_name: - param_values = marker.args[1] - if isinstance(param_values, (list, tuple)) and param_values: - model_id = param_values[0] # First model in parametrize - break - - # Extract backends from parametrize - backends: list[str] = [] - for marker in item.iter_markers("parametrize"): - if marker.args and len(marker.args) >= 2: - param_name = marker.args[0] - param_values = marker.args[1] - if param_name == PARAM_SETUP_BACKEND: - if isinstance(param_values, (list, tuple)): - backends.extend(param_values) - - # Check for workers marker (@pytest.mark.workers(...)) - workers_marker = item.get_closest_marker("workers") - prefill_count = 0 - decode_count = 0 - regular_count = 1 # Default to 1 regular worker - if workers_marker: - prefill_count = workers_marker.kwargs.get("prefill") or 0 - decode_count = workers_marker.kwargs.get("decode") or 0 - regular_count = workers_marker.kwargs.get("count") or 1 - - # Track if this test needs default model - is_e2e = item.get_closest_marker("e2e") is not None - if model_id is None and is_e2e: - _needs_default_model = True - model_id = DEFAULT_MODEL - - # Track worker requirements and calculate this test's GPU requirement - test_gpus = 0 - if model_id and backends: - for backend in backends: - # "pd" backend means PD workers - if backend == "pd": - mode = ConnectionMode.HTTP # PD uses HTTP mode - # Default to 1 prefill + 1 decode if not specified - p_count = prefill_count if prefill_count > 0 else 1 - d_count = decode_count if decode_count > 0 else 1 - track_worker(model_id, mode, WorkerType.PREFILL, p_count) - track_worker(model_id, mode, WorkerType.DECODE, d_count) - test_gpus = max( - test_gpus, calculate_test_gpus(model_id, p_count, d_count, 0) - ) - else: - try: - mode = ConnectionMode(backend) - except ValueError: - # Cloud backend (openai, xai, etc.) - skip - continue - - # Check if this backend also has PD workers - if prefill_count > 0 or decode_count > 0: - track_worker(model_id, mode, WorkerType.PREFILL, prefill_count) - track_worker(model_id, mode, WorkerType.DECODE, decode_count) - test_gpus = max( - test_gpus, - calculate_test_gpus( - model_id, prefill_count, decode_count, 0 - ), - ) - else: - # Regular worker - track_worker(model_id, mode, WorkerType.REGULAR, regular_count) - test_gpus = max( - test_gpus, - calculate_test_gpus(model_id, 0, 0, regular_count), - ) - - elif model_id and is_e2e: - # E2E test without explicit backend - will use HTTP by default - track_worker(model_id, ConnectionMode.HTTP, WorkerType.REGULAR, 1) - test_gpus = calculate_test_gpus(model_id, 0, 0, 1) - - # Track max GPU requirement across all tests - if test_gpus > _max_test_gpu_requirement: - _max_test_gpu_requirement = test_gpus - _max_test_name = item.nodeid - - # Log results - if _worker_counts: - summary = [] - for key in _first_seen_order: - model_id, mode, worker_type = key - count = _worker_counts[key] - if worker_type == WorkerType.REGULAR: - summary.append(f"{model_id}:{mode.value}x{count}") - else: - summary.append(f"{model_id}:{mode.value}:{worker_type.value}x{count}") - logger.info("Scanned worker requirements (in test order): %s", summary) - logger.info( - "Max GPU requirement for single test: %d (%s)", - _max_test_gpu_requirement, - _max_test_name, - ) - else: - logger.info("Scanned worker requirements: (none)") - - -def get_pool_requirements() -> list[WorkerIdentity]: - """Build pool requirements from scanned test markers. - - Returns: - List of WorkerIdentity objects to pre-launch. - Each WorkerIdentity has (model_id, mode, worker_type, index). - Requirements are ordered by first appearance in test collection order, - so workers needed by earlier tests are launched first. - - Note: - If a model's first test needs PD workers (prefill/decode), we skip - pre-launching regular workers for that model (they'd be evicted - immediately when PD workers are launched). - """ - # Track which models have PD workers as their first requirement - # These models shouldn't have regular workers pre-launched - models_with_pd_first: set[str] = set() - first_worker_type_per_model: dict[str, WorkerType] = {} - - for model_id, mode, worker_type in _first_seen_order: - if model_id not in first_worker_type_per_model: - first_worker_type_per_model[model_id] = worker_type - if worker_type in (WorkerType.PREFILL, WorkerType.DECODE): - models_with_pd_first.add(model_id) - logger.info( - "Model %s has PD test first - skipping regular worker pre-launch", - model_id, - ) - - # Generate individual WorkerIdentity objects in first-seen order - requirements: list[WorkerIdentity] = [] - for model_id, mode, worker_type in _first_seen_order: - # Skip regular workers for models that have PD first - if model_id in models_with_pd_first and worker_type == WorkerType.REGULAR: - continue - - count = _worker_counts.get((model_id, mode, worker_type), 1) - for i in range(count): - requirements.append(WorkerIdentity(model_id, mode, worker_type, i)) - - # Add default if no requirements - if not requirements: - requirements.append(WorkerIdentity(DEFAULT_MODEL, ConnectionMode.HTTP)) - - return requirements - - -def validate_gpu_requirements() -> tuple[int, int]: - """Check if there are enough GPUs for any single test. - - Returns: - Tuple of (max_required_gpus, available_gpus). - - Note: - We check the max requirement for any single test, not the sum. - Workers can be evicted between tests, so we only need enough GPUs - for the most demanding test. - """ - # Count available GPUs - available_gpus = 0 - try: - import torch - - if torch.cuda.is_available(): - available_gpus = torch.cuda.device_count() - except ImportError: - pass - - return _max_test_gpu_requirement, available_gpus - - -def pytest_collection_finish(session: pytest.Session) -> None: - """Validate GPU requirements after test collection. - - This runs after all tests are collected but before any tests execute. - Fails fast if any single test requires more GPUs than available. - """ - if not _worker_counts: - return - - # Skip validation if model pool is disabled - if os.environ.get(ENV_SKIP_MODEL_POOL, "").lower() in ("1", "true", "yes"): - return - - max_required, available_gpus = validate_gpu_requirements() - - if max_required > available_gpus: - sep = "=" * LOG_SEPARATOR_WIDTH - raise pytest.UsageError( - f"\n{sep}\n" - f"GPU REQUIREMENTS EXCEEDED\n" - f"{sep}\n" - f"Test '{_max_test_name}' requires {max_required} GPUs\n" - f"Available: {available_gpus} GPUs\n" - f"\nOptions:\n" - f" 1. Run tests that fit: pytest -k 'not {_max_test_name.split('::')[0]}'\n" - f" 2. Reduce workers: @pytest.mark.workers(prefill=1, decode=1)\n" - f" 3. Skip GPU tests: SKIP_MODEL_POOL=1 pytest\n" - f"{sep}" - ) - - logger.info( - "GPU validation passed: max %d required (by %s), %d available", - max_required, - _max_test_name, - available_gpus, - ) - - -# --------------------------------------------------------------------------- -# Custom pytest markers -# --------------------------------------------------------------------------- - - -def pytest_configure(config: pytest.Config) -> None: - """Register custom markers.""" - config.addinivalue_line( - "markers", - "model(name): mark test to use a specific model from MODEL_SPECS", - ) - config.addinivalue_line( - "markers", - "backend(name): mark test to use a specific backend (grpc, http, openai, etc.)", - ) - config.addinivalue_line( - "markers", - "workers(count=1, prefill=None, decode=None): " - "worker configuration - use count for regular workers, " - "or prefill/decode for PD disaggregation mode", - ) - config.addinivalue_line( - "markers", - "gateway(policy='round_robin', timeout=None, extra_args=None): " - "gateway/router configuration", - ) - config.addinivalue_line( - "markers", - "e2e: mark test as an end-to-end test requiring GPU workers", - ) - config.addinivalue_line( - "markers", - "slow: mark test as slow-running", - ) - - -# --------------------------------------------------------------------------- -# Model pool fixtures (session-scoped) -# --------------------------------------------------------------------------- - -# Global model pool instance -_model_pool: "ModelPool | None" = None - - -@pytest.fixture(scope="session") -def model_pool(request: pytest.FixtureRequest) -> "ModelPool": - """Session-scoped fixture that manages SGLang worker processes. - - Workers (sglang.launch_server) are expensive to start (~30-60s each due to - model loading). This fixture starts them ONCE per session and keeps them - running across all tests. The setup_backend fixture then launches cheap - routers (~1-2s) pointing to these workers. - - Startup behavior: - - Scans test markers to determine required workers (model, mode, type, count) - - Launches workers in test collection order - - Waits for all workers to become healthy before returning - - Test requirements are auto-detected from: - - @pytest.mark.parametrize("setup_backend", ["grpc", "http", "pd"]) - - @pytest.mark.model("model-name") - - @pytest.mark.workers(count=N) for regular workers - - @pytest.mark.workers(prefill=N, decode=N) for PD workers - - Environment variable overrides: - - E2E_MODELS: Comma-separated model IDs (e.g., "llama-8b,qwen-7b") - - E2E_BACKENDS: Comma-separated backends (e.g., "grpc,http") - - SKIP_MODEL_POOL: Set to "1" to skip worker startup - """ - global _model_pool - - from infra import MODEL_SPECS, GPUAllocator, ModelPool - - if _model_pool is not None: - return _model_pool - - # Check if we should skip model startup - if os.environ.get(ENV_SKIP_MODEL_POOL, "").lower() in ("1", "true", "yes"): - logger.info("%s is set, skipping model pool startup", ENV_SKIP_MODEL_POOL) - _model_pool = ModelPool(GPUAllocator(gpus=[])) - return _model_pool - - # Determine requirements from scanned tests or env vars - models_env = os.environ.get(ENV_MODELS, "") - backends_env = os.environ.get(ENV_BACKENDS, "") - - if models_env or backends_env: - # Use env var overrides - models = ( - {m.strip() for m in models_env.split(",") if m.strip()} - if models_env - else {DEFAULT_MODEL} - ) - - # Parse backend strings to ConnectionMode enums - backend_modes: set[ConnectionMode] = set() - if backends_env: - for b in backends_env.split(","): - b = b.strip() - if b: - try: - mode = ConnectionMode(b) - if mode in LOCAL_MODES: - backend_modes.add(mode) - except ValueError: - logger.warning("Unknown backend '%s', skipping", b) - - # Default to HTTP if no valid backends - if not backend_modes: - backend_modes = {ConnectionMode.HTTP} - - # Create WorkerIdentity objects (regular workers only from env vars) - requirements = [ - WorkerIdentity(m, b, WorkerType.REGULAR, 0) - for m in models - for b in backend_modes - ] - logger.info("Using env var requirements: %s", [str(r) for r in requirements]) - else: - # Use scanned requirements from test markers - requirements = get_pool_requirements() - logger.info("Using scanned requirements: %s", [str(r) for r in requirements]) - - # Filter to valid models - requirements = [r for r in requirements if r.model_id in MODEL_SPECS] - - if not requirements: - logger.warning("No valid requirements, model pool will be empty") - _model_pool = ModelPool(GPUAllocator(gpus=[])) - return _model_pool - - # Create and start the pool - allocator = GPUAllocator() - _model_pool = ModelPool(allocator) - - startup_timeout = int(os.environ.get(ENV_STARTUP_TIMEOUT, "300")) - _model_pool.startup( - requirements=requirements, - startup_timeout=startup_timeout, - ) - - # Log final GPU allocation summary - logger.info(_model_pool.allocator.summary()) - - # Register cleanup - request.addfinalizer(_model_pool.shutdown) - - return _model_pool - - -@pytest.fixture -def model_client(request: pytest.FixtureRequest, model_pool: "ModelPool"): - """Get OpenAI client for the model specified by @pytest.mark.model(). - - Usage: - @pytest.mark.model("llama-8b") - def test_chat(model_client): - response = model_client.chat.completions.create(...) - """ - marker = request.node.get_closest_marker(PARAM_MODEL) - if marker is None: - pytest.fail( - f"Test must be marked with @pytest.mark.{PARAM_MODEL}('model-id') to use model_client fixture" - ) - - model_id = marker.args[0] - - try: - return model_pool.get_client(model_id) - except KeyError: - pytest.skip(f"Model {model_id} not available in model pool") - - -@pytest.fixture -def model_base_url(request: pytest.FixtureRequest, model_pool: "ModelPool") -> str: - """Get the base URL for the model specified by @pytest.mark.model(). - - Usage: - @pytest.mark.model("llama-8b") - def test_direct_http(model_base_url): - response = httpx.get(f"{model_base_url}/health") - """ - marker = request.node.get_closest_marker(PARAM_MODEL) - if marker is None: - pytest.fail( - f"Test must be marked with @pytest.mark.{PARAM_MODEL}('model-id') to use model_base_url fixture" - ) - - model_id = marker.args[0] - - try: - return model_pool.get_base_url(model_id) - except KeyError: - pytest.skip(f"Model {model_id} not available in model pool") - - -# --------------------------------------------------------------------------- -# Backend fixtures -# --------------------------------------------------------------------------- - - -def _get_marker_value( - request: pytest.FixtureRequest, - marker_name: str, - arg_index: int = 0, - default: Any = None, -) -> Any: - """Get a value from a pytest marker. - - Args: - request: The pytest fixture request. - marker_name: Name of the marker to look for. - arg_index: Index of positional argument to extract. - default: Default value if marker not found. - - Returns: - The marker argument value or default. - """ - marker = request.node.get_closest_marker(marker_name) - if marker is None: - return default - if marker.args and len(marker.args) > arg_index: - return marker.args[arg_index] - return default - - -def _get_marker_kwargs( - request: pytest.FixtureRequest, - marker_name: str, - defaults: dict[str, Any] | None = None, -) -> dict[str, Any]: - """Get keyword arguments from a pytest marker. - - Args: - request: The pytest fixture request. - marker_name: Name of the marker to look for. - defaults: Default values if marker not found or missing kwargs. - - Returns: - Dict of keyword arguments merged with defaults. - """ - result = dict(defaults) if defaults else {} - marker = request.node.get_closest_marker(marker_name) - if marker is not None: - result.update(marker.kwargs) - return result - - -@pytest.fixture(scope="class") -def setup_backend(request: pytest.FixtureRequest, model_pool: "ModelPool"): - """Class-scoped fixture that launches a router for each test class. - - Routers are cheap to start (~1-2s) compared to workers (~30-60s), so we - launch a fresh router per test class for isolation while reusing the - expensive workers from model_pool. - - Backend types: - - "http", "grpc": Gets existing worker from model_pool, launches router - - "pd": Launches prefill/decode workers via model_pool, launches PD router - - "openai", "xai", etc.: Launches cloud router (no local workers) - - Configuration via markers: - - @pytest.mark.model("model-id"): Override default model - - @pytest.mark.workers(count=1): Number of regular workers behind router - - @pytest.mark.workers(prefill=1, decode=1): PD worker configuration - - @pytest.mark.gateway(policy="round_robin", timeout=60): Gateway configuration - - Returns: - Tuple of (backend_name, model_path, openai_client, gateway) - - Usage: - # Simple - uses defaults - @pytest.mark.parametrize("setup_backend", ["http"], indirect=True) - class TestBasic: - ... - - # With model override - @pytest.mark.model("qwen-7b") - @pytest.mark.parametrize("setup_backend", ["http"], indirect=True) - class TestWithModel: - ... - - # Load balancing with multiple workers - @pytest.mark.workers(count=3) - @pytest.mark.gateway(policy="round_robin") - @pytest.mark.parametrize("setup_backend", ["http"], indirect=True) - class TestLoadBalancing: - ... - - # PD with custom configuration - @pytest.mark.workers(prefill=2, decode=2) - @pytest.mark.gateway(policy="round_robin") - @pytest.mark.parametrize("setup_backend", ["pd"], indirect=True) - class TestPDScaling: - ... - """ - import openai - from infra import DEFAULT_ROUTER_TIMEOUT, Gateway, WorkerType - - backend_name = request.param - - # Skip if requested - if os.environ.get(ENV_SKIP_BACKEND_SETUP, "").lower() in ("1", "true", "yes"): - pytest.skip(f"{ENV_SKIP_BACKEND_SETUP} is set") - - # Get model from marker or env var or default - model_id = _get_marker_value(request, "model") - if model_id is None: - model_id = os.environ.get(ENV_MODEL, DEFAULT_MODEL) - - # Get worker configuration from marker - workers_config = _get_marker_kwargs( - request, "workers", defaults={"count": 1, "prefill": None, "decode": None} - ) - - # Get gateway configuration from marker - gateway_config = _get_marker_kwargs( - request, - "gateway", - defaults={ - "policy": "round_robin", - "timeout": DEFAULT_ROUTER_TIMEOUT, - "extra_args": None, - }, - ) - - # PD disaggregation backend - if backend_name == "pd": - # Check PD requirements - try: - import sgl_kernel # noqa: F401 - except ImportError: - pytest.skip("sgl_kernel not available, required for PD disaggregation") - - try: - import torch - except ImportError: - pytest.skip("torch not available") - - if not torch.cuda.is_available(): - pytest.skip("CUDA not available") - - # Get PD configuration from workers marker - num_prefill = workers_config.get("prefill") or 1 - num_decode = workers_config.get("decode") or 1 - - # Check GPU requirements - required_gpus = num_prefill + num_decode - gpu_count = torch.cuda.device_count() - if gpu_count < required_gpus: - pytest.skip( - f"PD tests require {required_gpus} GPUs " - f"({num_prefill} prefill + {num_decode} decode), found {gpu_count}" - ) - - # Try to use pre-launched PD workers, or launch additional ones if needed - existing_prefills = model_pool.get_workers_by_type(model_id, WorkerType.PREFILL) - existing_decodes = model_pool.get_workers_by_type(model_id, WorkerType.DECODE) - - # Calculate how many more we need (if any) - missing_prefill = max(0, num_prefill - len(existing_prefills)) - missing_decode = max(0, num_decode - len(existing_decodes)) - - if missing_prefill == 0 and missing_decode == 0: - # Use pre-launched workers (we have enough) - prefills = existing_prefills[:num_prefill] - decodes = existing_decodes[:num_decode] - logger.info( - "Using pre-launched PD workers: %d prefill, %d decode", - len(prefills), - len(decodes), - ) - else: - # Build WorkerIdentity list for missing workers - workers_to_launch: list[WorkerIdentity] = [] - for i in range(missing_prefill): - workers_to_launch.append( - WorkerIdentity( - model_id, - ConnectionMode.HTTP, - WorkerType.PREFILL, - len(existing_prefills) + i, - ) - ) - for i in range(missing_decode): - workers_to_launch.append( - WorkerIdentity( - model_id, - ConnectionMode.HTTP, - WorkerType.DECODE, - len(existing_decodes) + i, - ) - ) - - logger.info( - "Have %d/%d prefill, %d/%d decode. Launching %d more workers", - len(existing_prefills), - num_prefill, - len(existing_decodes), - num_decode, - len(workers_to_launch), - ) - new_instances = model_pool.launch_workers( - workers_to_launch, startup_timeout=300 - ) - - # Combine existing + newly launched - new_prefills = [ - w for w in new_instances if w.worker_type == WorkerType.PREFILL - ] - new_decodes = [ - w for w in new_instances if w.worker_type == WorkerType.DECODE - ] - prefills = existing_prefills + new_prefills - decodes = existing_decodes + new_decodes - - model_path = prefills[0].model_path if prefills else None - - # Launch PD gateway with configuration - gateway = Gateway() - gateway.start( - prefill_workers=prefills, - decode_workers=decodes, - policy=gateway_config["policy"], - timeout=gateway_config["timeout"], - extra_args=gateway_config["extra_args"], - ) - - client = openai.OpenAI( - base_url=f"{gateway.base_url}/v1", - api_key="not-used", - ) - - logger.info( - "Setup PD backend: model=%s, %d prefill + %d decode workers, " - "gateway=%s, policy=%s", - model_id, - len(prefills), - len(decodes), - gateway.base_url, - gateway_config["policy"], - ) - - try: - yield backend_name, model_path, client, gateway - finally: - logger.info("Tearing down PD gateway") - gateway.shutdown() - return - - # Check if this is a local backend (grpc, http) - try: - connection_mode = ConnectionMode(backend_name) - is_local = connection_mode in LOCAL_MODES - except ValueError: - is_local = False - connection_mode = None - - # Local backends: use worker from pool + launch gateway - if is_local: - # Get number of workers from marker - num_workers = workers_config.get("count") or 1 - - try: - if num_workers > 1: - # Check existing workers - existing = model_pool.get_workers_by_type(model_id, WorkerType.REGULAR) - existing_for_mode = [w for w in existing if w.mode == connection_mode] - - if len(existing_for_mode) >= num_workers: - instances = existing_for_mode[:num_workers] - else: - # Launch missing workers - missing = num_workers - len(existing_for_mode) - workers_to_launch = [ - WorkerIdentity( - model_id, - connection_mode, - WorkerType.REGULAR, - len(existing_for_mode) + i, - ) - for i in range(missing) - ] - new_instances = model_pool.launch_workers( - workers_to_launch, startup_timeout=300 - ) - instances = existing_for_mode + new_instances - - if not instances: - pytest.fail(f"Failed to get {num_workers} workers for {model_id}") - worker_urls = [inst.worker_url for inst in instances] - model_path = instances[0].model_path - else: - # Single worker - use existing get() method - instance = model_pool.get(model_id, connection_mode) - worker_urls = [instance.worker_url] - model_path = instance.model_path - except RuntimeError as e: - pytest.fail(str(e)) - - # Launch gateway with configuration - gateway = Gateway() - gateway.start( - worker_urls=worker_urls, - model_path=model_path, - policy=gateway_config["policy"], - timeout=gateway_config["timeout"], - extra_args=gateway_config["extra_args"], - ) - - client = openai.OpenAI( - base_url=f"{gateway.base_url}/v1", - api_key="not-used", - ) - - logger.info( - "Setup %s backend: model=%s, workers=%d, gateway=%s, policy=%s", - backend_name, - model_id, - num_workers, - gateway.base_url, - gateway_config["policy"], - ) - - try: - yield backend_name, model_path, client, gateway - finally: - logger.info("Tearing down gateway for %s backend", backend_name) - gateway.shutdown() - return - - # Cloud backends: launch cloud router - from backends import CLOUD_BACKENDS, launch_cloud_backend - - if backend_name not in CLOUD_BACKENDS: - pytest.fail(f"Unknown backend: {backend_name}") - - cfg = CLOUD_BACKENDS[backend_name] - api_key_env = cfg.get("api_key_env") - - if api_key_env and not os.environ.get(api_key_env): - pytest.skip(f"{api_key_env} not set, skipping {backend_name} tests") - - logger.info("Launching cloud backend: %s", backend_name) - router = launch_cloud_backend(backend_name) - - api_key = os.environ.get(api_key_env) if api_key_env else "not-used" - client = openai.OpenAI( - base_url=f"{router.base_url}/v1", - api_key=api_key, - ) - - try: - yield backend_name, cfg["model"], client - finally: - logger.info("Tearing down cloud backend: %s", backend_name) - router.shutdown() - - -@pytest.fixture -def backend_router(request: pytest.FixtureRequest, model_pool: "ModelPool"): - """Function-scoped fixture for launching a fresh router per test. - - This launches a new Gateway for each test, pointing to workers from the pool. - Use for tests that need isolated router state. - - Usage: - @pytest.mark.parametrize("backend_router", ["grpc", "http"], indirect=True) - def test_router_state(backend_router): - gateway = backend_router - # Test gateway-specific behavior - """ - from infra import Gateway - - backend_name = request.param - model_id = os.environ.get(ENV_MODEL, DEFAULT_MODEL) - - # Convert string to ConnectionMode - connection_mode = ConnectionMode(backend_name) - - try: - instance = model_pool.get(model_id, connection_mode) - except KeyError: - pytest.skip(f"Model {model_id}:{backend_name} not available in pool") - except RuntimeError as e: - pytest.fail(str(e)) - - gateway = Gateway() - gateway.start( - worker_urls=[instance.worker_url], - model_path=instance.model_path, - ) - - try: - yield gateway - finally: - gateway.shutdown() +# Re-export for pytest discovery +__all__ = [ + # Hooks + "pytest_runtest_logstart", + "pytest_collection_modifyitems", + "pytest_collection_finish", + "pytest_configure", + # Fixtures + "model_pool", + "model_client", + "model_base_url", + "setup_backend", + "backend_router", +] diff --git a/sgl-model-gateway/e2e_test/fixtures/__init__.py b/sgl-model-gateway/e2e_test/fixtures/__init__.py index 4ac754df8..b391945ab 100644 --- a/sgl-model-gateway/e2e_test/fixtures/__init__.py +++ b/sgl-model-gateway/e2e_test/fixtures/__init__.py @@ -1 +1,47 @@ -"""Shared fixtures for router integration tests.""" +"""Fixtures for E2E tests. + +This package contains modular pytest fixtures split by responsibility: +- hooks.py: Pytest collection hooks and marker registration +- pool.py: Model pool fixtures (session-scoped worker management) +- setup_backend.py: Backend setup fixtures (class/function-scoped) +- markers.py: Helper utilities for marker extraction + +Legacy modules (to be removed during e2e_response_api migration): +- ports.py: Use infra.get_open_port() instead +- router_manager.py: Use infra.Gateway instead +""" + +# Pytest hooks (imported by conftest.py via pytest_plugins) +from .hooks import ( + get_pool_requirements, + pytest_collection_finish, + pytest_collection_modifyitems, + pytest_configure, + validate_gpu_requirements, +) + +# Marker helpers +from .markers import get_marker_kwargs, get_marker_value + +# Fixtures (imported by conftest.py) +from .pool import model_base_url, model_client, model_pool +from .setup_backend import backend_router, setup_backend + +__all__ = [ + # Hooks + "pytest_collection_modifyitems", + "pytest_collection_finish", + "pytest_configure", + "get_pool_requirements", + "validate_gpu_requirements", + # Pool fixtures + "model_pool", + "model_client", + "model_base_url", + # Backend fixtures + "setup_backend", + "backend_router", + # Marker helpers + "get_marker_value", + "get_marker_kwargs", +] diff --git a/sgl-model-gateway/e2e_test/fixtures/hooks.py b/sgl-model-gateway/e2e_test/fixtures/hooks.py new file mode 100644 index 000000000..47cc36e42 --- /dev/null +++ b/sgl-model-gateway/e2e_test/fixtures/hooks.py @@ -0,0 +1,358 @@ +"""Pytest hooks for E2E test collection and validation. + +This module handles: +- Test collection: Scanning markers to determine required workers +- GPU validation: Ensuring sufficient GPUs for test requirements +- Marker registration: Defining custom pytest markers +""" + +from __future__ import annotations + +import logging +import os +from typing import TYPE_CHECKING + +import pytest + +if TYPE_CHECKING: + from infra import ConnectionMode, WorkerIdentity, WorkerType + +logger = logging.getLogger(__name__) + +# --------------------------------------------------------------------------- +# Test collection state +# --------------------------------------------------------------------------- + +# Track max worker counts: (model_id, mode, worker_type) -> max_count +_worker_counts: dict[tuple["ConnectionMode", "WorkerType"], int] = {} + +# Track first-seen order to preserve test collection order +_first_seen_order: list[tuple[str, "ConnectionMode", "WorkerType"]] = [] + +# Track max GPU requirement for any single test (for validation) +_max_test_gpu_requirement: int = 0 +_max_test_name: str = "" + +_needs_default_model: bool = False + + +def reset_collection_state() -> None: + """Reset collection state (useful for testing).""" + global _worker_counts, _first_seen_order + global _max_test_gpu_requirement, _max_test_name, _needs_default_model + _worker_counts = {} + _first_seen_order = [] + _max_test_gpu_requirement = 0 + _max_test_name = "" + _needs_default_model = False + + +def get_worker_counts() -> dict: + """Get the worker counts dictionary.""" + return _worker_counts + + +def get_first_seen_order() -> list: + """Get the first-seen order list.""" + return _first_seen_order + + +def get_max_gpu_requirement() -> tuple[int, str]: + """Get the max GPU requirement and test name.""" + return _max_test_gpu_requirement, _max_test_name + + +def needs_default_model() -> bool: + """Check if any test needs the default model.""" + return _needs_default_model + + +# --------------------------------------------------------------------------- +# Test collection hook +# --------------------------------------------------------------------------- + + +def pytest_collection_modifyitems( + session: pytest.Session, + config: pytest.Config, + items: list[pytest.Item], +) -> None: + """Scan collected tests to determine required workers. + + This runs after test collection but before tests execute. + It extracts worker requirements from markers in test collection order, + tracking the max count needed for each (model, mode, worker_type) combination. + """ + global _worker_counts, _first_seen_order, _needs_default_model + global _max_test_gpu_requirement, _max_test_name + + from infra import ( + DEFAULT_MODEL, + LOG_SEPARATOR_WIDTH, + MODEL_SPECS, + PARAM_MODEL, + PARAM_SETUP_BACKEND, + ConnectionMode, + WorkerType, + ) + + def track_worker( + model_id: str, mode: ConnectionMode, worker_type: WorkerType, count: int + ) -> None: + """Track a worker requirement, updating max count if needed.""" + key = (model_id, mode, worker_type) + if key not in _worker_counts: + _first_seen_order.append(key) + _worker_counts[key] = count + else: + _worker_counts[key] = max(_worker_counts[key], count) + + def calculate_test_gpus( + model_id: str, prefill: int, decode: int, regular: int + ) -> int: + """Calculate GPU requirement for a single test.""" + if model_id not in MODEL_SPECS: + return 0 + tp = MODEL_SPECS[model_id].get("tp", 1) + return tp * (prefill + decode + regular) + + for item in items: + # Extract model from marker or use default + model_marker = item.get_closest_marker(PARAM_MODEL) + model_id = model_marker.args[0] if model_marker and model_marker.args else None + + # Check parametrize for model + if model_id is None: + for marker in item.iter_markers("parametrize"): + if marker.args and len(marker.args) >= 2: + param_name = marker.args[0] + if param_name == PARAM_MODEL or PARAM_MODEL in param_name: + param_values = marker.args[1] + if isinstance(param_values, (list, tuple)) and param_values: + model_id = param_values[0] + break + + # Extract backends from parametrize + backends: list[str] = [] + for marker in item.iter_markers("parametrize"): + if marker.args and len(marker.args) >= 2: + param_name = marker.args[0] + param_values = marker.args[1] + if param_name == PARAM_SETUP_BACKEND: + if isinstance(param_values, (list, tuple)): + backends.extend(param_values) + + # Check for workers marker + workers_marker = item.get_closest_marker("workers") + prefill_count = 0 + decode_count = 0 + regular_count = 1 + if workers_marker: + prefill_count = workers_marker.kwargs.get("prefill") or 0 + decode_count = workers_marker.kwargs.get("decode") or 0 + regular_count = workers_marker.kwargs.get("count") or 1 + + # Track if this test needs default model + is_e2e = item.get_closest_marker("e2e") is not None + if model_id is None and is_e2e: + _needs_default_model = True + model_id = DEFAULT_MODEL + + # Track worker requirements + test_gpus = 0 + if model_id and backends: + for backend in backends: + if backend == "pd": + mode = ConnectionMode.HTTP + p_count = prefill_count if prefill_count > 0 else 1 + d_count = decode_count if decode_count > 0 else 1 + track_worker(model_id, mode, WorkerType.PREFILL, p_count) + track_worker(model_id, mode, WorkerType.DECODE, d_count) + test_gpus = max( + test_gpus, calculate_test_gpus(model_id, p_count, d_count, 0) + ) + else: + try: + mode = ConnectionMode(backend) + except ValueError: + continue + + if prefill_count > 0 or decode_count > 0: + track_worker(model_id, mode, WorkerType.PREFILL, prefill_count) + track_worker(model_id, mode, WorkerType.DECODE, decode_count) + test_gpus = max( + test_gpus, + calculate_test_gpus( + model_id, prefill_count, decode_count, 0 + ), + ) + else: + track_worker(model_id, mode, WorkerType.REGULAR, regular_count) + test_gpus = max( + test_gpus, + calculate_test_gpus(model_id, 0, 0, regular_count), + ) + + elif model_id and is_e2e: + track_worker(model_id, ConnectionMode.HTTP, WorkerType.REGULAR, 1) + test_gpus = calculate_test_gpus(model_id, 0, 0, 1) + + if test_gpus > _max_test_gpu_requirement: + _max_test_gpu_requirement = test_gpus + _max_test_name = item.nodeid + + # Log results + if _worker_counts: + summary = [] + for key in _first_seen_order: + model_id, mode, worker_type = key + count = _worker_counts[key] + if worker_type == WorkerType.REGULAR: + summary.append(f"{model_id}:{mode.value}x{count}") + else: + summary.append(f"{model_id}:{mode.value}:{worker_type.value}x{count}") + logger.info("Scanned worker requirements (in test order): %s", summary) + logger.info( + "Max GPU requirement for single test: %d (%s)", + _max_test_gpu_requirement, + _max_test_name, + ) + else: + logger.info("Scanned worker requirements: (none)") + + +# --------------------------------------------------------------------------- +# Pool requirements +# --------------------------------------------------------------------------- + + +def get_pool_requirements() -> list["WorkerIdentity"]: + """Build pool requirements from scanned test markers. + + Returns: + List of WorkerIdentity objects to pre-launch. + """ + from infra import DEFAULT_MODEL, ConnectionMode, WorkerIdentity, WorkerType + + # Track which models have PD workers as their first requirement + models_with_pd_first: set[str] = set() + first_worker_type_per_model: dict[str, WorkerType] = {} + + for model_id, mode, worker_type in _first_seen_order: + if model_id not in first_worker_type_per_model: + first_worker_type_per_model[model_id] = worker_type + if worker_type in (WorkerType.PREFILL, WorkerType.DECODE): + models_with_pd_first.add(model_id) + logger.info( + "Model %s has PD test first - skipping regular worker pre-launch", + model_id, + ) + + # Generate individual WorkerIdentity objects in first-seen order + requirements: list[WorkerIdentity] = [] + for model_id, mode, worker_type in _first_seen_order: + if model_id in models_with_pd_first and worker_type == WorkerType.REGULAR: + continue + + count = _worker_counts.get((model_id, mode, worker_type), 1) + for i in range(count): + requirements.append(WorkerIdentity(model_id, mode, worker_type, i)) + + if not requirements: + requirements.append(WorkerIdentity(DEFAULT_MODEL, ConnectionMode.HTTP)) + + return requirements + + +# --------------------------------------------------------------------------- +# GPU validation +# --------------------------------------------------------------------------- + + +def validate_gpu_requirements() -> tuple[int, int]: + """Check if there are enough GPUs for any single test. + + Returns: + Tuple of (max_required_gpus, available_gpus). + """ + available_gpus = 0 + try: + import torch + + if torch.cuda.is_available(): + available_gpus = torch.cuda.device_count() + except ImportError: + pass + + return _max_test_gpu_requirement, available_gpus + + +def pytest_collection_finish(session: pytest.Session) -> None: + """Validate GPU requirements after test collection.""" + from infra import ENV_SKIP_MODEL_POOL, LOG_SEPARATOR_WIDTH + + if not _worker_counts: + return + + if os.environ.get(ENV_SKIP_MODEL_POOL, "").lower() in ("1", "true", "yes"): + return + + max_required, available_gpus = validate_gpu_requirements() + + if max_required > available_gpus: + sep = "=" * LOG_SEPARATOR_WIDTH + raise pytest.UsageError( + f"\n{sep}\n" + f"GPU REQUIREMENTS EXCEEDED\n" + f"{sep}\n" + f"Test '{_max_test_name}' requires {max_required} GPUs\n" + f"Available: {available_gpus} GPUs\n" + f"\nOptions:\n" + f" 1. Run tests that fit: pytest -k 'not {_max_test_name.split('::')[0]}'\n" + f" 2. Reduce workers: @pytest.mark.workers(prefill=1, decode=1)\n" + f" 3. Skip GPU tests: SKIP_MODEL_POOL=1 pytest\n" + f"{sep}" + ) + + logger.info( + "GPU validation passed: max %d required (by %s), %d available", + max_required, + _max_test_name, + available_gpus, + ) + + +# --------------------------------------------------------------------------- +# Marker registration +# --------------------------------------------------------------------------- + + +def pytest_configure(config: pytest.Config) -> None: + """Register custom markers.""" + config.addinivalue_line( + "markers", + "model(name): mark test to use a specific model from MODEL_SPECS", + ) + config.addinivalue_line( + "markers", + "backend(name): mark test to use a specific backend (grpc, http, openai, etc.)", + ) + config.addinivalue_line( + "markers", + "workers(count=1, prefill=None, decode=None): " + "worker configuration - use count for regular workers, " + "or prefill/decode for PD disaggregation mode", + ) + config.addinivalue_line( + "markers", + "gateway(policy='round_robin', timeout=None, extra_args=None): " + "gateway/router configuration", + ) + config.addinivalue_line( + "markers", + "e2e: mark test as an end-to-end test requiring GPU workers", + ) + config.addinivalue_line( + "markers", + "slow: mark test as slow-running", + ) diff --git a/sgl-model-gateway/e2e_test/fixtures/markers.py b/sgl-model-gateway/e2e_test/fixtures/markers.py new file mode 100644 index 000000000..23c7735d7 --- /dev/null +++ b/sgl-model-gateway/e2e_test/fixtures/markers.py @@ -0,0 +1,57 @@ +"""Marker helper utilities for E2E tests. + +This module provides helper functions for extracting values from pytest markers. +""" + +from __future__ import annotations + +from typing import Any + +import pytest + + +def get_marker_value( + request: pytest.FixtureRequest, + marker_name: str, + arg_index: int = 0, + default: Any = None, +) -> Any: + """Get a value from a pytest marker. + + Args: + request: The pytest fixture request. + marker_name: Name of the marker to look for. + arg_index: Index of positional argument to extract. + default: Default value if marker not found. + + Returns: + The marker argument value or default. + """ + marker = request.node.get_closest_marker(marker_name) + if marker is None: + return default + if marker.args and len(marker.args) > arg_index: + return marker.args[arg_index] + return default + + +def get_marker_kwargs( + request: pytest.FixtureRequest, + marker_name: str, + defaults: dict[str, Any] | None = None, +) -> dict[str, Any]: + """Get keyword arguments from a pytest marker. + + Args: + request: The pytest fixture request. + marker_name: Name of the marker to look for. + defaults: Default values if marker not found or missing kwargs. + + Returns: + Dict of keyword arguments merged with defaults. + """ + result = dict(defaults) if defaults else {} + marker = request.node.get_closest_marker(marker_name) + if marker is not None: + result.update(marker.kwargs) + return result diff --git a/sgl-model-gateway/e2e_test/fixtures/pool.py b/sgl-model-gateway/e2e_test/fixtures/pool.py new file mode 100644 index 000000000..cd640f829 --- /dev/null +++ b/sgl-model-gateway/e2e_test/fixtures/pool.py @@ -0,0 +1,194 @@ +"""Model pool fixtures for E2E tests. + +This module provides session-scoped fixtures for managing SGLang worker processes. +Workers are expensive to start (~30-60s each), so they're kept running across tests. +""" + +from __future__ import annotations + +import logging +import os +from typing import TYPE_CHECKING + +import pytest + +if TYPE_CHECKING: + from infra import ModelPool + +from .hooks import get_pool_requirements + +logger = logging.getLogger(__name__) + +# Global model pool instance +_model_pool: "ModelPool | None" = None + + +@pytest.fixture(scope="session") +def model_pool(request: pytest.FixtureRequest) -> "ModelPool": + """Session-scoped fixture that manages SGLang worker processes. + + Workers (sglang.launch_server) are expensive to start (~30-60s each due to + model loading). This fixture starts them ONCE per session and keeps them + running across all tests. The setup_backend fixture then launches cheap + routers (~1-2s) pointing to these workers. + + Startup behavior: + - Scans test markers to determine required workers (model, mode, type, count) + - Launches workers in test collection order + - Waits for all workers to become healthy before returning + + Test requirements are auto-detected from: + - @pytest.mark.parametrize("setup_backend", ["grpc", "http", "pd"]) + - @pytest.mark.model("model-name") + - @pytest.mark.workers(count=N) for regular workers + - @pytest.mark.workers(prefill=N, decode=N) for PD workers + + Environment variable overrides: + - E2E_MODELS: Comma-separated model IDs (e.g., "llama-8b,qwen-7b") + - E2E_BACKENDS: Comma-separated backends (e.g., "grpc,http") + - SKIP_MODEL_POOL: Set to "1" to skip worker startup + """ + global _model_pool + + from infra import ( + DEFAULT_MODEL, + ENV_BACKENDS, + ENV_MODELS, + ENV_SKIP_MODEL_POOL, + ENV_STARTUP_TIMEOUT, + LOCAL_MODES, + MODEL_SPECS, + ConnectionMode, + GPUAllocator, + ModelPool, + WorkerIdentity, + WorkerType, + ) + + if _model_pool is not None: + return _model_pool + + # Check if we should skip model startup + if os.environ.get(ENV_SKIP_MODEL_POOL, "").lower() in ("1", "true", "yes"): + logger.info("%s is set, skipping model pool startup", ENV_SKIP_MODEL_POOL) + _model_pool = ModelPool(GPUAllocator(gpus=[])) + return _model_pool + + # Determine requirements from scanned tests or env vars + models_env = os.environ.get(ENV_MODELS, "") + backends_env = os.environ.get(ENV_BACKENDS, "") + + if models_env or backends_env: + # Use env var overrides + models = ( + {m.strip() for m in models_env.split(",") if m.strip()} + if models_env + else {DEFAULT_MODEL} + ) + + # Parse backend strings to ConnectionMode enums + backend_modes: set[ConnectionMode] = set() + if backends_env: + for b in backends_env.split(","): + b = b.strip() + if b: + try: + mode = ConnectionMode(b) + if mode in LOCAL_MODES: + backend_modes.add(mode) + except ValueError: + logger.warning("Unknown backend '%s', skipping", b) + + # Default to HTTP if no valid backends + if not backend_modes: + backend_modes = {ConnectionMode.HTTP} + + # Create WorkerIdentity objects (regular workers only from env vars) + requirements = [ + WorkerIdentity(m, b, WorkerType.REGULAR, 0) + for m in models + for b in backend_modes + ] + logger.info("Using env var requirements: %s", [str(r) for r in requirements]) + else: + # Use scanned requirements from test markers + requirements = get_pool_requirements() + logger.info("Using scanned requirements: %s", [str(r) for r in requirements]) + + # Filter to valid models + requirements = [r for r in requirements if r.model_id in MODEL_SPECS] + + if not requirements: + logger.warning("No valid requirements, model pool will be empty") + _model_pool = ModelPool(GPUAllocator(gpus=[])) + return _model_pool + + # Create and start the pool + allocator = GPUAllocator() + _model_pool = ModelPool(allocator) + + startup_timeout = int(os.environ.get(ENV_STARTUP_TIMEOUT, "300")) + _model_pool.startup( + requirements=requirements, + startup_timeout=startup_timeout, + ) + + # Log final GPU allocation summary + logger.info(_model_pool.allocator.summary()) + + # Register cleanup + request.addfinalizer(_model_pool.shutdown) + + return _model_pool + + +@pytest.fixture +def model_client(request: pytest.FixtureRequest, model_pool: "ModelPool"): + """Get OpenAI client for the model specified by @pytest.mark.model(). + + Usage: + @pytest.mark.model("llama-8b") + def test_chat(model_client): + response = model_client.chat.completions.create(...) + """ + from infra import PARAM_MODEL + + marker = request.node.get_closest_marker(PARAM_MODEL) + if marker is None: + pytest.fail( + f"Test must be marked with @pytest.mark.{PARAM_MODEL}('model-id') " + "to use model_client fixture" + ) + + model_id = marker.args[0] + + try: + return model_pool.get_client(model_id) + except KeyError: + pytest.skip(f"Model {model_id} not available in model pool") + + +@pytest.fixture +def model_base_url(request: pytest.FixtureRequest, model_pool: "ModelPool") -> str: + """Get the base URL for the model specified by @pytest.mark.model(). + + Usage: + @pytest.mark.model("llama-8b") + def test_direct_http(model_base_url): + response = httpx.get(f"{model_base_url}/health") + """ + from infra import PARAM_MODEL + + marker = request.node.get_closest_marker(PARAM_MODEL) + if marker is None: + pytest.fail( + f"Test must be marked with @pytest.mark.{PARAM_MODEL}('model-id') " + "to use model_base_url fixture" + ) + + model_id = marker.args[0] + + try: + return model_pool.get_base_url(model_id) + except KeyError: + pytest.skip(f"Model {model_id} not available in model pool") diff --git a/sgl-model-gateway/e2e_test/fixtures/ports.py b/sgl-model-gateway/e2e_test/fixtures/ports.py index d616cffa1..b04703e68 100644 --- a/sgl-model-gateway/e2e_test/fixtures/ports.py +++ b/sgl-model-gateway/e2e_test/fixtures/ports.py @@ -1,3 +1,9 @@ +"""Legacy port utilities. + +DEPRECATED: This module will be removed during e2e_response_api migration. +Use infra.get_open_port() instead. +""" + import socket diff --git a/sgl-model-gateway/e2e_test/fixtures/router_manager.py b/sgl-model-gateway/e2e_test/fixtures/router_manager.py index 23e98f579..ed7e7ace4 100644 --- a/sgl-model-gateway/e2e_test/fixtures/router_manager.py +++ b/sgl-model-gateway/e2e_test/fixtures/router_manager.py @@ -1,3 +1,9 @@ +"""Legacy router manager for integration tests. + +DEPRECATED: This module will be removed during e2e_response_api migration. +Use infra.Gateway instead for all router management. +""" + import subprocess import time from dataclasses import dataclass diff --git a/sgl-model-gateway/e2e_test/fixtures/setup_backend.py b/sgl-model-gateway/e2e_test/fixtures/setup_backend.py new file mode 100644 index 000000000..0c77c914a --- /dev/null +++ b/sgl-model-gateway/e2e_test/fixtures/setup_backend.py @@ -0,0 +1,394 @@ +"""Backend setup fixtures for E2E tests. + +This module provides fixtures for launching gateways/routers for different backends. +""" + +from __future__ import annotations + +import logging +import os +from typing import TYPE_CHECKING + +import pytest + +if TYPE_CHECKING: + from infra import ModelPool + +from .markers import get_marker_kwargs, get_marker_value + +logger = logging.getLogger(__name__) + + +@pytest.fixture(scope="class") +def setup_backend(request: pytest.FixtureRequest, model_pool: "ModelPool"): + """Class-scoped fixture that launches a router for each test class. + + Routers are cheap to start (~1-2s) compared to workers (~30-60s), so we + launch a fresh router per test class for isolation while reusing the + expensive workers from model_pool. + + Backend types: + - "http", "grpc": Gets existing worker from model_pool, launches router + - "pd": Launches prefill/decode workers via model_pool, launches PD router + - "openai", "xai", etc.: Launches cloud router (no local workers) + + Configuration via markers: + - @pytest.mark.model("model-id"): Override default model + - @pytest.mark.workers(count=1): Number of regular workers behind router + - @pytest.mark.workers(prefill=1, decode=1): PD worker configuration + - @pytest.mark.gateway(policy="round_robin", timeout=60): Gateway configuration + + Returns: + Tuple of (backend_name, model_path, openai_client, gateway) + + Usage: + @pytest.mark.parametrize("setup_backend", ["http"], indirect=True) + class TestBasic: + def test_chat(self, setup_backend): + backend, model, client, gateway = setup_backend + """ + import openai + from infra import ( + DEFAULT_MODEL, + DEFAULT_ROUTER_TIMEOUT, + ENV_MODEL, + ENV_SKIP_BACKEND_SETUP, + LOCAL_MODES, + ConnectionMode, + Gateway, + WorkerIdentity, + WorkerType, + ) + + backend_name = request.param + + # Skip if requested + if os.environ.get(ENV_SKIP_BACKEND_SETUP, "").lower() in ("1", "true", "yes"): + pytest.skip(f"{ENV_SKIP_BACKEND_SETUP} is set") + + # Get model from marker or env var or default + model_id = get_marker_value(request, "model") + if model_id is None: + model_id = os.environ.get(ENV_MODEL, DEFAULT_MODEL) + + # Get worker configuration from marker + workers_config = get_marker_kwargs( + request, "workers", defaults={"count": 1, "prefill": None, "decode": None} + ) + + # Get gateway configuration from marker + gateway_config = get_marker_kwargs( + request, + "gateway", + defaults={ + "policy": "round_robin", + "timeout": DEFAULT_ROUTER_TIMEOUT, + "extra_args": None, + }, + ) + + # PD disaggregation backend + if backend_name == "pd": + yield from _setup_pd_backend( + request, model_pool, model_id, workers_config, gateway_config + ) + return + + # Check if this is a local backend (grpc, http) + try: + connection_mode = ConnectionMode(backend_name) + is_local = connection_mode in LOCAL_MODES + except ValueError: + is_local = False + connection_mode = None + + # Local backends: use worker from pool + launch gateway + if is_local: + yield from _setup_local_backend( + request, + model_pool, + backend_name, + model_id, + connection_mode, + workers_config, + gateway_config, + ) + return + + # Cloud backends: launch cloud router + yield from _setup_cloud_backend(backend_name) + + +def _setup_pd_backend( + request: pytest.FixtureRequest, + model_pool: "ModelPool", + model_id: str, + workers_config: dict, + gateway_config: dict, +): + """Setup PD disaggregation backend.""" + import openai + from infra import ConnectionMode, Gateway, WorkerIdentity, WorkerType + + # Check PD requirements + try: + import sgl_kernel # noqa: F401 + except ImportError: + pytest.skip("sgl_kernel not available, required for PD disaggregation") + + try: + import torch + except ImportError: + pytest.skip("torch not available") + + if not torch.cuda.is_available(): + pytest.skip("CUDA not available") + + # Get PD configuration from workers marker + num_prefill = workers_config.get("prefill") or 1 + num_decode = workers_config.get("decode") or 1 + + # Check GPU requirements + required_gpus = num_prefill + num_decode + gpu_count = torch.cuda.device_count() + if gpu_count < required_gpus: + pytest.skip( + f"PD tests require {required_gpus} GPUs " + f"({num_prefill} prefill + {num_decode} decode), found {gpu_count}" + ) + + # Try to use pre-launched PD workers, or launch additional ones if needed + existing_prefills = model_pool.get_workers_by_type(model_id, WorkerType.PREFILL) + existing_decodes = model_pool.get_workers_by_type(model_id, WorkerType.DECODE) + + # Calculate how many more we need + missing_prefill = max(0, num_prefill - len(existing_prefills)) + missing_decode = max(0, num_decode - len(existing_decodes)) + + if missing_prefill == 0 and missing_decode == 0: + prefills = existing_prefills[:num_prefill] + decodes = existing_decodes[:num_decode] + logger.info( + "Using pre-launched PD workers: %d prefill, %d decode", + len(prefills), + len(decodes), + ) + else: + # Build WorkerIdentity list for missing workers + workers_to_launch: list[WorkerIdentity] = [] + for i in range(missing_prefill): + workers_to_launch.append( + WorkerIdentity( + model_id, + ConnectionMode.HTTP, + WorkerType.PREFILL, + len(existing_prefills) + i, + ) + ) + for i in range(missing_decode): + workers_to_launch.append( + WorkerIdentity( + model_id, + ConnectionMode.HTTP, + WorkerType.DECODE, + len(existing_decodes) + i, + ) + ) + + logger.info( + "Have %d/%d prefill, %d/%d decode. Launching %d more workers", + len(existing_prefills), + num_prefill, + len(existing_decodes), + num_decode, + len(workers_to_launch), + ) + new_instances = model_pool.launch_workers( + workers_to_launch, startup_timeout=300 + ) + + new_prefills = [w for w in new_instances if w.worker_type == WorkerType.PREFILL] + new_decodes = [w for w in new_instances if w.worker_type == WorkerType.DECODE] + prefills = existing_prefills + new_prefills + decodes = existing_decodes + new_decodes + + model_path = prefills[0].model_path if prefills else None + + # Launch PD gateway + gateway = Gateway() + gateway.start( + prefill_workers=prefills, + decode_workers=decodes, + policy=gateway_config["policy"], + timeout=gateway_config["timeout"], + extra_args=gateway_config["extra_args"], + ) + + client = openai.OpenAI( + base_url=f"{gateway.base_url}/v1", + api_key="not-used", + ) + + logger.info( + "Setup PD backend: model=%s, %d prefill + %d decode workers, " + "gateway=%s, policy=%s", + model_id, + len(prefills), + len(decodes), + gateway.base_url, + gateway_config["policy"], + ) + + try: + yield "pd", model_path, client, gateway + finally: + logger.info("Tearing down PD gateway") + gateway.shutdown() + + +def _setup_local_backend( + request: pytest.FixtureRequest, + model_pool: "ModelPool", + backend_name: str, + model_id: str, + connection_mode, + workers_config: dict, + gateway_config: dict, +): + """Setup local backend (grpc, http).""" + import openai + from infra import Gateway, WorkerIdentity, WorkerType + + num_workers = workers_config.get("count") or 1 + + try: + if num_workers > 1: + existing = model_pool.get_workers_by_type(model_id, WorkerType.REGULAR) + existing_for_mode = [w for w in existing if w.mode == connection_mode] + + if len(existing_for_mode) >= num_workers: + instances = existing_for_mode[:num_workers] + else: + missing = num_workers - len(existing_for_mode) + workers_to_launch = [ + WorkerIdentity( + model_id, + connection_mode, + WorkerType.REGULAR, + len(existing_for_mode) + i, + ) + for i in range(missing) + ] + new_instances = model_pool.launch_workers( + workers_to_launch, startup_timeout=300 + ) + instances = existing_for_mode + new_instances + + if not instances: + pytest.fail(f"Failed to get {num_workers} workers for {model_id}") + worker_urls = [inst.worker_url for inst in instances] + model_path = instances[0].model_path + else: + instance = model_pool.get(model_id, connection_mode) + worker_urls = [instance.worker_url] + model_path = instance.model_path + except RuntimeError as e: + pytest.fail(str(e)) + + # Launch gateway + gateway = Gateway() + gateway.start( + worker_urls=worker_urls, + model_path=model_path, + policy=gateway_config["policy"], + timeout=gateway_config["timeout"], + extra_args=gateway_config["extra_args"], + ) + + client = openai.OpenAI( + base_url=f"{gateway.base_url}/v1", + api_key="not-used", + ) + + logger.info( + "Setup %s backend: model=%s, workers=%d, gateway=%s, policy=%s", + backend_name, + model_id, + num_workers, + gateway.base_url, + gateway_config["policy"], + ) + + try: + yield backend_name, model_path, client, gateway + finally: + logger.info("Tearing down gateway for %s backend", backend_name) + gateway.shutdown() + + +def _setup_cloud_backend(backend_name: str): + """Setup cloud backend (openai, xai, etc.).""" + import openai + from backends import CLOUD_BACKENDS, launch_cloud_backend + + if backend_name not in CLOUD_BACKENDS: + pytest.fail(f"Unknown backend: {backend_name}") + + cfg = CLOUD_BACKENDS[backend_name] + api_key_env = cfg.get("api_key_env") + + if api_key_env and not os.environ.get(api_key_env): + pytest.skip(f"{api_key_env} not set, skipping {backend_name} tests") + + logger.info("Launching cloud backend: %s", backend_name) + router = launch_cloud_backend(backend_name) + + api_key = os.environ.get(api_key_env) if api_key_env else "not-used" + client = openai.OpenAI( + base_url=f"{router.base_url}/v1", + api_key=api_key, + ) + + try: + yield backend_name, cfg["model"], client, router + finally: + logger.info("Tearing down cloud backend: %s", backend_name) + router.shutdown() + + +@pytest.fixture +def backend_router(request: pytest.FixtureRequest, model_pool: "ModelPool"): + """Function-scoped fixture for launching a fresh router per test. + + This launches a new Gateway for each test, pointing to workers from the pool. + Use for tests that need isolated router state. + + Usage: + @pytest.mark.parametrize("backend_router", ["grpc", "http"], indirect=True) + def test_router_state(backend_router): + gateway = backend_router + """ + from infra import DEFAULT_MODEL, ENV_MODEL, ConnectionMode, Gateway + + backend_name = request.param + model_id = os.environ.get(ENV_MODEL, DEFAULT_MODEL) + + connection_mode = ConnectionMode(backend_name) + + try: + instance = model_pool.get(model_id, connection_mode) + except KeyError: + pytest.skip(f"Model {model_id}:{backend_name} not available in pool") + except RuntimeError as e: + pytest.fail(str(e)) + + gateway = Gateway() + gateway.start( + worker_urls=[instance.worker_url], + model_path=instance.model_path, + ) + + try: + yield gateway + finally: + gateway.shutdown() diff --git a/sgl-model-gateway/e2e_test/infra/gateway.py b/sgl-model-gateway/e2e_test/infra/gateway.py index 067e1e26e..b46fec184 100644 --- a/sgl-model-gateway/e2e_test/infra/gateway.py +++ b/sgl-model-gateway/e2e_test/infra/gateway.py @@ -40,10 +40,11 @@ class Gateway: - Worker management (list, add, remove) - Health and metrics endpoints - Three startup modes: + Four startup modes: 1. Regular mode: Start with worker URLs 2. PD mode: Start with prefill/decode workers 3. IGW mode: Start empty, add workers via API + 4. Cloud mode: Start with cloud backend (OpenAI, xAI) Example (regular mode): gateway = Gateway() @@ -71,6 +72,11 @@ class Gateway: # Cleanup gateway.shutdown() + + Example (cloud mode): + gateway = Gateway() + gateway.start(cloud_backend="openai") # or "xai" + # Requires OPENAI_API_KEY or XAI_API_KEY env var """ def __init__( @@ -97,7 +103,10 @@ class Gateway: self.policy: str = "round_robin" self.pd_mode: bool = False self.igw_mode: bool = False + self.cloud_mode: bool = False + self.cloud_backend: str | None = None self._started: bool = False + self._env: dict[str, str] | None = None # Custom env for subprocess @property def is_running(self) -> bool: @@ -115,6 +124,9 @@ class Gateway: decode_workers: list["ModelInstance"] | None = None, # IGW mode arguments igw_mode: bool = False, + # Cloud mode arguments + cloud_backend: str | None = None, + history_backend: str = "memory", # Common arguments policy: str = "round_robin", timeout: float = DEFAULT_ROUTER_TIMEOUT, @@ -123,10 +135,11 @@ class Gateway: ) -> None: """Start the gateway. - Can be started in three modes: + Can be started in four modes: 1. Regular mode: Provide worker_urls and model_path 2. PD mode: Provide prefill_workers and decode_workers 3. IGW mode: Set igw_mode=True, add workers later via add_worker() + 4. Cloud mode: Provide cloud_backend ("openai" or "xai") Args: worker_urls: List of worker URLs for regular mode. @@ -134,6 +147,8 @@ class Gateway: prefill_workers: List of prefill ModelInstance objects for PD mode. decode_workers: List of decode ModelInstance objects for PD mode. igw_mode: Start in IGW mode (no workers, add via API). + cloud_backend: Cloud backend type ("openai" or "xai"). + history_backend: History backend for cloud mode ("memory" or "oracle"). policy: Routing policy (round_robin, random, etc.) timeout: Startup timeout in seconds. show_output: Show subprocess output (env var override). @@ -150,19 +165,21 @@ class Gateway: is_pd_mode = prefill_workers is not None or decode_workers is not None is_regular_mode = worker_urls is not None is_igw_mode = igw_mode + is_cloud_mode = cloud_backend is not None # Validate mode exclusivity - modes_specified = sum([is_pd_mode, is_regular_mode, is_igw_mode]) + modes_specified = sum([is_pd_mode, is_regular_mode, is_igw_mode, is_cloud_mode]) if modes_specified > 1: raise ValueError( "Cannot specify multiple modes. Choose one of: " - "worker_urls (regular), prefill/decode_workers (PD), or igw_mode" + "worker_urls (regular), prefill/decode_workers (PD), " + "igw_mode, or cloud_backend" ) if modes_specified == 0: raise ValueError( "Must specify one mode: worker_urls (regular), " - "prefill/decode_workers (PD), or igw_mode=True" + "prefill/decode_workers (PD), igw_mode=True, or cloud_backend" ) if show_output is None: @@ -201,6 +218,47 @@ class Gateway: extra_args=extra_args, log_msg=f"PD gateway ({len(prefills)} prefill, {len(decodes)} decode)", ) + elif is_cloud_mode: + # Cloud mode: OpenAI/xAI backend + self.pd_mode = False + self.igw_mode = False + self.cloud_mode = True + self.cloud_backend = cloud_backend + + # Get worker URL and API key based on backend + if cloud_backend == "openai": + worker_url = "https://api.openai.com" + api_key = os.environ.get("OPENAI_API_KEY") + if not api_key: + raise ValueError("OPENAI_API_KEY environment variable required") + self._env = os.environ.copy() + self._env["OPENAI_API_KEY"] = api_key + elif cloud_backend == "xai": + worker_url = "https://api.x.ai" + api_key = os.environ.get("XAI_API_KEY") + if not api_key: + raise ValueError("XAI_API_KEY environment variable required") + self._env = os.environ.copy() + self._env["XAI_API_KEY"] = api_key + else: + raise ValueError(f"Unsupported cloud backend: {cloud_backend}") + + mode_args = [ + "--backend", + "openai", # Both OpenAI and xAI use openai backend type + "--worker-urls", + worker_url, + "--history-backend", + history_backend, + ] + + self._launch( + mode_args=mode_args, + timeout=timeout, + show_output=show_output, + extra_args=extra_args, + log_msg=f"{cloud_backend} cloud gateway", + ) else: # Regular mode: worker URLs if model_path is None: @@ -249,6 +307,7 @@ class Gateway: self.process = subprocess.Popen( cmd, + env=self._env, # Use custom env if set (e.g., for cloud mode API keys) stdout=None if show_output else subprocess.PIPE, stderr=None if show_output else subprocess.PIPE, start_new_session=True,