mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
IAM Join Method (backend implementation) (#10085)
This commit is contained in:
@@ -1,5 +1,5 @@
|
||||
/*
|
||||
Copyright 2020 Gravitational, Inc.
|
||||
Copyright 2020-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.
|
||||
@@ -25,6 +25,20 @@ import (
|
||||
"github.com/gravitational/trace"
|
||||
)
|
||||
|
||||
// JoinMethod is the method used for new nodes to join the cluster.
|
||||
type JoinMethod string
|
||||
|
||||
const (
|
||||
JoinMethodUnspecified JoinMethod = ""
|
||||
// JoinMethodToken is the default join method, nodes join the cluster by
|
||||
// presenting a secret token.
|
||||
JoinMethodToken JoinMethod = "token"
|
||||
// JoinMethodEC2 indicates that the node will join with the EC2 join method.
|
||||
JoinMethodEC2 JoinMethod = "ec2"
|
||||
// JoinMethodIAM indicates that the node will join with the IAM join method.
|
||||
JoinMethodIAM JoinMethod = "iam"
|
||||
)
|
||||
|
||||
// ProvisionToken is a provisioning token
|
||||
type ProvisionToken interface {
|
||||
Resource
|
||||
@@ -40,6 +54,8 @@ type ProvisionToken interface {
|
||||
GetAllowRules() []*TokenRule
|
||||
// GetAWSIIDTTL returns the TTL of EC2 IIDs
|
||||
GetAWSIIDTTL() Duration
|
||||
// GetJoinMethod returns joining method that must be used with this token.
|
||||
GetJoinMethod() JoinMethod
|
||||
// V1 returns V1 version of the resource
|
||||
V1() *ProvisionTokenV1
|
||||
// String returns user friendly representation of the resource
|
||||
@@ -98,8 +114,55 @@ func (p *ProvisionTokenV2) CheckAndSetDefaults() error {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
|
||||
if p.Spec.AWSIIDTTL == 0 {
|
||||
p.Spec.AWSIIDTTL = Duration(5 * time.Minute)
|
||||
hasAllowRules := len(p.Spec.Allow) > 0
|
||||
if p.Spec.JoinMethod == JoinMethodUnspecified {
|
||||
// Default to the ec2 join method if any allow rules were specified,
|
||||
// else default to the token method. These defaults are necessary for
|
||||
// backwards compatibility.
|
||||
if hasAllowRules {
|
||||
p.Spec.JoinMethod = JoinMethodEC2
|
||||
} else {
|
||||
p.Spec.JoinMethod = JoinMethodToken
|
||||
}
|
||||
}
|
||||
switch p.Spec.JoinMethod {
|
||||
case JoinMethodToken:
|
||||
if hasAllowRules {
|
||||
return trace.BadParameter("allow rules are not compatible with the %q join method", JoinMethodToken)
|
||||
}
|
||||
case JoinMethodEC2:
|
||||
if !hasAllowRules {
|
||||
return trace.BadParameter("the %q join method requires defined token allow rules", JoinMethodEC2)
|
||||
}
|
||||
for _, allowRule := range p.Spec.Allow {
|
||||
if allowRule.AWSARN != "" {
|
||||
return trace.BadParameter(`the %q join method does not support the "aws_arn" parameter`, JoinMethodEC2)
|
||||
}
|
||||
if allowRule.AWSAccount == "" && allowRule.AWSRole == "" {
|
||||
return trace.BadParameter(`allow rule for %q join method must set "aws_account" or "aws_role"`, JoinMethodEC2)
|
||||
}
|
||||
}
|
||||
if p.Spec.AWSIIDTTL == 0 {
|
||||
// default to 5 minute ttl if unspecified
|
||||
p.Spec.AWSIIDTTL = Duration(5 * time.Minute)
|
||||
}
|
||||
case JoinMethodIAM:
|
||||
if !hasAllowRules {
|
||||
return trace.BadParameter("the %q join method requires defined token allow rules", JoinMethodIAM)
|
||||
}
|
||||
for _, allowRule := range p.Spec.Allow {
|
||||
if allowRule.AWSRole != "" {
|
||||
return trace.BadParameter(`the %q join method does not support the "aws_role" parameter`, JoinMethodIAM)
|
||||
}
|
||||
if len(allowRule.AWSRegions) != 0 {
|
||||
return trace.BadParameter(`the %q join method does not support the "aws_regions" parameter`, JoinMethodIAM)
|
||||
}
|
||||
if allowRule.AWSAccount == "" && allowRule.AWSARN == "" {
|
||||
return trace.BadParameter(`allow rule for %q join method must set "aws_account" or "aws_arn"`, JoinMethodEC2)
|
||||
}
|
||||
}
|
||||
default:
|
||||
return trace.BadParameter("unknown join method %q", p.Spec.JoinMethod)
|
||||
}
|
||||
|
||||
return nil
|
||||
@@ -132,6 +195,11 @@ func (p *ProvisionTokenV2) GetAWSIIDTTL() Duration {
|
||||
return p.Spec.AWSIIDTTL
|
||||
}
|
||||
|
||||
// GetJoinMethod returns joining method that must be used with this token.
|
||||
func (p *ProvisionTokenV2) GetJoinMethod() JoinMethod {
|
||||
return p.Spec.JoinMethod
|
||||
}
|
||||
|
||||
// GetKind returns resource kind
|
||||
func (p *ProvisionTokenV2) GetKind() string {
|
||||
return p.Kind
|
||||
|
||||
@@ -0,0 +1,272 @@
|
||||
/*
|
||||
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 types
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gravitational/trace"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestProvisionTokenV2_CheckAndSetDefaults(t *testing.T) {
|
||||
testcases := []struct {
|
||||
desc string
|
||||
token *ProvisionTokenV2
|
||||
expected *ProvisionTokenV2
|
||||
expectedErr error
|
||||
}{
|
||||
{
|
||||
desc: "empty",
|
||||
token: &ProvisionTokenV2{},
|
||||
expectedErr: &trace.BadParameterError{},
|
||||
},
|
||||
{
|
||||
desc: "missing roles",
|
||||
token: &ProvisionTokenV2{
|
||||
Metadata: Metadata{
|
||||
Name: "test",
|
||||
},
|
||||
},
|
||||
expectedErr: &trace.BadParameterError{},
|
||||
},
|
||||
{
|
||||
desc: "invalid role",
|
||||
token: &ProvisionTokenV2{
|
||||
Metadata: Metadata{
|
||||
Name: "test",
|
||||
},
|
||||
Spec: ProvisionTokenSpecV2{
|
||||
Roles: []SystemRole{RoleNode, "not a role"},
|
||||
},
|
||||
},
|
||||
expectedErr: &trace.BadParameterError{},
|
||||
},
|
||||
{
|
||||
desc: "simple token",
|
||||
token: &ProvisionTokenV2{
|
||||
Metadata: Metadata{
|
||||
Name: "test",
|
||||
},
|
||||
Spec: ProvisionTokenSpecV2{
|
||||
Roles: []SystemRole{RoleNode},
|
||||
},
|
||||
},
|
||||
expected: &ProvisionTokenV2{
|
||||
Kind: "token",
|
||||
Version: "v2",
|
||||
Metadata: Metadata{
|
||||
Name: "test",
|
||||
Namespace: "default",
|
||||
},
|
||||
Spec: ProvisionTokenSpecV2{
|
||||
Roles: []SystemRole{RoleNode},
|
||||
JoinMethod: "token",
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
desc: "implicit ec2 method",
|
||||
token: &ProvisionTokenV2{
|
||||
Metadata: Metadata{
|
||||
Name: "test",
|
||||
},
|
||||
Spec: ProvisionTokenSpecV2{
|
||||
Roles: []SystemRole{RoleNode},
|
||||
Allow: []*TokenRule{
|
||||
&TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSRole: "1234/role",
|
||||
AWSRegions: []string{"us-west-2"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
expected: &ProvisionTokenV2{
|
||||
Kind: "token",
|
||||
Version: "v2",
|
||||
Metadata: Metadata{
|
||||
Name: "test",
|
||||
Namespace: "default",
|
||||
},
|
||||
Spec: ProvisionTokenSpecV2{
|
||||
Roles: []SystemRole{RoleNode},
|
||||
JoinMethod: "ec2",
|
||||
Allow: []*TokenRule{
|
||||
&TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSRole: "1234/role",
|
||||
AWSRegions: []string{"us-west-2"},
|
||||
},
|
||||
},
|
||||
AWSIIDTTL: Duration(5 * time.Minute),
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
desc: "explicit ec2 method",
|
||||
token: &ProvisionTokenV2{
|
||||
Metadata: Metadata{
|
||||
Name: "test",
|
||||
},
|
||||
Spec: ProvisionTokenSpecV2{
|
||||
Roles: []SystemRole{RoleNode},
|
||||
JoinMethod: "ec2",
|
||||
Allow: []*TokenRule{&TokenRule{AWSAccount: "1234"}},
|
||||
},
|
||||
},
|
||||
expected: &ProvisionTokenV2{
|
||||
Kind: "token",
|
||||
Version: "v2",
|
||||
Metadata: Metadata{
|
||||
Name: "test",
|
||||
Namespace: "default",
|
||||
},
|
||||
Spec: ProvisionTokenSpecV2{
|
||||
Roles: []SystemRole{RoleNode},
|
||||
JoinMethod: "ec2",
|
||||
Allow: []*TokenRule{&TokenRule{AWSAccount: "1234"}},
|
||||
AWSIIDTTL: Duration(5 * time.Minute),
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
desc: "ec2 method no allow rules",
|
||||
token: &ProvisionTokenV2{
|
||||
Metadata: Metadata{
|
||||
Name: "test",
|
||||
},
|
||||
Spec: ProvisionTokenSpecV2{
|
||||
Roles: []SystemRole{RoleNode},
|
||||
JoinMethod: "ec2",
|
||||
},
|
||||
},
|
||||
expectedErr: &trace.BadParameterError{},
|
||||
},
|
||||
{
|
||||
desc: "ec2 method with aws_arn",
|
||||
token: &ProvisionTokenV2{
|
||||
Metadata: Metadata{
|
||||
Name: "test",
|
||||
},
|
||||
Spec: ProvisionTokenSpecV2{
|
||||
Roles: []SystemRole{RoleNode},
|
||||
JoinMethod: "ec2",
|
||||
Allow: []*TokenRule{
|
||||
&TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSARN: "1234",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
expectedErr: &trace.BadParameterError{},
|
||||
},
|
||||
{
|
||||
desc: "ec2 method empty rule",
|
||||
token: &ProvisionTokenV2{
|
||||
Metadata: Metadata{
|
||||
Name: "test",
|
||||
},
|
||||
Spec: ProvisionTokenSpecV2{
|
||||
Roles: []SystemRole{RoleNode},
|
||||
JoinMethod: "ec2",
|
||||
Allow: []*TokenRule{&TokenRule{}},
|
||||
},
|
||||
},
|
||||
expectedErr: &trace.BadParameterError{},
|
||||
},
|
||||
{
|
||||
desc: "iam method",
|
||||
token: &ProvisionTokenV2{
|
||||
Metadata: Metadata{
|
||||
Name: "test",
|
||||
},
|
||||
Spec: ProvisionTokenSpecV2{
|
||||
Roles: []SystemRole{RoleNode},
|
||||
JoinMethod: "ec2",
|
||||
Allow: []*TokenRule{&TokenRule{AWSAccount: "1234"}},
|
||||
},
|
||||
},
|
||||
expected: &ProvisionTokenV2{
|
||||
Kind: "token",
|
||||
Version: "v2",
|
||||
Metadata: Metadata{
|
||||
Name: "test",
|
||||
Namespace: "default",
|
||||
},
|
||||
Spec: ProvisionTokenSpecV2{
|
||||
Roles: []SystemRole{RoleNode},
|
||||
JoinMethod: "ec2",
|
||||
Allow: []*TokenRule{&TokenRule{AWSAccount: "1234"}},
|
||||
AWSIIDTTL: Duration(5 * time.Minute),
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
desc: "iam method with aws_role",
|
||||
token: &ProvisionTokenV2{
|
||||
Metadata: Metadata{
|
||||
Name: "test",
|
||||
},
|
||||
Spec: ProvisionTokenSpecV2{
|
||||
Roles: []SystemRole{RoleNode},
|
||||
JoinMethod: "iam",
|
||||
Allow: []*TokenRule{
|
||||
&TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSRole: "1234/role",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
expectedErr: &trace.BadParameterError{},
|
||||
},
|
||||
{
|
||||
desc: "iam method with aws_regions",
|
||||
token: &ProvisionTokenV2{
|
||||
Metadata: Metadata{
|
||||
Name: "test",
|
||||
},
|
||||
Spec: ProvisionTokenSpecV2{
|
||||
Roles: []SystemRole{RoleNode},
|
||||
JoinMethod: "iam",
|
||||
Allow: []*TokenRule{
|
||||
&TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSRegions: []string{"us-west-2"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
expectedErr: &trace.BadParameterError{},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range testcases {
|
||||
t.Run(tc.desc, func(t *testing.T) {
|
||||
err := tc.token.CheckAndSetDefaults()
|
||||
if tc.expectedErr != nil {
|
||||
require.ErrorAs(t, err, &tc.expectedErr)
|
||||
return
|
||||
}
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tc.token, tc.expected)
|
||||
})
|
||||
}
|
||||
}
|
||||
+837
-682
File diff suppressed because it is too large
Load Diff
+24
-3
@@ -703,10 +703,20 @@ message ProvisionTokenV2List {
|
||||
repeated ProvisionTokenV2 ProvisionTokens = 1;
|
||||
}
|
||||
|
||||
// TokenRule is a rule that a joining node must match in order to use the
|
||||
// associated token.
|
||||
message TokenRule {
|
||||
// AWSAccount is the AWS account ID.
|
||||
string AWSAccount = 1 [ (gogoproto.jsontag) = "aws_account,omitempty" ];
|
||||
// AWSRegions is used for the EC2 join method and is a list of AWS regions a
|
||||
// node is allowed to join from.
|
||||
repeated string AWSRegions = 2 [ (gogoproto.jsontag) = "aws_regions,omitempty" ];
|
||||
// AWSRole is used for the EC2 join method and is the the ARN of the AWS
|
||||
// role that the auth server will assume in order to call the ec2 API.
|
||||
string AWSRole = 3 [ (gogoproto.jsontag) = "aws_role,omitempty" ];
|
||||
// AWSARN is used for the IAM join method, the AWS identity of joining nodes
|
||||
// must match this ARN. Supports wildcards "*" and "?".
|
||||
string AWSARN = 4 [ (gogoproto.jsontag) = "aws_arn,omitempty" ];
|
||||
}
|
||||
|
||||
// ProvisionTokenSpecV2 is a specification for V2 token
|
||||
@@ -716,9 +726,17 @@ message ProvisionTokenSpecV2 {
|
||||
// certificates issued to the user of the token
|
||||
repeated string Roles = 1
|
||||
[ (gogoproto.jsontag) = "roles", (gogoproto.casttype) = "SystemRole" ];
|
||||
repeated TokenRule allow = 2 [ (gogoproto.jsontag) = "allow,omitempty" ];
|
||||
// Allow is a list of TokenRules, nodes using this token must match one
|
||||
// allow rule to use this token.
|
||||
repeated TokenRule Allow = 2 [ (gogoproto.jsontag) = "allow,omitempty" ];
|
||||
// AWSIIDTTL is the TTL to use for AWS EC2 Instance Identity Documents used
|
||||
// to join the cluster with this token.
|
||||
int64 AWSIIDTTL = 3
|
||||
[ (gogoproto.jsontag) = "aws_iid_ttl,omitempty", (gogoproto.casttype) = "Duration" ];
|
||||
// JoinMethod is the joining method required in order to use this token.
|
||||
// Supported joining methods include "token", "ec2", and "iam".
|
||||
string JoinMethod = 4
|
||||
[ (gogoproto.jsontag) = "join_method", (gogoproto.casttype) = "JoinMethod" ];
|
||||
}
|
||||
|
||||
// StaticTokensV2 implements the StaticTokens interface.
|
||||
@@ -2712,9 +2730,12 @@ message RegisterUsingTokenRequest {
|
||||
// RemoteAddr is the remote address of the host requesting a host certificate.
|
||||
// It is used to replace 0.0.0.0 in the list of additional principals.
|
||||
string RemoteAddr = 9 [ (gogoproto.jsontag) = "remote_addr" ];
|
||||
// EC2IdentityDocument is used for Simplified Node Joining to prove the
|
||||
// identity of a joining EC2 instance.
|
||||
// EC2IdentityDocument is used for the EC2 join method to prove the identity
|
||||
// of a joining EC2 instance.
|
||||
bytes EC2IdentityDocument = 10 [ (gogoproto.jsontag) = "ec2_id" ];
|
||||
// STSIdentityRequest is used for the IAM join method to prove the AWS
|
||||
// identity of a joining node.
|
||||
bytes STSIdentityRequest = 11 [ (gogoproto.jsontag) = "-" ];
|
||||
}
|
||||
|
||||
// RecoveryCodes holds a user's recovery code information. Recovery codes allows users to regain
|
||||
|
||||
@@ -40,7 +40,7 @@ import (
|
||||
func newNodeConfig(t *testing.T, authAddr utils.NetAddr, awsTokenName string) *service.Config {
|
||||
config := service.MakeDefaultConfig()
|
||||
config.Token = awsTokenName
|
||||
config.JoinMethod = service.JoinMethodEC2
|
||||
config.JoinMethod = types.JoinMethodEC2
|
||||
config.SSH.Enabled = true
|
||||
config.SSH.Addr.Addr = net.JoinHostPort(Host, ports.Pop())
|
||||
config.Auth.Enabled = false
|
||||
|
||||
@@ -633,7 +633,7 @@ func (s *APIServer) validateTrustedCluster(auth ClientI, w http.ResponseWriter,
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
validateResponse, err := auth.ValidateTrustedCluster(validateRequest)
|
||||
validateResponse, err := auth.ValidateTrustedCluster(r.Context(), validateRequest)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -1033,7 +1033,7 @@ func (s *APIServer) registerUsingToken(auth ClientI, w http.ResponseWriter, r *h
|
||||
// Pass along the remote address the request came from to the registration function.
|
||||
req.RemoteAddr = r.RemoteAddr
|
||||
|
||||
certs, err := auth.RegisterUsingToken(req)
|
||||
certs, err := auth.RegisterUsingToken(r.Context(), &req)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
+3
-72
@@ -2262,9 +2262,8 @@ func (a *Server) GenerateHostCerts(ctx context.Context, req *proto.HostCertsRequ
|
||||
// ValidateToken takes a provisioning token value and finds if it's valid. Returns
|
||||
// a list of roles this token allows its owner to assume and token labels, or an error if the token
|
||||
// cannot be found.
|
||||
func (a *Server) ValidateToken(token string) (types.SystemRoles, map[string]string, error) {
|
||||
ctx := context.TODO()
|
||||
tkns, err := a.GetCache().GetStaticTokens()
|
||||
func (a *Server) ValidateToken(ctx context.Context, token string) (types.SystemRoles, map[string]string, error) {
|
||||
tkns, err := a.GetStaticTokens()
|
||||
if err != nil {
|
||||
return nil, nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -2279,7 +2278,7 @@ func (a *Server) ValidateToken(token string) (types.SystemRoles, map[string]stri
|
||||
|
||||
// If it's not a static token, check if it's a ephemeral token in the backend.
|
||||
// If a ephemeral token is found, make sure it's still valid.
|
||||
tok, err := a.GetCache().GetToken(ctx, token)
|
||||
tok, err := a.GetToken(ctx, token)
|
||||
if err != nil {
|
||||
return nil, nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -2307,74 +2306,6 @@ func (a *Server) checkTokenTTL(tok types.ProvisionToken) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// RegisterUsingToken adds a new node to the Teleport cluster using previously issued token.
|
||||
// A node must also request a specific role (and the role must match one of the roles
|
||||
// the token was generated for).
|
||||
//
|
||||
// If a token was generated with a TTL, it gets enforced (can't register new nodes after TTL expires)
|
||||
// If a token was generated with a TTL=0, it means it's a single-use token and it gets destroyed
|
||||
// after a successful registration.
|
||||
func (a *Server) RegisterUsingToken(req types.RegisterUsingTokenRequest) (*proto.Certs, error) {
|
||||
log.Infof("Node %q [%v] is trying to join with role: %v.", req.NodeName, req.HostID, req.Role)
|
||||
|
||||
if err := req.CheckAndSetDefaults(); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// If the request uses Simplified Node Joining check that the identity is
|
||||
// valid and matches the token allow rules.
|
||||
err := a.CheckEC2Request(context.Background(), req)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// make sure the token is valid
|
||||
roles, _, err := a.ValidateToken(req.Token)
|
||||
if err != nil {
|
||||
log.Warningf("%q [%v] can not join the cluster with role %s, token error: %v", req.NodeName, req.HostID, req.Role, err)
|
||||
return nil, trace.AccessDenied(fmt.Sprintf("%q [%v] can not join the cluster with role %s, the token is not valid", req.NodeName, req.HostID, req.Role))
|
||||
}
|
||||
|
||||
// make sure the caller is requested the role allowed by the token
|
||||
if !roles.Include(req.Role) {
|
||||
msg := fmt.Sprintf("node %q [%v] can not join the cluster, the token does not allow %q role", req.NodeName, req.HostID, req.Role)
|
||||
log.Warn(msg)
|
||||
return nil, trace.BadParameter(msg)
|
||||
}
|
||||
|
||||
// generate and return host certificate and keys
|
||||
certs, err := a.GenerateHostCerts(context.Background(),
|
||||
&proto.HostCertsRequest{
|
||||
HostID: req.HostID,
|
||||
NodeName: req.NodeName,
|
||||
Role: req.Role,
|
||||
AdditionalPrincipals: req.AdditionalPrincipals,
|
||||
PublicTLSKey: req.PublicTLSKey,
|
||||
PublicSSHKey: req.PublicSSHKey,
|
||||
RemoteAddr: req.RemoteAddr,
|
||||
DNSNames: req.DNSNames,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
log.Infof("Node %q [%v] has joined the cluster.", req.NodeName, req.HostID)
|
||||
return certs, nil
|
||||
}
|
||||
|
||||
func (a *Server) RegisterNewAuthServer(ctx context.Context, token string) error {
|
||||
tok, err := a.Provisioner.GetToken(ctx, token)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
if !tok.GetRoles().Include(types.RoleAuth) {
|
||||
return trace.AccessDenied("role does not match")
|
||||
}
|
||||
if err := a.DeleteToken(ctx, token); err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *Server) DeleteToken(ctx context.Context, token string) (err error) {
|
||||
tkns, err := a.GetStaticTokens()
|
||||
if err != nil {
|
||||
|
||||
+13
-95
@@ -1,5 +1,5 @@
|
||||
/*
|
||||
Copyright 2015-2019 Gravitational, Inc.
|
||||
Copyright 2015-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.
|
||||
@@ -39,7 +39,6 @@ import (
|
||||
apidefaults "github.com/gravitational/teleport/api/defaults"
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
apievents "github.com/gravitational/teleport/api/types/events"
|
||||
apiutils "github.com/gravitational/teleport/api/utils"
|
||||
"github.com/gravitational/teleport/api/utils/sshutils"
|
||||
"github.com/gravitational/teleport/lib/auth/testauthority"
|
||||
authority "github.com/gravitational/teleport/lib/auth/testauthority"
|
||||
@@ -155,6 +154,10 @@ func newTestPack(ctx context.Context, dataDir string) (testPack, error) {
|
||||
return p, trace.Wrap(err)
|
||||
}
|
||||
|
||||
if err := p.a.UpsertNamespace(types.DefaultNamespace()); err != nil {
|
||||
return p, trace.Wrap(err)
|
||||
}
|
||||
|
||||
p.mockEmitter = &events.MockEmitter{}
|
||||
p.a.emitter = p.mockEmitter
|
||||
return p, nil
|
||||
@@ -564,37 +567,18 @@ func (s *AuthSuite) TestTokensCRUD(c *C) {
|
||||
c.Assert(len(tokens), Equals, 1)
|
||||
c.Assert(tokens[0].GetName(), Equals, tok)
|
||||
|
||||
roles, _, err := s.a.ValidateToken(tok)
|
||||
roles, _, err := s.a.ValidateToken(ctx, tok)
|
||||
c.Assert(err, IsNil)
|
||||
c.Assert(roles.Include(types.RoleNode), Equals, true)
|
||||
c.Assert(roles.Include(types.RoleProxy), Equals, false)
|
||||
|
||||
priv, pub, err := s.a.GenerateKeyPair("")
|
||||
c.Assert(err, IsNil)
|
||||
|
||||
tlsPublicKey, err := PrivateKeyToPublicKeyTLS(priv)
|
||||
c.Assert(err, IsNil)
|
||||
|
||||
// unsuccessful registration (wrong role)
|
||||
certs, err := s.a.RegisterUsingToken(types.RegisterUsingTokenRequest{
|
||||
Token: tok,
|
||||
HostID: "bad-host-id",
|
||||
NodeName: "bad-node-name",
|
||||
Role: types.RoleProxy,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
PublicSSHKey: pub,
|
||||
})
|
||||
c.Assert(certs, IsNil)
|
||||
c.Assert(err, NotNil)
|
||||
c.Assert(err, ErrorMatches, `node "bad-node-name" \[bad-host-id\] can not join the cluster, the token does not allow "Proxy" role`)
|
||||
|
||||
// generate predefined token
|
||||
customToken := "custom-token"
|
||||
tok, err = s.a.GenerateToken(ctx, GenerateTokenRequest{Roles: types.SystemRoles{types.RoleNode}, Token: customToken})
|
||||
c.Assert(err, IsNil)
|
||||
c.Assert(tok, Equals, customToken)
|
||||
|
||||
roles, _, err = s.a.ValidateToken(tok)
|
||||
roles, _, err = s.a.ValidateToken(ctx, tok)
|
||||
c.Assert(err, IsNil)
|
||||
c.Assert(roles.Include(types.RoleNode), Equals, true)
|
||||
c.Assert(roles.Include(types.RoleProxy), Equals, false)
|
||||
@@ -602,56 +586,6 @@ func (s *AuthSuite) TestTokensCRUD(c *C) {
|
||||
err = s.a.DeleteToken(ctx, customToken)
|
||||
c.Assert(err, IsNil)
|
||||
|
||||
// generate multi-use token with long TTL:
|
||||
multiUseToken, err := s.a.GenerateToken(ctx, GenerateTokenRequest{Roles: types.SystemRoles{types.RoleProxy}, TTL: time.Hour})
|
||||
c.Assert(err, IsNil)
|
||||
_, _, err = s.a.ValidateToken(multiUseToken)
|
||||
c.Assert(err, IsNil)
|
||||
|
||||
// use it twice:
|
||||
certs, err = s.a.RegisterUsingToken(types.RegisterUsingTokenRequest{
|
||||
Token: multiUseToken,
|
||||
HostID: "once",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleProxy,
|
||||
AdditionalPrincipals: []string{"example.com"},
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
PublicSSHKey: pub,
|
||||
})
|
||||
c.Assert(err, IsNil)
|
||||
|
||||
// along the way, make sure that additional principals work
|
||||
hostCert, err := sshutils.ParseCertificate(certs.SSH)
|
||||
c.Assert(err, IsNil)
|
||||
comment := Commentf("can't find example.com in %v", hostCert.ValidPrincipals)
|
||||
c.Assert(apiutils.SliceContainsStr(hostCert.ValidPrincipals, "example.com"), Equals, true, comment)
|
||||
|
||||
_, err = s.a.RegisterUsingToken(types.RegisterUsingTokenRequest{
|
||||
Token: multiUseToken,
|
||||
HostID: "twice",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleProxy,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
PublicSSHKey: pub,
|
||||
})
|
||||
c.Assert(err, IsNil)
|
||||
|
||||
// try to use after TTL:
|
||||
s.a.SetClock(clockwork.NewFakeClockAt(time.Now().UTC().Add(time.Hour + 1)))
|
||||
_, err = s.a.RegisterUsingToken(types.RegisterUsingTokenRequest{
|
||||
Token: multiUseToken,
|
||||
HostID: "late.bird",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleProxy,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
PublicSSHKey: pub,
|
||||
})
|
||||
c.Assert(err, ErrorMatches, `"node-name" \[late.bird\] can not join the cluster with role Proxy, the token is not valid`)
|
||||
|
||||
// expired token should be gone now
|
||||
err = s.a.DeleteToken(ctx, multiUseToken)
|
||||
c.Assert(trace.IsNotFound(err), Equals, true, Commentf("%#v", err))
|
||||
|
||||
// lets use static tokens now
|
||||
roles = types.SystemRoles{types.RoleProxy}
|
||||
st, err := types.NewStaticTokens(types.StaticTokensSpecV2{
|
||||
@@ -662,27 +596,11 @@ func (s *AuthSuite) TestTokensCRUD(c *C) {
|
||||
}},
|
||||
})
|
||||
c.Assert(err, IsNil)
|
||||
|
||||
err = s.a.SetStaticTokens(st)
|
||||
c.Assert(err, IsNil)
|
||||
_, err = s.a.RegisterUsingToken(types.RegisterUsingTokenRequest{
|
||||
Token: "static-token-value",
|
||||
HostID: "static.host",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleProxy,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
PublicSSHKey: pub,
|
||||
})
|
||||
c.Assert(err, IsNil)
|
||||
_, err = s.a.RegisterUsingToken(types.RegisterUsingTokenRequest{
|
||||
Token: "static-token-value",
|
||||
HostID: "wrong.role",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleAuth,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
PublicSSHKey: pub,
|
||||
})
|
||||
c.Assert(err, NotNil)
|
||||
r, _, err := s.a.ValidateToken("static-token-value")
|
||||
|
||||
r, _, err := s.a.ValidateToken(ctx, "static-token-value")
|
||||
c.Assert(err, IsNil)
|
||||
c.Assert(r, DeepEquals, roles)
|
||||
|
||||
@@ -695,11 +613,11 @@ func (s *AuthSuite) TestTokensCRUD(c *C) {
|
||||
func (s *AuthSuite) TestBadTokens(c *C) {
|
||||
ctx := context.Background()
|
||||
// empty
|
||||
_, _, err := s.a.ValidateToken("")
|
||||
_, _, err := s.a.ValidateToken(ctx, "")
|
||||
c.Assert(err, NotNil)
|
||||
|
||||
// garbage
|
||||
_, _, err = s.a.ValidateToken("bla bla")
|
||||
_, _, err = s.a.ValidateToken(ctx, "bla bla")
|
||||
c.Assert(err, NotNil)
|
||||
|
||||
// tampered
|
||||
@@ -707,7 +625,7 @@ func (s *AuthSuite) TestBadTokens(c *C) {
|
||||
c.Assert(err, IsNil)
|
||||
|
||||
tampered := string(tok[0]+1) + tok[1:]
|
||||
_, _, err = s.a.ValidateToken(tampered)
|
||||
_, _, err = s.a.ValidateToken(ctx, tampered)
|
||||
c.Assert(err, NotNil)
|
||||
}
|
||||
|
||||
|
||||
@@ -456,9 +456,9 @@ func (a *ServerWithRoles) GenerateToken(ctx context.Context, req GenerateTokenRe
|
||||
return a.authServer.GenerateToken(ctx, req)
|
||||
}
|
||||
|
||||
func (a *ServerWithRoles) RegisterUsingToken(req types.RegisterUsingTokenRequest) (*proto.Certs, error) {
|
||||
func (a *ServerWithRoles) RegisterUsingToken(ctx context.Context, req *types.RegisterUsingTokenRequest) (*proto.Certs, error) {
|
||||
// tokens have authz mechanism on their own, no need to check
|
||||
return a.authServer.RegisterUsingToken(req)
|
||||
return a.authServer.RegisterUsingToken(ctx, req)
|
||||
}
|
||||
|
||||
func (a *ServerWithRoles) RegisterNewAuthServer(ctx context.Context, token string) error {
|
||||
@@ -466,6 +466,18 @@ func (a *ServerWithRoles) RegisterNewAuthServer(ctx context.Context, token strin
|
||||
return a.authServer.RegisterNewAuthServer(ctx, token)
|
||||
}
|
||||
|
||||
// RegisterUsingIAMMethod registers the caller using the IAM join method and
|
||||
// returns signed certs to join the cluster.
|
||||
//
|
||||
// See (*Server).RegisterUsingIAMMethod for further documentation.
|
||||
//
|
||||
// This wrapper does not do any extra authz checks, as the register method has
|
||||
// its own authz mechanism.
|
||||
func (a *ServerWithRoles) RegisterUsingIAMMethod(ctx context.Context, challengeResponse ChallengeResponseFunc) (*proto.Certs, error) {
|
||||
certs, err := a.authServer.RegisterUsingIAMMethod(ctx, challengeResponse)
|
||||
return certs, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// GenerateHostCerts generates new host certificates (signed
|
||||
// by the host certificate authority) for a node.
|
||||
func (a *ServerWithRoles) GenerateHostCerts(ctx context.Context, req *proto.HostCertsRequest) (*proto.Certs, error) {
|
||||
@@ -2716,9 +2728,9 @@ func (a *ServerWithRoles) UpsertTrustedCluster(ctx context.Context, tc types.Tru
|
||||
return a.authServer.UpsertTrustedCluster(ctx, tc)
|
||||
}
|
||||
|
||||
func (a *ServerWithRoles) ValidateTrustedCluster(validateRequest *ValidateTrustedClusterRequest) (*ValidateTrustedClusterResponse, error) {
|
||||
func (a *ServerWithRoles) ValidateTrustedCluster(ctx context.Context, validateRequest *ValidateTrustedClusterRequest) (*ValidateTrustedClusterResponse, error) {
|
||||
// the token provides it's own authorization and authentication
|
||||
return a.authServer.validateTrustedCluster(validateRequest)
|
||||
return a.authServer.validateTrustedCluster(ctx, validateRequest)
|
||||
}
|
||||
|
||||
// DeleteTrustedCluster deletes a trusted cluster by name.
|
||||
|
||||
+4
-4
@@ -542,7 +542,7 @@ func (c *Client) GenerateToken(ctx context.Context, req GenerateTokenRequest) (s
|
||||
|
||||
// RegisterUsingToken calls the auth service API to register a new node using a registration token
|
||||
// which was previously issued via GenerateToken.
|
||||
func (c *Client) RegisterUsingToken(req types.RegisterUsingTokenRequest) (*proto.Certs, error) {
|
||||
func (c *Client) RegisterUsingToken(ctx context.Context, req *types.RegisterUsingTokenRequest) (*proto.Certs, error) {
|
||||
if err := req.CheckAndSetDefaults(); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -1548,7 +1548,7 @@ func (c *Client) DeleteAllUsers() error {
|
||||
return trace.NotImplemented(notImplementedMessage)
|
||||
}
|
||||
|
||||
func (c *Client) ValidateTrustedCluster(validateRequest *ValidateTrustedClusterRequest) (*ValidateTrustedClusterResponse, error) {
|
||||
func (c *Client) ValidateTrustedCluster(ctx context.Context, validateRequest *ValidateTrustedClusterRequest) (*ValidateTrustedClusterResponse, error) {
|
||||
validateRequestRaw, err := validateRequest.ToRaw()
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -1872,7 +1872,7 @@ type ProvisioningService interface {
|
||||
|
||||
// RegisterUsingToken calls the auth service API to register a new node via registration token
|
||||
// which has been previously issued via GenerateToken
|
||||
RegisterUsingToken(req types.RegisterUsingTokenRequest) (*proto.Certs, error)
|
||||
RegisterUsingToken(ctx context.Context, req *types.RegisterUsingTokenRequest) (*proto.Certs, error)
|
||||
|
||||
// RegisterNewAuthServer is used to register new auth server with token
|
||||
RegisterNewAuthServer(ctx context.Context, token string) error
|
||||
@@ -1916,7 +1916,7 @@ type ClientI interface {
|
||||
// ValidateTrustedCluster validates trusted cluster token with
|
||||
// main cluster, in case if validation is successful, main cluster
|
||||
// adds remote cluster
|
||||
ValidateTrustedCluster(*ValidateTrustedClusterRequest) (*ValidateTrustedClusterResponse, error)
|
||||
ValidateTrustedCluster(context.Context, *ValidateTrustedClusterRequest) (*ValidateTrustedClusterResponse, error)
|
||||
|
||||
// GetDomainName returns auth server cluster name
|
||||
GetDomainName() (string, error)
|
||||
|
||||
@@ -0,0 +1,131 @@
|
||||
/*
|
||||
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 auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/gravitational/teleport/api/client/proto"
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/trace"
|
||||
)
|
||||
|
||||
// tokenJoinMethod returns the join method of the token with the given tokenName
|
||||
func (a *Server) tokenJoinMethod(ctx context.Context, tokenName string) types.JoinMethod {
|
||||
provisionToken, err := a.GetToken(ctx, tokenName)
|
||||
if err != nil {
|
||||
// could not find dynamic token, assume static token. If it does not
|
||||
// exist this will be caught later.
|
||||
return types.JoinMethodToken
|
||||
}
|
||||
return provisionToken.GetJoinMethod()
|
||||
}
|
||||
|
||||
// checkTokenJoinRequestCommon checks all token join rules that are common to
|
||||
// all join methods, including token existence, token TTL, and allowed roles.
|
||||
func (a *Server) checkTokenJoinRequestCommon(ctx context.Context, req *types.RegisterUsingTokenRequest) error {
|
||||
// make sure the token is valid
|
||||
roles, _, err := a.ValidateToken(ctx, req.Token)
|
||||
if err != nil {
|
||||
log.Warningf("%q [%v] can not join the cluster with role %s, token error: %v", req.NodeName, req.HostID, req.Role, err)
|
||||
return trace.AccessDenied(fmt.Sprintf("%q [%v] can not join the cluster with role %s, the token is not valid", req.NodeName, req.HostID, req.Role))
|
||||
}
|
||||
|
||||
// make sure the caller is requesting a role allowed by the token
|
||||
if !roles.Include(req.Role) {
|
||||
msg := fmt.Sprintf("node %q [%v] can not join the cluster, the token does not allow %q role", req.NodeName, req.HostID, req.Role)
|
||||
log.Warn(msg)
|
||||
return trace.BadParameter(msg)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// RegisterUsingToken returns credentials for a new node to join the Teleport
|
||||
// cluster using a previously issued token.
|
||||
//
|
||||
// A node must also request a specific role (and the role must match one of the roles
|
||||
// the token was generated for.)
|
||||
//
|
||||
// If a token was generated with a TTL, it gets enforced (can't register new
|
||||
// nodes after TTL expires.)
|
||||
//
|
||||
// If the token includes a specific join method, the rules for that join method
|
||||
// will be checked.
|
||||
func (a *Server) RegisterUsingToken(ctx context.Context, req *types.RegisterUsingTokenRequest) (*proto.Certs, error) {
|
||||
log.Infof("Node %q [%v] is trying to join with role: %v.", req.NodeName, req.HostID, req.Role)
|
||||
if err := req.CheckAndSetDefaults(); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
switch a.tokenJoinMethod(ctx, req.Token) {
|
||||
case types.JoinMethodEC2:
|
||||
if err := a.checkEC2JoinRequest(ctx, req); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
case types.JoinMethodIAM:
|
||||
// IAM join method must use the gRPC RegisterUsingIAMMethod
|
||||
return nil, trace.AccessDenied("this token is only valid for the IAM " +
|
||||
"join method but the node has connected to the wrong endpoint, make " +
|
||||
"sure your node is configured to use the IAM join method")
|
||||
case types.JoinMethodToken:
|
||||
// carry on to common token checking logic
|
||||
default:
|
||||
// this is a logic error, all valid join methods should be captured
|
||||
// above (empty join method will be set to JoinMethodToken by
|
||||
// CheckAndSetDefaults)
|
||||
return nil, trace.BadParameter("unrecognized token join method")
|
||||
}
|
||||
|
||||
// perform common token checks
|
||||
if err := a.checkTokenJoinRequestCommon(ctx, req); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// generate and return host certificate and keys
|
||||
certs, err := a.GenerateHostCerts(ctx,
|
||||
&proto.HostCertsRequest{
|
||||
HostID: req.HostID,
|
||||
NodeName: req.NodeName,
|
||||
Role: req.Role,
|
||||
AdditionalPrincipals: req.AdditionalPrincipals,
|
||||
PublicTLSKey: req.PublicTLSKey,
|
||||
PublicSSHKey: req.PublicSSHKey,
|
||||
RemoteAddr: req.RemoteAddr,
|
||||
DNSNames: req.DNSNames,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
log.Infof("Node %q [%v] has joined the cluster.", req.NodeName, req.HostID)
|
||||
return certs, nil
|
||||
}
|
||||
|
||||
func (a *Server) RegisterNewAuthServer(ctx context.Context, token string) error {
|
||||
tok, err := a.GetToken(ctx, token)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
if !tok.GetRoles().Include(types.RoleAuth) {
|
||||
return trace.AccessDenied("role does not match")
|
||||
}
|
||||
if err := a.DeleteToken(ctx, token); err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -68,13 +68,13 @@ func ec2ClientFromConfig(ctx context.Context, cfg aws.Config) ec2Client {
|
||||
func checkEC2AllowRules(ctx context.Context, iid *imds.InstanceIdentityDocument, provisionToken types.ProvisionToken) error {
|
||||
allowRules := provisionToken.GetAllowRules()
|
||||
for _, rule := range allowRules {
|
||||
// If this rule specifies and AWS account, the IID must match
|
||||
// if this rule specifies an AWS account, the IID must match
|
||||
if len(rule.AWSAccount) > 0 {
|
||||
if rule.AWSAccount != iid.AccountID {
|
||||
continue
|
||||
}
|
||||
}
|
||||
// If this rule specifies any AWS regions, the IID must match one of them
|
||||
// if this rule specifies any AWS regions, the IID must match one of them
|
||||
if len(rule.AWSRegions) > 0 {
|
||||
if !apiutils.SliceContainsStr(rule.AWSRegions, iid.Region) {
|
||||
continue
|
||||
@@ -258,7 +258,7 @@ func dbExists(ctx context.Context, presence services.Presence, hostID string) (b
|
||||
// only allow the roles which will actually be used by all expected instances so
|
||||
// that a stolen IID could not be used to join the cluster with a different
|
||||
// role.
|
||||
func (a *Server) checkInstanceUnique(ctx context.Context, req types.RegisterUsingTokenRequest, iid *imds.InstanceIdentityDocument) error {
|
||||
func (a *Server) checkInstanceUnique(ctx context.Context, req *types.RegisterUsingTokenRequest, iid *imds.InstanceIdentityDocument) error {
|
||||
requestedHostID := req.HostID
|
||||
expectedHostID := utils.NodeIDFromIID(iid)
|
||||
if requestedHostID != expectedHostID {
|
||||
@@ -294,7 +294,7 @@ func (a *Server) checkInstanceUnique(ctx context.Context, req types.RegisterUsin
|
||||
return nil
|
||||
}
|
||||
|
||||
// CheckEC2Request checks register requests which use EC2 Simplified Node
|
||||
// checkEC2JoinRequest checks register requests which use EC2 Simplified Node
|
||||
// Joining. This method checks that:
|
||||
// 1. The given Instance Identity Document has a valid signature (signed by AWS).
|
||||
// 2. A node has not already joined the cluster from this EC2 instance (to
|
||||
@@ -304,34 +304,21 @@ func (a *Server) checkInstanceUnique(ctx context.Context, req types.RegisterUsin
|
||||
// If the request does not include an Instance Identity Document, and the
|
||||
// token does not include any allow rules, this method returns nil and the
|
||||
// normal token checking logic resumes.
|
||||
func (a *Server) CheckEC2Request(ctx context.Context, req types.RegisterUsingTokenRequest) error {
|
||||
requestIncludesIID := req.EC2IdentityDocument != nil
|
||||
func (a *Server) checkEC2JoinRequest(ctx context.Context, req *types.RegisterUsingTokenRequest) error {
|
||||
tokenName := req.Token
|
||||
provisionToken, err := a.GetCache().GetToken(ctx, tokenName)
|
||||
provisionToken, err := a.GetToken(ctx, tokenName)
|
||||
if err != nil {
|
||||
if trace.IsNotFound(err) && !requestIncludesIID {
|
||||
// This is not a Simplified Node Joining request, pass on to the
|
||||
// regular token checking logic in case this is a static token.
|
||||
return nil
|
||||
}
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
tokenRequiresIID := len(provisionToken.GetAllowRules()) > 0
|
||||
|
||||
if !requestIncludesIID && !tokenRequiresIID {
|
||||
// not a simplified node joining request, pass on to the regular token
|
||||
// checking logic
|
||||
return nil
|
||||
}
|
||||
if tokenRequiresIID && !requestIncludesIID {
|
||||
return trace.AccessDenied("this token requires an EC2 Identity Document from the node")
|
||||
}
|
||||
if !tokenRequiresIID && requestIncludesIID {
|
||||
return trace.BadParameter("an EC2 Identity Document is included in a register request for a token which does not expect it")
|
||||
}
|
||||
|
||||
log.Debugf("Received Simplified Node Joining request for host %q", req.HostID)
|
||||
|
||||
if len(req.EC2IdentityDocument) == 0 {
|
||||
return trace.AccessDenied("this token is only valid for the EC2 join " +
|
||||
"method but the node has not included an EC2 Instance Identity " +
|
||||
"Document, make sure your node is configured to use the EC2 join method")
|
||||
}
|
||||
|
||||
iid, err := parseAndVerifyIID(req.EC2IdentityDocument)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
@@ -25,9 +25,6 @@ import (
|
||||
|
||||
"github.com/gravitational/teleport/api/defaults"
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/teleport/lib/auth/testauthority"
|
||||
"github.com/gravitational/teleport/lib/backend/lite"
|
||||
"github.com/gravitational/teleport/lib/services"
|
||||
"github.com/gravitational/trace"
|
||||
|
||||
"github.com/aws/aws-sdk-go-v2/service/ec2"
|
||||
@@ -141,44 +138,13 @@ func (c ec2ClientRunning) DescribeInstances(ctx context.Context, params *ec2.Des
|
||||
}, nil
|
||||
}
|
||||
|
||||
func newAuthServer(t *testing.T) *Server {
|
||||
b, err := lite.NewWithConfig(context.Background(), lite.Config{
|
||||
Path: t.TempDir(),
|
||||
PollStreamPeriod: 200 * time.Millisecond,
|
||||
})
|
||||
func TestAuth_RegisterUsingToken_EC2(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
p, err := newTestPack(ctx, t.TempDir())
|
||||
require.NoError(t, err)
|
||||
a := p.a
|
||||
|
||||
clusterName, err := services.NewClusterNameWithRandomID(types.ClusterNameSpecV2{
|
||||
ClusterName: "test-cluster",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
authConfig := &InitConfig{
|
||||
ClusterName: clusterName,
|
||||
Backend: b,
|
||||
Authority: testauthority.New(),
|
||||
SkipPeriodicOperations: true,
|
||||
}
|
||||
|
||||
a, err := NewServer(authConfig)
|
||||
require.NoError(t, err)
|
||||
|
||||
staticTokens, err := types.NewStaticTokens(types.StaticTokensSpecV2{
|
||||
StaticTokens: []types.ProvisionTokenV1{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
err = a.SetStaticTokens(staticTokens)
|
||||
require.NoError(t, err)
|
||||
|
||||
err = a.UpsertNamespace(types.DefaultNamespace())
|
||||
require.NoError(t, err)
|
||||
|
||||
return a
|
||||
}
|
||||
|
||||
func TestSimplifiedNodeJoin(t *testing.T) {
|
||||
a := newAuthServer(t)
|
||||
|
||||
// upsert a node to test duplicates
|
||||
node := &types.ServerV2{
|
||||
Kind: types.KindNode,
|
||||
Version: types.V2,
|
||||
@@ -187,7 +153,13 @@ func TestSimplifiedNodeJoin(t *testing.T) {
|
||||
Namespace: defaults.Namespace,
|
||||
},
|
||||
}
|
||||
_, err := a.UpsertNode(context.Background(), node)
|
||||
_, err = a.UpsertNode(ctx, node)
|
||||
require.NoError(t, err)
|
||||
|
||||
sshPrivateKey, sshPublicKey, err := a.GenerateKeyPair("")
|
||||
require.NoError(t, err)
|
||||
|
||||
tlsPublicKey, err := PrivateKeyToPublicKeyTLS(sshPrivateKey)
|
||||
require.NoError(t, err)
|
||||
|
||||
isNil := func(err error) bool {
|
||||
@@ -196,7 +168,6 @@ func TestSimplifiedNodeJoin(t *testing.T) {
|
||||
|
||||
testCases := []struct {
|
||||
desc string
|
||||
tokenRules []*types.TokenRule
|
||||
tokenSpec types.ProvisionTokenSpecV2
|
||||
ec2Client ec2Client
|
||||
request types.RegisterUsingTokenRequest
|
||||
@@ -360,6 +331,27 @@ func TestSimplifiedNodeJoin(t *testing.T) {
|
||||
expectError: trace.IsAccessDenied,
|
||||
clock: clockwork.NewFakeClockAt(instance1.pendingTime),
|
||||
},
|
||||
{
|
||||
desc: "no identity document",
|
||||
tokenSpec: types.ProvisionTokenSpecV2{
|
||||
Roles: []types.SystemRole{types.RoleNode},
|
||||
Allow: []*types.TokenRule{
|
||||
&types.TokenRule{
|
||||
AWSAccount: instance1.account,
|
||||
AWSRegions: []string{instance1.region},
|
||||
},
|
||||
},
|
||||
},
|
||||
ec2Client: ec2ClientRunning{},
|
||||
request: types.RegisterUsingTokenRequest{
|
||||
Token: "test_token",
|
||||
NodeName: "node_name",
|
||||
Role: types.RoleNode,
|
||||
HostID: instance1.account + "-" + instance1.instanceID,
|
||||
},
|
||||
expectError: trace.IsAccessDenied,
|
||||
clock: clockwork.NewFakeClockAt(instance1.pendingTime),
|
||||
},
|
||||
{
|
||||
desc: "bad identity document",
|
||||
tokenSpec: types.ProvisionTokenSpecV2{
|
||||
@@ -558,7 +550,12 @@ func TestSimplifiedNodeJoin(t *testing.T) {
|
||||
|
||||
ctx := context.WithValue(context.Background(), ec2ClientKey{}, tc.ec2Client)
|
||||
|
||||
err = a.CheckEC2Request(ctx, tc.request)
|
||||
// set common request values here to avoid setting them in every
|
||||
// testcase
|
||||
tc.request.PublicSSHKey = sshPublicKey
|
||||
tc.request.PublicTLSKey = tlsPublicKey
|
||||
|
||||
_, err = a.RegisterUsingToken(ctx, &tc.request)
|
||||
require.True(t, tc.expectError(err))
|
||||
|
||||
err = a.DeleteToken(context.Background(), token.GetName())
|
||||
@@ -576,16 +573,26 @@ func TestAWSCerts(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestHostUniqueCheck tests the uniqueness check used by CheckEC2Request
|
||||
// TestHostUniqueCheck tests the uniqueness check used by checkEC2JoinRequest
|
||||
func TestHostUniqueCheck(t *testing.T) {
|
||||
a := newAuthServer(t)
|
||||
ctx := context.Background()
|
||||
p, err := newTestPack(ctx, t.TempDir())
|
||||
require.NoError(t, err)
|
||||
a := p.a
|
||||
|
||||
a.clock = clockwork.NewFakeClockAt(instance1.pendingTime)
|
||||
|
||||
token, err := types.NewProvisionTokenFromSpec(
|
||||
"test_token",
|
||||
time.Now().Add(time.Minute),
|
||||
types.ProvisionTokenSpecV2{
|
||||
Roles: []types.SystemRole{types.RoleNode, types.RoleKube},
|
||||
Roles: []types.SystemRole{
|
||||
types.RoleNode,
|
||||
types.RoleProxy,
|
||||
types.RoleKube,
|
||||
types.RoleDatabase,
|
||||
types.RoleApp,
|
||||
},
|
||||
Allow: []*types.TokenRule{
|
||||
&types.TokenRule{
|
||||
AWSAccount: instance1.account,
|
||||
@@ -598,6 +605,12 @@ func TestHostUniqueCheck(t *testing.T) {
|
||||
err = a.UpsertToken(context.Background(), token)
|
||||
require.NoError(t, err)
|
||||
|
||||
sshPrivateKey, sshPublicKey, err := a.GenerateKeyPair("")
|
||||
require.NoError(t, err)
|
||||
|
||||
tlsPublicKey, err := PrivateKeyToPublicKeyTLS(sshPrivateKey)
|
||||
require.NoError(t, err)
|
||||
|
||||
testCases := []struct {
|
||||
role types.SystemRole
|
||||
upserter func(name string)
|
||||
@@ -692,7 +705,7 @@ func TestHostUniqueCheck(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
ctx := context.WithValue(context.Background(), ec2ClientKey{}, ec2ClientRunning{})
|
||||
ctx = context.WithValue(ctx, ec2ClientKey{}, ec2ClientRunning{})
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(string(tc.role), func(t *testing.T) {
|
||||
@@ -702,10 +715,12 @@ func TestHostUniqueCheck(t *testing.T) {
|
||||
Role: tc.role,
|
||||
HostID: instance1.account + "-" + instance1.instanceID,
|
||||
EC2IdentityDocument: instance1.iid,
|
||||
PublicSSHKey: sshPublicKey,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
}
|
||||
|
||||
// request works with no existing host
|
||||
err = a.CheckEC2Request(ctx, request)
|
||||
_, err = a.RegisterUsingToken(ctx, &request)
|
||||
require.NoError(t, err)
|
||||
|
||||
// add the server
|
||||
@@ -713,8 +728,9 @@ func TestHostUniqueCheck(t *testing.T) {
|
||||
tc.upserter(name)
|
||||
|
||||
// request should fail
|
||||
err = a.CheckEC2Request(ctx, request)
|
||||
require.Error(t, err)
|
||||
_, err = a.RegisterUsingToken(ctx, &request)
|
||||
expectedErr := &trace.AccessDeniedError{}
|
||||
require.ErrorAs(t, err, &expectedErr)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -0,0 +1,350 @@
|
||||
/*
|
||||
Copyright 2021-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 auth
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"github.com/gravitational/teleport/api/client/proto"
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/teleport/api/utils"
|
||||
"github.com/gravitational/teleport/lib/utils/aws"
|
||||
|
||||
"github.com/gravitational/trace"
|
||||
)
|
||||
|
||||
const (
|
||||
// Hardcoding the sts API version here may be more strict than necessary,
|
||||
// but this is set by the Teleport node and can only be changed when we
|
||||
// update our AWS SDK dependency. Since Auth should always be upgraded
|
||||
// before nodes, we will have a chance to update the check on Auth if we
|
||||
// ever have a need to allow a newer API version.
|
||||
expectedSTSIdentityRequestBody = "Action=GetCallerIdentity&Version=2011-06-15"
|
||||
|
||||
// Only allowing the global sts endpoint here, Teleport nodes will only send
|
||||
// requests for this endpoint. If we want to start using regional endpoints
|
||||
// we can update this check before updating the nodes.
|
||||
stsHost = "sts.amazonaws.com"
|
||||
|
||||
// AWS SignedHeaders will always be lowercase
|
||||
// https://docs.aws.amazon.com/AmazonS3/latest/API/sigv4-auth-using-authorization-header.html#sigv4-auth-header-overview
|
||||
challengeHeaderKey = "x-teleport-challenge"
|
||||
)
|
||||
|
||||
// validateSTSIdentityRequest checks that a received sts:GetCallerIdentity
|
||||
// request is valid and includes the challenge as a signed header. An example
|
||||
// valid request looks like:
|
||||
// ```
|
||||
// POST / HTTP/1.1
|
||||
// Host: sts.amazonaws.com
|
||||
// Accept: application/json
|
||||
// Authorization: AWS4-HMAC-SHA256 Credential=AAAAAAAAAAAAAAAAAAAA/20211108/us-east-1/sts/aws4_request, SignedHeaders=accept;content-length;content-type;host;x-amz-date;x-amz-security-token;x-teleport-challenge, Signature=999...
|
||||
// Content-Length: 43
|
||||
// Content-Type: application/x-www-form-urlencoded; charset=utf-8
|
||||
// User-Agent: aws-sdk-go/1.37.17 (go1.17.1; darwin; amd64)
|
||||
// X-Amz-Date: 20211108T190420Z
|
||||
// X-Amz-Security-Token: aaa...
|
||||
// X-Teleport-Challenge: 0ezlc3usTAkXeZTcfOazUq0BGrRaKmb4EwODk8U7J5A
|
||||
//
|
||||
// Action=GetCallerIdentity&Version=2011-06-15
|
||||
// ```
|
||||
func validateSTSIdentityRequest(req *http.Request, challenge string) error {
|
||||
if req.Host != stsHost {
|
||||
return trace.AccessDenied("sts identity request is for unknown host %q", req.Host)
|
||||
}
|
||||
|
||||
if req.Method != http.MethodPost {
|
||||
return trace.AccessDenied("sts identity request method %q does not match expected method %q", req.RequestURI, http.MethodPost)
|
||||
}
|
||||
|
||||
if req.Header.Get(challengeHeaderKey) != challenge {
|
||||
return trace.AccessDenied("sts identity request does not include challenge header or it does not match")
|
||||
}
|
||||
|
||||
authHeader := req.Header.Get(aws.AuthorizationHeader)
|
||||
|
||||
sigV4, err := aws.ParseSigV4(authHeader)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
if !utils.SliceContainsStr(sigV4.SignedHeaders, challengeHeaderKey) {
|
||||
return trace.AccessDenied("sts identity request auth header %q does not include "+
|
||||
challengeHeaderKey+" as a signed header", authHeader)
|
||||
}
|
||||
|
||||
body, err := aws.GetAndReplaceReqBody(req)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
if !bytes.Equal([]byte(expectedSTSIdentityRequestBody), body) {
|
||||
return trace.BadParameter("sts request body %q does not equal expected %q", string(body), expectedSTSIdentityRequestBody)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseSTSRequest(req []byte) (*http.Request, error) {
|
||||
httpReq, err := http.ReadRequest(bufio.NewReader(bytes.NewReader(req)))
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// Unset RequestURI and set req.URL instead (necessary quirk of sending a
|
||||
// request parsed by http.ReadRequest). Also, force https here.
|
||||
if httpReq.RequestURI != "/" {
|
||||
return nil, trace.AccessDenied("unexpected sts identity request URI: %q", httpReq.RequestURI)
|
||||
}
|
||||
httpReq.RequestURI = ""
|
||||
httpReq.URL = &url.URL{
|
||||
Scheme: "https",
|
||||
Host: stsHost,
|
||||
}
|
||||
return httpReq, nil
|
||||
}
|
||||
|
||||
// awsIdentity holds aws Account and Arn, used for JSON parsing
|
||||
type awsIdentity struct {
|
||||
Account string `json:"Account"`
|
||||
Arn string `json:"Arn"`
|
||||
}
|
||||
|
||||
// getCallerIdentityReponse is used for JSON parsing
|
||||
type getCallerIdentityResponse struct {
|
||||
GetCallerIdentityResult awsIdentity `json:"GetCallerIdentityResult"`
|
||||
}
|
||||
|
||||
// stsIdentityResponse is used for JSON parsing
|
||||
type stsIdentityResponse struct {
|
||||
GetCallerIdentityResponse getCallerIdentityResponse `json:"GetCallerIdentityResponse"`
|
||||
}
|
||||
|
||||
type stsClient interface {
|
||||
Do(*http.Request) (*http.Response, error)
|
||||
}
|
||||
|
||||
type stsClientKey struct{}
|
||||
|
||||
// stsClientFromContext allows the default http client to be overridden for tests
|
||||
func stsClientFromContext(ctx context.Context) stsClient {
|
||||
client, ok := ctx.Value(stsClientKey{}).(stsClient)
|
||||
if ok {
|
||||
return client
|
||||
}
|
||||
return http.DefaultClient
|
||||
}
|
||||
|
||||
// executeSTSIdentityRequest sends the sts:GetCallerIdentity HTTP request to the
|
||||
// AWS API, parses the response, and returns the awsIdentity
|
||||
func executeSTSIdentityRequest(ctx context.Context, req *http.Request) (*awsIdentity, error) {
|
||||
client := stsClientFromContext(ctx)
|
||||
|
||||
// set the http request context so it can be cancelled
|
||||
req = req.WithContext(ctx)
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, trace.AccessDenied("aws sts api returned status: %q body: %q",
|
||||
resp.Status, body)
|
||||
}
|
||||
|
||||
var identityResponse stsIdentityResponse
|
||||
if err := json.Unmarshal(body, &identityResponse); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
id := &identityResponse.GetCallerIdentityResponse.GetCallerIdentityResult
|
||||
if id.Account == "" {
|
||||
return nil, trace.BadParameter("received empty AWS account ID from sts API")
|
||||
}
|
||||
if id.Arn == "" {
|
||||
return nil, trace.BadParameter("received empty AWS identity ARN from sts API")
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// arnMatches returns true if arn matches the pattern.
|
||||
// Pattern should be an AWS ARN which may include "*" to match any combination
|
||||
// of zero or more characters and "?" to match any single character.
|
||||
// See https://docs.aws.amazon.com/IAM/latest/UserGuide/reference_policies_elements_resource.html
|
||||
func arnMatches(pattern, arn string) (bool, error) {
|
||||
pattern = regexp.QuoteMeta(pattern)
|
||||
pattern = strings.ReplaceAll(pattern, `\*`, ".*")
|
||||
pattern = strings.ReplaceAll(pattern, `\?`, ".")
|
||||
pattern = "^" + pattern + "$"
|
||||
matched, err := regexp.MatchString(pattern, arn)
|
||||
return matched, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// checkIAMAllowRules checks if the given identity matches any of the given
|
||||
// allowRules.
|
||||
func checkIAMAllowRules(identity *awsIdentity, allowRules []*types.TokenRule) error {
|
||||
for _, rule := range allowRules {
|
||||
// if this rule specifies an AWS account, the identity must match
|
||||
if len(rule.AWSAccount) > 0 {
|
||||
if rule.AWSAccount != identity.Account {
|
||||
// account doesn't match, continue to check the next rule
|
||||
continue
|
||||
}
|
||||
}
|
||||
// if this rule specifies an AWS ARN, the identity must match
|
||||
if len(rule.AWSARN) > 0 {
|
||||
matches, err := arnMatches(rule.AWSARN, identity.Arn)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
if !matches {
|
||||
// arn doesn't match, continue to check the next rule
|
||||
continue
|
||||
}
|
||||
}
|
||||
// node identity matches this allow rule
|
||||
return nil
|
||||
}
|
||||
return trace.AccessDenied("instance did not match any allow rules")
|
||||
}
|
||||
|
||||
// checkIAMRequest checks if the given request satisfies the token rules and
|
||||
// included the required challenge.
|
||||
func (a *Server) checkIAMRequest(ctx context.Context, challenge string, req *types.RegisterUsingTokenRequest) error {
|
||||
tokenName := req.Token
|
||||
provisionToken, err := a.GetToken(ctx, tokenName)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
if provisionToken.GetJoinMethod() != types.JoinMethodIAM {
|
||||
return trace.AccessDenied("this token does not support the IAM join method")
|
||||
}
|
||||
|
||||
// parse the incoming http request to the sts:GetCallerIdentity endpoint
|
||||
identityRequest, err := parseSTSRequest(req.STSIdentityRequest)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
|
||||
// validate that the host, method, and headers are correct and the expected
|
||||
// challenge is included in the signed portion of the request
|
||||
if err := validateSTSIdentityRequest(identityRequest, challenge); err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
|
||||
// send the signed request to the public AWS API and get the node identity
|
||||
// from the response
|
||||
identity, err := executeSTSIdentityRequest(ctx, identityRequest)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
|
||||
// check that the node identity matches an allow rule for this token
|
||||
if err := checkIAMAllowRules(identity, provisionToken.GetAllowRules()); err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func generateChallenge() (string, error) {
|
||||
// read 32 crypto-random bytes to generate the challenge
|
||||
challengeRawBytes := make([]byte, 32)
|
||||
if _, err := rand.Read(challengeRawBytes); err != nil {
|
||||
return "", trace.Wrap(err)
|
||||
}
|
||||
|
||||
// encode the challenge to base64 so it can be sent in an HTTP header
|
||||
return base64.RawStdEncoding.EncodeToString(challengeRawBytes), nil
|
||||
}
|
||||
|
||||
// ChallengeResponseFunc is a function type meant to be passed to
|
||||
// RegisterUsingIAMMethod. It must return a *types.RegisterUsingTokenRequest for
|
||||
// a given challenge, or an error.
|
||||
type ChallengeResponseFunc func(challenge string) (*types.RegisterUsingTokenRequest, error)
|
||||
|
||||
// RegisterUsingIAMMethod registers the caller using the IAM join method and
|
||||
// returns signed certs to join the cluster.
|
||||
//
|
||||
// The caller must provide a ChallengeResponseFunc which returns a
|
||||
// *types.RegisterUsingTokenRequest with a signed sts:GetCallerIdentity request
|
||||
// including the challenge as a signed header.
|
||||
func (a *Server) RegisterUsingIAMMethod(ctx context.Context, challengeResponse ChallengeResponseFunc) (*proto.Certs, error) {
|
||||
clientAddr, ok := ctx.Value(ContextClientAddr).(net.Addr)
|
||||
if !ok {
|
||||
return nil, trace.BadParameter("logic error: client address was not set")
|
||||
}
|
||||
|
||||
challenge, err := generateChallenge()
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
req, err := challengeResponse(challenge)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// fill in the client remote addr to the register request
|
||||
req.RemoteAddr = clientAddr.String()
|
||||
if err := req.CheckAndSetDefaults(); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// perform common token checks
|
||||
if err := a.checkTokenJoinRequestCommon(ctx, req); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// check that the GetCallerIdentity request is valid and matches the token
|
||||
if err := a.checkIAMRequest(ctx, challenge, req); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// generate and return host certificate and keys
|
||||
certs, err := a.GenerateHostCerts(ctx,
|
||||
&proto.HostCertsRequest{
|
||||
HostID: req.HostID,
|
||||
NodeName: req.NodeName,
|
||||
Role: req.Role,
|
||||
AdditionalPrincipals: req.AdditionalPrincipals,
|
||||
PublicTLSKey: req.PublicTLSKey,
|
||||
PublicSSHKey: req.PublicSSHKey,
|
||||
RemoteAddr: req.RemoteAddr,
|
||||
DNSNames: req.DNSNames,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
log.Infof("Node %q [%v] has joined the cluster.", req.NodeName, req.HostID)
|
||||
return certs, nil
|
||||
}
|
||||
@@ -0,0 +1,422 @@
|
||||
/*
|
||||
Copyright 2021-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 auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/trace"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func responseFromAWSIdentity(id awsIdentity) string {
|
||||
return fmt.Sprintf(`{
|
||||
"GetCallerIdentityResponse": {
|
||||
"GetCallerIdentityResult": {
|
||||
"Account": "%s",
|
||||
"Arn": "%s"
|
||||
}}}`, id.Account, id.Arn)
|
||||
}
|
||||
|
||||
type mockClient struct {
|
||||
respStatusCode int
|
||||
respBody string
|
||||
}
|
||||
|
||||
func (c *mockClient) Do(req *http.Request) (*http.Response, error) {
|
||||
return &http.Response{
|
||||
StatusCode: c.respStatusCode,
|
||||
Body: io.NopCloser(strings.NewReader(c.respBody)),
|
||||
}, nil
|
||||
}
|
||||
|
||||
const identityRequestTemplate = `POST / HTTP/1.1
|
||||
Host: sts.amazonaws.com
|
||||
User-Agent: aws-sdk-go/1.37.17 (go1.17.1; darwin; amd64)
|
||||
Content-Length: 43
|
||||
Accept: application/json
|
||||
Authorization: AWS4-HMAC-SHA256 Credential=AAAAAAAAAAAAAAAAAAAA/20211102/us-east-1/sts/aws4_request, SignedHeaders=accept;content-length;content-type;host;x-amz-date;x-amz-security-token;x-teleport-challenge, Signature=111
|
||||
Content-Type: application/x-www-form-urlencoded; charset=utf-8
|
||||
X-Amz-Date: 20211102T204300Z
|
||||
X-Amz-Security-Token: aaa
|
||||
X-Teleport-Challenge: %s
|
||||
|
||||
Action=GetCallerIdentity&Version=2011-06-15`
|
||||
|
||||
const wrongHostTemplate = `POST / HTTP/1.1
|
||||
Host: sts.example.com
|
||||
User-Agent: aws-sdk-go/1.37.17 (go1.17.1; darwin; amd64)
|
||||
Content-Length: 43
|
||||
Accept: application/json
|
||||
Authorization: AWS4-HMAC-SHA256 Credential=AAAAAAAAAAAAAAAAAAAA/20211102/us-east-1/sts/aws4_request, SignedHeaders=accept;content-length;content-type;host;x-amz-date;x-amz-security-token;x-teleport-challenge, Signature=111
|
||||
Content-Type: application/x-www-form-urlencoded; charset=utf-8
|
||||
X-Amz-Date: 20211102T204300Z
|
||||
X-Amz-Security-Token: aaa
|
||||
X-Teleport-Challenge: %s
|
||||
|
||||
Action=GetCallerIdentity&Version=2011-06-15`
|
||||
|
||||
const unsignedChallengeTemplate = `POST / HTTP/1.1
|
||||
Host: sts.amazonaws.com
|
||||
User-Agent: aws-sdk-go/1.37.17 (go1.17.1; darwin; amd64)
|
||||
Content-Length: 43
|
||||
Accept: application/json
|
||||
Authorization: AWS4-HMAC-SHA256 Credential=AAAAAAAAAAAAAAAAAAAA/20211102/us-east-1/sts/aws4_request, SignedHeaders=accept;content-length;content-type;host;x-amz-date;x-amz-security-token, Signature=111
|
||||
Content-Type: application/x-www-form-urlencoded; charset=utf-8
|
||||
X-Amz-Date: 20211102T204300Z
|
||||
X-Amz-Security-Token: aaa
|
||||
X-Teleport-Challenge: %s
|
||||
|
||||
Action=GetCallerIdentity&Version=2011-06-15`
|
||||
|
||||
func TestAuth_RegisterUsingIAMMethod(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
p, err := newTestPack(ctx, t.TempDir())
|
||||
require.NoError(t, err)
|
||||
a := p.a
|
||||
|
||||
sshPrivateKey, sshPublicKey, err := a.GenerateKeyPair("")
|
||||
require.NoError(t, err)
|
||||
|
||||
tlsPublicKey, err := PrivateKeyToPublicKeyTLS(sshPrivateKey)
|
||||
require.NoError(t, err)
|
||||
|
||||
isAccessDenied := func(t require.TestingT, err error, _ ...interface{}) {
|
||||
require.True(t, trace.IsAccessDenied(err), "expected Access Denied error, actual error: %v", err)
|
||||
}
|
||||
isBadParameter := func(t require.TestingT, err error, _ ...interface{}) {
|
||||
require.True(t, trace.IsBadParameter(err), "expected Bad Parameter error, actual error: %v", err)
|
||||
}
|
||||
|
||||
testCases := []struct {
|
||||
desc string
|
||||
tokenName string
|
||||
requestTokenName string
|
||||
tokenSpec types.ProvisionTokenSpecV2
|
||||
stsClient stsClient
|
||||
challengeResponseOverride string
|
||||
requestTemplate string
|
||||
challengeResponseErr error
|
||||
assertError require.ErrorAssertionFunc
|
||||
}{
|
||||
{
|
||||
desc: "basic passing case",
|
||||
tokenName: "test-token",
|
||||
requestTokenName: "test-token",
|
||||
tokenSpec: types.ProvisionTokenSpecV2{
|
||||
Roles: []types.SystemRole{types.RoleNode},
|
||||
Allow: []*types.TokenRule{
|
||||
&types.TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSARN: "arn:aws::1111",
|
||||
},
|
||||
},
|
||||
JoinMethod: types.JoinMethodIAM,
|
||||
},
|
||||
stsClient: &mockClient{
|
||||
respStatusCode: http.StatusOK,
|
||||
respBody: responseFromAWSIdentity(awsIdentity{
|
||||
Account: "1234",
|
||||
Arn: "arn:aws::1111",
|
||||
}),
|
||||
},
|
||||
requestTemplate: identityRequestTemplate,
|
||||
assertError: require.NoError,
|
||||
},
|
||||
{
|
||||
desc: "wildcard arn 1",
|
||||
tokenName: "test-token",
|
||||
requestTokenName: "test-token",
|
||||
tokenSpec: types.ProvisionTokenSpecV2{
|
||||
Roles: []types.SystemRole{types.RoleNode},
|
||||
Allow: []*types.TokenRule{
|
||||
&types.TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSARN: "arn:aws::role/admins-*",
|
||||
},
|
||||
},
|
||||
JoinMethod: types.JoinMethodIAM,
|
||||
},
|
||||
stsClient: &mockClient{
|
||||
respStatusCode: http.StatusOK,
|
||||
respBody: responseFromAWSIdentity(awsIdentity{
|
||||
Account: "1234",
|
||||
Arn: "arn:aws::role/admins-test",
|
||||
}),
|
||||
},
|
||||
requestTemplate: identityRequestTemplate,
|
||||
assertError: require.NoError,
|
||||
},
|
||||
{
|
||||
desc: "wildcard arn 2",
|
||||
tokenName: "test-token",
|
||||
requestTokenName: "test-token",
|
||||
tokenSpec: types.ProvisionTokenSpecV2{
|
||||
Roles: []types.SystemRole{types.RoleNode},
|
||||
Allow: []*types.TokenRule{
|
||||
&types.TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSARN: "arn:aws::role/admins-???",
|
||||
},
|
||||
},
|
||||
JoinMethod: types.JoinMethodIAM,
|
||||
},
|
||||
stsClient: &mockClient{
|
||||
respStatusCode: http.StatusOK,
|
||||
respBody: responseFromAWSIdentity(awsIdentity{
|
||||
Account: "1234",
|
||||
Arn: "arn:aws::role/admins-123",
|
||||
}),
|
||||
},
|
||||
requestTemplate: identityRequestTemplate,
|
||||
assertError: require.NoError,
|
||||
},
|
||||
{
|
||||
desc: "wrong token",
|
||||
tokenName: "test-token",
|
||||
requestTokenName: "wrong-token",
|
||||
tokenSpec: types.ProvisionTokenSpecV2{
|
||||
Roles: []types.SystemRole{types.RoleNode},
|
||||
Allow: []*types.TokenRule{
|
||||
&types.TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSARN: "arn:aws::1111",
|
||||
},
|
||||
},
|
||||
JoinMethod: types.JoinMethodIAM,
|
||||
},
|
||||
stsClient: &mockClient{
|
||||
respStatusCode: http.StatusOK,
|
||||
respBody: responseFromAWSIdentity(awsIdentity{
|
||||
Account: "1234",
|
||||
Arn: "arn:aws::1111",
|
||||
}),
|
||||
},
|
||||
requestTemplate: identityRequestTemplate,
|
||||
assertError: isAccessDenied,
|
||||
},
|
||||
{
|
||||
desc: "challenge response error",
|
||||
tokenName: "test-token",
|
||||
requestTokenName: "test-token",
|
||||
tokenSpec: types.ProvisionTokenSpecV2{
|
||||
Roles: []types.SystemRole{types.RoleNode},
|
||||
Allow: []*types.TokenRule{
|
||||
&types.TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSARN: "arn:aws::1111",
|
||||
},
|
||||
},
|
||||
JoinMethod: types.JoinMethodIAM,
|
||||
},
|
||||
stsClient: &mockClient{
|
||||
respStatusCode: http.StatusOK,
|
||||
respBody: responseFromAWSIdentity(awsIdentity{
|
||||
Account: "1234",
|
||||
Arn: "arn:aws::1111",
|
||||
}),
|
||||
},
|
||||
requestTemplate: identityRequestTemplate,
|
||||
challengeResponseErr: trace.BadParameter("test error"),
|
||||
assertError: isBadParameter,
|
||||
},
|
||||
{
|
||||
desc: "wrong arn",
|
||||
tokenName: "test-token",
|
||||
requestTokenName: "test-token",
|
||||
tokenSpec: types.ProvisionTokenSpecV2{
|
||||
Roles: []types.SystemRole{types.RoleNode},
|
||||
Allow: []*types.TokenRule{
|
||||
&types.TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSARN: "arn:aws::role/admins-???",
|
||||
},
|
||||
},
|
||||
JoinMethod: types.JoinMethodIAM,
|
||||
},
|
||||
stsClient: &mockClient{
|
||||
respStatusCode: http.StatusOK,
|
||||
respBody: responseFromAWSIdentity(awsIdentity{
|
||||
Account: "1234",
|
||||
Arn: "arn:aws::role/admins-1234",
|
||||
}),
|
||||
},
|
||||
requestTemplate: identityRequestTemplate,
|
||||
assertError: isAccessDenied,
|
||||
},
|
||||
{
|
||||
desc: "wrong challenge",
|
||||
tokenName: "test-token",
|
||||
requestTokenName: "test-token",
|
||||
tokenSpec: types.ProvisionTokenSpecV2{
|
||||
Roles: []types.SystemRole{types.RoleNode},
|
||||
Allow: []*types.TokenRule{
|
||||
&types.TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSARN: "arn:aws::1111",
|
||||
},
|
||||
},
|
||||
JoinMethod: types.JoinMethodIAM,
|
||||
},
|
||||
stsClient: &mockClient{
|
||||
respStatusCode: http.StatusOK,
|
||||
respBody: responseFromAWSIdentity(awsIdentity{
|
||||
Account: "1234",
|
||||
Arn: "arn:aws::1111",
|
||||
}),
|
||||
},
|
||||
challengeResponseOverride: "wrong-challenge",
|
||||
requestTemplate: identityRequestTemplate,
|
||||
assertError: isAccessDenied,
|
||||
},
|
||||
{
|
||||
desc: "wrong account",
|
||||
tokenName: "test-token",
|
||||
requestTokenName: "test-token",
|
||||
tokenSpec: types.ProvisionTokenSpecV2{
|
||||
Roles: []types.SystemRole{types.RoleNode},
|
||||
Allow: []*types.TokenRule{
|
||||
&types.TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSARN: "arn:aws::1111",
|
||||
},
|
||||
},
|
||||
JoinMethod: types.JoinMethodIAM,
|
||||
},
|
||||
stsClient: &mockClient{
|
||||
respStatusCode: http.StatusOK,
|
||||
respBody: responseFromAWSIdentity(awsIdentity{
|
||||
Account: "5678",
|
||||
Arn: "arn:aws::1111",
|
||||
}),
|
||||
},
|
||||
requestTemplate: identityRequestTemplate,
|
||||
assertError: isAccessDenied,
|
||||
},
|
||||
{
|
||||
desc: "sts api error",
|
||||
tokenName: "test-token",
|
||||
requestTokenName: "test-token",
|
||||
tokenSpec: types.ProvisionTokenSpecV2{
|
||||
Roles: []types.SystemRole{types.RoleNode},
|
||||
Allow: []*types.TokenRule{
|
||||
&types.TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSARN: "arn:aws::1111",
|
||||
},
|
||||
},
|
||||
JoinMethod: types.JoinMethodIAM,
|
||||
},
|
||||
stsClient: &mockClient{
|
||||
respStatusCode: http.StatusForbidden,
|
||||
respBody: "access denied",
|
||||
},
|
||||
requestTemplate: identityRequestTemplate,
|
||||
assertError: isAccessDenied,
|
||||
},
|
||||
{
|
||||
desc: "wrong sts host",
|
||||
tokenName: "test-token",
|
||||
requestTokenName: "test-token",
|
||||
tokenSpec: types.ProvisionTokenSpecV2{
|
||||
Roles: []types.SystemRole{types.RoleNode},
|
||||
Allow: []*types.TokenRule{
|
||||
&types.TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSARN: "arn:aws::1111",
|
||||
},
|
||||
},
|
||||
JoinMethod: types.JoinMethodIAM,
|
||||
},
|
||||
stsClient: &mockClient{
|
||||
respStatusCode: http.StatusOK,
|
||||
respBody: responseFromAWSIdentity(awsIdentity{
|
||||
Account: "1234",
|
||||
Arn: "arn:aws::1111",
|
||||
}),
|
||||
},
|
||||
requestTemplate: wrongHostTemplate,
|
||||
assertError: isAccessDenied,
|
||||
},
|
||||
{
|
||||
desc: "unsigned challenge header",
|
||||
tokenName: "test-token",
|
||||
requestTokenName: "test-token",
|
||||
tokenSpec: types.ProvisionTokenSpecV2{
|
||||
Roles: []types.SystemRole{types.RoleNode},
|
||||
Allow: []*types.TokenRule{
|
||||
&types.TokenRule{
|
||||
AWSAccount: "1234",
|
||||
AWSARN: "arn:aws::1111",
|
||||
},
|
||||
},
|
||||
JoinMethod: types.JoinMethodIAM,
|
||||
},
|
||||
stsClient: &mockClient{
|
||||
respStatusCode: http.StatusOK,
|
||||
respBody: responseFromAWSIdentity(awsIdentity{
|
||||
Account: "1234",
|
||||
Arn: "arn:aws::1111",
|
||||
}),
|
||||
},
|
||||
requestTemplate: unsignedChallengeTemplate,
|
||||
assertError: isAccessDenied,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.desc, func(t *testing.T) {
|
||||
// add token to auth server
|
||||
token, err := types.NewProvisionTokenFromSpec(
|
||||
tc.tokenName,
|
||||
time.Now().Add(time.Minute),
|
||||
tc.tokenSpec)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, a.UpsertToken(ctx, token))
|
||||
t.Cleanup(func() { require.NoError(t, a.DeleteToken(ctx, token.GetName())) })
|
||||
|
||||
requestContext := context.Background()
|
||||
requestContext = context.WithValue(requestContext, ContextClientAddr, &net.IPAddr{})
|
||||
requestContext = context.WithValue(requestContext, stsClientKey{}, tc.stsClient)
|
||||
|
||||
_, err = a.RegisterUsingIAMMethod(requestContext, func(challenge string) (*types.RegisterUsingTokenRequest, error) {
|
||||
if tc.challengeResponseOverride != "" {
|
||||
challenge = tc.challengeResponseOverride
|
||||
}
|
||||
identityRequest := []byte(fmt.Sprintf(tc.requestTemplate, challenge))
|
||||
req := &types.RegisterUsingTokenRequest{
|
||||
Token: tc.requestTokenName,
|
||||
HostID: "test-node",
|
||||
Role: types.RoleNode,
|
||||
PublicSSHKey: sshPublicKey,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
STSIdentityRequest: identityRequest,
|
||||
}
|
||||
return req, tc.challengeResponseErr
|
||||
})
|
||||
tc.assertError(t, err)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,251 @@
|
||||
/*
|
||||
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 auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gravitational/teleport/api/client/proto"
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/teleport/api/utils/sshutils"
|
||||
"github.com/gravitational/trace"
|
||||
"github.com/jonboulle/clockwork"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestAuth_RegisterUsingToken(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
p, err := newTestPack(ctx, t.TempDir())
|
||||
require.NoError(t, err)
|
||||
a := p.a
|
||||
|
||||
// create a static token
|
||||
staticToken := types.ProvisionTokenV1{
|
||||
Roles: []types.SystemRole{types.RoleNode},
|
||||
Token: "static_token",
|
||||
}
|
||||
staticTokens, err := types.NewStaticTokens(types.StaticTokensSpecV2{
|
||||
StaticTokens: []types.ProvisionTokenV1{staticToken},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
err = p.a.SetStaticTokens(staticTokens)
|
||||
require.NoError(t, err)
|
||||
|
||||
// create a dynamic token
|
||||
dynamicToken, err := a.GenerateToken(ctx, GenerateTokenRequest{
|
||||
Roles: types.SystemRoles{types.RoleNode},
|
||||
TTL: time.Hour,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, dynamicToken)
|
||||
|
||||
sshPrivateKey, sshPublicKey, err := a.GenerateKeyPair("")
|
||||
require.NoError(t, err)
|
||||
|
||||
tlsPublicKey, err := PrivateKeyToPublicKeyTLS(sshPrivateKey)
|
||||
require.NoError(t, err)
|
||||
|
||||
testcases := []struct {
|
||||
desc string
|
||||
req *types.RegisterUsingTokenRequest
|
||||
certsAssertion func(*proto.Certs)
|
||||
errorAssertion func(error) bool
|
||||
clock clockwork.Clock
|
||||
}{
|
||||
{
|
||||
desc: "reject empty",
|
||||
req: &types.RegisterUsingTokenRequest{},
|
||||
errorAssertion: trace.IsBadParameter,
|
||||
},
|
||||
{
|
||||
desc: "reject no token",
|
||||
req: &types.RegisterUsingTokenRequest{
|
||||
HostID: "localhost",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleNode,
|
||||
PublicSSHKey: sshPublicKey,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
},
|
||||
errorAssertion: trace.IsBadParameter,
|
||||
},
|
||||
{
|
||||
desc: "reject no HostID",
|
||||
req: &types.RegisterUsingTokenRequest{
|
||||
Token: staticToken.Token,
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleNode,
|
||||
PublicSSHKey: sshPublicKey,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
},
|
||||
errorAssertion: trace.IsBadParameter,
|
||||
},
|
||||
{
|
||||
desc: "allow no NodeName",
|
||||
req: &types.RegisterUsingTokenRequest{
|
||||
Token: staticToken.Token,
|
||||
HostID: "localhost",
|
||||
Role: types.RoleNode,
|
||||
PublicSSHKey: sshPublicKey,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
},
|
||||
},
|
||||
{
|
||||
desc: "reject no SSH pub",
|
||||
req: &types.RegisterUsingTokenRequest{
|
||||
Token: staticToken.Token,
|
||||
HostID: "localhost",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleNode,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
},
|
||||
errorAssertion: trace.IsBadParameter,
|
||||
},
|
||||
{
|
||||
desc: "reject no TLS pub",
|
||||
req: &types.RegisterUsingTokenRequest{
|
||||
Token: staticToken.Token,
|
||||
HostID: "localhost",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleNode,
|
||||
PublicSSHKey: sshPublicKey,
|
||||
},
|
||||
errorAssertion: trace.IsBadParameter,
|
||||
},
|
||||
{
|
||||
desc: "reject bad token",
|
||||
req: &types.RegisterUsingTokenRequest{
|
||||
Token: "not a token",
|
||||
HostID: "localhost",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleNode,
|
||||
PublicSSHKey: sshPublicKey,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
},
|
||||
errorAssertion: trace.IsAccessDenied,
|
||||
},
|
||||
{
|
||||
desc: "allow static token",
|
||||
req: &types.RegisterUsingTokenRequest{
|
||||
Token: staticToken.Token,
|
||||
HostID: "localhost",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleNode,
|
||||
PublicSSHKey: sshPublicKey,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
},
|
||||
},
|
||||
{
|
||||
desc: "reject wrong role static",
|
||||
req: &types.RegisterUsingTokenRequest{
|
||||
Token: staticToken.Token,
|
||||
HostID: "localhost",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleProxy,
|
||||
PublicSSHKey: sshPublicKey,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
},
|
||||
errorAssertion: trace.IsBadParameter,
|
||||
},
|
||||
{
|
||||
desc: "allow dynamic token",
|
||||
req: &types.RegisterUsingTokenRequest{
|
||||
Token: dynamicToken,
|
||||
HostID: "localhost",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleNode,
|
||||
PublicSSHKey: sshPublicKey,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
},
|
||||
},
|
||||
{
|
||||
desc: "reject wrong role dynamic",
|
||||
req: &types.RegisterUsingTokenRequest{
|
||||
Token: dynamicToken,
|
||||
HostID: "localhost",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleProxy,
|
||||
PublicSSHKey: sshPublicKey,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
},
|
||||
errorAssertion: trace.IsBadParameter,
|
||||
},
|
||||
{
|
||||
desc: "check additional pricipals",
|
||||
req: &types.RegisterUsingTokenRequest{
|
||||
Token: dynamicToken,
|
||||
HostID: "localhost",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleNode,
|
||||
PublicSSHKey: sshPublicKey,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
AdditionalPrincipals: []string{"example.com"},
|
||||
},
|
||||
certsAssertion: func(certs *proto.Certs) {
|
||||
hostCert, err := sshutils.ParseCertificate(certs.SSH)
|
||||
require.NoError(t, err)
|
||||
require.Contains(t, hostCert.ValidPrincipals, "example.com")
|
||||
},
|
||||
},
|
||||
{
|
||||
desc: "reject expired dynamic token",
|
||||
req: &types.RegisterUsingTokenRequest{
|
||||
Token: dynamicToken,
|
||||
HostID: "localhost",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleNode,
|
||||
PublicSSHKey: sshPublicKey,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
},
|
||||
clock: clockwork.NewFakeClockAt(time.Now().Add(time.Hour + 1)),
|
||||
errorAssertion: trace.IsAccessDenied,
|
||||
},
|
||||
{
|
||||
// relies on token being deleted during previous testcase
|
||||
desc: "expired token should be gone",
|
||||
req: &types.RegisterUsingTokenRequest{
|
||||
Token: dynamicToken,
|
||||
HostID: "localhost",
|
||||
NodeName: "node-name",
|
||||
Role: types.RoleNode,
|
||||
PublicSSHKey: sshPublicKey,
|
||||
PublicTLSKey: tlsPublicKey,
|
||||
},
|
||||
clock: clockwork.NewRealClock(),
|
||||
errorAssertion: trace.IsAccessDenied,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range testcases {
|
||||
t.Run(tc.desc, func(t *testing.T) {
|
||||
if tc.clock == nil {
|
||||
tc.clock = clockwork.NewRealClock()
|
||||
}
|
||||
a.SetClock(tc.clock)
|
||||
certs, err := a.RegisterUsingToken(ctx, tc.req)
|
||||
if tc.errorAssertion != nil {
|
||||
require.True(t, tc.errorAssertion(err))
|
||||
return
|
||||
}
|
||||
require.NoError(t, err)
|
||||
if tc.certsAssertion != nil {
|
||||
tc.certsAssertion(certs)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
+12
-11
@@ -227,17 +227,18 @@ func registerThroughAuth(token string, params RegisterParams) (*Identity, error)
|
||||
defer client.Close()
|
||||
|
||||
// Get the SSH and X509 certificates for a node.
|
||||
certs, err := client.RegisterUsingToken(types.RegisterUsingTokenRequest{
|
||||
Token: token,
|
||||
HostID: params.ID.HostUUID,
|
||||
NodeName: params.ID.NodeName,
|
||||
Role: params.ID.Role,
|
||||
AdditionalPrincipals: params.AdditionalPrincipals,
|
||||
DNSNames: params.DNSNames,
|
||||
PublicTLSKey: params.PublicTLSKey,
|
||||
PublicSSHKey: params.PublicSSHKey,
|
||||
EC2IdentityDocument: params.EC2IdentityDocument,
|
||||
})
|
||||
certs, err := client.RegisterUsingToken(context.Background(),
|
||||
&types.RegisterUsingTokenRequest{
|
||||
Token: token,
|
||||
HostID: params.ID.HostUUID,
|
||||
NodeName: params.ID.NodeName,
|
||||
Role: params.ID.Role,
|
||||
AdditionalPrincipals: params.AdditionalPrincipals,
|
||||
DNSNames: params.DNSNames,
|
||||
PublicTLSKey: params.PublicTLSKey,
|
||||
PublicSSHKey: params.PublicSSHKey,
|
||||
EC2IdentityDocument: params.EC2IdentityDocument,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
@@ -452,7 +452,7 @@ func (a *Server) GetRemoteClusters(opts ...services.MarshalOption) ([]types.Remo
|
||||
return remoteClusters, nil
|
||||
}
|
||||
|
||||
func (a *Server) validateTrustedCluster(validateRequest *ValidateTrustedClusterRequest) (resp *ValidateTrustedClusterResponse, err error) {
|
||||
func (a *Server) validateTrustedCluster(ctx context.Context, validateRequest *ValidateTrustedClusterRequest) (resp *ValidateTrustedClusterResponse, err error) {
|
||||
defer func() {
|
||||
if err != nil {
|
||||
log.WithError(err).Info("Trusted cluster validation failed")
|
||||
@@ -467,7 +467,7 @@ func (a *Server) validateTrustedCluster(validateRequest *ValidateTrustedClusterR
|
||||
}
|
||||
|
||||
// validate that we generated the token
|
||||
tokenLabels, err := a.validateTrustedClusterToken(validateRequest.Token)
|
||||
tokenLabels, err := a.validateTrustedClusterToken(ctx, validateRequest.Token)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -526,8 +526,8 @@ func (a *Server) validateTrustedCluster(validateRequest *ValidateTrustedClusterR
|
||||
return &validateResponse, nil
|
||||
}
|
||||
|
||||
func (a *Server) validateTrustedClusterToken(token string) (map[string]string, error) {
|
||||
roles, labels, err := a.ValidateToken(token)
|
||||
func (a *Server) validateTrustedClusterToken(ctx context.Context, token string) (map[string]string, error) {
|
||||
roles, labels, err := a.ValidateToken(ctx, token)
|
||||
if err != nil {
|
||||
return nil, trace.AccessDenied("the remote server denied access: invalid cluster token")
|
||||
}
|
||||
|
||||
@@ -1960,7 +1960,7 @@ func splitRoles(roles string) []string {
|
||||
// applyTokenConfig applies the auth_token and join_params to the config
|
||||
func applyTokenConfig(fc *FileConfig, cfg *service.Config) error {
|
||||
if fc.AuthToken != "" {
|
||||
cfg.JoinMethod = service.JoinMethodToken
|
||||
cfg.JoinMethod = types.JoinMethodToken
|
||||
_, err := cfg.ApplyToken(fc.AuthToken)
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
@@ -1972,7 +1972,7 @@ func applyTokenConfig(fc *FileConfig, cfg *service.Config) error {
|
||||
if fc.JoinParams.Method != "ec2" {
|
||||
return trace.BadParameter(`unknown value for join_params.method: %q, expected "ec2"`, fc.JoinParams.Method)
|
||||
}
|
||||
cfg.JoinMethod = service.JoinMethodEC2
|
||||
cfg.JoinMethod = types.JoinMethodEC2
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
+1
-12
@@ -87,7 +87,7 @@ type Config struct {
|
||||
Token string
|
||||
|
||||
// JoinMethod is the method the instance will use to join the auth server
|
||||
JoinMethod JoinMethod
|
||||
JoinMethod types.JoinMethod
|
||||
|
||||
// AuthServers is a list of auth servers, proxies and peer auth servers to
|
||||
// connect to. Yes, this is not just auth servers, the field name is
|
||||
@@ -1145,14 +1145,3 @@ func ApplyFIPSDefaults(cfg *Config) {
|
||||
// entire cluster is FedRAMP/FIPS 140-2 compliant.
|
||||
cfg.Auth.SessionRecordingConfig.SetMode(types.RecordAtNode)
|
||||
}
|
||||
|
||||
// JoinMethod is the method the instance will use to join the auth server.
|
||||
type JoinMethod int
|
||||
|
||||
const (
|
||||
// JoinMethodToken means the instance will use a basic token.
|
||||
JoinMethodToken JoinMethod = iota
|
||||
// JoinMethodEC2 means the instance will use Simplified Node Joining and send an
|
||||
// EC2 Instance Identity Document.
|
||||
JoinMethodEC2
|
||||
)
|
||||
|
||||
@@ -382,7 +382,7 @@ func (process *TeleportProcess) firstTimeConnect(role types.SystemRole) (*Connec
|
||||
}
|
||||
|
||||
var ec2IdentityDocument []byte
|
||||
if process.Config.JoinMethod == JoinMethodEC2 {
|
||||
if process.Config.JoinMethod == types.JoinMethodEC2 {
|
||||
ec2IdentityDocument, err = utils.GetEC2IdentityDocument()
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
|
||||
@@ -648,9 +648,9 @@ func NewTeleport(cfg *Config) (*TeleportProcess, error) {
|
||||
cfg.Log.Infof("Taking host UUID from first identity: %v.", cfg.HostUUID)
|
||||
} else {
|
||||
switch cfg.JoinMethod {
|
||||
case JoinMethodToken:
|
||||
case types.JoinMethodToken, types.JoinMethodUnspecified, types.JoinMethodIAM:
|
||||
cfg.HostUUID = uuid.New().String()
|
||||
case JoinMethodEC2:
|
||||
case types.JoinMethodEC2:
|
||||
cfg.HostUUID, err = utils.GetEC2NodeID()
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
|
||||
@@ -34,8 +34,8 @@ import (
|
||||
|
||||
"github.com/gravitational/teleport/lib/client"
|
||||
"github.com/gravitational/teleport/lib/srv/alpnproxy/common"
|
||||
appaws "github.com/gravitational/teleport/lib/srv/app/aws"
|
||||
"github.com/gravitational/teleport/lib/utils"
|
||||
"github.com/gravitational/teleport/lib/utils/aws"
|
||||
)
|
||||
|
||||
// LocalProxy allows upgrading incoming connection to TLS where custom TLS values are set SNI ALPN and
|
||||
@@ -323,7 +323,7 @@ func (l *LocalProxy) StartAWSAccessProxy(ctx context.Context) error {
|
||||
Transport: tr,
|
||||
}
|
||||
err := http.Serve(l.cfg.Listener, http.HandlerFunc(func(rw http.ResponseWriter, req *http.Request) {
|
||||
if err := appaws.VerifyAWSSignature(req, l.cfg.AWSCredentials); err != nil {
|
||||
if err := aws.VerifyAWSSignature(req, l.cfg.AWSCredentials); err != nil {
|
||||
log.WithError(err).Errorf("AWS signature verification failed.")
|
||||
rw.WriteHeader(http.StatusForbidden)
|
||||
return
|
||||
|
||||
@@ -38,6 +38,7 @@ import (
|
||||
"github.com/gravitational/teleport/lib/defaults"
|
||||
appcommon "github.com/gravitational/teleport/lib/srv/app/common"
|
||||
"github.com/gravitational/teleport/lib/tlsca"
|
||||
awsutils "github.com/gravitational/teleport/lib/utils/aws"
|
||||
)
|
||||
|
||||
// NewSigningService creates a new instance of SigningService.
|
||||
@@ -187,7 +188,7 @@ func (s *SigningService) formatForwardResponseError(rw http.ResponseWriter, r *h
|
||||
// resolveEndpoint extracts the aws-service on and aws-region from the request authorization header
|
||||
// and resolves the aws-service and aws-region to AWS endpoint.
|
||||
func resolveEndpoint(r *http.Request) (*endpoints.ResolvedEndpoint, error) {
|
||||
awsAuthHeader, err := ParseSigV4(r.Header.Get(authorizationHeader))
|
||||
awsAuthHeader, err := awsutils.ParseSigV4(r.Header.Get(awsutils.AuthorizationHeader))
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -201,7 +202,7 @@ func resolveEndpoint(r *http.Request) (*endpoints.ResolvedEndpoint, error) {
|
||||
// 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) {
|
||||
payload, err := GetAndReplaceReqBody(r)
|
||||
payload, err := awsutils.GetAndReplaceReqBody(r)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
@@ -35,6 +35,7 @@ import (
|
||||
|
||||
"github.com/gravitational/teleport/lib/auth"
|
||||
"github.com/gravitational/teleport/lib/tlsca"
|
||||
awsutils "github.com/gravitational/teleport/lib/utils/aws"
|
||||
)
|
||||
|
||||
// TestAWSSignerHandler test the AWS SigningService APP handler logic with mocked STS signing credentials.
|
||||
@@ -117,7 +118,7 @@ func TestAWSSignerHandler(t *testing.T) {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
handler := func(writer http.ResponseWriter, request *http.Request) {
|
||||
require.Equal(t, tc.wantHost, request.Host)
|
||||
awsAuthHeader, err := ParseSigV4(request.Header.Get(authorizationHeader))
|
||||
awsAuthHeader, err := awsutils.ParseSigV4(request.Header.Get(awsutils.AuthorizationHeader))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tc.wantAuthCredRegion, awsAuthHeader.Region)
|
||||
require.Equal(t, tc.wantAuthCredKeyID, awsAuthHeader.KeyID)
|
||||
|
||||
@@ -41,6 +41,7 @@ import (
|
||||
appaws "github.com/gravitational/teleport/lib/srv/app/aws"
|
||||
"github.com/gravitational/teleport/lib/tlsca"
|
||||
"github.com/gravitational/teleport/lib/utils"
|
||||
"github.com/gravitational/teleport/lib/utils/aws"
|
||||
|
||||
"github.com/gravitational/trace"
|
||||
|
||||
@@ -591,7 +592,7 @@ func (s *Server) serveHTTP(w http.ResponseWriter, r *http.Request) error {
|
||||
// access from AWS CLI where the request is already singed by the AWS Signature Version 4 algorithm.
|
||||
// AWS CLI, automatically use SigV4 for all services that support it (All services expect Amazon SimpleDB
|
||||
// but this AWS service has been deprecated)
|
||||
if appaws.IsSignedByAWSSigV4(r) && app.IsAWSConsole() {
|
||||
if aws.IsSignedByAWSSigV4(r) && app.IsAWSConsole() {
|
||||
// Sign the request based on RouteToApp.AWSRoleARN user identity and route signed request to the AWS API.
|
||||
s.awsSigner.Handle(w, r)
|
||||
return nil
|
||||
|
||||
@@ -1,61 +0,0 @@
|
||||
/*
|
||||
Copyright 2021 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 utils
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/aws/aws-sdk-go/aws/arn"
|
||||
)
|
||||
|
||||
// FilterAWSRoles returns role ARNs from the provided list that belong to the
|
||||
// specified AWS account ID.
|
||||
//
|
||||
// If AWS account ID is empty, all roles are returned.
|
||||
func FilterAWSRoles(arns []string, accountID string) (result []AWSRole) {
|
||||
for _, roleARN := range arns {
|
||||
parsed, err := arn.Parse(roleARN)
|
||||
if err != nil || (accountID != "" && parsed.AccountID != accountID) {
|
||||
continue
|
||||
}
|
||||
|
||||
// In AWS convention, the display of the role is the last
|
||||
// /-delineated substring.
|
||||
//
|
||||
// Example ARNs:
|
||||
// arn:aws:iam::1234567890:role/EC2FullAccess (display: EC2FullAccess)
|
||||
// arn:aws:iam::1234567890:role/path/to/customrole (display: customrole)
|
||||
parts := strings.Split(parsed.Resource, "/")
|
||||
numParts := len(parts)
|
||||
if numParts < 2 || parts[0] != "role" {
|
||||
continue
|
||||
}
|
||||
result = append(result, AWSRole{
|
||||
Display: parts[numParts-1],
|
||||
ARN: roleARN,
|
||||
})
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// AWSRole describes an AWS IAM role for AWS console access.
|
||||
type AWSRole struct {
|
||||
// Display is the role display name.
|
||||
Display string `json:"display"`
|
||||
// ARN is the full role ARN.
|
||||
ARN string `json:"arn"`
|
||||
}
|
||||
@@ -25,6 +25,7 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/aws/aws-sdk-go/aws/arn"
|
||||
"github.com/aws/aws-sdk-go/aws/credentials"
|
||||
v4 "github.com/aws/aws-sdk-go/aws/signer/v4"
|
||||
"github.com/gravitational/trace"
|
||||
@@ -45,7 +46,7 @@ const (
|
||||
// https://docs.aws.amazon.com/general/latest/gr/sigv4-date-handling.html
|
||||
AmzDateHeader = "X-Amz-Date"
|
||||
|
||||
authorizationHeader = "Authorization"
|
||||
AuthorizationHeader = "Authorization"
|
||||
credentialAuthHeaderElem = "Credential"
|
||||
signedHeaderAuthHeaderElem = "SignedHeaders"
|
||||
signatureAuthHeaderElem = "Signature"
|
||||
@@ -118,7 +119,7 @@ func ParseSigV4(header string) (*SigV4, error) {
|
||||
// IsSignedByAWSSigV4 checks is the request was signed by AWS Signature Version 4 algorithm.
|
||||
// https://docs.aws.amazon.com/general/latest/gr/signing_aws_api_requests.html
|
||||
func IsSignedByAWSSigV4(r *http.Request) bool {
|
||||
return strings.HasPrefix(r.Header.Get(authorizationHeader), AmazonSigV4AuthorizationPrefix)
|
||||
return strings.HasPrefix(r.Header.Get(AuthorizationHeader), AmazonSigV4AuthorizationPrefix)
|
||||
}
|
||||
|
||||
// GetAndReplaceReqBody returns the request and replace the drained body reader with io.NopCloser
|
||||
@@ -208,3 +209,41 @@ func filterHeaders(r *http.Request, headers []string) {
|
||||
}
|
||||
r.Header = out
|
||||
}
|
||||
|
||||
// FilterAWSRoles returns role ARNs from the provided list that belong to the
|
||||
// specified AWS account ID.
|
||||
//
|
||||
// If AWS account ID is empty, all roles are returned.
|
||||
func FilterAWSRoles(arns []string, accountID string) (result []AWSRole) {
|
||||
for _, roleARN := range arns {
|
||||
parsed, err := arn.Parse(roleARN)
|
||||
if err != nil || (accountID != "" && parsed.AccountID != accountID) {
|
||||
continue
|
||||
}
|
||||
|
||||
// In AWS convention, the display of the role is the last
|
||||
// /-delineated substring.
|
||||
//
|
||||
// Example ARNs:
|
||||
// arn:aws:iam::1234567890:role/EC2FullAccess (display: EC2FullAccess)
|
||||
// arn:aws:iam::1234567890:role/path/to/customrole (display: customrole)
|
||||
parts := strings.Split(parsed.Resource, "/")
|
||||
numParts := len(parts)
|
||||
if numParts < 2 || parts[0] != "role" {
|
||||
continue
|
||||
}
|
||||
result = append(result, AWSRole{
|
||||
Display: parts[numParts-1],
|
||||
ARN: roleARN,
|
||||
})
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// AWSRole describes an AWS IAM role for AWS console access.
|
||||
type AWSRole struct {
|
||||
// Display is the role display name.
|
||||
Display string `json:"display"`
|
||||
// ARN is the full role ARN.
|
||||
ARN string `json:"arn"`
|
||||
}
|
||||
@@ -89,3 +89,53 @@ func TestExtractCredFromAuthHeader(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestFilterAWSRoles verifies filtering AWS role ARNs by AWS account ID.
|
||||
func TestFilterAWSRoles(t *testing.T) {
|
||||
acc1ARN1 := AWSRole{
|
||||
ARN: "arn:aws:iam::1234567890:role/EC2FullAccess",
|
||||
Display: "EC2FullAccess",
|
||||
}
|
||||
acc1ARN2 := AWSRole{
|
||||
ARN: "arn:aws:iam::1234567890:role/EC2ReadOnly",
|
||||
Display: "EC2ReadOnly",
|
||||
}
|
||||
acc1ARN3 := AWSRole{
|
||||
ARN: "arn:aws:iam::1234567890:role/path/to/customrole",
|
||||
Display: "customrole",
|
||||
}
|
||||
acc2ARN1 := AWSRole{
|
||||
ARN: "arn:aws:iam::0987654321:role/test-role",
|
||||
Display: "test-role",
|
||||
}
|
||||
invalidARN := AWSRole{
|
||||
ARN: "invalid-arn",
|
||||
}
|
||||
allARNS := []string{
|
||||
acc1ARN1.ARN, acc1ARN2.ARN, acc1ARN3.ARN, acc2ARN1.ARN, invalidARN.ARN,
|
||||
}
|
||||
tests := []struct {
|
||||
name string
|
||||
accountID string
|
||||
outARNs []AWSRole
|
||||
}{
|
||||
{
|
||||
name: "first account roles",
|
||||
accountID: "1234567890",
|
||||
outARNs: []AWSRole{acc1ARN1, acc1ARN2, acc1ARN3},
|
||||
},
|
||||
{
|
||||
name: "second account roles",
|
||||
accountID: "0987654321",
|
||||
outARNs: []AWSRole{acc2ARN1},
|
||||
},
|
||||
{
|
||||
name: "all roles",
|
||||
accountID: "",
|
||||
outARNs: []AWSRole{acc1ARN1, acc1ARN2, acc1ARN3, acc2ARN1},
|
||||
},
|
||||
}
|
||||
for _, test := range tests {
|
||||
require.Equal(t, test.outARNs, FilterAWSRoles(allARNS, test.accountID))
|
||||
}
|
||||
}
|
||||
@@ -1,73 +0,0 @@
|
||||
/*
|
||||
Copyright 2021 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 utils
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestFilterAWSRoles verifies filtering AWS role ARNs by AWS account ID.
|
||||
func TestFilterAWSRoles(t *testing.T) {
|
||||
acc1ARN1 := AWSRole{
|
||||
ARN: "arn:aws:iam::1234567890:role/EC2FullAccess",
|
||||
Display: "EC2FullAccess",
|
||||
}
|
||||
acc1ARN2 := AWSRole{
|
||||
ARN: "arn:aws:iam::1234567890:role/EC2ReadOnly",
|
||||
Display: "EC2ReadOnly",
|
||||
}
|
||||
acc1ARN3 := AWSRole{
|
||||
ARN: "arn:aws:iam::1234567890:role/path/to/customrole",
|
||||
Display: "customrole",
|
||||
}
|
||||
acc2ARN1 := AWSRole{
|
||||
ARN: "arn:aws:iam::0987654321:role/test-role",
|
||||
Display: "test-role",
|
||||
}
|
||||
invalidARN := AWSRole{
|
||||
ARN: "invalid-arn",
|
||||
}
|
||||
allARNS := []string{
|
||||
acc1ARN1.ARN, acc1ARN2.ARN, acc1ARN3.ARN, acc2ARN1.ARN, invalidARN.ARN,
|
||||
}
|
||||
tests := []struct {
|
||||
name string
|
||||
accountID string
|
||||
outARNs []AWSRole
|
||||
}{
|
||||
{
|
||||
name: "first account roles",
|
||||
accountID: "1234567890",
|
||||
outARNs: []AWSRole{acc1ARN1, acc1ARN2, acc1ARN3},
|
||||
},
|
||||
{
|
||||
name: "second account roles",
|
||||
accountID: "0987654321",
|
||||
outARNs: []AWSRole{acc2ARN1},
|
||||
},
|
||||
{
|
||||
name: "all roles",
|
||||
accountID: "",
|
||||
outARNs: []AWSRole{acc1ARN1, acc1ARN2, acc1ARN3, acc2ARN1},
|
||||
},
|
||||
}
|
||||
for _, test := range tests {
|
||||
require.Equal(t, test.outARNs, FilterAWSRoles(allARNS, test.accountID))
|
||||
}
|
||||
}
|
||||
@@ -2318,7 +2318,7 @@ func (h *Handler) hostCredentials(w http.ResponseWriter, r *http.Request, p http
|
||||
}
|
||||
|
||||
authClient := h.cfg.ProxyClient
|
||||
certs, err := authClient.RegisterUsingToken(req)
|
||||
certs, err := authClient.RegisterUsingToken(r.Context(), &req)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -2397,7 +2397,7 @@ func (h *Handler) validateTrustedCluster(w http.ResponseWriter, r *http.Request,
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
validateResponse, err := h.auth.ValidateTrustedCluster(validateRequest)
|
||||
validateResponse, err := h.auth.ValidateTrustedCluster(r.Context(), validateRequest)
|
||||
if err != nil {
|
||||
h.log.WithError(err).Error("Failed validating trusted cluster")
|
||||
if trace.IsAccessDenied(err) {
|
||||
|
||||
+2
-2
@@ -666,8 +666,8 @@ func (s *sessionCache) Ping(ctx context.Context) (proto.PingResponse, error) {
|
||||
return s.proxyClient.Ping(ctx)
|
||||
}
|
||||
|
||||
func (s *sessionCache) ValidateTrustedCluster(validateRequest *auth.ValidateTrustedClusterRequest) (*auth.ValidateTrustedClusterResponse, error) {
|
||||
return s.proxyClient.ValidateTrustedCluster(validateRequest)
|
||||
func (s *sessionCache) ValidateTrustedCluster(ctx context.Context, validateRequest *auth.ValidateTrustedClusterRequest) (*auth.ValidateTrustedClusterResponse, error) {
|
||||
return s.proxyClient.ValidateTrustedCluster(ctx, validateRequest)
|
||||
}
|
||||
|
||||
// validateSession validates the session given with user and session ID.
|
||||
|
||||
+3
-3
@@ -22,7 +22,7 @@ import (
|
||||
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/teleport/lib/tlsca"
|
||||
"github.com/gravitational/teleport/lib/utils"
|
||||
"github.com/gravitational/teleport/lib/utils/aws"
|
||||
)
|
||||
|
||||
// App describes an application
|
||||
@@ -44,7 +44,7 @@ type App struct {
|
||||
// AWSConsole if true, indicates that the app represents AWS management console.
|
||||
AWSConsole bool `json:"awsConsole"`
|
||||
// AWSRoles is a list of AWS IAM roles for the application representing AWS console.
|
||||
AWSRoles []utils.AWSRole `json:"awsRoles,omitempty"`
|
||||
AWSRoles []aws.AWSRole `json:"awsRoles,omitempty"`
|
||||
}
|
||||
|
||||
// MakeAppsConfig contains parameters for converting apps to UI representation.
|
||||
@@ -88,7 +88,7 @@ func MakeApps(c MakeAppsConfig) []App {
|
||||
}
|
||||
|
||||
if teleApp.IsAWSConsole() {
|
||||
app.AWSRoles = utils.FilterAWSRoles(c.Identity.AWSRoleARNs,
|
||||
app.AWSRoles = aws.FilterAWSRoles(c.Identity.AWSRoleARNs,
|
||||
teleApp.GetAWSAccountID())
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user