diff --git a/miles/backends/sglang_utils/router_worker_client.py b/miles/backends/sglang_utils/router_worker_client.py new file mode 100644 index 0000000000..9a379af80a --- /dev/null +++ b/miles/backends/sglang_utils/router_worker_client.py @@ -0,0 +1,618 @@ +import logging +import time +from enum import StrEnum +from typing import NamedTuple, TypeVar +from urllib.parse import quote +from uuid import UUID + +import requests +from pydantic import BaseModel, ConfigDict, ValidationError + +logger = logging.getLogger(__name__) + +ROUTER_REQUEST_TIMEOUT_SECONDS = 10.0 +ROUTER_OPERATION_TIMEOUT_SECONDS = 240.0 +ROUTER_POLL_INTERVAL_SECONDS = 1.0 + + +class _AcceptedStatus(StrEnum): + ACCEPTED = "accepted" + + +class _JobState(StrEnum): + PENDING = "pending" + PROCESSING = "processing" + FAILED = "failed" + + +class RouterWorkerType(StrEnum): + """Worker roles accepted by the SGLang router worker API.""" + + REGULAR = "regular" + PREFILL = "prefill" + DECODE = "decode" + + +class RouterWorkerApi(StrEnum): + """Worker identity formats used by supported router APIs.""" + + URL = "url" + UUID = "uuid" + + +class RouterWorkerJobType(StrEnum): + """Worker lifecycle jobs reported by the router.""" + + ADD = "AddWorker" + REMOVE = "RemoveWorker" + + +class RouterWorkerSubmissionRejected(RuntimeError): + """The router definitively rejected a worker registration request.""" + + +class RouterWorkerSubmissionUnknown(RuntimeError): + """The router may have accepted a registration whose response was lost.""" + + +class RouterWorkerRegistrationFailed(RuntimeError): + """The router's accepted worker registration job failed.""" + + +class _RouterModel(BaseModel): + model_config = ConfigDict(extra="ignore") + + +_RouterModelT = TypeVar("_RouterModelT", bound=_RouterModel) + + +class _CreateWorkerRequest(_RouterModel): + url: str + worker_type: RouterWorkerType + bootstrap_port: int | None + + +class _CreateWorkerResponseV02(_RouterModel): + status: _AcceptedStatus + worker_id: str + + +class _CreateWorkerResponseV03(_RouterModel): + status: _AcceptedStatus + worker_id: UUID + url: str + location: str + + +class _DeleteWorkerResponse(_RouterModel): + status: _AcceptedStatus + worker_id: str + + +class _JobStatus(_RouterModel): + job_type: RouterWorkerJobType + worker_url: str + status: _JobState + message: str | None + + +class _WorkerInfo(_RouterModel): + id: str + url: str + is_healthy: bool + job_status: _JobStatus | None = None + + +class _ListedWorkerV03(_RouterModel): + id: UUID + url: str + + +class _ListWorkersResponseV03(_RouterModel): + workers: list[_ListedWorkerV03] + + +class RouterWorkerRegistration(NamedTuple): + """Identity of one worker accepted by the router. + + Args: + worker_id: Worker URL for Router 0.2.x or router-assigned UUID for + Router 0.3.x. + url: Canonical worker URL supplied during registration. + """ + + worker_id: str + url: str + + +class _RemovalReconciliation(NamedTuple): + is_absent: bool + registration_observed: bool + + +class RouterWorkerClient: + """Manage the asynchronous SGLang Router 0.2.2+ worker lifecycle. + + Args: + router_url: Base HTTP URL of the router. + worker_api: Worker identity format used by the router. + request_timeout_seconds: Timeout for each HTTP request. + operation_timeout_seconds: Overall deadline for one add or remove operation. + poll_interval_seconds: Delay between worker-state observations. + """ + + def __init__( + self, + *, + router_url: str, + worker_api: RouterWorkerApi, + request_timeout_seconds: float, + operation_timeout_seconds: float, + poll_interval_seconds: float, + ) -> None: + if request_timeout_seconds <= 0: + raise ValueError("request_timeout_seconds must be positive") + if operation_timeout_seconds <= 0: + raise ValueError("operation_timeout_seconds must be positive") + if poll_interval_seconds < 0: + raise ValueError("poll_interval_seconds must be nonnegative") + + self._router_url = router_url.rstrip("/") + self._worker_api = worker_api + self._request_timeout_seconds = request_timeout_seconds + self._operation_timeout_seconds = operation_timeout_seconds + self._poll_interval_seconds = poll_interval_seconds + + def submit_registration( + self, + *, + worker_url: str, + worker_type: RouterWorkerType, + bootstrap_port: int | None, + ) -> RouterWorkerRegistration: + """Submit one worker registration operation. + + Args: + worker_url: HTTP URL of the initialized worker. + worker_type: SGLang worker type. + bootstrap_port: Prefill bootstrap port, or None for other worker types. + + Returns: + The accepted router worker identity. + + Raises: + RuntimeError: The router rejects the operation or reports invalid state. + TimeoutError: Submission exceeds the operation deadline. + requests.RequestException: The submission request fails. + """ + request = _CreateWorkerRequest( + url=worker_url, + worker_type=worker_type, + bootstrap_port=bootstrap_port, + ) + deadline = time.monotonic() + self._operation_timeout_seconds + + with requests.Session() as session: + response = session.post( + f"{self._router_url}/workers", + json=request.model_dump(mode="json", exclude_none=True), + timeout=self._request_timeout(deadline=deadline, operation="register worker"), + ) + if 400 <= response.status_code < 500 and response.status_code not in {408, 429}: + try: + self._require_status(response, expected_status=202, operation="register worker") + except (requests.HTTPError, RuntimeError) as error: + raise RouterWorkerSubmissionRejected("Router rejected worker registration") from error + self._require_status(response, expected_status=202, operation="register worker") + + if self._worker_api is RouterWorkerApi.URL: + accepted_v02 = self._parse_model( + model_type=_CreateWorkerResponseV02, + response=response, + operation="register worker", + ) + if accepted_v02.worker_id != worker_url: + raise RuntimeError( + f"Router accepted worker identity {accepted_v02.worker_id!r}, expected {worker_url!r}" + ) + return RouterWorkerRegistration(worker_id=accepted_v02.worker_id, url=worker_url) + + accepted_v03 = self._parse_model( + model_type=_CreateWorkerResponseV03, + response=response, + operation="register worker", + ) + if accepted_v03.url != worker_url: + raise RuntimeError(f"Router accepted worker URL {accepted_v03.url!r}, expected {worker_url!r}") + + worker_id = str(accepted_v03.worker_id) + expected_location = f"/workers/{worker_id}" + if accepted_v03.location != expected_location: + raise RuntimeError( + f"Router returned worker location {accepted_v03.location!r}, expected {expected_location!r}" + ) + + return RouterWorkerRegistration(worker_id=worker_id, url=worker_url) + + def wait_until_active(self, *, registration: RouterWorkerRegistration) -> None: + """Wait until an accepted worker becomes healthy and routable. + + Args: + registration: Accepted worker identity to observe. + + Raises: + RuntimeError: The router reports failure or a different worker. + TimeoutError: Publication is not confirmed before the deadline. + """ + deadline = time.monotonic() + self._operation_timeout_seconds + with requests.Session() as session: + self._wait_until_active( + session=session, + location=self._worker_location(registration=registration), + registration=registration, + deadline=deadline, + ) + + def find_registration_by_url(self, *, worker_url: str) -> RouterWorkerRegistration | None: + """Wait for a worker URL to become observable in the router registry. + + This reconciles an accepted Router 0.3 registration when the POST + response, including its UUID, was lost. + + Args: + worker_url: Canonical worker URL supplied during registration. + + Returns: + The observed worker identity, or None if it remains absent through + the operation deadline. + """ + deadline = time.monotonic() + self._operation_timeout_seconds + with requests.Session() as session: + while True: + response = self._get_router_response( + session=session, + location="/workers", + deadline=deadline, + operation="list workers", + ) + self._require_status(response, expected_status=200, operation="list workers") + workers = self._parse_model( + model_type=_ListWorkersResponseV03, + response=response, + operation="list workers", + ) + matches = [worker for worker in workers.workers if worker.url == worker_url] + if len(matches) > 1: + raise RuntimeError(f"Router returned multiple workers for URL {worker_url!r}") + if matches: + return RouterWorkerRegistration(worker_id=str(matches[0].id), url=worker_url) + + remaining = deadline - time.monotonic() + if remaining <= 0: + logger.warning("Router worker URL %s did not become observable during cleanup", worker_url) + return None + time.sleep(min(self._poll_interval_seconds, remaining)) + + def remove( + self, + *, + registration: RouterWorkerRegistration, + registration_may_be_pending: bool, + ) -> None: + """Remove a worker and wait until the router no longer exposes it. + + Args: + registration: Accepted worker identity to remove. + registration_may_be_pending: Whether registration was accepted but + has not yet been observed as active. + + Raises: + RuntimeError: The router rejects the operation or reports invalid state. + TimeoutError: The router does not confirm removal before the deadline. + requests.RequestException: The removal request fails. + """ + location = self._worker_location(registration=registration) + deadline = time.monotonic() + self._operation_timeout_seconds + with requests.Session() as session: + while True: + try: + response = session.delete( + f"{self._router_url}{location}", + timeout=self._request_timeout(deadline=deadline, operation="remove worker"), + ) + except requests.RequestException as error: + logger.warning("Ambiguous router worker removal submission: %s", error) + else: + if response.status_code == 404: + if not registration_may_be_pending: + return + logger.warning( + "Router worker %s is not observable while registration may still be pending", + registration.worker_id, + ) + elif response.status_code in {408, 429} or response.status_code >= 500: + logger.warning( + "Ambiguous router worker removal response: status=%s", + response.status_code, + ) + else: + self._require_status(response, expected_status=202, operation="remove worker") + accepted = self._parse_model( + model_type=_DeleteWorkerResponse, + response=response, + operation="remove worker", + ) + if accepted.worker_id != registration.worker_id: + raise RuntimeError( + f"Router accepted removal for worker {accepted.worker_id}, " + f"expected {registration.worker_id}" + ) + self._wait_until_absent( + session=session, + location=location, + registration=registration, + deadline=deadline, + ) + return + + reconciliation = self._reconcile_ambiguous_removal( + session=session, + location=location, + registration=registration, + registration_may_be_pending=registration_may_be_pending, + deadline=deadline, + ) + if reconciliation.is_absent: + return + if reconciliation.registration_observed: + registration_may_be_pending = False + self._sleep_or_timeout(deadline=deadline, operation="remove worker") + + def _wait_until_active( + self, + *, + session: requests.Session, + location: str, + registration: RouterWorkerRegistration, + deadline: float, + ) -> None: + while True: + response = self._get_router_response( + session=session, + location=location, + deadline=deadline, + operation="observe worker", + ) + if response.status_code == 404: + self._sleep_or_timeout(deadline=deadline, operation="register worker") + continue + + self._require_status(response, expected_status=200, operation="observe registered worker") + worker = self._parse_model( + model_type=_WorkerInfo, + response=response, + operation="observe registered worker", + ) + self._validate_worker(worker=worker, registration=registration) + + if worker.job_status is not None: + self._validate_job( + job_status=worker.job_status, + expected_job_type=RouterWorkerJobType.ADD, + registration=registration, + ) + if worker.job_status.status is _JobState.FAILED: + detail = worker.job_status.message or "no failure message" + raise RouterWorkerRegistrationFailed( + f"Router failed to register worker {registration.worker_id}: {detail}" + ) + elif worker.is_healthy: + return + + self._sleep_or_timeout(deadline=deadline, operation="register worker") + + def _wait_until_absent( + self, + *, + session: requests.Session, + location: str, + registration: RouterWorkerRegistration, + deadline: float, + ) -> None: + while True: + response = self._get_router_response( + session=session, + location=location, + deadline=deadline, + operation="observe worker", + ) + if response.status_code == 404: + return + + self._require_status(response, expected_status=200, operation="observe removed worker") + worker = self._parse_model( + model_type=_WorkerInfo, + response=response, + operation="observe removed worker", + ) + self._validate_worker(worker=worker, registration=registration) + if worker.job_status is not None: + if worker.job_status.job_type is RouterWorkerJobType.ADD: + self._validate_job( + job_status=worker.job_status, + expected_job_type=RouterWorkerJobType.ADD, + registration=registration, + ) + self._sleep_or_timeout(deadline=deadline, operation="remove worker") + continue + self._validate_job( + job_status=worker.job_status, + expected_job_type=RouterWorkerJobType.REMOVE, + registration=registration, + ) + if worker.job_status.status is _JobState.FAILED: + detail = worker.job_status.message or "no failure message" + raise RuntimeError(f"Router failed to remove worker {registration.worker_id}: {detail}") + + self._sleep_or_timeout(deadline=deadline, operation="remove worker") + + def _reconcile_ambiguous_removal( + self, + *, + session: requests.Session, + location: str, + registration: RouterWorkerRegistration, + registration_may_be_pending: bool, + deadline: float, + ) -> _RemovalReconciliation: + registration_observed = False + while True: + response = self._get_router_response( + session=session, + location=location, + deadline=deadline, + operation="observe worker", + ) + if response.status_code == 404: + return _RemovalReconciliation( + is_absent=registration_observed or not registration_may_be_pending, + registration_observed=registration_observed, + ) + + self._require_status(response, expected_status=200, operation="reconcile removed worker") + worker = self._parse_model( + model_type=_WorkerInfo, + response=response, + operation="reconcile removed worker", + ) + self._validate_worker(worker=worker, registration=registration) + registration_observed = True + if worker.job_status is None: + return _RemovalReconciliation(is_absent=False, registration_observed=True) + + if worker.job_status.job_type is RouterWorkerJobType.ADD: + if worker.job_status.worker_url != registration.url: + raise RuntimeError( + f"Router reported a job for worker URL " + f"{worker.job_status.worker_url!r}, expected " + f"{registration.url!r}" + ) + if worker.job_status.status is _JobState.FAILED: + return _RemovalReconciliation(is_absent=False, registration_observed=True) + self._sleep_or_timeout(deadline=deadline, operation="remove worker") + continue + + self._validate_job( + job_status=worker.job_status, + expected_job_type=RouterWorkerJobType.REMOVE, + registration=registration, + ) + if worker.job_status.status is _JobState.FAILED: + detail = worker.job_status.message or "no failure message" + raise RuntimeError(f"Router failed to remove worker {registration.worker_id}: {detail}") + self._sleep_or_timeout(deadline=deadline, operation="remove worker") + + def _get_router_response( + self, + *, + session: requests.Session, + location: str, + deadline: float, + operation: str, + ) -> requests.Response: + last_error: requests.RequestException | None = None + while time.monotonic() < deadline: + try: + response = session.get( + f"{self._router_url}{location}", + timeout=self._request_timeout(deadline=deadline, operation=operation), + ) + if response.status_code in {408, 429} or response.status_code >= 500: + self._raise_for_status(response=response, operation=operation) + return response + except requests.RequestException as error: + last_error = error + logger.warning("Transient router %s failure: %s", operation, error) + self._sleep_or_timeout(deadline=deadline, operation=operation) + + raise TimeoutError(f"Timed out while attempting to {operation}") from last_error + + def _sleep_or_timeout(self, *, deadline: float, operation: str) -> None: + remaining = deadline - time.monotonic() + if remaining <= 0: + raise TimeoutError(f"Timed out after {self._operation_timeout_seconds}s while attempting to {operation}") + time.sleep(min(self._poll_interval_seconds, remaining)) + + def _request_timeout(self, *, deadline: float, operation: str) -> float: + remaining = deadline - time.monotonic() + if remaining <= 0: + raise TimeoutError(f"Timed out after {self._operation_timeout_seconds}s while attempting to {operation}") + return min(self._request_timeout_seconds, remaining) + + def _worker_location(self, *, registration: RouterWorkerRegistration) -> str: + if self._worker_api is RouterWorkerApi.URL: + return f"/workers/{quote(registration.worker_id, safe='')}" + return f"/workers/{registration.worker_id}" + + @staticmethod + def _validate_worker( + *, + worker: _WorkerInfo, + registration: RouterWorkerRegistration, + ) -> None: + if worker.id != registration.worker_id: + raise RuntimeError(f"Router returned worker {worker.id}, expected {registration.worker_id}") + if worker.url != registration.url: + raise RuntimeError(f"Router returned worker URL {worker.url!r}, expected {registration.url!r}") + + @staticmethod + def _validate_job( + *, + job_status: _JobStatus, + expected_job_type: RouterWorkerJobType, + registration: RouterWorkerRegistration, + ) -> None: + if job_status.job_type is not expected_job_type: + raise RuntimeError(f"Router reported {job_status.job_type} while waiting for {expected_job_type}") + if job_status.worker_url != registration.url: + raise RuntimeError( + f"Router reported a job for worker URL {job_status.worker_url!r}, " f"expected {registration.url!r}" + ) + + @staticmethod + def _parse_model( + *, + model_type: type[_RouterModelT], + response: requests.Response, + operation: str, + ) -> _RouterModelT: + try: + return model_type.model_validate_json(response.text) + except ValidationError as error: + raise RuntimeError( + f"Router returned an invalid response while attempting to {operation}: {response.text}" + ) from error + + @classmethod + def _require_status( + cls, + response: requests.Response, + *, + expected_status: int, + operation: str, + ) -> None: + if response.status_code == expected_status: + return + cls._raise_for_status(response=response, operation=operation) + raise RuntimeError( + f"Router returned HTTP {response.status_code} while attempting to {operation}; " + f"expected HTTP {expected_status}: {response.text}" + ) + + @staticmethod + def _raise_for_status(*, response: requests.Response, operation: str) -> None: + try: + response.raise_for_status() + except requests.HTTPError as error: + error.add_note(f"Router response while attempting to {operation}: {response.text}") + raise diff --git a/miles/backends/sglang_utils/sglang_engine.py b/miles/backends/sglang_utils/sglang_engine.py index 095ca2cd67..0ae2a74904 100644 --- a/miles/backends/sglang_utils/sglang_engine.py +++ b/miles/backends/sglang_utils/sglang_engine.py @@ -4,7 +4,6 @@ import multiprocessing import os import time -from urllib.parse import quote import requests import sglang_router @@ -14,6 +13,18 @@ from urllib3.exceptions import NewConnectionError from miles.backends.megatron_utils.lora_utils import convert_target_modules_to_hf, lora_base_cpu_backup_enabled +from miles.backends.sglang_utils.router_worker_client import ( + ROUTER_OPERATION_TIMEOUT_SECONDS, + ROUTER_POLL_INTERVAL_SECONDS, + ROUTER_REQUEST_TIMEOUT_SECONDS, + RouterWorkerApi, + RouterWorkerClient, + RouterWorkerRegistration, + RouterWorkerRegistrationFailed, + RouterWorkerSubmissionRejected, + RouterWorkerSubmissionUnknown, + RouterWorkerType, +) from miles.ray.ray_actor import RayActor from miles.utils.env_report import collect_and_print_node_env_report from miles.utils.http_utils import get_host_info @@ -138,6 +149,10 @@ def __init__( self.base_gpu_id = base_gpu_id self.sglang_overrides = sglang_overrides or {} self.num_gpus_per_engine = num_gpus_per_engine + self._is_registered_with_router = False + self._router_worker_id: str | None = None + self._router_registration_failed = False + self._router_registration_submission_unknown = False def get_topology_info(self) -> dict: """Placement facts for the dashboard timeline. ``base_gpu_id`` is @@ -168,7 +183,23 @@ def init( router_ip=None, router_port=None, engine_info_bootstrap_port=None, + *, + register_with_router: bool, ): + """Initialize an SGLang server actor. + + Args: + dist_init_addr: Distributed initialization address for the engine. + port: HTTP server port. + nccl_port: NCCL communication port. + host: HTTP server host, or None to use the current node address. + disaggregation_bootstrap_port: Bootstrap port for a prefill worker. + router_ip: Router address, or None to use the configured address. + router_port: Router port, or None to use the configured port. + engine_info_bootstrap_port: Port used to exchange engine metadata. + register_with_router: Whether the initialized engine can be routed + immediately. + """ if env_report := self.args.env_report: collect_and_print_node_env_report( role="rollout", @@ -213,11 +244,12 @@ def _format_v6_uri(addr): self.node_rank = server_args_dict["node_rank"] self.server_host = server_args_dict["host"] # with [] if ipv6 self.server_port = server_args_dict["port"] + self.disaggregation_bootstrap_port = server_args_dict.get("disaggregation_bootstrap_port") if self.args.rollout_external: self._init_external(server_args_dict, external_engine_need_check_fields=external_engine_need_check_fields) else: - self._init_normal(server_args_dict) + self._init_normal(server_args_dict, register_with_router=register_with_router) def _init_external(self, expect_server_args, external_engine_need_check_fields): logger.info(f"Use external SGLang engine (rank={self.rank}, expect_server_args={expect_server_args})") @@ -243,30 +275,104 @@ def _sanity_check_server_args(actual_server_args, expect_server_args): actual_server_args = _get_actual_server_args() _sanity_check_server_args(actual_server_args, expect_server_args) - def _init_normal(self, server_args_dict): + def _init_normal(self, server_args_dict, *, register_with_router: bool): logger.info(f"Launch HttpServerEngineAdapter at: {self.server_host}:{self.server_port}") self.process = launch_server_process(ServerArgs(**server_args_dict)) + if register_with_router: + self.register_with_router() - if self.node_rank == 0 and self.router_ip and self.router_port: - if parse(sglang_router.__version__) <= parse("0.2.1") or self.args.use_miles_router: - assert ( - self.worker_type == "regular" - ), "pd disaggregation is not supported in old router or miles router." + def register_with_router(self): + """Register this initialized engine as an available router worker. + + Recovery can defer this call until current actor weights have been + installed, preventing checkpoint-backed replacements from serving. + """ + if ( + self.args.rollout_external + or self.node_rank != 0 + or not self.router_ip + or not self.router_port + or self._is_registered_with_router + ): + return + + worker_url = f"http://{self.server_host}:{self.server_port}" + if parse(sglang_router.__version__) <= parse("0.2.1") or self.args.use_miles_router: + assert self.worker_type == "regular", "pd disaggregation is not supported in old router or miles router." + if self._router_registration_submission_unknown: + raise RouterWorkerSubmissionUnknown( + "The router registration submission result is unknown; " + "the engine generation must be replaced before retrying" + ) + self._router_worker_id = worker_url + self._router_registration_submission_unknown = True + try: response = requests.post( - f"http://{self.router_ip}:{self.router_port}/add_worker?url=http://{self.server_host}:{self.server_port}" + f"http://{self.router_ip}:{self.router_port}/add_worker?url={worker_url}", + timeout=ROUTER_REQUEST_TIMEOUT_SECONDS, + ) + response.raise_for_status() + except requests.HTTPError as error: + status_code = error.response.status_code if error.response is not None else None + if status_code is not None and 400 <= status_code < 500 and status_code not in {408, 429}: + self._router_worker_id = None + self._router_registration_submission_unknown = False + raise + raise RouterWorkerSubmissionUnknown( + "The router registration submission result is unknown; " + "the engine generation must be replaced before retrying" + ) from error + except requests.RequestException as error: + raise RouterWorkerSubmissionUnknown( + "The router registration submission result is unknown; " + "the engine generation must be replaced before retrying" + ) from error + self._router_registration_submission_unknown = False + else: + if self._router_registration_submission_unknown: + raise RouterWorkerSubmissionUnknown( + "The router registration submission result is unknown; " + "the engine generation must be replaced before retrying" ) + bootstrap_port = None + if self.worker_type == "prefill": + if self.disaggregation_bootstrap_port is None: + raise RuntimeError("Prefill worker requires disaggregation_bootstrap_port") + bootstrap_port = self.disaggregation_bootstrap_port + client = self._router_worker_client() + if self._router_worker_id is None: + if parse(sglang_router.__version__) < parse("0.3.0"): + self._router_worker_id = worker_url + self._router_registration_submission_unknown = True + try: + registration = client.submit_registration( + worker_url=worker_url, + worker_type=RouterWorkerType(self.worker_type), + bootstrap_port=bootstrap_port, + ) + except RouterWorkerSubmissionRejected: + self._router_worker_id = None + self._router_registration_submission_unknown = False + raise + except Exception as error: + raise RouterWorkerSubmissionUnknown( + "The router registration submission result is unknown; " + "the engine generation must be replaced before retrying" + ) from error + self._router_worker_id = registration.worker_id + self._router_registration_submission_unknown = False else: - payload = { - "url": f"http://{self.server_host}:{self.server_port}", - "worker_type": self.worker_type, - } - if self.worker_type == "prefill": - payload["bootstrap_port"] = server_args_dict["disaggregation_bootstrap_port"] - response = requests.post( - f"http://{self.router_ip}:{self.router_port}/workers", - json=payload, + registration = RouterWorkerRegistration( + worker_id=self._router_worker_id, + url=worker_url, ) - response.raise_for_status() + try: + client.wait_until_active(registration=registration) + except RouterWorkerRegistrationFailed: + self._router_registration_failed = True + raise + self._router_registration_failed = False + self._is_registered_with_router = True def _make_request(self, endpoint: str, payload: dict | None = None): """Make a POST request to the specified endpoint with the given payload. @@ -455,34 +561,59 @@ def shutdown(self): return logger.info(f"Shutdown engine {self.server_host}:{self.server_port}...") - if self.node_rank == 0: - worker_url = f"http://{self.server_host}:{self.server_port}" - response = None - if parse(sglang_router.__version__) <= parse("0.2.1") or self.args.use_miles_router: - response = requests.post( - f"http://{self.router_ip}:{self.router_port}/remove_worker?url=http://{self.server_host}:{self.server_port}" - ) - elif parse(sglang_router.__version__) < parse("0.3.0"): - worker_url = quote(worker_url, safe="") - response = requests.delete(f"http://{self.router_ip}:{self.router_port}/workers/{worker_url}") - else: - try: - all_workers = requests.get(f"http://{self.router_ip}:{self.router_port}/workers").json()["workers"] - for worker in all_workers: - if worker["url"] == worker_url: - worker_id = worker["id"] - response = requests.delete( - f"http://{self.router_ip}:{self.router_port}/workers/{worker_id}" - ) - break - else: - logger.warning(f"Worker {worker_url} not found in router during shutdown.") - except Exception as e: - logger.warning(f"Failed to fetch workers list or remove worker: {e}") - - if response is not None: - response.raise_for_status() - kill_process_tree(self.process.pid) + try: + if self.node_rank == 0 and ( + self._is_registered_with_router + or self._router_worker_id is not None + or self._router_registration_submission_unknown + ): + worker_url = f"http://{self.server_host}:{self.server_port}" + response = None + if parse(sglang_router.__version__) <= parse("0.2.1") or self.args.use_miles_router: + response = requests.post( + f"http://{self.router_ip}:{self.router_port}/remove_worker?url={worker_url}", + timeout=ROUTER_REQUEST_TIMEOUT_SECONDS, + ) + else: + client = self._router_worker_client() + registration_may_be_pending = ( + not self._is_registered_with_router + and not self._router_registration_failed + and (self._router_worker_id is not None or self._router_registration_submission_unknown) + ) + registration = ( + client.find_registration_by_url(worker_url=worker_url) + if self._router_worker_id is None + else RouterWorkerRegistration( + worker_id=self._router_worker_id, + url=worker_url, + ) + ) + if registration is not None: + client.remove( + registration=registration, + registration_may_be_pending=registration_may_be_pending, + ) + self._router_registration_submission_unknown = False + + if response is not None: + response.raise_for_status() + self._router_registration_submission_unknown = False + self._is_registered_with_router = False + self._router_worker_id = None + self._router_registration_failed = False + finally: + kill_process_tree(self.process.pid) + + def _router_worker_client(self) -> RouterWorkerClient: + worker_api = RouterWorkerApi.URL if parse(sglang_router.__version__) < parse("0.3.0") else RouterWorkerApi.UUID + return RouterWorkerClient( + router_url=f"http://{self.router_ip}:{self.router_port}", + worker_api=worker_api, + request_timeout_seconds=ROUTER_REQUEST_TIMEOUT_SECONDS, + operation_timeout_seconds=ROUTER_OPERATION_TIMEOUT_SECONDS, + poll_interval_seconds=ROUTER_POLL_INTERVAL_SECONDS, + ) def get_weight_version(self): if self.node_rank != 0: diff --git a/miles/ray/rollout/server_group.py b/miles/ray/rollout/server_group.py index 65b065963e..35e00a68c1 100644 --- a/miles/ray/rollout/server_group.py +++ b/miles/ray/rollout/server_group.py @@ -164,6 +164,7 @@ def start_engines( **addr_and_ports[index], router_ip=self.router_ip, router_port=self.router_port, + register_with_router=True, ) for index, engine in new_engines ] diff --git a/miles/utils/test_utils/mock_sglang_engine.py b/miles/utils/test_utils/mock_sglang_engine.py index 42bb858e7e..0621cf4d6b 100644 --- a/miles/utils/test_utils/mock_sglang_engine.py +++ b/miles/utils/test_utils/mock_sglang_engine.py @@ -39,6 +39,7 @@ "get_weight_version": "mock-v0", "get_parallelism_info": {"_mock": True}, "get_remote_instance_transfer_engine_info": {"_mock": True}, + "register_with_router": None, } @@ -97,6 +98,8 @@ def init(self, **kwargs): self._record("init", (), kwargs) self._maybe_fault("init") self.initialized = True + if kwargs["register_with_router"]: + self.register_with_router() return None def shutdown(self): diff --git a/tests/fast/backends/sglang_utils/test_router_worker_client.py b/tests/fast/backends/sglang_utils/test_router_worker_client.py new file mode 100644 index 0000000000..62f1dd93ae --- /dev/null +++ b/tests/fast/backends/sglang_utils/test_router_worker_client.py @@ -0,0 +1,751 @@ +import json +import time +from collections.abc import Iterator +from unittest.mock import MagicMock +from urllib.parse import quote + +import pytest +import requests +from tests.ci.ci_register import register_cpu_ci + +from miles.backends.sglang_utils.router_worker_client import ( + RouterWorkerApi, + RouterWorkerClient, + RouterWorkerJobType, + RouterWorkerRegistration, + RouterWorkerSubmissionRejected, + RouterWorkerType, +) + +register_cpu_ci(est_time=1, suite="stage-a-cpu", labels=[]) + +_WORKER_ID = "12345678-1234-5678-1234-567812345678" +_WORKER_URL = "http://127.0.0.1:30000" +_LOCATION = f"/workers/{_WORKER_ID}" + + +class _Response: + def __init__(self, status_code: int, payload: object) -> None: + self.status_code = status_code + self.text = json.dumps(payload) + + def raise_for_status(self) -> None: + if self.status_code >= 400: + raise requests.HTTPError( + f"HTTP {self.status_code}", + response=self, + ) + + +class _Session: + def __init__( + self, + *, + post_responses: list[_Response] | None = None, + get_responses: list[_Response] | None = None, + delete_responses: list[_Response | requests.RequestException] | None = None, + ) -> None: + self._post_responses: Iterator[_Response] = iter(post_responses or []) + self._get_responses: Iterator[_Response] = iter(get_responses or []) + self._delete_responses: Iterator[_Response | requests.RequestException] = iter(delete_responses or []) + self.calls: list[tuple[str, str, object | None, float]] = [] + + def __enter__(self) -> "_Session": + return self + + def __exit__(self, exc_type: object, exc_value: object, traceback: object) -> None: + return None + + def post(self, url: str, *, json: object, timeout: float) -> _Response: + self.calls.append(("POST", url, json, timeout)) + return next(self._post_responses) + + def get(self, url: str, *, timeout: float) -> _Response: + self.calls.append(("GET", url, None, timeout)) + return next(self._get_responses) + + def delete(self, url: str, *, timeout: float) -> _Response: + self.calls.append(("DELETE", url, None, timeout)) + response = next(self._delete_responses) + if isinstance(response, requests.RequestException): + raise response + return response + + +def _client( + *, + worker_api: RouterWorkerApi = RouterWorkerApi.UUID, + poll_interval_seconds: float = 0.0, +) -> RouterWorkerClient: + return RouterWorkerClient( + router_url="http://router:3000", + worker_api=worker_api, + request_timeout_seconds=2.0, + operation_timeout_seconds=10.0, + poll_interval_seconds=poll_interval_seconds, + ) + + +def _create_response() -> _Response: + return _Response( + 202, + { + "status": "accepted", + "worker_id": str(_WORKER_ID), + "url": _WORKER_URL, + "location": _LOCATION, + "message": "Worker addition queued for background processing", + }, + ) + + +def _worker_response( + *, + is_healthy: bool, + job_status: dict[str, object] | None, + job_type: RouterWorkerJobType = RouterWorkerJobType.ADD, +) -> _Response: + if job_status is not None: + job_status = { + "job_type": job_type, + "worker_url": _WORKER_URL, + **job_status, + } + return _Response( + 200, + { + "id": str(_WORKER_ID), + "url": _WORKER_URL, + "is_healthy": is_healthy, + "job_status": job_status, + }, + ) + + +def test_register_waits_for_confirmed_active_worker(monkeypatch: pytest.MonkeyPatch) -> None: + session = _Session( + post_responses=[_create_response()], + get_responses=[ + _worker_response( + is_healthy=False, + job_status={"status": "pending", "message": None}, + ), + _worker_response( + is_healthy=False, + job_status={"status": "processing", "message": None}, + ), + _worker_response(is_healthy=True, job_status=None), + ], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + client = _client() + result = client.submit_registration( + worker_url=_WORKER_URL, + worker_type=RouterWorkerType.REGULAR, + bootstrap_port=None, + ) + client.wait_until_active(registration=result) + + assert result == RouterWorkerRegistration(worker_id=_WORKER_ID, url=_WORKER_URL) + assert session.calls == [ + ( + "POST", + "http://router:3000/workers", + {"url": _WORKER_URL, "worker_type": "regular"}, + 2.0, + ), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ] + + +def test_register_reports_background_failure(monkeypatch: pytest.MonkeyPatch) -> None: + session = _Session( + post_responses=[_create_response()], + get_responses=[ + _worker_response( + is_healthy=False, + job_status={"status": "failed", "message": "Worker already exists"}, + ) + ], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + with pytest.raises(RuntimeError, match="Worker already exists"): + client = _client() + registration = client.submit_registration( + worker_url=_WORKER_URL, + worker_type=RouterWorkerType.REGULAR, + bootstrap_port=None, + ) + client.wait_until_active(registration=registration) + + +def test_register_rejects_colliding_removal_job(monkeypatch: pytest.MonkeyPatch) -> None: + session = _Session( + post_responses=[_create_response()], + get_responses=[ + _worker_response( + is_healthy=False, + job_status={"status": "processing", "message": None}, + job_type=RouterWorkerJobType.REMOVE, + ) + ], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + client = _client() + registration = client.submit_registration( + worker_url=_WORKER_URL, + worker_type=RouterWorkerType.REGULAR, + bootstrap_port=None, + ) + with pytest.raises(RuntimeError, match="RemoveWorker.*AddWorker"): + client.wait_until_active(registration=registration) + + +def test_register_rejects_malformed_acceptance(monkeypatch: pytest.MonkeyPatch) -> None: + session = _Session( + post_responses=[ + _Response( + 202, + { + "status": "accepted", + "worker_id": "not-a-uuid", + "url": _WORKER_URL, + "location": "/workers/not-a-uuid", + }, + ) + ] + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + with pytest.raises(RuntimeError, match="invalid response"): + _client().submit_registration( + worker_url=_WORKER_URL, + worker_type=RouterWorkerType.REGULAR, + bootstrap_port=None, + ) + + +def test_register_marks_client_rejection_as_retryable( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = _Session( + post_responses=[_Response(400, {"error": "invalid worker"})], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + with pytest.raises(RouterWorkerSubmissionRejected): + _client().submit_registration( + worker_url=_WORKER_URL, + worker_type=RouterWorkerType.REGULAR, + bootstrap_port=None, + ) + + +def test_register_preserves_ambiguous_server_failure( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = _Session( + post_responses=[_Response(503, {"error": "temporarily unavailable"})], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + with pytest.raises(requests.HTTPError): + _client().submit_registration( + worker_url=_WORKER_URL, + worker_type=RouterWorkerType.REGULAR, + bootstrap_port=None, + ) + + +def test_register_reconciles_transient_not_found_without_second_post( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = _Session( + post_responses=[_create_response()], + get_responses=[ + _Response(404, {"error": "not found"}), + _worker_response(is_healthy=True, job_status=None), + ], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + client = _client() + registration = client.submit_registration( + worker_url=_WORKER_URL, + worker_type=RouterWorkerType.REGULAR, + bootstrap_port=None, + ) + client.wait_until_active(registration=registration) + + assert [call[0] for call in session.calls] == ["POST", "GET", "GET"] + + +@pytest.mark.parametrize("status_code", [408, 429, 503]) +def test_register_retries_transient_observation_error( + monkeypatch: pytest.MonkeyPatch, + status_code: int, +) -> None: + session = _Session( + post_responses=[_create_response()], + get_responses=[ + _Response(status_code, {"error": "temporarily unavailable"}), + _worker_response(is_healthy=True, job_status=None), + ], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + client = _client() + registration = client.submit_registration( + worker_url=_WORKER_URL, + worker_type=RouterWorkerType.REGULAR, + bootstrap_port=None, + ) + client.wait_until_active(registration=registration) + + assert [call[0] for call in session.calls] == ["POST", "GET", "GET"] + + +def test_register_includes_prefill_bootstrap_port(monkeypatch: pytest.MonkeyPatch) -> None: + session = _Session( + post_responses=[_create_response()], + get_responses=[_worker_response(is_healthy=True, job_status=None)], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + client = _client() + registration = client.submit_registration( + worker_url=_WORKER_URL, + worker_type=RouterWorkerType.PREFILL, + bootstrap_port=40000, + ) + client.wait_until_active(registration=registration) + + assert session.calls[0] == ( + "POST", + "http://router:3000/workers", + { + "url": _WORKER_URL, + "worker_type": "prefill", + "bootstrap_port": 40000, + }, + 2.0, + ) + + +def test_find_registration_by_url_waits_until_worker_is_observable( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = _Session( + get_responses=[ + _Response(200, {"workers": [], "total": 0}), + _Response( + 200, + { + "workers": [ + { + "id": _WORKER_ID, + "url": _WORKER_URL, + "is_healthy": False, + "job_status": None, + } + ], + "total": 1, + }, + ), + ], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + registration = _client().find_registration_by_url(worker_url=_WORKER_URL) + + assert registration == RouterWorkerRegistration(worker_id=_WORKER_ID, url=_WORKER_URL) + assert session.calls == [ + ("GET", "http://router:3000/workers", None, 2.0), + ("GET", "http://router:3000/workers", None, 2.0), + ] + + +def test_remove_waits_until_worker_is_absent(monkeypatch: pytest.MonkeyPatch) -> None: + session = _Session( + delete_responses=[ + _Response( + 202, + { + "status": "accepted", + "worker_id": str(_WORKER_ID), + "message": "Worker removal queued for background processing", + }, + ) + ], + get_responses=[ + _worker_response( + is_healthy=False, + job_status={"status": "processing", "message": None}, + job_type=RouterWorkerJobType.REMOVE, + ), + _Response(404, {"error": "not found"}), + ], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + _client().remove( + registration=RouterWorkerRegistration(worker_id=_WORKER_ID, url=_WORKER_URL), + registration_may_be_pending=False, + ) + + assert session.calls == [ + ("DELETE", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ] + + +def test_remove_waits_for_registration_job_to_yield_to_removal( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = _Session( + delete_responses=[ + _Response( + 202, + { + "status": "accepted", + "worker_id": str(_WORKER_ID), + }, + ) + ], + get_responses=[ + _worker_response( + is_healthy=False, + job_status={"status": "processing", "message": None}, + job_type=RouterWorkerJobType.ADD, + ), + _worker_response( + is_healthy=False, + job_status={"status": "processing", "message": None}, + job_type=RouterWorkerJobType.REMOVE, + ), + _Response(404, {"error": "not found"}), + ], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + _client().remove( + registration=RouterWorkerRegistration(worker_id=_WORKER_ID, url=_WORKER_URL), + registration_may_be_pending=True, + ) + + assert session.calls == [ + ("DELETE", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ] + + +def test_remove_retries_not_found_while_registration_is_pending( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = _Session( + delete_responses=[ + _Response(404, {"error": "not found"}), + _Response( + 202, + { + "status": "accepted", + "worker_id": str(_WORKER_ID), + }, + ), + ], + get_responses=[ + _Response(404, {"error": "not found"}), + _worker_response( + is_healthy=False, + job_status={"status": "processing", "message": None}, + job_type=RouterWorkerJobType.ADD, + ), + _worker_response( + is_healthy=False, + job_status={"status": "processing", "message": None}, + job_type=RouterWorkerJobType.REMOVE, + ), + _Response(404, {"error": "not found"}), + ], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + _client().remove( + registration=RouterWorkerRegistration(worker_id=_WORKER_ID, url=_WORKER_URL), + registration_may_be_pending=True, + ) + + assert session.calls == [ + ("DELETE", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ("DELETE", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ] + + +def test_remove_accepts_absence_after_observing_pending_registration( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = _Session( + delete_responses=[_Response(404, {"error": "not found"})], + get_responses=[ + _worker_response( + is_healthy=False, + job_status={"status": "processing", "message": None}, + job_type=RouterWorkerJobType.ADD, + ), + _Response(404, {"error": "not found"}), + ], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + _client().remove( + registration=RouterWorkerRegistration(worker_id=_WORKER_ID, url=_WORKER_URL), + registration_may_be_pending=True, + ) + + assert session.calls == [ + ("DELETE", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ] + + +def test_remove_retries_delete_for_failed_registration_record( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = _Session( + delete_responses=[ + requests.Timeout("response lost"), + _Response( + 202, + { + "status": "accepted", + "worker_id": str(_WORKER_ID), + }, + ), + ], + get_responses=[ + _worker_response( + is_healthy=False, + job_status={"status": "failed", "message": "registration failed"}, + job_type=RouterWorkerJobType.ADD, + ), + _Response(404, {"error": "not found"}), + ], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + _client().remove( + registration=RouterWorkerRegistration(worker_id=_WORKER_ID, url=_WORKER_URL), + registration_may_be_pending=True, + ) + + assert session.calls == [ + ("DELETE", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ("DELETE", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ] + + +def test_remove_reconciles_ambiguous_submission( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = _Session( + delete_responses=[requests.Timeout("response lost")], + get_responses=[ + _worker_response( + is_healthy=False, + job_status={"status": "processing", "message": None}, + job_type=RouterWorkerJobType.REMOVE, + ), + _Response(404, {"error": "not found"}), + ], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + _client().remove( + registration=RouterWorkerRegistration( + worker_id=_WORKER_ID, + url=_WORKER_URL, + ), + registration_may_be_pending=False, + ) + + assert session.calls == [ + ("DELETE", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ] + + +def test_remove_throttles_retry_after_transient_response( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = _Session( + delete_responses=[ + _Response(503, {"error": "temporarily unavailable"}), + _Response( + 202, + { + "status": "accepted", + "worker_id": str(_WORKER_ID), + }, + ), + ], + get_responses=[ + _worker_response(is_healthy=True, job_status=None), + _Response(404, {"error": "not found"}), + ], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + sleep = MagicMock() + monkeypatch.setattr(time, "sleep", sleep) + + _client(poll_interval_seconds=1.0).remove( + registration=RouterWorkerRegistration(worker_id=_WORKER_ID, url=_WORKER_URL), + registration_may_be_pending=False, + ) + + assert session.calls == [ + ("DELETE", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ("DELETE", f"http://router:3000{_LOCATION}", None, 2.0), + ("GET", f"http://router:3000{_LOCATION}", None, 2.0), + ] + sleep.assert_called_once_with(1.0) + + +def test_remove_reports_background_failure(monkeypatch: pytest.MonkeyPatch) -> None: + session = _Session( + delete_responses=[ + _Response( + 202, + { + "status": "accepted", + "worker_id": str(_WORKER_ID), + }, + ) + ], + get_responses=[ + _worker_response( + is_healthy=False, + job_status={"status": "failed", "message": "remove failed"}, + job_type=RouterWorkerJobType.REMOVE, + ) + ], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + with pytest.raises(RuntimeError, match="remove failed"): + _client().remove( + registration=RouterWorkerRegistration(worker_id=_WORKER_ID, url=_WORKER_URL), + registration_may_be_pending=False, + ) + + +def test_remove_is_idempotent_when_worker_is_already_absent( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = _Session(delete_responses=[_Response(404, {"error": "not found"})]) + monkeypatch.setattr(requests, "Session", lambda: session) + + _client().remove( + registration=RouterWorkerRegistration(worker_id=_WORKER_ID, url=_WORKER_URL), + registration_may_be_pending=False, + ) + + assert session.calls == [ + ("DELETE", f"http://router:3000{_LOCATION}", None, 2.0), + ] + + +def test_remove_reports_definitive_client_rejection(monkeypatch: pytest.MonkeyPatch) -> None: + session = _Session(delete_responses=[_Response(400, {"error": "invalid worker"})]) + monkeypatch.setattr(requests, "Session", lambda: session) + + with pytest.raises(requests.HTTPError, match="400"): + _client().remove( + registration=RouterWorkerRegistration(worker_id=_WORKER_ID, url=_WORKER_URL), + registration_may_be_pending=False, + ) + + assert session.calls == [ + ("DELETE", f"http://router:3000{_LOCATION}", None, 2.0), + ] + + +def test_router_0_2_uses_encoded_worker_url_identity( + monkeypatch: pytest.MonkeyPatch, +) -> None: + encoded_worker_url = quote(_WORKER_URL, safe="") + location = f"/workers/{encoded_worker_url}" + session = _Session( + post_responses=[ + _Response( + 202, + { + "status": "accepted", + "worker_id": _WORKER_URL, + }, + ) + ], + get_responses=[ + _Response( + 200, + { + "id": _WORKER_URL, + "url": _WORKER_URL, + "is_healthy": True, + "job_status": None, + }, + ), + _Response(404, {"error": "not found"}), + ], + delete_responses=[ + _Response( + 202, + { + "status": "accepted", + "worker_id": _WORKER_URL, + }, + ) + ], + ) + monkeypatch.setattr(requests, "Session", lambda: session) + + client = _client(worker_api=RouterWorkerApi.URL) + registration = client.submit_registration( + worker_url=_WORKER_URL, + worker_type=RouterWorkerType.REGULAR, + bootstrap_port=None, + ) + client.wait_until_active(registration=registration) + client.remove(registration=registration, registration_may_be_pending=False) + + assert registration == RouterWorkerRegistration( + worker_id=_WORKER_URL, + url=_WORKER_URL, + ) + assert session.calls == [ + ( + "POST", + "http://router:3000/workers", + {"url": _WORKER_URL, "worker_type": "regular"}, + 2.0, + ), + ("GET", f"http://router:3000{location}", None, 2.0), + ("DELETE", f"http://router:3000{location}", None, 2.0), + ("GET", f"http://router:3000{location}", None, 2.0), + ] diff --git a/tests/fast/backends/sglang_utils/test_sglang_engine_router.py b/tests/fast/backends/sglang_utils/test_sglang_engine_router.py new file mode 100644 index 0000000000..1c3fc1b0f1 --- /dev/null +++ b/tests/fast/backends/sglang_utils/test_sglang_engine_router.py @@ -0,0 +1,437 @@ +from argparse import Namespace +from unittest.mock import MagicMock, call + +import pytest +import requests +import sglang_router +from tests.ci.ci_register import register_cpu_ci + +import miles.backends.sglang_utils.sglang_engine as sglang_engine_module +from miles.backends.sglang_utils.router_worker_client import ( + RouterWorkerRegistration, + RouterWorkerRegistrationFailed, + RouterWorkerSubmissionRejected, + RouterWorkerSubmissionUnknown, + RouterWorkerType, +) +from miles.backends.sglang_utils.sglang_engine import SGLangEngine + +register_cpu_ci(est_time=1, suite="stage-a-cpu", labels=[]) + +_WORKER_ID = "12345678-1234-5678-1234-567812345678" +_WORKER_URL = "http://127.0.0.1:30000" + + +def _engine() -> SGLangEngine: + engine = SGLangEngine.__new__(SGLangEngine) + engine.args = Namespace(rollout_external=False, use_miles_router=False) + engine.node_rank = 0 + engine.router_ip = "127.0.0.1" + engine.router_port = 3000 + engine.server_host = "127.0.0.1" + engine.server_port = 30000 + engine.worker_type = "regular" + engine.disaggregation_bootstrap_port = None + engine._is_registered_with_router = False + engine._router_worker_id = None + engine._router_registration_failed = False + engine._router_registration_submission_unknown = False + return engine + + +def test_registration_timeout_reconciles_same_worker_without_second_post( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(sglang_router, "__version__", "0.3.2") + client = MagicMock() + registration = RouterWorkerRegistration(worker_id=_WORKER_ID, url=_WORKER_URL) + client.submit_registration.return_value = registration + client.wait_until_active.side_effect = [TimeoutError("pending"), None] + engine = _engine() + monkeypatch.setattr(engine, "_router_worker_client", MagicMock(return_value=client)) + + with pytest.raises(TimeoutError, match="pending"): + engine.register_with_router() + assert engine._is_registered_with_router is False + assert engine._router_worker_id == _WORKER_ID + + engine.register_with_router() + + assert client.submit_registration.call_count == 1 + assert client.wait_until_active.call_count == 2 + assert engine._is_registered_with_router is True + + +def test_unknown_submission_result_blocks_duplicate_post( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(sglang_router, "__version__", "0.3.2") + client = MagicMock() + client.submit_registration.side_effect = TimeoutError("response lost") + engine = _engine() + monkeypatch.setattr(engine, "_router_worker_client", MagicMock(return_value=client)) + + with pytest.raises(RouterWorkerSubmissionUnknown, match="generation must be replaced"): + engine.register_with_router() + with pytest.raises(RouterWorkerSubmissionUnknown, match="generation must be replaced"): + engine.register_with_router() + + assert client.submit_registration.call_count == 1 + assert engine._is_registered_with_router is False + + registration = RouterWorkerRegistration(worker_id=_WORKER_ID, url=_WORKER_URL) + client.find_registration_by_url.return_value = registration + engine.process = Namespace(pid=123) + kill_process_tree = MagicMock() + monkeypatch.setattr(sglang_engine_module, "kill_process_tree", kill_process_tree) + + engine.shutdown() + + client.find_registration_by_url.assert_called_once_with(worker_url=_WORKER_URL) + client.remove.assert_called_once_with( + registration=registration, + registration_may_be_pending=True, + ) + assert engine._router_registration_submission_unknown is False + kill_process_tree.assert_called_once_with(123) + + +def test_router_0_2_unknown_submission_retains_url_identity_for_cleanup( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(sglang_router, "__version__", "0.2.2") + client = MagicMock() + client.submit_registration.side_effect = TimeoutError("response lost") + engine = _engine() + engine.process = Namespace(pid=123) + monkeypatch.setattr(engine, "_router_worker_client", MagicMock(return_value=client)) + kill_process_tree = MagicMock() + monkeypatch.setattr(sglang_engine_module, "kill_process_tree", kill_process_tree) + + with pytest.raises(RouterWorkerSubmissionUnknown, match="generation must be replaced"): + engine.register_with_router() + with pytest.raises(RouterWorkerSubmissionUnknown, match="generation must be replaced"): + engine.register_with_router() + + assert client.submit_registration.call_count == 1 + assert engine._router_worker_id == _WORKER_URL + + engine.shutdown() + + client.find_registration_by_url.assert_not_called() + client.remove.assert_called_once_with( + registration=RouterWorkerRegistration(worker_id=_WORKER_URL, url=_WORKER_URL), + registration_may_be_pending=True, + ) + kill_process_tree.assert_called_once_with(123) + + +def test_external_wrapper_recovery_does_not_require_router_publication( + monkeypatch: pytest.MonkeyPatch, +) -> None: + engine = _engine() + engine.args.rollout_external = True + engine.args.env_report = None + engine.args.sglang_router_ip = "127.0.0.1" + engine.args.sglang_router_port = 3000 + engine.rank = 0 + engine.base_gpu_id = 0 + engine.sglang_overrides = {} + engine.num_gpus_per_engine = 1 + server_args = { + "node_rank": 0, + "host": "127.0.0.1", + "port": 30000, + } + compute_server_args = MagicMock(return_value=(server_args, ["host", "port"])) + monkeypatch.setattr(sglang_engine_module, "_compute_server_args", compute_server_args) + init_external = MagicMock() + monkeypatch.setattr(engine, "_init_external", init_external) + + engine.init( + dist_init_addr="127.0.0.1:29500", + port=30000, + nccl_port=29501, + host="127.0.0.1", + register_with_router=False, + ) + + init_external.assert_called_once_with( + server_args, + external_engine_need_check_fields=["host", "port"], + ) + + +def test_normal_initialization_can_defer_router_registration( + monkeypatch: pytest.MonkeyPatch, +) -> None: + engine = _engine() + process = object() + server_args = object() + server_args_factory = MagicMock(return_value=server_args) + launch_server_process = MagicMock(return_value=process) + register_with_router = MagicMock() + monkeypatch.setattr(sglang_engine_module, "ServerArgs", server_args_factory) + monkeypatch.setattr(sglang_engine_module, "launch_server_process", launch_server_process) + monkeypatch.setattr(engine, "register_with_router", register_with_router) + + engine._init_normal({"model_path": "model"}, register_with_router=False) + + server_args_factory.assert_called_once_with(model_path="model") + launch_server_process.assert_called_once_with(server_args) + register_with_router.assert_not_called() + assert engine.process is process + + +def test_normal_initialization_registers_immediately( + monkeypatch: pytest.MonkeyPatch, +) -> None: + engine = _engine() + process = object() + server_args = object() + server_args_factory = MagicMock(return_value=server_args) + launch_server_process = MagicMock(return_value=process) + register_with_router = MagicMock() + monkeypatch.setattr(sglang_engine_module, "ServerArgs", server_args_factory) + monkeypatch.setattr(sglang_engine_module, "launch_server_process", launch_server_process) + monkeypatch.setattr(engine, "register_with_router", register_with_router) + + engine._init_normal({"model_path": "model"}, register_with_router=True) + + server_args_factory.assert_called_once_with(model_path="model") + launch_server_process.assert_called_once_with(server_args) + register_with_router.assert_called_once_with() + assert engine.process is process + + +def test_router_0_2_registration_uses_async_worker_client( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(sglang_router, "__version__", "0.2.2") + client = MagicMock() + registration = RouterWorkerRegistration(worker_id=_WORKER_URL, url=_WORKER_URL) + client.submit_registration.return_value = registration + engine = _engine() + monkeypatch.setattr(engine, "_router_worker_client", MagicMock(return_value=client)) + + engine.register_with_router() + + client.submit_registration.assert_called_once_with( + worker_url=_WORKER_URL, + worker_type=RouterWorkerType.REGULAR, + bootstrap_port=None, + ) + client.wait_until_active.assert_called_once_with(registration=registration) + assert engine._is_registered_with_router is True + + +def test_router_0_2_1_registration_keeps_synchronous_contract( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(sglang_router, "__version__", "0.2.1") + response = MagicMock() + post = MagicMock(return_value=response) + monkeypatch.setattr(requests, "post", post) + engine = _engine() + router_worker_client = MagicMock() + monkeypatch.setattr(engine, "_router_worker_client", router_worker_client) + + engine.register_with_router() + + post.assert_called_once_with( + "http://127.0.0.1:3000/add_worker?url=http://127.0.0.1:30000", + timeout=10.0, + ) + response.raise_for_status.assert_called_once_with() + router_worker_client.assert_not_called() + assert engine._is_registered_with_router is True + + +def test_router_0_2_1_unknown_submission_blocks_retry_and_allows_cleanup( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(sglang_router, "__version__", "0.2.1") + remove_response = MagicMock() + post = MagicMock(side_effect=[requests.Timeout("response lost"), remove_response]) + monkeypatch.setattr(requests, "post", post) + engine = _engine() + engine.process = Namespace(pid=123) + kill_process_tree = MagicMock() + monkeypatch.setattr(sglang_engine_module, "kill_process_tree", kill_process_tree) + + with pytest.raises(RouterWorkerSubmissionUnknown, match="generation must be replaced"): + engine.register_with_router() + with pytest.raises(RouterWorkerSubmissionUnknown, match="generation must be replaced"): + engine.register_with_router() + + assert engine._router_worker_id == _WORKER_URL + assert engine._is_registered_with_router is False + + engine.shutdown() + + assert post.call_args_list == [ + call( + "http://127.0.0.1:3000/add_worker?url=http://127.0.0.1:30000", + timeout=10.0, + ), + call( + "http://127.0.0.1:3000/remove_worker?url=http://127.0.0.1:30000", + timeout=10.0, + ), + ] + remove_response.raise_for_status.assert_called_once_with() + assert engine._router_worker_id is None + kill_process_tree.assert_called_once_with(123) + + +@pytest.mark.parametrize( + ("router_version", "worker_id"), + [ + ("0.2.2", _WORKER_URL), + ("0.3.2", _WORKER_ID), + ], +) +def test_definitive_submission_rejection_remains_retryable( + monkeypatch: pytest.MonkeyPatch, + router_version: str, + worker_id: str, +) -> None: + monkeypatch.setattr(sglang_router, "__version__", router_version) + client = MagicMock() + registration = RouterWorkerRegistration(worker_id=worker_id, url=_WORKER_URL) + client.submit_registration.side_effect = [ + RouterWorkerSubmissionRejected("registration rejected"), + registration, + ] + engine = _engine() + monkeypatch.setattr(engine, "_router_worker_client", MagicMock(return_value=client)) + + with pytest.raises(RouterWorkerSubmissionRejected, match="registration rejected"): + engine.register_with_router() + + assert engine._router_registration_submission_unknown is False + assert engine._router_worker_id is None + + engine.register_with_router() + + assert client.submit_registration.call_count == 2 + client.wait_until_active.assert_called_once_with(registration=registration) + assert engine._is_registered_with_router is True + + +def test_background_registration_failure_retains_identity_without_resubmission( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(sglang_router, "__version__", "0.3.2") + client = MagicMock() + registration = RouterWorkerRegistration(worker_id=_WORKER_ID, url=_WORKER_URL) + client.submit_registration.return_value = registration + client.wait_until_active.side_effect = [ + RouterWorkerRegistrationFailed("registration failed"), + RouterWorkerRegistrationFailed("registration failed"), + ] + engine = _engine() + monkeypatch.setattr(engine, "_router_worker_client", MagicMock(return_value=client)) + + with pytest.raises(RouterWorkerRegistrationFailed, match="registration failed"): + engine.register_with_router() + + assert engine._router_worker_id == _WORKER_ID + assert engine._is_registered_with_router is False + assert engine._router_registration_failed is True + + with pytest.raises(RouterWorkerRegistrationFailed, match="registration failed"): + engine.register_with_router() + + assert client.submit_registration.call_count == 1 + assert client.wait_until_active.call_count == 2 + assert engine._is_registered_with_router is False + assert engine._router_registration_failed is True + + engine.process = Namespace(pid=123) + kill_process_tree = MagicMock() + monkeypatch.setattr(sglang_engine_module, "kill_process_tree", kill_process_tree) + + engine.shutdown() + + client.remove.assert_called_once_with( + registration=registration, + registration_may_be_pending=False, + ) + kill_process_tree.assert_called_once_with(123) + + +def test_router_0_2_1_registration_failure_remains_retryable( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(sglang_router, "__version__", "0.2.1") + first_response = MagicMock() + first_response.status_code = 400 + first_response.raise_for_status.side_effect = requests.HTTPError( + "registration failed", + response=first_response, + ) + second_response = MagicMock() + post = MagicMock(side_effect=[first_response, second_response]) + monkeypatch.setattr(requests, "post", post) + engine = _engine() + + with pytest.raises(requests.HTTPError, match="registration failed"): + engine.register_with_router() + + assert engine._is_registered_with_router is False + + engine.register_with_router() + + assert post.call_count == 2 + assert second_response.raise_for_status.call_count == 1 + assert engine._is_registered_with_router is True + + +def test_shutdown_kills_process_when_router_removal_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(sglang_router, "__version__", "0.3.2") + client = MagicMock() + client.remove.side_effect = RuntimeError("remove failed") + engine = _engine() + engine._is_registered_with_router = True + engine._router_worker_id = _WORKER_ID + engine.process = Namespace(pid=123) + monkeypatch.setattr(engine, "_router_worker_client", MagicMock(return_value=client)) + kill_process_tree = MagicMock() + monkeypatch.setattr(sglang_engine_module, "kill_process_tree", kill_process_tree) + + with pytest.raises(RuntimeError, match="remove failed"): + engine.shutdown() + + client.remove.assert_called_once_with( + registration=RouterWorkerRegistration(worker_id=_WORKER_ID, url=_WORKER_URL), + registration_may_be_pending=False, + ) + kill_process_tree.assert_called_once_with(123) + + +def test_shutdown_removes_accepted_worker_before_activation( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(sglang_router, "__version__", "0.3.2") + client = MagicMock() + engine = _engine() + engine._router_worker_id = _WORKER_ID + engine.process = Namespace(pid=123) + monkeypatch.setattr(engine, "_router_worker_client", MagicMock(return_value=client)) + kill_process_tree = MagicMock() + monkeypatch.setattr(sglang_engine_module, "kill_process_tree", kill_process_tree) + + engine.shutdown() + + client.remove.assert_called_once_with( + registration=RouterWorkerRegistration( + worker_id=_WORKER_ID, + url=_WORKER_URL, + ), + registration_may_be_pending=True, + ) + assert engine._router_worker_id is None + kill_process_tree.assert_called_once_with(123) diff --git a/tests/fast/utils/test_utils/test_mock_sglang_engine.py b/tests/fast/utils/test_utils/test_mock_sglang_engine.py index a6d4475fb8..dc0b249b3c 100644 --- a/tests/fast/utils/test_utils/test_mock_sglang_engine.py +++ b/tests/fast/utils/test_utils/test_mock_sglang_engine.py @@ -107,7 +107,7 @@ def test_actor_construction_and_method_round_trip(self, ray_local_mode): num_gpus_per_engine=1, ) try: - ray.get(actor.init.remote(host="127.0.0.1", port=20000)) + ray.get(actor.init.remote(host="127.0.0.1", port=20000, register_with_router=True)) ray.get(actor.health_generate.remote(timeout=1.0)) ray.get(actor.release_memory_occupation.remote(tags=[GPU_MEMORY_TYPE_WEIGHTS])) ray.get(actor.resume_memory_occupation.remote(tags=[GPU_MEMORY_TYPE_WEIGHTS])) @@ -118,6 +118,7 @@ def test_actor_construction_and_method_round_trip(self, ray_local_mode): method_names = [name for name, _, _ in calls] assert method_names == [ "init", + "register_with_router", "health_generate", "release_memory_occupation", "resume_memory_occupation",