diff --git a/go.mod b/go.mod index 2a712339f..16b1620d2 100644 --- a/go.mod +++ b/go.mod @@ -7,6 +7,7 @@ require ( github.com/99designs/gqlgen v0.17.66 github.com/aws/aws-sdk-go-v2 v1.36.3 github.com/aws/aws-sdk-go-v2/credentials v1.17.62 + github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.16.30 github.com/aws/aws-sdk-go-v2/service/s3 v1.78.1 github.com/go-chi/chi/v5 v5.2.1 github.com/go-chi/cors v1.2.1 @@ -23,7 +24,6 @@ require ( require ( github.com/agnivade/levenshtein v1.2.1 // indirect github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.6.10 // indirect - github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.16.30 // indirect github.com/aws/aws-sdk-go-v2/internal/configsources v1.3.34 // indirect github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.6.34 // indirect github.com/aws/aws-sdk-go-v2/internal/v4a v1.3.34 // indirect diff --git a/go.sum b/go.sum index 6a50e905a..ac271d08c 100644 --- a/go.sum +++ b/go.sum @@ -148,8 +148,6 @@ github.com/xrash/smetrics v0.0.0-20240521201337-686a1a2994c1 h1:gEOO8jv9F4OT7lGC github.com/xrash/smetrics v0.0.0-20240521201337-686a1a2994c1/go.mod h1:Ohn+xnUBiLI6FVj/9LpzZWtj1/D6lUovWYBkxHVV3aM= go.gearno.de/crypto/uuid v0.1.0 h1:94BYg7GYItJ6yYZ1GJayb3VYhI9/FjxuR1nFaduR4hE= go.gearno.de/crypto/uuid v0.1.0/go.mod h1:fnIIvKO9QnsyLO3ZJLJT3r8KZv/p0FOeT5eZKilYWXg= -go.gearno.de/kit v0.0.0-20250308215532-082d0731efae h1:CrC8b2quMUZQ0ry16N40rvz7EktnEpOyBJrAId0Hdc8= -go.gearno.de/kit v0.0.0-20250308215532-082d0731efae/go.mod h1:RsqqVkwq+p4rmtOfYLX8NmQ+kIDi0tgyzAr2GsuLUnE= go.gearno.de/kit v0.0.0-20250313103045-779e525d954c h1:zriqh+c5QMPBL6fZtSxKNLgULooI9Oxnyh4IL9y0Fzs= go.gearno.de/kit v0.0.0-20250313103045-779e525d954c/go.mod h1:RsqqVkwq+p4rmtOfYLX8NmQ+kIDi0tgyzAr2GsuLUnE= go.gearno.de/x/panicf v0.1.1 h1:E3Cr9NB8Ry2EsvEG/1eHr7kplP3tEjTf5d56dTX64VQ= diff --git a/pkg/awsconfig/config.go b/pkg/awsconfig/config.go index b0672cf69..b20ecd7c6 100644 --- a/pkg/awsconfig/config.go +++ b/pkg/awsconfig/config.go @@ -15,12 +15,16 @@ package awsconfig import ( + "context" "net/http" + "os" "time" "github.com/aws/aws-sdk-go-v2/aws" "github.com/aws/aws-sdk-go-v2/credentials" "github.com/aws/aws-sdk-go-v2/credentials/ec2rolecreds" + "github.com/aws/aws-sdk-go-v2/credentials/endpointcreds" + "github.com/aws/aws-sdk-go-v2/feature/ec2/imds" "go.gearno.de/kit/httpclient" "go.gearno.de/kit/log" ) @@ -68,14 +72,35 @@ func NewConfig(logger *log.Logger, httpClient *http.Client, opts Options) aws.Co // cfg.Logger = logger TODO: add logger interface for aws if opts.AccessKeyID != "" && opts.SecretAccessKey != "" { + // Use static credentials if provided cfg.Credentials = credentials.NewStaticCredentialsProvider( opts.AccessKeyID, opts.SecretAccessKey, opts.SessionName, ) } else { + imdsClient := imds.New(imds.Options{HTTPClient: httpClient}) + + ec2Provider := ec2rolecreds.New(func(options *ec2rolecreds.Options) { + options.Client = imdsClient + }) + + ecsCredentialsURI := os.Getenv("AWS_CONTAINER_CREDENTIALS_RELATIVE_URI") + ecsProvider := endpointcreds.New("http://169.254.170.2"+ecsCredentialsURI, + func(options *endpointcreds.Options) { + options.HTTPClient = httpClient + }, + ) + cfg.Credentials = aws.NewCredentialsCache( - ec2rolecreds.New(), + aws.CredentialsProviderFunc(func(ctx context.Context) (aws.Credentials, error) { + creds, err := ecsProvider.Retrieve(ctx) + if err == nil { + return creds, nil + } + + return ec2Provider.Retrieve(ctx) + }), func(o *aws.CredentialsCacheOptions) { o.ExpiryWindow = 10 * time.Minute },