[model-gateway][e2e_test]: Create directory structure and backends config (#16469)

This commit is contained in:
Simo Lin
2026-01-05 08:01:53 -08:00
committed by GitHub
parent a3914e3b3f
commit 2b4d6d813b
9 changed files with 909 additions and 1 deletions

View File

@@ -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)

View File

@@ -45,6 +45,10 @@ def pytest_configure(config: pytest.Config) -> None:
"markers",
"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(
"markers",
"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)
except KeyError:
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()

View File

@@ -11,7 +11,21 @@ from .gpu_allocator import (
wait_for_gpu_memory_to_clear,
)
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__ = [
# GPU allocation
@@ -28,4 +42,18 @@ __all__ = [
"ModelInstance",
"ModelPool",
"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",
]

View File

@@ -107,3 +107,17 @@ CHAT_MODELS = get_models_with_feature("chat")
EMBEDDING_MODELS = get_models_with_feature("embedding")
REASONING_MODELS = get_models_with_feature("reasoning")
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"]

View File

@@ -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)")