Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand All @@ -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,
)
)
)
Expand All @@ -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__,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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")
Expand All @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)


Expand Down
Original file line number Diff line number Diff line change
@@ -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 ###
3 changes: 3 additions & 0 deletions src/dstack/_internal/server/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -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())
Expand Down
17 changes: 17 additions & 0 deletions src/dstack/_internal/server/services/gateways/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
3 changes: 3 additions & 0 deletions src/dstack/_internal/server/services/runs/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
3 changes: 3 additions & 0 deletions src/dstack/_internal/server/services/services/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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",
[
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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)
Expand All @@ -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,
Expand Down Expand Up @@ -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",
Expand Down
25 changes: 19 additions & 6 deletions src/tests/_internal/server/routers/test_runs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand All @@ -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",
Expand All @@ -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(
Expand Down
Loading