diff --git a/sgl-model-gateway/e2e_test/backends.py b/sgl-model-gateway/e2e_test/backends.py new file mode 100644 index 000000000..06651b80d --- /dev/null +++ b/sgl-model-gateway/e2e_test/backends.py @@ -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) diff --git a/sgl-model-gateway/e2e_test/benchmarks/__init__.py b/sgl-model-gateway/e2e_test/benchmarks/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/sgl-model-gateway/e2e_test/chat_completions/__init__.py b/sgl-model-gateway/e2e_test/chat_completions/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/sgl-model-gateway/e2e_test/conftest.py b/sgl-model-gateway/e2e_test/conftest.py index b41906fe0..e535781f3 100644 --- a/sgl-model-gateway/e2e_test/conftest.py +++ b/sgl-model-gateway/e2e_test/conftest.py @@ -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() diff --git a/sgl-model-gateway/e2e_test/embeddings/__init__.py b/sgl-model-gateway/e2e_test/embeddings/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/sgl-model-gateway/e2e_test/infra/__init__.py b/sgl-model-gateway/e2e_test/infra/__init__.py index 0c06d7c01..82912cdbc 100644 --- a/sgl-model-gateway/e2e_test/infra/__init__.py +++ b/sgl-model-gateway/e2e_test/infra/__init__.py @@ -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", ] diff --git a/sgl-model-gateway/e2e_test/infra/model_specs.py b/sgl-model-gateway/e2e_test/infra/model_specs.py index 519a3a7f2..090d4b468 100644 --- a/sgl-model-gateway/e2e_test/infra/model_specs.py +++ b/sgl-model-gateway/e2e_test/infra/model_specs.py @@ -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"] diff --git a/sgl-model-gateway/e2e_test/router/__init__.py b/sgl-model-gateway/e2e_test/router/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/sgl-model-gateway/e2e_test/utils.py b/sgl-model-gateway/e2e_test/utils.py new file mode 100644 index 000000000..9a8640528 --- /dev/null +++ b/sgl-model-gateway/e2e_test/utils.py @@ -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)")