[model-gateway][e2e_test]: Create directory structure and backends config (#16469)
This commit is contained in:
@@ -0,0 +1,504 @@
|
|||||||
|
"""Backend configurations for E2E tests.
|
||||||
|
|
||||||
|
This module defines the available backends for E2E testing:
|
||||||
|
- grpc: Local gRPC workers with SGLang router
|
||||||
|
- http: Local HTTP workers with SGLang router
|
||||||
|
- openai: OpenAI API backend
|
||||||
|
- xai: xAI API backend
|
||||||
|
|
||||||
|
Each backend configuration specifies:
|
||||||
|
- model: Model path or name
|
||||||
|
- launcher: Function to launch the backend
|
||||||
|
- launcher_kwargs: Arguments for the launcher
|
||||||
|
- needs_workers: Whether local GPU workers are needed
|
||||||
|
- api_key_env: Environment variable for API key (if needed)
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import signal
|
||||||
|
import socket
|
||||||
|
import subprocess
|
||||||
|
import time
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import TYPE_CHECKING, Any, Callable
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import openai
|
||||||
|
|
||||||
|
from infra.model_specs import _resolve_model_path
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
# Default ports for each backend type (can be overridden)
|
||||||
|
DEFAULT_PORTS = {
|
||||||
|
"grpc": 30030,
|
||||||
|
"grpc_harmony": 30031,
|
||||||
|
"http": 30020,
|
||||||
|
"openai": 30010,
|
||||||
|
"xai": 30011,
|
||||||
|
"oracle_store": 30040,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Prometheus port offset from main port
|
||||||
|
PROMETHEUS_PORT_OFFSET = 1000
|
||||||
|
|
||||||
|
|
||||||
|
def get_open_port() -> int:
|
||||||
|
"""Get an available port by binding to port 0."""
|
||||||
|
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||||
|
s.bind(("", 0))
|
||||||
|
s.listen(1)
|
||||||
|
return s.getsockname()[1]
|
||||||
|
|
||||||
|
|
||||||
|
def kill_process_tree(pid: int, sig: int = signal.SIGTERM) -> None:
|
||||||
|
"""Kill a process and all its children."""
|
||||||
|
try:
|
||||||
|
import psutil
|
||||||
|
|
||||||
|
parent = psutil.Process(pid)
|
||||||
|
children = parent.children(recursive=True)
|
||||||
|
for child in children:
|
||||||
|
try:
|
||||||
|
child.send_signal(sig)
|
||||||
|
except psutil.NoSuchProcess:
|
||||||
|
pass
|
||||||
|
parent.send_signal(sig)
|
||||||
|
except ImportError:
|
||||||
|
# Fallback if psutil not available
|
||||||
|
os.kill(pid, sig)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Failed to kill process tree for PID %d: %s", pid, e)
|
||||||
|
|
||||||
|
|
||||||
|
def wait_for_health(
|
||||||
|
url: str,
|
||||||
|
timeout: float = 60,
|
||||||
|
api_key: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Wait for a server's /health endpoint to return 200."""
|
||||||
|
start = time.time()
|
||||||
|
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
|
||||||
|
|
||||||
|
while time.time() - start < timeout:
|
||||||
|
try:
|
||||||
|
resp = requests.get(f"{url}/health", headers=headers, timeout=5)
|
||||||
|
if resp.status_code == 200:
|
||||||
|
return
|
||||||
|
except requests.RequestException:
|
||||||
|
pass
|
||||||
|
time.sleep(1)
|
||||||
|
|
||||||
|
raise TimeoutError(f"Server at {url} did not become healthy within {timeout}s")
|
||||||
|
|
||||||
|
|
||||||
|
def wait_for_workers_ready(
|
||||||
|
router_url: str,
|
||||||
|
expected_workers: int,
|
||||||
|
timeout: float = 300,
|
||||||
|
api_key: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Wait for router to have all workers connected."""
|
||||||
|
start = time.time()
|
||||||
|
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
|
||||||
|
|
||||||
|
while time.time() - start < timeout:
|
||||||
|
try:
|
||||||
|
resp = requests.get(f"{router_url}/workers", headers=headers, timeout=5)
|
||||||
|
if resp.status_code == 200:
|
||||||
|
data = resp.json()
|
||||||
|
if data.get("total", 0) >= expected_workers:
|
||||||
|
logger.info(
|
||||||
|
"All %d workers connected after %.1fs",
|
||||||
|
expected_workers,
|
||||||
|
time.time() - start,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
except requests.RequestException:
|
||||||
|
pass
|
||||||
|
time.sleep(2)
|
||||||
|
|
||||||
|
raise TimeoutError(
|
||||||
|
f"Router at {router_url} did not get {expected_workers} workers within {timeout}s"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ClusterInfo:
|
||||||
|
"""Information about a running cluster."""
|
||||||
|
|
||||||
|
base_url: str
|
||||||
|
router_process: subprocess.Popen
|
||||||
|
worker_processes: list[subprocess.Popen]
|
||||||
|
model: str
|
||||||
|
backend: str
|
||||||
|
|
||||||
|
def shutdown(self) -> None:
|
||||||
|
"""Shutdown the cluster."""
|
||||||
|
# Kill router first
|
||||||
|
if self.router_process.poll() is None:
|
||||||
|
kill_process_tree(self.router_process.pid)
|
||||||
|
|
||||||
|
# Kill workers
|
||||||
|
for proc in self.worker_processes:
|
||||||
|
if proc.poll() is None:
|
||||||
|
kill_process_tree(proc.pid)
|
||||||
|
|
||||||
|
|
||||||
|
def launch_grpc_cluster(
|
||||||
|
model: str,
|
||||||
|
base_url: str | None = None,
|
||||||
|
*,
|
||||||
|
num_workers: int = 1,
|
||||||
|
tp_size: int = 1,
|
||||||
|
policy: str = "round_robin",
|
||||||
|
api_key: str | None = None,
|
||||||
|
worker_args: list[str] | None = None,
|
||||||
|
router_args: list[str] | None = None,
|
||||||
|
timeout: float = 300,
|
||||||
|
show_output: bool | None = None,
|
||||||
|
) -> ClusterInfo:
|
||||||
|
"""Launch gRPC workers and router.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model: Model path
|
||||||
|
base_url: Base URL for router (auto-assigns port if None)
|
||||||
|
num_workers: Number of workers to launch
|
||||||
|
tp_size: Tensor parallelism size
|
||||||
|
policy: Routing policy
|
||||||
|
api_key: Optional API key for router auth
|
||||||
|
worker_args: Additional worker arguments
|
||||||
|
router_args: Additional router arguments
|
||||||
|
timeout: Startup timeout in seconds
|
||||||
|
show_output: Show subprocess output (default: SHOW_ROUTER_LOGS env var)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
ClusterInfo with running processes
|
||||||
|
"""
|
||||||
|
if show_output is None:
|
||||||
|
show_output = os.environ.get("SHOW_ROUTER_LOGS", "0") == "1"
|
||||||
|
|
||||||
|
# Determine router port
|
||||||
|
if base_url:
|
||||||
|
router_port = int(base_url.split(":")[-1])
|
||||||
|
else:
|
||||||
|
router_port = get_open_port()
|
||||||
|
base_url = f"http://127.0.0.1:{router_port}"
|
||||||
|
|
||||||
|
logger.info("Launching gRPC cluster: %d workers, tp=%d", num_workers, tp_size)
|
||||||
|
|
||||||
|
# Launch workers
|
||||||
|
workers = []
|
||||||
|
worker_urls = []
|
||||||
|
|
||||||
|
for i in range(num_workers):
|
||||||
|
worker_port = get_open_port()
|
||||||
|
worker_url = f"grpc://127.0.0.1:{worker_port}"
|
||||||
|
worker_urls.append(worker_url)
|
||||||
|
|
||||||
|
cmd = [
|
||||||
|
"python3",
|
||||||
|
"-m",
|
||||||
|
"sglang.launch_server",
|
||||||
|
"--model-path",
|
||||||
|
model,
|
||||||
|
"--host",
|
||||||
|
"127.0.0.1",
|
||||||
|
"--port",
|
||||||
|
str(worker_port),
|
||||||
|
"--grpc-mode",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.8",
|
||||||
|
"--log-level",
|
||||||
|
"warning",
|
||||||
|
]
|
||||||
|
|
||||||
|
if tp_size > 1:
|
||||||
|
cmd.extend(["--tp-size", str(tp_size)])
|
||||||
|
|
||||||
|
if worker_args:
|
||||||
|
cmd.extend(worker_args)
|
||||||
|
|
||||||
|
logger.info("Starting worker %d on port %d", i + 1, worker_port)
|
||||||
|
|
||||||
|
proc = subprocess.Popen(
|
||||||
|
cmd,
|
||||||
|
stdout=None if show_output else subprocess.PIPE,
|
||||||
|
stderr=None if show_output else subprocess.PIPE,
|
||||||
|
start_new_session=True,
|
||||||
|
)
|
||||||
|
workers.append(proc)
|
||||||
|
|
||||||
|
# Wait for workers to initialize
|
||||||
|
logger.info("Waiting for workers to initialize (20s)...")
|
||||||
|
time.sleep(20)
|
||||||
|
|
||||||
|
# Verify workers are alive
|
||||||
|
for i, worker in enumerate(workers):
|
||||||
|
if worker.poll() is not None:
|
||||||
|
# Cleanup
|
||||||
|
for w in workers:
|
||||||
|
try:
|
||||||
|
kill_process_tree(w.pid)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
raise RuntimeError(f"Worker {i + 1} died during startup")
|
||||||
|
|
||||||
|
# Launch router
|
||||||
|
router_cmd = [
|
||||||
|
"python3",
|
||||||
|
"-m",
|
||||||
|
"sglang_router.launch_router",
|
||||||
|
"--host",
|
||||||
|
"127.0.0.1",
|
||||||
|
"--port",
|
||||||
|
str(router_port),
|
||||||
|
"--prometheus-port",
|
||||||
|
str(router_port + PROMETHEUS_PORT_OFFSET),
|
||||||
|
"--policy",
|
||||||
|
policy,
|
||||||
|
"--model-path",
|
||||||
|
model,
|
||||||
|
"--log-level",
|
||||||
|
"warn",
|
||||||
|
"--worker-urls",
|
||||||
|
*worker_urls,
|
||||||
|
]
|
||||||
|
|
||||||
|
if api_key:
|
||||||
|
router_cmd.extend(["--api-key", api_key])
|
||||||
|
|
||||||
|
if router_args:
|
||||||
|
router_cmd.extend(router_args)
|
||||||
|
|
||||||
|
logger.info("Starting router on port %d", router_port)
|
||||||
|
|
||||||
|
router_proc = subprocess.Popen(
|
||||||
|
router_cmd,
|
||||||
|
stdout=None if show_output else subprocess.PIPE,
|
||||||
|
stderr=None if show_output else subprocess.PIPE,
|
||||||
|
start_new_session=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Wait for router to be ready with all workers
|
||||||
|
try:
|
||||||
|
wait_for_workers_ready(base_url, num_workers, timeout=timeout, api_key=api_key)
|
||||||
|
except TimeoutError:
|
||||||
|
# Cleanup on failure
|
||||||
|
kill_process_tree(router_proc.pid)
|
||||||
|
for w in workers:
|
||||||
|
kill_process_tree(w.pid)
|
||||||
|
raise
|
||||||
|
|
||||||
|
logger.info("gRPC cluster ready at %s with %d workers", base_url, num_workers)
|
||||||
|
|
||||||
|
return ClusterInfo(
|
||||||
|
base_url=base_url,
|
||||||
|
router_process=router_proc,
|
||||||
|
worker_processes=workers,
|
||||||
|
model=model,
|
||||||
|
backend="grpc",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def launch_openai_router(
|
||||||
|
backend: str, # "openai" or "xai"
|
||||||
|
base_url: str | None = None,
|
||||||
|
*,
|
||||||
|
history_backend: str = "memory",
|
||||||
|
router_args: list[str] | None = None,
|
||||||
|
timeout: float = 60,
|
||||||
|
show_output: bool | None = None,
|
||||||
|
) -> ClusterInfo:
|
||||||
|
"""Launch router with OpenAI/xAI backend.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
backend: "openai" or "xai"
|
||||||
|
base_url: Base URL for router (auto-assigns port if None)
|
||||||
|
history_backend: "memory" or "oracle"
|
||||||
|
router_args: Additional router arguments
|
||||||
|
timeout: Startup timeout in seconds
|
||||||
|
show_output: Show subprocess output
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
ClusterInfo with running router
|
||||||
|
"""
|
||||||
|
if show_output is None:
|
||||||
|
show_output = os.environ.get("SHOW_ROUTER_LOGS", "0") == "1"
|
||||||
|
|
||||||
|
# Determine port
|
||||||
|
if base_url:
|
||||||
|
router_port = int(base_url.split(":")[-1])
|
||||||
|
else:
|
||||||
|
router_port = get_open_port()
|
||||||
|
base_url = f"http://127.0.0.1:{router_port}"
|
||||||
|
|
||||||
|
# Get API key
|
||||||
|
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 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(router_port + PROMETHEUS_PORT_OFFSET),
|
||||||
|
"--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 ClusterInfo(
|
||||||
|
base_url=base_url,
|
||||||
|
router_process=router_proc,
|
||||||
|
worker_processes=[],
|
||||||
|
model="", # Cloud API - model specified per request
|
||||||
|
backend=backend,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# Backend configuration registry
|
||||||
|
BACKENDS: dict[str, dict[str, Any]] = {
|
||||||
|
"grpc": {
|
||||||
|
"description": "Local gRPC workers with SGLang router",
|
||||||
|
"model": _resolve_model_path("meta-llama/Llama-3.1-8B-Instruct"),
|
||||||
|
"launcher": launch_grpc_cluster,
|
||||||
|
"launcher_kwargs": {
|
||||||
|
"num_workers": 1,
|
||||||
|
"tp_size": 1,
|
||||||
|
"policy": "round_robin",
|
||||||
|
},
|
||||||
|
"needs_workers": True,
|
||||||
|
"api_key_env": None,
|
||||||
|
},
|
||||||
|
"grpc_harmony": {
|
||||||
|
"description": "Local gRPC workers with Harmony model",
|
||||||
|
"model": _resolve_model_path("openai/gpt-oss-20b"),
|
||||||
|
"launcher": launch_grpc_cluster,
|
||||||
|
"launcher_kwargs": {
|
||||||
|
"num_workers": 1,
|
||||||
|
"tp_size": 2,
|
||||||
|
"policy": "round_robin",
|
||||||
|
"worker_args": ["--reasoning-parser=gpt-oss"],
|
||||||
|
"router_args": ["--history-backend", "memory"],
|
||||||
|
},
|
||||||
|
"needs_workers": True,
|
||||||
|
"api_key_env": None,
|
||||||
|
},
|
||||||
|
"openai": {
|
||||||
|
"description": "OpenAI API backend",
|
||||||
|
"model": "gpt-4o-mini",
|
||||||
|
"launcher": launch_openai_router,
|
||||||
|
"launcher_kwargs": {
|
||||||
|
"backend": "openai",
|
||||||
|
"history_backend": "memory",
|
||||||
|
},
|
||||||
|
"needs_workers": False,
|
||||||
|
"api_key_env": "OPENAI_API_KEY",
|
||||||
|
},
|
||||||
|
"xai": {
|
||||||
|
"description": "xAI API backend",
|
||||||
|
"model": "grok-2-latest",
|
||||||
|
"launcher": launch_openai_router,
|
||||||
|
"launcher_kwargs": {
|
||||||
|
"backend": "xai",
|
||||||
|
"history_backend": "memory",
|
||||||
|
},
|
||||||
|
"needs_workers": False,
|
||||||
|
"api_key_env": "XAI_API_KEY",
|
||||||
|
},
|
||||||
|
"oracle_store": {
|
||||||
|
"description": "OpenAI API with Oracle history backend",
|
||||||
|
"model": "gpt-4o-mini",
|
||||||
|
"launcher": launch_openai_router,
|
||||||
|
"launcher_kwargs": {
|
||||||
|
"backend": "openai",
|
||||||
|
"history_backend": "oracle",
|
||||||
|
},
|
||||||
|
"needs_workers": False,
|
||||||
|
"api_key_env": "OPENAI_API_KEY",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def get_backend_config(backend: str) -> dict[str, Any]:
|
||||||
|
"""Get configuration for a backend."""
|
||||||
|
if backend not in BACKENDS:
|
||||||
|
raise KeyError(
|
||||||
|
f"Unknown backend: {backend}. Available: {list(BACKENDS.keys())}"
|
||||||
|
)
|
||||||
|
return BACKENDS[backend]
|
||||||
|
|
||||||
|
|
||||||
|
def launch_backend(backend: str, **kwargs: Any) -> ClusterInfo:
|
||||||
|
"""Launch a backend cluster.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
backend: Backend name from BACKENDS
|
||||||
|
**kwargs: Override launcher kwargs
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
ClusterInfo with running cluster
|
||||||
|
"""
|
||||||
|
cfg = get_backend_config(backend)
|
||||||
|
|
||||||
|
# Merge kwargs with defaults
|
||||||
|
launcher_kwargs = {**cfg["launcher_kwargs"], **kwargs}
|
||||||
|
|
||||||
|
# Add model for grpc backends
|
||||||
|
if cfg["needs_workers"]:
|
||||||
|
return cfg["launcher"](cfg["model"], **launcher_kwargs)
|
||||||
|
else:
|
||||||
|
return cfg["launcher"](**launcher_kwargs)
|
||||||
@@ -45,6 +45,10 @@ def pytest_configure(config: pytest.Config) -> None:
|
|||||||
"markers",
|
"markers",
|
||||||
"model(name): mark test to use a specific model from the model pool",
|
"model(name): mark test to use a specific model from the model pool",
|
||||||
)
|
)
|
||||||
|
config.addinivalue_line(
|
||||||
|
"markers",
|
||||||
|
"backend(name): mark test to use a specific backend (grpc, openai, etc.)",
|
||||||
|
)
|
||||||
config.addinivalue_line(
|
config.addinivalue_line(
|
||||||
"markers",
|
"markers",
|
||||||
"e2e: mark test as an end-to-end test requiring GPU workers",
|
"e2e: mark test as an end-to-end test requiring GPU workers",
|
||||||
@@ -189,3 +193,97 @@ def model_base_url(request: pytest.FixtureRequest, model_pool: "ModelPool") -> s
|
|||||||
return model_pool.get_base_url(model_id)
|
return model_pool.get_base_url(model_id)
|
||||||
except KeyError:
|
except KeyError:
|
||||||
pytest.skip(f"Model {model_id} not available in model pool")
|
pytest.skip(f"Model {model_id} not available in model pool")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Backend fixtures (class-scoped)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(scope="class")
|
||||||
|
def setup_backend(request: pytest.FixtureRequest):
|
||||||
|
"""Class-scoped fixture for launching backend clusters.
|
||||||
|
|
||||||
|
This fixture is used with pytest.mark.parametrize to run tests
|
||||||
|
against multiple backends.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
@pytest.mark.parametrize("setup_backend", ["grpc", "openai"], indirect=True)
|
||||||
|
class TestChatCompletions:
|
||||||
|
def test_basic(self, setup_backend):
|
||||||
|
backend, model, client = setup_backend
|
||||||
|
response = client.chat.completions.create(...)
|
||||||
|
|
||||||
|
Environment variables:
|
||||||
|
- SKIP_BACKEND_SETUP: Skip backend startup (for dry runs)
|
||||||
|
- SHOW_ROUTER_LOGS: Show subprocess output
|
||||||
|
"""
|
||||||
|
import openai
|
||||||
|
from backends import BACKENDS, ClusterInfo, launch_backend
|
||||||
|
|
||||||
|
backend_name = request.param
|
||||||
|
|
||||||
|
# Skip if requested
|
||||||
|
if os.environ.get("SKIP_BACKEND_SETUP", "").lower() in ("1", "true", "yes"):
|
||||||
|
pytest.skip("SKIP_BACKEND_SETUP is set")
|
||||||
|
|
||||||
|
# Check if backend requires API key
|
||||||
|
cfg = BACKENDS.get(backend_name)
|
||||||
|
if cfg is None:
|
||||||
|
pytest.fail(f"Unknown backend: {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("Setting up backend: %s", backend_name)
|
||||||
|
|
||||||
|
# Launch the backend
|
||||||
|
cluster: ClusterInfo = launch_backend(backend_name)
|
||||||
|
|
||||||
|
# Create OpenAI client
|
||||||
|
api_key = os.environ.get(api_key_env) if api_key_env else "not-used"
|
||||||
|
client = openai.OpenAI(
|
||||||
|
base_url=f"{cluster.base_url}/v1",
|
||||||
|
api_key=api_key,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Yield to test
|
||||||
|
try:
|
||||||
|
yield backend_name, cfg["model"], client
|
||||||
|
finally:
|
||||||
|
logger.info("Tearing down backend: %s", backend_name)
|
||||||
|
cluster.shutdown()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def backend_cluster(request: pytest.FixtureRequest):
|
||||||
|
"""Function-scoped fixture for launching a fresh backend per test.
|
||||||
|
|
||||||
|
Unlike setup_backend (class-scoped), this creates a new cluster
|
||||||
|
for each test function. Use for tests that modify cluster state.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
@pytest.mark.parametrize("backend_cluster", ["grpc"], indirect=True)
|
||||||
|
def test_add_worker(backend_cluster):
|
||||||
|
cluster = backend_cluster
|
||||||
|
# cluster.base_url, cluster.router_process, etc.
|
||||||
|
"""
|
||||||
|
from backends import BACKENDS, launch_backend
|
||||||
|
|
||||||
|
backend_name = request.param
|
||||||
|
|
||||||
|
cfg = BACKENDS.get(backend_name)
|
||||||
|
if cfg is None:
|
||||||
|
pytest.fail(f"Unknown backend: {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")
|
||||||
|
|
||||||
|
cluster = launch_backend(backend_name)
|
||||||
|
|
||||||
|
try:
|
||||||
|
yield cluster
|
||||||
|
finally:
|
||||||
|
cluster.shutdown()
|
||||||
|
|||||||
@@ -11,7 +11,21 @@ from .gpu_allocator import (
|
|||||||
wait_for_gpu_memory_to_clear,
|
wait_for_gpu_memory_to_clear,
|
||||||
)
|
)
|
||||||
from .model_pool import ModelInstance, ModelPool
|
from .model_pool import ModelInstance, ModelPool
|
||||||
from .model_specs import MODEL_SPECS
|
from .model_specs import ( # Default model paths; Model groups
|
||||||
|
CHAT_MODELS,
|
||||||
|
DEFAULT_EMBEDDING_MODEL_PATH,
|
||||||
|
DEFAULT_ENABLE_THINKING_MODEL_PATH,
|
||||||
|
DEFAULT_GPT_OSS_MODEL_PATH,
|
||||||
|
DEFAULT_MISTRAL_FUNCTION_CALLING_MODEL_PATH,
|
||||||
|
DEFAULT_MODEL_PATH,
|
||||||
|
DEFAULT_QWEN_FUNCTION_CALLING_MODEL_PATH,
|
||||||
|
DEFAULT_REASONING_MODEL_PATH,
|
||||||
|
DEFAULT_SMALL_MODEL_PATH,
|
||||||
|
EMBEDDING_MODELS,
|
||||||
|
FUNCTION_CALLING_MODELS,
|
||||||
|
MODEL_SPECS,
|
||||||
|
REASONING_MODELS,
|
||||||
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
# GPU allocation
|
# GPU allocation
|
||||||
@@ -28,4 +42,18 @@ __all__ = [
|
|||||||
"ModelInstance",
|
"ModelInstance",
|
||||||
"ModelPool",
|
"ModelPool",
|
||||||
"MODEL_SPECS",
|
"MODEL_SPECS",
|
||||||
|
# Default model paths
|
||||||
|
"DEFAULT_MODEL_PATH",
|
||||||
|
"DEFAULT_SMALL_MODEL_PATH",
|
||||||
|
"DEFAULT_REASONING_MODEL_PATH",
|
||||||
|
"DEFAULT_ENABLE_THINKING_MODEL_PATH",
|
||||||
|
"DEFAULT_QWEN_FUNCTION_CALLING_MODEL_PATH",
|
||||||
|
"DEFAULT_MISTRAL_FUNCTION_CALLING_MODEL_PATH",
|
||||||
|
"DEFAULT_GPT_OSS_MODEL_PATH",
|
||||||
|
"DEFAULT_EMBEDDING_MODEL_PATH",
|
||||||
|
# Model groups
|
||||||
|
"CHAT_MODELS",
|
||||||
|
"EMBEDDING_MODELS",
|
||||||
|
"REASONING_MODELS",
|
||||||
|
"FUNCTION_CALLING_MODELS",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -107,3 +107,17 @@ CHAT_MODELS = get_models_with_feature("chat")
|
|||||||
EMBEDDING_MODELS = get_models_with_feature("embedding")
|
EMBEDDING_MODELS = get_models_with_feature("embedding")
|
||||||
REASONING_MODELS = get_models_with_feature("reasoning")
|
REASONING_MODELS = get_models_with_feature("reasoning")
|
||||||
FUNCTION_CALLING_MODELS = get_models_with_feature("function_calling")
|
FUNCTION_CALLING_MODELS = get_models_with_feature("function_calling")
|
||||||
|
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# Default model path constants (for backward compatibility with existing tests)
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
DEFAULT_MODEL_PATH = MODEL_SPECS["llama-8b"]["model"]
|
||||||
|
DEFAULT_SMALL_MODEL_PATH = MODEL_SPECS["llama-1b"]["model"]
|
||||||
|
DEFAULT_REASONING_MODEL_PATH = MODEL_SPECS["deepseek-7b"]["model"]
|
||||||
|
DEFAULT_ENABLE_THINKING_MODEL_PATH = MODEL_SPECS["qwen-30b"]["model"]
|
||||||
|
DEFAULT_QWEN_FUNCTION_CALLING_MODEL_PATH = MODEL_SPECS["qwen-7b"]["model"]
|
||||||
|
DEFAULT_MISTRAL_FUNCTION_CALLING_MODEL_PATH = MODEL_SPECS["mistral-7b"]["model"]
|
||||||
|
DEFAULT_GPT_OSS_MODEL_PATH = MODEL_SPECS["gpt-oss"]["model"]
|
||||||
|
DEFAULT_EMBEDDING_MODEL_PATH = MODEL_SPECS["embedding"]["model"]
|
||||||
|
|||||||
@@ -0,0 +1,264 @@
|
|||||||
|
"""Consolidated utilities for E2E tests.
|
||||||
|
|
||||||
|
This module provides common utilities used across E2E tests:
|
||||||
|
- Tokenizer loading (get_tokenizer)
|
||||||
|
- Test base classes (CustomTestCase for unittest compatibility)
|
||||||
|
- Model path resolution
|
||||||
|
- Process management utilities
|
||||||
|
|
||||||
|
Import examples:
|
||||||
|
from utils import get_tokenizer, CustomTestCase
|
||||||
|
from utils import DEFAULT_MODEL_PATH, DEFAULT_TIMEOUT
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Re-export commonly used items from submodules
|
||||||
|
from backends import kill_process_tree # noqa: F401
|
||||||
|
from infra.model_specs import ( # noqa: F401; Default model paths
|
||||||
|
DEFAULT_EMBEDDING_MODEL_PATH,
|
||||||
|
DEFAULT_ENABLE_THINKING_MODEL_PATH,
|
||||||
|
DEFAULT_GPT_OSS_MODEL_PATH,
|
||||||
|
DEFAULT_MISTRAL_FUNCTION_CALLING_MODEL_PATH,
|
||||||
|
DEFAULT_MODEL_PATH,
|
||||||
|
DEFAULT_QWEN_FUNCTION_CALLING_MODEL_PATH,
|
||||||
|
DEFAULT_REASONING_MODEL_PATH,
|
||||||
|
DEFAULT_SMALL_MODEL_PATH,
|
||||||
|
MODEL_SPECS,
|
||||||
|
ROUTER_LOCAL_MODEL_PATH,
|
||||||
|
_resolve_model_path,
|
||||||
|
)
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# Constants
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
# Server startup timeout (seconds)
|
||||||
|
DEFAULT_TIMEOUT = 600
|
||||||
|
DEFAULT_STARTUP_TIMEOUT = 300
|
||||||
|
|
||||||
|
# Default test port range
|
||||||
|
DEFAULT_PORT_BASE = 20000
|
||||||
|
|
||||||
|
# File paths for test output
|
||||||
|
STDOUT_FILENAME = "/tmp/sglang_test_stdout.txt"
|
||||||
|
STDERR_FILENAME = "/tmp/sglang_test_stderr.txt"
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# Tokenizer Utilities
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
# Lazy import transformers to avoid import errors in environments without it
|
||||||
|
_transformers_available = None
|
||||||
|
_AutoTokenizer = None
|
||||||
|
_PreTrainedTokenizer = None
|
||||||
|
_PreTrainedTokenizerBase = None
|
||||||
|
_PreTrainedTokenizerFast = None
|
||||||
|
|
||||||
|
|
||||||
|
def _ensure_transformers():
|
||||||
|
"""Lazy load transformers module."""
|
||||||
|
global _transformers_available, _AutoTokenizer
|
||||||
|
global _PreTrainedTokenizer, _PreTrainedTokenizerBase, _PreTrainedTokenizerFast
|
||||||
|
|
||||||
|
if _transformers_available is not None:
|
||||||
|
return _transformers_available
|
||||||
|
|
||||||
|
try:
|
||||||
|
from transformers import (
|
||||||
|
AutoTokenizer,
|
||||||
|
PreTrainedTokenizer,
|
||||||
|
PreTrainedTokenizerBase,
|
||||||
|
PreTrainedTokenizerFast,
|
||||||
|
)
|
||||||
|
|
||||||
|
_AutoTokenizer = AutoTokenizer
|
||||||
|
_PreTrainedTokenizer = PreTrainedTokenizer
|
||||||
|
_PreTrainedTokenizerBase = PreTrainedTokenizerBase
|
||||||
|
_PreTrainedTokenizerFast = PreTrainedTokenizerFast
|
||||||
|
_transformers_available = True
|
||||||
|
except ImportError:
|
||||||
|
_transformers_available = False
|
||||||
|
|
||||||
|
return _transformers_available
|
||||||
|
|
||||||
|
|
||||||
|
def check_gguf_file(model_path: str) -> bool:
|
||||||
|
"""Check if the model path points to a GGUF file."""
|
||||||
|
if not isinstance(model_path, str):
|
||||||
|
return False
|
||||||
|
return model_path.endswith(".gguf")
|
||||||
|
|
||||||
|
|
||||||
|
def is_remote_url(path: str) -> bool:
|
||||||
|
"""Check if the path is a remote URL."""
|
||||||
|
if not isinstance(path, str):
|
||||||
|
return False
|
||||||
|
return path.startswith("http://") or path.startswith("https://")
|
||||||
|
|
||||||
|
|
||||||
|
def get_tokenizer(
|
||||||
|
tokenizer_name: str,
|
||||||
|
*args,
|
||||||
|
tokenizer_mode: str = "auto",
|
||||||
|
trust_remote_code: bool = False,
|
||||||
|
tokenizer_revision: str | None = None,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
"""Gets a tokenizer for the given model name via Huggingface.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
tokenizer_name: Name or path of the tokenizer
|
||||||
|
tokenizer_mode: Mode for tokenizer loading ("auto", "slow")
|
||||||
|
trust_remote_code: Whether to trust remote code
|
||||||
|
tokenizer_revision: Specific revision to use
|
||||||
|
**kwargs: Additional arguments passed to AutoTokenizer.from_pretrained
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Loaded tokenizer instance
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ImportError: If transformers is not installed
|
||||||
|
RuntimeError: If tokenizer loading fails
|
||||||
|
"""
|
||||||
|
if not _ensure_transformers():
|
||||||
|
raise ImportError(
|
||||||
|
"transformers is required for tokenizer utilities. "
|
||||||
|
"Install with: pip install transformers"
|
||||||
|
)
|
||||||
|
|
||||||
|
if tokenizer_mode == "slow":
|
||||||
|
if kwargs.get("use_fast", False):
|
||||||
|
raise ValueError("Cannot use the fast tokenizer in slow tokenizer mode.")
|
||||||
|
kwargs["use_fast"] = False
|
||||||
|
|
||||||
|
# Handle special model name mapping
|
||||||
|
if tokenizer_name == "mistralai/Devstral-Small-2505":
|
||||||
|
tokenizer_name = "mistralai/Mistral-Small-3.1-24B-Instruct-2503"
|
||||||
|
|
||||||
|
is_gguf = check_gguf_file(tokenizer_name)
|
||||||
|
if is_gguf:
|
||||||
|
kwargs["gguf_file"] = tokenizer_name
|
||||||
|
tokenizer_name = str(Path(tokenizer_name).parent)
|
||||||
|
|
||||||
|
try:
|
||||||
|
tokenizer = _AutoTokenizer.from_pretrained(
|
||||||
|
tokenizer_name,
|
||||||
|
*args,
|
||||||
|
trust_remote_code=trust_remote_code,
|
||||||
|
tokenizer_revision=tokenizer_revision,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
except TypeError as e:
|
||||||
|
err_msg = (
|
||||||
|
"Failed to load the tokenizer. If you are running a model with "
|
||||||
|
"a custom tokenizer, please set the --trust-remote-code flag."
|
||||||
|
)
|
||||||
|
raise RuntimeError(err_msg) from e
|
||||||
|
|
||||||
|
if not isinstance(tokenizer, _PreTrainedTokenizerFast):
|
||||||
|
logger.warning(
|
||||||
|
"Using a slow tokenizer. This might cause a performance "
|
||||||
|
"degradation. Consider using a fast tokenizer instead."
|
||||||
|
)
|
||||||
|
|
||||||
|
return tokenizer
|
||||||
|
|
||||||
|
|
||||||
|
def get_tokenizer_from_processor(processor):
|
||||||
|
"""Extract tokenizer from a processor object."""
|
||||||
|
if not _ensure_transformers():
|
||||||
|
raise ImportError("transformers is required for tokenizer utilities.")
|
||||||
|
|
||||||
|
if isinstance(processor, _PreTrainedTokenizerBase):
|
||||||
|
return processor
|
||||||
|
return processor.tokenizer
|
||||||
|
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# Pytest Utilities
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
def pytest_retry(max_retries: int = 3):
|
||||||
|
"""Decorator for pytest test functions with retry support.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
max_retries: Maximum number of retry attempts
|
||||||
|
|
||||||
|
Example:
|
||||||
|
@pytest_retry(max_retries=3)
|
||||||
|
def test_flaky_operation():
|
||||||
|
# Test that might occasionally fail
|
||||||
|
pass
|
||||||
|
"""
|
||||||
|
import functools
|
||||||
|
|
||||||
|
def decorator(func):
|
||||||
|
@functools.wraps(func)
|
||||||
|
def wrapper(*args, **kwargs):
|
||||||
|
last_exception = None
|
||||||
|
for attempt in range(max_retries + 1):
|
||||||
|
try:
|
||||||
|
return func(*args, **kwargs)
|
||||||
|
except Exception as e:
|
||||||
|
last_exception = e
|
||||||
|
if attempt < max_retries:
|
||||||
|
logger.info(
|
||||||
|
"Test %s failed on attempt %d/%d, retrying...",
|
||||||
|
func.__name__,
|
||||||
|
attempt + 1,
|
||||||
|
max_retries + 1,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
raise last_exception
|
||||||
|
|
||||||
|
return wrapper
|
||||||
|
|
||||||
|
return decorator
|
||||||
|
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# Environment Utilities
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
def is_ci_environment() -> bool:
|
||||||
|
"""Check if running in CI environment."""
|
||||||
|
ci_vars = ["CI", "GITHUB_ACTIONS", "JENKINS_URL", "GITLAB_CI", "CIRCLECI"]
|
||||||
|
return any(os.environ.get(var) for var in ci_vars)
|
||||||
|
|
||||||
|
|
||||||
|
def get_test_timeout() -> int:
|
||||||
|
"""Get test timeout from environment or default."""
|
||||||
|
return int(os.environ.get("E2E_TEST_TIMEOUT", str(DEFAULT_TIMEOUT)))
|
||||||
|
|
||||||
|
|
||||||
|
def skip_if_no_gpu():
|
||||||
|
"""Skip test if no GPU is available."""
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
try:
|
||||||
|
import torch
|
||||||
|
|
||||||
|
if not torch.cuda.is_available():
|
||||||
|
pytest.skip("No GPU available")
|
||||||
|
except ImportError:
|
||||||
|
# Try nvidia-ml-py
|
||||||
|
try:
|
||||||
|
import pynvml
|
||||||
|
|
||||||
|
pynvml.nvmlInit()
|
||||||
|
count = pynvml.nvmlDeviceGetCount()
|
||||||
|
pynvml.nvmlShutdown()
|
||||||
|
if count == 0:
|
||||||
|
pytest.skip("No GPU available")
|
||||||
|
except Exception:
|
||||||
|
pytest.skip("Cannot detect GPU (torch and pynvml not available)")
|
||||||
Reference in New Issue
Block a user