Skip to content

Commit c580cf0

Browse files
authored
Support replicated AWS gateways with ACM (#4071)
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 ```
1 parent ab915ee commit c580cf0

13 files changed

Lines changed: 1675 additions & 109 deletions

File tree

mkdocs/docs/concepts/gateways.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -221,7 +221,7 @@ $ dstack gateway list
221221
Replicated gateways are an experimental feature and currently have limitations:
222222

223223
- Changing the number of replicas or redeploying replicas is not supported.
224-
- HTTPS is not supported. Use an external load balancer for TLS termination.
224+
- HTTPS is only supported for AWS gateways with the `acm` [certificate type](#certificate). For other gateways, use an external load balancer for TLS termination.
225225
- An unavailable gateway replica prevents any new services or service replicas from being added.
226226
- All replicas are bound to the same backend and region.
227227
- At most 3 replicas are allowed per gateway.

src/dstack/_internal/core/backends/aws/compute.py

Lines changed: 142 additions & 39 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@
2525
ComputeTTLCache,
2626
ComputeWithAllOffersCached,
2727
ComputeWithCreateInstanceSupport,
28+
ComputeWithGatewayLoadBalancerSupport,
2829
ComputeWithGatewaySupport,
2930
ComputeWithInstanceVolumesSupport,
3031
ComputeWithMultinodeSupport,
@@ -57,6 +58,8 @@
5758
from dstack._internal.core.models.common import CoreModel
5859
from dstack._internal.core.models.gateways import (
5960
GatewayComputeConfiguration,
61+
GatewayLoadBalancerConfiguration,
62+
GatewayLoadBalancerData,
6063
GatewayProvisioningData,
6164
)
6265
from dstack._internal.core.models.instances import (
@@ -123,6 +126,7 @@ class AWSCompute(
123126
ComputeWithReservationSupport,
124127
ComputeWithPlacementGroupSupport,
125128
ComputeWithGatewaySupport,
129+
ComputeWithGatewayLoadBalancerSupport,
126130
ComputeWithPrivateGatewaySupport,
127131
ComputeWithVolumeSupport,
128132
Compute,
@@ -584,17 +588,51 @@ def create_gateway(
584588
instance = response[0]
585589
instance.wait_until_running()
586590
instance.reload() # populate instance.public_ip_address
587-
if configuration.certificate is None or configuration.certificate.type != "acm":
588-
ip_address = _get_instance_ip(instance, configuration.public_ip)
589-
return GatewayProvisioningData(
590-
instance_id=instance.instance_id,
591-
region=configuration.region,
592-
availability_zone=availability_zone,
593-
ip_address=ip_address,
594-
)
591+
ip_address = _get_instance_ip(instance, configuration.public_ip)
592+
return GatewayProvisioningData(
593+
instance_id=instance.instance_id,
594+
region=configuration.region,
595+
availability_zone=availability_zone,
596+
ip_address=ip_address,
597+
)
598+
599+
def create_gateway_load_balancer(
600+
self,
601+
configuration: GatewayLoadBalancerConfiguration,
602+
) -> GatewayLoadBalancerData:
603+
"""Creates an ALB, target group, and listeners for a gateway with an ACM certificate."""
604+
assert configuration.certificate is not None
605+
assert configuration.certificate.type == "acm"
595606

607+
ec2_client = self.session.client("ec2", region_name=configuration.region)
596608
elb_client = self.session.client("elbv2", region_name=configuration.region)
597609

610+
base_tags = {
611+
"owner": "dstack",
612+
"dstack_project": configuration.project_name,
613+
"dstack_name": configuration.gateway_name,
614+
}
615+
if settings.DSTACK_VERSION is not None:
616+
base_tags["dstack_version"] = settings.DSTACK_VERSION
617+
tags = merge_tags(
618+
base_tags=base_tags,
619+
backend_tags=self.config.tags,
620+
resource_tags=configuration.tags,
621+
)
622+
tags = aws_resources.filter_invalid_tags(tags)
623+
tags = aws_resources.make_tags(tags)
624+
625+
vpc_id, subnets_ids = self._get_vpc_id_subnets_ids_or_error(
626+
ec2_client=ec2_client,
627+
config=self.config,
628+
region=configuration.region,
629+
allocate_public_ip=configuration.public_ip,
630+
)
631+
security_group_id = aws_resources.create_gateway_security_group(
632+
ec2_client=ec2_client,
633+
project_id=configuration.project_name,
634+
vpc_id=vpc_id,
635+
)
598636
lb_subnets_ids = self._get_gateway_lb_subnets_ids(
599637
ec2_client=ec2_client, region=configuration.region, subnets_ids=subnets_ids
600638
)
@@ -606,7 +644,7 @@ def create_gateway(
606644
# Using short names as LB and target groups have length limit of 32.
607645
resources_name_prefix = generate_unique_short_backend_name()
608646

609-
logger.debug("Creating ALB for gateway %s...", configuration.instance_name)
647+
logger.debug("Creating ALB for gateway %s...", configuration.gateway_name)
610648
response = elb_client.create_load_balancer(
611649
Name=f"{resources_name_prefix}-lb",
612650
Subnets=lb_subnets_ids,
@@ -619,9 +657,9 @@ def create_gateway(
619657
lb = response["LoadBalancers"][0]
620658
lb_arn = lb["LoadBalancerArn"]
621659
lb_dns_name = lb["DNSName"]
622-
logger.debug("Created ALB for gateway %s.", configuration.instance_name)
660+
logger.debug("Created ALB for gateway %s.", configuration.gateway_name)
623661

624-
logger.debug("Creating Target Group for gateway %s...", configuration.instance_name)
662+
logger.debug("Creating Target Group for gateway %s...", configuration.gateway_name)
625663
response = elb_client.create_target_group(
626664
Name=f"{resources_name_prefix}-tg",
627665
Protocol="HTTP",
@@ -630,18 +668,9 @@ def create_gateway(
630668
TargetType="instance",
631669
)
632670
tg_arn = response["TargetGroups"][0]["TargetGroupArn"]
633-
logger.debug("Created Target Group for gateway %s", configuration.instance_name)
634-
635-
logger.debug("Registering ALB target for gateway %s...", configuration.instance_name)
636-
elb_client.register_targets(
637-
TargetGroupArn=tg_arn,
638-
Targets=[
639-
{"Id": instance.instance_id, "Port": 80},
640-
],
641-
)
642-
logger.debug("Registered ALB target for gateway %s", configuration.instance_name)
671+
logger.debug("Created Target Group for gateway %s", configuration.gateway_name)
643672

644-
logger.debug("Creating HTTPS ALB listener for gateway %s...", configuration.instance_name)
673+
logger.debug("Creating HTTPS ALB listener for gateway %s...", configuration.gateway_name)
645674
response = elb_client.create_listener(
646675
LoadBalancerArn=lb_arn,
647676
Protocol="HTTPS",
@@ -658,9 +687,9 @@ def create_gateway(
658687
],
659688
)
660689
listener_arn = response["Listeners"][0]["ListenerArn"]
661-
logger.debug("Created HTTPS ALB listener for gateway %s", configuration.instance_name)
690+
logger.debug("Created HTTPS ALB listener for gateway %s", configuration.gateway_name)
662691

663-
logger.debug("Creating HTTP ALB listener for gateway %s...", configuration.instance_name)
692+
logger.debug("Creating HTTP ALB listener for gateway %s...", configuration.gateway_name)
664693
response = elb_client.create_listener(
665694
LoadBalancerArn=lb_arn,
666695
Protocol="HTTP",
@@ -677,13 +706,9 @@ def create_gateway(
677706
],
678707
)
679708
http_listener_arn = response["Listeners"][0]["ListenerArn"]
680-
logger.debug("Created HTTP ALB listener for gateway %s", configuration.instance_name)
709+
logger.debug("Created HTTP ALB listener for gateway %s", configuration.gateway_name)
681710

682-
ip_address = _get_instance_ip(instance, configuration.public_ip)
683-
return GatewayProvisioningData(
684-
instance_id=instance.instance_id,
685-
region=configuration.region,
686-
ip_address=ip_address,
711+
return GatewayLoadBalancerData(
687712
hostname=lb_dns_name,
688713
backend_data=AWSGatewayBackendData(
689714
lb_arn=lb_arn,
@@ -704,34 +729,112 @@ def terminate_gateway(
704729
region=configuration.region,
705730
backend_data=None,
706731
)
707-
if configuration.certificate is None or configuration.certificate.type != "acm":
708-
return
709732

733+
def terminate_gateway_load_balancer(
734+
self,
735+
configuration: GatewayLoadBalancerConfiguration,
736+
backend_data: Optional[str],
737+
) -> None:
710738
if backend_data is None:
711739
logger.error(
712-
"Failed to terminate all gateway %s resources. backend_data is None.",
713-
configuration.instance_name,
740+
"Failed to terminate load balancer for gateway %s: backend_data is None.",
741+
configuration.gateway_name,
714742
)
715743
return
716-
717744
try:
718745
backend_data_parsed = AWSGatewayBackendData.__response__.parse_raw(backend_data)
719746
except ValidationError:
720747
logger.exception(
721-
"Failed to terminate all gateway %s resources. backend_data parsing error.",
722-
configuration.instance_name,
748+
"Failed to terminate load balancer for gateway %s: backend_data parsing error.",
749+
configuration.gateway_name,
723750
)
724751
return
725752

726753
elb_client = self.session.client("elbv2", region_name=configuration.region)
727754

728-
logger.debug("Deleting ALB resources for gateway %s...", configuration.instance_name)
755+
logger.debug("Deleting ALB resources for gateway %s...", configuration.gateway_name)
729756
if backend_data_parsed.http_listener_arn is not None:
730757
elb_client.delete_listener(ListenerArn=backend_data_parsed.http_listener_arn)
731758
elb_client.delete_listener(ListenerArn=backend_data_parsed.listener_arn)
732759
elb_client.delete_target_group(TargetGroupArn=backend_data_parsed.tg_arn)
733760
elb_client.delete_load_balancer(LoadBalancerArn=backend_data_parsed.lb_arn)
734-
logger.debug("Deleted ALB resources for gateway %s", configuration.instance_name)
761+
logger.debug("Deleted ALB resources for gateway %s.", configuration.gateway_name)
762+
763+
def register_gateway_replica_with_load_balancer(
764+
self,
765+
instance_id: str,
766+
configuration: GatewayLoadBalancerConfiguration,
767+
gateway_backend_data: Optional[str],
768+
) -> None:
769+
if gateway_backend_data is None:
770+
raise ComputeError(
771+
f"Cannot register gateway {configuration.gateway_name} replica with load balancer:"
772+
" gateway_backend_data is None"
773+
)
774+
try:
775+
gateway_backend_data_parsed = AWSGatewayBackendData.__response__.parse_raw(
776+
gateway_backend_data
777+
)
778+
except ValidationError as e:
779+
raise ComputeError(
780+
f"Cannot register gateway {configuration.gateway_name} replica with load balancer:"
781+
" gateway_backend_data parsing error"
782+
) from e
783+
784+
elb_client = self.session.client("elbv2", region_name=configuration.region)
785+
logger.debug(
786+
"Registering gateway %s replica %s with ALB target group %s...",
787+
configuration.gateway_name,
788+
instance_id,
789+
gateway_backend_data_parsed.tg_arn,
790+
)
791+
elb_client.register_targets(
792+
TargetGroupArn=gateway_backend_data_parsed.tg_arn,
793+
Targets=[{"Id": instance_id, "Port": 80}],
794+
)
795+
logger.debug(
796+
"Registered gateway %s replica %s with ALB target group.",
797+
configuration.gateway_name,
798+
instance_id,
799+
)
800+
801+
def deregister_gateway_replica_from_load_balancer(
802+
self,
803+
instance_id: str,
804+
configuration: GatewayLoadBalancerConfiguration,
805+
gateway_backend_data: Optional[str],
806+
) -> None:
807+
if gateway_backend_data is None:
808+
raise ComputeError(
809+
f"Cannot deregister gateway {configuration.gateway_name} replica from load balancer:"
810+
" gateway_backend_data is None"
811+
)
812+
try:
813+
gateway_backend_data_parsed = AWSGatewayBackendData.__response__.parse_raw(
814+
gateway_backend_data
815+
)
816+
except ValidationError as e:
817+
raise ComputeError(
818+
f"Cannot deregister gateway {configuration.gateway_name} replica from load balancer:"
819+
" gateway_backend_data parsing error",
820+
) from e
821+
822+
elb_client = self.session.client("elbv2", region_name=configuration.region)
823+
logger.debug(
824+
"Deregistering gateway %s replica %s from ALB target group %s...",
825+
configuration.gateway_name,
826+
instance_id,
827+
gateway_backend_data_parsed.tg_arn,
828+
)
829+
elb_client.deregister_targets(
830+
TargetGroupArn=gateway_backend_data_parsed.tg_arn,
831+
Targets=[{"Id": instance_id, "Port": 80}],
832+
)
833+
logger.debug(
834+
"Deregistered gateway %s replica %s from ALB target group.",
835+
configuration.gateway_name,
836+
instance_id,
837+
)
735838

736839
def register_volume(self, volume: Volume) -> VolumeProvisioningData:
737840
assert isinstance(volume.configuration, AWSVolumeConfiguration)

src/dstack/_internal/core/backends/base/compute.py

Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,8 @@
3030
from dstack._internal.core.models.compute_groups import ComputeGroup, ComputeGroupProvisioningData
3131
from dstack._internal.core.models.gateways import (
3232
GatewayComputeConfiguration,
33+
GatewayLoadBalancerConfiguration,
34+
GatewayLoadBalancerData,
3335
GatewayProvisioningData,
3436
)
3537
from dstack._internal.core.models.instances import (
@@ -578,6 +580,55 @@ def terminate_gateway(
578580
pass
579581

580582

583+
class ComputeWithGatewayLoadBalancerSupport(ABC):
584+
"""
585+
Must be subclassed and implemented to support gateways with a load balancer that fronts
586+
all replica instances.
587+
588+
Backends implementing this mixin must also implement `ComputeWithGatewaySupport`.
589+
"""
590+
591+
@abstractmethod
592+
def create_gateway_load_balancer(
593+
self,
594+
configuration: GatewayLoadBalancerConfiguration,
595+
) -> GatewayLoadBalancerData:
596+
"""Creates the load balancer for a gateway."""
597+
pass
598+
599+
@abstractmethod
600+
def terminate_gateway_load_balancer(
601+
self,
602+
configuration: GatewayLoadBalancerConfiguration,
603+
backend_data: Optional[str],
604+
) -> None:
605+
"""Deletes the load balancer."""
606+
pass
607+
608+
@abstractmethod
609+
def register_gateway_replica_with_load_balancer(
610+
self,
611+
instance_id: str,
612+
configuration: GatewayLoadBalancerConfiguration,
613+
gateway_backend_data: Optional[str],
614+
) -> None:
615+
"""Registers a gateway replica instance as a target of the load balancer."""
616+
pass
617+
618+
@abstractmethod
619+
def deregister_gateway_replica_from_load_balancer(
620+
self,
621+
instance_id: str,
622+
configuration: GatewayLoadBalancerConfiguration,
623+
gateway_backend_data: Optional[str],
624+
) -> None:
625+
"""Deregisters a gateway replica instance from the load balancer.
626+
627+
If the replica is not registered, it should not raise errors but return silently.
628+
"""
629+
pass
630+
631+
581632
class ComputeWithPrivateGatewaySupport:
582633
"""
583634
Must be subclassed to support private gateways.

src/dstack/_internal/core/models/gateways.py

Lines changed: 15 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -215,6 +215,20 @@ class GatewayProvisioningData(CoreModel):
215215
ip_address: str
216216
region: str
217217
availability_zone: Optional[str] = None
218-
hostname: Optional[str] = None
219218
backend_data: Optional[str] = None
220219
"""`backend_data` stores backend-specific data in JSON."""
220+
221+
222+
class GatewayLoadBalancerConfiguration(CoreModel):
223+
project_name: str
224+
gateway_name: str
225+
region: str
226+
public_ip: bool
227+
certificate: Optional[AnyGatewayCertificate] = None
228+
tags: Optional[Dict[str, str]] = None
229+
230+
231+
class GatewayLoadBalancerData(CoreModel):
232+
hostname: str
233+
backend_data: str
234+
"""Backend-specific JSON"""

0 commit comments

Comments
 (0)