mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
Support AWS external id (#14086)
This commit is contained in:
@@ -63,6 +63,8 @@ type Application interface {
|
||||
IsAWSConsole() bool
|
||||
// GetAWSAccountID returns value of label containing AWS account ID on this app.
|
||||
GetAWSAccountID() string
|
||||
// GetAWSExternalID returns the AWS External ID configured for this app.
|
||||
GetAWSExternalID() string
|
||||
// Copy returns a copy of this app resource.
|
||||
Copy() *AppV3
|
||||
}
|
||||
@@ -239,6 +241,14 @@ func (a *AppV3) GetAWSAccountID() string {
|
||||
return a.Metadata.Labels[constants.AWSAccountIDLabel]
|
||||
}
|
||||
|
||||
// GetAWSExternalID returns the AWS External ID configured for this app.
|
||||
func (a *AppV3) GetAWSExternalID() string {
|
||||
if a.Spec.AWS == nil {
|
||||
return ""
|
||||
}
|
||||
return a.Spec.AWS.ExternalID
|
||||
}
|
||||
|
||||
// String returns the app string representation.
|
||||
func (a *AppV3) String() string {
|
||||
return fmt.Sprintf("App(Name=%v, PublicAddr=%v, Labels=%v)",
|
||||
|
||||
@@ -22,6 +22,8 @@ import (
|
||||
|
||||
"github.com/gravitational/trace"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/gravitational/teleport/api/constants"
|
||||
)
|
||||
|
||||
// TestAppPublicAddrValidation tests PublicAddr field validation to make sure that
|
||||
@@ -164,3 +166,38 @@ func TestAppServerSorter(t *testing.T) {
|
||||
servers := makeServers(testValsUnordered, "does-not-matter")
|
||||
require.True(t, trace.IsNotImplemented(AppServers(servers).SortByCustom(sortBy)))
|
||||
}
|
||||
|
||||
func TestApplicationGetAWSExternalID(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
appAWS *AppAWS
|
||||
expectedExternalID string
|
||||
}{
|
||||
{
|
||||
name: "not configured",
|
||||
},
|
||||
{
|
||||
name: "configured",
|
||||
appAWS: &AppAWS{
|
||||
ExternalID: "default-external-id",
|
||||
},
|
||||
expectedExternalID: "default-external-id",
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
app, err := NewAppV3(Metadata{
|
||||
Name: "aws",
|
||||
}, AppSpecV3{
|
||||
URI: constants.AWSConsoleURL,
|
||||
AWS: test.appAWS,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Equal(t, test.expectedExternalID, app.GetAWSExternalID())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
+1405
-1173
File diff suppressed because it is too large
Load Diff
@@ -543,6 +543,8 @@ message AppSpecV3 {
|
||||
bool InsecureSkipVerify = 4 [ (gogoproto.jsontag) = "insecure_skip_verify" ];
|
||||
// Rewrite is a list of rewriting rules to apply to requests and responses.
|
||||
Rewrite Rewrite = 5 [ (gogoproto.jsontag) = "rewrite,omitempty" ];
|
||||
// AWS contains additional options for AWS applications.
|
||||
AppAWS AWS = 6 [ (gogoproto.jsontag) = "aws,omitempty" ];
|
||||
}
|
||||
|
||||
// App is a specific application that a server proxies.
|
||||
@@ -599,6 +601,12 @@ message CommandLabelV2 {
|
||||
string Result = 3 [ (gogoproto.jsontag) = "result" ];
|
||||
}
|
||||
|
||||
// AppAWS contains additional options for AWS applications.
|
||||
message AppAWS {
|
||||
// ExternalID is the AWS External ID used when assuming roles in this app.
|
||||
string ExternalID = 1 [ (gogoproto.jsontag) = "external_id,omitempty" ];
|
||||
}
|
||||
|
||||
// PrivateKeyType is the storage type of a private key.
|
||||
enum PrivateKeyType {
|
||||
// RAW is a plaintext private key.
|
||||
|
||||
@@ -1292,6 +1292,11 @@ func applyAppsConfig(fc *FileConfig, cfg *service.Config) error {
|
||||
Headers: headers,
|
||||
}
|
||||
}
|
||||
if application.AWS != nil {
|
||||
app.AWS = &service.AppAWS{
|
||||
ExternalID: application.AWS.ExternalID,
|
||||
}
|
||||
}
|
||||
if err := app.CheckAndSetDefaults(); err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
|
||||
@@ -1351,6 +1351,9 @@ type App struct {
|
||||
|
||||
// Rewrite defines a block that is used to rewrite requests and responses.
|
||||
Rewrite *Rewrite `yaml:"rewrite,omitempty"`
|
||||
|
||||
// AWS contains additional options for AWS applications.
|
||||
AWS *AppAWS `yaml:"aws,omitempty"`
|
||||
}
|
||||
|
||||
// Rewrite is a list of rewriting rules to apply to requests and responses.
|
||||
@@ -1361,6 +1364,12 @@ type Rewrite struct {
|
||||
Headers []string `yaml:"headers,omitempty"`
|
||||
}
|
||||
|
||||
// AppAWS contains additional options for AWS applications.
|
||||
type AppAWS struct {
|
||||
// ExternalID is the AWS External ID used when assuming roles in this app.
|
||||
ExternalID string `yaml:"external_id,omitempty"`
|
||||
}
|
||||
|
||||
// Proxy is a `proxy_service` section of the config file:
|
||||
type Proxy struct {
|
||||
// Service is a generic service configuration section
|
||||
|
||||
@@ -945,6 +945,9 @@ type App struct {
|
||||
|
||||
// Rewrite defines a block that is used to rewrite requests and responses.
|
||||
Rewrite *Rewrite
|
||||
|
||||
// AWS contains additional options for AWS applications.
|
||||
AWS *AppAWS `yaml:"aws,omitempty"`
|
||||
}
|
||||
|
||||
// CheckAndSetDefaults validates an application.
|
||||
@@ -1216,6 +1219,12 @@ func ParseHeaders(headers []string) (headersOut []Header, err error) {
|
||||
return headersOut, nil
|
||||
}
|
||||
|
||||
// AppAWS contains additional options for AWS applications.
|
||||
type AppAWS struct {
|
||||
// ExternalID is the AWS External ID used when assuming roles in this app.
|
||||
ExternalID string `yaml:"external_id,omitempty"`
|
||||
}
|
||||
|
||||
// MakeDefaultConfig creates a new Config structure and populates it with defaults
|
||||
func MakeDefaultConfig() (config *Config) {
|
||||
config = &Config{}
|
||||
|
||||
@@ -4190,6 +4190,13 @@ func (process *TeleportProcess) initApps() {
|
||||
}
|
||||
}
|
||||
|
||||
var aws *types.AppAWS
|
||||
if app.AWS != nil {
|
||||
aws = &types.AppAWS{
|
||||
ExternalID: app.AWS.ExternalID,
|
||||
}
|
||||
}
|
||||
|
||||
a, err := types.NewAppV3(types.Metadata{
|
||||
Name: app.Name,
|
||||
Description: app.Description,
|
||||
@@ -4200,6 +4207,7 @@ func (process *TeleportProcess) initApps() {
|
||||
DynamicLabels: types.LabelsToV2(app.DynamicLabels),
|
||||
InsecureSkipVerify: app.InsecureSkipVerify,
|
||||
Rewrite: rewrite,
|
||||
AWS: aws,
|
||||
})
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
|
||||
+15
-22
@@ -18,10 +18,10 @@ package aws
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
|
||||
"github.com/aws/aws-sdk-go/aws"
|
||||
"github.com/aws/aws-sdk-go/aws/client"
|
||||
"github.com/aws/aws-sdk-go/aws/credentials"
|
||||
"github.com/aws/aws-sdk-go/aws/credentials/stscreds"
|
||||
@@ -34,10 +34,9 @@ import (
|
||||
"github.com/jonboulle/clockwork"
|
||||
"github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/gravitational/teleport/lib/auth"
|
||||
"github.com/gravitational/teleport/lib/defaults"
|
||||
"github.com/gravitational/teleport/lib/srv/app/common"
|
||||
appcommon "github.com/gravitational/teleport/lib/srv/app/common"
|
||||
"github.com/gravitational/teleport/lib/tlsca"
|
||||
awsutils "github.com/gravitational/teleport/lib/utils/aws"
|
||||
)
|
||||
|
||||
@@ -142,7 +141,7 @@ func (s *SigningService) Handle(rw http.ResponseWriter, r *http.Request) {
|
||||
// 5) Sign HTTP request.
|
||||
// 6) Forward the signed HTTP request to the AWS API.
|
||||
func (s *SigningService) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
identity, err := getUserIdentityFromContext(req.Context())
|
||||
sessionCtx, err := common.GetSessionContext(req)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -150,7 +149,7 @@ func (s *SigningService) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
signedReq, err := s.prepareSignedRequest(req, resolvedEndpoint, identity)
|
||||
signedReq, err := s.prepareSignedRequest(req, resolvedEndpoint, sessionCtx)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -161,16 +160,6 @@ func (s *SigningService) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
func getUserIdentityFromContext(ctx context.Context) (*tlsca.Identity, error) {
|
||||
ctxUser := ctx.Value(auth.ContextUser)
|
||||
userI, ok := ctxUser.(auth.IdentityGetter)
|
||||
if !ok {
|
||||
return nil, trace.BadParameter("failed to get user identity")
|
||||
}
|
||||
identity := userI.GetIdentity()
|
||||
return &identity, nil
|
||||
}
|
||||
|
||||
func (s *SigningService) formatForwardResponseError(rw http.ResponseWriter, r *http.Request, err error) {
|
||||
switch trace.Unwrap(err).(type) {
|
||||
case *trace.BadParameterError:
|
||||
@@ -187,7 +176,7 @@ func (s *SigningService) formatForwardResponseError(rw http.ResponseWriter, r *h
|
||||
|
||||
// prepareSignedRequest creates a new HTTP request and rewrites the header from the original request and returns a new
|
||||
// HTTP request signed by STS AWS API.
|
||||
func (s *SigningService) prepareSignedRequest(r *http.Request, re *endpoints.ResolvedEndpoint, identity *tlsca.Identity) (*http.Request, error) {
|
||||
func (s *SigningService) prepareSignedRequest(r *http.Request, re *endpoints.ResolvedEndpoint, sessionCtx *common.SessionContext) (*http.Request, error) {
|
||||
payload, err := awsutils.GetAndReplaceReqBody(r)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -199,7 +188,7 @@ func (s *SigningService) prepareSignedRequest(r *http.Request, re *endpoints.Res
|
||||
}
|
||||
rewriteHeaders(r, reqCopy)
|
||||
// Sign the copy of the request.
|
||||
signer := v4.NewSigner(s.getSigningCredentials(s.Session, identity))
|
||||
signer := v4.NewSigner(s.getSigningCredentials(s.Session, sessionCtx))
|
||||
_, err = signer.Sign(reqCopy, bytes.NewReader(payload), re.SigningName, re.SigningRegion, s.Clock.Now())
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -220,13 +209,17 @@ func rewriteHeaders(r *http.Request, reqCopy *http.Request) {
|
||||
reqCopy.Header.Del("Content-Length")
|
||||
}
|
||||
|
||||
type getSigningCredentialsFunc func(c client.ConfigProvider, identity *tlsca.Identity) *credentials.Credentials
|
||||
type getSigningCredentialsFunc func(c client.ConfigProvider, sessionCtx *common.SessionContext) *credentials.Credentials
|
||||
|
||||
func getAWSCredentialsFromSTSAPI(provider client.ConfigProvider, identity *tlsca.Identity) *credentials.Credentials {
|
||||
return stscreds.NewCredentials(provider, identity.RouteToApp.AWSRoleARN,
|
||||
func getAWSCredentialsFromSTSAPI(provider client.ConfigProvider, sessionCtx *common.SessionContext) *credentials.Credentials {
|
||||
return stscreds.NewCredentials(provider, sessionCtx.Identity.RouteToApp.AWSRoleARN,
|
||||
func(cred *stscreds.AssumeRoleProvider) {
|
||||
cred.RoleSessionName = identity.Username
|
||||
cred.Expiry.SetExpiration(identity.Expires, 0)
|
||||
cred.RoleSessionName = sessionCtx.Identity.Username
|
||||
cred.Expiry.SetExpiration(sessionCtx.Identity.Expires, 0)
|
||||
|
||||
if externalID := sessionCtx.App.GetAWSExternalID(); externalID != "" {
|
||||
cred.ExternalID = aws.String(externalID)
|
||||
}
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
@@ -33,7 +33,10 @@ import (
|
||||
"github.com/jonboulle/clockwork"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/gravitational/teleport/api/constants"
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/teleport/lib/auth"
|
||||
"github.com/gravitational/teleport/lib/srv/app/common"
|
||||
"github.com/gravitational/teleport/lib/tlsca"
|
||||
awsutils "github.com/gravitational/teleport/lib/utils/aws"
|
||||
)
|
||||
@@ -137,11 +140,27 @@ func TestAWSSignerHandler(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func staticAWSCredentials(client.ConfigProvider, *tlsca.Identity) *credentials.Credentials {
|
||||
func staticAWSCredentials(client.ConfigProvider, *common.SessionContext) *credentials.Credentials {
|
||||
return credentials.NewStaticCredentials("AKIDl", "SECRET", "SESSION")
|
||||
}
|
||||
|
||||
func createSuite(t *testing.T, handler http.HandlerFunc) *httptest.Server {
|
||||
type suite struct {
|
||||
*httptest.Server
|
||||
|
||||
identity *tlsca.Identity
|
||||
app types.Application
|
||||
}
|
||||
|
||||
func createSuite(t *testing.T, handler http.HandlerFunc) *suite {
|
||||
user := auth.LocalUser{Username: "user"}
|
||||
app, err := types.NewAppV3(types.Metadata{
|
||||
Name: "awsconsole",
|
||||
}, types.AppSpecV3{
|
||||
URI: constants.AWSConsoleURL,
|
||||
PublicAddr: "test.local",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
awsAPIMock := httptest.NewUnstartedServer(handler)
|
||||
awsAPIMock.StartTLS()
|
||||
t.Cleanup(func() {
|
||||
@@ -168,12 +187,21 @@ func createSuite(t *testing.T, handler http.HandlerFunc) *httptest.Server {
|
||||
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/", func(writer http.ResponseWriter, request *http.Request) {
|
||||
ctx := context.WithValue(context.Background(), auth.ContextUser, auth.LocalUser{Username: "user"})
|
||||
svc.Handle(writer, request.WithContext(ctx))
|
||||
request = common.WithSessionContext(request, &common.SessionContext{
|
||||
Identity: &user.Identity,
|
||||
App: app,
|
||||
})
|
||||
|
||||
svc.Handle(writer, request)
|
||||
})
|
||||
server := httptest.NewServer(mux)
|
||||
t.Cleanup(func() {
|
||||
server.Close()
|
||||
})
|
||||
return server
|
||||
|
||||
return &suite{
|
||||
Server: server,
|
||||
identity: &user.Identity,
|
||||
app: app,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -26,6 +26,7 @@ import (
|
||||
|
||||
"github.com/gravitational/teleport/lib/tlsca"
|
||||
|
||||
"github.com/aws/aws-sdk-go/aws"
|
||||
"github.com/aws/aws-sdk-go/aws/arn"
|
||||
"github.com/aws/aws-sdk-go/aws/credentials/ec2rolecreds"
|
||||
"github.com/aws/aws-sdk-go/aws/credentials/ssocreds"
|
||||
@@ -52,6 +53,8 @@ type AWSSigninRequest struct {
|
||||
TargetURL string
|
||||
// Issuer is the application public URL.
|
||||
Issuer string
|
||||
// ExternalID is the AWS external ID.
|
||||
ExternalID string
|
||||
}
|
||||
|
||||
// CheckAndSetDefaults validates the request.
|
||||
@@ -193,6 +196,10 @@ func (c *cloud) getAWSSigninToken(req *AWSSigninRequest, endpoint string, option
|
||||
if temporarySession {
|
||||
creds.Duration = duration
|
||||
}
|
||||
|
||||
if req.ExternalID != "" {
|
||||
creds.ExternalID = aws.String(req.ExternalID)
|
||||
}
|
||||
})
|
||||
stsCredentials, err := stscreds.NewCredentials(c.cfg.Session, req.Identity.RouteToApp.AWSRoleARN, options...).Get()
|
||||
if err != nil {
|
||||
|
||||
@@ -0,0 +1,57 @@
|
||||
/*
|
||||
Copyright 2022 Gravitational, Inc.
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package common
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/teleport/lib/tlsca"
|
||||
"github.com/gravitational/trace"
|
||||
)
|
||||
|
||||
// SessionContext contains common context parameters for an App session.
|
||||
type SessionContext struct {
|
||||
// Identity is the requested identity.
|
||||
Identity *tlsca.Identity
|
||||
// App is the requested identity.
|
||||
App types.Application
|
||||
}
|
||||
|
||||
// WithSessionContext adds session context to provided request.
|
||||
func WithSessionContext(r *http.Request, sessionCtx *SessionContext) *http.Request {
|
||||
return r.WithContext(context.WithValue(
|
||||
r.Context(),
|
||||
contextSessionKey,
|
||||
sessionCtx,
|
||||
))
|
||||
}
|
||||
|
||||
// GetSessionContext retrieves the session context from a request.
|
||||
func GetSessionContext(r *http.Request) (*SessionContext, error) {
|
||||
sessionCtxValue := r.Context().Value(contextSessionKey)
|
||||
sessionCtx, ok := sessionCtxValue.(*SessionContext)
|
||||
if !ok {
|
||||
return nil, trace.BadParameter("failed to get session context")
|
||||
}
|
||||
return sessionCtx, nil
|
||||
}
|
||||
|
||||
type contextKey string
|
||||
|
||||
const (
|
||||
// contextSessionKey is the context key for the session context.
|
||||
contextSessionKey contextKey = "app-session-context"
|
||||
)
|
||||
+13
-4
@@ -40,6 +40,7 @@ import (
|
||||
"github.com/gravitational/teleport/lib/services"
|
||||
"github.com/gravitational/teleport/lib/srv"
|
||||
appaws "github.com/gravitational/teleport/lib/srv/app/aws"
|
||||
"github.com/gravitational/teleport/lib/srv/app/common"
|
||||
"github.com/gravitational/teleport/lib/tlsca"
|
||||
"github.com/gravitational/teleport/lib/utils"
|
||||
"github.com/gravitational/teleport/lib/utils/aws"
|
||||
@@ -617,8 +618,15 @@ func (s *Server) serveHTTP(w http.ResponseWriter, r *http.Request) error {
|
||||
// AWS CLI, automatically use SigV4 for all services that support it (All services expect Amazon SimpleDB
|
||||
// but this AWS service has been deprecated)
|
||||
if aws.IsSignedByAWSSigV4(r) && app.IsAWSConsole() {
|
||||
// TODO(greedy52) create a proper sessionChunk for AWS requests to
|
||||
// record audit events.
|
||||
sessionCtx := &common.SessionContext{
|
||||
Identity: identity,
|
||||
App: app,
|
||||
}
|
||||
|
||||
// Sign the request based on RouteToApp.AWSRoleARN user identity and route signed request to the AWS API.
|
||||
s.awsSigner.Handle(w, r)
|
||||
s.awsSigner.Handle(w, common.WithSessionContext(r, sessionCtx))
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -628,9 +636,10 @@ func (s *Server) serveHTTP(w http.ResponseWriter, r *http.Request) error {
|
||||
s.log.Debugf("Redirecting %v to AWS mananement console with role %v.",
|
||||
identity.Username, identity.RouteToApp.AWSRoleARN)
|
||||
url, err := s.c.Cloud.GetAWSSigninURL(AWSSigninRequest{
|
||||
Identity: identity,
|
||||
TargetURL: app.GetURI(),
|
||||
Issuer: app.GetPublicAddr(),
|
||||
Identity: identity,
|
||||
TargetURL: app.GetURI(),
|
||||
Issuer: app.GetPublicAddr(),
|
||||
ExternalID: app.GetAWSExternalID(),
|
||||
})
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
|
||||
+36
-1
@@ -22,6 +22,7 @@ import (
|
||||
"io"
|
||||
"net/http"
|
||||
"net/textproto"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -225,7 +226,7 @@ func filterHeaders(r *http.Request, headers []string) {
|
||||
// specified AWS account ID.
|
||||
//
|
||||
// If AWS account ID is empty, all roles are returned.
|
||||
func FilterAWSRoles(arns []string, accountID string) (result []Role) {
|
||||
func FilterAWSRoles(arns []string, accountID string) (result Roles) {
|
||||
for _, roleARN := range arns {
|
||||
parsed, err := arn.Parse(roleARN)
|
||||
if err != nil || (accountID != "" && parsed.AccountID != accountID) {
|
||||
@@ -244,6 +245,7 @@ func FilterAWSRoles(arns []string, accountID string) (result []Role) {
|
||||
continue
|
||||
}
|
||||
result = append(result, Role{
|
||||
Name: strings.Join(parts[1:], "/"),
|
||||
Display: parts[numParts-1],
|
||||
ARN: roleARN,
|
||||
})
|
||||
@@ -253,8 +255,41 @@ func FilterAWSRoles(arns []string, accountID string) (result []Role) {
|
||||
|
||||
// Role describes an AWS IAM role for AWS console access.
|
||||
type Role struct {
|
||||
// Name is the full role name with the entire path.
|
||||
Name string `json:"name"`
|
||||
// Display is the role display name.
|
||||
Display string `json:"display"`
|
||||
// ARN is the full role ARN.
|
||||
ARN string `json:"arn"`
|
||||
}
|
||||
|
||||
// Roles is a slice of roles.
|
||||
type Roles []Role
|
||||
|
||||
// Sort sorts the roles by their display names.
|
||||
func (roles Roles) Sort() {
|
||||
sort.SliceStable(roles, func(x, y int) bool {
|
||||
return strings.ToLower(roles[x].Display) < strings.ToLower(roles[y].Display)
|
||||
})
|
||||
}
|
||||
|
||||
// FindRoleByARN finds the role with the provided ARN.
|
||||
func (roles Roles) FindRoleByARN(arn string) (Role, bool) {
|
||||
for _, role := range roles {
|
||||
if role.ARN == arn {
|
||||
return role, true
|
||||
}
|
||||
}
|
||||
return Role{}, false
|
||||
}
|
||||
|
||||
// FindRolesByName finds all roles matching the provided name.
|
||||
func (roles Roles) FindRolesByName(name string) (result Roles) {
|
||||
for _, role := range roles {
|
||||
// Match either full name or the display name.
|
||||
if role.Display == name || role.Name == name {
|
||||
result = append(result, role)
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
@@ -95,18 +95,22 @@ func TestFilterAWSRoles(t *testing.T) {
|
||||
acc1ARN1 := Role{
|
||||
ARN: "arn:aws:iam::1234567890:role/EC2FullAccess",
|
||||
Display: "EC2FullAccess",
|
||||
Name: "EC2FullAccess",
|
||||
}
|
||||
acc1ARN2 := Role{
|
||||
ARN: "arn:aws:iam::1234567890:role/EC2ReadOnly",
|
||||
Display: "EC2ReadOnly",
|
||||
Name: "EC2ReadOnly",
|
||||
}
|
||||
acc1ARN3 := Role{
|
||||
ARN: "arn:aws:iam::1234567890:role/path/to/customrole",
|
||||
Display: "customrole",
|
||||
Name: "path/to/customrole",
|
||||
}
|
||||
acc2ARN1 := Role{
|
||||
ARN: "arn:aws:iam::0987654321:role/test-role",
|
||||
Display: "test-role",
|
||||
Name: "test-role",
|
||||
}
|
||||
invalidARN := Role{
|
||||
ARN: "invalid-arn",
|
||||
@@ -117,25 +121,78 @@ func TestFilterAWSRoles(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
accountID string
|
||||
outARNs []Role
|
||||
outARNs Roles
|
||||
}{
|
||||
{
|
||||
name: "first account roles",
|
||||
accountID: "1234567890",
|
||||
outARNs: []Role{acc1ARN1, acc1ARN2, acc1ARN3},
|
||||
outARNs: Roles{acc1ARN1, acc1ARN2, acc1ARN3},
|
||||
},
|
||||
{
|
||||
name: "second account roles",
|
||||
accountID: "0987654321",
|
||||
outARNs: []Role{acc2ARN1},
|
||||
outARNs: Roles{acc2ARN1},
|
||||
},
|
||||
{
|
||||
name: "all roles",
|
||||
accountID: "",
|
||||
outARNs: []Role{acc1ARN1, acc1ARN2, acc1ARN3, acc2ARN1},
|
||||
outARNs: Roles{acc1ARN1, acc1ARN2, acc1ARN3, acc2ARN1},
|
||||
},
|
||||
}
|
||||
for _, test := range tests {
|
||||
require.Equal(t, test.outARNs, FilterAWSRoles(allARNS, test.accountID))
|
||||
}
|
||||
}
|
||||
|
||||
func TestRoles(t *testing.T) {
|
||||
arns := []string{
|
||||
"arn:aws:iam::1234567890:role/test-role",
|
||||
"arn:aws:iam::1234567890:role/EC2FullAccess",
|
||||
"arn:aws:iam::1234567890:role/path/to/EC2FullAccess",
|
||||
}
|
||||
roles := FilterAWSRoles(arns, "1234567890")
|
||||
require.Len(t, roles, 3)
|
||||
|
||||
t.Run("Sort", func(t *testing.T) {
|
||||
roles.Sort()
|
||||
require.Equal(t, "arn:aws:iam::1234567890:role/EC2FullAccess", roles[0].ARN)
|
||||
require.Equal(t, "arn:aws:iam::1234567890:role/path/to/EC2FullAccess", roles[1].ARN)
|
||||
require.Equal(t, "arn:aws:iam::1234567890:role/test-role", roles[2].ARN)
|
||||
})
|
||||
|
||||
t.Run("FindRoleByARN", func(t *testing.T) {
|
||||
t.Run("found", func(t *testing.T) {
|
||||
for _, arn := range arns {
|
||||
role, found := roles.FindRoleByARN(arn)
|
||||
require.True(t, found)
|
||||
require.Equal(t, role.ARN, arn)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("not found", func(t *testing.T) {
|
||||
_, found := roles.FindRoleByARN("arn:aws:iam::1234567889:role/unknown")
|
||||
require.False(t, found)
|
||||
})
|
||||
})
|
||||
|
||||
t.Run("FindRolesByName", func(t *testing.T) {
|
||||
t.Run("found zero", func(t *testing.T) {
|
||||
rolesWithName := roles.FindRolesByName("unknown")
|
||||
require.Empty(t, rolesWithName)
|
||||
})
|
||||
|
||||
t.Run("found one", func(t *testing.T) {
|
||||
rolesWithName := roles.FindRolesByName("path/to/EC2FullAccess")
|
||||
require.Len(t, rolesWithName, 1)
|
||||
require.Equal(t, "path/to/EC2FullAccess", rolesWithName[0].Name)
|
||||
})
|
||||
|
||||
t.Run("found two", func(t *testing.T) {
|
||||
rolesWithName := roles.FindRolesByName("EC2FullAccess")
|
||||
require.Len(t, rolesWithName, 2)
|
||||
require.Equal(t, "EC2FullAccess", rolesWithName[0].Display)
|
||||
require.Equal(t, "EC2FullAccess", rolesWithName[1].Display)
|
||||
require.NotEqual(t, rolesWithName[0].ARN, rolesWithName[1].ARN)
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
+1
-1
@@ -61,7 +61,7 @@ func onAppLogin(cf *CLIConf) error {
|
||||
var arn string
|
||||
if app.IsAWSConsole() {
|
||||
var err error
|
||||
arn, err = getARNFromFlags(cf, profile)
|
||||
arn, err = getARNFromFlags(cf, profile, app)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
|
||||
+42
-51
@@ -24,7 +24,6 @@ import (
|
||||
"net"
|
||||
"os"
|
||||
"os/exec"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
@@ -32,12 +31,14 @@ import (
|
||||
"github.com/aws/aws-sdk-go/aws/credentials"
|
||||
"github.com/gravitational/trace"
|
||||
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/teleport/lib/asciitable"
|
||||
"github.com/gravitational/teleport/lib/client"
|
||||
"github.com/gravitational/teleport/lib/srv/alpnproxy"
|
||||
alpncommon "github.com/gravitational/teleport/lib/srv/alpnproxy/common"
|
||||
"github.com/gravitational/teleport/lib/tlsca"
|
||||
"github.com/gravitational/teleport/lib/utils"
|
||||
awsutils "github.com/gravitational/teleport/lib/utils/aws"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -363,69 +364,59 @@ func (a *awsApp) startLocalForwardProxy(port string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func printArrayAs(arr []string, columnName string) {
|
||||
sort.Strings(arr)
|
||||
if len(arr) == 0 {
|
||||
func printAWSRoles(roles awsutils.Roles) {
|
||||
if len(roles) == 0 {
|
||||
return
|
||||
}
|
||||
t := asciitable.MakeTable([]string{columnName})
|
||||
for _, v := range arr {
|
||||
t.AddRow([]string{v})
|
||||
|
||||
roles.Sort()
|
||||
|
||||
t := asciitable.MakeTable([]string{"Role Name", "Role ARN"})
|
||||
for _, role := range roles {
|
||||
// Use role.Display for role names to match what AWS web console shows.
|
||||
t.AddRow([]string{role.Display, role.ARN})
|
||||
}
|
||||
fmt.Println(t.AsBuffer().String())
|
||||
}
|
||||
|
||||
func getARNFromFlags(cf *CLIConf, profile *client.ProfileStatus) (string, error) {
|
||||
func getARNFromFlags(cf *CLIConf, profile *client.ProfileStatus, app types.Application) (string, error) {
|
||||
// Filter AWS roles by AWS account ID. If AWS account ID is empty, all
|
||||
// roles are returned.
|
||||
roles := awsutils.FilterAWSRoles(profile.AWSRolesARNs, app.GetAWSAccountID())
|
||||
|
||||
if cf.AWSRole == "" {
|
||||
printArrayAs(profile.AWSRolesARNs, "Available Role ARNs")
|
||||
if len(roles) == 1 {
|
||||
log.Infof("AWS Role %v is selected by default as it is the only role configured for this AWS app.", roles[0].Display)
|
||||
return roles[0].ARN, nil
|
||||
}
|
||||
|
||||
printAWSRoles(roles)
|
||||
return "", trace.BadParameter("--aws-role flag is required")
|
||||
}
|
||||
for _, v := range profile.AWSRolesARNs {
|
||||
if v == cf.AWSRole {
|
||||
return v, nil
|
||||
|
||||
// Match by role ARN.
|
||||
if awsarn.IsARN(cf.AWSRole) {
|
||||
if role, found := roles.FindRoleByARN(cf.AWSRole); found {
|
||||
return role.ARN, nil
|
||||
}
|
||||
|
||||
printAWSRoles(roles)
|
||||
return "", trace.NotFound("failed to find the %q role ARN", cf.AWSRole)
|
||||
}
|
||||
|
||||
roleNameToARN := make(map[string]string)
|
||||
for _, v := range profile.AWSRolesARNs {
|
||||
arn, err := awsarn.Parse(v)
|
||||
if err != nil {
|
||||
return "", trace.Wrap(err)
|
||||
}
|
||||
// Example of the ANR Resource: 'role/EC2FullAccess' or 'role/path/to/customrole'
|
||||
parts := strings.Split(arn.Resource, "/")
|
||||
if len(parts) < 1 || parts[0] != "role" {
|
||||
continue
|
||||
}
|
||||
roleName := strings.Join(parts[1:], "/")
|
||||
|
||||
if val, ok := roleNameToARN[roleName]; ok && cf.AWSRole == roleName {
|
||||
return "", trace.BadParameter(
|
||||
"provided role name %q is ambiguous between %q and %q ARNs, please specify full role ARN",
|
||||
cf.AWSRole, val, arn.String())
|
||||
}
|
||||
roleNameToARN[roleName] = arn.String()
|
||||
// Match by role name.
|
||||
rolesMatched := roles.FindRolesByName(cf.AWSRole)
|
||||
switch len(rolesMatched) {
|
||||
case 1:
|
||||
return rolesMatched[0].ARN, nil
|
||||
case 0:
|
||||
printAWSRoles(roles)
|
||||
return "", trace.NotFound("failed to find the %q role name", cf.AWSRole)
|
||||
default:
|
||||
// Print roles matched the provided role name.
|
||||
printAWSRoles(rolesMatched)
|
||||
return "", trace.BadParameter("provided role name %q is ambiguous, please specify full role ARN", cf.AWSRole)
|
||||
}
|
||||
|
||||
roleARN, ok := roleNameToARN[cf.AWSRole]
|
||||
if !ok {
|
||||
printArrayAs(profile.AWSRolesARNs, "Available Role ARNs")
|
||||
printArrayAs(mapKeysToSlice(roleNameToARN), "Available Role Names")
|
||||
inputType := "ARN"
|
||||
if !awsarn.IsARN(cf.AWSRole) {
|
||||
inputType = "name"
|
||||
}
|
||||
return "", trace.NotFound("failed to find the %q role %s", cf.AWSRole, inputType)
|
||||
}
|
||||
return roleARN, nil
|
||||
}
|
||||
|
||||
func mapKeysToSlice(m map[string]string) []string {
|
||||
out := make([]string, 0, len(m))
|
||||
for k := range m {
|
||||
out = append(out, k)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func pickActiveAWSApp(cf *CLIConf) (*awsApp, error) {
|
||||
|
||||
Reference in New Issue
Block a user