diff --git a/miles/ray/multi_lora/__init__.py b/miles/ray/multi_lora/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/miles/ray/multi_lora/backend.py b/miles/ray/multi_lora/backend.py new file mode 100644 index 0000000000..b8dd0eab27 --- /dev/null +++ b/miles/ray/multi_lora/backend.py @@ -0,0 +1,188 @@ +"""Multi-LoRA backend: the registry plus engine-facing aborts, shared by the +controller Ray actor and the HTTP server. Subclass via +``--multi-lora-backend-path``.""" + +import asyncio +import logging +from dataclasses import replace +from pathlib import Path +from typing import Any + +import httpx + +from miles.ray.multi_lora.registry import AdapterRegistry, AdapterState +from miles.utils.adapter_config import AdapterRunConfig +from miles.utils.multi_lora import RID_SEPARATOR, min_groups_per_dp_split + +logger = logging.getLogger(__name__) + + +class MultiLoRABackend: + """Registry + engine-facing aborts, shared by the Ray actor and HTTP server. + Subclass via --multi-lora-backend-path.""" + + def __init__(self, args: Any, router_url: str) -> None: + self.args = args + self.registry = AdapterRegistry(args.multi_lora_n_adapters) + self.router_url = router_url.rstrip("/") + self.client: httpx.AsyncClient | None = None + + async def init(self) -> None: + self.client = httpx.AsyncClient(timeout=httpx.Timeout(30.0)) + + async def close(self) -> None: + if self.client is not None: + await self.client.aclose() + self.client = None + + async def validate_adapter(self, name: str, config: Any) -> None: + """Override to reject adapter registrations (raise ValueError).""" + + def resolve_adapter_config(self, name: str, config: Any) -> Any: + """Resolve optional adapter-local values against process-wide defaults + and validate the batch shape against the trainer's DP layout. + + All batch-shape constraints are enforced here, at registration, so a + bad config fails immediately instead of crashing an arbitrary later + train batch. + """ + if config is None or not isinstance(config, AdapterRunConfig): + return config + + rank = config.rank if config.rank is not None else getattr(self.args, "lora_rank", 1) + alpha = config.alpha if config.alpha is not None else getattr(self.args, "lora_alpha", rank) + rollout_batch_size = ( + config.rollout_batch_size + if config.rollout_batch_size is not None + else getattr(self.args, "rollout_batch_size", None) + ) + n_samples_per_prompt = ( + config.n_samples_per_prompt + if config.n_samples_per_prompt is not None + else getattr(self.args, "n_samples_per_prompt", 1) + ) + + if type(rank) is not int or rank <= 0: + raise ValueError(f"Adapter '{name}' rank must be a positive integer") + if rank > getattr(self.args, "lora_rank", rank): + raise ValueError(f"Adapter '{name}' rank {rank} exceeds the allocated maximum rank {self.args.lora_rank}") + if alpha is None or alpha <= 0: + raise ValueError(f"Adapter '{name}' must have a positive alpha") + if type(rollout_batch_size) is not int or rollout_batch_size <= 0: + raise ValueError(f"Adapter '{name}' rollout_batch_size must be a positive integer (prompt groups)") + if type(n_samples_per_prompt) is not int or n_samples_per_prompt <= 0: + raise ValueError(f"Adapter '{name}' n_samples_per_prompt must be a positive integer") + if config.num_step is not None and (type(config.num_step) is not int or config.num_step <= 0): + raise ValueError(f"Adapter '{name}' num_step must be a positive integer") + if config.num_epoch is not None and (type(config.num_epoch) is not int or config.num_epoch <= 0): + raise ValueError(f"Adapter '{name}' num_epoch must be a positive integer") + if config.num_step is not None and config.num_epoch is not None: + logger.warning(f"Adapter '{name}' sets both num_step and num_epoch; num_step takes precedence") + + # A bad data path or unresolvable reward config does not fail at this + # API otherwise: the data path kills the shared rollout producer thread + # and an empty reward config burns every generated sample, either way + # stalling ALL adapters behind a misleading empty-batch timeout. + if not Path(config.data).expanduser().exists(): + raise ValueError( + f"Adapter '{name}' data path '{config.data}' does not exist " + "(checked from the controller process, which runs on the head node with the rollout data source)" + ) + if ( + config.custom_rm_path is None + and not (config.rm_type or "").strip() + and getattr(self.args, "custom_rm_path", None) is None + and not (getattr(self.args, "rm_type", None) or "").strip() + ): + raise ValueError( + f"Adapter '{name}' has no reward config: set rm_type or custom_rm_path in the adapter " + "config, or launch with --rm-type / --custom-rm-path" + ) + + adapter_global_batch_size = rollout_batch_size * n_samples_per_prompt + if (max_batch := getattr(self.args, "multi_lora_max_adapter_global_batch_size", None)) is not None: + if adapter_global_batch_size > max_batch: + raise ValueError( + f"Adapter '{name}' consumes {adapter_global_batch_size} samples per step " + f"(rollout_batch_size {rollout_batch_size} x n_samples_per_prompt {n_samples_per_prompt}), " + f"exceeding --multi-lora-max-adapter-global-batch-size {max_batch}" + ) + if (dp_size := getattr(self.args, "multi_lora_dp_size", None)) is not None: + try: + group_multiple = min_groups_per_dp_split(n_samples_per_prompt, dp_size) + except ValueError as e: + raise ValueError(f"Adapter '{name}': {e}") from None + if rollout_batch_size % group_multiple != 0: + raise ValueError( + f"Adapter '{name}' rollout_batch_size {rollout_batch_size} must be a multiple of " + f"its min_groups_per_dp_split ({group_multiple} at dp_size={dp_size}), so the " + f"adapter batch can complete from evenly-splitting takes" + ) + + save = Path(config.save) if config.save is not None else None + if save is None: + if getattr(self.args, "save", None) is None: + raise ValueError(f"Adapter '{name}' has no save dir: set 'save' in the adapter config or pass --save") + save = Path(self.args.save) / "adapters" / name + + return replace( + config, + rank=rank, + alpha=alpha, + rollout_batch_size=rollout_batch_size, + n_samples_per_prompt=n_samples_per_prompt, + save=save, + ) + + async def register(self, name: str, config: Any) -> dict: + config = self.resolve_adapter_config(name, config) + await self.validate_adapter(name, config) + result = self.registry.register(name, config) + resolved = getattr(config, "save", None) + if resolved is not None: + logger.info(f"Adapter '{name}' registered (slot {result['slot']}), checkpoints -> {resolved}") + return result + + async def deregister(self, name: str) -> None: + self.registry.deregister(name) + + async def retire_adapters(self) -> list[str]: + names = self.registry.retire_adapters() + for name in names: + await self.abort_adapter_requests(name) + return names + + async def free_slot(self, name: str) -> int: + """Free the adapter's slot after one final abort round: requests can survive the + ``retire_adapters`` abort (e.g. multi-turn groups), and must not leak to the slot's next tenant.""" + record = self.registry.records.get(name) + if record is not None and record.state is AdapterState.CLEANUP: + await self.abort_adapter_requests(name) + return self.registry.free_slot(name) + + async def worker_urls(self) -> list[str]: + assert self.client is not None + for endpoint, extract in ( + ("/list_workers", lambda body: body["urls"]), + ("/workers", lambda body: [worker["url"] for worker in body["workers"]]), + ): + try: + resp = await self.client.get(f"{self.router_url}{endpoint}") + if resp.status_code == 200: + return extract(resp.json()) + except Exception: + continue + return [] + + async def abort_adapter_requests(self, adapter_name: str) -> None: + prefix = f"{adapter_name}{RID_SEPARATOR}" + urls = await self.worker_urls() + if not urls: + logger.warning(f"Abort for adapter '{adapter_name}': no workers discovered at {self.router_url}") + return + results = await asyncio.gather( + *(self.client.post(f"{url}/abort_request", json={"rid": prefix, "prefix": True}) for url in urls), + return_exceptions=True, + ) + if failures := sum(isinstance(r, Exception) for r in results): + logger.warning(f"Abort for adapter '{adapter_name}': {failures}/{len(results)} posts failed") diff --git a/miles/ray/multi_lora/controller.py b/miles/ray/multi_lora/controller.py new file mode 100644 index 0000000000..7cbff2b5b9 --- /dev/null +++ b/miles/ray/multi_lora/controller.py @@ -0,0 +1,122 @@ +"""Named Ray actor wrapping the multi-LoRA backend + HTTP server.""" + +import time +from functools import cache +from typing import Any + +import ray + +from miles.ray.multi_lora.backend import MultiLoRABackend +from miles.ray.multi_lora.http_server import MultiLoRAHTTPServer +from miles.utils.adapter_config import AdapterRun +from miles.utils.misc import SingletonMeta, get_current_node_ip, load_function +from miles.utils.ray_utils import compute_ray_pin_head_options + +CONTROLLER_NAME = "miles_multi_lora_controller" +CONTROLLER_NAMESPACE = "miles" + + +@cache +def get_multi_lora_controller(): + return ray.get_actor(CONTROLLER_NAME, namespace=CONTROLLER_NAMESPACE) + + +class AdaptersCache(metaclass=SingletonMeta): + """TTL-cached controller snapshot; get/get_all expose the sampleable + projection (active + retiring).""" + + def __init__(self, ttl_s: float = 1.0) -> None: + self.ttl_s = ttl_s + self.snapshot: dict = {"pending": {}, "active": {}, "retiring": {}, "cleanup": []} + self.last_refresh: float | None = None + + async def get_snapshot(self) -> dict: + now = time.monotonic() + if self.last_refresh is None or now - self.last_refresh >= self.ttl_s: + try: + self.snapshot = await get_multi_lora_controller().snapshot.remote() + self.last_refresh = now + except Exception: + pass + return self.snapshot + + async def get_all(self) -> dict[str, "AdapterRun"]: + snapshot = await self.get_snapshot() + return {**snapshot["active"], **snapshot["retiring"]} + + async def get(self, adapter_name: str) -> "AdapterRun | None": + return (await self.get_all()).get(adapter_name) + + +def _load_subclass(path: str | None, base_cls): + if not path: + return base_cls + cls = load_function(path) + assert issubclass(cls, base_cls), f"{path} must point to a {base_cls.__name__} subclass, got {cls}" + return cls + + +@ray.remote(num_cpus=0) +class MultiLoRAController: + def __init__(self, args, router_url: str, host: str = "0.0.0.0") -> None: + backend_cls = _load_subclass(getattr(args, "multi_lora_backend_path", None), MultiLoRABackend) + server_cls = _load_subclass(getattr(args, "multi_lora_http_server_path", None), MultiLoRAHTTPServer) + self.backend = backend_cls(args, router_url) + self.server = server_cls(self.backend, host, api_port=getattr(args, "multi_lora_api_port", 0)) + + async def start(self) -> int: + await self.backend.init() + await self.server.start() + return self.server.actual_api_port + + async def stop(self) -> None: + await self.server.stop() + await self.backend.close() + + async def register_adapter(self, name: str, config: Any) -> dict: + return await self.backend.register(name, config) + + async def deregister_adapter(self, name: str) -> None: + await self.backend.deregister(name) + + async def retire_adapters(self) -> list[str]: + return await self.backend.retire_adapters() + + async def free_slot(self, name: str) -> int: + return await self.backend.free_slot(name) + + def record_weight_update(self, names: list[str]) -> None: + self.backend.registry.record_weight_update(names) + + def record_batch_adapters(self, rollout_id: int, groups: dict[str, int], step_names: list[str]) -> None: + self.backend.registry.record_batch_adapters(rollout_id, groups, step_names) + + def mark_batch_trained(self, rollout_id: int) -> list[str]: + return self.backend.registry.mark_batch_trained(rollout_id) + + def resolve_num_step(self, name: str, dataset_rows: int) -> None: + self.backend.registry.resolve_num_step(name, dataset_rows) + + def set_adapter_step(self, name: str, step: int) -> None: + self.backend.registry.set_step(name, step) + + def adapter_step(self, name: str) -> int: + return self.backend.registry.step_count(name) + + def snapshot(self) -> dict: + return self.backend.registry.snapshot() + + def http_host(self) -> str: + return get_current_node_ip() + + def api_port(self) -> int: + return self.server.actual_api_port + + +def create_multilora_controller(args, router_url: str, host: str = "0.0.0.0"): + # Pinned to the head node so the API sits at a port-forwardable address. + return MultiLoRAController.options( + name=CONTROLLER_NAME, + namespace=CONTROLLER_NAMESPACE, + **compute_ray_pin_head_options(), + ).remote(args, router_url, host) diff --git a/miles/ray/multi_lora/http_server.py b/miles/ray/multi_lora/http_server.py new file mode 100644 index 0000000000..b209142e1e --- /dev/null +++ b/miles/ray/multi_lora/http_server.py @@ -0,0 +1,129 @@ +"""Multi-LoRA control-plane HTTP API over a MultiLoRABackend. + +Subclass via ``--multi-lora-http-server-path`` (override add_routes / +create_app).""" + +import asyncio +from dataclasses import asdict +from pathlib import Path + +import uvicorn +from fastapi import FastAPI, HTTPException, Query, Request +from fastapi.responses import JSONResponse +from pydantic import BaseModel + +from miles.ray.multi_lora.registry import AdapterState +from miles.utils.adapter_config import AdapterRunConfig, parse_adapter_run_yaml + + +class RegisterAdapterRequest(BaseModel): + """Exactly one of ``config`` (inline) or ``yaml_path`` must be set.""" + + name: str + config: AdapterRunConfig | None = None + yaml_path: str | None = None + + +_NAMES_QUERY = Query(default_factory=list) + + +class MultiLoRAHTTPServer: + """Control-plane API over a MultiLoRABackend. Subclass via + --multi-lora-http-server-path (add_routes / create_app).""" + + def __init__(self, backend, host="127.0.0.1", api_port=0): + self.backend = backend + self.host = host + self.api_port = api_port + self.api_server: uvicorn.Server | None = None + self.api_task: asyncio.Task | None = None + + @property + def actual_api_port(self) -> int: + if self.api_server is not None and self.api_server.started: + return self.api_server.servers[0].sockets[0].getsockname()[1] + return self.api_port + + def create_app(self) -> FastAPI: + app = FastAPI(title="Miles Multi-LoRA Controller") + + @app.exception_handler(ValueError) + async def value_error_handler(request: Request, exc: ValueError): + return JSONResponse({"detail": str(exc)}, status_code=400) + + @app.exception_handler(RuntimeError) + async def runtime_error_handler(request: Request, exc: RuntimeError): + status = 409 if "No free adapter slots" in str(exc) else 500 + return JSONResponse({"detail": str(exc)}, status_code=status) + + return app + + def add_routes(self, app: FastAPI) -> None: + app.get("/health")(self.health) + app.get("/adapter_runs")(self.list_adapters) + app.get("/adapter_runs/state")(self.adapter_states) # before /adapter_runs/{name} + app.get("/adapter_runs/{name}")(self.get_adapter) + app.post("/adapter_runs")(self.register_adapter) + app.delete("/adapter_runs/{name}")(self.deregister_adapter) + + async def start(self) -> None: + app = self.create_app() + self.add_routes(app) + config = uvicorn.Config(app, host=self.host, port=self.api_port, log_level="warning", access_log=False) + self.api_server = uvicorn.Server(config) + self.api_task = asyncio.create_task(self.api_server.serve()) + while not self.api_server.started: + if self.api_task.done(): + self.api_task.result() + raise RuntimeError("uvicorn exited before startup completed") + await asyncio.sleep(0.01) + + async def stop(self) -> None: + if self.api_server is not None: + self.api_server.should_exit = True + await self.api_task + self.api_server = self.api_task = None + + async def health(self) -> dict: + return {"status": "healthy"} + + def adapter_statuses(self) -> list[dict]: + registry = self.backend.registry + statuses = [] + for record in registry.records.values(): + flat = asdict(registry.view(record)) + flat |= flat.pop("config") + flat["save"] = str(flat["save"]) + flat["state"] = record.state + if record.state is AdapterState.COMPLETED: + flat["version"] = None + statuses.append(flat) + return statuses + + async def list_adapters(self) -> dict: + return {"adapters": self.adapter_statuses()} + + async def adapter_states(self, names: list[str] = _NAMES_QUERY) -> dict: + return {"states": {name: self.backend.registry.adapter_state(name) for name in names}} + + async def get_adapter(self, name: str) -> dict: + for status in self.adapter_statuses(): + if status["name"] == name: + return status + raise HTTPException(status_code=404, detail=f"Adapter '{name}' not registered") + + async def register_adapter(self, request: RegisterAdapterRequest) -> dict: + if (request.config is None) == (request.yaml_path is None): + raise HTTPException(status_code=400, detail="Exactly one of 'config' or 'yaml_path' must be set") + if request.yaml_path is not None: + config = parse_adapter_run_yaml(Path(request.yaml_path)) + else: + config = request.config + return await self.backend.register(request.name, config) + + async def deregister_adapter(self, name: str) -> dict: + state = self.backend.registry.adapter_state(name) + if state is None: + raise HTTPException(status_code=404, detail=f"Adapter '{name}' not registered") + await self.backend.deregister(name) + return {"status": "ok", "name": name} diff --git a/miles/ray/multi_lora/registry.py b/miles/ray/multi_lora/registry.py new file mode 100644 index 0000000000..4c8723c29d --- /dev/null +++ b/miles/ray/multi_lora/registry.py @@ -0,0 +1,252 @@ +"""Multi-LoRA adapter registry: the controller-owned lifecycle state machine. + +One record per adapter name, walking PENDING -> ACTIVE -> RETIRING -> CLEANUP +-> COMPLETED. Slots are reused across registrations but ``slot_versions`` +never reset, so a (slot, version) pair never recurs. +""" + +import logging +import re +import uuid +from dataclasses import dataclass, field, replace +from enum import Enum +from pathlib import Path +from typing import Any + +from miles.utils.adapter_config import AdapterRun, AdapterRunConfig + +logger = logging.getLogger(__name__) + +VALID_ADAPTER_NAME = re.compile(r"^[A-Za-z0-9._-]+$") + + +class AdapterState(str, Enum): + PENDING = "PENDING" + ACTIVE = "ACTIVE" + RETIRING = "RETIRING" + CLEANUP = "CLEANUP" + COMPLETED = "COMPLETED" + + +# States that hold a slot. +LIVE_STATES = ( + AdapterState.PENDING, + AdapterState.ACTIVE, + AdapterState.RETIRING, + AdapterState.CLEANUP, +) + + +@dataclass +class AdapterRecord: + name: str + slot: int + config: Any + step: int = 0 + # Baseline step for relative num_step stopping (supports checkpoint resume). + start_step: int = 0 + # Committed prompt groups accumulated toward the current optimizer step. + # Only advanced by mark_batch_trained (after a successful train call). + accumulated_groups: int = 0 + state: AdapterState = AdapterState.PENDING + # Unique per registration: a re-registered name is a new tenant, and + # rollout-side state stamped by the previous tenant must not carry over. + registration_id: str = field(default_factory=lambda: uuid.uuid4().hex) + + +MAX_BATCH_RECORDS = 16 +MAX_COMPLETED_RECORDS = 1024 + + +class AdapterRegistry: + """One record per name; ``slot_versions`` never reset, so (slot, version) + never recurs across slot reuse.""" + + def __init__(self, max_adapters: int) -> None: + self.max_adapters = max_adapters + self.free_slots: set[int] = set(range(max_adapters)) + self.slot_versions: list[int] = [0] * max_adapters + self.records: dict[str, AdapterRecord] = {} + self.batch_records: dict[int, dict] = {} + + def in_state(self, *states: AdapterState) -> dict[str, AdapterRecord]: + return {name: r for name, r in self.records.items() if r.state in states} + + def find(self, name: str) -> AdapterRecord | None: + record = self.records.get(name) + return record if record is not None and record.state in LIVE_STATES else None + + def is_active(self, name: str) -> bool: + record = self.records.get(name) + return record is not None and record.state in (AdapterState.ACTIVE, AdapterState.RETIRING) + + def register(self, name: str, config: Any) -> dict: + if not VALID_ADAPTER_NAME.match(name) or name in (".", ".."): + raise ValueError(f"Adapter name '{name}' is invalid: use only letters, digits, '.', '_' and '-'") + if (existing := self.records.get(name)) is not None: + if existing.state in (AdapterState.PENDING, AdapterState.ACTIVE): + raise ValueError(f"Adapter '{name}' already registered") + if existing.state in (AdapterState.RETIRING, AdapterState.CLEANUP): + raise ValueError(f"Adapter '{name}' is still cleaning up; retry shortly") + if (save_dir := getattr(config, "save", None)) is not None: + for record in self.in_state(*LIVE_STATES).values(): + other_save = getattr(record.config, "save", None) + if other_save is not None and Path(other_save).resolve() == Path(save_dir).resolve(): + raise ValueError( + f"Adapter '{name}' save dir '{save_dir}' is already used by adapter '{record.name}'" + ) + if not self.free_slots: + raise RuntimeError(f"No free adapter slots (max {self.max_adapters})") + slot = min(self.free_slots) + self.free_slots.remove(slot) + self.records.pop(name, None) + self.records[name] = AdapterRecord(name=name, slot=slot, config=config) + return {"name": name, "slot": slot} + + def deregister(self, name: str) -> None: + record = self.records.get(name) + if record is not None and record.state in (AdapterState.PENDING, AdapterState.ACTIVE): + record.state = AdapterState.RETIRING + + def retire_adapters(self) -> list[str]: + retired = sorted(self.in_state(AdapterState.RETIRING)) + for name in retired: + self.records[name].state = AdapterState.CLEANUP + return retired + + def free_slot(self, name: str) -> int: + record = self.records.get(name) + if record is None or record.state is not AdapterState.CLEANUP: + return -1 + self.free_slots.add(record.slot) + record.state = AdapterState.COMPLETED + self.records[name] = self.records.pop(name) + completed = self.in_state(AdapterState.COMPLETED) + for oldest in list(completed)[: len(completed) - MAX_COMPLETED_RECORDS]: + self.records.pop(oldest) + return record.slot + + def adapter_state(self, name: str) -> AdapterState | None: + record = self.records.get(name) + if record is None: + return None + if record.state is AdapterState.COMPLETED: + self.records[name] = self.records.pop(name) + return record.state + + def record_weight_update(self, names: list[str]) -> None: + """A weight push landed: bump slot versions, promote PENDING to ACTIVE.""" + for name in names: + record = self.find(name) + if record is None: + continue + self.slot_versions[record.slot] += 1 + if record.state is AdapterState.PENDING: + record.state = AdapterState.ACTIVE + + def record_batch_adapters(self, rollout_id: int, groups: dict[str, int], step_names: list[str]) -> None: + """Register what a train batch contains before it trains. + + ``groups`` maps adapter name -> prompt groups riding in this batch; + ``step_names`` lists adapters whose adapter batch completes with + this batch (decided by the collection loop, which caps per-adapter + contributions at the adapter's remaining groups). + """ + unknown = set(step_names) - set(groups) + assert not unknown, f"step adapters {sorted(unknown)} not present in batch groups" + self.batch_records[rollout_id] = {"groups": dict(groups), "step_names": list(step_names)} + while len(self.batch_records) > MAX_BATCH_RECORDS: + self.batch_records.pop(next(iter(self.batch_records))) + + def mark_batch_trained(self, rollout_id: int) -> list[str]: + """Bank the batch's trained groups and fire steps; returns adapters that stepped. Only place + accumulation/step state advances, so a failed/retried train call leaves the registry untouched.""" + record_entry = self.batch_records.pop(rollout_id, None) + if record_entry is None: + return [] + stepped = [] + reached_num_step = [] + for name, n_groups in record_entry["groups"].items(): + record = self.records.get(name) + if record is None or record.state not in ( + AdapterState.ACTIVE, + AdapterState.RETIRING, + AdapterState.CLEANUP, + ): + continue + record.accumulated_groups += n_groups + if name in record_entry["step_names"]: + target = record.config.rollout_batch_size + if record.accumulated_groups != target: + logger.warning( + f"Adapter '{name}' stepped with accumulated_groups={record.accumulated_groups} " + f"!= rollout_batch_size={target}; adapter batch accounting drifted" + ) + record.step += 1 + record.accumulated_groups = 0 + stepped.append(name) + if ( + getattr(record.config, "num_step", None) is not None + and record.state is AdapterState.ACTIVE + and (record.step - record.start_step) >= record.config.num_step + ): + reached_num_step.append(name) + for name in reached_num_step: + logger.info( + f"Adapter '{name}' reached num_step={self.records[name].config.num_step} " + f"(start_step={self.records[name].start_step}, step={self.records[name].step}), deregistering" + ) + self.deregister(name) + return stepped + + def resolve_num_step(self, name: str, dataset_rows: int) -> None: + """Derive num_step from num_epoch once the data source knows the + post-filter dataset length. No-op when num_step was set explicitly.""" + record = self.find(name) + if record is None or not isinstance(record.config, AdapterRunConfig): + return + if record.config.num_step is not None: + return + num_epoch = record.config.num_epoch or 1 + num_step = max(1, num_epoch * dataset_rows // record.config.rollout_batch_size) + record.config = replace(record.config, num_step=num_step) + logger.info(f"Adapter '{name}': num_epoch={num_epoch} x {dataset_rows} rows -> num_step={num_step}") + + def set_step(self, name: str, step: int) -> None: + if (record := self.find(name)) is not None: + record.step = step + record.start_step = step + + def step_count(self, name: str) -> int: + record = self.find(name) + return record.step if record is not None else 0 + + def view(self, record: AdapterRecord) -> AdapterRun: + return AdapterRun( + name=record.name, + config=record.config, + slot=record.slot, + version=self.slot_versions[record.slot], + step=record.step, + accumulated_groups=record.accumulated_groups, + registration_id=record.registration_id, + ) + + def active_adapters(self) -> dict[str, AdapterRun]: + """Sampleable view: RETIRING keeps serving until retired.""" + return { + name: self.view(record) + for name, record in self.in_state(AdapterState.ACTIVE, AdapterState.RETIRING).items() + } + + def snapshot(self) -> dict: + def views(state: AdapterState) -> dict[str, AdapterRun]: + return {name: self.view(record) for name, record in self.in_state(state).items()} + + return { + "pending": views(AdapterState.PENDING), + "active": views(AdapterState.ACTIVE), + "retiring": views(AdapterState.RETIRING), + "cleanup": list(self.in_state(AdapterState.CLEANUP)), + "completed": list(self.in_state(AdapterState.COMPLETED)), + } diff --git a/tests/fast/ray/multi_lora/__init__.py b/tests/fast/ray/multi_lora/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/tests/fast/ray/multi_lora/test_controller_backend.py b/tests/fast/ray/multi_lora/test_controller_backend.py new file mode 100644 index 0000000000..fe75f985ef --- /dev/null +++ b/tests/fast/ray/multi_lora/test_controller_backend.py @@ -0,0 +1,392 @@ +"""Fast tests for AdapterRegistry + MultiLoRABackend validation +(no Ray, no HTTP I/O, no SGLang, no torch).""" + +from types import SimpleNamespace + +from tests.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=60, suite="stage-a-cpu") + +import pytest + +from miles.ray.multi_lora.backend import MultiLoRABackend +from miles.ray.multi_lora.registry import AdapterRegistry, AdapterState +from miles.utils.adapter_config import AdapterRunConfig +from miles.utils.multi_lora import make_rid, min_groups_per_dp_split, parse_adapter + + +# Registration validates that the data path exists; the test file itself is a +# convenient always-present stand-in. +DATA_FILE = __file__ + + +def make_args(max_adapters: int = 4, save: str | None = None, dp_size: int = 2) -> SimpleNamespace: + return SimpleNamespace( + multi_lora_n_adapters=max_adapters, + save=save, + lora_rank=32, + lora_alpha=32, + rollout_batch_size=16, + n_samples_per_prompt=4, + multi_lora_dp_size=dp_size, + multi_lora_max_adapter_global_batch_size=256, + ) + + +def make_backend(max_adapters: int = 4, save: str | None = None, dp_size: int = 2) -> MultiLoRABackend: + return MultiLoRABackend(make_args(max_adapters, save, dp_size), "http://unused") + + +def make_config(save: str | None = None, **overrides) -> AdapterRunConfig: + kwargs = dict( + rank=8, + alpha=16, + data=DATA_FILE, + rollout_batch_size=4, + n_samples_per_prompt=4, + save=save, + input_key="text", + label_key="label", + rm_type="math", + ) + kwargs.update(overrides) + return AdapterRunConfig(**kwargs) + + +def register_and_promote(registry: AdapterRegistry, name: str, config=None) -> None: + registry.register(name, config) + registry.record_weight_update([name]) + + +def test_rid_roundtrip_preserves_names_with_underscores(): + for name in ["a", "adapter_a", "weird__name", "x_y_z"]: + assert parse_adapter(make_rid(name)) == name + + +def test_register_starts_pending_and_push_promotes(): + registry = AdapterRegistry(max_adapters=4) + result = registry.register("A", config={"rm_type": "x"}) + assert result == {"name": "A", "slot": 0} + assert registry.active_adapters() == {} # pending: not sampleable + + registry.record_weight_update(["A"]) + assert registry.active_adapters()["A"].slot == 0 + view = registry.active_adapters()["A"] + assert view.slot == 0 + assert view.config == {"rm_type": "x"} + assert view.version == 1 + + +def test_snapshot_reports_sets_in_registry_vocabulary(): + registry = AdapterRegistry(max_adapters=4) + register_and_promote(registry, "A") + registry.register("B", None) + snapshot = registry.snapshot() + assert set(snapshot["active"]) == {"A"} + assert set(snapshot["pending"]) == {"B"} + assert snapshot["retiring"] == {} + assert snapshot["cleanup"] == [] + assert set(registry.active_adapters()) == {"A"} # only active adapters are sampleable + + +def test_slot_version_is_monotonic_across_slot_reuse(): + registry = AdapterRegistry(max_adapters=2) + register_and_promote(registry, "A") # slot 0, version 1 + registry.record_weight_update(["A"]) # version 2 + registry.deregister("A") + registry.retire_adapters() + registry.free_slot("A") + + registry.register("A2", None) # reuses slot 0 + assert registry.snapshot()["pending"]["A2"].version == 2 # inherits, not reset + registry.record_weight_update(["A2"]) + assert registry.active_adapters()["A2"].version == 3 + + +def test_record_weight_update_only_touches_reported_names(): + registry = AdapterRegistry(max_adapters=4) + register_and_promote(registry, "A") + register_and_promote(registry, "B") + registry.record_weight_update(["A"]) + assert registry.active_adapters()["A"].version == 2 + assert registry.active_adapters()["B"].version == 1 + + +def test_register_name_rejected_until_cleanup_done(): + registry = AdapterRegistry(max_adapters=4) + register_and_promote(registry, "A") + registry.deregister("A") + with pytest.raises(ValueError, match="cleaning up"): + registry.register("A", None) # retiring + registry.retire_adapters() + with pytest.raises(ValueError, match="cleaning up"): + registry.register("A", None) # cleanup + registry.free_slot("A") + assert registry.register("A", None) == {"name": "A", "slot": 0} + + +def test_deregister_retires_but_keeps_serving_until_demoted(): + registry = AdapterRegistry(max_adapters=4) + register_and_promote(registry, "A") + registry.deregister("A") + assert registry.adapter_state("A") == AdapterState.RETIRING + assert "A" in registry.active_adapters() # still sampleable this iteration + assert "A" in registry.snapshot()["retiring"] + assert registry.retire_adapters() == ["A"] + assert registry.active_adapters() == {} + assert registry.adapter_state("A") == AdapterState.CLEANUP + assert registry.retire_adapters() == [] # idempotent + + +# make_config(): rollout_batch_size=4 groups/step, n_samples_per_prompt=4. + + +def test_mark_batch_trained_accumulates_and_steps_on_completion(): + registry = AdapterRegistry(max_adapters=4) + register_and_promote(registry, "A", make_config()) + register_and_promote(registry, "B", make_config()) + + # Two partial batches accumulate; the third completes the adapter batch. + registry.record_batch_adapters(1, {"A": 1, "B": 2}, step_names=[]) + assert registry.mark_batch_trained(1) == [] + assert registry.records["A"].accumulated_groups == 1 + assert registry.records["B"].accumulated_groups == 2 + + registry.record_batch_adapters(2, {"A": 1}, step_names=[]) + assert registry.mark_batch_trained(2) == [] + assert registry.records["A"].accumulated_groups == 2 + + registry.record_batch_adapters(3, {"A": 2, "B": 2}, step_names=["A", "B"]) + assert registry.mark_batch_trained(3) == ["A", "B"] + assert registry.step_count("A") == 1 + assert registry.step_count("B") == 1 + assert registry.records["A"].accumulated_groups == 0 + assert registry.records["B"].accumulated_groups == 0 + + assert registry.mark_batch_trained(3) == [] # record consumed + + +def test_batch_trained_counts_deregistered_adapter_until_freed(): + registry = AdapterRegistry(max_adapters=4) + register_and_promote(registry, "A", make_config()) + registry.record_batch_adapters(3, {"A": 4}, step_names=["A"]) + registry.deregister("A") # deregistered while its batch is training + assert registry.mark_batch_trained(3) == ["A"] + assert registry.step_count("A") == 1 # final ckpt reads this + registry.retire_adapters() + assert registry.step_count("A") == 1 # cleanup record still holds it + registry.free_slot("A") + assert registry.step_count("A") == 0 + + +def test_set_step_on_resume(): + registry = AdapterRegistry(max_adapters=2) + registry.register("A", make_config()) + registry.set_step("A", 40) + registry.record_batch_adapters(1, {"A": 4}, step_names=["A"]) + registry.record_weight_update(["A"]) + registry.mark_batch_trained(1) + assert registry.step_count("A") == 41 + + +def test_num_step_deregisters_on_committed_steps(): + registry = AdapterRegistry(max_adapters=2) + register_and_promote(registry, "A", make_config(num_step=2)) + registry.record_batch_adapters(1, {"A": 4}, step_names=["A"]) + assert registry.mark_batch_trained(1) == ["A"] + assert registry.adapter_state("A") == AdapterState.ACTIVE + + registry.record_batch_adapters(2, {"A": 4}, step_names=["A"]) + assert registry.mark_batch_trained(2) == ["A"] + assert registry.step_count("A") == 2 + assert registry.adapter_state("A") == AdapterState.RETIRING + + +def test_num_step_is_relative_to_resume_step(): + registry = AdapterRegistry(max_adapters=2) + register_and_promote(registry, "A", make_config(num_step=2)) + registry.set_step("A", 40) + + registry.record_batch_adapters(1, {"A": 4}, step_names=["A"]) + registry.mark_batch_trained(1) + assert registry.step_count("A") == 41 + assert registry.adapter_state("A") == AdapterState.ACTIVE + + registry.record_batch_adapters(2, {"A": 4}, step_names=["A"]) + registry.mark_batch_trained(2) + assert registry.step_count("A") == 42 + assert registry.adapter_state("A") == AdapterState.RETIRING + + +def test_min_groups_per_dp_split(): + assert min_groups_per_dp_split(n_samples_per_prompt=4, dp_size=8) == 2 # divisor + assert min_groups_per_dp_split(n_samples_per_prompt=8, dp_size=8) == 1 # equal + assert min_groups_per_dp_split(n_samples_per_prompt=16, dp_size=8) == 1 # multiple + with pytest.raises(ValueError, match="divisor or a multiple"): + min_groups_per_dp_split(n_samples_per_prompt=6, dp_size=8) + + +@pytest.mark.asyncio +async def test_register_resolves_batch_shape_defaults(tmp_path): + backend = make_backend(save=str(tmp_path)) + await backend.register("A", AdapterRunConfig(data=DATA_FILE, rm_type="math")) + config = backend.registry.records["A"].config + assert config.rollout_batch_size == 16 # <- args.rollout_batch_size + assert config.n_samples_per_prompt == 4 # <- args.n_samples_per_prompt + assert config.rank == 32 and config.alpha == 32 + assert config.adapter_global_batch_size == 64 + + +@pytest.mark.asyncio +async def test_register_rejects_bad_batch_shapes(tmp_path): + backend = make_backend(save=str(tmp_path), dp_size=8) + with pytest.raises(ValueError, match="divisor or a multiple"): + await backend.register("B", make_config(n_samples_per_prompt=6, rollout_batch_size=4)) + with pytest.raises(ValueError, match="min_groups_per_dp_split"): + # dp=8, n_samples=4 -> multiple of 2 groups; 3 groups is not + await backend.register("C", make_config(rollout_batch_size=3)) + with pytest.raises(ValueError, match="exceeding"): + await backend.register("D", make_config(rollout_batch_size=128)) # 512 samples > cap 256 + with pytest.raises(ValueError, match="exceeds the allocated maximum rank"): + await backend.register("E", make_config(rank=64)) + with pytest.raises(ValueError, match="positive integer"): + await backend.register("F", make_config(rollout_batch_size=0)) + with pytest.raises(ValueError, match="num_step must be a positive integer"): + await backend.register("G", make_config(num_step=0)) + with pytest.raises(ValueError, match="num_epoch must be a positive integer"): + await backend.register("H", make_config(num_epoch=0)) + # A valid shape registers fine. + await backend.register("OK", make_config(rollout_batch_size=8)) + + +def test_deregister_holds_slot_until_free_slot(): + registry = AdapterRegistry(max_adapters=2) + register_and_promote(registry, "A") # slot 0 + register_and_promote(registry, "B") # slot 1 + registry.deregister("A") + registry.retire_adapters() + assert not registry.free_slots # slot 0 held until cleanup + with pytest.raises(RuntimeError, match="No free adapter slots"): + registry.register("C", None) + registry.free_slot("A") + assert registry.register("C", None) == {"name": "C", "slot": 0} + + +@pytest.mark.asyncio +async def test_free_slot_reaborts_before_releasing_slot(): + """Requests can survive the single retire-time abort (multi-turn groups + submitting between turns, engine tokenizer-adapter batch misses); free_slot must + fire one more abort round before the slot becomes reusable.""" + backend = make_backend() + aborted: list[str] = [] + + async def record_abort(name: str) -> None: + aborted.append(name) + + backend.abort_adapter_requests = record_abort + + register_and_promote(backend.registry, "A") + await backend.deregister("A") + await backend.retire_adapters() + assert aborted == ["A"] + + assert await backend.free_slot("A") == 0 + assert aborted == ["A", "A"] + assert backend.registry.free_slots == {0, 1, 2, 3} + + +@pytest.mark.asyncio +async def test_free_slot_skips_abort_when_not_in_cleanup(): + backend = make_backend() + aborted: list[str] = [] + + async def record_abort(name: str) -> None: + aborted.append(name) + + backend.abort_adapter_requests = record_abort + + register_and_promote(backend.registry, "A") # ACTIVE, not CLEANUP + assert await backend.free_slot("A") == -1 + assert await backend.free_slot("never-registered") == -1 + assert aborted == [] + + +@pytest.mark.asyncio +async def test_custom_backend_validation_rejects(): + class StrictBackend(MultiLoRABackend): + async def validate_adapter(self, name, config): + if not config: + raise ValueError("adapter config is required") + + backend = StrictBackend(make_args(), "http://unused") + with pytest.raises(ValueError, match="config is required"): + await backend.register("A", None) + assert backend.registry.active_adapters() == {} + + result = await backend.register("A", {"rm_type": "x"}) + assert result == {"name": "A", "slot": 0} + + +def test_register_rejects_unsafe_names(): + registry = AdapterRegistry(max_adapters=4) + for bad in ["a/b", "..", "a::b", "a b", ""]: + with pytest.raises(ValueError, match="invalid"): + registry.register(bad, None) + registry.register("ok-name_1.2", None) + + +def test_register_rejects_duplicate_save_dir(tmp_path): + registry = AdapterRegistry(max_adapters=4) + registry.register("A", make_config(save=tmp_path / "x")) + with pytest.raises(ValueError, match="already used by adapter 'A'"): + registry.register("B", make_config(save=tmp_path / "x")) + registry.register("C", make_config(save=tmp_path / "y")) + + +@pytest.mark.asyncio +async def test_save_dir_defaults_under_save_root(tmp_path): + backend = make_backend(save=str(tmp_path)) + await backend.register("A", make_config()) + saved = backend.registry.records["A"].config.save + assert saved == tmp_path / "adapters" / "A" + + +@pytest.mark.asyncio +async def test_explicit_save_dir_wins_over_root(tmp_path): + backend = make_backend(save=str(tmp_path)) + await backend.register("A", make_config(save=tmp_path / "custom")) + assert backend.registry.records["A"].config.save == tmp_path / "custom" + + +@pytest.mark.asyncio +async def test_register_fails_without_any_save_dir(): + backend = make_backend(save=None) + with pytest.raises(ValueError, match="no save dir"): + await backend.register("A", make_config()) + + +@pytest.mark.asyncio +async def test_register_rejects_missing_data_path(tmp_path): + # A nonexistent data path would otherwise kill the shared rollout producer + # thread at the first get_samples, stalling every adapter. + backend = make_backend(save=str(tmp_path)) + with pytest.raises(ValueError, match="data path"): + await backend.register("A", make_config(data=str(tmp_path / "missing.jsonl"))) + + +@pytest.mark.asyncio +async def test_register_rejects_unresolvable_reward_config(tmp_path): + # No adapter rm_type/custom_rm_path and no process-wide --rm-type: every + # sample would fail reward computation and be dropped. + backend = make_backend(save=str(tmp_path)) + with pytest.raises(ValueError, match="reward config"): + await backend.register("A", make_config(rm_type=None)) + + +@pytest.mark.asyncio +async def test_register_accepts_reward_config_from_global_args(tmp_path): + args = make_args(save=str(tmp_path)) + args.rm_type = "math" + backend = MultiLoRABackend(args, "http://unused") + await backend.register("A", make_config(rm_type=None)) + assert backend.registry.records["A"].config.rm_type is None # resolved at reward time via args diff --git a/tests/fast/ray/multi_lora/test_controller_http.py b/tests/fast/ray/multi_lora/test_controller_http.py new file mode 100644 index 0000000000..b270017597 --- /dev/null +++ b/tests/fast/ray/multi_lora/test_controller_http.py @@ -0,0 +1,239 @@ +"""HTTP tests for the MultiLoRAHTTPServer control plane with a mock router +(no Ray, no SGLang).""" + +import json +from contextlib import asynccontextmanager +from pathlib import Path +from types import SimpleNamespace + +import aiohttp +import pytest +from aiohttp import web + +from tests.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=60, suite="stage-a-cpu") + +from miles.ray.multi_lora.backend import MultiLoRABackend +from miles.ray.multi_lora.http_server import MultiLoRAHTTPServer +from miles.utils.adapter_config import AdapterRunConfig +from miles.utils.multi_lora import RID_SEPARATOR + + +# Registration validates that the data path exists; the test file itself is a +# convenient always-present stand-in. +DATA_FILE = __file__ + + +def minimal_config(name: str) -> dict: + return {"data": DATA_FILE, "rm_type": "math", "save": f"/tmp/adapters/{name}"} + + +class ControllerHarness: + """Running control plane (backend + API listener) against a mock router + that serves /list_workers and records /abort_request posts.""" + + def __init__(self, session: aiohttp.ClientSession, backend: MultiLoRABackend, srv: MultiLoRAHTTPServer): + self.session = session + self.backend = backend + self.srv = srv + self.aborts: list[dict] = [] + + @property + def api_base(self) -> str: + return f"http://127.0.0.1:{self.srv.actual_api_port}" + + async def api_post(self, path: str, payload: dict) -> tuple[int, dict]: + async with self.session.post(f"{self.api_base}{path}", json=payload) as resp: + return resp.status, await resp.json() + + async def api_get(self, path: str) -> tuple[int, dict, dict]: + async with self.session.get(f"{self.api_base}{path}") as resp: + headers = {k.lower(): v for k, v in resp.headers.items()} + return resp.status, await resp.json(), headers + + async def api_delete(self, path: str) -> tuple[int, dict]: + async with self.session.delete(f"{self.api_base}{path}") as resp: + return resp.status, await resp.json() + + async def register(self, name: str) -> tuple[int, dict]: + status, body = await self.api_post("/adapter_runs", {"name": name, "config": minimal_config(name)}) + # Registered adapters start pending; a weight push promotes them. + self.backend.registry.record_weight_update([name]) + return status, body + + async def deregister(self, name: str) -> tuple[int, dict]: + return await self.api_delete(f"/adapter_runs/{name}") + + async def active(self) -> dict: + _, body, _ = await self.api_get("/adapter_runs") + return { + s["name"]: {"slot": s["slot"], "version": s["version"], "step": s["step"]} + for s in body["adapters"] + if s["state"] == "ACTIVE" + } + + +@asynccontextmanager +async def running_controller(server_cls=MultiLoRAHTTPServer): + router_url = "" + harness: ControllerHarness | None = None + + async def router_handler(request): + if request.path == "/list_workers": + return web.json_response({"urls": [router_url]}) + if request.path == "/abort_request": + harness.aborts.append(json.loads(await request.read())) + return web.json_response({}) + return web.json_response({}, status=404) + + app = web.Application() + app.router.add_resource("/{tail:.*}").add_route("*", router_handler) + runner = web.AppRunner(app) + await runner.setup() + site = web.TCPSite(runner, "127.0.0.1", 0) + await site.start() + router_url = f"http://127.0.0.1:{site._server.sockets[0].getsockname()[1]}" + + backend = MultiLoRABackend( + SimpleNamespace( + multi_lora_n_adapters=4, + save=None, + lora_rank=32, + lora_alpha=32, + rollout_batch_size=16, + n_samples_per_prompt=4, + multi_lora_dp_size=2, + multi_lora_max_adapter_global_batch_size=256, + ), + router_url, + ) + srv = server_cls(backend) + await backend.init() + await srv.start() + try: + async with aiohttp.ClientSession() as session: + harness = ControllerHarness(session, backend, srv) + yield harness + finally: + await srv.stop() + await backend.close() + await runner.cleanup() + + +@pytest.mark.asyncio +async def test_register_and_active_view(): + async with running_controller() as ctl: + status, body = await ctl.register("A") + assert status == 200 + assert body["slot"] == 0 + assert await ctl.active() == {"A": {"slot": 0, "version": 1, "step": 0}} + + +@pytest.mark.asyncio +async def test_deregister_marks_and_retire_adapters_aborts(): + """Deregistration only marks; the driver-synced apply performs the + demotion and fans out one prefix abort per worker.""" + async with running_controller() as ctl: + await ctl.register("A") + status, _ = await ctl.deregister("A") + assert status == 200 + assert ctl.aborts == [] # still serving until the sync point + assert "A" in ctl.backend.registry.active_adapters() + + applied = await ctl.backend.retire_adapters() + assert applied == ["A"] + assert ctl.aborts == [{"rid": f"A{RID_SEPARATOR}", "prefix": True}] + assert ctl.backend.registry.active_adapters() == {} + + +@pytest.mark.asyncio +async def test_register_json_config_validates_to_adapter_config(): + """FastAPI validates the JSON body straight into AdapterRunConfig (422 on bad + payloads).""" + async with running_controller() as ctl: + config = { + "rank": 8, + "data": DATA_FILE, + "save": "/tmp/adapters/A", + "rm_type": "math", + } + status, _ = await ctl.api_post("/adapter_runs", {"name": "A", "config": config}) + assert status == 200 + record = ctl.backend.registry.find("A") + assert isinstance(record.config, AdapterRunConfig) + assert record.config.data == DATA_FILE + assert Path(record.config.save) == Path("/tmp/adapters/A") + assert record.config.input_key == "text" # dataclass default + + status, _ = await ctl.api_post("/adapter_runs", {"name": "B", "config": {"rank": 8}}) + assert status == 422 # data is required + + status, _ = await ctl.api_post("/adapter_runs", {"name": "C"}) + assert status == 400 # exactly one of config/yaml_path + + +@pytest.mark.asyncio +async def test_state_endpoint_reports_lifecycle_and_completed(): + """States walk PENDING -> ACTIVE -> RETIRING -> CLEANUP -> COMPLETED; + unknown names report null; COMPLETED is retained after free_slot.""" + async with running_controller() as ctl: + await ctl.api_post("/adapter_runs", {"name": "A", "config": minimal_config("A")}) + + async def state_of(name): + _, body, _ = await ctl.api_get(f"/adapter_runs/state?names={name}") + return body["states"][name] + + assert await state_of("A") == "PENDING" + ctl.backend.registry.record_weight_update(["A"]) + assert await state_of("A") == "ACTIVE" + + await ctl.deregister("A") + assert await state_of("A") == "RETIRING" + await ctl.backend.retire_adapters() + assert await state_of("A") == "CLEANUP" + + ctl.backend.registry.free_slot("A") + assert await state_of("A") == "COMPLETED" + assert await state_of("nope") is None + + # GET by name serves the completed record; DELETE of unknown 404s. + status, body, _ = await ctl.api_get("/adapter_runs/A") + assert status == 200 and body["state"] == "COMPLETED" + status, _ = await ctl.api_delete("/adapter_runs/nope") + assert status == 404 + + # Re-registration reclaims the name; the completed record is dropped. + status, _ = await ctl.api_post( + "/adapter_runs", + {"name": "A", "config": {"data": DATA_FILE, "rm_type": "math", "save": "/tmp/adapters/A2"}}, + ) + assert status == 200 + assert await state_of("A") == "PENDING" + + +@pytest.mark.asyncio +async def test_custom_server_subclass_adds_routes(): + class CustomServer(MultiLoRAHTTPServer): + def create_app(self): + app = super().create_app() + + @app.middleware("http") + async def tag_response(request, call_next): + response = await call_next(request) + response.headers["X-Custom-Server"] = "1" + return response + + return app + + def add_routes(self, app): + super().add_routes(app) + app.get("/custom_status")(self.custom_status) + + async def custom_status(self): + return {"custom": True, "active": sorted(self.backend.registry.active_adapters())} + + async with running_controller(server_cls=CustomServer) as ctl: + _, body, headers = await ctl.api_get("/custom_status") + assert headers.get("x-custom-server") == "1" + assert body == {"custom": True, "active": []}