From cb21af38ba76ff1b131f6d1c5fddb716d80afc5d Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Fri, 21 Aug 2026 23:16:08 +0800 Subject: [PATCH] feat: unify Docker serving configuration in YAML - Add server.py --config serve.yaml; explicit CLI flags override YAML - Add scripts/serve.sh and serve_runtime.py for the Compose lifecycle - Template server/cpu ports and param mounts in docker-compose.yml - Document schema in docs/developer/docker-serving.md and params guide - Add tests for runtime parsing and server CLI merge logic --- README.md | 3 + docker-compose.yml | 10 +- docs/developer/docker-serving.md | 108 ++++++++++++++++++ docs/guides/inference.md | 25 ++++ docs/guides/params.md | 17 +++ scripts/serve.sh | 189 +++++++++++++++++++++++++++++++ scripts/tools/serve_runtime.py | 154 +++++++++++++++++++++++++ scripts/tools/server.py | 130 ++++++++++++++++++++- tests/test_serve_cli.py | 105 +++++++++++++++++ tests/test_serve_runtime.py | 82 ++++++++++++++ 10 files changed, 817 insertions(+), 6 deletions(-) create mode 100644 docs/developer/docker-serving.md create mode 100644 scripts/serve.sh create mode 100644 scripts/tools/serve_runtime.py create mode 100644 tests/test_serve_cli.py create mode 100644 tests/test_serve_runtime.py diff --git a/README.md b/README.md index eafb123..98bc401 100644 --- a/README.md +++ b/README.md @@ -191,6 +191,9 @@ docker compose up -d # Docker Compose CPU server profile (CUDA-only generation scripts/demos are unavailable) docker compose --profile cpu up -d + +# YAML-driven serving (see serve.yaml; up/run/down/logs/status...) +bash scripts/serve.sh up ``` > **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`. diff --git a/docker-compose.yml b/docker-compose.yml index f61d9bd..98ed2c4 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -9,9 +9,11 @@ services: USER_GID: ${ASTRAI_GID:-1000} user: "${ASTRAI_UID:-1000}:${ASTRAI_GID:-1000}" ports: - - "8000:8000" + - "${SERVE_PORT:-8000}:${SERVE_CONTAINER_PORT:-8000}" volumes: - - ./params:/app/params:ro + - ${SERVE_PARAM_DIR:-./params}:/app/params:ro + environment: + - CUDA_VISIBLE_DEVICES command: python -m scripts.tools.server --port 8000 --device cuda deploy: resources: @@ -39,9 +41,9 @@ services: USER_GID: ${ASTRAI_GID:-1000} user: "${ASTRAI_UID:-1000}:${ASTRAI_GID:-1000}" ports: - - "8000:8000" + - "${SERVE_PORT:-8000}:${SERVE_CONTAINER_PORT:-8000}" volumes: - - ./params:/app/params:ro + - ${SERVE_PARAM_DIR:-./params}:/app/params:ro command: python -m scripts.tools.server --port 8000 --device cpu healthcheck: test: ["CMD", "curl", "-f", "http://localhost:8000/health"] diff --git a/docs/developer/docker-serving.md b/docs/developer/docker-serving.md new file mode 100644 index 0000000..6e0c917 --- /dev/null +++ b/docs/developer/docker-serving.md @@ -0,0 +1,108 @@ +# Containerized Serving Deployment + +AstrAI uses one serving YAML as the declaration for both host-side container +runtime settings and in-container server settings. `scripts/serve.sh` wraps the +Compose commands so preflight validation and container lifecycle stay +consistent with the trainer. + +## Architecture + +```text +serve.yaml + ├── runtime parsed on the host before Docker starts + └── server parsed by server.py inside the container + │ +scripts/serve.sh preflight, Compose wrapper, lifecycle + └── docker-compose.yml GPU passthrough, mounts, image, port mapping + └── server.py --config /run/astrai/serve.yaml +``` + +`scripts/tools/serve_runtime.py` reads `runtime:` plus the two container-side +values Compose needs (`server.port` for the port mapping, `server.device` for +the preflight GPU check). `scripts/tools/server.py --config` reads `server:`. +Explicit CLI arguments to `server.py` override `server:` YAML values. + +## Runtime Schema + +```yaml +runtime: + job_name: serve + port: 8000 + paths: + param: ./params + gpu: + enabled: true # false → cpu profile (server-cpu service) + devices: all # all | [0] + container: + cuda_tag: cu128 + # environment: + # TOKENIZERS_PARALLELISM: "false" + +server: + host: 0.0.0.0 + port: 8000 + device: cuda # cuda | cpu + dtype: bfloat16 # bfloat16 | float16 | float32 + max_batch_size: 16 + max_seq_len: null # falls back to model config +``` + +- Relative paths resolve from the YAML file's directory, not the current shell. +- `runtime.port` is the host publish port; `server.port` is the port the + container listens on. The Compose mapping is + `${SERVE_PORT}:${SERVE_CONTAINER_PORT}`. +- `runtime.gpu.enabled: true` (default) selects the `server` service with an + NVIDIA device reservation; `false` selects `server-cpu` (no GPU passthrough). + When disabled, `server.device` must be `cpu`. +- `runtime.gpu.devices` is either `all` or a single-device list such as `[0]`; + the list becomes `CUDA_VISIBLE_DEVICES`. Compose passes `count: 1`. +- `environment` values are explicitly passed to the serving container. Keep + host-specific settings here; they are not universal defaults. +- `server.device` must agree with `runtime.gpu.enabled`; `preflight` enforces it. + +## Fixed Container Paths + +| Runtime path | Container path | Access | +|---|---|---| +| `runtime.paths.param` | `/app/params` | read-only | +| the selected YAML | `/run/astrai/serve.yaml` | read-only | + +`server.param_path` is optional: the server default is +`project_root/params`, which is exactly `/app/params` inside the container +(the working directory is `/app`). Set it explicitly only when serving from a +different location; in Docker it must be a container path. + +## Operations + +The config argument defaults to `./serve.yaml`: + +```bash +bash scripts/serve.sh init [CONFIG] +bash scripts/serve.sh preflight [CONFIG] +bash scripts/serve.sh up [CONFIG] +bash scripts/serve.sh run [CONFIG] +bash scripts/serve.sh down [CONFIG] +bash scripts/serve.sh restart [CONFIG] +bash scripts/serve.sh logs [CONFIG] +bash scripts/serve.sh status [CONFIG] +``` + +`preflight` validates Docker, the model directory +(`config.json` + `model.safetensors`), GPU/device consistency, and the +rendered Compose configuration. `up` starts the container detached and +rebuilds the image when the code changed (`--build`); `run` keeps it in the +foreground. The wrapper manages a fixed container name +(`astrai-server` or `astrai-server-`); the plain +`docker compose up -d` / `docker compose --profile cpu up -d` path keeps +working with defaults (port 8000, `./params`). + +## Hard Rules + +1. Keep Docker settings in `runtime` and server settings in `server`. +2. Filter GPUs once: the `server` service reserves one device; a `devices` + list becomes `CUDA_VISIBLE_DEVICES`. +3. `runtime.gpu.enabled: false` requires `server.device: cpu`. +4. In Docker, `server.port` must match the published container port (default + `8000`); change `runtime.port` to publish on a different host port. +5. The image user is built with the host UID/GID so the mounted model + directory stays readable. diff --git a/docs/guides/inference.md b/docs/guides/inference.md index 8f9ab4d..bfc1a87 100644 --- a/docs/guides/inference.md +++ b/docs/guides/inference.md @@ -171,6 +171,31 @@ InferenceEngine `GenerateResult` uses `Condition` for non-streaming (`wait_completion()`) and `Event` for streaming (`wait()`). Stream callback is `cb(token)`. +## Launching the Server + +`scripts/tools/server.py` accepts every option as a CLI flag or from a YAML +config file (`--config serve.yaml`); explicit CLI flags override YAML values. +The YAML `server:` section mirrors the flags: + +```yaml +server: + host: 0.0.0.0 + port: 8000 + device: cuda + dtype: bfloat16 + max_batch_size: 16 + max_seq_len: null +``` + +```bash +python scripts/tools/server.py --config serve.yaml +python scripts/tools/server.py --config serve.yaml --port 9000 # CLI wins +``` + +In Docker, `scripts/serve.sh` drives the same YAML (a `runtime:` section +controls ports/GPU/mounts); see +[Docker Serving](../developer/docker-serving.md). + ## HTTP API ``` diff --git a/docs/guides/params.md b/docs/guides/params.md index 578be8c..0d7261a 100644 --- a/docs/guides/params.md +++ b/docs/guides/params.md @@ -203,6 +203,7 @@ nohup python scripts/tools/train.py \ | Parameter | Type | Default | Description | |-----------|------|---------|-------------| +| `--config`, `-c` | path | `None` | Serving YAML config. CLI flags override YAML values | | `--host` | str | `0.0.0.0` | Host address | | `--port` | int | `8000` | Port number | | `--param_path` | path | `project_root/params` | Path to model parameters | @@ -217,6 +218,22 @@ Usage: python scripts/tools/server.py --param_path ./params --device cuda --dtype bfloat16 ``` +YAML config (a `server:` section; explicit CLI flags override YAML values): +```bash +python scripts/tools/server.py --config serve.yaml +``` +```yaml +server: + host: 0.0.0.0 + port: 8000 + device: cuda + dtype: bfloat16 + max_batch_size: 16 + max_seq_len: null +``` +`serve.yaml` also carries a `runtime:` section for the Docker wrapper; see +[Docker Serving](../developer/docker-serving.md). + See [Inference Guide](inference.md) for HTTP API documentation. ## Generate (`generate.py`) diff --git a/scripts/serve.sh b/scripts/serve.sh new file mode 100644 index 0000000..0dade2d --- /dev/null +++ b/scripts/serve.sh @@ -0,0 +1,189 @@ +#!/usr/bin/env bash +set -euo pipefail + +ROOT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")/.." && pwd)" +source "${ROOT_DIR}/scripts/docker/lib/train-common.sh" + +COMPOSE_BASE=( + docker compose + --project-directory "${ROOT_DIR}" + --file "${ROOT_DIR}/docker-compose.yml" +) + +usage() { + cat <<'EOF' +Usage: scripts/serve.sh [CONFIG] [options] + +CONFIG defaults to ./serve.yaml. The same file declares host runtime settings +under `runtime:` and server settings under `server:`. + +Commands: + init [CONFIG] Create the model directory + preflight [CONFIG] Validate Docker, paths, GPU, and Compose + build [CONFIG] Build the serving image + up [CONFIG] Start the server container (detached) + run [CONFIG] Start the server container (foreground) + down [CONFIG] Stop and remove the server container + restart [CONFIG] Down, then up + logs [CONFIG] Follow server logs + status [CONFIG] Show container status +EOF +} + +resolve_path() { + if [[ "$1" = /* ]]; then + printf '%s\n' "$1" + else + printf '%s/%s\n' "${ROOT_DIR}" "${1#./}" + fi +} + +load_config() { + CONFIG_FILE="$(resolve_path "$1")" + [[ -f "${CONFIG_FILE}" ]] || die "Serving config not found: ${CONFIG_FILE}" + require_command python3 + python3 -c 'import yaml' >/dev/null 2>&1 || + die "PyYAML is required on the host (install python3-yaml)" + + local exports + exports="$(python3 "${ROOT_DIR}/scripts/tools/serve_runtime.py" exports "${CONFIG_FILE}")" || + die "Failed to load runtime configuration" + eval "${exports}" + if [[ -n "${SERVE_JOB_NAME}" ]]; then + validate_job_name "${SERVE_JOB_NAME}" + fi +} + +compose() { + ASTRAI_UID="$(id -u)" ASTRAI_GID="$(id -g)" "${COMPOSE_BASE[@]}" "$@" +} + +container_name() { + if [[ -n "${SERVE_JOB_NAME}" ]]; then + printf 'astrai-server-%s\n' "${SERVE_JOB_NAME}" + else + printf 'astrai-server\n' + fi +} + +service_name() { + if [[ "${SERVE_GPU_ENABLED:-true}" == "false" ]]; then + printf 'server-cpu\n' + else + printf 'server\n' + fi +} + +set_profile_args() { + PROFILE_ARGS=() + if [[ "${SERVE_GPU_ENABLED:-true}" == "false" ]]; then + PROFILE_ARGS=(--profile cpu) + fi +} + +init_environment() { + mkdir -p "${SERVE_PARAM_DIR}" + log_info "Model: ${SERVE_PARAM_DIR}" +} + +preflight() { + require_command docker + docker info >/dev/null 2>&1 || die "Docker daemon is unavailable" + [[ -d "${SERVE_PARAM_DIR}" ]] || die "Model directory not found: ${SERVE_PARAM_DIR}" + [[ -s "${SERVE_PARAM_DIR}/config.json" ]] || + die "Model config not found: ${SERVE_PARAM_DIR}/config.json" + [[ -s "${SERVE_PARAM_DIR}/model.safetensors" ]] || + die "Model weights not found: ${SERVE_PARAM_DIR}/model.safetensors" + + if [[ "${SERVE_GPU_ENABLED}" == "false" ]] && [[ "${SERVE_DEVICE}" != "cpu" ]]; then + die "runtime.gpu.enabled is false but server.device is '${SERVE_DEVICE}'; use server.device: cpu" + fi + + compose config --quiet + log_info "Preflight passed (service: $(service_name), device: ${SERVE_DEVICE})" +} + +runtime_environment_args() { + RUNTIME_ENV_ARGS=() + local pair + while IFS= read -r -d '' pair; do + RUNTIME_ENV_ARGS+=(--env "${pair}") + done < <(python3 "${ROOT_DIR}/scripts/tools/serve_runtime.py" environment "${CONFIG_FILE}") +} + +start_server() { + local foreground="$1" + shift + local container running + local -a run_options + preflight + runtime_environment_args + set_profile_args + container="$(container_name)" + running="$(docker inspect --format '{{.State.Running}}' "${container}" 2>/dev/null || true)" + [[ "${running}" != "true" ]] || die "Server is already running: ${container}" + docker rm "${container}" >/dev/null 2>&1 || true + + run_options=( + --volume "${CONFIG_FILE}:/run/astrai/serve.yaml:ro" + "${RUNTIME_ENV_ARGS[@]}" + ) + if [[ "${foreground}" == "true" ]]; then + compose "${PROFILE_ARGS[@]}" run --build --rm --service-ports \ + "${run_options[@]}" "$(service_name)" \ + python -m scripts.tools.server --config /run/astrai/serve.yaml "$@" + else + compose "${PROFILE_ARGS[@]}" run -d --build --service-ports \ + --name "${container}" "${run_options[@]}" "$(service_name)" \ + python -m scripts.tools.server --config /run/astrai/serve.yaml "$@" + log_info "Server started; run scripts/serve.sh logs ${CONFIG_FILE} to follow it" + fi +} + +stop_server() { + local container + container="$(container_name)" + docker stop --timeout 30 "${container}" >/dev/null 2>&1 || + log_warn "Server container is not running" + docker rm "${container}" >/dev/null 2>&1 || true +} + +show_status() { + docker ps -a --filter "name=^/$(container_name)$" +} + +main() { + local command="${1:-}" config="${SERVE_CONFIG_FILE:-${ROOT_DIR}/serve.yaml}" + [[ -n "${command}" ]] || { usage; exit 1; } + shift || true + + if [[ "${command}" =~ ^(help|-h|--help)$ ]]; then + usage + return + fi + + if [[ $# -gt 0 && "$1" != --* ]]; then + config="$1" + shift + fi + load_config "${config}" + + case "${command}" in + init) init_environment ;; + preflight) preflight ;; + build) + set_profile_args + preflight + compose "${PROFILE_ARGS[@]}" build "$(service_name)" + ;; + up) start_server false "$@" ;; + run) start_server true "$@" ;; + down) stop_server ;; + restart) stop_server; start_server false ;; + logs) docker logs -f --tail "${SERVE_LOG_TAIL:-200}" "$(container_name)" ;; + status) show_status ;; + *) die "Unknown command: ${command}" ;; + esac +} + +main "$@" diff --git a/scripts/tools/serve_runtime.py b/scripts/tools/serve_runtime.py new file mode 100644 index 0000000..5b6402e --- /dev/null +++ b/scripts/tools/serve_runtime.py @@ -0,0 +1,154 @@ +"""Parse the host-side runtime section of a serving configuration. + +The Compose wrapper needs a few container-side values on the host as well: +``server.port`` (the port the container listens on) and ``server.device`` +(used by the preflight GPU consistency check). Everything else under +``server:`` is owned by ``scripts/tools/server.py --config`` inside the +container. +""" + +import argparse +import re +import shlex +from pathlib import Path + +import yaml + +ENV_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") + + +def _mapping(value, name: str) -> dict: + if value is None: + return {} + if not isinstance(value, dict): + raise ValueError(f"{name} must be a mapping") + return value + + +def _path(value, name: str, config_dir: Path) -> str: + if not isinstance(value, str) or not value.strip(): + raise ValueError(f"runtime.paths.{name} is required") + path = Path(value).expanduser() + if not path.is_absolute(): + path = config_dir / path + return str(path.resolve()) + + +def _port(value, name: str) -> int: + if isinstance(value, bool) or not isinstance(value, int): + raise ValueError(f"{name} must be an integer") + if not 1 <= value <= 65535: + raise ValueError(f"{name} must be between 1 and 65535") + return value + + +def load_runtime(config_path: str) -> dict[str, str]: + path = Path(config_path).resolve() + with path.open(encoding="utf-8") as file: + config = yaml.safe_load(file) or {} + if not isinstance(config, dict): + raise ValueError("serving configuration must be a mapping") + + runtime = _mapping(config.get("runtime"), "runtime") + if not runtime: + raise ValueError("top-level runtime section is required") + paths = _mapping(runtime.get("paths"), "paths") + gpu = _mapping(runtime.get("gpu"), "gpu") + container = _mapping(runtime.get("container"), "container") + environment = _mapping(runtime.get("environment"), "environment") + server = _mapping(config.get("server"), "server") + + job_name = runtime.get("job_name", "") + if job_name and not isinstance(job_name, str): + raise ValueError("runtime.job_name must be a string") + if job_name and not re.fullmatch(r"[A-Za-z0-9][A-Za-z0-9._-]*", job_name): + raise ValueError( + "runtime.job_name must use letters, numbers, dot, underscore, or dash" + ) + + port = _port(runtime.get("port", 8000), "runtime.port") + container_port = _port(server.get("port", 8000), "server.port") + + device = server.get("device", "cuda") + if not isinstance(device, str) or not device.strip(): + raise ValueError("server.device must be a string") + + gpu_enabled = gpu.get("enabled", True) + if not isinstance(gpu_enabled, bool): + raise ValueError("runtime.gpu.enabled must be a boolean") + + devices = gpu.get("devices", "all") + if gpu_enabled: + if devices == "all": + visible_devices = "" + elif isinstance(devices, list) and len(devices) == 1: + text = str(devices[0]) + if not text.isdigit(): + raise ValueError( + "runtime.gpu.devices entries must be non-negative integers" + ) + visible_devices = text + else: + raise ValueError( + "runtime.gpu.devices must be 'all' or a single-device list such as [0]" + ) + else: + visible_devices = "" + if devices != "all": + raise ValueError( + "runtime.gpu.devices is ignored when runtime.gpu.enabled is false" + ) + if device != "cpu": + raise ValueError( + "server.device must be 'cpu' when runtime.gpu.enabled is false" + ) + + values = { + "SERVE_JOB_NAME": job_name, + "SERVE_PORT": str(port), + "SERVE_CONTAINER_PORT": str(container_port), + "SERVE_PARAM_DIR": _path(paths.get("param", "./params"), "param", path.parent), + "SERVE_GPU_ENABLED": "true" if gpu_enabled else "false", + "SERVE_DEVICE": device, + "CUDA_VISIBLE_DEVICES": visible_devices, + "CUDA_TAG": str(container.get("cuda_tag", "cu128")), + } + + for name, value in environment.items(): + if not isinstance(name, str) or not ENV_NAME.fullmatch(name): + raise ValueError(f"invalid runtime.environment name: {name!r}") + if value is not None and not isinstance(value, (str, int, float, bool)): + raise ValueError(f"runtime.environment.{name} must be a scalar") + values["environment"] = environment + return values + + +def shell_exports(runtime: dict[str, str]) -> str: + return "\n".join( + f"export {name}={shlex.quote(value)}" + for name, value in runtime.items() + if name != "environment" + ) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("command", choices=("exports", "environment")) + parser.add_argument("config") + args = parser.parse_args() + + try: + runtime = load_runtime(args.config) + except (OSError, ValueError, yaml.YAMLError) as exc: + parser.error(str(exc)) + + if args.command == "exports": + print(shell_exports(runtime)) + return + for name, value in runtime["environment"].items(): + rendered = "" if value is None else str(value) + print(f"{name}={rendered}", end="\0") + + +if __name__ == "__main__": + main() diff --git a/scripts/tools/server.py b/scripts/tools/server.py index 66c62cb..82b537a 100644 --- a/scripts/tools/server.py +++ b/scripts/tools/server.py @@ -2,16 +2,105 @@ from pathlib import Path import click import torch +import yaml +from click.core import ParameterSource from astrai.inference import run_server _DTYPES = ["bfloat16", "float16", "float32"] +_SERVER_KEYS = ( + "host", + "port", + "reload", + "param_path", + "device", + "dtype", + "max_batch_size", + "max_seq_len", +) + + +def _merge_yaml_into_kwargs( + config_path: str, + passed_kwargs: dict, + explicit_keys: set[str] | None = None, +) -> dict: + """Merge Click defaults, YAML server values, then explicit CLI values.""" + with open(config_path, encoding="utf-8") as file: + config = yaml.safe_load(file) or {} + if not isinstance(config, dict): + raise click.UsageError(f"Serving config must be a mapping: {config_path}") + server = config.get("server") or {} + if not isinstance(server, dict): + raise click.UsageError("top-level server section must be a mapping") + + unknown = sorted(set(server) - set(_SERVER_KEYS)) + if unknown: + click.echo( + f"Warning: ignoring unknown server config keys: {', '.join(unknown)}", + err=True, + ) + + merged = dict(passed_kwargs) + merged.update({key: server[key] for key in _SERVER_KEYS if key in server}) + if explicit_keys is None: + explicit_keys = set(passed_kwargs) + for key in explicit_keys: + if key in passed_kwargs: + merged[key] = passed_kwargs[key] + return merged + + +def _as_int(value, name: str) -> int | None: + if value is None: + return None + if isinstance(value, bool): + raise click.UsageError(f"{name} must be an integer") + try: + return int(value) + except (TypeError, ValueError): + raise click.UsageError(f"{name} must be an integer, got {value!r}") from None + + +def _resolve_server_config( + config_path: str, + passed_kwargs: dict, + explicit_keys: set[str] | None = None, +) -> dict: + """Merge YAML values, then coerce and validate the resolved settings. + + ``explicit_keys`` are CLI flags that win over YAML; when None, YAML values + win over Click defaults. + """ + merged = _merge_yaml_into_kwargs(config_path, passed_kwargs, explicit_keys or set()) + resolved = dict(merged) + resolved["port"] = _as_int(resolved["port"], "server.port") or 8000 + resolved["max_batch_size"] = ( + _as_int(resolved["max_batch_size"], "server.max_batch_size") or 16 + ) + resolved["max_seq_len"] = _as_int(resolved["max_seq_len"], "server.max_seq_len") + resolved["reload"] = bool(resolved["reload"]) + if resolved["dtype"] not in _DTYPES: + raise click.UsageError( + f"server.dtype must be one of {', '.join(_DTYPES)}, got {resolved['dtype']!r}" + ) + return resolved @click.command(name="serve", help="Launch inference server (OpenAI-compatible API).") +@click.option( + "--config", + "-c", + "config_path", + type=click.Path(exists=True, dir_okay=False), + default=None, + help="Serving YAML config. CLI flags override YAML values.", +) @click.option("--host", default="0.0.0.0", help="Host address.") @click.option("--port", type=int, default=8000, help="Port number.") -@click.option("--reload", is_flag=True, help="Enable auto-reload for development.") +@click.option( + "--reload", is_flag=True, default=False, help="Enable auto-reload for development." +) @click.option( "--param_path", type=click.Path(exists=True), @@ -37,10 +126,47 @@ _DTYPES = ["bfloat16", "float16", "float32"] default=None, help="Maximum sequence length (KV cache size + prompt truncation). Uses model config if not set.", ) +@click.pass_context def server_command( - host, port, reload, param_path, device, dtype, max_batch_size, max_seq_len + ctx, + config_path, + host, + port, + reload, + param_path, + device, + dtype, + max_batch_size, + max_seq_len, ): """Launch inference server (OpenAI-compatible API).""" + if config_path: + passed_kwargs = { + "host": host, + "port": port, + "reload": reload, + "param_path": param_path, + "device": device, + "dtype": dtype, + "max_batch_size": max_batch_size, + "max_seq_len": max_seq_len, + } + explicit_keys = { + key + for key in passed_kwargs + if ctx.get_parameter_source(key) is ParameterSource.COMMANDLINE + } + resolved = _resolve_server_config(config_path, passed_kwargs, explicit_keys) + host = resolved["host"] + port = resolved["port"] + reload = resolved["reload"] + param_path = resolved["param_path"] + device = resolved["device"] + dtype = resolved["dtype"] + max_batch_size = resolved["max_batch_size"] + max_seq_len = resolved["max_seq_len"] + click.echo(f"Config: {config_path}") + dtype_map = { "bfloat16": torch.bfloat16, "float16": torch.float16, diff --git a/tests/test_serve_cli.py b/tests/test_serve_cli.py new file mode 100644 index 0000000..46a8443 --- /dev/null +++ b/tests/test_serve_cli.py @@ -0,0 +1,105 @@ +"""Unit tests for the serving CLI YAML merge logic.""" + +import click +import pytest +import torch +from click.testing import CliRunner + +from scripts.tools.server import ( + _merge_yaml_into_kwargs, + _resolve_server_config, + server_command, +) + + +def _passed() -> dict: + return { + "host": "0.0.0.0", + "port": 8000, + "reload": False, + "param_path": None, + "device": "cuda", + "dtype": "bfloat16", + "max_batch_size": 16, + "max_seq_len": None, + } + + +def test_yaml_overrides_click_defaults_but_not_explicit_cli(tmp_path): + config_path = tmp_path / "serve.yaml" + config_path.write_text( + "server:\n device: cpu\n dtype: float16\n max_batch_size: 8\n", + encoding="utf-8", + ) + merged = _merge_yaml_into_kwargs( + str(config_path), _passed(), explicit_keys={"device"} + ) + + assert merged["device"] == "cuda" + assert merged["dtype"] == "float16" + assert merged["max_batch_size"] == 8 + + +def test_resolve_config_yaml_wins_by_default(tmp_path): + config_path = tmp_path / "serve.yaml" + config_path.write_text( + "server:\n port: 9000\n max_seq_len: 2048\n", + encoding="utf-8", + ) + resolved = _resolve_server_config(str(config_path), _passed()) + + assert resolved["port"] == 9000 + assert resolved["max_seq_len"] == 2048 + assert resolved["device"] == "cuda" + assert resolved["dtype"] == "bfloat16" + + +def test_resolve_config_rejects_bad_dtype(tmp_path): + config_path = tmp_path / "serve.yaml" + config_path.write_text("server:\n dtype: fp8\n", encoding="utf-8") + + with pytest.raises(click.UsageError, match="server.dtype"): + _resolve_server_config(str(config_path), _passed()) + + +def test_server_command_rejects_bad_yaml_dtype(tmp_path): + config_path = tmp_path / "serve.yaml" + config_path.write_text("server:\n dtype: fp8\n", encoding="utf-8") + + result = CliRunner().invoke(server_command, ["--config", str(config_path)]) + + assert result.exit_code == 2 + assert "server.dtype" in result.output + + +def test_server_command_merges_yaml_and_cli(tmp_path, monkeypatch): + """Full CLI path: YAML values apply, explicit CLI flags override, args reach run_server.""" + config_path = tmp_path / "serve.yaml" + config_path.write_text( + "server:\n device: cpu\n dtype: float16\n max_batch_size: 8\n", + encoding="utf-8", + ) + captured = {} + + def fake_run_server(**kwargs): + captured.update(kwargs) + + monkeypatch.setattr("scripts.tools.server.run_server", fake_run_server) + result = CliRunner().invoke( + server_command, + ["--config", str(config_path), "--max_batch_size", "32"], + ) + + assert result.exit_code == 0, result.output + assert captured["device"] == "cpu" + assert captured["dtype"] == torch.float16 + assert captured["max_batch_size"] == 32 + assert captured["port"] == 8000 + + +def test_config_option_rejects_missing_file(tmp_path): + result = CliRunner().invoke( + server_command, ["--config", str(tmp_path / "nope.yaml")] + ) + + assert result.exit_code == 2 diff --git a/tests/test_serve_runtime.py b/tests/test_serve_runtime.py new file mode 100644 index 0000000..6ada9f4 --- /dev/null +++ b/tests/test_serve_runtime.py @@ -0,0 +1,82 @@ +"""Unit tests for the serving runtime configuration parser.""" + +import pytest + +from scripts.tools.serve_runtime import load_runtime + + +def _write(tmp_path, body: str) -> str: + config_path = tmp_path / "serve.yaml" + config_path.write_text(body, encoding="utf-8") + return str(config_path) + + +def test_runtime_exports_defaults(tmp_path): + config_path = _write( + tmp_path, + "runtime:\n" + " port: 8000\n" + " paths:\n" + " param: ./params\n" + "server:\n" + " device: cuda\n", + ) + runtime = load_runtime(config_path) + assert runtime["SERVE_PORT"] == "8000" + assert runtime["SERVE_CONTAINER_PORT"] == "8000" + assert runtime["SERVE_PARAM_DIR"] == str((tmp_path / "params").resolve()) + assert runtime["SERVE_GPU_ENABLED"] == "true" + assert runtime["CUDA_VISIBLE_DEVICES"] == "" + assert runtime["SERVE_DEVICE"] == "cuda" + assert runtime["CUDA_TAG"] == "cu128" + assert runtime["SERVE_JOB_NAME"] == "" + + +def test_runtime_gpu_disabled_requires_cpu(tmp_path): + config_path = _write( + tmp_path, + "runtime:\n gpu:\n enabled: false\nserver:\n device: cuda\n", + ) + with pytest.raises(ValueError, match="server.device must be 'cpu'"): + load_runtime(config_path) + + +def test_runtime_gpu_devices_single_and_ports(tmp_path): + config_path = _write( + tmp_path, + "runtime:\n" + " gpu:\n" + " devices: [1]\n" + " port: 8080\n" + "server:\n" + " port: 9000\n" + " device: cuda\n", + ) + runtime = load_runtime(config_path) + assert runtime["SERVE_PORT"] == "8080" + assert runtime["SERVE_CONTAINER_PORT"] == "9000" + assert runtime["CUDA_VISIBLE_DEVICES"] == "1" + + +def test_runtime_gpu_devices_rejects_multi(tmp_path): + config_path = _write( + tmp_path, + "runtime:\n gpu:\n devices: [0, 1]\n", + ) + with pytest.raises(ValueError, match="single-device"): + load_runtime(config_path) + + +def test_runtime_port_out_of_range(tmp_path): + config_path = _write(tmp_path, "runtime:\n port: 70000\n") + with pytest.raises(ValueError, match="between 1 and 65535"): + load_runtime(config_path) + + +def test_runtime_environment_export(tmp_path): + config_path = _write( + tmp_path, + "runtime:\n environment:\n TOKENIZERS_PARALLELISM: 'false'\n", + ) + runtime = load_runtime(config_path) + assert runtime["environment"] == {"TOKENIZERS_PARALLELISM": "false"}