From 45d225756f72ef58e1397dfdb627a098786cd769 Mon Sep 17 00:00:00 2001 From: Jeffrey Chien Date: Wed, 21 Jan 2026 16:03:57 -0500 Subject: [PATCH] Migrate EC2 metadata to SDKv2 --- .../agenthealth/handler/stats/agent/flag.go | 7 + extension/entitystore/ec2Info_test.go | 10 +- extension/entitystore/extension.go | 53 +--- extension/entitystore/extension_test.go | 41 ++- extension/entitystore/retryer.go | 18 +- extension/entitystore/retryer_test.go | 4 +- extension/entitystore/serviceprovider.go | 3 +- extension/entitystore/serviceprovider_test.go | 4 +- go.mod | 4 +- .../ec2metadataprovider.go | 95 +++--- .../ec2metadataprovider_test.go | 300 +++++++++++++++--- internal/retryer/v2/imdsretryer.go | 53 ++++ internal/retryer/v2/imdsretryer_test.go | 142 +++++++++ plugins/processors/ec2tagger/ec2tagger.go | 142 +++++---- .../processors/ec2tagger/ec2tagger_test.go | 70 ++-- .../internal/volume/describevolumes.go | 33 +- .../internal/volume/describevolumes_test.go | 31 +- .../ec2tagger/internal/volume/host_linux.go | 3 +- .../internal/volume/host_linux_test.go | 6 +- .../internal/volume/host_nonlinux.go | 3 +- .../internal/volume/host_nonlinux_test.go | 2 +- .../ec2tagger/internal/volume/merge.go | 5 +- .../ec2tagger/internal/volume/merge_test.go | 5 +- .../ec2tagger/internal/volume/volume.go | 9 +- 24 files changed, 730 insertions(+), 313 deletions(-) create mode 100644 internal/retryer/v2/imdsretryer.go create mode 100644 internal/retryer/v2/imdsretryer_test.go diff --git a/extension/agenthealth/handler/stats/agent/flag.go b/extension/agenthealth/handler/stats/agent/flag.go index 65b77632e5b..38b8cc91363 100644 --- a/extension/agenthealth/handler/stats/agent/flag.go +++ b/extension/agenthealth/handler/stats/agent/flag.go @@ -108,6 +108,8 @@ type FlagSet interface { SetValues(flags map[Flag]any) // OnChange registers a callback that triggers on flag sets. OnChange(callback func()) + // Reset allows tests to clear the stored flags + Reset() } type flagSet struct { @@ -180,6 +182,11 @@ func (p *flagSet) notify() { } } +func (p *flagSet) Reset() { + p.m.Clear() + p.notify() +} + func UsageFlags() FlagSet { flagOnce.Do(func() { flagSingleton = &flagSet{} diff --git a/extension/entitystore/ec2Info_test.go b/extension/entitystore/ec2Info_test.go index 816e56a40df..69549939a9e 100644 --- a/extension/entitystore/ec2Info_test.go +++ b/extension/entitystore/ec2Info_test.go @@ -9,14 +9,14 @@ import ( "testing" "time" - "github.com/aws/aws-sdk-go/aws/ec2metadata" + "github.com/aws/aws-sdk-go-v2/feature/ec2/imds" "github.com/stretchr/testify/assert" "go.uber.org/zap" "github.com/aws/amazon-cloudwatch-agent/internal/ec2metadataprovider" ) -var mockedInstanceIdentityDoc = &ec2metadata.EC2InstanceIdentityDocument{ +var mockedInstanceIdentityDoc = &imds.InstanceIdentityDocument{ InstanceID: "i-01d2417c27a396e44", AccountID: "874389809020", Region: "us-east-1", @@ -24,7 +24,7 @@ var mockedInstanceIdentityDoc = &ec2metadata.EC2InstanceIdentityDocument{ ImageID: "ami-09edd32d9b0990d49", } -var mockedInstanceIdentityDocWithLargeInstanceId = &ec2metadata.EC2InstanceIdentityDocument{ +var mockedInstanceIdentityDocWithLargeInstanceID = &imds.InstanceIdentityDocument{ InstanceID: "i-01d2417c27a396e44394824728", AccountID: "874389809020", Region: "us-east-1", @@ -60,12 +60,12 @@ func TestSetInstanceIDAccountID(t *testing.T) { { name: "InstanceId too large", args: args{ - metadataProvider: &mockMetadataProvider{InstanceIdentityDocument: mockedInstanceIdentityDocWithLargeInstanceId}, + metadataProvider: &mockMetadataProvider{InstanceIdentityDocument: mockedInstanceIdentityDocWithLargeInstanceID}, }, wantErr: false, want: EC2Info{ InstanceID: "", - AccountID: mockedInstanceIdentityDocWithLargeInstanceId.AccountID, + AccountID: mockedInstanceIdentityDocWithLargeInstanceID.AccountID, }, }, } diff --git a/extension/entitystore/extension.go b/extension/entitystore/extension.go index ca5a1634ac9..c4f0af66511 100644 --- a/extension/entitystore/extension.go +++ b/extension/entitystore/extension.go @@ -7,20 +7,17 @@ import ( "context" "time" + "github.com/aws/aws-sdk-go-v2/aws" "github.com/aws/aws-sdk-go-v2/service/cloudwatchlogs/types" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/aws/client" - "github.com/aws/aws-sdk-go/service/ec2" - "github.com/aws/aws-sdk-go/service/ec2/ec2iface" "github.com/jellydator/ttlcache/v3" "go.opentelemetry.io/collector/component" "go.opentelemetry.io/collector/extension" "go.uber.org/atomic" "go.uber.org/zap" - configaws "github.com/aws/amazon-cloudwatch-agent/cfg/aws" + configaws "github.com/aws/amazon-cloudwatch-agent/cfg/aws/v2" "github.com/aws/amazon-cloudwatch-agent/internal/ec2metadataprovider" - "github.com/aws/amazon-cloudwatch-agent/internal/retryer" + "github.com/aws/amazon-cloudwatch-agent/internal/retryer/v2" "github.com/aws/amazon-cloudwatch-agent/plugins/processors/awsentity/entityattributes" "github.com/aws/amazon-cloudwatch-agent/translator/config" ) @@ -35,8 +32,6 @@ const ( podTerminationCheckInterval = 5 * time.Minute ) -type ec2ProviderType func(string, *configaws.CredentialConfig) ec2iface.EC2API - type serviceProviderInterface interface { startServiceProvider() addEntryForLogFile(LogFileGlob, ServiceAttribute) @@ -69,10 +64,6 @@ type EntityStore struct { // that we can attach to the entity serviceprovider serviceProviderInterface - // nativeCredential stores the credential config for agent's native - // component such as LogAgent - nativeCredential client.ConfigProvider - metadataprovider ec2metadataprovider.MetadataProvider podTerminationCheckInterval time.Duration @@ -80,20 +71,16 @@ type EntityStore struct { var _ extension.Extension = (*EntityStore)(nil) -func (e *EntityStore) Start(ctx context.Context, host component.Host) error { +func (e *EntityStore) Start(ctx context.Context, _ component.Host) error { // Get IMDS client and EC2 API client which requires region for authentication // These will be passed down to any object that requires access to IMDS or EC2 // API client so we have single source of truth for credential e.done = make(chan struct{}) - e.metadataprovider = getMetaDataProvider() + e.metadataprovider = getMetaDataProvider(ctx) e.mode = e.config.Mode e.kubernetesMode = e.config.KubernetesMode e.podTerminationCheckInterval = podTerminationCheckInterval - ec2CredentialConfig := &configaws.CredentialConfig{ - Profile: e.config.Profile, - Filename: e.config.Filename, - } - e.serviceprovider = newServiceProvider(e.mode, e.config.Region, &e.ec2Info, e.metadataprovider, getEC2Provider, ec2CredentialConfig, e.done, e.logger) + e.serviceprovider = newServiceProvider(e.mode, e.config.Region, &e.ec2Info, e.metadataprovider, e.done, e.logger) switch e.mode { case config.ModeEC2: e.ec2Info = *newEC2Info(e.metadataprovider, e.done, e.config.Region, e.logger) @@ -138,14 +125,6 @@ func (e *EntityStore) EC2Info() EC2Info { return e.ec2Info } -func (e *EntityStore) SetNativeCredential(client client.ConfigProvider) { - e.nativeCredential = client -} - -func (e *EntityStore) NativeCredentialExists() bool { - return e.nativeCredential != nil -} - // CreateLogFileEntity creates the entity for log events that are being uploaded from a log file in the environment. func (e *EntityStore) CreateLogFileEntity(logFileGlob LogFileGlob, logGroupName LogGroupName) *types.Entity { if e.serviceprovider == nil { @@ -260,19 +239,13 @@ func (e *EntityStore) createServiceKeyAttributes(serviceAttr ServiceAttribute) m return serviceKeyAttr } -var getMetaDataProvider = func() ec2metadataprovider.MetadataProvider { - mdCredentialConfig := &configaws.CredentialConfig{} - return ec2metadataprovider.NewMetadataProvider(mdCredentialConfig.Credentials(), retryer.GetDefaultRetryNumber()) -} - -var getEC2Provider = func(region string, ec2CredentialConfig *configaws.CredentialConfig) ec2iface.EC2API { - ec2CredentialConfig.Region = region - return ec2.New( - ec2CredentialConfig.Credentials(), - &aws.Config{ - LogLevel: configaws.SDKLogLevel(), - Logger: configaws.SDKLogger{}, - }) +var getMetaDataProvider = func(ctx context.Context) ec2metadataprovider.MetadataProvider { + mdCredentialConfig := &configaws.CredentialsConfig{} + cfg, err := mdCredentialConfig.LoadConfig(ctx) + if err != nil { + cfg = aws.Config{} + } + return ec2metadataprovider.NewMetadataProvider(cfg, retryer.GetDefaultRetryNumber()) } func addNonEmptyToMap(m map[string]string, key, value string) { diff --git a/extension/entitystore/extension_test.go b/extension/entitystore/extension_test.go index 046f220eaa9..40b87718c68 100644 --- a/extension/entitystore/extension_test.go +++ b/extension/entitystore/extension_test.go @@ -12,9 +12,8 @@ import ( "testing" "time" + "github.com/aws/aws-sdk-go-v2/feature/ec2/imds" "github.com/aws/aws-sdk-go-v2/service/cloudwatchlogs/types" - "github.com/aws/aws-sdk-go/aws/ec2metadata" - "github.com/aws/aws-sdk-go/aws/session" "github.com/jellydator/ttlcache/v3" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/mock" @@ -73,17 +72,19 @@ func (s *mockServiceProvider) setAutoScalingGroup(asg string) { } type mockMetadataProvider struct { - InstanceIdentityDocument *ec2metadata.EC2InstanceIdentityDocument + InstanceIdentityDocument *imds.InstanceIdentityDocument Tags map[string]string InstanceTagError bool } -func mockMetadataProviderFunc() ec2metadataprovider.MetadataProvider { +var _ ec2metadataprovider.MetadataProvider = (*mockMetadataProvider)(nil) + +func mockMetadataProviderFunc(context.Context) ec2metadataprovider.MetadataProvider { return &mockMetadataProvider{ Tags: map[string]string{ "aws:autoscaling:groupName": "ASG-1", }, - InstanceIdentityDocument: &ec2metadata.EC2InstanceIdentityDocument{ + InstanceIdentityDocument: &imds.InstanceIdentityDocument{ InstanceID: "i-123456789", }, } @@ -91,39 +92,39 @@ func mockMetadataProviderFunc() ec2metadataprovider.MetadataProvider { func mockMetadataProviderWithAccountId(accountId string) *mockMetadataProvider { return &mockMetadataProvider{ - InstanceIdentityDocument: &ec2metadata.EC2InstanceIdentityDocument{ + InstanceIdentityDocument: &imds.InstanceIdentityDocument{ AccountID: accountId, }, } } -func (m *mockMetadataProvider) Get(ctx context.Context) (ec2metadata.EC2InstanceIdentityDocument, error) { +func (m *mockMetadataProvider) Get(context.Context) (imds.InstanceIdentityDocument, error) { if m.InstanceIdentityDocument != nil { return *m.InstanceIdentityDocument, nil } - return ec2metadata.EC2InstanceIdentityDocument{}, errors.New("No instance identity document") + return imds.InstanceIdentityDocument{}, errors.New("no instance identity document") } -func (m *mockMetadataProvider) Hostname(ctx context.Context) (string, error) { +func (m *mockMetadataProvider) Hostname(context.Context) (string, error) { return "MockHostName", nil } -func (m *mockMetadataProvider) InstanceID(ctx context.Context) (string, error) { +func (m *mockMetadataProvider) InstanceID(context.Context) (string, error) { return "MockInstanceID", nil } -func (m *mockMetadataProvider) InstanceTags(_ context.Context) ([]string, error) { +func (m *mockMetadataProvider) InstanceTags(context.Context) ([]string, error) { if m.InstanceTagError { return nil, errors.New("an error occurred for instance tag retrieval") } return maps.Keys(m.Tags), nil } -func (m *mockMetadataProvider) ClientIAMRole(ctx context.Context) (string, error) { +func (m *mockMetadataProvider) ClientIAMRole(context.Context) (string, error) { return "TestRole", nil } -func (m *mockMetadataProvider) InstanceTagValue(ctx context.Context, tagKey string) (string, error) { +func (m *mockMetadataProvider) InstanceTagValue(_ context.Context, tagKey string) (string, error) { tag, ok := m.Tags[tagKey] if !ok { return "", errors.New("tag not found") @@ -333,10 +334,9 @@ func TestEntityStore_createLogFileRID(t *testing.T) { sp.On("logFileServiceAttribute", glob, group).Return(serviceAttr) sp.On("getAutoScalingGroup").Return("ASG-1") e := EntityStore{ - mode: config.ModeEC2, - ec2Info: EC2Info{InstanceID: instanceId, AccountID: accountId}, - serviceprovider: sp, - nativeCredential: &session.Session{}, + mode: config.ModeEC2, + ec2Info: EC2Info{InstanceID: instanceId, AccountID: accountId}, + serviceprovider: sp, } entity := e.CreateLogFileEntity(glob, group) @@ -364,9 +364,8 @@ func TestEntityStore_createLogFileRID_ServiceProviderIsEmpty(t *testing.T) { glob := LogFileGlob("glob") group := LogGroupName("group") e := EntityStore{ - mode: config.ModeEC2, - ec2Info: EC2Info{InstanceID: instanceId}, - nativeCredential: &session.Session{}, + mode: config.ModeEC2, + ec2Info: EC2Info{InstanceID: instanceId}, } entity := e.CreateLogFileEntity(glob, group) @@ -548,7 +547,6 @@ func TestEntityStore_GetMetricServiceNameSource(t *testing.T) { ec2Info: EC2Info{InstanceID: instanceId}, serviceprovider: sp, metadataprovider: mockMetadataProviderWithAccountId(accountId), - nativeCredential: &session.Session{}, } serviceName, serviceNameSource := e.GetMetricServiceNameAndSource() @@ -564,7 +562,6 @@ func TestEntityStore_GetMetricServiceNameSource_ServiceProviderEmpty(t *testing. mode: config.ModeEC2, ec2Info: EC2Info{InstanceID: instanceId}, metadataprovider: mockMetadataProviderWithAccountId(accountId), - nativeCredential: &session.Session{}, } serviceName, serviceNameSource := e.GetMetricServiceNameAndSource() diff --git a/extension/entitystore/retryer.go b/extension/entitystore/retryer.go index 65829f89702..bc158834521 100644 --- a/extension/entitystore/retryer.go +++ b/extension/entitystore/retryer.go @@ -4,16 +4,16 @@ package entitystore import ( + "errors" "math/rand" "time" - "github.com/aws/aws-sdk-go/aws/awserr" + "github.com/aws/smithy-go" "go.uber.org/zap" ) const ( - RequestLimitExceeded = "RequestLimitExceeded" - infRetry = -1 + infRetry = -1 ) var ( @@ -57,7 +57,7 @@ func (r *Retryer) refreshLoop(updateFunc func() error) int { err := updateFunc() if err == nil && r.oneTime { return retry - } else if awsErr, ok := err.(awserr.Error); ok && !r.retryAnyError && !retryableErrorMap[awsErr.Code()] { + } else if !r.retryAnyError && !isRetryableError(err) { return retry } @@ -83,7 +83,15 @@ func (r *Retryer) refreshLoop(updateFunc func() error) int { } } - return retry +} + +// isRetryableError checks if the error is a retryable API error code recognized by the extension. +func isRetryableError(err error) bool { + var apiErr smithy.APIError + if errors.As(err, &apiErr) { + return retryableErrorMap[apiErr.ErrorCode()] + } + return false } // calculateWaitTime returns different time based on whether if diff --git a/extension/entitystore/retryer_test.go b/extension/entitystore/retryer_test.go index 9c8c88951ed..a84f4914290 100644 --- a/extension/entitystore/retryer_test.go +++ b/extension/entitystore/retryer_test.go @@ -7,7 +7,7 @@ import ( "testing" "time" - "github.com/aws/aws-sdk-go/aws/ec2metadata" + "github.com/aws/aws-sdk-go-v2/feature/ec2/imds" "github.com/stretchr/testify/assert" "go.uber.org/zap" @@ -30,7 +30,7 @@ func TestRetryer_refreshLoop(t *testing.T) { name: "HappyPath_CorrectRefresh", fields: fields{ metadataProvider: &mockMetadataProvider{ - InstanceIdentityDocument: &ec2metadata.EC2InstanceIdentityDocument{ + InstanceIdentityDocument: &imds.InstanceIdentityDocument{ InstanceID: "i-123456789"}, }, iamRole: "original-role", diff --git a/extension/entitystore/serviceprovider.go b/extension/entitystore/serviceprovider.go index 188d69c5d31..17f7bc03cbd 100644 --- a/extension/entitystore/serviceprovider.go +++ b/extension/entitystore/serviceprovider.go @@ -10,7 +10,6 @@ import ( "go.uber.org/zap" - configaws "github.com/aws/amazon-cloudwatch-agent/cfg/aws" "github.com/aws/amazon-cloudwatch-agent/internal/ec2metadataprovider" "github.com/aws/amazon-cloudwatch-agent/plugins/processors/ec2tagger" "github.com/aws/amazon-cloudwatch-agent/translator/config" @@ -317,7 +316,7 @@ func toLowerKeyMap(values []string) map[string]string { return set } -func newServiceProvider(mode string, region string, ec2Info *EC2Info, metadataProvider ec2metadataprovider.MetadataProvider, providerType ec2ProviderType, ec2Credential *configaws.CredentialConfig, done chan struct{}, logger *zap.Logger) serviceProviderInterface { +func newServiceProvider(mode string, region string, ec2Info *EC2Info, metadataProvider ec2metadataprovider.MetadataProvider, done chan struct{}, logger *zap.Logger) serviceProviderInterface { return &serviceprovider{ mode: mode, region: region, diff --git a/extension/entitystore/serviceprovider_test.go b/extension/entitystore/serviceprovider_test.go index 21cdb1d9859..cc67d456514 100644 --- a/extension/entitystore/serviceprovider_test.go +++ b/extension/entitystore/serviceprovider_test.go @@ -8,7 +8,7 @@ import ( "testing" "time" - "github.com/aws/aws-sdk-go/aws/ec2metadata" + "github.com/aws/aws-sdk-go-v2/feature/ec2/imds" "github.com/stretchr/testify/assert" "go.uber.org/zap" @@ -26,7 +26,7 @@ func Test_serviceprovider_startServiceProvider(t *testing.T) { { name: "HappyPath_AllServiceNames", metadataProvider: &mockMetadataProvider{ - InstanceIdentityDocument: &ec2metadata.EC2InstanceIdentityDocument{ + InstanceIdentityDocument: &imds.InstanceIdentityDocument{ InstanceID: "i-123456789"}, Tags: map[string]string{"service": "test-service"}, }, diff --git a/go.mod b/go.mod index ca87c484303..279c5ab16cc 100644 --- a/go.mod +++ b/go.mod @@ -107,8 +107,10 @@ require ( github.com/aws/aws-sdk-go-v2 v1.41.1 github.com/aws/aws-sdk-go-v2/config v1.32.7 github.com/aws/aws-sdk-go-v2/credentials v1.19.7 + github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.17 github.com/aws/aws-sdk-go-v2/service/cloudwatch v1.53.1 github.com/aws/aws-sdk-go-v2/service/cloudwatchlogs v1.63.1 + github.com/aws/aws-sdk-go-v2/service/ec2 v1.211.2 github.com/aws/aws-sdk-go-v2/service/sts v1.41.6 github.com/aws/smithy-go v1.24.0 github.com/bigkevmcd/go-configparser v0.0.0-20200217161103-d137835d2579 @@ -278,11 +280,9 @@ require ( github.com/asaskevich/govalidator v0.0.0-20230301143203-a9d515a09cc2 // indirect github.com/aws/aws-msk-iam-sasl-signer-go v1.0.1 // indirect github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.4 // indirect - github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.17 // indirect github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.17 // indirect github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.17 // indirect github.com/aws/aws-sdk-go-v2/internal/ini v1.8.4 // indirect - github.com/aws/aws-sdk-go-v2/service/ec2 v1.211.2 // indirect github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.4 // indirect github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.17 // indirect github.com/aws/aws-sdk-go-v2/service/signin v1.0.5 // indirect diff --git a/internal/ec2metadataprovider/ec2metadataprovider.go b/internal/ec2metadataprovider/ec2metadataprovider.go index 5134c1ce5a6..1102f539c7d 100644 --- a/internal/ec2metadataprovider/ec2metadataprovider.go +++ b/internal/ec2metadataprovider/ec2metadataprovider.go @@ -5,20 +5,19 @@ package ec2metadataprovider import ( "context" + "io" "log" "strings" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/aws/client" - "github.com/aws/aws-sdk-go/aws/ec2metadata" + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/feature/ec2/imds" - configaws "github.com/aws/amazon-cloudwatch-agent/cfg/aws" "github.com/aws/amazon-cloudwatch-agent/extension/agenthealth/handler/stats/agent" - "github.com/aws/amazon-cloudwatch-agent/internal/retryer" + "github.com/aws/amazon-cloudwatch-agent/internal/retryer/v2" ) type MetadataProvider interface { - Get(ctx context.Context) (ec2metadata.EC2InstanceIdentityDocument, error) + Get(ctx context.Context) (imds.InstanceIdentityDocument, error) Hostname(ctx context.Context) (string, error) InstanceID(ctx context.Context) (string, error) InstanceTags(ctx context.Context) ([]string, error) @@ -27,51 +26,44 @@ type MetadataProvider interface { } type metadataClient struct { - metadataFallbackDisabled *ec2metadata.EC2Metadata - metadataFallbackEnabled *ec2metadata.EC2Metadata + v2Client *imds.Client + v1Client *imds.Client } var _ MetadataProvider = (*metadataClient)(nil) -func NewMetadataProvider(p client.ConfigProvider, retries int) MetadataProvider { - disableFallbackConfig := &aws.Config{ - LogLevel: configaws.SDKLogLevel(), - Logger: configaws.SDKLogger{}, - Retryer: retryer.NewIMDSRetryer(retries), - EC2MetadataEnableFallback: aws.Bool(false), - } - enableFallbackConfig := &aws.Config{ - LogLevel: configaws.SDKLogLevel(), - Logger: configaws.SDKLogger{}, - } +func NewMetadataProvider(cfg aws.Config, retries int) MetadataProvider { + return newMetadataProvider(cfg, retries) +} + +func newMetadataProvider(cfg aws.Config, retries int, optFns ...func(*imds.Options)) MetadataProvider { + v2Options := append(optFns, func(o *imds.Options) { + o.Retryer = retryer.NewIMDSRetryer(retries) + o.EnableFallback = aws.FalseTernary + }) + v1Options := append(optFns, func(o *imds.Options) { + o.EnableFallback = aws.TrueTernary + }) return &metadataClient{ - metadataFallbackDisabled: ec2metadata.New(p, disableFallbackConfig), - metadataFallbackEnabled: ec2metadata.New(p, enableFallbackConfig), + v2Client: imds.NewFromConfig(cfg, v2Options...), + v1Client: imds.NewFromConfig(cfg, v1Options...), } } func (c *metadataClient) InstanceID(ctx context.Context) (string, error) { - return withMetadataFallbackRetry(ctx, c, func(metadataClient *ec2metadata.EC2Metadata) (string, error) { - return metadataClient.GetMetadataWithContext(ctx, "instance-id") - }) + return c.getMetadata(ctx, "instance-id") } func (c *metadataClient) Hostname(ctx context.Context) (string, error) { - return withMetadataFallbackRetry(ctx, c, func(metadataClient *ec2metadata.EC2Metadata) (string, error) { - return metadataClient.GetMetadataWithContext(ctx, "hostname") - }) + return c.getMetadata(ctx, "hostname") } func (c *metadataClient) ClientIAMRole(ctx context.Context) (string, error) { - return withMetadataFallbackRetry(ctx, c, func(metadataClient *ec2metadata.EC2Metadata) (string, error) { - return metadataClient.GetMetadataWithContext(ctx, "iam/security-credentials") - }) + return c.getMetadata(ctx, "iam/security-credentials") } func (c *metadataClient) InstanceTags(ctx context.Context) ([]string, error) { - tags, err := withMetadataFallbackRetry(ctx, c, func(metadataClient *ec2metadata.EC2Metadata) (string, error) { - return metadataClient.GetMetadataWithContext(ctx, "tags/instance") - }) + tags, err := c.getMetadata(ctx, "tags/instance") if err != nil { return nil, err } @@ -79,23 +71,40 @@ func (c *metadataClient) InstanceTags(ctx context.Context) ([]string, error) { } func (c *metadataClient) InstanceTagValue(ctx context.Context, tagKey string) (string, error) { - path := "tags/instance/" + tagKey - return withMetadataFallbackRetry(ctx, c, func(metadataClient *ec2metadata.EC2Metadata) (string, error) { - return metadataClient.GetMetadataWithContext(ctx, path) + return c.getMetadata(ctx, "tags/instance/"+tagKey) +} + +func (c *metadataClient) Get(ctx context.Context) (imds.InstanceIdentityDocument, error) { + return withMetadataFallbackRetry(c, func(client *imds.Client) (imds.InstanceIdentityDocument, error) { + out, err := client.GetInstanceIdentityDocument(ctx, &imds.GetInstanceIdentityDocumentInput{}) + if err != nil { + return imds.InstanceIdentityDocument{}, err + } + return out.InstanceIdentityDocument, nil }) } -func (c *metadataClient) Get(ctx context.Context) (ec2metadata.EC2InstanceIdentityDocument, error) { - return withMetadataFallbackRetry(ctx, c, func(metadataClient *ec2metadata.EC2Metadata) (ec2metadata.EC2InstanceIdentityDocument, error) { - return metadataClient.GetInstanceIdentityDocumentWithContext(ctx) +func (c *metadataClient) getMetadata(ctx context.Context, path string) (string, error) { + return withMetadataFallbackRetry(c, func(client *imds.Client) (string, error) { + out, err := client.GetMetadata(ctx, &imds.GetMetadataInput{ + Path: path, + }) + if err != nil { + return "", err + } + content, err := io.ReadAll(out.Content) + if err != nil { + return "", err + } + return string(content), nil }) } -func withMetadataFallbackRetry[T any](ctx context.Context, c *metadataClient, operation func(*ec2metadata.EC2Metadata) (T, error)) (T, error) { - result, err := operation(c.metadataFallbackDisabled) +func withMetadataFallbackRetry[T any](c *metadataClient, fn func(*imds.Client) (T, error)) (T, error) { + result, err := fn(c.v2Client) if err != nil { - log.Printf("D! could not perform operation without imds v1 fallback enable thus enable fallback") - result, err = operation(c.metadataFallbackEnabled) + log.Printf("D! Could not perform operation without IMDS v1 fallback enabled. Enabling fallback.") + result, err = fn(c.v1Client) if err == nil { agent.UsageFlags().Set(agent.FlagIMDSFallbackSuccess) } diff --git a/internal/ec2metadataprovider/ec2metadataprovider_test.go b/internal/ec2metadataprovider/ec2metadataprovider_test.go index 5252ce6c1d4..458a90fe519 100644 --- a/internal/ec2metadataprovider/ec2metadataprovider_test.go +++ b/internal/ec2metadataprovider/ec2metadataprovider_test.go @@ -4,69 +4,277 @@ package ec2metadataprovider import ( - "context" - "os" - "reflect" + "fmt" + "net/http" + "net/http/httptest" + "strings" "testing" - "github.com/aws/aws-sdk-go/aws/ec2metadata" - "github.com/aws/aws-sdk-go/aws/session" - "github.com/aws/aws-sdk-go/awstesting/mock" + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/feature/ec2/imds" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/aws/amazon-cloudwatch-agent/extension/agenthealth/handler/stats/agent" ) +func mockIMDSServer(t *testing.T, v2Enabled bool, responses map[string]string) *httptest.Server { + return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + t.Helper() + if r.URL.Path == "/latest/api/token" { + if !v2Enabled { + w.WriteHeader(http.StatusForbidden) + return + } + w.WriteHeader(http.StatusOK) + w.Write([]byte("test-token")) + return + } + + path := strings.TrimPrefix(r.URL.Path, "/latest/meta-data/") + if strings.HasPrefix(r.URL.Path, "/latest/dynamic/instance-identity/document") { + path = "instance-identity/document" + } + + if response, ok := responses[path]; ok { + w.WriteHeader(http.StatusOK) + w.Write([]byte(response)) + return + } + + w.WriteHeader(http.StatusNotFound) + })) +} + +func createTestProvider(serverURL string, retries int) MetadataProvider { + return newMetadataProvider(aws.Config{}, retries, func(o *imds.Options) { + o.Endpoint = serverURL + }) +} + func TestMetadataProvider_Get(t *testing.T) { - tests := []struct { - name string - ctx context.Context - sess *session.Session - expectDoc ec2metadata.EC2InstanceIdentityDocument + instanceDoc := `{ + "instanceId": "i-1234567890abcdef0", + "region": "us-west-2", + "availabilityZone": "us-west-2a", + "instanceType": "t3.micro" + }` + + testCases := map[string]struct { + v2Enabled bool + wantFallback bool + }{ + "v2_enabled": { + v2Enabled: true, + wantFallback: false, + }, + "v2_disabled_fallback_to_v1": { + v2Enabled: false, + wantFallback: true, + }, + } + + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + server := mockIMDSServer(t, testCase.v2Enabled, map[string]string{ + "instance-identity/document": instanceDoc, + }) + defer server.Close() + + provider := createTestProvider(server.URL, 1) + doc, err := provider.Get(t.Context()) + + require.NoError(t, err) + assert.Equal(t, "i-1234567890abcdef0", doc.InstanceID) + assert.Equal(t, "us-west-2", doc.Region) + assert.Equal(t, "us-west-2a", doc.AvailabilityZone) + assert.Equal(t, "t3.micro", doc.InstanceType) + + if testCase.wantFallback { + assert.True(t, agent.UsageFlags().IsSet(agent.FlagIMDSFallbackSuccess)) + } + }) + } +} + +func TestMetadataProvider_InstanceID(t *testing.T) { + testCases := map[string]struct { + v2Enabled bool + instanceID string + wantFallback bool + }{ + "v2_enabled": { + v2Enabled: true, + instanceID: "i-1234567890abcdef0", + wantFallback: false, + }, + "v2_disabled_fallback_to_v1": { + v2Enabled: false, + instanceID: "i-0987654321fedcba0", + wantFallback: true, + }, + } + + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + server := mockIMDSServer(t, testCase.v2Enabled, map[string]string{ + "instance-id": testCase.instanceID, + }) + defer server.Close() + + provider := createTestProvider(server.URL, 1) + instanceID, err := provider.InstanceID(t.Context()) + + require.NoError(t, err) + assert.Equal(t, testCase.instanceID, instanceID) + + if testCase.wantFallback { + assert.True(t, agent.UsageFlags().IsSet(agent.FlagIMDSFallbackSuccess)) + } + }) + } +} + +func TestMetadataProvider_Hostname(t *testing.T) { + testCases := map[string]struct { + v2Enabled bool + hostname string }{ - { - name: "mock session", - ctx: context.Background(), - sess: mock.Session, - expectDoc: ec2metadata.EC2InstanceIdentityDocument{}, + "v2_enabled": { + v2Enabled: true, + hostname: "ip-10-0-0-1.us-west-2.compute.internal", + }, + "v2_disabled_fallback_to_v1": { + v2Enabled: false, + hostname: "ip-10-0-0-2.us-west-2.compute.internal", }, } - for _, tc := range tests { - t.Run(tc.name, func(t *testing.T) { - c := NewMetadataProvider(tc.sess, 0) - gotDoc, err := c.Get(tc.ctx) - assert.NotNil(t, err) - assert.Truef(t, reflect.DeepEqual(gotDoc, tc.expectDoc), "get() gotDoc: %v, expected: %v", gotDoc, tc.expectDoc) + + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + server := mockIMDSServer(t, testCase.v2Enabled, map[string]string{ + "hostname": testCase.hostname, + }) + defer server.Close() + + provider := createTestProvider(server.URL, 1) + hostname, err := provider.Hostname(t.Context()) + + require.NoError(t, err) + assert.Equal(t, testCase.hostname, hostname) }) } } -func TestMetadataProvider_available(t *testing.T) { - tests := []struct { - name string - ctx context.Context - sess *session.Session - want error +func TestMetadataProvider_InstanceTags(t *testing.T) { + testCases := map[string]struct { + v2Enabled bool + tagsString string + wantTags []string }{ - { - name: "mock session", - ctx: context.Background(), - sess: mock.Session, - want: nil, + "v2_enabled_multiple_tags": { + v2Enabled: true, + tagsString: "Name\nEnvironment\nApplication", + wantTags: []string{"Name", "Environment", "Application"}, + }, + "v2_disabled_fallback_to_v1": { + v2Enabled: false, + tagsString: "Tag1\nTag2", + wantTags: []string{"Tag1", "Tag2"}, }, } - // For build environments where IMDS is disabled via environment variable, explicitly re-enable it. Otherwise the - // call to c.InstanceId() fails before even contacting the mock session. - // See https://docs.aws.amazon.com/cli/latest/userguide/cli-configure-envvars.html#envvars-list-AWS_EC2_METADATA_DISABLED - const awsEc2MetadataDisabledEnvVar = "AWS_EC2_METADATA_DISABLED" - val := os.Getenv(awsEc2MetadataDisabledEnvVar) - defer func() { assert.NoError(t, os.Setenv(awsEc2MetadataDisabledEnvVar, val)) }() - assert.NoError(t, os.Setenv(awsEc2MetadataDisabledEnvVar, "false")) - - for _, tc := range tests { - t.Run(tc.name, func(t *testing.T) { - c := NewMetadataProvider(tc.sess, 0) - _, err := c.InstanceID(tc.ctx) - assert.ErrorIs(t, err, tc.want) + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + server := mockIMDSServer(t, testCase.v2Enabled, map[string]string{ + "tags/instance": testCase.tagsString, + }) + defer server.Close() + + provider := createTestProvider(server.URL, 1) + tags, err := provider.InstanceTags(t.Context()) + + require.NoError(t, err) + assert.Equal(t, testCase.wantTags, tags) }) } } + +func TestMetadataProvider_ClientIAMRole(t *testing.T) { + testCases := map[string]struct { + v2Enabled bool + roleName string + }{ + "v2_enabled": { + v2Enabled: true, + roleName: "MyInstanceRole", + }, + "v2_disabled_fallback_to_v1": { + v2Enabled: false, + roleName: "AnotherRole", + }, + } + + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + server := mockIMDSServer(t, testCase.v2Enabled, map[string]string{ + "iam/security-credentials": testCase.roleName, + }) + defer server.Close() + + provider := createTestProvider(server.URL, 1) + roleName, err := provider.ClientIAMRole(t.Context()) + + require.NoError(t, err) + assert.Equal(t, testCase.roleName, roleName) + }) + } +} + +func TestMetadataProvider_InstanceTagValue(t *testing.T) { + testCases := map[string]struct { + v2Enabled bool + tagKey string + tagValue string + }{ + "v2_enabled": { + v2Enabled: true, + tagKey: "Name", + tagValue: "my-instance", + }, + "v2_disabled_fallback_to_v1": { + v2Enabled: false, + tagKey: "Environment", + tagValue: "production", + }, + } + + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + server := mockIMDSServer(t, testCase.v2Enabled, map[string]string{ + fmt.Sprintf("tags/instance/%s", testCase.tagKey): testCase.tagValue, + }) + defer server.Close() + + provider := createTestProvider(server.URL, 1) + tagValue, err := provider.InstanceTagValue(t.Context(), testCase.tagKey) + + require.NoError(t, err) + assert.Equal(t, testCase.tagValue, tagValue) + }) + } +} + +func TestMetadataProvider_ErrorHandling(t *testing.T) { + t.Run("both_v2_and_v1_fail", func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusNotFound) + })) + defer server.Close() + + provider := createTestProvider(server.URL, 0) + _, err := provider.InstanceID(t.Context()) + + assert.Error(t, err) + }) +} diff --git a/internal/retryer/v2/imdsretryer.go b/internal/retryer/v2/imdsretryer.go new file mode 100644 index 00000000000..6ae3ac96041 --- /dev/null +++ b/internal/retryer/v2/imdsretryer.go @@ -0,0 +1,53 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: MIT + +package retryer + +import ( + "errors" + "fmt" + "os" + "strconv" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/aws/retry" + smithyhttp "github.com/aws/smithy-go/transport/http" + + "github.com/aws/amazon-cloudwatch-agent/cfg/envconfig" +) + +const ( + DefaultMetadataRetries = 1 +) + +type IMDSRetryer struct { + *retry.Standard +} + +var _ aws.RetryerV2 = (*IMDSRetryer)(nil) + +// NewIMDSRetryer allows us to retry IMDS errors +func NewIMDSRetryer(retries int) *IMDSRetryer { + fmt.Printf("I! IMDS retry client will retry %d times", retries) + return &IMDSRetryer{ + Standard: retry.NewStandard(func(options *retry.StandardOptions) { + options.MaxAttempts = retries + 1 // MaxAttempts include the first attempt + }), + } +} + +func (r *IMDSRetryer) IsErrorRetryable(err error) bool { + // SDKv2 returns a ResponseError on request failure + // https://github.com/aws/aws-sdk-go-v2/blob/dcbed91b6c6235022f15eda6ea526dbb91e1cb81/feature/ec2/imds/request_middleware.go#L185-L191 + var responseErr *smithyhttp.ResponseError + return errors.As(err, &responseErr) || r.Standard.IsErrorRetryable(err) +} + +func GetDefaultRetryNumber() int { + imdsRetryEnv := os.Getenv(envconfig.IMDS_NUMBER_RETRY) + imdsRetry, err := strconv.Atoi(imdsRetryEnv) + if err == nil && imdsRetry >= 0 { + return imdsRetry + } + return DefaultMetadataRetries +} diff --git a/internal/retryer/v2/imdsretryer_test.go b/internal/retryer/v2/imdsretryer_test.go new file mode 100644 index 00000000000..c9d38112b5d --- /dev/null +++ b/internal/retryer/v2/imdsretryer_test.go @@ -0,0 +1,142 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: MIT + +package retryer + +import ( + "errors" + "net/http" + "testing" + + smithyhttp "github.com/aws/smithy-go/transport/http" + "github.com/stretchr/testify/assert" + + "github.com/aws/amazon-cloudwatch-agent/cfg/envconfig" +) + +func TestIMDSRetryer_IsErrorRetryable(t *testing.T) { + testCases := map[string]struct { + err error + want bool + }{ + "ErrorIsNilNotRetryable": { + err: nil, + want: false, + }, + "ErrorIsIMDSResponseErrorRetryable": { + err: &smithyhttp.ResponseError{ + Response: &smithyhttp.Response{ + Response: &http.Response{ + StatusCode: 404, + }, + }, + Err: errors.New("request to EC2 IMDS failed"), + }, + want: true, + }, + "ErrorIsIMDSResponseError5xxRetryable": { + err: &smithyhttp.ResponseError{ + Response: &smithyhttp.Response{ + Response: &http.Response{ + StatusCode: 500, + }, + }, + Err: errors.New("request to EC2 IMDS failed"), + }, + want: true, + }, + "ErrorIsWrappedIMDSResponseErrorRetryable": { + err: errors.Join( + errors.New("outer error"), + &smithyhttp.ResponseError{ + Response: &smithyhttp.Response{ + Response: &http.Response{ + StatusCode: 503, + }, + }, + Err: errors.New("request to EC2 IMDS failed"), + }, + ), + want: true, + }, + "ErrorIsGenericErrorNotRetryableByDefault": { + err: errors.New("some other error"), + want: false, // Standard retryer doesn't treat generic errors as retryable by default + }, + } + + retryer := NewIMDSRetryer(DefaultMetadataRetries) + + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + got := retryer.IsErrorRetryable(testCase.err) + assert.Equal(t, testCase.want, got) + }) + } +} + +func TestIMDSRetryer_MaxAttempts(t *testing.T) { + testCases := map[string]struct { + retries int + want int + }{ + "DefaultRetries": { + retries: DefaultMetadataRetries, + want: DefaultMetadataRetries + 1, + }, + "TwoRetries": { + retries: 2, + want: 3, + }, + "ZeroRetries": { + retries: 0, + want: 1, + }, + "FiveRetries": { + retries: 5, + want: 6, + }, + } + + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + retryer := NewIMDSRetryer(testCase.retries) + assert.Equal(t, testCase.want, retryer.MaxAttempts()) + }) + } +} + +func TestGetDefaultRetryNumber(t *testing.T) { + testCases := map[string]struct { + envValue string + expectedRetries int + }{ + "EmptyEnvUsesDefault": { + expectedRetries: DefaultMetadataRetries, + }, + "NegativeEnvUsesDefault": { + envValue: "-1", + expectedRetries: DefaultMetadataRetries, + }, + "InvalidEnvUsesDefault": { + envValue: "not an int", + expectedRetries: DefaultMetadataRetries, + }, + "ZeroEnvValue": { + envValue: "0", + expectedRetries: 0, + }, + "ValidEnvValue": { + envValue: "2", + expectedRetries: 2, + }, + } + + for name, testCase := range testCases { + t.Run(name, func(t *testing.T) { + t.Setenv(envconfig.IMDS_NUMBER_RETRY, testCase.envValue) + + assert.Equal(t, testCase.expectedRetries, GetDefaultRetryNumber()) + }) + } +} diff --git a/plugins/processors/ec2tagger/ec2tagger.go b/plugins/processors/ec2tagger/ec2tagger.go index c93cc7cd98a..81c6f7394cb 100644 --- a/plugins/processors/ec2tagger/ec2tagger.go +++ b/plugins/processors/ec2tagger/ec2tagger.go @@ -11,15 +11,15 @@ import ( "time" "github.com/amazon-contributing/opentelemetry-collector-contrib/extension/awsmiddleware" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/service/ec2" - "github.com/aws/aws-sdk-go/service/ec2/ec2iface" + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/ec2" + "github.com/aws/aws-sdk-go-v2/service/ec2/types" "go.opentelemetry.io/collector/component" "go.opentelemetry.io/collector/pdata/pcommon" "go.opentelemetry.io/collector/pdata/pmetric" "go.uber.org/zap" - configaws "github.com/aws/amazon-cloudwatch-agent/cfg/aws" + configaws "github.com/aws/amazon-cloudwatch-agent/cfg/aws/v2" "github.com/aws/amazon-cloudwatch-agent/internal/ec2metadataprovider" "github.com/aws/amazon-cloudwatch-agent/plugins/processors/ec2tagger/internal/volume" translatorCtx "github.com/aws/amazon-cloudwatch-agent/translator/context" @@ -38,7 +38,13 @@ type ec2MetadataRespondType struct { region string } -type ec2ProviderType func(*configaws.CredentialConfig) ec2iface.EC2API +// EC2APIClient defines the interface for EC2 API operations needed by the tagger +type EC2APIClient interface { + ec2.DescribeTagsAPIClient + ec2.DescribeVolumesAPIClient +} + +type ec2ProviderType func(ctx context.Context, host component.Host, credentialConfig *configaws.CredentialsConfig) EC2APIClient type Tagger struct { *Config @@ -53,8 +59,8 @@ type Tagger struct { started bool ec2MetadataLookup ec2MetadataLookupType ec2MetadataRespond ec2MetadataRespondType - tagFilters []*ec2.Filter - ec2API ec2iface.EC2API + tagFilters []types.Filter + ec2API EC2APIClient volumeSerialCache volume.Cache Configurer *awsmiddleware.Configurer @@ -64,52 +70,37 @@ type Tagger struct { // newTagger returns a new EC2 Tagger processor. func newTagger(config *Config, logger *zap.Logger) *Tagger { _, cancel := context.WithCancel(context.Background()) - mdCredentialConfig := &configaws.CredentialConfig{} + mdCredentialConfig := &configaws.CredentialsConfig{} + + mdCfg, err := mdCredentialConfig.LoadConfig(context.Background()) + if err != nil { + logger.Error("ec2tagger: Failed to load AWS config for metadata provider", zap.Error(err)) + } + p := &Tagger{ Config: config, logger: logger, cancelFunc: cancel, - metadataProvider: ec2metadataprovider.NewMetadataProvider(mdCredentialConfig.Credentials(), config.IMDSRetries), - ec2Provider: func(ec2CredentialConfig *configaws.CredentialConfig) ec2iface.EC2API { - return ec2.New( - ec2CredentialConfig.Credentials(), - &aws.Config{ - LogLevel: configaws.SDKLogLevel(), - Logger: configaws.SDKLogger{}, - }) - }, + metadataProvider: ec2metadataprovider.NewMetadataProvider(mdCfg, config.IMDSRetries), } + p.ec2Provider = p.createEC2Client return p } -func getOtelAttributes(m pmetric.Metric) []pcommon.Map { - attributes := []pcommon.Map{} - switch m.Type() { - case pmetric.MetricTypeGauge: - dps := m.Gauge().DataPoints() - for i := 0; i < dps.Len(); i++ { - attributes = append(attributes, dps.At(i).Attributes()) - } - case pmetric.MetricTypeSum: - dps := m.Sum().DataPoints() - for i := 0; i < dps.Len(); i++ { - attributes = append(attributes, dps.At(i).Attributes()) - } - case pmetric.MetricTypeHistogram: - dps := m.Histogram().DataPoints() - for i := 0; i < dps.Len(); i++ { - attributes = append(attributes, dps.At(i).Attributes()) - } - case pmetric.MetricTypeExponentialHistogram: - dps := m.ExponentialHistogram().DataPoints() - for i := 0; i < dps.Len(); i++ { - attributes = append(attributes, dps.At(i).Attributes()) - } +func (t *Tagger) createEC2Client(ctx context.Context, host component.Host, credentialConfig *configaws.CredentialsConfig) EC2APIClient { + cfg, err := credentialConfig.LoadConfig(ctx) + if err != nil { + cfg = aws.Config{} } - return attributes + + if t.MiddlewareID != nil { + awsmiddleware.TryConfigure(t.logger, host, *t.MiddlewareID, awsmiddleware.SDKv2(&cfg)) + } + + return ec2.NewFromConfig(cfg) } -func (t *Tagger) processMetrics(ctx context.Context, md pmetric.Metrics) (pmetric.Metrics, error) { +func (t *Tagger) processMetrics(_ context.Context, md pmetric.Metrics) (pmetric.Metrics, error) { // grab the pointer to the map in case it gets refreshed while we're applying this round of metrics. At least // this batch then will all get the same tags. t.RLock() @@ -133,6 +124,34 @@ func (t *Tagger) processMetrics(ctx context.Context, md pmetric.Metrics) (pmetri return md, nil } +func getOtelAttributes(m pmetric.Metric) []pcommon.Map { + var attributes []pcommon.Map + switch m.Type() { + case pmetric.MetricTypeGauge: + dps := m.Gauge().DataPoints() + for i := 0; i < dps.Len(); i++ { + attributes = append(attributes, dps.At(i).Attributes()) + } + case pmetric.MetricTypeSum: + dps := m.Sum().DataPoints() + for i := 0; i < dps.Len(); i++ { + attributes = append(attributes, dps.At(i).Attributes()) + } + case pmetric.MetricTypeHistogram: + dps := m.Histogram().DataPoints() + for i := 0; i < dps.Len(); i++ { + attributes = append(attributes, dps.At(i).Attributes()) + } + case pmetric.MetricTypeExponentialHistogram: + dps := m.ExponentialHistogram().DataPoints() + for i := 0; i < dps.Len(); i++ { + attributes = append(attributes, dps.At(i).Attributes()) + } + default: + } + return attributes +} + // updateOtelAttributes adds tags and the requested dimensions to the attributes of each // DataPoint. We add and remove at the DataPoint level instead of resource level because this is // where the receiver/adapter does. @@ -175,14 +194,15 @@ func (t *Tagger) updateOtelAttributes(attributes []pcommon.Map) { } // updateTags calls EC2 Describe Tags and replaces the Tagger's tagCache with the newly retrieved values -func (t *Tagger) updateTags() error { +func (t *Tagger) updateTags(ctx context.Context) error { tags := make(map[string]string) - input := &ec2.DescribeTagsInput{ + + paginator := ec2.NewDescribeTagsPaginator(t.ec2API, &ec2.DescribeTagsInput{ Filters: t.tagFilters, - } + }) - for { - result, err := t.ec2API.DescribeTags(input) + for paginator.HasMorePages() { + result, err := paginator.NextPage(ctx) if err != nil { return err } @@ -194,11 +214,8 @@ func (t *Tagger) updateTags() error { } tags[key] = *tag.Value } - if result.NextToken == nil { - break - } - input.SetNextToken(*result.NextToken) } + t.Lock() defer t.Unlock() t.ec2TagCache = tags @@ -234,7 +251,7 @@ func (t *Tagger) refreshLoopTags(refreshInterval time.Duration, stopAfterFirstSu } if refreshTags { - if err := t.updateTags(); err != nil { + if err := t.updateTags(context.Background()); err != nil { t.logger.Warn("ec2tagger: Error refreshing EC2 tags, keeping old values", zap.Error(err)) } } @@ -324,14 +341,14 @@ func (t *Tagger) Start(ctx context.Context, host component.Host) error { if err := t.deriveEC2MetadataFromIMDS(ctx); err != nil { return err } - t.tagFilters = []*ec2.Filter{ + t.tagFilters = []types.Filter{ { Name: aws.String("resource-type"), - Values: aws.StringSlice([]string{"instance"}), + Values: []string{"instance"}, }, { Name: aws.String("resource-id"), - Values: aws.StringSlice([]string{t.ec2MetadataRespond.instanceId}), + Values: []string{t.ec2MetadataRespond.instanceId}, }, } // if the customer said 'AutoScalingGroupName' (the CW dimension), do what they mean not what they said @@ -345,13 +362,13 @@ func (t *Tagger) Start(ctx context.Context, host component.Host) error { } } - t.tagFilters = append(t.tagFilters, &ec2.Filter{ + t.tagFilters = append(t.tagFilters, types.Filter{ Name: aws.String("key"), - Values: aws.StringSlice(t.EC2InstanceTagKeys), + Values: t.EC2InstanceTagKeys, }) } if len(t.EC2InstanceTagKeys) > 0 || len(t.EBSDeviceKeys) > 0 { - ec2CredentialConfig := &configaws.CredentialConfig{ + ec2CredentialConfig := &configaws.CredentialsConfig{ AccessKey: t.AccessKey, SecretKey: t.SecretKey, RoleARN: t.RoleARN, @@ -360,13 +377,8 @@ func (t *Tagger) Start(ctx context.Context, host component.Host) error { Token: t.Token, Region: t.ec2MetadataRespond.region, } - t.ec2API = t.ec2Provider(ec2CredentialConfig) - if client, ok := t.ec2API.(*ec2.EC2); ok { - if t.MiddlewareID != nil { - awsmiddleware.TryConfigure(t.logger, host, *t.MiddlewareID, awsmiddleware.SDKv1(&client.Handlers)) - } - } + t.ec2API = t.ec2Provider(ctx, host, ec2CredentialConfig) go func() { //Async start of initial retrieval to prevent block of agent start t.initialRetrievalOfTagsAndVolumes() @@ -528,7 +540,7 @@ func (t *Tagger) initialRetrievalOfTagsAndVolumes() { } if !tagsRetrieved { - if err := t.updateTags(); err != nil { + if err := t.updateTags(context.Background()); err != nil { t.logger.Warn("ec2tagger: Unable to describe ec2 tags for initial retrieval", zap.Error(err)) } else { tagsRetrieved = true diff --git a/plugins/processors/ec2tagger/ec2tagger_test.go b/plugins/processors/ec2tagger/ec2tagger_test.go index 96362272635..fbd12dec47a 100644 --- a/plugins/processors/ec2tagger/ec2tagger_test.go +++ b/plugins/processors/ec2tagger/ec2tagger_test.go @@ -10,9 +10,9 @@ import ( "testing" "time" - "github.com/aws/aws-sdk-go/aws/ec2metadata" - "github.com/aws/aws-sdk-go/service/ec2" - "github.com/aws/aws-sdk-go/service/ec2/ec2iface" + "github.com/aws/aws-sdk-go-v2/feature/ec2/imds" + "github.com/aws/aws-sdk-go-v2/service/ec2" + "github.com/aws/aws-sdk-go-v2/service/ec2/types" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "go.opentelemetry.io/collector/component" @@ -22,11 +22,10 @@ import ( "go.opentelemetry.io/collector/processor/processortest" "golang.org/x/exp/maps" - configaws "github.com/aws/amazon-cloudwatch-agent/cfg/aws" + configaws "github.com/aws/amazon-cloudwatch-agent/cfg/aws/v2" ) type mockEC2Client struct { - ec2iface.EC2API //The following fields are used to control how the mocked DescribeTags api behave: //tagsCallCount records how many times DescribeTags has been called //if tagsCallCount <= tagsFailLimit, DescribeTags call fails @@ -39,49 +38,50 @@ type mockEC2Client struct { UseUpdatedTags bool } -// construct the return results for the mocked DescribeTags api +var _ EC2APIClient = (*mockEC2Client)(nil) + var ( tagKey1 = "tagKey1" tagVal1 = "tagVal1" - tagDes1 = ec2.TagDescription{Key: &tagKey1, Value: &tagVal1} + tagDes1 = types.TagDescription{Key: &tagKey1, Value: &tagVal1} ) var ( tagKey2 = "tagKey2" tagVal2 = "tagVal2" - tagDes2 = ec2.TagDescription{Key: &tagKey2, Value: &tagVal2} + tagDes2 = types.TagDescription{Key: &tagKey2, Value: &tagVal2} ) var ( tagKey3 = "aws:autoscaling:groupName" tagVal3 = "ASG-1" - tagDes3 = ec2.TagDescription{Key: &tagKey3, Value: &tagVal3} + tagDes3 = types.TagDescription{Key: &tagKey3, Value: &tagVal3} ) var ( updatedTagVal2 = "updated-tagVal2" - updatedTagDes2 = ec2.TagDescription{Key: &tagKey2, Value: &updatedTagVal2} + updatedTagDes2 = types.TagDescription{Key: &tagKey2, Value: &updatedTagVal2} ) -func (m *mockEC2Client) DescribeTags(*ec2.DescribeTagsInput) (*ec2.DescribeTagsOutput, error) { +func (m *mockEC2Client) DescribeTags(context.Context, *ec2.DescribeTagsInput, ...func(*ec2.Options)) (*ec2.DescribeTagsOutput, error) { //partial tags returned when the DescribeTags api are called initially //some tags are not returned because customer just attach them to the ec2 instance //and the api doesn't know about them yet partialTags := ec2.DescribeTagsOutput{ NextToken: nil, - Tags: []*ec2.TagDescription{&tagDes1}, + Tags: []types.TagDescription{tagDes1}, } //all tags are returned when the ec2 metadata service knows about all tags allTags := ec2.DescribeTagsOutput{ NextToken: nil, - Tags: []*ec2.TagDescription{&tagDes1, &tagDes2, &tagDes3}, + Tags: []types.TagDescription{tagDes1, tagDes2, tagDes3}, } //later customer changes the value of the second tag and DescribeTags api returns updated tags allTagsUpdated := ec2.DescribeTagsOutput{ NextToken: nil, - Tags: []*ec2.TagDescription{&tagDes1, &updatedTagDes2, &tagDes3}, + Tags: []types.TagDescription{tagDes1, updatedTagDes2, tagDes3}, } //return error initially to simulate the case @@ -111,6 +111,10 @@ func (m *mockEC2Client) DescribeTags(*ec2.DescribeTagsInput) (*ec2.DescribeTagsO return nil, nil } +func (m *mockEC2Client) DescribeVolumes(context.Context, *ec2.DescribeVolumesInput, ...func(*ec2.Options)) (*ec2.DescribeVolumesOutput, error) { + return &ec2.DescribeVolumesOutput{}, nil +} + // construct the return results for the mocked DescribeTags api var ( device1 = "xvdc" @@ -127,37 +131,37 @@ var ( ) type mockMetadataProvider struct { - InstanceIdentityDocument *ec2metadata.EC2InstanceIdentityDocument + InstanceIdentityDocument *imds.InstanceIdentityDocument } -func (m *mockMetadataProvider) Get(ctx context.Context) (ec2metadata.EC2InstanceIdentityDocument, error) { +func (m *mockMetadataProvider) Get(context.Context) (imds.InstanceIdentityDocument, error) { if m.InstanceIdentityDocument != nil { return *m.InstanceIdentityDocument, nil } - return ec2metadata.EC2InstanceIdentityDocument{}, errors.New("No instance identity document") + return imds.InstanceIdentityDocument{}, errors.New("No instance identity document") } -func (m *mockMetadataProvider) Hostname(ctx context.Context) (string, error) { +func (m *mockMetadataProvider) Hostname(context.Context) (string, error) { return "MockHostName", nil } -func (m *mockMetadataProvider) InstanceID(ctx context.Context) (string, error) { +func (m *mockMetadataProvider) InstanceID(context.Context) (string, error) { return "MockInstanceID", nil } -func (m *mockMetadataProvider) InstanceTags(_ context.Context) ([]string, error) { +func (m *mockMetadataProvider) InstanceTags(context.Context) ([]string, error) { return []string{"MockInstanceTag"}, nil } -func (m *mockMetadataProvider) InstanceTagValue(ctx context.Context, tagKey string) (string, error) { +func (m *mockMetadataProvider) InstanceTagValue(context.Context, string) (string, error) { return "MockInstanceValue", nil } -func (m *mockMetadataProvider) ClientIAMRole(ctx context.Context) (string, error) { +func (m *mockMetadataProvider) ClientIAMRole(context.Context) (string, error) { return "MockIAMRole", nil } -var mockedInstanceIdentityDoc = &ec2metadata.EC2InstanceIdentityDocument{ +var mockedInstanceIdentityDoc = &imds.InstanceIdentityDocument{ InstanceID: "i-01d2417c27a396e44", Region: "us-east-1", InstanceType: "m5ad.large", @@ -295,7 +299,7 @@ func TestStartSuccessWithNoTagsVolumesUpdate(t *testing.T) { tagsPartialLimit: 1, UseUpdatedTags: false, } - ec2Provider := func(*configaws.CredentialConfig) ec2iface.EC2API { + ec2Provider := func(context.Context, component.Host, *configaws.CredentialsConfig) EC2APIClient { return ec2Client } volumeCache := &mockVolumeCache{cache: make(map[string]string)} @@ -340,7 +344,7 @@ func TestStartSuccessWithTagsVolumesUpdate(t *testing.T) { tagsPartialLimit: 2, UseUpdatedTags: false, } - ec2Provider := func(*configaws.CredentialConfig) ec2iface.EC2API { + ec2Provider := func(context.Context, component.Host, *configaws.CredentialsConfig) EC2APIClient { return ec2Client } volumeCache := &mockVolumeCache{cache: make(map[string]string)} @@ -397,7 +401,7 @@ func TestStartSuccessWithWildcardTagVolumeKey(t *testing.T) { tagsPartialLimit: 1, UseUpdatedTags: false, } - ec2Provider := func(*configaws.CredentialConfig) ec2iface.EC2API { + ec2Provider := func(context.Context, component.Host, *configaws.CredentialsConfig) EC2APIClient { return ec2Client } volumeCache := &mockVolumeCache{cache: make(map[string]string)} @@ -443,7 +447,7 @@ func TestApplyWithTagsVolumesUpdate(t *testing.T) { tagsPartialLimit: 1, UseUpdatedTags: false, } - ec2Provider := func(*configaws.CredentialConfig) ec2iface.EC2API { + ec2Provider := func(context.Context, component.Host, *configaws.CredentialsConfig) EC2APIClient { return ec2Client } volumeCache := &mockVolumeCache{cache: make(map[string]string)} @@ -546,7 +550,7 @@ func TestMetricsDroppedBeforeStarted(t *testing.T) { tagsPartialLimit: 1, UseUpdatedTags: false, } - ec2Provider := func(*configaws.CredentialConfig) ec2iface.EC2API { + ec2Provider := func(context.Context, component.Host, *configaws.CredentialsConfig) EC2APIClient { return ec2Client } volumeCache := &mockVolumeCache{cache: make(map[string]string)} @@ -562,13 +566,13 @@ func TestMetricsDroppedBeforeStarted(t *testing.T) { } md := createTestMetrics([]map[string]string{ - map[string]string{ + { "host": "example.org", }, - map[string]string{ + { "device": device1, }, - map[string]string{ + { "device": device2, }, }) @@ -612,7 +616,7 @@ func TestTaggerStartDoesNotBlock(t *testing.T) { tagsPartialLimit: 1, UseUpdatedTags: false, } - ec2Provider := func(*configaws.CredentialConfig) ec2iface.EC2API { + ec2Provider := func(context.Context, component.Host, *configaws.CredentialsConfig) EC2APIClient { return ec2Client } BackoffSleepArray = []time.Duration{1 * time.Minute, 1 * time.Minute, 1 * time.Minute, 3 * time.Minute, 3 * time.Minute, 3 * time.Minute, 10 * time.Minute} @@ -657,7 +661,7 @@ func TestExistingAttributesNotOverwritten(t *testing.T) { tagsPartialLimit: 1, UseUpdatedTags: false, } - ec2Provider := func(*configaws.CredentialConfig) ec2iface.EC2API { + ec2Provider := func(context.Context, component.Host, *configaws.CredentialsConfig) EC2APIClient { return ec2Client } volumeCache := &mockVolumeCache{cache: make(map[string]string)} diff --git a/plugins/processors/ec2tagger/internal/volume/describevolumes.go b/plugins/processors/ec2tagger/internal/volume/describevolumes.go index 11108c996b6..bc06658779f 100644 --- a/plugins/processors/ec2tagger/internal/volume/describevolumes.go +++ b/plugins/processors/ec2tagger/internal/volume/describevolumes.go @@ -4,48 +4,47 @@ package volume import ( + "context" "fmt" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/service/ec2" - "github.com/aws/aws-sdk-go/service/ec2/ec2iface" + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/ec2" + "github.com/aws/aws-sdk-go-v2/service/ec2/types" ) type describeVolumesProvider struct { - ec2Client ec2iface.EC2API + ec2Client ec2.DescribeVolumesAPIClient instanceID string } -func newDescribeVolumesProvider(ec2Client ec2iface.EC2API, instanceID string) Provider { +func newDescribeVolumesProvider(ec2Client ec2.DescribeVolumesAPIClient, instanceID string) Provider { return &describeVolumesProvider{ec2Client: ec2Client, instanceID: instanceID} } -func (p *describeVolumesProvider) DeviceToSerialMap() (map[string]string, error) { +func (p *describeVolumesProvider) DeviceToSerialMap(ctx context.Context) (map[string]string, error) { result := map[string]string{} - input := &ec2.DescribeVolumesInput{ - Filters: []*ec2.Filter{ + paginator := ec2.NewDescribeVolumesPaginator(p.ec2Client, &ec2.DescribeVolumesInput{ + Filters: []types.Filter{ { Name: aws.String("attachment.instance-id"), - Values: aws.StringSlice([]string{p.instanceID}), + Values: []string{p.instanceID}, }, }, - } - for { - output, err := p.ec2Client.DescribeVolumes(input) + }) + + for paginator.HasMorePages() { + output, err := paginator.NextPage(ctx) if err != nil { return nil, fmt.Errorf("unable to describe volumes: %w", err) } for _, volume := range output.Volumes { for _, attachment := range volume.Attachments { if attachment.Device != nil && attachment.VolumeId != nil { - result[aws.StringValue(attachment.Device)] = aws.StringValue(attachment.VolumeId) + result[*attachment.Device] = *attachment.VolumeId } } } - if output.NextToken == nil { - break - } - input.SetNextToken(*output.NextToken) } + return result, nil } diff --git a/plugins/processors/ec2tagger/internal/volume/describevolumes_test.go b/plugins/processors/ec2tagger/internal/volume/describevolumes_test.go index d83452f1a77..aa181f6228b 100644 --- a/plugins/processors/ec2tagger/internal/volume/describevolumes_test.go +++ b/plugins/processors/ec2tagger/internal/volume/describevolumes_test.go @@ -4,11 +4,12 @@ package volume import ( + "context" "errors" "testing" - "github.com/aws/aws-sdk-go/service/ec2" - "github.com/aws/aws-sdk-go/service/ec2/ec2iface" + "github.com/aws/aws-sdk-go-v2/service/ec2" + "github.com/aws/aws-sdk-go-v2/service/ec2/types" "github.com/stretchr/testify/assert" ) @@ -16,10 +17,10 @@ import ( var ( device1 = "/dev/xvdc" volumeId1 = "vol-0303a1cc896c42d28" - volumeAttachment1 = ec2.VolumeAttachment{Device: &device1, VolumeId: &volumeId1} + volumeAttachment1 = types.VolumeAttachment{Device: &device1, VolumeId: &volumeId1} availabilityZone = "us-east-1a" - volume1 = ec2.Volume{ - Attachments: []*ec2.VolumeAttachment{&volumeAttachment1}, + volume1 = types.Volume{ + Attachments: []types.VolumeAttachment{volumeAttachment1}, AvailabilityZone: &availabilityZone, } ) @@ -27,21 +28,21 @@ var ( var ( device2 = "/dev/xvdf" volumeId2 = "vol-0c241693efb58734a" - volumeAttachment2 = ec2.VolumeAttachment{Device: &device2, VolumeId: &volumeId2} - volume2 = ec2.Volume{ - Attachments: []*ec2.VolumeAttachment{&volumeAttachment2}, + volumeAttachment2 = types.VolumeAttachment{Device: &device2, VolumeId: &volumeId2} + volume2 = types.Volume{ + Attachments: []types.VolumeAttachment{volumeAttachment2}, AvailabilityZone: &availabilityZone, } ) type mockEC2Client struct { - ec2iface.EC2API - callCount int err error } -func (m *mockEC2Client) DescribeVolumes(input *ec2.DescribeVolumesInput) (*ec2.DescribeVolumesOutput, error) { +var _ ec2.DescribeVolumesAPIClient = (*mockEC2Client)(nil) + +func (m *mockEC2Client) DescribeVolumes(_ context.Context, input *ec2.DescribeVolumesInput, _ ...func(*ec2.Options)) (*ec2.DescribeVolumesOutput, error) { m.callCount++ if m.err != nil { @@ -51,26 +52,26 @@ func (m *mockEC2Client) DescribeVolumes(input *ec2.DescribeVolumesInput) (*ec2.D if input.NextToken == nil { return &ec2.DescribeVolumesOutput{ NextToken: &device2, - Volumes: []*ec2.Volume{&volume1}, + Volumes: []types.Volume{volume1}, }, nil } return &ec2.DescribeVolumesOutput{ NextToken: nil, - Volumes: []*ec2.Volume{&volume2}, + Volumes: []types.Volume{volume2}, }, nil } func TestDescribeVolumesProvider(t *testing.T) { ec2Client := &mockEC2Client{} p := newDescribeVolumesProvider(ec2Client, "") - got, err := p.DeviceToSerialMap() + got, err := p.DeviceToSerialMap(t.Context()) assert.NoError(t, err) assert.Equal(t, 2, ec2Client.callCount) want := map[string]string{device1: volumeId1, device2: volumeId2} assert.Equal(t, want, got) ec2Client.err = errors.New("test") ec2Client.callCount = 0 - got, err = p.DeviceToSerialMap() + got, err = p.DeviceToSerialMap(t.Context()) assert.Error(t, err) assert.Equal(t, 1, ec2Client.callCount) assert.Nil(t, got) diff --git a/plugins/processors/ec2tagger/internal/volume/host_linux.go b/plugins/processors/ec2tagger/internal/volume/host_linux.go index 170fa56acdb..5f57a3b21cd 100644 --- a/plugins/processors/ec2tagger/internal/volume/host_linux.go +++ b/plugins/processors/ec2tagger/internal/volume/host_linux.go @@ -7,6 +7,7 @@ package volume import ( "bytes" + "context" "errors" "fmt" "os" @@ -36,7 +37,7 @@ func newHostProvider() Provider { } } -func (p *hostProvider) DeviceToSerialMap() (map[string]string, error) { +func (p *hostProvider) DeviceToSerialMap(context.Context) (map[string]string, error) { result := map[string]string{} dirs, err := p.osReadDir(sysBlockPath) if err != nil { diff --git a/plugins/processors/ec2tagger/internal/volume/host_linux_test.go b/plugins/processors/ec2tagger/internal/volume/host_linux_test.go index 0e6b34e5892..5c3c1178f8b 100644 --- a/plugins/processors/ec2tagger/internal/volume/host_linux_test.go +++ b/plugins/processors/ec2tagger/internal/volume/host_linux_test.go @@ -63,15 +63,15 @@ func TestHostProvider(t *testing.T) { p := newHostProvider().(*hostProvider) p.osReadDir = m.ReadDir p.osReadFile = m.ReadFile - got, err := p.DeviceToSerialMap() + got, err := p.DeviceToSerialMap(t.Context()) assert.Error(t, err) assert.Nil(t, got) m.errDir = nil - got, err = p.DeviceToSerialMap() + got, err = p.DeviceToSerialMap(t.Context()) assert.Error(t, err) assert.Nil(t, got) m.serialMap = testSerialMap - got, err = p.DeviceToSerialMap() + got, err = p.DeviceToSerialMap(t.Context()) assert.NoError(t, err) assert.Equal(t, map[string]string{ "xvdc": "vol-0303a1cc896c42d28", diff --git a/plugins/processors/ec2tagger/internal/volume/host_nonlinux.go b/plugins/processors/ec2tagger/internal/volume/host_nonlinux.go index 24c00ba2e28..3968f4ad80d 100644 --- a/plugins/processors/ec2tagger/internal/volume/host_nonlinux.go +++ b/plugins/processors/ec2tagger/internal/volume/host_nonlinux.go @@ -6,6 +6,7 @@ package volume import ( + "context" "errors" ) @@ -16,6 +17,6 @@ func newHostProvider() Provider { return &hostProvider{} } -func (*hostProvider) DeviceToSerialMap() (map[string]string, error) { +func (*hostProvider) DeviceToSerialMap(context.Context) (map[string]string, error) { return nil, errors.New("local block device retrieval only supported on linux") } diff --git a/plugins/processors/ec2tagger/internal/volume/host_nonlinux_test.go b/plugins/processors/ec2tagger/internal/volume/host_nonlinux_test.go index 43c1df013c6..2a106b4c65f 100644 --- a/plugins/processors/ec2tagger/internal/volume/host_nonlinux_test.go +++ b/plugins/processors/ec2tagger/internal/volume/host_nonlinux_test.go @@ -13,7 +13,7 @@ import ( func TestHostProvider(t *testing.T) { p := newHostProvider() - got, err := p.DeviceToSerialMap() + got, err := p.DeviceToSerialMap(t.Context()) assert.Error(t, err) assert.Nil(t, got) } diff --git a/plugins/processors/ec2tagger/internal/volume/merge.go b/plugins/processors/ec2tagger/internal/volume/merge.go index ccce0cecdd1..bed9e83145b 100644 --- a/plugins/processors/ec2tagger/internal/volume/merge.go +++ b/plugins/processors/ec2tagger/internal/volume/merge.go @@ -4,6 +4,7 @@ package volume import ( + "context" "errors" "github.com/aws/amazon-cloudwatch-agent/internal/util/collections" @@ -17,11 +18,11 @@ func newMergeProvider(providers []Provider) Provider { return &mergeProvider{providers: providers} } -func (p *mergeProvider) DeviceToSerialMap() (map[string]string, error) { +func (p *mergeProvider) DeviceToSerialMap(ctx context.Context) (map[string]string, error) { var errs error results := make([]map[string]string, 0, len(p.providers)) for _, provider := range p.providers { - if result, err := provider.DeviceToSerialMap(); err != nil { + if result, err := provider.DeviceToSerialMap(ctx); err != nil { errs = errors.Join(errs, err) } else { results = append(results, result) diff --git a/plugins/processors/ec2tagger/internal/volume/merge_test.go b/plugins/processors/ec2tagger/internal/volume/merge_test.go index a6e09054d2c..e85baf74270 100644 --- a/plugins/processors/ec2tagger/internal/volume/merge_test.go +++ b/plugins/processors/ec2tagger/internal/volume/merge_test.go @@ -4,6 +4,7 @@ package volume import ( + "context" "errors" "testing" @@ -15,7 +16,7 @@ type mockProvider struct { err error } -func (m *mockProvider) DeviceToSerialMap() (map[string]string, error) { +func (m *mockProvider) DeviceToSerialMap(context.Context) (map[string]string, error) { return m.serialMap, m.err } @@ -66,7 +67,7 @@ func TestMergeProvider(t *testing.T) { for name, testCase := range testCases { t.Run(name, func(t *testing.T) { p := newMergeProvider(testCase.providers) - got, err := p.DeviceToSerialMap() + got, err := p.DeviceToSerialMap(t.Context()) assert.ErrorIs(t, err, testCase.wantErr) assert.Equal(t, testCase.wantSerialMap, got) }) diff --git a/plugins/processors/ec2tagger/internal/volume/volume.go b/plugins/processors/ec2tagger/internal/volume/volume.go index b19f3df81d0..abc30077a04 100644 --- a/plugins/processors/ec2tagger/internal/volume/volume.go +++ b/plugins/processors/ec2tagger/internal/volume/volume.go @@ -4,6 +4,7 @@ package volume import ( + "context" "errors" "fmt" "os" @@ -11,7 +12,7 @@ import ( "strings" "sync" - "github.com/aws/aws-sdk-go/service/ec2/ec2iface" + "github.com/aws/aws-sdk-go-v2/service/ec2" "golang.org/x/exp/maps" ) @@ -21,10 +22,10 @@ var ( type Provider interface { // DeviceToSerialMap provides a map with device name keys and serial number values. - DeviceToSerialMap() (map[string]string, error) + DeviceToSerialMap(context.Context) (map[string]string, error) } -func NewProvider(ec2Client ec2iface.EC2API, instanceID string) Provider { +func NewProvider(ec2Client ec2.DescribeVolumesAPIClient, instanceID string) Provider { return newMergeProvider([]Provider{ newHostProvider(), newDescribeVolumesProvider(ec2Client, instanceID), @@ -71,7 +72,7 @@ func (c *cache) Refresh() error { if c.provider == nil { return errNoProviders } - result, err := c.provider.DeviceToSerialMap() + result, err := c.provider.DeviceToSerialMap(context.Background()) if err != nil { return fmt.Errorf("unable to refresh volume cache: %w", err) }