diff --git a/Makefile b/Makefile index d16753d..24cab7e 100644 --- a/Makefile +++ b/Makefile @@ -1,7 +1,10 @@ -MOCKGEN := go run go.uber.org/mock/mockgen +# go.modのバージョンを使うと、missing go.sum entry for module providing package...エラーが出る +MOCKGEN := go run go.uber.org/mock/mockgen@v0.6.0 .PHONY: test test: go test ./... -coverprofile=coverage.txt -covermode=count +test/cli: + go test ./cli/... -coverprofile=coverage.txt -covermode=count test-container: docker build -t canarycage/test-container test-container push-test-container: test-container @@ -13,16 +16,23 @@ mocks: go.sum \ mocks/mock_awsiface/iface.go \ mocks/mock_types/iface.go \ mocks/mock_upgrade/upgrade.go \ + mocks/mock_audit/scanner.go \ + mocks/mock_audit/printer.go \ mocks/mock_task/task.go \ mocks/mock_taskset/taskset.go \ mocks/mock_task/factory.go \ - mocks/mock_rollout/executor.go + mocks/mock_rollout/executor.go \ + mocks/mock_logger/logger.go mocks/mock_awsiface/iface.go: awsiface/iface.go $(MOCKGEN) -source=./awsiface/iface.go > mocks/mock_awsiface/iface.go mocks/mock_types/iface.go: types/iface.go $(MOCKGEN) -source=./types/iface.go > mocks/mock_types/iface.go mocks/mock_upgrade/upgrade.go: cli/cage/upgrade/upgrade.go $(MOCKGEN) -source=./cli/cage/upgrade/upgrade.go > mocks/mock_upgrade/upgrade.go +mocks/mock_audit/scanner.go: cli/cage/audit/scanner.go + $(MOCKGEN) -source=./cli/cage/audit/scanner.go > mocks/mock_audit/scanner.go +mocks/mock_audit/printer.go: cli/cage/audit/printer.go + $(MOCKGEN) -source=./cli/cage/audit/printer.go > mocks/mock_audit/printer.go mocks/mock_task/task.go: task/task.go $(MOCKGEN) -source=./task/task.go > mocks/mock_task/task.go mocks/mock_taskset/taskset.go: taskset/taskset.go @@ -31,4 +41,6 @@ mocks/mock_task/factory.go: task/factory.go $(MOCKGEN) -source=./task/factory.go > mocks/mock_task/factory.go mocks/mock_rollout/executor.go: rollout/executor.go $(MOCKGEN) -source=./rollout/executor.go > mocks/mock_rollout/executor.go +mocks/mock_logger/logger.go: logger/logger.go + $(MOCKGEN) -source=./logger/logger.go > mocks/mock_logger/logger.go .PHONY: mocks diff --git a/awsiface/conf.go b/awsiface/conf.go new file mode 100644 index 0000000..ef8c931 --- /dev/null +++ b/awsiface/conf.go @@ -0,0 +1,17 @@ +package awsiface + +import ( + "context" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/config" +) + +// coverage cheat: always use MustLoadConfig to avoid error handling repetition +func MustLoadConfig(ctx context.Context, opts ...func(*config.LoadOptions) error) aws.Config { + cfg, err := config.LoadDefaultConfig(ctx, opts...) + if err != nil { + panic(err) + } + return cfg +} diff --git a/awsiface/conf_test.go b/awsiface/conf_test.go new file mode 100644 index 0000000..fd75907 --- /dev/null +++ b/awsiface/conf_test.go @@ -0,0 +1,59 @@ +package awsiface + +import ( + "context" + "errors" + "testing" + + "github.com/aws/aws-sdk-go-v2/config" +) + +func TestMustLoadConfig_Success(t *testing.T) { + ctx := context.Background() + + // This should not panic in normal circumstances + defer func() { + if r := recover(); r != nil { + t.Errorf("MustLoadConfig panicked unexpectedly: %v", r) + } + }() + + cfg := MustLoadConfig(ctx) + + if cfg.Region == "" && cfg.Credentials == nil { + t.Log("Config loaded (region or credentials may be empty in test environment)") + } +} + +func TestMustLoadConfig_WithOptions(t *testing.T) { + ctx := context.Background() + + defer func() { + if r := recover(); r != nil { + t.Errorf("MustLoadConfig with options panicked unexpectedly: %v", r) + } + }() + + cfg := MustLoadConfig(ctx, config.WithRegion("us-west-2")) + + if cfg.Region != "us-west-2" { + t.Errorf("Expected region us-west-2, got %s", cfg.Region) + } +} + +func TestMustLoadConfig_Panic(t *testing.T) { + ctx := context.Background() + + defer func() { + if r := recover(); r == nil { + t.Error("Expected MustLoadConfig to panic with invalid option, but it didn't") + } + }() + + // Pass an option that returns an error to trigger panic + invalidOpt := func(*config.LoadOptions) error { + return errors.New("forced error") + } + + MustLoadConfig(ctx, invalidOpt) +} diff --git a/awsiface/iface.go b/awsiface/iface.go index 94cd650..f106afb 100644 --- a/awsiface/iface.go +++ b/awsiface/iface.go @@ -4,6 +4,7 @@ import ( "context" "github.com/aws/aws-sdk-go-v2/service/ec2" + "github.com/aws/aws-sdk-go-v2/service/ecr" "github.com/aws/aws-sdk-go-v2/service/ecs" elbv2 "github.com/aws/aws-sdk-go-v2/service/elasticloadbalancingv2" ) @@ -23,6 +24,10 @@ type ( StopTask(ctx context.Context, params *ecs.StopTaskInput, optFns ...func(*ecs.Options)) (*ecs.StopTaskOutput, error) DescribeTaskDefinition(ctx context.Context, params *ecs.DescribeTaskDefinitionInput, optFns ...func(*ecs.Options)) (*ecs.DescribeTaskDefinitionOutput, error) } + EcrClient interface { + BatchGetImage(ctx context.Context, params *ecr.BatchGetImageInput, optFns ...func(*ecr.Options)) (*ecr.BatchGetImageOutput, error) + DescribeImageScanFindings(ctx context.Context, params *ecr.DescribeImageScanFindingsInput, optFns ...func(*ecr.Options)) (*ecr.DescribeImageScanFindingsOutput, error) + } AlbClient interface { DescribeTargetGroups(ctx context.Context, params *elbv2.DescribeTargetGroupsInput, optFns ...func(*elbv2.Options)) (*elbv2.DescribeTargetGroupsOutput, error) DescribeTargetHealth(ctx context.Context, params *elbv2.DescribeTargetHealthInput, optFns ...func(*elbv2.Options)) (*elbv2.DescribeTargetHealthOutput, error) diff --git a/cli/cage/audit/aggregator.go b/cli/cage/audit/aggregator.go new file mode 100644 index 0000000..f4b4f78 --- /dev/null +++ b/cli/cage/audit/aggregator.go @@ -0,0 +1,162 @@ +package audit + +import ( + "fmt" + + ecrtypes "github.com/aws/aws-sdk-go-v2/service/ecr/types" + "github.com/loilo-inc/canarycage/logger" +) + +type aggregater struct { + cves map[string]ecrtypes.ImageScanFinding + cveToSeverity map[string]string + cveToContainers map[string][]string + // container name to summaries + summaries map[string][]*ScanResultSummary +} + +func NewAggregater() *aggregater { + return &aggregater{ + cves: make(map[string]ecrtypes.ImageScanFinding), + cveToSeverity: make(map[string]string), + cveToContainers: make(map[string][]string), + summaries: make(map[string][]*ScanResultSummary)} +} + +func (a *aggregater) Add(r *ScanResult) { + container := r.ContainerName + if r.Err != nil { + a.summaries[container] = append(a.summaries[container], &ScanResultSummary{ + ContainerName: container, + Status: "ERROR", + }) + return + } else if r.ImageScanFindings == nil { + a.summaries[container] = append(a.summaries[container], &ScanResultSummary{ + ContainerName: container, + Status: "N/A", + }) + return + } + summary := summaryScanResult(r) + a.summaries[container] = append(a.summaries[container], summary) + for _, f := range r.ImageScanFindings.Findings { + if _, exists := a.cves[*f.Name]; !exists { + a.cves[*f.Name] = f + a.cveToSeverity[*f.Name] = string(f.Severity) + a.cveToContainers[*f.Name] = append(a.cveToContainers[*f.Name], container) + } + } +} + +type AggregateResult struct { + CriticalCount int32 + HighCount int32 + MediumCount int32 + LowCount int32 + InfoCount int32 + TotalCount int32 + HighestSeverity ecrtypes.FindingSeverity +} + +func (a *aggregater) SummarizeTotal() *AggregateResult { + result := &AggregateResult{} + highest := ecrtypes.FindingSeverityInformational + for cve := range a.cves { + severity := a.cveToSeverity[cve] + switch severity { + case string(ecrtypes.FindingSeverityCritical): + result.CriticalCount++ + case string(ecrtypes.FindingSeverityHigh): + result.HighCount++ + case string(ecrtypes.FindingSeverityMedium): + result.MediumCount++ + case string(ecrtypes.FindingSeverityLow): + result.LowCount++ + case string(ecrtypes.FindingSeverityInformational): + result.InfoCount++ + } + } + if result.CriticalCount > 0 { + highest = ecrtypes.FindingSeverityCritical + } else if result.HighCount > 0 { + highest = ecrtypes.FindingSeverityHigh + } else if result.MediumCount > 0 { + highest = ecrtypes.FindingSeverityMedium + } else if result.LowCount > 0 { + highest = ecrtypes.FindingSeverityLow + } else { + highest = ecrtypes.FindingSeverityInformational + } + result.HighestSeverity = highest + result.TotalCount = int32(len(a.cves)) + return result +} + +type SeverityCount struct { + Severity ecrtypes.FindingSeverity + Count int +} + +func (a *AggregateResult) SeverityCounts() []SeverityCount { + return []SeverityCount{ + {Severity: ecrtypes.FindingSeverityInformational, Count: int(a.InfoCount)}, + {Severity: ecrtypes.FindingSeverityLow, Count: int(a.LowCount)}, + {Severity: ecrtypes.FindingSeverityMedium, Count: int(a.MediumCount)}, + {Severity: ecrtypes.FindingSeverityHigh, Count: int(a.HighCount)}, + {Severity: ecrtypes.FindingSeverityCritical, Count: int(a.CriticalCount)}, + } +} + +func (a *aggregater) TotalCVECount() int { + return len(a.cves) +} + +func (a *aggregater) CriticalCves() []ecrtypes.ImageScanFinding { + return a.filterCvesBySeverity(ecrtypes.FindingSeverityCritical) +} + +func (a *aggregater) HighCves() []ecrtypes.ImageScanFinding { + return a.filterCvesBySeverity(ecrtypes.FindingSeverityHigh) +} + +func (a *aggregater) MediumCves() []ecrtypes.ImageScanFinding { + return a.filterCvesBySeverity(ecrtypes.FindingSeverityMedium) +} + +func (a *aggregater) filterCvesBySeverity(severity ecrtypes.FindingSeverity) []ecrtypes.ImageScanFinding { + var cves []ecrtypes.ImageScanFinding + for cve, sev := range a.cveToSeverity { + if sev == string(severity) { + cves = append(cves, a.cves[cve]) + } + } + return cves +} + +func (a *aggregater) GetVulnContainers(cveName string) []string { + containersSet := a.cveToContainers[cveName] + return containersSet +} + +type severityPrinter struct { + severity ecrtypes.FindingSeverity + color logger.Color +} + +func (s *severityPrinter) Sprintf(format string, a ...any) string { + switch s.severity { + case ecrtypes.FindingSeverityCritical: + return s.color.Magentaf(format, a...) + case ecrtypes.FindingSeverityHigh: + return s.color.Redf(format, a...) + case ecrtypes.FindingSeverityMedium: + return s.color.Yellowf(format, a...) + default: + return fmt.Sprintf(format, a...) + } +} + +func (s *severityPrinter) BSprintf(format string, a ...any) string { + return s.color.Bold(s.Sprintf(format, a...)) +} diff --git a/cli/cage/audit/aggregator_test.go b/cli/cage/audit/aggregator_test.go new file mode 100644 index 0000000..ee80e70 --- /dev/null +++ b/cli/cage/audit/aggregator_test.go @@ -0,0 +1,265 @@ +package audit + +import ( + "testing" + + ecrtypes "github.com/aws/aws-sdk-go-v2/service/ecr/types" + "github.com/stretchr/testify/assert" +) + +func TestSeverityPrinter_Sprintf(t *testing.T) { + tests := []struct { + name string + severity ecrtypes.FindingSeverity + format string + args []any + want string + }{ + { + name: "critical severity formats with magenta", + severity: ecrtypes.FindingSeverityCritical, + format: "test %s", + args: []any{"critical"}, + want: "\x1b[35mtest critical\x1b[0m", + }, + { + name: "high severity formats with red", + severity: ecrtypes.FindingSeverityHigh, + format: "test %s", + args: []any{"high"}, + want: "\x1b[31mtest high\x1b[0m", + }, + { + name: "medium severity formats with yellow", + severity: ecrtypes.FindingSeverityMedium, + format: "test %s", + args: []any{"medium"}, + want: "\x1b[33mtest medium\x1b[0m", + }, + { + name: "low severity formats without color", + severity: ecrtypes.FindingSeverityLow, + format: "test %s", + args: []any{"low"}, + want: "test low", + }, + { + name: "informational severity formats without color", + severity: ecrtypes.FindingSeverityInformational, + format: "test %s", + args: []any{"info"}, + want: "test info", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + s := &severityPrinter{ + severity: tt.severity, + } + got := s.Sprintf(tt.format, tt.args...) + assert.Equal(t, tt.want, got) + }) + } +} + +func TestSeverityPrinter_BSprintf(t *testing.T) { + s := &severityPrinter{ + severity: ecrtypes.FindingSeverityCritical, + } + got := s.BSprintf("test %s", "critical") + want := "\x1b[1m\x1b[35mtest critical\x1b[0m\x1b[0m" + assert.Equal(t, want, got) +} + +func TestNewAggregater(t *testing.T) { + agg := NewAggregater() + assert.NotNil(t, agg) + assert.NotNil(t, agg.cves) + assert.NotNil(t, agg.cveToSeverity) + assert.NotNil(t, agg.summaries) + assert.Equal(t, 0, len(agg.cves)) + assert.Equal(t, 0, len(agg.cveToSeverity)) + assert.Equal(t, 0, len(agg.summaries)) +} + +func TestAggregater_Add(t *testing.T) { + tests := []struct { + name string + scanResult *ScanResult + wantStatus string + wantCVECount int + }{ + { + name: "add result with error", + scanResult: &ScanResult{ + Err: assert.AnError, + }, + wantStatus: "ERROR", + wantCVECount: 0, + }, + { + name: "add result with nil findings", + scanResult: &ScanResult{ + ImageScanFindings: nil, + }, + wantStatus: "N/A", + wantCVECount: 0, + }, + { + name: "add result with findings", + scanResult: &ScanResult{ + ImageInfo: ImageInfo{ + ContainerName: "test-container", + }, + ImageScanFindings: &ecrtypes.ImageScanFindings{ + Findings: []ecrtypes.ImageScanFinding{ + { + Name: stringPtr("CVE-2021-1234"), + Severity: ecrtypes.FindingSeverityCritical, + }, + { + Name: stringPtr("CVE-2021-5678"), + Severity: ecrtypes.FindingSeverityHigh, + }, + }, + }, + }, + wantStatus: "VULNERABLE", + wantCVECount: 2, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + agg := NewAggregater() + agg.Add(tt.scanResult) + assert.Equal(t, 1, len(agg.summaries)) + assert.Equal(t, tt.wantStatus, agg.summaries[tt.scanResult.ImageInfo.ContainerName][0].Status) + assert.Equal(t, tt.wantCVECount, len(agg.cves)) + }) + + } +} + +func TestAggregater_SummarizeTotal(t *testing.T) { + t.Run("should summarize total counts and highest severity", func(t *testing.T) { + agg := NewAggregater() + agg.cves = map[string]ecrtypes.ImageScanFinding{ + "CVE-2021-1": {Name: stringPtr("CVE-2021-1"), Severity: ecrtypes.FindingSeverityCritical}, + "CVE-2021-2": {Name: stringPtr("CVE-2021-2"), Severity: ecrtypes.FindingSeverityHigh}, + "CVE-2021-3": {Name: stringPtr("CVE-2021-3"), Severity: ecrtypes.FindingSeverityMedium}, + "CVE-2021-4": {Name: stringPtr("CVE-2021-4"), Severity: ecrtypes.FindingSeverityLow}, + "CVE-2021-5": {Name: stringPtr("CVE-2021-5"), Severity: ecrtypes.FindingSeverityInformational}, + } + agg.cveToSeverity = map[string]string{ + "CVE-2021-1": string(ecrtypes.FindingSeverityCritical), + "CVE-2021-2": string(ecrtypes.FindingSeverityHigh), + "CVE-2021-3": string(ecrtypes.FindingSeverityMedium), + "CVE-2021-4": string(ecrtypes.FindingSeverityLow), + "CVE-2021-5": string(ecrtypes.FindingSeverityInformational), + } + + result := agg.SummarizeTotal() + assert.Equal(t, int32(1), result.CriticalCount) + assert.Equal(t, int32(1), result.HighCount) + assert.Equal(t, int32(1), result.MediumCount) + assert.Equal(t, int32(1), result.LowCount) + assert.Equal(t, int32(1), result.InfoCount) + assert.Equal(t, int32(5), result.TotalCount) + assert.Equal(t, ecrtypes.FindingSeverityCritical, result.HighestSeverity) + }) + t.Run("highest", func(t *testing.T) { + tests := []struct { + severity ecrtypes.FindingSeverity + }{ + {severity: ecrtypes.FindingSeverityHigh}, + {severity: ecrtypes.FindingSeverityMedium}, + {severity: ecrtypes.FindingSeverityLow}, + {severity: ecrtypes.FindingSeverityInformational}, + } + for _, tt := range tests { + t.Run(string(tt.severity), func(t *testing.T) { + agg := NewAggregater() + agg.cves = map[string]ecrtypes.ImageScanFinding{ + "CVE-2021-1": {Name: stringPtr("CVE-2021-1"), Severity: tt.severity}, + } + agg.cveToSeverity = map[string]string{ + "CVE-2021-1": string(tt.severity), + } + result := agg.SummarizeTotal() + assert.Equal(t, tt.severity, result.HighestSeverity) + }) + } + }) +} + +func TestAggregater_FilterCvesBySeverity(t *testing.T) { + agg := NewAggregater() + agg.cves = map[string]ecrtypes.ImageScanFinding{ + "CVE-2021-1": {Name: stringPtr("CVE-2021-1"), Severity: ecrtypes.FindingSeverityCritical}, + "CVE-2021-2": {Name: stringPtr("CVE-2021-2"), Severity: ecrtypes.FindingSeverityHigh}, + "CVE-2021-3": {Name: stringPtr("CVE-2021-3"), Severity: ecrtypes.FindingSeverityCritical}, + } + agg.cveToSeverity = map[string]string{ + "CVE-2021-1": string(ecrtypes.FindingSeverityCritical), + "CVE-2021-2": string(ecrtypes.FindingSeverityHigh), + "CVE-2021-3": string(ecrtypes.FindingSeverityCritical), + } + + critical := agg.CriticalCves() + assert.Equal(t, 2, len(critical)) + + high := agg.HighCves() + assert.Equal(t, 1, len(high)) + + medium := agg.MediumCves() + assert.Equal(t, 0, len(medium)) +} + +func TestAggregateResult_SeverityCounts(t *testing.T) { + result := &AggregateResult{ + CriticalCount: 1, + HighCount: 2, + MediumCount: 3, + LowCount: 4, + InfoCount: 5, + } + + counts := result.SeverityCounts() + assert.Equal(t, 5, len(counts)) + assert.Equal(t, 5, counts[0].Count) + assert.Equal(t, 4, counts[1].Count) + assert.Equal(t, 3, counts[2].Count) + assert.Equal(t, 2, counts[3].Count) + assert.Equal(t, 1, counts[4].Count) +} + +func TestAggregater_GetVulnContainers(t *testing.T) { + t.Run("returns containers affected by CVE", func(t *testing.T) { + agg := NewAggregater() + agg.cveToContainers = map[string][]string{ + "CVE-2021-1234": {"container1", "container2"}, + "CVE-2021-5678": {"container3"}, + } + + containers := agg.GetVulnContainers("CVE-2021-1234") + assert.Equal(t, 2, len(containers)) + assert.Contains(t, containers, "container1") + assert.Contains(t, containers, "container2") + }) + + t.Run("returns nil for non-existent CVE", func(t *testing.T) { + agg := NewAggregater() + agg.cveToContainers = map[string][]string{ + "CVE-2021-1234": {"container1"}, + } + + containers := agg.GetVulnContainers("CVE-9999-9999") + assert.Nil(t, containers) + }) +} + +func stringPtr(s string) *string { + return &s +} diff --git a/cli/cage/audit/command.go b/cli/cage/audit/command.go new file mode 100644 index 0000000..88eb445 --- /dev/null +++ b/cli/cage/audit/command.go @@ -0,0 +1,61 @@ +package audit + +import ( + "context" + "time" + + "github.com/loilo-inc/canarycage/cli/cage/cageapp" + "github.com/loilo-inc/canarycage/key" + "github.com/loilo-inc/canarycage/logger" + "github.com/loilo-inc/canarycage/types" + "github.com/loilo-inc/logos/di" +) + +type command struct { + di *di.D + input *cageapp.AuditCmdInput + spinInterval time.Duration +} + +func NewCommand(di *di.D, input *cageapp.AuditCmdInput) *command { + return &command{ + di: di, + input: input, + spinInterval: 100 * time.Millisecond, + } +} + +func (a *command) Run(ctx context.Context) error { + results, err := a.doScan(ctx) + if err != nil { + return err + } + p := a.di.Get(key.Printer).(Printer) + p.Print(results) + return nil +} + +func (a *command) doScan(ctx context.Context) (results []*ScanResult, err error) { + l := a.di.Get(key.Logger).(logger.Logger) + t := a.di.Get(key.Time).(types.Time) + defer l.Printf("\r") + waiter := make(chan struct{}, 1) + spinner := logger.NewSpinner() + go func() { + defer close(waiter) + scanner := a.di.Get(key.Scanner).(Scanner) + results, err = scanner.Scan(ctx, a.input.Cluster, a.input.Service) + waiter <- struct{}{} + }() + for { + timer := t.NewTimer(a.spinInterval) + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-waiter: + return + case <-timer.C: + l.Printf("\r%s", spinner.Next()) + } + } +} diff --git a/cli/cage/audit/command_test.go b/cli/cage/audit/command_test.go new file mode 100644 index 0000000..36d6bf4 --- /dev/null +++ b/cli/cage/audit/command_test.go @@ -0,0 +1,100 @@ +package audit_test + +import ( + "context" + "testing" + + "github.com/loilo-inc/canarycage/cli/cage/audit" + "github.com/loilo-inc/canarycage/cli/cage/cageapp" + "github.com/loilo-inc/canarycage/key" + "github.com/loilo-inc/canarycage/mocks/mock_audit" + "github.com/loilo-inc/canarycage/mocks/mock_logger" + "github.com/loilo-inc/canarycage/test" + "github.com/loilo-inc/logos/di" + "github.com/stretchr/testify/assert" + "go.uber.org/mock/gomock" +) + +func TestAuditCommandRun(t *testing.T) { + setup := func(t *testing.T) (*mock_audit.MockScanner, *mock_logger.MockLogger, *mock_audit.MockPrinter) { + t.Helper() + ctrl := gomock.NewController(t) + mockScanner := mock_audit.NewMockScanner(ctrl) + mockPrinter := mock_audit.NewMockPrinter(ctrl) + mockLogger := mock_logger.NewMockLogger(ctrl) + return mockScanner, mockLogger, mockPrinter + } + t.Run("should return error from scanner", func(t *testing.T) { + ctx := context.Background() + mockScanner, mockLogger, _ := setup(t) + mockDI := di.NewDomain(func(b *di.B) { + b.Set(key.Logger, mockLogger) + b.Set(key.Scanner, mockScanner) + b.Set(key.Time, test.NewFakeNeverTimer()) + }) + gomock.InOrder( + mockScanner.EXPECT().Scan(ctx, "cluster", "service").Return(nil, test.Err), + mockLogger.EXPECT().Printf("\r"), + ) + input := cageapp.NewAuditCmdInput() + input.Cluster = "cluster" + input.Service = "service" + cmd := audit.NewCommand(mockDI, input) + err := cmd.Run(ctx) + assert.Equal(t, test.Err, err) + }) + + t.Run("should return nil on successful scan", func(t *testing.T) { + ctx := context.Background() + + mockScanner, mockLogger, mockPrinter := setup(t) + mockDI := di.NewDomain(func(b *di.B) { + b.Set(key.Logger, mockLogger) + b.Set(key.Scanner, mockScanner) + b.Set(key.Printer, mockPrinter) + b.Set(key.Time, test.NewFakeNeverTimer()) + }) + + var results []*audit.ScanResult + gomock.InOrder( + mockScanner.EXPECT().Scan(ctx, "cluster", "service").Return(results, nil), + mockLogger.EXPECT().Printf("\r"), + mockPrinter.EXPECT().Print(results), + ) + + cmd := audit.NewCommand(mockDI, &cageapp.AuditCmdInput{ + Cluster: "cluster", + Service: "service", + App: &cageapp.App{}, + }) + + err := cmd.Run(ctx) + assert.NoError(t, err) + }) + + t.Run("should return context error when context is cancelled", func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + mockScanner, mockLogger, _ := setup(t) + mockDI := di.NewDomain(func(b *di.B) { + b.Set(key.Logger, mockLogger) + b.Set(key.Scanner, mockScanner) + b.Set(key.Time, test.NewFakeNeverTimer()) + }) + gomock.InOrder( + mockScanner.EXPECT().Scan(ctx, "cluster", "service").DoAndReturn(func(context.Context, string, string) ([]audit.ScanResult, error) { + cancel() + return nil, nil + }), + mockLogger.EXPECT().Printf("\r"), + ) + + cmd := audit.NewCommand(mockDI, &cageapp.AuditCmdInput{ + Cluster: "cluster", + Service: "service", + App: &cageapp.App{}, + }) + + err := cmd.Run(ctx) + assert.Equal(t, context.Canceled, err) + }) +} diff --git a/cli/cage/audit/deps.go b/cli/cage/audit/deps.go new file mode 100644 index 0000000..c1f65b5 --- /dev/null +++ b/cli/cage/audit/deps.go @@ -0,0 +1,35 @@ +package audit + +import ( + "context" + "os" + + "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/service/ecr" + "github.com/aws/aws-sdk-go-v2/service/ecs" + "github.com/loilo-inc/canarycage/awsiface" + "github.com/loilo-inc/canarycage/cli/cage/cageapp" + "github.com/loilo-inc/canarycage/key" + "github.com/loilo-inc/canarycage/logger" + "github.com/loilo-inc/canarycage/timeout" + "github.com/loilo-inc/canarycage/types" + "github.com/loilo-inc/logos/di" +) + +func ProvideAuditCmd(ctx context.Context, input *cageapp.AuditCmdInput) (types.Audit, error) { + conf := awsiface.MustLoadConfig( + ctx, + config.WithRegion(input.Region), + ) + d := di.NewDomain(func(b *di.B) { + ecsCli := ecs.NewFromConfig(conf) + ecrCli := ecr.NewFromConfig(conf) + l := logger.DefaultLogger(os.Stdout) + p := NewPrinter(l, input.NoColor, input.LogDetail) + b.Set(key.Scanner, NewScanner(ecsCli, ecrCli)) + b.Set(key.Logger, l) + b.Set(key.Printer, p) + b.Set(key.Time, &timeout.Time{}) + }) + return NewCommand(d, input), nil +} diff --git a/cli/cage/audit/deps_test.go b/cli/cage/audit/deps_test.go new file mode 100644 index 0000000..ecc5c1a --- /dev/null +++ b/cli/cage/audit/deps_test.go @@ -0,0 +1,43 @@ +package audit + +import ( + "context" + "testing" + + "github.com/loilo-inc/canarycage/cli/cage/cageapp" +) + +func TestProvideAuditCmd(t *testing.T) { + ctx := context.Background() + input := cageapp.NewAuditCmdInput() + input.Region = "us-east-1" + audit, err := ProvideAuditCmd(ctx, input) + if err != nil { + t.Fatalf("ProvideAuditCmd() error = %v, want nil", err) + } + + if audit == nil { + t.Fatal("ProvideAuditCmd() returned nil audit") + } +} + +func TestProvideAuditCmd_WithDifferentRegions(t *testing.T) { + regions := []string{"us-east-1", "eu-west-1", "ap-northeast-1"} + + for _, region := range regions { + t.Run(region, func(t *testing.T) { + ctx := context.Background() + input := cageapp.NewAuditCmdInput() + input.Region = region + + audit, err := ProvideAuditCmd(ctx, input) + if err != nil { + t.Fatalf("ProvideAuditCmd() error = %v, want nil", err) + } + + if audit == nil { + t.Fatalf("ProvideAuditCmd() returned nil audit for region %s", region) + } + }) + } +} diff --git a/cli/cage/audit/ecr.go b/cli/cage/audit/ecr.go new file mode 100644 index 0000000..00e16ed --- /dev/null +++ b/cli/cage/audit/ecr.go @@ -0,0 +1,107 @@ +package audit + +import ( + "context" + "encoding/json" + "fmt" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/ecr" + ecrtypes "github.com/aws/aws-sdk-go-v2/service/ecr/types" + ecstypes "github.com/aws/aws-sdk-go-v2/service/ecs/types" + "github.com/loilo-inc/canarycage/awsiface" +) + +const dockerManifestListMediaType = "application/vnd.docker.distribution.manifest.list.v2+json" + +type ecrTool struct { + Ecr awsiface.EcrClient +} + +type EcrTool interface { + GetActualImageIdentifier(ctx context.Context, info *ImageInfo) (*ecrtypes.ImageIdentifier, error) + GetImageScanFindings(ctx context.Context, info *ImageInfo, imageID *ecrtypes.ImageIdentifier) (*ecrtypes.ImageScanFindings, error) +} + +func newEcrTool(ecrClient awsiface.EcrClient) EcrTool { + return &ecrTool{Ecr: ecrClient} +} + +func (t *ecrTool) GetActualImageIdentifier(ctx context.Context, info *ImageInfo) (*ecrtypes.ImageIdentifier, error) { + res, err := t.Ecr.BatchGetImage(ctx, &ecr.BatchGetImageInput{ + RepositoryName: aws.String(info.Repository), + ImageIds: []ecrtypes.ImageIdentifier{{ImageTag: aws.String(info.Tag)}}, + }) + if err != nil { + return nil, err + } + if len(res.Images) == 0 || res.Images[0].ImageManifest == nil { + return nil, fmt.Errorf("image manifest not found for %s:%s", info.Repository, info.Tag) + } + + var manifest dockerSchema + if err := json.Unmarshal([]byte(*res.Images[0].ImageManifest), &manifest); err != nil { + return nil, fmt.Errorf("parse image manifest for %s:%s: %w", info.Repository, info.Tag, err) + } + + if manifest.MediaType == dockerManifestListMediaType { + for _, candidate := range manifest.Manifests { + if candidate.Platform == nil { + continue + } + if toCPUArchitecture(candidate.Platform.Architecture) == info.PlatformArch { + return &ecrtypes.ImageIdentifier{ImageDigest: aws.String(candidate.Digest)}, nil + } + } + return nil, fmt.Errorf("no image found for architecture: %s in %s:%s", info.PlatformArch, info.Repository, info.Tag) + } + + return &ecrtypes.ImageIdentifier{ImageTag: aws.String(info.Tag)}, nil +} + +func (t *ecrTool) GetImageScanFindings(ctx context.Context, info *ImageInfo, imageID *ecrtypes.ImageIdentifier) (*ecrtypes.ImageScanFindings, error) { + res, err := t.Ecr.DescribeImageScanFindings(ctx, &ecr.DescribeImageScanFindingsInput{ + RepositoryName: aws.String(info.Repository), + ImageId: imageID, + }) + if err != nil { + return nil, err + } + if res.ImageScanFindings == nil { + return nil, fmt.Errorf("image scan findings missing for %s:%s", info.Repository, info.Tag) + } + return res.ImageScanFindings, nil +} + +var _ awsiface.EcrClient = (*ecr.Client)(nil) + +type dockerSchema struct { + SchemaVersion int `json:"schemaVersion"` + MediaType string `json:"mediaType"` + Manifests []dockerManifest `json:"manifests,omitempty"` + Config *dockerManifest `json:"config,omitempty"` + Layers []dockerManifest `json:"layers,omitempty"` +} + +type dockerManifest struct { + MediaType string `json:"mediaType"` + Size int64 `json:"size"` + Digest string `json:"digest"` + Platform *dockerPlatform `json:"platform,omitempty"` +} + +type dockerPlatform struct { + Architecture string `json:"architecture"` + OS string `json:"os"` +} + +func toCPUArchitecture(arch string) ecstypes.CPUArchitecture { + switch arch { + case "amd64": + return ecstypes.CPUArchitectureX8664 + case "arm64": + return ecstypes.CPUArchitectureArm64 + default: + return "" + } +} diff --git a/cli/cage/audit/ecr_test.go b/cli/cage/audit/ecr_test.go new file mode 100644 index 0000000..4fe12d9 --- /dev/null +++ b/cli/cage/audit/ecr_test.go @@ -0,0 +1,373 @@ +package audit + +import ( + "context" + "encoding/json" + "errors" + "testing" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/ecr" + ecrtypes "github.com/aws/aws-sdk-go-v2/service/ecr/types" + ecstypes "github.com/aws/aws-sdk-go-v2/service/ecs/types" + "github.com/loilo-inc/canarycage/mocks/mock_awsiface" + "github.com/stretchr/testify/assert" + "go.uber.org/mock/gomock" +) + +func TestGetActualImageIdentifier(t *testing.T) { + ctx := context.Background() + + t.Run("single architecture image returns tag", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcrClient(ctrl) + tool := newEcrTool(mockClient) + + info := &ImageInfo{ + Repository: "my-repo", + Tag: "v1.0.0", + PlatformArch: ecstypes.CPUArchitectureX8664, + } + + manifest := dockerSchema{ + SchemaVersion: 2, + MediaType: "application/vnd.docker.distribution.manifest.v2+json", + } + manifestJSON, _ := json.Marshal(manifest) + manifestStr := string(manifestJSON) + + mockClient.EXPECT().BatchGetImage(ctx, gomock.AssignableToTypeOf(&ecr.BatchGetImageInput{})).Return(&ecr.BatchGetImageOutput{ + Images: []ecrtypes.Image{ + {ImageManifest: &manifestStr}, + }, + }, nil) + + result, err := tool.GetActualImageIdentifier(ctx, info) + + assert.NoError(t, err) + assert.NotNil(t, result) + assert.Equal(t, "v1.0.0", *result.ImageTag) + }) + + t.Run("multi-arch image returns digest for matching architecture", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcrClient(ctrl) + tool := newEcrTool(mockClient) + + info := &ImageInfo{ + Repository: "my-repo", + Tag: "v1.0.0", + PlatformArch: ecstypes.CPUArchitectureArm64, + } + + manifest := dockerSchema{ + SchemaVersion: 2, + MediaType: dockerManifestListMediaType, + Manifests: []dockerManifest{ + { + Digest: "sha256:amd64digest", + Platform: &dockerPlatform{ + Architecture: "amd64", + OS: "linux", + }, + }, + { + Digest: "sha256:arm64digest", + Platform: &dockerPlatform{ + Architecture: "arm64", + OS: "linux", + }, + }, + }, + } + manifestJSON, _ := json.Marshal(manifest) + manifestStr := string(manifestJSON) + + mockClient.EXPECT().BatchGetImage(ctx, gomock.AssignableToTypeOf(&ecr.BatchGetImageInput{})).Return(&ecr.BatchGetImageOutput{ + Images: []ecrtypes.Image{ + {ImageManifest: &manifestStr}, + }, + }, nil) + + result, err := tool.GetActualImageIdentifier(ctx, info) + + assert.NoError(t, err) + assert.NotNil(t, result) + assert.Equal(t, "sha256:arm64digest", *result.ImageDigest) + }) + + t.Run("multi-arch image with no matching architecture returns error", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcrClient(ctrl) + tool := newEcrTool(mockClient) + + info := &ImageInfo{ + Repository: "my-repo", + Tag: "v1.0.0", + PlatformArch: ecstypes.CPUArchitectureArm64, + } + + manifest := dockerSchema{ + SchemaVersion: 2, + MediaType: dockerManifestListMediaType, + Manifests: []dockerManifest{ + { + Digest: "sha256:amd64digest", + Platform: &dockerPlatform{ + Architecture: "amd64", + OS: "linux", + }, + }, + }, + } + manifestJSON, _ := json.Marshal(manifest) + manifestStr := string(manifestJSON) + + mockClient.EXPECT().BatchGetImage(ctx, gomock.AssignableToTypeOf(&ecr.BatchGetImageInput{})).Return(&ecr.BatchGetImageOutput{ + Images: []ecrtypes.Image{ + {ImageManifest: &manifestStr}, + }, + }, nil) + + result, err := tool.GetActualImageIdentifier(ctx, info) + + assert.Error(t, err) + assert.Nil(t, result) + assert.Contains(t, err.Error(), "no image found for architecture") + }) + + t.Run("BatchGetImage error returns error", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcrClient(ctrl) + tool := newEcrTool(mockClient) + + info := &ImageInfo{ + Repository: "my-repo", + Tag: "v1.0.0", + PlatformArch: ecstypes.CPUArchitectureX8664, + } + + mockClient.EXPECT().BatchGetImage(ctx, gomock.AssignableToTypeOf(&ecr.BatchGetImageInput{})).Return(nil, errors.New("API error")) + + result, err := tool.GetActualImageIdentifier(ctx, info) + + assert.Error(t, err) + assert.Nil(t, result) + assert.Contains(t, err.Error(), "API error") + }) + + t.Run("empty images returns error", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcrClient(ctrl) + tool := newEcrTool(mockClient) + + info := &ImageInfo{ + Repository: "my-repo", + Tag: "v1.0.0", + PlatformArch: ecstypes.CPUArchitectureX8664, + } + + mockClient.EXPECT().BatchGetImage(ctx, gomock.AssignableToTypeOf(&ecr.BatchGetImageInput{})).Return(&ecr.BatchGetImageOutput{ + Images: []ecrtypes.Image{}, + }, nil) + + result, err := tool.GetActualImageIdentifier(ctx, info) + + assert.Error(t, err) + assert.Nil(t, result) + assert.Contains(t, err.Error(), "image manifest not found") + }) + + t.Run("nil manifest returns error", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcrClient(ctrl) + tool := newEcrTool(mockClient) + + info := &ImageInfo{ + Repository: "my-repo", + Tag: "v1.0.0", + PlatformArch: ecstypes.CPUArchitectureX8664, + } + + mockClient.EXPECT().BatchGetImage(ctx, gomock.AssignableToTypeOf(&ecr.BatchGetImageInput{})).Return(&ecr.BatchGetImageOutput{ + Images: []ecrtypes.Image{ + {ImageManifest: nil}, + }, + }, nil) + + result, err := tool.GetActualImageIdentifier(ctx, info) + + assert.Error(t, err) + assert.Nil(t, result) + assert.Contains(t, err.Error(), "image manifest not found") + }) + + t.Run("invalid JSON manifest returns error", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcrClient(ctrl) + tool := newEcrTool(mockClient) + + info := &ImageInfo{ + Repository: "my-repo", + Tag: "v1.0.0", + PlatformArch: ecstypes.CPUArchitectureX8664, + } + + invalidJSON := "invalid json" + mockClient.EXPECT().BatchGetImage(ctx, gomock.AssignableToTypeOf(&ecr.BatchGetImageInput{})).Return(&ecr.BatchGetImageOutput{ + Images: []ecrtypes.Image{ + {ImageManifest: &invalidJSON}, + }, + }, nil) + + result, err := tool.GetActualImageIdentifier(ctx, info) + + assert.Error(t, err) + assert.Nil(t, result) + assert.Contains(t, err.Error(), "parse image manifest") + }) + + t.Run("multi-arch image skips manifests without platform", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcrClient(ctrl) + tool := newEcrTool(mockClient) + + info := &ImageInfo{ + Repository: "my-repo", + Tag: "v1.0.0", + PlatformArch: ecstypes.CPUArchitectureArm64, + } + + manifest := dockerSchema{ + SchemaVersion: 2, + MediaType: dockerManifestListMediaType, + Manifests: []dockerManifest{ + { + Digest: "sha256:noplatform", + Platform: nil, + }, + { + Digest: "sha256:arm64digest", + Platform: &dockerPlatform{ + Architecture: "arm64", + OS: "linux", + }, + }, + }, + } + manifestJSON, _ := json.Marshal(manifest) + manifestStr := string(manifestJSON) + + mockClient.EXPECT().BatchGetImage(ctx, gomock.AssignableToTypeOf(&ecr.BatchGetImageInput{})).Return(&ecr.BatchGetImageOutput{ + Images: []ecrtypes.Image{ + {ImageManifest: &manifestStr}, + }, + }, nil) + + result, err := tool.GetActualImageIdentifier(ctx, info) + + assert.NoError(t, err) + assert.NotNil(t, result) + assert.Equal(t, "sha256:arm64digest", *result.ImageDigest) + }) +} +func TestGetImageScanFindings(t *testing.T) { + ctx := context.Background() + + t.Run("successfully returns scan findings", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcrClient(ctrl) + tool := newEcrTool(mockClient) + + info := &ImageInfo{ + Repository: "my-repo", + Tag: "v1.0.0", + } + imageID := &ecrtypes.ImageIdentifier{ + ImageDigest: aws.String("sha256:abc123"), + } + + expectedFindings := &ecrtypes.ImageScanFindings{ + FindingSeverityCounts: map[string]int32{ + "CRITICAL": 1, + "HIGH": 2, + }, + } + + mockClient.EXPECT().DescribeImageScanFindings( + ctx, gomock.AssignableToTypeOf(&ecr.DescribeImageScanFindingsInput{})). + DoAndReturn(func(ctx context.Context, + input *ecr.DescribeImageScanFindingsInput, + opts ...func(*ecr.Options)) (*ecr.DescribeImageScanFindingsOutput, error) { + assert.Equal(t, "my-repo", *input.RepositoryName) + assert.Equal(t, "sha256:abc123", *input.ImageId.ImageDigest) + return &ecr.DescribeImageScanFindingsOutput{ImageScanFindings: expectedFindings}, nil + }) + + result, err := tool.GetImageScanFindings(ctx, info, imageID) + + assert.NoError(t, err) + assert.NotNil(t, result) + assert.Equal(t, expectedFindings, result) + }) + + t.Run("DescribeImageScanFindings error returns error", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcrClient(ctrl) + tool := newEcrTool(mockClient) + + info := &ImageInfo{ + Repository: "my-repo", + Tag: "v1.0.0", + } + imageID := &ecrtypes.ImageIdentifier{ + ImageTag: aws.String("v1.0.0"), + } + + mockClient.EXPECT().DescribeImageScanFindings(ctx, gomock.AssignableToTypeOf(&ecr.DescribeImageScanFindingsInput{})).Return(nil, errors.New("API error")) + + result, err := tool.GetImageScanFindings(ctx, info, imageID) + + assert.Error(t, err) + assert.Nil(t, result) + assert.Contains(t, err.Error(), "API error") + }) + + t.Run("nil ImageScanFindings returns error", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcrClient(ctrl) + tool := newEcrTool(mockClient) + + info := &ImageInfo{ + Repository: "my-repo", + Tag: "v1.0.0", + } + imageID := &ecrtypes.ImageIdentifier{ + ImageTag: aws.String("v1.0.0"), + } + + mockClient.EXPECT().DescribeImageScanFindings(ctx, gomock.AssignableToTypeOf(&ecr.DescribeImageScanFindingsInput{})).Return(&ecr.DescribeImageScanFindingsOutput{ + ImageScanFindings: nil, + }, nil) + + result, err := tool.GetImageScanFindings(ctx, info, imageID) + + assert.Error(t, err) + assert.Nil(t, result) + assert.Contains(t, err.Error(), "image scan findings missing for my-repo:v1.0.0") + }) +} + +func TestToCPUArchitecture(t *testing.T) { + t.Run("amd64 maps to x86_64", func(t *testing.T) { + assert.Equal(t, ecstypes.CPUArchitectureX8664, toCPUArchitecture("amd64")) + }) + + t.Run("arm64 maps to arm64", func(t *testing.T) { + assert.Equal(t, ecstypes.CPUArchitectureArm64, toCPUArchitecture("arm64")) + }) + + t.Run("unknown architecture maps to empty", func(t *testing.T) { + assert.Equal(t, ecstypes.CPUArchitecture(""), toCPUArchitecture("riscv")) + }) +} diff --git a/cli/cage/audit/ecs.go b/cli/cage/audit/ecs.go new file mode 100644 index 0000000..5817c39 --- /dev/null +++ b/cli/cage/audit/ecs.go @@ -0,0 +1,106 @@ +package audit + +import ( + "context" + "fmt" + "strings" + + "github.com/aws/aws-sdk-go-v2/service/ecs" + ecstypes "github.com/aws/aws-sdk-go-v2/service/ecs/types" + "github.com/loilo-inc/canarycage/awsiface" +) + +type ecsTool struct { + Ecs awsiface.EcsClient +} + +type EcsTool interface { + GetServiceImageInfos(ctx context.Context, cluster string, service string) ([]ImageInfo, error) +} + +func newEcsTool(ecsClient awsiface.EcsClient) EcsTool { + return &ecsTool{Ecs: ecsClient} +} + +func (t *ecsTool) GetServiceImageInfos(ctx context.Context, cluster string, service string) ([]ImageInfo, error) { + res, err := t.Ecs.DescribeServices(ctx, &ecs.DescribeServicesInput{ + Cluster: &cluster, + Services: []string{service}, + }) + if err != nil { + return nil, err + } + if len(res.Services) == 0 || res.Services[0].TaskDefinition == nil { + return nil, fmt.Errorf("service not found: %s/%s", cluster, service) + } + taskDefinition := *res.Services[0].TaskDefinition + + tdRes, err := t.Ecs.DescribeTaskDefinition(ctx, &ecs.DescribeTaskDefinitionInput{ + TaskDefinition: &taskDefinition, + }) + if err != nil { + return nil, err + } + td := tdRes.TaskDefinition + if td == nil { + return nil, fmt.Errorf("task definition not found: %s", taskDefinition) + } + + arch := ecstypes.CPUArchitectureX8664 + if td.RuntimePlatform != nil && td.RuntimePlatform.CpuArchitecture != "" { + arch = td.RuntimePlatform.CpuArchitecture + } + + if len(td.ContainerDefinitions) == 0 { + return nil, fmt.Errorf("no container definitions found for task definition: %s", taskDefinition) + } + + images := make([]ImageInfo, 0, len(td.ContainerDefinitions)) + for _, cd := range td.ContainerDefinitions { + if cd.Name == nil || cd.Image == nil { + return nil, fmt.Errorf("container definition is missing name or image: %s", taskDefinition) + } + parsed := ParseImageInfo(*cd.Image) + images = append(images, ImageInfo{ + ContainerName: *cd.Name, + PlatformArch: arch, + Registry: parsed.Registry, + Repository: parsed.Repository, + Tag: parsed.Tag, + }) + } + + return images, nil +} + +type ParsedImageInfo struct { + Registry string + Repository string + Tag string +} + +func ParseImageInfo(image string) ParsedImageInfo { + parts := strings.Split(image, "/") + if len(parts) == 1 { + repository, tag := splitRepoTag(image) + return ParsedImageInfo{Repository: repository, Tag: tag} + } + + registry := parts[0] + repoAndTag := strings.Join(parts[1:], "/") + repository, tag := splitRepoTag(repoAndTag) + return ParsedImageInfo{Registry: registry, Repository: repository, Tag: tag} +} + +func splitRepoTag(value string) (string, string) { + repository := value + tag := "latest" + if strings.Contains(value, ":") { + parts := strings.SplitN(value, ":", 2) + repository = parts[0] + if parts[1] != "" { + tag = parts[1] + } + } + return repository, tag +} diff --git a/cli/cage/audit/ecs_test.go b/cli/cage/audit/ecs_test.go new file mode 100644 index 0000000..7ff1c92 --- /dev/null +++ b/cli/cage/audit/ecs_test.go @@ -0,0 +1,280 @@ +package audit + +import ( + "context" + "errors" + "testing" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/ecs" + ecstypes "github.com/aws/aws-sdk-go-v2/service/ecs/types" + "github.com/loilo-inc/canarycage/mocks/mock_awsiface" + "github.com/stretchr/testify/assert" + "go.uber.org/mock/gomock" +) + +func TestGetServiceImageInfos(t *testing.T) { + ctx := context.Background() + + t.Run("returns image infos with default architecture", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcsClient(ctrl) + tool := newEcsTool(mockClient) + + mockClient.EXPECT().DescribeServices(ctx, gomock.AssignableToTypeOf(&ecs.DescribeServicesInput{})). + DoAndReturn(func(ctx context.Context, input *ecs.DescribeServicesInput, opts ...func(*ecs.Options)) (*ecs.DescribeServicesOutput, error) { + assert.Equal(t, "cluster-a", *input.Cluster) + assert.Equal(t, []string{"service-a"}, input.Services) + return &ecs.DescribeServicesOutput{ + Services: []ecstypes.Service{{TaskDefinition: aws.String("td:1")}}, + }, nil + }) + + mockClient.EXPECT().DescribeTaskDefinition(ctx, gomock.AssignableToTypeOf(&ecs.DescribeTaskDefinitionInput{})). + DoAndReturn(func(ctx context.Context, input *ecs.DescribeTaskDefinitionInput, opts ...func(*ecs.Options)) (*ecs.DescribeTaskDefinitionOutput, error) { + assert.Equal(t, "td:1", *input.TaskDefinition) + return &ecs.DescribeTaskDefinitionOutput{ + TaskDefinition: &ecstypes.TaskDefinition{ + ContainerDefinitions: []ecstypes.ContainerDefinition{ + { + Name: aws.String("app"), + Image: aws.String("123456789012.dkr.ecr.us-west-2.amazonaws.com/my-repo:1.2.3"), + }, + { + Name: aws.String("sidecar"), + Image: aws.String("nginx:latest"), + }, + }, + }, + }, nil + }) + + result, err := tool.GetServiceImageInfos(ctx, "cluster-a", "service-a") + + assert.NoError(t, err) + if assert.Len(t, result, 2) { + assert.Equal(t, "app", result[0].ContainerName) + assert.Equal(t, ecstypes.CPUArchitectureX8664, result[0].PlatformArch) + assert.Equal(t, "123456789012.dkr.ecr.us-west-2.amazonaws.com", result[0].Registry) + assert.Equal(t, "my-repo", result[0].Repository) + assert.Equal(t, "1.2.3", result[0].Tag) + assert.Equal(t, "sidecar", result[1].ContainerName) + assert.Equal(t, "", result[1].Registry) + assert.Equal(t, "nginx", result[1].Repository) + assert.Equal(t, "latest", result[1].Tag) + } + }) + + t.Run("uses runtime platform architecture when provided", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcsClient(ctrl) + tool := newEcsTool(mockClient) + + mockClient.EXPECT().DescribeServices(ctx, gomock.AssignableToTypeOf(&ecs.DescribeServicesInput{})). + Return(&ecs.DescribeServicesOutput{ + Services: []ecstypes.Service{{TaskDefinition: aws.String("td:2")}}, + }, nil) + + mockClient.EXPECT().DescribeTaskDefinition(ctx, gomock.AssignableToTypeOf(&ecs.DescribeTaskDefinitionInput{})). + Return(&ecs.DescribeTaskDefinitionOutput{ + TaskDefinition: &ecstypes.TaskDefinition{ + RuntimePlatform: &ecstypes.RuntimePlatform{CpuArchitecture: ecstypes.CPUArchitectureArm64}, + ContainerDefinitions: []ecstypes.ContainerDefinition{ + { + Name: aws.String("app"), + Image: aws.String("repo:v2"), + }, + }, + }, + }, nil) + + result, err := tool.GetServiceImageInfos(ctx, "cluster-a", "service-a") + + assert.NoError(t, err) + if assert.Len(t, result, 1) { + assert.Equal(t, ecstypes.CPUArchitectureArm64, result[0].PlatformArch) + } + }) + + t.Run("DescribeServices error returns error", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcsClient(ctrl) + tool := newEcsTool(mockClient) + + mockClient.EXPECT().DescribeServices(ctx, gomock.AssignableToTypeOf(&ecs.DescribeServicesInput{})). + Return(nil, errors.New("API error")) + + result, err := tool.GetServiceImageInfos(ctx, "cluster-a", "service-a") + + assert.Error(t, err) + assert.Nil(t, result) + assert.Contains(t, err.Error(), "API error") + }) + + t.Run("service not found returns error", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcsClient(ctrl) + tool := newEcsTool(mockClient) + + mockClient.EXPECT().DescribeServices(ctx, gomock.AssignableToTypeOf(&ecs.DescribeServicesInput{})). + Return(&ecs.DescribeServicesOutput{Services: []ecstypes.Service{}}, nil) + + result, err := tool.GetServiceImageInfos(ctx, "cluster-a", "service-a") + + assert.Error(t, err) + assert.Nil(t, result) + assert.Contains(t, err.Error(), "service not found") + }) + + t.Run("missing task definition returns error", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcsClient(ctrl) + tool := newEcsTool(mockClient) + + mockClient.EXPECT().DescribeServices(ctx, gomock.AssignableToTypeOf(&ecs.DescribeServicesInput{})). + Return(&ecs.DescribeServicesOutput{ + Services: []ecstypes.Service{{TaskDefinition: nil}}, + }, nil) + + result, err := tool.GetServiceImageInfos(ctx, "cluster-a", "service-a") + + assert.Error(t, err) + assert.Nil(t, result) + assert.Contains(t, err.Error(), "service not found") + }) + + t.Run("DescribeTaskDefinition error returns error", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcsClient(ctrl) + tool := newEcsTool(mockClient) + + mockClient.EXPECT().DescribeServices(ctx, gomock.AssignableToTypeOf(&ecs.DescribeServicesInput{})). + Return(&ecs.DescribeServicesOutput{ + Services: []ecstypes.Service{{TaskDefinition: aws.String("td:1")}}, + }, nil) + + mockClient.EXPECT().DescribeTaskDefinition(ctx, gomock.AssignableToTypeOf(&ecs.DescribeTaskDefinitionInput{})). + Return(nil, errors.New("API error")) + + result, err := tool.GetServiceImageInfos(ctx, "cluster-a", "service-a") + + assert.Error(t, err) + assert.Nil(t, result) + assert.Contains(t, err.Error(), "API error") + }) + + t.Run("nil task definition returns error", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcsClient(ctrl) + tool := newEcsTool(mockClient) + + mockClient.EXPECT().DescribeServices(ctx, gomock.AssignableToTypeOf(&ecs.DescribeServicesInput{})). + Return(&ecs.DescribeServicesOutput{ + Services: []ecstypes.Service{{TaskDefinition: aws.String("td:1")}}, + }, nil) + + mockClient.EXPECT().DescribeTaskDefinition(ctx, gomock.AssignableToTypeOf(&ecs.DescribeTaskDefinitionInput{})). + Return(&ecs.DescribeTaskDefinitionOutput{TaskDefinition: nil}, nil) + + result, err := tool.GetServiceImageInfos(ctx, "cluster-a", "service-a") + + assert.Error(t, err) + assert.Nil(t, result) + assert.Contains(t, err.Error(), "task definition not found") + }) + + t.Run("no container definitions returns error", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcsClient(ctrl) + tool := newEcsTool(mockClient) + + mockClient.EXPECT().DescribeServices(ctx, gomock.AssignableToTypeOf(&ecs.DescribeServicesInput{})). + Return(&ecs.DescribeServicesOutput{ + Services: []ecstypes.Service{{TaskDefinition: aws.String("td:1")}}, + }, nil) + + mockClient.EXPECT().DescribeTaskDefinition(ctx, gomock.AssignableToTypeOf(&ecs.DescribeTaskDefinitionInput{})). + Return(&ecs.DescribeTaskDefinitionOutput{ + TaskDefinition: &ecstypes.TaskDefinition{ + ContainerDefinitions: []ecstypes.ContainerDefinition{}, + }, + }, nil) + + result, err := tool.GetServiceImageInfos(ctx, "cluster-a", "service-a") + + assert.Error(t, err) + assert.Nil(t, result) + assert.Contains(t, err.Error(), "no container definitions") + }) + + t.Run("container definition missing name or image returns error", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := mock_awsiface.NewMockEcsClient(ctrl) + tool := newEcsTool(mockClient) + + mockClient.EXPECT().DescribeServices(ctx, gomock.AssignableToTypeOf(&ecs.DescribeServicesInput{})). + Return(&ecs.DescribeServicesOutput{ + Services: []ecstypes.Service{{TaskDefinition: aws.String("td:1")}}, + }, nil) + + mockClient.EXPECT().DescribeTaskDefinition(ctx, gomock.AssignableToTypeOf(&ecs.DescribeTaskDefinitionInput{})). + Return(&ecs.DescribeTaskDefinitionOutput{ + TaskDefinition: &ecstypes.TaskDefinition{ + ContainerDefinitions: []ecstypes.ContainerDefinition{ + { + Name: nil, + Image: aws.String("repo:v1"), + }, + }, + }, + }, nil) + + result, err := tool.GetServiceImageInfos(ctx, "cluster-a", "service-a") + + assert.Error(t, err) + assert.Nil(t, result) + assert.Contains(t, err.Error(), "container definition is missing name or image") + }) +} + +func TestParseImageInfo(t *testing.T) { + t.Run("parses registry, repository, and tag", func(t *testing.T) { + parsed := ParseImageInfo("123456789012.dkr.ecr.us-west-2.amazonaws.com/my-repo:1.2.3") + assert.Equal(t, "123456789012.dkr.ecr.us-west-2.amazonaws.com", parsed.Registry) + assert.Equal(t, "my-repo", parsed.Repository) + assert.Equal(t, "1.2.3", parsed.Tag) + }) + + t.Run("parses repository without registry", func(t *testing.T) { + parsed := ParseImageInfo("nginx:latest") + assert.Equal(t, "", parsed.Registry) + assert.Equal(t, "nginx", parsed.Repository) + assert.Equal(t, "latest", parsed.Tag) + }) + + t.Run("defaults tag to latest", func(t *testing.T) { + parsed := ParseImageInfo("nginx") + assert.Equal(t, "nginx", parsed.Repository) + assert.Equal(t, "latest", parsed.Tag) + }) +} + +func TestSplitRepoTag(t *testing.T) { + t.Run("returns latest when tag is missing", func(t *testing.T) { + repo, tag := splitRepoTag("repo") + assert.Equal(t, "repo", repo) + assert.Equal(t, "latest", tag) + }) + + t.Run("returns tag when present", func(t *testing.T) { + repo, tag := splitRepoTag("repo:v1") + assert.Equal(t, "repo", repo) + assert.Equal(t, "v1", tag) + }) + + t.Run("returns latest when tag is empty", func(t *testing.T) { + repo, tag := splitRepoTag("repo:") + assert.Equal(t, "repo", repo) + assert.Equal(t, "latest", tag) + }) +} diff --git a/cli/cage/audit/printer.go b/cli/cage/audit/printer.go new file mode 100644 index 0000000..83340b7 --- /dev/null +++ b/cli/cage/audit/printer.go @@ -0,0 +1,137 @@ +package audit + +import ( + "fmt" + "strings" + + ecrtypes "github.com/aws/aws-sdk-go-v2/service/ecr/types" + "github.com/loilo-inc/canarycage/logger" +) + +type printer struct { + logger logger.Logger + color logger.Color + logDetail bool +} + +type Printer interface { + Print(result []*ScanResult) +} + +func NewPrinter(l logger.Logger, noColor, logDetail bool) *printer { + return &printer{ + logger: l, + color: logger.Color{NoColor: noColor}, + logDetail: logDetail, + } +} + +func (p *printer) Print(result []*ScanResult) { + containerMax, imageMax := MaxHeaderWidth(result) + // |container|status|critical|high|medium|low|info|image| + headerFmt := fmt.Sprintf("|%%-%ds|%%-10s|%%-8s|%%-5s|%%-6s|%%-4s|%%-4s|%%-%ds|\n", containerMax, imageMax) + p.logger.Printf(headerFmt, "CONTAINER", "STATUS", "CRITICAL", "HIGH", "MEDIUM", "LOW", "INFO", "IMAGE") + bodyFmt := fmt.Sprintf("|%%-%ds|%%-10s|%%-8d|%%-5d|%%-6d|%%-4d|%%-4d|%%-%ds|\n", containerMax, imageMax) + agg := NewAggregater() + for _, r := range result { + agg.Add(r) + } + for _, summaries := range agg.summaries { + for _, summary := range summaries { + p.logger.Printf( + bodyFmt, + summary.ContainerName, + summary.Status, + summary.CriticalCount, + summary.HighCount, + summary.MediumCount, + summary.LowCount, + summary.InfoCount, + summary.ImageURI, + ) + } + } + p.logImageScanFindings("CRITICAL", agg.CriticalCves(), agg) + p.logImageScanFindings("HIGH", agg.HighCves(), agg) + p.logImageScanFindings("MEDIUM", agg.MediumCves(), agg) + total := agg.TotalCVECount() + color := p.color + if total == 0 { + p.logger.Printf("%s\n", color.Greenf("No CVEs found")) + return + } + summary := agg.SummarizeTotal() + highest := &severityPrinter{ + severity: summary.HighestSeverity, + color: p.color, + } + var list []string + for _, v := range summary.SeverityCounts() { + if v.Count == 0 { + continue + } + sp := &severityPrinter{severity: v.Severity, color: p.color} + list = append(list, fmt.Sprintf("%d %s", v.Count, sp.BSprintf("%s", v.Severity))) + } + + p.logger.Printf( + "\nTotal: %s (%s)\n", + highest.BSprintf("%d", summary.TotalCount), + strings.Join(list, ", "), + ) +} + +func (p *printer) logImageScanFindings( + severity ecrtypes.FindingSeverity, + findings []ecrtypes.ImageScanFinding, + aggregater *aggregater, +) { + if len(findings) == 0 { + return + } + sp := &severityPrinter{severity: severity, color: p.color} + color := p.color + p.logger.Printf("\n=== %s ===\n", sp.BSprintf("%s", severity)) + for _, cve := range findings { + containers := aggregater.GetVulnContainers(*cve.Name) + var containerList []string + for _, c := range containers { + containerList = append(containerList, color.Bold(c)) + } + var packageName string = "unknown" + var packageVersion string = "unknown" + for _, attr := range cve.Attributes { + switch *attr.Key { + case "package_name": + packageName = *attr.Value + case "package_version": + packageVersion = *attr.Value + } + } + p.logger.Printf("- %s %s \n", *cve.Name, strings.Join(containerList, ", ")) + p.logger.Printf(" %s::%s (%s)\n", + packageName, packageVersion, *cve.Uri) + if p.logDetail { + p.logger.Printf("\n%s\n", *cve.Description) + } + } +} + +func (i *ImageInfo) formatImageLabel() string { + return fmt.Sprintf("%s/%s:%s", i.Registry, i.Repository, i.Tag) +} + +func MaxHeaderWidth(imageInfos []*ScanResult) (int, int) { + containerMax := len("CONTAINER") + imageMax := len("IMAGE") + for _, info := range imageInfos { + if l := len(info.ImageInfo.ContainerName); l > containerMax { + containerMax = l + } + imageLabel := info.formatImageLabel() + if l := len(imageLabel); l > imageMax { + imageMax = l + } + } + return containerMax, imageMax +} diff --git a/cli/cage/audit/printer_test.go b/cli/cage/audit/printer_test.go new file mode 100644 index 0000000..e38a326 --- /dev/null +++ b/cli/cage/audit/printer_test.go @@ -0,0 +1,353 @@ +package audit + +import ( + "fmt" + "testing" + + "github.com/aws/aws-sdk-go-v2/aws" + ecrtypes "github.com/aws/aws-sdk-go-v2/service/ecr/types" +) + +type mockLogger struct { + logs []string +} + +func (m *mockLogger) Printf(format string, args ...any) { + m.logs = append(m.logs, fmt.Sprintf(format, args...)) +} + +func makeScanResult( + list ...ecrtypes.FindingSeverity) []*ScanResult { + return []*ScanResult{ + { + ImageInfo: ImageInfo{ + ContainerName: "test-container", + Registry: "test-registry", + Repository: "test-repo", + Tag: "latest", + }, + ImageScanFindings: &ecrtypes.ImageScanFindings{ + Findings: makeFindings(list), + }, + }, + } +} + +func makeFindings(severities []ecrtypes.FindingSeverity) []ecrtypes.ImageScanFinding { + findings := make([]ecrtypes.ImageScanFinding, len(severities)) + for i, sev := range severities { + findings[i] = ecrtypes.ImageScanFinding{ + Severity: sev, + Name: aws.String(fmt.Sprintf("CVE-2023-%04d", i+1)), + Uri: aws.String("http://example.com"), + Description: aws.String("Test vulnerability description"), + } + } + return findings +} + +func TestPrinter_Print(t *testing.T) { + t.Run("prints no CVEs message when no findings", func(t *testing.T) { + logger := &mockLogger{} + printer := NewPrinter(logger, true, false) + + result := makeScanResult() + printer.Print(result) + + // Check that "No CVEs found" message is present + found := false + for _, log := range logger.logs { + if log == "No CVEs found\n" { + found = true + break + } + } + if !found { + t.Error("Expected 'No CVEs found' message") + } + }) + + t.Run("prints table header", func(t *testing.T) { + logger := &mockLogger{} + printer := NewPrinter(logger, true, false) + + result := makeScanResult(ecrtypes.FindingSeverityCritical) + printer.Print(result) + + // Check that header contains expected columns + if len(logger.logs) == 0 { + t.Fatal("Expected logs to be generated") + } + header := logger.logs[0] + expectedCols := []string{"CONTAINER", "STATUS", "CRITICAL", "HIGH", "MEDIUM", "LOW", "INFO", "IMAGE"} + for _, col := range expectedCols { + if !containsString(header, col) { + t.Errorf("Header missing column: %s", col) + } + } + }) + + t.Run("prints findings by severity", func(t *testing.T) { + logger := &mockLogger{} + printer := NewPrinter(logger, true, false) + + result := makeScanResult( + ecrtypes.FindingSeverityCritical, + ecrtypes.FindingSeverityHigh, + ecrtypes.FindingSeverityMedium, + ) + printer.Print(result) + + // Should have CRITICAL, HIGH, MEDIUM sections + criticalFound := false + highFound := false + mediumFound := false + for _, log := range logger.logs { + if containsString(log, "CRITICAL") && containsString(log, "===") { + criticalFound = true + } + if containsString(log, "HIGH") && containsString(log, "===") { + highFound = true + } + if containsString(log, "MEDIUM") && containsString(log, "===") { + mediumFound = true + } + } + if !criticalFound { + t.Error("Expected CRITICAL section") + } + if !highFound { + t.Error("Expected HIGH section") + } + if !mediumFound { + t.Error("Expected MEDIUM section") + } + }) + + t.Run("prints total summary with counts", func(t *testing.T) { + logger := &mockLogger{} + printer := NewPrinter(logger, true, false) + + result := makeScanResult( + ecrtypes.FindingSeverityCritical, + ecrtypes.FindingSeverityHigh, + ) + printer.Print(result) + + // Check for total line + totalFound := false + for _, log := range logger.logs { + if containsString(log, "Total:") { + totalFound = true + break + } + } + if !totalFound { + t.Error("Expected Total summary line") + } + }) +} + +func TestPrinter_logImageScanFindings(t *testing.T) { + t.Run("returns early when no findings", func(t *testing.T) { + logger := &mockLogger{} + printer := NewPrinter(logger, true, false) + agg := NewAggregater() + + printer.logImageScanFindings(ecrtypes.FindingSeverityCritical, []ecrtypes.ImageScanFinding{}, agg) + + if len(logger.logs) != 0 { + t.Errorf("Expected no logs, got %d", len(logger.logs)) + } + }) + + t.Run("prints severity header", func(t *testing.T) { + logger := &mockLogger{} + printer := NewPrinter(logger, true, false) + agg := NewAggregater() + + findings := []ecrtypes.ImageScanFinding{ + { + Name: aws.String("CVE-2023-0001"), + Uri: aws.String("http://example.com"), + Description: aws.String("Test description"), + Attributes: []ecrtypes.Attribute{}, + }, + } + + printer.logImageScanFindings(ecrtypes.FindingSeverityCritical, findings, agg) + + headerFound := false + for _, log := range logger.logs { + if containsString(log, "CRITICAL") && containsString(log, "===") { + headerFound = true + break + } + } + if !headerFound { + t.Error("Expected severity header with CRITICAL") + } + }) + + t.Run("prints CVE name and URI", func(t *testing.T) { + logger := &mockLogger{} + printer := NewPrinter(logger, true, false) + agg := NewAggregater() + + findings := []ecrtypes.ImageScanFinding{ + { + Name: aws.String("CVE-2023-0001"), + Uri: aws.String("http://example.com/cve"), + Description: aws.String("Test description"), + Attributes: []ecrtypes.Attribute{}, + }, + } + + printer.logImageScanFindings(ecrtypes.FindingSeverityHigh, findings, agg) + + cveFound := false + uriFound := false + for _, log := range logger.logs { + if containsString(log, "CVE-2023-0001") { + cveFound = true + } + if containsString(log, "http://example.com/cve") { + uriFound = true + } + } + if !cveFound { + t.Error("Expected CVE name in output") + } + if !uriFound { + t.Error("Expected CVE URI in output") + } + }) + + t.Run("extracts package name and version from attributes", func(t *testing.T) { + logger := &mockLogger{} + printer := NewPrinter(logger, true, false) + agg := NewAggregater() + + findings := []ecrtypes.ImageScanFinding{ + { + Name: aws.String("CVE-2023-0001"), + Uri: aws.String("http://example.com"), + Description: aws.String("Test description"), + Attributes: []ecrtypes.Attribute{ + {Key: aws.String("package_name"), Value: aws.String("test-package")}, + {Key: aws.String("package_version"), Value: aws.String("1.2.3")}, + }, + }, + } + + printer.logImageScanFindings(ecrtypes.FindingSeverityMedium, findings, agg) + + packageFound := false + for _, log := range logger.logs { + if containsString(log, "test-package::1.2.3") { + packageFound = true + break + } + } + if !packageFound { + t.Error("Expected package name and version in output") + } + }) + + t.Run("uses unknown for missing package info", func(t *testing.T) { + logger := &mockLogger{} + printer := NewPrinter(logger, true, false) + agg := NewAggregater() + + findings := []ecrtypes.ImageScanFinding{ + { + Name: aws.String("CVE-2023-0001"), + Uri: aws.String("http://example.com"), + Description: aws.String("Test description"), + Attributes: []ecrtypes.Attribute{}, + }, + } + + printer.logImageScanFindings(ecrtypes.FindingSeverityLow, findings, agg) + + unknownFound := false + for _, log := range logger.logs { + if containsString(log, "unknown::unknown") { + unknownFound = true + break + } + } + if !unknownFound { + t.Error("Expected unknown::unknown for missing package info") + } + }) + + t.Run("prints description when logDetail is true", func(t *testing.T) { + logger := &mockLogger{} + printer := NewPrinter(logger, true, true) // logDetail = true + agg := NewAggregater() + + findings := []ecrtypes.ImageScanFinding{ + { + Name: aws.String("CVE-2023-0001"), + Uri: aws.String("http://example.com"), + Description: aws.String("Detailed vulnerability description"), + Attributes: []ecrtypes.Attribute{}, + }, + } + + printer.logImageScanFindings(ecrtypes.FindingSeverityHigh, findings, agg) + + descFound := false + for _, log := range logger.logs { + if containsString(log, "Detailed vulnerability description") { + descFound = true + break + } + } + if !descFound { + t.Error("Expected description in output when logDetail is true") + } + }) + + t.Run("does not print description when logDetail is false", func(t *testing.T) { + logger := &mockLogger{} + printer := NewPrinter(logger, true, false) // logDetail = false + agg := NewAggregater() + + findings := []ecrtypes.ImageScanFinding{ + { + Name: aws.String("CVE-2023-0001"), + Uri: aws.String("http://example.com"), + Description: aws.String("Detailed vulnerability description"), + Attributes: []ecrtypes.Attribute{}, + }, + } + + printer.logImageScanFindings(ecrtypes.FindingSeverityHigh, findings, agg) + + descFound := false + for _, log := range logger.logs { + if containsString(log, "Detailed vulnerability description") { + descFound = true + break + } + } + if descFound { + t.Error("Expected no description in output when logDetail is false") + } + }) +} + +func containsString(s, substr string) bool { + return len(s) >= len(substr) && (s == substr || len(s) > len(substr) && stringContains(s, substr)) +} + +func stringContains(s, substr string) bool { + for i := 0; i <= len(s)-len(substr); i++ { + if s[i:i+len(substr)] == substr { + return true + } + } + return false +} diff --git a/cli/cage/audit/scanner.go b/cli/cage/audit/scanner.go new file mode 100644 index 0000000..6d8b029 --- /dev/null +++ b/cli/cage/audit/scanner.go @@ -0,0 +1,55 @@ +package audit + +import ( + "context" + "fmt" + + "github.com/loilo-inc/canarycage/awsiface" +) + +type scanner struct { + ecs awsiface.EcsClient + ecr awsiface.EcrClient +} + +type Scanner interface { + Scan(ctx context.Context, cluster string, service string) ([]*ScanResult, error) +} + +func NewScanner(ecs awsiface.EcsClient, ecr awsiface.EcrClient) Scanner { + return &scanner{ecs: ecs, ecr: ecr} +} + +func (s *scanner) Scan( + ctx context.Context, + cluster string, + service string, +) (results []*ScanResult, err error) { + ecsTool := newEcsTool(s.ecs) + ecrTool := newEcrTool(s.ecr) + var imageInfos []ImageInfo + if imageInfos, err = ecsTool.GetServiceImageInfos(ctx, cluster, service); err != nil { + return nil, err + } + findingsList := make([]*ScanResult, len(imageInfos)) + for i, info := range imageInfos { + if info.IsECRImage() { + findingsList[i] = scanImage(ctx, ecrTool, info) + } else { + findingsList[i] = &ScanResult{ImageInfo: info, Err: ErrNonEcrImage} + } + } + return findingsList, nil +} + +var ErrNonEcrImage = fmt.Errorf("non-ECR image") + +func scanImage(ctx context.Context, ecrTool EcrTool, info ImageInfo) *ScanResult { + if imageID, err := ecrTool.GetActualImageIdentifier(ctx, &info); err != nil { + return &ScanResult{ImageInfo: info, Err: err} + } else if findings, err := ecrTool.GetImageScanFindings(ctx, &info, imageID); err != nil { + return &ScanResult{ImageInfo: info, Err: err} + } else { + return &ScanResult{ImageInfo: info, ImageScanFindings: findings} + } +} diff --git a/cli/cage/audit/scanner_test.go b/cli/cage/audit/scanner_test.go new file mode 100644 index 0000000..7ceca8e --- /dev/null +++ b/cli/cage/audit/scanner_test.go @@ -0,0 +1,157 @@ +package audit + +import ( + "context" + "encoding/json" + "errors" + "testing" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/ecr" + ecrtypes "github.com/aws/aws-sdk-go-v2/service/ecr/types" + "github.com/aws/aws-sdk-go-v2/service/ecs" + ecstypes "github.com/aws/aws-sdk-go-v2/service/ecs/types" + "github.com/loilo-inc/canarycage/mocks/mock_awsiface" + "github.com/stretchr/testify/assert" + "go.uber.org/mock/gomock" +) + +type stubEcrTool struct { + imageID *ecrtypes.ImageIdentifier + findings *ecrtypes.ImageScanFindings + errID error + errScan error +} + +func (s *stubEcrTool) GetActualImageIdentifier(ctx context.Context, info *ImageInfo) (*ecrtypes.ImageIdentifier, error) { + if s.errID != nil { + return nil, s.errID + } + return s.imageID, nil +} + +func (s *stubEcrTool) GetImageScanFindings(ctx context.Context, info *ImageInfo, imageID *ecrtypes.ImageIdentifier) (*ecrtypes.ImageScanFindings, error) { + if s.errScan != nil { + return nil, s.errScan + } + return s.findings, nil +} + +func TestScanner_Scan(t *testing.T) { + ctx := context.Background() + + t.Run("returns scan results for each image", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockEcs := mock_awsiface.NewMockEcsClient(ctrl) + mockEcr := mock_awsiface.NewMockEcrClient(ctrl) + scanner := NewScanner(mockEcs, mockEcr) + + mockEcs.EXPECT().DescribeServices(ctx, gomock.AssignableToTypeOf(&ecs.DescribeServicesInput{})). + Return(&ecs.DescribeServicesOutput{ + Services: []ecstypes.Service{{TaskDefinition: aws.String("td:1")}}, + }, nil) + + mockEcs.EXPECT().DescribeTaskDefinition(ctx, gomock.AssignableToTypeOf(&ecs.DescribeTaskDefinitionInput{})). + Return(&ecs.DescribeTaskDefinitionOutput{ + TaskDefinition: &ecstypes.TaskDefinition{ + ContainerDefinitions: []ecstypes.ContainerDefinition{ + { + Name: aws.String("app"), + Image: aws.String("123456789012.dkr.ecr.us-west-2.amazonaws.com/my-repo:1.2.3"), + }, + { + Name: aws.String("sidecar"), + Image: aws.String("nginx:latest"), + }, + }, + }, + }, nil) + + manifestJSON, _ := json.Marshal(dockerSchema{ + SchemaVersion: 2, + MediaType: "application/vnd.docker.distribution.manifest.v2+json", + }) + manifestStr := string(manifestJSON) + + mockEcr.EXPECT().BatchGetImage(ctx, gomock.AssignableToTypeOf(&ecr.BatchGetImageInput{})). + DoAndReturn(func(ctx context.Context, input *ecr.BatchGetImageInput, opts ...func(*ecr.Options)) (*ecr.BatchGetImageOutput, error) { + assert.Equal(t, "my-repo", *input.RepositoryName) + assert.Equal(t, "1.2.3", *input.ImageIds[0].ImageTag) + return &ecr.BatchGetImageOutput{ + Images: []ecrtypes.Image{{ImageManifest: &manifestStr}}, + }, nil + }) + + mockEcr.EXPECT().DescribeImageScanFindings(ctx, gomock.AssignableToTypeOf(&ecr.DescribeImageScanFindingsInput{})). + DoAndReturn(func(ctx context.Context, input *ecr.DescribeImageScanFindingsInput, opts ...func(*ecr.Options)) (*ecr.DescribeImageScanFindingsOutput, error) { + assert.Equal(t, "my-repo", *input.RepositoryName) + assert.Equal(t, "1.2.3", *input.ImageId.ImageTag) + return &ecr.DescribeImageScanFindingsOutput{ + ImageScanFindings: &ecrtypes.ImageScanFindings{}, + }, nil + }) + + results, err := scanner.Scan(ctx, "cluster-a", "service-a") + + assert.NoError(t, err) + if assert.Len(t, results, 2) { + assert.Equal(t, "app", results[0].ImageInfo.ContainerName) + assert.NoError(t, results[0].Err) + assert.Equal(t, "sidecar", results[1].ImageInfo.ContainerName) + assert.EqualError(t, results[1].Err, "non-ECR image") + } + }) + + t.Run("ecs error returns error", func(t *testing.T) { + ctrl := gomock.NewController(t) + mockEcs := mock_awsiface.NewMockEcsClient(ctrl) + mockEcr := mock_awsiface.NewMockEcrClient(ctrl) + scanner := NewScanner(mockEcs, mockEcr) + + mockEcs.EXPECT().DescribeServices(ctx, gomock.AssignableToTypeOf(&ecs.DescribeServicesInput{})). + Return(nil, errors.New("ecs error")) + + results, err := scanner.Scan(ctx, "cluster-a", "service-a") + + assert.EqualError(t, err, "ecs error") + assert.Nil(t, results) + }) +} + +func TestScanImage(t *testing.T) { + ctx := context.Background() + + t.Run("GetActualImageIdentifier error returns error", func(t *testing.T) { + tool := &stubEcrTool{ + errID: errors.New("id error"), + } + + result := scanImage(ctx, tool, ImageInfo{Repository: "repo"}) + + assert.EqualError(t, result.Err, "id error") + }) + + t.Run("GetImageScanFindings error returns error", func(t *testing.T) { + tool := &stubEcrTool{ + imageID: &ecrtypes.ImageIdentifier{ImageTag: aws.String("v1")}, + errScan: errors.New("scan error"), + } + + result := scanImage(ctx, tool, ImageInfo{Repository: "repo"}) + + assert.EqualError(t, result.Err, "scan error") + }) + + t.Run("success returns findings", func(t *testing.T) { + findings := &ecrtypes.ImageScanFindings{} + tool := &stubEcrTool{ + imageID: &ecrtypes.ImageIdentifier{ImageTag: aws.String("v1")}, + findings: findings, + } + + result := scanImage(ctx, tool, ImageInfo{Repository: "repo"}) + + assert.NoError(t, result.Err) + assert.Equal(t, findings, result.ImageScanFindings) + }) +} diff --git a/cli/cage/audit/types.go b/cli/cage/audit/types.go new file mode 100644 index 0000000..741fc17 --- /dev/null +++ b/cli/cage/audit/types.go @@ -0,0 +1,81 @@ +package audit + +import ( + "regexp" + + ecrtypes "github.com/aws/aws-sdk-go-v2/service/ecr/types" + ecstypes "github.com/aws/aws-sdk-go-v2/service/ecs/types" +) + +type ImageInfo struct { + ContainerName string + Registry string + Repository string + Tag string + PlatFormOS ecstypes.OSFamily + PlatformArch ecstypes.CPUArchitecture +} + +func (i *ImageInfo) IsECRImage() bool { + return i.Registry == "public.ecr.aws" || i.registryHasECRSuffix() +} + +var ecrURLPattern = regexp.MustCompile(`^[0-9]{12}\.dkr\.ecr\.[a-zA-Z0-9-]+\.amazonaws\.com$`) + +func (i *ImageInfo) registryHasECRSuffix() bool { + return ecrURLPattern.MatchString(i.Registry) +} + +type ScanResult struct { + ImageInfo + ImageScanFindings *ecrtypes.ImageScanFindings + Err error +} + +type ScanResultSummary struct { + ContainerName string + Status string + CriticalCount int32 + HighCount int32 + MediumCount int32 + LowCount int32 + InfoCount int32 + ImageURI string +} + +func summaryScanResult(result *ScanResult) *ScanResultSummary { + var status = "OK" + var critical, high, medium, low, info int32 + findings := result.ImageScanFindings + for _, f := range findings.Findings { + switch f.Severity { + case "CRITICAL": + critical++ + case "HIGH": + high++ + case "MEDIUM": + medium++ + case "LOW": + low++ + case "INFORMATIONAL": + info++ + } + } + if len(result.ImageScanFindings.Findings) == 0 { + status = "NONE" + } else if critical > 0 || high > 0 { + status = "VULNERABLE" + } else if medium > 0 { + status = "WARNING" + } + return &ScanResultSummary{ + ContainerName: result.ContainerName, + Status: status, + CriticalCount: critical, + HighCount: high, + MediumCount: medium, + LowCount: low, + InfoCount: info, + ImageURI: result.formatImageLabel(), + } +} diff --git a/cli/cage/audit/types_test.go b/cli/cage/audit/types_test.go new file mode 100644 index 0000000..3abd2a0 --- /dev/null +++ b/cli/cage/audit/types_test.go @@ -0,0 +1,245 @@ +package audit + +import ( + "testing" + + ecrtypes "github.com/aws/aws-sdk-go-v2/service/ecr/types" +) + +func TestImageInfo_IsECRImage(t *testing.T) { + tests := []struct { + name string + registry string + want bool + }{ + { + name: "public ECR registry", + registry: "public.ecr.aws", + want: true, + }, + { + name: "private ECR registry with standard suffix", + registry: "123456789012.dkr.ecr.us-east-1.amazonaws.com", + want: true, + }, + { + name: "private ECR registry with different region", + registry: "123456789012.dkr.ecr.eu-west-1.amazonaws.com", + want: true, + }, + { + name: "Docker Hub registry", + registry: "docker.io", + want: false, + }, + { + name: "empty registry", + registry: "", + want: false, + }, + { + name: "non-ECR AWS registry", + registry: "amazonaws.com", + want: false, + }, + { + name: "registry with partial ECR suffix", + registry: "example.com", + want: false, + }, + { + name: "registry with ECR substring but not suffix", + registry: ".dkr.ecr.amazonaws.com.example.com", + want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + i := ImageInfo{ + Registry: tt.registry, + } + if got := i.IsECRImage(); got != tt.want { + t.Errorf("ImageInfo.IsECRImage() = %v, want %v", got, tt.want) + } + }) + } +} +func Test_summaryScanResult(t *testing.T) { + tests := []struct { + name string + result *ScanResult + want *ScanResultSummary + }{ + { + name: "no findings - NONE status", + result: &ScanResult{ + ImageInfo: ImageInfo{ + ContainerName: "test-container", + Registry: "123456789012.dkr.ecr.us-east-1.amazonaws.com", + Repository: "test-repo", + Tag: "latest", + }, + ImageScanFindings: &ecrtypes.ImageScanFindings{ + Findings: []ecrtypes.ImageScanFinding{}, + }, + }, + want: &ScanResultSummary{ + ContainerName: "test-container", + Status: "NONE", + CriticalCount: 0, + HighCount: 0, + MediumCount: 0, + LowCount: 0, + InfoCount: 0, + ImageURI: "123456789012.dkr.ecr.us-east-1.amazonaws.com/test-repo:latest", + }, + }, + { + name: "critical findings - VULNERABLE status", + result: &ScanResult{ + ImageInfo: ImageInfo{ + ContainerName: "test-container", + }, + ImageScanFindings: &ecrtypes.ImageScanFindings{ + Findings: []ecrtypes.ImageScanFinding{ + {Severity: ecrtypes.FindingSeverityCritical}, + {Severity: ecrtypes.FindingSeverityHigh}, + }, + }, + }, + want: &ScanResultSummary{ + ContainerName: "test-container", + Status: "VULNERABLE", + CriticalCount: 1, + HighCount: 1, + MediumCount: 0, + LowCount: 0, + InfoCount: 0, + }, + }, + { + name: "high findings - VULNERABLE status", + result: &ScanResult{ + ImageInfo: ImageInfo{ + ContainerName: "test-container", + }, + ImageScanFindings: &ecrtypes.ImageScanFindings{ + Findings: []ecrtypes.ImageScanFinding{ + {Severity: ecrtypes.FindingSeverityHigh}, + {Severity: ecrtypes.FindingSeverityHigh}, + }, + }, + }, + want: &ScanResultSummary{ + ContainerName: "test-container", + Status: "VULNERABLE", + CriticalCount: 0, + HighCount: 2, + MediumCount: 0, + LowCount: 0, + InfoCount: 0, + }, + }, + { + name: "medium findings - WARNING status", + result: &ScanResult{ + ImageInfo: ImageInfo{ + ContainerName: "test-container", + }, + ImageScanFindings: &ecrtypes.ImageScanFindings{ + Findings: []ecrtypes.ImageScanFinding{ + {Severity: ecrtypes.FindingSeverityMedium}, + {Severity: ecrtypes.FindingSeverityLow}, + }, + }, + }, + want: &ScanResultSummary{ + ContainerName: "test-container", + Status: "WARNING", + CriticalCount: 0, + HighCount: 0, + MediumCount: 1, + LowCount: 1, + InfoCount: 0, + }, + }, + { + name: "low and informational findings - empty status", + result: &ScanResult{ + ImageInfo: ImageInfo{ + ContainerName: "test-container", + }, + ImageScanFindings: &ecrtypes.ImageScanFindings{ + Findings: []ecrtypes.ImageScanFinding{ + {Severity: ecrtypes.FindingSeverityLow}, + {Severity: ecrtypes.FindingSeverityInformational}, + }, + }, + }, + want: &ScanResultSummary{ + ContainerName: "test-container", + Status: "OK", + CriticalCount: 0, + HighCount: 0, + MediumCount: 0, + LowCount: 1, + InfoCount: 1, + }, + }, + { + name: "mixed severity findings", + result: &ScanResult{ + ImageInfo: ImageInfo{ + ContainerName: "mixed-container", + }, + ImageScanFindings: &ecrtypes.ImageScanFindings{ + Findings: []ecrtypes.ImageScanFinding{ + {Severity: ecrtypes.FindingSeverityCritical}, + {Severity: ecrtypes.FindingSeverityCritical}, + {Severity: ecrtypes.FindingSeverityHigh}, + {Severity: ecrtypes.FindingSeverityMedium}, + {Severity: ecrtypes.FindingSeverityLow}, + {Severity: ecrtypes.FindingSeverityInformational}, + }, + }, + }, + want: &ScanResultSummary{ + ContainerName: "mixed-container", + Status: "VULNERABLE", + CriticalCount: 2, + HighCount: 1, + MediumCount: 1, + LowCount: 1, + InfoCount: 1, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := summaryScanResult(tt.result) + if got.ContainerName != tt.want.ContainerName { + t.Errorf("ContainerName = %v, want %v", got.ContainerName, tt.want.ContainerName) + } + if got.Status != tt.want.Status { + t.Errorf("Status = %v, want %v", got.Status, tt.want.Status) + } + if got.CriticalCount != tt.want.CriticalCount { + t.Errorf("CriticalCount = %v, want %v", got.CriticalCount, tt.want.CriticalCount) + } + if got.HighCount != tt.want.HighCount { + t.Errorf("HighCount = %v, want %v", got.HighCount, tt.want.HighCount) + } + if got.MediumCount != tt.want.MediumCount { + t.Errorf("MediumCount = %v, want %v", got.MediumCount, tt.want.MediumCount) + } + if got.LowCount != tt.want.LowCount { + t.Errorf("LowCount = %v, want %v", got.LowCount, tt.want.LowCount) + } + if got.InfoCount != tt.want.InfoCount { + t.Errorf("InfoCount = %v, want %v", got.InfoCount, tt.want.InfoCount) + } + }) + } +} diff --git a/cli/cage/commands/flags.go b/cli/cage/cageapp/flags.go similarity index 97% rename from cli/cage/commands/flags.go rename to cli/cage/cageapp/flags.go index eb991bb..8e3f6e3 100644 --- a/cli/cage/commands/flags.go +++ b/cli/cage/cageapp/flags.go @@ -1,10 +1,15 @@ -package commands +package cageapp import ( "github.com/loilo-inc/canarycage/env" "github.com/urfave/cli/v2" ) +type App struct { + CI bool + NoColor bool +} + func RegionFlag(dest *string) *cli.StringFlag { return &cli.StringFlag{ Name: "region", diff --git a/cli/cage/cageapp/types.go b/cli/cage/cageapp/types.go new file mode 100644 index 0000000..c865550 --- /dev/null +++ b/cli/cage/cageapp/types.go @@ -0,0 +1,47 @@ +package cageapp + +import ( + "context" + "io" + + "github.com/loilo-inc/canarycage/env" + "github.com/loilo-inc/canarycage/types" +) + +type CageCmdInput struct { + *env.Envars + *App + Stdin io.Reader +} + +func NewCageCmdInput(stdin io.Reader, opts ...func(*CageCmdInput)) *CageCmdInput { + input := &CageCmdInput{ + Envars: &env.Envars{}, + App: &App{}, + Stdin: stdin, + } + for _, opt := range opts { + opt(input) + } + return input +} + +type CageCmdProvider = func(ctx context.Context, input *CageCmdInput) (types.Cage, error) + +type AuditCmdInput struct { + *App + Region string + Cluster string + Service string + LogDetail bool +} + +type AuditCmdProvider = func(ctx context.Context, input *AuditCmdInput) (types.Audit, error) + +func NewAuditCmdInput(opts ...func(*AuditCmdInput)) *AuditCmdInput { + input := &AuditCmdInput{App: &App{}} + for _, opt := range opts { + opt(input) + } + return input +} diff --git a/cli/cage/cageapp/types_test.go b/cli/cage/cageapp/types_test.go new file mode 100644 index 0000000..31e8a7d --- /dev/null +++ b/cli/cage/cageapp/types_test.go @@ -0,0 +1,43 @@ +package cageapp + +import ( + "strings" + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestNewCageCmdInput(t *testing.T) { + t.Run("basic", func(t *testing.T) { + input := NewCageCmdInput(strings.NewReader("test")) + assert := assert.New(t) + assert.NotNil(input.App) + assert.NotNil(input.Envars) + assert.NotNil(input.Stdin) + }) + t.Run("with options", func(t *testing.T) { + input := NewCageCmdInput(nil, func(c *CageCmdInput) { + c.Envars.Region = "us-west-2" + }) + assert := assert.New(t) + assert.NotNil(input.App) + assert.NotNil(input.Envars) + assert.Equal("us-west-2", input.Envars.Region) + }) +} + +func TestNewAuditCmdInput(t *testing.T) { + t.Run("basic", func(t *testing.T) { + input := NewAuditCmdInput() + assert := assert.New(t) + assert.NotNil(input.App) + }) + t.Run("with options", func(t *testing.T) { + input := NewAuditCmdInput(func(a *AuditCmdInput) { + a.Region = "us-west-2" + }) + assert := assert.New(t) + assert.NotNil(input.App) + assert.Equal("us-west-2", input.Region) + }) +} diff --git a/cli/cage/commands/audit.go b/cli/cage/commands/audit.go new file mode 100644 index 0000000..997affe --- /dev/null +++ b/cli/cage/commands/audit.go @@ -0,0 +1,57 @@ +package commands + +import ( + "errors" + + "github.com/loilo-inc/canarycage/cli/cage/cageapp" + "github.com/loilo-inc/canarycage/env" + "github.com/urfave/cli/v2" +) + +func Audit(app *cageapp.App, provider cageapp.AuditCmdProvider) *cli.Command { + input := cageapp.NewAuditCmdInput() + input.App = app + return &cli.Command{ + Name: "audit", + Usage: "Audit container images used in an ECS service", + ArgsUsage: "[directory path of service.json and task-definition.json]", + Flags: []cli.Flag{ + cageapp.RegionFlag(&input.Region), + cageapp.ClusterFlag(&input.Cluster), + cageapp.ServiceFlag(&input.Service), + &cli.BoolFlag{ + Name: "detail", + Usage: "By default, only the name and URI of the finding are logged.", + Value: false, + Destination: &input.LogDetail, + }, + }, + Action: func(ctx *cli.Context) error { + dir, _, err := RequireArgs(ctx, 0, 1) + if err != nil { + return err + } + if input.Region == "" { + return errors.New("--region flag is required") + } + if dir != "" { + srv, err := env.LoadServiceDefinition(dir) + if err != nil { + return err + } + if srv.ServiceName == nil || srv.Cluster == nil { + return errors.New("service.json must contain ServiceName and Cluster") + } + input.Service = *srv.ServiceName + input.Cluster = *srv.Cluster + } else if input.Cluster == "" || input.Service == "" { + return errors.New("either directory argument or both --cluster and --service flags must be provided") + } + cmd, err := provider(ctx.Context, input) + if err != nil { + return err + } + return cmd.Run(ctx.Context) + }, + } +} diff --git a/cli/cage/commands/audit_test.go b/cli/cage/commands/audit_test.go new file mode 100644 index 0000000..b5646c8 --- /dev/null +++ b/cli/cage/commands/audit_test.go @@ -0,0 +1,134 @@ +package commands + +import ( + "context" + "errors" + "testing" + + "github.com/loilo-inc/canarycage/cli/cage/cageapp" + "github.com/loilo-inc/canarycage/mocks/mock_types" + "github.com/loilo-inc/canarycage/types" + "github.com/stretchr/testify/assert" + "github.com/urfave/cli/v2" + "go.uber.org/mock/gomock" +) + +func TestAudit(t *testing.T) { + t.Run("returns error when region is missing", func(t *testing.T) { + app := setupAuditApp(t, nil) + err := app.Run([]string{"cage", "audit", "--region", ""}) + assert.Error(t, err) + assert.Contains(t, err.Error(), "--region flag is required") + }) + t.Run("return errors when too many arguments", func(t *testing.T) { + app := setupAuditApp(t, nil) + err := app.Run([]string{"cage", "audit", "--region", "us-east-1", "arg1", "arg2"}) + assert.Error(t, err) + assert.Contains(t, err.Error(), "invalid number of arguments. expected at most 1") + }) + t.Run("returns error when both directory and flags are missing", func(t *testing.T) { + app := setupAuditApp(t, nil) + + err := app.Run([]string{"cage", "audit", "--region", "us-east-1"}) + assert.Error(t, err) + assert.Contains(t, err.Error(), "either directory argument or both --cluster and --service flags must be provided") + }) + + t.Run("returns error when only cluster flag is provided", func(t *testing.T) { + app := setupAuditApp(t, nil) + + err := app.Run([]string{"cage", "audit", "--region", "us-east-1", "--cluster", "test-cluster"}) + assert.Error(t, err) + assert.Contains(t, err.Error(), "either directory argument or both --cluster and --service flags must be provided") + }) + + t.Run("returns error when only service flag is provided", func(t *testing.T) { + app := setupAuditApp(t, nil) + + err := app.Run([]string{"cage", "audit", "--region", "us-east-1", "--service", "test-service"}) + assert.Error(t, err) + assert.Contains(t, err.Error(), "either directory argument or both --cluster and --service flags must be provided") + }) + + t.Run("returns error when diProvider fails", func(t *testing.T) { + expectedErr := errors.New("di provider error") + app := setupAuditApp(t, func(ctx context.Context, input *cageapp.AuditCmdInput) (types.Audit, error) { + assert.Equal(t, "us-east-1", input.Region) + return nil, expectedErr + }) + + err := app.Run([]string{ + "cage", "audit", "--region", "us-east-1", "--cluster", "test-cluster", "--service", "test-service", + }) + assert.Error(t, err) + assert.Equal(t, expectedErr, err) + }) + setupBase := func(t *testing.T) (*cli.App, *mock_types.MockAudit) { + t.Helper() + ctrl := gomock.NewController(t) + mockAudit := mock_types.NewMockAudit(ctrl) + + app := setupAuditApp(t, func(ctx context.Context, input *cageapp.AuditCmdInput) (types.Audit, error) { + assert.Equal(t, "us-east-1", input.Region) + return mockAudit, nil + }) + return app, mockAudit + } + t.Run("Succcess", func(t *testing.T) { + setup := func(t *testing.T) *cli.App { + t.Helper() + app, mockAudit := setupBase(t) + mockAudit.EXPECT(). + Run(gomock.Any()). + Return(nil) + return app + } + t.Run("executes scan with directory argument", func(t *testing.T) { + app := setup(t) + err := app.Run([]string{"cage", "audit", + "--region", "us-east-1", "../../../fixtures"}) + assert.NoError(t, err) + }) + t.Run("executes scan with flags", func(t *testing.T) { + app := setup(t) + err := app.Run([]string{"cage", "audit", + "--region", "us-east-1", + "--cluster", "cluster", + "--service", "service"}) + assert.NoError(t, err) + }) + }) + t.Run("Error", func(t *testing.T) { + t.Run("error on scanner.Scan()", func(t *testing.T) { + app, mockAudit := setupBase(t) + mockAudit.EXPECT(). + Run(gomock.Any()). + Return(errors.New("scan error")) + + err := app.Run([]string{"cage", "audit", + "--region", "us-east-1", + "--cluster", "cluster", + "--service", "service"}) + assert.Error(t, err) + assert.Contains(t, err.Error(), "scan error") + }) + t.Run("error on loading service definition", func(t *testing.T) { + app := setupAuditApp(t, nil) + err := app.Run([]string{"cage", "audit", + "--region", "us-east-1", "../../../fixtures/invalid-service"}) + assert.Error(t, err) + assert.Contains(t, err.Error(), "no 'service.json' found") + }) + }) +} + +func setupAuditApp(t *testing.T, provider cageapp.AuditCmdProvider) *cli.App { + t.Helper() + conf := &cageapp.App{} + app := cli.NewApp() + app.Name = "cage" + app.Commands = []*cli.Command{ + Audit(conf, provider), + } + return app +} diff --git a/cli/cage/commands/command.go b/cli/cage/commands/command.go index c8c162e..bdf26a5 100644 --- a/cli/cage/commands/command.go +++ b/cli/cage/commands/command.go @@ -1,10 +1,10 @@ package commands import ( - "io" + "context" "github.com/aws/aws-sdk-go-v2/service/ecs" - "github.com/loilo-inc/canarycage/cli/cage/prompt" + "github.com/loilo-inc/canarycage/cli/cage/cageapp" "github.com/loilo-inc/canarycage/env" "github.com/loilo-inc/canarycage/types" "github.com/urfave/cli/v2" @@ -12,23 +12,17 @@ import ( ) type CageCommands struct { - Prompt *prompt.Prompter - cageCliProvier cageCliProvier + cageCliProvider cageapp.CageCmdProvider } func NewCageCommands( - stdin io.Reader, - cageCliProvier cageCliProvier, + cageCliProvider cageapp.CageCmdProvider, ) *CageCommands { - return &CageCommands{ - Prompt: prompt.NewPrompter(stdin), - cageCliProvier: cageCliProvier, - } + cmds := &CageCommands{cageCliProvider: cageCliProvider} + return cmds } -type cageCliProvier = func(envars *env.Envars) (types.Cage, error) - -func (c *CageCommands) requireArgs( +func RequireArgs( ctx *cli.Context, minArgs int, maxArgs int, @@ -44,7 +38,7 @@ func (c *CageCommands) requireArgs( } func (c *CageCommands) setupCage( - envars *env.Envars, + input *cageapp.CageCmdInput, dir string, ) (types.Cage, error) { var service *ecs.CreateServiceInput @@ -54,23 +48,23 @@ func (c *CageCommands) setupCage( } else { service = srv } - if envars.TaskDefinitionArn == "" { + if input.TaskDefinitionArn == "" { if td, err := env.LoadTaskDefinition(dir); err != nil { return nil, err } else { taskDefinition = td } } - env.MergeEnvars(envars, &env.Envars{ + env.MergeEnvars(input.Envars, &env.Envars{ Cluster: *service.Cluster, Service: *service.ServiceName, TaskDefinitionInput: taskDefinition, ServiceDefinitionInput: service, }) - if err := env.EnsureEnvars(envars); err != nil { + if err := env.EnsureEnvars(input.Envars); err != nil { return nil, err } - cagecli, err := c.cageCliProvier(envars) + cagecli, err := c.cageCliProvider(context.TODO(), input) if err != nil { return nil, err } diff --git a/cli/cage/commands/command_test.go b/cli/cage/commands/command_test.go index 346b084..d67180c 100644 --- a/cli/cage/commands/command_test.go +++ b/cli/cage/commands/command_test.go @@ -1,148 +1,60 @@ package commands import ( - "fmt" - "strings" + "context" "testing" - "github.com/loilo-inc/canarycage/env" + "github.com/loilo-inc/canarycage/cli/cage/cageapp" "github.com/loilo-inc/canarycage/mocks/mock_types" "github.com/loilo-inc/canarycage/test" "github.com/loilo-inc/canarycage/types" "github.com/stretchr/testify/assert" - "github.com/urfave/cli/v2" "go.uber.org/mock/gomock" ) -func TestCommands(t *testing.T) { - region := "ap-notheast-1" - cluster := "cluster" - service := "service" - stdinService := fmt.Sprintf("%s\n%s\n%s\n%s\n", region, cluster, service, "yes") - stdinTask := fmt.Sprintf("%s\n%s\n%s\n", region, cluster, "yes") - setup := func(t *testing.T, input string) (*cli.App, *mock_types.MockCage) { - ctrl := gomock.NewController(t) - stdin := strings.NewReader(input) - cagecli := mock_types.NewMockCage(ctrl) - app := cli.NewApp() - cmds := NewCageCommands(stdin, func(envars *env.Envars) (types.Cage, error) { - return cagecli, nil - }) - envars := env.Envars{CI: input == ""} - app.Commands = []*cli.Command{ - cmds.Up(&envars), - cmds.RollOut(&envars), - cmds.Run(&envars), - } - return app, cagecli - } - t.Run("rollout", func(t *testing.T) { - t.Run("basic", func(t *testing.T) { - app, cagecli := setup(t, stdinService) - cagecli.EXPECT().RollOut(gomock.Any(), &types.RollOutInput{}).Return(&types.RollOutResult{}, nil) - err := app.Run([]string{"cage", "rollout", "--region", "ap-notheast-1", "../../../fixtures"}) - assert.NoError(t, err) - }) - t.Run("basic/ci", func(t *testing.T) { - app, cagecli := setup(t, "") - cagecli.EXPECT().RollOut(gomock.Any(), &types.RollOutInput{}).Return(&types.RollOutResult{}, nil) - err := app.Run([]string{"cage", "rollout", "--region", "ap-notheast-1", "../../../fixtures"}) - assert.NoError(t, err) - }) - t.Run("basic/udate-service", func(t *testing.T) { - app, cagecli := setup(t, stdinService) - cagecli.EXPECT().RollOut(gomock.Any(), &types.RollOutInput{UpdateService: true}).Return(&types.RollOutResult{}, nil) - err := app.Run([]string{"cage", "rollout", "--region", "ap-notheast-1", "--updateService", "../../../fixtures"}) - assert.NoError(t, err) - }) - t.Run("error", func(t *testing.T) { - app, cagecli := setup(t, stdinService) - cagecli.EXPECT().RollOut(gomock.Any(), &types.RollOutInput{}).Return(&types.RollOutResult{}, fmt.Errorf("error")) - err := app.Run([]string{"cage", "rollout", "--region", "ap-notheast-1", "../../../fixtures"}) - assert.EqualError(t, err, "error") - }) - }) - t.Run("up", func(t *testing.T) { - t.Run("basic", func(t *testing.T) { - app, cagecli := setup(t, stdinService) - cagecli.EXPECT().Up(gomock.Any()).Return(&types.UpResult{}, nil) - err := app.Run([]string{"cage", "up", "--region", "ap-notheast-1", "../../../fixtures"}) - assert.NoError(t, err) - }) - t.Run("basic/ci", func(t *testing.T) { - app, cagecli := setup(t, "") - cagecli.EXPECT().Up(gomock.Any()).Return(&types.UpResult{}, nil) - err := app.Run([]string{"cage", "up", "--region", "ap-notheast-1", "../../../fixtures"}) - assert.NoError(t, err) - }) - t.Run("error", func(t *testing.T) { - app, cagecli := setup(t, stdinService) - cagecli.EXPECT().Up(gomock.Any()).Return(nil, fmt.Errorf("error")) - err := app.Run([]string{"cage", "up", "--region", "ap-notheast-1", "../../../fixtures"}) - assert.EqualError(t, err, "error") - }) - }) - t.Run("run", func(t *testing.T) { - t.Run("basic", func(t *testing.T) { - app, cagecli := setup(t, stdinTask) - cagecli.EXPECT().Run(gomock.Any(), gomock.Any()).Return(&types.RunResult{}, nil) - err := app.Run([]string{"cage", "run", "--region", "ap-notheast-1", "../../../fixtures", "container", "exec"}) - assert.NoError(t, err) - }) - t.Run("basic/ci", func(t *testing.T) { - app, cagecli := setup(t, "") - cagecli.EXPECT().Run(gomock.Any(), gomock.Any()).Return(&types.RunResult{}, nil) - err := app.Run([]string{"cage", "run", "--region", "ap-notheast-1", "../../../fixtures", "container", "exec"}) - assert.NoError(t, err) - }) - t.Run("error", func(t *testing.T) { - app, cagecli := setup(t, stdinTask) - cagecli.EXPECT().Run(gomock.Any(), gomock.Any()).Return(nil, fmt.Errorf("error")) - err := app.Run([]string{"cage", "run", "--region", "ap-notheast-1", "../../../fixtures", "container", "exec"}) - assert.EqualError(t, err, "error") - }) - }) -} - func TestSetupCage(t *testing.T) { t.Run("basic", func(t *testing.T) { - envars := &env.Envars{Region: "us-west-2"} + input := cageapp.NewCageCmdInput(nil) + input.Region = "us-west-2" cageCli := mock_types.NewMockCage(gomock.NewController(t)) - cmd := NewCageCommands(nil, func(envars *env.Envars) (types.Cage, error) { + cmd := NewCageCommands(func(ctx context.Context, envars *cageapp.CageCmdInput) (types.Cage, error) { return cageCli, nil }) - v, err := cmd.setupCage(envars, "../../../fixtures") + v, err := cmd.setupCage(input, "../../../fixtures") if err != nil { t.Fatal(err) } assert.Equal(t, v, cageCli) - assert.Equal(t, envars.Service, "service") - assert.Equal(t, envars.Cluster, "cluster") - assert.NotNil(t, envars.ServiceDefinitionInput) - assert.NotNil(t, envars.TaskDefinitionInput) + assert.Equal(t, input.Service, "service") + assert.Equal(t, input.Cluster, "cluster") + assert.NotNil(t, input.ServiceDefinitionInput) + assert.NotNil(t, input.TaskDefinitionInput) }) t.Run("should skip load task definition if --taskDefinitionArn provided", func(t *testing.T) { - envars := &env.Envars{Region: "us-west-2", TaskDefinitionArn: "arn"} + input := cageapp.NewCageCmdInput(nil) + input.Region = "us-west-2" + input.TaskDefinitionArn = "arn" cageCli := mock_types.NewMockCage(gomock.NewController(t)) - cmd := NewCageCommands(nil, func(envars *env.Envars) (types.Cage, error) { + cmd := NewCageCommands(func(ctx context.Context, input *cageapp.CageCmdInput) (types.Cage, error) { return cageCli, nil }) - v, err := cmd.setupCage(envars, "../../../fixtures") + v, err := cmd.setupCage(input, "../../../fixtures") if err != nil { t.Fatal(err) } assert.Equal(t, v, cageCli) - assert.Equal(t, envars.Service, "service") - assert.Equal(t, envars.Cluster, "cluster") - assert.NotNil(t, envars.ServiceDefinitionInput) - assert.Nil(t, envars.TaskDefinitionInput) + assert.Equal(t, input.Service, "service") + assert.Equal(t, input.Cluster, "cluster") + assert.NotNil(t, input.ServiceDefinitionInput) + assert.Nil(t, input.TaskDefinitionInput) }) t.Run("should error if error returned from NewCage", func(t *testing.T) { - envars := &env.Envars{Region: "us-west-2"} - cmd := NewCageCommands(nil, func(envars *env.Envars) (types.Cage, error) { + input := cageapp.NewCageCmdInput(nil) + input.Region = "us-west-2" + cmd := NewCageCommands(func(ctx context.Context, input *cageapp.CageCmdInput) (types.Cage, error) { return nil, test.Err }) - _, err := cmd.setupCage(envars, "../../../fixtures") + _, err := cmd.setupCage(input, "../../../fixtures") assert.EqualError(t, err, "error") }) } diff --git a/cli/cage/commands/provider.go b/cli/cage/commands/provider.go new file mode 100644 index 0000000..9b2fdab --- /dev/null +++ b/cli/cage/commands/provider.go @@ -0,0 +1,37 @@ +package commands + +import ( + "context" + + "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/service/ec2" + "github.com/aws/aws-sdk-go-v2/service/ecr" + "github.com/aws/aws-sdk-go-v2/service/ecs" + "github.com/aws/aws-sdk-go-v2/service/elasticloadbalancingv2" + cage "github.com/loilo-inc/canarycage" + "github.com/loilo-inc/canarycage/awsiface" + "github.com/loilo-inc/canarycage/cli/cage/cageapp" + "github.com/loilo-inc/canarycage/key" + "github.com/loilo-inc/canarycage/task" + "github.com/loilo-inc/canarycage/timeout" + "github.com/loilo-inc/canarycage/types" + "github.com/loilo-inc/logos/di" +) + +func ProvideCageCli(ctx context.Context, input *cageapp.CageCmdInput) (types.Cage, error) { + conf := awsiface.MustLoadConfig( + ctx, + config.WithRegion(input.Region), + ) + d := di.NewDomain(func(b *di.B) { + b.Set(key.Env, input.Envars) + b.Set(key.EcsCli, ecs.NewFromConfig(conf)) + b.Set(key.EcrCli, ecr.NewFromConfig(conf)) + b.Set(key.Ec2Cli, ec2.NewFromConfig(conf)) + b.Set(key.AlbCli, elasticloadbalancingv2.NewFromConfig(conf)) + b.Set(key.TaskFactory, task.NewFactory(b.Future())) + b.Set(key.Time, &timeout.Time{}) + }) + cagecli := cage.NewCage(d) + return cagecli, nil +} diff --git a/cli/cage/commands/provider_test.go b/cli/cage/commands/provider_test.go new file mode 100644 index 0000000..0867a4a --- /dev/null +++ b/cli/cage/commands/provider_test.go @@ -0,0 +1,45 @@ +package commands + +import ( + "context" + "testing" + + "github.com/loilo-inc/canarycage/cli/cage/cageapp" + "github.com/stretchr/testify/assert" +) + +func TestProvideCageCli(t *testing.T) { + t.Run("successfully creates cage cli with valid region", func(t *testing.T) { + input := cageapp.NewCageCmdInput(nil) + input.Envars.Region = "us-west-2" + + cage, err := ProvideCageCli(context.TODO(), input) + assert.NoError(t, err) + assert.NotNil(t, cage) + }) + + t.Run("returns error with invalid region", func(t *testing.T) { + input := cageapp.NewCageCmdInput(nil) + input.Envars.Region = "" + + cage, err := ProvideCageCli(context.TODO(), input) + if err != nil { + assert.Nil(t, cage, "expected cage to be nil when error occurs") + return + } + assert.NotNil(t, cage, "expected cage to be non-nil when no error") + }) + + t.Run("handles nil envars", func(t *testing.T) { + defer func() { + if r := recover(); r != nil { + return + } + }() + + cage, err := ProvideCageCli(context.TODO(), nil) + if err == nil { + assert.NotNil(t, cage, "expected cage to be non-nil when no error") + } + }) +} diff --git a/cli/cage/commands/rollout.go b/cli/cage/commands/rollout.go index 0b2151c..d4c4c9b 100644 --- a/cli/cage/commands/rollout.go +++ b/cli/cage/commands/rollout.go @@ -4,14 +4,14 @@ import ( "context" "github.com/apex/log" + "github.com/loilo-inc/canarycage/cli/cage/cageapp" + "github.com/loilo-inc/canarycage/cli/cage/prompt" "github.com/loilo-inc/canarycage/env" "github.com/loilo-inc/canarycage/types" "github.com/urfave/cli/v2" ) -func (c *CageCommands) RollOut( - envars *env.Envars, -) *cli.Command { +func (c *CageCommands) RollOut(input *cageapp.CageCmdInput) *cli.Command { var updateServiceConf bool return &cli.Command{ Name: "rollout", @@ -20,16 +20,16 @@ func (c *CageCommands) RollOut( Args: true, ArgsUsage: "[directory path of service.json and task-definition.json]", Flags: []cli.Flag{ - RegionFlag(&envars.Region), - ClusterFlag(&envars.Cluster), - ServiceFlag(&envars.Service), - TaskDefinitionArnFlag(&envars.TaskDefinitionArn), - CanaryTaskIdleDurationFlag(&envars.CanaryTaskIdleDuration), + cageapp.RegionFlag(&input.Region), + cageapp.ClusterFlag(&input.Cluster), + cageapp.ServiceFlag(&input.Service), + cageapp.TaskDefinitionArnFlag(&input.TaskDefinitionArn), + cageapp.CanaryTaskIdleDurationFlag(&input.CanaryTaskIdleDuration), &cli.StringFlag{ Name: "canaryInstanceArn", EnvVars: []string{env.CanaryInstanceArnKey}, Usage: "EC2 instance ARN for placing canary task. required only when LaunchType is EC2", - Destination: &envars.CanaryInstanceArn, + Destination: &input.CanaryInstanceArn, }, &cli.BoolFlag{ Name: "updateService", @@ -37,29 +37,32 @@ func (c *CageCommands) RollOut( Usage: "Update service configurations except for task definiton. Default is false.", Destination: &updateServiceConf, }, - TaskRunningWaitFlag(&envars.CanaryTaskRunningWait), - TaskHealthCheckWaitFlag(&envars.CanaryTaskHealthCheckWait), - TaskStoppedWaitFlag(&envars.CanaryTaskStoppedWait), - ServiceStableWaitFlag(&envars.ServiceStableWait), + cageapp.TaskRunningWaitFlag(&input.CanaryTaskRunningWait), + cageapp.TaskHealthCheckWaitFlag(&input.CanaryTaskHealthCheckWait), + cageapp.TaskStoppedWaitFlag(&input.CanaryTaskStoppedWait), + cageapp.ServiceStableWaitFlag(&input.ServiceStableWait), }, Action: func(ctx *cli.Context) error { - dir, _, err := c.requireArgs(ctx, 1, 1) + dir, _, err := RequireArgs(ctx, 1, 1) if err != nil { return err } - cagecli, err := c.setupCage(envars, dir) + cagecli, err := c.setupCage(input, dir) if err != nil { return err } - if err := c.Prompt.ConfirmService(envars); err != nil { - return err + if !input.CI { + prompter := prompt.NewPrompter(input.Stdin) + if err := prompter.ConfirmService(input.Envars); err != nil { + return err + } } result, err := cagecli.RollOut(context.Background(), &types.RollOutInput{UpdateService: updateServiceConf}) if err != nil { if !result.ServiceUpdated { - log.Errorf("🤕 failed to roll out new tasks but service '%s' is not changed", envars.Service) + log.Errorf("🤕 failed to roll out new tasks but service '%s' is not changed", input.Service) } else { - log.Errorf("😭 failed to roll out new tasks and service '%s' might be changed. CHECK ECS CONSOLE NOW!", envars.Service) + log.Errorf("😭 failed to roll out new tasks and service '%s' might be changed. CHECK ECS CONSOLE NOW!", input.Service) } return err } diff --git a/cli/cage/commands/rollout_test.go b/cli/cage/commands/rollout_test.go new file mode 100644 index 0000000..4201f4f --- /dev/null +++ b/cli/cage/commands/rollout_test.go @@ -0,0 +1,48 @@ +package commands_test + +import ( + "fmt" + "strings" + "testing" + + "github.com/loilo-inc/canarycage/types" + "github.com/stretchr/testify/assert" + "go.uber.org/mock/gomock" +) + +func TestRollOut(t *testing.T) { + t.Run("basic", func(t *testing.T) { + app, cagecli := setup(t, strings.NewReader(stdinService)) + cagecli.EXPECT().RollOut(gomock.Any(), &types.RollOutInput{}).Return(&types.RollOutResult{}, nil) + err := app.Run([]string{"cage", "rollout", "--region", "ap-notheast-1", "../../../fixtures"}) + assert.NoError(t, err) + }) + t.Run("basic/ci", func(t *testing.T) { + app, cagecli := setup(t, nil) + cagecli.EXPECT().RollOut(gomock.Any(), &types.RollOutInput{}).Return(&types.RollOutResult{}, nil) + err := app.Run([]string{"cage", "--ci", "rollout", "--region", "ap-notheast-1", "../../../fixtures"}) + assert.NoError(t, err) + }) + t.Run("basic/update-service", func(t *testing.T) { + app, cagecli := setup(t, strings.NewReader(stdinService)) + cagecli.EXPECT().RollOut(gomock.Any(), &types.RollOutInput{UpdateService: true}).Return(&types.RollOutResult{}, nil) + err := app.Run([]string{"cage", "rollout", "--region", "ap-notheast-1", "--updateService", "../../../fixtures"}) + assert.NoError(t, err) + }) + t.Run("missing args", func(t *testing.T) { + app, _ := setup(t, nil) + err := app.Run([]string{"cage", "rollout", "--region", "ap-notheast-1"}) + assert.EqualError(t, err, "invalid number of arguments. expected at least 1") + }) + t.Run("reading stdin error", func(t *testing.T) { + app, _ := setup(t, &errorReader{}) + err := app.Run([]string{"cage", "rollout", "--region", "ap-notheast-1", "../../../fixtures"}) + assert.EqualError(t, err, "failed to read from stdin: EOF") + }) + t.Run("error", func(t *testing.T) { + app, cagecli := setup(t, strings.NewReader(stdinService)) + cagecli.EXPECT().RollOut(gomock.Any(), &types.RollOutInput{}).Return(&types.RollOutResult{}, fmt.Errorf("error")) + err := app.Run([]string{"cage", "rollout", "--region", "ap-notheast-1", "../../../fixtures"}) + assert.EqualError(t, err, "error") + }) +} diff --git a/cli/cage/commands/run.go b/cli/cage/commands/run.go index c830f5b..7d4377e 100644 --- a/cli/cage/commands/run.go +++ b/cli/cage/commands/run.go @@ -5,14 +5,13 @@ import ( "github.com/apex/log" ecstypes "github.com/aws/aws-sdk-go-v2/service/ecs/types" - "github.com/loilo-inc/canarycage/env" + "github.com/loilo-inc/canarycage/cli/cage/cageapp" + "github.com/loilo-inc/canarycage/cli/cage/prompt" "github.com/loilo-inc/canarycage/types" "github.com/urfave/cli/v2" ) -func (c *CageCommands) Run( - envars *env.Envars, -) *cli.Command { +func (c *CageCommands) Run(input *cageapp.CageCmdInput) *cli.Command { return &cli.Command{ Name: "run", Usage: "run task with specified task definition", @@ -20,22 +19,25 @@ func (c *CageCommands) Run( Args: true, ArgsUsage: " ...", Flags: []cli.Flag{ - RegionFlag(&envars.Region), - ClusterFlag(&envars.Cluster), - TaskRunningWaitFlag(&envars.CanaryTaskRunningWait), - TaskStoppedWaitFlag(&envars.CanaryTaskStoppedWait), + cageapp.RegionFlag(&input.Region), + cageapp.ClusterFlag(&input.Cluster), + cageapp.TaskRunningWaitFlag(&input.CanaryTaskRunningWait), + cageapp.TaskStoppedWaitFlag(&input.CanaryTaskStoppedWait), }, Action: func(ctx *cli.Context) error { - dir, rest, err := c.requireArgs(ctx, 3, 100) + dir, rest, err := RequireArgs(ctx, 3, 100) if err != nil { return err } - cagecli, err := c.setupCage(envars, dir) + cagecli, err := c.setupCage(input, dir) if err != nil { return err } - if err := c.Prompt.ConfirmTask(envars); err != nil { - return err + if !input.CI { + prompter := prompt.NewPrompter(input.Stdin) + if err := prompter.ConfirmTask(input.Envars); err != nil { + return err + } } container := rest[0] commands := rest[1:] diff --git a/cli/cage/commands/run_test.go b/cli/cage/commands/run_test.go new file mode 100644 index 0000000..0ab397c --- /dev/null +++ b/cli/cage/commands/run_test.go @@ -0,0 +1,42 @@ +package commands_test + +import ( + "fmt" + "strings" + "testing" + + "github.com/loilo-inc/canarycage/types" + "github.com/stretchr/testify/assert" + "go.uber.org/mock/gomock" +) + +func TestRun(t *testing.T) { + t.Run("basic", func(t *testing.T) { + app, cagecli := setup(t, strings.NewReader(stdinTask)) + cagecli.EXPECT().Run(gomock.Any(), gomock.Any()).Return(&types.RunResult{}, nil) + err := app.Run([]string{"cage", "run", "--region", "ap-notheast-1", "../../../fixtures", "container", "exec"}) + assert.NoError(t, err) + }) + t.Run("basic/ci", func(t *testing.T) { + app, cagecli := setup(t, nil) + cagecli.EXPECT().Run(gomock.Any(), gomock.Any()).Return(&types.RunResult{}, nil) + err := app.Run([]string{"cage", "--ci", "run", "--region", "ap-notheast-1", "../../../fixtures", "container", "exec"}) + assert.NoError(t, err) + }) + t.Run("missing args", func(t *testing.T) { + app, _ := setup(t, nil) + err := app.Run([]string{"cage", "run", "--region", "ap-notheast-1"}) + assert.EqualError(t, err, "invalid number of arguments. expected at least 3") + }) + t.Run("reading stdin error", func(t *testing.T) { + app, _ := setup(t, &errorReader{}) + err := app.Run([]string{"cage", "run", "--region", "ap-notheast-1", "../../../fixtures", "container", "exec"}) + assert.EqualError(t, err, "failed to read from stdin: EOF") + }) + t.Run("error", func(t *testing.T) { + app, cagecli := setup(t, strings.NewReader(stdinTask)) + cagecli.EXPECT().Run(gomock.Any(), gomock.Any()).Return(nil, fmt.Errorf("error")) + err := app.Run([]string{"cage", "run", "--region", "ap-notheast-1", "../../../fixtures", "container", "exec"}) + assert.EqualError(t, err, "error") + }) +} diff --git a/cli/cage/commands/tools_test.go b/cli/cage/commands/tools_test.go new file mode 100644 index 0000000..eaacd3a --- /dev/null +++ b/cli/cage/commands/tools_test.go @@ -0,0 +1,46 @@ +package commands_test + +import ( + "context" + "io" + "testing" + + "github.com/loilo-inc/canarycage/cli/cage/cageapp" + "github.com/loilo-inc/canarycage/cli/cage/commands" + "github.com/loilo-inc/canarycage/mocks/mock_types" + "github.com/loilo-inc/canarycage/types" + "github.com/urfave/cli/v2" + "go.uber.org/mock/gomock" +) + +var stdinService = "ap-notheast-1\ncluster\nservice\nyes\n" +var stdinTask = "ap-notheast-1\ncluster\nyes\n" + +func setup(t *testing.T, stdin io.Reader) (*cli.App, *mock_types.MockCage) { + ctrl := gomock.NewController(t) + cagecli := mock_types.NewMockCage(ctrl) + input := cageapp.NewCageCmdInput(stdin) + app := cli.NewApp() + cmds := commands.NewCageCommands(func(ctx context.Context, input *cageapp.CageCmdInput) (types.Cage, error) { + return cagecli, nil + }) + app.Commands = []*cli.Command{ + cmds.Up(input), + cmds.RollOut(input), + cmds.Run(input), + } + app.Flags = []cli.Flag{ + &cli.BoolFlag{ + Name: "ci", + Destination: &input.CI, + Value: false, + }, + } + return app, cagecli +} + +type errorReader struct{} + +func (e *errorReader) Read(p []byte) (n int, err error) { + return 0, io.EOF +} diff --git a/cli/cage/commands/up.go b/cli/cage/commands/up.go index 3ecd906..b1a7379 100644 --- a/cli/cage/commands/up.go +++ b/cli/cage/commands/up.go @@ -3,13 +3,12 @@ package commands import ( "context" - "github.com/loilo-inc/canarycage/env" + "github.com/loilo-inc/canarycage/cli/cage/cageapp" + "github.com/loilo-inc/canarycage/cli/cage/prompt" "github.com/urfave/cli/v2" ) -func (c *CageCommands) Up( - envars *env.Envars, -) *cli.Command { +func (c *CageCommands) Up(input *cageapp.CageCmdInput) *cli.Command { return &cli.Command{ Name: "up", Usage: "create new ECS service with specified task definition", @@ -17,24 +16,27 @@ func (c *CageCommands) Up( Args: true, ArgsUsage: "[directory path of service.json and task-definition.json]", Flags: []cli.Flag{ - RegionFlag(&envars.Region), - ClusterFlag(&envars.Cluster), - ServiceFlag(&envars.Service), - TaskDefinitionArnFlag(&envars.TaskDefinitionArn), - CanaryTaskIdleDurationFlag(&envars.CanaryTaskIdleDuration), - ServiceStableWaitFlag(&envars.ServiceStableWait), + cageapp.RegionFlag(&input.Region), + cageapp.ClusterFlag(&input.Cluster), + cageapp.ServiceFlag(&input.Service), + cageapp.TaskDefinitionArnFlag(&input.TaskDefinitionArn), + cageapp.CanaryTaskIdleDurationFlag(&input.CanaryTaskIdleDuration), + cageapp.ServiceStableWaitFlag(&input.ServiceStableWait), }, Action: func(ctx *cli.Context) error { - dir, _, err := c.requireArgs(ctx, 1, 1) + dir, _, err := RequireArgs(ctx, 1, 1) if err != nil { return err } - cagecli, err := c.setupCage(envars, dir) + cagecli, err := c.setupCage(input, dir) if err != nil { return err } - if err := c.Prompt.ConfirmService(envars); err != nil { - return err + if !input.CI { + prompter := prompt.NewPrompter(input.Stdin) + if err := prompter.ConfirmService(input.Envars); err != nil { + return err + } } _, err = cagecli.Up(context.Background()) return err diff --git a/cli/cage/commands/up_test.go b/cli/cage/commands/up_test.go new file mode 100644 index 0000000..f6e81d4 --- /dev/null +++ b/cli/cage/commands/up_test.go @@ -0,0 +1,42 @@ +package commands_test + +import ( + "fmt" + "strings" + "testing" + + "github.com/loilo-inc/canarycage/types" + "github.com/stretchr/testify/assert" + "go.uber.org/mock/gomock" +) + +func TestUp(t *testing.T) { + t.Run("basic", func(t *testing.T) { + app, cagecli := setup(t, strings.NewReader(stdinService)) + cagecli.EXPECT().Up(gomock.Any()).Return(&types.UpResult{}, nil) + err := app.Run([]string{"cage", "up", "--region", "ap-notheast-1", "../../../fixtures"}) + assert.NoError(t, err) + }) + t.Run("basic/ci", func(t *testing.T) { + app, cagecli := setup(t, nil) + cagecli.EXPECT().Up(gomock.Any()).Return(&types.UpResult{}, nil) + err := app.Run([]string{"cage", "--ci", "up", "--region", "ap-notheast-1", "../../../fixtures"}) + assert.NoError(t, err) + }) + t.Run("missing args", func(t *testing.T) { + app, _ := setup(t, nil) + err := app.Run([]string{"cage", "up", "--region", "ap-notheast-1"}) + assert.EqualError(t, err, "invalid number of arguments. expected at least 1") + }) + t.Run("reading stdin error", func(t *testing.T) { + app, _ := setup(t, &errorReader{}) + err := app.Run([]string{"cage", "up", "--region", "ap-notheast-1", "../../../fixtures"}) + assert.EqualError(t, err, "failed to read from stdin: EOF") + }) + t.Run("error", func(t *testing.T) { + app, cagecli := setup(t, strings.NewReader(stdinService)) + cagecli.EXPECT().Up(gomock.Any()).Return(nil, fmt.Errorf("error")) + err := app.Run([]string{"cage", "up", "--region", "ap-notheast-1", "../../../fixtures"}) + assert.EqualError(t, err, "error") + }) +} diff --git a/cli/cage/commands/upgrade.go b/cli/cage/commands/upgrade.go index f590226..08aa3ba 100644 --- a/cli/cage/commands/upgrade.go +++ b/cli/cage/commands/upgrade.go @@ -5,9 +5,7 @@ import ( "github.com/urfave/cli/v2" ) -func (c *CageCommands) Upgrade( - upgrader upgrade.Upgrader, -) *cli.Command { +func Upgrade(upgrader upgrade.Upgrader) *cli.Command { var preRelease bool return &cli.Command{ Name: "upgrade", diff --git a/cli/cage/commands/upgrade_test.go b/cli/cage/commands/upgrade_test.go index e7122a8..5921891 100644 --- a/cli/cage/commands/upgrade_test.go +++ b/cli/cage/commands/upgrade_test.go @@ -19,9 +19,8 @@ func TestUpgrade(t *testing.T) { u.EXPECT().Upgrade( gomock.Eq(&upgrade.Input{}), ).Return(nil) - cmds := commands.NewCageCommands(nil, nil) app.Commands = []*cli.Command{ - cmds.Upgrade(u), + commands.Upgrade(u), } err := app.Run([]string{"cage", "upgrade"}) assert.NoError(t, err) @@ -33,9 +32,8 @@ func TestUpgrade(t *testing.T) { u.EXPECT().Upgrade( gomock.Eq(&upgrade.Input{PreRelease: true}), ).Return(nil) - cmds := commands.NewCageCommands(nil, nil) app.Commands = []*cli.Command{ - cmds.Upgrade(u), + commands.Upgrade(u), } err := app.Run([]string{"cage", "upgrade", "--pre-release"}) assert.NoError(t, err) diff --git a/cli/cage/main.go b/cli/cage/main.go index 5dcb4e4..bba8a47 100644 --- a/cli/cage/main.go +++ b/cli/cage/main.go @@ -1,26 +1,15 @@ package main import ( - "context" "fmt" "log" "os" - "github.com/aws/aws-sdk-go-v2/config" - "github.com/aws/aws-sdk-go-v2/service/ec2" - "github.com/aws/aws-sdk-go-v2/service/ecs" - "github.com/aws/aws-sdk-go-v2/service/elasticloadbalancingv2" - cage "github.com/loilo-inc/canarycage" + "github.com/loilo-inc/canarycage/cli/cage/audit" + "github.com/loilo-inc/canarycage/cli/cage/cageapp" "github.com/loilo-inc/canarycage/cli/cage/commands" "github.com/loilo-inc/canarycage/cli/cage/upgrade" - "github.com/loilo-inc/canarycage/env" - "github.com/loilo-inc/canarycage/key" - "github.com/loilo-inc/canarycage/task" - "github.com/loilo-inc/canarycage/timeout" - "github.com/loilo-inc/canarycage/types" - "github.com/loilo-inc/logos/di" "github.com/urfave/cli/v2" - "golang.org/x/xerrors" ) // set by goreleaser @@ -31,48 +20,39 @@ var ( ) func main() { + appConf := &cageapp.App{} + configCmdInput := func(input *cageapp.CageCmdInput) { + input.App = appConf + } app := cli.NewApp() app.Name = "canarycage" app.HelpName = "cage" app.Version = fmt.Sprintf("%s (commit: %s, date: %s)", version, commit, date) app.Usage = "A deployment tool for AWS ECS" app.Description = "A deployment tool for AWS ECS" - envars := env.Envars{} - cmds := commands.NewCageCommands(os.Stdin, provideCageCli) + cmds := commands.NewCageCommands(commands.ProvideCageCli) app.Commands = []*cli.Command{ - cmds.Up(&envars), - cmds.RollOut(&envars), - cmds.Run(&envars), - cmds.Upgrade(upgrade.NewUpgrader(version)), + cmds.Up(cageapp.NewCageCmdInput(os.Stdin, configCmdInput)), + cmds.RollOut(cageapp.NewCageCmdInput(os.Stdin, configCmdInput)), + cmds.Run(cageapp.NewCageCmdInput(os.Stdin, configCmdInput)), + commands.Upgrade(upgrade.NewUpgrader(version)), + commands.Audit(appConf, audit.ProvideAuditCmd), } app.Flags = []cli.Flag{ &cli.BoolFlag{ Name: "ci", Usage: "CI mode. Skip all confirmations and use default values.", EnvVars: []string{"CI"}, - Destination: &envars.CI, + Destination: &appConf.CI, + }, + &cli.BoolFlag{ + Name: "no-color", + Usage: "Disable colored output", + EnvVars: []string{"NO_COLOR"}, + Destination: &appConf.NoColor, }, } if err := app.Run(os.Args); err != nil { log.Fatal(err) } } - -func provideCageCli(envars *env.Envars) (types.Cage, error) { - conf, err := config.LoadDefaultConfig( - context.Background(), - config.WithRegion(envars.Region)) - if err != nil { - return nil, xerrors.Errorf("failed to load aws config: %w", err) - } - d := di.NewDomain(func(b *di.B) { - b.Set(key.Env, envars) - b.Set(key.EcsCli, ecs.NewFromConfig(conf)) - b.Set(key.Ec2Cli, ec2.NewFromConfig(conf)) - b.Set(key.AlbCli, elasticloadbalancingv2.NewFromConfig(conf)) - b.Set(key.TaskFactory, task.NewFactory(b.Future())) - b.Set(key.Time, &timeout.Time{}) - }) - cagecli := cage.NewCage(d) - return cagecli, nil -} diff --git a/cli/cage/prompt/prompt.go b/cli/cage/prompt/prompt.go index 05b8e84..bfeb458 100644 --- a/cli/cage/prompt/prompt.go +++ b/cli/cage/prompt/prompt.go @@ -47,10 +47,6 @@ func (s *Prompter) confirmStackChange( envars *env.Envars, service bool, ) error { - // Skip confirmation if running in CI - if envars.CI { - return nil - } if err := s.Confirm("region", envars.Region); err != nil { return err } diff --git a/env/env.go b/env/env.go index 0922204..a078612 100644 --- a/env/env.go +++ b/env/env.go @@ -13,7 +13,6 @@ import ( type Envars struct { _ struct{} `type:"struct"` - CI bool `json:"ci" type:"bool"` Region string `json:"region" type:"string"` Cluster string `json:"cluster" type:"string" required:"true"` Service string `json:"service" type:"string" required:"true"` diff --git a/go.mod b/go.mod index c6e09a6..97f0e3e 100644 --- a/go.mod +++ b/go.mod @@ -4,9 +4,10 @@ go 1.25.5 require ( github.com/apex/log v1.9.0 - github.com/aws/aws-sdk-go-v2 v1.41.0 + github.com/aws/aws-sdk-go-v2 v1.41.1 github.com/aws/aws-sdk-go-v2/config v1.32.6 github.com/aws/aws-sdk-go-v2/service/ec2 v1.279.0 + github.com/aws/aws-sdk-go-v2/service/ecr v1.55.1 github.com/aws/aws-sdk-go-v2/service/ecs v1.70.0 github.com/aws/aws-sdk-go-v2/service/elasticloadbalancingv2 v1.54.5 github.com/google/go-github/v62 v62.0.0 @@ -15,6 +16,7 @@ require ( github.com/stretchr/testify v1.11.1 github.com/urfave/cli/v2 v2.27.7 go.uber.org/mock v0.6.0 + golang.org/x/sync v0.16.0 ) require ( @@ -26,8 +28,8 @@ require ( github.com/Masterminds/semver/v3 v3.4.0 github.com/aws/aws-sdk-go-v2/credentials v1.19.6 // indirect github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.16 // indirect - github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.16 // indirect - github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.16 // indirect + github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.17 // indirect + github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.17 // indirect github.com/aws/aws-sdk-go-v2/internal/ini v1.8.4 // indirect github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.4 // indirect github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.16 // indirect @@ -42,7 +44,6 @@ require ( github.com/pmezard/go-difflib v1.0.0 // indirect github.com/russross/blackfriday/v2 v2.1.0 // indirect github.com/xrash/smetrics v0.0.0-20250705151800-55b8f293f342 // indirect - golang.org/x/sync v0.19.0 golang.org/x/xerrors v0.0.0-20240903120638-7835f813f4da gopkg.in/yaml.v3 v3.0.1 // indirect ) diff --git a/go.sum b/go.sum index 5d225bc..257a196 100644 --- a/go.sum +++ b/go.sum @@ -6,22 +6,24 @@ github.com/apex/logs v1.0.0/go.mod h1:XzxuLZ5myVHDy9SAmYpamKKRNApGj54PfYLcFrXqDw github.com/aphistic/golf v0.0.0-20180712155816-02c07f170c5a/go.mod h1:3NqKYiepwy8kCu4PNA+aP7WUV72eXWJeP9/r3/K9aLE= github.com/aphistic/sweet v0.2.0/go.mod h1:fWDlIh/isSE9n6EPsRmC0det+whmX6dJid3stzu0Xys= github.com/aws/aws-sdk-go v1.20.6/go.mod h1:KmX6BPdI08NWTb3/sm4ZGu5ShLoqVDhKgpiN924inxo= -github.com/aws/aws-sdk-go-v2 v1.41.0 h1:tNvqh1s+v0vFYdA1xq0aOJH+Y5cRyZ5upu6roPgPKd4= -github.com/aws/aws-sdk-go-v2 v1.41.0/go.mod h1:MayyLB8y+buD9hZqkCW3kX1AKq07Y5pXxtgB+rRFhz0= +github.com/aws/aws-sdk-go-v2 v1.41.1 h1:ABlyEARCDLN034NhxlRUSZr4l71mh+T5KAeGh6cerhU= +github.com/aws/aws-sdk-go-v2 v1.41.1/go.mod h1:MayyLB8y+buD9hZqkCW3kX1AKq07Y5pXxtgB+rRFhz0= github.com/aws/aws-sdk-go-v2/config v1.32.6 h1:hFLBGUKjmLAekvi1evLi5hVvFQtSo3GYwi+Bx4lpJf8= github.com/aws/aws-sdk-go-v2/config v1.32.6/go.mod h1:lcUL/gcd8WyjCrMnxez5OXkO3/rwcNmvfno62tnXNcI= github.com/aws/aws-sdk-go-v2/credentials v1.19.6 h1:F9vWao2TwjV2MyiyVS+duza0NIRtAslgLUM0vTA1ZaE= github.com/aws/aws-sdk-go-v2/credentials v1.19.6/go.mod h1:SgHzKjEVsdQr6Opor0ihgWtkWdfRAIwxYzSJ8O85VHY= github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.16 h1:80+uETIWS1BqjnN9uJ0dBUaETh+P1XwFy5vwHwK5r9k= github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.16/go.mod h1:wOOsYuxYuB/7FlnVtzeBYRcjSRtQpAW0hCP7tIULMwo= -github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.16 h1:rgGwPzb82iBYSvHMHXc8h9mRoOUBZIGFgKb9qniaZZc= -github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.16/go.mod h1:L/UxsGeKpGoIj6DxfhOWHWQ/kGKcd4I1VncE4++IyKA= -github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.16 h1:1jtGzuV7c82xnqOVfx2F0xmJcOw5374L7N6juGW6x6U= -github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.16/go.mod h1:M2E5OQf+XLe+SZGmmpaI2yy+J326aFf6/+54PoxSANc= +github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.17 h1:xOLELNKGp2vsiteLsvLPwxC+mYmO6OZ8PYgiuPJzF8U= +github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.17/go.mod h1:5M5CI3D12dNOtH3/mk6minaRwI2/37ifCURZISxA/IQ= +github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.17 h1:WWLqlh79iO48yLkj1v3ISRNiv+3KdQoZ6JWyfcsyQik= +github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.17/go.mod h1:EhG22vHRrvF8oXSTYStZhJc1aUgKtnJe+aOiFEV90cM= github.com/aws/aws-sdk-go-v2/internal/ini v1.8.4 h1:WKuaxf++XKWlHWu9ECbMlha8WOEGm0OUEZqm4K/Gcfk= github.com/aws/aws-sdk-go-v2/internal/ini v1.8.4/go.mod h1:ZWy7j6v1vWGmPReu0iSGvRiise4YI5SkR3OHKTZ6Wuc= github.com/aws/aws-sdk-go-v2/service/ec2 v1.279.0 h1:o7eJKe6VYAnqERPlLAvDW5VKXV6eTKv1oxTpMoDP378= github.com/aws/aws-sdk-go-v2/service/ec2 v1.279.0/go.mod h1:Wg68QRgy2gEGGdmTPU/UbVpdv8sM14bUZmF64KFwAsY= +github.com/aws/aws-sdk-go-v2/service/ecr v1.55.1 h1:B7f9R99lCF83XlolTg6d6Lvghyto+/VU83ZrneAVfK8= +github.com/aws/aws-sdk-go-v2/service/ecr v1.55.1/go.mod h1:cpYRXx5BkmS3mwWRKPbWSPKmyAUNL7aLWAPiiinwk/U= github.com/aws/aws-sdk-go-v2/service/ecs v1.70.0 h1:IZpZatHsscdOKjwmDXC6idsCXmm3F/obutAUNjnX+OM= github.com/aws/aws-sdk-go-v2/service/ecs v1.70.0/go.mod h1:LQMlcWBoiFVD3vUVEz42ST0yTiaDujv2dRE6sXt1yPE= github.com/aws/aws-sdk-go-v2/service/elasticloadbalancingv2 v1.54.5 h1:JjKuK9zbAVv6X44ia/OZrRS8ngOx3QfvtQTN0poJdPw= @@ -119,8 +121,8 @@ golang.org/x/net v0.0.0-20180906233101-161cd47e91fd/go.mod h1:mL1N/T3taQHkDXs73r golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= -golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4= -golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI= +golang.org/x/sync v0.16.0 h1:ycBJEhp9p4vXvUZNszeOq0kGTPghopOL8q0fq3vstxw= +golang.org/x/sync v0.16.0/go.mod h1:1dzgHSNfp02xaA81J2MS99Qcpr2w7fw1gpm99rleRqA= golang.org/x/sys v0.0.0-20180909124046-d0be0721c37e/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190222072716-a9d3bda3a223/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= diff --git a/key/keys.go b/key/keys.go index 7b32dda..2fef017 100644 --- a/key/keys.go +++ b/key/keys.go @@ -4,8 +4,12 @@ type DepsKey string const ( EcsCli DepsKey = "ecs" + EcrCli DepsKey = "ecr" Ec2Cli DepsKey = "ec2" AlbCli DepsKey = "alb" + Logger DepsKey = "logger" + Scanner DepsKey = "scanner" + Printer DepsKey = "printer" Env DepsKey = "env" Time DepsKey = "time" TaskFactory DepsKey = "task-factory" diff --git a/logger/color.go b/logger/color.go new file mode 100644 index 0000000..118703d --- /dev/null +++ b/logger/color.go @@ -0,0 +1,49 @@ +package logger + +import "fmt" + +type Color struct { + NoColor bool +} + +func (c *Color) sprintf(prefix, s, suffix string, args ...any) string { + if c.NoColor { + return fmt.Sprintf(s, args...) + } + return prefix + fmt.Sprintf(s, args...) + suffix +} + +func (c *Color) Red(s string) string { + return c.Redf("%s", s) +} +func (c *Color) Redf(s string, args ...any) string { + return c.sprintf("\033[31m", s, "\033[0m", args...) +} + +func (c *Color) Green(s string) string { + return c.Greenf("%s", s) +} +func (c *Color) Greenf(s string, args ...any) string { + return c.sprintf("\033[32m", s, "\033[0m", args...) +} + +func (c *Color) Yellow(s string) string { + return c.Yellowf("%s", s) +} +func (c *Color) Yellowf(s string, args ...any) string { + return c.sprintf("\033[33m", s, "\033[0m", args...) +} + +func (c *Color) Magenta(s string) string { + return c.Magentaf("%s", s) +} +func (c *Color) Magentaf(s string, args ...any) string { + return c.sprintf("\033[35m", s, "\033[0m", args...) +} + +func (c *Color) Bold(s string) string { + return c.Boldf("%s", s) +} +func (c *Color) Boldf(s string, args ...any) string { + return c.sprintf("\033[1m", s, "\033[0m", args...) +} diff --git a/logger/color_test.go b/logger/color_test.go new file mode 100644 index 0000000..7d9553f --- /dev/null +++ b/logger/color_test.go @@ -0,0 +1,118 @@ +package logger + +import "testing" + +func TestColor_Red(t *testing.T) { + c := &Color{NoColor: false} + result := c.Red("error") + expected := "\033[31merror\033[0m" + if result != expected { + t.Errorf("Red() = %q, want %q", result, expected) + } +} + +func TestColor_Redf(t *testing.T) { + c := &Color{NoColor: false} + result := c.Redf("error: %s", "message") + expected := "\033[31merror: message\033[0m" + if result != expected { + t.Errorf("Redf() = %q, want %q", result, expected) + } +} + +func TestColor_Green(t *testing.T) { + c := &Color{NoColor: false} + result := c.Green("success") + expected := "\033[32msuccess\033[0m" + if result != expected { + t.Errorf("Green() = %q, want %q", result, expected) + } +} + +func TestColor_Greenf(t *testing.T) { + c := &Color{NoColor: false} + result := c.Greenf("success: %d", 100) + expected := "\033[32msuccess: 100\033[0m" + if result != expected { + t.Errorf("Greenf() = %q, want %q", result, expected) + } +} + +func TestColor_Yellow(t *testing.T) { + c := &Color{NoColor: false} + result := c.Yellow("warning") + expected := "\033[33mwarning\033[0m" + if result != expected { + t.Errorf("Yellow() = %q, want %q", result, expected) + } +} + +func TestColor_Yellowf(t *testing.T) { + c := &Color{NoColor: false} + result := c.Yellowf("warning: %s", "test") + expected := "\033[33mwarning: test\033[0m" + if result != expected { + t.Errorf("Yellowf() = %q, want %q", result, expected) + } +} + +func TestColor_Magenta(t *testing.T) { + c := &Color{NoColor: false} + result := c.Magenta("info") + expected := "\033[35minfo\033[0m" + if result != expected { + t.Errorf("Magenta() = %q, want %q", result, expected) + } +} + +func TestColor_Magentaf(t *testing.T) { + c := &Color{NoColor: false} + result := c.Magentaf("info: %v", true) + expected := "\033[35minfo: true\033[0m" + if result != expected { + t.Errorf("Magentaf() = %q, want %q", result, expected) + } +} + +func TestColor_Bold(t *testing.T) { + c := &Color{NoColor: false} + result := c.Bold("bold") + expected := "\033[1mbold\033[0m" + if result != expected { + t.Errorf("Bold() = %q, want %q", result, expected) + } +} + +func TestColor_Boldf(t *testing.T) { + c := &Color{NoColor: false} + result := c.Boldf("bold: %s", "text") + expected := "\033[1mbold: text\033[0m" + if result != expected { + t.Errorf("Boldf() = %q, want %q", result, expected) + } +} + +func TestColor_NoColor(t *testing.T) { + c := &Color{NoColor: true} + tests := []struct { + name string + fn func() string + expected string + }{ + {"Red", func() string { return c.Red("text") }, "text"}, + {"Redf", func() string { return c.Redf("text: %s", "arg") }, "text: arg"}, + {"Green", func() string { return c.Green("text") }, "text"}, + {"Yellow", func() string { return c.Yellow("text") }, "text"}, + {"Magenta", func() string { return c.Magenta("text") }, "text"}, + {"Boldf", func() string { return c.Boldf("text: %d", 42) }, "text: 42"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := tt.fn() + if result != tt.expected { + t.Errorf("%s with NoColor = %q, want %q", tt.name, result, tt.expected) + } + }) + } +} diff --git a/logger/logger.go b/logger/logger.go new file mode 100644 index 0000000..95b56a1 --- /dev/null +++ b/logger/logger.go @@ -0,0 +1,22 @@ +package logger + +import ( + "fmt" + "io" +) + +type Logger interface { + Printf(format string, args ...any) +} + +func DefaultLogger(stdout io.Writer) Logger { + return &defaultLogger{stdout: stdout} +} + +type defaultLogger struct { + stdout io.Writer +} + +func (l *defaultLogger) Printf(format string, args ...any) { + fmt.Fprintf(l.stdout, format, args...) +} diff --git a/logger/logger_test.go b/logger/logger_test.go new file mode 100644 index 0000000..3307c7c --- /dev/null +++ b/logger/logger_test.go @@ -0,0 +1,17 @@ +package logger_test + +import ( + "bytes" + "testing" + + "github.com/loilo-inc/canarycage/logger" + "github.com/stretchr/testify/assert" +) + +func TestDefaultLogger_Printf(t *testing.T) { + var bin bytes.Buffer + logger := logger.DefaultLogger(&bin) + logger.Printf("Hello, %s!", "world") + output := bin.String() + assert.Equal(t, "Hello, world!", output) +} diff --git a/logger/spinner.go b/logger/spinner.go new file mode 100644 index 0000000..fae384c --- /dev/null +++ b/logger/spinner.go @@ -0,0 +1,19 @@ +package logger + +type spinner struct { + frames []string + index int +} + +func NewSpinner() *spinner { + return &spinner{ + frames: []string{"⠋", "⠙", "⠹", "⠸", "⠼", "⠴", "⠦", "⠧", "⠇", "⠏"}, + index: 0, + } +} + +func (s *spinner) Next() string { + frame := s.frames[s.index] + s.index = (s.index + 1) % len(s.frames) + return frame +} diff --git a/logger/spinner_test.go b/logger/spinner_test.go new file mode 100644 index 0000000..86f2c5e --- /dev/null +++ b/logger/spinner_test.go @@ -0,0 +1,35 @@ +package logger + +import "testing" + +func TestNewSpinner(t *testing.T) { + s := NewSpinner() + if s == nil { + t.Fatal("NewSpinner() returned nil") + } + if len(s.frames) != 10 { + t.Errorf("expected 10 frames, got %d", len(s.frames)) + } + if s.index != 0 { + t.Errorf("expected initial index 0, got %d", s.index) + } +} + +func TestSpinnerNext(t *testing.T) { + s := NewSpinner() + expectedFrames := []string{"⠋", "⠙", "⠹", "⠸", "⠼", "⠴", "⠦", "⠧", "⠇", "⠏"} + + // Test that Next() returns frames in order + for i := 0; i < len(expectedFrames); i++ { + frame := s.Next() + if frame != expectedFrames[i] { + t.Errorf("expected frame %q at index %d, got %q", expectedFrames[i], i, frame) + } + } + + // Test that it wraps around + frame := s.Next() + if frame != expectedFrames[0] { + t.Errorf("expected frame to wrap around to %q, got %q", expectedFrames[0], frame) + } +} diff --git a/mocks/mock_audit/printer.go b/mocks/mock_audit/printer.go new file mode 100644 index 0000000..787517c --- /dev/null +++ b/mocks/mock_audit/printer.go @@ -0,0 +1,53 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: ./cli/cage/audit/printer.go +// +// Generated by this command: +// +// mockgen -source=./cli/cage/audit/printer.go +// + +// Package mock_audit is a generated GoMock package. +package mock_audit + +import ( + reflect "reflect" + + audit "github.com/loilo-inc/canarycage/cli/cage/audit" + gomock "go.uber.org/mock/gomock" +) + +// MockPrinter is a mock of Printer interface. +type MockPrinter struct { + ctrl *gomock.Controller + recorder *MockPrinterMockRecorder + isgomock struct{} +} + +// MockPrinterMockRecorder is the mock recorder for MockPrinter. +type MockPrinterMockRecorder struct { + mock *MockPrinter +} + +// NewMockPrinter creates a new mock instance. +func NewMockPrinter(ctrl *gomock.Controller) *MockPrinter { + mock := &MockPrinter{ctrl: ctrl} + mock.recorder = &MockPrinterMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockPrinter) EXPECT() *MockPrinterMockRecorder { + return m.recorder +} + +// Print mocks base method. +func (m *MockPrinter) Print(result []*audit.ScanResult) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "Print", result) +} + +// Print indicates an expected call of Print. +func (mr *MockPrinterMockRecorder) Print(result any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Print", reflect.TypeOf((*MockPrinter)(nil).Print), result) +} diff --git a/mocks/mock_audit/scanner.go b/mocks/mock_audit/scanner.go new file mode 100644 index 0000000..957baa7 --- /dev/null +++ b/mocks/mock_audit/scanner.go @@ -0,0 +1,57 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: ./cli/cage/audit/scanner.go +// +// Generated by this command: +// +// mockgen -source=./cli/cage/audit/scanner.go +// + +// Package mock_audit is a generated GoMock package. +package mock_audit + +import ( + context "context" + reflect "reflect" + + audit "github.com/loilo-inc/canarycage/cli/cage/audit" + gomock "go.uber.org/mock/gomock" +) + +// MockScanner is a mock of Scanner interface. +type MockScanner struct { + ctrl *gomock.Controller + recorder *MockScannerMockRecorder + isgomock struct{} +} + +// MockScannerMockRecorder is the mock recorder for MockScanner. +type MockScannerMockRecorder struct { + mock *MockScanner +} + +// NewMockScanner creates a new mock instance. +func NewMockScanner(ctrl *gomock.Controller) *MockScanner { + mock := &MockScanner{ctrl: ctrl} + mock.recorder = &MockScannerMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockScanner) EXPECT() *MockScannerMockRecorder { + return m.recorder +} + +// Scan mocks base method. +func (m *MockScanner) Scan(ctx context.Context, cluster, service string) ([]*audit.ScanResult, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Scan", ctx, cluster, service) + ret0, _ := ret[0].([]*audit.ScanResult) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// Scan indicates an expected call of Scan. +func (mr *MockScannerMockRecorder) Scan(ctx, cluster, service any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Scan", reflect.TypeOf((*MockScanner)(nil).Scan), ctx, cluster, service) +} diff --git a/mocks/mock_awsiface/iface.go b/mocks/mock_awsiface/iface.go index a7dc97d..b3fdbda 100644 --- a/mocks/mock_awsiface/iface.go +++ b/mocks/mock_awsiface/iface.go @@ -1,5 +1,10 @@ // Code generated by MockGen. DO NOT EDIT. // Source: ./awsiface/iface.go +// +// Generated by this command: +// +// mockgen -source=./awsiface/iface.go +// // Package mock_awsiface is a generated GoMock package. package mock_awsiface @@ -9,6 +14,7 @@ import ( reflect "reflect" ec2 "github.com/aws/aws-sdk-go-v2/service/ec2" + ecr "github.com/aws/aws-sdk-go-v2/service/ecr" ecs "github.com/aws/aws-sdk-go-v2/service/ecs" elasticloadbalancingv2 "github.com/aws/aws-sdk-go-v2/service/elasticloadbalancingv2" gomock "go.uber.org/mock/gomock" @@ -18,6 +24,7 @@ import ( type MockEcsClient struct { ctrl *gomock.Controller recorder *MockEcsClientMockRecorder + isgomock struct{} } // MockEcsClientMockRecorder is the mock recorder for MockEcsClient. @@ -40,7 +47,7 @@ func (m *MockEcsClient) EXPECT() *MockEcsClientMockRecorder { // CreateService mocks base method. func (m *MockEcsClient) CreateService(ctx context.Context, params *ecs.CreateServiceInput, optFns ...func(*ecs.Options)) (*ecs.CreateServiceOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -51,16 +58,16 @@ func (m *MockEcsClient) CreateService(ctx context.Context, params *ecs.CreateSer } // CreateService indicates an expected call of CreateService. -func (mr *MockEcsClientMockRecorder) CreateService(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockEcsClientMockRecorder) CreateService(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CreateService", reflect.TypeOf((*MockEcsClient)(nil).CreateService), varargs...) } // DeleteService mocks base method. func (m *MockEcsClient) DeleteService(ctx context.Context, params *ecs.DeleteServiceInput, optFns ...func(*ecs.Options)) (*ecs.DeleteServiceOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -71,16 +78,16 @@ func (m *MockEcsClient) DeleteService(ctx context.Context, params *ecs.DeleteSer } // DeleteService indicates an expected call of DeleteService. -func (mr *MockEcsClientMockRecorder) DeleteService(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockEcsClientMockRecorder) DeleteService(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteService", reflect.TypeOf((*MockEcsClient)(nil).DeleteService), varargs...) } // DescribeContainerInstances mocks base method. func (m *MockEcsClient) DescribeContainerInstances(ctx context.Context, params *ecs.DescribeContainerInstancesInput, optFns ...func(*ecs.Options)) (*ecs.DescribeContainerInstancesOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -91,16 +98,16 @@ func (m *MockEcsClient) DescribeContainerInstances(ctx context.Context, params * } // DescribeContainerInstances indicates an expected call of DescribeContainerInstances. -func (mr *MockEcsClientMockRecorder) DescribeContainerInstances(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockEcsClientMockRecorder) DescribeContainerInstances(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DescribeContainerInstances", reflect.TypeOf((*MockEcsClient)(nil).DescribeContainerInstances), varargs...) } // DescribeServices mocks base method. func (m *MockEcsClient) DescribeServices(ctx context.Context, params *ecs.DescribeServicesInput, optFns ...func(*ecs.Options)) (*ecs.DescribeServicesOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -111,16 +118,16 @@ func (m *MockEcsClient) DescribeServices(ctx context.Context, params *ecs.Descri } // DescribeServices indicates an expected call of DescribeServices. -func (mr *MockEcsClientMockRecorder) DescribeServices(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockEcsClientMockRecorder) DescribeServices(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DescribeServices", reflect.TypeOf((*MockEcsClient)(nil).DescribeServices), varargs...) } // DescribeTaskDefinition mocks base method. func (m *MockEcsClient) DescribeTaskDefinition(ctx context.Context, params *ecs.DescribeTaskDefinitionInput, optFns ...func(*ecs.Options)) (*ecs.DescribeTaskDefinitionOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -131,16 +138,16 @@ func (m *MockEcsClient) DescribeTaskDefinition(ctx context.Context, params *ecs. } // DescribeTaskDefinition indicates an expected call of DescribeTaskDefinition. -func (mr *MockEcsClientMockRecorder) DescribeTaskDefinition(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockEcsClientMockRecorder) DescribeTaskDefinition(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DescribeTaskDefinition", reflect.TypeOf((*MockEcsClient)(nil).DescribeTaskDefinition), varargs...) } // DescribeTasks mocks base method. func (m *MockEcsClient) DescribeTasks(ctx context.Context, params *ecs.DescribeTasksInput, optFns ...func(*ecs.Options)) (*ecs.DescribeTasksOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -151,16 +158,16 @@ func (m *MockEcsClient) DescribeTasks(ctx context.Context, params *ecs.DescribeT } // DescribeTasks indicates an expected call of DescribeTasks. -func (mr *MockEcsClientMockRecorder) DescribeTasks(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockEcsClientMockRecorder) DescribeTasks(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DescribeTasks", reflect.TypeOf((*MockEcsClient)(nil).DescribeTasks), varargs...) } // ListTasks mocks base method. func (m *MockEcsClient) ListTasks(ctx context.Context, params *ecs.ListTasksInput, optFns ...func(*ecs.Options)) (*ecs.ListTasksOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -171,16 +178,16 @@ func (m *MockEcsClient) ListTasks(ctx context.Context, params *ecs.ListTasksInpu } // ListTasks indicates an expected call of ListTasks. -func (mr *MockEcsClientMockRecorder) ListTasks(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockEcsClientMockRecorder) ListTasks(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListTasks", reflect.TypeOf((*MockEcsClient)(nil).ListTasks), varargs...) } // RegisterTaskDefinition mocks base method. func (m *MockEcsClient) RegisterTaskDefinition(ctx context.Context, params *ecs.RegisterTaskDefinitionInput, optFns ...func(*ecs.Options)) (*ecs.RegisterTaskDefinitionOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -191,16 +198,16 @@ func (m *MockEcsClient) RegisterTaskDefinition(ctx context.Context, params *ecs. } // RegisterTaskDefinition indicates an expected call of RegisterTaskDefinition. -func (mr *MockEcsClientMockRecorder) RegisterTaskDefinition(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockEcsClientMockRecorder) RegisterTaskDefinition(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RegisterTaskDefinition", reflect.TypeOf((*MockEcsClient)(nil).RegisterTaskDefinition), varargs...) } // RunTask mocks base method. func (m *MockEcsClient) RunTask(ctx context.Context, params *ecs.RunTaskInput, optFns ...func(*ecs.Options)) (*ecs.RunTaskOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -211,16 +218,16 @@ func (m *MockEcsClient) RunTask(ctx context.Context, params *ecs.RunTaskInput, o } // RunTask indicates an expected call of RunTask. -func (mr *MockEcsClientMockRecorder) RunTask(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockEcsClientMockRecorder) RunTask(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RunTask", reflect.TypeOf((*MockEcsClient)(nil).RunTask), varargs...) } // StartTask mocks base method. func (m *MockEcsClient) StartTask(ctx context.Context, params *ecs.StartTaskInput, optFns ...func(*ecs.Options)) (*ecs.StartTaskOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -231,16 +238,16 @@ func (m *MockEcsClient) StartTask(ctx context.Context, params *ecs.StartTaskInpu } // StartTask indicates an expected call of StartTask. -func (mr *MockEcsClientMockRecorder) StartTask(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockEcsClientMockRecorder) StartTask(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "StartTask", reflect.TypeOf((*MockEcsClient)(nil).StartTask), varargs...) } // StopTask mocks base method. func (m *MockEcsClient) StopTask(ctx context.Context, params *ecs.StopTaskInput, optFns ...func(*ecs.Options)) (*ecs.StopTaskOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -251,16 +258,16 @@ func (m *MockEcsClient) StopTask(ctx context.Context, params *ecs.StopTaskInput, } // StopTask indicates an expected call of StopTask. -func (mr *MockEcsClientMockRecorder) StopTask(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockEcsClientMockRecorder) StopTask(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "StopTask", reflect.TypeOf((*MockEcsClient)(nil).StopTask), varargs...) } // UpdateService mocks base method. func (m *MockEcsClient) UpdateService(ctx context.Context, params *ecs.UpdateServiceInput, optFns ...func(*ecs.Options)) (*ecs.UpdateServiceOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -271,16 +278,81 @@ func (m *MockEcsClient) UpdateService(ctx context.Context, params *ecs.UpdateSer } // UpdateService indicates an expected call of UpdateService. -func (mr *MockEcsClientMockRecorder) UpdateService(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockEcsClientMockRecorder) UpdateService(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateService", reflect.TypeOf((*MockEcsClient)(nil).UpdateService), varargs...) } +// MockEcrClient is a mock of EcrClient interface. +type MockEcrClient struct { + ctrl *gomock.Controller + recorder *MockEcrClientMockRecorder + isgomock struct{} +} + +// MockEcrClientMockRecorder is the mock recorder for MockEcrClient. +type MockEcrClientMockRecorder struct { + mock *MockEcrClient +} + +// NewMockEcrClient creates a new mock instance. +func NewMockEcrClient(ctrl *gomock.Controller) *MockEcrClient { + mock := &MockEcrClient{ctrl: ctrl} + mock.recorder = &MockEcrClientMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockEcrClient) EXPECT() *MockEcrClientMockRecorder { + return m.recorder +} + +// BatchGetImage mocks base method. +func (m *MockEcrClient) BatchGetImage(ctx context.Context, params *ecr.BatchGetImageInput, optFns ...func(*ecr.Options)) (*ecr.BatchGetImageOutput, error) { + m.ctrl.T.Helper() + varargs := []any{ctx, params} + for _, a := range optFns { + varargs = append(varargs, a) + } + ret := m.ctrl.Call(m, "BatchGetImage", varargs...) + ret0, _ := ret[0].(*ecr.BatchGetImageOutput) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// BatchGetImage indicates an expected call of BatchGetImage. +func (mr *MockEcrClientMockRecorder) BatchGetImage(ctx, params any, optFns ...any) *gomock.Call { + mr.mock.ctrl.T.Helper() + varargs := append([]any{ctx, params}, optFns...) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "BatchGetImage", reflect.TypeOf((*MockEcrClient)(nil).BatchGetImage), varargs...) +} + +// DescribeImageScanFindings mocks base method. +func (m *MockEcrClient) DescribeImageScanFindings(ctx context.Context, params *ecr.DescribeImageScanFindingsInput, optFns ...func(*ecr.Options)) (*ecr.DescribeImageScanFindingsOutput, error) { + m.ctrl.T.Helper() + varargs := []any{ctx, params} + for _, a := range optFns { + varargs = append(varargs, a) + } + ret := m.ctrl.Call(m, "DescribeImageScanFindings", varargs...) + ret0, _ := ret[0].(*ecr.DescribeImageScanFindingsOutput) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// DescribeImageScanFindings indicates an expected call of DescribeImageScanFindings. +func (mr *MockEcrClientMockRecorder) DescribeImageScanFindings(ctx, params any, optFns ...any) *gomock.Call { + mr.mock.ctrl.T.Helper() + varargs := append([]any{ctx, params}, optFns...) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DescribeImageScanFindings", reflect.TypeOf((*MockEcrClient)(nil).DescribeImageScanFindings), varargs...) +} + // MockAlbClient is a mock of AlbClient interface. type MockAlbClient struct { ctrl *gomock.Controller recorder *MockAlbClientMockRecorder + isgomock struct{} } // MockAlbClientMockRecorder is the mock recorder for MockAlbClient. @@ -303,7 +375,7 @@ func (m *MockAlbClient) EXPECT() *MockAlbClientMockRecorder { // DeregisterTargets mocks base method. func (m *MockAlbClient) DeregisterTargets(ctx context.Context, params *elasticloadbalancingv2.DeregisterTargetsInput, optFns ...func(*elasticloadbalancingv2.Options)) (*elasticloadbalancingv2.DeregisterTargetsOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -314,16 +386,16 @@ func (m *MockAlbClient) DeregisterTargets(ctx context.Context, params *elasticlo } // DeregisterTargets indicates an expected call of DeregisterTargets. -func (mr *MockAlbClientMockRecorder) DeregisterTargets(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockAlbClientMockRecorder) DeregisterTargets(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeregisterTargets", reflect.TypeOf((*MockAlbClient)(nil).DeregisterTargets), varargs...) } // DescribeTargetGroupAttributes mocks base method. func (m *MockAlbClient) DescribeTargetGroupAttributes(ctx context.Context, params *elasticloadbalancingv2.DescribeTargetGroupAttributesInput, optFns ...func(*elasticloadbalancingv2.Options)) (*elasticloadbalancingv2.DescribeTargetGroupAttributesOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -334,16 +406,16 @@ func (m *MockAlbClient) DescribeTargetGroupAttributes(ctx context.Context, param } // DescribeTargetGroupAttributes indicates an expected call of DescribeTargetGroupAttributes. -func (mr *MockAlbClientMockRecorder) DescribeTargetGroupAttributes(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockAlbClientMockRecorder) DescribeTargetGroupAttributes(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DescribeTargetGroupAttributes", reflect.TypeOf((*MockAlbClient)(nil).DescribeTargetGroupAttributes), varargs...) } // DescribeTargetGroups mocks base method. func (m *MockAlbClient) DescribeTargetGroups(ctx context.Context, params *elasticloadbalancingv2.DescribeTargetGroupsInput, optFns ...func(*elasticloadbalancingv2.Options)) (*elasticloadbalancingv2.DescribeTargetGroupsOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -354,16 +426,16 @@ func (m *MockAlbClient) DescribeTargetGroups(ctx context.Context, params *elasti } // DescribeTargetGroups indicates an expected call of DescribeTargetGroups. -func (mr *MockAlbClientMockRecorder) DescribeTargetGroups(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockAlbClientMockRecorder) DescribeTargetGroups(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DescribeTargetGroups", reflect.TypeOf((*MockAlbClient)(nil).DescribeTargetGroups), varargs...) } // DescribeTargetHealth mocks base method. func (m *MockAlbClient) DescribeTargetHealth(ctx context.Context, params *elasticloadbalancingv2.DescribeTargetHealthInput, optFns ...func(*elasticloadbalancingv2.Options)) (*elasticloadbalancingv2.DescribeTargetHealthOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -374,16 +446,16 @@ func (m *MockAlbClient) DescribeTargetHealth(ctx context.Context, params *elasti } // DescribeTargetHealth indicates an expected call of DescribeTargetHealth. -func (mr *MockAlbClientMockRecorder) DescribeTargetHealth(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockAlbClientMockRecorder) DescribeTargetHealth(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DescribeTargetHealth", reflect.TypeOf((*MockAlbClient)(nil).DescribeTargetHealth), varargs...) } // RegisterTargets mocks base method. func (m *MockAlbClient) RegisterTargets(ctx context.Context, params *elasticloadbalancingv2.RegisterTargetsInput, optFns ...func(*elasticloadbalancingv2.Options)) (*elasticloadbalancingv2.RegisterTargetsOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -394,9 +466,9 @@ func (m *MockAlbClient) RegisterTargets(ctx context.Context, params *elasticload } // RegisterTargets indicates an expected call of RegisterTargets. -func (mr *MockAlbClientMockRecorder) RegisterTargets(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockAlbClientMockRecorder) RegisterTargets(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RegisterTargets", reflect.TypeOf((*MockAlbClient)(nil).RegisterTargets), varargs...) } @@ -404,6 +476,7 @@ func (mr *MockAlbClientMockRecorder) RegisterTargets(ctx, params interface{}, op type MockEc2Client struct { ctrl *gomock.Controller recorder *MockEc2ClientMockRecorder + isgomock struct{} } // MockEc2ClientMockRecorder is the mock recorder for MockEc2Client. @@ -426,7 +499,7 @@ func (m *MockEc2Client) EXPECT() *MockEc2ClientMockRecorder { // DescribeInstances mocks base method. func (m *MockEc2Client) DescribeInstances(ctx context.Context, params *ec2.DescribeInstancesInput, optFns ...func(*ec2.Options)) (*ec2.DescribeInstancesOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -437,16 +510,16 @@ func (m *MockEc2Client) DescribeInstances(ctx context.Context, params *ec2.Descr } // DescribeInstances indicates an expected call of DescribeInstances. -func (mr *MockEc2ClientMockRecorder) DescribeInstances(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockEc2ClientMockRecorder) DescribeInstances(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DescribeInstances", reflect.TypeOf((*MockEc2Client)(nil).DescribeInstances), varargs...) } // DescribeSubnets mocks base method. func (m *MockEc2Client) DescribeSubnets(ctx context.Context, params *ec2.DescribeSubnetsInput, optFns ...func(*ec2.Options)) (*ec2.DescribeSubnetsOutput, error) { m.ctrl.T.Helper() - varargs := []interface{}{ctx, params} + varargs := []any{ctx, params} for _, a := range optFns { varargs = append(varargs, a) } @@ -457,8 +530,8 @@ func (m *MockEc2Client) DescribeSubnets(ctx context.Context, params *ec2.Describ } // DescribeSubnets indicates an expected call of DescribeSubnets. -func (mr *MockEc2ClientMockRecorder) DescribeSubnets(ctx, params interface{}, optFns ...interface{}) *gomock.Call { +func (mr *MockEc2ClientMockRecorder) DescribeSubnets(ctx, params any, optFns ...any) *gomock.Call { mr.mock.ctrl.T.Helper() - varargs := append([]interface{}{ctx, params}, optFns...) + varargs := append([]any{ctx, params}, optFns...) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DescribeSubnets", reflect.TypeOf((*MockEc2Client)(nil).DescribeSubnets), varargs...) } diff --git a/mocks/mock_logger/logger.go b/mocks/mock_logger/logger.go new file mode 100644 index 0000000..f1610ad --- /dev/null +++ b/mocks/mock_logger/logger.go @@ -0,0 +1,57 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: ./logger/logger.go +// +// Generated by this command: +// +// mockgen -source=./logger/logger.go +// + +// Package mock_logger is a generated GoMock package. +package mock_logger + +import ( + reflect "reflect" + + gomock "go.uber.org/mock/gomock" +) + +// MockLogger is a mock of Logger interface. +type MockLogger struct { + ctrl *gomock.Controller + recorder *MockLoggerMockRecorder + isgomock struct{} +} + +// MockLoggerMockRecorder is the mock recorder for MockLogger. +type MockLoggerMockRecorder struct { + mock *MockLogger +} + +// NewMockLogger creates a new mock instance. +func NewMockLogger(ctrl *gomock.Controller) *MockLogger { + mock := &MockLogger{ctrl: ctrl} + mock.recorder = &MockLoggerMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockLogger) EXPECT() *MockLoggerMockRecorder { + return m.recorder +} + +// Printf mocks base method. +func (m *MockLogger) Printf(format string, args ...any) { + m.ctrl.T.Helper() + varargs := []any{format} + for _, a := range args { + varargs = append(varargs, a) + } + m.ctrl.Call(m, "Printf", varargs...) +} + +// Printf indicates an expected call of Printf. +func (mr *MockLoggerMockRecorder) Printf(format any, args ...any) *gomock.Call { + mr.mock.ctrl.T.Helper() + varargs := append([]any{format}, args...) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Printf", reflect.TypeOf((*MockLogger)(nil).Printf), varargs...) +} diff --git a/mocks/mock_rollout/executor.go b/mocks/mock_rollout/executor.go index 4ee039e..f86d285 100644 --- a/mocks/mock_rollout/executor.go +++ b/mocks/mock_rollout/executor.go @@ -1,5 +1,10 @@ // Code generated by MockGen. DO NOT EDIT. // Source: ./rollout/executor.go +// +// Generated by this command: +// +// mockgen -source=./rollout/executor.go +// // Package mock_rollout is a generated GoMock package. package mock_rollout @@ -8,14 +13,15 @@ import ( context "context" reflect "reflect" - gomock "go.uber.org/mock/gomock" types "github.com/loilo-inc/canarycage/types" + gomock "go.uber.org/mock/gomock" ) // MockExecutor is a mock of Executor interface. type MockExecutor struct { ctrl *gomock.Controller recorder *MockExecutorMockRecorder + isgomock struct{} } // MockExecutorMockRecorder is the mock recorder for MockExecutor. @@ -44,7 +50,7 @@ func (m *MockExecutor) RollOut(ctx context.Context, input *types.RollOutInput) e } // RollOut indicates an expected call of RollOut. -func (mr *MockExecutorMockRecorder) RollOut(ctx, input interface{}) *gomock.Call { +func (mr *MockExecutorMockRecorder) RollOut(ctx, input any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RollOut", reflect.TypeOf((*MockExecutor)(nil).RollOut), ctx, input) } diff --git a/mocks/mock_task/factory.go b/mocks/mock_task/factory.go index eb99313..08cdfa0 100644 --- a/mocks/mock_task/factory.go +++ b/mocks/mock_task/factory.go @@ -1,5 +1,10 @@ // Code generated by MockGen. DO NOT EDIT. // Source: ./task/factory.go +// +// Generated by this command: +// +// mockgen -source=./task/factory.go +// // Package mock_task is a generated GoMock package. package mock_task @@ -8,14 +13,15 @@ import ( reflect "reflect" types "github.com/aws/aws-sdk-go-v2/service/ecs/types" - gomock "go.uber.org/mock/gomock" task "github.com/loilo-inc/canarycage/task" + gomock "go.uber.org/mock/gomock" ) // MockFactory is a mock of Factory interface. type MockFactory struct { ctrl *gomock.Controller recorder *MockFactoryMockRecorder + isgomock struct{} } // MockFactoryMockRecorder is the mock recorder for MockFactory. @@ -44,7 +50,7 @@ func (m *MockFactory) NewAlbTask(input *task.Input, lb *types.LoadBalancer) task } // NewAlbTask indicates an expected call of NewAlbTask. -func (mr *MockFactoryMockRecorder) NewAlbTask(input, lb interface{}) *gomock.Call { +func (mr *MockFactoryMockRecorder) NewAlbTask(input, lb any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "NewAlbTask", reflect.TypeOf((*MockFactory)(nil).NewAlbTask), input, lb) } @@ -58,7 +64,7 @@ func (m *MockFactory) NewSimpleTask(input *task.Input) task.Task { } // NewSimpleTask indicates an expected call of NewSimpleTask. -func (mr *MockFactoryMockRecorder) NewSimpleTask(input interface{}) *gomock.Call { +func (mr *MockFactoryMockRecorder) NewSimpleTask(input any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "NewSimpleTask", reflect.TypeOf((*MockFactory)(nil).NewSimpleTask), input) } diff --git a/mocks/mock_task/task.go b/mocks/mock_task/task.go index d991706..ff4490a 100644 --- a/mocks/mock_task/task.go +++ b/mocks/mock_task/task.go @@ -1,5 +1,10 @@ // Code generated by MockGen. DO NOT EDIT. // Source: ./task/task.go +// +// Generated by this command: +// +// mockgen -source=./task/task.go +// // Package mock_task is a generated GoMock package. package mock_task @@ -15,6 +20,7 @@ import ( type MockTask struct { ctrl *gomock.Controller recorder *MockTaskMockRecorder + isgomock struct{} } // MockTaskMockRecorder is the mock recorder for MockTask. @@ -43,7 +49,7 @@ func (m *MockTask) Start(ctx context.Context) error { } // Start indicates an expected call of Start. -func (mr *MockTaskMockRecorder) Start(ctx interface{}) *gomock.Call { +func (mr *MockTaskMockRecorder) Start(ctx any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Start", reflect.TypeOf((*MockTask)(nil).Start), ctx) } @@ -57,7 +63,7 @@ func (m *MockTask) Stop(ctx context.Context) error { } // Stop indicates an expected call of Stop. -func (mr *MockTaskMockRecorder) Stop(ctx interface{}) *gomock.Call { +func (mr *MockTaskMockRecorder) Stop(ctx any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Stop", reflect.TypeOf((*MockTask)(nil).Stop), ctx) } @@ -71,7 +77,7 @@ func (m *MockTask) Wait(ctx context.Context) error { } // Wait indicates an expected call of Wait. -func (mr *MockTaskMockRecorder) Wait(ctx interface{}) *gomock.Call { +func (mr *MockTaskMockRecorder) Wait(ctx any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Wait", reflect.TypeOf((*MockTask)(nil).Wait), ctx) } diff --git a/mocks/mock_taskset/taskset.go b/mocks/mock_taskset/taskset.go index 3cd1520..057f771 100644 --- a/mocks/mock_taskset/taskset.go +++ b/mocks/mock_taskset/taskset.go @@ -1,5 +1,10 @@ // Code generated by MockGen. DO NOT EDIT. // Source: ./taskset/taskset.go +// +// Generated by this command: +// +// mockgen -source=./taskset/taskset.go +// // Package mock_taskset is a generated GoMock package. package mock_taskset @@ -15,6 +20,7 @@ import ( type MockSet struct { ctrl *gomock.Controller recorder *MockSetMockRecorder + isgomock struct{} } // MockSetMockRecorder is the mock recorder for MockSet. @@ -43,7 +49,7 @@ func (m *MockSet) Cleanup(ctx context.Context) error { } // Cleanup indicates an expected call of Cleanup. -func (mr *MockSetMockRecorder) Cleanup(ctx interface{}) *gomock.Call { +func (mr *MockSetMockRecorder) Cleanup(ctx any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Cleanup", reflect.TypeOf((*MockSet)(nil).Cleanup), ctx) } @@ -57,7 +63,7 @@ func (m *MockSet) Exec(ctx context.Context) error { } // Exec indicates an expected call of Exec. -func (mr *MockSetMockRecorder) Exec(ctx interface{}) *gomock.Call { +func (mr *MockSetMockRecorder) Exec(ctx any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Exec", reflect.TypeOf((*MockSet)(nil).Exec), ctx) } diff --git a/mocks/mock_types/iface.go b/mocks/mock_types/iface.go index a119993..de84332 100644 --- a/mocks/mock_types/iface.go +++ b/mocks/mock_types/iface.go @@ -1,5 +1,10 @@ // Code generated by MockGen. DO NOT EDIT. // Source: ./types/iface.go +// +// Generated by this command: +// +// mockgen -source=./types/iface.go +// // Package mock_types is a generated GoMock package. package mock_types @@ -9,14 +14,15 @@ import ( reflect "reflect" time "time" - gomock "go.uber.org/mock/gomock" types "github.com/loilo-inc/canarycage/types" + gomock "go.uber.org/mock/gomock" ) // MockCage is a mock of Cage interface. type MockCage struct { ctrl *gomock.Controller recorder *MockCageMockRecorder + isgomock struct{} } // MockCageMockRecorder is the mock recorder for MockCage. @@ -46,7 +52,7 @@ func (m *MockCage) RollOut(ctx context.Context, input *types.RollOutInput) (*typ } // RollOut indicates an expected call of RollOut. -func (mr *MockCageMockRecorder) RollOut(ctx, input interface{}) *gomock.Call { +func (mr *MockCageMockRecorder) RollOut(ctx, input any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RollOut", reflect.TypeOf((*MockCage)(nil).RollOut), ctx, input) } @@ -61,7 +67,7 @@ func (m *MockCage) Run(ctx context.Context, input *types.RunInput) (*types.RunRe } // Run indicates an expected call of Run. -func (mr *MockCageMockRecorder) Run(ctx, input interface{}) *gomock.Call { +func (mr *MockCageMockRecorder) Run(ctx, input any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Run", reflect.TypeOf((*MockCage)(nil).Run), ctx, input) } @@ -76,15 +82,54 @@ func (m *MockCage) Up(ctx context.Context) (*types.UpResult, error) { } // Up indicates an expected call of Up. -func (mr *MockCageMockRecorder) Up(ctx interface{}) *gomock.Call { +func (mr *MockCageMockRecorder) Up(ctx any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Up", reflect.TypeOf((*MockCage)(nil).Up), ctx) } +// MockAudit is a mock of Audit interface. +type MockAudit struct { + ctrl *gomock.Controller + recorder *MockAuditMockRecorder + isgomock struct{} +} + +// MockAuditMockRecorder is the mock recorder for MockAudit. +type MockAuditMockRecorder struct { + mock *MockAudit +} + +// NewMockAudit creates a new mock instance. +func NewMockAudit(ctrl *gomock.Controller) *MockAudit { + mock := &MockAudit{ctrl: ctrl} + mock.recorder = &MockAuditMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockAudit) EXPECT() *MockAuditMockRecorder { + return m.recorder +} + +// Run mocks base method. +func (m *MockAudit) Run(ctx context.Context) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Run", ctx) + ret0, _ := ret[0].(error) + return ret0 +} + +// Run indicates an expected call of Run. +func (mr *MockAuditMockRecorder) Run(ctx any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Run", reflect.TypeOf((*MockAudit)(nil).Run), ctx) +} + // MockTime is a mock of Time interface. type MockTime struct { ctrl *gomock.Controller recorder *MockTimeMockRecorder + isgomock struct{} } // MockTimeMockRecorder is the mock recorder for MockTime. @@ -113,7 +158,7 @@ func (m *MockTime) NewTimer(arg0 time.Duration) *time.Timer { } // NewTimer indicates an expected call of NewTimer. -func (mr *MockTimeMockRecorder) NewTimer(arg0 interface{}) *gomock.Call { +func (mr *MockTimeMockRecorder) NewTimer(arg0 any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "NewTimer", reflect.TypeOf((*MockTime)(nil).NewTimer), arg0) } diff --git a/mocks/mock_upgrade/upgrade.go b/mocks/mock_upgrade/upgrade.go index 70de7a7..6f4faff 100644 --- a/mocks/mock_upgrade/upgrade.go +++ b/mocks/mock_upgrade/upgrade.go @@ -1,5 +1,10 @@ // Code generated by MockGen. DO NOT EDIT. // Source: ./cli/cage/upgrade/upgrade.go +// +// Generated by this command: +// +// mockgen -source=./cli/cage/upgrade/upgrade.go +// // Package mock_upgrade is a generated GoMock package. package mock_upgrade @@ -7,14 +12,15 @@ package mock_upgrade import ( reflect "reflect" - gomock "go.uber.org/mock/gomock" upgrade "github.com/loilo-inc/canarycage/cli/cage/upgrade" + gomock "go.uber.org/mock/gomock" ) // MockUpgrader is a mock of Upgrader interface. type MockUpgrader struct { ctrl *gomock.Controller recorder *MockUpgraderMockRecorder + isgomock struct{} } // MockUpgraderMockRecorder is the mock recorder for MockUpgrader. @@ -43,7 +49,7 @@ func (m *MockUpgrader) Upgrade(p *upgrade.Input) error { } // Upgrade indicates an expected call of Upgrade. -func (mr *MockUpgraderMockRecorder) Upgrade(p interface{}) *gomock.Call { +func (mr *MockUpgraderMockRecorder) Upgrade(p any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Upgrade", reflect.TypeOf((*MockUpgrader)(nil).Upgrade), p) } diff --git a/test/fake_timer.go b/test/fake_timer.go index 111c5d8..a758da4 100644 --- a/test/fake_timer.go +++ b/test/fake_timer.go @@ -11,12 +11,12 @@ func newTimer(_ time.Duration) *time.Timer { go func() { ch <- time.Now() }() - return &time.Timer{ - C: ch, - } + return &time.Timer{C: ch} } -type timeImpl struct{} +type timeImpl struct { + never bool +} func (t *timeImpl) Now() time.Time { return time.Now() @@ -24,6 +24,18 @@ func (t *timeImpl) Now() time.Time { func (t *timeImpl) NewTimer(d time.Duration) *time.Timer { return newTimer(d) } + func NewFakeTime() types.Time { return &timeImpl{} } + +type neverTimer struct{} + +func (t *neverTimer) NewTimer(d time.Duration) *time.Timer { + ch := make(chan time.Time) + return &time.Timer{C: ch} +} + +func NewFakeNeverTimer() types.Time { + return &neverTimer{} +} diff --git a/types/iface.go b/types/iface.go index e4ca1f0..37c6707 100644 --- a/types/iface.go +++ b/types/iface.go @@ -13,6 +13,10 @@ type Cage interface { RollOut(ctx context.Context, input *RollOutInput) (*RollOutResult, error) } +type Audit interface { + Run(ctx context.Context) error +} + type Time interface { NewTimer(time.Duration) *time.Timer }