diff --git a/.mockery.yaml b/.mockery.yaml index 93145a4a..c6b93a78 100644 --- a/.mockery.yaml +++ b/.mockery.yaml @@ -6,5 +6,6 @@ pkgname: mocks recursive: true packages: go.woodpecker-ci.org/autoscaler/providers/hetznercloud/hcapi: + go.woodpecker-ci.org/autoscaler/providers/aws/ec2api: go.woodpecker-ci.org/autoscaler/engine/types: go.woodpecker-ci.org/autoscaler/server: diff --git a/go.mod b/go.mod index b91ba13b..a0022776 100644 --- a/go.mod +++ b/go.mod @@ -7,7 +7,10 @@ toolchain go1.26.5 require ( github.com/aws/aws-sdk-go-v2 v1.42.1 github.com/aws/aws-sdk-go-v2/config v1.32.30 + github.com/aws/aws-sdk-go-v2/credentials v1.19.29 github.com/aws/aws-sdk-go-v2/service/ec2 v1.316.1 + github.com/aws/aws-sdk-go-v2/service/sts v1.44.1 + github.com/aws/smithy-go v1.27.3 github.com/docker/go-units v0.5.0 github.com/equinix/equinix-sdk-go v0.66.0 github.com/gophercloud/gophercloud/v2 v2.13.0 @@ -26,7 +29,6 @@ require ( ) require ( - github.com/aws/aws-sdk-go-v2/credentials v1.19.29 // indirect github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.30 // indirect github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.30 // indirect github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.30 // indirect @@ -36,8 +38,6 @@ require ( github.com/aws/aws-sdk-go-v2/service/signin v1.4.1 // indirect github.com/aws/aws-sdk-go-v2/service/sso v1.32.1 // indirect github.com/aws/aws-sdk-go-v2/service/ssooidc v1.37.1 // indirect - github.com/aws/aws-sdk-go-v2/service/sts v1.44.1 // indirect - github.com/aws/smithy-go v1.27.3 // indirect github.com/beorn7/perks v1.0.1 // indirect github.com/cespare/xxhash/v2 v2.3.0 // indirect github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect diff --git a/providers/aws/credentials_test.go b/providers/aws/credentials_test.go new file mode 100644 index 00000000..b1db8adf --- /dev/null +++ b/providers/aws/credentials_test.go @@ -0,0 +1,125 @@ +package aws + +import ( + "bytes" + "context" + "errors" + "testing" + + "github.com/aws/aws-sdk-go-v2/aws" + awsconfig "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/service/sts" + "github.com/rs/zerolog" + "github.com/rs/zerolog/log" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/urfave/cli/v3" +) + +type fakeIdentityClient struct { + output *sts.GetCallerIdentityOutput + err error +} + +func (f fakeIdentityClient) GetCallerIdentity(context.Context, *sts.GetCallerIdentityInput, ...func(*sts.Options)) (*sts.GetCallerIdentityOutput, error) { + return f.output, f.err +} + +func TestCredentialFlagsUseStandardAWSEnvironment(t *testing.T) { + t.Setenv("AWS_ACCESS_KEY_ID", "test-access-key") + t.Setenv("AWS_SECRET_ACCESS_KEY", "test-secret-key") + + cmd := &cli.Command{ + Flags: ProviderFlags, + Action: func(_ context.Context, cmd *cli.Command) error { + assert.Equal(t, "test-access-key", cmd.String("aws-access-key-id")) + assert.Equal(t, "test-secret-key", cmd.String("aws-secret-access-key")) + return nil + }, + } + + require.NoError(t, cmd.Run(t.Context(), []string{"autoscaler"})) +} + +func TestStaticCredentialsOption(t *testing.T) { + t.Run("omitted", func(t *testing.T) { + option, err := staticCredentialsOption("", "") + require.NoError(t, err) + assert.Nil(t, option) + }) + + t.Run("missing access key ID", func(t *testing.T) { + _, err := staticCredentialsOption("", "secret") + assert.ErrorContains(t, err, "aws-access-key-id") + }) + + t.Run("missing secret access key", func(t *testing.T) { + _, err := staticCredentialsOption("access", "") + assert.ErrorContains(t, err, "aws-secret-access-key") + }) + + t.Run("forwarded to SDK", func(t *testing.T) { + option, err := staticCredentialsOption("test-access-key", "test-secret-key") + require.NoError(t, err) + require.NotNil(t, option) + + var options awsconfig.LoadOptions + require.NoError(t, option(&options)) + require.NotNil(t, options.Credentials) + + credentials, err := options.Credentials.Retrieve(t.Context()) + require.NoError(t, err) + assert.Equal(t, "test-access-key", credentials.AccessKeyID) + assert.Equal(t, "test-secret-key", credentials.SecretAccessKey) + }) +} + +func TestAuthenticateAWS(t *testing.T) { + t.Run("success", func(t *testing.T) { + var output bytes.Buffer + originalLogger := log.Logger + log.Logger = zerolog.New(&output) + t.Cleanup(func() { + log.Logger = originalLogger + }) + + err := authenticateAWS(t.Context(), fakeIdentityClient{output: &sts.GetCallerIdentityOutput{ + Account: aws.String("123456789012"), + Arn: aws.String("arn:aws:iam::123456789012:user/autoscaler"), + }}) + require.NoError(t, err) + assert.JSONEq(t, `{ + "level":"info", + "account":"123456789012", + "arn":"arn:aws:iam::123456789012:user/autoscaler", + "message":"authenticated with AWS" + }`, output.String()) + }) + + t.Run("failure", func(t *testing.T) { + err := authenticateAWS(t.Context(), fakeIdentityClient{err: errors.New("invalid credentials")}) + assert.ErrorContains(t, err, "authenticate with AWS: invalid credentials") + }) +} + +func TestAuthenticationRegion(t *testing.T) { + t.Run("default region", func(t *testing.T) { + p := provider{name: "aws", region: "eu-central-1"} + region, err := p.authenticationRegion(nil) + require.NoError(t, err) + assert.Equal(t, "eu-central-1", region) + }) + + t.Run("qualified instance type", func(t *testing.T) { + p := provider{name: "aws"} + region, err := p.authenticationRegion([]string{"t3.micro:us-east-1"}) + require.NoError(t, err) + assert.Equal(t, "us-east-1", region) + }) + + t.Run("missing region", func(t *testing.T) { + p := provider{name: "aws"} + _, err := p.authenticationRegion([]string{"t3.micro"}) + assert.ErrorIs(t, err, ErrRegionNotSet) + }) +} diff --git a/providers/aws/ec2api/api.go b/providers/aws/ec2api/api.go new file mode 100644 index 00000000..77365224 --- /dev/null +++ b/providers/aws/ec2api/api.go @@ -0,0 +1,19 @@ +package ec2api + +import ( + "context" + + "github.com/aws/aws-sdk-go-v2/service/ec2" +) + +// Client is the subset of the EC2 API the aws provider uses, so it can be +// mocked in tests. *ec2.Client satisfies it. +type Client interface { + DescribeImages(ctx context.Context, params *ec2.DescribeImagesInput, optFns ...func(*ec2.Options)) (*ec2.DescribeImagesOutput, error) + DescribeInstances(ctx context.Context, params *ec2.DescribeInstancesInput, optFns ...func(*ec2.Options)) (*ec2.DescribeInstancesOutput, error) + DescribeInstanceTypes(ctx context.Context, params *ec2.DescribeInstanceTypesInput, optFns ...func(*ec2.Options)) (*ec2.DescribeInstanceTypesOutput, error) + DescribeSecurityGroups(ctx context.Context, params *ec2.DescribeSecurityGroupsInput, optFns ...func(*ec2.Options)) (*ec2.DescribeSecurityGroupsOutput, error) + DescribeSubnets(ctx context.Context, params *ec2.DescribeSubnetsInput, optFns ...func(*ec2.Options)) (*ec2.DescribeSubnetsOutput, error) + RunInstances(ctx context.Context, params *ec2.RunInstancesInput, optFns ...func(*ec2.Options)) (*ec2.RunInstancesOutput, error) + TerminateInstances(ctx context.Context, params *ec2.TerminateInstancesInput, optFns ...func(*ec2.Options)) (*ec2.TerminateInstancesOutput, error) +} diff --git a/providers/aws/ec2api/mocks/mock_Client.go b/providers/aws/ec2api/mocks/mock_Client.go new file mode 100644 index 00000000..0da9bb12 --- /dev/null +++ b/providers/aws/ec2api/mocks/mock_Client.go @@ -0,0 +1,620 @@ +// Code generated by mockery; DO NOT EDIT. +// github.com/vektra/mockery +// template: testify + +package mocks + +import ( + "context" + + "github.com/aws/aws-sdk-go-v2/service/ec2" + mock "github.com/stretchr/testify/mock" +) + +// NewMockClient creates a new instance of MockClient. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func NewMockClient(t interface { + mock.TestingT + Cleanup(func()) +}) *MockClient { + mock := &MockClient{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} + +// MockClient is an autogenerated mock type for the Client type +type MockClient struct { + mock.Mock +} + +type MockClient_Expecter struct { + mock *mock.Mock +} + +func (_m *MockClient) EXPECT() *MockClient_Expecter { + return &MockClient_Expecter{mock: &_m.Mock} +} + +// DescribeImages provides a mock function for the type MockClient +func (_mock *MockClient) DescribeImages(ctx context.Context, params *ec2.DescribeImagesInput, optFns ...func(*ec2.Options)) (*ec2.DescribeImagesOutput, error) { + var tmpRet mock.Arguments + if len(optFns) > 0 { + tmpRet = _mock.Called(ctx, params, optFns) + } else { + tmpRet = _mock.Called(ctx, params) + } + ret := tmpRet + + if len(ret) == 0 { + panic("no return value specified for DescribeImages") + } + + var r0 *ec2.DescribeImagesOutput + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context, *ec2.DescribeImagesInput, ...func(*ec2.Options)) (*ec2.DescribeImagesOutput, error)); ok { + return returnFunc(ctx, params, optFns...) + } + if returnFunc, ok := ret.Get(0).(func(context.Context, *ec2.DescribeImagesInput, ...func(*ec2.Options)) *ec2.DescribeImagesOutput); ok { + r0 = returnFunc(ctx, params, optFns...) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*ec2.DescribeImagesOutput) + } + } + if returnFunc, ok := ret.Get(1).(func(context.Context, *ec2.DescribeImagesInput, ...func(*ec2.Options)) error); ok { + r1 = returnFunc(ctx, params, optFns...) + } else { + r1 = ret.Error(1) + } + return r0, r1 +} + +// MockClient_DescribeImages_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DescribeImages' +type MockClient_DescribeImages_Call struct { + *mock.Call +} + +// DescribeImages is a helper method to define mock.On call +// - ctx context.Context +// - params *ec2.DescribeImagesInput +// - optFns ...func(*ec2.Options) +func (_e *MockClient_Expecter) DescribeImages(ctx interface{}, params interface{}, optFns ...interface{}) *MockClient_DescribeImages_Call { + return &MockClient_DescribeImages_Call{Call: _e.mock.On("DescribeImages", + append([]interface{}{ctx, params}, optFns...)...)} +} + +func (_c *MockClient_DescribeImages_Call) Run(run func(ctx context.Context, params *ec2.DescribeImagesInput, optFns ...func(*ec2.Options))) *MockClient_DescribeImages_Call { + _c.Call.Run(func(args mock.Arguments) { + var arg0 context.Context + if args[0] != nil { + arg0 = args[0].(context.Context) + } + var arg1 *ec2.DescribeImagesInput + if args[1] != nil { + arg1 = args[1].(*ec2.DescribeImagesInput) + } + var arg2 []func(*ec2.Options) + var variadicArgs []func(*ec2.Options) + if len(args) > 2 { + variadicArgs = args[2].([]func(*ec2.Options)) + } + arg2 = variadicArgs + run( + arg0, + arg1, + arg2..., + ) + }) + return _c +} + +func (_c *MockClient_DescribeImages_Call) Return(describeImagesOutput *ec2.DescribeImagesOutput, err error) *MockClient_DescribeImages_Call { + _c.Call.Return(describeImagesOutput, err) + return _c +} + +func (_c *MockClient_DescribeImages_Call) RunAndReturn(run func(ctx context.Context, params *ec2.DescribeImagesInput, optFns ...func(*ec2.Options)) (*ec2.DescribeImagesOutput, error)) *MockClient_DescribeImages_Call { + _c.Call.Return(run) + return _c +} + +// DescribeInstanceTypes provides a mock function for the type MockClient +func (_mock *MockClient) DescribeInstanceTypes(ctx context.Context, params *ec2.DescribeInstanceTypesInput, optFns ...func(*ec2.Options)) (*ec2.DescribeInstanceTypesOutput, error) { + var tmpRet mock.Arguments + if len(optFns) > 0 { + tmpRet = _mock.Called(ctx, params, optFns) + } else { + tmpRet = _mock.Called(ctx, params) + } + ret := tmpRet + + if len(ret) == 0 { + panic("no return value specified for DescribeInstanceTypes") + } + + var r0 *ec2.DescribeInstanceTypesOutput + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context, *ec2.DescribeInstanceTypesInput, ...func(*ec2.Options)) (*ec2.DescribeInstanceTypesOutput, error)); ok { + return returnFunc(ctx, params, optFns...) + } + if returnFunc, ok := ret.Get(0).(func(context.Context, *ec2.DescribeInstanceTypesInput, ...func(*ec2.Options)) *ec2.DescribeInstanceTypesOutput); ok { + r0 = returnFunc(ctx, params, optFns...) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*ec2.DescribeInstanceTypesOutput) + } + } + if returnFunc, ok := ret.Get(1).(func(context.Context, *ec2.DescribeInstanceTypesInput, ...func(*ec2.Options)) error); ok { + r1 = returnFunc(ctx, params, optFns...) + } else { + r1 = ret.Error(1) + } + return r0, r1 +} + +// MockClient_DescribeInstanceTypes_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DescribeInstanceTypes' +type MockClient_DescribeInstanceTypes_Call struct { + *mock.Call +} + +// DescribeInstanceTypes is a helper method to define mock.On call +// - ctx context.Context +// - params *ec2.DescribeInstanceTypesInput +// - optFns ...func(*ec2.Options) +func (_e *MockClient_Expecter) DescribeInstanceTypes(ctx interface{}, params interface{}, optFns ...interface{}) *MockClient_DescribeInstanceTypes_Call { + return &MockClient_DescribeInstanceTypes_Call{Call: _e.mock.On("DescribeInstanceTypes", + append([]interface{}{ctx, params}, optFns...)...)} +} + +func (_c *MockClient_DescribeInstanceTypes_Call) Run(run func(ctx context.Context, params *ec2.DescribeInstanceTypesInput, optFns ...func(*ec2.Options))) *MockClient_DescribeInstanceTypes_Call { + _c.Call.Run(func(args mock.Arguments) { + var arg0 context.Context + if args[0] != nil { + arg0 = args[0].(context.Context) + } + var arg1 *ec2.DescribeInstanceTypesInput + if args[1] != nil { + arg1 = args[1].(*ec2.DescribeInstanceTypesInput) + } + var arg2 []func(*ec2.Options) + var variadicArgs []func(*ec2.Options) + if len(args) > 2 { + variadicArgs = args[2].([]func(*ec2.Options)) + } + arg2 = variadicArgs + run( + arg0, + arg1, + arg2..., + ) + }) + return _c +} + +func (_c *MockClient_DescribeInstanceTypes_Call) Return(describeInstanceTypesOutput *ec2.DescribeInstanceTypesOutput, err error) *MockClient_DescribeInstanceTypes_Call { + _c.Call.Return(describeInstanceTypesOutput, err) + return _c +} + +func (_c *MockClient_DescribeInstanceTypes_Call) RunAndReturn(run func(ctx context.Context, params *ec2.DescribeInstanceTypesInput, optFns ...func(*ec2.Options)) (*ec2.DescribeInstanceTypesOutput, error)) *MockClient_DescribeInstanceTypes_Call { + _c.Call.Return(run) + return _c +} + +// DescribeInstances provides a mock function for the type MockClient +func (_mock *MockClient) DescribeInstances(ctx context.Context, params *ec2.DescribeInstancesInput, optFns ...func(*ec2.Options)) (*ec2.DescribeInstancesOutput, error) { + var tmpRet mock.Arguments + if len(optFns) > 0 { + tmpRet = _mock.Called(ctx, params, optFns) + } else { + tmpRet = _mock.Called(ctx, params) + } + ret := tmpRet + + if len(ret) == 0 { + panic("no return value specified for DescribeInstances") + } + + var r0 *ec2.DescribeInstancesOutput + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context, *ec2.DescribeInstancesInput, ...func(*ec2.Options)) (*ec2.DescribeInstancesOutput, error)); ok { + return returnFunc(ctx, params, optFns...) + } + if returnFunc, ok := ret.Get(0).(func(context.Context, *ec2.DescribeInstancesInput, ...func(*ec2.Options)) *ec2.DescribeInstancesOutput); ok { + r0 = returnFunc(ctx, params, optFns...) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*ec2.DescribeInstancesOutput) + } + } + if returnFunc, ok := ret.Get(1).(func(context.Context, *ec2.DescribeInstancesInput, ...func(*ec2.Options)) error); ok { + r1 = returnFunc(ctx, params, optFns...) + } else { + r1 = ret.Error(1) + } + return r0, r1 +} + +// MockClient_DescribeInstances_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DescribeInstances' +type MockClient_DescribeInstances_Call struct { + *mock.Call +} + +// DescribeInstances is a helper method to define mock.On call +// - ctx context.Context +// - params *ec2.DescribeInstancesInput +// - optFns ...func(*ec2.Options) +func (_e *MockClient_Expecter) DescribeInstances(ctx interface{}, params interface{}, optFns ...interface{}) *MockClient_DescribeInstances_Call { + return &MockClient_DescribeInstances_Call{Call: _e.mock.On("DescribeInstances", + append([]interface{}{ctx, params}, optFns...)...)} +} + +func (_c *MockClient_DescribeInstances_Call) Run(run func(ctx context.Context, params *ec2.DescribeInstancesInput, optFns ...func(*ec2.Options))) *MockClient_DescribeInstances_Call { + _c.Call.Run(func(args mock.Arguments) { + var arg0 context.Context + if args[0] != nil { + arg0 = args[0].(context.Context) + } + var arg1 *ec2.DescribeInstancesInput + if args[1] != nil { + arg1 = args[1].(*ec2.DescribeInstancesInput) + } + var arg2 []func(*ec2.Options) + var variadicArgs []func(*ec2.Options) + if len(args) > 2 { + variadicArgs = args[2].([]func(*ec2.Options)) + } + arg2 = variadicArgs + run( + arg0, + arg1, + arg2..., + ) + }) + return _c +} + +func (_c *MockClient_DescribeInstances_Call) Return(describeInstancesOutput *ec2.DescribeInstancesOutput, err error) *MockClient_DescribeInstances_Call { + _c.Call.Return(describeInstancesOutput, err) + return _c +} + +func (_c *MockClient_DescribeInstances_Call) RunAndReturn(run func(ctx context.Context, params *ec2.DescribeInstancesInput, optFns ...func(*ec2.Options)) (*ec2.DescribeInstancesOutput, error)) *MockClient_DescribeInstances_Call { + _c.Call.Return(run) + return _c +} + +// DescribeSecurityGroups provides a mock function for the type MockClient +func (_mock *MockClient) DescribeSecurityGroups(ctx context.Context, params *ec2.DescribeSecurityGroupsInput, optFns ...func(*ec2.Options)) (*ec2.DescribeSecurityGroupsOutput, error) { + var tmpRet mock.Arguments + if len(optFns) > 0 { + tmpRet = _mock.Called(ctx, params, optFns) + } else { + tmpRet = _mock.Called(ctx, params) + } + ret := tmpRet + + if len(ret) == 0 { + panic("no return value specified for DescribeSecurityGroups") + } + + var r0 *ec2.DescribeSecurityGroupsOutput + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context, *ec2.DescribeSecurityGroupsInput, ...func(*ec2.Options)) (*ec2.DescribeSecurityGroupsOutput, error)); ok { + return returnFunc(ctx, params, optFns...) + } + if returnFunc, ok := ret.Get(0).(func(context.Context, *ec2.DescribeSecurityGroupsInput, ...func(*ec2.Options)) *ec2.DescribeSecurityGroupsOutput); ok { + r0 = returnFunc(ctx, params, optFns...) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*ec2.DescribeSecurityGroupsOutput) + } + } + if returnFunc, ok := ret.Get(1).(func(context.Context, *ec2.DescribeSecurityGroupsInput, ...func(*ec2.Options)) error); ok { + r1 = returnFunc(ctx, params, optFns...) + } else { + r1 = ret.Error(1) + } + return r0, r1 +} + +// MockClient_DescribeSecurityGroups_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DescribeSecurityGroups' +type MockClient_DescribeSecurityGroups_Call struct { + *mock.Call +} + +// DescribeSecurityGroups is a helper method to define mock.On call +// - ctx context.Context +// - params *ec2.DescribeSecurityGroupsInput +// - optFns ...func(*ec2.Options) +func (_e *MockClient_Expecter) DescribeSecurityGroups(ctx interface{}, params interface{}, optFns ...interface{}) *MockClient_DescribeSecurityGroups_Call { + return &MockClient_DescribeSecurityGroups_Call{Call: _e.mock.On("DescribeSecurityGroups", + append([]interface{}{ctx, params}, optFns...)...)} +} + +func (_c *MockClient_DescribeSecurityGroups_Call) Run(run func(ctx context.Context, params *ec2.DescribeSecurityGroupsInput, optFns ...func(*ec2.Options))) *MockClient_DescribeSecurityGroups_Call { + _c.Call.Run(func(args mock.Arguments) { + var arg0 context.Context + if args[0] != nil { + arg0 = args[0].(context.Context) + } + var arg1 *ec2.DescribeSecurityGroupsInput + if args[1] != nil { + arg1 = args[1].(*ec2.DescribeSecurityGroupsInput) + } + var arg2 []func(*ec2.Options) + var variadicArgs []func(*ec2.Options) + if len(args) > 2 { + variadicArgs = args[2].([]func(*ec2.Options)) + } + arg2 = variadicArgs + run( + arg0, + arg1, + arg2..., + ) + }) + return _c +} + +func (_c *MockClient_DescribeSecurityGroups_Call) Return(describeSecurityGroupsOutput *ec2.DescribeSecurityGroupsOutput, err error) *MockClient_DescribeSecurityGroups_Call { + _c.Call.Return(describeSecurityGroupsOutput, err) + return _c +} + +func (_c *MockClient_DescribeSecurityGroups_Call) RunAndReturn(run func(ctx context.Context, params *ec2.DescribeSecurityGroupsInput, optFns ...func(*ec2.Options)) (*ec2.DescribeSecurityGroupsOutput, error)) *MockClient_DescribeSecurityGroups_Call { + _c.Call.Return(run) + return _c +} + +// DescribeSubnets provides a mock function for the type MockClient +func (_mock *MockClient) DescribeSubnets(ctx context.Context, params *ec2.DescribeSubnetsInput, optFns ...func(*ec2.Options)) (*ec2.DescribeSubnetsOutput, error) { + var tmpRet mock.Arguments + if len(optFns) > 0 { + tmpRet = _mock.Called(ctx, params, optFns) + } else { + tmpRet = _mock.Called(ctx, params) + } + ret := tmpRet + + if len(ret) == 0 { + panic("no return value specified for DescribeSubnets") + } + + var r0 *ec2.DescribeSubnetsOutput + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context, *ec2.DescribeSubnetsInput, ...func(*ec2.Options)) (*ec2.DescribeSubnetsOutput, error)); ok { + return returnFunc(ctx, params, optFns...) + } + if returnFunc, ok := ret.Get(0).(func(context.Context, *ec2.DescribeSubnetsInput, ...func(*ec2.Options)) *ec2.DescribeSubnetsOutput); ok { + r0 = returnFunc(ctx, params, optFns...) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*ec2.DescribeSubnetsOutput) + } + } + if returnFunc, ok := ret.Get(1).(func(context.Context, *ec2.DescribeSubnetsInput, ...func(*ec2.Options)) error); ok { + r1 = returnFunc(ctx, params, optFns...) + } else { + r1 = ret.Error(1) + } + return r0, r1 +} + +// MockClient_DescribeSubnets_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DescribeSubnets' +type MockClient_DescribeSubnets_Call struct { + *mock.Call +} + +// DescribeSubnets is a helper method to define mock.On call +// - ctx context.Context +// - params *ec2.DescribeSubnetsInput +// - optFns ...func(*ec2.Options) +func (_e *MockClient_Expecter) DescribeSubnets(ctx interface{}, params interface{}, optFns ...interface{}) *MockClient_DescribeSubnets_Call { + return &MockClient_DescribeSubnets_Call{Call: _e.mock.On("DescribeSubnets", + append([]interface{}{ctx, params}, optFns...)...)} +} + +func (_c *MockClient_DescribeSubnets_Call) Run(run func(ctx context.Context, params *ec2.DescribeSubnetsInput, optFns ...func(*ec2.Options))) *MockClient_DescribeSubnets_Call { + _c.Call.Run(func(args mock.Arguments) { + var arg0 context.Context + if args[0] != nil { + arg0 = args[0].(context.Context) + } + var arg1 *ec2.DescribeSubnetsInput + if args[1] != nil { + arg1 = args[1].(*ec2.DescribeSubnetsInput) + } + var arg2 []func(*ec2.Options) + var variadicArgs []func(*ec2.Options) + if len(args) > 2 { + variadicArgs = args[2].([]func(*ec2.Options)) + } + arg2 = variadicArgs + run( + arg0, + arg1, + arg2..., + ) + }) + return _c +} + +func (_c *MockClient_DescribeSubnets_Call) Return(describeSubnetsOutput *ec2.DescribeSubnetsOutput, err error) *MockClient_DescribeSubnets_Call { + _c.Call.Return(describeSubnetsOutput, err) + return _c +} + +func (_c *MockClient_DescribeSubnets_Call) RunAndReturn(run func(ctx context.Context, params *ec2.DescribeSubnetsInput, optFns ...func(*ec2.Options)) (*ec2.DescribeSubnetsOutput, error)) *MockClient_DescribeSubnets_Call { + _c.Call.Return(run) + return _c +} + +// RunInstances provides a mock function for the type MockClient +func (_mock *MockClient) RunInstances(ctx context.Context, params *ec2.RunInstancesInput, optFns ...func(*ec2.Options)) (*ec2.RunInstancesOutput, error) { + var tmpRet mock.Arguments + if len(optFns) > 0 { + tmpRet = _mock.Called(ctx, params, optFns) + } else { + tmpRet = _mock.Called(ctx, params) + } + ret := tmpRet + + if len(ret) == 0 { + panic("no return value specified for RunInstances") + } + + var r0 *ec2.RunInstancesOutput + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context, *ec2.RunInstancesInput, ...func(*ec2.Options)) (*ec2.RunInstancesOutput, error)); ok { + return returnFunc(ctx, params, optFns...) + } + if returnFunc, ok := ret.Get(0).(func(context.Context, *ec2.RunInstancesInput, ...func(*ec2.Options)) *ec2.RunInstancesOutput); ok { + r0 = returnFunc(ctx, params, optFns...) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*ec2.RunInstancesOutput) + } + } + if returnFunc, ok := ret.Get(1).(func(context.Context, *ec2.RunInstancesInput, ...func(*ec2.Options)) error); ok { + r1 = returnFunc(ctx, params, optFns...) + } else { + r1 = ret.Error(1) + } + return r0, r1 +} + +// MockClient_RunInstances_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RunInstances' +type MockClient_RunInstances_Call struct { + *mock.Call +} + +// RunInstances is a helper method to define mock.On call +// - ctx context.Context +// - params *ec2.RunInstancesInput +// - optFns ...func(*ec2.Options) +func (_e *MockClient_Expecter) RunInstances(ctx interface{}, params interface{}, optFns ...interface{}) *MockClient_RunInstances_Call { + return &MockClient_RunInstances_Call{Call: _e.mock.On("RunInstances", + append([]interface{}{ctx, params}, optFns...)...)} +} + +func (_c *MockClient_RunInstances_Call) Run(run func(ctx context.Context, params *ec2.RunInstancesInput, optFns ...func(*ec2.Options))) *MockClient_RunInstances_Call { + _c.Call.Run(func(args mock.Arguments) { + var arg0 context.Context + if args[0] != nil { + arg0 = args[0].(context.Context) + } + var arg1 *ec2.RunInstancesInput + if args[1] != nil { + arg1 = args[1].(*ec2.RunInstancesInput) + } + var arg2 []func(*ec2.Options) + var variadicArgs []func(*ec2.Options) + if len(args) > 2 { + variadicArgs = args[2].([]func(*ec2.Options)) + } + arg2 = variadicArgs + run( + arg0, + arg1, + arg2..., + ) + }) + return _c +} + +func (_c *MockClient_RunInstances_Call) Return(runInstancesOutput *ec2.RunInstancesOutput, err error) *MockClient_RunInstances_Call { + _c.Call.Return(runInstancesOutput, err) + return _c +} + +func (_c *MockClient_RunInstances_Call) RunAndReturn(run func(ctx context.Context, params *ec2.RunInstancesInput, optFns ...func(*ec2.Options)) (*ec2.RunInstancesOutput, error)) *MockClient_RunInstances_Call { + _c.Call.Return(run) + return _c +} + +// TerminateInstances provides a mock function for the type MockClient +func (_mock *MockClient) TerminateInstances(ctx context.Context, params *ec2.TerminateInstancesInput, optFns ...func(*ec2.Options)) (*ec2.TerminateInstancesOutput, error) { + var tmpRet mock.Arguments + if len(optFns) > 0 { + tmpRet = _mock.Called(ctx, params, optFns) + } else { + tmpRet = _mock.Called(ctx, params) + } + ret := tmpRet + + if len(ret) == 0 { + panic("no return value specified for TerminateInstances") + } + + var r0 *ec2.TerminateInstancesOutput + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context, *ec2.TerminateInstancesInput, ...func(*ec2.Options)) (*ec2.TerminateInstancesOutput, error)); ok { + return returnFunc(ctx, params, optFns...) + } + if returnFunc, ok := ret.Get(0).(func(context.Context, *ec2.TerminateInstancesInput, ...func(*ec2.Options)) *ec2.TerminateInstancesOutput); ok { + r0 = returnFunc(ctx, params, optFns...) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*ec2.TerminateInstancesOutput) + } + } + if returnFunc, ok := ret.Get(1).(func(context.Context, *ec2.TerminateInstancesInput, ...func(*ec2.Options)) error); ok { + r1 = returnFunc(ctx, params, optFns...) + } else { + r1 = ret.Error(1) + } + return r0, r1 +} + +// MockClient_TerminateInstances_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'TerminateInstances' +type MockClient_TerminateInstances_Call struct { + *mock.Call +} + +// TerminateInstances is a helper method to define mock.On call +// - ctx context.Context +// - params *ec2.TerminateInstancesInput +// - optFns ...func(*ec2.Options) +func (_e *MockClient_Expecter) TerminateInstances(ctx interface{}, params interface{}, optFns ...interface{}) *MockClient_TerminateInstances_Call { + return &MockClient_TerminateInstances_Call{Call: _e.mock.On("TerminateInstances", + append([]interface{}{ctx, params}, optFns...)...)} +} + +func (_c *MockClient_TerminateInstances_Call) Run(run func(ctx context.Context, params *ec2.TerminateInstancesInput, optFns ...func(*ec2.Options))) *MockClient_TerminateInstances_Call { + _c.Call.Run(func(args mock.Arguments) { + var arg0 context.Context + if args[0] != nil { + arg0 = args[0].(context.Context) + } + var arg1 *ec2.TerminateInstancesInput + if args[1] != nil { + arg1 = args[1].(*ec2.TerminateInstancesInput) + } + var arg2 []func(*ec2.Options) + var variadicArgs []func(*ec2.Options) + if len(args) > 2 { + variadicArgs = args[2].([]func(*ec2.Options)) + } + arg2 = variadicArgs + run( + arg0, + arg1, + arg2..., + ) + }) + return _c +} + +func (_c *MockClient_TerminateInstances_Call) Return(terminateInstancesOutput *ec2.TerminateInstancesOutput, err error) *MockClient_TerminateInstances_Call { + _c.Call.Return(terminateInstancesOutput, err) + return _c +} + +func (_c *MockClient_TerminateInstances_Call) RunAndReturn(run func(ctx context.Context, params *ec2.TerminateInstancesInput, optFns ...func(*ec2.Options)) (*ec2.TerminateInstancesOutput, error)) *MockClient_TerminateInstances_Call { + _c.Call.Return(run) + return _c +} diff --git a/providers/aws/flags.go b/providers/aws/flags.go index ed970854..00add895 100644 --- a/providers/aws/flags.go +++ b/providers/aws/flags.go @@ -6,15 +6,17 @@ const Category = "AWS" var ProviderFlags = []cli.Flag{ // aws - &cli.StringFlag{ + &cli.StringSliceFlag{ Name: "aws-instance-type", - Usage: "EC2 instance type", + Usage: "EC2 instance types, optionally with region as 'type:region'; tried in order as deploy fallbacks", + Value: []string{"t3.medium", "t3.micro"}, Sources: cli.EnvVars("WOODPECKER_AWS_INSTANCE_TYPE"), Category: Category, }, &cli.StringFlag{ Name: "aws-ami-id", - Usage: "AMI id", + Usage: "AMI ID or alias (ubuntu--server, amazon, suse[-], debian-); architecture and region come from the instance type", + Value: "ubuntu-26.04-server", Sources: cli.EnvVars("WOODPECKER_AWS_AMI_ID"), Category: Category, }, @@ -24,15 +26,28 @@ var ProviderFlags = []cli.Flag{ Sources: cli.EnvVars("WOODPECKER_AWS_TAGS"), Category: Category, }, + &cli.StringFlag{ + Name: "aws-access-key-id", + Usage: "AWS access key ID", + Sources: cli.EnvVars("WOODPECKER_AWS_ACCESS_KEY_ID", "AWS_ACCESS_KEY_ID"), + Category: Category, + }, + &cli.StringFlag{ + Name: "aws-secret-access-key", + Usage: "AWS secret access key", + Sources: cli.EnvVars("WOODPECKER_AWS_SECRET_ACCESS_KEY", "AWS_SECRET_ACCESS_KEY"), + Category: Category, + }, &cli.StringFlag{ Name: "aws-region", - Usage: "AWS region", + Usage: "default AWS region for unqualified instance types and resources", + Value: "us-east-1", Sources: cli.EnvVars("WOODPECKER_AWS_REGION"), Category: Category, }, &cli.StringSliceFlag{ Name: "aws-subnets", - Usage: "VPC subnets IDs, e.g. subnet-0987a87c8b37348ef", + Usage: "VPC subnet IDs, optionally with region as 'subnet:region'; default subnets are used when omitted", Sources: cli.EnvVars("WOODPECKER_AWS_SUBNETS"), Category: Category, }, @@ -44,7 +59,7 @@ var ProviderFlags = []cli.Flag{ }, &cli.StringSliceFlag{ Name: "aws-security-groups", - Usage: "security groups attached to EC2 instances", + Usage: "security group IDs, optionally with region as 'group:region'", Sources: cli.EnvVars("WOODPECKER_AWS_SECURITY_GROUPS"), Category: Category, }, diff --git a/providers/aws/helper.go b/providers/aws/helper.go new file mode 100644 index 00000000..711e61c5 --- /dev/null +++ b/providers/aws/helper.go @@ -0,0 +1,313 @@ +package aws + +import ( + "context" + "errors" + "fmt" + "slices" + "sort" + "strings" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/ec2" + ec2_types "github.com/aws/aws-sdk-go-v2/service/ec2/types" + "github.com/aws/smithy-go" + "github.com/rs/zerolog/log" + + "go.woodpecker-ci.org/woodpecker/v3/woodpecker-go/woodpecker" +) + +// resolveDeployCandidates resolves the ordered instance type fallbacks. An +// instance type has the form value or value:region; unqualified values use +// aws-region. The AMI reference is resolved separately in each candidate's +// region, while subnets and security groups remain region-qualified resources. +func (p *provider) resolveDeployCandidates(ctx context.Context, instanceTypes []string, image string, subnets, securityGroups []string) error { + if len(instanceTypes) == 0 { + return fmt.Errorf("%s: %w", p.name, ErrNoDeployCandidates) + } + + subnetsByRegion, err := p.valuesByRegion("aws-subnets", subnets) + if err != nil { + return err + } + securityGroupsByRegion, err := p.valuesByRegion("aws-security-groups", securityGroups) + if err != nil { + return err + } + + type regionConfigKey struct { + region string + architecture ec2_types.ArchitectureValues + } + configs := map[regionConfigKey]regionConfig{} + imageValidated := false + for _, raw := range instanceTypes { + instanceType, region, err := p.valueRegion("aws-instance-type", raw) + if err != nil { + return err + } + + it, err := p.resolveInstanceType(ctx, instanceType, region) + if err != nil { + return err + } + if !imageValidated { + if err := validateImageReference(image); err != nil { + return fmt.Errorf("%s: %w", p.name, err) + } + imageValidated = true + } + + architecture, err := instanceTypeArchitecture(it) + if err != nil { + return fmt.Errorf("%s: %w", p.name, err) + } + key := regionConfigKey{region: region, architecture: architecture} + config, ok := configs[key] + if !ok { + config, err = p.resolveRegionConfig(ctx, region, image, architecture, subnetsByRegion[region], securityGroupsByRegion[region]) + if err != nil { + return err + } + configs[key] = config + } + + if !instanceTypeSupportsArch(it, config.image.Architecture) { + return fmt.Errorf("%s: %w: %s needs one of %v, AMI is %s", + p.name, ErrArchMismatch, it.InstanceType, + it.ProcessorInfo.SupportedArchitectures, config.image.Architecture) + } + + p.deployCandidates = append(p.deployCandidates, deployCandidate{ + instanceType: it, + regionConfig: config, + }) + if !slices.Contains(p.regions, region) { + p.regions = append(p.regions, region) + } + } + + return nil +} + +// valuesByRegion groups a flag's value or value:region entries by their +// effective region. +func (p *provider) valuesByRegion(option string, values []string) (map[string][]string, error) { + byRegion := make(map[string][]string, len(values)) + for _, raw := range values { + value, region, err := p.valueRegion(option, raw) + if err != nil { + return nil, err + } + byRegion[region] = append(byRegion[region], value) + } + return byRegion, nil +} + +func (p *provider) valueRegion(option, raw string) (string, string, error) { + value, region, qualified := strings.Cut(raw, ":") + if !qualified { + region = p.region + } + if value == "" { + return "", "", fmt.Errorf("%s: empty %s value", p.name, option) + } + if region == "" { + return "", "", fmt.Errorf("%s: %w: %s", p.name, ErrRegionNotSet, option) + } + return value, region, nil +} + +func (p *provider) resolveRegionConfig(ctx context.Context, region, image string, architecture ec2_types.ArchitectureValues, subnetIDs, securityGroupIDs []string) (regionConfig, error) { + resolvedImage, err := p.resolveImage(ctx, image, region, architecture) + if err != nil { + return regionConfig{}, err + } + + subnetIDs, err = p.resolveSubnets(ctx, region, subnetIDs) + if err != nil { + return regionConfig{}, err + } + + if len(securityGroupIDs) > 0 { + groups, err := p.client.DescribeSecurityGroups(ctx, &ec2.DescribeSecurityGroupsInput{ + GroupIds: securityGroupIDs, + }, regionOpt(region)) + if err != nil { + return regionConfig{}, fmt.Errorf("%s: DescribeSecurityGroups %v in %q: %w", + p.name, securityGroupIDs, region, err) + } + if len(groups.SecurityGroups) != len(securityGroupIDs) { + return regionConfig{}, fmt.Errorf("%s: %w: got %d of %d security groups in %q", + p.name, ErrSecurityGroupNotFound, len(groups.SecurityGroups), len(securityGroupIDs), region) + } + } + + return regionConfig{ + region: region, + image: resolvedImage, + subnets: subnetIDs, + securityGroups: securityGroupIDs, + }, nil +} + +func (p *provider) resolveSubnets(ctx context.Context, region string, subnetIDs []string) ([]string, error) { + input := &ec2.DescribeSubnetsInput{SubnetIds: subnetIDs} + useDefault := len(subnetIDs) == 0 + if useDefault { + input.Filters = []ec2_types.Filter{ + {Name: aws.String("default-for-az"), Values: []string{"true"}}, + {Name: aws.String("state"), Values: []string{"available"}}, + } + } + + resolved, err := p.client.DescribeSubnets(ctx, input, regionOpt(region)) + if err != nil { + return nil, fmt.Errorf("%s: DescribeSubnets %v in %q: %w", p.name, subnetIDs, region, err) + } + if !useDefault { + if len(resolved.Subnets) != len(subnetIDs) { + return nil, fmt.Errorf("%s: %w: got %d of %d subnets in %q", + p.name, ErrSubnetNotFound, len(resolved.Subnets), len(subnetIDs), region) + } + return subnetIDs, nil + } + + for _, subnet := range resolved.Subnets { + if subnet.SubnetId != nil && *subnet.SubnetId != "" { + subnetIDs = append(subnetIDs, *subnet.SubnetId) + } + } + if len(subnetIDs) == 0 { + return nil, fmt.Errorf("%s: %w: no default subnets in region %q", p.name, ErrSubnetsNotSet, region) + } + sort.Strings(subnetIDs) + log.Info(). + Str("region", region). + Strs("subnets", subnetIDs). + Msg("resolved default AWS subnets") + return subnetIDs, nil +} + +func regionOpt(region string) func(*ec2.Options) { + return func(o *ec2.Options) { + o.Region = region + } +} + +func (p *provider) resolveInstanceType(ctx context.Context, instanceType, region string) (ec2_types.InstanceTypeInfo, error) { + out, err := p.client.DescribeInstanceTypes(ctx, &ec2.DescribeInstanceTypesInput{ + InstanceTypes: []ec2_types.InstanceType{ec2_types.InstanceType(instanceType)}, + }, regionOpt(region)) + if err != nil { + // DescribeInstanceTypes only knows types offered in the queried + // region and rejects others with InvalidInstanceType. + var apiErr smithy.APIError + if errors.As(err, &apiErr) && apiErr.ErrorCode() == "InvalidInstanceType" { + return ec2_types.InstanceTypeInfo{}, fmt.Errorf("%s: %w: %s in %q", p.name, ErrTypeNotInRegion, instanceType, region) + } + return ec2_types.InstanceTypeInfo{}, fmt.Errorf("%s: DescribeInstanceTypes %q in %q: %w", p.name, instanceType, region, err) + } + if len(out.InstanceTypes) == 0 { + return ec2_types.InstanceTypeInfo{}, fmt.Errorf("%s: %w: %s in %q", p.name, ErrInstanceTypeNotFound, instanceType, region) + } + resolved := out.InstanceTypes[0] + var architectures []ec2_types.ArchitectureType + if resolved.ProcessorInfo != nil { + architectures = resolved.ProcessorInfo.SupportedArchitectures + } + if len(architectures) != 1 { + return ec2_types.InstanceTypeInfo{}, fmt.Errorf( + "%s: %w: %s in %q reports %v; only instance types with exactly one architecture are supported", + p.name, ErrInstanceTypeArchitecture, instanceType, region, architectures, + ) + } + return resolved, nil +} + +func instanceTypeArchitecture(instanceType ec2_types.InstanceTypeInfo) (ec2_types.ArchitectureValues, error) { + if instanceType.ProcessorInfo == nil || len(instanceType.ProcessorInfo.SupportedArchitectures) != 1 { + return "", fmt.Errorf("%w: %s", ErrInstanceTypeArchitecture, instanceType.InstanceType) + } + return ec2_types.ArchitectureValues(instanceType.ProcessorInfo.SupportedArchitectures[0]), nil +} + +func instanceTypeSupportsArch(it ec2_types.InstanceTypeInfo, arch ec2_types.ArchitectureValues) bool { + if it.ProcessorInfo == nil { + return false + } + for _, a := range it.ProcessorInfo.SupportedArchitectures { + if string(a) == string(arch) { + return true + } + } + return false +} + +// instancesByTag returns all pending or running instances in the given region +// that carry the tag, with the reservation layer flattened away. Terminated +// instances keep their tags for a while and must not show up as agents. +func (p *provider) instancesByTag(ctx context.Context, region, tag, value string) ([]ec2_types.Instance, error) { + out, err := p.client.DescribeInstances(ctx, &ec2.DescribeInstancesInput{ + Filters: []ec2_types.Filter{ + {Name: aws.String("tag:" + tag), Values: []string{value}}, + {Name: aws.String("instance-state-name"), Values: []string{ + string(ec2_types.InstanceStateNamePending), + string(ec2_types.InstanceStateNameRunning), + }}, + }, + }, regionOpt(region)) + if err != nil { + return nil, err + } + var instances []ec2_types.Instance + for _, r := range out.Reservations { + instances = append(instances, r.Instances...) + } + return instances, nil +} + +// getAgent finds the agent's instance and the region it runs in. +func (p *provider) getAgent(ctx context.Context, agent *woodpecker.Agent) (*ec2_types.Instance, string, error) { + for _, region := range p.regions { + instances, err := p.instancesByTag(ctx, region, "Name", agent.Name) + if err != nil { + return nil, "", err + } + if len(instances) > 1 { + return nil, "", fmt.Errorf("expected 1 instance with tag:Name=%s, got %d", agent.Name, len(instances)) + } + if len(instances) == 1 { + return &instances[0], region, nil + } + } + return nil, "", fmt.Errorf("no instance with tag:Name=%s in any deploy region", agent.Name) +} + +// capacityErrorCodes are the RunInstances error codes that mean the requested +// capacity is not available right now, for which deploying the next fallback +// candidate is worthwhile. +// See https://docs.aws.amazon.com/AWSEC2/latest/APIReference/errors-overview.html +var capacityErrorCodes = map[string]bool{ + "InsufficientInstanceCapacity": true, + "Server.InsufficientInstanceCapacity": true, + // The instance type is not supported in the requested availability zone. + "Unsupported": true, + // Spot-specific unavailability. + "MaxSpotInstanceCountExceeded": true, + "SpotMaxPriceTooLow": true, + // Not enough spare capacity to fulfill the Spot request right now. + "UnfulfillableCapacity": true, + // Instance type rejected for the account, e.g. not Free Tier eligible. + "InvalidParameterCombination": true, +} + +// isCapacityError reports whether err is an AWS capacity error, for which +// deploying the next fallback candidate is worthwhile. +func isCapacityError(err error) bool { + var apiErr smithy.APIError + if errors.As(err, &apiErr) { + return capacityErrorCodes[apiErr.ErrorCode()] + } + return false +} diff --git a/providers/aws/helper_test.go b/providers/aws/helper_test.go new file mode 100644 index 00000000..7e32e0ec --- /dev/null +++ b/providers/aws/helper_test.go @@ -0,0 +1,422 @@ +package aws + +import ( + "bytes" + "errors" + "fmt" + "testing" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/ec2" + ec2_types "github.com/aws/aws-sdk-go-v2/service/ec2/types" + "github.com/aws/smithy-go" + "github.com/rs/zerolog" + "github.com/rs/zerolog/log" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" + + "go.woodpecker-ci.org/autoscaler/providers/aws/ec2api/mocks" + "go.woodpecker-ci.org/woodpecker/v3/woodpecker-go/woodpecker" +) + +func TestResolveSubnets(t *testing.T) { + t.Run("configured subnets", func(t *testing.T) { + client := mocks.NewMockClient(t) + client.On("DescribeSubnets", mock.Anything, mock.MatchedBy(func(input *ec2.DescribeSubnetsInput) bool { + return assert.ObjectsAreEqual([]string{"subnet-2", "subnet-1"}, input.SubnetIds) && len(input.Filters) == 0 + }), mock.MatchedBy(regionOptions("eu-central-1"))). + Return(&ec2.DescribeSubnetsOutput{Subnets: []ec2_types.Subnet{{}, {}}}, nil).Once() + + p := newTestProvider(client) + subnets, err := p.resolveSubnets(t.Context(), "eu-central-1", []string{"subnet-2", "subnet-1"}) + require.NoError(t, err) + assert.Equal(t, []string{"subnet-2", "subnet-1"}, subnets) + }) + + t.Run("default subnets", func(t *testing.T) { + var output bytes.Buffer + originalLogger := log.Logger + log.Logger = zerolog.New(&output) + t.Cleanup(func() { + log.Logger = originalLogger + }) + + client := mocks.NewMockClient(t) + client.On("DescribeSubnets", mock.Anything, mock.MatchedBy(func(input *ec2.DescribeSubnetsInput) bool { + return len(input.SubnetIds) == 0 && + assert.ObjectsAreEqual([]string{"true"}, filterValues(input.Filters, "default-for-az")) && + assert.ObjectsAreEqual([]string{"available"}, filterValues(input.Filters, "state")) + }), mock.MatchedBy(regionOptions("us-east-1"))). + Return(&ec2.DescribeSubnetsOutput{Subnets: []ec2_types.Subnet{ + {SubnetId: aws.String("subnet-z")}, + {SubnetId: aws.String("subnet-a")}, + }}, nil).Once() + + p := newTestProvider(client) + subnets, err := p.resolveSubnets(t.Context(), "us-east-1", nil) + require.NoError(t, err) + assert.Equal(t, []string{"subnet-a", "subnet-z"}, subnets) + assert.JSONEq(t, `{ + "level":"info", + "region":"us-east-1", + "subnets":["subnet-a","subnet-z"], + "message":"resolved default AWS subnets" + }`, output.String()) + }) + + t.Run("no default subnets", func(t *testing.T) { + client := mocks.NewMockClient(t) + client.On("DescribeSubnets", mock.Anything, mock.Anything, mock.MatchedBy(regionOptions("us-east-1"))). + Return(&ec2.DescribeSubnetsOutput{}, nil).Once() + + p := newTestProvider(client) + _, err := p.resolveSubnets(t.Context(), "us-east-1", nil) + assert.ErrorIs(t, err, ErrSubnetsNotSet) + assert.ErrorContains(t, err, "no default subnets") + }) +} + +var ( + testImage = ec2_types.Image{ + ImageId: aws.String("ami-x86"), + Architecture: ec2_types.ArchitectureValuesX8664, + } + testArmImage = ec2_types.Image{ + ImageId: aws.String("ami-arm"), + Architecture: ec2_types.ArchitectureValuesArm64, + } + testTypeInfo = ec2_types.InstanceTypeInfo{ + InstanceType: ec2_types.InstanceType("m6i.large"), + CurrentGeneration: aws.Bool(true), + ProcessorInfo: &ec2_types.ProcessorInfo{ + SupportedArchitectures: []ec2_types.ArchitectureType{ec2_types.ArchitectureTypeX8664}, + }, + } + testArmTypeInfo = ec2_types.InstanceTypeInfo{ + InstanceType: ec2_types.InstanceType("t4g.micro"), + ProcessorInfo: &ec2_types.ProcessorInfo{ + SupportedArchitectures: []ec2_types.ArchitectureType{ec2_types.ArchitectureTypeArm64}, + }, + } +) + +func newTestProvider(client *mocks.MockClient) *provider { + return &provider{ + name: "aws", + region: "eu-central-1", + client: client, + } +} + +func regionOptions(region string) func([]func(*ec2.Options)) bool { + return func(options []func(*ec2.Options)) bool { + var got ec2.Options + for _, option := range options { + option(&got) + } + return got.Region == region + } +} + +func mockRegionResources(client *mocks.MockClient, region string, image ec2_types.Image, subnet, securityGroup string) { + client.On("DescribeImages", mock.Anything, mock.MatchedBy(func(in *ec2.DescribeImagesInput) bool { + return assert.ObjectsAreEqual([]string{aws.ToString(image.ImageId)}, in.ImageIds) + }), mock.MatchedBy(regionOptions(region))). + Return(&ec2.DescribeImagesOutput{Images: []ec2_types.Image{image}}, nil).Once() + client.On("DescribeSubnets", mock.Anything, mock.MatchedBy(func(in *ec2.DescribeSubnetsInput) bool { + return assert.ObjectsAreEqual([]string{subnet}, in.SubnetIds) + }), mock.MatchedBy(regionOptions(region))). + Return(&ec2.DescribeSubnetsOutput{Subnets: []ec2_types.Subnet{{}}}, nil).Once() + client.On("DescribeSecurityGroups", mock.Anything, mock.MatchedBy(func(in *ec2.DescribeSecurityGroupsInput) bool { + return assert.ObjectsAreEqual([]string{securityGroup}, in.GroupIds) + }), mock.MatchedBy(regionOptions(region))). + Return(&ec2.DescribeSecurityGroupsOutput{SecurityGroups: []ec2_types.SecurityGroup{{}}}, nil).Once() +} + +func mockAliasRegionResources(client *mocks.MockClient, region string, image ec2_types.Image, subnet, securityGroup string) { + nameArchitecture := "amd64" + if image.Architecture == ec2_types.ArchitectureValuesArm64 { + nameArchitecture = "arm64" + } + image.Name = aws.String("debian-13-" + nameArchitecture + "-20260714-1") + image.CreationDate = aws.String("2026-07-14T00:00:00Z") + image.ImageOwnerAlias = aws.String("amazon") + image.Public = aws.Bool(true) + image.RootDeviceType = ec2_types.DeviceTypeEbs + image.State = ec2_types.ImageStateAvailable + image.VirtualizationType = ec2_types.VirtualizationTypeHvm + client.On("DescribeImages", mock.Anything, mock.MatchedBy(func(input *ec2.DescribeImagesInput) bool { + return assert.ObjectsAreEqual([]string{string(image.Architecture)}, filterValues(input.Filters, "architecture")) && + assert.ObjectsAreEqual([]string{"debian-13-" + nameArchitecture + "-*"}, filterValues(input.Filters, "name")) + }), mock.MatchedBy(regionOptions(region))). + Return(&ec2.DescribeImagesOutput{Images: []ec2_types.Image{image}}, nil).Once() + client.On("DescribeSubnets", mock.Anything, mock.MatchedBy(func(input *ec2.DescribeSubnetsInput) bool { + return assert.ObjectsAreEqual([]string{subnet}, input.SubnetIds) + }), mock.MatchedBy(regionOptions(region))). + Return(&ec2.DescribeSubnetsOutput{Subnets: []ec2_types.Subnet{{}}}, nil).Once() + client.On("DescribeSecurityGroups", mock.Anything, mock.MatchedBy(func(input *ec2.DescribeSecurityGroupsInput) bool { + return assert.ObjectsAreEqual([]string{securityGroup}, input.GroupIds) + }), mock.MatchedBy(regionOptions(region))). + Return(&ec2.DescribeSecurityGroupsOutput{SecurityGroups: []ec2_types.SecurityGroup{{}}}, nil).Once() +} + +func TestResolveDeployCandidates(t *testing.T) { + t.Run("RegionalCandidates", func(t *testing.T) { + client := mocks.NewMockClient(t) + mockAliasRegionResources(client, "eu-central-1", testArmImage, "subnet-arm", "sg-arm") + mockAliasRegionResources(client, "us-east-1", testImage, "subnet-x86", "sg-x86") + client.On("DescribeInstanceTypes", mock.Anything, mock.MatchedBy(func(in *ec2.DescribeInstanceTypesInput) bool { + return in.InstanceTypes[0] == "t4g.micro" + }), mock.MatchedBy(regionOptions("eu-central-1"))). + Return(&ec2.DescribeInstanceTypesOutput{InstanceTypes: []ec2_types.InstanceTypeInfo{testArmTypeInfo}}, nil).Once() + client.On("DescribeInstanceTypes", mock.Anything, mock.MatchedBy(func(in *ec2.DescribeInstanceTypesInput) bool { + return in.InstanceTypes[0] == "m6i.large" + }), mock.MatchedBy(regionOptions("us-east-1"))). + Return(&ec2.DescribeInstanceTypesOutput{InstanceTypes: []ec2_types.InstanceTypeInfo{testTypeInfo}}, nil).Once() + + p := newTestProvider(client) + err := p.resolveDeployCandidates( + t.Context(), + []string{"t4g.micro:eu-central-1", "m6i.large:us-east-1"}, + "debian-13", + []string{"subnet-arm:eu-central-1", "subnet-x86:us-east-1"}, + []string{"sg-arm:eu-central-1", "sg-x86:us-east-1"}, + ) + assert.NoError(t, err) + assert.Equal(t, []string{"eu-central-1", "us-east-1"}, p.regions) + assert.Equal(t, "t4g.micro", string(p.deployCandidates[0].instanceType.InstanceType)) + assert.Equal(t, "eu-central-1", p.deployCandidates[0].regionConfig.region) + assert.Equal(t, "ami-arm", aws.ToString(p.deployCandidates[0].regionConfig.image.ImageId)) + assert.Equal(t, []string{"subnet-arm"}, p.deployCandidates[0].regionConfig.subnets) + assert.Equal(t, []string{"sg-arm"}, p.deployCandidates[0].regionConfig.securityGroups) + assert.Equal(t, "m6i.large", string(p.deployCandidates[1].instanceType.InstanceType)) + assert.Equal(t, "us-east-1", p.deployCandidates[1].regionConfig.region) + assert.Equal(t, "ami-x86", aws.ToString(p.deployCandidates[1].regionConfig.image.ImageId)) + }) + + t.Run("UnqualifiedValuesUseDefaultRegion", func(t *testing.T) { + client := mocks.NewMockClient(t) + mockRegionResources(client, "eu-central-1", testImage, "subnet-1", "sg-1") + client.On("DescribeInstanceTypes", mock.Anything, mock.Anything, mock.MatchedBy(regionOptions("eu-central-1"))). + Return(&ec2.DescribeInstanceTypesOutput{InstanceTypes: []ec2_types.InstanceTypeInfo{testTypeInfo}}, nil).Once() + + p := newTestProvider(client) + err := p.resolveDeployCandidates( + t.Context(), + []string{"m6i.large"}, "ami-x86", []string{"subnet-1"}, []string{"sg-1"}, + ) + assert.NoError(t, err) + assert.Equal(t, []string{"eu-central-1"}, p.regions) + }) + + t.Run("FullyQualifiedValuesDoNotNeedDefaultRegion", func(t *testing.T) { + client := mocks.NewMockClient(t) + mockRegionResources(client, "us-east-1", testImage, "subnet-1", "sg-1") + client.On("DescribeInstanceTypes", mock.Anything, mock.Anything, mock.MatchedBy(regionOptions("us-east-1"))). + Return(&ec2.DescribeInstanceTypesOutput{InstanceTypes: []ec2_types.InstanceTypeInfo{testTypeInfo}}, nil).Once() + + p := newTestProvider(client) + p.region = "" + err := p.resolveDeployCandidates( + t.Context(), + []string{"m6i.large:us-east-1"}, "ami-x86", + []string{"subnet-1:us-east-1"}, []string{"sg-1:us-east-1"}, + ) + assert.NoError(t, err) + }) + + t.Run("ImageAliasUsesInstanceTypeRegionAndArchitecture", func(t *testing.T) { + client := mocks.NewMockClient(t) + client.On("DescribeInstanceTypes", mock.Anything, mock.Anything, mock.MatchedBy(regionOptions("us-east-1"))). + Return(&ec2.DescribeInstanceTypesOutput{InstanceTypes: []ec2_types.InstanceTypeInfo{testArmTypeInfo}}, nil).Once() + client.On("DescribeImages", mock.Anything, mock.MatchedBy(func(input *ec2.DescribeImagesInput) bool { + return assert.ObjectsAreEqual( + []string{"debian-13-arm64-*"}, + filterValues(input.Filters, "name"), + ) + }), mock.MatchedBy(regionOptions("us-east-1"))).Return(&ec2.DescribeImagesOutput{Images: []ec2_types.Image{{ + ImageId: aws.String("ami-debian-arm"), + Name: aws.String("debian-13-arm64-20260714-1"), + Architecture: ec2_types.ArchitectureValuesArm64, + CreationDate: aws.String("2026-07-14T00:00:00Z"), + ImageOwnerAlias: aws.String("amazon"), + Public: aws.Bool(true), + RootDeviceType: ec2_types.DeviceTypeEbs, + State: ec2_types.ImageStateAvailable, + VirtualizationType: ec2_types.VirtualizationTypeHvm, + }}}, nil).Once() + client.On("DescribeSubnets", mock.Anything, mock.Anything, mock.MatchedBy(regionOptions("us-east-1"))). + Return(&ec2.DescribeSubnetsOutput{Subnets: []ec2_types.Subnet{{}}}, nil).Once() + + p := newTestProvider(client) + p.region = "" + err := p.resolveDeployCandidates( + t.Context(), + []string{"t4g.micro:us-east-1"}, "debian-13", + []string{"subnet-1:us-east-1"}, nil, + ) + require.NoError(t, err) + assert.Equal(t, "ami-debian-arm", aws.ToString(p.deployCandidates[0].regionConfig.image.ImageId)) + }) + + t.Run("AmbiguousInstanceTypeArchitecture", func(t *testing.T) { + client := mocks.NewMockClient(t) + client.On("DescribeInstanceTypes", mock.Anything, mock.Anything, mock.MatchedBy(regionOptions("us-east-1"))). + Return(&ec2.DescribeInstanceTypesOutput{InstanceTypes: []ec2_types.InstanceTypeInfo{{ + InstanceType: ec2_types.InstanceType("c3.large"), + ProcessorInfo: &ec2_types.ProcessorInfo{SupportedArchitectures: []ec2_types.ArchitectureType{ + ec2_types.ArchitectureTypeI386, + ec2_types.ArchitectureTypeX8664, + }}, + }}}, nil).Once() + + p := newTestProvider(client) + err := p.resolveDeployCandidates( + t.Context(), + []string{"c3.large:us-east-1"}, "debian-13", + []string{"subnet-1:us-east-1"}, nil, + ) + assert.ErrorIs(t, err, ErrInstanceTypeArchitecture) + assert.ErrorContains(t, err, "[i386 x86_64]") + }) + + t.Run("ImageMustNotSpecifyRegion", func(t *testing.T) { + client := mocks.NewMockClient(t) + client.On("DescribeInstanceTypes", mock.Anything, mock.Anything, mock.MatchedBy(regionOptions("us-east-1"))). + Return(&ec2.DescribeInstanceTypesOutput{InstanceTypes: []ec2_types.InstanceTypeInfo{testTypeInfo}}, nil).Once() + + p := newTestProvider(client) + err := p.resolveDeployCandidates( + t.Context(), + []string{"m6i.large:us-east-1"}, "ami-x86:us-east-1", + []string{"subnet-1:us-east-1"}, nil, + ) + assert.ErrorContains(t, err, "aws-ami-id must not specify a region") + }) + + t.Run("UnqualifiedValueNeedsDefaultRegion", func(t *testing.T) { + client := mocks.NewMockClient(t) + p := newTestProvider(client) + p.region = "" + err := p.resolveDeployCandidates( + t.Context(), + []string{"m6i.large"}, "ami-x86", + []string{"subnet-1:us-east-1"}, []string{"sg-1:us-east-1"}, + ) + assert.ErrorIs(t, err, ErrRegionNotSet) + }) + + t.Run("NoCandidates", func(t *testing.T) { + client := mocks.NewMockClient(t) + p := newTestProvider(client) + err := p.resolveDeployCandidates(t.Context(), nil, "", nil, nil) + assert.ErrorIs(t, err, ErrNoDeployCandidates) + }) + + t.Run("MissingAMIPerCandidateRegion", func(t *testing.T) { + client := mocks.NewMockClient(t) + client.On("DescribeInstanceTypes", mock.Anything, mock.Anything, mock.MatchedBy(regionOptions("us-east-1"))). + Return(&ec2.DescribeInstanceTypesOutput{InstanceTypes: []ec2_types.InstanceTypeInfo{testTypeInfo}}, nil).Once() + p := newTestProvider(client) + err := p.resolveDeployCandidates( + t.Context(), + []string{"m6i.large:us-east-1"}, "", + []string{"subnet-1:us-east-1"}, nil, + ) + assert.ErrorIs(t, err, ErrAMINotFound) + }) + + t.Run("TypeNotOfferedInRegion", func(t *testing.T) { + client := mocks.NewMockClient(t) + client.On("DescribeInstanceTypes", mock.Anything, mock.Anything, mock.MatchedBy(regionOptions("us-east-1"))). + Return(nil, &apiError{code: "InvalidInstanceType"}).Once() + + p := newTestProvider(client) + err := p.resolveDeployCandidates( + t.Context(), + []string{"m6i.large:us-east-1"}, "ami-x86", + []string{"subnet-1:us-east-1"}, []string{"sg-1:us-east-1"}, + ) + assert.ErrorIs(t, err, ErrTypeNotInRegion) + }) + + t.Run("ArchMismatch", func(t *testing.T) { + client := mocks.NewMockClient(t) + mockRegionResources(client, "eu-central-1", testImage, "subnet-1", "sg-1") + client.On("DescribeInstanceTypes", mock.Anything, mock.Anything, mock.MatchedBy(regionOptions("eu-central-1"))). + Return(&ec2.DescribeInstanceTypesOutput{InstanceTypes: []ec2_types.InstanceTypeInfo{testArmTypeInfo}}, nil).Once() + + p := newTestProvider(client) + err := p.resolveDeployCandidates( + t.Context(), + []string{"t4g.micro"}, "ami-x86", []string{"subnet-1"}, []string{"sg-1"}, + ) + assert.ErrorIs(t, err, ErrArchMismatch) + }) +} + +type apiError struct{ code string } + +func (e *apiError) Error() string { return e.code } +func (e *apiError) ErrorCode() string { return e.code } +func (e *apiError) ErrorMessage() string { return e.code } +func (e *apiError) ErrorFault() smithy.ErrorFault { return smithy.FaultServer } + +func TestIsCapacityError(t *testing.T) { + for _, code := range []string{ + "InsufficientInstanceCapacity", + "Server.InsufficientInstanceCapacity", + "Unsupported", + "MaxSpotInstanceCountExceeded", + "SpotMaxPriceTooLow", + "UnfulfillableCapacity", + "InvalidParameterCombination", + } { + assert.True(t, isCapacityError(fmt.Errorf("wrapped: %w", &apiError{code: code})), code) + } + + assert.False(t, isCapacityError(&apiError{code: "InvalidSubnetID.NotFound"})) + assert.False(t, isCapacityError(errors.New("no api error"))) +} + +func TestGetAgent(t *testing.T) { + instance := ec2_types.Instance{InstanceId: aws.String("i-1")} + + t.Run("FoundInFallbackRegion", func(t *testing.T) { + client := mocks.NewMockClient(t) + client.On("DescribeInstances", mock.Anything, mock.Anything, mock.MatchedBy(regionOptions("eu-central-1"))). + Return(&ec2.DescribeInstancesOutput{}, nil).Once() + client.On("DescribeInstances", mock.Anything, mock.Anything, mock.MatchedBy(regionOptions("us-east-1"))). + Return(&ec2.DescribeInstancesOutput{ + Reservations: []ec2_types.Reservation{{Instances: []ec2_types.Instance{instance}}}, + }, nil).Once() + + p := newTestProvider(client) + p.regions = []string{"eu-central-1", "us-east-1"} + got, region, err := p.getAgent(t.Context(), &woodpecker.Agent{Name: "pool-1-agent-abcd"}) + assert.NoError(t, err) + assert.Equal(t, "us-east-1", region) + assert.Equal(t, "i-1", aws.ToString(got.InstanceId)) + }) + + t.Run("QueriesOnlyNonTerminatedStates", func(t *testing.T) { + client := mocks.NewMockClient(t) + client.On("DescribeInstances", mock.Anything, mock.MatchedBy(func(in *ec2.DescribeInstancesInput) bool { + for _, f := range in.Filters { + if aws.ToString(f.Name) == "instance-state-name" { + return assert.ObjectsAreEqual([]string{"pending", "running"}, f.Values) + } + } + return false + }), mock.MatchedBy(regionOptions("eu-central-1"))). + Return(&ec2.DescribeInstancesOutput{ + Reservations: []ec2_types.Reservation{{Instances: []ec2_types.Instance{instance}}}, + }, nil) + + p := newTestProvider(client) + p.regions = []string{"eu-central-1"} + _, _, err := p.getAgent(t.Context(), &woodpecker.Agent{Name: "pool-1-agent-abcd"}) + assert.NoError(t, err) + }) +} diff --git a/providers/aws/image.go b/providers/aws/image.go new file mode 100644 index 00000000..d27ad7f8 --- /dev/null +++ b/providers/aws/image.go @@ -0,0 +1,237 @@ +package aws + +import ( + "context" + "fmt" + "strings" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/ec2" + ec2_types "github.com/aws/aws-sdk-go-v2/service/ec2/types" + "github.com/rs/zerolog/log" +) + +func validateImageReference(reference string) error { + if reference == "" { + return ErrAMINotFound + } + if strings.Contains(reference, ":") { + return fmt.Errorf("aws-ami-id must not specify a region: %s", reference) + } + if isImageAlias(reference) || strings.HasPrefix(reference, "ami-") { + return nil + } + return fmt.Errorf("unsupported aws-ami-id alias: %s", reference) +} + +func isImageAlias(reference string) bool { + _, found, _ := imageAliasQuery(reference, ec2_types.ArchitectureValuesX8664) + return found +} + +func (p *provider) resolveImage(ctx context.Context, reference, region string, architecture ec2_types.ArchitectureValues) (ec2_types.Image, error) { + if isImageAlias(reference) { + return p.resolveImageAlias(ctx, reference, region, architecture) + } + + images, err := p.client.DescribeImages(ctx, &ec2.DescribeImagesInput{ + ImageIds: []string{reference}, + }, regionOpt(region)) + if err != nil { + return ec2_types.Image{}, fmt.Errorf("%s: DescribeImages %q in %q: %w", p.name, reference, region, err) + } + if len(images.Images) != 1 { + return ec2_types.Image{}, fmt.Errorf("%s: %w: %s in %q", p.name, ErrAMINotFound, reference, region) + } + return images.Images[0], nil +} + +type aliasQuery struct { + namePattern string + matchesName func(string) bool +} + +func (p *provider) resolveImageAlias(ctx context.Context, alias, region string, architecture ec2_types.ArchitectureValues) (ec2_types.Image, error) { + query, found, err := imageAliasQuery(alias, architecture) + if err != nil { + return ec2_types.Image{}, fmt.Errorf("%s: %w", p.name, err) + } + if !found { + return ec2_types.Image{}, fmt.Errorf("%s: unsupported aws-ami-id alias: %s", p.name, alias) + } + images, err := p.client.DescribeImages(ctx, &ec2.DescribeImagesInput{ + Filters: []ec2_types.Filter{ + {Name: aws.String("architecture"), Values: []string{string(architecture)}}, + {Name: aws.String("is-public"), Values: []string{"true"}}, + {Name: aws.String("name"), Values: []string{query.namePattern}}, + {Name: aws.String("owner-alias"), Values: []string{"amazon"}}, + {Name: aws.String("root-device-type"), Values: []string{string(ec2_types.DeviceTypeEbs)}}, + {Name: aws.String("state"), Values: []string{string(ec2_types.ImageStateAvailable)}}, + {Name: aws.String("virtualization-type"), Values: []string{string(ec2_types.VirtualizationTypeHvm)}}, + }, + }, regionOpt(region)) + if err != nil { + return ec2_types.Image{}, fmt.Errorf("%s: DescribeImages alias %q in %q: %w", p.name, alias, region, err) + } + + var newest ec2_types.Image + for _, image := range images.Images { + if !isVerifiedAliasImage(image, architecture, query.matchesName) { + continue + } + if newest.ImageId == nil || newerImage(image, newest) { + newest = image + } + } + if newest.ImageId == nil { + return ec2_types.Image{}, fmt.Errorf("%s: %w: verified provider alias %q for %s in %q", + p.name, ErrAMINotFound, alias, architecture, region) + } + log.Info(). + Str("alias", alias). + Str("ami_id", aws.ToString(newest.ImageId)). + Str("region", region). + Str("architecture", string(architecture)). + Msg("resolved AWS AMI alias") + return newest, nil +} + +func isVerifiedAliasImage(image ec2_types.Image, architecture ec2_types.ArchitectureValues, matchesName func(string) bool) bool { + return image.Architecture == architecture && + matchesName(aws.ToString(image.Name)) && + aws.ToString(image.ImageOwnerAlias) == "amazon" && + aws.ToBool(image.Public) && + image.RootDeviceType == ec2_types.DeviceTypeEbs && + image.State == ec2_types.ImageStateAvailable && + image.VirtualizationType == ec2_types.VirtualizationTypeHvm +} + +func newerImage(left, right ec2_types.Image) bool { + leftCreationDate := aws.ToString(left.CreationDate) + rightCreationDate := aws.ToString(right.CreationDate) + if leftCreationDate != rightCreationDate { + return leftCreationDate > rightCreationDate + } + return aws.ToString(left.ImageId) > aws.ToString(right.ImageId) +} + +func imageAliasQuery(alias string, architecture ec2_types.ArchitectureValues) (aliasQuery, bool, error) { + debianArch, awsArch, err := imageArchitectures(architecture) + if err != nil { + return aliasQuery{}, true, err + } + + if version, found := strings.CutPrefix(alias, "debian-"); found && onlyDigits(version) { + prefix := alias + "-" + debianArch + "-" + return aliasQuery{ + namePattern: prefix + "*", + matchesName: func(name string) bool { return strings.HasPrefix(name, prefix) }, + }, true, nil + } + + if version, found := strings.CutPrefix(alias, "ubuntu-"); found { + version, found = strings.CutSuffix(version, "-server") + if found && numericVersion(version) { + // hvm-ssd covers both the hvm-ssd and hvm-ssd-gp3 storage generations + prefix := "ubuntu/images/hvm-ssd" + marker := "-" + version + "-" + debianArch + "-server-" + return aliasQuery{ + namePattern: prefix + "*/ubuntu-*" + marker + "*", + matchesName: func(name string) bool { + return strings.HasPrefix(name, prefix) && strings.Contains(name, "/ubuntu-") && strings.Contains(name, marker) + }, + }, true, nil + } + } + + if alias == "amazon" || alias == "amazon_linux" || alias == "amazon-linux" { + // pinned to the current major release; bump when a new one ships + alias = "amazon-linux-2023" + } + if version, found := strings.CutPrefix(alias, "amazon-linux-"); found && onlyDigits(version) { + prefix := "al" + version + "-ami-" + version + "." + suffix := "-" + awsArch + return aliasQuery{ + namePattern: prefix + "*-kernel-*" + suffix, + matchesName: func(name string) bool { + return strings.HasPrefix(name, prefix) && strings.Contains(name, "-kernel-") && strings.HasSuffix(name, suffix) + }, + }, true, nil + } + + if alias == "suse" { + // pinned to the current major release; bump when a new one ships + alias = "suse-16" + } + if version, found := strings.CutPrefix(alias, "suse-"); found && onlyDigits(version) { + prefix := "suse-sles-" + version + "-" + suffix := "-hvm-ssd-" + awsArch + return aliasQuery{ + namePattern: prefix + "*" + suffix, + matchesName: func(name string) bool { + middle, found := strings.CutPrefix(name, prefix) + if !found { + return false + } + middle, found = strings.CutSuffix(middle, suffix) + // vDATE for 15.x service packs (sp7-v20260630) and 16.x + // minor releases (0-v20260625); anything else is a vendor + // variant such as sapcal, chost-byos or ecs. + return found && suseVersion(middle) + }, + }, true, nil + } + + return aliasQuery{}, false, nil +} + +// suseVersion reports whether the middle part of a SLES image name is a plain +// release, i.e. an optional service pack or minor release followed by the +// vDATE stamp, without any vendor variant in between. +func suseVersion(middle string) bool { + release, rest, found := strings.Cut(middle, "-") + if found { + release = strings.TrimPrefix(release, "sp") + if !onlyDigits(release) { + return false + } + middle = rest + } + date, found := strings.CutPrefix(middle, "v") + return found && onlyDigits(date) +} + +func imageArchitectures(architecture ec2_types.ArchitectureValues) (string, string, error) { + switch architecture { + case ec2_types.ArchitectureValuesX8664: + return "amd64", "x86_64", nil + case ec2_types.ArchitectureValuesArm64: + return "arm64", "arm64", nil + default: + return "", "", fmt.Errorf("unsupported image alias architecture: %s", architecture) + } +} + +func numericVersion(version string) bool { + if version == "" || strings.HasPrefix(version, ".") || strings.HasSuffix(version, ".") { + return false + } + for part := range strings.SplitSeq(version, ".") { + if !onlyDigits(part) { + return false + } + } + return true +} + +func onlyDigits(value string) bool { + if value == "" { + return false + } + for _, char := range value { + if char < '0' || char > '9' { + return false + } + } + return true +} diff --git a/providers/aws/image_test.go b/providers/aws/image_test.go new file mode 100644 index 00000000..02c72a90 --- /dev/null +++ b/providers/aws/image_test.go @@ -0,0 +1,215 @@ +package aws + +import ( + "bytes" + "testing" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/ec2" + ec2_types "github.com/aws/aws-sdk-go-v2/service/ec2/types" + "github.com/rs/zerolog" + "github.com/rs/zerolog/log" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" + "github.com/urfave/cli/v3" + + "go.woodpecker-ci.org/autoscaler/providers/aws/ec2api/mocks" +) + +func TestValidateImageReference(t *testing.T) { + t.Run("aliases", func(t *testing.T) { + for _, alias := range []string{ + "ubuntu-26.04-server", + "ubuntu-20.04-server", + "amazon-linux", + "amazon_linux", + "amazon-linux-2023", + "suse", + "suse-15", + "suse-16", + "debian-13", + } { + require.NoError(t, validateImageReference(alias), alias) + } + }) + + t.Run("broad alias is rejected", func(t *testing.T) { + for _, alias := range []string{"ubuntu", "ubuntu-server", "debian", "suse-", "amazon-linux-"} { + err := validateImageReference(alias) + assert.ErrorContains(t, err, "unsupported aws-ami-id alias", alias) + } + }) + + t.Run("alias region is rejected", func(t *testing.T) { + err := validateImageReference("debian-13:us-east-1") + assert.ErrorContains(t, err, "aws-ami-id must not specify a region") + }) + + t.Run("explicit image", func(t *testing.T) { + require.NoError(t, validateImageReference("ami-1")) + }) + + t.Run("explicit image region is rejected", func(t *testing.T) { + err := validateImageReference("ami-1:us-east-1") + assert.ErrorContains(t, err, "aws-ami-id must not specify a region") + }) +} + +func TestDefaultImageReferenceIsValid(t *testing.T) { + for _, flag := range ProviderFlags { + imageFlag, ok := flag.(*cli.StringFlag) + if !ok || imageFlag.Name != "aws-ami-id" { + continue + } + require.NoError(t, validateImageReference(imageFlag.Value)) + return + } + t.Fatal("aws-ami-id flag not found") +} + +func TestResolveImageAliases(t *testing.T) { + tests := []struct { + name string + alias string + nameFilter string + selectedName string + rejectedName string + }{ + { + name: "Ubuntu", + alias: "ubuntu-26.04-server", + nameFilter: "ubuntu/images/hvm-ssd*/ubuntu-*-26.04-amd64-server-*", + selectedName: "ubuntu/images/hvm-ssd-gp3/ubuntu-resolute-26.04-amd64-server-20260714", + rejectedName: "ubuntu-pro-server/images/hvm-ssd-gp3/ubuntu-resolute-26.04-amd64-pro-server-20260714", + }, + { + // releases before 22.04 are published under hvm-ssd instead of hvm-ssd-gp3 + name: "UbuntuOldStorageGeneration", + alias: "ubuntu-20.04-server", + nameFilter: "ubuntu/images/hvm-ssd*/ubuntu-*-20.04-amd64-server-*", + selectedName: "ubuntu/images/hvm-ssd/ubuntu-focal-20.04-amd64-server-20250624", + rejectedName: "ubuntu-minimal/images/hvm-ssd/ubuntu-focal-20.04-amd64-minimal-server-20250624", + }, + { + name: "AmazonLinux", + alias: "amazon_linux", + nameFilter: "al2023-ami-2023.*-kernel-*-x86_64", + selectedName: "al2023-ami-2023.12.20260710.0-kernel-6.18-x86_64", + rejectedName: "al2023-ami-minimal-2023.12.20260710.0-kernel-6.18-x86_64", + }, + { + name: "AmazonLinuxVersioned", + alias: "amazon-linux-2023", + nameFilter: "al2023-ami-2023.*-kernel-*-x86_64", + selectedName: "al2023-ami-2023.12.20260710.0-kernel-6.18-x86_64", + rejectedName: "al2023-ami-minimal-2023.12.20260710.0-kernel-6.18-x86_64", + }, + { + name: "SUSE", + alias: "suse", + nameFilter: "suse-sles-16-*-hvm-ssd-x86_64", + selectedName: "suse-sles-16-0-v20260703-hvm-ssd-x86_64", + rejectedName: "suse-sles-16-0-v20260703-ecs-hvm-ssd-x86_64", + }, + { + // 15.x releases are published with service pack naming (spN) + name: "SUSEServicePackNaming", + alias: "suse-15", + nameFilter: "suse-sles-15-*-hvm-ssd-x86_64", + selectedName: "suse-sles-15-sp7-v20260630-hvm-ssd-x86_64", + rejectedName: "suse-sles-15-sp4-sapcal-v20260703-hvm-ssd-x86_64", + }, + { + name: "Debian", + alias: "debian-13", + nameFilter: "debian-13-amd64-*", + selectedName: "debian-13-amd64-20260712-2537", + rejectedName: "debian-13-backports-amd64-20260712-2537", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + client := mocks.NewMockClient(t) + client.On("DescribeImages", mock.Anything, mock.MatchedBy(func(input *ec2.DescribeImagesInput) bool { + return len(input.Owners) == 0 && + assert.ObjectsAreEqual([]string{test.nameFilter}, filterValues(input.Filters, "name")) && + assert.ObjectsAreEqual([]string{"x86_64"}, filterValues(input.Filters, "architecture")) && + assert.ObjectsAreEqual([]string{"amazon"}, filterValues(input.Filters, "owner-alias")) && + assert.ObjectsAreEqual([]string{"ebs"}, filterValues(input.Filters, "root-device-type")) && + assert.ObjectsAreEqual([]string{"hvm"}, filterValues(input.Filters, "virtualization-type")) + }), mock.MatchedBy(regionOptions("eu-central-1"))).Return(&ec2.DescribeImagesOutput{Images: []ec2_types.Image{ + verifiedImage("ami-rejected-newest", test.rejectedName, "2026-07-15T00:00:00Z", "amazon"), + verifiedImage("ami-unverified", test.selectedName, "2026-07-15T00:00:00Z", ""), + verifiedImage("ami-selected", test.selectedName, "2026-07-14T00:00:00Z", "amazon"), + verifiedImage("ami-old", test.selectedName, "2026-07-13T00:00:00Z", "amazon"), + }}, nil).Once() + + p := newTestProvider(client) + image, err := p.resolveImage(t.Context(), test.alias, "eu-central-1", ec2_types.ArchitectureValuesX8664) + require.NoError(t, err) + assert.Equal(t, "ami-selected", aws.ToString(image.ImageId)) + }) + } +} + +func TestResolveImageAliasLogsAMI(t *testing.T) { + var output bytes.Buffer + originalLogger := log.Logger + log.Logger = zerolog.New(&output) + t.Cleanup(func() { + log.Logger = originalLogger + }) + + client := mocks.NewMockClient(t) + client.On("DescribeImages", mock.Anything, mock.MatchedBy(func(input *ec2.DescribeImagesInput) bool { + return len(input.Owners) == 0 && + assert.ObjectsAreEqual([]string{"debian-13-amd64-*"}, filterValues(input.Filters, "name")) && + assert.ObjectsAreEqual([]string{"x86_64"}, filterValues(input.Filters, "architecture")) && + assert.ObjectsAreEqual([]string{"amazon"}, filterValues(input.Filters, "owner-alias")) && + assert.ObjectsAreEqual([]string{"ebs"}, filterValues(input.Filters, "root-device-type")) && + assert.ObjectsAreEqual([]string{"hvm"}, filterValues(input.Filters, "virtualization-type")) + }), mock.MatchedBy(regionOptions("eu-central-1"))).Return(&ec2.DescribeImagesOutput{Images: []ec2_types.Image{ + verifiedImage("ami-unverified-newest", "debian-13-amd64-20260714-1", "2026-07-14T00:00:00Z", ""), + verifiedImage("ami-marketplace", "debian-13-amd64-20260714-prod", "2026-07-14T00:00:00Z", "aws-marketplace"), + verifiedImage("ami-verified-old", "debian-13-amd64-20260712-1", "2026-07-12T00:00:00Z", "amazon"), + verifiedImage("ami-verified-new", "debian-13-amd64-20260713-1", "2026-07-13T00:00:00Z", "amazon"), + }}, nil).Once() + + p := newTestProvider(client) + image, err := p.resolveImage(t.Context(), "debian-13", "eu-central-1", ec2_types.ArchitectureValuesX8664) + require.NoError(t, err) + assert.Equal(t, "ami-verified-new", aws.ToString(image.ImageId)) + assert.JSONEq(t, `{ + "level":"info", + "alias":"debian-13", + "ami_id":"ami-verified-new", + "region":"eu-central-1", + "architecture":"x86_64", + "message":"resolved AWS AMI alias" + }`, output.String()) +} + +func verifiedImage(id, name, created, ownerAlias string) ec2_types.Image { + return ec2_types.Image{ + ImageId: aws.String(id), + Name: aws.String(name), + Architecture: ec2_types.ArchitectureValuesX8664, + CreationDate: aws.String(created), + ImageOwnerAlias: aws.String(ownerAlias), + Public: aws.Bool(true), + RootDeviceType: ec2_types.DeviceTypeEbs, + State: ec2_types.ImageStateAvailable, + VirtualizationType: ec2_types.VirtualizationTypeHvm, + } +} + +func filterValues(filters []ec2_types.Filter, name string) []string { + for _, filter := range filters { + if aws.ToString(filter.Name) == name { + return filter.Values + } + } + return nil +} diff --git a/providers/aws/provider.go b/providers/aws/provider.go index 85cb644d..20290a95 100644 --- a/providers/aws/provider.go +++ b/providers/aws/provider.go @@ -10,8 +10,10 @@ import ( "github.com/aws/aws-sdk-go-v2/aws" awsconfig "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/credentials" "github.com/aws/aws-sdk-go-v2/service/ec2" ec2_types "github.com/aws/aws-sdk-go-v2/service/ec2/types" + "github.com/aws/aws-sdk-go-v2/service/sts" "github.com/rs/zerolog/log" "github.com/urfave/cli/v3" @@ -19,52 +21,145 @@ import ( "go.woodpecker-ci.org/autoscaler/engine" "go.woodpecker-ci.org/autoscaler/engine/inits/cloudinit" "go.woodpecker-ci.org/autoscaler/engine/types" + "go.woodpecker-ci.org/autoscaler/providers/aws/ec2api" "go.woodpecker-ci.org/woodpecker/v3/woodpecker-go/woodpecker" ) type provider struct { name string config *config.Config - instanceType string - amiID string tags []string region string - subnets []string - securityGroups []string iamInstanceProfileArn string useSpotInstances bool - client *ec2.Client + client ec2api.Client lock sync.Mutex subnetRR int sshKeyName string + // resolved config + deployCandidates []deployCandidate + regions []string } func New(ctx context.Context, c *cli.Command, config *config.Config) (types.Provider, error) { - if len(c.StringSlice("aws-subnets")) == 0 { - return nil, fmt.Errorf("aws-subnets must be set") - } p := &provider{ name: "aws", config: config, - instanceType: c.String("aws-instance-type"), - amiID: c.String("aws-ami-id"), tags: c.StringSlice("aws-tags"), region: c.String("aws-region"), - subnets: c.StringSlice("aws-subnets"), iamInstanceProfileArn: c.String("aws-iam-instance-profile-arn"), - securityGroups: c.StringSlice("aws-security-groups"), useSpotInstances: c.Bool("aws-use-spot-instances"), sshKeyName: c.String("aws-ssh-key-name"), } - cfg, err := awsconfig.LoadDefaultConfig(ctx, awsconfig.WithRegion(p.region)) + loadOptions := []func(*awsconfig.LoadOptions) error{} + credentialsOption, err := staticCredentialsOption( + c.String("aws-access-key-id"), + c.String("aws-secret-access-key"), + ) + if err != nil { + return nil, fmt.Errorf("%s: %w", p.name, err) + } + if credentialsOption != nil { + loadOptions = append(loadOptions, credentialsOption) + } + if p.region != "" { + loadOptions = append(loadOptions, awsconfig.WithRegion(p.region)) + } + cfg, err := awsconfig.LoadDefaultConfig(ctx, loadOptions...) if err != nil { return nil, fmt.Errorf("failed to load configuration, %w", err) } + if p.region == "" { + // Fall back to the region the SDK resolved from its default chain + // (AWS_REGION, shared config, IMDS) so unqualified values keep + // working without the aws-region flag. + p.region = cfg.Region + } + instanceTypes := c.StringSlice("aws-instance-type") + authenticationRegion, err := p.authenticationRegion(instanceTypes) + if err != nil { + return nil, err + } + identityClient := sts.NewFromConfig(cfg, func(options *sts.Options) { + options.Region = authenticationRegion + }) + if err := authenticateAWS(ctx, identityClient); err != nil { + return nil, fmt.Errorf("%s: %w", p.name, err) + } p.client = ec2.NewFromConfig(cfg) + if err := p.resolveDeployCandidates( + ctx, + instanceTypes, + c.String("aws-ami-id"), + c.StringSlice("aws-subnets"), + c.StringSlice("aws-security-groups"), + ); err != nil { + return nil, err + } + + p.printResolvedConfig() + return p, nil } +func (p *provider) authenticationRegion(instanceTypes []string) (string, error) { + if p.region != "" { + return p.region, nil + } + if len(instanceTypes) == 0 { + return "", fmt.Errorf("%s: %w", p.name, ErrNoDeployCandidates) + } + + _, region, err := p.valueRegion("aws-instance-type", instanceTypes[0]) + return region, err +} + +type identityClient interface { + GetCallerIdentity(context.Context, *sts.GetCallerIdentityInput, ...func(*sts.Options)) (*sts.GetCallerIdentityOutput, error) +} + +func authenticateAWS(ctx context.Context, client identityClient) error { + identity, err := client.GetCallerIdentity(ctx, &sts.GetCallerIdentityInput{}) + if err != nil { + return fmt.Errorf("authenticate with AWS: %w", err) + } + + log.Info(). + Str("account", aws.ToString(identity.Account)). + Str("arn", aws.ToString(identity.Arn)). + Msg("authenticated with AWS") + return nil +} + +func staticCredentialsOption(accessKeyID, secretAccessKey string) (func(*awsconfig.LoadOptions) error, error) { + if accessKeyID == "" && secretAccessKey == "" { + return nil, nil + } + if accessKeyID == "" { + return nil, fmt.Errorf("aws-access-key-id must be set when aws-secret-access-key is set") + } + if secretAccessKey == "" { + return nil, fmt.Errorf("aws-secret-access-key must be set when aws-access-key-id is set") + } + + return awsconfig.WithCredentialsProvider( + credentials.NewStaticCredentialsProvider(accessKeyID, secretAccessKey, ""), + ), nil +} + +func (p *provider) printResolvedConfig() { + for _, c := range p.deployCandidates { + log.Info(). + Str("type", string(c.instanceType.InstanceType)). + Str("region", c.regionConfig.region). + Str("ami", aws.ToString(c.regionConfig.image.ImageId)). + Str("ami_arch", string(c.regionConfig.image.Architecture)). + Bool("current_gen", aws.ToBool(c.instanceType.CurrentGeneration)). + Msg("deploy candidate") + } +} + func (p *provider) DeployAgent(ctx context.Context, agent *woodpecker.Agent) error { userData, err := cloudinit.RenderUserDataTemplate(p.config, agent, cloudinit.RenderOption{}) if err != nil { @@ -103,16 +198,13 @@ func (p *provider) DeployAgent(ctx context.Context, agent *woodpecker.Agent) err IamInstanceProfile: &ec2_types.IamInstanceProfileSpecification{ Arn: aws.String(p.iamInstanceProfileArn), }, - ImageId: aws.String(p.amiID), - InstanceType: ec2_types.InstanceType(p.instanceType), MetadataOptions: &ec2_types.InstanceMetadataOptionsRequest{ HttpEndpoint: ec2_types.InstanceMetadataEndpointStateEnabled, HttpPutResponseHopLimit: aws.Int32(1), HttpTokens: ec2_types.HttpTokensStateRequired, }, - SecurityGroupIds: p.securityGroups, - MinCount: aws.Int32(1), - MaxCount: aws.Int32(1), + MinCount: aws.Int32(1), + MaxCount: aws.Int32(1), TagSpecifications: []ec2_types.TagSpecification{ { ResourceType: "instance", @@ -125,12 +217,6 @@ func (p *provider) DeployAgent(ctx context.Context, agent *woodpecker.Agent) err }, } - // When multiple subnets are given, assign agent to a subnet in a round-robin fashion. - p.lock.Lock() - runInstancesInput.SubnetId = aws.String(p.subnets[p.subnetRR]) - p.subnetRR = (p.subnetRR + 1) % len(p.subnets) - p.lock.Unlock() - if p.useSpotInstances { runInstancesInput.InstanceMarketOptions = &ec2_types.InstanceMarketOptionsRequest{ MarketType: ec2_types.MarketTypeSpot, @@ -142,9 +228,50 @@ func (p *provider) DeployAgent(ctx context.Context, agent *woodpecker.Agent) err } runInstancesInput.UserData = aws.String(b64.StdEncoding.EncodeToString([]byte(userData))) - result, err := p.client.RunInstances(ctx, &runInstancesInput) - if err != nil { - return fmt.Errorf("%s: RunInstances: %w", p.name, err) + + var result *ec2.RunInstancesOutput + for i, c := range p.deployCandidates { + runInstancesInput.InstanceType = c.instanceType.InstanceType + runInstancesInput.ImageId = c.regionConfig.image.ImageId + runInstancesInput.SecurityGroupIds = c.regionConfig.securityGroups + + // When multiple subnets are given, assign agent to a subnet in a round-robin fashion. + p.lock.Lock() + runInstancesInput.SubnetId = aws.String(c.regionConfig.subnets[p.subnetRR%len(c.regionConfig.subnets)]) + p.subnetRR = (p.subnetRR + 1) % len(c.regionConfig.subnets) + p.lock.Unlock() + + log.Info(). + Str("type", string(c.instanceType.InstanceType)). + Str("region", c.regionConfig.region). + Msg("create agent") + + result, err = p.client.RunInstances(ctx, &runInstancesInput, regionOpt(c.regionConfig.region)) + if err == nil { + break + } + + // Continue to next fallback entry only if capacity is unavailable. + if !isCapacityError(err) { + return fmt.Errorf("%s: RunInstances: %w", p.name, err) + } + + // Only log and continue if there are more candidates left. + if i < len(p.deployCandidates)-1 { + log.Warn().Msgf( + "create agent failed: type = %s region = %s: %s", + c.instanceType.InstanceType, c.regionConfig.region, err, + ) + continue + } + + // Last candidate failed: the whole fallback chain is exhausted. + return fmt.Errorf("%s: all %d deploy candidates out of capacity, last: RunInstances: %w", + p.name, len(p.deployCandidates), err) + } + + if result == nil || len(result.Instances) == 0 { + return fmt.Errorf("%s: RunInstances returned no instances", p.name) } // Wait until instance is available. Sometimes it can take a second or two for the tag based @@ -153,7 +280,7 @@ func (p *provider) DeployAgent(ctx context.Context, agent *woodpecker.Agent) err for range 5 { agents, err := p.ListDeployedAgentNames(ctx) if err != nil { - return fmt.Errorf("failed to return list for agents") + return fmt.Errorf("%s: ListDeployedAgentNames: %w", p.name, err) } for _, a := range agents { @@ -169,36 +296,18 @@ func (p *provider) DeployAgent(ctx context.Context, agent *woodpecker.Agent) err return fmt.Errorf("instance did not resolve in agent list: %s", *result.Instances[0].InstanceId) } -func (p *provider) getAgent(ctx context.Context, agent *woodpecker.Agent) (*ec2_types.Instance, error) { - instances, err := p.client.DescribeInstances(ctx, &ec2.DescribeInstancesInput{ - Filters: []ec2_types.Filter{ - { - Name: aws.String("tag:Name"), - Values: []string{agent.Name}, - }, - }, - }) - if err != nil { - return nil, err - } - if len(instances.Reservations) != 1 { - return nil, fmt.Errorf("expected 1 reservation with tag:Name=%s, got %d", agent.Name, len(instances.Reservations)) - } - if len(instances.Reservations[0].Instances) != 1 { - return nil, fmt.Errorf("expected 1 instance with tag:Name=%s, got %d", agent.Name, len(instances.Reservations[0].Instances)) - } - return &instances.Reservations[0].Instances[0], nil -} - func (p *provider) RemoveAgent(ctx context.Context, agent *woodpecker.Agent) error { - instance, err := p.getAgent(ctx, agent) + instance, region, err := p.getAgent(ctx, agent) if err != nil { return err } _, err = p.client.TerminateInstances(ctx, &ec2.TerminateInstancesInput{ InstanceIds: []string{*instance.InstanceId}, - }) + // Skip the graceful OS shutdown so an unresponsive or hung guest + // cannot block termination and leave a dangling instance behind. + SkipOsShutdown: aws.Bool(true), + }, regionOpt(region)) return err } @@ -206,23 +315,12 @@ func (p *provider) ListDeployedAgentNames(ctx context.Context) ([]string, error) log.Debug().Msgf("list deployed agent names") var names []string - instances, err := p.client.DescribeInstances(ctx, &ec2.DescribeInstancesInput{ - Filters: []ec2_types.Filter{ - { - Name: aws.String(fmt.Sprintf("tag:%s", engine.LabelPool)), - Values: []string{p.config.PoolID}, - }, - }, - }) - if err != nil { - return nil, err - } - for _, reservation := range instances.Reservations { - for _, instance := range reservation.Instances { - if instance.State.Name != ec2_types.InstanceStateNamePending && - instance.State.Name != ec2_types.InstanceStateNameRunning { - continue - } + for _, region := range p.regions { + instances, err := p.instancesByTag(ctx, region, engine.LabelPool, p.config.PoolID) + if err != nil { + return nil, err + } + for _, instance := range instances { for _, tag := range instance.Tags { if *tag.Key == "Name" { log.Debug().Msgf("found agent %s", *tag.Value) diff --git a/providers/aws/provider_test.go b/providers/aws/provider_test.go new file mode 100644 index 00000000..e919e126 --- /dev/null +++ b/providers/aws/provider_test.go @@ -0,0 +1,150 @@ +package aws + +import ( + "slices" + "testing" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/ec2" + ec2_types "github.com/aws/aws-sdk-go-v2/service/ec2/types" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + + "go.woodpecker-ci.org/autoscaler/config" + "go.woodpecker-ci.org/autoscaler/providers/aws/ec2api/mocks" + "go.woodpecker-ci.org/woodpecker/v3/woodpecker-go/woodpecker" +) + +func newDeployTestProvider(client *mocks.MockClient, candidates []deployCandidate) *provider { + p := newTestProvider(client) + p.config = &config.Config{PoolID: "1"} + p.deployCandidates = candidates + for _, candidate := range candidates { + if !slices.Contains(p.regions, candidate.regionConfig.region) { + p.regions = append(p.regions, candidate.regionConfig.region) + } + } + return p +} + +func testCandidates() []deployCandidate { + return []deployCandidate{ + { + instanceType: testArmTypeInfo, + regionConfig: regionConfig{ + region: "eu-central-1", + image: testArmImage, + subnets: []string{"subnet-arm"}, + securityGroups: []string{"sg-arm"}, + }, + }, + { + instanceType: testTypeInfo, + regionConfig: regionConfig{ + region: "us-east-1", + image: testImage, + subnets: []string{"subnet-x86"}, + securityGroups: []string{"sg-x86"}, + }, + }, + } +} + +func mockAgentVisible(client *mocks.MockClient, agentName string) { + client.On("DescribeInstances", mock.Anything, mock.Anything, mock.Anything). + Return(&ec2.DescribeInstancesOutput{ + Reservations: []ec2_types.Reservation{{Instances: []ec2_types.Instance{{ + InstanceId: aws.String("i-1"), + State: &ec2_types.InstanceState{Name: ec2_types.InstanceStateNameRunning}, + Tags: []ec2_types.Tag{{Key: aws.String("Name"), Value: aws.String(agentName)}}, + }}}}, + }, nil) +} + +func runRequest(instanceType, imageID, subnet, securityGroup string) func(*ec2.RunInstancesInput) bool { + return func(in *ec2.RunInstancesInput) bool { + return in.InstanceType == ec2_types.InstanceType(instanceType) && + aws.ToString(in.ImageId) == imageID && + aws.ToString(in.SubnetId) == subnet && + assert.ObjectsAreEqual([]string{securityGroup}, in.SecurityGroupIds) + } +} + +func TestDeployAgentFallback(t *testing.T) { + agent := &woodpecker.Agent{Name: "pool-1-agent-abcd"} + runOut := &ec2.RunInstancesOutput{Instances: []ec2_types.Instance{{InstanceId: aws.String("i-1")}}} + + t.Run("FirstCandidateSucceeds", func(t *testing.T) { + client := mocks.NewMockClient(t) + client.On("RunInstances", mock.Anything, mock.Anything, mock.Anything). + Return(runOut, nil).Once() + mockAgentVisible(client, agent.Name) + + p := newDeployTestProvider(client, testCandidates()) + assert.NoError(t, p.DeployAgent(t.Context(), agent)) + }) + + t.Run("CapacityErrorFallsBackToRegionalCandidate", func(t *testing.T) { + client := mocks.NewMockClient(t) + client.On("RunInstances", mock.Anything, + mock.MatchedBy(runRequest("t4g.micro", "ami-arm", "subnet-arm", "sg-arm")), + mock.MatchedBy(regionOptions("eu-central-1"))). + Return(nil, &apiError{code: "InsufficientInstanceCapacity"}).Once() + client.On("RunInstances", mock.Anything, + mock.MatchedBy(runRequest("m6i.large", "ami-x86", "subnet-x86", "sg-x86")), + mock.MatchedBy(regionOptions("us-east-1"))). + Return(runOut, nil).Once() + mockAgentVisible(client, agent.Name) + + p := newDeployTestProvider(client, testCandidates()) + assert.NoError(t, p.DeployAgent(t.Context(), agent)) + }) + + t.Run("NonCapacityErrorAborts", func(t *testing.T) { + client := mocks.NewMockClient(t) + client.On("RunInstances", mock.Anything, mock.Anything, mock.Anything). + Return(nil, &apiError{code: "InvalidSubnetID.NotFound"}).Once() + + p := newDeployTestProvider(client, testCandidates()) + err := p.DeployAgent(t.Context(), agent) + assert.ErrorContains(t, err, "InvalidSubnetID.NotFound") + }) + + t.Run("EmptyRunInstancesResult", func(t *testing.T) { + client := mocks.NewMockClient(t) + client.On("RunInstances", mock.Anything, mock.Anything, mock.Anything). + Return(&ec2.RunInstancesOutput{}, nil).Once() + + p := newDeployTestProvider(client, testCandidates()) + err := p.DeployAgent(t.Context(), agent) + assert.ErrorContains(t, err, "returned no instances") + }) + + t.Run("AllCandidatesOutOfCapacity", func(t *testing.T) { + client := mocks.NewMockClient(t) + client.On("RunInstances", mock.Anything, mock.Anything, mock.Anything). + Return(nil, &apiError{code: "InsufficientInstanceCapacity"}).Twice() + + p := newDeployTestProvider(client, testCandidates()) + err := p.DeployAgent(t.Context(), agent) + assert.ErrorContains(t, err, "all 2 deploy candidates out of capacity") + }) +} + +func TestRemoveAgentSkipsOSShutdown(t *testing.T) { + agent := &woodpecker.Agent{Name: "pool-1-agent-abcd"} + + client := mocks.NewMockClient(t) + mockAgentVisible(client, agent.Name) + client.On("TerminateInstances", mock.Anything, + mock.MatchedBy(func(in *ec2.TerminateInstancesInput) bool { + return assert.ObjectsAreEqual([]string{"i-1"}, in.InstanceIds) && + aws.ToBool(in.SkipOsShutdown) + }), + mock.Anything). + Return(&ec2.TerminateInstancesOutput{}, nil).Once() + + p := newTestProvider(client) + p.regions = []string{"eu-central-1"} + assert.NoError(t, p.RemoveAgent(t.Context(), agent)) +} diff --git a/providers/aws/types.go b/providers/aws/types.go new file mode 100644 index 00000000..e63d7142 --- /dev/null +++ b/providers/aws/types.go @@ -0,0 +1,35 @@ +package aws + +import ( + "errors" + + ec2_types "github.com/aws/aws-sdk-go-v2/service/ec2/types" +) + +var ( + ErrInstanceTypeNotFound = errors.New("instance type not found") + ErrAMINotFound = errors.New("AMI not found") + ErrSubnetsNotSet = errors.New("aws-subnets must be set") + ErrSubnetNotFound = errors.New("subnet not found") + ErrSecurityGroupNotFound = errors.New("security group not found") + ErrArchMismatch = errors.New("instance type architecture not supported by AMI") + ErrInstanceTypeArchitecture = errors.New("instance type must report exactly one architecture") + ErrTypeNotInRegion = errors.New("instance type not offered in region") + ErrNoDeployCandidates = errors.New("no deploy candidates resolved") + ErrRegionNotSet = errors.New("aws-region must be set for unqualified values") +) + +// regionConfig contains the resources that exist together in an AWS region. +type regionConfig struct { + region string + image ec2_types.Image + subnets []string + securityGroups []string +} + +// deployCandidate is one deploy specification the provider tries in order. +// It owns its instance type and all region-scoped resources it needs. +type deployCandidate struct { + instanceType ec2_types.InstanceTypeInfo + regionConfig regionConfig +}