diff --git a/api/holodeck/v1alpha1/types.go b/api/holodeck/v1alpha1/types.go index d40bb02de..076cb2c36 100644 --- a/api/holodeck/v1alpha1/types.go +++ b/api/holodeck/v1alpha1/types.go @@ -97,6 +97,12 @@ type Instance struct { Type string `json:"type"` Region string `json:"region"` + // AvailabilityZone places the instance in a specific zone of Region + // (e.g., "us-west-2a"). When unset, Holodeck picks a zone that offers + // the instance type. + // +optional + AvailabilityZone string `json:"availabilityZone,omitempty"` + // OS specifies the operating system by ID (e.g., "ubuntu-22.04"). // When set, the AMI is automatically resolved for the region and // architecture. Takes precedence over Image.ImageId if both are specified. @@ -169,6 +175,12 @@ type ClusterSpec struct { // +required Region string `json:"region"` + // AvailabilityZone places all cluster nodes in a specific zone of Region + // (e.g., "us-west-2a"). When unset, Holodeck picks a zone that offers + // both the control-plane and worker instance types. + // +optional + AvailabilityZone string `json:"availabilityZone,omitempty"` + // ControlPlane defines the control-plane node configuration. // +required ControlPlane ControlPlaneSpec `json:"controlPlane"` diff --git a/cmd/cli/describe/describe.go b/cmd/cli/describe/describe.go index 067585ebc..940db3949 100644 --- a/cmd/cli/describe/describe.go +++ b/cmd/cli/describe/describe.go @@ -27,6 +27,7 @@ import ( "github.com/NVIDIA/holodeck/internal/logger" "github.com/NVIDIA/holodeck/pkg/jyaml" "github.com/NVIDIA/holodeck/pkg/output" + "github.com/NVIDIA/holodeck/pkg/provider/aws" cli "github.com/urfave/cli/v3" ) @@ -59,10 +60,11 @@ type InstanceInfo struct { // ProviderInfo contains provider configuration type ProviderInfo struct { - Type string `json:"type" yaml:"type"` - Region string `json:"region,omitempty" yaml:"region,omitempty"` - Username string `json:"username" yaml:"username"` - KeyName string `json:"keyName" yaml:"keyName"` + Type string `json:"type" yaml:"type"` + Region string `json:"region,omitempty" yaml:"region,omitempty"` + AvailabilityZone string `json:"availabilityZone,omitempty" yaml:"availabilityZone,omitempty"` + Username string `json:"username" yaml:"username"` + KeyName string `json:"keyName" yaml:"keyName"` } // ClusterInfo contains cluster configuration @@ -308,6 +310,11 @@ func (m command) buildDescribeOutput(instance *instances.Instance, env *v1alpha1 } else { output.Provider.Region = env.Spec.Region } + for _, property := range env.Status.Properties { + if property.Name == aws.AvailabilityZone { + output.Provider.AvailabilityZone = property.Value + } + } // Cluster info if env.Spec.Cluster != nil { @@ -555,6 +562,9 @@ func (m command) printTableFormat(d *DescribeOutput) error { if d.Provider.Region != "" { fmt.Printf("Region: %s\n", d.Provider.Region) } + if d.Provider.AvailabilityZone != "" { + fmt.Printf("Zone: %s\n", d.Provider.AvailabilityZone) + } fmt.Printf("Username: %s\n", d.Provider.Username) fmt.Printf("Key Name: %s\n", d.Provider.KeyName) diff --git a/cmd/cli/describe/describe_test.go b/cmd/cli/describe/describe_test.go index 3bc69fa35..a2cd90187 100644 --- a/cmd/cli/describe/describe_test.go +++ b/cmd/cli/describe/describe_test.go @@ -19,6 +19,9 @@ package describe import ( "testing" "time" + + "github.com/NVIDIA/holodeck/api/holodeck/v1alpha1" + "github.com/NVIDIA/holodeck/internal/instances" ) func TestDescribeOutput_InstanceInfo(t *testing.T) { @@ -198,3 +201,46 @@ func TestAWSResourcesInfo(t *testing.T) { t.Errorf("expected vpc-123, got %s", output.AWSResources.VpcID) } } + +func TestBuildDescribeOutput_AvailabilityZone(t *testing.T) { + singleNodeSpec := v1alpha1.EnvironmentSpec{ + Provider: v1alpha1.ProviderAWS, + Instance: v1alpha1.Instance{Region: "us-west-2"}, + } + clusterSpec := v1alpha1.EnvironmentSpec{ + Provider: v1alpha1.ProviderAWS, + Cluster: &v1alpha1.ClusterSpec{Region: "us-west-2"}, + } + propertiesWithZone := []v1alpha1.Properties{ + {Name: "vpc-id", Value: "vpc-123"}, + {Name: "availability-zone", Value: "us-west-2c"}, + } + propertiesWithoutZone := []v1alpha1.Properties{ + {Name: "vpc-id", Value: "vpc-123"}, + } + + tests := []struct { + name string + spec v1alpha1.EnvironmentSpec + properties []v1alpha1.Properties + want string + }{ + {"single node", singleNodeSpec, propertiesWithZone, "us-west-2c"}, + {"cluster", clusterSpec, propertiesWithZone, "us-west-2c"}, + {"cache written before the zone was recorded", singleNodeSpec, propertiesWithoutZone, ""}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + env := &v1alpha1.Environment{ + Spec: tt.spec, + Status: v1alpha1.EnvironmentStatus{Properties: tt.properties}, + } + + output := command{}.buildDescribeOutput(&instances.Instance{}, env, time.Hour) + + if output.Provider.AvailabilityZone != tt.want { + t.Errorf("expected availability zone %q, got %q", tt.want, output.Provider.AvailabilityZone) + } + }) + } +} diff --git a/docs/commands/create.md b/docs/commands/create.md index 3817fc2b4..13db93385 100644 --- a/docs/commands/create.md +++ b/docs/commands/create.md @@ -91,6 +91,28 @@ spec: See the [Multinode Clusters Guide](../guides/multinode-clusters.md) for detailed configuration options and examples. +### Availability Zone + +Not every instance type is offered in every Availability Zone of a region. +Holodeck creates the environment's subnets in a zone that offers all of the +requested instance types, and fails before creating any resources if no zone +does. To pin a zone, set `availabilityZone` under `instance` (or `cluster`): + +```yaml + instance: + type: g5g.xlarge + region: us-west-2 + availabilityZone: +``` + +Zone names differ between AWS accounts, so leave `availabilityZone` unset to let +Holodeck choose, or pick a zone that offers the instance type; if it does not, +the pre-flight error lists the zones that do. + +Choosing a zone needs the `ec2:DescribeAvailabilityZones` and +`ec2:DescribeInstanceTypeOfferings` permissions; without them Holodeck lets AWS +choose the zone and rejects a pinned `availabilityZone`. + ## Automated IP Detection Holodeck now automatically detects your public IP address when creating AWS diff --git a/docs/guides/multinode-clusters.md b/docs/guides/multinode-clusters.md index d0b53ec96..e0473fd96 100644 --- a/docs/guides/multinode-clusters.md +++ b/docs/guides/multinode-clusters.md @@ -90,6 +90,7 @@ holodeck create -f cluster.yaml --provision -k kubeconfig.yaml | Field | Type | Description | |-------|------|-------------| | `region` | string | AWS region for all nodes (required) | +| `availabilityZone` | string | Zone for all nodes (optional; by default a zone offering every instance type is picked). Zone names differ between AWS accounts | | `controlPlane` | ControlPlaneSpec | Control plane node configuration | | `workers` | WorkerPoolSpec | Worker node pool configuration | | `highAvailability` | HAConfig | HA settings (optional) | diff --git a/docs/prerequisites.md b/docs/prerequisites.md index bf1b939d9..51ab29c39 100644 --- a/docs/prerequisites.md +++ b/docs/prerequisites.md @@ -29,6 +29,13 @@ To use the AWS provider, you need: - VPC configuration - Security group management - IAM role management + - Pre-flight checks run by `holodeck create` and `holodeck dryrun`: + `ec2:DescribeInstanceTypes` + - Recommended: `ec2:DescribeInstanceTypeOfferings` and + `ec2:DescribeAvailabilityZones`, used to find an Availability Zone that + offers the requested instance types. Without them Holodeck logs a + warning and lets AWS choose the zone, and rejects a pinned + `availabilityZone` ### SSH Provider diff --git a/internal/aws/awsfake/awsfake_test.go b/internal/aws/awsfake/awsfake_test.go index 385f40f7f..f4d0745a4 100644 --- a/internal/aws/awsfake/awsfake_test.go +++ b/internal/aws/awsfake/awsfake_test.go @@ -944,3 +944,101 @@ func TestSetInstanceTypeCatalog(t *testing.T) { t.Fatalf("empty catalog must return no types, got %+v", empty.InstanceTypes) } } + +// The provider's zone-selection tests depend on these seeding controls and on following NextToken. +func TestDescribeInstanceTypeOfferings(t *testing.T) { + f := New() + f.Store.SeedInstanceTypeZones("g5g.xlarge", "us-west-2b", "us-west-2c") + f.Store.SeedInstanceTypeAbsent("t99.nonexistent") + f.Store.SeedAvailabilityZone(ec2types.AvailabilityZone{ + ZoneName: aws.String("us-west-2-lax-1a"), + ZoneType: aws.String("local-zone"), + State: ec2types.AvailabilityZoneStateAvailable, + }) + + zonesByInstanceType := map[string][]string{} + input := &ec2.DescribeInstanceTypeOfferingsInput{ + LocationType: ec2types.LocationTypeAvailabilityZone, + Filters: []ec2types.Filter{{ + Name: aws.String("instance-type"), + Values: []string{"g5g.xlarge", "t3.medium", "t99.nonexistent"}, + }}, + } + pageCount := 0 + for { + out, err := f.EC2.DescribeInstanceTypeOfferings(ctx, input) + if err != nil { + t.Fatalf("DescribeInstanceTypeOfferings: %v", err) + } + pageCount++ + for _, offering := range out.InstanceTypeOfferings { + instanceType := string(offering.InstanceType) + zonesByInstanceType[instanceType] = append(zonesByInstanceType[instanceType], aws.ToString(offering.Location)) + } + if out.NextToken == nil { + break + } + input.NextToken = out.NextToken + } + + if pageCount < 2 { + t.Fatalf("expected paginated results, got %d page(s)", pageCount) + } + if got, want := zonesByInstanceType["g5g.xlarge"], []string{"us-west-2b", "us-west-2c"}; !slices.Equal(got, want) { + t.Fatalf("g5g.xlarge zones = %v, want %v", got, want) + } + if got, want := zonesByInstanceType["t3.medium"], []string{"us-west-2a", "us-west-2b", "us-west-2c", "us-west-2d"}; !slices.Equal(got, want) { + t.Fatalf("t3.medium zones = %v, want %v (local zones excluded by default)", got, want) + } + if got := zonesByInstanceType["t99.nonexistent"]; len(got) != 0 { + t.Fatalf("absent type must be offered nowhere, got %v", got) + } +} + +// The provider filters on zone-type to exclude Local Zones, which sort before the region's own zones. +func TestDescribeAvailabilityZonesFiltersByZoneType(t *testing.T) { + f := New() + f.Store.SeedAvailabilityZone(ec2types.AvailabilityZone{ + ZoneName: aws.String("us-west-2-lax-1a"), + ZoneType: aws.String("local-zone"), + State: ec2types.AvailabilityZoneStateAvailable, + }) + + zones, err := f.EC2.DescribeAvailabilityZones(ctx, &ec2.DescribeAvailabilityZonesInput{ + Filters: []ec2types.Filter{{Name: aws.String("zone-type"), Values: []string{"availability-zone"}}}, + }) + if err != nil { + t.Fatalf("DescribeAvailabilityZones: %v", err) + } + if len(zones.AvailabilityZones) != 4 { + t.Fatalf("zone-type filter must drop the local zone, got %+v", zones.AvailabilityZones) + } +} + +// The mock e2e test compares each stored subnet's zone with the zone Create records. +func TestCreateSubnetRecordsAvailabilityZone(t *testing.T) { + f := New() + subnet, err := f.EC2.CreateSubnet(ctx, &ec2.CreateSubnetInput{ + VpcId: aws.String("vpc-x"), + CidrBlock: aws.String("10.0.0.0/24"), + AvailabilityZone: aws.String("us-west-2c"), + }) + if err != nil { + t.Fatalf("CreateSubnet: %v", err) + } + if got := aws.ToString(subnet.Subnet.AvailabilityZone); got != "us-west-2c" { + t.Fatalf("subnet zone = %q, want us-west-2c", got) + } +} + +// A negative NextToken parses cleanly, so it must be range-checked before it +// is used as a slice index. +func TestDescribeInstanceTypeOfferingsRejectsNegativeNextToken(t *testing.T) { + f := New() + _, err := f.EC2.DescribeInstanceTypeOfferings(ctx, &ec2.DescribeInstanceTypeOfferingsInput{ + NextToken: aws.String("-1"), + }) + if err == nil { + t.Fatal("expected an error for a negative NextToken") + } +} diff --git a/internal/aws/awsfake/ec2.go b/internal/aws/awsfake/ec2.go index 7c65cc7e8..e23a04b8f 100644 --- a/internal/aws/awsfake/ec2.go +++ b/internal/aws/awsfake/ec2.go @@ -19,6 +19,9 @@ package awsfake import ( "context" "fmt" + "maps" + "slices" + "strconv" "github.com/aws/aws-sdk-go-v2/aws" "github.com/aws/aws-sdk-go-v2/service/ec2" @@ -114,11 +117,12 @@ func (f *FakeEC2) CreateSubnet(ctx context.Context, params *ec2.CreateSubnetInpu } id := f.store.nextID("subnet") sn := ec2types.Subnet{ - SubnetId: aws.String(id), - VpcId: params.VpcId, - CidrBlock: params.CidrBlock, - State: ec2types.SubnetStateAvailable, - Tags: tagsFromSpecs(params.TagSpecifications), + SubnetId: aws.String(id), + VpcId: params.VpcId, + CidrBlock: params.CidrBlock, + AvailabilityZone: params.AvailabilityZone, + State: ec2types.SubnetStateAvailable, + Tags: tagsFromSpecs(params.TagSpecifications), } f.store.Subnets[id] = &sn return &ec2.CreateSubnetOutput{Subnet: &sn}, nil @@ -620,6 +624,75 @@ func (f *FakeEC2) DescribeInstanceTypes(ctx context.Context, params *ec2.Describ return &ec2.DescribeInstanceTypesOutput{InstanceTypes: infos, NextToken: nil}, nil } +// instanceTypeOfferingsPageSize is deliberately small so callers must follow +// NextToken to see every offering, as they must against the real API. +const instanceTypeOfferingsPageSize = 2 + +// DescribeInstanceTypeOfferings models only the availability-zone location +// type, honouring the "instance-type" filter. +func (f *FakeEC2) DescribeInstanceTypeOfferings(_ context.Context, params *ec2.DescribeInstanceTypeOfferingsInput, _ ...func(*ec2.Options)) (*ec2.DescribeInstanceTypeOfferingsOutput, error) { + f.store.mu.Lock() + defer f.store.mu.Unlock() + f.store.record("DescribeInstanceTypeOfferings", params) + if err := f.store.failure("DescribeInstanceTypeOfferings"); err != nil { + return nil, err + } + instanceTypes := filterValues(params.Filters, "instance-type") + if len(instanceTypes) == 0 { + instanceTypes = slices.Sorted(maps.Keys(f.store.InstanceTypes)) + } + var offerings []ec2types.InstanceTypeOffering + for _, instanceType := range instanceTypes { + for _, zoneName := range f.store.zonesOffering(instanceType) { + offerings = append(offerings, ec2types.InstanceTypeOffering{ + InstanceType: ec2types.InstanceType(instanceType), + LocationType: ec2types.LocationTypeAvailabilityZone, + Location: aws.String(zoneName), + }) + } + } + + start := 0 + if params.NextToken != nil { + var err error + if start, err = strconv.Atoi(*params.NextToken); err != nil || start < 0 || start > len(offerings) { + return nil, fmt.Errorf("InvalidNextToken: %q", *params.NextToken) + } + } + end := min(start+instanceTypeOfferingsPageSize, len(offerings)) + out := &ec2.DescribeInstanceTypeOfferingsOutput{InstanceTypeOfferings: offerings[start:end]} + if end < len(offerings) { + out.NextToken = aws.String(strconv.Itoa(end)) + } + return out, nil +} + +// ---- Availability Zones ---- + +// DescribeAvailabilityZones honours only the "zone-type" and "state" filters; +// ZoneNames, ZoneIds and AllAvailabilityZones are ignored. +func (f *FakeEC2) DescribeAvailabilityZones(_ context.Context, params *ec2.DescribeAvailabilityZonesInput, _ ...func(*ec2.Options)) (*ec2.DescribeAvailabilityZonesOutput, error) { + f.store.mu.Lock() + defer f.store.mu.Unlock() + f.store.record("DescribeAvailabilityZones", params) + if err := f.store.failure("DescribeAvailabilityZones"); err != nil { + return nil, err + } + zoneTypes := filterValues(params.Filters, "zone-type") + states := filterValues(params.Filters, "state") + var zones []ec2types.AvailabilityZone + for _, zone := range f.store.AvailabilityZones { + if len(zoneTypes) > 0 && !slices.Contains(zoneTypes, aws.ToString(zone.ZoneType)) { + continue + } + if len(states) > 0 && !slices.Contains(states, string(zone.State)) { + continue + } + zones = append(zones, zone) + } + return &ec2.DescribeAvailabilityZonesOutput{AvailabilityZones: zones}, nil +} + func instanceTypeInfo(name string) ec2types.InstanceTypeInfo { return ec2types.InstanceTypeInfo{ InstanceType: ec2types.InstanceType(name), @@ -914,6 +987,16 @@ func (f *FakeEC2) ModifySubnetAttribute(ctx context.Context, params *ec2.ModifyS return &ec2.ModifySubnetAttributeOutput{}, nil } +// filterValues returns every value of the named filter, or nil if absent. +func filterValues(filters []ec2types.Filter, name string) []string { + for _, filter := range filters { + if aws.ToString(filter.Name) == name { + return filter.Values + } + } + return nil +} + // filterValue returns the first value of the named filter, or "" if absent. func filterValue(filters []ec2types.Filter, name string) string { for _, filter := range filters { diff --git a/internal/aws/awsfake/store.go b/internal/aws/awsfake/store.go index 63292b427..95b757378 100644 --- a/internal/aws/awsfake/store.go +++ b/internal/aws/awsfake/store.go @@ -85,8 +85,9 @@ type Store struct { Parameters map[string]string // Seed data (excluded from ResourceCounts/Empty). - Images []ec2types.Image - InstanceTypes map[string][]ec2types.ArchitectureType + Images []ec2types.Image + InstanceTypes map[string][]ec2types.ArchitectureType + AvailabilityZones []ec2types.AvailabilityZone // Per-instance-type overrides for filtered DescribeInstanceTypes queries: // explicit architectures (bypassing the prefix heuristic) and types marked @@ -94,6 +95,11 @@ type Store struct { instanceTypeArchs map[string][]ec2types.ArchitectureType absentInstanceTypes map[string]bool + // Per-instance-type Availability Zone offerings for + // DescribeInstanceTypeOfferings. Types without an entry are offered in + // every standard zone of AvailabilityZones, unless seeded absent. + instanceTypeZones map[string][]string + // Recorder + fault injection + id generator. calls map[string]int inputs map[string][]any @@ -142,6 +148,7 @@ func newStore() *Store { InstanceTypes: map[string][]ec2types.ArchitectureType{}, instanceTypeArchs: map[string][]ec2types.ArchitectureType{}, absentInstanceTypes: map[string]bool{}, + instanceTypeZones: map[string][]string{}, calls: map[string]int{}, inputs: map[string][]any{}, failures: map[string][]error{}, @@ -203,6 +210,15 @@ func (s *Store) seed() { } { s.InstanceTypes[t] = archsFor(t) } + + for _, zoneName := range []string{"us-west-2a", "us-west-2b", "us-west-2c", "us-west-2d"} { + s.AvailabilityZones = append(s.AvailabilityZones, ec2types.AvailabilityZone{ + ZoneName: aws.String(zoneName), + ZoneType: aws.String("availability-zone"), + State: ec2types.AvailabilityZoneStateAvailable, + RegionName: aws.String("us-west-2"), + }) + } } // nextID returns a unique, deterministic id for the given prefix @@ -347,6 +363,54 @@ func (s *Store) SeedInstanceTypeAbsent(name string) { s.absentInstanceTypes[name] = true } +// SeedInstanceTypeZones restricts the Availability Zones in which +// DescribeInstanceTypeOfferings reports an instance type as offered. Passing +// no zones models a type that no zone in the region offers. +func (s *Store) SeedInstanceTypeZones(name string, zoneNames ...string) { + s.mu.Lock() + defer s.mu.Unlock() + s.instanceTypeZones[name] = zoneNames +} + +// SeedAvailabilityZone adds a zone to the DescribeAvailabilityZones catalog, +// e.g. a Local Zone (ZoneType "local-zone") or a zone in a non-available state. +func (s *Store) SeedAvailabilityZone(zone ec2types.AvailabilityZone) { + s.mu.Lock() + defer s.mu.Unlock() + s.AvailabilityZones = append(s.AvailabilityZones, zone) +} + +// SeedAvailabilityZoneState changes the state DescribeAvailabilityZones +// reports for an already-seeded zone. DescribeInstanceTypeOfferings ignores +// zone state, so the zone keeps being listed as offering instance types. +func (s *Store) SeedAvailabilityZoneState(zoneName string, state ec2types.AvailabilityZoneState) { + s.mu.Lock() + defer s.mu.Unlock() + for i := range s.AvailabilityZones { + if aws.ToString(s.AvailabilityZones[i].ZoneName) == zoneName { + s.AvailabilityZones[i].State = state + } + } +} + +// zonesOffering returns the zones that offer an instance type. Callers must +// hold mu. +func (s *Store) zonesOffering(instanceType string) []string { + if zoneNames, ok := s.instanceTypeZones[instanceType]; ok { + return zoneNames + } + if s.absentInstanceTypes[instanceType] { + return nil + } + var zoneNames []string + for _, zone := range s.AvailabilityZones { + if aws.ToString(zone.ZoneType) == "availability-zone" { + zoneNames = append(zoneNames, aws.ToString(zone.ZoneName)) + } + } + return zoneNames +} + // SetImages replaces the DescribeImages catalog (clearing the seeded default // images) so a test can present exactly the AMIs a resolution path should see. func (s *Store) SetImages(imgs ...ec2types.Image) { diff --git a/internal/aws/ec2_client.go b/internal/aws/ec2_client.go index 5d56ec54f..0c52363d3 100644 --- a/internal/aws/ec2_client.go +++ b/internal/aws/ec2_client.go @@ -115,6 +115,15 @@ type EC2Client interface { DescribeInstanceTypes(ctx context.Context, params *ec2.DescribeInstanceTypesInput, optFns ...func(*ec2.Options)) (*ec2.DescribeInstanceTypesOutput, error) + DescribeInstanceTypeOfferings(ctx context.Context, + params *ec2.DescribeInstanceTypeOfferingsInput, + optFns ...func(*ec2.Options)) (*ec2.DescribeInstanceTypeOfferingsOutput, + error) + + // Availability Zone operations + DescribeAvailabilityZones(ctx context.Context, + params *ec2.DescribeAvailabilityZonesInput, + optFns ...func(*ec2.Options)) (*ec2.DescribeAvailabilityZonesOutput, error) // Image operations DescribeImages(ctx context.Context, params *ec2.DescribeImagesInput, diff --git a/pkg/provider/aws/aws.go b/pkg/provider/aws/aws.go index e590854af..c5bab77cd 100644 --- a/pkg/provider/aws/aws.go +++ b/pkg/provider/aws/aws.go @@ -46,6 +46,7 @@ const ( SecurityGroupID string = "security-group-id" InstanceID string = "instance-id" PublicDnsName string = "public-dns-name" + AvailabilityZone string = "availability-zone" // Cluster networking cache keys PublicSubnetID string = "public-subnet-id" @@ -82,6 +83,7 @@ type AWS struct { SecurityGroupid string Instanceid string PublicDnsName string + AvailabilityZone string // Cluster networking fields PublicSubnetid string @@ -102,6 +104,9 @@ type Provider struct { cacheFile string sleep func(time.Duration) + selectedAvailabilityZone string + letAWSChooseAvailabilityZone bool + *v1alpha1.Environment log *logger.FunLogger } @@ -262,6 +267,8 @@ func (p *Provider) unmarsalCache() (*AWS, error) { aws.Instanceid = p.Value case PublicDnsName: aws.PublicDnsName = p.Value + case AvailabilityZone: + aws.AvailabilityZone = p.Value case PublicSubnetID: aws.PublicSubnetid = p.Value case NatGatewayID: diff --git a/pkg/provider/aws/aws_ginkgo_test.go b/pkg/provider/aws/aws_ginkgo_test.go index bfc8e363a..b50402ef5 100644 --- a/pkg/provider/aws/aws_ginkgo_test.go +++ b/pkg/provider/aws/aws_ginkgo_test.go @@ -83,6 +83,7 @@ var _ = Describe("AWS Provider", func() { Expect(SecurityGroupID).To(Equal("security-group-id")) Expect(InstanceID).To(Equal("instance-id")) Expect(PublicDnsName).To(Equal("public-dns-name")) + Expect(AvailabilityZone).To(Equal("availability-zone")) }) }) @@ -164,6 +165,7 @@ status: Expect(aws.SecurityGroupid).To(Equal("sg-44444")) Expect(aws.Instanceid).To(Equal("i-55555")) Expect(aws.PublicDnsName).To(Equal("ec2-1-2-3-4.compute.amazonaws.com")) + Expect(aws.AvailabilityZone).To(BeEmpty()) }) }) @@ -325,7 +327,7 @@ status: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", }, }, } @@ -464,7 +466,7 @@ status: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", }, }, } @@ -499,6 +501,15 @@ status: Expect(err.Error()).To(ContainSubstring("not supported")) }) + It("should fail when no availability zone offers the instance type", func() { + f.Store.SeedInstanceTypeZones("t3.medium") + + err := provider.DryRun() + Expect(err).To(HaveOccurred()) + Expect(err.Error()).To(ContainSubstring("no availability zone")) + Expect(err.Error()).To(ContainSubstring("t3.medium")) + }) + It("should fail when image check fails (triggers fail())", func() { // t3.medium is valid, but the image describe fails. f.Store.FailNext("DescribeImages", ErrMockDescribeImages) diff --git a/pkg/provider/aws/aws_test.go b/pkg/provider/aws/aws_test.go index ccf5fda02..675a701b8 100644 --- a/pkg/provider/aws/aws_test.go +++ b/pkg/provider/aws/aws_test.go @@ -481,7 +481,7 @@ status: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", }, }, } @@ -795,7 +795,7 @@ spec: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "x86_64"}, }, Auth: v1alpha1.Auth{ @@ -824,7 +824,7 @@ spec: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "x86_64"}, }, Auth: v1alpha1.Auth{ @@ -855,7 +855,7 @@ spec: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "x86_64"}, }, Auth: v1alpha1.Auth{ @@ -886,7 +886,7 @@ spec: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "x86_64"}, }, Auth: v1alpha1.Auth{ @@ -916,7 +916,7 @@ spec: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "x86_64"}, }, Auth: v1alpha1.Auth{ @@ -948,7 +948,7 @@ spec: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "x86_64"}, }, Auth: v1alpha1.Auth{ @@ -977,7 +977,7 @@ spec: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "x86_64"}, }, Auth: v1alpha1.Auth{ @@ -1006,7 +1006,7 @@ spec: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "x86_64"}, }, Auth: v1alpha1.Auth{ @@ -1037,7 +1037,7 @@ spec: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "x86_64"}, }, Auth: v1alpha1.Auth{ @@ -1068,7 +1068,7 @@ spec: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "x86_64"}, }, Auth: v1alpha1.Auth{ @@ -1100,7 +1100,7 @@ spec: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "x86_64"}, }, Auth: v1alpha1.Auth{ @@ -1136,7 +1136,7 @@ spec: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{ImageId: &imageID}, }, Auth: v1alpha1.Auth{ @@ -1174,7 +1174,7 @@ spec: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "x86_64"}, }, Auth: v1alpha1.Auth{ @@ -1207,7 +1207,7 @@ spec: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "x86_64"}, }, Auth: v1alpha1.Auth{ @@ -1236,7 +1236,7 @@ spec: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "x86_64"}, }, Auth: v1alpha1.Auth{ @@ -1376,7 +1376,7 @@ status: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "x86_64"}, }, Auth: v1alpha1.Auth{ @@ -1410,7 +1410,7 @@ status: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t4g.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "arm64"}, }, Auth: v1alpha1.Auth{ @@ -1437,7 +1437,7 @@ status: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{Architecture: "invalid_arch"}, }, Auth: v1alpha1.Auth{ @@ -1471,7 +1471,7 @@ status: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{ImageId: &imageID}, }, Auth: v1alpha1.Auth{ @@ -1504,7 +1504,7 @@ status: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{ImageId: &imageID}, }, Auth: v1alpha1.Auth{ @@ -1539,7 +1539,7 @@ status: Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{ Architecture: "x86_64", OwnerId: &ownerID, diff --git a/pkg/provider/aws/cache_test.go b/pkg/provider/aws/cache_test.go index af56fd269..8d22e03c8 100644 --- a/pkg/provider/aws/cache_test.go +++ b/pkg/provider/aws/cache_test.go @@ -66,6 +66,7 @@ func TestCacheRoundTrip(t *testing.T) { SecurityGroupid: "sg-mno345", Instanceid: "i-pqr678", PublicDnsName: "ec2-1-2-3-4.compute.amazonaws.com", + AvailabilityZone: "us-west-2c", // New cluster networking fields PublicSubnetid: "subnet-pub789", NatGatewayid: "nat-abc123", @@ -101,6 +102,7 @@ func TestCacheRoundTrip(t *testing.T) { {"SecurityGroupid", restored.SecurityGroupid, original.SecurityGroupid}, {"Instanceid", restored.Instanceid, original.Instanceid}, {"PublicDnsName", restored.PublicDnsName, original.PublicDnsName}, + {"AvailabilityZone", restored.AvailabilityZone, original.AvailabilityZone}, {"PublicSubnetid", restored.PublicSubnetid, original.PublicSubnetid}, {"NatGatewayid", restored.NatGatewayid, original.NatGatewayid}, {"PublicRouteTable", restored.PublicRouteTable, original.PublicRouteTable}, @@ -174,6 +176,9 @@ func TestCacheRoundTripSingleNode(t *testing.T) { if restored.NatGatewayid != "" { t.Errorf("NatGatewayid should be empty, got %q", restored.NatGatewayid) } + if restored.AvailabilityZone != "" { + t.Errorf("AvailabilityZone should be empty, got %q", restored.AvailabilityZone) + } // Original fields must survive if restored.Vpcid != original.Vpcid { diff --git a/pkg/provider/aws/cluster.go b/pkg/provider/aws/cluster.go index bb22efb85..cdc2204f1 100644 --- a/pkg/provider/aws/cluster.go +++ b/pkg/provider/aws/cluster.go @@ -112,7 +112,7 @@ func (p *Provider) CreateCluster() error { return fmt.Errorf("pre-flight check failed: %w", err) } - cache := &ClusterCache{} + cache := &ClusterCache{AWS: AWS{AvailabilityZone: p.selectedAvailabilityZone}} _ = p.updateProgressingCondition(*p.DeepCopy(), &cache.AWS, "v1alpha1.Creating", "Creating multinode cluster resources") diff --git a/pkg/provider/aws/cluster_test.go b/pkg/provider/aws/cluster_test.go index 69dc0c6d2..c50a2c047 100644 --- a/pkg/provider/aws/cluster_test.go +++ b/pkg/provider/aws/cluster_test.go @@ -17,6 +17,7 @@ package aws import ( + "errors" "strings" "testing" @@ -212,7 +213,8 @@ func TestPublicSubnetCreatedInCorrectCIDR(t *testing.T) { provider := newTestProvider(f.EC2) cache := &AWS{ - Vpcid: "vpc-test", + Vpcid: "vpc-test", + AvailabilityZone: "us-west-2a", } if err := provider.createPublicSubnet(cache); err != nil { @@ -243,6 +245,159 @@ func TestPublicSubnetCreatedInCorrectCIDR(t *testing.T) { } } +// Each type alone is offered in other zones; only us-west-2c offers both. +func TestCreateClusterPlacesSubnetsInZoneOfferingAllInstanceTypes(t *testing.T) { + f := awsfake.New() + f.Store.SeedInstanceTypeZones("m5.xlarge", "us-west-2a", "us-west-2c") + f.Store.SeedInstanceTypeZones("g5g.xlarge", "us-west-2b", "us-west-2c", "us-west-2d") + f.Store.FailNext("CreateRouteTable", errors.New("stop after subnets")) + + provider := newTestProvider(f.EC2) + provider.Spec.Cluster = &v1alpha1.ClusterSpec{ + Region: "us-west-2", + ControlPlane: v1alpha1.ControlPlaneSpec{Count: 1, InstanceType: "m5.xlarge"}, + Workers: &v1alpha1.WorkerPoolSpec{Count: 1, InstanceType: "g5g.xlarge"}, + } + + err := provider.CreateCluster() + if err == nil || !strings.Contains(err.Error(), "stop after subnets") { + t.Fatalf("expected CreateCluster() to stop at the public route table, got: %v", err) + } + + subnetZones := requestedSubnetZones(f) + if len(subnetZones) != 2 || subnetZones[0] != "us-west-2c" || subnetZones[1] != "us-west-2c" { + t.Errorf("expected the private and public subnets in us-west-2c, got zones %q", subnetZones) + } +} + +// Unseeded types are offered in every zone, so us-west-2d offers both; the default would be us-west-2a. +func TestCreateClusterHonoursRequestedAvailabilityZone(t *testing.T) { + f := awsfake.New() + f.Store.FailNext("CreateRouteTable", errors.New("stop after subnets")) + + provider := newTestProvider(f.EC2) + provider.Spec.Cluster = &v1alpha1.ClusterSpec{ + Region: "us-west-2", + AvailabilityZone: "us-west-2d", + ControlPlane: v1alpha1.ControlPlaneSpec{Count: 1, InstanceType: "m5.xlarge"}, + Workers: &v1alpha1.WorkerPoolSpec{Count: 1, InstanceType: "g5g.xlarge"}, + } + + err := provider.CreateCluster() + if err == nil || !strings.Contains(err.Error(), "stop after subnets") { + t.Fatalf("expected CreateCluster() to stop at the public route table, got: %v", err) + } + + subnetZones := requestedSubnetZones(f) + if len(subnetZones) != 2 || subnetZones[0] != "us-west-2d" || subnetZones[1] != "us-west-2d" { + t.Errorf("expected the private and public subnets in us-west-2d, got zones %q", subnetZones) + } +} + +// Each type is offered somewhere, just never in the same zone. +func TestCreateClusterRejectsWhenNoZoneOffersAllInstanceTypes(t *testing.T) { + f := awsfake.New() + f.Store.SeedInstanceTypeZones("m5.xlarge", "us-west-2a") + f.Store.SeedInstanceTypeZones("g5g.xlarge", "us-west-2b") + + provider := newTestProvider(f.EC2) + provider.Spec.Cluster = &v1alpha1.ClusterSpec{ + Region: "us-west-2", + ControlPlane: v1alpha1.ControlPlaneSpec{Count: 1, InstanceType: "m5.xlarge"}, + Workers: &v1alpha1.WorkerPoolSpec{Count: 1, InstanceType: "g5g.xlarge"}, + } + + err := provider.CreateCluster() + if err == nil { + t.Fatal("expected CreateCluster() to fail when no zone offers every instance type") + } + for _, want := range []string{"pre-flight check failed", "g5g.xlarge", "m5.xlarge", "us-west-2"} { + if !strings.Contains(err.Error(), want) { + t.Errorf("expected error to mention %q, got: %v", want, err) + } + } + if !f.Store.Empty() { + t.Errorf("resources were created despite the failed pre-flight: %v", f.Store.ResourceCounts()) + } +} + +// A zero-count worker pool still names an instance type, but CreateCluster() launches none of it. +func TestCreateClusterIgnoresWorkerInstanceTypeWithoutWorkers(t *testing.T) { + f := awsfake.New() + f.Store.SeedInstanceTypeZones("m5.xlarge", "us-west-2a") + f.Store.SeedInstanceTypeZones("g5g.xlarge", "us-west-2b") + f.Store.FailNext("CreateRouteTable", errors.New("stop after subnets")) + + provider := newTestProvider(f.EC2) + provider.Spec.Cluster = &v1alpha1.ClusterSpec{ + Region: "us-west-2", + ControlPlane: v1alpha1.ControlPlaneSpec{Count: 1, InstanceType: "m5.xlarge"}, + Workers: &v1alpha1.WorkerPoolSpec{Count: 0, InstanceType: "g5g.xlarge"}, + } + + err := provider.CreateCluster() + if err == nil || !strings.Contains(err.Error(), "stop after subnets") { + t.Fatalf("expected CreateCluster() to stop at the public route table, got: %v", err) + } + + subnetZones := requestedSubnetZones(f) + if len(subnetZones) != 2 || subnetZones[0] != "us-west-2a" || subnetZones[1] != "us-west-2a" { + t.Errorf("expected the private and public subnets in us-west-2a, got zones %q", subnetZones) + } +} + +// createPublicSubnet builds its own CreateSubnetInput, separate from createSubnet. +func TestCreateClusterLetsAWSChooseZoneWithoutZoneDiscoveryPermissions(t *testing.T) { + f := awsfake.New() + f.Store.FailNext("DescribeAvailabilityZones", &apiError{code: "UnauthorizedOperation", message: "not authorized"}) + f.Store.FailNext("CreateRouteTable", errors.New("stop after subnets")) + + provider := newTestProvider(f.EC2) + provider.Spec.Cluster = &v1alpha1.ClusterSpec{ + Region: "us-west-2", + ControlPlane: v1alpha1.ControlPlaneSpec{Count: 1, InstanceType: "m5.xlarge"}, + Workers: &v1alpha1.WorkerPoolSpec{Count: 1, InstanceType: "g5g.xlarge"}, + } + + err := provider.CreateCluster() + if err == nil || !strings.Contains(err.Error(), "stop after subnets") { + t.Fatalf("expected CreateCluster() to stop at the public route table, got: %v", err) + } + + if subnetCount := len(f.Store.Inputs("CreateSubnet")); subnetCount != 2 { + t.Fatalf("expected the private and public subnets to be created, got %d CreateSubnet calls", subnetCount) + } + if subnetRequestsWithZoneCount(f) != 0 { + t.Errorf("expected both subnets to be created without an AvailabilityZone, got zones %q", requestedSubnetZones(f)) + } +} + +func TestCreateClusterRejectsRequestedZoneWithoutZoneDiscoveryPermissions(t *testing.T) { + f := awsfake.New() + f.Store.FailNext("DescribeInstanceTypeOfferings", &apiError{code: "UnauthorizedOperation", message: "not authorized"}) + + provider := newTestProvider(f.EC2) + provider.Spec.Cluster = &v1alpha1.ClusterSpec{ + Region: "us-west-2", + AvailabilityZone: "us-west-2b", + ControlPlane: v1alpha1.ControlPlaneSpec{Count: 1, InstanceType: "m5.xlarge"}, + Workers: &v1alpha1.WorkerPoolSpec{Count: 1, InstanceType: "g5g.xlarge"}, + } + + err := provider.CreateCluster() + if err == nil { + t.Fatal("expected CreateCluster() to fail when the requested zone cannot be validated") + } + for _, want := range []string{"pre-flight check failed", "us-west-2b", "ec2:DescribeInstanceTypeOfferings"} { + if !strings.Contains(err.Error(), want) { + t.Errorf("expected error to mention %q, got: %v", want, err) + } + } + if !f.Store.Empty() { + t.Errorf("resources were created despite the failed pre-flight: %v", f.Store.ResourceCounts()) + } +} + // TestNATGatewayCreatedInPublicSubnet verifies that createNATGateway // places the NAT gateway in the public subnet. func TestNATGatewayCreatedInPublicSubnet(t *testing.T) { diff --git a/pkg/provider/aws/create.go b/pkg/provider/aws/create.go index 477dc7082..5c322d971 100644 --- a/pkg/provider/aws/create.go +++ b/pkg/provider/aws/create.go @@ -61,7 +61,7 @@ func (p *Provider) Create() error { return fmt.Errorf("pre-flight check failed: %w", err) } - cache := new(AWS) + cache := &AWS{AvailabilityZone: p.selectedAvailabilityZone} var cleanupStack []cleanupFunc var err error @@ -215,6 +215,10 @@ func (p *Provider) createVPC(cache *AWS) error { // createSubnet creates a subnet for the VPC func (p *Provider) createSubnet(cache *AWS) error { + if cache.AvailabilityZone == "" && !p.letAWSChooseAvailabilityZone { + return fmt.Errorf("no availability zone selected for subnet; the pre-flight must run first") + } + cancelLoading := p.log.Loading("Creating subnet") subnetInput := &ec2.CreateSubnetInput{ @@ -227,6 +231,9 @@ func (p *Provider) createSubnet(cache *AWS) error { }, }, } + if cache.AvailabilityZone != "" { + subnetInput.AvailabilityZone = aws.String(cache.AvailabilityZone) + } ctx, cancel := context.WithTimeout(context.Background(), defaultSubnetTimeout) defer cancel() @@ -568,6 +575,10 @@ func (p *Provider) createEC2Instance(cache *AWS) error { // createPublicSubnet creates a public subnet (10.0.1.0/24) for NAT gateway and NLB. // The subnet is configured with MapPublicIpOnLaunch enabled. func (p *Provider) createPublicSubnet(cache *AWS) error { + if cache.AvailabilityZone == "" && !p.letAWSChooseAvailabilityZone { + return fmt.Errorf("no availability zone selected for public subnet; the pre-flight must run first") + } + cancelLoading := p.log.Loading("Creating public subnet") // Build tags with a public-specific Name tag @@ -593,6 +604,9 @@ func (p *Provider) createPublicSubnet(cache *AWS) error { }, }, } + if cache.AvailabilityZone != "" { + subnetInput.AvailabilityZone = aws.String(cache.AvailabilityZone) + } ctx, cancel := context.WithTimeout(context.Background(), defaultSubnetTimeout) defer cancel() diff --git a/pkg/provider/aws/create_test.go b/pkg/provider/aws/create_test.go index 8dd9dd402..185886095 100644 --- a/pkg/provider/aws/create_test.go +++ b/pkg/provider/aws/create_test.go @@ -255,7 +255,8 @@ func TestCreateSubnet_Success(t *testing.T) { f := awsfake.New() provider := createTestProvider(f.EC2) cache := &AWS{ - Vpcid: "vpc-test-123", + Vpcid: "vpc-test-123", + AvailabilityZone: "us-west-2a", } err := provider.createSubnet(cache) @@ -272,6 +273,9 @@ func TestCreateSubnet_Success(t *testing.T) { if aws.ToString(call.VpcId) != "vpc-test-123" { t.Errorf("Expected VpcId 'vpc-test-123', got %v", call.VpcId) } + if aws.ToString(call.AvailabilityZone) != "us-west-2a" { + t.Errorf("Expected AvailabilityZone 'us-west-2a', got %v", call.AvailabilityZone) + } // Verify subnet ID was set in cache subnetID := onlyID(t, f.Store.Subnets, "subnet") @@ -286,7 +290,8 @@ func TestCreateSubnet_Error(t *testing.T) { f.Store.FailNext("CreateSubnet", expectedErr) provider := createTestProvider(f.EC2) cache := &AWS{ - Vpcid: "vpc-test-123", + Vpcid: "vpc-test-123", + AvailabilityZone: "us-west-2a", } err := provider.createSubnet(cache) @@ -299,6 +304,20 @@ func TestCreateSubnet_Error(t *testing.T) { } } +func TestCreateSubnet_FailsWithoutSelectedAvailabilityZone(t *testing.T) { + f := awsfake.New() + provider := createTestProvider(f.EC2) + cache := &AWS{Vpcid: "vpc-test-123"} + + err := provider.createSubnet(cache) + if err == nil || !contains(err.Error(), "no availability zone selected") { + t.Fatalf("Expected createSubnet to fail without a selected availability zone, got: %v", err) + } + if f.Store.CallsTo("CreateSubnet") != 0 || len(f.Store.Subnets) != 0 { + t.Errorf("A subnet was created without a selected availability zone: %v", f.Store.ResourceCounts()) + } +} + func TestCreateInternetGateway_Success(t *testing.T) { f := awsfake.New() provider := createTestProvider(f.EC2) @@ -591,7 +610,7 @@ func TestCreate_RejectsUnsupportedInstanceType(t *testing.T) { // (the fake seeds only the types the test configs use), so // checkInstanceTypes must reject it before any resource is created. Type: "p4d.24xlarge", - Region: "us-west-1", + Region: "us-west-2", Image: v1alpha1.Image{ ImageId: aws.String("ami-123"), }, @@ -622,11 +641,261 @@ func TestCreate_RejectsUnsupportedInstanceType(t *testing.T) { } } +// newSingleNodeProvider returns a provider for a single-node environment with +// the given instance spec, backed by the fake. +func newSingleNodeProvider(f *awsfake.Fake, instance v1alpha1.Instance) *Provider { + return &Provider{ + ec2: f.EC2, + log: mockLogger(), + sleep: noopSleep, + Environment: &v1alpha1.Environment{ + ObjectMeta: metav1.ObjectMeta{Name: "test-env"}, + Spec: v1alpha1.EnvironmentSpec{ + Auth: v1alpha1.Auth{KeyName: "test-key"}, + Instance: instance, + }, + }, + Tags: []types.Tag{ + {Key: aws.String("Name"), Value: aws.String("test")}, + }, + } +} + +// requestedSubnetZones returns the AvailabilityZone requested by each +// CreateSubnet call, in call order ("" when none was requested). +func requestedSubnetZones(f *awsfake.Fake) []string { + var subnetZones []string + for _, input := range f.Store.Inputs("CreateSubnet") { + subnetZones = append(subnetZones, aws.ToString(input.(*ec2.CreateSubnetInput).AvailabilityZone)) + } + return subnetZones +} + +// requestedSubnetZones reports an omitted AvailabilityZone as "", so it cannot +// show that the field was left unset. +func subnetRequestsWithZoneCount(f *awsfake.Fake) int { + count := 0 + for _, input := range f.Store.Inputs("CreateSubnet") { + if input.(*ec2.CreateSubnetInput).AvailabilityZone != nil { + count++ + } + } + return count +} + +func TestCreate_PlacesSubnetInZoneOfferingInstanceType(t *testing.T) { + f := awsfake.New() + f.Store.SeedInstanceTypeZones("g5g.xlarge", "us-west-2b", "us-west-2c") + f.Store.FailNext("CreateInternetGateway", errors.New("stop after subnet")) + provider := newSingleNodeProvider(f, v1alpha1.Instance{Type: "g5g.xlarge", Region: "us-west-2"}) + + err := provider.Create() + if err == nil || !contains(err.Error(), "stop after subnet") { + t.Fatalf("Expected Create() to stop at the Internet Gateway, got: %v", err) + } + + subnetZones := requestedSubnetZones(f) + if len(subnetZones) != 1 || subnetZones[0] != "us-west-2b" { + t.Errorf("Expected one subnet in us-west-2b, got zones %q", subnetZones) + } +} + +func TestCreate_RejectsInstanceTypeOfferedInNoZone(t *testing.T) { + f := awsfake.New() + f.Store.SeedInstanceTypeZones("g5g.xlarge") + provider := newSingleNodeProvider(f, v1alpha1.Instance{Type: "g5g.xlarge", Region: "us-west-2"}) + + err := provider.Create() + if err == nil { + t.Fatal("Expected Create() to fail when no zone offers the instance type") + } + for _, expected := range []string{"pre-flight check failed", "g5g.xlarge", "us-west-2"} { + if !contains(err.Error(), expected) { + t.Errorf("Expected error to mention %q, got: %v", expected, err) + } + } + if f.Store.CallsTo("CreateVpc") != 0 || !f.Store.Empty() { + t.Errorf("Resources were created despite the failed pre-flight: %v", f.Store.ResourceCounts()) + } +} + +func TestCreate_FailsPreflightWhenZoneLookupFails(t *testing.T) { + for _, operation := range []string{"DescribeAvailabilityZones", "DescribeInstanceTypeOfferings"} { + for _, testCase := range []struct { + name string + injectedErr error + }{ + {name: "plain error", injectedErr: errors.New(operation + " failed")}, + {name: "API error", injectedErr: &apiError{code: "InternalError", message: operation + " failed"}}, + } { + t.Run(operation+"/"+testCase.name, func(t *testing.T) { + f := awsfake.New() + f.Store.FailNext(operation, testCase.injectedErr) + provider := newSingleNodeProvider(f, v1alpha1.Instance{Type: "t3.medium", Region: "us-west-2"}) + + err := provider.Create() + if err == nil || !contains(err.Error(), "pre-flight check failed") || !contains(err.Error(), testCase.injectedErr.Error()) { + t.Fatalf("Expected Create() to fail the pre-flight with the injected failure, got: %v", err) + } + if f.Store.CallsTo("CreateVpc") != 0 || !f.Store.Empty() { + t.Errorf("Resources were created despite the failed pre-flight: %v", f.Store.ResourceCounts()) + } + }) + } + } +} + +func TestCreate_LetsAWSChooseZoneWithoutZoneDiscoveryPermissions(t *testing.T) { + for _, operation := range []string{"DescribeAvailabilityZones", "DescribeInstanceTypeOfferings"} { + t.Run(operation, func(t *testing.T) { + f := awsfake.New() + f.Store.FailNext(operation, &apiError{code: "UnauthorizedOperation", message: "not authorized"}) + f.Store.FailNext("CreateInternetGateway", errors.New("stop after subnet")) + provider := newSingleNodeProvider(f, v1alpha1.Instance{Type: "t3.medium", Region: "us-west-2"}) + + err := provider.Create() + if err == nil || !contains(err.Error(), "stop after subnet") { + t.Fatalf("Expected Create() to stop at the Internet Gateway, got: %v", err) + } + + if subnetCount := len(f.Store.Inputs("CreateSubnet")); subnetCount != 1 { + t.Fatalf("Expected one CreateSubnet call, got %d", subnetCount) + } + if subnetRequestsWithZoneCount(f) != 0 { + t.Errorf("Expected the subnet to be created without an AvailabilityZone, got zones %q", requestedSubnetZones(f)) + } + }) + } +} + +func TestCreate_RejectsRequestedZoneWithoutZoneDiscoveryPermissions(t *testing.T) { + for _, operation := range []string{"DescribeAvailabilityZones", "DescribeInstanceTypeOfferings"} { + t.Run(operation, func(t *testing.T) { + f := awsfake.New() + f.Store.FailNext(operation, &apiError{code: "UnauthorizedOperation", message: "not authorized"}) + provider := newSingleNodeProvider(f, v1alpha1.Instance{ + Type: "t3.medium", + Region: "us-west-2", + AvailabilityZone: "us-west-2b", + }) + + err := provider.Create() + if err == nil { + t.Fatal("Expected Create() to fail when the requested zone cannot be validated") + } + for _, expected := range []string{ + "pre-flight check failed", "us-west-2b", + "ec2:DescribeAvailabilityZones", "ec2:DescribeInstanceTypeOfferings", + } { + if !contains(err.Error(), expected) { + t.Errorf("Expected error to mention %q, got: %v", expected, err) + } + } + if f.Store.CallsTo("CreateVpc") != 0 || !f.Store.Empty() { + t.Errorf("Resources were created despite the failed pre-flight: %v", f.Store.ResourceCounts()) + } + }) + } +} + +func TestCreate_HonoursRequestedAvailabilityZone(t *testing.T) { + f := awsfake.New() + f.Store.SeedInstanceTypeZones("g5g.xlarge", "us-west-2b", "us-west-2c") + f.Store.FailNext("CreateInternetGateway", errors.New("stop after subnet")) + provider := newSingleNodeProvider(f, v1alpha1.Instance{ + Type: "g5g.xlarge", + Region: "us-west-2", + AvailabilityZone: "us-west-2c", + }) + + err := provider.Create() + if err == nil || !contains(err.Error(), "stop after subnet") { + t.Fatalf("Expected Create() to stop at the Internet Gateway, got: %v", err) + } + + subnetZones := requestedSubnetZones(f) + if len(subnetZones) != 1 || subnetZones[0] != "us-west-2c" { + t.Errorf("Expected one subnet in us-west-2c, got zones %q", subnetZones) + } +} + +func TestCreate_RejectsRequestedZoneNotOfferingInstanceType(t *testing.T) { + f := awsfake.New() + f.Store.SeedInstanceTypeZones("g5g.xlarge", "us-west-2a", "us-west-2b", "us-west-2c") + provider := newSingleNodeProvider(f, v1alpha1.Instance{ + Type: "g5g.xlarge", + Region: "us-west-2", + AvailabilityZone: "us-west-2d", + }) + + err := provider.Create() + if err == nil { + t.Fatal("Expected Create() to fail when the requested zone does not offer the instance type") + } + for _, expected := range []string{"us-west-2d", "g5g.xlarge", "us-west-2a, us-west-2b, us-west-2c"} { + if !contains(err.Error(), expected) { + t.Errorf("Expected error to mention %q, got: %v", expected, err) + } + } + if f.Store.CallsTo("CreateVpc") != 0 || !f.Store.Empty() { + t.Errorf("Resources were created despite the failed pre-flight: %v", f.Store.ResourceCounts()) + } +} + +func TestCreate_SkipsLocalZones(t *testing.T) { + f := awsfake.New() + // us-west-2-lax-1a sorts before us-west-2c, so it would win if not filtered out. + f.Store.SeedAvailabilityZone(types.AvailabilityZone{ + ZoneName: aws.String("us-west-2-lax-1a"), + ZoneType: aws.String("local-zone"), + State: types.AvailabilityZoneStateAvailable, + }) + f.Store.SeedInstanceTypeZones("t3.medium", "us-west-2-lax-1a", "us-west-2c") + f.Store.FailNext("CreateInternetGateway", errors.New("stop after subnet")) + provider := newSingleNodeProvider(f, v1alpha1.Instance{Type: "t3.medium", Region: "us-west-2"}) + + err := provider.Create() + if err == nil || !contains(err.Error(), "stop after subnet") { + t.Fatalf("Expected Create() to stop at the Internet Gateway, got: %v", err) + } + + subnetZones := requestedSubnetZones(f) + if len(subnetZones) != 1 || subnetZones[0] != "us-west-2c" { + t.Errorf("Expected one subnet in us-west-2c, got zones %q", subnetZones) + } +} + +func TestCreate_SkipsZonesThatAreNotAvailable(t *testing.T) { + for _, state := range []types.AvailabilityZoneState{ + types.AvailabilityZoneStateUnavailable, + types.AvailabilityZoneStateConstrained, + } { + t.Run(string(state), func(t *testing.T) { + f := awsfake.New() + f.Store.SeedAvailabilityZoneState("us-west-2a", state) + f.Store.SeedInstanceTypeZones("t3.medium", "us-west-2a", "us-west-2b") + f.Store.FailNext("CreateInternetGateway", errors.New("stop after subnet")) + provider := newSingleNodeProvider(f, v1alpha1.Instance{Type: "t3.medium", Region: "us-west-2"}) + + err := provider.Create() + if err == nil || !contains(err.Error(), "stop after subnet") { + t.Fatalf("Expected Create() to stop at the Internet Gateway, got: %v", err) + } + + subnetZones := requestedSubnetZones(f) + if len(subnetZones) != 1 || subnetZones[0] != "us-west-2b" { + t.Errorf("Expected one subnet in us-west-2b, got zones %q", subnetZones) + } + }) + } +} + func TestCreatePublicSubnet_Success(t *testing.T) { f := awsfake.New() provider := createTestProvider(f.EC2) cache := &AWS{ - Vpcid: "vpc-test-123", + Vpcid: "vpc-test-123", + AvailabilityZone: "us-west-2a", } err := provider.createPublicSubnet(cache) @@ -646,6 +915,9 @@ func TestCreatePublicSubnet_Success(t *testing.T) { if aws.ToString(call.VpcId) != "vpc-test-123" { t.Errorf("Expected VpcId 'vpc-test-123', got %v", call.VpcId) } + if aws.ToString(call.AvailabilityZone) != "us-west-2a" { + t.Errorf("Expected AvailabilityZone 'us-west-2a', got %v", call.AvailabilityZone) + } // Verify public subnet ID was set in cache subnetID := onlyID(t, f.Store.Subnets, "subnet") @@ -666,7 +938,8 @@ func TestCreatePublicSubnet_Error(t *testing.T) { f.Store.FailNext("CreateSubnet", expectedErr) provider := createTestProvider(f.EC2) cache := &AWS{ - Vpcid: "vpc-test-123", + Vpcid: "vpc-test-123", + AvailabilityZone: "us-west-2a", } err := provider.createPublicSubnet(cache) @@ -679,6 +952,20 @@ func TestCreatePublicSubnet_Error(t *testing.T) { } } +func TestCreatePublicSubnet_FailsWithoutSelectedAvailabilityZone(t *testing.T) { + f := awsfake.New() + provider := createTestProvider(f.EC2) + cache := &AWS{Vpcid: "vpc-test-123"} + + err := provider.createPublicSubnet(cache) + if err == nil || !contains(err.Error(), "no availability zone selected") { + t.Fatalf("Expected createPublicSubnet to fail without a selected availability zone, got: %v", err) + } + if f.Store.CallsTo("CreateSubnet") != 0 || len(f.Store.Subnets) != 0 { + t.Errorf("A subnet was created without a selected availability zone: %v", f.Store.ResourceCounts()) + } +} + func TestCreateNATGateway_Success(t *testing.T) { f := awsfake.New() provider := createTestProvider(f.EC2) diff --git a/pkg/provider/aws/helpers_test.go b/pkg/provider/aws/helpers_test.go index 8be2b8117..8ff994188 100644 --- a/pkg/provider/aws/helpers_test.go +++ b/pkg/provider/aws/helpers_test.go @@ -31,3 +31,16 @@ func strPtr(s string) *string { // ErrMockDescribeImages is a sentinel error injected for DescribeImages failures. var ErrMockDescribeImages = fmt.Errorf("mock describe images error") + +type apiError struct { + code string + message string +} + +func (e *apiError) Error() string { + return fmt.Sprintf("api error %s: %s", e.code, e.message) +} + +func (e *apiError) ErrorCode() string { + return e.code +} diff --git a/pkg/provider/aws/image.go b/pkg/provider/aws/image.go index 6d6a08bc2..fa87cdf49 100644 --- a/pkg/provider/aws/image.go +++ b/pkg/provider/aws/image.go @@ -20,6 +20,8 @@ import ( "context" "errors" "fmt" + "maps" + "slices" "sort" "strings" @@ -340,7 +342,7 @@ func (p *Provider) checkInstanceTypes() error { if t := p.Spec.Cluster.ControlPlane.InstanceType; t != "" { needed[t] = false } - if p.Spec.Cluster.Workers != nil { + if p.Spec.Cluster.Workers != nil && p.Spec.Cluster.Workers.Count > 0 { if t := p.Spec.Cluster.Workers.InstanceType; t != "" { needed[t] = false } @@ -375,7 +377,7 @@ func (p *Provider) checkInstanceTypes() error { } } if allFound { - return nil + break } if resp.NextToken != nil { @@ -395,9 +397,119 @@ func (p *Provider) checkInstanceTypes() error { return fmt.Errorf("instance type %s is not supported in the current region %s", instanceType, region) } } + + availabilityZone, err := p.selectAvailabilityZone(slices.Sorted(maps.Keys(needed)), region) + if err != nil { + return err + } + p.selectedAvailabilityZone = availabilityZone + if availabilityZone != "" { + p.log.Info("Using availability zone %s", availabilityZone) + } return nil } +// selectAvailabilityZone returns the zone to create the environment's subnets +// in. Left to itself, AWS picks a subnet's zone without regard to instance +// types, and not every zone of a region offers every type, so RunInstances can +// fail with "Unsupported" after the VPC and its networking already exist. +func (p *Provider) selectAvailabilityZone(instanceTypes []string, region string) (string, error) { + requestedZone := p.Spec.AvailabilityZone + if p.Spec.Cluster != nil { + requestedZone = p.Spec.Cluster.AvailabilityZone + } + joinedInstanceTypes := strings.Join(instanceTypes, ", ") + + offeringZones, err := p.zonesOfferingAllInstanceTypes(instanceTypes) + if isUnauthorizedOperation(err) { + if requestedZone != "" { + return "", fmt.Errorf("cannot validate availability zone %s without the %s permissions: %w", + requestedZone, zoneDiscoveryPermissions, err) + } + p.log.Warning("Cannot pick an availability zone without the %s permissions; AWS will choose the subnet's zone, which may not offer instance type(s) %s", + zoneDiscoveryPermissions, joinedInstanceTypes) + p.letAWSChooseAvailabilityZone = true + return "", nil + } + if err != nil { + return "", err + } + + if len(offeringZones) == 0 { + return "", fmt.Errorf("no availability zone in region %s offers instance type(s) %s", region, joinedInstanceTypes) + } + + if requestedZone == "" { + return offeringZones[0], nil + } + if !slices.Contains(offeringZones, requestedZone) { + return "", fmt.Errorf("availability zone %s does not offer instance type(s) %s; zones in region %s that do: %s", + requestedZone, joinedInstanceTypes, region, strings.Join(offeringZones, ", ")) + } + return requestedZone, nil +} + +// zonesOfferingAllInstanceTypes returns, sorted, the available standard zones of +// the region that offer every one of the given instance types. Local and +// Wavelength Zones are excluded: they sort before the region's own zones and +// support only a subset of AWS services. +func (p *Provider) zonesOfferingAllInstanceTypes(instanceTypes []string) ([]string, error) { + zonesOutput, err := p.ec2.DescribeAvailabilityZones(context.TODO(), &ec2.DescribeAvailabilityZonesInput{ + Filters: []types.Filter{ + {Name: aws.String("zone-type"), Values: []string{"availability-zone"}}, + {Name: aws.String("state"), Values: []string{string(types.AvailabilityZoneStateAvailable)}}, + }, + }) + if err != nil { + return nil, fmt.Errorf("failed to describe availability zones: %w", err) + } + + offeredTypesByZone := make(map[string]map[types.InstanceType]bool) + offeringsInput := &ec2.DescribeInstanceTypeOfferingsInput{ + LocationType: types.LocationTypeAvailabilityZone, + Filters: []types.Filter{ + {Name: aws.String("instance-type"), Values: instanceTypes}, + }, + } + for { + offeringsOutput, err := p.ec2.DescribeInstanceTypeOfferings(context.TODO(), offeringsInput) + if err != nil { + return nil, fmt.Errorf("failed to describe instance type offerings: %w", err) + } + for _, offering := range offeringsOutput.InstanceTypeOfferings { + zoneName := aws.ToString(offering.Location) + if offeredTypesByZone[zoneName] == nil { + offeredTypesByZone[zoneName] = make(map[types.InstanceType]bool) + } + offeredTypesByZone[zoneName][offering.InstanceType] = true + } + if offeringsOutput.NextToken == nil { + break + } + offeringsInput.NextToken = offeringsOutput.NextToken + } + + var offeringZones []string + for _, zone := range zonesOutput.AvailabilityZones { + zoneName := aws.ToString(zone.ZoneName) + if len(offeredTypesByZone[zoneName]) == len(instanceTypes) { + offeringZones = append(offeringZones, zoneName) + } + } + slices.Sort(offeringZones) + return offeringZones, nil +} + +const zoneDiscoveryPermissions = "ec2:DescribeAvailabilityZones and ec2:DescribeInstanceTypeOfferings" + +// isUnauthorizedOperation reports whether err is an AWS API error with the +// UnauthorizedOperation code. It matches smithy.APIError's ErrorCode method +// rather than the type so that smithy-go stays an indirect dependency. +func isUnauthorizedOperation(err error) bool { + var apiErr interface{ ErrorCode() string } + return errors.As(err, &apiErr) && apiErr.ErrorCode() == "UnauthorizedOperation" +} + // normalizeArchToEC2 converts architecture aliases to EC2 canonical form. // EC2 APIs use "x86_64" and "arm64", but users and other systems may use // "amd64" (Debian convention) or "aarch64" (kernel convention). diff --git a/pkg/provider/aws/image_test.go b/pkg/provider/aws/image_test.go index dbcb5c2b7..ea852be94 100644 --- a/pkg/provider/aws/image_test.go +++ b/pkg/provider/aws/image_test.go @@ -746,7 +746,7 @@ func TestDryRun_ArchitectureMismatch(t *testing.T) { Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t3.medium", // x86_64 only - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{ ImageId: aws.String("ami-arm64-image"), Architecture: "arm64", // Mismatched! @@ -791,7 +791,7 @@ func TestDryRun_ArchitectureMatch(t *testing.T) { Provider: v1alpha1.ProviderAWS, Instance: v1alpha1.Instance{ Type: "t4g.medium", // arm64 - Region: "us-east-1", + Region: "us-west-2", Image: v1alpha1.Image{ ImageId: aws.String("ami-arm64-image"), Architecture: "arm64", // Matches! @@ -813,6 +813,44 @@ func TestDryRun_ArchitectureMatch(t *testing.T) { require.NoError(t, err) } +func TestDryRun_LetsAWSChooseZoneWithoutZoneDiscoveryPermissions(t *testing.T) { + f := awsfake.New() + f.Store.FailNext("DescribeAvailabilityZones", &apiError{code: "UnauthorizedOperation", message: "not authorized"}) + f.Store.SetImages(types.Image{ + ImageId: aws.String("ami-x86-image"), + CreationDate: aws.String("2026-01-01T00:00:00.000Z"), + Architecture: types.ArchitectureValuesX8664, + }) + + env := v1alpha1.Environment{ + ObjectMeta: metav1.ObjectMeta{Name: "test-env"}, + Spec: v1alpha1.EnvironmentSpec{ + Provider: v1alpha1.ProviderAWS, + Instance: v1alpha1.Instance{ + Type: "t3.medium", + Region: "us-west-2", + Image: v1alpha1.Image{ + ImageId: aws.String("ami-x86-image"), + Architecture: "x86_64", + }, + }, + Auth: v1alpha1.Auth{ + KeyName: "test-key", + }, + }, + } + + p := &Provider{ + Environment: &env, + ec2: f.EC2, + log: mockLogger(), + } + + err := p.DryRun() + require.NoError(t, err) + assert.Empty(t, p.selectedAvailabilityZone) +} + func TestInferArchFromInstanceType(t *testing.T) { tests := []struct { name string diff --git a/pkg/provider/aws/status.go b/pkg/provider/aws/status.go index f4dc2784a..03a13ee1a 100644 --- a/pkg/provider/aws/status.go +++ b/pkg/provider/aws/status.go @@ -93,6 +93,7 @@ func (p *Provider) updateStatus(env v1alpha1.Environment, cache *AWS, condition {Name: SecurityGroupID, Value: cache.SecurityGroupid}, {Name: InstanceID, Value: cache.Instanceid}, {Name: PublicDnsName, Value: cache.PublicDnsName}, + {Name: AvailabilityZone, Value: cache.AvailabilityZone}, {Name: PublicSubnetID, Value: cache.PublicSubnetid}, {Name: NatGatewayID, Value: cache.NatGatewayid}, {Name: PublicRouteTable, Value: cache.PublicRouteTable}, @@ -145,6 +146,11 @@ func (p *Provider) updateStatus(env v1alpha1.Environment, cache *AWS, condition properties.Value = cache.PublicDnsName modified = true } + case AvailabilityZone: + if properties.Value != cache.AvailabilityZone { + properties.Value = cache.AvailabilityZone + modified = true + } case PublicSubnetID: if properties.Value != cache.PublicSubnetid { properties.Value = cache.PublicSubnetid diff --git a/pkg/testutil/mocks/aws.go b/pkg/testutil/mocks/aws.go index 406f3d027..be2f3efbc 100644 --- a/pkg/testutil/mocks/aws.go +++ b/pkg/testutil/mocks/aws.go @@ -153,9 +153,11 @@ type MockEC2Client struct { DescribeTagsFunc func(ctx context.Context, params *ec2.DescribeTagsInput, optFns ...func(*ec2.Options)) (*ec2.DescribeTagsOutput, error) // Additional operations required by internal/aws.EC2Client - DescribeInternetGatewaysFunc func(ctx context.Context, params *ec2.DescribeInternetGatewaysInput, optFns ...func(*ec2.Options)) (*ec2.DescribeInternetGatewaysOutput, error) - DescribeInstanceTypesFunc func(ctx context.Context, params *ec2.DescribeInstanceTypesInput, optFns ...func(*ec2.Options)) (*ec2.DescribeInstanceTypesOutput, error) - ReplaceRouteTableAssociationFunc func(ctx context.Context, params *ec2.ReplaceRouteTableAssociationInput, optFns ...func(*ec2.Options)) (*ec2.ReplaceRouteTableAssociationOutput, error) + DescribeInternetGatewaysFunc func(ctx context.Context, params *ec2.DescribeInternetGatewaysInput, optFns ...func(*ec2.Options)) (*ec2.DescribeInternetGatewaysOutput, error) + DescribeInstanceTypesFunc func(ctx context.Context, params *ec2.DescribeInstanceTypesInput, optFns ...func(*ec2.Options)) (*ec2.DescribeInstanceTypesOutput, error) + DescribeInstanceTypeOfferingsFunc func(ctx context.Context, params *ec2.DescribeInstanceTypeOfferingsInput, optFns ...func(*ec2.Options)) (*ec2.DescribeInstanceTypeOfferingsOutput, error) + DescribeAvailabilityZonesFunc func(ctx context.Context, params *ec2.DescribeAvailabilityZonesInput, optFns ...func(*ec2.Options)) (*ec2.DescribeAvailabilityZonesOutput, error) + ReplaceRouteTableAssociationFunc func(ctx context.Context, params *ec2.ReplaceRouteTableAssociationInput, optFns ...func(*ec2.Options)) (*ec2.ReplaceRouteTableAssociationOutput, error) // Security Group Revoke operations RevokeSecurityGroupIngressFunc func(ctx context.Context, params *ec2.RevokeSecurityGroupIngressInput, optFns ...func(*ec2.Options)) (*ec2.RevokeSecurityGroupIngressOutput, error) @@ -442,6 +444,20 @@ func (m *MockEC2Client) DescribeInstanceTypes(ctx context.Context, params *ec2.D return &ec2.DescribeInstanceTypesOutput{}, nil } +func (m *MockEC2Client) DescribeInstanceTypeOfferings(ctx context.Context, params *ec2.DescribeInstanceTypeOfferingsInput, optFns ...func(*ec2.Options)) (*ec2.DescribeInstanceTypeOfferingsOutput, error) { //nolint:revive + if m.DescribeInstanceTypeOfferingsFunc != nil { + return m.DescribeInstanceTypeOfferingsFunc(ctx, params, optFns...) + } + return &ec2.DescribeInstanceTypeOfferingsOutput{}, nil +} + +func (m *MockEC2Client) DescribeAvailabilityZones(ctx context.Context, params *ec2.DescribeAvailabilityZonesInput, optFns ...func(*ec2.Options)) (*ec2.DescribeAvailabilityZonesOutput, error) { //nolint:revive + if m.DescribeAvailabilityZonesFunc != nil { + return m.DescribeAvailabilityZonesFunc(ctx, params, optFns...) + } + return &ec2.DescribeAvailabilityZonesOutput{}, nil +} + func (m *MockEC2Client) ReplaceRouteTableAssociation(ctx context.Context, params *ec2.ReplaceRouteTableAssociationInput, optFns ...func(*ec2.Options)) (*ec2.ReplaceRouteTableAssociationOutput, error) { if m.ReplaceRouteTableAssociationFunc != nil { return m.ReplaceRouteTableAssociationFunc(ctx, params, optFns...) diff --git a/tests/e2e_mock_test.go b/tests/e2e_mock_test.go index 5428204dd..4cec269d0 100644 --- a/tests/e2e_mock_test.go +++ b/tests/e2e_mock_test.go @@ -18,7 +18,6 @@ package e2e import ( "fmt" - "os" "path/filepath" "time" @@ -50,6 +49,12 @@ func newMockProvider(cfgFile string) (provider.Provider, *awsfake.Fake, *v1alpha cfg, err := jyaml.UnmarshalFromFile[v1alpha1.Environment](cfgPath) Expect(err).NotTo(HaveOccurred(), "failed to read config %s", cfgPath) cfg.Name += "-" + common.GenerateUID() + // The fake only seeds us-west-2 zones; keep the region consistent with the zone Create records. + if cfg.Spec.Cluster != nil { + cfg.Spec.Cluster.Region = "us-west-2" + } else { + cfg.Spec.Region = "us-west-2" + } cacheFile := filepath.Join(GinkgoT().TempDir(), "cache.yaml") @@ -94,8 +99,19 @@ var _ = Describe("AWS Mock E2E", Label("mock"), func() { "non-HA topology must not provision an NLB: %v", counts) } - _, err := os.Stat(cacheFile) + env, err := jyaml.UnmarshalFromFile[v1alpha1.Environment](cacheFile) Expect(err).NotTo(HaveOccurred(), "Create must write the cache file") + var availabilityZone string + for _, property := range env.Status.Properties { + if property.Name == aws.AvailabilityZone { + availabilityZone = property.Value + } + } + Expect(availabilityZone).NotTo(BeEmpty(), "Create must record the availability zone") + for _, subnet := range fake.Store.Subnets { + Expect(subnet.AvailabilityZone).To(HaveValue(Equal(availabilityZone)), + "the recorded zone must be the one the subnets were created in") + } Expect(p.Delete()).To(Succeed()) Expect(fake.Store.Empty()).To(BeTrue(),