Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 12 additions & 12 deletions cmd/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import (
"fmt"
"log"
"regexp"
"sort"
"strings"
"time"

Expand Down Expand Up @@ -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) {
Expand Down Expand Up @@ -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.
Expand Down
108 changes: 83 additions & 25 deletions cmd/main_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)

Expand All @@ -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",
Expand All @@ -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",
Expand All @@ -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)
})
}
Expand Down Expand Up @@ -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) {
Expand Down Expand Up @@ -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)
}
})
}
}
5 changes: 4 additions & 1 deletion cmd/multi_service.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
}
Expand Down
6 changes: 3 additions & 3 deletions cmd/multi_service_helpers.go
Original file line number Diff line number Diff line change
Expand Up @@ -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).
Expand Down
17 changes: 15 additions & 2 deletions cmd/multi_service_helpers_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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)
Expand Down Expand Up @@ -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)
}
17 changes: 16 additions & 1 deletion cmd/validators.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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 {
Expand Down
Loading