Migrate AWS session to SDK v2 (#51626)

Only DynamoDB and AWS MongoDB Atlas depended on GetAWSSession, and the
migration for these packages was trivial.
Since this was the last AWS method in lib/cloud/clients, the vast majority
of the changes are to remove dead code.
This commit is contained in:
Gavin Frazar
2025-01-30 16:51:33 +00:00
committed by GitHub
parent 476c9e40f1
commit 5fb00a4aee
45 changed files with 99 additions and 771 deletions
-7
View File
@@ -92,7 +92,6 @@ import (
"github.com/gravitational/teleport/lib/bitbucket"
"github.com/gravitational/teleport/lib/cache"
"github.com/gravitational/teleport/lib/circleci"
"github.com/gravitational/teleport/lib/cloud"
"github.com/gravitational/teleport/lib/cryptosuites"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/devicetrust/assertserver"
@@ -373,12 +372,6 @@ func NewServer(cfg *InitConfig, opts ...ServerOption) (*Server, error) {
return nil, trace.Wrap(err)
}
}
if cfg.CloudClients == nil {
cfg.CloudClients, err = cloud.NewClients()
if err != nil {
return nil, trace.Wrap(err)
}
}
if cfg.Notifications == nil {
cfg.Notifications, err = local.NewNotificationsService(cfg.Backend, cfg.Clock)
if err != nil {
-4
View File
@@ -57,7 +57,6 @@ import (
"github.com/gravitational/teleport/lib/auth/migration"
"github.com/gravitational/teleport/lib/auth/state"
"github.com/gravitational/teleport/lib/backend"
"github.com/gravitational/teleport/lib/cloud"
"github.com/gravitational/teleport/lib/cryptosuites"
"github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/modules"
@@ -302,9 +301,6 @@ type InitConfig struct {
// AccessMonitoringRules is a service that manages access monitoring rules.
AccessMonitoringRules services.AccessMonitoringRules
// CloudClients provides clients for various cloud providers.
CloudClients cloud.Clients
// KubeWaitingContainers is a service that manages
// Kubernetes ephemeral containers that are waiting
// to be created until moderated session conditions are met.
+4 -384
View File
@@ -21,9 +21,7 @@ package cloud
import (
"context"
"io"
"log/slog"
"sync"
"time"
gcpcredentials "cloud.google.com/go/iam/credentials/apiv1"
"github.com/Azure/azure-sdk-for-go/sdk/azcore"
@@ -33,29 +31,17 @@ import (
"github.com/Azure/azure-sdk-for-go/sdk/resourcemanager/mysql/armmysql"
"github.com/Azure/azure-sdk-for-go/sdk/resourcemanager/postgresql/armpostgresql"
"github.com/Azure/azure-sdk-for-go/sdk/resourcemanager/subscription/armsubscription"
"github.com/aws/aws-sdk-go/aws"
"github.com/aws/aws-sdk-go/aws/credentials"
"github.com/aws/aws-sdk-go/aws/credentials/stscreds"
"github.com/aws/aws-sdk-go/aws/endpoints"
"github.com/aws/aws-sdk-go/aws/request"
awssession "github.com/aws/aws-sdk-go/aws/session"
"github.com/aws/aws-sdk-go/service/sts"
"github.com/aws/aws-sdk-go/service/sts/stsiface"
"github.com/gravitational/trace"
"google.golang.org/api/option"
"google.golang.org/grpc"
"google.golang.org/grpc/credentials/insecure"
"github.com/gravitational/teleport/api/types"
libcloudaws "github.com/gravitational/teleport/lib/cloud/aws"
"github.com/gravitational/teleport/lib/cloud/azure"
"github.com/gravitational/teleport/lib/cloud/gcp"
"github.com/gravitational/teleport/lib/cloud/imds"
awsimds "github.com/gravitational/teleport/lib/cloud/imds/aws"
azureimds "github.com/gravitational/teleport/lib/cloud/imds/azure"
gcpimds "github.com/gravitational/teleport/lib/cloud/imds/gcp"
"github.com/gravitational/teleport/lib/modules"
"github.com/gravitational/teleport/lib/utils"
)
// Clients provides interface for obtaining cloud provider clients.
@@ -65,8 +51,6 @@ type Clients interface {
GetInstanceMetadataClient(ctx context.Context) (imds.Client, error)
// GCPClients is an interface for providing GCP API clients.
GCPClients
// AWSClients is an interface for providing AWS API clients.
AWSClients
// AzureClients is an interface for Azure-specific API clients
AzureClients
// Closer closes all initialized clients.
@@ -87,12 +71,6 @@ type GCPClients interface {
GetGCPInstancesClient(context.Context) (gcp.InstancesClient, error)
}
// AWSClients is an interface for providing AWS API clients.
type AWSClients interface {
// GetAWSSession returns AWS session for the specified region and any role(s).
GetAWSSession(ctx context.Context, region string, opts ...AWSOptionsFn) (*awssession.Session, error)
}
// AzureClients is an interface for Azure-specific API clients
type AzureClients interface {
// GetAzureCredential returns Azure default token credential chain.
@@ -199,23 +177,16 @@ type ClientsOption func(cfg *cloudClients)
// NewClients returns a new instance of cloud clients retriever.
func NewClients(opts ...ClientsOption) (Clients, error) {
awsSessionsCache, err := utils.NewFnCache(utils.FnCacheConfig{
TTL: 15 * time.Minute,
})
if err != nil {
return nil, trace.Wrap(err)
}
azClients, err := newAzureClients()
if err != nil {
return nil, trace.Wrap(err)
}
cloudClients := &cloudClients{
awsSessionsCache: awsSessionsCache,
gcpClients: gcpClients{
gcpSQLAdmin: newClientCache[gcp.SQLAdminClient](gcp.NewSQLAdminClient),
gcpGKE: newClientCache[gcp.GKEClient](gcp.NewGKEClient),
gcpProjects: newClientCache[gcp.ProjectsClient](gcp.NewProjectsClient),
gcpInstances: newClientCache[gcp.InstancesClient](gcp.NewInstancesClient),
gcpSQLAdmin: newClientCache(gcp.NewSQLAdminClient),
gcpGKE: newClientCache(gcp.NewGKEClient),
gcpProjects: newClientCache(gcp.NewProjectsClient),
gcpInstances: newClientCache(gcp.NewInstancesClient),
},
azureClients: azClients,
}
@@ -230,31 +201,7 @@ func NewClients(opts ...ClientsOption) (Clients, error) {
// cloudClients implements Clients
var _ Clients = (*cloudClients)(nil)
// WithAWSIntegrationSessionProvider sets an integration session generator for AWS apis.
// If a client is requested for a specific Integration, instead of using the ambient credentials, this generator is used to fetch the AWS Session.
func WithAWSIntegrationSessionProvider(sessionProvider AWSIntegrationSessionProvider) func(*cloudClients) {
return func(cc *cloudClients) {
cc.awsIntegrationSessionProviderFn = sessionProvider
}
}
// AWSIntegrationSessionProvider defines a function that creates an [awssession.Session] from a Region and an Integration.
// This is used to generate aws sessions for clients that must use an Integration instead of ambient credentials.
type AWSIntegrationSessionProvider func(ctx context.Context, region string, integration string) (*awssession.Session, error)
type awsSessionCacheKey struct {
region string
integration string
roleARN string
externalID string
}
type cloudClients struct {
// awsSessionsCache is a cache of AWS sessions, where the cache key is
// an instance of awsSessionCacheKey.
awsSessionsCache *utils.FnCache
// awsIntegrationSessionProviderFn is a AWS Session Generator that uses an Integration to generate an AWS Session.
awsIntegrationSessionProviderFn AWSIntegrationSessionProvider
// instanceMetadata is the cached instance metadata client.
instanceMetadata imds.Client
// gcpClients contains GCP-specific clients.
@@ -316,156 +263,6 @@ type azureClients struct {
azureRoleAssignmentsClients azure.ClientMap[azure.RoleAssignmentsClient]
}
// credentialsSource defines where the credentials must come from.
type credentialsSource int
const (
// credentialsSourceAmbient uses the default Cloud SDK method to load the credentials.
credentialsSourceAmbient = iota + 1
// credentialsSourceIntegration uses an Integration to load the credentials.
credentialsSourceIntegration
)
// awsOptions a struct of additional options for assuming an AWS role
// when construction an underlying AWS session.
type awsOptions struct {
// baseSession is a session to use instead of the default session for an
// AWS region, which is used to enable role chaining.
baseSession *awssession.Session
// assumeRoleARN is the AWS IAM Role ARN to assume.
assumeRoleARN string
// assumeRoleExternalID is used to assume an external AWS IAM Role.
assumeRoleExternalID string
// credentialsSource describes which source to use to fetch credentials.
credentialsSource credentialsSource
// integration is the name of the integration to be used to fetch the credentials.
integration string
// customRetryer is a custom retryer to use for the session.
customRetryer request.Retryer
// maxRetries is the maximum number of retries to use for the session.
maxRetries *int
// withoutSessionCache disables the session cache for the AWS session.
withoutSessionCache bool
}
func (a *awsOptions) checkAndSetDefaults() error {
switch a.credentialsSource {
case credentialsSourceAmbient:
if a.integration != "" {
return trace.BadParameter("integration and ambient credentials cannot be used at the same time")
}
case credentialsSourceIntegration:
if a.integration == "" {
return trace.BadParameter("missing integration name")
}
default:
return trace.BadParameter("missing credentials source (ambient or integration)")
}
return nil
}
// AWSOptionsFn is an option function for setting additional options
// when getting an AWS session.
type AWSOptionsFn func(*awsOptions)
// WithAssumeRole configures options needed for assuming an AWS role.
func WithAssumeRole(roleARN, externalID string) AWSOptionsFn {
return func(options *awsOptions) {
options.assumeRoleARN = roleARN
options.assumeRoleExternalID = externalID
}
}
// WithoutSessionCache disables the session cache for the AWS session.
func WithoutSessionCache() AWSOptionsFn {
return func(options *awsOptions) {
options.withoutSessionCache = true
}
}
// WithAssumeRoleFromAWSMeta extracts options needed from AWS metadata for
// assuming an AWS role.
func WithAssumeRoleFromAWSMeta(meta types.AWS) AWSOptionsFn {
return WithAssumeRole(meta.AssumeRoleARN, meta.ExternalID)
}
// WithChainedAssumeRole sets a role to assume with a base session to use
// for assuming the role, which enables role chaining.
func WithChainedAssumeRole(session *awssession.Session, roleARN, externalID string) AWSOptionsFn {
return func(options *awsOptions) {
options.baseSession = session
options.assumeRoleARN = roleARN
options.assumeRoleExternalID = externalID
}
}
// WithRetryer sets a custom retryer for the session.
func WithRetryer(retryer request.Retryer) AWSOptionsFn {
return func(options *awsOptions) {
options.customRetryer = retryer
}
}
// WithMaxRetries sets the maximum allowed value for the sdk to keep retrying.
func WithMaxRetries(maxRetries int) AWSOptionsFn {
return func(options *awsOptions) {
options.maxRetries = &maxRetries
}
}
// WithCredentialsMaybeIntegration sets the credential source to be
// - ambient if the integration is an empty string
// - integration, otherwise
func WithCredentialsMaybeIntegration(integration string) AWSOptionsFn {
if integration != "" {
return withIntegrationCredentials(integration)
}
return WithAmbientCredentials()
}
// withIntegrationCredentials configures options with an Integration that must be used to fetch Credentials to assume a role.
// This prevents the usage of AWS environment credentials.
func withIntegrationCredentials(integration string) AWSOptionsFn {
return func(options *awsOptions) {
options.credentialsSource = credentialsSourceIntegration
options.integration = integration
}
}
// WithAmbientCredentials configures options to use the ambient credentials.
func WithAmbientCredentials() AWSOptionsFn {
return func(options *awsOptions) {
options.credentialsSource = credentialsSourceAmbient
}
}
// GetAWSSession returns AWS session for the specified region, optionally
// assuming AWS IAM Roles.
func (c *cloudClients) GetAWSSession(ctx context.Context, region string, opts ...AWSOptionsFn) (*awssession.Session, error) {
var options awsOptions
for _, opt := range opts {
opt(&options)
}
var err error
if options.baseSession == nil {
options.baseSession, err = c.getAWSSessionForRegion(ctx, region, options)
if err != nil {
return nil, trace.Wrap(err)
}
}
if options.assumeRoleARN == "" {
return options.baseSession, nil
}
return c.getAWSSessionForRole(ctx, region, options)
}
// GetGCPIAMClient returns GCP IAM client.
func (c *cloudClients) GetGCPIAMClient(ctx context.Context) (*gcpcredentials.IamCredentialsClient, error) {
c.mtx.RLock()
@@ -627,105 +424,6 @@ func (c *cloudClients) Close() (err error) {
return trace.Wrap(err)
}
// awsAmbientSessionProvider loads a new session using the environment variables.
// Describe in detail here: https://docs.aws.amazon.com/sdk-for-go/v1/developer-guide/configuring-sdk.html#specifying-credentials
func awsAmbientSessionProvider(ctx context.Context, region string) (*awssession.Session, error) {
awsSessionOptions := buildAWSSessionOptions(region, nil /* credentials */)
session, err := awssession.NewSessionWithOptions(awsSessionOptions)
return session, trace.Wrap(err)
}
// getAWSSessionForRegion returns AWS session for the specified region.
func (c *cloudClients) getAWSSessionForRegion(ctx context.Context, region string, opts awsOptions) (*awssession.Session, error) {
if err := opts.checkAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
createSession := func(ctx context.Context) (*awssession.Session, error) {
if opts.credentialsSource == credentialsSourceIntegration {
if c.awsIntegrationSessionProviderFn == nil {
return nil, trace.BadParameter("missing aws integration session provider")
}
slog.DebugContext(ctx, "Initializing AWS session",
"region", region,
"integration", opts.integration,
)
session, err := c.awsIntegrationSessionProviderFn(ctx, region, opts.integration)
return session, trace.Wrap(err)
}
slog.DebugContext(ctx, "Initializing AWS session using environment credentials",
"region", region,
)
session, err := awsAmbientSessionProvider(ctx, region)
return session, trace.Wrap(err)
}
if opts.withoutSessionCache {
sess, err := createSession(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
if opts.customRetryer != nil || opts.maxRetries != nil {
return sess.Copy(&aws.Config{
Retryer: opts.customRetryer,
MaxRetries: opts.maxRetries,
}), nil
}
return sess, trace.Wrap(err)
}
cacheKey := awsSessionCacheKey{
region: region,
integration: opts.integration,
}
sess, err := utils.FnCacheGet(ctx, c.awsSessionsCache, cacheKey, func(ctx context.Context) (*awssession.Session, error) {
session, err := createSession(ctx)
return session, trace.Wrap(err)
})
if err != nil {
return nil, trace.Wrap(err)
}
if opts.customRetryer != nil || opts.maxRetries != nil {
return sess.Copy(&aws.Config{
Retryer: opts.customRetryer,
MaxRetries: opts.maxRetries,
}), nil
}
return sess, err
}
// getAWSSessionForRole returns AWS session for the specified region and role.
func (c *cloudClients) getAWSSessionForRole(ctx context.Context, region string, options awsOptions) (*awssession.Session, error) {
if err := options.checkAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
createSession := func(ctx context.Context) (*awssession.Session, error) {
stsClient := sts.New(options.baseSession)
return newSessionWithRole(ctx, stsClient, region, options.assumeRoleARN, options.assumeRoleExternalID)
}
if options.withoutSessionCache {
session, err := createSession(ctx)
return session, trace.Wrap(err)
}
cacheKey := awsSessionCacheKey{
region: region,
integration: options.integration,
roleARN: options.assumeRoleARN,
externalID: options.assumeRoleExternalID,
}
return utils.FnCacheGet(ctx, c.awsSessionsCache, cacheKey, func(ctx context.Context) (*awssession.Session, error) {
session, err := createSession(ctx)
return session, trace.Wrap(err)
})
}
func (c *cloudClients) initGCPIAMClient(ctx context.Context) (*gcpcredentials.IamCredentialsClient, error) {
c.mtx.Lock()
defer c.mtx.Unlock()
@@ -891,7 +589,6 @@ var _ Clients = (*TestCloudClients)(nil)
// TestCloudClients are used in tests.
type TestCloudClients struct {
STS stsiface.STSAPI
GCPSQL gcp.SQLAdminClient
GCPGKE gcp.GKEClient
GCPProjects gcp.ProjectsClient
@@ -916,43 +613,6 @@ type TestCloudClients struct {
AzureRoleAssignments azure.RoleAssignmentsClient
}
// GetAWSSession returns AWS session for the specified region, optionally
// assuming AWS IAM Roles.
func (c *TestCloudClients) GetAWSSession(ctx context.Context, region string, opts ...AWSOptionsFn) (*awssession.Session, error) {
var options awsOptions
for _, opt := range opts {
opt(&options)
}
var err error
if options.baseSession == nil {
options.baseSession, err = c.getAWSSessionForRegion(region)
if err != nil {
return nil, trace.Wrap(err)
}
}
if options.assumeRoleARN == "" {
return options.baseSession, nil
}
return newSessionWithRole(ctx, c.STS, region, options.assumeRoleARN, options.assumeRoleExternalID)
}
// GetAWSSession returns AWS session for the specified region.
func (c *TestCloudClients) getAWSSessionForRegion(region string) (*awssession.Session, error) {
useFIPSEndpoint := endpoints.FIPSEndpointStateUnset
if modules.GetModules().IsBoringBinary() {
useFIPSEndpoint = endpoints.FIPSEndpointStateEnabled
}
return awssession.NewSession(&aws.Config{
Credentials: credentials.NewCredentials(&credentials.StaticProvider{Value: credentials.Value{
AccessKeyID: "fakeClientKeyID",
SecretAccessKey: "fakeClientSecret",
}}),
Region: aws.String(region),
UseFIPSEndpoint: useFIPSEndpoint,
})
}
// GetGCPIAMClient returns GCP IAM client.
func (c *TestCloudClients) GetGCPIAMClient(ctx context.Context) (*gcpcredentials.IamCredentialsClient, error) {
return gcpcredentials.NewIamCredentialsClient(ctx,
@@ -1075,43 +735,3 @@ func (c *TestCloudClients) GetAzureRoleAssignmentsClient(subscription string) (a
func (c *TestCloudClients) Close() error {
return nil
}
// newSessionWithRole assumes a given AWS IAM Role, passing an external ID if given,
// and returns a new AWS session with the assumed role in the given region.
func newSessionWithRole(ctx context.Context, svc stscreds.AssumeRoler, region, roleARN, externalID string) (*awssession.Session, error) {
slog.DebugContext(ctx, "Initializing AWS session for assumed role",
"assumed_role", roleARN,
"region", region,
)
// Make a credentials with AssumeRoleProvider and test it out.
cred := stscreds.NewCredentialsWithClient(svc, roleARN, func(p *stscreds.AssumeRoleProvider) {
if externalID != "" {
p.ExternalID = aws.String(externalID)
}
})
if _, err := cred.GetWithContext(ctx); err != nil {
return nil, trace.Wrap(libcloudaws.ConvertRequestFailureError(err))
}
awsSessionOptions := buildAWSSessionOptions(region, cred)
// Create a new session with the credentials.
roleSession, err := awssession.NewSessionWithOptions(awsSessionOptions)
return roleSession, trace.Wrap(err)
}
func buildAWSSessionOptions(region string, cred *credentials.Credentials) awssession.Options {
useFIPSEndpoint := endpoints.FIPSEndpointStateUnset
if modules.GetModules().IsBoringBinary() {
useFIPSEndpoint = endpoints.FIPSEndpointStateEnabled
}
return awssession.Options{
SharedConfigState: awssession.SharedConfigEnable,
Config: aws.Config{
Region: aws.String(region),
Credentials: cred,
UseFIPSEndpoint: useFIPSEndpoint,
},
}
}
-114
View File
@@ -1,114 +0,0 @@
/*
* Teleport
* Copyright (C) 2023 Gravitational, Inc.
*
* This program is free software: you can redistribute it and/or modify
* it under the terms of the GNU Affero General Public License as published by
* the Free Software Foundation, either version 3 of the License, or
* (at your option) any later version.
*
* This program is distributed in the hope that it will be useful,
* but WITHOUT ANY WARRANTY; without even the implied warranty of
* MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
* GNU Affero General Public License for more details.
*
* You should have received a copy of the GNU Affero General Public License
* along with this program. If not, see <http://www.gnu.org/licenses/>.
*/
package cloud
import (
"context"
"testing"
"github.com/aws/aws-sdk-go/aws"
awssession "github.com/aws/aws-sdk-go/aws/session"
"github.com/gravitational/trace"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestClientGetAWSSessionIntegration(t *testing.T) {
dummyIntegration := "integration-test"
dummyRegion := "test-region-123"
t.Run("without an integration session provider, must return a missing aws integration session provider error", func(t *testing.T) {
ctx := context.Background()
clients, err := NewClients()
require.NoError(t, err)
t.Cleanup(func() { require.NoError(t, clients.Close()) })
_, err = clients.GetAWSSession(ctx, "us-region-2", WithCredentialsMaybeIntegration("integration-test"))
require.True(t, trace.IsBadParameter(err), "expected err to be BadParameter, got %+v", err)
require.ErrorContains(t, err, "missing aws integration session provider")
})
t.Run("with an integration session provider, must return the session", func(t *testing.T) {
ctx := context.Background()
dummySession := &awssession.Session{
Config: &aws.Config{
Region: &dummyRegion,
},
}
clients, err := NewClients(WithAWSIntegrationSessionProvider(func(ctx context.Context, region, integration string) (*awssession.Session, error) {
assert.Equal(t, dummyIntegration, integration)
assert.Equal(t, dummyRegion, region)
return dummySession, nil
}))
require.NoError(t, err)
t.Cleanup(func() { require.NoError(t, clients.Close()) })
sess, err := clients.GetAWSSession(ctx, dummyRegion, WithCredentialsMaybeIntegration("integration-test"))
require.NoError(t, err)
require.Equal(t, dummySession, sess)
})
t.Run("with an integration session provider, but using an empty integration falls back to ambient credentials, must not call the integration session provider", func(t *testing.T) {
ctx := context.Background()
clients, err := NewClients(WithAWSIntegrationSessionProvider(func(ctx context.Context, region, integration string) (*awssession.Session, error) {
assert.Fail(t, "should not be called")
return nil, nil
}))
require.NoError(t, err)
t.Cleanup(func() { require.NoError(t, clients.Close()) })
sess, err := clients.GetAWSSession(ctx, dummyRegion, WithCredentialsMaybeIntegration(""))
require.NoError(t, err)
require.NotNil(t, sess)
})
t.Run("with an integration session provider, but using ambient credentials, must not call the integration session provider", func(t *testing.T) {
ctx := context.Background()
clients, err := NewClients(WithAWSIntegrationSessionProvider(func(ctx context.Context, region, integration string) (*awssession.Session, error) {
assert.Fail(t, "should not be called")
return nil, nil
}))
require.NoError(t, err)
t.Cleanup(func() { require.NoError(t, clients.Close()) })
sess, err := clients.GetAWSSession(ctx, dummyRegion, WithAmbientCredentials())
require.NoError(t, err)
require.NotNil(t, sess)
})
t.Run("with an integration session provider, but no credential source defined", func(t *testing.T) {
ctx := context.Background()
clients, err := NewClients(WithAWSIntegrationSessionProvider(func(ctx context.Context, region, integration string) (*awssession.Session, error) {
assert.Fail(t, "should not be called")
return nil, nil
}))
require.NoError(t, err)
t.Cleanup(func() { require.NoError(t, clients.Close()) })
_, err = clients.GetAWSSession(ctx, dummyRegion)
require.Error(t, err)
require.ErrorContains(t, err, "missing credentials source")
})
}
+4 -1
View File
@@ -63,7 +63,10 @@ func (f *FakeOIDCIntegrationClient) GetIntegration(ctx context.Context, name str
if f.Unauth {
return nil, trace.AccessDenied("unauthorized")
}
return f.Integration, nil
if f.Integration.GetName() == name {
return f.Integration, nil
}
return nil, trace.NotFound("integration %q not found", name)
}
func (f *FakeOIDCIntegrationClient) GenerateAWSOIDCToken(ctx context.Context, integrationName string) (string, error) {
-6
View File
@@ -2136,11 +2136,6 @@ func (process *TeleportProcess) initAuthService() error {
return trace.Wrap(err)
}
cloudClients, err := cloud.NewClients()
if err != nil {
return trace.Wrap(err)
}
logger := process.logger.With(teleport.ComponentKey, teleport.Component(teleport.ComponentAuth, process.id))
// first, create the AuthServer
@@ -2187,7 +2182,6 @@ func (process *TeleportProcess) initAuthService() error {
Clock: cfg.Clock,
HTTPClientForAWSSTS: cfg.Auth.HTTPClientForAWSSTS,
Tracer: process.TracingProvider.Tracer(teleport.ComponentAuth),
CloudClients: cloudClients,
Logger: logger,
}, func(as *auth.Server) error {
if !process.Config.CachePolicy.Enabled {
-2
View File
@@ -107,7 +107,6 @@ func TestMain(m *testing.M) {
registerTestSnowflakeEngine()
registerTestElasticsearchEngine()
registerTestSQLServerEngine()
registerTestDynamoDBEngine()
os.Exit(m.Run())
}
@@ -2483,7 +2482,6 @@ func (p *agentParams) setDefaults(c *testContext) {
if p.CloudClients == nil {
p.CloudClients = &clients.TestCloudClients{
STS: &mocks.STSClientV1{},
GCPSQL: p.GCPSQL,
}
}
-6
View File
@@ -30,7 +30,6 @@ import (
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/cloud"
awslib "github.com/gravitational/teleport/lib/cloud/aws"
"github.com/gravitational/teleport/lib/cloud/awsconfig"
dbiam "github.com/gravitational/teleport/lib/srv/db/common/iam"
@@ -40,8 +39,6 @@ import (
type awsConfig struct {
// awsConfigProvider provides [aws.Config] for AWS SDK service clients.
awsConfigProvider awsconfig.Provider
// clients is an interface for creating AWS clients.
clients cloud.Clients
// identity is AWS identity this database agent is running as.
identity awslib.Identity
// database is the database instance to configure.
@@ -55,9 +52,6 @@ type awsConfig struct {
// Check validates the config.
func (c *awsConfig) Check() error {
if c.clients == nil {
return trace.BadParameter("missing parameter clients")
}
if c.identity == nil {
return trace.BadParameter("missing parameter identity")
}
-11
View File
@@ -33,7 +33,6 @@ import (
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/retryutils"
"github.com/gravitational/teleport/lib/auth/authclient"
"github.com/gravitational/teleport/lib/cloud"
awslib "github.com/gravitational/teleport/lib/cloud/aws"
"github.com/gravitational/teleport/lib/cloud/awsconfig"
"github.com/gravitational/teleport/lib/services"
@@ -48,8 +47,6 @@ type IAMConfig struct {
AccessPoint authclient.DatabaseAccessPoint
// AWSConfigProvider provides [aws.Config] for AWS SDK service clients.
AWSConfigProvider awsconfig.Provider
// Clients is an interface for retrieving cloud clients.
Clients cloud.Clients
// HostID is the host identified where this agent is running.
// DELETE IN 11.0.
HostID string
@@ -70,13 +67,6 @@ func (c *IAMConfig) Check() error {
if c.AWSConfigProvider == nil {
return trace.BadParameter("missing AWSConfigProvider")
}
if c.Clients == nil {
cloudClients, err := cloud.NewClients()
if err != nil {
return trace.Wrap(err)
}
c.Clients = cloudClients
}
if c.HostID == "" {
return trace.BadParameter("missing HostID")
}
@@ -245,7 +235,6 @@ func (c *IAM) getAWSConfigurator(ctx context.Context, database types.Database) (
}
return newAWS(ctx, awsConfig{
awsConfigProvider: c.cfg.AWSConfigProvider,
clients: c.cfg.Clients,
database: database,
identity: identity,
policyName: policyName,
+13 -23
View File
@@ -33,7 +33,6 @@ import (
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/auth/authclient"
clients "github.com/gravitational/teleport/lib/cloud"
"github.com/gravitational/teleport/lib/cloud/mocks"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/services"
@@ -156,8 +155,7 @@ func TestAWSIAM(t *testing.T) {
AWSConfigProvider: &mocks.AWSConfigProvider{
STSClient: stsClient,
},
Clients: &clients.TestCloudClients{},
HostID: "host-id",
HostID: "host-id",
onProcessedTask: func(iamTask, error) {
taskChan <- struct{}{}
},
@@ -298,13 +296,11 @@ func TestAWSIAMNoPermissions(t *testing.T) {
tests := []struct {
name string
meta types.AWS
clients clients.Clients
awsClients awsClientProvider
}{
{
name: "RDS database",
meta: types.AWS{Region: "localhost", AccountID: "123456789012", RDS: types.RDS{InstanceID: "postgres-rds", ResourceID: "postgres-rds-resource-id"}},
clients: &clients.TestCloudClients{},
name: "RDS database",
meta: types.AWS{Region: "localhost", AccountID: "123456789012", RDS: types.RDS{InstanceID: "postgres-rds", ResourceID: "postgres-rds-resource-id"}},
awsClients: fakeAWSClients{
iamClient: &mocks.IAMMock{Unauth: true},
rdsClient: &mocks.RDSClient{Unauth: true},
@@ -312,9 +308,8 @@ func TestAWSIAMNoPermissions(t *testing.T) {
},
},
{
name: "Aurora cluster",
meta: types.AWS{Region: "localhost", AccountID: "123456789012", RDS: types.RDS{ClusterID: "postgres-aurora", ResourceID: "postgres-aurora-resource-id"}},
clients: &clients.TestCloudClients{},
name: "Aurora cluster",
meta: types.AWS{Region: "localhost", AccountID: "123456789012", RDS: types.RDS{ClusterID: "postgres-aurora", ResourceID: "postgres-aurora-resource-id"}},
awsClients: fakeAWSClients{
iamClient: &mocks.IAMMock{Unauth: true},
rdsClient: &mocks.RDSClient{Unauth: true},
@@ -322,9 +317,8 @@ func TestAWSIAMNoPermissions(t *testing.T) {
},
},
{
name: "RDS database missing metadata",
meta: types.AWS{Region: "localhost", RDS: types.RDS{ClusterID: "postgres-aurora"}},
clients: &clients.TestCloudClients{},
name: "RDS database missing metadata",
meta: types.AWS{Region: "localhost", RDS: types.RDS{ClusterID: "postgres-aurora"}},
awsClients: fakeAWSClients{
iamClient: &mocks.IAMMock{Unauth: true},
rdsClient: &mocks.RDSClient{Unauth: true},
@@ -332,27 +326,24 @@ func TestAWSIAMNoPermissions(t *testing.T) {
},
},
{
name: "Redshift cluster",
meta: types.AWS{Region: "localhost", AccountID: "123456789012", Redshift: types.Redshift{ClusterID: "redshift-cluster-1"}},
clients: &clients.TestCloudClients{},
name: "Redshift cluster",
meta: types.AWS{Region: "localhost", AccountID: "123456789012", Redshift: types.Redshift{ClusterID: "redshift-cluster-1"}},
awsClients: fakeAWSClients{
iamClient: &mocks.IAMMock{Unauth: true},
stsClient: stsClient,
},
},
{
name: "ElastiCache",
meta: types.AWS{Region: "localhost", AccountID: "123456789012", ElastiCache: types.ElastiCache{ReplicationGroupID: "some-group"}},
clients: &clients.TestCloudClients{},
name: "ElastiCache",
meta: types.AWS{Region: "localhost", AccountID: "123456789012", ElastiCache: types.ElastiCache{ReplicationGroupID: "some-group"}},
awsClients: fakeAWSClients{
iamClient: &mocks.IAMMock{Unauth: true},
stsClient: stsClient,
},
},
{
name: "IAM UnmodifiableEntityException",
meta: types.AWS{Region: "localhost", AccountID: "123456789012", Redshift: types.Redshift{ClusterID: "redshift-cluster-1"}},
clients: &clients.TestCloudClients{},
name: "IAM UnmodifiableEntityException",
meta: types.AWS{Region: "localhost", AccountID: "123456789012", Redshift: types.Redshift{ClusterID: "redshift-cluster-1"}},
awsClients: fakeAWSClients{
iamClient: &mocks.IAMMock{
Error: &iamtypes.UnmodifiableEntityException{
@@ -369,7 +360,6 @@ func TestAWSIAMNoPermissions(t *testing.T) {
// Make configurator.
configurator, err := NewIAM(ctx, IAMConfig{
AccessPoint: &mockAccessPoint{},
Clients: test.clients,
HostID: "host-id",
AWSConfigProvider: &mocks.AWSConfigProvider{
STSClient: stsClient,
-10
View File
@@ -41,7 +41,6 @@ import (
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/cloud"
"github.com/gravitational/teleport/lib/cloud/awsconfig"
"github.com/gravitational/teleport/lib/srv/db/common"
discoverycommon "github.com/gravitational/teleport/lib/srv/discovery/common"
@@ -147,8 +146,6 @@ func (defaultAWSClients) getSTSClient(cfg aws.Config, optFns ...func(*sts.Option
// MetadataConfig is the cloud metadata service config.
type MetadataConfig struct {
// Clients is an interface for retrieving cloud clients.
Clients cloud.Clients
// AWSConfigProvider provides [aws.Config] for AWS SDK service clients.
AWSConfigProvider awsconfig.Provider
@@ -158,13 +155,6 @@ type MetadataConfig struct {
// Check validates the metadata service config.
func (c *MetadataConfig) Check() error {
if c.Clients == nil {
cloudClients, err := cloud.NewClients()
if err != nil {
return trace.Wrap(err)
}
c.Clients = cloudClients
}
if c.AWSConfigProvider == nil {
return trace.BadParameter("missing AWSConfigProvider")
}
-7
View File
@@ -39,7 +39,6 @@ import (
"github.com/stretchr/testify/require"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/cloud"
"github.com/gravitational/teleport/lib/cloud/mocks"
"github.com/gravitational/teleport/lib/defaults"
)
@@ -137,9 +136,6 @@ func TestAWSMetadata(t *testing.T) {
// Create metadata fetcher.
metadata, err := NewMetadata(MetadataConfig{
Clients: &cloud.TestCloudClients{
STS: &fakeSTS.STSClientV1,
},
AWSConfigProvider: &mocks.AWSConfigProvider{
STSClient: fakeSTS,
},
@@ -420,9 +416,6 @@ func TestAWSMetadataNoPermissions(t *testing.T) {
// Create metadata fetcher.
metadata, err := NewMetadata(MetadataConfig{
Clients: &cloud.TestCloudClients{
STS: &fakeSTS.STSClientV1,
},
AWSConfigProvider: &mocks.AWSConfigProvider{
STSClient: fakeSTS,
},
+4 -4
View File
@@ -45,8 +45,8 @@ type DiscoveryResourceCheckerConfig struct {
AWSConfigProvider awsconfig.Provider
// ResourceMatchers is a list of database resource matchers.
ResourceMatchers []services.ResourceMatcher
// Clients is an interface for retrieving cloud clients.
Clients cloud.Clients
// AzureClients is an interface for retrieving Azure cloud clients.
AzureClients cloud.AzureClients
// Context is the database server close context.
Context context.Context
// Logger is used for logging.
@@ -55,12 +55,12 @@ type DiscoveryResourceCheckerConfig struct {
// CheckAndSetDefaults validates the config and sets default values.
func (c *DiscoveryResourceCheckerConfig) CheckAndSetDefaults() error {
if c.Clients == nil {
if c.AzureClients == nil {
cloudClients, err := cloud.NewClients()
if err != nil {
return trace.Wrap(err)
}
c.Clients = cloudClients
c.AzureClients = cloudClients
}
if c.AWSConfigProvider == nil {
return trace.BadParameter("missing AWSConfigProvider")
@@ -61,7 +61,7 @@ func newCredentialsChecker(cfg DiscoveryResourceCheckerConfig) (*credentialsChec
return &credentialsChecker{
awsConfigProvider: cfg.AWSConfigProvider,
awsClients: defaultAWSClients{},
azureClients: cfg.Clients,
azureClients: cfg.AzureClients,
resourceMatchers: cfg.ResourceMatchers,
logger: cfg.Logger,
cache: cache,
-3
View File
@@ -33,7 +33,6 @@ import (
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils"
apiawsutils "github.com/gravitational/teleport/api/utils/aws"
"github.com/gravitational/teleport/lib/cloud"
"github.com/gravitational/teleport/lib/cloud/awsconfig"
)
@@ -44,7 +43,6 @@ type urlChecker struct {
// awsClients is an SDK client provider.
awsClients awsClientProvider
clients cloud.Clients
logger *slog.Logger
warnOnError bool
@@ -60,7 +58,6 @@ func newURLChecker(cfg DiscoveryResourceCheckerConfig) *urlChecker {
return &urlChecker{
awsConfigProvider: cfg.AWSConfigProvider,
awsClients: defaultAWSClients{},
clients: cfg.Clients,
logger: cfg.Logger,
warnOnError: getWarnOnError(),
}
@@ -32,7 +32,6 @@ import (
"github.com/gravitational/teleport/api/types"
apiawsutils "github.com/gravitational/teleport/api/utils/aws"
"github.com/gravitational/teleport/lib/cloud"
"github.com/gravitational/teleport/lib/cloud/awsconfig"
"github.com/gravitational/teleport/lib/cloud/mocks"
"github.com/gravitational/teleport/lib/srv/discovery/common"
@@ -119,26 +118,16 @@ func TestURLChecker_AWS(t *testing.T) {
require.Len(t, docdbClusterDBs, 2) // Primary, reader.
testCases = append(testCases, docdbClusterDBs...)
// Mock cloud clients.
mockClients := &cloud.TestCloudClients{
STS: &mocks.STSClientV1{},
}
mockClientsUnauth := &cloud.TestCloudClients{
STS: &mocks.STSClientV1{},
}
// Test both check methods.
// Note that "No permissions" logs should only be printed during the second
// group ("basic endpoint check").
methods := []struct {
name string
clients cloud.Clients
awsConfigProvider awsconfig.Provider
awsClients awsClientProvider
}{
{
name: "API check",
clients: mockClients,
awsConfigProvider: &mocks.AWSConfigProvider{},
awsClients: fakeAWSClients{
ecClient: &mocks.ElastiCacheClient{
@@ -167,7 +156,6 @@ func TestURLChecker_AWS(t *testing.T) {
},
{
name: "basic endpoint check",
clients: mockClientsUnauth,
awsConfigProvider: &mocks.AWSConfigProvider{},
awsClients: fakeAWSClients{
ecClient: &mocks.ElastiCacheClient{Unauth: true},
@@ -183,7 +171,6 @@ func TestURLChecker_AWS(t *testing.T) {
for _, method := range methods {
t.Run(method.name, func(t *testing.T) {
c := newURLChecker(DiscoveryResourceCheckerConfig{
Clients: method.clients,
AWSConfigProvider: method.awsConfigProvider,
Logger: utils.NewSlogLoggerForTests(),
})
-10
View File
@@ -33,7 +33,6 @@ import (
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/retryutils"
"github.com/gravitational/teleport/lib/cloud"
"github.com/gravitational/teleport/lib/cloud/awsconfig"
"github.com/gravitational/teleport/lib/srv/db/secrets"
"github.com/gravitational/teleport/lib/utils/interval"
@@ -43,8 +42,6 @@ import (
type Config struct {
// AWSConfigProvider provides [aws.Config] for AWS SDK service clients.
AWSConfigProvider awsconfig.Provider
// Clients is an interface for retrieving cloud clients.
Clients cloud.Clients
// Clock is used to control time.
Clock clockwork.Clock
// Interval is the interval between user updates. Interval is also used as
@@ -93,13 +90,6 @@ func (c *Config) CheckAndSetDefaults() error {
if c.UpdateMeta == nil {
return trace.BadParameter("missing UpdateMeta")
}
if c.Clients == nil {
cloudClients, err := cloud.NewClients()
if err != nil {
return trace.Wrap(err)
}
c.Clients = cloudClients
}
if c.Clock == nil {
c.Clock = clockwork.NewRealClock()
}
-2
View File
@@ -36,7 +36,6 @@ import (
"github.com/stretchr/testify/require"
"github.com/gravitational/teleport/api/types"
clients "github.com/gravitational/teleport/lib/cloud"
libaws "github.com/gravitational/teleport/lib/cloud/aws"
"github.com/gravitational/teleport/lib/cloud/mocks"
"github.com/gravitational/teleport/lib/defaults"
@@ -85,7 +84,6 @@ func TestUsers(t *testing.T) {
users, err := NewUsers(Config{
AWSConfigProvider: &mocks.AWSConfigProvider{},
Clients: &clients.TestCloudClients{},
Clock: clock,
UpdateMeta: func(_ context.Context, database types.Database) error {
// Update db1 to group3 when setupAllDatabases.
+7 -15
View File
@@ -1180,26 +1180,18 @@ func (a *dbAuth) GetAWSIAMCreds(ctx context.Context, database types.Database, da
return "", "", "", trace.Wrap(err)
}
baseSession, err := a.cfg.Clients.GetAWSSession(ctx, dbAWS.Region,
cloud.WithAssumeRoleFromAWSMeta(dbAWS),
cloud.WithAmbientCredentials(),
awsCfg, err := a.cfg.AWSConfigProvider.GetConfig(ctx, dbAWS.Region,
awsconfig.WithAssumeRole(dbAWS.AssumeRoleARN, dbAWS.ExternalID),
// ExternalID should only be used once. If the baseSession assumes a role,
// the chained sessions should have an empty external ID.
awsconfig.WithAssumeRole(arn, externalIDForChainedAssumeRole(dbAWS)),
awsconfig.WithAmbientCredentials(),
)
if err != nil {
return "", "", "", trace.Wrap(err)
}
// ExternalID should only be used once. If the baseSession assumes a role,
// the chained sessions should have an empty external ID.
sess, err := a.cfg.Clients.GetAWSSession(ctx, dbAWS.Region,
cloud.WithChainedAssumeRole(baseSession, arn, externalIDForChainedAssumeRole(dbAWS)),
cloud.WithAmbientCredentials(),
)
if err != nil {
return "", "", "", trace.Wrap(err)
}
creds, err := sess.Config.Credentials.Get()
creds, err := awsCfg.Credentials.Retrieve(ctx)
if err != nil {
return "", "", "", trace.Wrap(err)
}
+4 -6
View File
@@ -609,9 +609,7 @@ func TestAuthGetAWSTokenWithAssumedRole(t *testing.T) {
Clock: clock,
AuthClient: new(authClientMock),
AccessPoint: new(accessPointMock),
Clients: &cloud.TestCloudClients{
STS: &fakeSTS.STSClientV1,
},
Clients: &cloud.TestCloudClients{},
AWSConfigProvider: &mocks.AWSConfigProvider{
STSClient: fakeSTS,
},
@@ -701,10 +699,10 @@ func TestGetAWSIAMCreds(t *testing.T) {
Clock: clock,
AuthClient: new(authClientMock),
AccessPoint: new(accessPointMock),
Clients: &cloud.TestCloudClients{
STS: &tt.stsMock.STSClientV1,
Clients: &cloud.TestCloudClients{},
AWSConfigProvider: &mocks.AWSConfigProvider{
STSClient: tt.stsMock,
},
AWSConfigProvider: &mocks.AWSConfigProvider{},
awsClients: fakeAWSClients{
stsClient: tt.stsMock,
},
+4 -4
View File
@@ -106,8 +106,8 @@ type EngineConfig struct {
AuthClient *authclient.Client
// AWSConfigProvider provides [aws.Config] for AWS SDK service clients.
AWSConfigProvider awsconfig.Provider
// CloudClients provides access to cloud API clients.
CloudClients cloud.Clients
// GCPClients provides access to Google Cloud API clients.
GCPClients cloud.GCPClients
// Context is the database server close context.
Context context.Context
// Clock is the clock interface.
@@ -141,8 +141,8 @@ func (c *EngineConfig) CheckAndSetDefaults() error {
if c.AWSConfigProvider == nil {
return trace.BadParameter("missing AWSConfigProvider")
}
if c.CloudClients == nil {
return trace.BadParameter("engine config CloudClients are missing")
if c.GCPClients == nil {
return trace.BadParameter("engine config GCPClients are missing")
}
if c.Context == nil {
c.Context = context.Background()
+1 -1
View File
@@ -51,7 +51,7 @@ func TestRegisterEngine(t *testing.T) {
Audit: &testAudit{},
AuthClient: &authclient.Client{},
AWSConfigProvider: &mocks.AWSConfigProvider{},
CloudClients: cloudClients,
GCPClients: cloudClients,
}
require.NoError(t, ec.CheckAndSetDefaults())
+9 -19
View File
@@ -40,7 +40,6 @@ import (
"github.com/gravitational/teleport"
apievents "github.com/gravitational/teleport/api/types/events"
apiaws "github.com/gravitational/teleport/api/utils/aws"
"github.com/gravitational/teleport/lib/cloud"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/modules"
@@ -71,8 +70,6 @@ type Engine struct {
// RoundTrippers is a cache of RoundTrippers, mapped by service endpoint.
// It is not guarded by a mutex, since requests are processed serially.
RoundTrippers map[string]http.RoundTripper
// CredentialsGetter is used to obtain STS credentials.
CredentialsGetter libaws.CredentialsGetter
// UseFIPS will ensure FIPS endpoint resolution.
UseFIPS bool
}
@@ -141,18 +138,9 @@ func (e *Engine) HandleConnection(ctx context.Context, _ *common.Session) error
}
defer e.Audit.OnSessionEnd(e.Context, e.sessionCtx)
meta := e.sessionCtx.Database.GetAWS()
awsSession, err := e.CloudClients.GetAWSSession(ctx, meta.Region,
cloud.WithAssumeRoleFromAWSMeta(meta),
cloud.WithAmbientCredentials(),
)
if err != nil {
return trace.Wrap(err)
}
signer, err := libaws.NewSigningService(libaws.SigningServiceConfig{
Clock: e.Clock,
SessionProvider: libaws.StaticAWSSessionProvider(awsSession),
CredentialsGetter: e.CredentialsGetter,
AWSConfigProvider: e.AWSConfigProvider,
})
if err != nil {
return trace.Wrap(err)
@@ -223,12 +211,14 @@ func (e *Engine) process(ctx context.Context, req *http.Request, signer *libaws.
return trace.Wrap(err)
}
signingCtx := &libaws.SigningCtx{
SigningName: re.SigningName,
SigningRegion: re.SigningRegion,
Expiry: e.sessionCtx.Identity.Expires,
SessionName: e.sessionCtx.Identity.Username,
AWSRoleArn: roleArn,
SessionTags: e.sessionCtx.Database.GetAWS().SessionTags,
SigningName: re.SigningName,
SigningRegion: re.SigningRegion,
Expiry: e.sessionCtx.Identity.Expires,
SessionName: e.sessionCtx.Identity.Username,
BaseAWSRoleARN: meta.AssumeRoleARN,
BaseAWSExternalID: meta.ExternalID,
AWSRoleArn: roleArn,
SessionTags: e.sessionCtx.Database.GetAWS().SessionTags,
}
if meta.AssumeRoleARN == "" {
signingCtx.AWSExternalID = meta.ExternalID
+3 -1
View File
@@ -101,7 +101,9 @@ func NewTestServer(config common.TestServerConfig, opts ...TestServerOption) (*T
mux := http.NewServeMux()
mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
err := awsutils.VerifyAWSSignatureV2(r, credentials.NewStaticCredentialsProvider("AKIDl", "SECRET", "SESSION"))
err := awsutils.VerifyAWSSignatureV2(r,
credentials.NewStaticCredentialsProvider("FAKEACCESSKEYID", "secret", "token"),
)
if err != nil {
code := trace.ErrorToCode(err)
body, _ := json.Marshal(jsonErr{
+11 -27
View File
@@ -22,48 +22,29 @@ import (
"context"
"crypto/tls"
"net"
"net/http"
"testing"
"github.com/aws/aws-sdk-go-v2/credentials"
awsdynamodb "github.com/aws/aws-sdk-go-v2/service/dynamodb"
"github.com/stretchr/testify/require"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/cloud/mocks"
"github.com/gravitational/teleport/lib/defaults"
libevents "github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/srv/db/common"
"github.com/gravitational/teleport/lib/srv/db/dynamodb"
awsutils "github.com/gravitational/teleport/lib/utils/aws"
"github.com/gravitational/teleport/lib/utils/aws/migration"
)
func registerTestDynamoDBEngine() {
// Override DynamoDB engine that is used normally with the test one
// with custom HTTP client.
common.RegisterEngine(newTestDynamoDBEngine, defaults.ProtocolDynamoDB)
}
func newTestDynamoDBEngine(ec common.EngineConfig) common.Engine {
return &dynamodb.Engine{
EngineConfig: ec,
RoundTrippers: make(map[string]http.RoundTripper),
// inject mock AWS credentials.
CredentialsGetter: awsutils.NewStaticCredentialsGetter(
migration.NewCredentialsAdapter(
credentials.NewStaticCredentialsProvider("AKIDl", "SECRET", "SESSION"),
),
),
}
}
func TestAccessDynamoDB(t *testing.T) {
t.Parallel()
ctx := context.Background()
mockTables := []string{"table-one", "table-two"}
testCtx := setupTestContext(ctx, t,
withDynamoDB("DynamoDB"))
testCtx := setupTestContext(ctx, t)
testCtx.server = testCtx.setupDatabaseServer(ctx, t, agentParams{
AWSConfigProvider: &mocks.AWSConfigProvider{},
Databases: []types.Database{withDynamoDB("DynamoDB")(t, ctx, testCtx)},
})
go testCtx.startHandlingConnections()
tests := []struct {
@@ -143,8 +124,11 @@ func TestAccessDynamoDB(t *testing.T) {
func TestAuditDynamoDB(t *testing.T) {
ctx := context.Background()
testCtx := setupTestContext(ctx, t,
withDynamoDB("DynamoDB"))
testCtx := setupTestContext(ctx, t)
testCtx.server = testCtx.setupDatabaseServer(ctx, t, agentParams{
AWSConfigProvider: &mocks.AWSConfigProvider{},
Databases: []types.Database{withDynamoDB("DynamoDB")(t, ctx, testCtx)},
})
go testCtx.startHandlingConnections()
testCtx.createUserAndRole(ctx, t, "alice", "admin", []string{"admin"}, []string{types.Wildcard})
+1 -1
View File
@@ -237,7 +237,7 @@ func (e *Engine) connect(ctx context.Context, sessionCtx *common.Session) (*clie
}
case sessionCtx.Database.IsCloudSQL():
// Get the client once for subsequent calls (it acquires a read lock).
gcpClient, err := e.CloudClients.GetGCPSQLAdminClient(ctx)
gcpClient, err := e.GCPClients.GetGCPSQLAdminClient(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
+4 -4
View File
@@ -38,10 +38,10 @@ type ImportRulesReader interface {
// ObjectFetcherConfig provides static object fetcher configuration.
type ObjectFetcherConfig struct {
ImportRules ImportRulesReader
Auth common.Auth
CloudClients libcloud.Clients
Log *slog.Logger
ImportRules ImportRulesReader
Auth common.Auth
GCPClients libcloud.GCPClients
Log *slog.Logger
}
// ObjectFetcher defines an interface for retrieving database objects.
+4 -4
View File
@@ -51,10 +51,10 @@ func startDatabaseImporter(ctx context.Context, cfg Config, database types.Datab
cfg.Log = cfg.Log.With("database", database.GetName(), "protocol", database.GetProtocol())
fetcher, err := GetObjectFetcher(ctx, database, ObjectFetcherConfig{
ImportRules: cfg.ImportRules,
Auth: cfg.Auth,
CloudClients: cfg.CloudClients,
Log: cfg.Log,
ImportRules: cfg.ImportRules,
Auth: cfg.Auth,
GCPClients: cfg.GCPClients,
Log: cfg.Log,
})
if err != nil {
return nil, trace.Wrap(err)
+3 -3
View File
@@ -42,7 +42,7 @@ type Config struct {
DatabaseObjectClient *databaseobject.Client
ImportRules ImportRulesReader
Auth common.Auth
CloudClients cloud.Clients
GCPClients cloud.GCPClients
// ScanInterval specifies how often the database is scanned.
// A higher ScanInterval reduces the load on the database and database agent,
@@ -113,8 +113,8 @@ func (c *Config) CheckAndSetDefaults(ctx context.Context) error {
if c.Auth == nil {
return trace.BadParameter("missing parameter Auth")
}
if c.CloudClients == nil {
return trace.BadParameter("missing parameter CloudClients")
if c.GCPClients == nil {
return trace.BadParameter("missing parameter GCPClients")
}
if c.Log == nil {
c.Log = slog.Default().With(teleport.ComponentKey, "db:obj_importer")
+4 -4
View File
@@ -34,9 +34,9 @@ import (
)
type connector struct {
auth common.Auth
cloudClients libcloud.Clients
log *slog.Logger
auth common.Auth
gcpClients libcloud.GCPClients
log *slog.Logger
certExpiry time.Time
database types.Database
@@ -91,7 +91,7 @@ func (c *connector) getConnectConfig(ctx context.Context) (*pgconn.Config, error
return nil, trace.Wrap(err)
}
// Get the client once for subsequent calls (it acquires a read lock).
gcpClient, err := c.cloudClients.GetGCPSQLAdminClient(ctx)
gcpClient, err := c.gcpClients.GetGCPSQLAdminClient(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
+3 -6
View File
@@ -505,9 +505,9 @@ func (e *Engine) receiveFromServer(serverConn *pgconn.PgConn, serverErrCh chan<-
func (e *Engine) newConnector(sessionCtx *common.Session) *connector {
conn := &connector{
auth: e.Auth,
cloudClients: e.CloudClients,
log: e.Log,
auth: e.Auth,
gcpClients: e.GCPClients,
log: e.Log,
certExpiry: sessionCtx.GetExpiry(),
database: sessionCtx.Database,
@@ -533,9 +533,6 @@ func (e *Engine) handleCancelRequest(ctx context.Context, sessionCtx *common.Ses
// Instead, use the pgconn config string parser for convenience and dial
// db host:port ourselves.
network, address := pgconn.NetworkAddress(config.Host, config.Port)
if err != nil {
return trace.Wrap(err)
}
dialer := net.Dialer{Timeout: defaults.DefaultIOTimeout}
conn, err := dialer.DialContext(ctx, network, address)
if err != nil {
+3 -3
View File
@@ -119,9 +119,9 @@ func (f *objectFetcher) getDatabaseNames(ctx context.Context) ([]string, error)
func (f *objectFetcher) connectAsAdmin(ctx context.Context, databaseName string) (*pgx.Conn, error) {
conn := &connector{
auth: f.cfg.Auth,
cloudClients: f.cfg.CloudClients,
log: f.cfg.Log,
auth: f.cfg.Auth,
gcpClients: f.cfg.GCPClients,
log: f.cfg.Log,
certExpiry: time.Now().Add(time.Hour),
database: f.db,
+4 -4
View File
@@ -204,10 +204,10 @@ func (e *Engine) applyPermissions(ctx context.Context, sessionCtx *common.Sessio
}
fetcher, err := objects.GetObjectFetcher(ctx, sessionCtx.Database, objects.ObjectFetcherConfig{
ImportRules: e.AuthClient,
Auth: e.Auth,
CloudClients: e.CloudClients,
Log: e.Log,
ImportRules: e.AuthClient,
Auth: e.Auth,
GCPClients: e.GCPClients,
Log: e.Log,
})
if err != nil {
return trace.Wrap(err)
+3 -7
View File
@@ -210,7 +210,6 @@ func (c *Config) CheckAndSetDefaults(ctx context.Context) (err error) {
}
if c.AWSDatabaseFetcherFactory == nil {
factory, err := db.NewAWSFetcherFactory(db.AWSFetcherFactoryConfig{
CloudClients: c.CloudClients,
AWSConfigProvider: c.AWSConfigProvider,
})
if err != nil {
@@ -253,7 +252,6 @@ func (c *Config) CheckAndSetDefaults(ctx context.Context) (err error) {
}
if c.CloudMeta == nil {
c.CloudMeta, err = cloud.NewMetadata(cloud.MetadataConfig{
Clients: c.CloudClients,
AWSConfigProvider: c.AWSConfigProvider,
})
if err != nil {
@@ -264,7 +262,6 @@ func (c *Config) CheckAndSetDefaults(ctx context.Context) (err error) {
c.CloudIAM, err = cloud.NewIAM(ctx, cloud.IAMConfig{
AccessPoint: c.AccessPoint,
AWSConfigProvider: c.AWSConfigProvider,
Clients: c.CloudClients,
HostID: c.HostID,
})
if err != nil {
@@ -289,7 +286,6 @@ func (c *Config) CheckAndSetDefaults(ctx context.Context) (err error) {
}
c.CloudUsers, err = users.NewUsers(users.Config{
AWSConfigProvider: c.AWSConfigProvider,
Clients: c.CloudClients,
UpdateMeta: c.CloudMeta.Update,
ClusterName: clusterName.GetClusterName(),
})
@@ -303,7 +299,7 @@ func (c *Config) CheckAndSetDefaults(ctx context.Context) (err error) {
DatabaseObjectClient: c.AuthClient.DatabaseObjectsClient(),
ImportRules: c.AuthClient,
Auth: c.Auth,
CloudClients: c.CloudClients,
GCPClients: c.CloudClients,
})
if err != nil {
return trace.Wrap(err)
@@ -313,7 +309,7 @@ func (c *Config) CheckAndSetDefaults(ctx context.Context) (err error) {
if c.discoveryResourceChecker == nil {
c.discoveryResourceChecker, err = cloud.NewDiscoveryResourceChecker(cloud.DiscoveryResourceCheckerConfig{
ResourceMatchers: c.ResourceMatchers,
Clients: c.CloudClients,
AzureClients: c.CloudClients,
AWSConfigProvider: c.AWSConfigProvider,
Context: ctx,
})
@@ -1204,7 +1200,7 @@ func (s *Server) createEngine(sessionCtx *common.Session, audit common.Audit) (c
Audit: audit,
AuthClient: s.cfg.AuthClient,
AWSConfigProvider: s.cfg.AWSConfigProvider,
CloudClients: s.cfg.CloudClients,
GCPClients: s.cfg.CloudClients,
Context: s.connContext,
Clock: s.cfg.Clock,
Log: sessionCtx.Log,
-8
View File
@@ -335,15 +335,8 @@ func TestWatcherCloudFetchers(t *testing.T) {
ctx := context.Background()
testCtx := setupTestContext(ctx, t)
testCloudClients := &clients.TestCloudClients{
AzureSQLServer: azure.NewSQLClientByAPI(&azure.ARMSQLServerMock{
AllServers: []*armsql.Server{azSQLServer},
}),
AzureManagedSQLServer: azure.NewManagedSQLClientByAPI(&azure.ARMSQLManagedServerMock{}),
}
dbFetcherFactory, err := db.NewAWSFetcherFactory(db.AWSFetcherFactoryConfig{
AWSConfigProvider: &mocks.AWSConfigProvider{},
CloudClients: testCloudClients,
AWSClients: fakeAWSClients{
rdsClient: &mocks.RDSClient{Unauth: true}, // Access denied error should not affect other fetchers.
rssClient: &mocks.RedshiftServerlessClient{
@@ -376,7 +369,6 @@ func TestWatcherCloudFetchers(t *testing.T) {
},
}},
CloudClients: &clients.TestCloudClients{
STS: &mocks.STSClientV1{},
AzureSQLServer: azure.NewSQLClientByAPI(&azure.ARMSQLServerMock{
AllServers: []*armsql.Server{azSQLServer},
}),
-1
View File
@@ -505,7 +505,6 @@ func (s *Server) accessGraphAWSFetchersFromMatchers(ctx context.Context, matcher
ctx,
aws_sync.Config{
AWSConfigProvider: s.AWSConfigProvider,
CloudClients: s.CloudClients,
GetEKSClient: s.GetAWSSyncEKSClient,
GetEC2Client: s.GetEC2Client,
AssumeRole: assumeRole,
+2 -2
View File
@@ -190,11 +190,11 @@ func TestServer_updateDiscoveryConfigStatus(t *testing.T) {
name: "merge two errors",
args: args{
fetchers: []*fakeFetcher{
&fakeFetcher{
{
discoveryConfigName: "test1",
err: fmt.Errorf("error in fetcher 1"),
},
&fakeFetcher{
{
discoveryConfigName: "test1",
err: fmt.Errorf("error in fetcher 2"),
},
+1 -9
View File
@@ -36,7 +36,6 @@ import (
"github.com/aws/aws-sdk-go-v2/service/ssm"
ssmtypes "github.com/aws/aws-sdk-go-v2/service/ssm/types"
"github.com/aws/aws-sdk-go-v2/service/sts"
"github.com/aws/aws-sdk-go/aws/session"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"google.golang.org/protobuf/types/known/timestamppb"
@@ -57,7 +56,6 @@ import (
"github.com/gravitational/teleport/lib/cloud/awsconfig"
gcpimds "github.com/gravitational/teleport/lib/cloud/imds/gcp"
"github.com/gravitational/teleport/lib/cryptosuites"
"github.com/gravitational/teleport/lib/integrations/awsoidc"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/services/readonly"
"github.com/gravitational/teleport/lib/srv/discovery/common"
@@ -240,12 +238,7 @@ func (c *Config) CheckAndSetDefaults() error {
kubernetes matchers are present.`)
}
if c.CloudClients == nil {
awsIntegrationSessionProvider := func(ctx context.Context, region, integration string) (*session.Session, error) {
return awsoidc.NewSessionV1(ctx, c.AccessPoint, region, integration)
}
cloudClients, err := cloud.NewClients(
cloud.WithAWSIntegrationSessionProvider(awsIntegrationSessionProvider),
)
cloudClients, err := cloud.NewClients()
if err != nil {
return trace.Wrap(err, "unable to create cloud clients")
}
@@ -264,7 +257,6 @@ kubernetes matchers are present.`)
}
if c.AWSDatabaseFetcherFactory == nil {
factory, err := db.NewAWSFetcherFactory(db.AWSFetcherFactoryConfig{
CloudClients: c.CloudClients,
AWSConfigProvider: c.AWSConfigProvider,
})
if err != nil {
-9
View File
@@ -365,7 +365,6 @@ func TestDiscoveryServer(t *testing.T) {
wantInstalledInstances []string
wantDiscoveryConfigStatus *discoveryconfig.Status
userTasksDiscoverCheck require.ValueAssertionFunc
cloudClients cloud.Clients
ssmRunError error
}{
{
@@ -947,7 +946,6 @@ func TestDiscoveryServer(t *testing.T) {
Emitter: tc.emitter,
Log: logger,
DiscoveryGroup: defaultDiscoveryGroup,
CloudClients: tc.cloudClients,
clock: fakeClock,
})
require.NoError(t, err)
@@ -2174,7 +2172,6 @@ func TestDiscoveryDatabase(t *testing.T) {
}
testCloudClients := &cloud.TestCloudClients{
STS: &mocks.STSClientV1{},
AzureRedis: azure.NewRedisClientByAPI(&azure.ARMRedisMock{
Servers: []*armredis.ResourceInfo{azRedisResource},
}),
@@ -2550,7 +2547,6 @@ func TestDiscoveryDatabase(t *testing.T) {
}
dbFetcherFactory, err := db.NewAWSFetcherFactory(db.AWSFetcherFactoryConfig{
AWSConfigProvider: fakeConfigProvider,
CloudClients: testCloudClients,
AWSClients: fakeAWSClients{
ecClient: &mocks.ElastiCacheClient{},
mdbClient: &mocks.MemoryDBClient{},
@@ -2685,12 +2681,8 @@ func TestDiscoveryDatabaseRemovingDiscoveryConfigs(t *testing.T) {
fakeConfigProvider := &mocks.AWSConfigProvider{
STSClient: &mocks.STSClient{},
}
testCloudClients := &cloud.TestCloudClients{
STS: &fakeConfigProvider.STSClient.STSClientV1,
}
dbFetcherFactory, err := db.NewAWSFetcherFactory(db.AWSFetcherFactoryConfig{
AWSConfigProvider: fakeConfigProvider,
CloudClients: testCloudClients,
AWSClients: fakeAWSClients{
rdsClient: &mocks.RDSClient{
DBInstances: []rdstypes.DBInstance{*awsRDSInstance},
@@ -2730,7 +2722,6 @@ func TestDiscoveryDatabaseRemovingDiscoveryConfigs(t *testing.T) {
&Config{
AWSConfigProvider: fakeConfigProvider,
AWSDatabaseFetcherFactory: dbFetcherFactory,
CloudClients: testCloudClients,
ClusterFeatures: func() proto.Features { return proto.Features{} },
KubernetesClient: fake.NewSimpleClientset(),
AccessPoint: getDiscoveryAccessPoint(tlsServer.Auth(), authClient),
@@ -35,7 +35,6 @@ import (
usageeventsv1 "github.com/gravitational/teleport/api/gen/proto/go/usageevents/v1"
accessgraphv1alpha "github.com/gravitational/teleport/gen/proto/go/accessgraph/v1alpha"
"github.com/gravitational/teleport/lib/cloud"
"github.com/gravitational/teleport/lib/cloud/awsconfig"
"github.com/gravitational/teleport/lib/srv/server"
)
@@ -48,8 +47,6 @@ const pageSize int32 = 500
type Config struct {
// AWSConfigProvider provides [aws.Config] for AWS SDK service clients.
AWSConfigProvider awsconfig.Provider
// CloudClients is the cloud clients to use when fetching AWS resources.
CloudClients cloud.Clients
// GetEKSClient gets an AWS EKS client for the given region.
GetEKSClient EKSClientGetter
// GetEC2Client gets an AWS EC2 client for the given region.
-6
View File
@@ -27,7 +27,6 @@ import (
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/cloud"
"github.com/gravitational/teleport/lib/cloud/awsconfig"
"github.com/gravitational/teleport/lib/srv/discovery/common"
)
@@ -49,8 +48,6 @@ type awsFetcherPlugin interface {
// awsFetcherConfig is the AWS database fetcher configuration.
type awsFetcherConfig struct {
// AWSClients are the AWS API clients.
AWSClients cloud.AWSClients
// AWSConfigProvider provides [aws.Config] for AWS SDK service clients.
AWSConfigProvider awsconfig.Provider
// Type is the type of DB matcher, for example "rds", "redshift", etc.
@@ -78,9 +75,6 @@ type awsFetcherConfig struct {
// CheckAndSetDefaults validates the config and sets defaults.
func (cfg *awsFetcherConfig) CheckAndSetDefaults(component string) error {
if cfg.AWSClients == nil {
return trace.BadParameter("missing parameter AWSClients")
}
if cfg.AWSConfigProvider == nil {
return trace.BadParameter("missing AWSConfigProvider")
}
@@ -26,7 +26,6 @@ import (
"github.com/stretchr/testify/require"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/cloud"
"github.com/gravitational/teleport/lib/cloud/awstesthelpers"
"github.com/gravitational/teleport/lib/cloud/mocks"
"github.com/gravitational/teleport/lib/srv/discovery/common"
@@ -51,8 +50,7 @@ func TestRedshiftServerlessFetcher(t *testing.T) {
tests := []awsFetcherTest{
{
name: "fetch all",
inputClients: &cloud.TestCloudClients{},
name: "fetch all",
fetcherCfg: AWSFetcherFactoryConfig{
AWSClients: fakeAWSClients{
rssClient: &mocks.RedshiftServerlessClient{
-6
View File
@@ -121,14 +121,9 @@ type AWSFetcherFactoryConfig struct {
AWSConfigProvider awsconfig.Provider
// AWSClients provides AWS SDK clients.
AWSClients AWSClientProvider
// CloudClients is an interface for retrieving AWS SDK v1 cloud clients.
CloudClients cloud.AWSClients
}
func (c *AWSFetcherFactoryConfig) checkAndSetDefaults() error {
if c.CloudClients == nil {
return trace.BadParameter("missing CloudClients")
}
if c.AWSConfigProvider == nil {
return trace.BadParameter("missing AWSConfigProvider")
}
@@ -173,7 +168,6 @@ func (f *AWSFetcherFactory) MakeFetchers(ctx context.Context, matchers []types.A
for _, makeFetcher := range makeFetchers {
for _, region := range matcher.Regions {
fetcher, err := makeFetcher(awsFetcherConfig{
AWSClients: f.cfg.CloudClients,
Type: matcherType,
AssumeRole: assumeRole,
Labels: matcher.Tags,
@@ -112,7 +112,6 @@ var testAssumeRole = types.AssumeRole{
// awsFetcherTest is a common test struct for AWS fetchers.
type awsFetcherTest struct {
name string
inputClients *cloud.TestCloudClients
fetcherCfg AWSFetcherFactoryConfig
inputMatchers []types.AWSMatcher
wantDatabases types.Databases
@@ -125,11 +124,6 @@ func testAWSFetchers(t *testing.T, tests ...awsFetcherTest) {
for _, test := range tests {
test := test
fakeSTS := &mocks.STSClient{}
if test.inputClients != nil {
require.Nil(t, test.inputClients.STS, "testAWSFetchers injects an STS mock itself, but test input had already configured it. This is a test configuration error.")
test.inputClients.STS = &fakeSTS.STSClientV1
}
test.fetcherCfg.CloudClients = test.inputClients
require.Nil(t, test.fetcherCfg.AWSConfigProvider, "testAWSFetchers injects a fake AWSConfigProvider, but the test input had already configured it. This is a test configuration error.")
test.fetcherCfg.AWSConfigProvider = &mocks.AWSConfigProvider{
STSClient: fakeSTS,
+1 -1
View File
@@ -33,7 +33,7 @@ import (
)
type mockClients struct {
cloud.Clients
cloud.AzureClients
azureClient azure.VirtualMachinesClient
}