2525 ComputeTTLCache ,
2626 ComputeWithAllOffersCached ,
2727 ComputeWithCreateInstanceSupport ,
28+ ComputeWithGatewayLoadBalancerSupport ,
2829 ComputeWithGatewaySupport ,
2930 ComputeWithInstanceVolumesSupport ,
3031 ComputeWithMultinodeSupport ,
5758from dstack ._internal .core .models .common import CoreModel
5859from dstack ._internal .core .models .gateways import (
5960 GatewayComputeConfiguration ,
61+ GatewayLoadBalancerConfiguration ,
62+ GatewayLoadBalancerData ,
6063 GatewayProvisioningData ,
6164)
6265from 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 )
0 commit comments