From 99da49115317bc0973d5cf94ed771cc157a00380 Mon Sep 17 00:00:00 2001 From: Jvst Me Date: Mon, 27 Jul 2026 07:41:44 +0200 Subject: [PATCH 1/3] Support replicated AWS gateways with ACM Gateways with an ACM certificate can now have more than one replica. ```yaml type: gateway backend: aws region: eu-west-1 domain: example.com certificate: type: acm arn: arn:aws:acm:eu-west-1:164099421079:certificate/3670388f-f43b-4872-aaf8-907b107a170d replicas: 2 ``` Load balancing across gateway replicas is performed by a single ALB associated with the gateway. ```shell $ dstack gateway list NAME BACKEND HOSTNAME DOMAIN DEFAULT STATUS little-sloth dstack-qe1na76o-lb-187858581.eu-west-1.elb.amazonaws.com example.com running replica=0 aws (eu-west-1) 18.202.25.65 running replica=1 aws (eu-west-1) 3.255.100.238 running ``` --- mkdocs/docs/concepts/gateways.md | 2 +- .../_internal/core/backends/aws/compute.py | 183 +++++++++++---- .../_internal/core/backends/base/compute.py | 51 ++++ src/dstack/_internal/core/models/gateways.py | 16 +- .../pipeline_tasks/gateway_replicas.py | 199 ++++++++++++++-- .../background/pipeline_tasks/gateways.py | 220 ++++++++++++++++-- ...bfac_gateways_hostname_and_backend_data.py | 80 +++++++ src/dstack/_internal/server/models.py | 14 +- .../server/services/gateways/__init__.py | 32 ++- .../_internal/server/routers/test_gateways.py | 23 +- 10 files changed, 711 insertions(+), 109 deletions(-) create mode 100644 src/dstack/_internal/server/migrations/versions/2026/07_22_1953_ecc9e8a0bfac_gateways_hostname_and_backend_data.py diff --git a/mkdocs/docs/concepts/gateways.md b/mkdocs/docs/concepts/gateways.md index a959ab104..64e197fed 100644 --- a/mkdocs/docs/concepts/gateways.md +++ b/mkdocs/docs/concepts/gateways.md @@ -221,7 +221,7 @@ $ dstack gateway list Replicated gateways are an experimental feature and currently have limitations: - Changing the number of replicas or redeploying replicas is not supported. - - HTTPS is not supported. Use an external load balancer for TLS termination. + - HTTPS is only supported for AWS gateways with the `acm` [certificate type](#certificate). For other gateways, use an external load balancer for TLS termination. - An unavailable gateway replica prevents any new services or service replicas from being added. - All replicas are bound to the same backend and region. - At most 3 replicas are allowed per gateway. diff --git a/src/dstack/_internal/core/backends/aws/compute.py b/src/dstack/_internal/core/backends/aws/compute.py index b28cb0e10..0c93080f3 100644 --- a/src/dstack/_internal/core/backends/aws/compute.py +++ b/src/dstack/_internal/core/backends/aws/compute.py @@ -25,6 +25,7 @@ ComputeTTLCache, ComputeWithAllOffersCached, ComputeWithCreateInstanceSupport, + ComputeWithGatewayLoadBalancerSupport, ComputeWithGatewaySupport, ComputeWithInstanceVolumesSupport, ComputeWithMultinodeSupport, @@ -57,6 +58,8 @@ from dstack._internal.core.models.common import CoreModel from dstack._internal.core.models.gateways import ( GatewayComputeConfiguration, + GatewayLoadBalancerConfiguration, + GatewayLoadBalancerData, GatewayProvisioningData, ) from dstack._internal.core.models.instances import ( @@ -123,6 +126,7 @@ class AWSCompute( ComputeWithReservationSupport, ComputeWithPlacementGroupSupport, ComputeWithGatewaySupport, + ComputeWithGatewayLoadBalancerSupport, ComputeWithPrivateGatewaySupport, ComputeWithVolumeSupport, Compute, @@ -584,17 +588,51 @@ def create_gateway( instance = response[0] instance.wait_until_running() instance.reload() # populate instance.public_ip_address - if configuration.certificate is None or configuration.certificate.type != "acm": - ip_address = _get_instance_ip(instance, configuration.public_ip) - return GatewayProvisioningData( - instance_id=instance.instance_id, - region=configuration.region, - availability_zone=availability_zone, - ip_address=ip_address, - ) + ip_address = _get_instance_ip(instance, configuration.public_ip) + return GatewayProvisioningData( + instance_id=instance.instance_id, + region=configuration.region, + availability_zone=availability_zone, + ip_address=ip_address, + ) + + def create_gateway_load_balancer( + self, + configuration: GatewayLoadBalancerConfiguration, + ) -> GatewayLoadBalancerData: + """Creates an ALB, target group, and listeners for a gateway with an ACM certificate.""" + assert configuration.certificate is not None + assert configuration.certificate.type == "acm" + ec2_client = self.session.client("ec2", region_name=configuration.region) elb_client = self.session.client("elbv2", region_name=configuration.region) + base_tags = { + "owner": "dstack", + "dstack_project": configuration.project_name, + "dstack_name": configuration.gateway_name, + } + if settings.DSTACK_VERSION is not None: + base_tags["dstack_version"] = settings.DSTACK_VERSION + tags = merge_tags( + base_tags=base_tags, + backend_tags=self.config.tags, + resource_tags=configuration.tags, + ) + tags = aws_resources.filter_invalid_tags(tags) + tags = aws_resources.make_tags(tags) + + vpc_id, subnets_ids = self._get_vpc_id_subnets_ids_or_error( + ec2_client=ec2_client, + config=self.config, + region=configuration.region, + allocate_public_ip=configuration.public_ip, + ) + security_group_id = aws_resources.create_gateway_security_group( + ec2_client=ec2_client, + project_id=configuration.project_name, + vpc_id=vpc_id, + ) lb_subnets_ids = self._get_gateway_lb_subnets_ids( ec2_client=ec2_client, region=configuration.region, subnets_ids=subnets_ids ) @@ -606,7 +644,7 @@ def create_gateway( # Using short names as LB and target groups have length limit of 32. resources_name_prefix = generate_unique_short_backend_name() - logger.debug("Creating ALB for gateway %s...", configuration.instance_name) + logger.debug("Creating ALB for gateway %s...", configuration.gateway_name) response = elb_client.create_load_balancer( Name=f"{resources_name_prefix}-lb", Subnets=lb_subnets_ids, @@ -619,9 +657,9 @@ def create_gateway( lb = response["LoadBalancers"][0] lb_arn = lb["LoadBalancerArn"] lb_dns_name = lb["DNSName"] - logger.debug("Created ALB for gateway %s.", configuration.instance_name) + logger.debug("Created ALB for gateway %s.", configuration.gateway_name) - logger.debug("Creating Target Group for gateway %s...", configuration.instance_name) + logger.debug("Creating Target Group for gateway %s...", configuration.gateway_name) response = elb_client.create_target_group( Name=f"{resources_name_prefix}-tg", Protocol="HTTP", @@ -630,18 +668,9 @@ def create_gateway( TargetType="instance", ) tg_arn = response["TargetGroups"][0]["TargetGroupArn"] - logger.debug("Created Target Group for gateway %s", configuration.instance_name) - - logger.debug("Registering ALB target for gateway %s...", configuration.instance_name) - elb_client.register_targets( - TargetGroupArn=tg_arn, - Targets=[ - {"Id": instance.instance_id, "Port": 80}, - ], - ) - logger.debug("Registered ALB target for gateway %s", configuration.instance_name) + logger.debug("Created Target Group for gateway %s", configuration.gateway_name) - logger.debug("Creating HTTPS ALB listener for gateway %s...", configuration.instance_name) + logger.debug("Creating HTTPS ALB listener for gateway %s...", configuration.gateway_name) response = elb_client.create_listener( LoadBalancerArn=lb_arn, Protocol="HTTPS", @@ -658,9 +687,9 @@ def create_gateway( ], ) listener_arn = response["Listeners"][0]["ListenerArn"] - logger.debug("Created HTTPS ALB listener for gateway %s", configuration.instance_name) + logger.debug("Created HTTPS ALB listener for gateway %s", configuration.gateway_name) - logger.debug("Creating HTTP ALB listener for gateway %s...", configuration.instance_name) + logger.debug("Creating HTTP ALB listener for gateway %s...", configuration.gateway_name) response = elb_client.create_listener( LoadBalancerArn=lb_arn, Protocol="HTTP", @@ -677,13 +706,9 @@ def create_gateway( ], ) http_listener_arn = response["Listeners"][0]["ListenerArn"] - logger.debug("Created HTTP ALB listener for gateway %s", configuration.instance_name) + logger.debug("Created HTTP ALB listener for gateway %s", configuration.gateway_name) - ip_address = _get_instance_ip(instance, configuration.public_ip) - return GatewayProvisioningData( - instance_id=instance.instance_id, - region=configuration.region, - ip_address=ip_address, + return GatewayLoadBalancerData( hostname=lb_dns_name, backend_data=AWSGatewayBackendData( lb_arn=lb_arn, @@ -704,34 +729,112 @@ def terminate_gateway( region=configuration.region, backend_data=None, ) - if configuration.certificate is None or configuration.certificate.type != "acm": - return + def terminate_gateway_load_balancer( + self, + configuration: GatewayLoadBalancerConfiguration, + backend_data: Optional[str], + ) -> None: if backend_data is None: logger.error( - "Failed to terminate all gateway %s resources. backend_data is None.", - configuration.instance_name, + "Failed to terminate load balancer for gateway %s: backend_data is None.", + configuration.gateway_name, ) return - try: - backend_data_parsed = AWSGatewayBackendData.parse_raw(backend_data) + backend_data_parsed = AWSGatewayBackendData.__response__.parse_raw(backend_data) except ValidationError: logger.exception( - "Failed to terminate all gateway %s resources. backend_data parsing error.", - configuration.instance_name, + "Failed to terminate load balancer for gateway %s: backend_data parsing error.", + configuration.gateway_name, ) return elb_client = self.session.client("elbv2", region_name=configuration.region) - logger.debug("Deleting ALB resources for gateway %s...", configuration.instance_name) + logger.debug("Deleting ALB resources for gateway %s...", configuration.gateway_name) if backend_data_parsed.http_listener_arn is not None: elb_client.delete_listener(ListenerArn=backend_data_parsed.http_listener_arn) elb_client.delete_listener(ListenerArn=backend_data_parsed.listener_arn) elb_client.delete_target_group(TargetGroupArn=backend_data_parsed.tg_arn) elb_client.delete_load_balancer(LoadBalancerArn=backend_data_parsed.lb_arn) - logger.debug("Deleted ALB resources for gateway %s", configuration.instance_name) + logger.debug("Deleted ALB resources for gateway %s.", configuration.gateway_name) + + def register_gateway_replica_with_load_balancer( + self, + instance_id: str, + configuration: GatewayLoadBalancerConfiguration, + gateway_backend_data: Optional[str], + ) -> None: + if gateway_backend_data is None: + raise ComputeError( + f"Cannot register gateway {configuration.gateway_name} replica with load balancer:" + " gateway_backend_data is None" + ) + try: + gateway_backend_data_parsed = AWSGatewayBackendData.__response__.parse_raw( + gateway_backend_data + ) + except ValidationError as e: + raise ComputeError( + f"Cannot register gateway {configuration.gateway_name} replica with load balancer:" + " gateway_backend_data parsing error" + ) from e + + elb_client = self.session.client("elbv2", region_name=configuration.region) + logger.debug( + "Registering gateway %s replica %s with ALB target group %s...", + configuration.gateway_name, + instance_id, + gateway_backend_data_parsed.tg_arn, + ) + elb_client.register_targets( + TargetGroupArn=gateway_backend_data_parsed.tg_arn, + Targets=[{"Id": instance_id, "Port": 80}], + ) + logger.debug( + "Registered gateway %s replica %s with ALB target group.", + configuration.gateway_name, + instance_id, + ) + + def deregister_gateway_replica_from_load_balancer( + self, + instance_id: str, + configuration: GatewayLoadBalancerConfiguration, + gateway_backend_data: Optional[str], + ) -> None: + if gateway_backend_data is None: + raise ComputeError( + f"Cannot deregister gateway {configuration.gateway_name} replica from load balancer:" + " gateway_backend_data is None" + ) + try: + gateway_backend_data_parsed = AWSGatewayBackendData.__response__.parse_raw( + gateway_backend_data + ) + except ValidationError as e: + raise ComputeError( + f"Cannot deregister gateway {configuration.gateway_name} replica from load balancer:" + " gateway_backend_data parsing error", + ) from e + + elb_client = self.session.client("elbv2", region_name=configuration.region) + logger.debug( + "Deregistering gateway %s replica %s from ALB target group %s...", + configuration.gateway_name, + instance_id, + gateway_backend_data_parsed.tg_arn, + ) + elb_client.deregister_targets( + TargetGroupArn=gateway_backend_data_parsed.tg_arn, + Targets=[{"Id": instance_id, "Port": 80}], + ) + logger.debug( + "Deregistered gateway %s replica %s from ALB target group.", + configuration.gateway_name, + instance_id, + ) def register_volume(self, volume: Volume) -> VolumeProvisioningData: assert isinstance(volume.configuration, AWSVolumeConfiguration) diff --git a/src/dstack/_internal/core/backends/base/compute.py b/src/dstack/_internal/core/backends/base/compute.py index 14455f601..33013abdf 100644 --- a/src/dstack/_internal/core/backends/base/compute.py +++ b/src/dstack/_internal/core/backends/base/compute.py @@ -30,6 +30,8 @@ from dstack._internal.core.models.compute_groups import ComputeGroup, ComputeGroupProvisioningData from dstack._internal.core.models.gateways import ( GatewayComputeConfiguration, + GatewayLoadBalancerConfiguration, + GatewayLoadBalancerData, GatewayProvisioningData, ) from dstack._internal.core.models.instances import ( @@ -578,6 +580,55 @@ def terminate_gateway( pass +class ComputeWithGatewayLoadBalancerSupport(ABC): + """ + Must be subclassed and implemented to support gateways with a load balancer that fronts + all replica instances. + + Backends implementing this mixin must also implement `ComputeWithGatewaySupport`. + """ + + @abstractmethod + def create_gateway_load_balancer( + self, + configuration: GatewayLoadBalancerConfiguration, + ) -> GatewayLoadBalancerData: + """Creates the load balancer for a gateway.""" + pass + + @abstractmethod + def terminate_gateway_load_balancer( + self, + configuration: GatewayLoadBalancerConfiguration, + backend_data: Optional[str], + ) -> None: + """Deletes the load balancer.""" + pass + + @abstractmethod + def register_gateway_replica_with_load_balancer( + self, + instance_id: str, + configuration: GatewayLoadBalancerConfiguration, + gateway_backend_data: Optional[str], + ) -> None: + """Registers a gateway replica instance as a target of the load balancer.""" + pass + + @abstractmethod + def deregister_gateway_replica_from_load_balancer( + self, + instance_id: str, + configuration: GatewayLoadBalancerConfiguration, + gateway_backend_data: Optional[str], + ) -> None: + """Deregisters a gateway replica instance from the load balancer. + + If the replica is not registered, it should not raise errors but return silently. + """ + pass + + class ComputeWithPrivateGatewaySupport: """ Must be subclassed to support private gateways. diff --git a/src/dstack/_internal/core/models/gateways.py b/src/dstack/_internal/core/models/gateways.py index 191bfc81e..ed51d3f71 100644 --- a/src/dstack/_internal/core/models/gateways.py +++ b/src/dstack/_internal/core/models/gateways.py @@ -214,6 +214,20 @@ class GatewayProvisioningData(CoreModel): ip_address: str region: str availability_zone: Optional[str] = None - hostname: Optional[str] = None backend_data: Optional[str] = None """`backend_data` stores backend-specific data in JSON.""" + + +class GatewayLoadBalancerConfiguration(CoreModel): + project_name: str + gateway_name: str + region: str + public_ip: bool + certificate: Optional[AnyGatewayCertificate] = None + tags: Optional[Dict[str, str]] = None + + +class GatewayLoadBalancerData(CoreModel): + hostname: str + backend_data: str + """Backend-specific JSON""" 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 e8fb62291..cfc058793 100644 --- a/src/dstack/_internal/server/background/pipeline_tasks/gateway_replicas.py +++ b/src/dstack/_internal/server/background/pipeline_tasks/gateway_replicas.py @@ -8,7 +8,10 @@ from sqlalchemy.orm import InstrumentedAttribute, joinedload, load_only from sqlalchemy.sql.base import ExecutableOption -from dstack._internal.core.backends.base.compute import ComputeWithGatewaySupport +from dstack._internal.core.backends.base.compute import ( + ComputeWithGatewayLoadBalancerSupport, + ComputeWithGatewaySupport, +) from dstack._internal.core.errors import BackendError, BackendNotAvailable from dstack._internal.core.models.gateways import GatewayReplicaStatus, GatewayStatus from dstack._internal.server.background.pipeline_tasks.base import ( @@ -33,7 +36,11 @@ ) from dstack._internal.server.services import backends as backends_services from dstack._internal.server.services import gateways as gateways_services -from dstack._internal.server.services.gateways import get_gateway_compute_configuration +from dstack._internal.server.services.gateways import ( + get_gateway_compute_configuration, + get_gateway_configuration, + get_gateway_lb_configuration, +) from dstack._internal.server.services.gateways.pool import gateway_connections_pool from dstack._internal.server.services.locking import get_locker from dstack._internal.server.services.logging import fmt @@ -248,7 +255,6 @@ class _GatewayReplicaUpdateMap(ItemUpdateMap, total=False): instance_id: Optional[str] ip_address: Optional[str] region: Optional[str] - hostname: Optional[str] backend_data: Optional[str] @@ -475,7 +481,6 @@ async def _provision_gateway_replica( instance_id=gpd.instance_id, ip_address=gpd.ip_address, region=gpd.region, - hostname=gpd.hostname, backend_data=gpd.backend_data, ) @@ -488,8 +493,20 @@ async def _process_provisioning_item(item: GatewayReplicaPipelineItem): GatewayComputeModel.ip_address, GatewayComputeModel.ssh_private_key, GatewayComputeModel.scale_in, + GatewayComputeModel.instance_id, + GatewayComputeModel.backend_id, + GatewayComputeModel.configuration, ], - gateway_fields=_GATEWAY_FIELDS_MIN, + gateway_fields=_GATEWAY_FIELDS_MIN + + [ + GatewayModel.configuration, + GatewayModel.region, + GatewayModel.wildcard_domain, + GatewayModel.hostname, + GatewayModel.backend_data, + ], + load_backends=True, + load_gateway_backend_type=True, ) if replica_model is None: return @@ -500,25 +517,95 @@ async def _process_provisioning_item(item: GatewayReplicaPipelineItem): if update_map := _mark_terminating_if_needed(gateway_model, replica_model): await _commit_update(item, replica_model, update_map=update_map) return + if _is_legacy_aws_acm_gateway_with_pending_migration(gateway_model): + await _commit_update(item, replica_model, update_map={}) + return error = await _connect_and_configure_gateway_replica(gateway_model, replica_model) - if error is None: - logger.info( - "%s replica %d: running", - fmt(gateway_model), - replica_model.replica_num, - ) - update_map = _GatewayReplicaUpdateMap(status=GatewayReplicaStatus.RUNNING, active=True) - else: + if error is not None: logger.warning( "%s replica %d: provisioning failed: %s", fmt(gateway_model), replica_model.replica_num, error, ) - update_map = _GatewayReplicaUpdateMap( - status=GatewayReplicaStatus.TERMINATING, status_message=error, active=False + await _commit_update( + item, + replica_model, + _GatewayReplicaUpdateMap( + status=GatewayReplicaStatus.TERMINATING, status_message=error, active=False + ), ) - await _commit_update(item, replica_model, update_map) + return + + if gateway_model.hostname is not None: + reg_error = await _register_replica_with_load_balancer(gateway_model, replica_model) + if reg_error is not None: + logger.warning( + "%s replica %d: failed to register with load balancer: %s", + fmt(gateway_model), + replica_model.replica_num, + reg_error, + ) + await _commit_update( + item, + replica_model, + _GatewayReplicaUpdateMap( + status=GatewayReplicaStatus.TERMINATING, + status_message=reg_error, + active=False, + ), + ) + return + + logger.info("%s replica %d: running", fmt(gateway_model), replica_model.replica_num) + await _commit_update( + item, + replica_model, + _GatewayReplicaUpdateMap(status=GatewayReplicaStatus.RUNNING, active=True), + ) + + +async def _register_replica_with_load_balancer( + gateway_model: GatewayModel, + replica_model: GatewayComputeModel, +) -> Optional[str]: + """Registers the replica instance with the gateway's load balancer. + Returns an error message on failure, None on success. + """ + if replica_model.instance_id is None: + return "instance_id is None, cannot register with load balancer" + try: + if replica_model.backend_id is None: + raise BackendNotAvailable() + (_, backend) = await backends_services.get_project_backend_with_model_by_id_or_error( + project=gateway_model.project, backend_id=replica_model.backend_id + ) + except BackendNotAvailable: + return "Backend not available" + compute = backend.compute() + if not isinstance(compute, ComputeWithGatewayLoadBalancerSupport): + return "Backend does not support load balancer operations" + lb_configuration = get_gateway_lb_configuration(gateway_model) + try: + await run_async( + compute.register_gateway_replica_with_load_balancer, + replica_model.instance_id, + lb_configuration, + gateway_model.backend_data, + ) + except Exception: + logger.exception( + "%s replica %d: error registering with load balancer", + fmt(gateway_model), + replica_model.replica_num, + ) + return "Error registering with load balancer" + logger.info( + "%s replica %d: registered with load balancer", + fmt(gateway_model), + replica_model.replica_num, + ) + return None async def _connect_and_configure_gateway_replica( @@ -604,6 +691,8 @@ async def _process_terminating_item(item: GatewayReplicaPipelineItem): GatewayModel.configuration, GatewayModel.region, GatewayModel.wildcard_domain, + GatewayModel.hostname, + GatewayModel.backend_data, ], load_backends=True, load_gateway_backend_type=True, @@ -643,6 +732,12 @@ async def _process_terminating_item(item: GatewayReplicaPipelineItem): ) await _commit_update(item, replica_model, mark_terminated_update_map) return + if _is_legacy_aws_acm_gateway_with_pending_migration(gateway_model): + await _commit_update(item, replica_model, update_map={}) + return + + if gateway_model.hostname is not None: + await _deregister_gateway_replica_from_load_balancer(compute, gateway_model, replica_model) logger.debug( "%s replica %d: terminating gateway compute", @@ -675,3 +770,75 @@ async def _process_terminating_item(item: GatewayReplicaPipelineItem): await gateway_connections_pool.remove(replica_model.ip_address) await _commit_update(item, replica_model, mark_terminated_update_map) + + +async def _deregister_gateway_replica_from_load_balancer( + compute: ComputeWithGatewaySupport, + gateway_model: GatewayModel, + replica_model: GatewayComputeModel, +) -> None: + if not isinstance(compute, ComputeWithGatewayLoadBalancerSupport): + logger.error( + ( + "%s replica %d: cannot deregister from load balancer," + " backend does not support load balancer operations" + ), + fmt(gateway_model), + replica_model.replica_num, + ) + return + if replica_model.instance_id is None: + logger.error( + "%s replica %d: cannot deregister from load balancer, instance_id is None", + fmt(gateway_model), + replica_model.replica_num, + ) + return + logger.debug( + "%s replica %d: deregistering from load balancer", + fmt(gateway_model), + replica_model.replica_num, + ) + try: + await run_async( + compute.deregister_gateway_replica_from_load_balancer, + replica_model.instance_id, + get_gateway_lb_configuration(gateway_model), + gateway_model.backend_data, + ) + logger.info( + "%s replica %d: deregistered from load balancer", + fmt(gateway_model), + replica_model.replica_num, + ) + except Exception: + logger.exception( + ( + "%s replica %d: error deregistering from load balancer." + " Proceeding with gateway replica termination," + " relying on automatic deregistration by the load balancer" + ), + fmt(gateway_model), + replica_model.replica_num, + ) + + +def _is_legacy_aws_acm_gateway_with_pending_migration(gateway_model: GatewayModel) -> bool: + """ + If `True`, the gateway cannot be used for replica (de)register operations until the migration + completes, since its `backend_data` does not yet have the relevant load balancer details. + """ + configuration = get_gateway_configuration(gateway_model) + if ( + configuration.certificate is not None + and configuration.certificate.type == "acm" + and gateway_model.hostname is None + ): + logger.warning( + "Found AWS ACM gateway %s without a hostname, which should indicate a pre-0.20.30" + " gateway not yet migrated to the 0.20.30 format. Waiting for the gateway pipeline to" + " perform the migration", + gateway_model.id, + ) + return True + return False diff --git a/src/dstack/_internal/server/background/pipeline_tasks/gateways.py b/src/dstack/_internal/server/background/pipeline_tasks/gateways.py index 8e40c405f..3a50add43 100644 --- a/src/dstack/_internal/server/background/pipeline_tasks/gateways.py +++ b/src/dstack/_internal/server/background/pipeline_tasks/gateways.py @@ -3,12 +3,14 @@ import uuid from dataclasses import dataclass, field from datetime import timedelta -from typing import Sequence +from typing import Optional, Sequence from sqlalchemy import ColumnElement, and_, delete, func, or_, select, update from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import joinedload, load_only, selectinload +from dstack._internal.core.backends.base.compute import ComputeWithGatewayLoadBalancerSupport +from dstack._internal.core.errors import BackendError, BackendNotAvailable from dstack._internal.core.models.gateways import ( GATEWAY_REPLICAS_DEFAULT, GatewayReplicaStatus, @@ -36,17 +38,19 @@ GatewayModel, ProjectModel, ) +from dstack._internal.server.services import backends as backends_services from dstack._internal.server.services import events from dstack._internal.server.services import gateways as gateways_services from dstack._internal.server.services.gateways import ( emit_gateway_status_change_event, get_gateway_compute_models, + get_gateway_lb_configuration, ) 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.utils import tracing -from dstack._internal.utils.common import get_current_datetime, get_lowest_unused_nums +from dstack._internal.utils.common import get_current_datetime, get_lowest_unused_nums, run_async from dstack._internal.utils.logging import get_logger logger = get_logger(__name__) @@ -153,6 +157,18 @@ async def fetch(self, limit: int) -> list[GatewayPipelineItem]: .correlate(GatewayModel) .scalar_subquery() ) + unmigrated_hostname_exists_subquery = ( + select(GatewayComputeModel.id) + .where( + or_( + GatewayComputeModel.gateway_id == GatewayModel.id, + GatewayComputeModel.id == GatewayModel.gateway_compute_id, + ), + GatewayComputeModel.hostname_deprecated_readonly.is_not(None), + ) + .correlate(GatewayModel) + .exists() + ) res = await session.execute( select(GatewayModel) .where( @@ -162,13 +178,23 @@ async def fetch(self, limit: int) -> list[GatewayPipelineItem]: ), and_( GatewayModel.status == GatewayStatus.RUNNING, - GatewayModel.desired_replica_count.is_not(None), or_( - # fetch to reconcile replica count - GatewayModel.desired_replica_count - != active_replica_count_subquery, - # fetch to potentially reset attempts - GatewayModel.replica_scale_attempt > 0, + and_( + GatewayModel.desired_replica_count.is_not(None), + or_( + # fetch to reconcile replica count + GatewayModel.desired_replica_count + != active_replica_count_subquery, + # fetch to potentially reset attempts + GatewayModel.replica_scale_attempt > 0, + ), + ), + # fetch pre-0.20.30 AWS ACM gateways to migrate their hostname + # and backend_data onto GatewayModel + and_( + GatewayModel.hostname.is_(None), + unmigrated_hostname_exists_subquery, + ), ), ), GatewayModel.to_be_deleted == True, @@ -250,9 +276,11 @@ async def process(self, item: GatewayPipelineItem): class _GatewayUpdateMap(ItemUpdateMap, total=False): status: GatewayStatus - status_message: str + status_message: Optional[str] replica_scale_attempt: int last_replica_scale_attempt_at: UpdateMapDateTime + hostname: Optional[str] + backend_data: Optional[str] @dataclass @@ -273,6 +301,7 @@ async def _process_submitted_item(item: GatewayPipelineItem): GatewayModel.lock_token == item.lock_token, ) .options(joinedload(GatewayModel.project).load_only(ProjectModel.name)) + .options(joinedload(GatewayModel.project).selectinload(ProjectModel.backends)) .options(joinedload(GatewayModel.backend).load_only(BackendModel.type)) ) gateway_model = res.unique().scalar_one_or_none() @@ -318,10 +347,69 @@ class _SubmittedResult: async def _process_submitted_gateway(gateway_model: GatewayModel) -> _SubmittedResult: - # NOTE: On a later stage of #3959, the SUBMITTED status may also be responsible for - # setting up the load balancer (e.g., AWS ALB) before replicas are created. + configuration = gateways_services.get_gateway_configuration(gateway_model) + update_map: _GatewayUpdateMap = {} + if configuration.certificate is not None and configuration.certificate.type == "acm": + try: + ( + _, + backend, + ) = await backends_services.get_project_backend_with_model_by_type_or_error( + project=gateway_model.project, backend_type=configuration.backend + ) + except BackendNotAvailable: + return _SubmittedResult( + update_map={ + "status": GatewayStatus.FAILED, + "status_message": "Backend not available", + } + ) + compute = backend.compute() + if not isinstance(compute, ComputeWithGatewayLoadBalancerSupport): + logger.error( + "%s: backend does not support load balancer operations", fmt(gateway_model) + ) + return _SubmittedResult( + update_map={ + "status": GatewayStatus.FAILED, + "status_message": "Backend does not support load balancer operations", + } + ) + lb_configuration = get_gateway_lb_configuration(gateway_model) + logger.info("%s: creating gateway load balancer...", fmt(gateway_model)) + try: + backend_data = await run_async(compute.create_gateway_load_balancer, lb_configuration) + except BackendError as e: + status_message = f"Backend error: {repr(e)}" + if len(e.args) > 0: + status_message = str(e.args[0]) + logger.warning( + "%s: failed to create gateway load balancer: %s", + fmt(gateway_model), + status_message, + ) + return _SubmittedResult( + update_map={ + "status": GatewayStatus.FAILED, + "status_message": status_message, + } + ) + except Exception: + logger.exception( + "%s: unexpected error when creating gateway load balancer", fmt(gateway_model) + ) + return _SubmittedResult( + update_map={ + "status": GatewayStatus.FAILED, + "status_message": "Unexpected error when creating load balancer", + } + ) + logger.info("%s: gateway load balancer created.", fmt(gateway_model)) + update_map["hostname"] = backend_data.hostname + update_map["backend_data"] = backend_data.backend_data + scale_result = _reconcile_gateway_replica_count(gateway_model, gateway_replicas=[]) - update_map = _GatewayUpdateMap(status=GatewayStatus.PROVISIONING) + update_map["status"] = GatewayStatus.PROVISIONING update_map.update(scale_result.gateway_update_map) return _SubmittedResult( update_map=update_map, @@ -347,6 +435,8 @@ async def _process_provisioning_item(item: GatewayPipelineItem): GatewayComputeModel.replica_num, GatewayComputeModel.created_at, GatewayComputeModel.scale_in, + GatewayComputeModel.hostname_deprecated_readonly, + GatewayComputeModel.backend_data, ) ) ) @@ -403,17 +493,22 @@ def _process_provisioning_gateway(gateway_model: GatewayModel) -> _ProvisioningR for gc in gateway_computes if not gc.scale_in and gc.id not in scale_result.scale_in_replica_ids } + update_map = _migrate_hostname_and_backend_data_from_legacy_replica( + gateway_model, gateway_computes + ) if statuses & {GatewayReplicaStatus.TERMINATING, GatewayReplicaStatus.TERMINATED}: - return _ProvisioningResult( - gateway_update_map={ + update_map.update( + { "status": GatewayStatus.FAILED, "status_message": "Failed to provision gateway replica", - }, + } + ) + return _ProvisioningResult( + gateway_update_map=update_map, scale_result=_ReplicaScalingResult(), # do not scale, gateway failed ) - update_map = _GatewayUpdateMap() update_map.update(scale_result.gateway_update_map) if statuses == {GatewayReplicaStatus.RUNNING} and not scale_result.needs_more_replicas: @@ -448,6 +543,8 @@ async def _process_running_item(item: GatewayPipelineItem): GatewayComputeModel.replica_num, GatewayComputeModel.created_at, GatewayComputeModel.scale_in, + GatewayComputeModel.hostname_deprecated_readonly, + GatewayComputeModel.backend_data, ) ) ) @@ -460,6 +557,9 @@ async def _process_running_item(item: GatewayPipelineItem): scale_result = _reconcile_gateway_replica_count(gateway_model, gateway_computes) update_map = _GatewayUpdateMap() + update_map.update( + _migrate_hostname_and_backend_data_from_legacy_replica(gateway_model, gateway_computes) + ) update_map.update(scale_result.gateway_update_map) set_processed_update_map_fields(update_map) set_unlock_update_map_fields(update_map) @@ -490,10 +590,15 @@ async def _process_to_be_deleted_item(item: GatewayPipelineItem): GatewayModel.id == item.id, GatewayModel.lock_token == item.lock_token, ) + .options(joinedload(GatewayModel.project).joinedload(ProjectModel.backends)) + .options(joinedload(GatewayModel.backend).load_only(BackendModel.type)) .options(joinedload(GatewayModel.gateway_compute)) .options( selectinload(GatewayModel.gateway_computes).load_only( - GatewayComputeModel.id, GatewayComputeModel.status + GatewayComputeModel.id, + GatewayComputeModel.status, + GatewayComputeModel.hostname_deprecated_readonly, + GatewayComputeModel.backend_data, ) ) ) @@ -502,7 +607,8 @@ async def _process_to_be_deleted_item(item: GatewayPipelineItem): log_lock_token_mismatch(logger, item) return - result = _process_to_be_deleted_gateway(gateway_model) + result = await _process_to_be_deleted_gateway(gateway_model) + async with get_session_ctx() as session: if result.delete_gateway: res = await session.execute( @@ -529,7 +635,7 @@ async def _process_to_be_deleted_item(item: GatewayPipelineItem): targets=[events.Target.from_model(gateway_model)], ) else: - update_map = _GatewayUpdateMap() + update_map = result.update_map set_processed_update_map_fields(update_map) set_unlock_update_map_fields(update_map) resolve_now_placeholders(update_map, now=get_current_datetime()) @@ -551,12 +657,54 @@ async def _process_to_be_deleted_item(item: GatewayPipelineItem): @dataclass class _ProcessToBeDeletedResult: delete_gateway: bool + update_map: _GatewayUpdateMap = field(default_factory=_GatewayUpdateMap) -def _process_to_be_deleted_gateway(gateway_model: GatewayModel) -> _ProcessToBeDeletedResult: +async def _process_to_be_deleted_gateway(gateway_model: GatewayModel) -> _ProcessToBeDeletedResult: gateway_computes = get_gateway_compute_models(gateway_model) - all_terminated = all(gc.status == GatewayReplicaStatus.TERMINATED for gc in gateway_computes) - return _ProcessToBeDeletedResult(delete_gateway=all_terminated) + if update_map := _migrate_hostname_and_backend_data_from_legacy_replica( + gateway_model, gateway_computes + ): + return _ProcessToBeDeletedResult( + delete_gateway=False, + update_map=update_map, + ) + all_replicas_terminated = all( + gc.status == GatewayReplicaStatus.TERMINATED for gc in gateway_computes + ) + lb_terminated = True + if all_replicas_terminated and gateway_model.hostname is not None: + lb_terminated = await _terminate_gateway_load_balancer(gateway_model) + return _ProcessToBeDeletedResult(delete_gateway=all_replicas_terminated and lb_terminated) + + +async def _terminate_gateway_load_balancer(gateway_model: GatewayModel) -> bool: + """Terminates the gateway's load balancer. Returns True on success, False on failure.""" + configuration = gateways_services.get_gateway_configuration(gateway_model) + try: + (_, backend) = await backends_services.get_project_backend_with_model_by_type_or_error( + project=gateway_model.project, backend_type=configuration.backend + ) + except BackendNotAvailable: + logger.error( + "%s: backend not available, cannot terminate load balancer", fmt(gateway_model) + ) + return False + compute = backend.compute() + if not isinstance(compute, ComputeWithGatewayLoadBalancerSupport): + logger.error("%s: backend does not support load balancer operations", fmt(gateway_model)) + return False + lb_configuration = get_gateway_lb_configuration(gateway_model) + logger.info("%s: terminating gateway load balancer...", fmt(gateway_model)) + try: + await run_async( + compute.terminate_gateway_load_balancer, lb_configuration, gateway_model.backend_data + ) + except Exception: + logger.exception("%s: error when terminating gateway load balancer", fmt(gateway_model)) + return False + logger.info("%s: gateway load balancer terminated.", fmt(gateway_model)) + return True REPLICA_SCALE_IN_PRIORITY: dict[GatewayReplicaStatus, int] = { @@ -736,3 +884,31 @@ async def _apply_replica_scaling( actor=events.SystemActor(), targets=[events.Target.from_model(gateway_model)], ) + + +def _migrate_hostname_and_backend_data_from_legacy_replica( + gateway_model: GatewayModel, + gateway_computes: list[GatewayComputeModel], +) -> _GatewayUpdateMap: + """ + Move `hostname` and `backend_data` from pre-0.20.30 GatewayComputeModel onto GatewayModel. + + Alembic migration ecc9e8a0bfac does the same thing. This function is a fallback in case any + gateways are created by an older server replica after the migration passes. + """ + if gateway_model.hostname is not None: + return {} + for gateway_compute in gateway_computes: + if gateway_compute.hostname_deprecated_readonly is not None: + update_map: _GatewayUpdateMap = { + "hostname": gateway_compute.hostname_deprecated_readonly, + # Pre-0.20.30 AWS ACM gateways used GatewayComputeModel.backend_data exclusively + # for load-balancer related fields, and not for gateway replica instance fields. + # So GatewayComputeModel.backend_data is copied entirely. + "backend_data": gateway_compute.backend_data, + } + logger.info( + "%s: migrating hostname and backend_data onto GatewayModel", fmt(gateway_model) + ) + return update_map + return {} diff --git a/src/dstack/_internal/server/migrations/versions/2026/07_22_1953_ecc9e8a0bfac_gateways_hostname_and_backend_data.py b/src/dstack/_internal/server/migrations/versions/2026/07_22_1953_ecc9e8a0bfac_gateways_hostname_and_backend_data.py new file mode 100644 index 000000000..9e4b3b07d --- /dev/null +++ b/src/dstack/_internal/server/migrations/versions/2026/07_22_1953_ecc9e8a0bfac_gateways_hostname_and_backend_data.py @@ -0,0 +1,80 @@ +"""Gateways hostname and backend_data + +Revision ID: ecc9e8a0bfac +Revises: dd83c131e78f +Create Date: 2026-07-22 19:53:45.438752+00:00 + +""" + +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision = "ecc9e8a0bfac" +down_revision = "dd83c131e78f" +branch_labels = None +depends_on = None + +# Partial table definitions for queries +gateways = sa.table( + "gateways", + sa.column("id"), + sa.column("gateway_compute_id"), + sa.column("hostname", sa.String(255)), + sa.column("backend_data", sa.Text()), +) +gateway_computes = sa.table( + "gateway_computes", + sa.column("id"), + sa.column("gateway_id"), + sa.column("hostname", sa.String(100)), + sa.column("backend_data", sa.Text()), +) + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + with op.batch_alter_table("gateways", schema=None) as batch_op: + batch_op.add_column(sa.Column("hostname", sa.String(length=255), nullable=True)) + batch_op.add_column(sa.Column("backend_data", sa.Text(), nullable=True)) + + # ### end Alembic commands ### + + # Backfill from gateway_computes. + # Only for AWS ACM gateways, as indicated by `hostname` being present. + related_compute_with_hostname_filter = sa.and_( + gateway_computes.c.hostname.is_not(None), + sa.or_( + gateway_computes.c.gateway_id == gateways.c.id, + gateway_computes.c.id == gateways.c.gateway_compute_id, + ), + ) + op.execute( + sa.update(gateways) + .where( + sa.select(gateway_computes.c.id).where(related_compute_with_hostname_filter).exists() + ) + .values( + hostname=( + sa.select(gateway_computes.c.hostname) + .where(related_compute_with_hostname_filter) + .limit(1) + .scalar_subquery() + ), + backend_data=( + sa.select(gateway_computes.c.backend_data) + .where(related_compute_with_hostname_filter) + .limit(1) + .scalar_subquery() + ), + ) + ) + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + with op.batch_alter_table("gateways", schema=None) as batch_op: + batch_op.drop_column("backend_data") + batch_op.drop_column("hostname") + + # ### end Alembic commands ### diff --git a/src/dstack/_internal/server/models.py b/src/dstack/_internal/server/models.py index 68903caec..01abeea86 100644 --- a/src/dstack/_internal/server/models.py +++ b/src/dstack/_internal/server/models.py @@ -632,6 +632,14 @@ class GatewayModel(PipelineModelMixin, BaseModel): backend_id: Mapped[uuid.UUID] = mapped_column(ForeignKey("backends.id", ondelete="CASCADE")) backend: Mapped["BackendModel"] = relationship() + hostname: Mapped[Optional[str]] = mapped_column(String(255)) + """Hostname of the gateway's load balancer (e.g. ALB domain name for AWS ACM gateways). + Unset when there is no load balancer. + """ + backend_data: Mapped[Optional[str]] = mapped_column(Text) + """Backend-specific load balancer resource data in JSON. + """ + gateway_compute_id: Mapped[Optional[uuid.UUID]] = mapped_column( ForeignKey("gateway_computes.id", ondelete="CASCADE") ) @@ -682,10 +690,8 @@ class GatewayComputeModel(PipelineModelMixin, BaseModel): """Gateway replica IP address or domain name (e.g., k8s can use domain names). **TODO**: rename. """ - hostname: Mapped[Optional[str]] = mapped_column(String(100)) - """Hostname of the gateway's load balancer. - **TODO**: move to `GatewayModel`. - """ + hostname_deprecated_readonly: Mapped[Optional[str]] = mapped_column("hostname", String(100)) + """Replaced by GatewayModel.hostname since 0.20.30""" configuration: Mapped[Optional[str]] = mapped_column(Text) """`configuration` is optional for compatibility with pre-0.18.2 gateways. Use `get_gateway_compute_configuration` to construct `configuration` for old gateways. diff --git a/src/dstack/_internal/server/services/gateways/__init__.py b/src/dstack/_internal/server/services/gateways/__init__.py index 36378f30d..4cf2f3073 100644 --- a/src/dstack/_internal/server/services/gateways/__init__.py +++ b/src/dstack/_internal/server/services/gateways/__init__.py @@ -37,6 +37,7 @@ Gateway, GatewayComputeConfiguration, GatewayConfiguration, + GatewayLoadBalancerConfiguration, GatewayPlan, GatewayReplica, GatewayReplicaStatus, @@ -884,6 +885,20 @@ def get_gateway_compute_configuration( ) +def get_gateway_lb_configuration( + gateway_model: GatewayModel, +) -> GatewayLoadBalancerConfiguration: + configuration = get_gateway_configuration(gateway_model) + return GatewayLoadBalancerConfiguration( + project_name=gateway_model.project.name, + gateway_name=gateway_model.name, + region=configuration.region, + public_ip=configuration.public_ip, + certificate=configuration.certificate, + tags=configuration.tags, + ) + + def gateway_model_to_gateway( gateway_model: GatewayModel, default_gateway_id: Optional[uuid.UUID] ) -> Gateway: @@ -905,7 +920,6 @@ def gateway_model_to_gateway( all_compute_models, key=lambda c: c.replica_num ): relevant_compute_models.append(max(compute_models_for_num, key=lambda c: c.created_at)) - gateway_hostname = None replicas = [] for compute in relevant_compute_models: replicas.append( @@ -919,13 +933,12 @@ def gateway_model_to_gateway( status_message=compute.status_message, ) ) - gateway_hostname = compute.hostname return Gateway( id=gateway_model.id, name=gateway_model.name, project_name=gateway_model.project.name, - hostname=gateway_hostname, + hostname=gateway_model.hostname, backend=gateway_model.backend.type, region=gateway_model.region, wildcard_domain=gateway_model.wildcard_domain, @@ -1129,11 +1142,16 @@ def _validate_gateway_configuration(configuration: GatewayConfiguration): ) if configuration.certificate.type == "acm" and configuration.backend != BackendType.AWS: raise ServerClientError("acm certificate type is supported for aws backend only") - if replicas > 1: - raise ServerClientError( - "Replicated gateways do not support certificates." - " Set either `certificate: null` or `replicas: 1` in the gateway configuration" + if configuration.certificate.type == "lets-encrypt" and replicas > 1: + err = ( + "The `lets-encrypt` certificate type is not supported for gateways with `replicas`" + " greater than `1`. To create a replicated gateway, set the `certificate`" + " configuration property to one of the supported values, such as" + " `certificate: null` (no HTTPS)" ) + if configuration.backend == BackendType.AWS: + err += " or `certificate: { type: acm, arn: }` (AWS ACM)" + raise ServerClientError(err) if configuration.router is not None and replicas > 1: raise ServerClientError( diff --git a/src/tests/_internal/server/routers/test_gateways.py b/src/tests/_internal/server/routers/test_gateways.py index 26b0efe55..c9c7c9638 100644 --- a/src/tests/_internal/server/routers/test_gateways.py +++ b/src/tests/_internal/server/routers/test_gateways.py @@ -730,22 +730,6 @@ async def test_create_gateway_with_invalid_domain_interpolation( "Cannot interpolate gateway domain name: Failed to interpolate due to missing vars: ['run.unknown_variable']", id="invalid-domain-interpolation", ), - pytest.param( - { - "type": "gateway", - "name": "test", - "backend": "aws", - "region": "us", - "certificate": { - "type": "acm", - "arn": "arn:aws:acm:us-east-1:123456789:certificate/abc", - }, - "replicas": 2, - }, - "Replicated gateways do not support certificates." - " Set either `certificate: null` or `replicas: 1` in the gateway configuration", - id="multi-replica-with-acm-cert", - ), pytest.param( { "type": "gateway", @@ -755,8 +739,11 @@ async def test_create_gateway_with_invalid_domain_interpolation( "certificate": {"type": "lets-encrypt"}, "replicas": 2, }, - "Replicated gateways do not support certificates." - " Set either `certificate: null` or `replicas: 1` in the gateway configuration", + "The `lets-encrypt` certificate type is not supported for gateways with `replicas`" + " greater than `1`. To create a replicated gateway, set the `certificate`" + " configuration property to one of the supported values, such as" + " `certificate: null` (no HTTPS)" + " or `certificate: { type: acm, arn: }` (AWS ACM)", id="multi-replica-with-letsencrypt-cert", ), pytest.param( From bb63aefa8049628a5d6dbb3335ee9f7df117fcff Mon Sep 17 00:00:00 2001 From: Jvst Me Date: Wed, 29 Jul 2026 16:36:11 +0200 Subject: [PATCH 2/3] Add unit tests --- src/dstack/_internal/server/testing/common.py | 14 + .../pipeline_tasks/test_gateway_replicas.py | 409 +++++++++++++ .../pipeline_tasks/test_gateways.py | 543 +++++++++++++++++- 3 files changed, 965 insertions(+), 1 deletion(-) diff --git a/src/dstack/_internal/server/testing/common.py b/src/dstack/_internal/server/testing/common.py index 4371175e5..09da5af3b 100644 --- a/src/dstack/_internal/server/testing/common.py +++ b/src/dstack/_internal/server/testing/common.py @@ -14,6 +14,7 @@ from dstack._internal.core.backends.base.compute import ( Compute, ComputeWithCreateInstanceSupport, + ComputeWithGatewayLoadBalancerSupport, ComputeWithGatewaySupport, ComputeWithGroupProvisioningSupport, ComputeWithInstanceVolumesSupport, @@ -46,10 +47,12 @@ ) from dstack._internal.core.models.gateways import ( GATEWAY_REPLICAS_DEFAULT, + AnyGatewayCertificate, GatewayComputeConfiguration, GatewayConfiguration, GatewayReplicaStatus, GatewayStatus, + LetsEncryptGatewayCertificate, ) from dstack._internal.core.models.health import HealthStatus from dstack._internal.core.models.instances import ( @@ -649,6 +652,9 @@ async def create_gateway( last_processed_at: datetime = datetime(2023, 1, 2, 3, 4, tzinfo=timezone.utc), forbid_new_services: bool = False, populate_configuration: bool = True, + certificate: Optional[AnyGatewayCertificate] = LetsEncryptGatewayCertificate(), + hostname: Optional[str] = None, + backend_data: Optional[str] = None, ) -> GatewayModel: """ Args: @@ -666,6 +672,7 @@ async def create_gateway( region=region, domain=wildcard_domain, replicas=replicas, + certificate=certificate, ).json() gateway = GatewayModel( project_id=project_id, @@ -678,6 +685,8 @@ async def create_gateway( desired_replica_count=replicas if replicas is not None else GATEWAY_REPLICAS_DEFAULT, last_processed_at=last_processed_at, forbid_new_services=forbid_new_services, + hostname=hostname, + backend_data=backend_data, ) session.add(gateway) await session.commit() @@ -699,6 +708,8 @@ async def create_gateway_compute( active: bool = True, configuration: Optional[str] = None, populate_configuration: bool = True, + hostname_deprecated_readonly: Optional[str] = None, + backend_data: Optional[str] = None, ) -> GatewayComputeModel: """ Args: @@ -735,6 +746,8 @@ async def create_gateway_compute( replica_num=replica_num, active=active, configuration=configuration, + hostname_deprecated_readonly=hostname_deprecated_readonly, + backend_data=backend_data, ) session.add(gateway_compute) await session.commit() @@ -1415,6 +1428,7 @@ class ComputeMockSpec( ComputeWithReservationSupport, ComputeWithPlacementGroupSupport, ComputeWithGatewaySupport, + ComputeWithGatewayLoadBalancerSupport, ComputeWithPrivateGatewaySupport, ComputeWithVolumeSupport, ): 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 638c5ff8d..6665d243d 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 @@ -6,8 +6,10 @@ import pytest from sqlalchemy.ext.asyncio import AsyncSession +from dstack._internal.core.backends.base.compute import ComputeWithGatewaySupport from dstack._internal.core.errors import BackendError from dstack._internal.core.models.gateways import ( + ACMGatewayCertificate, GatewayProvisioningData, GatewayReplicaStatus, GatewayStatus, @@ -716,6 +718,223 @@ async def test_provisioning_to_running( assert compute.status == GatewayReplicaStatus.RUNNING assert compute.active is True + async def test_provisioning_to_running_registers_with_load_balancer( + self, test_db, session: AsyncSession, worker: GatewayReplicaWorker + ): + 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.PROVISIONING, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname="gateway-lb.example.com", + backend_data="lb-backend-data", + ) + compute = await create_gateway_compute( + session=session, + gateway_id=gateway.id, + backend_id=backend.id, + status=GatewayReplicaStatus.PROVISIONING, + ) + _lock_compute(compute) + await session.commit() + + with ( + patch( + "dstack._internal.server.services.gateways.gateway_connections_pool.get_or_add" + ) as pool_add, + patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock, + ): + pool_add.return_value = MagicMock() + pool_add.return_value.client.return_value = MagicMock(AsyncContextManager()) + backend_mock = Mock() + backend_mock.compute.return_value = Mock(spec=ComputeMockSpec) + get_backends_mock.return_value = [(backend, backend_mock)] + + await worker.process(_compute_to_pipeline_item(compute)) + + register_mock = ( + backend_mock.compute.return_value.register_gateway_replica_with_load_balancer + ) + register_mock.assert_called_once() + call_args = register_mock.call_args.args + assert call_args[0] == compute.instance_id + assert call_args[1].gateway_name == gateway.name + assert call_args[2] == "lb-backend-data" + + await session.refresh(compute) + assert compute.status == GatewayReplicaStatus.RUNNING + assert compute.active is True + + async def test_provisioning_skips_load_balancer_registration_without_hostname( + self, test_db, session: AsyncSession, worker: GatewayReplicaWorker + ): + 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.PROVISIONING, + ) + compute = await create_gateway_compute( + session=session, + gateway_id=gateway.id, + backend_id=backend.id, + status=GatewayReplicaStatus.PROVISIONING, + ) + _lock_compute(compute) + await session.commit() + + with ( + patch( + "dstack._internal.server.services.gateways.gateway_connections_pool.get_or_add" + ) as pool_add, + patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock, + ): + pool_add.return_value = MagicMock() + pool_add.return_value.client.return_value = MagicMock(AsyncContextManager()) + + await worker.process(_compute_to_pipeline_item(compute)) + + get_backends_mock.assert_not_called() + + await session.refresh(compute) + assert compute.status == GatewayReplicaStatus.RUNNING + assert compute.active is True + + async def test_provisioning_to_terminating_when_load_balancer_registration_fails( + self, test_db, session: AsyncSession, worker: GatewayReplicaWorker + ): + 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.PROVISIONING, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname="gateway-lb.example.com", + backend_data="lb-backend-data", + ) + compute = await create_gateway_compute( + session=session, + gateway_id=gateway.id, + backend_id=backend.id, + status=GatewayReplicaStatus.PROVISIONING, + ) + _lock_compute(compute) + await session.commit() + + with ( + patch( + "dstack._internal.server.services.gateways.gateway_connections_pool.get_or_add" + ) as pool_add, + patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock, + ): + pool_add.return_value = MagicMock() + pool_add.return_value.client.return_value = MagicMock(AsyncContextManager()) + backend_mock = Mock() + backend_mock.compute.return_value = Mock(spec=ComputeMockSpec) + backend_mock.compute.return_value.register_gateway_replica_with_load_balancer.side_effect = Exception( + "boom" + ) + get_backends_mock.return_value = [(backend, backend_mock)] + + await worker.process(_compute_to_pipeline_item(compute)) + + await session.refresh(compute) + assert compute.status == GatewayReplicaStatus.TERMINATING + assert compute.active is False + assert compute.status_message == "Error registering with load balancer" + + async def test_provisioning_to_terminating_when_backend_does_not_support_load_balancer( + self, test_db, session: AsyncSession, worker: GatewayReplicaWorker + ): + 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.PROVISIONING, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname="gateway-lb.example.com", + backend_data="lb-backend-data", + ) + compute = await create_gateway_compute( + session=session, + gateway_id=gateway.id, + backend_id=backend.id, + status=GatewayReplicaStatus.PROVISIONING, + ) + _lock_compute(compute) + await session.commit() + + with ( + patch( + "dstack._internal.server.services.gateways.gateway_connections_pool.get_or_add" + ) as pool_add, + patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock, + ): + pool_add.return_value = MagicMock() + pool_add.return_value.client.return_value = MagicMock(AsyncContextManager()) + backend_mock = Mock() + backend_mock.compute.return_value = Mock(spec=ComputeWithGatewaySupport) + get_backends_mock.return_value = [(backend, backend_mock)] + + await worker.process(_compute_to_pipeline_item(compute)) + + await session.refresh(compute) + assert compute.status == GatewayReplicaStatus.TERMINATING + assert compute.active is False + assert compute.status_message == "Backend does not support load balancer operations" + + async def test_provisioning_waits_for_pending_acm_gateway_migration( + self, test_db, session: AsyncSession, worker: GatewayReplicaWorker + ): + 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, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname=None, # migration not yet performed + ) + compute = await create_gateway_compute( + session=session, + gateway_id=gateway.id, + backend_id=backend.id, + status=GatewayReplicaStatus.PROVISIONING, + hostname_deprecated_readonly="legacy-lb.example.com", + ) + _lock_compute(compute) + original_last_processed_at = compute.last_processed_at + await session.commit() + + with patch( + "dstack._internal.server.services.gateways.gateway_connections_pool.get_or_add" + ) as pool_add: + await worker.process(_compute_to_pipeline_item(compute)) + pool_add.assert_not_called() + + await session.refresh(compute) + assert compute.status == GatewayReplicaStatus.PROVISIONING + assert compute.last_processed_at > original_last_processed_at + assert compute.lock_token is None + @pytest.mark.parametrize("legacy_compute", [False, True]) async def test_provisioning_to_terminating_if_connect_fails( self, test_db, session: AsyncSession, worker: GatewayReplicaWorker, legacy_compute: bool @@ -917,6 +1136,196 @@ async def test_terminating_to_terminated( assert compute.active is False assert compute.deleted is True + async def test_terminating_deregisters_from_load_balancer_before_terminating( + self, test_db, session: AsyncSession, worker: GatewayReplicaWorker + ): + 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, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname="gateway-lb.example.com", + backend_data="lb-backend-data", + ) + compute = await create_gateway_compute( + session=session, + gateway_id=gateway.id, + backend_id=backend.id, + status=GatewayReplicaStatus.TERMINATING, + active=False, + ) + _lock_compute(compute) + await session.commit() + + with ( + patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock, + patch( + "dstack._internal.server.background.pipeline_tasks.gateway_replicas.gateway_connections_pool.remove" + ), + ): + backend_mock = Mock() + backend_mock.compute.return_value = Mock(spec=ComputeMockSpec) + get_backends_mock.return_value = [(backend, backend_mock)] + + await worker.process(_compute_to_pipeline_item(compute)) + + deregister_mock = ( + backend_mock.compute.return_value.deregister_gateway_replica_from_load_balancer + ) + deregister_mock.assert_called_once() + call_args = deregister_mock.call_args.args + assert call_args[0] == compute.instance_id + assert call_args[1].gateway_name == gateway.name + assert call_args[2] == "lb-backend-data" + backend_mock.compute.return_value.terminate_gateway.assert_called_once() + + await session.refresh(compute) + assert compute.status == GatewayReplicaStatus.TERMINATED + assert compute.active is False + assert compute.deleted is True + + async def test_terminating_proceeds_when_load_balancer_deregistration_raises( + self, test_db, session: AsyncSession, worker: GatewayReplicaWorker + ): + 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, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname="gateway-lb.example.com", + backend_data="lb-backend-data", + ) + compute = await create_gateway_compute( + session=session, + gateway_id=gateway.id, + backend_id=backend.id, + status=GatewayReplicaStatus.TERMINATING, + active=False, + ) + _lock_compute(compute) + await session.commit() + + with ( + patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock, + patch( + "dstack._internal.server.background.pipeline_tasks.gateway_replicas.gateway_connections_pool.remove" + ), + ): + backend_mock = Mock() + backend_mock.compute.return_value = Mock(spec=ComputeMockSpec) + backend_mock.compute.return_value.deregister_gateway_replica_from_load_balancer.side_effect = Exception( + "boom" + ) + get_backends_mock.return_value = [(backend, backend_mock)] + + await worker.process(_compute_to_pipeline_item(compute)) + + backend_mock.compute.return_value.terminate_gateway.assert_called_once() + deregister_mock = ( + backend_mock.compute.return_value.deregister_gateway_replica_from_load_balancer + ) + deregister_mock.assert_called_once() + + await session.refresh(compute) + # Deregistration failures do not block termination: the load balancer is expected + # to eventually deregister the (now-terminated) target automatically. + assert compute.status == GatewayReplicaStatus.TERMINATED + assert compute.active is False + assert compute.deleted is True + + async def test_terminating_skips_deregistration_when_gateway_has_no_hostname( + self, test_db, session: AsyncSession, worker: GatewayReplicaWorker + ): + 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, + ) + compute = await create_gateway_compute( + session=session, + gateway_id=gateway.id, + backend_id=backend.id, + status=GatewayReplicaStatus.TERMINATING, + active=False, + ) + _lock_compute(compute) + await session.commit() + + with ( + patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock, + patch( + "dstack._internal.server.background.pipeline_tasks.gateway_replicas.gateway_connections_pool.remove" + ), + ): + backend_mock = Mock() + backend_mock.compute.return_value = Mock(spec=ComputeMockSpec) + get_backends_mock.return_value = [(backend, backend_mock)] + + await worker.process(_compute_to_pipeline_item(compute)) + + backend_mock.compute.return_value.deregister_gateway_replica_from_load_balancer.assert_not_called() + backend_mock.compute.return_value.terminate_gateway.assert_called_once() + + await session.refresh(compute) + assert compute.status == GatewayReplicaStatus.TERMINATED + + async def test_terminating_waits_for_pending_acm_gateway_migration( + self, test_db, session: AsyncSession, worker: GatewayReplicaWorker + ): + 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, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname=None, # migration not yet performed by the gateway pipeline + ) + compute = await create_gateway_compute( + session=session, + gateway_id=gateway.id, + backend_id=backend.id, + status=GatewayReplicaStatus.TERMINATING, + active=False, + hostname_deprecated_readonly="legacy-lb.example.com", + ) + _lock_compute(compute) + original_last_processed_at = compute.last_processed_at + await session.commit() + + with patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock: + backend_mock = Mock() + backend_mock.compute.return_value = Mock(spec=ComputeMockSpec) + get_backends_mock.return_value = [(backend, backend_mock)] + + await worker.process(_compute_to_pipeline_item(compute)) + + backend_mock.compute.return_value.terminate_gateway.assert_not_called() + backend_mock.compute.return_value.deregister_gateway_replica_from_load_balancer.assert_not_called() + + await session.refresh(compute) + assert compute.status == GatewayReplicaStatus.TERMINATING + assert compute.last_processed_at > original_last_processed_at + assert compute.lock_token is None + @pytest.mark.parametrize("legacy_compute", [False, True]) async def test_terminating_to_terminated_if_backend_not_available( self, test_db, session: AsyncSession, worker: GatewayReplicaWorker, legacy_compute: bool diff --git a/src/tests/_internal/server/background/pipeline_tasks/test_gateways.py b/src/tests/_internal/server/background/pipeline_tasks/test_gateways.py index 452e68628..9de2f6acb 100644 --- a/src/tests/_internal/server/background/pipeline_tasks/test_gateways.py +++ b/src/tests/_internal/server/background/pipeline_tasks/test_gateways.py @@ -1,14 +1,19 @@ import asyncio import uuid from datetime import datetime, timedelta, timezone -from unittest.mock import Mock +from unittest.mock import Mock, patch import pytest from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import selectinload +from dstack._internal.core.backends.base.compute import ComputeWithGatewaySupport +from dstack._internal.core.errors import BackendError +from dstack._internal.core.models.backends.base import BackendType from dstack._internal.core.models.gateways import ( + ACMGatewayCertificate, + GatewayLoadBalancerData, GatewayReplicaStatus, GatewayStatus, ) @@ -22,6 +27,7 @@ from dstack._internal.server.models import GatewayComputeModel, GatewayModel from dstack._internal.server.services.gateways import get_gateway_compute_models from dstack._internal.server.testing.common import ( + ComputeMockSpec, create_backend, create_gateway, create_gateway_compute, @@ -349,6 +355,99 @@ async def test_fetch_includes_running_gateway_when_replica_count_not_matches( items = await fetcher.fetch(limit=10) assert {item.id for item in items} == {gateway.id} + @pytest.mark.parametrize("legacy_compute", [False, True]) + async def test_fetch_includes_running_gateway_with_unmigrated_legacy_hostname( + self, test_db, session: AsyncSession, fetcher: GatewayFetcher, legacy_compute: bool + ): + project = await create_project(session=session) + backend = await create_backend(session=session, project_id=project.id) + stale = get_current_datetime() - timedelta(minutes=1) + gateway = await create_gateway( + session=session, + project_id=project.id, + backend_id=backend.id, + status=GatewayStatus.RUNNING, + replicas=1, + last_processed_at=stale, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname=None, # not yet migrated + ) + if legacy_compute: + compute = await create_gateway_compute( + session=session, + backend_id=backend.id, + status=GatewayReplicaStatus.RUNNING, + hostname_deprecated_readonly="legacy-lb.example.com", + ) + gateway.gateway_compute_id = compute.id + else: + await create_gateway_compute( + session=session, + gateway_id=gateway.id, + status=GatewayReplicaStatus.RUNNING, + hostname_deprecated_readonly="legacy-lb.example.com", + ) + await session.commit() + + # Desired replica count matches the active replica count, so absent the + # pending-migration condition this gateway would not be fetched. + items = await fetcher.fetch(limit=10) + assert {item.id for item in items} == {gateway.id} + + async def test_fetch_excludes_running_gateway_without_legacy_hostname_to_migrate( + self, test_db, session: AsyncSession, fetcher: GatewayFetcher + ): + project = await create_project(session=session) + backend = await create_backend(session=session, project_id=project.id) + stale = get_current_datetime() - timedelta(minutes=1) + gateway = await create_gateway( + session=session, + project_id=project.id, + backend_id=backend.id, + status=GatewayStatus.RUNNING, + replicas=1, + last_processed_at=stale, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname=None, + ) + await create_gateway_compute( + session=session, + gateway_id=gateway.id, + status=GatewayReplicaStatus.RUNNING, + hostname_deprecated_readonly=None, + ) + await session.commit() + + items = await fetcher.fetch(limit=10) + assert items == [] + + async def test_fetch_excludes_already_migrated_gateway( + self, test_db, session: AsyncSession, fetcher: GatewayFetcher + ): + project = await create_project(session=session) + backend = await create_backend(session=session, project_id=project.id) + stale = get_current_datetime() - timedelta(minutes=1) + gateway = await create_gateway( + session=session, + project_id=project.id, + backend_id=backend.id, + status=GatewayStatus.RUNNING, + replicas=1, + last_processed_at=stale, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname="gateway-lb.example.com", # already migrated + ) + await create_gateway_compute( + session=session, + gateway_id=gateway.id, + status=GatewayReplicaStatus.RUNNING, + hostname_deprecated_readonly="legacy-lb.example.com", + ) + await session.commit() + + items = await fetcher.fetch(limit=10) + assert items == [] + @pytest.mark.asyncio @pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True) @@ -389,6 +488,144 @@ async def test_submitted_to_provisioning( assert computes[1].replica_num == 1 assert all(c.ip_address is None for c in computes) + async def test_submitted_to_provisioning_creates_load_balancer_for_acm_gateway( + self, test_db, session: AsyncSession, worker: GatewayWorker + ): + 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.SUBMITTED, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + ) + gateway.lock_token = uuid.uuid4() + gateway.lock_expires_at = datetime(2025, 1, 2, 3, 4, tzinfo=timezone.utc) + await session.commit() + + with patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock: + backend_mock = Mock() + backend_mock.TYPE = BackendType.AWS + backend_mock.compute.return_value = Mock(spec=ComputeMockSpec) + backend_mock.compute.return_value.create_gateway_load_balancer.return_value = ( + GatewayLoadBalancerData( + hostname="gateway-lb.example.com", backend_data="lb-backend-data" + ) + ) + get_backends_mock.return_value = [(backend, backend_mock)] + + await worker.process(_gateway_to_pipeline_item(gateway)) + + create_lb_mock = backend_mock.compute.return_value.create_gateway_load_balancer + create_lb_mock.assert_called_once() + assert create_lb_mock.call_args.args[0].gateway_name == gateway.name + + await session.refresh(gateway) + assert gateway.status == GatewayStatus.PROVISIONING + assert gateway.hostname == "gateway-lb.example.com" + assert gateway.backend_data == "lb-backend-data" + + async def test_submitted_skips_load_balancer_creation_for_non_acm_gateway( + self, test_db, session: AsyncSession, worker: GatewayWorker + ): + 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.SUBMITTED, + ) + gateway.lock_token = uuid.uuid4() + gateway.lock_expires_at = datetime(2025, 1, 2, 3, 4, tzinfo=timezone.utc) + await session.commit() + + with patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock: + await worker.process(_gateway_to_pipeline_item(gateway)) + get_backends_mock.assert_not_called() + + await session.refresh(gateway) + assert gateway.status == GatewayStatus.PROVISIONING + assert gateway.hostname is None + assert gateway.backend_data is None + + async def test_submitted_to_failed_when_backend_does_not_support_load_balancer( + self, test_db, session: AsyncSession, worker: GatewayWorker + ): + 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.SUBMITTED, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + ) + gateway.lock_token = uuid.uuid4() + gateway.lock_expires_at = datetime(2025, 1, 2, 3, 4, tzinfo=timezone.utc) + await session.commit() + + with patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock: + backend_mock = Mock() + backend_mock.TYPE = BackendType.AWS + backend_mock.compute.return_value = Mock(spec=ComputeWithGatewaySupport) + get_backends_mock.return_value = [(backend, backend_mock)] + await worker.process(_gateway_to_pipeline_item(gateway)) + + await session.refresh(gateway) + assert gateway.status == GatewayStatus.FAILED + assert gateway.status_message == "Backend does not support load balancer operations" + + @pytest.mark.parametrize( + "exception,expected_status_message", + [ + (BackendError("Quota exceeded"), "Quota exceeded"), + (RuntimeError("boom"), "Unexpected error when creating load balancer"), + ], + ) + async def test_submitted_to_failed_when_load_balancer_creation_raises( + self, + test_db, + session: AsyncSession, + worker: GatewayWorker, + exception: Exception, + expected_status_message: str, + ): + 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.SUBMITTED, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + ) + gateway.lock_token = uuid.uuid4() + gateway.lock_expires_at = datetime(2025, 1, 2, 3, 4, tzinfo=timezone.utc) + await session.commit() + + with patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock: + backend_mock = Mock() + backend_mock.TYPE = BackendType.AWS + backend_mock.compute.return_value = Mock(spec=ComputeMockSpec) + backend_mock.compute.return_value.create_gateway_load_balancer.side_effect = exception + get_backends_mock.return_value = [(backend, backend_mock)] + await worker.process(_gateway_to_pipeline_item(gateway)) + + await session.refresh(gateway) + assert gateway.status == GatewayStatus.FAILED + assert gateway.status_message == expected_status_message + assert gateway.hostname is None + @pytest.mark.asyncio @pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True) @@ -439,6 +676,69 @@ async def test_provisioning_to_running( assert len(events) == 1 assert events[0].message == "Gateway status changed PROVISIONING -> RUNNING" + async def test_provisioning_migrates_hostname_and_backend_data_from_legacy_replica( + self, test_db, session: AsyncSession, worker: GatewayWorker + ): + 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.PROVISIONING, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname=None, # migration not yet performed + ) + await create_gateway_compute( + session, + gateway_id=gateway.id, + status=GatewayReplicaStatus.RUNNING, + hostname_deprecated_readonly="legacy-lb.example.com", + backend_data="legacy-backend-data", + ) + gateway.lock_token = uuid.uuid4() + gateway.lock_expires_at = datetime(2025, 1, 2, 3, 4, tzinfo=timezone.utc) + await session.commit() + + await worker.process(_gateway_to_pipeline_item(gateway)) + + await session.refresh(gateway) + assert gateway.status == GatewayStatus.RUNNING + assert gateway.hostname == "legacy-lb.example.com" + assert gateway.backend_data == "legacy-backend-data" + + async def test_provisioning_does_not_overwrite_already_migrated_hostname( + self, test_db, session: AsyncSession, worker: GatewayWorker + ): + 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.PROVISIONING, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname="already-migrated.example.com", + backend_data="current-backend-data", + ) + await create_gateway_compute( + session, + gateway_id=gateway.id, + status=GatewayReplicaStatus.RUNNING, + hostname_deprecated_readonly="stale-legacy-lb.example.com", + backend_data="stale-legacy-backend-data", + ) + gateway.lock_token = uuid.uuid4() + gateway.lock_expires_at = datetime(2025, 1, 2, 3, 4, tzinfo=timezone.utc) + await session.commit() + + await worker.process(_gateway_to_pipeline_item(gateway)) + + await session.refresh(gateway) + assert gateway.status == GatewayStatus.RUNNING + assert gateway.hostname == "already-migrated.example.com" + assert gateway.backend_data == "current-backend-data" + async def test_provisioning_to_running_with_multiple_replicas( self, test_db, session: AsyncSession, worker: GatewayWorker ): @@ -787,6 +1087,37 @@ async def test_no_scaling_when_replica_count_matches( assert computes[0].scale_in is False assert gateway.replica_scale_attempt == 0 # The desired count is met, reset counter + async def test_running_migrates_hostname_and_backend_data_from_legacy_replica( + self, test_db, session: AsyncSession, worker: GatewayWorker + ): + 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, + replicas=1, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname=None, # migration not yet performed + ) + await create_gateway_compute( + session, + gateway_id=gateway.id, + status=GatewayReplicaStatus.RUNNING, + hostname_deprecated_readonly="legacy-lb.example.com", + backend_data="legacy-backend-data", + ) + gateway.lock_token = uuid.uuid4() + gateway.lock_expires_at = datetime(2025, 1, 2, 3, 4, tzinfo=timezone.utc) + await session.commit() + + await worker.process(_gateway_to_pipeline_item(gateway)) + + await session.refresh(gateway) + assert gateway.hostname == "legacy-lb.example.com" + assert gateway.backend_data == "legacy-backend-data" + @pytest.mark.parametrize("legacy_compute", [False, True]) @pytest.mark.parametrize("populate_configuration", [True, False]) async def test_scales_out_when_desired_replica_count_increased( @@ -1243,6 +1574,216 @@ async def test_deletes_gateway_when_all_replicas_terminated( assert len(events) == 1 assert events[0].message == "Gateway deleted" + async def test_deletes_gateway_and_terminates_load_balancer_when_hostname_set( + self, test_db, session: AsyncSession, worker: GatewayWorker + ): + 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, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname="gateway-lb.example.com", + backend_data="lb-backend-data", + ) + await create_gateway_compute( + session=session, + backend_id=backend.id, + gateway_id=gateway.id, + status=GatewayReplicaStatus.TERMINATED, + active=False, + ) + gateway.lock_token = uuid.uuid4() + gateway.lock_expires_at = datetime(2025, 1, 2, 3, 4, tzinfo=timezone.utc) + gateway.to_be_deleted = True + await session.commit() + + with patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock: + backend_mock = Mock() + backend_mock.TYPE = BackendType.AWS + backend_mock.compute.return_value = Mock(spec=ComputeMockSpec) + get_backends_mock.return_value = [(backend, backend_mock)] + + await worker.process(_gateway_to_pipeline_item(gateway)) + + terminate_lb_mock = backend_mock.compute.return_value.terminate_gateway_load_balancer + terminate_lb_mock.assert_called_once() + assert terminate_lb_mock.call_args.args[0].gateway_name == gateway.name + assert terminate_lb_mock.call_args.args[1] == "lb-backend-data" + + res = await session.execute(select(GatewayModel.id).where(GatewayModel.id == gateway.id)) + assert res.scalar_one_or_none() is None + events = await list_events(session) + assert len(events) == 1 + assert events[0].message == "Gateway deleted" + + async def test_delete_skips_load_balancer_termination_when_hostname_not_set( + self, test_db, session: AsyncSession, worker: GatewayWorker + ): + 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, + ) + await create_gateway_compute( + session=session, + backend_id=backend.id, + gateway_id=gateway.id, + status=GatewayReplicaStatus.TERMINATED, + active=False, + ) + gateway.lock_token = uuid.uuid4() + gateway.lock_expires_at = datetime(2025, 1, 2, 3, 4, tzinfo=timezone.utc) + gateway.to_be_deleted = True + await session.commit() + + with patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock: + await worker.process(_gateway_to_pipeline_item(gateway)) + get_backends_mock.assert_not_called() + + res = await session.execute(select(GatewayModel.id).where(GatewayModel.id == gateway.id)) + assert res.scalar_one_or_none() is None + + async def test_delete_deferred_when_load_balancer_termination_fails( + self, test_db, session: AsyncSession, worker: GatewayWorker + ): + 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, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname="gateway-lb.example.com", + backend_data="lb-backend-data", + ) + await create_gateway_compute( + session=session, + backend_id=backend.id, + gateway_id=gateway.id, + status=GatewayReplicaStatus.TERMINATED, + active=False, + ) + gateway.lock_token = uuid.uuid4() + gateway.lock_expires_at = datetime(2025, 1, 2, 3, 4, tzinfo=timezone.utc) + gateway.to_be_deleted = True + original_last_processed_at = gateway.last_processed_at + await session.commit() + + with patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock: + backend_mock = Mock() + backend_mock.TYPE = BackendType.AWS + backend_mock.compute.return_value = Mock(spec=ComputeMockSpec) + backend_mock.compute.return_value.terminate_gateway_load_balancer.side_effect = ( + Exception("boom") + ) + get_backends_mock.return_value = [(backend, backend_mock)] + + await worker.process(_gateway_to_pipeline_item(gateway)) + + res = await session.execute(select(GatewayModel.id).where(GatewayModel.id == gateway.id)) + assert res.scalar_one_or_none() is not None + await session.refresh(gateway) + assert gateway.status == GatewayStatus.RUNNING + assert gateway.to_be_deleted is True + assert gateway.last_processed_at > original_last_processed_at + assert gateway.lock_token is None + events = await list_events(session) + assert len(events) == 0 + + async def test_delete_deferred_when_backend_does_not_support_load_balancer( + self, test_db, session: AsyncSession, worker: GatewayWorker + ): + 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, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname="gateway-lb.example.com", + backend_data="lb-backend-data", + ) + await create_gateway_compute( + session=session, + backend_id=backend.id, + gateway_id=gateway.id, + status=GatewayReplicaStatus.TERMINATED, + active=False, + ) + gateway.lock_token = uuid.uuid4() + gateway.lock_expires_at = datetime(2025, 1, 2, 3, 4, tzinfo=timezone.utc) + gateway.to_be_deleted = True + await session.commit() + + with patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock: + backend_mock = Mock() + backend_mock.TYPE = BackendType.AWS + backend_mock.compute.return_value = Mock(spec=ComputeWithGatewaySupport) + get_backends_mock.return_value = [(backend, backend_mock)] + + await worker.process(_gateway_to_pipeline_item(gateway)) + + res = await session.execute(select(GatewayModel.id).where(GatewayModel.id == gateway.id)) + assert res.scalar_one_or_none() is not None + + async def test_delete_migrates_hostname_before_evaluating_termination( + self, test_db, session: AsyncSession, worker: GatewayWorker + ): + 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, + certificate=ACMGatewayCertificate(arn="arn:aws:acm:us:1:certificate/x"), + hostname=None, # migration not yet performed + ) + await create_gateway_compute( + session=session, + backend_id=backend.id, + gateway_id=gateway.id, + status=GatewayReplicaStatus.TERMINATED, + active=False, + hostname_deprecated_readonly="legacy-lb.example.com", + backend_data="legacy-backend-data", + ) + gateway.lock_token = uuid.uuid4() + gateway.lock_expires_at = datetime(2025, 1, 2, 3, 4, tzinfo=timezone.utc) + gateway.to_be_deleted = True + await session.commit() + + with patch( + "dstack._internal.server.services.backends.get_project_backends_with_models" + ) as get_backends_mock: + await worker.process(_gateway_to_pipeline_item(gateway)) + # Migration happens before the load balancer would be evaluated for termination, + # so no backend lookup occurs on this tick. + get_backends_mock.assert_not_called() + + res = await session.execute(select(GatewayModel.id).where(GatewayModel.id == gateway.id)) + assert res.scalar_one_or_none() is not None + await session.refresh(gateway) + assert gateway.to_be_deleted is True + assert gateway.hostname == "legacy-lb.example.com" + assert gateway.backend_data == "legacy-backend-data" + @pytest.mark.parametrize( "replica_status", [ From b51cc09d78a466278c080adee0dd9f77d5456196 Mon Sep 17 00:00:00 2001 From: Jvst Me Date: Wed, 29 Jul 2026 19:16:43 +0200 Subject: [PATCH 3/3] Rename migration --- ..._ecc9e8a0bfac_add_gatewaymodel_hostname_and_backend_data.py} | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) rename src/dstack/_internal/server/migrations/versions/2026/{07_22_1953_ecc9e8a0bfac_gateways_hostname_and_backend_data.py => 07_22_1953_ecc9e8a0bfac_add_gatewaymodel_hostname_and_backend_data.py} (98%) diff --git a/src/dstack/_internal/server/migrations/versions/2026/07_22_1953_ecc9e8a0bfac_gateways_hostname_and_backend_data.py b/src/dstack/_internal/server/migrations/versions/2026/07_22_1953_ecc9e8a0bfac_add_gatewaymodel_hostname_and_backend_data.py similarity index 98% rename from src/dstack/_internal/server/migrations/versions/2026/07_22_1953_ecc9e8a0bfac_gateways_hostname_and_backend_data.py rename to src/dstack/_internal/server/migrations/versions/2026/07_22_1953_ecc9e8a0bfac_add_gatewaymodel_hostname_and_backend_data.py index 9e4b3b07d..971e215dd 100644 --- a/src/dstack/_internal/server/migrations/versions/2026/07_22_1953_ecc9e8a0bfac_gateways_hostname_and_backend_data.py +++ b/src/dstack/_internal/server/migrations/versions/2026/07_22_1953_ecc9e8a0bfac_add_gatewaymodel_hostname_and_backend_data.py @@ -1,4 +1,4 @@ -"""Gateways hostname and backend_data +"""Add GatewayModel.hostname and backend_data Revision ID: ecc9e8a0bfac Revises: dd83c131e78f