From 0eb560686fa77b7f070b6833f18ff62e7acba2f6 Mon Sep 17 00:00:00 2001 From: Jvst Me Date: Wed, 5 Aug 2026 08:46:22 +0200 Subject: [PATCH] Distinguish router and non-router replicas **Problem**: Non-router replicas in a service with a router should not be registered on the gateway. To determine whether a particular replica should be registered, one has to parse its job spec and run spec and see if it is a router replica or not. This does not work well for the upcoming gateway state sync mechanism, which needs to quickly fetch all jobs to be registered from the database, without parsing their specs. **Solution**: - Do **not** set `JobModel.registered` for non-router replicas in services with a router. That way, `JobModel.registered` can be reliably used to identify the jobs that should be registered on the gateway, without parsing the job spec. - Since the router worker sync pipeline needs both router and non-router replicas with successful probes, introduce `JobModel.ready`, set for any replica once probes pass. ```shell $ dstack event --within-run prefill-decode | grep -E "ready|register" [2026-08-05 10:30:29] [job prefill-decode-0-0] Service replica ready to receive requests [2026-08-05 10:30:29] [job prefill-decode-0-0] Service replica registered to receive requests [2026-08-05 10:44:45] [job prefill-decode-0-1] Service replica ready to receive requests [2026-08-05 10:45:00] [job prefill-decode-0-2] Service replica ready to receive requests [2026-08-05 11:15:22] [job prefill-decode-0-0] Service replica unregistered from receiving requests ``` --- .../background/pipeline_tasks/jobs_running.py | 68 ++++--- .../service_router_worker_sync.py | 3 +- ...04_0702_72cfa56364ad_add_jobmodel_ready.py | 119 +++++++++++++ src/dstack/_internal/server/models.py | 9 +- .../_internal/server/services/proxy/repo.py | 4 - .../services/runs/router_worker_sync.py | 2 +- src/dstack/_internal/server/testing/common.py | 10 +- .../pipeline_tasks/test_running_jobs.py | 167 +++++++++++++++++- 8 files changed, 347 insertions(+), 35 deletions(-) create mode 100644 src/dstack/_internal/server/migrations/versions/2026/08_04_0702_72cfa56364ad_add_jobmodel_ready.py diff --git a/src/dstack/_internal/server/background/pipeline_tasks/jobs_running.py b/src/dstack/_internal/server/background/pipeline_tasks/jobs_running.py index 720b0141b..b1af848ff 100644 --- a/src/dstack/_internal/server/background/pipeline_tasks/jobs_running.py +++ b/src/dstack/_internal/server/background/pipeline_tasks/jobs_running.py @@ -20,7 +20,6 @@ ) from dstack._internal.core.models.configurations import ( DevEnvironmentConfiguration, - ServiceConfiguration, ) from dstack._internal.core.models.files import FileArchiveMapping from dstack._internal.core.models.instances import InstanceStatus, SSHConnectionParams @@ -34,7 +33,6 @@ Job, JobProvisioningData, JobRuntimeData, - JobSpec, JobStatus, JobSubmission, JobTerminationReason, @@ -366,6 +364,7 @@ class _JobUpdateMap(ItemUpdateMap, total=False): disconnected_at: Optional[datetime] inactivity_secs: Optional[int] exit_status: Optional[int] + ready: bool registered: bool image_pull_progress: Optional[str] skip_min_processing_interval: bool @@ -1083,6 +1082,12 @@ def _emit_result_events( job_model.disconnected_at, ), ) + _emit_readiness_change_event( + session=session, + job_model=job_model, + old_ready=job_model.ready, + new_ready=result.job_update_map.get("ready", job_model.ready), + ) if result.replica_registration is not None: targets = [events.Target.from_model(job_model)] if result.replica_registration.gateway_target is not None: @@ -1175,13 +1180,32 @@ async def _maybe_register_replica( ) -> None: if ( context.run.run_spec.configuration.type != "service" - or _get_result_registered(context.job_model, result) or context.job_model.job_num != 0 or result.new_probe_models - or not is_job_ready(context.job_model.probes, context.job.job_spec.probes) ): return + is_ready = is_job_ready(context.job_model.probes, context.job.job_spec.probes) + if is_ready and not context.job_model.ready: + result.job_update_map["ready"] = True + + router_group = next( + (g for g in context.run.run_spec.configuration.replica_groups if g.router is not None), + None, + ) + is_router_replica = ( + router_group is not None and context.job.job_spec.replica_group == router_group.name + ) + # non-router replicas aren't registered if the service has a router + if router_group is not None and not is_router_replica: + if context.job_model.registered: + # migration edge case: a pre-0.21.0 server replica incorrectly set registered=True + result.job_update_map["registered"] = False + return + + if not is_ready or _get_result_registered(context.job_model, result): + return + ssh_head_proxy: Optional[SSHConnectionParams] = None ssh_head_proxy_private_key: Optional[str] = None instance = get_or_error(context.job_model.instance) @@ -1220,23 +1244,6 @@ async def _register_service_replica( ) -> Optional[events.Target]: if context.run_model.gateway_id is None: return None - - job_spec = validate_json_extra_ignore(JobSpec, context.job_model.job_spec_data) - - # For router-based services (e.g. PD disaggregation), only router replicas should be - # registered with the gateway. Worker replicas are discovered by the router-worker - # sync pipeline and should not be routed to directly by the gateway. - config = context.run.run_spec.configuration - assert isinstance(config, ServiceConfiguration) - router_group = next((g for g in config.replica_groups if g.router is not None), None) - is_router_replica = router_group is not None and job_spec.replica_group == router_group.name - if router_group is not None and not is_router_replica: - logger.debug( - "%s: skipping gateway replica registration (non-router replica)", - fmt(context.job_model), - ) - return None - async with get_session_ctx() as session: gateway_model, connections = await get_or_add_gateway_connections( session, context.run_model.gateway_id @@ -1261,7 +1268,7 @@ async def _register_service_replica( async with conn.client() as gateway_client: await gateway_client.register_replica( run=context.run, - job_spec=job_spec, + job_spec=context.job.job_spec, job_submission=job_submission, instance_project_ssh_private_key=instance_project_ssh_private_key, ssh_head_proxy=ssh_head_proxy, @@ -1877,6 +1884,23 @@ def _emit_reachability_change_event( ) +def _emit_readiness_change_event( + session: AsyncSession, + job_model: JobModel, + old_ready: bool, + new_ready: bool, +) -> None: + # ready: False -> True + if not old_ready and new_ready: + events.emit( + session, + "Service replica ready to receive requests", + actor=events.SystemActor(), + targets=[events.Target.from_model(job_model)], + ) + # ready: True -> False is not possible as of this writing + + def _terminate_job( job_model: JobModel, job_update_map: _JobUpdateMap, 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 5ded45afe..48fe52210 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 @@ -233,13 +233,12 @@ async def process(self, item: ServiceRouterWorkerSyncPipelineItem) -> None: selectinload( RunModel.jobs.and_( JobModel.status == JobStatus.RUNNING, - JobModel.registered == True, + JobModel.ready == True, ) ) .load_only( JobModel.id, JobModel.status, - JobModel.registered, JobModel.job_spec_data, JobModel.job_provisioning_data, JobModel.job_runtime_data, diff --git a/src/dstack/_internal/server/migrations/versions/2026/08_04_0702_72cfa56364ad_add_jobmodel_ready.py b/src/dstack/_internal/server/migrations/versions/2026/08_04_0702_72cfa56364ad_add_jobmodel_ready.py new file mode 100644 index 000000000..97d1ce394 --- /dev/null +++ b/src/dstack/_internal/server/migrations/versions/2026/08_04_0702_72cfa56364ad_add_jobmodel_ready.py @@ -0,0 +1,119 @@ +"""Add JobModel.ready + +Revision ID: 72cfa56364ad +Revises: ecc9e8a0bfac +Create Date: 2026-08-04 07:02:02.041000+00:00 + +""" + +import json +import uuid +from typing import Optional + +import sqlalchemy as sa +from alembic import op +from sqlalchemy_utils import UUIDType + +# revision identifiers, used by Alembic. +revision = "72cfa56364ad" +down_revision = "ecc9e8a0bfac" +branch_labels = None +depends_on = None + +# Partial table descriptions - only columns needed for the data migration below. +jobs_table = sa.Table( + "jobs", + sa.MetaData(), + sa.Column("id", UUIDType(binary=False), primary_key=True, default=uuid.uuid4), + sa.Column("run_id", UUIDType(binary=False)), + sa.Column("status", sa.String(100)), + sa.Column("registered", sa.Boolean()), + sa.Column("ready", sa.Boolean()), + sa.Column("job_spec_data", sa.Text()), +) +runs_table = sa.Table( + "runs", + sa.MetaData(), + sa.Column("id", UUIDType(binary=False), primary_key=True, default=uuid.uuid4), + sa.Column("run_spec", sa.Text()), +) +service_router_worker_sync_table = sa.Table( + "service_router_worker_sync", + sa.MetaData(), + sa.Column("run_id", UUIDType(binary=False)), +) + + +def _get_router_group_name(run_spec_data: str) -> Optional[str]: + configuration = json.loads(run_spec_data).get("configuration") or {} + if configuration.get("type") != "service": + return None + replica_groups = configuration.get("replicas") + if not isinstance(replica_groups, list): + return None + for group in replica_groups: + if isinstance(group, dict) and group.get("router") is not None: + return group.get("name") + return None + + +def _get_job_replica_group(job_spec_data: str) -> str: + return json.loads(job_spec_data).get("replica_group", "0") + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + with op.batch_alter_table("jobs", schema=None) as batch_op: + batch_op.add_column( + sa.Column("ready", sa.Boolean(), server_default=sa.false(), nullable=False) + ) + + # ### end Alembic commands ### + + bind = op.get_bind() + + # backfill ready=True for existing registered replicas + + bind.execute(jobs_table.update().where(jobs_table.c.registered == True).values(ready=True)) + + # set registered=False for non-router replicas in services with a router + + router_run_ids_subq = sa.select(service_router_worker_sync_table.c.run_id) + candidate_jobs = bind.execute( + sa.select(jobs_table.c.id, jobs_table.c.run_id, jobs_table.c.job_spec_data).where( + jobs_table.c.registered == True, + jobs_table.c.status == "RUNNING", + jobs_table.c.run_id.in_(router_run_ids_subq), + ) + ).all() + + run_ids = {row.run_id for row in candidate_jobs} + router_group_name_by_run_id = {} + if run_ids: + for row in bind.execute( + sa.select(runs_table.c.id, runs_table.c.run_spec).where(runs_table.c.id.in_(run_ids)) + ).all(): + router_group_name_by_run_id[row.id] = _get_router_group_name(row.run_spec) + + non_router_job_ids = [] + for row in candidate_jobs: + router_group_name = router_group_name_by_run_id.get(row.run_id) + if router_group_name is None: + continue + if _get_job_replica_group(row.job_spec_data) != router_group_name: + non_router_job_ids.append(row.id) + + if non_router_job_ids: + bind.execute( + jobs_table.update() + .where(jobs_table.c.id.in_(non_router_job_ids)) + .values(registered=False) + ) + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + with op.batch_alter_table("jobs", schema=None) as batch_op: + batch_op.drop_column("ready") + + # ### end Alembic commands ### diff --git a/src/dstack/_internal/server/models.py b/src/dstack/_internal/server/models.py index 72c974542..ef1f4d838 100644 --- a/src/dstack/_internal/server/models.py +++ b/src/dstack/_internal/server/models.py @@ -574,9 +574,14 @@ class JobModel(PipelineModelMixin, BaseModel): probes: Mapped[list["ProbeModel"]] = relationship( back_populates="job", order_by="ProbeModel.probe_num" ) + ready: Mapped[bool] = mapped_column(Boolean, server_default=false()) + """Whether the replica is ready to receive service requests based on probe statuses. + Always `False` for non-service runs. + """ registered: Mapped[bool] = mapped_column(Boolean, server_default=false()) - """`registered` shows whether the replica is registered to receive service requests. - It is always `False` for non-service runs. + """Whether the replica is registered to receive service requests from dstack-proxy. + Always `False` for non-service runs or jobs that shouldn't be registered + (e.g., non-router replicas for services with routers). """ waiting_master_job: Mapped[Optional[bool]] = mapped_column(Boolean) """`waiting_master_job` is `True` for non-master jobs that have to wait for master processing before diff --git a/src/dstack/_internal/server/services/proxy/repo.py b/src/dstack/_internal/server/services/proxy/repo.py index 5a0caffd7..d0afd41de 100644 --- a/src/dstack/_internal/server/services/proxy/repo.py +++ b/src/dstack/_internal/server/services/proxy/repo.py @@ -113,10 +113,6 @@ async def get_service(self, project_name: str, run_name: str) -> Optional[Servic ssh_head_proxy = rci.ssh_proxy ssh_head_proxy_private_key = get_or_error(rci.ssh_proxy_keys)[0].private job_spec = get_job_spec(job) - if router_group is not None and job_spec.replica_group != router_group.name: - # Strict router-only: when a router is configured, the proxy should only be aware - # of router replicas. - continue replica = Replica( id=job.id.hex, app_port=get_service_port(job_spec, run_spec.configuration), 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 4b9b8af65..9876d2415 100644 --- a/src/dstack/_internal/server/services/runs/router_worker_sync.py +++ b/src/dstack/_internal/server/services/runs/router_worker_sync.py @@ -1,4 +1,4 @@ -"""Reconcile SGLang router /workers with dstack's registered worker replicas (async, SSH-tunneled).""" +"""Reconcile SGLang router /workers with dstack's ready worker replicas (async, SSH-tunneled).""" import json from typing import Any, List, Literal, Optional, TypedDict diff --git a/src/dstack/_internal/server/testing/common.py b/src/dstack/_internal/server/testing/common.py index 46b51a189..e892512df 100644 --- a/src/dstack/_internal/server/testing/common.py +++ b/src/dstack/_internal/server/testing/common.py @@ -442,13 +442,20 @@ async def create_job( instance_assigned: bool = False, disconnected_at: Optional[datetime] = None, registered: bool = False, + ready: bool = False, waiting_master_job: Optional[bool] = None, + replica_group_name: Optional[str] = None, ) -> JobModel: if deployment_num is None: deployment_num = run.deployment_num run_spec = validate_json_extra_ignore(RunSpec, run.run_spec) job_spec = ( - await get_job_specs_from_run_spec(run_spec=run_spec, secrets={}, replica_num=replica_num) + await get_job_specs_from_run_spec( + run_spec=run_spec, + secrets={}, + replica_num=replica_num, + replica_group_name=replica_group_name, + ) )[0] job_spec.job_num = job_num job = JobModel( @@ -476,6 +483,7 @@ async def create_job( disconnected_at=disconnected_at, probes=[], registered=registered, + ready=ready, waiting_master_job=waiting_master_job, ) session.add(job) diff --git a/src/tests/_internal/server/background/pipeline_tasks/test_running_jobs.py b/src/tests/_internal/server/background/pipeline_tasks/test_running_jobs.py index ffb7d1802..e3978428d 100644 --- a/src/tests/_internal/server/background/pipeline_tasks/test_running_jobs.py +++ b/src/tests/_internal/server/background/pipeline_tasks/test_running_jobs.py @@ -69,6 +69,7 @@ from dstack._internal.server.services.runs.replicas import RouterEnvStatus from dstack._internal.server.services.volumes import volume_model_to_volume from dstack._internal.server.testing.common import ( + clear_events, create_backend, create_code, create_export, @@ -1971,6 +1972,7 @@ async def test_registers_service_replica_immediately_if_no_probes( events = await list_events(session) assert {event.message for event in events} == { "Job status changed PULLING -> RUNNING", + "Service replica ready to receive requests", "Service replica registered to receive requests", } @@ -2060,8 +2062,11 @@ async def test_registers_service_replica_only_after_probes_pass( events = await list_events(session) if expect_to_register: assert job.registered - assert len(events) == 1 - assert events[0].message == "Service replica registered to receive requests" + assert len(events) == 2 + assert {event.message for event in events} == { + "Service replica ready to receive requests", + "Service replica registered to receive requests", + } else: assert not job.registered assert not events @@ -2132,6 +2137,7 @@ async def test_registers_service_replica_in_gateway( events = await list_events(session) assert {event.message for event in events} == { "Job status changed PULLING -> RUNNING", + "Service replica ready to receive requests", "Service replica registered to receive requests", } mock_gateway_connection.return_value.client.return_value.__aenter__.return_value.register_replica.assert_called_once_with( @@ -2218,6 +2224,7 @@ async def test_registers_service_replica_in_gateway_when_running_on_imported_ins events = await list_events(session) assert {event.message for event in events} == { "Job status changed PULLING -> RUNNING", + "Service replica ready to receive requests", "Service replica registered to receive requests", } mock_gateway_connection.return_value.client.return_value.__aenter__.return_value.register_replica.assert_called_once_with( @@ -2458,13 +2465,167 @@ async def test_provisioning_shim_uses_server_default_registry( assert call_kwargs["registry_username"] == "server-user" assert call_kwargs["registry_password"] == "server-pass" + async def test_registers_router_replica_but_not_worker_replica_in_gateway( + self, + test_db, + session: AsyncSession, + worker: JobRunningWorker, + ssh_tunnel_mock: Mock, + runner_client_mock: Mock, + mock_gateway_connection: AsyncMock, + ): + user = await create_user(session=session) + project = await create_project(session=session, owner=user) + repo = await create_repo(session=session, project_id=project.id) + backend = await create_backend(session=session, project_id=project.id) + gateway = await create_gateway( + session=session, + project_id=project.id, + backend_id=backend.id, + status=GatewayStatus.RUNNING, + name="test-gateway", + wildcard_domain="example.com", + ) + await create_gateway_compute( + session=session, + backend_id=backend.id, + gateway_id=gateway.id, + ) + run = await create_run( + session=session, + project=project, + repo=repo, + user=user, + run_spec=get_run_spec( + run_name="test", + repo_id=repo.name, + configuration=_router_service_configuration("sglang", gateway="test-gateway"), + ), + gateway=gateway, + ) + instance = await create_instance( + session=session, + project=project, + status=InstanceStatus.BUSY, + ) + router_job = await create_job( + session=session, + run=run, + replica_num=0, + replica_group_name="router", + status=JobStatus.RUNNING, + job_provisioning_data=get_job_provisioning_data(dockerized=True), + instance=instance, + instance_assigned=True, + ) + worker_job = await create_job( + session=session, + run=run, + replica_num=1, + replica_group_name="worker", + status=JobStatus.RUNNING, + job_provisioning_data=get_job_provisioning_data(dockerized=True), + instance=instance, + instance_assigned=True, + ) + runner_client_mock.pull.return_value = PullResponse( + job_states=[], job_logs=[], runner_logs=[], last_updated=0 + ) + + await _process_job(session, worker, router_job) + + await session.refresh(router_job) + assert router_job.registered + assert router_job.ready + events = await list_events(session) + assert {event.message for event in events} == { + "Service replica ready to receive requests", + "Service replica registered to receive requests", + } + + await clear_events(session) + + await _process_job(session, worker, worker_job) + + await session.refresh(worker_job) + assert not worker_job.registered + assert worker_job.ready + events = await list_events(session) + assert {event.message for event in events} == { + "Service replica ready to receive requests", + } + + async def test_resets_stale_registered_flag_for_non_router_replica( + self, + test_db, + session: AsyncSession, + worker: JobRunningWorker, + ssh_tunnel_mock: Mock, + runner_client_mock: Mock, + ): + """Migration edge case: a pre-0.21.0 server may have incorrectly marked + a non-router replica as registered. Should be corrected. + """ + 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, + run_spec=get_run_spec( + run_name="test", + repo_id=repo.name, + configuration=_router_service_configuration( + "sglang", + probes=[ProbeConfig(type="http", url="/health", ready_after=1)], + ), + ), + ) + instance = await create_instance( + session=session, + project=project, + status=InstanceStatus.BUSY, + ) + worker_job = await create_job( + session=session, + run=run, + replica_num=1, + replica_group_name="worker", + status=JobStatus.RUNNING, + job_provisioning_data=get_job_provisioning_data(dockerized=True), + instance=instance, + instance_assigned=True, + registered=True, + ) + await create_probe(session=session, job=worker_job, probe_num=0, success_streak=0) + runner_client_mock.pull.return_value = PullResponse( + job_states=[], job_logs=[], runner_logs=[], last_updated=0 + ) + + await _process_job(session, worker, worker_job) + + await session.refresh(worker_job) + assert worker_job.status == JobStatus.RUNNING + assert not worker_job.registered + events = await list_events(session) + assert events == [] + -def _router_service_configuration(router_type: str) -> ServiceConfiguration: +def _router_service_configuration( + router_type: str, + *, + gateway: Optional[str] = None, + probes: Optional[list[ProbeConfig]] = None, +) -> ServiceConfiguration: return ServiceConfiguration.model_validate( { "type": "service", "port": 8000, "image": "ubuntu", + "gateway": gateway, + "probes": probes, "replicas": [ {"name": "worker", "commands": ["echo worker"], "count": 1}, {