From 6deb6c303d9a8e2f04d38a5d7ea3dfe9e701f2aa Mon Sep 17 00:00:00 2001 From: Dmitry Meyer Date: Fri, 9 Oct 2026 10:10:23 +0000 Subject: [PATCH] Rework router worker sync around job ID labels The sync matched SMG router workers to jobs by URL and removed every worker that wasn't a ready dstack worker, so users couldn't register workers of their own. It also probed every worker on each sync and deregistered the ones that failed, so a flaky connection from the server made a worker flap even when the router, which reaches it within the cluster, had no problem with it. Probes ran one at a time, and only for the connection mode and runtime type of the workers already registered, which broke services that mix them or switch between them in a rolling deployment. Now each worker that dstack registers is labeled with its job ID (`dstack.ai/job-id`), and the sync works in terms of jobs: * A worker is probed only while some router lacks it. Once it's registered, its health is left to the router, which checks it anyway. A worker that is truly unreachable has its job terminated, which deregisters it. * A worker is removed once its job is no longer running. Workers that dstack can't map to a job are left alone, so users can register their own. * Workers registered before the label existed are mapped to jobs by address. The ones whose job has already stopped are no longer removed: the router marks them unhealthy, and a rolling deployment of the router drops them. * Unregistered workers are probed for HTTP, SGLang gRPC and vLLM gRPC at once, with no hints from the registered ones. * Routers and workers are processed in parallel, up to 8 replicas at a time. An unexpected error in one of them is logged and no longer aborts the sync for the others. Also: * `gather_async()` is added, and `gather_map_async()` is built on it. Both accept `max_concurrency`, and cancel the remaining awaitables when one fails unless `return_exceptions` is set. * Router responses are validated with pydantic. Co-authored-by: Claude Opus 5.5 (1M context) --- .../service_router_worker_sync.py | 8 +- .../services/jobs/job_replica_grpc_client.py | 2 +- .../services/runs/router_worker_sync.py | 1126 ++++++------ src/dstack/_internal/server/utils/common.py | 105 +- .../test_service_router_worker_sync.py | 44 +- .../services/runs/test_router_worker_sync.py | 1538 +++++++++++------ .../_internal/server/utils/test_common.py | 118 ++ 7 files changed, 1879 insertions(+), 1062 deletions(-) diff --git a/src/dstack/_internal/server/background/pipeline_tasks/service_router_worker_sync.py b/src/dstack/_internal/server/background/pipeline_tasks/service_router_worker_sync.py index f946688a16..3d66280bf9 100644 --- a/src/dstack/_internal/server/background/pipeline_tasks/service_router_worker_sync.py +++ b/src/dstack/_internal/server/background/pipeline_tasks/service_router_worker_sync.py @@ -31,6 +31,7 @@ ServiceRouterWorkerSyncModel, ) from dstack._internal.server.services.locking import get_locker +from dstack._internal.server.services.logging import fmt from dstack._internal.server.services.pipelines import PipelineHinterProtocol from dstack._internal.server.services.runs.router_worker_sync import ( run_model_has_sglang_router_replica_group, @@ -287,7 +288,12 @@ async def process(self, item: ServiceRouterWorkerSyncPipelineItem) -> None: await _update_sync_row_or_log_lock_token_changed(session, item, cleanup_update_map) return - await sync_router_workers_for_run_model(run_for_sync) + try: + await sync_router_workers_for_run_model(run_for_sync) + except Exception: + logger.exception( + "%s: unexpected error when syncing workers with router", fmt(run_for_sync) + ) update_map: _SyncRowUpdateMap = {} set_processed_update_map_fields(update_map) diff --git a/src/dstack/_internal/server/services/jobs/job_replica_grpc_client.py b/src/dstack/_internal/server/services/jobs/job_replica_grpc_client.py index 1827274a10..167bc8bf89 100644 --- a/src/dstack/_internal/server/services/jobs/job_replica_grpc_client.py +++ b/src/dstack/_internal/server/services/jobs/job_replica_grpc_client.py @@ -21,7 +21,7 @@ @asynccontextmanager async def get_service_replica_grpc_channel_over_uds( uds_path: Path, -) -> AsyncGenerator[Any, None]: +) -> AsyncGenerator[grpc.aio.Channel, None]: target = f"unix://{uds_path}" channel = grpc.aio.insecure_channel(target, options=_GRPC_CHANNEL_OPTIONS) try: diff --git a/src/dstack/_internal/server/services/runs/router_worker_sync.py b/src/dstack/_internal/server/services/runs/router_worker_sync.py index d1e8a45c19..83e78e582b 100644 --- a/src/dstack/_internal/server/services/runs/router_worker_sync.py +++ b/src/dstack/_internal/server/services/runs/router_worker_sync.py @@ -1,25 +1,32 @@ """Reconcile SGLang router /workers with dstack's ready worker replicas (async, SSH-tunneled).""" +import asyncio import json -from dataclasses import dataclass -from typing import Any, List, Literal, Optional, TypedDict -from urllib.parse import urlsplit, urlunsplit +import logging +from collections.abc import Awaitable +from contextlib import AsyncExitStack +from typing import Annotated, Any, Literal, Optional +from urllib.parse import urlsplit +from uuid import UUID import grpc from google.protobuf.json_format import MessageToDict from httpx import AsyncClient, RequestError, Response +from pydantic import FailFast, Field, OnErrorOmit, TypeAdapter, ValidationError from smg_grpc_proto import ( sglang_scheduler_pb2, sglang_scheduler_pb2_grpc, vllm_engine_pb2, vllm_engine_pb2_grpc, ) -from typing_extensions import NotRequired -from dstack._internal.core.errors import SSHError -from dstack._internal.core.models.common import validate_json_extra_ignore -from dstack._internal.core.models.configurations import ReplicaGroup, ServiceConfiguration -from dstack._internal.core.models.runs import JobStatus, RunSpec, get_service_port +# > Because of runtime limitations, Pydantic will require using the TypedDict type from +# > typing_extensions when using Python 3.12 and lower. +from typing_extensions import NotRequired, TypedDict + +from dstack._internal.core.errors import DstackError, SSHError +from dstack._internal.core.models.configurations import ServiceConfiguration +from dstack._internal.core.models.runs import JobStatus, get_service_port from dstack._internal.server.models import JobModel, RunModel from dstack._internal.server.services.jobs import get_job_provisioning_data, get_job_spec from dstack._internal.server.services.jobs.job_replica_grpc_client import ( @@ -31,6 +38,8 @@ ) from dstack._internal.server.services.jobs.job_replica_tunnel import get_service_replica_tunnel from dstack._internal.server.services.logging import fmt +from dstack._internal.server.services.runs import get_run_spec +from dstack._internal.server.utils.common import gather_async, gather_map_async from dstack._internal.utils.logging import get_logger from .replicas import job_belongs_to_group @@ -38,67 +47,325 @@ logger = get_logger(__name__) -# Requests are made over a UDS tunnel to the replica, so the authority is a placeholder. -_HTTP_BASE_URL = "http://dstack" -_HTTP_TIMEOUT = 10.0 -_MAX_SERVER_INFO_RESPONSE_BYTES = 256 * 1024 -_MAX_WORKERS_RESPONSE_BYTES = 2 * 1024 * 1024 -_MAX_WORKERS_COMMAND_ACK_BYTES = 64 * 1024 -_MAX_WORKERS_LIST_ITEMS = 8192 -_GRPC_TIMEOUT = 30.0 +def run_model_has_sglang_router_replica_group(run_model: RunModel) -> bool: + run_spec = get_run_spec(run_model) + return run_spec_has_sglang_router_replica_group(run_spec) -class _ResponseTooLargeError(Exception): - pass +async def sync_router_workers_for_run_model(run_model: RunModel) -> None: + run_spec = get_run_spec(run_model) + config = run_spec.configuration + if not isinstance(config, ServiceConfiguration): + return -async def _stream_response_body_bytes(resp: Response, max_bytes: int) -> bytes: - buf = bytearray() - async for chunk in resp.aiter_bytes(): - buf.extend(chunk) - if len(buf) > max_bytes: - raise _ResponseTooLargeError() - return bytes(buf) + router_groups = [g for g in config.replica_groups if g.router is not None] + if not router_groups: + logger.debug("%s: no router replica group, skipping worker sync", fmt(run_model)) + return + if len(router_groups) > 1: + logger.warning( + "%s: more than one router replica group, skipping worker sync", fmt(run_model) + ) + return + router_group = router_groups[0] + assert router_group.name is not None, "Replica group name is set by validation" + router_jobs = [ + j + for j in run_model.jobs + if job_belongs_to_group(j, router_group.name) and j.status == JobStatus.RUNNING + ] + if not router_jobs: + logger.debug( + "%s: no running router job in group %s, skipping worker sync", + fmt(run_model), + router_group.name, + ) + return -async def _request_json_limited( - client: AsyncClient, - method: str, - url: str, - *, - max_response_bytes: int, - ok_statuses: set[int], - json_body: Optional[dict] = None, - timeout: float = _HTTP_TIMEOUT, -) -> Any: - kwargs: dict[str, Any] = {"timeout": timeout} - if json_body is not None: - kwargs["json"] = json_body - endpoint = f"{method} {url}" - async with client.stream(method, url, **kwargs) as resp: - if resp.status_code not in ok_statuses: - logger.warning( - "router_http unexpected status endpoint=%s status_code=%s expected=%s", - endpoint, - resp.status_code, - sorted(ok_statuses), + worker_jobs_with_addresses = _get_worker_jobs_with_addresses( + jobs=run_model.jobs, configuration=config, router_group_name=router_group.name + ) + worker_address_to_job_id_map = {address: job.id for job, address in worker_jobs_with_addresses} + + # Probing the workers may take minutes, so no router tunnel is held meanwhile. Each router + # is read first, and connected to again only if it needs updating, which in most syncs it + # doesn't. + router_jobs_with_worker_id_maps: list[tuple[JobModel, dict[UUID, str]]] = [] + for router_job, current_workers in await gather_map_async( + router_jobs, _get_router_replica_workers, max_concurrency=_MAX_REPLICA_CONCURRENCY + ): + if current_workers is None: + continue + job_id_to_current_worker_id_map: dict[UUID, str] = {} + for current_worker in current_workers: + job_id = _get_current_worker_job_id(current_worker) + if job_id is not None: + job_id_to_current_worker_id_map[job_id] = current_worker["id"] + else: + # For backward compatibility with workers registered before + # _DSTACK_JOB_ID_WORKER_LABEL was introduced: look up a job by worker's URL + # TODO: remove eventually + address = urlsplit(current_worker["url"]).netloc + job_id = worker_address_to_job_id_map.get(address) + if job_id is not None: + job_id_to_current_worker_id_map[job_id] = current_worker["id"] + router_jobs_with_worker_id_maps.append((router_job, job_id_to_current_worker_id_map)) + if not router_jobs_with_worker_id_maps: + logger.debug( + "%s: no router in group %s returned its workers, skipping worker sync", + fmt(run_model), + router_group.name, + ) + return + + # A set of worker job ids registered in _all_ router jobs. + # If at least one router is missing a worker job, we will probe that job -- we don't reuse + # worker info reported by other routers on purpose -- to avoid propagating possibly incomplete + # (e.g., collected by an older dstack server) or stale (e.g., a dead worker servicer inside + # a still running job) info infinitely. + registered_worker_job_ids = set.intersection( + *(set(id_map.keys()) for _, id_map in router_jobs_with_worker_id_maps) + ) + job_id_to_target_worker_map: dict[UUID, _TargetWorker] = {} + worker_jobs_to_probe: list[JobModel] = [] + probe_coros: list[Awaitable[Optional[_TargetWorker]]] = [] + for worker_job, address in worker_jobs_with_addresses: + if worker_job.id in registered_worker_job_ids: + continue + worker_jobs_to_probe.append(worker_job) + probe_coros.append(_probe_worker_replica(worker_job, address=address)) + if probe_coros: + for worker_job, target_worker in zip( + worker_jobs_to_probe, + await gather_async(probe_coros, max_concurrency=_MAX_REPLICA_CONCURRENCY), + ): + if target_worker is None: + continue + _set_target_worker_job_id(target_worker, worker_job.id) + job_id_to_target_worker_map[worker_job.id] = target_worker + + running_worker_job_ids = {job.id for job, _ in worker_jobs_with_addresses} + sync_coros: list[Awaitable[None]] = [] + for router_job, job_id_to_current_worker_id_map in router_jobs_with_worker_id_maps: + to_add = [ + target_worker + for job_id, target_worker in job_id_to_target_worker_map.items() + if job_id not in job_id_to_current_worker_id_map + ] + to_remove = [ + worker_id + for job_id, worker_id in job_id_to_current_worker_id_map.items() + if job_id not in running_worker_job_ids + ] + if not to_add and not to_remove: + logger.debug("Router %s: workers in sync", fmt(router_job)) + continue + sync_coros.append( + _sync_router_replica_workers(router_job, to_add=to_add, to_remove=to_remove) + ) + if sync_coros: + await gather_async(sync_coros, max_concurrency=_MAX_REPLICA_CONCURRENCY) + + +async def _get_router_replica_workers(job: JobModel) -> Optional[list["_CurrentWorker"]]: + log_prefix = f"Router {fmt(job)}" + try: + async with get_service_replica_client(job) as client: + return await _get_router_workers(client, log_prefix=log_prefix) + except SSHError as e: + _log_unreachable_replica(log_prefix=log_prefix, error=e) + return None + except Exception: + logger.exception("%s: unexpected error when getting workers", log_prefix) + return None + + +async def _sync_router_replica_workers( + job: JobModel, *, to_add: list["_TargetWorker"], to_remove: list[str] +) -> None: + log_prefix = f"Router {fmt(job)}" + try: + async with get_service_replica_client(job) as client: + for worker in to_add: + await _add_worker_to_router(client, log_prefix=log_prefix, worker=worker) + for worker_id in to_remove: + await _remove_worker_from_router( + client, log_prefix=log_prefix, worker_id=worker_id + ) + except SSHError as e: + _log_unreachable_replica(log_prefix=log_prefix, error=e) + except Exception: + logger.exception("%s: unexpected error when syncing workers", log_prefix) + + +async def _probe_worker_replica(job: JobModel, *, address: str) -> Optional["_TargetWorker"]: + log_prefix = f"Worker {fmt(job)}" + try: + async with AsyncExitStack() as exit_stack: + uds_path = await exit_stack.enter_async_context(get_service_replica_tunnel(job)) + http_client = await exit_stack.enter_async_context( + get_service_replica_http_client_over_uds(uds_path) ) - return None - cl = resp.headers.get("content-length") - if cl is not None: + grpc_channel = await exit_stack.enter_async_context( + get_service_replica_grpc_channel_over_uds(uds_path) + ) + tasks = [ + asyncio.create_task( + _probe_http_worker(http_client, log_prefix=log_prefix, address=address) + ), + asyncio.create_task( + _probe_sglang_grpc_worker(grpc_channel, log_prefix=log_prefix, address=address) + ), + asyncio.create_task( + _probe_vllm_grpc_worker(grpc_channel, log_prefix=log_prefix, address=address) + ), + ] try: - if int(cl) > max_response_bytes: - raise _ResponseTooLargeError() - except ValueError: - pass - raw = await _stream_response_body_bytes(resp, max_response_bytes) + for future in asyncio.as_completed(tasks): + try: + worker = await future + except Exception: + # Don't let one probe failure prevent other probes from succeeding + logger.exception("%s: probe failed unexpectedly", log_prefix) + continue + if worker is not None: + logger.debug("%s: probe succeeded: %s", log_prefix, worker) + return worker + finally: + for task in tasks: + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + logger.debug("%s: all probes failed", log_prefix) + except SSHError as e: + _log_unreachable_replica(log_prefix=log_prefix, error=e) + return None + except Exception: + logger.exception("%s: unexpected error when probing", log_prefix) + return None + + +def _log_unreachable_replica(log_prefix: str, error: SSHError) -> None: + # Warning is the right level: a job only reaches `RUNNING` after the server has talked to + # its runner over SSH, so an unreachable replica is always a regression, never a replica + # that has not started yet. + logger.warning("%s: failed to connect: %r", log_prefix, error) + + +def _get_worker_jobs_with_addresses( + jobs: list[JobModel], configuration: ServiceConfiguration, router_group_name: str +) -> list[tuple[JobModel, str]]: + """ + Returns a list of running worker jobs with their addresses. + + Address format is "{internal_ip or hostname}:{port}", no protocol component. + + Returns: + A list of (JobModel, address) pairs. + """ + worker_jobs_with_addresses: list[tuple[JobModel, str]] = [] + for job in jobs: + job_spec = get_job_spec(job) + if job_spec.replica_group == router_group_name: + continue + + if job.status != JobStatus.RUNNING: + logger.debug("%s: not running, skipping", fmt(job)) + continue + + jpd = get_job_provisioning_data(job) + if jpd is None: + logger.debug("%s: no job provisioning data, skipping", fmt(job)) + continue + + hostname = jpd.internal_ip + if not hostname: + if jpd.hostname: + hostname = jpd.hostname + logger.debug("%s: internal_ip is not set, using hostname as a fallback", fmt(job)) + else: + logger.debug("%s: neither internal_ip nor hostname is set, skipping", fmt(job)) + continue + + port = get_service_port(job_spec, configuration) + address = f"{hostname}:{port}" + worker_jobs_with_addresses.append((job, address)) + + return worker_jobs_with_addresses + + +def _get_current_worker_job_id(worker: "_CurrentWorker") -> Optional[UUID]: + labels = worker.get("labels") + if not labels: + return None + job_id = labels.get(_WORKER_LABEL_DSTACK_JOB_ID) + if not job_id: + return None try: - return json.loads(raw) - except json.JSONDecodeError: - logger.warning("router_http JSON parse failed endpoint=%s", endpoint) + return UUID(job_id) + except ValueError as e: + logger.warning("Unparsable dstack job id worker label: %r: %s", job_id, e) return None +def _set_target_worker_job_id(worker: "_TargetWorker", job_id: UUID) -> None: + labels = worker.setdefault("labels", {}) + labels[_WORKER_LABEL_DSTACK_JOB_ID] = str(job_id) + + +_DEFAULT_HTTP_TIMEOUT = 10.0 +_MAX_LOGGED_RESPONSE_BYTES = 1024 +_MAX_REPLICA_CONCURRENCY = 8 + +# Requests are made over a UDS tunnel to the replica, so the authority is a placeholder +# for debugging purposes only +_ROUTER_BASE_URL = "http://router" +_ROUTER_MAX_WORKERS_RESPONSE_BYTES = 2 * 1024 * 1024 +_ROUTER_MAX_WORKERS_COMMAND_ACK_BYTES = 64 * 1024 +_ROUTER_MAX_WORKERS_LIST_ITEMS = 8192 + +_WORKER_BASE_URL = "http://worker" +_WORKER_PROBE_TIMEOUT = 30.0 +_WORKER_MAX_SERVER_INFO_RESPONSE_BYTES = 256 * 1024 +_WORKER_LABEL_DSTACK_JOB_ID = "dstack.ai/job-id" + + +# https://github.com/smg-project/smg/blob/3be823a700fabaff3add8a390cf78f163479d686/crates/protocols/src/worker.rs#L903 +class _CurrentWorker(TypedDict): + id: str + url: str + labels: NotRequired[dict[str, str]] + + +class _WorkersStats(TypedDict): + regular_count: NotRequired[OnErrorOmit[int]] + prefill_count: NotRequired[OnErrorOmit[int]] + decode_count: NotRequired[OnErrorOmit[int]] + + +class _WorkerListResponse(TypedDict): + workers: Annotated[ + list[_CurrentWorker], Field(max_length=_ROUTER_MAX_WORKERS_LIST_ITEMS), FailFast() + ] + stats: NotRequired[_WorkersStats] + + +class _WorkerErrorResponse(TypedDict): + error: str + code: str + + +class _CreateDeleteWorkerResponse(TypedDict): + status: str + worker_id: str + + +_worker_list_response_ta = TypeAdapter(_WorkerListResponse) +_worker_error_response_ta = TypeAdapter(_WorkerErrorResponse) +_create_delete_worker_response_ta = TypeAdapter(_CreateDeleteWorkerResponse) + + # https://github.com/smg-project/smg/blob/3be823a700fabaff3add8a390cf78f163479d686/crates/protocols/src/worker.rs#L601 # Only fields used to register a worker (POST /workers) are included class _TargetWorker(TypedDict): @@ -109,6 +376,7 @@ class _TargetWorker(TypedDict): bootstrap_port: NotRequired[int] kv_connector: NotRequired[str] kv_role: NotRequired[str] + labels: NotRequired[dict[str, str]] # https://github.com/smg-project/smg/blob/3be823a700fabaff3add8a390cf78f163479d686/crates/protocols/src/worker.rs#L29 @@ -118,318 +386,312 @@ class _TargetWorker(TypedDict): # https://github.com/smg-project/smg/blob/3be823a700fabaff3add8a390cf78f163479d686/crates/protocols/src/worker.rs#L76 # Only modes we support are included _ConnectionMode = Literal["http", "grpc"] -# The order does matter -- we discover connection modes in the specified order -_CONNECTION_MODES: tuple[_ConnectionMode, ...] = ("http", "grpc") # https://github.com/smg-project/smg/blob/3be823a700fabaff3add8a390cf78f163479d686/crates/protocols/src/worker.rs#L227 # Only types we support are included _RuntimeType = Literal["sglang", "vllm"] -# The order does matter -- we discover runtime types in the specified order -_RUNTIME_TYPES: tuple[_RuntimeType, ...] = ("sglang", "vllm") -def run_model_has_sglang_router_replica_group(run_model: RunModel) -> bool: - run_spec = validate_json_extra_ignore(RunSpec, run_model.run_spec) - return run_spec_has_sglang_router_replica_group(run_spec) +class _ResponseTooLargeError(DstackError): + pass -def _get_router_jobs(run_model: RunModel, router_group: ReplicaGroup) -> List[JobModel]: - group_name = router_group.name - assert group_name is not None, "Replica group name is set by validation" - # The router group is validated to have `replicas: 1`, but a rolling deployment runs the - # replacement router alongside the old one until the old one is scaled down. Every running - # router is synced, otherwise the replacement has no workers until the old one is gone. - return [ - j - for j in run_model.jobs - if job_belongs_to_group(j, group_name) and j.status == JobStatus.RUNNING - ] +async def _stream_response_body_bytes(resp: Response, max_bytes: int) -> bytes: + buf = bytearray() + async for chunk in resp.aiter_bytes(): + buf.extend(chunk) + if len(buf) > max_bytes: + raise _ResponseTooLargeError() + return bytes(buf) -def _normalize_worker_url(url: str) -> str: - url = url.strip() - parts = urlsplit(url) - path = (parts.path or "").rstrip("/") - return urlunsplit((parts.scheme, parts.netloc, path, parts.query, parts.fragment)) - - -def _get_connection_mode_from_workers( - current_workers: List[dict], -) -> Optional[_ConnectionMode]: - # PD services register multiple workers (e.g. prefill and decode). We expect - # every listed worker to use the same connection_mode (all grpc or all http), - # not a mix of protocols on one router. - modes: set[str] = set() - for worker in current_workers: - mode = worker.get("connection_mode") - if isinstance(mode, str) and mode in ("http", "grpc"): - modes.add(mode) - if modes == {"grpc"}: - return "grpc" - if modes == {"http"}: - return "http" - return None - - -def _get_runtime_type_from_workers( - current_workers: List[dict], -) -> Optional[_RuntimeType]: - # We expect every listed gRPC worker to share the same runtime_type - # (all sglang or all vllm), not a mix of runtimes on one router. - runtimes: set[str] = set() - for worker in current_workers: - # For HTTP workers,there is no “pick vLLM vs SGLang gRPC stub” step, - # so runtime_type is irrelevant for HTTP workers. - if worker.get("connection_mode") != "grpc": - continue - runtime_type = worker.get("runtime_type") - if isinstance(runtime_type, str) and runtime_type in _RUNTIME_TYPES: - runtimes.add(runtime_type) - if runtimes == {"sglang"}: - return "sglang" - if runtimes == {"vllm"}: - return "vllm" - return None +async def _request_json_limited( + client: AsyncClient, + method: str, + url: str, + *, + max_response_bytes: int, + json_body: Optional[dict] = None, + timeout: float = _DEFAULT_HTTP_TIMEOUT, +) -> tuple[int, bytes]: + kwargs: dict[str, Any] = {"timeout": timeout} + if json_body is not None: + kwargs["json"] = json_body + async with client.stream(method, url, **kwargs) as resp: + cl = resp.headers.get("content-length") + if cl is not None: + try: + if int(cl) > max_response_bytes: + raise _ResponseTooLargeError() + except ValueError: + pass + body = await _stream_response_body_bytes(resp, max_response_bytes) + return resp.status_code, body + +def _fmt_response_body(body: bytes) -> str: + if len(body) <= _MAX_LOGGED_RESPONSE_BYTES: + return repr(body) + return f"{body[:_MAX_LOGGED_RESPONSE_BYTES]!r}... ({len(body)} bytes total)" -async def _get_router_workers(client: AsyncClient) -> Optional[List[dict]]: + +async def _get_router_workers( + client: AsyncClient, *, log_prefix: str +) -> Optional[list[_CurrentWorker]]: try: - data = await _request_json_limited( + status_code, raw_response = await _request_json_limited( client, "GET", - f"{_HTTP_BASE_URL}/workers", - max_response_bytes=_MAX_WORKERS_RESPONSE_BYTES, - ok_statuses={200}, + f"{_ROUTER_BASE_URL}/workers", + max_response_bytes=_ROUTER_MAX_WORKERS_RESPONSE_BYTES, ) - if not isinstance(data, dict): - # Non-200 status or unparsable response or response is not a JSON object - return None - workers = data.get("workers") - if not isinstance(workers, list): - # Unexpected response structure -- `workers` is missing or is not an array - return None - # TODO: Truncating a long list and/or dropping unexpectedly shaped items doesn't seem - # right. We should add validation (an item must be a dict with some required fields, see - # _get_workers_diff) and decide what to do with a partially valid list - if len(workers) > _MAX_WORKERS_LIST_ITEMS: - logger.warning( - "Router /workers list exceeds %s items, truncating", - _MAX_WORKERS_LIST_ITEMS, - ) - workers = workers[:_MAX_WORKERS_LIST_ITEMS] - return [w for w in workers if isinstance(w, dict)] - except _ResponseTooLargeError: - logger.warning("Router /workers response exceeded size limit") except RequestError as e: - logger.debug("Router /workers not ready yet: %r", e) - return None + logger.debug("%s: GET /workers: request failed: %r", log_prefix, e) + return None + except _ResponseTooLargeError: + logger.warning("%s: GET /workers: response too large", log_prefix) + return None + + if status_code != 200: + logger.warning( + "%s: GET /workers: unexpected status code: %d: %s", + log_prefix, + status_code, + _fmt_response_body(raw_response), + ) + return None + + try: + response = _worker_list_response_ta.validate_json(raw_response) + except ValidationError as e: + logger.warning("%s: GET /workers: response validation failed: %s", log_prefix, e) + return None + + logger.debug("%s: workers stats: %s", log_prefix, response.get("stats")) + return response["workers"] -async def _add_worker_to_router(client: AsyncClient, worker: _TargetWorker) -> bool: +async def _add_worker_to_router( + client: AsyncClient, *, log_prefix: str, worker: _TargetWorker +) -> None: url = worker["url"] try: - body = await _request_json_limited( + status_code, raw_response = await _request_json_limited( client, "POST", - f"{_HTTP_BASE_URL}/workers", - max_response_bytes=_MAX_WORKERS_COMMAND_ACK_BYTES, - ok_statuses={202}, + f"{_ROUTER_BASE_URL}/workers", + max_response_bytes=_ROUTER_MAX_WORKERS_COMMAND_ACK_BYTES, json_body=dict(worker), ) - added = isinstance(body, dict) and body.get("status") == "accepted" - if not added: - logger.warning("Unexpected add-worker response for %s: %s", url, body) - return added - except _ResponseTooLargeError: - logger.warning("Router add-worker response exceeded size limit for %s", url) except RequestError as e: - logger.warning("Error adding worker %s: %r", url, e) - return False - - -async def _remove_worker_from_router_by_id( - client: AsyncClient, worker_id: str, *, worker_url: str -) -> bool: - try: - body = await _request_json_limited( - client, - "DELETE", - f"{_HTTP_BASE_URL}/workers/{worker_id}", - max_response_bytes=_MAX_WORKERS_COMMAND_ACK_BYTES, - ok_statuses={202}, - ) - removed = isinstance(body, dict) and body.get("status") == "accepted" - if not removed: - logger.warning("Unexpected remove-worker response for %s: %s", worker_url, body) - return removed + logger.warning("%s: POST /workers: %s: request failed: %r", log_prefix, url, e) + return except _ResponseTooLargeError: - logger.warning("Router remove-worker response exceeded size limit for %s", worker_url) - except RequestError as e: - logger.warning("Error removing worker %s: %r", worker_url, e) - return False - + logger.warning("%s: POST /workers: %s: response too large", log_prefix, url) + return -@dataclass -class _WorkersDiff: - to_add: List[_TargetWorker] - to_remove: dict[str, Optional[str]] - """Normalized worker URL to the router's worker id, `None` if the router reported none""" + if status_code == 202: + try: + response = _create_delete_worker_response_ta.validate_json(raw_response) + except ValidationError as e: + logger.warning( + "%s: POST /workers: %s: accepted response validation failed: %s", + log_prefix, + url, + e, + ) + return - def is_empty(self) -> bool: - return not self.to_add and not self.to_remove + if response["status"] != "accepted": + logger.warning( + "%s: POST /workers: %s: unexpected accepted status: %s", + log_prefix, + url, + response["status"], + ) + else: + logger.debug( + "%s: POST /workers: %s: accepted: %s", log_prefix, url, response["worker_id"] + ) + elif status_code == 409: + try: + response = _worker_error_response_ta.validate_json(raw_response) + except ValidationError as e: + logger.warning( + "%s: POST /workers: %s: conflict response validation failed: %s", + log_prefix, + url, + e, + ) + return -def _get_workers_diff( - target_workers: List[_TargetWorker], current_workers: List[dict] -) -> _WorkersDiff: - current_ids_by_norm_url: dict[str, Optional[str]] = {} - for w in current_workers: - u = w.get("url") - if not isinstance(u, str) or not u: - continue - norm_u = _normalize_worker_url(u) - wid = w.get("id") - if isinstance(wid, str) and wid: - current_ids_by_norm_url[norm_u] = wid + if response["code"] not in ["WORKER_CREATE_IN_PROGRESS", "WORKER_ALREADY_EXISTS"]: + logger.warning( + "%s: POST /workers: %s: unexpected conflict code: %s", + log_prefix, + url, + response["code"], + ) else: - current_ids_by_norm_url.setdefault(norm_u, None) - target_by_norm_url = {_normalize_worker_url(t["url"]): t for t in target_workers} - to_add = sorted(target_by_norm_url.keys() - current_ids_by_norm_url.keys()) - to_remove = sorted(current_ids_by_norm_url.keys() - target_by_norm_url.keys()) - return _WorkersDiff( - to_add=[target_by_norm_url[u] for u in to_add], - to_remove={u: current_ids_by_norm_url[u] for u in to_remove}, - ) + logger.debug("%s: POST /workers: %s: conflict: %s", log_prefix, url, response["code"]) + else: + logger.warning( + "%s: POST /workers: %s: unexpected status code: %d: %s", + log_prefix, + url, + status_code, + _fmt_response_body(raw_response), + ) -async def _apply_workers_diff( - client: AsyncClient, diff: _WorkersDiff, *, router_job: JobModel -) -> None: - for worker in diff.to_add: - if not await _add_worker_to_router(client, worker): - logger.debug( - "%s: failed to add worker %s, continuing with others", - fmt(router_job), - worker["url"], - ) - for url, worker_id in diff.to_remove.items(): - if worker_id is None: - logger.error("%s: no worker id found for url %s", fmt(router_job), url) - continue - if not await _remove_worker_from_router_by_id(client, worker_id, worker_url=url): - logger.debug( - "%s: failed to remove worker %s, continuing with others", fmt(router_job), url - ) +async def _remove_worker_from_router( + client: AsyncClient, *, log_prefix: str, worker_id: str +) -> None: + try: + status_code, raw_response = await _request_json_limited( + client, + "DELETE", + f"{_ROUTER_BASE_URL}/workers/{worker_id}", + max_response_bytes=_ROUTER_MAX_WORKERS_COMMAND_ACK_BYTES, + ) + except RequestError as e: + logger.warning("%s: DELETE /workers/%s: request failed: %r", log_prefix, worker_id, e) + return + except _ResponseTooLargeError: + logger.warning("%s: DELETE /workers/%s: response too large", log_prefix, worker_id) + return -def _vllm_kv_role_to_worker_type(kv_role: str) -> _WorkerType: - if kv_role == "kv_producer": - return "prefill" - if kv_role == "kv_consumer": - return "decode" - return "regular" + if status_code != 202: + logger.warning( + "%s: DELETE /workers/%s: unexpected status code: %d: %s", + log_prefix, + worker_id, + status_code, + _fmt_response_body(raw_response), + ) + return + try: + response = _create_delete_worker_response_ta.validate_json(raw_response) + except ValidationError as e: + logger.warning( + "%s: DELETE /workers/%s: response validation failed: %s", log_prefix, worker_id, e + ) + return -def _is_expected_grpc_error(error: grpc.aio.AioRpcError) -> bool: - """Expected while a gRPC worker is still starting or the wrong stub is probed.""" - return error.code() in ( - grpc.StatusCode.UNAVAILABLE, - grpc.StatusCode.DEADLINE_EXCEEDED, - grpc.StatusCode.UNIMPLEMENTED, - ) + if response["status"] != "accepted": + logger.warning( + "%s: DELETE /workers/%s: unexpected accepted status: %s", + log_prefix, + worker_id, + response["status"], + ) + else: + logger.debug("%s: DELETE /workers/%s: accepted", log_prefix, worker_id) -async def _probe_http_worker(client: AsyncClient, *, address: str) -> Optional[_TargetWorker]: +async def _probe_http_worker( + client: AsyncClient, *, log_prefix: str, address: str +) -> Optional[_TargetWorker]: # The request goes over the tunnel, `worker_url` is the address the router itself dials. worker_url = f"http://{address}" + logger.debug("%s: %s: probing", log_prefix, worker_url) + try: - data = await _request_json_limited( + status_code, raw_response = await _request_json_limited( client, "GET", - f"{_HTTP_BASE_URL}/server_info", - max_response_bytes=_MAX_SERVER_INFO_RESPONSE_BYTES, - ok_statuses={200}, + f"{_WORKER_BASE_URL}/server_info", + max_response_bytes=_WORKER_MAX_SERVER_INFO_RESPONSE_BYTES, + timeout=_WORKER_PROBE_TIMEOUT, ) - if isinstance(data, dict): - if data.get("status") != "ready": - return None - mode = data.get("disaggregation_mode", "") - if mode == "prefill": - bootstrap_port = data.get("disaggregation_bootstrap_port") - worker: _TargetWorker = { - "url": worker_url, - "worker_type": "prefill", - "connection_mode": "http", - "runtime_type": "sglang", - } - if bootstrap_port is not None: - worker["bootstrap_port"] = bootstrap_port - return worker - if mode == "decode": - return { - "url": worker_url, - "worker_type": "decode", - "connection_mode": "http", - "runtime_type": "sglang", - } + except RequestError as e: + logger.debug("%s: %s: GET /server_info: request failed: %r", log_prefix, worker_url, e) + return None + except _ResponseTooLargeError: + logger.warning("%s: %s: GET /server_info: response too large", log_prefix, worker_url) + return None + + if status_code != 200: + logger.warning( + "%s: %s: GET /server_info: unexpected status code: %s", + log_prefix, + worker_url, + status_code, + ) + return None + + # ValueError, not JSONDecodeError: a body that is not valid UTF-8 raises UnicodeDecodeError + try: + data = json.loads(raw_response) + except ValueError as e: + logger.warning( + "%s: %s: GET /server_info: response parsing failed: %s", log_prefix, worker_url, e + ) + return None + if isinstance(data, dict): + if data.get("status") != "ready": + return None + mode = data.get("disaggregation_mode", "") + if mode == "prefill": + bootstrap_port = data.get("disaggregation_bootstrap_port") + worker: _TargetWorker = { + "url": worker_url, + "worker_type": "prefill", + "connection_mode": "http", + "runtime_type": "sglang", + } + if bootstrap_port is not None: + worker["bootstrap_port"] = bootstrap_port + return worker + if mode == "decode": return { "url": worker_url, - "worker_type": "regular", + "worker_type": "decode", "connection_mode": "http", "runtime_type": "sglang", } - except _ResponseTooLargeError: - logger.warning("server_info response too large for worker %s", worker_url) - except RequestError as e: - logger.debug("Could not fetch server_info for worker %s: %r", worker_url, e) - return None - - -async def _get_grpc_server_info( - channel: grpc.aio.Channel, - runtime_type: _RuntimeType, -) -> Any: - if runtime_type == "sglang": - stub = sglang_scheduler_pb2_grpc.SglangSchedulerStub(channel) - request = sglang_scheduler_pb2.GetServerInfoRequest() - else: - stub = vllm_engine_pb2_grpc.VllmEngineStub(channel) - request = vllm_engine_pb2.GetServerInfoRequest() - return await stub.GetServerInfo(request, timeout=_GRPC_TIMEOUT) - - -def _grpc_server_info_to_worker( - worker_url: str, - runtime_type: _RuntimeType, - response: Any, -) -> _TargetWorker: - if runtime_type == "vllm": - kv_role = response.kv_role or "" - kv_connector = response.kv_connector or "" - worker: _TargetWorker = { + return { "url": worker_url, - "connection_mode": "grpc", - "runtime_type": runtime_type, - "worker_type": _vllm_kv_role_to_worker_type(kv_role), + "worker_type": "regular", + "connection_mode": "http", + "runtime_type": "sglang", } - if kv_connector: - worker["kv_connector"] = kv_connector - if kv_role: - worker["kv_role"] = kv_role - return worker - - server_args = ( - MessageToDict(response.server_args, preserving_proto_field_name=True) - if response.server_args is not None - else {} - ) - mode = server_args.get("disaggregation_mode") - worker_type = mode if mode in ("prefill", "decode") else "regular" - worker = { + + +async def _probe_sglang_grpc_worker( + channel: grpc.aio.Channel, *, log_prefix: str, address: str +) -> Optional[_TargetWorker]: + runtime_type: _RuntimeType = "sglang" + # The RPC goes over the tunnel, `worker_url` is the address the router itself dials. + worker_url = f"grpc://{address}" + logger.debug("%s: %s: probing %s", log_prefix, worker_url, runtime_type) + + stub = sglang_scheduler_pb2_grpc.SglangSchedulerStub(channel) + request = sglang_scheduler_pb2.GetServerInfoRequest() + response: sglang_scheduler_pb2.GetServerInfoResponse + try: + response = await stub.GetServerInfo(request, timeout=_WORKER_PROBE_TIMEOUT) + except grpc.aio.AioRpcError as e: + _log_failed_grpc_worker_probe( + log_prefix=log_prefix, worker_url=worker_url, runtime_type=runtime_type, error=e + ) + return None + + server_args = MessageToDict(response.server_args, preserving_proto_field_name=True) + worker_type: _WorkerType + match disaggregation_mode := server_args.get("disaggregation_mode"): + case "prefill" | "decode": + worker_type = disaggregation_mode + case _: + worker_type = "regular" + worker: _TargetWorker = { "url": worker_url, + "worker_type": worker_type, "connection_mode": "grpc", "runtime_type": runtime_type, - "worker_type": worker_type, } if worker_type == "prefill": bootstrap_port = server_args.get("disaggregation_bootstrap_port") @@ -438,200 +700,66 @@ def _grpc_server_info_to_worker( return worker -async def _probe_grpc_worker( - channel: grpc.aio.Channel, - *, - address: str, - runtime_type: Optional[_RuntimeType] = None, +async def _probe_vllm_grpc_worker( + channel: grpc.aio.Channel, *, log_prefix: str, address: str ) -> Optional[_TargetWorker]: + runtime_type: _RuntimeType = "vllm" # The RPC goes over the tunnel, `worker_url` is the address the router itself dials. worker_url = f"grpc://{address}" - runtime_types: tuple[_RuntimeType, ...] - if runtime_type is None: - # Bootstrap only: router workers list has no runtime_type yet, should try all - runtime_types = _RUNTIME_TYPES - else: - runtime_types = (runtime_type,) - for runtime_type in runtime_types: - try: - response = await _get_grpc_server_info(channel, runtime_type) - break - except grpc.aio.AioRpcError as e: - if _is_expected_grpc_error(e): - continue - raise - else: - logger.debug("gRPC worker %s not ready (GetServerInfo)", worker_url) - return None - return _grpc_server_info_to_worker(worker_url, runtime_type, response) + logger.debug("%s: %s: probing %s", log_prefix, worker_url, runtime_type) - -async def _get_worker( - job_model: JobModel, - *, - address: str, - connection_mode: Optional[_ConnectionMode] = None, - runtime_type: Optional[_RuntimeType] = None, -) -> Optional[_TargetWorker]: - connection_modes: tuple[_ConnectionMode, ...] - if connection_mode is None: - # No connection_mode discovered -- should probe all - connection_modes = _CONNECTION_MODES - else: - connection_modes = (connection_mode,) + stub = vllm_engine_pb2_grpc.VllmEngineStub(channel) + request = vllm_engine_pb2.GetServerInfoRequest() + response: vllm_engine_pb2.GetServerInfoResponse try: - async with get_service_replica_tunnel(job_model) as uds_path: - for connection_mode in connection_modes: - if connection_mode == "grpc": - async with get_service_replica_grpc_channel_over_uds(uds_path) as channel: - worker = await _probe_grpc_worker( - channel, address=address, runtime_type=runtime_type - ) - elif connection_mode == "http": - async with get_service_replica_http_client_over_uds(uds_path) as client: - worker = await _probe_http_worker(client, address=address) - if worker is not None: - return worker - except SSHError as e: - # An unreachable worker is reported as not ready rather than aborting the sync, so that - # one dead replica cannot hold back registration of the healthy ones. The cost is that a - # transient failure deregisters a healthy worker until the next sync re-adds it. - # TODO: `_get_workers_diff` cannot tell "not serving" from "could not be - # reached" -- both mean "absent from the target list", hence "remove". A third, unknown - # outcome should be excluded from both `to_add` and `to_remove`, leaving an unreachable - # worker as the router last saw it. - logger.warning("%s: failed to connect to worker replica: %r", fmt(job_model), e) - return None - - -async def _build_target_workers( - run_model: RunModel, - run_spec: RunSpec, - replica_groups: list[ReplicaGroup], - *, - connection_mode: Optional[_ConnectionMode] = None, - runtime_type: Optional[_RuntimeType] = None, -) -> List[_TargetWorker]: - workers: List[_TargetWorker] = [] - config = run_spec.configuration - if not isinstance(config, ServiceConfiguration): - return workers - - for group in replica_groups: - if group.router is not None: - continue - assert group.name is not None, "Replica group name is set by validation" - group_name = group.name - for job in run_model.jobs: - if not job_belongs_to_group(job, group_name): - continue - if job.status != JobStatus.RUNNING: - continue - jpd = get_job_provisioning_data(job) - if jpd is None: - continue - hostname = jpd.internal_ip or jpd.hostname - if not hostname: - continue - job_spec = get_job_spec(job) - port = get_service_port(job_spec, config) - worker = await _get_worker( - job, - address=f"{hostname}:{port}", - connection_mode=connection_mode, - runtime_type=runtime_type, - ) - if worker is not None: - workers.append(worker) - else: - logger.debug("%s: worker replica not ready", fmt(job)) - return workers - - -async def sync_router_workers_for_run_model(run_model: RunModel) -> None: - run_spec = validate_json_extra_ignore(RunSpec, run_model.run_spec) - config = run_spec.configuration - if not isinstance(config, ServiceConfiguration): - return - replica_groups = config.replica_groups - router_group = next((g for g in replica_groups if g.router is not None), None) - if router_group is None: - return - - router_jobs = _get_router_jobs(run_model, router_group) - if not router_jobs: - logger.debug( - "%s: no running router job in group %s, skipping worker sync", - fmt(run_model), - router_group.name, + response = await stub.GetServerInfo(request, timeout=_WORKER_PROBE_TIMEOUT) + except grpc.aio.AioRpcError as e: + _log_failed_grpc_worker_probe( + log_prefix=log_prefix, worker_url=worker_url, runtime_type=runtime_type, error=e ) - return - # Probing the workers may take minutes, so no router tunnel is held meanwhile. Each router - # is read first, and connected to again only if it needs updating, which in most syncs it - # doesn't. An unreachable router is skipped, like an unreachable worker, see `_get_worker`. - try: - current_workers_by_router: List[tuple[JobModel, List[dict]]] = [] - for router_job in router_jobs: - current_workers = await _get_router_replica_workers(router_job) - if current_workers is not None: - current_workers_by_router.append((router_job, current_workers)) - if not current_workers_by_router: - logger.debug( - "%s: no router in group %s returned its workers, skipping worker sync", - fmt(run_model), - router_group.name, - ) - return - # The hints spare probing connection modes and runtime types that no registered worker - # uses. They are taken from all routers, as a router started by a rolling deployment has - # no workers yet, which alone would mean "probe everything". - all_current_workers = [w for _, workers in current_workers_by_router for w in workers] - target_workers = await _build_target_workers( - run_model, - run_spec, - replica_groups, - connection_mode=_get_connection_mode_from_workers(all_current_workers), - runtime_type=_get_runtime_type_from_workers(all_current_workers), - ) - for router_job, current_workers in current_workers_by_router: - if _get_workers_diff(target_workers, current_workers).is_empty(): - continue - await _update_workers_in_router_replica(router_job, target_workers) - except Exception: - logger.exception("%s: unexpected error when syncing workers with router", fmt(run_model)) - - -async def _get_router_replica_workers(router_job: JobModel) -> Optional[List[dict]]: - try: - async with get_service_replica_client(router_job) as client: - current_workers = await _get_router_workers(client) - except SSHError as e: - _log_router_replica_unreachable(router_job, e) return None - if current_workers is None: - logger.debug("%s: failed to get current workers from the router", fmt(router_job)) - return current_workers + worker_type: _WorkerType + match response.kv_role: + case "kv_producer": + worker_type = "prefill" + case "kv_consumer": + worker_type = "decode" + case _: + worker_type = "regular" + worker: _TargetWorker = { + "url": worker_url, + "worker_type": worker_type, + "connection_mode": "grpc", + "runtime_type": "vllm", + } + if response.kv_connector: + worker["kv_connector"] = response.kv_connector + if response.kv_role: + worker["kv_role"] = response.kv_role + return worker -async def _update_workers_in_router_replica( - router_job: JobModel, target_workers: List[_TargetWorker] -) -> None: - try: - async with get_service_replica_client(router_job) as client: - # Read again: the list this update was decided on was fetched before the workers - # were probed, possibly minutes ago. - current_workers = await _get_router_workers(client) - if current_workers is None: - logger.debug("%s: failed to get current workers from the router", fmt(router_job)) - return - diff = _get_workers_diff(target_workers, current_workers) - await _apply_workers_diff(client, diff, router_job=router_job) - except SSHError as e: - _log_router_replica_unreachable(router_job, e) +def _log_failed_grpc_worker_probe( + log_prefix: str, worker_url: str, runtime_type: _RuntimeType, error: grpc.aio.AioRpcError +): + log_level = logging.DEBUG if _is_expected_grpc_error(error) else logging.WARNING + logger.log( + log_level, + "%s: %s: %s probe failed: %s: %s", + log_prefix, + worker_url, + runtime_type, + error.code(), + error.details(), + ) -def _log_router_replica_unreachable(router_job: JobModel, error: SSHError) -> None: - # Warning is the right level: a job only reaches `RUNNING` after the server has talked to - # its runner over SSH, so an unreachable replica is always a regression, never a replica - # that has not started yet. - logger.warning("%s: failed to sync workers with router: %r", fmt(router_job), error) + +def _is_expected_grpc_error(error: grpc.aio.AioRpcError) -> bool: + """Expected while a gRPC worker is still starting or the wrong stub is probed.""" + return error.code() in ( + grpc.StatusCode.UNAVAILABLE, + grpc.StatusCode.DEADLINE_EXCEEDED, + grpc.StatusCode.UNIMPLEMENTED, + grpc.StatusCode.CANCELLED, + ) diff --git a/src/dstack/_internal/server/utils/common.py b/src/dstack/_internal/server/utils/common.py index 624e957476..b3981ce173 100644 --- a/src/dstack/_internal/server/utils/common.py +++ b/src/dstack/_internal/server/utils/common.py @@ -1,45 +1,123 @@ import asyncio +import inspect from typing import ( Awaitable, Callable, Iterable, List, + Literal, Optional, Sequence, Tuple, TypeVar, Union, + overload, ) ItemT = TypeVar("ItemT") ResultT = TypeVar("ResultT") +@overload +async def gather_map_async( + items: Sequence[ItemT], + func: Callable[[ItemT], Awaitable[ResultT]], + *, + return_exceptions: Literal[False] = False, + max_concurrency: Optional[int] = None, +) -> List[Tuple[ItemT, ResultT]]: ... + + +@overload +async def gather_map_async( + items: Sequence[ItemT], + func: Callable[[ItemT], Awaitable[ResultT]], + *, + return_exceptions: bool, + max_concurrency: Optional[int] = None, +) -> List[Tuple[ItemT, Union[ResultT, BaseException]]]: ... + + async def gather_map_async( items: Sequence[ItemT], func: Callable[[ItemT], Awaitable[ResultT]], *, return_exceptions: bool = False, -) -> List[Tuple[ItemT, Union[ResultT, BaseException]]]: + max_concurrency: Optional[int] = None, +) -> Union[List[Tuple[ItemT, ResultT]], List[Tuple[ItemT, Union[ResultT, BaseException]]]]: """ A parallel wrapper around asyncio.gather that returns a list of tuples (item, result). Args: items: list of items to be processed func: function to be applied to each item, return awaitable coroutine - return_exceptions: passed to asyncio.gather + return_exceptions: passed to gather_async + max_concurrency: passed to gather_async Returns: list of tuples (item, result) or (item, exception) if return_exceptions is True """ - return [ - (item, result) - for item, result in zip( - items, - await asyncio.gather( - *(func(item) for item in items), return_exceptions=return_exceptions - ), - ) - ] + results = await gather_async( + [func(item) for item in items], + return_exceptions=return_exceptions, + max_concurrency=max_concurrency, + ) + return list(zip(items, results)) + + +@overload +async def gather_async( + aws: Sequence[Awaitable[ResultT]], + *, + return_exceptions: Literal[False] = False, + max_concurrency: Optional[int] = None, +) -> List[ResultT]: ... + + +@overload +async def gather_async( + aws: Sequence[Awaitable[ResultT]], + *, + return_exceptions: bool, + max_concurrency: Optional[int] = None, +) -> List[Union[ResultT, BaseException]]: ... + + +async def gather_async( + aws: Sequence[Awaitable[ResultT]], + *, + return_exceptions: bool = False, + max_concurrency: Optional[int] = None, +) -> Union[List[ResultT], List[Union[ResultT, BaseException]]]: + """ + Like asyncio.gather, but takes a sequence of awaitables and also: + - Runs at most `max_concurrency` awaitables at once, unlimited if None. Only awaitables that + start when awaited, such as coroutines, can be limited. Tasks and futures already run. + - Cancels the remaining awaitables when one fails and `return_exceptions` is False. + """ + tasks: List[asyncio.Future[ResultT]] = [] + try: + semaphore = None + if max_concurrency is not None: + if max_concurrency < 1: + raise ValueError("max_concurrency must be >= 1") + semaphore = asyncio.BoundedSemaphore(max_concurrency) + for aw in aws: + if semaphore is not None: + aw = _await_with_semaphore(aw, semaphore) + tasks.append(asyncio.ensure_future(aw)) + return await asyncio.gather(*tasks, return_exceptions=return_exceptions) + finally: + # asyncio.gather doesn't cancel the remaining tasks when one of them fails + pending = [task for task in tasks if not task.done()] + for task in pending: + task.cancel() + if pending: + await asyncio.gather(*pending, return_exceptions=True) + # Coroutines that never started (queued when cancelled, or rejected by validation) + # would warn that they were never awaited + for aw in aws: + if inspect.iscoroutine(aw): + aw.close() def join_byte_stream_checked(stream: Iterable[bytes], max_size: int) -> Optional[bytes]: @@ -61,3 +139,8 @@ def join_byte_stream_checked(stream: Iterable[bytes], max_size: int) -> Optional def is_background_task_name(name: str) -> bool: return name.startswith(SCHEDULED_TASKS_PREFIX) or name.startswith(PIPELINE_TASKS_PREFIX) + + +async def _await_with_semaphore(aw: Awaitable[ResultT], semaphore: asyncio.Semaphore) -> ResultT: + async with semaphore: + return await aw diff --git a/src/tests/_internal/server/background/pipeline_tasks/test_service_router_worker_sync.py b/src/tests/_internal/server/background/pipeline_tasks/test_service_router_worker_sync.py index 98a4c3998b..a62d50b01d 100644 --- a/src/tests/_internal/server/background/pipeline_tasks/test_service_router_worker_sync.py +++ b/src/tests/_internal/server/background/pipeline_tasks/test_service_router_worker_sync.py @@ -586,9 +586,49 @@ async def test_process_logs_router_job_when_router_connection_fails( await worker.process(item) assert ( - f"job({router_job.id.hex[:6]}){router_job.job_name}:" - f" failed to sync workers with router: {ssh_error!r}" in caplog.text + f"Router job({router_job.id.hex[:6]}){router_job.job_name}:" + f" failed to connect: {ssh_error!r}" in caplog.text ) await session.refresh(sync_row) assert sync_row.deleted is False assert sync_row.lock_token is None + + async def test_process_unlocks_when_sync_fails_unexpectedly( + self, + test_db, + session: AsyncSession, + worker: ServiceRouterWorkerSyncWorker, + caplog: pytest.LogCaptureFixture, + ): + project = await create_project(session=session) + user = await create_user(session=session) + repo = await create_repo(session=session, project_id=project.id) + run = await create_run( + session=session, + project=project, + repo=repo, + user=user, + status=RunStatus.RUNNING, + run_spec=_router_service_run_spec(repo.name), + ) + sync_row = await _add_service_router_worker_sync_row(session, run.id) + sync_row.lock_token = uuid.uuid4() + sync_row.lock_expires_at = get_current_datetime() + timedelta(seconds=30) + sync_row.lock_owner = ServiceRouterWorkerSyncPipeline.__name__ + await session.commit() + item = _sync_row_to_pipeline_item(sync_row) + + with patch( + "dstack._internal.server.background.pipeline_tasks.service_router_worker_sync" + ".sync_router_workers_for_run_model", + new_callable=AsyncMock, + side_effect=RuntimeError("boom"), + ): + await worker.process(item) + + assert "unexpected error when syncing workers with router" in caplog.text + await session.refresh(sync_row) + assert sync_row.deleted is False + assert sync_row.lock_token is None + assert sync_row.lock_expires_at is None + assert sync_row.lock_owner is None diff --git a/src/tests/_internal/server/services/runs/test_router_worker_sync.py b/src/tests/_internal/server/services/runs/test_router_worker_sync.py index 1277f040a4..9e4c583206 100644 --- a/src/tests/_internal/server/services/runs/test_router_worker_sync.py +++ b/src/tests/_internal/server/services/runs/test_router_worker_sync.py @@ -1,33 +1,41 @@ +import asyncio +import copy import json import logging import uuid +from collections.abc import AsyncIterator, Callable, Iterator, Mapping from contextlib import asynccontextmanager, contextmanager from pathlib import Path -from typing import Optional -from unittest.mock import AsyncMock, MagicMock, patch +from typing import Optional, Union +from unittest.mock import AsyncMock, patch import grpc import httpx import pytest +from google.protobuf.message import Message from httpx import AsyncClient +from smg_grpc_proto import sglang_scheduler_pb2, vllm_engine_pb2 from sqlalchemy.ext.asyncio import AsyncSession from dstack._internal.core.errors import SSHError -from dstack._internal.core.models.configurations import parse_run_configuration +from dstack._internal.core.models.configurations import ( + ServiceConfiguration, + parse_run_configuration, +) from dstack._internal.core.models.runs import JobStatus, RunStatus from dstack._internal.server.models import JobModel, RunModel from dstack._internal.server.services.logging import fmt from dstack._internal.server.services.runs import router_worker_sync from dstack._internal.server.services.runs.router_worker_sync import ( _add_worker_to_router, - _get_connection_mode_from_workers, + _get_current_worker_job_id, _get_router_workers, - _get_runtime_type_from_workers, - _get_worker, - _get_workers_diff, - _grpc_server_info_to_worker, - _probe_grpc_worker, + _get_worker_jobs_with_addresses, _probe_http_worker, + _probe_sglang_grpc_worker, + _probe_vllm_grpc_worker, + _probe_worker_replica, + _remove_worker_from_router, _TargetWorker, sync_router_workers_for_run_model, ) @@ -37,79 +45,422 @@ create_repo, create_run, create_user, + get_job_provisioning_data, get_run_spec, ) +# Hardcoded rather than imported: routers keep the label of every worker registered so far, +# so changing the key would break matching workers to jobs. +_JOB_ID_LABEL = "dstack.ai/job-id" -class TestGetConnectionModeFromWorkers: - def test_grpc(self): - current = [{"connection_mode": "grpc"}] - assert _get_connection_mode_from_workers(current) == "grpc" +# Hardcoded rather than taken from the stubs: these are the services the workers serve. +_SGLANG_GET_SERVER_INFO = "/sglang.grpc.scheduler.SglangScheduler/GetServerInfo" +_VLLM_GET_SERVER_INFO = "/vllm.grpc.engine.VllmEngine/GetServerInfo" - def test_http(self): - current = [{"connection_mode": "http"}] - assert _get_connection_mode_from_workers(current) == "http" - def test_mixed(self): - current = [{"connection_mode": "grpc"}, {"connection_mode": "http"}] - assert _get_connection_mode_from_workers(current) is None +@pytest.mark.asyncio +@pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True) +class TestSyncRouterWorkersForRunModel: + async def test_registers_ready_workers_labeled_with_job_id( + self, test_db, session: AsyncSession + ): + run, [router_job], worker_jobs = await _create_router_service_run(session, worker_count=3) + router = _FakeRouter() + # The second worker is not ready yet + probe_results = {_worker_address(0): _http_worker(0), _worker_address(2): _http_worker(2)} + with ( + _fake_router_replicas({router_job.id: router}), + _fake_worker_probes(probe_results) as probe_mock, + ): + await sync_router_workers_for_run_model(run) -class TestRuntimeTypeFromRouterWorkers: - def test_vllm_grpc_workers(self): - current = [{"connection_mode": "grpc", "runtime_type": "vllm"}] - assert _get_runtime_type_from_workers(current) == "vllm" + assert [call.kwargs["address"] for call in probe_mock.await_args_list] == [ + _worker_address(0), + _worker_address(1), + _worker_address(2), + ] + assert router.added == [ + {**_http_worker(0), "labels": _job_id_label(worker_jobs[0])}, + {**_http_worker(2), "labels": _job_id_label(worker_jobs[2])}, + ] + assert router.removed_ids == [] - def test_sglang_grpc_workers(self): - current = [{"connection_mode": "grpc", "runtime_type": "sglang"}] - assert _get_runtime_type_from_workers(current) == "sglang" + async def test_does_not_probe_registered_workers(self, test_db, session: AsyncSession): + # A registered worker is left to the router's own health checks + run, [router_job], [worker_job] = await _create_router_service_run(session, worker_count=1) + router = _FakeRouter( + [_router_entry("w0", f"http://{_worker_address(0)}", _job_id_label(worker_job))] + ) - def test_ignores_http_workers(self): - current = [{"connection_mode": "http", "runtime_type": "sglang"}] - assert _get_runtime_type_from_workers(current) is None + with ( + _fake_router_replicas({router_job.id: router}), + _fake_worker_probes({}) as probe_mock, + ): + await sync_router_workers_for_run_model(run) - def test_mixed_runtimes(self): - current = [ - {"connection_mode": "grpc", "runtime_type": "vllm"}, - {"connection_mode": "grpc", "runtime_type": "sglang"}, - ] - assert _get_runtime_type_from_workers(current) is None - - -class TestGrpcServerInfoToWorker: - def test_vllm_prefill(self): - response = MagicMock(kv_role="kv_producer", kv_connector="NixlConnector") - worker = _grpc_server_info_to_worker("grpc://10.0.0.1:50051", "vllm", response) - assert worker["worker_type"] == "prefill" - assert worker.get("runtime_type") == "vllm" - assert worker.get("kv_role") == "kv_producer" - - def test_sglang_prefill(self): - server_args = MagicMock() - response = MagicMock(server_args=server_args) - with patch( - "dstack._internal.server.services.runs.router_worker_sync.MessageToDict", - return_value={ - "disaggregation_mode": "prefill", - "disaggregation_bootstrap_port": 8998, - }, + probe_mock.assert_not_awaited() + assert router.added == [] + assert router.removed_ids == [] + # Read once, not connected to again, as there is nothing to update + assert router.connections == 1 + + async def test_removes_workers_of_jobs_no_longer_running(self, test_db, session: AsyncSession): + run, [router_job], [running_job] = await _create_router_service_run( + session, worker_count=1 + ) + stopped_job = await create_job( + session=session, + run=run, + status=JobStatus.TERMINATING, + replica_num=2, + replica_group_name="worker", + job_provisioning_data=get_job_provisioning_data(internal_ip="10.0.0.2"), + ) + await session.refresh(run, attribute_names=["jobs"]) + router = _FakeRouter( + [ + _router_entry( + "running", f"http://{_worker_address(0)}", _job_id_label(running_job) + ), + _router_entry( + "stopped", f"http://{_worker_address(1)}", _job_id_label(stopped_job) + ), + ] + ) + + with _fake_router_replicas({router_job.id: router}), _fake_worker_probes({}): + await sync_router_workers_for_run_model(run) + + assert router.removed_ids == ["stopped"] + assert router.added == [] + + async def test_keeps_workers_registered_by_others(self, test_db, session: AsyncSession): + run, [router_job], [worker_job] = await _create_router_service_run(session, worker_count=1) + router = _FakeRouter( + [ + _router_entry("dstack", f"http://{_worker_address(0)}", _job_id_label(worker_job)), + _router_entry("unlabeled", "http://10.0.9.1:8000"), + _router_entry("labeled", "http://10.0.9.2:8000", {"team": "ml"}), + ] + ) + + with _fake_router_replicas({router_job.id: router}), _fake_worker_probes({}): + await sync_router_workers_for_run_model(run) + + assert router.removed_ids == [] + assert router.connections == 1 + + @pytest.mark.parametrize("scheme", ["http", "grpc"]) + async def test_matches_unlabeled_worker_by_address( + self, test_db, session: AsyncSession, scheme: str + ): + # Workers registered before the job id label was introduced have no labels + run, [router_job], _ = await _create_router_service_run(session, worker_count=1) + router = _FakeRouter([_router_entry("legacy", f"{scheme}://{_worker_address(0)}")]) + + with ( + _fake_router_replicas({router_job.id: router}), + _fake_worker_probes({}) as probe_mock, + ): + await sync_router_workers_for_run_model(run) + + probe_mock.assert_not_awaited() + assert router.added == [] + assert router.removed_ids == [] + + async def test_registers_workers_in_replacement_router_only( + self, test_db, session: AsyncSession + ): + # A rolling deployment runs the replacement router alongside the old one + run, [old_router_job, new_router_job], [worker_job] = await _create_router_service_run( + session, router_count=2, worker_count=1 + ) + old_router = _FakeRouter( + [_router_entry("w0", f"http://{_worker_address(0)}", _job_id_label(worker_job))] + ) + new_router = _FakeRouter() + routers = {old_router_job.id: old_router, new_router_job.id: new_router} + + with ( + _fake_router_replicas(routers), + _fake_worker_probes({_worker_address(0): _http_worker(0)}) as probe_mock, + ): + await sync_router_workers_for_run_model(run) + + # Probed anew rather than copied from the old router, see the comment in the code + probe_mock.assert_awaited_once() + assert new_router.added == [{**_http_worker(0), "labels": _job_id_label(worker_job)}] + assert old_router.added == [] + # Read once, connected to again only if it needs updating + assert old_router.connections == 1 + assert new_router.connections == 2 + + @pytest.mark.parametrize( + ["error", "message"], + [ + pytest.param( + SSHError("connection refused"), + "failed to connect: SSHError('connection refused')", + id="unreachable", + ), + pytest.param( + RuntimeError("boom"), "unexpected error when getting workers", id="unexpected" + ), + ], + ) + async def test_failing_router_does_not_block_others( + self, + test_db, + session: AsyncSession, + caplog: pytest.LogCaptureFixture, + error: Exception, + message: str, + ): + caplog.set_level(level=logging.WARNING, logger=router_worker_sync.__name__) + run, [failing_router_job, router_job], [worker_job] = await _create_router_service_run( + session, router_count=2, worker_count=1 + ) + failing_router = _FakeRouter(connect_error=error) + router = _FakeRouter() + routers = {failing_router_job.id: failing_router, router_job.id: router} + + with ( + _fake_router_replicas(routers), + _fake_worker_probes({_worker_address(0): _http_worker(0)}), + ): + await sync_router_workers_for_run_model(run) + + assert router.added == [{**_http_worker(0), "labels": _job_id_label(worker_job)}] + assert failing_router.connections == 1 + assert f"Router {fmt(failing_router_job)}: {message}" in caplog.text + + async def test_failing_router_update_does_not_block_others( + self, test_db, session: AsyncSession, caplog: pytest.LogCaptureFixture + ): + caplog.set_level(level=logging.WARNING, logger=router_worker_sync.__name__) + run, [failing_router_job, router_job], [worker_job] = await _create_router_service_run( + session, router_count=2, worker_count=1 + ) + failing_router = _FakeRouter(update_error=RuntimeError("boom")) + router = _FakeRouter() + routers = {failing_router_job.id: failing_router, router_job.id: router} + + with ( + _fake_router_replicas(routers), + _fake_worker_probes({_worker_address(0): _http_worker(0)}), + ): + await sync_router_workers_for_run_model(run) + + assert router.added == [{**_http_worker(0), "labels": _job_id_label(worker_job)}] + assert ( + f"Router {fmt(failing_router_job)}: unexpected error when syncing workers" + in caplog.text + ) + + async def test_skips_probing_when_no_router_returns_workers( + self, test_db, session: AsyncSession + ): + run, [router_job], _ = await _create_router_service_run(session, worker_count=1) + router = _FakeRouter(connect_error=SSHError("connection refused")) + + with ( + _fake_router_replicas({router_job.id: router}), + _fake_worker_probes({_worker_address(0): _http_worker(0)}) as probe_mock, + ): + await sync_router_workers_for_run_model(run) + + probe_mock.assert_not_awaited() + + +@pytest.mark.asyncio +class TestProbeWorkerReplica: + async def test_detects_http_sglang_worker(self): + with _fake_replica_transports( + http_handler=lambda _: httpx.Response(200, json={"status": "ready"}), + # The HTTP/1.1 server answers the HTTP/2 preface with an error + grpc_channel=_FakeGrpcChannel(unknown_method_code=grpc.StatusCode.UNAVAILABLE), + ): + worker = await _probe_worker_replica(_job(), address="10.0.0.1:8000") + + assert worker == { + "url": "http://10.0.0.1:8000", + "worker_type": "regular", + "connection_mode": "http", + "runtime_type": "sglang", + } + + async def test_detects_sglang_grpc_worker(self): + with _fake_replica_transports( + http_handler=_grpc_server_http_handler, + grpc_channel=_FakeGrpcChannel( + {_SGLANG_GET_SERVER_INFO: sglang_scheduler_pb2.GetServerInfoResponse()} + ), ): - worker = _grpc_server_info_to_worker("grpc://10.0.0.1:8000", "sglang", response) + worker = await _probe_worker_replica(_job(), address="10.0.0.1:8000") + assert worker == { "url": "grpc://10.0.0.1:8000", - "worker_type": "prefill", + "worker_type": "regular", "connection_mode": "grpc", "runtime_type": "sglang", - "bootstrap_port": 8998, } + async def test_detects_vllm_grpc_worker(self): + with _fake_replica_transports( + http_handler=_grpc_server_http_handler, + grpc_channel=_FakeGrpcChannel( + {_VLLM_GET_SERVER_INFO: vllm_engine_pb2.GetServerInfoResponse()} + ), + ): + worker = await _probe_worker_replica(_job(), address="10.0.0.1:8000") + + assert worker == { + "url": "grpc://10.0.0.1:8000", + "worker_type": "regular", + "connection_mode": "grpc", + "runtime_type": "vllm", + } + + async def test_returns_none_when_no_probe_succeeds(self): + with _fake_replica_transports( + http_handler=lambda _: httpx.Response(200, json={"status": "starting"}), + grpc_channel=_FakeGrpcChannel(unknown_method_code=grpc.StatusCode.UNAVAILABLE), + ): + assert await _probe_worker_replica(_job(), address="10.0.0.1:8000") is None + + async def test_cancels_pending_probes_after_success(self): + grpc_channel = _FakeGrpcChannel(hang=True) + with _fake_replica_transports( + http_handler=lambda _: httpx.Response(200, json={"status": "ready"}), + grpc_channel=grpc_channel, + ): + worker = await _probe_worker_replica(_job(), address="10.0.0.1:8000") + + assert worker is not None + assert worker["connection_mode"] == "http" + assert sorted(grpc_channel.cancelled) == [_SGLANG_GET_SERVER_INFO, _VLLM_GET_SERVER_INFO] + + async def test_failing_probe_does_not_hide_successful_one( + self, caplog: pytest.LogCaptureFixture + ): + def http_handler(request: httpx.Request) -> httpx.Response: + raise RuntimeError("boom") + + job = _job() + with _fake_replica_transports( + http_handler=http_handler, + grpc_channel=_FakeGrpcChannel( + {_VLLM_GET_SERVER_INFO: vllm_engine_pb2.GetServerInfoResponse()} + ), + ): + worker = await _probe_worker_replica(job, address="10.0.0.1:8000") + + assert worker is not None + assert worker["runtime_type"] == "vllm" + assert f"Worker {fmt(job)}: probe failed unexpectedly" in caplog.text + + @pytest.mark.parametrize( + ["error", "message"], + [ + pytest.param( + SSHError("connection refused"), + "failed to connect: SSHError('connection refused')", + id="unreachable", + ), + pytest.param(RuntimeError("boom"), "unexpected error when probing", id="unexpected"), + ], + ) + async def test_returns_none_on_tunnel_error( + self, caplog: pytest.LogCaptureFixture, error: Exception, message: str + ): + caplog.set_level(level=logging.WARNING, logger=router_worker_sync.__name__) + job = _job() + with _fake_replica_transports(tunnel_error=error): + assert await _probe_worker_replica(job, address="10.0.0.1:8000") is None + assert f"Worker {fmt(job)}: {message}" in caplog.text + + +@pytest.mark.asyncio +@pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True) +class TestGetWorkerJobsWithAddresses: + async def test_returns_running_worker_jobs_with_addresses( + self, test_db, session: AsyncSession + ): + run, _, _ = await _create_router_service_run(session) + with_internal_ip = await create_job( + session=session, + run=run, + status=JobStatus.RUNNING, + replica_num=1, + replica_group_name="worker", + job_provisioning_data=get_job_provisioning_data( + internal_ip="10.0.0.1", hostname="203.0.113.1" + ), + ) + without_internal_ip = await create_job( + session=session, + run=run, + status=JobStatus.RUNNING, + replica_num=2, + replica_group_name="worker", + job_provisioning_data=get_job_provisioning_data( + internal_ip=None, hostname="worker.example" + ), + ) + # Skipped: not provisioned + await create_job( + session=session, + run=run, + status=JobStatus.RUNNING, + replica_num=3, + replica_group_name="worker", + ) + # Skipped: not running + await create_job( + session=session, + run=run, + status=JobStatus.TERMINATING, + replica_num=4, + replica_group_name="worker", + job_provisioning_data=get_job_provisioning_data(internal_ip="10.0.0.4"), + ) + await session.refresh(run, attribute_names=["jobs"]) + + result = _get_worker_jobs_with_addresses( + run.jobs, configuration=_service_configuration(), router_group_name="router" + ) + + assert [(job.id, address) for job, address in result] == [ + (with_internal_ip.id, "10.0.0.1:8000"), + (without_internal_ip.id, "worker.example:8000"), + ] -def _router_client(handler) -> AsyncClient: - return AsyncClient(transport=httpx.MockTransport(handler)) +class TestGetCurrentWorkerJobId: + def test_returns_job_id_from_label(self): + job_id = uuid.uuid4() + worker = {"id": "1", "url": "http://10.0.0.1:8000", "labels": {_JOB_ID_LABEL: str(job_id)}} + assert _get_current_worker_job_id(worker) == job_id + + @pytest.mark.parametrize( + "labels", + [ + pytest.param(None, id="no-labels"), + pytest.param({}, id="empty-labels"), + pytest.param({"team": "ml"}, id="other-labels"), + pytest.param({_JOB_ID_LABEL: ""}, id="empty-label"), + ], + ) + def test_returns_none_without_label(self, labels: Optional[dict[str, str]]): + assert ( + _get_current_worker_job_id(_router_entry("1", "http://10.0.0.1:8000", labels)) is None + ) -def _json_response(status_code: int, payload) -> httpx.Response: - return httpx.Response(status_code, content=json.dumps(payload).encode()) + def test_returns_none_on_unparsable_label(self, caplog: pytest.LogCaptureFixture): + worker = _router_entry("1", "http://10.0.0.1:8000", {_JOB_ID_LABEL: "not-a-uuid"}) + assert _get_current_worker_job_id(worker) is None + assert "Unparsable dstack job id worker label: 'not-a-uuid'" in caplog.text @pytest.mark.asyncio @@ -122,44 +473,89 @@ class TestGetRouterWorkers: """ async def test_returns_workers(self): - payload = {"workers": [{"id": "1", "url": "http://10.0.0.1:8000"}]} - async with _router_client(lambda _: _json_response(200, payload)) as client: - assert await _get_router_workers(client) == payload["workers"] + payload = { + "workers": [ + { + "id": "1", + "url": "http://10.0.0.1:8000", + "labels": {_JOB_ID_LABEL: "a7d1c5a4-5d53-4d7e-a4ad-2a6c4e1f3e0b"}, + "worker_type": "regular", + "is_healthy": True, + }, + # The router omits empty labels + {"id": "2", "url": "grpc://10.0.0.2:8000", "worker_type": "regular"}, + ], + "total": 2, + "stats": {"prefill_count": 0, "decode_count": 0, "regular_count": 2}, + } + async with _mock_client(lambda _: httpx.Response(200, json=payload)) as client: + assert await _get_router_workers(client, log_prefix="Router") == [ + { + "id": "1", + "url": "http://10.0.0.1:8000", + "labels": {_JOB_ID_LABEL: "a7d1c5a4-5d53-4d7e-a4ad-2a6c4e1f3e0b"}, + }, + {"id": "2", "url": "grpc://10.0.0.2:8000"}, + ] async def test_returns_empty_list_when_router_has_no_workers(self): - async with _router_client(lambda _: _json_response(200, {"workers": []})) as client: - assert await _get_router_workers(client) == [] - - async def test_returns_none_on_unexpected_status(self): - async with _router_client(lambda _: _json_response(500, {"workers": []})) as client: - assert await _get_router_workers(client) is None - - async def test_returns_none_on_unparsable_body(self): - async with _router_client(lambda _: httpx.Response(200, content=b"not json")) as client: - assert await _get_router_workers(client) is None + async with _mock_client(lambda _: httpx.Response(200, json={"workers": []})) as client: + assert await _get_router_workers(client, log_prefix="Router") == [] - async def test_returns_none_when_body_is_not_an_object(self): - async with _router_client(lambda _: _json_response(200, ["a"])) as client: - assert await _get_router_workers(client) is None + @pytest.mark.parametrize( + "response", + [ + pytest.param(httpx.Response(500, json={"workers": []}), id="unexpected-status"), + pytest.param(httpx.Response(200, content=b"not json"), id="unparsable"), + pytest.param(httpx.Response(200, json=["a"]), id="not-an-object"), + pytest.param(httpx.Response(200, json={}), id="no-workers"), + pytest.param(httpx.Response(200, json={"workers": "oops"}), id="not-an-array"), + pytest.param( + httpx.Response(200, json={"workers": [{"url": "http://10.0.0.1:8000"}]}), + id="worker-without-id", + ), + pytest.param( + httpx.Response( + 200, + json={"workers": [{"id": "1", "url": "http://10.0.0.1", "labels": {"a": 1}}]}, + ), + id="worker-with-invalid-labels", + ), + pytest.param( + httpx.Response( + 200, json={"workers": [{"id": str(i), "url": "u"} for i in range(8193)]} + ), + id="too-many-workers", + ), + ], + ) + async def test_returns_none_on_unusable_response(self, response: httpx.Response): + async with _mock_client(lambda _: response) as client: + assert await _get_router_workers(client, log_prefix="Router") is None - async def test_returns_none_when_workers_key_is_missing(self): - async with _router_client(lambda _: _json_response(200, {})) as client: - assert await _get_router_workers(client) is None + @pytest.mark.parametrize("with_content_length", [True, False]) + async def test_returns_none_when_response_too_large( + self, caplog: pytest.LogCaptureFixture, with_content_length: bool + ): + def handler(request: httpx.Request) -> httpx.Response: + if with_content_length: + return httpx.Response(200, content=b" " * (2 * 1024 * 1024 + 1)) + return httpx.Response(200, content=_chunks(b" " * 1024 * 1024, count=3)) - async def test_returns_none_when_workers_is_not_an_array(self): - async with _router_client(lambda _: _json_response(200, {"workers": "oops"})) as client: - assert await _get_router_workers(client) is None + async with _mock_client(handler) as client: + assert await _get_router_workers(client, log_prefix="Router") is None + assert "Router: GET /workers: response too large" in caplog.text async def test_returns_none_on_request_error(self, caplog: pytest.LogCaptureFixture): - def handler(_): + def handler(request: httpx.Request) -> httpx.Response: # What an SSH tunnel to a port nobody listens on yet produces. raise httpx.ReadError("") caplog.set_level(level=logging.DEBUG, logger=router_worker_sync.__name__) - async with _router_client(handler) as client: - assert await _get_router_workers(client) is None + async with _mock_client(handler) as client: + assert await _get_router_workers(client, log_prefix="Router") is None # A router that has not bound its port yet is expected, not an error. - assert "Router /workers not ready yet" in caplog.text + assert "Router: GET /workers: request failed" in caplog.text assert not [r for r in caplog.records if r.levelno > logging.DEBUG] @@ -180,6 +576,7 @@ class TestAddWorkerToRouter: "worker_type": "regular", "connection_mode": "http", "runtime_type": "sglang", + "labels": {_JOB_ID_LABEL: "a7d1c5a4-5d53-4d7e-a4ad-2a6c4e1f3e0b"}, }, id="http-regular", ), @@ -192,448 +589,360 @@ class TestAddWorkerToRouter: "bootstrap_port": 8998, "kv_connector": "NixlConnector", "kv_role": "kv_producer", + "labels": {_JOB_ID_LABEL: "a7d1c5a4-5d53-4d7e-a4ad-2a6c4e1f3e0b"}, }, id="grpc-prefill", ), ], ) - async def test_posts_worker_as_is(self, worker: _TargetWorker): + async def test_posts_worker_as_is( + self, caplog: pytest.LogCaptureFixture, worker: _TargetWorker + ): requests: list[httpx.Request] = [] def handler(request: httpx.Request) -> httpx.Response: requests.append(request) - return _json_response(202, {"status": "accepted"}) + return httpx.Response(202, json={"status": "accepted", "worker_id": "1"}) + + async with _mock_client(handler) as client: + await _add_worker_to_router(client, log_prefix="Router", worker=worker) - async with _router_client(handler) as client: - assert await _add_worker_to_router(client, worker) is True assert len(requests) == 1 assert requests[0].method == "POST" assert requests[0].url.path == "/workers" assert json.loads(requests[0].content) == worker + assert not [r for r in caplog.records if r.levelno > logging.DEBUG] - async def test_returns_false_when_not_accepted(self, caplog: pytest.LogCaptureFixture): - worker: _TargetWorker = { - "url": "http://10.0.0.1:8000", - "worker_type": "regular", - "connection_mode": "http", - "runtime_type": "sglang", - } - async with _router_client(lambda _: _json_response(202, {"status": "rejected"})) as client: - assert await _add_worker_to_router(client, worker) is False - assert "Unexpected add-worker response for http://10.0.0.1:8000" in caplog.text - + @pytest.mark.parametrize("code", ["WORKER_CREATE_IN_PROGRESS", "WORKER_ALREADY_EXISTS"]) + async def test_accepts_conflict_with_same_worker( + self, caplog: pytest.LogCaptureFixture, code: str + ): + # The router registers workers asynchronously, so a sync may add one again before + # the router lists it + response = httpx.Response(409, json={"error": "conflict", "code": code}) + async with _mock_client(lambda _: response) as client: + await _add_worker_to_router(client, log_prefix="Router", worker=_http_worker(0)) + assert not [r for r in caplog.records if r.levelno > logging.DEBUG] -@pytest.mark.asyncio -class TestProbeHttpWorker: - """ - The probe talks to the replica over the tunnel, but the `url` it reports is the address the - router dials. It must stay byte-identical to what the router echoes back in `/workers`, - otherwise `_get_workers_diff` re-registers every worker on each sync. - """ + @pytest.mark.parametrize( + ["response", "message"], + [ + pytest.param( + httpx.Response(202, json={"status": "rejected", "worker_id": "1"}), + "unexpected accepted status: rejected", + id="not-accepted", + ), + pytest.param( + httpx.Response(202, json={"status": "accepted"}), + "accepted response validation failed", + id="invalid-accepted", + ), + pytest.param( + httpx.Response(409, json={"error": "conflict", "code": "OTHER"}), + "unexpected conflict code: OTHER", + id="unexpected-conflict", + ), + pytest.param( + httpx.Response(409, json={}), + "conflict response validation failed", + id="invalid-conflict", + ), + pytest.param( + httpx.Response(500, content=b"boom"), + "unexpected status code: 500: b'boom'", + id="unexpected-status", + ), + ], + ) + async def test_warns_on_unexpected_response( + self, caplog: pytest.LogCaptureFixture, response: httpx.Response, message: str + ): + async with _mock_client(lambda _: response) as client: + await _add_worker_to_router(client, log_prefix="Router", worker=_http_worker(0)) + assert f"Router: POST /workers: http://{_worker_address(0)}: {message}" in caplog.text - async def test_regular_worker(self): - async with _router_client(lambda _: _json_response(200, {"status": "ready"})) as client: - assert await _probe_http_worker(client, address="10.0.0.1:8000") == { - "url": "http://10.0.0.1:8000", - "worker_type": "regular", - "connection_mode": "http", - "runtime_type": "sglang", - } - async def test_prefill_worker(self): - payload = { - "status": "ready", - "disaggregation_mode": "prefill", - "disaggregation_bootstrap_port": 8998, - } - async with _router_client(lambda _: _json_response(200, payload)) as client: - assert await _probe_http_worker(client, address="10.0.0.1:8000") == { - "url": "http://10.0.0.1:8000", - "worker_type": "prefill", - "connection_mode": "http", - "runtime_type": "sglang", - "bootstrap_port": 8998, - } - - async def test_decode_worker(self): - payload = {"status": "ready", "disaggregation_mode": "decode"} - async with _router_client(lambda _: _json_response(200, payload)) as client: - worker = await _probe_http_worker(client, address="10.0.0.1:8000") - assert worker is not None - assert worker["worker_type"] == "decode" +@pytest.mark.asyncio +class TestRemoveWorkerFromRouter: + async def test_deletes_worker_by_id(self, caplog: pytest.LogCaptureFixture): + requests: list[httpx.Request] = [] - async def test_returns_none_when_not_ready(self): - payload = {"status": "starting"} - async with _router_client(lambda _: _json_response(200, payload)) as client: - assert await _probe_http_worker(client, address="10.0.0.1:8000") is None + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(202, json={"status": "accepted", "worker_id": "w1"}) - async def test_returns_none_on_request_error(self, caplog: pytest.LogCaptureFixture): - def handler(_): - # A gRPC-only worker, or one that has not bound its port yet. - raise httpx.ReadError("") + async with _mock_client(handler) as client: + await _remove_worker_from_router(client, log_prefix="Router", worker_id="w1") - caplog.set_level(level=logging.DEBUG, logger=router_worker_sync.__name__) - async with _router_client(handler) as client: - assert await _probe_http_worker(client, address="10.0.0.1:8000") is None - # Repeats every sync until the worker is up, so it must not be logged as an error, - # see https://github.com/dstackai/dstack/issues/4300. - assert "Could not fetch server_info for worker http://10.0.0.1:8000" in caplog.text + assert [(r.method, r.url.path) for r in requests] == [("DELETE", "/workers/w1")] assert not [r for r in caplog.records if r.levelno > logging.DEBUG] - -@contextmanager -def _fake_vllm_grpc_proto(*, server_info=None, error: Optional[Exception] = None): - stub = MagicMock() - stub.GetServerInfo = AsyncMock(return_value=server_info, side_effect=error) - pb2 = MagicMock(GetServerInfoRequest=MagicMock(return_value="req")) - pb2_grpc = MagicMock(VllmEngineStub=MagicMock(return_value=stub)) - with ( - patch( - "dstack._internal.server.services.runs.router_worker_sync.vllm_engine_pb2", - pb2, - ), - patch( - "dstack._internal.server.services.runs.router_worker_sync.vllm_engine_pb2_grpc", - pb2_grpc, - ), - ): - yield - - -@contextmanager -def _fake_sglang_grpc_proto(*, server_info=None, error: Optional[Exception] = None): - stub = MagicMock() - stub.GetServerInfo = AsyncMock(return_value=server_info, side_effect=error) - pb2 = MagicMock(GetServerInfoRequest=MagicMock(return_value="req")) - pb2_grpc = MagicMock(SglangSchedulerStub=MagicMock(return_value=stub)) - with ( - patch( - "dstack._internal.server.services.runs.router_worker_sync.sglang_scheduler_pb2", - pb2, - ), - patch( - "dstack._internal.server.services.runs.router_worker_sync.sglang_scheduler_pb2_grpc", - pb2_grpc, - ), + async def test_logs_truncated_body_on_unexpected_status( + self, caplog: pytest.LogCaptureFixture ): - yield - - -def _rpc_error(code: grpc.StatusCode) -> grpc.aio.AioRpcError: - return grpc.aio.AioRpcError(code, grpc.aio.Metadata(), grpc.aio.Metadata(), details=code.name) + response = httpx.Response(500, content=b"x" * 2000) + async with _mock_client(lambda _: response) as client: + await _remove_worker_from_router(client, log_prefix="Router", worker_id="w1") + assert "Router: DELETE /workers/w1: unexpected status code: 500:" in caplog.text + assert f"b'{'x' * 1024}'... (2000 bytes total)" in caplog.text + assert "x" * 1025 not in caplog.text @pytest.mark.asyncio -class TestProbeGrpcWorker: - async def test_known_runtime_type(self): - server_info = MagicMock(kv_role="kv_producer", kv_connector="NixlConnector") - with _fake_vllm_grpc_proto(server_info=server_info): - worker = await _probe_grpc_worker( - MagicMock(), address="10.0.0.1:50051", runtime_type="vllm" - ) - assert worker == { - "url": "grpc://10.0.0.1:50051", - "worker_type": "prefill", - "connection_mode": "grpc", - "runtime_type": "vllm", - "kv_connector": "NixlConnector", - "kv_role": "kv_producer", - } +class TestProbeHttpWorker: + """ + The probe talks to the replica over the tunnel, but the `url` it reports is the address the + router dials. + """ - async def test_bootstrap_tries_sglang_first(self): - with ( - _fake_sglang_grpc_proto(server_info=MagicMock(server_args=MagicMock())), - patch( - "dstack._internal.server.services.runs.router_worker_sync.MessageToDict", - return_value={ + @pytest.mark.parametrize( + ["server_info", "expected"], + [ + pytest.param( + {"status": "ready"}, + {"worker_type": "regular"}, + id="regular", + ), + pytest.param( + { + "status": "ready", "disaggregation_mode": "prefill", "disaggregation_bootstrap_port": 8998, }, + {"worker_type": "prefill", "bootstrap_port": 8998}, + id="prefill", ), - ): - worker = await _probe_grpc_worker(MagicMock(), address="10.0.0.1:8000") + pytest.param( + {"status": "ready", "disaggregation_mode": "decode"}, + {"worker_type": "decode"}, + id="decode", + ), + ], + ) + async def test_reports_worker(self, server_info: dict, expected: dict): + async with _mock_client(lambda _: httpx.Response(200, json=server_info)) as client: + worker = await _probe_http_worker(client, log_prefix="Worker", address="10.0.0.1:8000") assert worker == { - "url": "grpc://10.0.0.1:8000", - "worker_type": "prefill", - "connection_mode": "grpc", + "url": "http://10.0.0.1:8000", + "connection_mode": "http", "runtime_type": "sglang", - "bootstrap_port": 8998, + **expected, } - async def test_bootstrap_falls_back_to_vllm(self): - # A vLLM worker does not implement the SGLang scheduler service. - with ( - _fake_sglang_grpc_proto(error=_rpc_error(grpc.StatusCode.UNIMPLEMENTED)), - _fake_vllm_grpc_proto( - server_info=MagicMock(kv_role="kv_consumer", kv_connector="NixlConnector") - ), - ): - worker = await _probe_grpc_worker(MagicMock(), address="10.0.0.1:8000") - assert worker is not None - assert worker["runtime_type"] == "vllm" - assert worker["worker_type"] == "decode" + async def test_returns_none_when_not_ready(self): + response = httpx.Response(200, json={"status": "starting"}) + async with _mock_client(lambda _: response) as client: + assert ( + await _probe_http_worker(client, log_prefix="Worker", address="10.0.0.1:8000") + is None + ) @pytest.mark.parametrize( - "code", + "body", [ - # No listener yet, or an SSH tunnel to a dead remote port, or an HTTP-only worker. - grpc.StatusCode.UNAVAILABLE, - grpc.StatusCode.DEADLINE_EXCEEDED, - grpc.StatusCode.UNIMPLEMENTED, + pytest.param(b"not json", id="not-json"), + # Detected as UTF-16 by its BOM, but not valid UTF-16 + pytest.param(b"\xff\xfe\xfa", id="not-text"), ], ) - async def test_returns_none_on_expected_error(self, code: grpc.StatusCode): - with _fake_vllm_grpc_proto(error=_rpc_error(code)): - worker = await _probe_grpc_worker( - MagicMock(), address="10.0.0.1:8000", runtime_type="vllm" + async def test_returns_none_on_unparsable_body( + self, caplog: pytest.LogCaptureFixture, body: bytes + ): + async with _mock_client(lambda _: httpx.Response(200, content=body)) as client: + assert ( + await _probe_http_worker(client, log_prefix="Worker", address="10.0.0.1:8000") + is None ) - assert worker is None - - async def test_reraises_unexpected_error(self): - error = _rpc_error(grpc.StatusCode.PERMISSION_DENIED) - with _fake_vllm_grpc_proto(error=error), pytest.raises(grpc.aio.AioRpcError) as exc_info: - await _probe_grpc_worker(MagicMock(), address="10.0.0.1:8000", runtime_type="vllm") - assert exc_info.value is error - - async def test_returns_none_when_no_runtime_type_matches(self): - with ( - _fake_sglang_grpc_proto(error=_rpc_error(grpc.StatusCode.UNAVAILABLE)), - _fake_vllm_grpc_proto(error=_rpc_error(grpc.StatusCode.UNAVAILABLE)), - ): - worker = await _probe_grpc_worker(MagicMock(), address="10.0.0.1:8000") - assert worker is None - - -_HTTP_WORKER = { - "url": "http://10.0.0.1:8000", - "worker_type": "regular", - "connection_mode": "http", - "runtime_type": "sglang", -} -_GRPC_WORKER = { - "url": "grpc://10.0.0.1:8000", - "worker_type": "prefill", - "connection_mode": "grpc", - "runtime_type": "vllm", -} - - -def _async_cm(enter_value=None, enter_error: Optional[Exception] = None) -> AsyncMock: - cm = AsyncMock() - cm.__aenter__ = AsyncMock(return_value=enter_value, side_effect=enter_error) - cm.__aexit__ = AsyncMock(return_value=False) - return cm + assert ( + "Worker: http://10.0.0.1:8000: GET /server_info: response parsing failed" + in caplog.text + ) + async def test_returns_none_on_unexpected_status(self, caplog: pytest.LogCaptureFixture): + async with _mock_client(lambda _: httpx.Response(404)) as client: + assert ( + await _probe_http_worker(client, log_prefix="Worker", address="10.0.0.1:8000") + is None + ) + assert "GET /server_info: unexpected status code: 404" in caplog.text -@contextmanager -def _fake_replica_transports( - *, - uds_path: Path = Path("/tmp/replica.sock"), - tunnel_error: Optional[Exception] = None, - http_worker=None, - grpc_worker=None, -): - """Patch the tunnel and both probes, leaving `_get_worker`'s own logic intact.""" - mocks = MagicMock() - with ( - patch( - "dstack._internal.server.services.runs.router_worker_sync.get_service_replica_tunnel", - return_value=_async_cm(uds_path, tunnel_error), - ) as mocks.tunnel, - patch( - "dstack._internal.server.services.runs.router_worker_sync" - ".get_service_replica_http_client_over_uds", - return_value=_async_cm(MagicMock()), - ) as mocks.http_client, - patch( - "dstack._internal.server.services.runs.router_worker_sync" - ".get_service_replica_grpc_channel_over_uds", - return_value=_async_cm(MagicMock()), - ) as mocks.grpc_channel, - patch( - "dstack._internal.server.services.runs.router_worker_sync._probe_http_worker", - new_callable=AsyncMock, - return_value=http_worker, - ) as mocks.http_probe, - patch( - "dstack._internal.server.services.runs.router_worker_sync._probe_grpc_worker", - new_callable=AsyncMock, - return_value=grpc_worker, - ) as mocks.grpc_probe, - ): - yield mocks + async def test_returns_none_on_request_error(self, caplog: pytest.LogCaptureFixture): + caplog.set_level(level=logging.DEBUG, logger=router_worker_sync.__name__) + async with _mock_client(_grpc_server_http_handler) as client: + assert ( + await _probe_http_worker(client, log_prefix="Worker", address="10.0.0.1:8000") + is None + ) + # Repeats every sync until the worker is up, or forever if it is a gRPC worker, so it + # must not be logged as an error, see https://github.com/dstackai/dstack/issues/4300. + assert "GET /server_info: request failed" in caplog.text + assert not [r for r in caplog.records if r.levelno > logging.DEBUG] @pytest.mark.asyncio -class TestGetWorker: - async def test_connection_mode_grpc_skips_http(self): - with _fake_replica_transports(grpc_worker=_GRPC_WORKER) as mocks: - worker = await _get_worker( - MagicMock(), - address="10.0.0.1:8000", - connection_mode="grpc", - ) - assert worker == _GRPC_WORKER - mocks.grpc_probe.assert_awaited_once() - mocks.http_probe.assert_not_awaited() - - async def test_connection_mode_http_skips_grpc(self): - with _fake_replica_transports(http_worker=_HTTP_WORKER) as mocks: - worker = await _get_worker( - MagicMock(), - address="10.0.0.1:8000", - connection_mode="http", - ) - assert worker == _HTTP_WORKER - mocks.http_probe.assert_awaited_once() - mocks.grpc_probe.assert_not_awaited() - - async def test_bootstrap_probes_http_first(self): - # An HTTP worker must not pay for two gRPC `GetServerInfo` timeouts first. - with _fake_replica_transports(http_worker=_HTTP_WORKER, grpc_worker=_GRPC_WORKER) as mocks: - worker = await _get_worker( - MagicMock(), - address="10.0.0.1:8000", - ) - assert worker == _HTTP_WORKER - mocks.http_probe.assert_awaited_once() - mocks.grpc_probe.assert_not_awaited() - - async def test_bootstrap_falls_back_to_grpc_over_one_tunnel(self): - job = MagicMock() - uds_path = Path("/tmp/replica.sock") - with _fake_replica_transports(uds_path=uds_path, grpc_worker=_GRPC_WORKER) as mocks: - worker = await _get_worker( - job, - address="10.0.0.1:8000", - ) - assert worker == _GRPC_WORKER - mocks.tunnel.assert_called_once_with(job) - mocks.http_client.assert_called_once_with(uds_path) - mocks.grpc_channel.assert_called_once_with(uds_path) - mocks.http_probe.assert_awaited_once() - mocks.grpc_probe.assert_awaited_once() - - async def test_returns_none_when_no_mode_reports_ready(self): - with _fake_replica_transports() as mocks: - worker = await _get_worker( - MagicMock(), - address="10.0.0.1:8000", - ) - assert worker is None - mocks.http_probe.assert_awaited_once() - mocks.grpc_probe.assert_awaited_once() +class TestProbeSglangGrpcWorker: + @pytest.mark.parametrize( + ["server_args", "expected"], + [ + pytest.param( + {"disaggregation_mode": "null"}, + {"worker_type": "regular"}, + id="regular", + ), + pytest.param( + {"disaggregation_mode": "prefill", "disaggregation_bootstrap_port": 8998}, + {"worker_type": "prefill", "bootstrap_port": 8998}, + id="prefill", + ), + pytest.param( + {"disaggregation_mode": "decode"}, + {"worker_type": "decode"}, + id="decode", + ), + ], + ) + async def test_reports_worker(self, server_args: dict, expected: dict): + response = sglang_scheduler_pb2.GetServerInfoResponse(server_args=server_args) + channel = _FakeGrpcChannel({_SGLANG_GET_SERVER_INFO: response}) - async def test_unreachable_worker_is_skipped_not_raised( - self, caplog: pytest.LogCaptureFixture - ): - # One dead replica must not abort the sync for the healthy ones. - caplog.set_level(level=logging.WARNING, logger=router_worker_sync.__name__) - ssh_error = SSHError("connection refused") - with _fake_replica_transports(tunnel_error=ssh_error) as mocks: - worker = await _get_worker( - MagicMock(), - address="10.0.0.1:8000", - ) - assert worker is None - mocks.http_probe.assert_not_awaited() - mocks.grpc_probe.assert_not_awaited() - assert f"failed to connect to worker replica: {ssh_error!r}" in caplog.text - - async def test_unexpected_tunnel_error_propagates(self): - # Only `SSHError` means "unreachable"; anything else is a bug and must not be swallowed. - with _fake_replica_transports(tunnel_error=RuntimeError("boom")): - with pytest.raises(RuntimeError, match="boom"): - await _get_worker( - MagicMock(), - address="10.0.0.1:8000", - ) - - -class TestGetWorkersDiff: - def test_adds_missing_and_removes_extra_workers(self): - kept: _TargetWorker = { - "url": "http://10.0.0.1:8000", - "worker_type": "regular", - "connection_mode": "http", - "runtime_type": "sglang", - } - added: _TargetWorker = {**kept, "url": "http://10.0.0.2:8000"} - current = [ - # The router may echo the URL back with a trailing slash - {"id": "1", "url": "http://10.0.0.1:8000/"}, - {"id": "2", "url": "http://10.0.0.3:8000"}, - ] - diff = _get_workers_diff([kept, added], current) - assert diff.to_add == [added] - assert diff.to_remove == {"http://10.0.0.3:8000": "2"} - assert not diff.is_empty() + worker = await _probe_sglang_grpc_worker( + channel, log_prefix="Worker", address="10.0.0.1:8000" + ) - def test_is_empty_when_router_is_in_sync(self): - worker: _TargetWorker = { - "url": "http://10.0.0.1:8000", - "worker_type": "regular", - "connection_mode": "http", + expected = { + "url": "grpc://10.0.0.1:8000", + "connection_mode": "grpc", "runtime_type": "sglang", + **expected, } - diff = _get_workers_diff([worker], [{"id": "1", "url": "http://10.0.0.1:8000"}]) - assert diff.is_empty() + # Compared as sent to the router: `server_args` is a `Struct`, which has only float + # numbers, and `8998.0 == 8998`, but the router expects an integer port + assert json.dumps(worker, sort_keys=True) == json.dumps(expected, sort_keys=True) - def test_removed_worker_without_id(self): - diff = _get_workers_diff([], [{"url": "http://10.0.0.1:8000"}]) - assert diff.to_remove == {"http://10.0.0.1:8000": None} - - -class _FakeRouter: - """An in-memory router `/workers` API that counts the connections made to it.""" + @pytest.mark.parametrize( + "code", + [ + # Not listening yet, or an HTTP worker + grpc.StatusCode.UNAVAILABLE, + grpc.StatusCode.DEADLINE_EXCEEDED, + # A vLLM worker + grpc.StatusCode.UNIMPLEMENTED, + grpc.StatusCode.CANCELLED, + ], + ) + async def test_returns_none_on_expected_error( + self, caplog: pytest.LogCaptureFixture, code: grpc.StatusCode + ): + caplog.set_level(level=logging.DEBUG, logger=router_worker_sync.__name__) + channel = _FakeGrpcChannel({_SGLANG_GET_SERVER_INFO: _rpc_error(code)}) + assert ( + await _probe_sglang_grpc_worker(channel, log_prefix="Worker", address="10.0.0.1:8000") + is None + ) + assert not [r for r in caplog.records if r.levelno > logging.DEBUG] - def __init__(self, workers: Optional[list[dict]] = None, *, unreachable: bool = False): - self.workers = list(workers or []) - self.unreachable = unreachable - self.connections = 0 - self.added: list[dict] = [] - self.removed_ids: list[str] = [] + async def test_warns_on_unexpected_error(self, caplog: pytest.LogCaptureFixture): + channel = _FakeGrpcChannel( + {_SGLANG_GET_SERVER_INFO: _rpc_error(grpc.StatusCode.PERMISSION_DENIED)} + ) + assert ( + await _probe_sglang_grpc_worker(channel, log_prefix="Worker", address="10.0.0.1:8000") + is None + ) + assert ( + "Worker: grpc://10.0.0.1:8000: sglang probe failed: StatusCode.PERMISSION_DENIED" + in caplog.text + ) - def handle(self, request: httpx.Request) -> httpx.Response: - if request.method == "GET" and request.url.path == "/workers": - return _json_response(200, {"workers": self.workers}) - if request.method == "POST" and request.url.path == "/workers": - worker = json.loads(request.content) - self.added.append(worker) - self.workers.append({"id": f"id-{len(self.added)}", **worker}) - return _json_response(202, {"status": "accepted"}) - if request.method == "DELETE" and request.url.path.startswith("/workers/"): - worker_id = request.url.path.removeprefix("/workers/") - self.removed_ids.append(worker_id) - self.workers = [w for w in self.workers if w["id"] != worker_id] - return _json_response(202, {"status": "accepted"}) - return httpx.Response(404) +@pytest.mark.asyncio +class TestProbeVllmGrpcWorker: + @pytest.mark.parametrize( + ["kv_role", "kv_connector", "expected"], + [ + pytest.param("", "", {"worker_type": "regular"}, id="regular"), + pytest.param( + "kv_both", + "NixlConnector", + {"worker_type": "regular", "kv_role": "kv_both", "kv_connector": "NixlConnector"}, + id="kv-both", + ), + pytest.param( + "kv_producer", + "NixlConnector", + { + "worker_type": "prefill", + "kv_role": "kv_producer", + "kv_connector": "NixlConnector", + }, + id="prefill", + ), + pytest.param( + "kv_consumer", + "NixlConnector", + { + "worker_type": "decode", + "kv_role": "kv_consumer", + "kv_connector": "NixlConnector", + }, + id="decode", + ), + ], + ) + async def test_reports_worker(self, kv_role: str, kv_connector: str, expected: dict): + response = vllm_engine_pb2.GetServerInfoResponse( + kv_role=kv_role, kv_connector=kv_connector + ) + channel = _FakeGrpcChannel({_VLLM_GET_SERVER_INFO: response}) -@contextmanager -def _fake_router_replicas(routers: dict[uuid.UUID, _FakeRouter]): - """Route each router job's client to its fake router, keyed by the job id.""" + worker = await _probe_vllm_grpc_worker( + channel, log_prefix="Worker", address="10.0.0.1:8000" + ) - @asynccontextmanager - async def get_service_replica_client(job: JobModel): - router = routers[job.id] - router.connections += 1 - if router.unreachable: - raise SSHError("connection refused") - async with AsyncClient(transport=httpx.MockTransport(router.handle)) as client: - yield client + assert worker == { + "url": "grpc://10.0.0.1:8000", + "connection_mode": "grpc", + "runtime_type": "vllm", + **expected, + } - with patch( - "dstack._internal.server.services.runs.router_worker_sync.get_service_replica_client", - get_service_replica_client, + @pytest.mark.parametrize( + "code", + [ + # Not listening yet, or an HTTP worker + grpc.StatusCode.UNAVAILABLE, + grpc.StatusCode.DEADLINE_EXCEEDED, + # An SGLang worker + grpc.StatusCode.UNIMPLEMENTED, + grpc.StatusCode.CANCELLED, + ], + ) + async def test_returns_none_on_expected_error( + self, caplog: pytest.LogCaptureFixture, code: grpc.StatusCode ): - yield + caplog.set_level(level=logging.DEBUG, logger=router_worker_sync.__name__) + channel = _FakeGrpcChannel({_VLLM_GET_SERVER_INFO: _rpc_error(code)}) + assert ( + await _probe_vllm_grpc_worker(channel, log_prefix="Worker", address="10.0.0.1:8000") + is None + ) + assert not [r for r in caplog.records if r.levelno > logging.DEBUG] + async def test_warns_on_unexpected_error(self, caplog: pytest.LogCaptureFixture): + channel = _FakeGrpcChannel( + {_VLLM_GET_SERVER_INFO: _rpc_error(grpc.StatusCode.PERMISSION_DENIED)} + ) + assert ( + await _probe_vllm_grpc_worker(channel, log_prefix="Worker", address="10.0.0.1:8000") + is None + ) + assert ( + "Worker: grpc://10.0.0.1:8000: vllm probe failed: StatusCode.PERMISSION_DENIED" + in caplog.text + ) -async def _create_router_service_run(session: AsyncSession, router_count: int) -> RunModel: - project = await create_project(session=session) - user = await create_user(session=session) - repo = await create_repo(session=session, project_id=project.id) + +def _service_configuration() -> ServiceConfiguration: conf = parse_run_configuration( { "type": "service", @@ -650,16 +959,35 @@ async def _create_router_service_run(session: AsyncSession, router_count: int) - ], } ) + assert isinstance(conf, ServiceConfiguration) + return conf + + +async def _create_router_service_run( + session: AsyncSession, *, router_count: int = 1, worker_count: int = 0 +) -> tuple[RunModel, list[JobModel], list[JobModel]]: + """ + Creates a running service run with running router and worker jobs. Worker `i` listens on + `_worker_address(i)`. + + Returns: + The run, its router jobs, and its worker jobs. + """ + project = await create_project(session=session) + user = await create_user(session=session) + repo = await create_repo(session=session, project_id=project.id) run = await create_run( session=session, project=project, repo=repo, user=user, status=RunStatus.RUNNING, - run_spec=get_run_spec(repo_id=repo.name, run_name="test-run", configuration=conf), + run_spec=get_run_spec( + repo_id=repo.name, run_name="test-run", configuration=_service_configuration() + ), ) # More than one running router means a rolling deployment is replacing the router - for replica_num in range(router_count): + router_jobs = [ await create_job( session=session, run=run, @@ -667,102 +995,216 @@ async def _create_router_service_run(session: AsyncSession, router_count: int) - replica_num=replica_num, replica_group_name="router", ) + for replica_num in range(router_count) + ] + worker_jobs = [ + await create_job( + session=session, + run=run, + status=JobStatus.RUNNING, + replica_num=router_count + i, + replica_group_name="worker", + job_provisioning_data=get_job_provisioning_data(internal_ip=f"10.0.0.{i + 1}"), + ) + for i in range(worker_count) + ] await session.refresh(run, attribute_names=["jobs"]) - return run + return run, router_jobs, worker_jobs -_VLLM_WORKER: _TargetWorker = { - "url": "grpc://10.0.0.1:8000", - "worker_type": "regular", - "connection_mode": "grpc", - "runtime_type": "vllm", -} +def _worker_address(index: int) -> str: + return f"10.0.0.{index + 1}:8000" -def _patch_build_target_workers(**kwargs): - return patch( - "dstack._internal.server.services.runs.router_worker_sync._build_target_workers", - new_callable=AsyncMock, - **kwargs, - ) +def _http_worker(index: int) -> _TargetWorker: + return { + "url": f"http://{_worker_address(index)}", + "worker_type": "regular", + "connection_mode": "http", + "runtime_type": "sglang", + } -@pytest.mark.asyncio -@pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True) -class TestSyncRouterWorkersForRunModel: - async def test_syncs_replacement_router_without_touching_up_to_date_one( - self, test_db, session: AsyncSession +def _job_id_label(job: JobModel) -> dict[str, str]: + return {_JOB_ID_LABEL: str(job.id)} + + +def _router_entry(worker_id: str, url: str, labels: Optional[dict[str, str]] = None) -> dict: + """A worker as listed by the router `GET /workers`, with the fields the sync reads""" + entry: dict = {"id": worker_id, "url": url} + if labels is not None: + entry["labels"] = labels + return entry + + +class _FakeRouter: + """An in-memory router `/workers` API that counts the connections made to it""" + + def __init__( + self, + workers: Optional[list[dict]] = None, + *, + connect_error: Optional[Exception] = None, + update_error: Optional[Exception] = None, ): - run = await _create_router_service_run(session, router_count=2) - old_router = _FakeRouter([{"id": "1", **_VLLM_WORKER}]) - new_router = _FakeRouter() - routers = {run.jobs[0].id: old_router, run.jobs[1].id: new_router} + self.workers = list(workers or []) + self.connect_error = connect_error + self.update_error = update_error + self.connections = 0 + self.added: list[dict] = [] + self.removed_ids: list[str] = [] - with ( - _fake_router_replicas(routers), - _patch_build_target_workers(return_value=[_VLLM_WORKER]) as build_mock, - ): - await sync_router_workers_for_run_model(run) + def handle(self, request: httpx.Request) -> httpx.Response: + if request.method == "GET" and request.url.path == "/workers": + return httpx.Response(200, json={"workers": self.workers, "total": len(self.workers)}) + if self.update_error is not None: + raise self.update_error + if request.method == "POST" and request.url.path == "/workers": + worker = json.loads(request.content) + worker_id = f"added-{len(self.added)}" + self.added.append(worker) + self.workers.append({"id": worker_id, **worker}) + return httpx.Response(202, json={"status": "accepted", "worker_id": worker_id}) + if request.method == "DELETE" and request.url.path.startswith("/workers/"): + worker_id = request.url.path.removeprefix("/workers/") + self.removed_ids.append(worker_id) + self.workers = [w for w in self.workers if w["id"] != worker_id] + return httpx.Response(202, json={"status": "accepted", "worker_id": worker_id}) + return httpx.Response(404) - assert new_router.added == [_VLLM_WORKER] - assert old_router.added == [] - assert old_router.removed_ids == [] - # Read once; reconnected only for the router that needed an update - assert old_router.connections == 1 - assert new_router.connections == 2 - # Workers are probed once for all routers, with hints taken from all of them: the - # replacement router's empty list alone would mean "probe everything" - build_mock.assert_awaited_once() - assert build_mock.await_args is not None - assert build_mock.await_args.kwargs["connection_mode"] == "grpc" - assert build_mock.await_args.kwargs["runtime_type"] == "vllm" - - async def test_unreachable_router_does_not_block_others( - self, test_db, session: AsyncSession, caplog: pytest.LogCaptureFixture + +@contextmanager +def _fake_router_replicas(routers: dict[uuid.UUID, _FakeRouter]) -> Iterator[None]: + """Route each router job's client to its fake router, keyed by the job id""" + + @asynccontextmanager + async def get_service_replica_client(job: JobModel) -> AsyncIterator[AsyncClient]: + router = routers[job.id] + router.connections += 1 + if router.connect_error is not None: + raise router.connect_error + async with _mock_client(router.handle) as client: + yield client + + with patch.object( + router_worker_sync, "get_service_replica_client", get_service_replica_client ): - caplog.set_level(level=logging.WARNING, logger=router_worker_sync.__name__) - run = await _create_router_service_run(session, router_count=2) - unreachable_router = _FakeRouter(unreachable=True) - reachable_router = _FakeRouter() - routers = {run.jobs[0].id: unreachable_router, run.jobs[1].id: reachable_router} + yield - with ( - _fake_router_replicas(routers), - _patch_build_target_workers(return_value=[_VLLM_WORKER]), - ): - await sync_router_workers_for_run_model(run) - assert reachable_router.added == [_VLLM_WORKER] - assert unreachable_router.connections == 1 - assert f"{fmt(run.jobs[0])}: failed to sync workers with router" in caplog.text +@contextmanager +def _fake_worker_probes(workers: Mapping[str, _TargetWorker]) -> Iterator[AsyncMock]: + """Patch probing so that the worker replica at each address reports the given worker""" + + async def probe(job: JobModel, *, address: str) -> Optional[_TargetWorker]: + # A copy, as the sync labels the reported worker in place + return copy.deepcopy(workers.get(address)) - async def test_skips_probing_workers_when_no_router_is_reachable( - self, test_db, session: AsyncSession + with patch.object( + router_worker_sync, "_probe_worker_replica", new_callable=AsyncMock, side_effect=probe + ) as probe_mock: + yield probe_mock + + +def _job() -> JobModel: + return JobModel(id=uuid.uuid4(), job_name="test-run-0-1") + + +@contextmanager +def _fake_replica_transports( + *, + http_handler: Optional[Callable[[httpx.Request], httpx.Response]] = None, + grpc_channel: Optional["_FakeGrpcChannel"] = None, + tunnel_error: Optional[Exception] = None, +) -> Iterator[None]: + """Patch the tunnel to a worker replica and the HTTP client and gRPC channel over it""" + + @asynccontextmanager + async def get_tunnel(job: JobModel) -> AsyncIterator[Path]: + if tunnel_error is not None: + raise tunnel_error + yield Path("replica.sock") + + @asynccontextmanager + async def get_http_client(uds_path: Path) -> AsyncIterator[AsyncClient]: + assert http_handler is not None + async with _mock_client(http_handler) as client: + yield client + + @asynccontextmanager + async def get_grpc_channel(uds_path: Path) -> AsyncIterator["_FakeGrpcChannel"]: + assert grpc_channel is not None + yield grpc_channel + + with ( + patch.object(router_worker_sync, "get_service_replica_tunnel", get_tunnel), + patch.object( + router_worker_sync, "get_service_replica_http_client_over_uds", get_http_client + ), + patch.object( + router_worker_sync, "get_service_replica_grpc_channel_over_uds", get_grpc_channel + ), ): - run = await _create_router_service_run(session, router_count=1) - routers = {run.jobs[0].id: _FakeRouter(unreachable=True)} + yield - with _fake_router_replicas(routers), _patch_build_target_workers() as build_mock: - await sync_router_workers_for_run_model(run) - build_mock.assert_not_awaited() +def _grpc_server_http_handler(request: httpx.Request) -> httpx.Response: + # What an HTTP/1.1 request to an HTTP/2-only gRPC server produces + raise httpx.RemoteProtocolError("illegal request line") - async def test_rereads_router_workers_before_updating(self, test_db, session: AsyncSession): - run = await _create_router_service_run(session, router_count=1) - router = _FakeRouter([{"id": "1", **_VLLM_WORKER}]) - routers = {run.jobs[0].id: router} - def reregister_worker(*args, **kwargs): - # Worker ids are assigned by the router. Here, the worker got a new one while the - # workers were being probed, so the id read before probing is stale. - router.workers = [{"id": "2", **_VLLM_WORKER}] - return [] +class _FakeGrpcChannel: + """ + A channel to a gRPC worker replica that serves the given unary methods. The probes create + real stubs on it, so a stub of a service the worker doesn't serve fails as UNIMPLEMENTED, + as against a real worker. + """ + + def __init__( + self, + responses: Optional[Mapping[str, Union[Message, grpc.aio.AioRpcError]]] = None, + *, + unknown_method_code: grpc.StatusCode = grpc.StatusCode.UNIMPLEMENTED, + hang: bool = False, + ): + self.responses = responses or {} + self.unknown_method_code = unknown_method_code + self.hang = hang + """Whether calls never complete, only cancellation ends them""" + self.cancelled: list[str] = [] + """Methods whose calls were cancelled""" + + def unary_unary(self, method: str, request_serializer, response_deserializer, **kwargs): + async def call(request, **kwargs): + if self.hang: + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + self.cancelled.append(method) + raise + response = self.responses.get(method) + if response is None: + raise _rpc_error(self.unknown_method_code) + if isinstance(response, grpc.aio.AioRpcError): + raise response + # Through the wire format, as the real stub would receive it + return response_deserializer(response.SerializeToString()) + + return call + + def unary_stream(self, *args, **kwargs): + # Stubs create all their methods, but the probes only call unary ones + return None + + +def _rpc_error(code: grpc.StatusCode) -> grpc.aio.AioRpcError: + return grpc.aio.AioRpcError(code, grpc.aio.Metadata(), grpc.aio.Metadata(), details=code.name) + + +def _mock_client(handler: Callable[[httpx.Request], httpx.Response]) -> AsyncClient: + return AsyncClient(transport=httpx.MockTransport(handler)) - with ( - _fake_router_replicas(routers), - _patch_build_target_workers(side_effect=reregister_worker), - ): - await sync_router_workers_for_run_model(run) - assert router.removed_ids == ["2"] - assert router.workers == [] +async def _chunks(chunk: bytes, *, count: int) -> AsyncIterator[bytes]: + for _ in range(count): + yield chunk diff --git a/src/tests/_internal/server/utils/test_common.py b/src/tests/_internal/server/utils/test_common.py index 175fb042b8..71b73ec3cd 100644 --- a/src/tests/_internal/server/utils/test_common.py +++ b/src/tests/_internal/server/utils/test_common.py @@ -1,6 +1,11 @@ +import asyncio +import inspect + import pytest from dstack._internal.server.utils.common import ( + gather_async, + gather_map_async, join_byte_stream_checked, ) @@ -33,3 +38,116 @@ def generator(stream): raise RuntimeError("Stream end reached, but next value was requested") assert join_byte_stream_checked(generator(stream), max_size) is None + + +class TestGatherMapAsync: + @pytest.mark.asyncio + async def test_returns_items_with_results_in_order(self): + async def double(x: int) -> int: + return x * 2 + + assert await gather_map_async([3, 1, 2], double) == [(3, 6), (1, 2), (2, 4)] + + @pytest.mark.asyncio + async def test_returns_exceptions(self): + error = ValueError("odd") + + async def even_only(x: int) -> int: + if x % 2: + raise error + return x + + result = await gather_map_async([1, 2], even_only, return_exceptions=True) + assert result == [(1, error), (2, 2)] + + @pytest.mark.asyncio + async def test_limits_concurrency(self): + tracker = _ConcurrencyTracker() + items = list(range(5)) + result = await gather_map_async(items, tracker.call, max_concurrency=2) + assert result == [(x, x) for x in items] + assert tracker.peak == 2 + + +class TestGatherAsync: + @pytest.mark.asyncio + async def test_returns_results_in_order(self): + async def add(x: int, *, y: int) -> int: + return x + y + + assert await gather_async([add(1, y=2), add(3, y=4)]) == [3, 7] + + @pytest.mark.asyncio + async def test_returns_exceptions(self): + error = ValueError("odd") + + async def even_only(x: int) -> int: + if x % 2: + raise error + return x + + result = await gather_async([even_only(1), even_only(2)], return_exceptions=True) + assert result == [error, 2] + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ["max_concurrency", "expected_peak"], + [ + [None, 5], + [1, 1], + [2, 2], + [10, 5], + ], + ) + async def test_limits_concurrency(self, max_concurrency, expected_peak): + tracker = _ConcurrencyTracker() + result = await gather_async( + [tracker.call(x) for x in range(5)], max_concurrency=max_concurrency + ) + assert result == list(range(5)) + assert tracker.peak == expected_peak + + @pytest.mark.asyncio + @pytest.mark.parametrize("max_concurrency", [None, 1, 2]) + async def test_cancels_remaining_calls_on_error(self, max_concurrency): + tracker = _ConcurrencyTracker() + coros = [ + tracker.call(0, fail=True), + tracker.call(1, block=True), + tracker.call(2, block=True), + ] + with pytest.raises(ValueError, match="0"): + await gather_async(coros, max_concurrency=max_concurrency) + assert tracker.running == 0 + assert all(inspect.getcoroutinestate(c) == inspect.CORO_CLOSED for c in coros) + + @pytest.mark.asyncio + @pytest.mark.parametrize("max_concurrency", [0, -1]) + async def test_rejects_max_concurrency_below_one(self, max_concurrency): + tracker = _ConcurrencyTracker() + coro = tracker.call(1) + with pytest.raises(ValueError, match="max_concurrency"): + await gather_async([coro], max_concurrency=max_concurrency) + assert inspect.getcoroutinestate(coro) == inspect.CORO_CLOSED + + +class _ConcurrencyTracker: + def __init__(self) -> None: + self.running = 0 + self.peak = 0 + + async def call(self, x: int, *, fail: bool = False, block: bool = False) -> int: + self.running += 1 + self.peak = max(self.peak, self.running) + try: + # Yield to the event loop so other calls can start if the limit allows + for _ in range(3): + await asyncio.sleep(0) + if fail: + raise ValueError(x) + if block: + # Never set, only cancellation ends the call + await asyncio.Event().wait() + return x + finally: + self.running -= 1