From 0e6af06395169ffaf3028935689c0c6eab404627 Mon Sep 17 00:00:00 2001 From: Dmitry Meyer Date: Mon, 28 Sep 2026 13:34:51 +0000 Subject: [PATCH] Sync workers with every running router During a rolling deployment, the replacement router runs alongside the old one, but the worker sync only picked the first running router job, which is always the old one. The replacement had no workers until the old router was scaled down and could not serve requests routed to it. The sync now reads the workers of every running router, probes the workers once, and updates each router that is out of sync, so an unreachable router no longer blocks the others. Router tunnels are no longer held while workers are probed, which may take minutes: each router is read first, and connected to again only if it needs an update, re-reading its workers before applying the changes. Also: * `_Worker` is renamed back to `_TargetWorker`, its fields are typed after SMG's `WorkerSpec`, and `_add_worker_to_router()` sends it as the `POST /workers` body as is. Fixes: https://github.com/dstackai/dstack/issues/4315 Co-Authored-By: Claude Opus 5.5 (1M context) --- .../services/runs/router_worker_sync.py | 264 ++++++++------- .../services/runs/test_router_worker_sync.py | 300 +++++++++++++++++- 2 files changed, 441 insertions(+), 123 deletions(-) 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 100597dc0..d1e8a45c1 100644 --- a/src/dstack/_internal/server/services/runs/router_worker_sync.py +++ b/src/dstack/_internal/server/services/runs/router_worker_sync.py @@ -1,6 +1,7 @@ """Reconcile SGLang router /workers with dstack's ready worker replicas (async, SSH-tunneled).""" import json +from dataclasses import dataclass from typing import Any, List, Literal, Optional, TypedDict from urllib.parse import urlsplit, urlunsplit @@ -98,20 +99,30 @@ async def _request_json_limited( return None -class _Worker(TypedDict): +# 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): url: str - worker_type: str - bootstrap_port: NotRequired[Optional[int]] - connection_mode: NotRequired[str] - runtime_type: NotRequired[str] + worker_type: "_WorkerType" + connection_mode: "_ConnectionMode" + runtime_type: "_RuntimeType" + bootstrap_port: NotRequired[int] kv_connector: NotRequired[str] kv_role: NotRequired[str] -_ConnectionMode = Literal["grpc", "http"] +# https://github.com/smg-project/smg/blob/3be823a700fabaff3add8a390cf78f163479d686/crates/protocols/src/worker.rs#L29 +# Only types we support are included +_WorkerType = Literal["regular", "prefill", "decode"] + +# 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") @@ -122,20 +133,17 @@ def run_model_has_sglang_router_replica_group(run_model: RunModel) -> bool: return run_spec_has_sglang_router_replica_group(run_spec) -def _get_router_job(run_model: RunModel, router_group: ReplicaGroup) -> Optional[JobModel]: +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" - router_jobs = [ + # 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 ] - if not router_jobs: - return None - # Router replica group is currently validated to have count=1, so we assume a single active - # router job here. When we support multiple router replicas for HA, this should be updated - # to handle syncing across all active router jobs. - return router_jobs[0] def _normalize_worker_url(url: str) -> str: @@ -202,7 +210,7 @@ async def _get_router_workers(client: AsyncClient) -> Optional[List[dict]]: 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 - # _update_workers_in_router_replica) and decide what to do with a partially valid list + # _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", @@ -217,36 +225,16 @@ async def _get_router_workers(client: AsyncClient) -> Optional[List[dict]]: return None -async def _add_worker_to_router( - client: AsyncClient, - url: str, - worker_type: str = "regular", - bootstrap_port: Optional[int] = None, - *, - connection_mode: Optional[str] = None, - runtime_type: Optional[str] = None, - kv_connector: Optional[str] = None, - kv_role: Optional[str] = None, -) -> bool: +async def _add_worker_to_router(client: AsyncClient, worker: _TargetWorker) -> bool: + url = worker["url"] try: - payload: dict = {"url": url, "worker_type": worker_type} - if bootstrap_port is not None: - payload["bootstrap_port"] = bootstrap_port - if connection_mode is not None: - payload["connection_mode"] = connection_mode - if runtime_type is not None: - payload["runtime_type"] = runtime_type - if kv_connector is not None: - payload["kv_connector"] = kv_connector - if kv_role is not None: - payload["kv_role"] = kv_role body = await _request_json_limited( client, "POST", f"{_HTTP_BASE_URL}/workers", max_response_bytes=_MAX_WORKERS_COMMAND_ACK_BYTES, ok_statuses={202}, - json_body=payload, + json_body=dict(worker), ) added = isinstance(body, dict) and body.get("status") == "accepted" if not added: @@ -281,53 +269,60 @@ async def _remove_worker_from_router_by_id( return False -async def _update_workers_in_router_replica( - client: AsyncClient, - target_workers: List[_Worker], - *, - current_workers: List[dict], -) -> None: - current_urls: set[str] = set() - current_ids_by_norm_url: dict[str, str] = {} +@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""" + + def is_empty(self) -> bool: + return not self.to_add and not self.to_remove + + +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) - current_urls.add(norm_u) wid = w.get("id") if isinstance(wid, str) and wid: current_ids_by_norm_url[norm_u] = wid - target_by_norm = {_normalize_worker_url(t["url"]): t for t in target_workers} - target_urls = set(target_by_norm.keys()) - to_add = sorted(target_urls - current_urls) - to_remove = sorted(current_urls - target_urls) - for norm_url in to_add: - tw = target_by_norm[norm_url] - ok = await _add_worker_to_router( - client, - tw["url"], - tw["worker_type"], - tw.get("bootstrap_port"), - connection_mode=tw.get("connection_mode"), - runtime_type=tw.get("runtime_type"), - kv_connector=tw.get("kv_connector"), - kv_role=tw.get("kv_role"), - ) - if not ok: - logger.debug("Failed to add worker %s, continuing with others", tw["url"]) - for url in to_remove: - wid = current_ids_by_norm_url.get(url) - if not wid: - logger.error("No worker id found for url %s", url) - ok = False else: - ok = await _remove_worker_from_router_by_id(client, wid, worker_url=url) - if not ok: - logger.debug("Failed to remove worker %s, continuing with others", url) + 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}, + ) -def _vllm_kv_role_to_worker_type(kv_role: str) -> str: +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 + ) + + +def _vllm_kv_role_to_worker_type(kv_role: str) -> _WorkerType: if kv_role == "kv_producer": return "prefill" if kv_role == "kv_consumer": @@ -344,7 +339,7 @@ def _is_expected_grpc_error(error: grpc.aio.AioRpcError) -> bool: ) -async def _probe_http_worker(client: AsyncClient, *, address: str) -> Optional[_Worker]: +async def _probe_http_worker(client: AsyncClient, *, address: str) -> Optional[_TargetWorker]: # The request goes over the tunnel, `worker_url` is the address the router itself dials. worker_url = f"http://{address}" try: @@ -361,7 +356,7 @@ async def _probe_http_worker(client: AsyncClient, *, address: str) -> Optional[_ mode = data.get("disaggregation_mode", "") if mode == "prefill": bootstrap_port = data.get("disaggregation_bootstrap_port") - worker: _Worker = { + worker: _TargetWorker = { "url": worker_url, "worker_type": "prefill", "connection_mode": "http", @@ -407,11 +402,11 @@ def _grpc_server_info_to_worker( worker_url: str, runtime_type: _RuntimeType, response: Any, -) -> _Worker: +) -> _TargetWorker: if runtime_type == "vllm": kv_role = response.kv_role or "" kv_connector = response.kv_connector or "" - worker: _Worker = { + worker: _TargetWorker = { "url": worker_url, "connection_mode": "grpc", "runtime_type": runtime_type, @@ -448,7 +443,7 @@ async def _probe_grpc_worker( *, address: str, runtime_type: Optional[_RuntimeType] = None, -) -> Optional[_Worker]: +) -> Optional[_TargetWorker]: # The RPC goes over the tunnel, `worker_url` is the address the router itself dials. worker_url = f"grpc://{address}" runtime_types: tuple[_RuntimeType, ...] @@ -477,7 +472,7 @@ async def _get_worker( address: str, connection_mode: Optional[_ConnectionMode] = None, runtime_type: Optional[_RuntimeType] = None, -) -> Optional[_Worker]: +) -> Optional[_TargetWorker]: connection_modes: tuple[_ConnectionMode, ...] if connection_mode is None: # No connection_mode discovered -- should probe all @@ -501,7 +496,7 @@ async def _get_worker( # 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: `_update_workers_in_router_replica` cannot tell "not serving" from "could not be + # 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. @@ -516,8 +511,8 @@ async def _build_target_workers( *, connection_mode: Optional[_ConnectionMode] = None, runtime_type: Optional[_RuntimeType] = None, -) -> List[_Worker]: - workers: List[_Worker] = [] +) -> List[_TargetWorker]: + workers: List[_TargetWorker] = [] config = run_spec.configuration if not isinstance(config, ServiceConfiguration): return workers @@ -563,53 +558,80 @@ async def sync_router_workers_for_run_model(run_model: RunModel) -> None: if router_group is None: return - router_job = _get_router_job(run_model, router_group) - if router_job is None: + 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, ) return - # A tunnel is opened here for the router, and inside `_build_target_workers` for every - # worker. Only the router being unreachable aborts the sync -- without a client there is - # nothing to reconcile against. An unreachable worker is skipped instead, see `_get_worker`. + # 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 + + +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 in group %s," - " skipping worker sync", - fmt(run_model), - router_group.name, - ) + logger.debug("%s: failed to get current workers from the router", fmt(router_job)) return - # connection_mode can be grpc or http, runtime_type can be sglang or vllm. - connection_mode = _get_connection_mode_from_workers(current_workers) - runtime_type = _get_runtime_type_from_workers(current_workers) - # Empty current_workers on first sync is expected. First syncprobes both connection_mode and - # runtime_type. Subsequent syncs don't need to probe again because connection_mode and runtime_type - # is already set in current_workers. - target_workers = await _build_target_workers( - run_model, - run_spec, - replica_groups, - connection_mode=connection_mode, - runtime_type=runtime_type, - ) - await _update_workers_in_router_replica( - client, target_workers, current_workers=current_workers - ) + diff = _get_workers_diff(target_workers, current_workers) + await _apply_workers_diff(client, diff, router_job=router_job) except SSHError as e: - # Only the router's own tunnel reaches here, worker tunnels are handled in - # `_get_worker`. 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), - e, - ) - except Exception: - logger.exception("%s: unexpected error when syncing workers with router", fmt(run_model)) + _log_router_replica_unreachable(router_job, e) + + +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) 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 80a9137e0..1277f040a 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,6 +1,7 @@ import json import logging -from contextlib import contextmanager +import uuid +from contextlib import asynccontextmanager, contextmanager from pathlib import Path from typing import Optional from unittest.mock import AsyncMock, MagicMock, patch @@ -9,17 +10,34 @@ import httpx import pytest from httpx import AsyncClient +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.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_router_workers, _get_runtime_type_from_workers, _get_worker, + _get_workers_diff, _grpc_server_info_to_worker, _probe_grpc_worker, _probe_http_worker, + _TargetWorker, + sync_router_workers_for_run_model, +) +from dstack._internal.server.testing.common import ( + create_job, + create_project, + create_repo, + create_run, + create_user, + get_run_spec, ) @@ -145,12 +163,72 @@ def handler(_): assert not [r for r in caplog.records if r.levelno > logging.DEBUG] +@pytest.mark.asyncio +class TestAddWorkerToRouter: + """ + `_TargetWorker` is sent as the `POST /workers` body as is. `TypedDict` does not reject extra + keys at runtime, and the router silently ignores unknown ones, so this test is what catches a + field that should not go over the wire. + """ + + @pytest.mark.parametrize( + "worker", + [ + pytest.param( + { + "url": "http://10.0.0.1:8000", + "worker_type": "regular", + "connection_mode": "http", + "runtime_type": "sglang", + }, + id="http-regular", + ), + pytest.param( + { + "url": "grpc://10.0.0.1:8000", + "worker_type": "prefill", + "connection_mode": "grpc", + "runtime_type": "vllm", + "bootstrap_port": 8998, + "kv_connector": "NixlConnector", + "kv_role": "kv_producer", + }, + id="grpc-prefill", + ), + ], + ) + async def test_posts_worker_as_is(self, worker: _TargetWorker): + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return _json_response(202, {"status": "accepted"}) + + 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 + + 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.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 `_update_workers_in_router_replica` re-registers every worker on each sync. + otherwise `_get_workers_diff` re-registers every worker on each sync. """ async def test_regular_worker(self): @@ -470,3 +548,221 @@ async def test_unexpected_tunnel_error_propagates(self): 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() + + 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", + "runtime_type": "sglang", + } + diff = _get_workers_diff([worker], [{"id": "1", "url": "http://10.0.0.1:8000"}]) + assert diff.is_empty() + + 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.""" + + 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] = [] + + 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) + + +@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.""" + + @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 + + with patch( + "dstack._internal.server.services.runs.router_worker_sync.get_service_replica_client", + get_service_replica_client, + ): + yield + + +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) + conf = parse_run_configuration( + { + "type": "service", + "port": 8000, + "gateway": False, + "groups": [ + { + "name": "router", + "replicas": 1, + "commands": ["smg launch"], + "router": {"type": "sglang"}, + }, + {"name": "worker", "replicas": 1, "commands": ["worker"]}, + ], + } + ) + 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), + ) + # More than one running router means a rolling deployment is replacing the router + for replica_num in range(router_count): + await create_job( + session=session, + run=run, + status=JobStatus.RUNNING, + replica_num=replica_num, + replica_group_name="router", + ) + await session.refresh(run, attribute_names=["jobs"]) + return run + + +_VLLM_WORKER: _TargetWorker = { + "url": "grpc://10.0.0.1:8000", + "worker_type": "regular", + "connection_mode": "grpc", + "runtime_type": "vllm", +} + + +def _patch_build_target_workers(**kwargs): + return patch( + "dstack._internal.server.services.runs.router_worker_sync._build_target_workers", + new_callable=AsyncMock, + **kwargs, + ) + + +@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 + ): + 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} + + with ( + _fake_router_replicas(routers), + _patch_build_target_workers(return_value=[_VLLM_WORKER]) as build_mock, + ): + await sync_router_workers_for_run_model(run) + + 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 + ): + 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} + + 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 + + async def test_skips_probing_workers_when_no_router_is_reachable( + self, test_db, session: AsyncSession + ): + run = await _create_router_service_run(session, router_count=1) + routers = {run.jobs[0].id: _FakeRouter(unreachable=True)} + + 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() + + 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 [] + + 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 == []