diff --git a/pkg/asset/installconfig/gcp/client.go b/pkg/asset/installconfig/gcp/client.go new file mode 100644 index 00000000000..3eaf5218f2b --- /dev/null +++ b/pkg/asset/installconfig/gcp/client.go @@ -0,0 +1,153 @@ +package gcp + +import ( + "context" + "fmt" + "strings" + "time" + + "github.com/pkg/errors" + compute "google.golang.org/api/compute/v1" + dns "google.golang.org/api/dns/v1" + "google.golang.org/api/option" +) + +//go:generate mockgen -source=./client.go -destination=.mock/gcpclient_generated.go -package=mock + +// API represents the calls made to the API. +type API interface { + GetNetwork(ctx context.Context, network, project string) (*compute.Network, error) + GetPublicDomains(ctx context.Context, project string) ([]string, error) + GetPublicDNSZone(ctx context.Context, baseDomain, project string) (*dns.ManagedZone, error) + GetSubnetworks(ctx context.Context, network, project, region string) ([]*compute.Subnetwork, error) +} + +// Client makes calls to the GCP API. +type Client struct { + ssn *Session +} + +// NewClient initializes a client with a session. +func NewClient(ctx context.Context) (*Client, error) { + ctx, cancel := context.WithTimeout(ctx, 1*time.Minute) + defer cancel() + + ssn, err := GetSession(ctx) + if err != nil { + return nil, errors.Wrap(err, "failed to get session") + } + + client := &Client{ + ssn: ssn, + } + return client, 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) { + ctx, cancel := context.WithTimeout(ctx, 1*time.Minute) + defer cancel() + + svc, err := c.getComputeService(ctx) + if err != nil { + return nil, err + } + 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 res, nil +} + +// GetPublicDomains returns all of the domains from among the project's public DNS zones. +func (c *Client) GetPublicDomains(ctx context.Context, project string) ([]string, error) { + ctx, cancel := context.WithTimeout(context.TODO(), 1*time.Minute) + defer cancel() + + svc, err := c.getDNSService(ctx) + if err != nil { + return []string{}, err + } + + var publicZones []string + req := svc.ManagedZones.List(project).Context(ctx) + if err := req.Pages(ctx, func(page *dns.ManagedZonesListResponse) error { + for _, v := range page.ManagedZones { + if v.Visibility != "private" { + publicZones = append(publicZones, strings.TrimSuffix(v.DnsName, ".")) + } + } + return nil + }); err != nil { + return publicZones, err + } + return publicZones, nil +} + +// GetPublicDNSZone returns a public DNS zone for a basedomain. +func (c *Client) GetPublicDNSZone(ctx context.Context, project, baseDomain string) (*dns.ManagedZone, error) { + ctx, cancel := context.WithTimeout(context.TODO(), 1*time.Minute) + defer cancel() + + svc, err := c.getDNSService(ctx) + if err != nil { + return nil, err + } + + req := svc.ManagedZones.List(project).DnsName(baseDomain).Context(ctx) + var res *dns.ManagedZone + if err := req.Pages(ctx, func(page *dns.ManagedZonesListResponse) error { + for idx, v := range page.ManagedZones { + if v.Visibility != "private" { + res = page.ManagedZones[idx] + } + } + return nil + }); err != nil { + return nil, errors.Wrap(err, "failed to list DNS Zones") + } + if res == nil { + return nil, errors.New("no matching public DNS Zone found") + } + return res, nil +} + +// GetSubnetworks uses the GCP Compute Service API to retrieve all subnetworks in a given network. +func (c *Client) GetSubnetworks(ctx context.Context, network, project, region string) ([]*compute.Subnetwork, error) { + ctx, cancel := context.WithTimeout(ctx, 1*time.Minute) + defer cancel() + + svc, err := c.getComputeService(ctx) + if err != nil { + return nil, err + } + + filter := fmt.Sprintf("network eq .*%s", network) + req := svc.Subnetworks.List(project, region).Filter(filter) + var res []*compute.Subnetwork + if err := req.Pages(ctx, func(page *compute.SubnetworkList) error { + for _, subnet := range page.Items { + res = append(res, subnet) + } + return nil + }); err != nil { + return nil, err + } + return res, nil +} + +func (c *Client) getComputeService(ctx context.Context) (*compute.Service, error) { + svc, err := compute.NewService(ctx, option.WithCredentials(c.ssn.Credentials)) + if err != nil { + return nil, errors.Wrap(err, "failed to create compute service") + } + return svc, nil +} + +func (c *Client) getDNSService(ctx context.Context) (*dns.Service, error) { + svc, err := dns.NewService(ctx, option.WithCredentials(c.ssn.Credentials)) + if err != nil { + return nil, errors.Wrap(err, "failed to create dns service") + } + return svc, nil +} diff --git a/pkg/asset/installconfig/gcp/dns.go b/pkg/asset/installconfig/gcp/dns.go index 3c41dc25f2a..590c86ab9c3 100644 --- a/pkg/asset/installconfig/gcp/dns.go +++ b/pkg/asset/installconfig/gcp/dns.go @@ -10,77 +10,42 @@ import ( "github.com/pkg/errors" dns "google.golang.org/api/dns/v1" "google.golang.org/api/googleapi" - "google.golang.org/api/option" survey "gopkg.in/AlecAivazis/survey.v1" ) -func getDNSService(ctx context.Context) (*dns.Service, error) { - ssn, err := GetSession(ctx) - if err != nil { - return nil, errors.Wrap(err, "failed to get session") - } - - svc, err := dns.NewService(ctx, option.WithCredentials(ssn.Credentials)) - if err != nil { - return nil, errors.Wrap(err, "failed to create compute service") - } - return svc, nil -} - // GetPublicZone returns a DNS managed zone from the provided project which matches the baseDomain // If multiple zones match the basedomain, it uses the last public zone in the list as provided by the GCP API. func GetPublicZone(ctx context.Context, project, baseDomain string) (*dns.ManagedZone, error) { - ctx, cancel := context.WithTimeout(ctx, 1*time.Minute) - defer cancel() - - svc, err := getDNSService(ctx) + client, err := NewClient(context.TODO()) if err != nil { return nil, err } + ctx, cancel := context.WithTimeout(context.Background(), 1*time.Minute) + defer cancel() if !strings.HasSuffix(baseDomain, ".") { baseDomain = fmt.Sprintf("%s.", baseDomain) } - req := svc.ManagedZones.List(project).DnsName(baseDomain).Context(ctx) - var res *dns.ManagedZone - if err := req.Pages(ctx, func(page *dns.ManagedZonesListResponse) error { - for idx, v := range page.ManagedZones { - if v.Visibility != "private" { - res = page.ManagedZones[idx] - } - } - return nil - }); err != nil { - return nil, errors.Wrap(err, "failed to list DNS Zones") - } - if res == nil { - return nil, errors.New("no matching public DNS Zone found") + dnsZone, err := client.GetPublicDNSZone(ctx, project, baseDomain) + if err != nil { + return nil, err } - return res, nil + return dnsZone, nil } // GetBaseDomain returns a base domain chosen from among the project's public DNS zones. func GetBaseDomain(project string) (string, error) { - ctx, cancel := context.WithTimeout(context.TODO(), 1*time.Minute) - defer cancel() - - svc, err := getDNSService(ctx) + client, err := NewClient(context.TODO()) if err != nil { return "", err } + ctx, cancel := context.WithTimeout(context.Background(), 1*time.Minute) + defer cancel() - var publicZones []string - req := svc.ManagedZones.List(project).Context(ctx) - if err := req.Pages(ctx, func(page *dns.ManagedZonesListResponse) error { - for _, v := range page.ManagedZones { - if v.Visibility != "private" { - publicZones = append(publicZones, strings.TrimSuffix(v.DnsName, ".")) - } - } - return nil - }); err != nil { - return "", err + publicZones, err := client.GetPublicDomains(ctx, project) + if err != nil { + return "", errors.Wrap(err, "could not retrieve base domains") } if len(publicZones) == 0 { return "", errors.New("no domain names found in project") diff --git a/pkg/asset/installconfig/gcp/mock/gcpclient_generated.go b/pkg/asset/installconfig/gcp/mock/gcpclient_generated.go new file mode 100644 index 00000000000..35041058b1d --- /dev/null +++ b/pkg/asset/installconfig/gcp/mock/gcpclient_generated.go @@ -0,0 +1,96 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: pkg/asset/installconfig/gcp/client.go + +// Package mock is a generated GoMock package. +package mock + +import ( + context "context" + gomock "github.com/golang/mock/gomock" + v1 "google.golang.org/api/compute/v1" + v10 "google.golang.org/api/dns/v1" + reflect "reflect" +) + +// MockAPI is a mock of API interface +type MockAPI struct { + ctrl *gomock.Controller + recorder *MockAPIMockRecorder +} + +// MockAPIMockRecorder is the mock recorder for MockAPI +type MockAPIMockRecorder struct { + mock *MockAPI +} + +// NewMockAPI creates a new mock instance +func NewMockAPI(ctrl *gomock.Controller) *MockAPI { + mock := &MockAPI{ctrl: ctrl} + mock.recorder = &MockAPIMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use +func (m *MockAPI) EXPECT() *MockAPIMockRecorder { + return m.recorder +} + +// GetNetwork mocks base method +func (m *MockAPI) GetNetwork(ctx context.Context, network, project string) (*v1.Network, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetNetwork", ctx, network, project) + ret0, _ := ret[0].(*v1.Network) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetNetwork indicates an expected call of GetNetwork +func (mr *MockAPIMockRecorder) GetNetwork(ctx, network, project interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetNetwork", reflect.TypeOf((*MockAPI)(nil).GetNetwork), ctx, network, project) +} + +// GetPublicDomains mocks base method +func (m *MockAPI) GetPublicDomains(ctx context.Context, project string) ([]string, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetPublicDomains", ctx, project) + ret0, _ := ret[0].([]string) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetPublicDomains indicates an expected call of GetPublicDomains +func (mr *MockAPIMockRecorder) GetPublicDomains(ctx, project interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPublicDomains", reflect.TypeOf((*MockAPI)(nil).GetPublicDomains), ctx, project) +} + +// GetPublicDNSZone mocks base method +func (m *MockAPI) GetPublicDNSZone(ctx context.Context, baseDomain, project string) (*v10.ManagedZone, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetPublicDNSZone", ctx, baseDomain, project) + ret0, _ := ret[0].(*v10.ManagedZone) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetPublicDNSZone indicates an expected call of GetPublicDNSZone +func (mr *MockAPIMockRecorder) GetPublicDNSZone(ctx, baseDomain, project interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPublicDNSZone", reflect.TypeOf((*MockAPI)(nil).GetPublicDNSZone), ctx, baseDomain, project) +} + +// GetSubnetworks mocks base method +func (m *MockAPI) GetSubnetworks(ctx context.Context, network, project, region string) ([]*v1.Subnetwork, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetSubnetworks", ctx, network, project, region) + ret0, _ := ret[0].([]*v1.Subnetwork) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetSubnetworks indicates an expected call of GetSubnetworks +func (mr *MockAPIMockRecorder) GetSubnetworks(ctx, network, project, region interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetSubnetworks", reflect.TypeOf((*MockAPI)(nil).GetSubnetworks), ctx, network, project, region) +} diff --git a/pkg/asset/installconfig/gcp/validation.go b/pkg/asset/installconfig/gcp/validation.go new file mode 100644 index 00000000000..0e53f9931f3 --- /dev/null +++ b/pkg/asset/installconfig/gcp/validation.go @@ -0,0 +1,74 @@ +package gcp + +import ( + "context" + "fmt" + "net" + + compute "google.golang.org/api/compute/v1" + "k8s.io/apimachinery/pkg/util/validation/field" + + "github.com/openshift/installer/pkg/types" +) + +// Validate executes platform-specific validation. +func Validate(client API, ic *types.InstallConfig) error { + allErrs := field.ErrorList{} + + allErrs = append(allErrs, validateNetworks(client, ic, field.NewPath("platform").Child("gcp"))...) + + return allErrs.ToAggregate() +} + +// validateNetworks checks that the user-provided VPC is in the project and the provided subnets are valid. +func validateNetworks(client API, ic *types.InstallConfig, fieldPath *field.Path) field.ErrorList { + allErrs := field.ErrorList{} + + if ic.GCP.Network != "" { + _, err := client.GetNetwork(context.TODO(), ic.GCP.Network, ic.GCP.ProjectID) + if err != nil { + return append(allErrs, field.Invalid(fieldPath.Child("network"), ic.GCP.Network, err.Error())) + } + + subnets, err := client.GetSubnetworks(context.TODO(), ic.GCP.Network, ic.GCP.ProjectID, ic.GCP.Region) + if err != nil { + return append(allErrs, field.Invalid(fieldPath.Child("network"), ic.GCP.Network, "failed to retrieve subnets")) + } + + allErrs = append(allErrs, validateSubnet(client, ic, fieldPath.Child("computeSubnet"), subnets, ic.GCP.ComputeSubnet)...) + allErrs = append(allErrs, validateSubnet(client, ic, fieldPath.Child("controlPlaneSubnet"), subnets, ic.GCP.ControlPlaneSubnet)...) + } + + return allErrs +} + +func validateSubnet(client API, ic *types.InstallConfig, fieldPath *field.Path, subnets []*compute.Subnetwork, name string) field.ErrorList { + allErrs := field.ErrorList{} + machineCIDR := ic.Networking.MachineCIDR + + subnet, errMsg := findSubnet(subnets, name, ic.GCP.Network, ic.GCP.Region) + if subnet == nil { + return append(allErrs, field.Invalid(fieldPath, name, errMsg)) + } + + subnetIP, _, err := net.ParseCIDR(subnet.IpCidrRange) + if err != nil { + return append(allErrs, field.Invalid(fieldPath, name, "unable to parse subnet CIDR")) + } + + if !machineCIDR.Contains(subnetIP) { + errMsg := fmt.Sprintf("subnet %v has an IP address range %v outside of the MachineCIDR %v", name, subnet.IpCidrRange, machineCIDR) + return append(allErrs, field.Invalid(fieldPath, name, errMsg)) + } + return nil +} + +// findSubnet checks that the subnets are in the provided VPC and region. +func findSubnet(subnets []*compute.Subnetwork, userSubnet, network, region string) (*compute.Subnetwork, string) { + for _, vpcSubnet := range subnets { + if userSubnet == vpcSubnet.Name { + return vpcSubnet, "" + } + } + return nil, fmt.Sprintf("could not find subnet %s in network %s and region %s", userSubnet, network, region) +} diff --git a/pkg/asset/installconfig/gcp/validation_test.go b/pkg/asset/installconfig/gcp/validation_test.go new file mode 100644 index 00000000000..6c7558c2adf --- /dev/null +++ b/pkg/asset/installconfig/gcp/validation_test.go @@ -0,0 +1,173 @@ +package gcp + +import ( + "fmt" + "net" + "testing" + + "github.com/golang/mock/gomock" + "github.com/stretchr/testify/assert" + compute "google.golang.org/api/compute/v1" + + "github.com/openshift/installer/pkg/asset/installconfig/gcp/mock" + "github.com/openshift/installer/pkg/ipnet" + "github.com/openshift/installer/pkg/types" + "github.com/openshift/installer/pkg/types/gcp" +) + +type editFunctions []func(ic *types.InstallConfig) + +var ( + validNetworkName = "valid-vpc" + validProjectName = "valid-project" + validRegion = "us-east1" + validComputeSubnet = "valid-compute-subnet" + validCPSubnet = "valid-controlplane-subnet" + validCIDR = "10.0.0.0/16" + + invalidateMachineCIDR = func(ic *types.InstallConfig) { + _, newCidr, _ := net.ParseCIDR("192.168.111.0/24") + ic.MachineCIDR = &ipnet.IPNet{IPNet: *newCidr} + } + + invalidateNetwork = func(ic *types.InstallConfig) { ic.GCP.Network = "invalid-vpc" } + invalidateComputeSubnet = func(ic *types.InstallConfig) { ic.GCP.ComputeSubnet = "invalid-compute-subnet" } + invalidateCPSubnet = func(ic *types.InstallConfig) { ic.GCP.ControlPlaneSubnet = "invalid-cp-subnet" } + invalidateRegion = func(ic *types.InstallConfig) { ic.GCP.Region = "us-east4" } + invalidateProject = func(ic *types.InstallConfig) { ic.GCP.ProjectID = "invalid-project" } + removeVPC = func(ic *types.InstallConfig) { ic.GCP.Network = "" } + removeSubnets = func(ic *types.InstallConfig) { ic.GCP.ComputeSubnet, ic.GCP.ControlPlaneSubnet = "", "" } + + subnetAPIResult = []*compute.Subnetwork{ + { + Name: validCPSubnet, + IpCidrRange: validCIDR, + }, + { + Name: validComputeSubnet, + IpCidrRange: validCIDR, + }, + } +) + +func validInstallConfig() *types.InstallConfig { + return &types.InstallConfig{ + Networking: &types.Networking{ + MachineCIDR: ipnet.MustParseCIDR(validCIDR), + }, + Platform: types.Platform{ + GCP: &gcp.Platform{ + ProjectID: validProjectName, + Region: validRegion, + Network: validNetworkName, + ComputeSubnet: validComputeSubnet, + ControlPlaneSubnet: validCPSubnet, + }, + }, + } +} + +func TestGCPInstallConfigValidation(t *testing.T) { + cases := []struct { + name string + edits editFunctions + expectedError bool + expectedErrMsg string + }{ + { + name: "Valid network & subnets", + edits: editFunctions{}, + expectedError: false, + expectedErrMsg: "", + }, + { + name: "Valid install config without network & subnets", + edits: editFunctions{removeVPC, removeSubnets}, + expectedError: false, + expectedErrMsg: "", + }, + { + name: "Invalid subnet range", + edits: editFunctions{invalidateMachineCIDR}, + expectedError: true, + expectedErrMsg: "computeSubnet: Invalid value.*MachineCIDR", + }, + { + name: "Invalid network", + edits: editFunctions{invalidateNetwork}, + expectedError: true, + expectedErrMsg: "network: Invalid value", + }, + { + name: "Invalid compute subnet", + edits: editFunctions{invalidateComputeSubnet}, + expectedError: true, + expectedErrMsg: "computeSubnet: Invalid value", + }, + { + name: "Invalid control plane subnet", + edits: editFunctions{invalidateCPSubnet}, + expectedError: true, + expectedErrMsg: "controlPlaneSubnet: Invalid value", + }, + { + name: "Invalid both subnets", + edits: editFunctions{invalidateCPSubnet, invalidateComputeSubnet}, + expectedError: true, + expectedErrMsg: "computeSubnet: Invalid value.*controlPlaneSubnet: Invalid value", + }, + { + name: "Invalid region", + edits: editFunctions{invalidateRegion}, + expectedError: true, + expectedErrMsg: "could not find subnet valid-compute-subnet in network valid-vpc and region us-east4", + }, + { + name: "Invalid project", + edits: editFunctions{invalidateProject}, + expectedError: true, + expectedErrMsg: "network: Invalid value", + }, + { + name: "Invalid project & region", + edits: editFunctions{invalidateRegion, invalidateProject}, + expectedError: true, + expectedErrMsg: "network: Invalid value", + }, + } + mockCtrl := gomock.NewController(t) + defer mockCtrl.Finish() + + gcpClient := mock.NewMockAPI(mockCtrl) + // 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() + + // When passed an incorrect network or incorrect project, the API returns nil + gcpClient.EXPECT().GetNetwork(gomock.Any(), gomock.Not(validNetworkName), gomock.Any()).Return(nil, fmt.Errorf("404")).AnyTimes() + gcpClient.EXPECT().GetNetwork(gomock.Any(), gomock.Any(), gomock.Not(validProjectName)).Return(nil, fmt.Errorf("404")).AnyTimes() + + // When passed a correct network, project, & region, returns valid subnets. + // We will test incorrect subnets, by changing the install config. + gcpClient.EXPECT().GetSubnetworks(gomock.Any(), validNetworkName, validProjectName, validRegion).Return(subnetAPIResult, nil).AnyTimes() + + // When passed an incorrect network, project or region, return empty list. + gcpClient.EXPECT().GetSubnetworks(gomock.Any(), gomock.Not(validNetworkName), gomock.Any(), gomock.Any()).Return([]*compute.Subnetwork{}, nil).AnyTimes() + gcpClient.EXPECT().GetSubnetworks(gomock.Any(), gomock.Any(), gomock.Not(validProjectName), gomock.Any()).Return([]*compute.Subnetwork{}, nil).AnyTimes() + gcpClient.EXPECT().GetSubnetworks(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Not(validRegion)).Return([]*compute.Subnetwork{}, nil).AnyTimes() + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + editedInstallConfig := validInstallConfig() + for _, edit := range tc.edits { + edit(editedInstallConfig) + } + + errs := Validate(gcpClient, editedInstallConfig) + if tc.expectedError { + assert.Regexp(t, tc.expectedErrMsg, errs) + } else { + assert.Empty(t, errs) + } + }) + } +} diff --git a/pkg/asset/installconfig/installconfig.go b/pkg/asset/installconfig/installconfig.go index e5167602095..997df70fcac 100644 --- a/pkg/asset/installconfig/installconfig.go +++ b/pkg/asset/installconfig/installconfig.go @@ -1,14 +1,17 @@ package installconfig import ( + "context" "os" "github.com/ghodss/yaml" "github.com/pkg/errors" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/util/validation/field" "github.com/openshift/installer/pkg/asset" "github.com/openshift/installer/pkg/asset/installconfig/aws" + icgcp "github.com/openshift/installer/pkg/asset/installconfig/gcp" "github.com/openshift/installer/pkg/types" "github.com/openshift/installer/pkg/types/conversion" "github.com/openshift/installer/pkg/types/defaults" @@ -135,6 +138,10 @@ func (a *InstallConfig) finish(filename string) error { return errors.Wrapf(err, "invalid %q file", filename) } + if err := a.platformValidation(); err != nil { + return err + } + data, err := yaml.Marshal(a.Config) if err != nil { return errors.Wrap(err, "failed to Marshal InstallConfig") @@ -145,3 +152,14 @@ func (a *InstallConfig) finish(filename string) error { } return nil } + +func (a *InstallConfig) platformValidation() error { + if a.Config.Platform.GCP != nil { + client, err := icgcp.NewClient(context.TODO()) + if err != nil { + return err + } + return icgcp.Validate(client, a.Config) + } + return field.ErrorList{}.ToAggregate() +} diff --git a/pkg/types/gcp/validation/platform.go b/pkg/types/gcp/validation/platform.go index 1b7fc46119e..e50ce2e0dd0 100644 --- a/pkg/types/gcp/validation/platform.go +++ b/pkg/types/gcp/validation/platform.go @@ -55,5 +55,17 @@ func ValidatePlatform(p *gcp.Platform, fldPath *field.Path) field.ErrorList { if p.DefaultMachinePlatform != nil { allErrs = append(allErrs, ValidateMachinePool(p, p.DefaultMachinePlatform, fldPath.Child("defaultMachinePlatform"))...) } + if p.Network != "" { + if p.ComputeSubnet == "" { + allErrs = append(allErrs, field.Required(fldPath.Child("computeSubnet"), "must provide a compute subnet when a network is specified")) + } + if p.ControlPlaneSubnet == "" { + allErrs = append(allErrs, field.Required(fldPath.Child("controlPlaneSubnet"), "must provide a control plane subnet when a network is specified")) + } + } + if (p.ComputeSubnet != "" || p.ControlPlaneSubnet != "") && p.Network == "" { + allErrs = append(allErrs, field.Required(fldPath.Child("network"), "must provide a VPC network when supplying subnets")) + } + return allErrs } diff --git a/pkg/types/gcp/validation/platform_test.go b/pkg/types/gcp/validation/platform_test.go index 2048d854771..f5e681e3e34 100644 --- a/pkg/types/gcp/validation/platform_test.go +++ b/pkg/types/gcp/validation/platform_test.go @@ -37,6 +37,32 @@ func TestValidatePlatform(t *testing.T) { }, valid: true, }, + { + name: "valid subnets & network", + platform: &gcp.Platform{ + Region: "us-east1", + Network: "valid-vpc", + ComputeSubnet: "valid-compute-subnet", + ControlPlaneSubnet: "valid-cp-subnet", + }, + valid: true, + }, + { + name: "missing subnets", + platform: &gcp.Platform{ + Region: "us-east1", + Network: "valid-vpc", + }, + valid: false, + }, + { + name: "subnets missing network", + platform: &gcp.Platform{ + Region: "us-east1", + ComputeSubnet: "valid-compute-subnet", + }, + valid: false, + }, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) {