implemented access token refreshing - untested
This commit is contained in:
@@ -8,7 +8,7 @@
|
|||||||
- [x] move auth state stuff to database (out of cache)
|
- [x] move auth state stuff to database (out of cache)
|
||||||
- [ ] only generate a sign up link IF they click the link on the accounts page
|
- [ ] only generate a sign up link IF they click the link on the accounts page
|
||||||
- [ ] get api key approved
|
- [ ] get api key approved
|
||||||
- [ ] Get new access token using refresh token flow
|
- [o] NEEDS TESTING - Get new access token using refresh token flow
|
||||||
- [ ] Get new refresh token flow
|
- [ ] Get new refresh token flow
|
||||||
- [ ] make a FK between the etsy_store_events table and etsy_users table (store_id columns don't match types)
|
- [ ] make a FK between the etsy_store_events table and etsy_users table (store_id columns don't match types)
|
||||||
- [ ] Auth0
|
- [ ] Auth0
|
||||||
|
|||||||
@@ -95,19 +95,14 @@ func (a *Authenticator) Exchange(ctx context.Context, state, code string) (acces
|
|||||||
|
|
||||||
// obtain token and profile
|
// obtain token and profile
|
||||||
|
|
||||||
token, err := a.Config.Exchange(ctx, code)
|
tkn, err := a.Config.Exchange(ctx, code)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", time.Time{}, fmt.Errorf("failed to exchange an authorization code for a token: %w", err)
|
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 {
|
if err != nil {
|
||||||
return "", time.Time{}, fmt.Errorf("failed to verify ID Token: %w", err)
|
return "", time.Time{}, 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)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
claimsJSON, _ := json.Marshal(claims)
|
claimsJSON, _ := json.Marshal(claims)
|
||||||
@@ -171,10 +166,10 @@ func (a *Authenticator) Exchange(ctx context.Context, state, code string) (acces
|
|||||||
the_user
|
the_user
|
||||||
`,
|
`,
|
||||||
pgx.NamedArgs{
|
pgx.NamedArgs{
|
||||||
"access_token": token.AccessToken,
|
"access_token": tkn.AccessToken,
|
||||||
"token_type": token.TokenType,
|
"token_type": tkn.TokenType,
|
||||||
"refresh_token": token.RefreshToken,
|
"refresh_token": tkn.RefreshToken,
|
||||||
"expiry": token.Expiry,
|
"expiry": tkn.Expiry,
|
||||||
|
|
||||||
"id_token_issuer": idToken.Issuer,
|
"id_token_issuer": idToken.Issuer,
|
||||||
"id_token_audience": pgtype.FlatArray[string](idToken.Audience),
|
"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 "", 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.
|
// VerifyIDToken verifies that an *oauth2.Token is a valid *oidc.IDToken.
|
||||||
@@ -223,3 +218,142 @@ func (a *Authenticator) GetLogoutURL(requestHost string) *url.URL {
|
|||||||
}.Encode(),
|
}.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
|
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
@@ -1,9 +1,11 @@
|
|||||||
package site
|
package site
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -100,10 +102,31 @@ func (s *Server) authenticateAndAddIdentity(f response.HandlerFunc, assertions .
|
|||||||
return nil, response.Errorf("failed to authenticate: %w", err)
|
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
|
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)
|
user, acct, err := s.accts.GetUserAndAccountByAccessToken(ctx, accessToken)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, response.Errorf("failed to authorize: %w", err)
|
return nil, response.Errorf("failed to authorize: %w", err)
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"ruben/inventory2/internal/site/response"
|
"ruben/inventory2/internal/site/response"
|
||||||
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
// GET /login
|
// 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
|
// set access_token cookie and redirect to a reasonable place
|
||||||
|
|
||||||
return response.TemporaryRedirect("/").
|
return response.TemporaryRedirect("/").
|
||||||
Cookie(http.Cookie{
|
Cookie(newAccessTokenCookie(accessToken, expiration)), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAccessTokenCookie(tkn string, expiration time.Time) http.Cookie {
|
||||||
|
return http.Cookie{
|
||||||
Name: "access_token",
|
Name: "access_token",
|
||||||
Value: accessToken,
|
Value: tkn,
|
||||||
Path: "/",
|
Path: "/",
|
||||||
Expires: expiration,
|
Expires: expiration,
|
||||||
MaxAge: 0, // using Expiration instead
|
MaxAge: 0, // using Expiration instead
|
||||||
Secure: true,
|
Secure: true,
|
||||||
}), nil
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// GET /logout
|
// GET /logout
|
||||||
|
|||||||
Reference in New Issue
Block a user