diff --git a/cmd/main.go b/cmd/main.go index b3690cc56..d269360ef 100644 --- a/cmd/main.go +++ b/cmd/main.go @@ -5,6 +5,7 @@ import ( "fmt" "log" "regexp" + "sort" "strings" "time" @@ -161,16 +162,9 @@ var toolCfg = Config{} // validateFlags is now defined in validators.go -// parseServices converts service names to ServiceType. The legacy -// "savingsplans" / "sp" aliases fan out to all four per-plan-type SP slugs so -// existing CLI scripts that pass --services savingsplans keep covering every -// plan type. Specific slugs (savingsplans-compute, etc.) are also accepted for -// targeted runs. -// -// Duplicates are silently dropped via a `seen` set so combinations like -// `--services savingsplans,savingsplans-compute` don't double-process Compute -// SP through both the fan-out path and the explicit-slug path. -func parseServices(serviceNames []string) []common.ServiceType { +// Legacy Savings Plans aliases fan out to every plan type for existing CLI scripts. +// Deduplication prevents overlapping aliases from processing a service twice. +func parseServices(serviceNames []string) ([]common.ServiceType, error) { var result []common.ServiceType seen := make(map[common.ServiceType]struct{}) add := func(service common.ServiceType) { @@ -215,11 +209,17 @@ func parseServices(serviceNames []string) []common.ServiceType { if service, ok := serviceMap[key]; ok { add(service) } else { - log.Printf("Warning: Unknown service '%s', skipping", name) + validNames := make([]string, 0, len(serviceMap)+3) + validNames = append(validNames, "savingsplans", "savings-plans", "sp") + for supported := range serviceMap { + validNames = append(validNames, supported) + } + sort.Strings(validNames) + return nil, fmt.Errorf("unknown service %q; valid services: %s", name, strings.Join(validNames, ", ")) } } - return result + return result, nil } // getAllServices returns all supported services. diff --git a/cmd/main_test.go b/cmd/main_test.go index 3473675da..bb92fb449 100644 --- a/cmd/main_test.go +++ b/cmd/main_test.go @@ -9,6 +9,7 @@ import ( "github.com/LeanerCloud/cloud-commitments-go/pkg/common" "github.com/aws/aws-sdk-go-v2/aws" + "github.com/spf13/cobra" "github.com/stretchr/testify/assert" ) @@ -31,9 +32,10 @@ func TestMain(m *testing.M) { func TestParseServices(t *testing.T) { tests := []struct { - name string - input []string - expected []common.ServiceType + name string + input []string + expected []common.ServiceType + errorContains string }{ { name: "Valid services", @@ -54,17 +56,16 @@ func TestParseServices(t *testing.T) { }, }, { - name: "Invalid services", - input: []string{"invalid", "unknown"}, - expected: nil, + name: "Invalid services", + errorContains: "invalid", + input: []string{"invalid", "unknown"}, + expected: nil, }, { - name: "Mix of valid and invalid", - input: []string{"rds", "invalid", "ec2"}, - expected: []common.ServiceType{ - common.ServiceRDS, - common.ServiceEC2, - }, + name: "Mix of valid and invalid", + input: []string{"rds", "invalid", "ec2"}, + errorContains: "invalid", + expected: nil, }, { name: "All supported services", @@ -89,7 +90,12 @@ func TestParseServices(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - result := parseServices(tt.input) + result, err := parseServices(tt.input) + if tt.errorContains != "" { + assert.ErrorContains(t, err, tt.errorContains) + } else { + assert.NoError(t, err) + } assert.Equal(t, tt.expected, result) }) } @@ -680,18 +686,30 @@ func TestEffectiveSizingPct(t *testing.T) { } func TestParseServicesWithEmptyAndNil(t *testing.T) { - // Empty slice - result := parseServices([]string{}) - assert.Empty(t, result) - - // Slice with empty strings - result = parseServices([]string{"", "rds", ""}) - assert.Len(t, result, 1) - assert.Equal(t, common.ServiceRDS, result[0]) - - // All invalid - result = parseServices([]string{"foo", "bar", "baz"}) - assert.Empty(t, result) + for _, input := range [][]string{nil, {}} { + result, err := parseServices(input) + assert.NoError(t, err) + assert.Nil(t, result) + } + for _, input := range [][]string{{"", "rds", ""}, {"rds", ""}, {"foo", "bar", "baz"}, {" rds"}} { + result, err := parseServices(input) + assert.Error(t, err) + assert.Nil(t, result) + } +} + +func TestParseServicesSavingsPlanAliases(t *testing.T) { + expected := []common.ServiceType{common.ServiceSavingsPlansCompute, common.ServiceSavingsPlansEC2Instance, common.ServiceSavingsPlansSageMaker, common.ServiceSavingsPlansDatabase} + for _, alias := range []string{"savingsplans", "savings-plans", "SP"} { + result, err := parseServices([]string{"ec2", alias, "savingsplans-compute", "rds", "EC2"}) + assert.NoError(t, err) + assert.Equal(t, append(append([]common.ServiceType{common.ServiceEC2}, expected...), common.ServiceRDS), result) + } + for i, name := range []string{"compute", "ec2instance", "sagemaker", "database"} { + result, err := parseServices([]string{"savingsplans-" + name, "savings-plans-" + name}) + assert.NoError(t, err) + assert.Equal(t, []common.ServiceType{expected[i]}, result) + } } func TestFilterFlagValidation(t *testing.T) { @@ -1256,3 +1274,43 @@ func TestValidateInstanceTypes(t *testing.T) { }) } } + +func TestServicesCommandRejectsUnknownBeforeRun(t *testing.T) { + original := toolCfg + t.Cleanup(func() { toolCfg = original }) + tests := []struct { + name string + args []string + errorContains string + }{ + {"mixed", []string{"--services=rds,elasticahe"}, "elasticahe"}, + {"unknown", []string{"--services=elasticahe"}, "elasticahe"}, + {"empty", []string{"--services="}, "--services must contain"}, + {"empty element", []string{"--services=rds,"}, "unknown service \"\""}, + {"all-services", []string{"--all-services", "--services=rds,elasticahe"}, "elasticahe"}, + {"CSV", []string{"--input-csv=missing.csv", "--services=rds,elasticahe"}, "elasticahe"}, + {"empty all-services", []string{"--all-services", "--services="}, "--services must contain"}, + {"empty CSV", []string{"--input-csv=missing.csv", "--services="}, "--services must contain"}, + {"default", nil, ""}, + {"valid", []string{"--services=RDS,SP,elasticsearch"}, ""}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + toolCfg = Config{Coverage: 80, CoverageLookbackDays: 30, PaymentOption: "partial-upfront", TermYears: 1, RecLookbackPeriod: "30d", IdempotencyWindow: "24h"} + ran := false + cmd := &cobra.Command{Use: "test", PreRunE: validateFlags, Run: func(*cobra.Command, []string) { ran = true }, SilenceUsage: true, SilenceErrors: true} + cmd.Flags().StringSliceVar(&toolCfg.Services, "services", []string{"rds"}, "") + cmd.Flags().BoolVar(&toolCfg.AllServices, "all-services", false, "") + cmd.Flags().StringVar(&toolCfg.CSVInput, "input-csv", "", "") + cmd.SetArgs(tt.args) + err := cmd.Execute() + if tt.errorContains != "" { + assert.ErrorContains(t, err, tt.errorContains) + assert.False(t, ran, "invalid services must stop before Run") + } else { + assert.NoError(t, err) + assert.True(t, ran) + } + }) + } +} diff --git a/cmd/multi_service.go b/cmd/multi_service.go index ed7887e2a..b00c0f328 100644 --- a/cmd/multi_service.go +++ b/cmd/multi_service.go @@ -121,7 +121,10 @@ func runToolMultiService(ctx context.Context, cfg Config) { return } - servicesToProcess := determineServicesToProcess(cfg) + servicesToProcess, serviceErr := determineServicesToProcess(cfg) + if serviceErr != nil { + log.Fatalf("Invalid services: %v", serviceErr) + } if len(servicesToProcess) == 0 { log.Fatalf("No valid services specified") } diff --git a/cmd/multi_service_helpers.go b/cmd/multi_service_helpers.go index 8a04076a3..bd8d8c32b 100644 --- a/cmd/multi_service_helpers.go +++ b/cmd/multi_service_helpers.go @@ -145,15 +145,15 @@ func applyCommonCoverage(recs []common.Recommendation, coverage float64) []commo } // determineServicesToProcess returns the list of services to process based on flags. -func determineServicesToProcess(cfg Config) []common.ServiceType { +func determineServicesToProcess(cfg Config) ([]common.ServiceType, error) { if cfg.AllServices { - return getAllServices() + return getAllServices(), nil } if len(cfg.Services) > 0 { return parseServices(cfg.Services) } // Default to RDS only for backward compatibility - return []common.ServiceType{common.ServiceRDS} + return []common.ServiceType{common.ServiceRDS}, nil } // printRunMode prints the current run mode (dry run or purchase). diff --git a/cmd/multi_service_helpers_test.go b/cmd/multi_service_helpers_test.go index af0cfdd9f..494e2bf5e 100644 --- a/cmd/multi_service_helpers_test.go +++ b/cmd/multi_service_helpers_test.go @@ -614,7 +614,8 @@ func TestDetermineServicesToProcess_AllServices(t *testing.T) { AllServices: true, } - result := determineServicesToProcess(cfg) + result, err := determineServicesToProcess(cfg) + assert.NoError(t, err) // Should contain all supported services assert.Contains(t, result, common.ServiceRDS) @@ -631,7 +632,8 @@ func TestDetermineServicesToProcess_SpecificServices(t *testing.T) { Services: []string{"rds", "elasticache"}, } - result := determineServicesToProcess(cfg) + result, err := determineServicesToProcess(cfg) + assert.NoError(t, err) assert.Equal(t, 2, len(result)) assert.Contains(t, result, common.ServiceRDS) @@ -844,3 +846,14 @@ func TestPopulateAccountNamesLogic(t *testing.T) { mockOrg.AssertExpectations(t) }) } + +func TestDetermineServicesToProcess_DefaultAndInvalid(t *testing.T) { + for _, services := range [][]string{nil, {}} { + result, err := determineServicesToProcess(Config{Services: services}) + assert.NoError(t, err) + assert.Equal(t, []common.ServiceType{common.ServiceRDS}, result) + } + result, err := determineServicesToProcess(Config{Services: []string{"rds", "elasticahe"}}) + assert.ErrorContains(t, err, "elasticahe") + assert.Nil(t, result) +} diff --git a/cmd/validators.go b/cmd/validators.go index 21bbbc58e..6bf6e58d1 100644 --- a/cmd/validators.go +++ b/cmd/validators.go @@ -14,6 +14,10 @@ import ( // validateFlags performs validation on command line flags before execution. func validateFlags(cmd *cobra.Command, args []string) error { + if err := validateServices(cmd); err != nil { + return err + } + if err := validateNumericRanges(cmd); err != nil { return err } @@ -45,6 +49,14 @@ func validateFlags(cmd *cobra.Command, args []string) error { return nil } +func validateServices(cmd *cobra.Command) error { + if cmd != nil && cmd.Flags().Changed("services") && len(toolCfg.Services) == 0 { + return fmt.Errorf("--services must contain at least one service") + } + _, err := parseServices(toolCfg.Services) + return err +} + // validateIdempotencyWindow parses --idempotency-window into the whole hours // the duplicate check's lookback is measured in. Anything it cannot represent // exactly is rejected rather than rounded, since a shorter window than asked @@ -180,7 +192,10 @@ func warnRDS3YearNoUpfront() error { return nil } - services := determineServicesToProcess(toolCfg) + services, err := determineServicesToProcess(toolCfg) + if err != nil { + return err + } hasRDS := toolCfg.AllServices || containsService(services, common.ServiceRDS) if hasRDS {