implemented access token refreshing - untested
This commit is contained in:
@@ -95,19 +95,14 @@ func (a *Authenticator) Exchange(ctx context.Context, state, code string) (acces
|
||||
|
||||
// obtain token and profile
|
||||
|
||||
token, err := a.Config.Exchange(ctx, code)
|
||||
tkn, err := a.Config.Exchange(ctx, code)
|
||||
if err != nil {
|
||||
return "", time.Time{}, fmt.Errorf("failed to exchange an authorization code for a token: %w", err)
|
||||
}
|
||||
|
||||
idToken, err := a.VerifyIDToken(ctx, token)
|
||||
idToken, claims, err := a.verifyIDTokenAndClaimsFromToken(ctx, tkn)
|
||||
if err != nil {
|
||||
return "", time.Time{}, fmt.Errorf("failed to verify ID Token: %w", err)
|
||||
}
|
||||
|
||||
var claims map[string]any
|
||||
if err := idToken.Claims(&claims); err != nil {
|
||||
return "", time.Time{}, fmt.Errorf("Failed to obtain id token claims: %w", err)
|
||||
return "", time.Time{}, err
|
||||
}
|
||||
|
||||
claimsJSON, _ := json.Marshal(claims)
|
||||
@@ -171,10 +166,10 @@ func (a *Authenticator) Exchange(ctx context.Context, state, code string) (acces
|
||||
the_user
|
||||
`,
|
||||
pgx.NamedArgs{
|
||||
"access_token": token.AccessToken,
|
||||
"token_type": token.TokenType,
|
||||
"refresh_token": token.RefreshToken,
|
||||
"expiry": token.Expiry,
|
||||
"access_token": tkn.AccessToken,
|
||||
"token_type": tkn.TokenType,
|
||||
"refresh_token": tkn.RefreshToken,
|
||||
"expiry": tkn.Expiry,
|
||||
|
||||
"id_token_issuer": idToken.Issuer,
|
||||
"id_token_audience": pgtype.FlatArray[string](idToken.Audience),
|
||||
@@ -190,7 +185,7 @@ func (a *Authenticator) Exchange(ctx context.Context, state, code string) (acces
|
||||
return "", time.Time{}, fmt.Errorf("failed to perform query to save tokens: %w", err)
|
||||
}
|
||||
|
||||
return token.AccessToken, token.Expiry.UTC(), nil
|
||||
return tkn.AccessToken, tkn.Expiry.UTC(), nil
|
||||
}
|
||||
|
||||
// VerifyIDToken verifies that an *oauth2.Token is a valid *oidc.IDToken.
|
||||
@@ -223,3 +218,142 @@ func (a *Authenticator) GetLogoutURL(requestHost string) *url.URL {
|
||||
}.Encode(),
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Authenticator) RefreshAccessToken(
|
||||
ctx context.Context,
|
||||
oldAccessToken string,
|
||||
) (
|
||||
accessToken string,
|
||||
expiration time.Time,
|
||||
err error,
|
||||
) {
|
||||
|
||||
refreshToken, tokenType, err := a.getRefreshTokenForAccessToken(ctx, accessToken)
|
||||
if err != nil {
|
||||
return "", time.Time{}, fmt.Errorf("failed to load refresh token: %w", err)
|
||||
}
|
||||
|
||||
tkn, err := a.TokenSource(ctx, &oauth2.Token{
|
||||
// AccessToken is the token that authorizes and authenticates
|
||||
// the requests.
|
||||
AccessToken: oldAccessToken,
|
||||
|
||||
// TokenType is the type of token.
|
||||
// The Type method returns either this or "Bearer", the default.
|
||||
TokenType: tokenType,
|
||||
|
||||
// RefreshToken is a token that's used by the application
|
||||
// (as opposed to the user) to refresh the access token
|
||||
// if it expires.
|
||||
RefreshToken: refreshToken,
|
||||
|
||||
/*
|
||||
// Expiry is the optional expiration time of the access token.
|
||||
//
|
||||
// If zero, TokenSource implementations will reuse the same
|
||||
// token forever and RefreshToken or equivalent
|
||||
// mechanisms for that TokenSource will not be used.
|
||||
Expiry time.Time `json:"expiry,omitempty"`
|
||||
*/
|
||||
}).Token()
|
||||
if err != nil {
|
||||
return "", time.Time{}, fmt.Errorf("failed to fetch refresh token: %w", err)
|
||||
}
|
||||
|
||||
idToken, claims, err := a.verifyIDTokenAndClaimsFromToken(ctx, tkn)
|
||||
if err != nil {
|
||||
return "", time.Time{}, err
|
||||
}
|
||||
|
||||
claimsJSON, _ := json.Marshal(claims)
|
||||
|
||||
if _, err = a.db.Exec(
|
||||
ctx,
|
||||
`
|
||||
WITH deleted_tokens AS (
|
||||
DELETE FROM
|
||||
oauth_tokens
|
||||
WHERE
|
||||
access_token = @old_access_token
|
||||
RETURNING
|
||||
access_token AS old_access_token
|
||||
), new_tokens AS (
|
||||
INSERT INTO oauth_tokens (
|
||||
access_token,
|
||||
token_type,
|
||||
refresh_token,
|
||||
expiry,
|
||||
|
||||
id_token_issuer,
|
||||
id_token_audience,
|
||||
id_token_subject,
|
||||
id_token_expiry,
|
||||
id_token_issued_at,
|
||||
id_token_nonce,
|
||||
id_token_access_token_hash,
|
||||
|
||||
claims -- might want to open this up
|
||||
)
|
||||
VALUES (
|
||||
@new_access_token,
|
||||
@token_type,
|
||||
@refresh_token,
|
||||
@expiry,
|
||||
|
||||
@id_token_issuer,
|
||||
@id_token_audience,
|
||||
@id_token_subject,
|
||||
@id_token_expiry,
|
||||
@id_token_issued_at,
|
||||
@id_token_nonce,
|
||||
@id_token_access_token_hash,
|
||||
|
||||
@claims
|
||||
)
|
||||
RETURNING
|
||||
access_token AS new_access_token
|
||||
)
|
||||
SELECT
|
||||
old_access_token,
|
||||
new_access_token
|
||||
FROM
|
||||
deleted_tokens
|
||||
FULL JOIN
|
||||
new_tokens
|
||||
`,
|
||||
pgx.NamedArgs{
|
||||
"old_access_token": oldAccessToken,
|
||||
|
||||
"new_access_token": tkn.AccessToken,
|
||||
"token_type": tkn.TokenType,
|
||||
"refresh_token": tkn.RefreshToken,
|
||||
"expiry": tkn.Expiry,
|
||||
|
||||
"id_token_issuer": idToken.Issuer,
|
||||
"id_token_audience": pgtype.FlatArray[string](idToken.Audience),
|
||||
"id_token_subject": idToken.Subject,
|
||||
"id_token_expiry": idToken.Expiry,
|
||||
"id_token_issued_at": idToken.IssuedAt,
|
||||
"id_token_nonce": idToken.Nonce,
|
||||
"id_token_access_token_hash": idToken.AccessTokenHash,
|
||||
|
||||
"claims": json.RawMessage(claimsJSON),
|
||||
},
|
||||
); err != nil {
|
||||
return "", time.Time{}, fmt.Errorf("failed to save new access token and delete old access token: %w", err)
|
||||
}
|
||||
|
||||
return tkn.AccessToken, tkn.Expiry.UTC(), nil
|
||||
}
|
||||
|
||||
func (a *Authenticator) verifyIDTokenAndClaimsFromToken(ctx context.Context, tkn *oauth2.Token) (idToken *oidc.IDToken, claims map[string]any, err error) {
|
||||
if idToken, err = a.VerifyIDToken(ctx, tkn); err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to verify id Token: %w", err)
|
||||
}
|
||||
|
||||
if err := idToken.Claims(&claims); err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to obtain id token claims: %w", err)
|
||||
}
|
||||
|
||||
return idToken, claims, nil
|
||||
}
|
||||
|
||||
@@ -113,3 +113,39 @@ func (a *Authenticator) GetAccessTokenClaimsAndExpiration(ctx context.Context, a
|
||||
|
||||
return claims, r.Expiry, nil
|
||||
}
|
||||
|
||||
func (a *Authenticator) getRefreshTokenForAccessToken(ctx context.Context, accessToken string) (refreshToken, tokenType string, err error) {
|
||||
rows, err := a.db.Query(
|
||||
ctx,
|
||||
`
|
||||
SELECT
|
||||
refresh_token,
|
||||
token_type
|
||||
FROM
|
||||
oauth_tokens
|
||||
WHERE
|
||||
access_token = @access_token
|
||||
`,
|
||||
pgx.NamedArgs{
|
||||
"access_token": accessToken,
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("failed to perform query: %w", err)
|
||||
}
|
||||
|
||||
type Row struct {
|
||||
Refresh_token string
|
||||
Token_type string
|
||||
}
|
||||
|
||||
r, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByNameLax[Row])
|
||||
if err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return "", "", consts.ErrNotFound
|
||||
}
|
||||
return "", "", fmt.Errorf("failed to scan row: %w", err)
|
||||
}
|
||||
|
||||
return r.Refresh_token, r.Token_type, nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user