diff --git a/lib/integrations/awsoidc/access_graph_aws_sync.go b/lib/integrations/awsoidc/access_graph_aws_sync.go index cf996ae605a..046ad3e7dd4 100644 --- a/lib/integrations/awsoidc/access_graph_aws_sync.go +++ b/lib/integrations/awsoidc/access_graph_aws_sync.go @@ -64,13 +64,13 @@ func (r *AccessGraphAWSIAMConfigureRequest) CheckAndSetDefaults() error { // AccessGraphIAMConfigureClient describes the required methods to create the IAM Policies // required for enrolling Access Graph AWS Sync into Teleport. type AccessGraphIAMConfigureClient interface { - callerIdentityGetter + CallerIdentityGetter // PutRolePolicy creates or replaces a Policy by its name in a IAM Role. PutRolePolicy(ctx context.Context, params *iam.PutRolePolicyInput, optFns ...func(*iam.Options)) (*iam.PutRolePolicyOutput, error) } type defaultTAGIAMConfigureClient struct { - callerIdentityGetter + CallerIdentityGetter *iam.Client } @@ -82,7 +82,7 @@ func NewAccessGraphIAMConfigureClient(ctx context.Context) (AccessGraphIAMConfig } return &defaultTAGIAMConfigureClient{ - callerIdentityGetter: sts.NewFromConfig(cfg), + CallerIdentityGetter: sts.NewFromConfig(cfg), Client: iam.NewFromConfig(cfg), }, nil } @@ -96,7 +96,7 @@ func ConfigureAccessGraphSyncIAM(ctx context.Context, clt AccessGraphIAMConfigur return trace.Wrap(err) } - if err := checkAccountID(ctx, clt, req.AccountID); err != nil { + if err := CheckAccountID(ctx, clt, req.AccountID); err != nil { return trace.Wrap(err) } diff --git a/lib/integrations/awsoidc/access_graph_aws_sync_test.go b/lib/integrations/awsoidc/access_graph_aws_sync_test.go index ac82c3709f9..e55e2f25baa 100644 --- a/lib/integrations/awsoidc/access_graph_aws_sync_test.go +++ b/lib/integrations/awsoidc/access_graph_aws_sync_test.go @@ -129,7 +129,7 @@ func TestAccessGraphAWSIAMConfig(t *testing.T) { } { t.Run(tt.name, func(t *testing.T) { clt := mockAccessGraphAWSAMConfigClient{ - callerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, + CallerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, existingRoles: tt.mockExistingRoles, } @@ -140,7 +140,7 @@ func TestAccessGraphAWSIAMConfig(t *testing.T) { } type mockAccessGraphAWSAMConfigClient struct { - callerIdentityGetter + CallerIdentityGetter existingRoles []string } diff --git a/lib/integrations/awsoidc/aws_app_access_iam_config.go b/lib/integrations/awsoidc/aws_app_access_iam_config.go index 5fa5dc6f802..09e5e269726 100644 --- a/lib/integrations/awsoidc/aws_app_access_iam_config.go +++ b/lib/integrations/awsoidc/aws_app_access_iam_config.go @@ -70,14 +70,14 @@ func (r *AWSAppAccessConfigureRequest) CheckAndSetDefaults() error { // AWSAppAccessConfigureClient describes the required methods to create the IAM Policies required for AWS App Access. type AWSAppAccessConfigureClient interface { - callerIdentityGetter + CallerIdentityGetter // PutRolePolicy creates or replaces a Policy by its name in a IAM Role. PutRolePolicy(ctx context.Context, params *iam.PutRolePolicyInput, optFns ...func(*iam.Options)) (*iam.PutRolePolicyOutput, error) } type defaultAWSAppAccessConfigureClient struct { *iam.Client - callerIdentityGetter + CallerIdentityGetter } // NewAWSAppAccessConfigureClient creates a new AWSAppAccessConfigureClient. @@ -101,7 +101,7 @@ func NewAWSAppAccessConfigureClient(ctx context.Context) (AWSAppAccessConfigureC return &defaultAWSAppAccessConfigureClient{ Client: iam.NewFromConfig(cfg), - callerIdentityGetter: sts.NewFromConfig(cfg), + CallerIdentityGetter: sts.NewFromConfig(cfg), }, nil } @@ -116,7 +116,7 @@ func ConfigureAWSAppAccess(ctx context.Context, awsClient AWSAppAccessConfigureC return trace.Wrap(err) } - if err := checkAccountID(ctx, awsClient, req.AccountID); err != nil { + if err := CheckAccountID(ctx, awsClient, req.AccountID); err != nil { return trace.Wrap(err) } diff --git a/lib/integrations/awsoidc/aws_app_access_iam_config_test.go b/lib/integrations/awsoidc/aws_app_access_iam_config_test.go index fb9c6b37003..95e26547653 100644 --- a/lib/integrations/awsoidc/aws_app_access_iam_config_test.go +++ b/lib/integrations/awsoidc/aws_app_access_iam_config_test.go @@ -129,7 +129,7 @@ func TestAWSAppAccessConfig(t *testing.T) { } { t.Run(tt.name, func(t *testing.T) { awsClient := &mockAWSAppAccessConfigClient{ - callerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, + CallerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, existingRoles: tt.mockExistingRoles, } @@ -140,7 +140,7 @@ func TestAWSAppAccessConfig(t *testing.T) { } type mockAWSAppAccessConfigClient struct { - callerIdentityGetter + CallerIdentityGetter existingRoles []string } diff --git a/lib/integrations/awsoidc/clients.go b/lib/integrations/awsoidc/clients.go index 43b2a72c1f4..82c84136ecb 100644 --- a/lib/integrations/awsoidc/clients.go +++ b/lib/integrations/awsoidc/clients.go @@ -177,15 +177,15 @@ func (j IdentityToken) GetIdentityToken() ([]byte, error) { return []byte(j), nil } -type callerIdentityGetter interface { +// CallerIdentityGetter is a subset of [sts.Client] that can be used to information about the caller identity. +type CallerIdentityGetter interface { // GetCallerIdentity returns information about the caller identity. GetCallerIdentity(ctx context.Context, params *sts.GetCallerIdentityInput, optFns ...func(*sts.Options)) (*sts.GetCallerIdentityOutput, error) } -// checkAccountID is a helper func that check if the current caller account ID -// matches the expected account ID, in order to check that the command was run -// in the expected AWS account. -func checkAccountID(ctx context.Context, clt callerIdentityGetter, wantAccountID string) error { +// CheckAccountID is a helper func that check if the current caller account ID +// matches the expected account ID. +func CheckAccountID(ctx context.Context, clt CallerIdentityGetter, wantAccountID string) error { if wantAccountID == "" { return nil } diff --git a/lib/integrations/awsoidc/clients_test.go b/lib/integrations/awsoidc/clients_test.go index f1c4c75e428..ade39eeecdf 100644 --- a/lib/integrations/awsoidc/clients_test.go +++ b/lib/integrations/awsoidc/clients_test.go @@ -19,6 +19,7 @@ package awsoidc import ( + "context" "testing" "github.com/gravitational/trace" @@ -56,3 +57,42 @@ func TestCheckAndSetDefaults(t *testing.T) { require.NoError(t, err) }) } + +func TestCheckAccountID(t *testing.T) { + t.Parallel() + tests := []struct { + desc string + clt CallerIdentityGetter + accountID string + wantErrContains string + }{ + { + desc: "caller identity matches expected account ID", + clt: mockSTSClient{accountID: "123456789012"}, + accountID: "123456789012", + }, + { + desc: "empty expected account ID is always valid", + clt: mockSTSClient{accountID: "123456789012"}, + }, + { + desc: "caller identity does not match expected account ID", + clt: mockSTSClient{accountID: "123456789012"}, + accountID: "222222222222", + wantErrContains: "expected account ID 222222222222 but current account ID is 123456789012", + }, + } + + ctx := context.Background() + for _, test := range tests { + t.Run(test.desc, func(t *testing.T) { + err := CheckAccountID(ctx, test.clt, test.accountID) + if test.wantErrContains != "" { + require.Error(t, err) + require.ErrorContains(t, err, test.wantErrContains) + return + } + require.NoError(t, err) + }) + } +} diff --git a/lib/integrations/awsoidc/deployservice.go b/lib/integrations/awsoidc/deployservice.go index fbd7e9a1d22..31c86ce6b9a 100644 --- a/lib/integrations/awsoidc/deployservice.go +++ b/lib/integrations/awsoidc/deployservice.go @@ -322,11 +322,11 @@ type DeployServiceClient interface { // Before deploying the service, it must ensure that the token exists and has the appropriate token rul. TokenService - callerIdentityGetter + CallerIdentityGetter } type defaultDeployServiceClient struct { - callerIdentityGetter + CallerIdentityGetter *ecs.Client tokenServiceClient TokenService } @@ -355,7 +355,7 @@ func NewDeployServiceClient(ctx context.Context, clientReq *AWSClientRequest, to return &defaultDeployServiceClient{ Client: ecsClient, - callerIdentityGetter: stsClient, + CallerIdentityGetter: stsClient, tokenServiceClient: tokenServiceClient, }, nil } diff --git a/lib/integrations/awsoidc/deployservice_iam_config.go b/lib/integrations/awsoidc/deployservice_iam_config.go index e639c20a47d..c520ddeded2 100644 --- a/lib/integrations/awsoidc/deployservice_iam_config.go +++ b/lib/integrations/awsoidc/deployservice_iam_config.go @@ -113,7 +113,7 @@ func (r *DeployServiceIAMConfigureRequest) CheckAndSetDefaults() error { // DeployServiceIAMConfigureClient describes the required methods to create the IAM Roles/Policies required for the DeployService action. type DeployServiceIAMConfigureClient interface { - callerIdentityGetter + CallerIdentityGetter // CreateRole creates a new IAM Role. CreateRole(ctx context.Context, params *iam.CreateRoleInput, optFns ...func(*iam.Options)) (*iam.CreateRoleOutput, error) @@ -124,7 +124,7 @@ type DeployServiceIAMConfigureClient interface { type defaultDeployServiceIAMConfigureClient struct { *iam.Client - callerIdentityGetter + CallerIdentityGetter } // NewDeployServiceIAMConfigureClient creates a new DeployServiceIAMConfigureClient. @@ -140,7 +140,7 @@ func NewDeployServiceIAMConfigureClient(ctx context.Context, region string) (Dep return &defaultDeployServiceIAMConfigureClient{ Client: iam.NewFromConfig(cfg), - callerIdentityGetter: sts.NewFromConfig(cfg), + CallerIdentityGetter: sts.NewFromConfig(cfg), }, nil } @@ -169,7 +169,7 @@ func ConfigureDeployServiceIAM(ctx context.Context, clt DeployServiceIAMConfigur return trace.Wrap(err) } req.AccountID = aws.ToString(callerIdentity.Account) - } else if err := checkAccountID(ctx, clt, req.AccountID); err != nil { + } else if err := CheckAccountID(ctx, clt, req.AccountID); err != nil { return trace.Wrap(err) } diff --git a/lib/integrations/awsoidc/deployservice_iam_config_test.go b/lib/integrations/awsoidc/deployservice_iam_config_test.go index b72dfcafa4d..718fb6de71a 100644 --- a/lib/integrations/awsoidc/deployservice_iam_config_test.go +++ b/lib/integrations/awsoidc/deployservice_iam_config_test.go @@ -206,7 +206,7 @@ func TestDeployServiceIAMConfig(t *testing.T) { } { t.Run(tt.name, func(t *testing.T) { clt := mockDeployServiceIAMConfigClient{ - callerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, + CallerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, existingRoles: tt.mockExistingRoles, } @@ -217,7 +217,7 @@ func TestDeployServiceIAMConfig(t *testing.T) { } type mockDeployServiceIAMConfigClient struct { - callerIdentityGetter + CallerIdentityGetter existingRoles []string } diff --git a/lib/integrations/awsoidc/ec2_ssm_iam_config.go b/lib/integrations/awsoidc/ec2_ssm_iam_config.go index d86a61f06e0..e1be3382693 100644 --- a/lib/integrations/awsoidc/ec2_ssm_iam_config.go +++ b/lib/integrations/awsoidc/ec2_ssm_iam_config.go @@ -108,7 +108,7 @@ func (r *EC2SSMIAMConfigureRequest) CheckAndSetDefaults() error { // EC2SSMConfigureClient describes the required methods to create the IAM Policies and SSM Document required for installing Teleport in EC2 instances. type EC2SSMConfigureClient interface { - callerIdentityGetter + CallerIdentityGetter // PutRolePolicy creates or replaces a Policy by its name in a IAM Role. PutRolePolicy(ctx context.Context, params *iam.PutRolePolicyInput, optFns ...func(*iam.Options)) (*iam.PutRolePolicyOutput, error) @@ -120,7 +120,7 @@ type EC2SSMConfigureClient interface { type defaultEC2SSMConfigureClient struct { *iam.Client ssmClient *ssm.Client - callerIdentityGetter + CallerIdentityGetter } // CreateDocument creates a Amazon Web Services Systems Manager (SSM document). @@ -142,7 +142,7 @@ func NewEC2SSMConfigureClient(ctx context.Context, region string) (EC2SSMConfigu return &defaultEC2SSMConfigureClient{ Client: iam.NewFromConfig(cfg), ssmClient: ssm.NewFromConfig(cfg), - callerIdentityGetter: sts.NewFromConfig(cfg), + CallerIdentityGetter: sts.NewFromConfig(cfg), }, nil } @@ -167,7 +167,7 @@ func ConfigureEC2SSM(ctx context.Context, clt EC2SSMConfigureClient, req EC2SSMI return trace.Wrap(err) } - if err := checkAccountID(ctx, clt, req.AccountID); err != nil { + if err := CheckAccountID(ctx, clt, req.AccountID); err != nil { return trace.Wrap(err) } diff --git a/lib/integrations/awsoidc/ec2_ssm_iam_config_test.go b/lib/integrations/awsoidc/ec2_ssm_iam_config_test.go index aa6c40b2fd7..6f206bbb08c 100644 --- a/lib/integrations/awsoidc/ec2_ssm_iam_config_test.go +++ b/lib/integrations/awsoidc/ec2_ssm_iam_config_test.go @@ -209,7 +209,7 @@ func TestEC2SSMIAMConfig(t *testing.T) { } { t.Run(tt.name, func(t *testing.T) { clt := mockEC2SSMIAMConfigClient{ - callerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, + CallerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, existingRoles: tt.mockExistingRoles, } @@ -228,7 +228,7 @@ func TestEC2SSMIAMConfig(t *testing.T) { } type mockEC2SSMIAMConfigClient struct { - callerIdentityGetter + CallerIdentityGetter existingRoles []string existingDocs map[string][]ssmtypes.Tag } diff --git a/lib/integrations/awsoidc/eice_iam_config.go b/lib/integrations/awsoidc/eice_iam_config.go index 1f17bedb38a..43d2e43ae58 100644 --- a/lib/integrations/awsoidc/eice_iam_config.go +++ b/lib/integrations/awsoidc/eice_iam_config.go @@ -71,13 +71,13 @@ func (r *EICEIAMConfigureRequest) CheckAndSetDefaults() error { // EICEIAMConfigureClient describes the required methods to create the IAM Policies required for accessing EC2 instances usine EICE. type EICEIAMConfigureClient interface { - callerIdentityGetter + CallerIdentityGetter // PutRolePolicy creates or replaces a Policy by its name in a IAM Role. PutRolePolicy(ctx context.Context, params *iam.PutRolePolicyInput, optFns ...func(*iam.Options)) (*iam.PutRolePolicyOutput, error) } type defaultEICEIAMConfigureClient struct { - callerIdentityGetter + CallerIdentityGetter *iam.Client } @@ -93,7 +93,7 @@ func NewEICEIAMConfigureClient(ctx context.Context, region string) (EICEIAMConfi } return &defaultEICEIAMConfigureClient{ - callerIdentityGetter: sts.NewFromConfig(cfg), + CallerIdentityGetter: sts.NewFromConfig(cfg), Client: iam.NewFromConfig(cfg), }, nil } @@ -130,7 +130,7 @@ func ConfigureEICEIAM(ctx context.Context, clt EICEIAMConfigureClient, req EICEI return trace.Wrap(err) } - if err := checkAccountID(ctx, clt, req.AccountID); err != nil { + if err := CheckAccountID(ctx, clt, req.AccountID); err != nil { return trace.Wrap(err) } diff --git a/lib/integrations/awsoidc/eice_iam_config_test.go b/lib/integrations/awsoidc/eice_iam_config_test.go index 0f03a57bdd6..94570882efd 100644 --- a/lib/integrations/awsoidc/eice_iam_config_test.go +++ b/lib/integrations/awsoidc/eice_iam_config_test.go @@ -142,7 +142,7 @@ func TestEICEIAMConfig(t *testing.T) { } { t.Run(tt.name, func(t *testing.T) { clt := mockEICEIAMConfigClient{ - callerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, + CallerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, existingRoles: tt.mockExistingRoles, } @@ -153,7 +153,7 @@ func TestEICEIAMConfig(t *testing.T) { } type mockEICEIAMConfigClient struct { - callerIdentityGetter + CallerIdentityGetter existingRoles []string } diff --git a/lib/integrations/awsoidc/eks_iam_config.go b/lib/integrations/awsoidc/eks_iam_config.go index 7cac117916a..04a70fab0f1 100644 --- a/lib/integrations/awsoidc/eks_iam_config.go +++ b/lib/integrations/awsoidc/eks_iam_config.go @@ -71,13 +71,13 @@ func (r *EKSIAMConfigureRequest) CheckAndSetDefaults() error { // EKSIAMConfigureClient describes the required methods to create the IAM Policies required for enrolling EKS clusters into Teleport. type EKSIAMConfigureClient interface { - callerIdentityGetter + CallerIdentityGetter // PutRolePolicy creates or replaces a Policy by its name in a IAM Role. PutRolePolicy(ctx context.Context, params *iam.PutRolePolicyInput, optFns ...func(*iam.Options)) (*iam.PutRolePolicyOutput, error) } type defaultEKSEIAMConfigureClient struct { - callerIdentityGetter + CallerIdentityGetter *iam.Client } @@ -94,7 +94,7 @@ func NewEKSIAMConfigureClient(ctx context.Context, region string) (EKSIAMConfigu return &defaultEKSEIAMConfigureClient{ Client: iam.NewFromConfig(cfg), - callerIdentityGetter: sts.NewFromConfig(cfg), + CallerIdentityGetter: sts.NewFromConfig(cfg), }, nil } @@ -118,7 +118,7 @@ func ConfigureEKSIAM(ctx context.Context, clt EKSIAMConfigureClient, req EKSIAMC return trace.Wrap(err) } - if err := checkAccountID(ctx, clt, req.AccountID); err != nil { + if err := CheckAccountID(ctx, clt, req.AccountID); err != nil { return trace.Wrap(err) } diff --git a/lib/integrations/awsoidc/eks_iam_config_test.go b/lib/integrations/awsoidc/eks_iam_config_test.go index 9a0a05e7dfc..444f4c53397 100644 --- a/lib/integrations/awsoidc/eks_iam_config_test.go +++ b/lib/integrations/awsoidc/eks_iam_config_test.go @@ -140,7 +140,7 @@ func TestEKSAMConfig(t *testing.T) { } { t.Run(tt.name, func(t *testing.T) { clt := mockEKSIAMConfigClient{ - callerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, + CallerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, existingRoles: tt.mockExistingRoles, } @@ -151,7 +151,7 @@ func TestEKSAMConfig(t *testing.T) { } type mockEKSIAMConfigClient struct { - callerIdentityGetter + CallerIdentityGetter existingRoles []string } diff --git a/lib/integrations/awsoidc/externalauditstorage_iam_config.go b/lib/integrations/awsoidc/externalauditstorage_iam_config.go index 79168e01732..aad6c5f37ef 100644 --- a/lib/integrations/awsoidc/externalauditstorage_iam_config.go +++ b/lib/integrations/awsoidc/externalauditstorage_iam_config.go @@ -74,6 +74,7 @@ func ConfigureExternalAuditStorage( policyCfg := &awslib.ExternalAuditStoragePolicyConfig{ Partition: params.Partition, Region: params.Region, + Account: params.AccountID, AthenaWorkgroupName: params.AthenaWorkgroup, GlueDatabaseName: params.GlueDatabase, GlueTableName: params.GlueTable, @@ -97,11 +98,13 @@ func ConfigureExternalAuditStorage( policyCfg.S3ARNs = append(policyCfg.S3ARNs, bucketARN, wildcardARN) policyCfg.S3ARNs = utils.Deduplicate(policyCfg.S3ARNs) - stsResp, err := clt.GetCallerIdentity(ctx, nil) - if err != nil { - return trace.Wrap(err, "attempting to find caller's AWS account ID: call to sts:GetCallerIdentity failed") + if policyCfg.Account == "" { + stsResp, err := clt.GetCallerIdentity(ctx, nil) + if err != nil { + return trace.Wrap(err, "attempting to find caller's AWS account ID: call to sts:GetCallerIdentity failed") + } + policyCfg.Account = aws.ToString(stsResp.Account) } - policyCfg.Account = aws.ToString(stsResp.Account) policyDoc, err := awslib.PolicyDocumentForExternalAuditStorage(policyCfg) if err != nil { diff --git a/lib/integrations/awsoidc/idp_iam_config.go b/lib/integrations/awsoidc/idp_iam_config.go index 4a6b345092b..24d14ffdcb2 100644 --- a/lib/integrations/awsoidc/idp_iam_config.go +++ b/lib/integrations/awsoidc/idp_iam_config.go @@ -108,7 +108,7 @@ func (r *IdPIAMConfigureRequest) CheckAndSetDefaults() error { // IdPIAMConfigureClient describes the required methods to create the AWS OIDC IdP and a Role that trusts that identity provider. // There is no guarantee that the client is thread safe. type IdPIAMConfigureClient interface { - callerIdentityGetter + CallerIdentityGetter // CreateOpenIDConnectProvider creates an IAM OIDC IdP. CreateOpenIDConnectProvider(ctx context.Context, params *iam.CreateOpenIDConnectProviderInput, optFns ...func(*iam.Options)) (*iam.CreateOpenIDConnectProviderOutput, error) @@ -131,7 +131,7 @@ type defaultIdPIAMConfigureClient struct { *iam.Client awsConfig aws.Config - callerIdentityGetter + CallerIdentityGetter } // NewIdPIAMConfigureClient creates a new IdPIAMConfigureClient. @@ -155,7 +155,7 @@ func NewIdPIAMConfigureClient(ctx context.Context) (IdPIAMConfigureClient, error httpClient: httpClient, awsConfig: cfg, Client: iam.NewFromConfig(cfg), - callerIdentityGetter: sts.NewFromConfig(cfg), + CallerIdentityGetter: sts.NewFromConfig(cfg), }, nil } diff --git a/lib/integrations/awsoidc/idp_iam_config_test.go b/lib/integrations/awsoidc/idp_iam_config_test.go index 70012936ac7..03bfcca1373 100644 --- a/lib/integrations/awsoidc/idp_iam_config_test.go +++ b/lib/integrations/awsoidc/idp_iam_config_test.go @@ -276,7 +276,7 @@ func TestConfigureIdPIAM(t *testing.T) { } { t.Run(tt.name, func(t *testing.T) { clt := mockIdPIAMConfigClient{ - callerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, + CallerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, existingRoles: tt.mockExistingRoles, existingIDPUrl: tt.mockExistingIdPUrl, } @@ -297,7 +297,7 @@ type mockRole struct { } type mockIdPIAMConfigClient struct { - callerIdentityGetter + CallerIdentityGetter existingIDPUrl []string existingRoles map[string]mockRole } diff --git a/lib/integrations/awsoidc/listdatabases_iam_config.go b/lib/integrations/awsoidc/listdatabases_iam_config.go index b4aa8201efb..9d3a13a35f8 100644 --- a/lib/integrations/awsoidc/listdatabases_iam_config.go +++ b/lib/integrations/awsoidc/listdatabases_iam_config.go @@ -63,14 +63,14 @@ func (r *ConfigureIAMListDatabasesRequest) CheckAndSetDefaults() error { // ListDatabasesIAMConfigureClient describes the required methods to create the IAM Policies required for Listing Databases. type ListDatabasesIAMConfigureClient interface { - callerIdentityGetter + CallerIdentityGetter // PutRolePolicy creates or replaces a Policy by its name in a IAM Role. PutRolePolicy(ctx context.Context, params *iam.PutRolePolicyInput, optFns ...func(*iam.Options)) (*iam.PutRolePolicyOutput, error) } type defaultListDatabasesIAMConfigureClient struct { *iam.Client - callerIdentityGetter + CallerIdentityGetter } // NewListDatabasesIAMConfigureClient creates a new ListDatabasesIAMConfigureClient. @@ -86,7 +86,7 @@ func NewListDatabasesIAMConfigureClient(ctx context.Context, region string) (Lis return &defaultListDatabasesIAMConfigureClient{ Client: iam.NewFromConfig(cfg), - callerIdentityGetter: sts.NewFromConfig(cfg), + CallerIdentityGetter: sts.NewFromConfig(cfg), }, nil } @@ -102,7 +102,7 @@ func ConfigureListDatabasesIAM(ctx context.Context, clt ListDatabasesIAMConfigur return trace.Wrap(err) } - if err := checkAccountID(ctx, clt, req.AccountID); err != nil { + if err := CheckAccountID(ctx, clt, req.AccountID); err != nil { return trace.Wrap(err) } diff --git a/lib/integrations/awsoidc/listdatabases_iam_config_test.go b/lib/integrations/awsoidc/listdatabases_iam_config_test.go index 36c0163c30f..83d7456ebad 100644 --- a/lib/integrations/awsoidc/listdatabases_iam_config_test.go +++ b/lib/integrations/awsoidc/listdatabases_iam_config_test.go @@ -116,7 +116,7 @@ func TestListDatabasesIAMConfig(t *testing.T) { } { t.Run(tt.name, func(t *testing.T) { clt := mockListDatabasesIAMConfigClient{ - callerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, + CallerIdentityGetter: mockSTSClient{accountID: tt.mockAccountID}, existingRoles: tt.mockExistingRoles, } @@ -127,7 +127,7 @@ func TestListDatabasesIAMConfig(t *testing.T) { } type mockListDatabasesIAMConfigClient struct { - callerIdentityGetter + CallerIdentityGetter existingRoles []string } diff --git a/lib/integrations/externalauditstorage/easconfig/externalauditstroageconfig.go b/lib/integrations/externalauditstorage/easconfig/externalauditstroageconfig.go index aa5113acfec..c6bc31ec70b 100644 --- a/lib/integrations/externalauditstorage/easconfig/externalauditstroageconfig.go +++ b/lib/integrations/externalauditstorage/easconfig/externalauditstroageconfig.go @@ -49,4 +49,6 @@ type ExternalAuditStorageConfiguration struct { GlueTable string // Partition is the AWS partition to use (default: aws). Partition string + // AccountID is the AWS account ID. + AccountID string } diff --git a/tool/teleport/common/integration_configure.go b/tool/teleport/common/integration_configure.go index 20b839155a0..dceb1bd2ca5 100644 --- a/tool/teleport/common/integration_configure.go +++ b/tool/teleport/common/integration_configure.go @@ -178,6 +178,15 @@ func onIntegrationConfExternalAuditCmd(ctx context.Context, params easconfig.Ext if err != nil { return trace.Wrap(err) } + + if params.AccountID != "" { + stsClient := sts.NewFromConfig(cfg) + err = awsoidc.CheckAccountID(ctx, stsClient, params.AccountID) + if err != nil { + return trace.Wrap(err) + } + } + if params.Bootstrap { err = externalauditstorage.BootstrapInfra(ctx, externalauditstorage.BootstrapInfraParams{ Athena: athena.NewFromConfig(cfg), diff --git a/tool/teleport/common/teleport.go b/tool/teleport/common/teleport.go index 1c891d3c5aa..7047f1eb627 100644 --- a/tool/teleport/common/teleport.go +++ b/tool/teleport/common/teleport.go @@ -536,6 +536,7 @@ func Run(options Options) (app *kingpin.Application, executedCommand string, con integrationConfExternalAuditCmd.Flag("glue-database", "The name of the Glue database used.").Required().StringVar(&ccf.IntegrationConfExternalAuditStorageArguments.GlueDatabase) integrationConfExternalAuditCmd.Flag("glue-table", "The name of the Glue table used.").Required().StringVar(&ccf.IntegrationConfExternalAuditStorageArguments.GlueTable) integrationConfExternalAuditCmd.Flag("aws-partition", "AWS partition (default: aws).").Default("aws").StringVar(&ccf.IntegrationConfExternalAuditStorageArguments.Partition) + integrationConfExternalAuditCmd.Flag("aws-account-id", "The AWS account ID.").StringVar(&ccf.IntegrationConfExternalAuditStorageArguments.AccountID) integrationConfAzureOIDCCmd := integrationConfigureCmd.Command("azure-oidc", "Configures Azure / Entra ID OIDC integration.") integrationConfAzureOIDCCmd.Flag("proxy-public-addr", "The public address of Teleport Proxy Service").Required().StringVar(&ccf.IntegrationConfAzureOIDCArguments.ProxyPublicAddr)