Skip to content
Draft
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
2 changes: 1 addition & 1 deletion mkdocs/docs/concepts/gateways.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
183 changes: 143 additions & 40 deletions src/dstack/_internal/core/backends/aws/compute.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@
ComputeTTLCache,
ComputeWithAllOffersCached,
ComputeWithCreateInstanceSupport,
ComputeWithGatewayLoadBalancerSupport,
ComputeWithGatewaySupport,
ComputeWithInstanceVolumesSupport,
ComputeWithMultinodeSupport,
Expand Down Expand Up @@ -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 (
Expand Down Expand Up @@ -123,6 +126,7 @@ class AWSCompute(
ComputeWithReservationSupport,
ComputeWithPlacementGroupSupport,
ComputeWithGatewaySupport,
ComputeWithGatewayLoadBalancerSupport,
ComputeWithPrivateGatewaySupport,
ComputeWithVolumeSupport,
Compute,
Expand Down Expand Up @@ -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
)
Expand All @@ -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,
Expand All @@ -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",
Expand All @@ -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",
Expand All @@ -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",
Expand All @@ -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,
Expand All @@ -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)
Expand Down
51 changes: 51 additions & 0 deletions src/dstack/_internal/core/backends/base/compute.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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.
Expand Down
16 changes: 15 additions & 1 deletion src/dstack/_internal/core/models/gateways.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"""
Loading
Loading