diff --git a/src/dstack/_internal/server/background/pipeline_tasks/gateway_replicas.py b/src/dstack/_internal/server/background/pipeline_tasks/gateway_replicas.py index 0ba64dc8a..6cf455812 100644 --- a/src/dstack/_internal/server/background/pipeline_tasks/gateway_replicas.py +++ b/src/dstack/_internal/server/background/pipeline_tasks/gateway_replicas.py @@ -189,6 +189,7 @@ async def fetch(self, limit: int) -> list[GatewayReplicaPipelineItem]: <= now - self._min_processing_interval, GatewayReplicaModel.last_processed_at == GatewayReplicaModel.created_at, + GatewayReplicaModel.skip_min_processing_interval == True, ), or_( GatewayReplicaModel.lock_expires_at.is_(None), @@ -208,6 +209,7 @@ async def fetch(self, limit: int) -> list[GatewayReplicaPipelineItem]: GatewayReplicaModel.lock_token, GatewayReplicaModel.lock_expires_at, GatewayReplicaModel.status, + GatewayReplicaModel.skip_min_processing_interval, ) ) ) @@ -220,6 +222,7 @@ async def fetch(self, limit: int) -> list[GatewayReplicaPipelineItem]: replica_model.lock_expires_at = lock_expires_at replica_model.lock_token = lock_token replica_model.lock_owner = GatewayReplicaPipeline.__name__ + replica_model.skip_min_processing_interval = False items.append( GatewayReplicaPipelineItem( __tablename__=GatewayReplicaModel.__tablename__, 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 01ec6ba50..a89ef1ede 100644 --- a/src/dstack/_internal/server/background/pipeline_tasks/jobs_running.py +++ b/src/dstack/_internal/server/background/pipeline_tasks/jobs_running.py @@ -79,7 +79,10 @@ get_instance_specific_mounts, resolve_provisioning_image, ) -from dstack._internal.server.services.gateways import get_gateway_replica_models +from dstack._internal.server.services.gateways import ( + get_gateway_replica_models, + skip_gateway_replicas_min_processing_interval, +) from dstack._internal.server.services.instances import ( get_instance_remote_connection_info, get_instance_ssh_private_keys, @@ -331,6 +334,7 @@ async def process(self, item: JobRunningPipelineItem): await _apply_process_result( item=item, job_model=context.job_model, + run_model=context.run_model, result=result, ) new_status = result.job_update_map.get("status") @@ -339,6 +343,8 @@ async def process(self, item: JobRunningPipelineItem): # Hint run pipeline for fast run transition to RUNNING status. if new_status == JobStatus.RUNNING and context.job_model.run.status != RunStatus.RUNNING: self._pipeline_hinter.hint_fetch(RunModel.__name__) + if context.run_model.gateway_id is not None and result.job_update_map.get("registered"): + self._pipeline_hinter.hint_fetch(GatewayReplicaModel.__name__) @dataclass @@ -1084,6 +1090,7 @@ def _server_access_enabled(context: _ProcessContext) -> bool: async def _apply_process_result( item: JobRunningPipelineItem, job_model: JobModel, + run_model: RunModel, result: _ProcessResult, ) -> None: set_processed_update_map_fields(result.job_update_map) @@ -1121,6 +1128,9 @@ async def _apply_process_result( .values(skip_min_processing_interval=True) ) + if run_model.gateway_id is not None and result.job_update_map.get("registered"): + await skip_gateway_replicas_min_processing_interval(session, run_model.gateway_id) + _emit_result_events(session=session, job_model=job_model, result=result) diff --git a/src/dstack/_internal/server/migrations/versions/2026/08_18_1758_04126c7ea0c8_add_gatewayreplicamodel_skip_min_.py b/src/dstack/_internal/server/migrations/versions/2026/08_18_1758_04126c7ea0c8_add_gatewayreplicamodel_skip_min_.py new file mode 100644 index 000000000..14e7e9447 --- /dev/null +++ b/src/dstack/_internal/server/migrations/versions/2026/08_18_1758_04126c7ea0c8_add_gatewayreplicamodel_skip_min_.py @@ -0,0 +1,37 @@ +"""Add GatewayReplicaModel.skip_min_processing_interval + +Revision ID: 04126c7ea0c8 +Revises: 3d4f69210528 +Create Date: 2026-08-18 17:58:33.045050+00:00 + +""" + +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision = "04126c7ea0c8" +down_revision = "3d4f69210528" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + with op.batch_alter_table("gateway_computes", schema=None) as batch_op: + batch_op.add_column( + sa.Column( + "skip_min_processing_interval", + sa.Boolean(), + server_default=sa.false(), + nullable=False, + ) + ) + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + with op.batch_alter_table("gateway_computes", schema=None) as batch_op: + batch_op.drop_column("skip_min_processing_interval") + # ### end Alembic commands ### diff --git a/src/dstack/_internal/server/models.py b/src/dstack/_internal/server/models.py index a1225fb8e..dcc5707f4 100644 --- a/src/dstack/_internal/server/models.py +++ b/src/dstack/_internal/server/models.py @@ -694,6 +694,9 @@ class GatewayReplicaModel(PipelineModelMixin, BaseModel): ) created_at: Mapped[datetime] = mapped_column(NaiveDateTime, default=get_current_datetime) last_processed_at: Mapped[datetime] = mapped_column(NaiveDateTime) + skip_min_processing_interval: Mapped[bool] = mapped_column( + Boolean, default=False, server_default=false() + ) status: Mapped[GatewayReplicaStatus] = mapped_column(EnumAsString(GatewayReplicaStatus, 100)) status_message: Mapped[Optional[str]] = mapped_column(Text) scale_in: Mapped[bool] = mapped_column(Boolean, server_default=false()) diff --git a/src/dstack/_internal/server/services/gateways/__init__.py b/src/dstack/_internal/server/services/gateways/__init__.py index 4c9a69635..a9fec7aca 100644 --- a/src/dstack/_internal/server/services/gateways/__init__.py +++ b/src/dstack/_internal/server/services/gateways/__init__.py @@ -868,6 +868,23 @@ def get_gateway_replica_models(gateway_model: GatewayModel) -> List[GatewayRepli return replicas +async def skip_gateway_replicas_min_processing_interval( + session: AsyncSession, gateway_id: uuid.UUID +) -> None: + await session.execute( + update(GatewayReplicaModel) + .where( + or_( + GatewayReplicaModel.gateway_id == gateway_id, + GatewayReplicaModel.id.in_( + select(GatewayModel.gateway_replica_id).where(GatewayModel.id == gateway_id) + ), + ) + ) + .values(skip_min_processing_interval=True) + ) + + def get_gateway_configuration(gateway_model: GatewayModel) -> GatewayConfiguration: if gateway_model.configuration is not None: return validate_json_extra_ignore(GatewayConfiguration, gateway_model.configuration) diff --git a/src/dstack/_internal/server/services/runs/__init__.py b/src/dstack/_internal/server/services/runs/__init__.py index 47903412e..a541b3904 100644 --- a/src/dstack/_internal/server/services/runs/__init__.py +++ b/src/dstack/_internal/server/services/runs/__init__.py @@ -45,6 +45,7 @@ from dstack._internal.server.db import get_db, is_db_postgres, is_db_sqlite from dstack._internal.server.models import ( FleetModel, + GatewayReplicaModel, JobModel, MemberModel, ProbeModel, @@ -850,6 +851,8 @@ async def submit_run( if pipeline_hinter is not None: pipeline_hinter.hint_fetch(JobModel.__name__) pipeline_hinter.hint_fetch(RunModel.__name__) + if run_model.gateway is not None or run_model.gateway_id is not None: + pipeline_hinter.hint_fetch(GatewayReplicaModel.__name__) await session.refresh(run_model) run = await get_run_by_id(session, project, run_model.id) diff --git a/src/dstack/_internal/server/services/services/__init__.py b/src/dstack/_internal/server/services/services/__init__.py index a08ed1a8e..ff6ce177c 100644 --- a/src/dstack/_internal/server/services/services/__init__.py +++ b/src/dstack/_internal/server/services/services/__init__.py @@ -73,6 +73,9 @@ async def register_service(session: AsyncSession, run_model: RunModel, run_spec: if gateway is not None: service_spec = await _register_service_in_gateway(session, run_model, run_spec, gateway) run_model.gateway = gateway + # For faster registration + for replica_model in get_gateway_replica_models(gateway): + replica_model.skip_min_processing_interval = True elif not settings.FORBID_SERVICES_WITHOUT_GATEWAY: service_spec = _register_service_in_server(session, run_model, run_spec) else: diff --git a/src/tests/_internal/server/background/pipeline_tasks/test_gateway_replicas.py b/src/tests/_internal/server/background/pipeline_tasks/test_gateway_replicas.py index c9c3f66d9..9c2e0356e 100644 --- a/src/tests/_internal/server/background/pipeline_tasks/test_gateway_replicas.py +++ b/src/tests/_internal/server/background/pipeline_tasks/test_gateway_replicas.py @@ -197,6 +197,33 @@ async def test_fetch_selects_eligible_replicas_and_sets_lock_fields( assert recent.lock_owner is None assert locked.lock_owner == "OtherPipeline" + async def test_fetch_includes_recent_replica_with_skip_min_processing_interval( + self, test_db, session: AsyncSession, fetcher: GatewayReplicaFetcher + ): + project = await create_project(session=session) + 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, + ) + now = get_current_datetime() + replica = await create_gateway_replica( + session=session, + gateway_id=gateway.id, + status=GatewayReplicaStatus.RUNNING, + last_processed_at=now, + ) + replica.skip_min_processing_interval = True + await session.commit() + + items = await fetcher.fetch(limit=10) + + assert [item.id for item in items] == [replica.id] + await session.refresh(replica) + assert not replica.skip_min_processing_interval + @pytest.mark.parametrize( "gateway_status,to_be_deleted", [ 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 1e854148a..560e3c201 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 @@ -2421,6 +2421,7 @@ async def test_registers_service_replica_only_after_probes_pass( assert not job.registered assert not events + @pytest.mark.parametrize("legacy_replica", [False, True]) async def test_registers_service_replica_in_gateway( self, test_db, @@ -2429,6 +2430,7 @@ async def test_registers_service_replica_in_gateway( ssh_tunnel_mock: Mock, shim_client_mock: Mock, runner_client_mock: Mock, + legacy_replica: bool, ): user = await create_user(session=session) project = await create_project(session=session, owner=user) @@ -2442,11 +2444,19 @@ async def test_registers_service_replica_in_gateway( name="test-gateway", wildcard_domain="example.com", ) - await create_gateway_replica( - session=session, - backend_id=backend.id, - gateway_id=gateway.id, - ) + if legacy_replica: + gateway_replica = await create_gateway_replica( + session=session, + backend_id=backend.id, + ) + gateway.gateway_replica_id = gateway_replica.id + await session.commit() + else: + gateway_replica = await create_gateway_replica( + session=session, + backend_id=backend.id, + gateway_id=gateway.id, + ) run = await create_run( session=session, project=project, @@ -2483,6 +2493,8 @@ async def test_registers_service_replica_in_gateway( await session.refresh(job) assert job.status == JobStatus.RUNNING assert job.registered + await session.refresh(gateway_replica) + assert gateway_replica.skip_min_processing_interval events = await list_events(session) assert {event.message for event in events} == { "Job status changed PULLING -> RUNNING", diff --git a/src/tests/_internal/server/routers/test_runs.py b/src/tests/_internal/server/routers/test_runs.py index 6010d7c7d..832633db9 100644 --- a/src/tests/_internal/server/routers/test_runs.py +++ b/src/tests/_internal/server/routers/test_runs.py @@ -4015,12 +4015,14 @@ async def test_submit_to_correct_proxy( @pytest.mark.asyncio @pytest.mark.parametrize("populate_configuration", [True, False]) + @pytest.mark.parametrize("legacy_replica", [False, True]) async def test_submit_to_gateway_by_name( self, test_db, session: AsyncSession, client: AsyncClient, populate_configuration: bool, + legacy_replica: bool, ) -> None: user = await create_user(session=session, global_role=GlobalRole.USER) project = await create_project(session=session, owner=user, name="test-project") @@ -4038,12 +4040,21 @@ async def test_submit_to_gateway_by_name( wildcard_domain="my-gateway.example", populate_configuration=populate_configuration, ) - await create_gateway_replica( - session=session, - backend_id=backend.id, - gateway_id=gateway.id, - populate_configuration=populate_configuration, - ) + if legacy_replica: + gateway_replica = await create_gateway_replica( + session=session, + backend_id=backend.id, + populate_configuration=populate_configuration, + ) + gateway.gateway_replica_id = gateway_replica.id + await session.commit() + else: + gateway_replica = await create_gateway_replica( + session=session, + backend_id=backend.id, + gateway_id=gateway.id, + populate_configuration=populate_configuration, + ) run_spec = get_service_run_spec( repo_id=repo.name, run_name="test-service", @@ -4065,6 +4076,8 @@ async def test_submit_to_gateway_by_name( res = await session.execute(select(RunModel)) run = res.scalar_one() assert run.gateway_id is not None + await session.refresh(gateway_replica) + assert gateway_replica.skip_min_processing_interval @pytest.mark.asyncio async def test_return_error_if_specified_gateway_not_exists(