implemented access token refreshing - untested

This commit is contained in:
2026-01-04 21:34:24 -07:00
parent e36c759dc2
commit 8ccc795812
5 changed files with 221 additions and 23 deletions
+147 -13
View File
@@ -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
}