mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
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:
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
})
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(),
|
||||
})
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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())
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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})
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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},
|
||||
}),
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"),
|
||||
},
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -33,7 +33,7 @@ import (
|
||||
)
|
||||
|
||||
type mockClients struct {
|
||||
cloud.Clients
|
||||
cloud.AzureClients
|
||||
|
||||
azureClient azure.VirtualMachinesClient
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user