diff --git a/README.md b/README.md index 84db2ba..cd7b40e 100644 --- a/README.md +++ b/README.md @@ -8,7 +8,7 @@ - [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 - [ ] 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 - [ ] make a FK between the etsy_store_events table and etsy_users table (store_id columns don't match types) - [ ] Auth0 diff --git a/internal/domains/authentication/auth.go b/internal/domains/authentication/auth.go index ffa3157..bf09109 100644 --- a/internal/domains/authentication/auth.go +++ b/internal/domains/authentication/auth.go @@ -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 +} diff --git a/internal/domains/authentication/store.go b/internal/domains/authentication/store.go index 20e7897..90cda2c 100644 --- a/internal/domains/authentication/store.go +++ b/internal/domains/authentication/store.go @@ -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 +} diff --git a/internal/site/auth.go b/internal/site/auth.go index 5744a24..3dca43a 100644 --- a/internal/site/auth.go +++ b/internal/site/auth.go @@ -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) diff --git a/internal/site/login.go b/internal/site/login.go index a29910d..1c2ba17 100644 --- a/internal/site/login.go +++ b/internal/site/login.go @@ -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