diff --git a/pkg/apis/compliance/v1alpha1/compliancescan_types.go b/pkg/apis/compliance/v1alpha1/compliancescan_types.go index e1d6531307..a0e65855e8 100644 --- a/pkg/apis/compliance/v1alpha1/compliancescan_types.go +++ b/pkg/apis/compliance/v1alpha1/compliancescan_types.go @@ -369,11 +369,11 @@ func (cs *ComplianceScan) NeedsTimeoutRescan() bool { // GetScanTypeIfValid returns scan type if the scan has a valid one, else it returns // an error func (cs *ComplianceScan) GetScanTypeIfValid() (ComplianceScanType, error) { - if strings.ToLower(string(cs.Spec.ScanType)) == strings.ToLower(string(ScanTypePlatform)) { + if strings.EqualFold(string(cs.Spec.ScanType), string(ScanTypePlatform)) { return ScanTypePlatform, nil } - if strings.ToLower(string(cs.Spec.ScanType)) == strings.ToLower(string(ScanTypeNode)) { + if strings.EqualFold(string(cs.Spec.ScanType), string(ScanTypeNode)) { return ScanTypeNode, nil } return "", ErrUnkownScanType @@ -382,36 +382,16 @@ func (cs *ComplianceScan) GetScanTypeIfValid() (ComplianceScanType, error) { // GetScannerTypeIfValid returns scaner type we will be using if the scan has a valid one, else it returns // an error func (cs *ComplianceScan) GetScannerTypeIfValid() (ScannerType, error) { - if strings.ToLower(string(cs.Spec.ScannerType)) == strings.ToLower(string(ScannerTypeOpenSCAP)) { + if strings.EqualFold(string(cs.Spec.ScannerType), string(ScannerTypeOpenSCAP)) { return ScannerTypeOpenSCAP, nil } - if strings.ToLower(string(cs.Spec.ScannerType)) == strings.ToLower(string(ScannerTypeCEL)) { + if strings.EqualFold(string(cs.Spec.ScannerType), string(ScannerTypeCEL)) { return ScannerTypeCEL, nil } return "", ErrUnkownScanerType } -// GetScanType get's the scan type for a scan -func (cs *ComplianceScan) GetScanType() ComplianceScanType { - scantype, err := cs.GetScanTypeIfValid() - if err != nil { - // This shouldn't happen - panic(err) - } - return scantype -} - -// GetScannerType will get the scanner type for a scan -func (cs *ComplianceScan) GetScannerType() ScannerType { - scannertype, err := cs.GetScannerTypeIfValid() - if err != nil { - // This shouldn't happen - panic(err) - } - return scannertype -} - // Returns whether remediation enforcement is off or not func (cs *ComplianceScan) RemediationEnforcementIsOff() bool { return (strings.EqualFold(cs.Spec.RemediationEnforcement, RemediationEnforcementEmpty) || diff --git a/pkg/apis/compliance/v1alpha1/compliancescan_types_test.go b/pkg/apis/compliance/v1alpha1/compliancescan_types_test.go new file mode 100644 index 0000000000..00f675620a --- /dev/null +++ b/pkg/apis/compliance/v1alpha1/compliancescan_types_test.go @@ -0,0 +1,150 @@ +package v1alpha1 + +import ( + . "github.com/onsi/ginkgo" + . "github.com/onsi/gomega" +) + +var _ = Describe("Testing ComplianceScan type validation", func() { + Context("GetScanTypeIfValid", func() { + It("should accept valid Platform scan type", func() { + scan := &ComplianceScan{ + Spec: ComplianceScanSpec{ + ScanType: ScanTypePlatform, + }, + } + scanType, err := scan.GetScanTypeIfValid() + Expect(err).To(BeNil()) + Expect(scanType).To(Equal(ScanTypePlatform)) + }) + + It("should accept valid Node scan type", func() { + scan := &ComplianceScan{ + Spec: ComplianceScanSpec{ + ScanType: ScanTypeNode, + }, + } + scanType, err := scan.GetScanTypeIfValid() + Expect(err).To(BeNil()) + Expect(scanType).To(Equal(ScanTypeNode)) + }) + + It("should accept Platform in lowercase", func() { + scan := &ComplianceScan{ + Spec: ComplianceScanSpec{ + ScanType: "platform", + }, + } + scanType, err := scan.GetScanTypeIfValid() + Expect(err).To(BeNil()) + Expect(scanType).To(Equal(ScanTypePlatform)) + }) + + It("should accept Node in uppercase", func() { + scan := &ComplianceScan{ + Spec: ComplianceScanSpec{ + ScanType: "NODE", + }, + } + scanType, err := scan.GetScanTypeIfValid() + Expect(err).To(BeNil()) + Expect(scanType).To(Equal(ScanTypeNode)) + }) + + It("should reject invalid scan type", func() { + scan := &ComplianceScan{ + Spec: ComplianceScanSpec{ + ScanType: "InvalidType", + }, + } + _, err := scan.GetScanTypeIfValid() + Expect(err).To(Equal(ErrUnkownScanType)) + }) + + It("should reject empty scan type", func() { + scan := &ComplianceScan{ + Spec: ComplianceScanSpec{ + ScanType: "", + }, + } + _, err := scan.GetScanTypeIfValid() + Expect(err).To(Equal(ErrUnkownScanType)) + }) + + It("should reject typo in Platform", func() { + scan := &ComplianceScan{ + Spec: ComplianceScanSpec{ + ScanType: "Plattform", + }, + } + _, err := scan.GetScanTypeIfValid() + Expect(err).To(Equal(ErrUnkownScanType)) + }) + }) + + Context("GetScannerTypeIfValid", func() { + It("should accept valid OpenSCAP scanner type", func() { + scan := &ComplianceScan{ + Spec: ComplianceScanSpec{ + ScannerType: ScannerTypeOpenSCAP, + }, + } + scannerType, err := scan.GetScannerTypeIfValid() + Expect(err).To(BeNil()) + Expect(scannerType).To(Equal(ScannerTypeOpenSCAP)) + }) + + It("should accept valid CEL scanner type", func() { + scan := &ComplianceScan{ + Spec: ComplianceScanSpec{ + ScannerType: ScannerTypeCEL, + }, + } + scannerType, err := scan.GetScannerTypeIfValid() + Expect(err).To(BeNil()) + Expect(scannerType).To(Equal(ScannerTypeCEL)) + }) + + It("should accept OpenSCAP in lowercase", func() { + scan := &ComplianceScan{ + Spec: ComplianceScanSpec{ + ScannerType: "openscap", + }, + } + scannerType, err := scan.GetScannerTypeIfValid() + Expect(err).To(BeNil()) + Expect(scannerType).To(Equal(ScannerTypeOpenSCAP)) + }) + + It("should accept CEL in mixed case", func() { + scan := &ComplianceScan{ + Spec: ComplianceScanSpec{ + ScannerType: "Cel", + }, + } + scannerType, err := scan.GetScannerTypeIfValid() + Expect(err).To(BeNil()) + Expect(scannerType).To(Equal(ScannerTypeCEL)) + }) + + It("should reject invalid scanner type", func() { + scan := &ComplianceScan{ + Spec: ComplianceScanSpec{ + ScannerType: "InvalidScanner", + }, + } + _, err := scan.GetScannerTypeIfValid() + Expect(err).To(Equal(ErrUnkownScanerType)) + }) + + It("should reject empty scanner type", func() { + scan := &ComplianceScan{ + Spec: ComplianceScanSpec{ + ScannerType: "", + }, + } + _, err := scan.GetScannerTypeIfValid() + Expect(err).To(Equal(ErrUnkownScanerType)) + }) + }) +}) diff --git a/pkg/controller/compliancescan/scantype.go b/pkg/controller/compliancescan/scantype.go index f04fbb1603..1890ad7cb7 100644 --- a/pkg/controller/compliancescan/scantype.go +++ b/pkg/controller/compliancescan/scantype.go @@ -82,7 +82,12 @@ func (nh *nodeScanTypeHandler) getScan() *compv1alpha1.ComplianceScan { func (nh *nodeScanTypeHandler) getTargetNodes() ([]corev1.Node, error) { var nodes corev1.NodeList - switch nh.scan.GetScanType() { + scanType, err := nh.scan.GetScanTypeIfValid() + if err != nil { + return nil, fmt.Errorf("invalid scan type: %w", err) + } + + switch scanType { case compv1alpha1.ScanTypePlatform: return nodes.Items, nil // Nodes are only relevant to the node scan type. Return the empty node list otherwise. case compv1alpha1.ScanTypeNode: