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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 8 additions & 16 deletions lib/srv/server/ec2_watcher.go
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,6 @@ import (
"sync"

"github.com/aws/aws-sdk-go-v2/aws"
"github.com/aws/aws-sdk-go-v2/aws/arn"
"github.com/aws/aws-sdk-go-v2/service/ec2"
ec2types "github.com/aws/aws-sdk-go-v2/service/ec2/types"
"github.com/aws/aws-sdk-go-v2/service/sts"
Expand All @@ -38,6 +37,7 @@ import (
awsregions "github.com/gravitational/teleport/lib/cloud/aws/regions"
"github.com/gravitational/teleport/lib/cloud/awsconfig"
"github.com/gravitational/teleport/lib/labels"
awsutils "github.com/gravitational/teleport/lib/utils/aws"
"github.com/gravitational/teleport/lib/utils/aws/organizations"
)

Expand Down Expand Up @@ -601,7 +601,7 @@ func (f *ec2InstanceFetcher) matcherRegions(ctx context.Context, params matcherR
return regions, nil
}

func (f *ec2InstanceFetcher) fetchAccountIDsUnderOrganization(ctx context.Context) ([]string, error) {
func (f *ec2InstanceFetcher) fetchAccountsUnderOrganization(ctx context.Context) (*organizations.Accounts, error) {
awsOpts := []awsconfig.OptionsFn{
awsconfig.WithCredentialsMaybeIntegration(awsconfig.IntegrationMetadata{Name: f.Matcher.Integration}),
}
Expand All @@ -625,7 +625,7 @@ func (f *ec2InstanceFetcher) fetchAccountIDsUnderOrganization(ctx context.Contex
})
}

accountIDs, err := organizations.MatchingAccounts(ctx, f.Logger, orgsClient, organizations.MatchingAccountsFilter{
accounts, err := organizations.MatchingAccounts(ctx, f.Logger, orgsClient, organizations.MatchingAccountsFilter{
IncludeOUs: includeOUs,
ExcludeOUs: excludeOUs,
OrganizationID: organizationID,
Expand All @@ -638,7 +638,7 @@ func (f *ec2InstanceFetcher) fetchAccountIDsUnderOrganization(ctx context.Contex
})
}

return accountIDs, nil
return accounts, nil
}

type assumeRoleWithExternalID struct {
Expand Down Expand Up @@ -668,23 +668,15 @@ func (f *ec2InstanceFetcher) allAssumeRoles(ctx context.Context) ([]assumeRoleWi
return nil, trace.BadParameter("assume role name is required when using AWS organization discovery")
}

accountIDs, err := f.fetchAccountIDsUnderOrganization(ctx)
accounts, err := f.fetchAccountsUnderOrganization(ctx)
if err != nil {
return nil, trace.Wrap(err)
}

var allAssumeRoles []assumeRoleWithExternalID
for _, accountID := range accountIDs {
assumeRoleARN := arn.ARN{
Partition: "aws",
Service: "iam",
Region: "",
AccountID: accountID,
Resource: "role/" + f.Matcher.AssumeRole.RoleName,
}

allAssumeRoles := make([]assumeRoleWithExternalID, 0, len(accounts.IDs))
for _, accountID := range accounts.IDs {
allAssumeRoles = append(allAssumeRoles, assumeRoleWithExternalID{
RoleARN: assumeRoleARN.String(),
RoleARN: awsutils.RoleARN(accounts.Partition, accountID, f.Matcher.AssumeRole.RoleName),
ExternalID: f.Matcher.AssumeRole.ExternalID,
})
}
Expand Down
75 changes: 62 additions & 13 deletions lib/srv/server/ec2_watcher_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
package server

import (
"cmp"
"context"
"errors"
"fmt"
Expand All @@ -34,7 +35,7 @@ import (
organizationtypes "github.com/aws/aws-sdk-go-v2/service/organizations/types"
"github.com/aws/aws-sdk-go-v2/service/sts"
ststypes "github.com/aws/aws-sdk-go-v2/service/sts/types"
"github.com/google/go-cmp/cmp"
gocmp "github.com/google/go-cmp/cmp"
"github.com/gravitational/trace"
"github.com/stretchr/testify/require"

Expand Down Expand Up @@ -101,14 +102,15 @@ func (m *mockAWSSTSClient) GetCallerIdentity(ctx context.Context, params *sts.Ge
type mockOrganizationsClient struct {
organizationID string
rootOUID string
rootARN string
ouItems map[string]ouItem
responseError error
}

type ouItem struct {
innerOUs []string
innerAccounts []string
innerNotActiveAccounts []string
innerOUs []string
innerAccounts []string
innerInactiveAccounts []string
}

func TestEC2WatcherResolveCallerIdentity(t *testing.T) {
Expand Down Expand Up @@ -208,7 +210,7 @@ func (m *mockOrganizationsClient) ListRoots(ctx context.Context, input *organiza
if m.responseError != nil {
return nil, m.responseError
}
rootARN := fmt.Sprintf("arn:aws:organizations::0000000000:root/%s/%s", m.organizationID, m.rootOUID)
rootARN := cmp.Or(m.rootARN, fmt.Sprintf("arn:aws:organizations::0000000000:root/%s/%s", m.organizationID, m.rootOUID))
return &organizations.ListRootsOutput{
Roots: []organizationtypes.Root{
{
Expand All @@ -230,19 +232,15 @@ func (m *mockOrganizationsClient) ListAccountsForParent(ctx context.Context, inp

var accounts []organizationtypes.Account
for _, accountID := range ouItem.innerAccounts {
accountARN := fmt.Sprintf("arn:aws:organizations::0000000000:account/%s/%s", m.organizationID, accountID)
accounts = append(accounts, organizationtypes.Account{
Id: aws.String(accountID),
State: organizationtypes.AccountStateActive,
Arn: aws.String(accountARN),
})
}
for _, accountID := range ouItem.innerNotActiveAccounts {
accountARN := fmt.Sprintf("arn:aws:organizations::0000000000:account/%s/%s", m.organizationID, accountID)
for _, accountID := range ouItem.innerInactiveAccounts {
accounts = append(accounts, organizationtypes.Account{
Id: aws.String(accountID),
State: organizationtypes.AccountStateSuspended,
Arn: aws.String(accountARN),
})
}
return &organizations.ListAccountsForParentOutput{
Expand Down Expand Up @@ -842,9 +840,9 @@ func TestEC2WatcherCallerIdentityFailureDoesNotBlockOrganizationDiscovery(t *tes
Logger: logtest.NewLogger(),
})

accountIDs, err := fetcher.fetchAccountIDsUnderOrganization(t.Context())
accounts, err := fetcher.fetchAccountsUnderOrganization(t.Context())
require.Error(t, err)
require.Empty(t, accountIDs)
require.Empty(t, accounts)

permissionErrors := EC2IAMPermissionErrors(err)
require.Len(t, permissionErrors, 1)
Expand All @@ -853,6 +851,57 @@ func TestEC2WatcherCallerIdentityFailureDoesNotBlockOrganizationDiscovery(t *tes
require.True(t, trace.IsAccessDenied(permissionErrors[0].Err))
}

func TestEC2WatcherAssumeRolesUseOrganizationPartition(t *testing.T) {
t.Parallel()

for _, tt := range []struct {
name string
rootARN string
expectedRoleARN string
}{
{
name: "aws-us-gov",
rootARN: "arn:aws-us-gov:organizations::000000000000:root/o-1/r-1",
expectedRoleARN: "arn:aws-us-gov:iam::000000000001:role/MyRole",
},
{
name: "aws-cn",
rootARN: "arn:aws-cn:organizations::000000000000:root/o-1/r-1",
expectedRoleARN: "arn:aws-cn:iam::000000000001:role/MyRole",
},
} {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()

fetcher := newEC2InstanceFetcher(ec2FetcherConfig{
Matcher: types.AWSMatcher{
AssumeRole: &types.AssumeRole{RoleName: "MyRole", ExternalID: "ext-1"},
Organization: &types.AWSOrganizationMatcher{
OrganizationID: "o-1",
OrganizationalUnits: &types.AWSOrganizationUnitsMatcher{Include: []string{types.Wildcard}},
},
},
AWSOrganizationsGetter: func(ctx context.Context, opts ...awsconfig.OptionsFn) (liborganizations.OrganizationsClient, error) {
return &mockOrganizationsClient{
organizationID: "o-1",
rootOUID: "r-1",
rootARN: tt.rootARN,
ouItems: map[string]ouItem{"r-1": {innerAccounts: []string{"000000000001"}}},
}, nil
},
Logger: logtest.NewLogger(),
})

assumeRoles, err := fetcher.allAssumeRoles(t.Context())
require.NoError(t, err)
require.Equal(t, []assumeRoleWithExternalID{{
RoleARN: tt.expectedRoleARN,
ExternalID: "ext-1",
}}, assumeRoles)
})
}
}

func TestEC2WatcherOrganizationClientCreationPermissionErrorReturnsPermissionError(t *testing.T) {
t.Parallel()

Expand Down Expand Up @@ -1600,7 +1649,7 @@ func TestConvertEC2InstancesToServerInfos(t *testing.T) {
require.NoError(t, err)
require.Len(t, serverInfos, 1)

require.Empty(t, cmp.Diff(expected, serverInfos[0]))
require.Empty(t, gocmp.Diff(expected, serverInfos[0]))
}

func TestMakeEvents(t *testing.T) {
Expand Down
29 changes: 12 additions & 17 deletions lib/utils/aws/organizations/arn.go
Original file line number Diff line number Diff line change
Expand Up @@ -28,28 +28,23 @@ import (
// OrganizationIDFromAccountARN extracts the organization ID from an account ARN.
// Example ARN: arn:aws:organizations::<org-master-account-id>:account/<org-id>/<account-id>
func OrganizationIDFromAccountARN(accountARN string) (string, error) {
return organizationIDFromARN(accountARN, "account")
}

// organizationIDFromRootOUARN extracts the organization ID from an root Organizational Unit ARN.
// Example ARN: arn:aws:organizations::<org-master-account-id>:root/<org-id>/<root-ou-id>
func organizationIDFromRootOUARN(rootOUARN string) (string, error) {
return organizationIDFromARN(rootOUARN, "root")
}

func organizationIDFromARN(orgARN string, resourceType string) (string, error) {
arnParsed, err := arn.Parse(orgARN)
parsed, err := arn.Parse(accountARN)
if err != nil {
return "", trace.Wrap(err)
}
resourceSplitted := strings.Split(arnParsed.Resource, "/")
if len(resourceSplitted) != 3 {
return organizationIDFromARN(parsed, "account")
}

// organizationIDFromARN extracts the organization ID from an AWS Organizations ARN
// whose resource is resourceType/<org-id>/<resource-id>.
func organizationIDFromARN(orgARN arn.ARN, resourceType string) (string, error) {
parts := strings.Split(orgARN.Resource, "/")
if len(parts) != 3 {
return "", trace.BadParameter("unexpected resource received in ARN from organizations API call: %s", orgARN)
}
if resourceSplitted[0] != resourceType {
return "", trace.BadParameter("expected resource type %s but received unexpected resource type %s in ARN from organizations API call: %s", resourceType, resourceSplitted[0], orgARN)
if parts[0] != resourceType {
return "", trace.BadParameter("expected resource type %s but received unexpected resource type %s in ARN from organizations API call: %s", resourceType, parts[0], orgARN)
}
organizationID := resourceSplitted[1]

return organizationID, nil
return parts[1], nil
}
46 changes: 30 additions & 16 deletions lib/utils/aws/organizations/arn_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ package organizations
import (
"testing"

"github.com/aws/aws-sdk-go-v2/aws/arn"
"github.com/stretchr/testify/require"
)

Expand Down Expand Up @@ -56,32 +57,45 @@ func TestOrganizationIDFromAccountARN(t *testing.T) {
}
}

func TestOrganizationIDFromRootOUARN(t *testing.T) {
func errContains(substr string) require.ErrorAssertionFunc {
return func(t require.TestingT, err error, msgAndArgs ...any) {
require.ErrorContains(t, err, substr, msgAndArgs...)
}
}

func TestOrganizationIDFromARN(t *testing.T) {
for _, tt := range []struct {
name string
accountARN string
expectedOrg string
errCheck require.ErrorAssertionFunc
name string
orgARN string
resourceType string
expectedOrg string
errCheck require.ErrorAssertionFunc
}{
{
name: "valid account ARN",
accountARN: "arn:aws:organizations::123456789012:root/o-exampleorgid/111111111111",
expectedOrg: "o-exampleorgid",
errCheck: require.NoError,
name: "root ARN",
orgARN: "arn:aws:organizations::123456789012:root/o-exampleorgid/r-exampleroot",
resourceType: "root",
expectedOrg: "o-exampleorgid",
errCheck: require.NoError,
},
{
name: "invalid ARN format",
accountARN: "invalid-arn-format",
errCheck: require.Error,
name: "too few resource parts",
orgARN: "arn:aws:organizations::123456789012:root/o-exampleorgid",
resourceType: "root",
errCheck: errContains("unexpected resource received"),
},
{
name: "wrong resource type",
accountARN: "arn:aws:organizations::123456789012:account/o-exampleorgid/111111111111",
errCheck: require.Error,
name: "wrong resource type",
orgARN: "arn:aws:organizations::123456789012:account/o-exampleorgid/111111111111",
resourceType: "root",
errCheck: errContains("expected resource type root"),
},
} {
t.Run(tt.name, func(t *testing.T) {
gotOrg, err := organizationIDFromRootOUARN(tt.accountARN)
parsed, err := arn.Parse(tt.orgARN)
require.NoError(t, err)

gotOrg, err := organizationIDFromARN(parsed, tt.resourceType)
tt.errCheck(t, err)
require.Equal(t, tt.expectedOrg, gotOrg)
})
Expand Down
Loading
Loading