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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .mockery.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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:
6 changes: 3 additions & 3 deletions go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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
Expand Down
125 changes: 125 additions & 0 deletions providers/aws/credentials_test.go
Original file line number Diff line number Diff line change
@@ -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)
})
}
19 changes: 19 additions & 0 deletions providers/aws/ec2api/api.go
Original file line number Diff line number Diff line change
@@ -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)
}
Loading