Unify auth server receiver names (#33968)

There was a mix of `a` and `s` being used as the receiver names
for auth.Server. This converts them all to use `a` to eliminate
warnings in editors.
This commit is contained in:
rosstimothy
2023-10-27 15:19:22 +00:00
committed by GitHub
parent 9e4ca3cd9f
commit 5d89de6024
11 changed files with 338 additions and 338 deletions
+65 -65
View File
@@ -70,8 +70,8 @@ const (
var fakeRecoveryCodeHash = []byte(`$2a$10$c2.h4pF9AA25lbrWo6U0D.ZmnYpFDaNzN3weNNYNC3jAkYEX9kpzu`)
// StartAccountRecovery implements AuthService.StartAccountRecovery.
func (s *Server) StartAccountRecovery(ctx context.Context, req *proto.StartAccountRecoveryRequest) (types.UserToken, error) {
if err := s.isAccountRecoveryAllowed(ctx); err != nil {
func (a *Server) StartAccountRecovery(ctx context.Context, req *proto.StartAccountRecoveryRequest) (types.UserToken, error) {
if err := a.isAccountRecoveryAllowed(ctx); err != nil {
return nil, trace.Wrap(err)
}
@@ -81,17 +81,17 @@ func (s *Server) StartAccountRecovery(ctx context.Context, req *proto.StartAccou
return nil, trace.AccessDenied(startRecoveryGenericErrMsg)
}
if err := s.verifyCodeWithRecoveryLock(ctx, req.GetUsername(), req.GetRecoveryCode()); err != nil {
if err := a.verifyCodeWithRecoveryLock(ctx, req.GetUsername(), req.GetRecoveryCode()); err != nil {
return nil, trace.Wrap(err)
}
// Remove any other existing tokens for this user before creating a token.
if err := s.deleteUserTokens(ctx, req.Username); err != nil {
if err := a.deleteUserTokens(ctx, req.Username); err != nil {
log.Error(trace.DebugReport(err))
return nil, trace.AccessDenied(startRecoveryGenericErrMsg)
}
token, err := s.createRecoveryToken(ctx, req.GetUsername(), UserTokenTypeRecoveryStart, req.GetRecoverType())
token, err := a.createRecoveryToken(ctx, req.GetUsername(), UserTokenTypeRecoveryStart, req.GetRecoverType())
if err != nil {
log.Error(trace.DebugReport(err))
return nil, trace.AccessDenied(startRecoveryGenericErrMsg)
@@ -103,26 +103,26 @@ func (s *Server) StartAccountRecovery(ctx context.Context, req *proto.StartAccou
// verifyCodeWithRecoveryLock counts number of failed attempts at providing a valid recovery code.
// After MaxAccountRecoveryAttempts, user is temporarily locked from further attempts at recovering and also
// locked from logging in. Modeled after existing function WithUserLock.
func (s *Server) verifyCodeWithRecoveryLock(ctx context.Context, username string, recoveryCode []byte) error {
user, err := s.Services.GetUser(ctx, username, false)
func (a *Server) verifyCodeWithRecoveryLock(ctx context.Context, username string, recoveryCode []byte) error {
user, err := a.Services.GetUser(ctx, username, false)
switch {
case trace.IsNotFound(err):
// If user is not found, still authenticate. It should always return an error.
// This prevents username oracles and timing attacks.
return s.verifyRecoveryCode(ctx, username, recoveryCode)
return a.verifyRecoveryCode(ctx, username, recoveryCode)
case err != nil:
log.Error(trace.DebugReport(err))
return trace.AccessDenied(startRecoveryGenericErrMsg)
}
status := user.GetStatus()
if status.IsLocked && status.RecoveryAttemptLockExpires.After(s.clock.Now().UTC()) {
if status.IsLocked && status.RecoveryAttemptLockExpires.After(a.clock.Now().UTC()) {
log.Debugf("%v exceeds %v failed account recovery attempts, locked until %v",
user.GetName(), defaults.MaxAccountRecoveryAttempts, apiutils.HumanTimeFormat(status.RecoveryAttemptLockExpires))
return trace.AccessDenied(startRecoveryMaxFailedAttemptsErrMsg)
}
verifyCodeErr := s.verifyRecoveryCode(ctx, username, recoveryCode)
verifyCodeErr := a.verifyRecoveryCode(ctx, username, recoveryCode)
switch {
case trace.IsConnectionProblem(verifyCodeErr):
return trace.Wrap(verifyCodeErr)
@@ -130,7 +130,7 @@ func (s *Server) verifyCodeWithRecoveryLock(ctx context.Context, username string
return nil
}
lockedUntil, maxedAttempts, err := s.recordFailedRecoveryAttempt(ctx, username)
lockedUntil, maxedAttempts, err := a.recordFailedRecoveryAttempt(ctx, username)
switch {
case err != nil:
log.Error(trace.DebugReport(err))
@@ -141,7 +141,7 @@ func (s *Server) verifyCodeWithRecoveryLock(ctx context.Context, username string
// Temp lock both user login and recovery attempts.
user.SetRecoveryAttemptLockExpires(lockedUntil, accountLockedMsg)
_, err = s.UpsertUser(ctx, user)
_, err = a.UpsertUser(ctx, user)
if err != nil {
log.Error(trace.DebugReport(err))
return trace.Wrap(verifyCodeErr)
@@ -150,8 +150,8 @@ func (s *Server) verifyCodeWithRecoveryLock(ctx context.Context, username string
return trace.AccessDenied(MaxFailedAttemptsFromStartRecoveryErrMsg)
}
func (s *Server) verifyRecoveryCode(ctx context.Context, user string, givenCode []byte) error {
recovery, err := s.GetRecoveryCodes(ctx, user, true /* withSecrets */)
func (a *Server) verifyRecoveryCode(ctx context.Context, user string, givenCode []byte) error {
recovery, err := a.GetRecoveryCodes(ctx, user, true /* withSecrets */)
if err != nil && !trace.IsNotFound(err) {
return trace.Wrap(err)
}
@@ -179,7 +179,7 @@ func (s *Server) verifyRecoveryCode(ctx context.Context, user string, givenCode
codeMatch = true
// Mark matched token as used in backend, so it can't be used again.
recovery.GetCodes()[i].IsUsed = true
if err := s.UpsertRecoveryCodes(ctx, user, recovery); err != nil {
if err := a.UpsertRecoveryCodes(ctx, user, recovery); err != nil {
log.Error(trace.DebugReport(err))
return trace.AccessDenied(startRecoveryGenericErrMsg)
}
@@ -209,14 +209,14 @@ func (s *Server) verifyRecoveryCode(ctx context.Context, user string, givenCode
event.Status.Error = traceErr.Error()
event.Status.UserMessage = traceErr.Error()
if err := s.emitter.EmitAuditEvent(s.closeCtx, event); err != nil {
if err := a.emitter.EmitAuditEvent(a.closeCtx, event); err != nil {
log.WithFields(logrus.Fields{"user": user}).Warn("Failed to emit account recovery code used failed event.")
}
return trace.AccessDenied(startRecoveryBadAuthnErrMsg)
}
if err := s.emitter.EmitAuditEvent(s.closeCtx, event); err != nil {
if err := a.emitter.EmitAuditEvent(a.closeCtx, event); err != nil {
log.WithFields(logrus.Fields{"user": user}).Warn("Failed to emit account recovery code used event.")
}
@@ -224,12 +224,12 @@ func (s *Server) verifyRecoveryCode(ctx context.Context, user string, givenCode
}
// VerifyAccountRecovery implements AuthService.VerifyAccountRecovery.
func (s *Server) VerifyAccountRecovery(ctx context.Context, req *proto.VerifyAccountRecoveryRequest) (types.UserToken, error) {
if err := s.isAccountRecoveryAllowed(ctx); err != nil {
func (a *Server) VerifyAccountRecovery(ctx context.Context, req *proto.VerifyAccountRecoveryRequest) (types.UserToken, error) {
if err := a.isAccountRecoveryAllowed(ctx); err != nil {
return nil, trace.Wrap(err)
}
startToken, err := s.GetUserToken(ctx, req.GetRecoveryStartTokenID())
startToken, err := a.GetUserToken(ctx, req.GetRecoveryStartTokenID())
switch {
case err != nil:
return nil, trace.AccessDenied(verifyRecoveryGenericErrMsg)
@@ -237,7 +237,7 @@ func (s *Server) VerifyAccountRecovery(ctx context.Context, req *proto.VerifyAcc
return nil, trace.AccessDenied(verifyRecoveryBadAuthnErrMsg)
}
if err := s.verifyUserToken(startToken, UserTokenTypeRecoveryStart); err != nil {
if err := a.verifyUserToken(startToken, UserTokenTypeRecoveryStart); err != nil {
return nil, trace.Wrap(err)
}
@@ -249,8 +249,8 @@ func (s *Server) VerifyAccountRecovery(ctx context.Context, req *proto.VerifyAcc
return nil, trace.AccessDenied(verifyRecoveryBadAuthnErrMsg)
}
if err := s.verifyAuthnWithRecoveryLock(ctx, startToken, func() error {
return s.checkPasswordWOToken(startToken.GetUser(), req.GetPassword())
if err := a.verifyAuthnWithRecoveryLock(ctx, startToken, func() error {
return a.checkPasswordWOToken(startToken.GetUser(), req.GetPassword())
}); err != nil {
return nil, trace.Wrap(err)
}
@@ -261,8 +261,8 @@ func (s *Server) VerifyAccountRecovery(ctx context.Context, req *proto.VerifyAcc
return nil, trace.AccessDenied(verifyRecoveryBadAuthnErrMsg)
}
if err := s.verifyAuthnWithRecoveryLock(ctx, startToken, func() error {
_, _, err := s.validateMFAAuthResponse(
if err := a.verifyAuthnWithRecoveryLock(ctx, startToken, func() error {
_, _, err := a.validateMFAAuthResponse(
ctx, req.GetMFAAuthenticateResponse(), startToken.GetUser(), false /* passwordless */)
return err
}); err != nil {
@@ -273,13 +273,13 @@ func (s *Server) VerifyAccountRecovery(ctx context.Context, req *proto.VerifyAcc
return nil, trace.AccessDenied("unsupported authentication method")
}
approvedToken, err := s.createRecoveryToken(ctx, startToken.GetUser(), UserTokenTypeRecoveryApproved, startToken.GetUsage())
approvedToken, err := a.createRecoveryToken(ctx, startToken.GetUser(), UserTokenTypeRecoveryApproved, startToken.GetUsage())
if err != nil {
return nil, trace.AccessDenied(verifyRecoveryGenericErrMsg)
}
// Delete start token to invalidate the recovery link sent to users.
if err := s.DeleteUserToken(ctx, startToken.GetName()); err != nil {
if err := a.DeleteUserToken(ctx, startToken.GetName()); err != nil {
log.Error(trace.DebugReport(err))
}
@@ -289,10 +289,10 @@ func (s *Server) VerifyAccountRecovery(ctx context.Context, req *proto.VerifyAcc
// verifyAuthnWithRecoveryLock counts number of failed attempts at providing a valid password or second factor.
// After MaxAccountRecoveryAttempts, user's account is temporarily locked from logging in, recovery attempts are reset,
// and all user's tokens are deleted. Modeled after existing function WithUserLock.
func (s *Server) verifyAuthnWithRecoveryLock(ctx context.Context, startToken types.UserToken, authenticateFn func() error) error {
func (a *Server) verifyAuthnWithRecoveryLock(ctx context.Context, startToken types.UserToken, authenticateFn func() error) error {
// Determine user exists first since an existence of token
// does not guarantee the user defined in token exists anymore.
user, err := s.Services.GetUser(ctx, startToken.GetUser(), false)
user, err := a.Services.GetUser(ctx, startToken.GetUser(), false)
if err != nil {
log.Error(trace.DebugReport(err))
return trace.AccessDenied(verifyRecoveryGenericErrMsg)
@@ -308,7 +308,7 @@ func (s *Server) verifyAuthnWithRecoveryLock(ctx context.Context, startToken typ
case verifyAuthnErr == nil:
// Reset attempt counter.
if err := s.DeleteUserRecoveryAttempts(ctx, startToken.GetUser()); err != nil {
if err := a.DeleteUserRecoveryAttempts(ctx, startToken.GetUser()); err != nil {
log.Error(trace.DebugReport(err))
}
@@ -317,7 +317,7 @@ func (s *Server) verifyAuthnWithRecoveryLock(ctx context.Context, startToken typ
log.Error(trace.DebugReport(verifyAuthnErr))
lockedUntil, maxedAttempts, err := s.recordFailedRecoveryAttempt(ctx, startToken.GetUser())
lockedUntil, maxedAttempts, err := a.recordFailedRecoveryAttempt(ctx, startToken.GetUser())
switch {
case err != nil:
log.Error(trace.DebugReport(err))
@@ -327,20 +327,20 @@ func (s *Server) verifyAuthnWithRecoveryLock(ctx context.Context, startToken typ
}
// Delete all tokens related to this user, to force user to restart the recovery flow.
if err := s.deleteUserTokens(ctx, startToken.GetUser()); err != nil {
if err := a.deleteUserTokens(ctx, startToken.GetUser()); err != nil {
log.Error(trace.DebugReport(err))
return trace.AccessDenied(verifyRecoveryGenericErrMsg)
}
// Restart the attempt counter, to not block users from trying again with another recovery code.
if err := s.DeleteUserRecoveryAttempts(ctx, startToken.GetUser()); err != nil {
if err := a.DeleteUserRecoveryAttempts(ctx, startToken.GetUser()); err != nil {
log.Error(trace.DebugReport(err))
return trace.AccessDenied(verifyRecoveryGenericErrMsg)
}
// Lock the user from logging in.
user.SetLocked(lockedUntil, accountLockedMsg)
_, err = s.UpsertUser(ctx, user)
_, err = a.UpsertUser(ctx, user)
if err != nil {
log.Error(trace.DebugReport(err))
return trace.AccessDenied(verifyRecoveryBadAuthnErrMsg)
@@ -351,18 +351,18 @@ func (s *Server) verifyAuthnWithRecoveryLock(ctx context.Context, startToken typ
// recordFailedRecoveryAttempt creates and inserts a recovery attempt and if user has reached max failed attempts,
// returns the locked until time. The boolean determines if user reached maxed failed attempts (true) or not (false).
func (s *Server) recordFailedRecoveryAttempt(ctx context.Context, username string) (time.Time, bool, error) {
func (a *Server) recordFailedRecoveryAttempt(ctx context.Context, username string) (time.Time, bool, error) {
maxedAttempts := true
// Record and log failed attempt.
now := s.clock.Now().UTC()
now := a.clock.Now().UTC()
attempt := &types.RecoveryAttempt{Time: now, Expires: now.Add(defaults.AttemptTTL)}
if err := s.CreateUserRecoveryAttempt(ctx, username, attempt); err != nil {
if err := a.CreateUserRecoveryAttempt(ctx, username, attempt); err != nil {
return time.Time{}, !maxedAttempts, trace.Wrap(err)
}
// Collect all attempts.
attempts, err := s.GetUserRecoveryAttempts(ctx, username)
attempts, err := a.GetUserRecoveryAttempts(ctx, username)
if err != nil {
return time.Time{}, !maxedAttempts, trace.Wrap(err)
}
@@ -373,7 +373,7 @@ func (s *Server) recordFailedRecoveryAttempt(ctx context.Context, username strin
}
// At this point, user has reached max attempts.
lockUntil := s.clock.Now().UTC().Add(defaults.AccountLockInterval)
lockUntil := a.clock.Now().UTC().Add(defaults.AccountLockInterval)
log.Debugf("%v exceeds %v failed account recovery attempts, account locked until %v and an email has been sent",
username, defaults.MaxAccountRecoveryAttempts, apiutils.HumanTimeFormat(lockUntil))
@@ -381,18 +381,18 @@ func (s *Server) recordFailedRecoveryAttempt(ctx context.Context, username strin
}
// CompleteAccountRecovery implements AuthService.CompleteAccountRecovery.
func (s *Server) CompleteAccountRecovery(ctx context.Context, req *proto.CompleteAccountRecoveryRequest) error {
if err := s.isAccountRecoveryAllowed(ctx); err != nil {
func (a *Server) CompleteAccountRecovery(ctx context.Context, req *proto.CompleteAccountRecoveryRequest) error {
if err := a.isAccountRecoveryAllowed(ctx); err != nil {
return trace.Wrap(err)
}
approvedToken, err := s.GetUserToken(ctx, req.GetRecoveryApprovedTokenID())
approvedToken, err := a.GetUserToken(ctx, req.GetRecoveryApprovedTokenID())
if err != nil {
log.Error(trace.DebugReport(err))
return trace.AccessDenied(completeRecoveryGenericErrMsg)
}
if err := s.verifyUserToken(approvedToken, UserTokenTypeRecoveryApproved); err != nil {
if err := a.verifyUserToken(approvedToken, UserTokenTypeRecoveryApproved); err != nil {
return trace.Wrap(err)
}
@@ -408,7 +408,7 @@ func (s *Server) CompleteAccountRecovery(ctx context.Context, req *proto.Complet
return trace.Wrap(err)
}
if err := s.UpsertPassword(approvedToken.GetUser(), req.GetNewPassword()); err != nil {
if err := a.UpsertPassword(approvedToken.GetUser(), req.GetNewPassword()); err != nil {
log.Error(trace.DebugReport(err))
return trace.AccessDenied(completeRecoveryGenericErrMsg)
}
@@ -419,7 +419,7 @@ func (s *Server) CompleteAccountRecovery(ctx context.Context, req *proto.Complet
return trace.AccessDenied(completeRecoveryGenericErrMsg)
}
_, err = s.verifyMFARespAndAddDevice(ctx, &newMFADeviceFields{
_, err = a.verifyMFARespAndAddDevice(ctx, &newMFADeviceFields{
username: approvedToken.GetUser(),
newDeviceName: req.GetNewDeviceName(),
tokenID: approvedToken.GetName(),
@@ -434,7 +434,7 @@ func (s *Server) CompleteAccountRecovery(ctx context.Context, req *proto.Complet
}
// Check and remove user locks so user can immediately sign in after finishing recovering.
user, err := s.Services.GetUser(ctx, approvedToken.GetUser(), false /* without secrets */)
user, err := a.Services.GetUser(ctx, approvedToken.GetUser(), false /* without secrets */)
if err != nil {
log.Error(trace.DebugReport(err))
return trace.AccessDenied(completeRecoveryGenericErrMsg)
@@ -442,13 +442,13 @@ func (s *Server) CompleteAccountRecovery(ctx context.Context, req *proto.Complet
if user.GetStatus().IsLocked {
user.ResetLocks()
_, err = s.UpsertUser(ctx, user)
_, err = a.UpsertUser(ctx, user)
if err != nil {
log.Error(trace.DebugReport(err))
return trace.AccessDenied(completeRecoveryGenericErrMsg)
}
if err := s.DeleteUserLoginAttempts(approvedToken.GetUser()); err != nil {
if err := a.DeleteUserLoginAttempts(approvedToken.GetUser()); err != nil {
log.Error(trace.DebugReport(err))
return trace.AccessDenied(completeRecoveryGenericErrMsg)
}
@@ -458,14 +458,14 @@ func (s *Server) CompleteAccountRecovery(ctx context.Context, req *proto.Complet
}
// CreateAccountRecoveryCodes implements AuthService.CreateAccountRecoveryCodes.
func (s *Server) CreateAccountRecoveryCodes(ctx context.Context, req *proto.CreateAccountRecoveryCodesRequest) (*proto.RecoveryCodes, error) {
func (a *Server) CreateAccountRecoveryCodes(ctx context.Context, req *proto.CreateAccountRecoveryCodesRequest) (*proto.RecoveryCodes, error) {
const unableToCreateCodesMsg = "unable to create new recovery codes, please contact your system administrator"
if err := s.isAccountRecoveryAllowed(ctx); err != nil {
if err := a.isAccountRecoveryAllowed(ctx); err != nil {
return nil, trace.Wrap(err)
}
token, err := s.GetUserToken(ctx, req.GetTokenID())
token, err := a.GetUserToken(ctx, req.GetTokenID())
if err != nil {
log.Error(trace.DebugReport(err))
return nil, trace.AccessDenied(unableToCreateCodesMsg)
@@ -476,11 +476,11 @@ func (s *Server) CreateAccountRecoveryCodes(ctx context.Context, req *proto.Crea
return nil, trace.AccessDenied(unableToCreateCodesMsg)
}
if err := s.verifyUserToken(token, UserTokenTypeRecoveryApproved, UserTokenTypePrivilege); err != nil {
if err := a.verifyUserToken(token, UserTokenTypeRecoveryApproved, UserTokenTypePrivilege); err != nil {
return nil, trace.Wrap(err)
}
newRecovery, err := s.generateAndUpsertRecoveryCodes(ctx, token.GetUser())
newRecovery, err := a.generateAndUpsertRecoveryCodes(ctx, token.GetUser())
if err != nil {
log.Error(trace.DebugReport(err))
return nil, trace.AccessDenied(unableToCreateCodesMsg)
@@ -488,7 +488,7 @@ func (s *Server) CreateAccountRecoveryCodes(ctx context.Context, req *proto.Crea
// If used as part of the recovery flow, getting new recovery codes marks the end of the flow in the UI.
if token.GetSubKind() == UserTokenTypeRecoveryApproved {
if err := s.deleteUserTokens(ctx, token.GetUser()); err != nil {
if err := a.deleteUserTokens(ctx, token.GetUser()); err != nil {
log.Error(trace.DebugReport(err))
}
}
@@ -497,14 +497,14 @@ func (s *Server) CreateAccountRecoveryCodes(ctx context.Context, req *proto.Crea
}
// GetAccountRecoveryToken implements AuthService.GetAccountRecoveryToken.
func (s *Server) GetAccountRecoveryToken(ctx context.Context, req *proto.GetAccountRecoveryTokenRequest) (types.UserToken, error) {
token, err := s.GetUserToken(ctx, req.GetRecoveryTokenID())
func (a *Server) GetAccountRecoveryToken(ctx context.Context, req *proto.GetAccountRecoveryTokenRequest) (types.UserToken, error) {
token, err := a.GetUserToken(ctx, req.GetRecoveryTokenID())
if err != nil {
log.Error(trace.DebugReport(err))
return nil, trace.AccessDenied("access denied")
}
if err := s.verifyUserToken(token, UserTokenTypeRecoveryStart, UserTokenTypeRecoveryApproved); err != nil {
if err := a.verifyUserToken(token, UserTokenTypeRecoveryStart, UserTokenTypeRecoveryApproved); err != nil {
return nil, trace.Wrap(err)
}
@@ -512,13 +512,13 @@ func (s *Server) GetAccountRecoveryToken(ctx context.Context, req *proto.GetAcco
}
// GetAccountRecoveryCodes implements AuthService.GetAccountRecoveryCodes.
func (s *Server) GetAccountRecoveryCodes(ctx context.Context, req *proto.GetAccountRecoveryCodesRequest) (*proto.RecoveryCodes, error) {
func (a *Server) GetAccountRecoveryCodes(ctx context.Context, req *proto.GetAccountRecoveryCodesRequest) (*proto.RecoveryCodes, error) {
username, err := authz.GetClientUsername(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
rc, err := s.GetRecoveryCodes(ctx, username, false /* without secrets */)
rc, err := a.GetRecoveryCodes(ctx, username, false /* without secrets */)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -528,7 +528,7 @@ func (s *Server) GetAccountRecoveryCodes(ctx context.Context, req *proto.GetAcco
}, nil
}
func (s *Server) generateAndUpsertRecoveryCodes(ctx context.Context, username string) (*proto.RecoveryCodes, error) {
func (a *Server) generateAndUpsertRecoveryCodes(ctx context.Context, username string) (*proto.RecoveryCodes, error) {
codes, err := generateRecoveryCodes()
if err != nil {
return nil, trace.Wrap(err)
@@ -544,16 +544,16 @@ func (s *Server) generateAndUpsertRecoveryCodes(ctx context.Context, username st
hashedCodes[i].HashedCode = hashedCode
}
rc, err := types.NewRecoveryCodes(hashedCodes, s.GetClock().Now().UTC(), username)
rc, err := types.NewRecoveryCodes(hashedCodes, a.GetClock().Now().UTC(), username)
if err != nil {
return nil, trace.Wrap(err)
}
if err := s.UpsertRecoveryCodes(ctx, username, rc); err != nil {
if err := a.UpsertRecoveryCodes(ctx, username, rc); err != nil {
return nil, trace.Wrap(err)
}
if err := s.emitter.EmitAuditEvent(s.closeCtx, &apievents.RecoveryCodeGenerate{
if err := a.emitter.EmitAuditEvent(a.closeCtx, &apievents.RecoveryCodeGenerate{
Metadata: apievents.Metadata{
Type: events.RecoveryCodeGeneratedEvent,
Code: events.RecoveryCodesGenerateCode,
@@ -571,12 +571,12 @@ func (s *Server) generateAndUpsertRecoveryCodes(ctx context.Context, username st
// isAccountRecoveryAllowed gets cluster auth configuration and check if cloud, local auth
// and second factor is allowed, which are required for account recovery.
func (s *Server) isAccountRecoveryAllowed(ctx context.Context) error {
func (a *Server) isAccountRecoveryAllowed(ctx context.Context) error {
if !modules.GetModules().Features().RecoveryCodes {
return trace.AccessDenied("account recovery is only available for Teleport enterprise")
}
authPref, err := s.GetAuthPreference(ctx)
authPref, err := a.GetAuthPreference(ctx)
if err != nil {
return trace.Wrap(err)
}
+40 -40
View File
@@ -145,7 +145,7 @@ func ulsFromUser(user types.User) (*userloginstate.UserLoginState, error) {
}
// createBot creates a new certificate renewal bot from a bot request.
func (s *Server) createBot(ctx context.Context, req *proto.CreateBotRequest) (*proto.CreateBotResponse, error) {
func (a *Server) createBot(ctx context.Context, req *proto.CreateBotRequest) (*proto.CreateBotResponse, error) {
if req.Name == "" {
return nil, trace.BadParameter("bot name must not be empty")
}
@@ -155,14 +155,14 @@ func (s *Server) createBot(ctx context.Context, req *proto.CreateBotRequest) (*p
// Ensure conflicting resources don't already exist.
// We skip the cache here to allow for bot recreation shortly after bot
// deletion.
_, err := s.Services.GetRole(ctx, resourceName)
_, err := a.Services.GetRole(ctx, resourceName)
if err != nil && !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
if roleExists := (err == nil); roleExists {
return nil, trace.AlreadyExists("cannot add bot: role %q already exists", resourceName)
}
_, err = s.Services.GetUser(ctx, resourceName, false)
_, err = a.Services.GetUser(ctx, resourceName, false)
if err != nil && !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
@@ -177,23 +177,23 @@ func (s *Server) createBot(ctx context.Context, req *proto.CreateBotRequest) (*p
// Ensure all requested roles exist.
for _, roleName := range req.Roles {
_, err := s.GetRole(ctx, roleName)
_, err := a.GetRole(ctx, roleName)
if err != nil {
return nil, trace.Wrap(err)
}
}
provisionToken, err := s.checkOrCreateBotToken(ctx, req)
provisionToken, err := a.checkOrCreateBotToken(ctx, req)
if err != nil {
return nil, trace.Wrap(err)
}
// Create the resources.
if _, err := createBotRole(ctx, s, req.Name, resourceName, req.Roles); err != nil {
if _, err := createBotRole(ctx, a, req.Name, resourceName, req.Roles); err != nil {
return nil, trace.Wrap(err)
}
if _, err := createBotUser(ctx, s, req.Name, resourceName, req.Traits); err != nil {
if _, err := createBotUser(ctx, a, req.Name, resourceName, req.Traits); err != nil {
return nil, trace.Wrap(err)
}
@@ -203,7 +203,7 @@ func (s *Server) createBot(ctx context.Context, req *proto.CreateBotRequest) (*p
}
// Emit usage analytics event for bot creation.
s.AnonymizeAndSubmit(&usagereporter.BotCreateEvent{
a.AnonymizeAndSubmit(&usagereporter.BotCreateEvent{
UserName: authz.ClientUsername(ctx),
BotUserName: resourceName,
RoleName: resourceName,
@@ -223,8 +223,8 @@ func (s *Server) createBot(ctx context.Context, req *proto.CreateBotRequest) (*p
// deleteBotUser removes an existing bot user, ensuring that it has bot labels
// matching the bot before deleting anything.
func (s *Server) deleteBotUser(ctx context.Context, botName, resourceName string) error {
user, err := s.GetUser(ctx, resourceName, false)
func (a *Server) deleteBotUser(ctx context.Context, botName, resourceName string) error {
user, err := a.GetUser(ctx, resourceName, false)
if err != nil {
return trace.Wrap(err, "could not fetch expected bot user %s", resourceName)
}
@@ -235,7 +235,7 @@ func (s *Server) deleteBotUser(ctx context.Context, botName, resourceName string
} else if label != botName {
err = trace.Errorf("will not delete user %s with mismatched label %s = %s", resourceName, types.BotLabel, label)
} else {
err = s.DeleteUser(ctx, resourceName)
err = a.DeleteUser(ctx, resourceName)
}
return err
@@ -243,8 +243,8 @@ func (s *Server) deleteBotUser(ctx context.Context, botName, resourceName string
// deleteBotRole removes an existing bot role, ensuring that it has bot labels
// matching the bot before deleting anything.
func (s *Server) deleteBotRole(ctx context.Context, botName, resourceName string) error {
role, err := s.GetRole(ctx, resourceName)
func (a *Server) deleteBotRole(ctx context.Context, botName, resourceName string) error {
role, err := a.GetRole(ctx, resourceName)
if err != nil {
return trace.Wrap(err, "could not fetch expected bot role %s", resourceName)
}
@@ -255,13 +255,13 @@ func (s *Server) deleteBotRole(ctx context.Context, botName, resourceName string
} else if label != botName {
err = trace.Errorf("will not delete role %s with mismatched label %s = %s", resourceName, types.BotLabel, label)
} else {
err = s.DeleteRole(ctx, resourceName)
err = a.DeleteRole(ctx, resourceName)
}
return err
}
func (s *Server) deleteBot(ctx context.Context, botName string) error {
func (a *Server) deleteBot(ctx context.Context, botName string) error {
// Note: this does not remove any locks for the bot's user / role. That
// might be convenient in case of accidental bot locking but there doesn't
// seem to be any automatic deletion of locks in teleport today (other
@@ -269,17 +269,17 @@ func (s *Server) deleteBot(ctx context.Context, botName string) error {
// but we can revisit this if desired.
resourceName := BotResourceName(botName)
userErr := s.deleteBotUser(ctx, botName, resourceName)
roleErr := s.deleteBotRole(ctx, botName, resourceName)
userErr := a.deleteBotUser(ctx, botName, resourceName)
roleErr := a.deleteBotRole(ctx, botName, resourceName)
return trace.NewAggregate(userErr, roleErr)
}
// getBotUsers fetches all Users with the BotLabel field set. Users are fetched
// without secrets.
func (s *Server) getBotUsers(ctx context.Context) ([]types.User, error) {
func (a *Server) getBotUsers(ctx context.Context) ([]types.User, error) {
var botUsers []types.User
users, err := s.GetUsers(ctx, false)
users, err := a.GetUsers(ctx, false)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -309,12 +309,12 @@ var supportedBotJoinMethods = []types.JoinMethod{
// checkOrCreateBotToken checks the existing token if given, or creates a new
// random dynamic provision token which allows bots to join with the given
// botName. Returns the token and any error.
func (s *Server) checkOrCreateBotToken(ctx context.Context, req *proto.CreateBotRequest) (types.ProvisionToken, error) {
func (a *Server) checkOrCreateBotToken(ctx context.Context, req *proto.CreateBotRequest) (types.ProvisionToken, error) {
botName := req.Name
// if the request includes a TokenID it should already exist
if req.TokenID != "" {
provisionToken, err := s.GetToken(ctx, req.TokenID)
provisionToken, err := a.GetToken(ctx, req.TokenID)
if err != nil {
if trace.IsNotFound(err) {
return nil, trace.NotFound("token with name %q not found, create the token or do not set TokenName: %v",
@@ -358,12 +358,12 @@ func (s *Server) checkOrCreateBotToken(ctx context.Context, req *proto.CreateBot
JoinMethod: types.JoinMethodToken,
BotName: botName,
}
token, err := types.NewProvisionTokenFromSpec(tokenName, s.clock.Now().Add(ttl), tokenSpec)
token, err := types.NewProvisionTokenFromSpec(tokenName, a.clock.Now().Add(ttl), tokenSpec)
if err != nil {
return nil, trace.Wrap(err)
}
if err := s.UpsertToken(ctx, token); err != nil {
if err := a.UpsertToken(ctx, token); err != nil {
return nil, trace.Wrap(err)
}
@@ -373,10 +373,10 @@ func (s *Server) checkOrCreateBotToken(ctx context.Context, req *proto.CreateBot
}
// validateGenerationLabel validates and updates a generation label.
func (s *Server) validateGenerationLabel(ctx context.Context, username string, certReq *certRequest, currentIdentityGeneration uint64) error {
func (a *Server) validateGenerationLabel(ctx context.Context, username string, certReq *certRequest, currentIdentityGeneration uint64) error {
// Fetch the user, bypassing the cache. We might otherwise fetch a stale
// value in case of a rapid certificate renewal.
user, err := s.Services.GetUser(ctx, username, false)
user, err := a.Services.GetUser(ctx, username, false)
if err != nil {
return trace.Wrap(err)
}
@@ -431,14 +431,14 @@ func (s *Server) validateGenerationLabel(ctx context.Context, username string, c
// Note: we bypass the RBAC check on purpose as bot users should not
// have user update permissions.
if err := s.CompareAndSwapUser(ctx, newUser, user); err != nil {
if err := a.CompareAndSwapUser(ctx, newUser, user); err != nil {
// If this fails it's likely to be some miscellaneous competing
// write. The request should be tried again - if it's malicious,
// someone will get a generation mismatch and trigger a lock.
return trace.CompareFailed("Database comparison failed, try the request again")
}
uls, err := s.GetUserLoginState(ctx, user.GetName())
uls, err := a.GetUserLoginState(ctx, user.GetName())
if err != nil && !trace.IsNotFound(err) {
return trace.Wrap(err)
}
@@ -450,7 +450,7 @@ func (s *Server) validateGenerationLabel(ctx context.Context, username string, c
}
uls.ResourceHeader.Metadata.Labels[types.BotGenerationLabel] = generation
if _, err := s.UpsertUserLoginState(ctx, uls); err != nil {
if _, err := a.UpsertUserLoginState(ctx, uls); err != nil {
return trace.Wrap(err)
}
@@ -472,13 +472,13 @@ func (s *Server) validateGenerationLabel(ctx context.Context, username string, c
if err != nil {
return trace.Wrap(err)
}
if err := s.UpsertLock(ctx, lock); err != nil {
if err := a.UpsertLock(ctx, lock); err != nil {
return trace.Wrap(err)
}
// Emit an audit event.
userMetadata := authz.ClientUserMetadata(ctx)
if err := s.emitter.EmitAuditEvent(s.closeCtx, &apievents.RenewableCertificateGenerationMismatch{
if err := a.emitter.EmitAuditEvent(a.closeCtx, &apievents.RenewableCertificateGenerationMismatch{
Metadata: apievents.Metadata{
Type: events.RenewableCertificateGenerationMismatchEvent,
Code: events.RenewableCertificateGenerationMismatchCode,
@@ -498,7 +498,7 @@ func (s *Server) validateGenerationLabel(ctx context.Context, username string, c
newGeneration := currentIdentityGeneration + 1
// As above, commit some crimes to clone the User.
newUser, err := s.Services.GetUser(ctx, user.GetName(), false)
newUser, err := a.Services.GetUser(ctx, user.GetName(), false)
if err != nil {
return trace.Wrap(err)
}
@@ -506,7 +506,7 @@ func (s *Server) validateGenerationLabel(ctx context.Context, username string, c
metadata.Labels[types.BotGenerationLabel] = fmt.Sprint(newGeneration)
newUser.SetMetadata(metadata)
if err := s.CompareAndSwapUser(ctx, newUser, user); err != nil {
if err := a.CompareAndSwapUser(ctx, newUser, user); err != nil {
// If this fails it's likely to be some miscellaneous competing
// write. The request should be tried again - if it's malicious,
// someone will get a generation mismatch and trigger a lock.
@@ -526,7 +526,7 @@ func (s *Server) validateGenerationLabel(ctx context.Context, username string, c
// care if the current identity is Nop. This function does not validate the
// current identity at all; the caller is expected to validate that the client
// is allowed to issue the (possibly renewable) certificates.
func (s *Server) generateInitialBotCerts(ctx context.Context, username string, pubKey []byte, expires time.Time, renewable bool) (*proto.Certs, error) {
func (a *Server) generateInitialBotCerts(ctx context.Context, username string, pubKey []byte, expires time.Time, renewable bool) (*proto.Certs, error) {
var err error
// Extract the user and role set for whom the certificate will be generated.
@@ -535,7 +535,7 @@ func (s *Server) generateInitialBotCerts(ctx context.Context, username string, p
// This call bypasses RBAC check for users read on purpose.
// Users who are allowed to impersonate other users might not have
// permissions to read user data.
userState, err := s.GetUserOrLoginState(ctx, username)
userState, err := a.GetUserOrLoginState(ctx, username)
if err != nil {
log.WithError(err).Debugf("Could not impersonate user %v. The user could not be fetched from local store.", username)
return nil, trace.AccessDenied("access denied")
@@ -548,17 +548,17 @@ func (s *Server) generateInitialBotCerts(ctx context.Context, username string, p
}
// Cap the cert TTL to the MaxRenewableCertTTL.
if max := s.GetClock().Now().Add(defaults.MaxRenewableCertTTL); expires.After(max) {
if max := a.GetClock().Now().Add(defaults.MaxRenewableCertTTL); expires.After(max) {
expires = max
}
// Inherit the user's roles and traits verbatim.
accessInfo := services.AccessInfoFromUserState(userState)
clusterName, err := s.GetClusterName()
clusterName, err := a.GetClusterName()
if err != nil {
return nil, trace.Wrap(err)
}
checker, err := services.NewAccessChecker(accessInfo, clusterName.GetClusterName(), s)
checker, err := services.NewAccessChecker(accessInfo, clusterName.GetClusterName(), a)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -572,7 +572,7 @@ func (s *Server) generateInitialBotCerts(ctx context.Context, username string, p
// Generate certificate
certReq := certRequest{
user: userState,
ttl: expires.Sub(s.GetClock().Now()),
ttl: expires.Sub(a.GetClock().Now()),
publicKey: pubKey,
checker: checker,
traits: accessInfo.Traits,
@@ -581,11 +581,11 @@ func (s *Server) generateInitialBotCerts(ctx context.Context, username string, p
generation: generation,
}
if err := s.validateGenerationLabel(ctx, userState.GetName(), &certReq, 0); err != nil {
if err := a.validateGenerationLabel(ctx, userState.GetName(), &certReq, 0); err != nil {
return nil, trace.Wrap(err)
}
certs, err := s.generateUserCert(certReq)
certs, err := a.generateUserCert(certReq)
if err != nil {
return nil, trace.Wrap(err)
}
+20 -20
View File
@@ -45,16 +45,16 @@ import (
// GenerateDatabaseCert generates client certificate used by a database
// service to authenticate with the database instance.
func (s *Server) GenerateDatabaseCert(ctx context.Context, req *proto.DatabaseCertRequest) (*proto.DatabaseCertResponse, error) {
func (a *Server) GenerateDatabaseCert(ctx context.Context, req *proto.DatabaseCertRequest) (*proto.DatabaseCertResponse, error) {
csr, err := tlsca.ParseCertificateRequestPEM(req.CSR)
if err != nil {
return nil, trace.Wrap(err)
}
clusterName, err := s.GetClusterName()
clusterName, err := a.GetClusterName()
if err != nil {
return nil, trace.Wrap(err)
}
databaseCA, err := s.GetCertAuthority(ctx, types.CertAuthID{
databaseCA, err := a.GetCertAuthority(ctx, types.CertAuthID{
Type: types.DatabaseCA,
DomainName: clusterName.GetClusterName(),
}, true)
@@ -62,7 +62,7 @@ func (s *Server) GenerateDatabaseCert(ctx context.Context, req *proto.DatabaseCe
if trace.IsNotFound(err) {
// Database CA doesn't exist. Fallback to Host CA.
// https://github.com/gravitational/teleport/issues/5029
databaseCA, err = s.GetCertAuthority(ctx, types.CertAuthID{
databaseCA, err = a.GetCertAuthority(ctx, types.CertAuthID{
Type: types.HostCA,
DomainName: clusterName.GetClusterName(),
}, true)
@@ -71,7 +71,7 @@ func (s *Server) GenerateDatabaseCert(ctx context.Context, req *proto.DatabaseCe
return nil, trace.Wrap(err)
}
}
caCert, signer, err := getCAandSigner(ctx, s.GetKeyStore(), databaseCA, req)
caCert, signer, err := getCAandSigner(ctx, a.GetKeyStore(), databaseCA, req)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -80,10 +80,10 @@ func (s *Server) GenerateDatabaseCert(ctx context.Context, req *proto.DatabaseCe
return nil, trace.Wrap(err)
}
certReq := tlsca.CertificateRequest{
Clock: s.clock,
Clock: a.clock,
PublicKey: csr.PublicKey,
Subject: csr.Subject,
NotAfter: s.clock.Now().UTC().Add(req.TTL.Get()),
NotAfter: a.clock.Now().UTC().Add(req.TTL.Get()),
}
if req.CertificateExtensions == proto.DatabaseCertRequest_WINDOWS_SMARTCARD {
// Pass through ExtKeyUsage (which we need for Smartcard Logon usage)
@@ -134,7 +134,7 @@ func getServerNames(req *proto.DatabaseCertRequest) []string {
// SignDatabaseCSR generates a client certificate used by proxy when talking
// to a remote database service.
func (s *Server) SignDatabaseCSR(ctx context.Context, req *proto.DatabaseCSRRequest) (*proto.DatabaseCSRResponse, error) {
func (a *Server) SignDatabaseCSR(ctx context.Context, req *proto.DatabaseCSRRequest) (*proto.DatabaseCSRResponse, error) {
if !modules.GetModules().Features().DB {
return nil, trace.AccessDenied(
"this Teleport cluster is not licensed for database access, please contact the cluster administrator")
@@ -142,12 +142,12 @@ func (s *Server) SignDatabaseCSR(ctx context.Context, req *proto.DatabaseCSRRequ
log.Debugf("Signing database CSR for cluster %v.", req.ClusterName)
clusterName, err := s.GetClusterName()
clusterName, err := a.GetClusterName()
if err != nil {
return nil, trace.Wrap(err)
}
hostCA, err := s.GetCertAuthority(ctx, types.CertAuthID{
hostCA, err := a.GetCertAuthority(ctx, types.CertAuthID{
Type: types.HostCA,
DomainName: req.ClusterName,
}, false)
@@ -180,7 +180,7 @@ func (s *Server) SignDatabaseCSR(ctx context.Context, req *proto.DatabaseCSRRequ
}
// Extract user roles from the identity.
roles, err := services.FetchRoles(id.Groups, s, id.Traits)
roles, err := services.FetchRoles(id.Groups, a, id.Traits)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -189,7 +189,7 @@ func (s *Server) SignDatabaseCSR(ctx context.Context, req *proto.DatabaseCSRRequ
ttl := roles.AdjustSessionTTL(apidefaults.CertDuration)
// Generate the TLS certificate.
ca, err := s.GetCertAuthority(ctx, types.CertAuthID{
ca, err := a.GetCertAuthority(ctx, types.CertAuthID{
Type: types.DatabaseCA,
DomainName: clusterName.GetClusterName(),
}, true)
@@ -197,7 +197,7 @@ func (s *Server) SignDatabaseCSR(ctx context.Context, req *proto.DatabaseCSRRequ
return nil, trace.Wrap(err)
}
cert, signer, err := s.GetKeyStore().GetTLSCertAndSigner(ctx, ca)
cert, signer, err := a.GetKeyStore().GetTLSCertAndSigner(ctx, ca)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -207,10 +207,10 @@ func (s *Server) SignDatabaseCSR(ctx context.Context, req *proto.DatabaseCSRRequ
}
tlsCert, err := tlsAuthority.GenerateCertificate(tlsca.CertificateRequest{
Clock: s.clock,
Clock: a.clock,
PublicKey: csr.PublicKey,
Subject: subject,
NotAfter: s.clock.Now().UTC().Add(ttl),
NotAfter: a.clock.Now().UTC().Add(ttl),
})
if err != nil {
return nil, trace.Wrap(err)
@@ -223,17 +223,17 @@ func (s *Server) SignDatabaseCSR(ctx context.Context, req *proto.DatabaseCSRRequ
}
// GenerateSnowflakeJWT generates JWT in the format required by Snowflake.
func (s *Server) GenerateSnowflakeJWT(ctx context.Context, req *proto.SnowflakeJWTRequest) (*proto.SnowflakeJWTResponse, error) {
func (a *Server) GenerateSnowflakeJWT(ctx context.Context, req *proto.SnowflakeJWTRequest) (*proto.SnowflakeJWTResponse, error) {
if !modules.GetModules().Features().DB {
return nil, trace.AccessDenied(
"this Teleport cluster is not licensed for database access, please contact the cluster administrator")
}
clusterName, err := s.GetClusterName()
clusterName, err := a.GetClusterName()
if err != nil {
return nil, trace.Wrap(err)
}
ca, err := s.GetCertAuthority(ctx, types.CertAuthID{
ca, err := a.GetCertAuthority(ctx, types.CertAuthID{
Type: types.DatabaseCA,
DomainName: clusterName.GetClusterName(),
}, true)
@@ -264,11 +264,11 @@ func (s *Server) GenerateSnowflakeJWT(ctx context.Context, req *proto.SnowflakeJ
subject, issuer := getSnowflakeJWTParams(req.AccountName, req.UserName, pubKey)
_, signer, err := s.GetKeyStore().GetTLSCertAndSigner(ctx, ca)
_, signer, err := a.GetKeyStore().GetTLSCertAndSigner(ctx, ca)
if err != nil {
return nil, trace.Wrap(err)
}
privateKey, err := services.GetJWTSigner(signer, ca.GetClusterName(), s.clock)
privateKey, err := services.GetJWTSigner(signer, ca.GetClusterName(), a.clock)
if err != nil {
return nil, trace.Wrap(err)
}
+7 -7
View File
@@ -32,7 +32,7 @@ import (
// GenerateWindowsDesktopCert generates client certificate for Windows RDP
// authentication.
func (s *Server) GenerateWindowsDesktopCert(ctx context.Context, req *proto.WindowsDesktopCertRequest) (*proto.WindowsDesktopCertResponse, error) {
func (a *Server) GenerateWindowsDesktopCert(ctx context.Context, req *proto.WindowsDesktopCertRequest) (*proto.WindowsDesktopCertResponse, error) {
if !modules.GetModules().Features().Desktop {
return nil, trace.AccessDenied(
"this Teleport cluster is not licensed for desktop access, please contact the cluster administrator")
@@ -41,18 +41,18 @@ func (s *Server) GenerateWindowsDesktopCert(ctx context.Context, req *proto.Wind
if err != nil {
return nil, trace.Wrap(err)
}
clusterName, err := s.GetClusterName()
clusterName, err := a.GetClusterName()
if err != nil {
return nil, trace.Wrap(err)
}
userCA, err := s.GetCertAuthority(ctx, types.CertAuthID{
userCA, err := a.GetCertAuthority(ctx, types.CertAuthID{
Type: types.UserCA,
DomainName: clusterName.GetClusterName(),
}, true)
if err != nil {
return nil, trace.Wrap(err)
}
caCert, signer, err := s.GetKeyStore().GetTLSCertAndSigner(ctx, userCA)
caCert, signer, err := a.GetKeyStore().GetTLSCertAndSigner(ctx, userCA)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -63,17 +63,17 @@ func (s *Server) GenerateWindowsDesktopCert(ctx context.Context, req *proto.Wind
// See https://docs.microsoft.com/en-us/troubleshoot/windows-server/windows-security/enabling-smart-card-logon-third-party-certification-authorities
// for cert requirements for Windows authn.
certReq := tlsca.CertificateRequest{
Clock: s.clock,
Clock: a.clock,
PublicKey: csr.PublicKey,
Subject: csr.Subject,
NotAfter: s.clock.Now().UTC().Add(req.TTL.Get()),
NotAfter: a.clock.Now().UTC().Add(req.TTL.Get()),
ExtraExtensions: csr.Extensions,
KeyUsage: x509.KeyUsageDigitalSignature,
// CRL is required for Windows smartcard certs.
CRLDistributionPoints: []string{req.CRLEndpoint},
}
limitExceeded, err := s.desktopsLimitExceeded(ctx)
limitExceeded, err := a.desktopsLimitExceeded(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
+9 -9
View File
@@ -61,7 +61,7 @@ type KubeCSRResponse struct {
// ProcessKubeCSR processes CSR request against Kubernetes CA, returns
// signed certificate if successful.
func (s *Server) ProcessKubeCSR(req KubeCSR) (*KubeCSRResponse, error) {
func (a *Server) ProcessKubeCSR(req KubeCSR) (*KubeCSRResponse, error) {
ctx := context.TODO()
if err := enforceLicense(types.KindKubernetesCluster); err != nil {
return nil, trace.Wrap(err)
@@ -70,7 +70,7 @@ func (s *Server) ProcessKubeCSR(req KubeCSR) (*KubeCSRResponse, error) {
return nil, trace.Wrap(err)
}
clusterName, err := s.GetClusterName()
clusterName, err := a.GetClusterName()
if err != nil {
return nil, trace.Wrap(err)
}
@@ -79,7 +79,7 @@ func (s *Server) ProcessKubeCSR(req KubeCSR) (*KubeCSRResponse, error) {
// with special provisions.
log.Debugf("Generating certificate to access remote Kubernetes clusters.")
hostCA, err := s.GetCertAuthority(ctx, types.CertAuthID{
hostCA, err := a.GetCertAuthority(ctx, types.CertAuthID{
Type: types.HostCA,
DomainName: req.ClusterName,
}, false)
@@ -109,7 +109,7 @@ func (s *Server) ProcessKubeCSR(req KubeCSR) (*KubeCSRResponse, error) {
roleNames := id.Groups
// This is a remote user, map roles to local roles first.
if id.TeleportCluster != clusterName.GetClusterName() {
ca, err := s.GetCertAuthority(ctx, types.CertAuthID{Type: types.UserCA, DomainName: id.TeleportCluster}, false)
ca, err := a.GetCertAuthority(ctx, types.CertAuthID{Type: types.UserCA, DomainName: id.TeleportCluster}, false)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -123,14 +123,14 @@ func (s *Server) ProcessKubeCSR(req KubeCSR) (*KubeCSRResponse, error) {
}
// Extract user roles from the identity (from the CSR Subject).
roles, err := services.FetchRoles(roleNames, s, id.Traits)
roles, err := services.FetchRoles(roleNames, a, id.Traits)
if err != nil {
return nil, trace.Wrap(err)
}
// Get the correct cert TTL based on roles.
ttl := roles.AdjustSessionTTL(apidefaults.CertDuration)
userCA, err := s.GetCertAuthority(ctx, types.CertAuthID{
userCA, err := a.GetCertAuthority(ctx, types.CertAuthID{
Type: types.UserCA,
DomainName: clusterName.GetClusterName(),
}, true)
@@ -138,7 +138,7 @@ func (s *Server) ProcessKubeCSR(req KubeCSR) (*KubeCSRResponse, error) {
return nil, trace.Wrap(err)
}
// generate TLS certificate
cert, signer, err := s.GetKeyStore().GetTLSCertAndSigner(ctx, userCA)
cert, signer, err := a.GetKeyStore().GetTLSCertAndSigner(ctx, userCA)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -148,14 +148,14 @@ func (s *Server) ProcessKubeCSR(req KubeCSR) (*KubeCSRResponse, error) {
}
certRequest := tlsca.CertificateRequest{
Clock: s.clock,
Clock: a.clock,
PublicKey: csr.PublicKey,
// Always trust the Subject sent by the proxy (minus the Usage field).
// A user may have received temporary extra roles via workflow API, we
// must preserve those. The storage backend doesn't record temporary
// granted roles.
Subject: subject,
NotAfter: s.clock.Now().UTC().Add(ttl),
NotAfter: a.clock.Now().UTC().Add(ttl),
}
tlsCert, err := tlsAuthority.GenerateCertificate(certRequest)
if err != nil {
+56 -56
View File
@@ -113,13 +113,13 @@ type SessionCreds struct {
// AuthenticateUser authenticates user based on the request type.
// Returns the username of the authenticated user.
func (s *Server) AuthenticateUser(ctx context.Context, req AuthenticateUserRequest) (services.UserState, services.AccessChecker, error) {
func (a *Server) AuthenticateUser(ctx context.Context, req AuthenticateUserRequest) (services.UserState, services.AccessChecker, error) {
username := req.Username
mfaDev, actualUsername, err := s.authenticateUser(ctx, req)
mfaDev, actualUsername, err := a.authenticateUser(ctx, req)
if err != nil {
// Log event after authentication failure
if err := s.emitAuthAuditEvent(ctx, authAuditProps{
if err := a.emitAuthAuditEvent(ctx, authAuditProps{
username: req.Username,
clientMetadata: req.ClientMetadata,
authErr: err,
@@ -137,7 +137,7 @@ func (s *Server) AuthenticateUser(ctx context.Context, req AuthenticateUserReque
username = actualUsername
}
user, err := s.GetUser(ctx, username, false /* withSecrets */)
user, err := a.GetUser(ctx, username, false /* withSecrets */)
if err != nil {
return nil, nil, trace.Wrap(err)
}
@@ -145,27 +145,27 @@ func (s *Server) AuthenticateUser(ctx context.Context, req AuthenticateUserReque
// After we're sure that the user has been logged in successfully, we should call
// the registered login hooks. Login hooks can be registered by other processes to
// execute arbitrary operations after a successful login.
if err := s.CallLoginHooks(ctx, user); err != nil {
if err := a.CallLoginHooks(ctx, user); err != nil {
return nil, nil, trace.Wrap(err)
}
userState, err := s.GetUserOrLoginState(ctx, user.GetName())
userState, err := a.GetUserOrLoginState(ctx, user.GetName())
if err != nil {
return nil, nil, trace.Wrap(err)
}
clusterName, err := s.GetClusterName()
clusterName, err := a.GetClusterName()
if err != nil {
return nil, nil, trace.Wrap(err)
}
accessInfo := services.AccessInfoFromUserState(userState)
checker, err := services.NewAccessChecker(accessInfo, clusterName.GetClusterName(), s)
checker, err := services.NewAccessChecker(accessInfo, clusterName.GetClusterName(), a)
if err != nil {
return nil, nil, trace.Wrap(err)
}
// Log event after authentication success
if err := s.emitAuthAuditEvent(ctx, authAuditProps{
if err := a.emitAuthAuditEvent(ctx, authAuditProps{
username: username,
clientMetadata: req.ClientMetadata,
mfaDevice: mfaDev,
@@ -185,7 +185,7 @@ type authAuditProps struct {
authErr error
}
func (s *Server) emitAuthAuditEvent(ctx context.Context, props authAuditProps) error {
func (a *Server) emitAuthAuditEvent(ctx context.Context, props authAuditProps) error {
event := &apievents.UserLogin{
Metadata: apievents.Metadata{
Type: events.UserLoginEvent,
@@ -222,7 +222,7 @@ func (s *Server) emitAuthAuditEvent(ctx context.Context, props authAuditProps) e
// Add required key policy to the event.
if props.checker != nil {
authPref, err := s.GetAuthPreference(ctx)
authPref, err := a.GetAuthPreference(ctx)
if err != nil {
return trace.Wrap(err)
}
@@ -233,7 +233,7 @@ func (s *Server) emitAuthAuditEvent(ctx context.Context, props authAuditProps) e
event.RequiredPrivateKeyPolicy = string(privateKeyPolicy)
}
return trace.Wrap(s.emitter.EmitAuditEvent(s.closeCtx, event))
return trace.Wrap(a.emitter.EmitAuditEvent(a.closeCtx, event))
}
var (
@@ -260,7 +260,7 @@ func IsInvalidLocalCredentialError(err error) bool {
// authenticateUser authenticates a user through various methods (password, MFA,
// passwordless)
// Returns the device used to authenticate (if applicable) and the username.
func (s *Server) authenticateUser(ctx context.Context, req AuthenticateUserRequest) (*types.MFADevice, string, error) {
func (a *Server) authenticateUser(ctx context.Context, req AuthenticateUserRequest) (*types.MFADevice, string, error) {
if err := req.CheckAndSetDefaults(); err != nil {
return nil, "", trace.Wrap(err)
}
@@ -269,7 +269,7 @@ func (s *Server) authenticateUser(ctx context.Context, req AuthenticateUserReque
// Only one path if passwordless, other variants shouldn't see an empty user.
if passwordless {
return s.authenticatePasswordless(ctx, req)
return a.authenticatePasswordless(ctx, req)
}
// Try 2nd-factor-enabled authentication schemes first.
@@ -280,7 +280,7 @@ func (s *Server) authenticateUser(ctx context.Context, req AuthenticateUserReque
case req.HeadlessAuthenticationID != "":
// handle authentication before the user lock to prevent locking out users
// due to timed-out/canceled headless authentication attempts.
mfaDevice, err := s.authenticateHeadless(ctx, req)
mfaDevice, err := a.authenticateHeadless(ctx, req)
if err != nil {
log.Debugf("Headless Authentication for user %q failed while waiting for approval: %v", user, err)
return nil, "", trace.Wrap(authenticateHeadlessError)
@@ -296,7 +296,7 @@ func (s *Server) authenticateUser(ctx context.Context, req AuthenticateUserReque
Webauthn: wantypes.CredentialAssertionResponseToProto(req.Webauthn),
},
}
dev, _, err := s.validateMFAAuthResponse(ctx, mfaResponse, user, passwordless)
dev, _, err := a.validateMFAAuthResponse(ctx, mfaResponse, user, passwordless)
return dev, trace.Wrap(err)
}
authErr = authenticateWebauthnError
@@ -304,7 +304,7 @@ func (s *Server) authenticateUser(ctx context.Context, req AuthenticateUserReque
authenticateFn = func() (*types.MFADevice, error) {
// OTP cannot be validated by validateMFAAuthResponse because we need to
// check the user's password too.
res, err := s.checkPassword(user, req.OTP.Password, req.OTP.Token)
res, err := a.checkPassword(user, req.OTP.Password, req.OTP.Token)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -314,7 +314,7 @@ func (s *Server) authenticateUser(ctx context.Context, req AuthenticateUserReque
}
if authenticateFn != nil {
var dev *types.MFADevice
err := s.WithUserLock(ctx, user, func() error {
err := a.WithUserLock(ctx, user, func() error {
var err error
dev, err = authenticateFn()
return err
@@ -342,7 +342,7 @@ func (s *Server) authenticateUser(ctx context.Context, req AuthenticateUserReque
return nil, "", trace.AccessDenied("unsupported authentication method")
}
authPreference, err := s.GetAuthPreference(ctx)
authPreference, err := a.GetAuthPreference(ctx)
if err != nil {
return nil, "", trace.Wrap(err)
}
@@ -355,7 +355,7 @@ func (s *Server) authenticateUser(ctx context.Context, req AuthenticateUserReque
case constants.SecondFactorOptional:
// 2FA is optional. Make sure that a user does not have MFA devices
// registered.
devs, err := s.Services.GetMFADevices(ctx, user, false /* withSecrets */)
devs, err := a.Services.GetMFADevices(ctx, user, false /* withSecrets */)
if err != nil && !trace.IsNotFound(err) {
return nil, "", trace.Wrap(err)
}
@@ -370,8 +370,8 @@ func (s *Server) authenticateUser(ctx context.Context, req AuthenticateUserReque
log.Warningf("MFA bypass attempt by user %q, access denied.", user)
return nil, "", trace.AccessDenied("missing second factor")
}
if err = s.WithUserLock(ctx, user, func() error {
return s.checkPasswordWOToken(user, req.Pass.Password)
if err = a.WithUserLock(ctx, user, func() error {
return a.checkPasswordWOToken(user, req.Pass.Password)
}); err != nil {
if fieldErr := getErrorByTraceField(err); fieldErr != nil {
return nil, "", trace.Wrap(fieldErr)
@@ -384,13 +384,13 @@ func (s *Server) authenticateUser(ctx context.Context, req AuthenticateUserReque
return nil, user, nil
}
func (s *Server) authenticatePasswordless(ctx context.Context, req AuthenticateUserRequest) (*types.MFADevice, string, error) {
func (a *Server) authenticatePasswordless(ctx context.Context, req AuthenticateUserRequest) (*types.MFADevice, string, error) {
mfaResponse := &proto.MFAAuthenticateResponse{
Response: &proto.MFAAuthenticateResponse_Webauthn{
Webauthn: wantypes.CredentialAssertionResponseToProto(req.Webauthn),
},
}
dev, user, err := s.validateMFAAuthResponse(ctx, mfaResponse, "", true /* passwordless */)
dev, user, err := a.validateMFAAuthResponse(ctx, mfaResponse, "", true /* passwordless */)
if err != nil {
log.Debugf("Passwordless authentication failed: %v", err)
return nil, "", trace.Wrap(authenticateWebauthnError)
@@ -399,7 +399,7 @@ func (s *Server) authenticatePasswordless(ctx context.Context, req AuthenticateU
// A distinction between passwordless and "plain" MFA is that we can't
// acquire the user lock beforehand (or at all on failures!)
// We do grab it here so successful logins go through the regular process.
if err := s.WithUserLock(ctx, user, func() error { return nil }); err != nil {
if err := a.WithUserLock(ctx, user, func() error { return nil }); err != nil {
log.Debugf("WithUserLock for user %q failed during passwordless authentication: %v", user, err)
return nil, user, trace.Wrap(authenticateWebauthnError)
}
@@ -407,11 +407,11 @@ func (s *Server) authenticatePasswordless(ctx context.Context, req AuthenticateU
return dev, user, nil
}
func (s *Server) authenticateHeadless(ctx context.Context, req AuthenticateUserRequest) (mfa *types.MFADevice, err error) {
func (a *Server) authenticateHeadless(ctx context.Context, req AuthenticateUserRequest) (mfa *types.MFADevice, err error) {
// Delete the headless authentication upon failure.
defer func() {
if err != nil {
if err := s.DeleteHeadlessAuthentication(s.CloseContext(), req.Username, req.HeadlessAuthenticationID); err != nil && !trace.IsNotFound(err) {
if err := a.DeleteHeadlessAuthentication(a.CloseContext(), req.Username, req.HeadlessAuthenticationID); err != nil && !trace.IsNotFound(err) {
log.Debugf("Failed to delete headless authentication: %v", err)
}
}
@@ -423,7 +423,7 @@ func (s *Server) authenticateHeadless(ctx context.Context, req AuthenticateUserR
defer cancel()
// Headless Authentication should expire when the callback expires.
expires := s.clock.Now().Add(defaults.CallbackTimeout)
expires := a.clock.Now().Add(defaults.CallbackTimeout)
// Create the headless authentication and validate request details.
ha, err := types.NewHeadlessAuthentication(req.Username, req.HeadlessAuthenticationID, expires)
@@ -438,22 +438,22 @@ func (s *Server) authenticateHeadless(ctx context.Context, req AuthenticateUserR
return nil, trace.Wrap(err)
}
emitHeadlessLoginEvent(ctx, events.UserHeadlessLoginRequestedCode, s.emitter, ha, nil)
emitHeadlessLoginEvent(ctx, events.UserHeadlessLoginRequestedCode, a.emitter, ha, nil)
// Headless authentication requests are made without any prior authentication. To avoid DDos
// attacks on the Auth server's backend, we don't create the headless authentication in the
// backend until an authenticated client creates a headless authentication stub. This serves
// as indirect authorization to insert the full headless authentication details into the backend.
if _, err := s.waitForHeadlessStub(ctx, ha); err != nil {
if _, err := a.waitForHeadlessStub(ctx, ha); err != nil {
return nil, trace.Wrap(err)
}
if err := s.UpsertHeadlessAuthentication(ctx, ha); err != nil {
if err := a.UpsertHeadlessAuthentication(ctx, ha); err != nil {
return nil, trace.Wrap(err)
}
// Wait for the request to be approved/denied.
approvedHeadlessAuthn, err := s.waitForHeadlessApproval(ctx, req.Username, req.HeadlessAuthenticationID)
approvedHeadlessAuthn, err := a.waitForHeadlessApproval(ctx, req.Username, req.HeadlessAuthenticationID)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -469,8 +469,8 @@ func (s *Server) authenticateHeadless(ctx context.Context, req AuthenticateUserR
return approvedHeadlessAuthn.MfaDevice, nil
}
func (s *Server) waitForHeadlessStub(ctx context.Context, ha *types.HeadlessAuthentication) (*types.HeadlessAuthentication, error) {
sub, err := s.headlessAuthenticationWatcher.Subscribe(ctx, ha.User, services.HeadlessAuthenticationUserStubID)
func (a *Server) waitForHeadlessStub(ctx context.Context, ha *types.HeadlessAuthentication) (*types.HeadlessAuthentication, error) {
sub, err := a.headlessAuthenticationWatcher.Subscribe(ctx, ha.User, services.HeadlessAuthenticationUserStubID)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -485,8 +485,8 @@ func (s *Server) waitForHeadlessStub(ctx context.Context, ha *types.HeadlessAuth
return stub, nil
}
func (s *Server) waitForHeadlessApproval(ctx context.Context, username, reqID string) (*types.HeadlessAuthentication, error) {
sub, err := s.headlessAuthenticationWatcher.Subscribe(ctx, username, reqID)
func (a *Server) waitForHeadlessApproval(ctx context.Context, username, reqID string) (*types.HeadlessAuthentication, error) {
sub, err := a.headlessAuthenticationWatcher.Subscribe(ctx, username, reqID)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -514,10 +514,10 @@ func (s *Server) waitForHeadlessApproval(ctx context.Context, username, reqID st
// AuthenticateWebUser authenticates web user, creates and returns a web session
// if authentication is successful. In case the existing session ID is used to authenticate,
// returns the existing session instead of creating a new one
func (s *Server) AuthenticateWebUser(ctx context.Context, req AuthenticateUserRequest) (types.WebSession, error) {
func (a *Server) AuthenticateWebUser(ctx context.Context, req AuthenticateUserRequest) (types.WebSession, error) {
username := req.Username // Empty if passwordless.
authPref, err := s.GetAuthPreference(ctx)
authPref, err := a.GetAuthPreference(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -527,12 +527,12 @@ func (s *Server) AuthenticateWebUser(ctx context.Context, req AuthenticateUserRe
// This condition uses Session as a blanket check, because any new method added
// to the local auth will be disabled by default.
if !authPref.GetAllowLocalAuth() && req.Session == nil {
s.emitNoLocalAuthEvent(username)
a.emitNoLocalAuthEvent(username)
return nil, trace.AccessDenied(noLocalAuth)
}
if req.Session != nil {
session, err := s.GetWebSession(ctx, types.GetWebSessionRequest{
session, err := a.GetWebSession(ctx, types.GetWebSessionRequest{
User: username,
SessionID: req.Session.ID,
})
@@ -542,7 +542,7 @@ func (s *Server) AuthenticateWebUser(ctx context.Context, req AuthenticateUserRe
return session, nil
}
user, _, err := s.AuthenticateUser(ctx, req)
user, _, err := a.AuthenticateUser(ctx, req)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -555,12 +555,12 @@ func (s *Server) AuthenticateWebUser(ctx context.Context, req AuthenticateUserRe
}
}
sess, err := s.CreateWebSessionFromReq(ctx, types.NewWebSessionRequest{
sess, err := a.CreateWebSessionFromReq(ctx, types.NewWebSessionRequest{
User: user.GetName(),
LoginIP: loginIP,
Roles: user.GetRoles(),
Traits: user.GetTraits(),
LoginTime: s.clock.Now().UTC(),
LoginTime: a.clock.Now().UTC(),
AttestWebSession: true,
})
if err != nil {
@@ -658,32 +658,32 @@ func AuthoritiesToTrustedCerts(authorities []types.CertAuthority) []TrustedCerts
// AuthenticateSSHUser authenticates an SSH user and returns SSH and TLS
// certificates for the public key in req.
func (s *Server) AuthenticateSSHUser(ctx context.Context, req AuthenticateSSHRequest) (*SSHLoginResponse, error) {
func (a *Server) AuthenticateSSHUser(ctx context.Context, req AuthenticateSSHRequest) (*SSHLoginResponse, error) {
username := req.Username // Empty if passwordless.
authPref, err := s.GetAuthPreference(ctx)
authPref, err := a.GetAuthPreference(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
if !authPref.GetAllowLocalAuth() {
s.emitNoLocalAuthEvent(username)
a.emitNoLocalAuthEvent(username)
return nil, trace.AccessDenied(noLocalAuth)
}
clusterName, err := s.GetClusterName()
clusterName, err := a.GetClusterName()
if err != nil {
return nil, trace.Wrap(err)
}
// It's safe to extract the roles and traits directly from services.User as
// this endpoint is only used for local accounts.
user, checker, err := s.AuthenticateUser(ctx, req.AuthenticateUserRequest)
user, checker, err := a.AuthenticateUser(ctx, req.AuthenticateUserRequest)
if err != nil {
return nil, trace.Wrap(err)
}
// Return the host CA for this cluster only.
authority, err := s.GetCertAuthority(ctx, types.CertAuthID{
authority, err := a.GetCertAuthority(ctx, types.CertAuthID{
Type: types.HostCA,
DomainName: clusterName.GetClusterName(),
}, false)
@@ -721,7 +721,7 @@ func (s *Server) AuthenticateSSHUser(ctx context.Context, req AuthenticateSSHReq
// For headless authentication, a short-lived mfa-verified cert should be generated.
if req.HeadlessAuthenticationID != "" {
ha, err := s.GetHeadlessAuthentication(ctx, req.Username, req.HeadlessAuthenticationID)
ha, err := a.GetHeadlessAuthentication(ctx, req.Username, req.HeadlessAuthenticationID)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -732,7 +732,7 @@ func (s *Server) AuthenticateSSHUser(ctx context.Context, req AuthenticateSSHReq
certReq.ttl = time.Minute
}
certs, err := s.generateUserCert(certReq)
certs, err := a.generateUserCert(certReq)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -746,8 +746,8 @@ func (s *Server) AuthenticateSSHUser(ctx context.Context, req AuthenticateSSHReq
}
// emitNoLocalAuthEvent creates and emits a local authentication is disabled message.
func (s *Server) emitNoLocalAuthEvent(username string) {
if err := s.emitter.EmitAuditEvent(s.closeCtx, &apievents.AuthAttempt{
func (a *Server) emitNoLocalAuthEvent(username string) {
if err := a.emitter.EmitAuditEvent(a.closeCtx, &apievents.AuthAttempt{
Metadata: apievents.Metadata{
Type: events.AuthAttemptEvent,
Code: events.AuthAttemptFailureCode,
@@ -764,15 +764,15 @@ func (s *Server) emitNoLocalAuthEvent(username string) {
}
}
func (s *Server) createUserWebSession(ctx context.Context, user services.UserState, loginIP string) (types.WebSession, error) {
func (a *Server) createUserWebSession(ctx context.Context, user services.UserState, loginIP string) (types.WebSession, error) {
// It's safe to extract the roles and traits directly from services.User as this method
// is only used for local accounts.
return s.CreateWebSessionFromReq(ctx, types.NewWebSessionRequest{
return a.CreateWebSessionFromReq(ctx, types.NewWebSessionRequest{
User: user.GetName(),
LoginIP: loginIP,
Roles: user.GetRoles(),
Traits: user.GetTraits(),
LoginTime: s.clock.Now().UTC(),
LoginTime: a.clock.Now().UTC(),
})
}
+37 -37
View File
@@ -42,8 +42,8 @@ import (
var fakePasswordHash = []byte(`$2a$10$Yy.e6BmS2SrGbBDsyDLVkOANZmvjjMR890nUGSXFJHBXWzxe7T44m`)
// ChangeUserAuthentication implements AuthService.ChangeUserAuthentication.
func (s *Server) ChangeUserAuthentication(ctx context.Context, req *proto.ChangeUserAuthenticationRequest) (*proto.ChangeUserAuthenticationResponse, error) {
user, err := s.changeUserAuthentication(ctx, req)
func (a *Server) ChangeUserAuthentication(ctx context.Context, req *proto.ChangeUserAuthenticationRequest) (*proto.ChangeUserAuthenticationResponse, error) {
user, err := a.changeUserAuthentication(ctx, req)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -52,18 +52,18 @@ func (s *Server) ChangeUserAuthentication(ctx context.Context, req *proto.Change
_, emailErr := mail.ParseAddress(user.GetName())
hasEmail := emailErr == nil
hasMFA := req.GetNewMFARegisterResponse() != nil
recoveryAllowed := s.isAccountRecoveryAllowed(ctx) == nil
recoveryAllowed := a.isAccountRecoveryAllowed(ctx) == nil
createRecoveryCodes := hasEmail && hasMFA && recoveryAllowed
var newRecovery *proto.RecoveryCodes
if createRecoveryCodes {
newRecovery, err = s.generateAndUpsertRecoveryCodes(ctx, user.GetName())
newRecovery, err = a.generateAndUpsertRecoveryCodes(ctx, user.GetName())
if err != nil {
return nil, trace.Wrap(err)
}
}
webSession, err := s.createUserWebSession(ctx, user, req.LoginIP)
webSession, err := a.createUserWebSession(ctx, user, req.LoginIP)
if err != nil {
if keys.IsPrivateKeyPolicyError(err) {
// Do not return an error, otherwise
@@ -93,8 +93,8 @@ func (s *Server) ChangeUserAuthentication(ctx context.Context, req *proto.Change
// ResetPassword securely generates a new random password and assigns it to user.
// This method is used to invalidate existing user password during password
// reset process.
func (s *Server) ResetPassword(ctx context.Context, username string) (string, error) {
user, err := s.GetUser(ctx, username, false)
func (a *Server) ResetPassword(ctx context.Context, username string) (string, error) {
user, err := a.GetUser(ctx, username, false)
if err != nil {
return "", trace.Wrap(err)
}
@@ -104,7 +104,7 @@ func (s *Server) ResetPassword(ctx context.Context, username string) (string, er
return "", trace.Wrap(err)
}
err = s.UpsertPassword(user.GetName(), []byte(password))
err = a.UpsertPassword(user.GetName(), []byte(password))
if err != nil {
return "", trace.Wrap(err)
}
@@ -113,7 +113,7 @@ func (s *Server) ResetPassword(ctx context.Context, username string) (string, er
}
// ChangePassword updates users password based on the old password.
func (s *Server) ChangePassword(ctx context.Context, req *proto.ChangePasswordRequest) error {
func (a *Server) ChangePassword(ctx context.Context, req *proto.ChangePasswordRequest) error {
// validate new password
if err := services.VerifyPassword(req.NewPassword); err != nil {
return trace.Wrap(err)
@@ -136,15 +136,15 @@ func (s *Server) ChangePassword(ctx context.Context, req *proto.ChangePasswordRe
Token: req.SecondFactorToken,
}
}
if _, _, err := s.authenticateUser(ctx, authReq); err != nil {
if _, _, err := a.authenticateUser(ctx, authReq); err != nil {
return trace.Wrap(err)
}
if err := s.UpsertPassword(user, req.NewPassword); err != nil {
if err := a.UpsertPassword(user, req.NewPassword); err != nil {
return trace.Wrap(err)
}
if err := s.emitter.EmitAuditEvent(s.closeCtx, &apievents.UserPasswordChange{
if err := a.emitter.EmitAuditEvent(a.closeCtx, &apievents.UserPasswordChange{
Metadata: apievents.Metadata{
Type: events.UserPasswordChangeEvent,
Code: events.UserPasswordChangeCode,
@@ -158,7 +158,7 @@ func (s *Server) ChangePassword(ctx context.Context, req *proto.ChangePasswordRe
// checkPasswordWOToken checks just password without checking OTP tokens
// used in case of SSH authentication, when token has been validated.
func (s *Server) checkPasswordWOToken(user string, password []byte) error {
func (a *Server) checkPasswordWOToken(user string, password []byte) error {
const errMsg = "invalid username or password"
err := services.VerifyPassword(password)
@@ -166,7 +166,7 @@ func (s *Server) checkPasswordWOToken(user string, password []byte) error {
return trace.BadParameter(errMsg)
}
hash, err := s.GetPasswordHash(user)
hash, err := a.GetPasswordHash(user)
if err != nil && !trace.IsNotFound(err) {
return trace.Wrap(err)
}
@@ -196,13 +196,13 @@ type checkPasswordResult struct {
}
// checkPassword checks the password and OTP token. Called by tsh or lib/web/*.
func (s *Server) checkPassword(user string, password []byte, otpToken string) (*checkPasswordResult, error) {
err := s.checkPasswordWOToken(user, password)
func (a *Server) checkPassword(user string, password []byte, otpToken string) (*checkPasswordResult, error) {
err := a.checkPasswordWOToken(user, password)
if err != nil {
return nil, trace.Wrap(err)
}
mfaDev, err := s.checkOTP(user, otpToken)
mfaDev, err := a.checkOTP(user, otpToken)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -210,9 +210,9 @@ func (s *Server) checkPassword(user string, password []byte, otpToken string) (*
}
// checkOTP checks if the OTP token is valid.
func (s *Server) checkOTP(user string, otpToken string) (*types.MFADevice, error) {
func (a *Server) checkOTP(user string, otpToken string) (*types.MFADevice, error) {
// get the previously used token to mitigate token replay attacks
usedToken, err := s.GetUsedTOTPToken(user)
usedToken, err := a.GetUsedTOTPToken(user)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -222,7 +222,7 @@ func (s *Server) checkOTP(user string, otpToken string) (*types.MFADevice, error
}
ctx := context.TODO()
devs, err := s.Services.GetMFADevices(ctx, user, true)
devs, err := a.Services.GetMFADevices(ctx, user, true)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -233,7 +233,7 @@ func (s *Server) checkOTP(user string, otpToken string) (*types.MFADevice, error
continue
}
if err := s.checkTOTP(ctx, user, otpToken, dev); err != nil {
if err := a.checkTOTP(ctx, user, otpToken, dev); err != nil {
log.WithError(err).Errorf("Using TOTP device %q", dev.GetName())
continue
}
@@ -243,13 +243,13 @@ func (s *Server) checkOTP(user string, otpToken string) (*types.MFADevice, error
}
// checkTOTP checks if the TOTP token is valid.
func (s *Server) checkTOTP(ctx context.Context, user, otpToken string, dev *types.MFADevice) error {
func (a *Server) checkTOTP(ctx context.Context, user, otpToken string, dev *types.MFADevice) error {
if dev.GetTotp() == nil {
return trace.BadParameter("checkTOTP called with non-TOTP MFADevice %T", dev.Device)
}
// we use totp.ValidateCustom over totp.Validate so we can use
// a fake clock in tests to get reliable results
valid, err := totp.ValidateCustom(otpToken, dev.GetTotp().Key, s.clock.Now(), totp.ValidateOpts{
valid, err := totp.ValidateCustom(otpToken, dev.GetTotp().Key, a.clock.Now(), totp.ValidateOpts{
Period: teleport.TOTPValidityPeriod,
Skew: teleport.TOTPSkew,
Digits: otp.DigitsSix,
@@ -262,21 +262,21 @@ func (s *Server) checkTOTP(ctx context.Context, user, otpToken string, dev *type
return trace.AccessDenied("invalid one time token, please check if the token has expired and try again")
}
// if we have a valid token, update the previously used token
if err := s.UpsertUsedTOTPToken(user, otpToken); err != nil {
if err := a.UpsertUsedTOTPToken(user, otpToken); err != nil {
return trace.Wrap(err)
}
// Update LastUsed timestamp on the device.
dev.LastUsed = s.clock.Now()
if err := s.UpsertMFADevice(ctx, user, dev); err != nil {
dev.LastUsed = a.clock.Now()
if err := a.UpsertMFADevice(ctx, user, dev); err != nil {
return trace.Wrap(err)
}
return nil
}
func (s *Server) changeUserAuthentication(ctx context.Context, req *proto.ChangeUserAuthenticationRequest) (types.User, error) {
func (a *Server) changeUserAuthentication(ctx context.Context, req *proto.ChangeUserAuthenticationRequest) (types.User, error) {
// Get cluster configuration and check if local auth is allowed.
authPref, err := s.GetAuthPreference(ctx)
authPref, err := a.GetAuthPreference(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -297,16 +297,16 @@ func (s *Server) changeUserAuthentication(ctx context.Context, req *proto.Change
}
// Check if token exists.
token, err := s.getResetPasswordToken(ctx, req.TokenID)
token, err := a.getResetPasswordToken(ctx, req.TokenID)
if err != nil {
return nil, trace.Wrap(err)
}
if token.Expiry().Before(s.clock.Now().UTC()) {
if token.Expiry().Before(a.clock.Now().UTC()) {
return nil, trace.BadParameter("expired token")
}
err = s.changeUserSecondFactor(ctx, req, token)
err = a.changeUserSecondFactor(ctx, req, token)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -314,18 +314,18 @@ func (s *Server) changeUserAuthentication(ctx context.Context, req *proto.Change
username := token.GetUser()
// Delete this token first to minimize the chances
// of partially updated user with still valid token.
err = s.deleteUserTokens(ctx, username)
err = a.deleteUserTokens(ctx, username)
if err != nil {
return nil, trace.Wrap(err)
}
if !reqPasswordless {
if err := s.UpsertPassword(username, req.GetNewPassword()); err != nil {
if err := a.UpsertPassword(username, req.GetNewPassword()); err != nil {
return nil, trace.Wrap(err)
}
}
user, err := s.GetUser(ctx, username, false)
user, err := a.GetUser(ctx, username, false)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -333,9 +333,9 @@ func (s *Server) changeUserAuthentication(ctx context.Context, req *proto.Change
return user, nil
}
func (s *Server) changeUserSecondFactor(ctx context.Context, req *proto.ChangeUserAuthenticationRequest, token types.UserToken) error {
func (a *Server) changeUserSecondFactor(ctx context.Context, req *proto.ChangeUserAuthenticationRequest, token types.UserToken) error {
username := token.GetUser()
cap, err := s.GetAuthPreference(ctx)
cap, err := a.GetAuthPreference(ctx)
if err != nil {
return trace.Wrap(err)
}
@@ -378,7 +378,7 @@ func (s *Server) changeUserSecondFactor(ctx context.Context, req *proto.ChangeUs
deviceUsage = proto.DeviceUsage_DEVICE_USAGE_PASSWORDLESS
}
_, err = s.verifyMFARespAndAddDevice(ctx, &newMFADeviceFields{
_, err = a.verifyMFARespAndAddDevice(ctx, &newMFADeviceFields{
username: token.GetUser(),
newDeviceName: deviceName,
tokenID: token.GetName(),
+9 -9
View File
@@ -30,21 +30,21 @@ import (
// ReconcileServerInfos periodically reconciles the labels of ServerInfo
// resources with their corresponding Teleport SSH servers.
func (s *Server) ReconcileServerInfos(ctx context.Context) error {
func (a *Server) ReconcileServerInfos(ctx context.Context) error {
const batchSize = 100
const timeBetweenBatches = 10 * time.Second
const timeBetweenReconciliationLoops = 10 * time.Minute
clock := s.GetClock()
clock := a.GetClock()
for {
var failedUpdates int
// Iterate over nodes in batches.
nodeStream := s.GetNodeStream(ctx, defaults.Namespace)
nodeStream := a.GetNodeStream(ctx, defaults.Namespace)
var nodes []types.Server
for moreNodes := true; moreNodes; {
nodes, moreNodes = stream.Take(nodeStream, batchSize)
updates, err := s.setCloudLabelsOnNodes(ctx, nodes)
updates, err := a.setCloudLabelsOnNodes(ctx, nodes)
if err != nil {
return trace.Wrap(err)
}
@@ -70,13 +70,13 @@ func (s *Server) ReconcileServerInfos(ctx context.Context) error {
}
}
func (s *Server) setCloudLabelsOnNodes(ctx context.Context, nodes []types.Server) (failedUpdates int, err error) {
func (a *Server) setCloudLabelsOnNodes(ctx context.Context, nodes []types.Server) (failedUpdates int, err error) {
for _, node := range nodes {
meta := node.GetCloudMetadata()
if meta != nil && meta.AWS != nil {
si, err := s.GetServerInfo(ctx, meta.AWS.GetServerInfoName())
si, err := a.GetServerInfo(ctx, meta.AWS.GetServerInfoName())
if err == nil {
err := s.updateLabelsOnNode(ctx, node, si)
err := a.updateLabelsOnNode(ctx, node, si)
// Didn't find control stream for node, save count for logging.
if trace.IsNotFound(err) {
failedUpdates++
@@ -91,8 +91,8 @@ func (s *Server) setCloudLabelsOnNodes(ctx context.Context, nodes []types.Server
return failedUpdates, nil
}
func (s *Server) updateLabelsOnNode(ctx context.Context, node types.Server, si types.ServerInfo) error {
err := s.UpdateLabels(ctx, proto.InventoryUpdateLabelsRequest{
func (a *Server) updateLabelsOnNode(ctx context.Context, node types.Server, si types.ServerInfo) error {
err := a.UpdateLabels(ctx, proto.InventoryUpdateLabelsRequest{
ServerID: node.GetName(),
Kind: proto.LabelUpdateKind_SSHServerCloudLabels,
Labels: si.GetStaticLabels(),
+26 -26
View File
@@ -40,7 +40,7 @@ import (
// backend with the identity of the caller used to generate the certificate.
// The certificate is used for all access requests, which is where access
// control is enforced.
func (s *Server) CreateAppSession(ctx context.Context, req types.CreateAppSessionRequest, user services.UserState, identity tlsca.Identity, checker services.AccessChecker) (types.WebSession, error) {
func (a *Server) CreateAppSession(ctx context.Context, req types.CreateAppSessionRequest, user services.UserState, identity tlsca.Identity, checker services.AccessChecker) (types.WebSession, error) {
if !modules.GetModules().Features().App {
return nil, trace.AccessDenied(
"this Teleport cluster is not licensed for application access, please contact the cluster administrator")
@@ -53,11 +53,11 @@ func (s *Server) CreateAppSession(ctx context.Context, req types.CreateAppSessio
// sessionID which will be used to fetch services.WebSession which
// contains a certificate whose life matches the life of the session
// that will be used to establish the connection.
ttl := checker.AdjustSessionTTL(identity.Expires.Sub(s.clock.Now()))
ttl := checker.AdjustSessionTTL(identity.Expires.Sub(a.clock.Now()))
// Encode user traits in the app access certificate. This will allow to
// pass user traits when talking to app servers in leaf clusters.
_, traits, err := services.ExtractFromIdentity(ctx, s, identity)
_, traits, err := services.ExtractFromIdentity(ctx, a, identity)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -67,7 +67,7 @@ func (s *Server) CreateAppSession(ctx context.Context, req types.CreateAppSessio
if err != nil {
return nil, trace.Wrap(err)
}
certs, err := s.generateUserCert(certRequest{
certs, err := a.generateUserCert(certRequest{
user: user,
loginIP: identity.LoginIP,
publicKey: publicKey,
@@ -108,13 +108,13 @@ func (s *Server) CreateAppSession(ctx context.Context, req types.CreateAppSessio
Priv: privateKey,
Pub: certs.SSH,
TLSCert: certs.TLS,
Expires: s.clock.Now().Add(ttl),
Expires: a.clock.Now().Add(ttl),
BearerToken: bearer,
})
if err != nil {
return nil, trace.Wrap(err)
}
if err = s.UpsertAppSession(ctx, session); err != nil {
if err = a.UpsertAppSession(ctx, session); err != nil {
return nil, trace.Wrap(err)
}
log.Debugf("Generated application web session for %v with TTL %v.", req.Username, ttl)
@@ -202,13 +202,13 @@ func waitForWebSession(ctx context.Context, sessionID, user string, evenSubKind
// generateAppToken generates an JWT token that will be passed along with every
// application request.
func (s *Server) generateAppToken(ctx context.Context, username string, roles []string, traits map[string][]string, uri string, expires time.Time) (string, error) {
func (a *Server) generateAppToken(ctx context.Context, username string, roles []string, traits map[string][]string, uri string, expires time.Time) (string, error) {
// Get the clusters CA.
clusterName, err := s.GetDomainName()
clusterName, err := a.GetDomainName()
if err != nil {
return "", trace.Wrap(err)
}
ca, err := s.GetCertAuthority(ctx, types.CertAuthID{
ca, err := a.GetCertAuthority(ctx, types.CertAuthID{
Type: types.JWTSigner,
DomainName: clusterName,
}, true)
@@ -226,11 +226,11 @@ func (s *Server) generateAppToken(ctx context.Context, username string, roles []
}
// Extract the JWT signing key and sign the claims.
signer, err := s.GetKeyStore().GetJWTSigner(ctx, ca)
signer, err := a.GetKeyStore().GetJWTSigner(ctx, ca)
if err != nil {
return "", trace.Wrap(err)
}
privateKey, err := services.GetJWTSigner(signer, ca.GetClusterName(), s.clock)
privateKey, err := services.GetJWTSigner(signer, ca.GetClusterName(), a.clock)
if err != nil {
return "", trace.Wrap(err)
}
@@ -248,13 +248,13 @@ func (s *Server) generateAppToken(ctx context.Context, username string, roles []
return token, nil
}
func (s *Server) CreateWebSessionFromReq(ctx context.Context, req types.NewWebSessionRequest) (types.WebSession, error) {
session, err := s.NewWebSession(ctx, req)
func (a *Server) CreateWebSessionFromReq(ctx context.Context, req types.NewWebSessionRequest) (types.WebSession, error) {
session, err := a.NewWebSession(ctx, req)
if err != nil {
return nil, trace.Wrap(err)
}
err = s.upsertWebSession(ctx, req.User, session)
err = a.upsertWebSession(ctx, req.User, session)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -262,29 +262,29 @@ func (s *Server) CreateWebSessionFromReq(ctx context.Context, req types.NewWebSe
return session, nil
}
func (s *Server) CreateSessionCert(user services.UserState, sessionTTL time.Duration, publicKey []byte, compatibility, routeToCluster, kubernetesCluster, loginIP string, attestationReq *keys.AttestationStatement) ([]byte, []byte, error) {
func (a *Server) CreateSessionCert(user services.UserState, sessionTTL time.Duration, publicKey []byte, compatibility, routeToCluster, kubernetesCluster, loginIP string, attestationReq *keys.AttestationStatement) ([]byte, []byte, error) {
// It's safe to extract the access info directly from services.User because
// this occurs during the initial login before the first certs have been
// generated, so there's no possibility of any active access requests.
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
userState, err := s.GetUserOrLoginState(ctx, user.GetName())
userState, err := a.GetUserOrLoginState(ctx, user.GetName())
if err != nil {
return nil, nil, trace.Wrap(err)
}
accessInfo := services.AccessInfoFromUserState(userState)
clusterName, err := s.GetClusterName()
clusterName, err := a.GetClusterName()
if err != nil {
return nil, nil, trace.Wrap(err)
}
checker, err := services.NewAccessChecker(accessInfo, clusterName.GetClusterName(), s)
checker, err := services.NewAccessChecker(accessInfo, clusterName.GetClusterName(), a)
if err != nil {
return nil, nil, trace.Wrap(err)
}
certs, err := s.generateUserCert(certRequest{
certs, err := a.generateUserCert(certRequest{
user: userState,
ttl: sessionTTL,
publicKey: publicKey,
@@ -303,7 +303,7 @@ func (s *Server) CreateSessionCert(user services.UserState, sessionTTL time.Dura
return certs.SSH, certs.TLS, nil
}
func (s *Server) CreateSnowflakeSession(ctx context.Context, req types.CreateSnowflakeSessionRequest,
func (a *Server) CreateSnowflakeSession(ctx context.Context, req types.CreateSnowflakeSessionRequest,
identity tlsca.Identity, checker services.AccessChecker,
) (types.WebSession, error) {
if !modules.GetModules().Features().DB {
@@ -318,7 +318,7 @@ func (s *Server) CreateSnowflakeSession(ctx context.Context, req types.CreateSno
// sessionID which will be used to fetch services.WebSession which
// contains a certificate whose life matches the life of the session
// that will be used to establish the connection.
ttl := checker.AdjustSessionTTL(identity.Expires.Sub(s.clock.Now()))
ttl := checker.AdjustSessionTTL(identity.Expires.Sub(a.clock.Now()))
// Create services.WebSession for this session.
sessionID, err := utils.CryptoRandomHex(SessionTokenBytes)
@@ -327,14 +327,14 @@ func (s *Server) CreateSnowflakeSession(ctx context.Context, req types.CreateSno
}
session, err := types.NewWebSession(sessionID, types.KindSnowflakeSession, types.WebSessionSpecV2{
User: req.Username,
Expires: s.clock.Now().Add(ttl),
Expires: a.clock.Now().Add(ttl),
BearerToken: req.SessionToken,
BearerTokenExpires: s.clock.Now().Add(req.TokenTTL),
BearerTokenExpires: a.clock.Now().Add(req.TokenTTL),
})
if err != nil {
return nil, trace.Wrap(err)
}
if err = s.UpsertSnowflakeSession(ctx, session); err != nil {
if err = a.UpsertSnowflakeSession(ctx, session); err != nil {
return nil, trace.Wrap(err)
}
log.Debugf("Generated Snowflake web session for %v with TTL %v.", req.Username, ttl)
@@ -342,7 +342,7 @@ func (s *Server) CreateSnowflakeSession(ctx context.Context, req types.CreateSno
return session, nil
}
func (s *Server) CreateSAMLIdPSession(ctx context.Context, req types.CreateSAMLIdPSessionRequest,
func (a *Server) CreateSAMLIdPSession(ctx context.Context, req types.CreateSAMLIdPSessionRequest,
identity tlsca.Identity, checker services.AccessChecker,
) (types.WebSession, error) {
// TODO(mdwn): implement a module.Features() check.
@@ -360,7 +360,7 @@ func (s *Server) CreateSAMLIdPSession(ctx context.Context, req types.CreateSAMLI
if err != nil {
return nil, trace.Wrap(err)
}
if err = s.UpsertSAMLIdPSession(ctx, session); err != nil {
if err = a.UpsertSAMLIdPSession(ctx, session); err != nil {
return nil, trace.Wrap(err)
}
log.Debugf("Generated SAML IdP web session for %v.", req.Username)
+26 -26
View File
@@ -37,15 +37,15 @@ import (
)
// CreateUser inserts a new user entry in a backend.
func (s *Server) CreateUser(ctx context.Context, user types.User) (types.User, error) {
func (a *Server) CreateUser(ctx context.Context, user types.User) (types.User, error) {
if user.GetCreatedBy().IsEmpty() {
user.SetCreatedBy(types.CreatedBy{
User: types.UserRef{Name: authz.ClientUsername(ctx)},
Time: s.GetClock().Now().UTC(),
Time: a.GetClock().Now().UTC(),
})
}
created, err := s.Services.CreateUser(ctx, user)
created, err := a.Services.CreateUser(ctx, user)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -57,7 +57,7 @@ func (s *Server) CreateUser(ctx context.Context, user types.User) (types.User, e
connectorName = created.GetCreatedBy().Connector.ID
}
if err := s.emitter.EmitAuditEvent(ctx, &apievents.UserCreate{
if err := a.emitter.EmitAuditEvent(ctx, &apievents.UserCreate{
Metadata: apievents.Metadata{
Type: events.UserCreateEvent,
Code: events.UserCreateCode,
@@ -73,14 +73,14 @@ func (s *Server) CreateUser(ctx context.Context, user types.User) (types.User, e
log.WithError(err).Warn("Failed to emit user create event.")
}
usagereporter.EmitEditorChangeEvent(created.GetName(), nil, created.GetRoles(), s.AnonymizeAndSubmit)
usagereporter.EmitEditorChangeEvent(created.GetName(), nil, created.GetRoles(), a.AnonymizeAndSubmit)
return created, nil
}
// UpdateUser updates an existing user in a backend.
func (s *Server) UpdateUser(ctx context.Context, user types.User) (types.User, error) {
prevUser, err := s.GetUser(ctx, user.GetName(), false)
func (a *Server) UpdateUser(ctx context.Context, user types.User) (types.User, error) {
prevUser, err := a.GetUser(ctx, user.GetName(), false)
var omitEditorEvent bool
if err != nil {
// don't return error here since this call is for event emitting purposes only
@@ -88,7 +88,7 @@ func (s *Server) UpdateUser(ctx context.Context, user types.User) (types.User, e
omitEditorEvent = true
}
updated, err := s.Services.UpdateUser(ctx, user)
updated, err := a.Services.UpdateUser(ctx, user)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -100,7 +100,7 @@ func (s *Server) UpdateUser(ctx context.Context, user types.User) (types.User, e
connectorName = updated.GetCreatedBy().Connector.ID
}
if err := s.emitter.EmitAuditEvent(ctx, &apievents.UserUpdate{
if err := a.emitter.EmitAuditEvent(ctx, &apievents.UserUpdate{
Metadata: apievents.Metadata{
Type: events.UserUpdatedEvent,
Code: events.UserUpdateCode,
@@ -117,15 +117,15 @@ func (s *Server) UpdateUser(ctx context.Context, user types.User) (types.User, e
}
if !omitEditorEvent {
usagereporter.EmitEditorChangeEvent(updated.GetName(), prevUser.GetRoles(), updated.GetRoles(), s.AnonymizeAndSubmit)
usagereporter.EmitEditorChangeEvent(updated.GetName(), prevUser.GetRoles(), updated.GetRoles(), a.AnonymizeAndSubmit)
}
return updated, nil
}
// UpsertUser updates a user.
func (s *Server) UpsertUser(ctx context.Context, user types.User) (types.User, error) {
prevUser, err := s.GetUser(ctx, user.GetName(), false)
func (a *Server) UpsertUser(ctx context.Context, user types.User) (types.User, error) {
prevUser, err := a.GetUser(ctx, user.GetName(), false)
var omitEditorEvent bool
if err != nil {
if trace.IsNotFound(err) {
@@ -137,7 +137,7 @@ func (s *Server) UpsertUser(ctx context.Context, user types.User) (types.User, e
}
}
upserted, err := s.Services.UpsertUser(ctx, user)
upserted, err := a.Services.UpsertUser(ctx, user)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -149,7 +149,7 @@ func (s *Server) UpsertUser(ctx context.Context, user types.User) (types.User, e
connectorName = upserted.GetCreatedBy().Connector.ID
}
if err := s.emitter.EmitAuditEvent(s.closeCtx, &apievents.UserCreate{
if err := a.emitter.EmitAuditEvent(a.closeCtx, &apievents.UserCreate{
Metadata: apievents.Metadata{
Type: events.UserCreateEvent,
Code: events.UserCreateCode,
@@ -172,7 +172,7 @@ func (s *Server) UpsertUser(ctx context.Context, user types.User) (types.User, e
prevRoles = prevUser.GetRoles()
}
if !omitEditorEvent {
usagereporter.EmitEditorChangeEvent(upserted.GetName(), prevRoles, upserted.GetRoles(), s.AnonymizeAndSubmit)
usagereporter.EmitEditorChangeEvent(upserted.GetName(), prevRoles, upserted.GetRoles(), a.AnonymizeAndSubmit)
}
return upserted, nil
@@ -180,8 +180,8 @@ func (s *Server) UpsertUser(ctx context.Context, user types.User) (types.User, e
// CompareAndSwapUser updates a user but fails if the value on the backend does
// not match the expected value.
func (s *Server) CompareAndSwapUser(ctx context.Context, new, existing types.User) error {
err := s.Services.CompareAndSwapUser(ctx, new, existing)
func (a *Server) CompareAndSwapUser(ctx context.Context, new, existing types.User) error {
err := a.Services.CompareAndSwapUser(ctx, new, existing)
if err != nil {
return trace.Wrap(err)
}
@@ -193,7 +193,7 @@ func (s *Server) CompareAndSwapUser(ctx context.Context, new, existing types.Use
connectorName = new.GetCreatedBy().Connector.ID
}
if err := s.emitter.EmitAuditEvent(ctx, &apievents.UserUpdate{
if err := a.emitter.EmitAuditEvent(ctx, &apievents.UserUpdate{
Metadata: apievents.Metadata{
Type: events.UserUpdatedEvent,
Code: events.UserUpdateCode,
@@ -209,14 +209,14 @@ func (s *Server) CompareAndSwapUser(ctx context.Context, new, existing types.Use
log.WithError(err).Warn("Failed to emit user update event.")
}
usagereporter.EmitEditorChangeEvent(new.GetName(), existing.GetRoles(), new.GetRoles(), s.AnonymizeAndSubmit)
usagereporter.EmitEditorChangeEvent(new.GetName(), existing.GetRoles(), new.GetRoles(), a.AnonymizeAndSubmit)
return nil
}
// DeleteUser deletes an existing user in a backend by username.
func (s *Server) DeleteUser(ctx context.Context, user string) error {
prevUser, err := s.GetUser(ctx, user, false)
func (a *Server) DeleteUser(ctx context.Context, user string) error {
prevUser, err := a.GetUser(ctx, user, false)
var omitEditorEvent bool
if err != nil && !trace.IsNotFound(err) {
// don't return error here, delete may still succeed
@@ -225,26 +225,26 @@ func (s *Server) DeleteUser(ctx context.Context, user string) error {
omitEditorEvent = true
}
role, err := s.Services.GetRole(ctx, services.RoleNameForUser(user))
role, err := a.Services.GetRole(ctx, services.RoleNameForUser(user))
if err != nil {
if !trace.IsNotFound(err) {
return trace.Wrap(err)
}
} else {
if err := s.DeleteRole(ctx, role.GetName()); err != nil {
if err := a.DeleteRole(ctx, role.GetName()); err != nil {
if !trace.IsNotFound(err) {
return trace.Wrap(err)
}
}
}
err = s.Services.DeleteUser(ctx, user)
err = a.Services.DeleteUser(ctx, user)
if err != nil {
return trace.Wrap(err)
}
// If the user was successfully deleted, emit an event.
if err := s.emitter.EmitAuditEvent(s.closeCtx, &apievents.UserDelete{
if err := a.emitter.EmitAuditEvent(a.closeCtx, &apievents.UserDelete{
Metadata: apievents.Metadata{
Type: events.UserDeleteEvent,
Code: events.UserDeleteCode,
@@ -258,7 +258,7 @@ func (s *Server) DeleteUser(ctx context.Context, user string) error {
}
if !omitEditorEvent {
usagereporter.EmitEditorChangeEvent(user, prevUser.GetRoles(), nil, s.AnonymizeAndSubmit)
usagereporter.EmitEditorChangeEvent(user, prevUser.GetRoles(), nil, a.AnonymizeAndSubmit)
}
return nil
+43 -43
View File
@@ -135,7 +135,7 @@ func (r *CreateUserTokenRequest) CheckAndSetDefaults() error {
}
// CreateResetPasswordToken creates a reset password token
func (s *Server) CreateResetPasswordToken(ctx context.Context, req CreateUserTokenRequest) (types.UserToken, error) {
func (a *Server) CreateResetPasswordToken(ctx context.Context, req CreateUserTokenRequest) (types.UserToken, error) {
err := req.CheckAndSetDefaults()
if err != nil {
return nil, trace.Wrap(err)
@@ -145,32 +145,32 @@ func (s *Server) CreateResetPasswordToken(ctx context.Context, req CreateUserTok
return nil, trace.BadParameter("invalid reset password token request type")
}
_, err = s.ResetPassword(ctx, req.Name)
_, err = a.ResetPassword(ctx, req.Name)
if err != nil {
return nil, trace.Wrap(err)
}
if err := s.resetMFA(ctx, req.Name); err != nil {
if err := a.resetMFA(ctx, req.Name); err != nil {
return nil, trace.Wrap(err)
}
token, err := s.newUserToken(req)
token, err := a.newUserToken(req)
if err != nil {
return nil, trace.Wrap(err)
}
// remove any other existing tokens for this user
err = s.deleteUserTokens(ctx, req.Name)
err = a.deleteUserTokens(ctx, req.Name)
if err != nil {
return nil, trace.Wrap(err)
}
_, err = s.CreateUserToken(ctx, token)
_, err = a.CreateUserToken(ctx, token)
if err != nil {
return nil, trace.Wrap(err)
}
if err := s.emitter.EmitAuditEvent(ctx, &apievents.UserTokenCreate{
if err := a.emitter.EmitAuditEvent(ctx, &apievents.UserTokenCreate{
Metadata: apievents.Metadata{
Type: events.ResetPasswordTokenCreateEvent,
Code: events.ResetPasswordTokenCreateCode,
@@ -179,23 +179,23 @@ func (s *Server) CreateResetPasswordToken(ctx context.Context, req CreateUserTok
ResourceMetadata: apievents.ResourceMetadata{
Name: req.Name,
TTL: req.TTL.String(),
Expires: s.GetClock().Now().UTC().Add(req.TTL),
Expires: a.GetClock().Now().UTC().Add(req.TTL),
},
}); err != nil {
log.WithError(err).Warn("Failed to emit create reset password token event.")
}
return s.GetUserToken(ctx, token.GetName())
return a.GetUserToken(ctx, token.GetName())
}
func (s *Server) resetMFA(ctx context.Context, user string) error {
devs, err := s.Services.GetMFADevices(ctx, user, false)
func (a *Server) resetMFA(ctx context.Context, user string) error {
devs, err := a.Services.GetMFADevices(ctx, user, false)
if err != nil {
return trace.Wrap(err)
}
var errs []error
for _, d := range devs {
errs = append(errs, s.DeleteMFADevice(ctx, user, d.Id))
errs = append(errs, a.DeleteMFADevice(ctx, user, d.Id))
}
return trace.NewAggregate(errs...)
}
@@ -243,7 +243,7 @@ func formatAccountName(s proxyDomainGetter, username string, authHostname string
}
// createTOTPUserTokenSecrets creates new UserTokenSecrets resource for the given token.
func (s *Server) createTOTPUserTokenSecrets(ctx context.Context, token types.UserToken, otpKey *otp.Key) (types.UserTokenSecrets, error) {
func (a *Server) createTOTPUserTokenSecrets(ctx context.Context, token types.UserToken, otpKey *otp.Key) (types.UserTokenSecrets, error) {
// Create QR code.
var otpQRBuf bytes.Buffer
otpImage, err := otpKey.Image(456, 456)
@@ -261,7 +261,7 @@ func (s *Server) createTOTPUserTokenSecrets(ctx context.Context, token types.Use
secrets.SetOTPKey(otpKey.Secret())
secrets.SetQRCode(otpQRBuf.Bytes())
secrets.SetExpiry(token.Expiry())
err = s.UpsertUserTokenSecrets(ctx, secrets)
err = a.UpsertUserTokenSecrets(ctx, secrets)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -269,13 +269,13 @@ func (s *Server) createTOTPUserTokenSecrets(ctx context.Context, token types.Use
return secrets, nil
}
func (s *Server) newTOTPKey(user string) (*otp.Key, *totp.GenerateOpts, error) {
func (a *Server) newTOTPKey(user string) (*otp.Key, *totp.GenerateOpts, error) {
// Fetch account name to display in OTP apps.
accountName, err := formatAccountName(s, user, s.AuthServiceName)
accountName, err := formatAccountName(a, user, a.AuthServiceName)
if err != nil {
return nil, nil, trace.Wrap(err)
}
clusterName, err := s.GetClusterName()
clusterName, err := a.GetClusterName()
if err != nil {
return nil, nil, trace.Wrap(err)
}
@@ -294,7 +294,7 @@ func (s *Server) newTOTPKey(user string) (*otp.Key, *totp.GenerateOpts, error) {
return key, &opts, nil
}
func (s *Server) newUserToken(req CreateUserTokenRequest) (types.UserToken, error) {
func (a *Server) newUserToken(req CreateUserTokenRequest) (types.UserToken, error) {
var err error
var proxyHost string
@@ -310,7 +310,7 @@ func (s *Server) newUserToken(req CreateUserTokenRequest) (types.UserToken, erro
// Get the list of proxies and try and guess the address of the proxy. If
// failed to guess public address, use "<proxyhost>:3080" as a fallback.
proxies, err := s.GetProxies()
proxies, err := a.GetProxies()
if err != nil {
return nil, trace.Wrap(err)
}
@@ -334,9 +334,9 @@ func (s *Server) newUserToken(req CreateUserTokenRequest) (types.UserToken, erro
}
token.SetSubKind(req.Type)
token.SetExpiry(s.clock.Now().UTC().Add(req.TTL))
token.SetExpiry(a.clock.Now().UTC().Add(req.TTL))
token.SetUser(req.Name)
token.SetCreated(s.clock.Now().UTC())
token.SetCreated(a.clock.Now().UTC())
token.SetURL(url)
return token, nil
@@ -364,8 +364,8 @@ func formatUserTokenURL(proxyHost string, tokenID string, reqType string) (strin
}
// deleteUserTokens deletes all user tokens for the specified user.
func (s *Server) deleteUserTokens(ctx context.Context, username string) error {
tokens, err := s.GetUserTokens(ctx)
func (a *Server) deleteUserTokens(ctx context.Context, username string) error {
tokens, err := a.GetUserTokens(ctx)
if err != nil {
return trace.Wrap(err)
}
@@ -375,7 +375,7 @@ func (s *Server) deleteUserTokens(ctx context.Context, username string) error {
continue
}
err = s.DeleteUserToken(ctx, token.GetName())
err = a.DeleteUserToken(ctx, token.GetName())
if err != nil {
return trace.Wrap(err)
}
@@ -386,8 +386,8 @@ func (s *Server) deleteUserTokens(ctx context.Context, username string) error {
// getResetPasswordToken returns user token with subkind set to reset or invite, both
// types which allows users to change their password and set new second factors (if enabled).
func (s *Server) getResetPasswordToken(ctx context.Context, tokenID string) (types.UserToken, error) {
token, err := s.GetUserToken(ctx, tokenID)
func (a *Server) getResetPasswordToken(ctx context.Context, tokenID string) (types.UserToken, error) {
token, err := a.GetUserToken(ctx, tokenID)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -400,7 +400,7 @@ func (s *Server) getResetPasswordToken(ctx context.Context, tokenID string) (typ
}
// createRecoveryToken creates a user token for account recovery.
func (s *Server) createRecoveryToken(ctx context.Context, username, tokenType string, usage types.UserTokenUsage) (types.UserToken, error) {
func (a *Server) createRecoveryToken(ctx context.Context, username, tokenType string, usage types.UserTokenUsage) (types.UserToken, error) {
if tokenType != UserTokenTypeRecoveryStart && tokenType != UserTokenTypeRecoveryApproved {
return nil, trace.BadParameter("invalid recovery token type: %s", tokenType)
}
@@ -418,7 +418,7 @@ func (s *Server) createRecoveryToken(ctx context.Context, username, tokenType st
return nil, trace.Wrap(err)
}
newToken, err := s.newUserToken(req)
newToken, err := a.newUserToken(req)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -426,11 +426,11 @@ func (s *Server) createRecoveryToken(ctx context.Context, username, tokenType st
// Mark what recover type user requested.
newToken.SetUsage(usage)
if _, err := s.CreateUserToken(ctx, newToken); err != nil {
if _, err := a.CreateUserToken(ctx, newToken); err != nil {
return nil, trace.Wrap(err)
}
if err := s.emitter.EmitAuditEvent(ctx, &apievents.UserTokenCreate{
if err := a.emitter.EmitAuditEvent(ctx, &apievents.UserTokenCreate{
Metadata: apievents.Metadata{
Type: events.RecoveryTokenCreateEvent,
Code: events.RecoveryTokenCreateCode,
@@ -439,7 +439,7 @@ func (s *Server) createRecoveryToken(ctx context.Context, username, tokenType st
ResourceMetadata: apievents.ResourceMetadata{
Name: req.Name,
TTL: req.TTL.String(),
Expires: s.GetClock().Now().UTC().Add(req.TTL),
Expires: a.GetClock().Now().UTC().Add(req.TTL),
},
}); err != nil {
log.WithError(err).Warn("Failed to emit create recovery token event.")
@@ -449,13 +449,13 @@ func (s *Server) createRecoveryToken(ctx context.Context, username, tokenType st
}
// CreatePrivilegeToken implements AuthService.CreatePrivilegeToken.
func (s *Server) CreatePrivilegeToken(ctx context.Context, req *proto.CreatePrivilegeTokenRequest) (*types.UserTokenV3, error) {
func (a *Server) CreatePrivilegeToken(ctx context.Context, req *proto.CreatePrivilegeTokenRequest) (*types.UserTokenV3, error) {
username, err := authz.GetClientUsername(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
authPref, err := s.GetAuthPreference(ctx)
authPref, err := a.GetAuthPreference(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -468,7 +468,7 @@ func (s *Server) CreatePrivilegeToken(ctx context.Context, req *proto.CreatePriv
}
tokenKind := UserTokenTypePrivilege
switch hasDevices, err := s.validateMFAAuthResponseForRegister(
switch hasDevices, err := a.validateMFAAuthResponseForRegister(
ctx, req.GetExistingMFAResponse(), username, false /* passwordless */); {
case err != nil:
return nil, trace.Wrap(err)
@@ -477,15 +477,15 @@ func (s *Server) CreatePrivilegeToken(ctx context.Context, req *proto.CreatePriv
}
// Delete any existing user tokens for user before creating.
if err := s.deleteUserTokens(ctx, username); err != nil {
if err := a.deleteUserTokens(ctx, username); err != nil {
return nil, trace.Wrap(err)
}
token, err := s.createPrivilegeToken(ctx, username, tokenKind)
token, err := a.createPrivilegeToken(ctx, username, tokenKind)
return token, trace.Wrap(err)
}
func (s *Server) createPrivilegeToken(ctx context.Context, username, tokenKind string) (*types.UserTokenV3, error) {
func (a *Server) createPrivilegeToken(ctx context.Context, username, tokenKind string) (*types.UserTokenV3, error) {
if tokenKind != UserTokenTypePrivilege && tokenKind != UserTokenTypePrivilegeException {
return nil, trace.BadParameter("invalid privilege token type")
}
@@ -499,17 +499,17 @@ func (s *Server) createPrivilegeToken(ctx context.Context, username, tokenKind s
return nil, trace.Wrap(err)
}
newToken, err := s.newUserToken(req)
newToken, err := a.newUserToken(req)
if err != nil {
return nil, trace.Wrap(err)
}
token, err := s.CreateUserToken(ctx, newToken)
token, err := a.CreateUserToken(ctx, newToken)
if err != nil {
return nil, trace.Wrap(err)
}
if err := s.emitter.EmitAuditEvent(ctx, &apievents.UserTokenCreate{
if err := a.emitter.EmitAuditEvent(ctx, &apievents.UserTokenCreate{
Metadata: apievents.Metadata{
Type: events.PrivilegeTokenCreateEvent,
Code: events.PrivilegeTokenCreateCode,
@@ -518,7 +518,7 @@ func (s *Server) createPrivilegeToken(ctx context.Context, username, tokenKind s
ResourceMetadata: apievents.ResourceMetadata{
Name: req.Name,
TTL: req.TTL.String(),
Expires: s.GetClock().Now().UTC().Add(req.TTL),
Expires: a.GetClock().Now().UTC().Add(req.TTL),
},
}); err != nil {
log.WithError(err).Warn("Failed to emit create privilege token event.")
@@ -533,8 +533,8 @@ func (s *Server) createPrivilegeToken(ctx context.Context, username, tokenKind s
}
// verifyUserToken verifies that the token is not expired and is of the allowed kinds.
func (s *Server) verifyUserToken(token types.UserToken, allowedKinds ...string) error {
if token.Expiry().Before(s.clock.Now().UTC()) {
func (a *Server) verifyUserToken(token types.UserToken, allowedKinds ...string) error {
if token.Expiry().Before(a.clock.Now().UTC()) {
// Provide obscure message on purpose, while logging the real error server side.
log.Debugf("Expired token(%s) type(%s)", token.GetName(), token.GetSubKind())
return trace.AccessDenied("invalid token")