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
}
+36
View File
@@ -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
}
+24 -1
View File
@@ -1,9 +1,11 @@
package site
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"net/http"
"strconv"
"strings"
@@ -100,10 +102,31 @@ func (s *Server) authenticateAndAddIdentity(f response.HandlerFunc, assertions .
return nil, response.Errorf("failed to authenticate: %w", err)
}
if expiration.Before(time.Now()) {
now := time.Now()
if expiration.Before(now) {
return response.TemporaryRedirect("/").Cookie(getExpiredCookie("access_token")), nil
}
// refresh tokens, when the access token is "old enough"
const idTokenLifetime = 10 * time.Hour
if refreshFloor := expiration.Add(-idTokenLifetime); refreshFloor.Before(now) {
accessToken, expiration, err = s.auth.RefreshAccessToken(ctx, accessToken)
if err != nil {
fmt.Println("failed to refresh access token:", err)
return response.TemporaryRedirect("/").
Body(io.NopCloser(bytes.NewBuffer([]byte(fmt.Sprintf("failed to refresh access token: %v", err))))).
Cookie(getExpiredCookie("access_token")), nil
}
// 'redirect' to same url, to set the new access_token cookie
return response.TemporaryRedirect(r.URL.String()).
Cookie(newAccessTokenCookie(accessToken, expiration)), nil
}
// add identity info to request context
user, acct, err := s.accts.GetUserAndAccountByAccessToken(ctx, accessToken)
if err != nil {
return nil, response.Errorf("failed to authorize: %w", err)
+13 -8
View File
@@ -4,6 +4,7 @@ import (
"fmt"
"net/http"
"ruben/inventory2/internal/site/response"
"time"
)
// GET /login
@@ -53,14 +54,18 @@ func (s *Server) loginCallback(r *http.Request) (response.Response, error) {
// set access_token cookie and redirect to a reasonable place
return response.TemporaryRedirect("/").
Cookie(http.Cookie{
Name: "access_token",
Value: accessToken,
Path: "/",
Expires: expiration,
MaxAge: 0, // using Expiration instead
Secure: true,
}), nil
Cookie(newAccessTokenCookie(accessToken, expiration)), nil
}
func newAccessTokenCookie(tkn string, expiration time.Time) http.Cookie {
return http.Cookie{
Name: "access_token",
Value: tkn,
Path: "/",
Expires: expiration,
MaxAge: 0, // using Expiration instead
Secure: true,
}
}
// GET /logout