[release-12.2.1] Auth: Fix render user OAuth passthrough (#112092)

Auth: Fix render user OAuth passthrough (#111636)

* devenv: fix volumes section when sources don't contain one

* wip

* Working correctly with improvedExternalSessionHandling on

* Remove not needed lines

* Working with the old flow, tests

* Handle compatibility with the feature toggle, tests wip

* Tests

* Cleanup

* Address feedback

* Align tests

* Add comment

* Fix issue with session removal after the invalidation of tokens

* Remove commented out code

* clean up

(cherry picked from commit 53f4803e98)

Co-authored-by: Misi <mgyongyosi@users.noreply.github.com>
This commit is contained in:
grafana-delivery-bot[bot]
2025-10-09 15:47:18 +02:00
committed by GitHub
co-authored by Misi
parent b3551d6d0d
commit 05e3befe99
12 changed files with 1911 additions and 322 deletions
+142 -89
View File
@@ -10,6 +10,7 @@ import (
"github.com/go-jose/go-jose/v3/jwt"
"github.com/prometheus/client_golang/prometheus"
"go.opentelemetry.io/otel/attribute"
"go.opentelemetry.io/otel/codes"
"go.opentelemetry.io/otel/trace"
"golang.org/x/oauth2"
@@ -55,8 +56,14 @@ var _ OAuthTokenService = (*Service)(nil)
type OAuthTokenService interface {
GetCurrentOAuthToken(context.Context, identity.Requester, *auth.UserToken) *oauth2.Token
IsOAuthPassThruEnabled(*datasources.DataSource) bool
TryTokenRefresh(context.Context, identity.Requester, *auth.UserToken) (*oauth2.Token, error)
InvalidateOAuthTokens(context.Context, identity.Requester, *auth.UserToken) error
TryTokenRefresh(context.Context, identity.Requester, *TokenRefreshMetadata) (*oauth2.Token, error)
InvalidateOAuthTokens(context.Context, identity.Requester, *TokenRefreshMetadata) error
}
type TokenRefreshMetadata struct {
ExternalSessionID int64
AuthModule string
AuthID string
}
func ProvideService(socialService social.Service, authInfoService login.AuthInfoService, cfg *setting.Cfg, registerer prometheus.Registerer,
@@ -100,51 +107,71 @@ func (o *Service) GetCurrentOAuthToken(ctx context.Context, usr identity.Request
ctxLogger = ctxLogger.New("userID", userID)
if !strings.HasPrefix(usr.GetAuthenticatedBy(), "oauth_") {
tokenRefreshMetadata := &TokenRefreshMetadata{
ExternalSessionID: 0,
}
var persistedToken *oauth2.Token
// Find the external session associated with the user and session token
// regardless of the improvedExternalSessionHandling feature toggle,
// because Grafana writes and updates both tables to make the switch
// to the new session handling smoother.
externalSession, err := o.getExternalSession(ctx, usr, userID, sessionToken)
if err != nil && !errors.Is(err, auth.ErrExternalSessionNotFound) {
ctxLogger.Error("Failed to get external session", "error", err)
return nil
}
// If the feature toggle is enabled, an external session is required.
if o.features.IsEnabledGlobally(featuremgmt.FlagImprovedExternalSessionHandling) && (externalSession == nil || errors.Is(err, auth.ErrExternalSessionNotFound)) {
ctxLogger.Error("No external session found for user", "userID", userID)
return nil
}
// externalSession can be nil if Grafana was updated from a version where the
// external session table was not used yet (did not exist) and the user has not logged in since
// the version update (therefore no external session was created for the user yet).
if externalSession != nil {
tokenRefreshMetadata.ExternalSessionID = externalSession.ID
}
authInfo, err := o.AuthInfoService.GetAuthInfo(ctx, &login.GetAuthInfoQuery{
UserId: userID,
})
if err != nil {
if errors.Is(err, user.ErrUserNotFound) {
ctxLogger.Warn("No AuthInfo found for user", "userID", userID)
return nil
}
ctxLogger.Error("Failed to fetch AuthInfo for user", "userID", userID, "error", err)
return nil
}
tokenRefreshMetadata.AuthID = authInfo.AuthId
tokenRefreshMetadata.AuthModule = authInfo.AuthModule
if !strings.HasPrefix(tokenRefreshMetadata.AuthModule, "oauth_") {
ctxLogger.Warn("The specified user's auth provider is not oauth",
"authmodule", usr.GetAuthenticatedBy())
"authmodule", tokenRefreshMetadata.AuthModule)
return nil
}
var persistedToken *oauth2.Token
if o.features.IsEnabledGlobally(featuremgmt.FlagImprovedExternalSessionHandling) {
externalSession, err := o.sessionService.GetExternalSession(ctx, sessionToken.ExternalSessionId)
if err != nil {
if errors.Is(err, auth.ErrExternalSessionNotFound) {
return nil
}
ctxLogger.Error("Failed to fetch external session", "error", err)
return nil
}
persistedToken = buildOAuthTokenFromExternalSession(externalSession)
if persistedToken.RefreshToken == "" {
return persistedToken
}
} else {
authInfo, ok, _ := o.hasOAuthEntry(ctx, usr)
if !ok {
return nil
}
if err := checkOAuthRefreshToken(authInfo); err != nil {
if errors.Is(err, ErrNoRefreshTokenFound) {
return buildOAuthTokenFromAuthInfo(authInfo)
}
return nil
}
persistedToken = buildOAuthTokenFromAuthInfo(authInfo)
}
if persistedToken.RefreshToken == "" {
return persistedToken
}
refreshNeeded := needTokenRefresh(ctx, persistedToken)
if !refreshNeeded {
return persistedToken
}
token, err := o.TryTokenRefresh(ctx, usr, sessionToken)
token, err := o.TryTokenRefresh(ctx, usr, tokenRefreshMetadata)
if err != nil {
if errors.Is(err, ErrNoRefreshTokenFound) {
return persistedToken
@@ -212,7 +239,7 @@ func (o *Service) hasOAuthEntry(ctx context.Context, usr identity.Requester) (*l
// TryTokenRefresh returns an error in case the OAuth token refresh was unsuccessful
// It uses a server lock to prevent getting the Refresh Token multiple times for a given User
func (o *Service) TryTokenRefresh(ctx context.Context, usr identity.Requester, sessionToken *auth.UserToken) (*oauth2.Token, error) {
func (o *Service) TryTokenRefresh(ctx context.Context, usr identity.Requester, tokenRefreshMetadata *TokenRefreshMetadata) (*oauth2.Token, error) {
ctx, span := o.tracer.Start(ctx, "oauthtoken.TryTokenRefresh")
defer span.End()
@@ -237,14 +264,13 @@ func (o *Service) TryTokenRefresh(ctx context.Context, usr identity.Requester, s
ctxLogger = ctxLogger.New("userID", userID)
// get the token's auth provider (f.e. azuread)
currAuthenticator := usr.GetAuthenticatedBy()
if !strings.HasPrefix(currAuthenticator, "oauth") {
ctxLogger.Warn("The specified user's auth provider is not OAuth", "authmodule", currAuthenticator)
if !strings.HasPrefix(tokenRefreshMetadata.AuthModule, "oauth_") {
ctxLogger.Warn("The specified user's auth provider is not oauth",
"authmodule", tokenRefreshMetadata.AuthModule)
return nil, nil
}
provider := strings.TrimPrefix(currAuthenticator, "oauth_")
provider := strings.TrimPrefix(tokenRefreshMetadata.AuthModule, "oauth_")
currentOAuthInfo := o.SocialService.GetOAuthInfoProvider(provider)
if currentOAuthInfo == nil {
ctxLogger.Warn("OAuth provider not found", "provider", provider)
@@ -259,7 +285,7 @@ func (o *Service) TryTokenRefresh(ctx context.Context, usr identity.Requester, s
lockKey := fmt.Sprintf("oauth-refresh-token-%d", userID)
if o.features.IsEnabledGlobally(featuremgmt.FlagImprovedExternalSessionHandling) {
lockKey = fmt.Sprintf("oauth-refresh-token-%d-%d", userID, sessionToken.ExternalSessionId)
lockKey = fmt.Sprintf("oauth-refresh-token-%d-%d", userID, tokenRefreshMetadata.ExternalSessionID)
}
lockTimeConfig := serverlock.LockTimeConfig{
@@ -288,7 +314,7 @@ func (o *Service) TryTokenRefresh(ctx context.Context, usr identity.Requester, s
var persistedToken *oauth2.Token
var externalSession *auth.ExternalSession
if o.features.IsEnabledGlobally(featuremgmt.FlagImprovedExternalSessionHandling) {
externalSession, err = o.sessionService.GetExternalSession(ctx, sessionToken.ExternalSessionId)
externalSession, err = o.sessionService.GetExternalSession(ctx, tokenRefreshMetadata.ExternalSessionID)
if err != nil {
if errors.Is(err, auth.ErrExternalSessionNotFound) {
ctxLogger.Error("External session was not found for user", "error", err)
@@ -319,7 +345,7 @@ func (o *Service) TryTokenRefresh(ctx context.Context, usr identity.Requester, s
return
}
newToken, cmdErr = o.tryGetOrRefreshOAuthToken(ctx, persistedToken, usr, sessionToken)
newToken, cmdErr = o.tryGetOrRefreshOAuthToken(ctx, persistedToken, usr, tokenRefreshMetadata)
}, retryOpt)
if lockErr != nil {
ctxLogger.Error("Failed to obtain token refresh lock", "error", lockErr)
@@ -328,14 +354,14 @@ func (o *Service) TryTokenRefresh(ctx context.Context, usr identity.Requester, s
// Silence ErrNoRefreshTokenFound
if errors.Is(cmdErr, ErrNoRefreshTokenFound) {
return nil, nil
return nil, ErrNoRefreshTokenFound
}
return newToken, cmdErr
}
// InvalidateOAuthTokens invalidates the OAuth tokens (access_token, refresh_token) and sets the Expiry to default/zero
func (o *Service) InvalidateOAuthTokens(ctx context.Context, usr identity.Requester, sessionToken *auth.UserToken) error {
func (o *Service) InvalidateOAuthTokens(ctx context.Context, usr identity.Requester, tokenRefreshMetadata *TokenRefreshMetadata) error {
userID, err := usr.GetInternalID()
if err != nil {
logger.Error("Failed to convert user id to int", "id", usr.GetID(), "error", err)
@@ -345,7 +371,7 @@ func (o *Service) InvalidateOAuthTokens(ctx context.Context, usr identity.Reques
ctxLogger := logger.FromContext(ctx).New("userID", userID)
if o.features.IsEnabledGlobally(featuremgmt.FlagImprovedExternalSessionHandling) {
err := o.sessionService.UpdateExternalSession(ctx, sessionToken.ExternalSessionId, &auth.UpdateExternalSessionCommand{
err := o.sessionService.UpdateExternalSession(ctx, tokenRefreshMetadata.ExternalSessionID, &auth.UpdateExternalSessionCommand{
Token: &oauth2.Token{},
})
if err != nil {
@@ -356,8 +382,8 @@ func (o *Service) InvalidateOAuthTokens(ctx context.Context, usr identity.Reques
return o.AuthInfoService.UpdateAuthInfo(ctx, &login.UpdateAuthInfoCommand{
UserId: userID,
AuthModule: usr.GetAuthenticatedBy(),
AuthId: usr.GetAuthID(),
AuthModule: tokenRefreshMetadata.AuthModule,
AuthId: tokenRefreshMetadata.AuthID,
OAuthToken: &oauth2.Token{
AccessToken: "",
RefreshToken: "",
@@ -366,13 +392,14 @@ func (o *Service) InvalidateOAuthTokens(ctx context.Context, usr identity.Reques
})
}
func (o *Service) tryGetOrRefreshOAuthToken(ctx context.Context, persistedToken *oauth2.Token, usr identity.Requester, sessionToken *auth.UserToken) (*oauth2.Token, error) {
func (o *Service) tryGetOrRefreshOAuthToken(ctx context.Context, persistedToken *oauth2.Token, usr identity.Requester, tokenRefreshMetadata *TokenRefreshMetadata) (*oauth2.Token, error) {
ctx, span := o.tracer.Start(ctx, "oauthtoken.tryGetOrRefreshOAuthToken")
defer span.End()
userID, err := usr.GetInternalID()
if err != nil {
logger.Error("Failed to convert user id to int", "id", usr.GetID(), "error", err)
span.SetStatus(codes.Error, "Failed to convert user id to int")
return nil, err
}
@@ -380,8 +407,11 @@ func (o *Service) tryGetOrRefreshOAuthToken(ctx context.Context, persistedToken
ctxLogger := logger.FromContext(ctx).New("userID", userID)
// tryGetOrRefreshOAuthToken assumes that the AuthModule has RefreshToken enabled
// which is checked by the caller (TryTokenRefresh)
if persistedToken.RefreshToken == "" {
ctxLogger.Warn("No refresh token available", "authmodule", usr.GetAuthenticatedBy())
ctxLogger.Error("No refresh token available", "authmodule", tokenRefreshMetadata.AuthModule)
span.SetStatus(codes.Error, ErrNoRefreshTokenFound.Error())
return nil, ErrNoRefreshTokenFound
}
@@ -390,50 +420,44 @@ func (o *Service) tryGetOrRefreshOAuthToken(ctx context.Context, persistedToken
return persistedToken, nil
}
authProvider := usr.GetAuthenticatedBy()
connect, err := o.SocialService.GetConnector(authProvider)
connect, err := o.SocialService.GetConnector(tokenRefreshMetadata.AuthModule)
if err != nil {
ctxLogger.Error("Failed to get oauth connector", "provider", authProvider, "error", err)
ctxLogger.Error("Failed to get oauth connector", "provider", tokenRefreshMetadata.AuthModule, "error", err)
span.SetStatus(codes.Error, "Failed to get oauth connector: "+err.Error())
return nil, err
}
client, err := o.SocialService.GetOAuthHttpClient(authProvider)
client, err := o.SocialService.GetOAuthHttpClient(tokenRefreshMetadata.AuthModule)
if err != nil {
ctxLogger.Error("Failed to get oauth http client", "provider", authProvider, "error", err)
ctxLogger.Error("Failed to get oauth http client", "provider", tokenRefreshMetadata.AuthModule, "error", err)
span.SetStatus(codes.Error, "Failed to get oauth http client")
return nil, err
}
ctx = context.WithValue(ctx, oauth2.HTTPClient, client)
start := time.Now()
// TokenSource handles refreshing the token if it has expired
token, err := connect.TokenSource(ctx, persistedToken).Token()
token, refreshErr := connect.TokenSource(ctx, persistedToken).Token()
duration := time.Since(start)
o.tokenRefreshDuration.WithLabelValues(authProvider, fmt.Sprintf("%t", err == nil)).Observe(duration.Seconds())
o.tokenRefreshDuration.WithLabelValues(tokenRefreshMetadata.AuthModule, fmt.Sprintf("%t", err == nil)).Observe(duration.Seconds())
if err != nil {
if refreshErr != nil {
span.SetAttributes(attribute.Bool("token_refreshed", false))
ctxLogger.Error("Failed to retrieve oauth access token",
"provider", usr.GetAuthenticatedBy(), "error", err)
"provider", tokenRefreshMetadata.AuthModule, "error", refreshErr)
// token refresh failed, invalidate the old token
if err := o.InvalidateOAuthTokens(ctx, usr, sessionToken); err != nil {
ctxLogger.Warn("Failed to invalidate OAuth tokens", "authID", usr.GetAuthID(), "error", err)
if err := o.InvalidateOAuthTokens(ctx, usr, tokenRefreshMetadata); err != nil {
ctxLogger.Warn("Failed to invalidate OAuth tokens", "authID", tokenRefreshMetadata.AuthID, "error", err)
}
return nil, err
return nil, refreshErr
}
span.SetAttributes(attribute.Bool("token_refreshed", true))
// If the tokens are not the same, update the entry in the DB
if !tokensEq(persistedToken, token) {
updateAuthCommand := &login.UpdateAuthInfoCommand{
UserId: userID,
AuthModule: usr.GetAuthenticatedBy(),
AuthId: usr.GetAuthID(),
OAuthToken: token,
}
if o.Cfg.Env == setting.Dev {
ctxLogger.Debug("Oauth got token",
"auth_module", usr.GetAuthenticatedBy(),
@@ -444,17 +468,32 @@ func (o *Service) tryGetOrRefreshOAuthToken(ctx context.Context, persistedToken
}
if !o.features.IsEnabledGlobally(featuremgmt.FlagImprovedExternalSessionHandling) {
updateAuthCommand := &login.UpdateAuthInfoCommand{
UserId: userID,
AuthModule: tokenRefreshMetadata.AuthModule,
AuthId: tokenRefreshMetadata.AuthID,
OAuthToken: token,
}
if err := o.AuthInfoService.UpdateAuthInfo(ctx, updateAuthCommand); err != nil {
ctxLogger.Error("Failed to update auth info during token refresh", "authID", usr.GetAuthID(), "error", err)
ctxLogger.Error("Failed to update auth info during token refresh", "authID", tokenRefreshMetadata.AuthID, "error", err)
span.SetStatus(codes.Error, "Failed to update auth info during token refresh")
return nil, err
}
}
if err := o.sessionService.UpdateExternalSession(ctx, sessionToken.ExternalSessionId, &auth.UpdateExternalSessionCommand{
Token: token,
}); err != nil {
ctxLogger.Error("Failed to update external session during token refresh", "error", err)
return nil, err
// Update the external session with the new token if we the user has an external session,
// regardless of the feature flag state to keep the `user_external_session` table in sync.
// ExternalSessionID should always be set except for some edge cases:
// - when Grafana was updated to a version where the `improvedExternalSessionHandling` feature flag
// was enabled after the user logged in
if tokenRefreshMetadata.ExternalSessionID != 0 {
if err := o.sessionService.UpdateExternalSession(ctx, tokenRefreshMetadata.ExternalSessionID, &auth.UpdateExternalSessionCommand{
Token: token,
}); err != nil {
ctxLogger.Error("Failed to update external session during token refresh", "error", err)
span.SetStatus(codes.Error, "Failed to update external session during token refresh")
return nil, err
}
}
ctxLogger.Debug("Updated oauth info for user")
@@ -500,6 +539,11 @@ func needTokenRefresh(ctx context.Context, persistedToken *oauth2.Token) bool {
ctxLogger := logger.FromContext(ctx)
if persistedToken.AccessToken == "" {
ctxLogger.Debug("Access token has been cleared, need to refresh")
return true
}
idTokenExp, err := GetIDTokenExpiry(persistedToken)
if err != nil {
ctxLogger.Warn("Could not get ID Token expiry", "error", err)
@@ -550,22 +594,6 @@ func buildOAuthTokenFromExternalSession(externalSession *auth.ExternalSession) *
return token
}
func checkOAuthRefreshToken(authInfo *login.UserAuth) error {
if !strings.Contains(authInfo.AuthModule, "oauth") {
logger.Warn("The specified user's auth provider is not oauth",
"authmodule", authInfo.AuthModule, "userid", authInfo.UserId)
return ErrNotAnOAuthProvider
}
if authInfo.OAuthRefreshToken == "" {
logger.Warn("No refresh token available",
"authmodule", authInfo.AuthModule, "userid", authInfo.UserId)
return ErrNoRefreshTokenFound
}
return nil
}
// GetIDTokenExpiry extracts the expiry time from the ID token
func GetIDTokenExpiry(token *oauth2.Token) (time.Time, error) {
idToken, ok := token.Extra("id_token").(string)
@@ -598,3 +626,28 @@ func getExpiryWithSkew(expiry time.Time) (adjustedExpiry time.Time, hasTokenExpi
hasTokenExpired = adjustedExpiry.Before(time.Now())
return
}
// getExternalSession fetches the external session based on the user and session token.
// When using the render module, it fetches the most recent external session for the user
// since the session token ID is not available.
// For regular users, it uses the session token ID to fetch the external session.
func (o *Service) getExternalSession(ctx context.Context, usr identity.Requester, userID int64, sessionToken *auth.UserToken) (*auth.ExternalSession, error) {
if usr.GetAuthenticatedBy() == login.RenderModule {
// When using render module, we don't have the session token ID, so we need to fetch the most recent session
// entry for the user (as it is done with the old flow).
// In the future, we might want to consider passing the session token ID to the render module to make this more robust.
externalSessions, err := o.sessionService.FindExternalSessions(ctx, &auth.ListExternalSessionQuery{UserID: userID})
if err != nil {
return nil, err
}
if len(externalSessions) == 0 || externalSessions[0] == nil {
return nil, auth.ErrExternalSessionNotFound
}
return externalSessions[0], nil
}
// For regular users, we use the session token ID to fetch the external session
return o.sessionService.GetExternalSession(ctx, sessionToken.ExternalSessionId)
}
File diff suppressed because it is too large Load Diff
@@ -8,13 +8,14 @@ import (
"github.com/grafana/grafana/pkg/apimachinery/identity"
"github.com/grafana/grafana/pkg/services/auth"
"github.com/grafana/grafana/pkg/services/datasources"
"github.com/grafana/grafana/pkg/services/oauthtoken"
)
type MockOauthTokenService struct {
GetCurrentOauthTokenFunc func(ctx context.Context, usr identity.Requester, sessionToken *auth.UserToken) *oauth2.Token
IsOAuthPassThruEnabledFunc func(ds *datasources.DataSource) bool
InvalidateOAuthTokensFunc func(ctx context.Context, usr identity.Requester, sessionToken *auth.UserToken) error
TryTokenRefreshFunc func(ctx context.Context, usr identity.Requester, sessionToken *auth.UserToken) (*oauth2.Token, error)
InvalidateOAuthTokensFunc func(ctx context.Context, usr identity.Requester, metadata *oauthtoken.TokenRefreshMetadata) error
TryTokenRefreshFunc func(ctx context.Context, usr identity.Requester, metadata *oauthtoken.TokenRefreshMetadata) (*oauth2.Token, error)
}
func (m *MockOauthTokenService) GetCurrentOAuthToken(ctx context.Context, usr identity.Requester, sessionToken *auth.UserToken) *oauth2.Token {
@@ -31,16 +32,16 @@ func (m *MockOauthTokenService) IsOAuthPassThruEnabled(ds *datasources.DataSourc
return false
}
func (m *MockOauthTokenService) InvalidateOAuthTokens(ctx context.Context, usr identity.Requester, sessionToken *auth.UserToken) error {
func (m *MockOauthTokenService) InvalidateOAuthTokens(ctx context.Context, usr identity.Requester, metadata *oauthtoken.TokenRefreshMetadata) error {
if m.InvalidateOAuthTokensFunc != nil {
return m.InvalidateOAuthTokensFunc(ctx, usr, sessionToken)
return m.InvalidateOAuthTokensFunc(ctx, usr, metadata)
}
return nil
}
func (m *MockOauthTokenService) TryTokenRefresh(ctx context.Context, usr identity.Requester, sessionToken *auth.UserToken) (*oauth2.Token, error) {
func (m *MockOauthTokenService) TryTokenRefresh(ctx context.Context, usr identity.Requester, metadata *oauthtoken.TokenRefreshMetadata) (*oauth2.Token, error) {
if m.TryTokenRefreshFunc != nil {
return m.TryTokenRefreshFunc(ctx, usr, sessionToken)
return m.TryTokenRefreshFunc(ctx, usr, metadata)
}
return nil, nil
}
@@ -29,10 +29,10 @@ func (s *Service) IsOAuthPassThruEnabled(ds *datasources.DataSource) bool {
return oauthtoken.IsOAuthPassThruEnabled(ds)
}
func (s *Service) TryTokenRefresh(context.Context, identity.Requester, *auth.UserToken) (*oauth2.Token, error) {
func (s *Service) TryTokenRefresh(context.Context, identity.Requester, *oauthtoken.TokenRefreshMetadata) (*oauth2.Token, error) {
return s.Token, nil
}
func (s *Service) InvalidateOAuthTokens(context.Context, identity.Requester, *auth.UserToken) error {
func (s *Service) InvalidateOAuthTokens(context.Context, identity.Requester, *oauthtoken.TokenRefreshMetadata) error {
return nil
}