diff --git a/pkg/asset/installconfig/azure/mock/azureclient_generated.go b/pkg/asset/installconfig/azure/mock/azureclient_generated.go index 18eb90b533c..996a85b69bf 100644 --- a/pkg/asset/installconfig/azure/mock/azureclient_generated.go +++ b/pkg/asset/installconfig/azure/mock/azureclient_generated.go @@ -1,9 +1,9 @@ // Code generated by MockGen. DO NOT EDIT. -// Source: client.go +// Source: ./client.go // // Generated by this command: // -// mockgen -source=client.go -destination=mock/azureclient_generated.go -package=mock +// mockgen -source=./client.go -destination=mock/azureclient_generated.go -package=mock // // Package mock is a generated GoMock package. diff --git a/pkg/asset/installconfig/gcp/client.go b/pkg/asset/installconfig/gcp/client.go index a7adb72e732..90d00ab2343 100644 --- a/pkg/asset/installconfig/gcp/client.go +++ b/pkg/asset/installconfig/gcp/client.go @@ -2,6 +2,7 @@ package gcp import ( "context" + "errors" "fmt" "net/http" "strings" @@ -9,7 +10,6 @@ import ( kms "cloud.google.com/go/kms/apiv1" "cloud.google.com/go/kms/apiv1/kmspb" - "github.com/pkg/errors" googleoauth "golang.org/x/oauth2/google" "google.golang.org/api/cloudresourcemanager/v3" compute "google.golang.org/api/compute/v1" @@ -39,6 +39,7 @@ type API interface { GetNetwork(ctx context.Context, network, project string) (*compute.Network, error) GetMachineType(ctx context.Context, project, zone, machineType string) (*compute.MachineType, error) GetMachineTypeWithZones(ctx context.Context, project, region, machineType string) (*compute.MachineType, sets.Set[string], error) + GetDiskTypeWithZones(ctx context.Context, project, region, diskType string) (*compute.DiskType, sets.Set[string], error) GetPublicDomains(ctx context.Context, project string) ([]string, error) GetDNSZone(ctx context.Context, project, baseDomain string, isPublic bool) (*dns.ManagedZone, error) GetDNSZoneFromParams(ctx context.Context, params gcptypes.DNSZoneParams) (*dns.ManagedZone, error) @@ -72,7 +73,7 @@ type Client struct { func NewClient(ctx context.Context, endpoint *gcptypes.PSCEndpoint) (*Client, error) { ssn, err := GetSession(ctx) if err != nil { - return nil, errors.Wrap(err, "failed to get session") + return nil, fmt.Errorf("failed to get session: %w", err) } endpointName := "" @@ -223,6 +224,50 @@ func (c *Client) GetMachineTypeWithZones(ctx context.Context, project, region, m return machines[0], zones, nil } +// GetDiskTypeWithZones retrieves the specified disk type and the zones in which it is available. +// It queries each zone individually because DiskTypes.AggregatedList may not +// return results on sovereign clouds. +func (c *Client) GetDiskTypeWithZones(ctx context.Context, project, region, diskType string) (*compute.DiskType, sets.Set[string], error) { + svc, err := c.getComputeService(ctx) + if err != nil { + return nil, nil, err + } + + pz, err := GetZones(ctx, svc, project, region) + if err != nil { + return nil, nil, err + } + if len(pz) == 0 { + return nil, nil, fmt.Errorf("failed to find zones in project %s region %s", project, region) + } + + ctx, cancel := context.WithTimeout(ctx, defaultTimeout) + defer cancel() + + var found *compute.DiskType + zones := sets.New[string]() + for _, zone := range pz { + dt, err := svc.DiskTypes.Get(project, zone.Name, diskType).Context(ctx).Do() + if err != nil { + var gerr *googleapi.Error + if errors.As(err, &gerr) && gerr.Code == http.StatusNotFound { + continue + } + return nil, nil, err + } + if found == nil { + found = dt + } + zones.Insert(zone.Name) + } + + if found == nil { + return nil, nil, nil + } + + return found, zones, nil +} + // GetNetwork uses the GCP Compute Service API to get a network by name from a project. func (c *Client) GetNetwork(ctx context.Context, network, project string) (*compute.Network, error) { svc, err := c.getComputeService(ctx) @@ -234,7 +279,7 @@ func (c *Client) GetNetwork(ctx context.Context, network, project string) (*comp defer cancel() res, err := svc.Networks.Get(project, network).Context(ctx).Do() if err != nil { - return nil, errors.Wrapf(err, "failed to get network %s", network) + return nil, fmt.Errorf("failed to get network %s: %w", network, err) } return res, nil } @@ -526,7 +571,7 @@ func (c *Client) GetRegions(ctx context.Context, project string) ([]string, erro defer cancel() gcpRegionsList, err := svc.Regions.List(project).Context(ctx).Do() if err != nil { - return nil, errors.Wrapf(err, "failed to get regions for project") + return nil, fmt.Errorf("failed to get regions for project: %w", err) } computeRegions := make([]string, 0, len(gcpRegionsList.Items)) @@ -551,7 +596,7 @@ func GetZones(ctx context.Context, svc *compute.Service, project, region string) } return nil }); err != nil { - return nil, errors.Wrapf(err, "failed to get zones from project %s", project) + return nil, fmt.Errorf("failed to get zones from project %s: %w", project, err) } return zones, nil } @@ -601,7 +646,7 @@ func (c *Client) GetServiceAccount(ctx context.Context, project, serviceAccount } svc, err := GetIAMService(ctx, opts...) if err != nil { - return "", errors.Wrapf(err, "failed create IAM service") + return "", fmt.Errorf("failed create IAM service: %w", err) } ctx, cancel := context.WithTimeout(ctx, 1*time.Minute) @@ -610,7 +655,7 @@ func (c *Client) GetServiceAccount(ctx context.Context, project, serviceAccount fullServiceAccountPath := fmt.Sprintf("projects/%s/serviceAccounts/%s", project, serviceAccount) rsp, err := svc.Projects.ServiceAccounts.Get(fullServiceAccountPath).Context(ctx).Do() if err != nil { - return "", errors.Wrapf(err, "failed to find resource %s", fullServiceAccountPath) + return "", fmt.Errorf("failed to find resource %s: %w", fullServiceAccountPath, err) } return rsp.Name, nil } @@ -639,14 +684,14 @@ func (c *Client) getPermissions(ctx context.Context, project string, permissions service, err := c.getCloudResourceService(ctx) if err != nil { - return nil, errors.Wrapf(err, "failed to get cloud resource manager service") + return nil, fmt.Errorf("failed to get cloud resource manager service: %w", err) } projectsService := cloudresourcemanager.NewProjectsService(service) rb := &cloudresourcemanager.TestIamPermissionsRequest{Permissions: permissions} response, err := projectsService.TestIamPermissions(fmt.Sprintf(gcpconsts.ProjectNameFmt, project), rb).Context(ctx).Do() if err != nil { - return nil, errors.Wrapf(err, "failed to get Iam permissions") + return nil, fmt.Errorf("failed to get Iam permissions: %w", err) } return response.Permissions, nil diff --git a/pkg/asset/installconfig/gcp/mock/gcpclient_generated.go b/pkg/asset/installconfig/gcp/mock/gcpclient_generated.go index 205cacb7483..04e4a889fcd 100644 --- a/pkg/asset/installconfig/gcp/mock/gcpclient_generated.go +++ b/pkg/asset/installconfig/gcp/mock/gcpclient_generated.go @@ -106,6 +106,22 @@ func (mr *MockAPIMockRecorder) GetDNSZoneFromParams(ctx, params any) *gomock.Cal return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetDNSZoneFromParams", reflect.TypeOf((*MockAPI)(nil).GetDNSZoneFromParams), ctx, params) } +// GetDiskTypeWithZones mocks base method. +func (m *MockAPI) GetDiskTypeWithZones(ctx context.Context, project, region, diskType string) (*compute.DiskType, sets.Set[string], error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetDiskTypeWithZones", ctx, project, region, diskType) + ret0, _ := ret[0].(*compute.DiskType) + ret1, _ := ret[1].(sets.Set[string]) + ret2, _ := ret[2].(error) + return ret0, ret1, ret2 +} + +// GetDiskTypeWithZones indicates an expected call of GetDiskTypeWithZones. +func (mr *MockAPIMockRecorder) GetDiskTypeWithZones(ctx, project, region, diskType any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetDiskTypeWithZones", reflect.TypeOf((*MockAPI)(nil).GetDiskTypeWithZones), ctx, project, region, diskType) +} + // GetEnabledServices mocks base method. func (m *MockAPI) GetEnabledServices(ctx context.Context, project string) ([]string, error) { m.ctrl.T.Helper() diff --git a/pkg/asset/installconfig/gcp/validation.go b/pkg/asset/installconfig/gcp/validation.go index 12c268b9a17..39393e8d66a 100644 --- a/pkg/asset/installconfig/gcp/validation.go +++ b/pkg/asset/installconfig/gcp/validation.go @@ -3,12 +3,12 @@ package gcp import ( "context" "encoding/json" + "errors" "fmt" "net" "slices" "strings" - "github.com/pkg/errors" "github.com/sirupsen/logrus" compute "google.golang.org/api/compute/v1" "google.golang.org/api/dns/v1" @@ -140,7 +140,8 @@ func ValidateInstanceType(client API, fieldPath *field.Path, project, region str typeMeta, typeZones, err := client.GetMachineTypeWithZones(context.TODO(), project, region, instanceType) if err != nil { - if _, ok := err.(*googleapi.Error); ok { + var gerr *googleapi.Error + if errors.As(err, &gerr) { return append(allErrs, field.Invalid(fieldPath.Child("type"), instanceType, err.Error())) } return append(allErrs, field.InternalError(nil, err)) @@ -162,6 +163,9 @@ func ValidateInstanceType(client API, fieldPath *field.Path, project, region str if len(userZones) == 0 { userZones = typeZones } + + allErrs = append(allErrs, validateDiskTypeAvailability(client, fieldPath, project, region, userZones, diskType)...) + if diff := userZones.Difference(typeZones); len(diff) > 0 { errMsg := fmt.Sprintf("instance type not available in zones: %v", sets.List(diff)) allErrs = append(allErrs, field.Invalid(fieldPath.Child("type"), instanceType, errMsg)) @@ -186,6 +190,37 @@ func ValidateInstanceType(client API, fieldPath *field.Path, project, region str return allErrs } +func validateDiskTypeAvailability(client API, fieldPath *field.Path, project, region string, zones sets.Set[string], diskType string) field.ErrorList { + allErrs := field.ErrorList{} + + if diskType == "" { + return allErrs + } + + dt, dtZones, err := client.GetDiskTypeWithZones(context.TODO(), project, region, diskType) + if err != nil { + var gerr *googleapi.Error + if errors.As(err, &gerr) { + return append(allErrs, field.Invalid(fieldPath.Child("diskType"), diskType, err.Error())) + } + return append(allErrs, field.InternalError(fieldPath.Child("diskType"), err)) + } + + if dt == nil { + errMsg := fmt.Sprintf("disk type %s is not available in region %s", diskType, region) + return append(allErrs, field.Invalid(fieldPath.Child("diskType"), diskType, errMsg)) + } + + if len(zones) > 0 { + if diff := zones.Difference(dtZones); len(diff) > 0 { + errMsg := fmt.Sprintf("disk type %s is not available in zones: %v", diskType, sets.List(diff)) + allErrs = append(allErrs, field.Invalid(fieldPath.Child("diskType"), diskType, errMsg)) + } + } + + return allErrs +} + func validateServiceAccountPresent(client API, ic *types.InstallConfig) field.ErrorList { allErrs := field.ErrorList{} @@ -662,7 +697,7 @@ func ValidateEnabledServices(ctx context.Context, client API, project string) er projectServices, err := client.GetEnabledServices(ctx, project) if err != nil { if IsForbidden(err) { - return errors.Wrap(err, "unable to fetch enabled services for project. Make sure 'serviceusage.googleapis.com' is enabled") + return fmt.Errorf("unable to fetch enabled services for project. Make sure 'serviceusage.googleapis.com' is enabled: %w", err) } return err } diff --git a/pkg/asset/installconfig/gcp/validation_test.go b/pkg/asset/installconfig/gcp/validation_test.go index c58cd3a7df5..b5747f88cb0 100644 --- a/pkg/asset/installconfig/gcp/validation_test.go +++ b/pkg/asset/installconfig/gcp/validation_test.go @@ -489,6 +489,10 @@ func TestGCPInstallConfigValidation(t *testing.T) { // When passed incorrect machine type, the API returns nil. gcpClient.EXPECT().GetMachineTypeWithZones(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil, nil, fmt.Errorf("404")).AnyTimes() + // Mock disk type availability - all disk types available in valid zones + gcpClient.EXPECT().GetDiskTypeWithZones(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()). + Return(&compute.DiskType{Name: "available"}, sets.New(validZone), nil).AnyTimes() + // When passed the correct network & project, return an empty network, which should be enough to validate ok. gcpClient.EXPECT().GetNetwork(gomock.Any(), validNetworkName, validProjectName).Return(&compute.Network{}, nil).AnyTimes() @@ -1061,7 +1065,7 @@ func TestValidateInstanceType(t *testing.T) { onHostMaintenance: "Migrate", confidentialCompute: "Disabled", expectedError: true, - expectedErrMsg: `\[instance.type: Invalid value: "n1\-standard\-4": instance type not available in zones: \[x y\]\]$`, + expectedErrMsg: `\[instance.diskType: Invalid value: "pd\-ssd": disk type pd\-ssd is not available in zones: \[x y\] instance.type: Invalid value: "n1\-standard\-4": instance type not available in zones: \[x y\]\]$`, }, { name: "Valid instance fails min requirements and no zones specified", @@ -1369,6 +1373,10 @@ func TestValidateInstanceType(t *testing.T) { // When passed incorrect machine type, the API returns nil. gcpClient.EXPECT().GetMachineTypeWithZones(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil, nil, fmt.Errorf("404")).AnyTimes() + // Mock disk type availability - all disk types available in all test zones + gcpClient.EXPECT().GetDiskTypeWithZones(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()). + Return(&compute.DiskType{Name: "available"}, sets.New("a", "b", "c", "d"), nil).AnyTimes() + for _, test := range cases { t.Run(test.name, func(t *testing.T) { errs := ValidateInstanceType(gcpClient, field.NewPath("instance"), "project-id", "region", test.zones, test.diskType, test.instanceType, controlPlaneReq, test.arch, test.onHostMaintenance, test.confidentialCompute) @@ -1381,6 +1389,94 @@ func TestValidateInstanceType(t *testing.T) { } } +func TestValidateDiskTypeAvailability(t *testing.T) { + cases := []struct { + name string + diskType string + zones sets.Set[string] + mockDiskType *compute.DiskType + mockZones sets.Set[string] + mockErr error + expectedError bool + expectedErrMsg string + }{ + { + name: "Empty disk type is a no-op", + diskType: "", + expectedError: false, + }, + { + name: "Disk type available in all zones", + diskType: "pd-ssd", + zones: sets.New("us-east1-b", "us-east1-c"), + mockDiskType: &compute.DiskType{Name: "pd-ssd"}, + mockZones: sets.New("us-east1-b", "us-east1-c", "us-east1-d"), + expectedError: false, + }, + { + name: "Disk type not available in region", + diskType: "hyperdisk-balanced", + zones: sets.New[string](), + mockDiskType: nil, + mockZones: nil, + expectedError: true, + expectedErrMsg: `disk type hyperdisk-balanced is not available in region us-east1`, + }, + { + name: "Disk type not available in specific zones", + diskType: "pd-ssd", + zones: sets.New("us-east1-b", "us-east1-x"), + mockDiskType: &compute.DiskType{Name: "pd-ssd"}, + mockZones: sets.New("us-east1-b", "us-east1-c"), + expectedError: true, + expectedErrMsg: `disk type pd-ssd is not available in zones: \[us-east1-x\]`, + }, + { + name: "Disk type available in fewer zones than requested", + diskType: "pd-ssd", + zones: sets.New("us-east1-b", "us-east1-c", "us-east1-d"), + mockDiskType: &compute.DiskType{Name: "pd-ssd"}, + mockZones: sets.New("us-east1-b", "us-east1-c"), + expectedError: true, + expectedErrMsg: `disk type pd-ssd is not available in zones: \[us-east1-d\]`, + }, + { + name: "GCP API error returns field error", + diskType: "pd-ssd", + mockErr: &googleapi.Error{Code: http.StatusForbidden, Message: "forbidden"}, + expectedError: true, + expectedErrMsg: `forbidden`, + }, + { + name: "Non-API error returns internal error", + diskType: "pd-ssd", + mockErr: fmt.Errorf("network timeout"), + expectedError: true, + expectedErrMsg: `network timeout`, + }, + } + + for _, test := range cases { + t.Run(test.name, func(t *testing.T) { + mockCtrl := gomock.NewController(t) + defer mockCtrl.Finish() + gcpClient := mock.NewMockAPI(mockCtrl) + + if test.diskType != "" { + gcpClient.EXPECT().GetDiskTypeWithZones(gomock.Any(), "project-id", "us-east1", test.diskType). + Return(test.mockDiskType, test.mockZones, test.mockErr).AnyTimes() + } + + errs := validateDiskTypeAvailability(gcpClient, field.NewPath("test"), "project-id", "us-east1", test.zones, test.diskType) + if test.expectedError { + assert.Regexp(t, test.expectedErrMsg, errs.ToAggregate().Error()) + } else { + assert.Empty(t, errs) + } + }) + } +} + func TestValidateMarketplaceImages(t *testing.T) { var ( validImage = "valid-image"