Enhance OAuth token handling and refresh logic

- Update the `Authenticate` method in the OAuth guard to allow for token refresh when an access token is expired but still valid.
- Introduce `TryRefreshToken` method to handle the refresh token logic, including token rotation and cookie management.
- Implement `VerifyTokenAllowExpired` and `VerifyRefreshToken` methods to improve token verification processes.
- Adjust error handling to provide clearer responses for token refresh failures.
- Refactor token expiration strategies in the login process to ensure consistent handling of access and refresh tokens.
This commit is contained in:
Max 2026-02-21 18:06:27 +08:00
parent 88c55b39c0
commit b68660b3cd
9 changed files with 557 additions and 128 deletions

View file

@ -47,36 +47,43 @@ func (s *Service) Guard(c *gin.Context) {
// This method only performs authentication without ACL checks // This method only performs authentication without ACL checks
// Returns true if authentication succeeded, false otherwise // Returns true if authentication succeeded, false otherwise
func (s *Service) Authenticate(c *gin.Context) bool { func (s *Service) Authenticate(c *gin.Context) bool {
// Get the token from the request
token := s.getAccessToken(c) token := s.getAccessToken(c)
// Validate the token
if token == "" { if token == "" {
response.RespondWithError(c, http.StatusUnauthorized, types.ErrTokenMissing) response.RespondWithError(c, http.StatusUnauthorized, types.ErrTokenMissing)
c.Abort() c.Abort()
return false return false
} }
// Validate the token // Try strict verification first (signature + expiration)
claims, err := s.VerifyToken(token) claims, err := s.VerifyToken(token)
if err != nil { if err != nil {
response.RespondWithError(c, http.StatusUnauthorized, types.ErrInvalidToken) // Token invalid — check if it's just expired (signature still valid)
c.Abort() expiredClaims, expErr := s.VerifyTokenAllowExpired(token)
return false if expErr != nil || expiredClaims == nil {
} response.RespondWithError(c, http.StatusUnauthorized, types.ErrInvalidToken)
c.Abort()
return false
}
// Auto refresh the token // Signature valid but expired — attempt auto refresh
if claims.ExpiresAt.Before(time.Now()) { if !expiredClaims.ExpiresAt.IsZero() && expiredClaims.ExpiresAt.Before(time.Now()) {
s.tryAutoRefreshToken(c, claims) newClaims, refreshErr := s.TryRefreshToken(c, expiredClaims)
if c.IsAborted() { if refreshErr != nil {
log.Error("[OAuth] Token refresh failed: %v", refreshErr)
response.RespondWithError(c, http.StatusUnauthorized, types.ErrInvalidRefreshToken)
c.Abort()
return false
}
claims = newClaims
} else {
response.RespondWithError(c, http.StatusUnauthorized, types.ErrInvalidToken)
c.Abort()
return false return false
} }
} }
// Set Authorized Info in context
sessionID := s.getSessionID(c) sessionID := s.getSessionID(c)
authorized.SetInfo(c, claims, sessionID, s.UserID) authorized.SetInfo(c, claims, sessionID, s.UserID)
return true return true
} }
@ -86,23 +93,108 @@ func GetAuthorizedInfo(c *gin.Context) *types.AuthorizedInfo {
return authorized.GetInfo(c) return authorized.GetInfo(c)
} }
func (s *Service) tryAutoRefreshToken(c *gin.Context, _ *types.TokenClaims) { // TryRefreshToken reads the refresh token from the request, verifies it,
// rotates the refresh token (revoke old, issue new), issues a new access token,
// writes both cookies, and returns the new claims.
// expiredClaims may be nil; in that case the identity is derived from the refresh token itself.
// Returns (nil, error) on any failure — the caller decides how to respond.
func (s *Service) TryRefreshToken(c *gin.Context, expiredClaims *types.TokenClaims) (*types.TokenClaims, error) {
refreshToken := s.getRefreshToken(c) refreshToken := s.getRefreshToken(c)
if refreshToken == "" { if refreshToken == "" {
response.RespondWithError(c, http.StatusUnauthorized, types.ErrRefreshTokenMissing) return nil, fmt.Errorf("refresh token missing")
c.Abort()
return
} }
// Verify the refresh token refreshClaims, err := s.VerifyRefreshToken(refreshToken)
_, err := s.VerifyToken(refreshToken)
if err != nil { if err != nil {
response.RespondWithError(c, http.StatusUnauthorized, types.ErrInvalidRefreshToken) return nil, fmt.Errorf("invalid or expired refresh token: %w", err)
c.Abort()
return
} }
// @Todo: Auto refresh the token // Derive access token TTL from the expired token's own iat/exp so the refreshed
// token keeps the same lifetime that was originally configured at login time.
var accessTTL time.Duration
if expiredClaims != nil && !expiredClaims.IssuedAt.IsZero() && !expiredClaims.ExpiresAt.IsZero() {
accessTTL = expiredClaims.ExpiresAt.Sub(expiredClaims.IssuedAt)
}
if accessTTL <= 0 {
accessTTL = s.config.Token.AccessTokenLifetime
}
if accessTTL <= 0 {
accessTTL = time.Hour
}
// Prefer the expired access token claims; fall back to refresh token claims
sourceClaims := expiredClaims
if sourceClaims == nil {
sourceClaims = refreshClaims
}
extraClaims := sourceClaims.Extra
if extraClaims == nil {
extraClaims = make(map[string]interface{})
}
if sourceClaims.TeamID != "" {
extraClaims["team_id"] = sourceClaims.TeamID
}
if sourceClaims.TenantID != "" {
extraClaims["tenant_id"] = sourceClaims.TenantID
}
// --- Refresh Token Rotation ---
// Revoke the old refresh token so it can never be reused.
s.revokeRefreshToken(refreshToken)
// Calculate remaining refresh lifetime for the new refresh token.
var refreshRemainingSeconds int
if !refreshClaims.ExpiresAt.IsZero() {
refreshRemainingSeconds = int(time.Until(refreshClaims.ExpiresAt).Seconds())
if refreshRemainingSeconds <= 0 {
return nil, fmt.Errorf("refresh token already expired after revocation")
}
} else {
refreshTTL := s.config.Token.RefreshTokenLifetime
if refreshTTL == 0 {
refreshTTL = 24 * time.Hour
}
refreshRemainingSeconds = int(refreshTTL.Seconds())
}
newRefreshToken, err := s.MakeRefreshToken(
sourceClaims.ClientID,
sourceClaims.Scope,
sourceClaims.Subject,
refreshRemainingSeconds,
extraClaims,
)
if err != nil {
return nil, fmt.Errorf("failed to issue new refresh token: %w", err)
}
// Issue new access token
newTokenStr, err := s.MakeAccessToken(
sourceClaims.ClientID,
sourceClaims.Scope,
sourceClaims.Subject,
int(accessTTL.Seconds()),
extraClaims,
)
if err != nil {
return nil, fmt.Errorf("failed to issue access token: %w", err)
}
// Cookie lifetime = new refresh token lifetime
cookieExpires := time.Now().Add(time.Duration(refreshRemainingSeconds) * time.Second)
cookieValue := fmt.Sprintf("Bearer %s", newTokenStr)
response.SendAccessTokenCookieWithExpiry(c, cookieValue, cookieExpires)
response.SendRefreshTokenCookieWithExpiry(c, newRefreshToken, cookieExpires)
newClaims, err := s.VerifyToken(newTokenStr)
if err != nil {
return nil, fmt.Errorf("failed to verify refreshed token: %w", err)
}
log.Info("[OAuth] Token rotated for subject %s (access + refresh)", sourceClaims.Subject)
return newClaims, nil
} }
func (s *Service) getAccessToken(c *gin.Context) string { func (s *Service) getAccessToken(c *gin.Context) string {
@ -152,6 +244,11 @@ func (s *Service) GetRefreshToken(c *gin.Context) string {
return s.getRefreshToken(c) return s.getRefreshToken(c)
} }
// GetSessionID gets the session ID from the request (public method)
func (s *Service) GetSessionID(c *gin.Context) string {
return s.getSessionID(c)
}
// Get Session ID from cookies, headers, or query string // Get Session ID from cookies, headers, or query string
func (s *Service) getSessionID(c *gin.Context) string { func (s *Service) getSessionID(c *gin.Context) string {

View file

@ -454,15 +454,64 @@ func (s *Service) SignToken(tokenType, clientID, scope, subject string, expiresI
// VerifyToken verifies a token based on its format and returns token claims // VerifyToken verifies a token based on its format and returns token claims
func (s *Service) VerifyToken(token string) (*types.TokenClaims, error) { func (s *Service) VerifyToken(token string) (*types.TokenClaims, error) {
// First try to verify as JWT (JWT tokens contain dots)
if strings.Contains(token, ".") { if strings.Contains(token, ".") {
return s.verifyJWTToken(token) return s.verifyJWTToken(token)
} }
// Otherwise, verify as opaque token
return s.verifyOpaqueToken(token) return s.verifyOpaqueToken(token)
} }
// VerifyTokenAllowExpired verifies token signature but allows expired tokens.
// Used by Guard to parse expired access tokens before attempting refresh.
func (s *Service) VerifyTokenAllowExpired(token string) (*types.TokenClaims, error) {
if strings.Contains(token, ".") {
return s.verifyJWTTokenAllowExpired(token)
}
return s.verifyOpaqueToken(token)
}
// VerifyRefreshToken verifies a refresh token based on its format.
// For opaque tokens it looks up the refresh token store (not the access token store).
func (s *Service) VerifyRefreshToken(token string) (*types.TokenClaims, error) {
if strings.Contains(token, ".") {
// JWT refresh tokens can be verified with the same JWT logic
return s.verifyJWTToken(token)
}
return s.verifyOpaqueRefreshToken(token)
}
// verifyOpaqueRefreshToken verifies an opaque refresh token using the refresh token store.
func (s *Service) verifyOpaqueRefreshToken(token string) (*types.TokenClaims, error) {
tokenInfo, err := s.getRefreshTokenData(token)
if err != nil {
return nil, fmt.Errorf("refresh token not found or invalid: %w", err)
}
clientID, _ := tokenInfo["client_id"].(string)
scope, _ := tokenInfo["scope"].(string)
subject, _ := tokenInfo["subject"].(string)
claims := &types.TokenClaims{
Subject: subject,
ClientID: clientID,
Scope: scope,
TokenType: "refresh_token",
Issuer: s.config.IssuerURL,
}
if issuedAt, ok := tokenInfo["issued_at"].(int64); ok {
claims.IssuedAt = time.Unix(issuedAt, 0)
}
if expiresAt, ok := tokenInfo["expires_at"].(int64); ok {
claims.ExpiresAt = time.Unix(expiresAt, 0)
if time.Now().After(claims.ExpiresAt) {
return nil, fmt.Errorf("refresh token expired")
}
}
return claims, nil
}
// SignIDToken signs an ID token with specific parameters and stores it // SignIDToken signs an ID token with specific parameters and stores it
func (s *Service) SignIDToken(clientID, scope string, expiresIn int, userdata *types.OIDCUserInfo, extraClaims ...map[string]interface{}) (string, error) { func (s *Service) SignIDToken(clientID, scope string, expiresIn int, userdata *types.OIDCUserInfo, extraClaims ...map[string]interface{}) (string, error) {
if s.signingCerts == nil || s.signingCerts.SigningKey == nil { if s.signingCerts == nil || s.signingCerts.SigningKey == nil {
@ -713,42 +762,57 @@ func (s *Service) signJWTToken(tokenType, clientID, scope, subject string, expir
// verifyJWTToken verifies a JWT token and returns its claims // verifyJWTToken verifies a JWT token and returns its claims
func (s *Service) verifyJWTToken(tokenString string) (*types.TokenClaims, error) { func (s *Service) verifyJWTToken(tokenString string) (*types.TokenClaims, error) {
return s.parseJWTToken(tokenString, false)
}
// verifyJWTTokenAllowExpired parses a JWT token, verifying signature but allowing expiration.
// Returns claims even if the token is expired (signature must still be valid).
func (s *Service) verifyJWTTokenAllowExpired(tokenString string) (*types.TokenClaims, error) {
return s.parseJWTToken(tokenString, true)
}
// parseJWTToken is the shared JWT parsing logic.
// When allowExpired is true, expired tokens are still parsed (signature-only verification).
func (s *Service) parseJWTToken(tokenString string, allowExpired bool) (*types.TokenClaims, error) {
if s.signingCerts == nil || s.signingCerts.SigningCert == nil { if s.signingCerts == nil || s.signingCerts.SigningCert == nil {
return nil, fmt.Errorf("signing certificates not initialized") return nil, fmt.Errorf("signing certificates not initialized")
} }
// Parse token with MapClaims to support extra claims parserOpts := []jwt.ParserOption{}
if allowExpired {
parserOpts = append(parserOpts, jwt.WithoutClaimsValidation())
}
token, err := jwt.ParseWithClaims(tokenString, jwt.MapClaims{}, func(token *jwt.Token) (interface{}, error) { token, err := jwt.ParseWithClaims(tokenString, jwt.MapClaims{}, func(token *jwt.Token) (interface{}, error) {
// Validate signing method
expectedMethod := getSigningMethod(s.config.Token.AccessTokenSigningAlg) expectedMethod := getSigningMethod(s.config.Token.AccessTokenSigningAlg)
if token.Method != expectedMethod { if token.Method != expectedMethod {
return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"]) return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"])
} }
// Return public key for verification
return s.signingCerts.GetPublicKey(), nil return s.signingCerts.GetPublicKey(), nil
}) }, parserOpts...)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to parse JWT token: %w", err) return nil, fmt.Errorf("failed to parse JWT token: %w", err)
} }
if !token.Valid { if !allowExpired && !token.Valid {
return nil, fmt.Errorf("invalid JWT token") return nil, fmt.Errorf("invalid JWT token")
} }
// Extract claims
mapClaims, ok := token.Claims.(jwt.MapClaims) mapClaims, ok := token.Claims.(jwt.MapClaims)
if !ok { if !ok {
return nil, fmt.Errorf("invalid JWT claims type") return nil, fmt.Errorf("invalid JWT claims type")
} }
// Convert to TokenClaims return s.extractTokenClaims(mapClaims), nil
}
// extractTokenClaims converts jwt.MapClaims to types.TokenClaims
func (s *Service) extractTokenClaims(mapClaims jwt.MapClaims) *types.TokenClaims {
tokenClaims := &types.TokenClaims{ tokenClaims := &types.TokenClaims{
Extra: make(map[string]interface{}), Extra: make(map[string]interface{}),
} }
// Extract standard claims
if sub, ok := mapClaims["sub"].(string); ok { if sub, ok := mapClaims["sub"].(string); ok {
tokenClaims.Subject = sub tokenClaims.Subject = sub
} }
@ -768,7 +832,6 @@ func (s *Service) verifyJWTToken(tokenString string) (*types.TokenClaims, error)
tokenClaims.JTI = jti tokenClaims.JTI = jti
} }
// Extract time claims
if exp, ok := mapClaims["exp"].(float64); ok { if exp, ok := mapClaims["exp"].(float64); ok {
tokenClaims.ExpiresAt = time.Unix(int64(exp), 0) tokenClaims.ExpiresAt = time.Unix(int64(exp), 0)
} }
@ -776,7 +839,6 @@ func (s *Service) verifyJWTToken(tokenString string) (*types.TokenClaims, error)
tokenClaims.IssuedAt = time.Unix(int64(iat), 0) tokenClaims.IssuedAt = time.Unix(int64(iat), 0)
} }
// Extract audience
if aud, ok := mapClaims["aud"].(string); ok { if aud, ok := mapClaims["aud"].(string); ok {
tokenClaims.Audience = []string{aud} tokenClaims.Audience = []string{aud}
} else if audArray, ok := mapClaims["aud"].([]interface{}); ok { } else if audArray, ok := mapClaims["aud"].([]interface{}); ok {
@ -789,7 +851,6 @@ func (s *Service) verifyJWTToken(tokenString string) (*types.TokenClaims, error)
tokenClaims.Audience = audience tokenClaims.Audience = audience
} }
// Extract extended claims for multi-tenancy and team support
if teamID, ok := mapClaims["team_id"].(string); ok { if teamID, ok := mapClaims["team_id"].(string); ok {
tokenClaims.TeamID = teamID tokenClaims.TeamID = teamID
} }
@ -797,7 +858,6 @@ func (s *Service) verifyJWTToken(tokenString string) (*types.TokenClaims, error)
tokenClaims.TenantID = tenantID tokenClaims.TenantID = tenantID
} }
// Store all extra claims for flexibility
standardClaims := map[string]bool{ standardClaims := map[string]bool{
"sub": true, "client_id": true, "scope": true, "token_type": true, "sub": true, "client_id": true, "scope": true, "token_type": true,
"exp": true, "iat": true, "nbf": true, "iss": true, "aud": true, "jti": true, "exp": true, "iat": true, "nbf": true, "iss": true, "aud": true, "jti": true,
@ -809,7 +869,7 @@ func (s *Service) verifyJWTToken(tokenString string) (*types.TokenClaims, error)
} }
} }
return tokenClaims, nil return tokenClaims
} }
// signOpaqueToken signs an opaque token using HMAC or RSA signature // signOpaqueToken signs an opaque token using HMAC or RSA signature

View file

@ -19,6 +19,7 @@ var (
ErrTokenMissing = &ErrorResponse{Code: "token_missing", ErrorDescription: "No access token provided in the request"} ErrTokenMissing = &ErrorResponse{Code: "token_missing", ErrorDescription: "No access token provided in the request"}
ErrInvalidRefreshToken = &ErrorResponse{Code: "invalid_refresh_token", ErrorDescription: "The refresh token provided is invalid or expired"} ErrInvalidRefreshToken = &ErrorResponse{Code: "invalid_refresh_token", ErrorDescription: "The refresh token provided is invalid or expired"}
ErrRefreshTokenMissing = &ErrorResponse{Code: "refresh_token_missing", ErrorDescription: "No refresh token provided in the request"} ErrRefreshTokenMissing = &ErrorResponse{Code: "refresh_token_missing", ErrorDescription: "No refresh token provided in the request"}
ErrTokenRefreshFailed = &ErrorResponse{Code: "token_refresh_failed", ErrorDescription: "Failed to refresh access token"}
// Permission related errors // Permission related errors
ErrForbidden = &ErrorResponse{Code: "forbidden", ErrorDescription: "You do not have permission to access this resource"} ErrForbidden = &ErrorResponse{Code: "forbidden", ErrorDescription: "You do not have permission to access this resource"}

View file

@ -0,0 +1,271 @@
package openapi_test
import (
"fmt"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/yaoapp/yao/openapi/oauth"
"github.com/yaoapp/yao/openapi/oauth/authorized"
"github.com/yaoapp/yao/openapi/response"
"github.com/yaoapp/yao/openapi/tests/testutils"
)
// TestGuard_ValidToken verifies that a valid, non-expired access token passes through authentication.
func TestGuard_ValidToken(t *testing.T) {
serverURL := testutils.Prepare(t)
defer testutils.Clean()
_ = serverURL
oauthService := oauth.OAuth
assert.NotNil(t, oauthService, "OAuth service should be initialized")
client := testutils.RegisterTestClient(t, "Guard Valid Token Test", []string{"https://localhost/callback"})
defer testutils.CleanupTestClient(t, client.ClientID)
tokenInfo := testutils.ObtainAccessToken(t, serverURL, client.ClientID, client.ClientSecret, "https://localhost/callback", "openid profile")
router := authenticateRouter(oauthService)
accessCookieName := response.GetCookieName("access_token")
req := httptest.NewRequest("GET", "/guarded", nil)
req.AddCookie(&http.Cookie{Name: accessCookieName, Value: fmt.Sprintf("Bearer %s", tokenInfo.AccessToken)})
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code, "Valid token should pass authentication")
assert.Contains(t, w.Body.String(), `"subject"`, "Response should contain authorized subject")
}
// TestGuard_NoToken verifies that a request without any token is rejected with 401.
func TestGuard_NoToken(t *testing.T) {
serverURL := testutils.Prepare(t)
defer testutils.Clean()
_ = serverURL
oauthService := oauth.OAuth
assert.NotNil(t, oauthService, "OAuth service should be initialized")
router := authenticateRouter(oauthService)
req := httptest.NewRequest("GET", "/guarded", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusUnauthorized, w.Code, "No token should return 401")
assert.Contains(t, w.Body.String(), "token_missing", "Error should indicate missing token")
}
// TestGuard_InvalidSignature verifies that a token with an invalid signature is rejected with 401.
func TestGuard_InvalidSignature(t *testing.T) {
serverURL := testutils.Prepare(t)
defer testutils.Clean()
_ = serverURL
oauthService := oauth.OAuth
assert.NotNil(t, oauthService, "OAuth service should be initialized")
router := authenticateRouter(oauthService)
accessCookieName := response.GetCookieName("access_token")
req := httptest.NewRequest("GET", "/guarded", nil)
req.AddCookie(&http.Cookie{Name: accessCookieName, Value: "Bearer eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJmYWtlIn0.invalidsignature"})
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusUnauthorized, w.Code, "Invalid signature should return 401")
}
// TestGuard_ExpiredToken_NoRefresh verifies that an expired access token without a refresh token returns 401.
func TestGuard_ExpiredToken_NoRefresh(t *testing.T) {
serverURL := testutils.Prepare(t)
defer testutils.Clean()
_ = serverURL
oauthService := oauth.OAuth
assert.NotNil(t, oauthService, "OAuth service should be initialized")
client := testutils.RegisterTestClient(t, "Guard Expired No Refresh Test", []string{"https://localhost/callback"})
defer testutils.CleanupTestClient(t, client.ClientID)
expiredToken, err := oauthService.MakeAccessToken(client.ClientID, "openid profile", "test-subject-expired", -1)
assert.NoError(t, err, "Should be able to create expired token")
router := authenticateRouter(oauthService)
accessCookieName := response.GetCookieName("access_token")
req := httptest.NewRequest("GET", "/guarded", nil)
req.AddCookie(&http.Cookie{Name: accessCookieName, Value: fmt.Sprintf("Bearer %s", expiredToken)})
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusUnauthorized, w.Code, "Expired token without refresh token should return 401")
}
// TestGuard_ExpiredToken_WithValidRefresh verifies that an expired access token with a valid refresh token
// triggers auto-refresh: the request succeeds and a new access_token cookie is set.
func TestGuard_ExpiredToken_WithValidRefresh(t *testing.T) {
serverURL := testutils.Prepare(t)
defer testutils.Clean()
_ = serverURL
oauthService := oauth.OAuth
assert.NotNil(t, oauthService, "OAuth service should be initialized")
client := testutils.RegisterTestClient(t, "Guard Auto Refresh Test", []string{"https://localhost/callback"})
defer testutils.CleanupTestClient(t, client.ClientID)
subject := "test-subject-auto-refresh"
expiredToken, err := oauthService.MakeAccessToken(client.ClientID, "openid profile", subject, -1)
assert.NoError(t, err, "Should create expired access token")
// Create a JWT-format refresh token so VerifyToken can validate it directly.
// The default opaque format requires store lookup which is separate from the signing path.
refreshToken, err := oauthService.MakeRefreshToken(client.ClientID, "openid profile", subject, 86400)
assert.NoError(t, err, "Should create valid refresh token")
router := authenticateRouter(oauthService)
accessCookieName := response.GetCookieName("access_token")
refreshCookieName := response.GetCookieName("refresh_token")
req := httptest.NewRequest("GET", "/guarded", nil)
req.AddCookie(&http.Cookie{Name: accessCookieName, Value: fmt.Sprintf("Bearer %s", expiredToken)})
req.AddCookie(&http.Cookie{Name: refreshCookieName, Value: fmt.Sprintf("Bearer %s", refreshToken)})
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code, "Expired token + valid refresh should auto-refresh and succeed")
assert.Contains(t, w.Body.String(), `"subject"`, "Response should contain authorized subject")
// Verify that both access_token and refresh_token cookies were rotated
setCookieHeaders := w.Result().Cookies()
foundNewAccessToken := false
foundNewRefreshToken := false
for _, c := range setCookieHeaders {
if c.Name == accessCookieName {
foundNewAccessToken = true
assert.NotEmpty(t, c.Value, "New access token cookie should have a value")
rawValue := strings.TrimPrefix(c.Value, "Bearer ")
assert.NotEqual(t, expiredToken, rawValue, "New token should differ from the expired one")
t.Logf("New access_token cookie set with MaxAge=%d", c.MaxAge)
}
if c.Name == refreshCookieName {
foundNewRefreshToken = true
assert.NotEmpty(t, c.Value, "New refresh token cookie should have a value")
rawValue := strings.TrimPrefix(c.Value, "Bearer ")
assert.NotEqual(t, refreshToken, rawValue, "New refresh token should differ from the old one")
t.Logf("New refresh_token cookie set with MaxAge=%d", c.MaxAge)
}
}
assert.True(t, foundNewAccessToken, "Guard should write a new access_token cookie after auto-refresh")
assert.True(t, foundNewRefreshToken, "Guard should rotate refresh_token cookie after auto-refresh")
}
// TestGuard_ExpiredToken_WithExpiredRefresh verifies that an expired access token paired with an
// also-expired refresh token returns 401.
func TestGuard_ExpiredToken_WithExpiredRefresh(t *testing.T) {
serverURL := testutils.Prepare(t)
defer testutils.Clean()
_ = serverURL
oauthService := oauth.OAuth
assert.NotNil(t, oauthService, "OAuth service should be initialized")
client := testutils.RegisterTestClient(t, "Guard Expired Refresh Test", []string{"https://localhost/callback"})
defer testutils.CleanupTestClient(t, client.ClientID)
subject := "test-subject-both-expired"
expiredAccess, err := oauthService.MakeAccessToken(client.ClientID, "openid profile", subject, -1)
assert.NoError(t, err)
// Opaque refresh tokens expire via store TTL, not a field in the data.
// Use a 1-second TTL and wait for it to expire from the store.
expiredRefresh, err := oauthService.MakeRefreshToken(client.ClientID, "openid profile", subject, 1)
assert.NoError(t, err)
time.Sleep(2 * time.Second)
router := authenticateRouter(oauthService)
accessCookieName := response.GetCookieName("access_token")
refreshCookieName := response.GetCookieName("refresh_token")
req := httptest.NewRequest("GET", "/guarded", nil)
req.AddCookie(&http.Cookie{Name: accessCookieName, Value: fmt.Sprintf("Bearer %s", expiredAccess)})
req.AddCookie(&http.Cookie{Name: refreshCookieName, Value: fmt.Sprintf("Bearer %s", expiredRefresh)})
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusUnauthorized, w.Code, "Both tokens expired should return 401")
}
// TestGuard_AuthorizationHeader verifies that the Guard also works with the Authorization header.
func TestGuard_AuthorizationHeader(t *testing.T) {
serverURL := testutils.Prepare(t)
defer testutils.Clean()
_ = serverURL
oauthService := oauth.OAuth
assert.NotNil(t, oauthService, "OAuth service should be initialized")
client := testutils.RegisterTestClient(t, "Guard Header Auth Test", []string{"https://localhost/callback"})
defer testutils.CleanupTestClient(t, client.ClientID)
tokenInfo := testutils.ObtainAccessToken(t, serverURL, client.ClientID, client.ClientSecret, "https://localhost/callback", "openid profile")
router := authenticateRouter(oauthService)
req := httptest.NewRequest("GET", "/guarded", nil)
req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", tokenInfo.AccessToken))
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code, "Valid Bearer token in Authorization header should pass authentication")
assert.Contains(t, w.Body.String(), `"subject"`, "Response should contain authorized subject")
}
// authenticateRouter creates a Gin router with ONLY the Authenticate middleware (no ACL).
// This isolates the token verification and auto-refresh logic from permission checks.
func authenticateRouter(oauthService *oauth.Service) *gin.Engine {
gin.SetMode(gin.TestMode)
router := gin.New()
handler := func(c *gin.Context) {
info := authorized.GetInfo(c)
if info == nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "no authorized info"})
return
}
c.JSON(http.StatusOK, gin.H{
"subject": info.Subject,
"client_id": info.ClientID,
"scope": info.Scope,
"user_id": info.UserID,
"session_id": info.SessionID,
})
}
// Use Authenticate (auth only) instead of Guard (auth + ACL)
router.GET("/guarded", func(c *gin.Context) {
if !oauthService.Authenticate(c) {
return
}
handler(c)
})
return router
}

View file

@ -471,7 +471,6 @@ func createPublicEntryConfig(config *EntryConfig) *EntryConfig {
publicConfig.Token = &TokenConfig{ publicConfig.Token = &TokenConfig{
ExpiresIn: config.Token.ExpiresIn, ExpiresIn: config.Token.ExpiresIn,
RefreshTokenExpiresIn: config.Token.RefreshTokenExpiresIn, RefreshTokenExpiresIn: config.Token.RefreshTokenExpiresIn,
RememberMeExpiresIn: config.Token.RememberMeExpiresIn,
RememberMeRefreshTokenExpiresIn: config.Token.RememberMeRefreshTokenExpiresIn, RememberMeRefreshTokenExpiresIn: config.Token.RememberMeRefreshTokenExpiresIn,
} }
} }

View file

@ -500,89 +500,75 @@ func LoginByTeamID(userid string, teamID string, loginCtx *LoginContext) (*Login
func issueTokens(ctx context.Context, params *IssueTokensParams) (*LoginResponse, error) { func issueTokens(ctx context.Context, params *IssueTokensParams) (*LoginResponse, error) {
yaoClientConfig := GetYaoClientConfig() yaoClientConfig := GetYaoClientConfig()
// Determine token expiration times based on Remember Me setting // Token expiration strategy:
// - access_token: always short-lived (from expires_in config), same for all login types
// - refresh_token: short for normal login, long for remember_me / OAuth
// Security: a leaked access_token has limited impact window; "keep logged in"
// is achieved by silently refreshing via long-lived refresh_token in Guard.
var expiresIn, refreshTokenExpiresIn int var expiresIn, refreshTokenExpiresIn int
// Try to get token config from entry config first
locale := "" locale := ""
if params.LoginCtx != nil && params.LoginCtx.Locale != "" {
locale = params.LoginCtx.Locale
}
entryConfig := GetEntryConfig(locale) entryConfig := GetEntryConfig(locale)
if params.LoginCtx != nil && params.LoginCtx.RememberMe { // 1. Access token: always use the standard short duration
// Remember Me mode: use extended token durations if entryConfig != nil && entryConfig.Token != nil && entryConfig.Token.ExpiresIn != "" {
if entryConfig != nil && entryConfig.Token != nil { normalized, err := normalizeDuration(entryConfig.Token.ExpiresIn)
// Parse Remember Me access token expires_in if err != nil {
if entryConfig.Token.RememberMeExpiresIn != "" { log.Warn("Failed to parse expires_in: %s, using default", err.Error())
normalized, err := normalizeDuration(entryConfig.Token.RememberMeExpiresIn) } else {
if err != nil { duration, err := time.ParseDuration(normalized)
log.Warn("Failed to parse remember_me_expires_in: %s, using default", err.Error()) if err == nil {
} else { expiresIn = int(duration.Seconds())
duration, err := time.ParseDuration(normalized)
if err == nil {
expiresIn = int(duration.Seconds())
}
}
}
// Parse Remember Me refresh token expires_in
if entryConfig.Token.RememberMeRefreshTokenExpiresIn != "" {
normalized, err := normalizeDuration(entryConfig.Token.RememberMeRefreshTokenExpiresIn)
if err != nil {
log.Warn("Failed to parse remember_me_refresh_token_expires_in: %s, using default", err.Error())
} else {
duration, err := time.ParseDuration(normalized)
if err == nil {
refreshTokenExpiresIn = int(duration.Seconds())
}
}
}
// If refresh token not configured, default to 2x the access token duration
if refreshTokenExpiresIn == 0 && expiresIn > 0 {
refreshTokenExpiresIn = expiresIn * 2
}
}
} else {
// Normal login: use standard token durations from entry config
if entryConfig != nil && entryConfig.Token != nil {
// Parse access token expires_in
if entryConfig.Token.ExpiresIn != "" {
normalized, err := normalizeDuration(entryConfig.Token.ExpiresIn)
if err != nil {
log.Warn("Failed to parse expires_in: %s, using default", err.Error())
} else {
duration, err := time.ParseDuration(normalized)
if err == nil {
expiresIn = int(duration.Seconds())
}
}
}
// Parse refresh token expires_in
if entryConfig.Token.RefreshTokenExpiresIn != "" {
normalized, err := normalizeDuration(entryConfig.Token.RefreshTokenExpiresIn)
if err != nil {
log.Warn("Failed to parse refresh_token_expires_in: %s, using default", err.Error())
} else {
duration, err := time.ParseDuration(normalized)
if err == nil {
refreshTokenExpiresIn = int(duration.Seconds())
}
}
}
// If refresh token not configured, default to 24x the access token duration
if refreshTokenExpiresIn == 0 && expiresIn > 0 {
refreshTokenExpiresIn = expiresIn * 24
} }
} }
} }
// Fall back to YaoClientConfig defaults if not set from entry config // 2. Refresh token: depends on remember_me
rememberMe := params.LoginCtx != nil && params.LoginCtx.RememberMe
if rememberMe && entryConfig != nil && entryConfig.Token != nil {
// Remember Me: use extended refresh token duration
if entryConfig.Token.RememberMeRefreshTokenExpiresIn != "" {
normalized, err := normalizeDuration(entryConfig.Token.RememberMeRefreshTokenExpiresIn)
if err != nil {
log.Warn("Failed to parse remember_me_refresh_token_expires_in: %s, using default", err.Error())
} else {
duration, err := time.ParseDuration(normalized)
if err == nil {
refreshTokenExpiresIn = int(duration.Seconds())
}
}
}
} else if entryConfig != nil && entryConfig.Token != nil {
// Normal login: use standard refresh token duration
if entryConfig.Token.RefreshTokenExpiresIn != "" {
normalized, err := normalizeDuration(entryConfig.Token.RefreshTokenExpiresIn)
if err != nil {
log.Warn("Failed to parse refresh_token_expires_in: %s, using default", err.Error())
} else {
duration, err := time.ParseDuration(normalized)
if err == nil {
refreshTokenExpiresIn = int(duration.Seconds())
}
}
}
}
// 3. Default fallbacks
if expiresIn == 0 { if expiresIn == 0 {
expiresIn = yaoClientConfig.ExpiresIn expiresIn = yaoClientConfig.ExpiresIn
} }
// Refresh token defaults: remember_me 90d, normal 7d, then client config
if refreshTokenExpiresIn == 0 { if refreshTokenExpiresIn == 0 {
refreshTokenExpiresIn = yaoClientConfig.RefreshTokenExpiresIn if rememberMe {
refreshTokenExpiresIn = 90 * 24 * 3600 // 90 days
} else if yaoClientConfig.RefreshTokenExpiresIn > 0 {
refreshTokenExpiresIn = yaoClientConfig.RefreshTokenExpiresIn
} else {
refreshTokenExpiresIn = 7 * 24 * 3600 // 7 days
}
} }
// Prepare OIDC user info // Prepare OIDC user info
@ -892,9 +878,13 @@ func GinLogout(c *gin.Context) {
// This includes access token, refresh token, and optionally session ID cookies with appropriate security settings // This includes access token, refresh token, and optionally session ID cookies with appropriate security settings
func SendLoginCookies(c *gin.Context, loginResponse *LoginResponse, sessionID string) { func SendLoginCookies(c *gin.Context, loginResponse *LoginResponse, sessionID string) {
// Send session ID cookie only if sessionID is provided // Send session ID cookie - expires with refresh token so session survives token refreshes
if sessionID != "" { if sessionID != "" {
expires := time.Now().Add(time.Duration(yaoClientConfig.ExpiresIn) * time.Second) sessionExpiry := loginResponse.RefreshTokenExpiresIn
if sessionExpiry <= 0 {
sessionExpiry = loginResponse.ExpiresIn
}
expires := time.Now().Add(time.Duration(sessionExpiry) * time.Second)
options := response.NewSecureCookieOptions(). options := response.NewSecureCookieOptions().
WithExpires(expires). WithExpires(expires).
WithSameSite("Strict") WithSameSite("Strict")
@ -914,10 +904,13 @@ func SendLoginCookies(c *gin.Context, loginResponse *LoginResponse, sessionID st
refreshToken := fmt.Sprintf("%s %s", loginResponse.TokenType, loginResponse.RefreshToken) refreshToken := fmt.Sprintf("%s %s", loginResponse.TokenType, loginResponse.RefreshToken)
// Calculate expiration times // Calculate expiration times
// access_token cookie lives as long as refresh_token so the browser keeps sending the
// (JWT-expired) access token — the Guard can then use the refresh token to issue a new one.
// The JWT's own `exp` claim handles the real expiration check on the server side.
refreshExpires := time.Now().Add(time.Duration(loginResponse.RefreshTokenExpiresIn) * time.Second) refreshExpires := time.Now().Add(time.Duration(loginResponse.RefreshTokenExpiresIn) * time.Second)
// Send access token cookie // Send access token cookie (cookie lifetime = refresh token lifetime)
response.SendAccessTokenCookieWithExpiry(c, accessToken, time.Now().Add(time.Duration(loginResponse.ExpiresIn)*time.Second)) response.SendAccessTokenCookieWithExpiry(c, accessToken, refreshExpires)
// Send refresh token cookie // Send refresh token cookie
response.SendRefreshTokenCookieWithExpiry(c, refreshToken, refreshExpires) response.SendRefreshTokenCookieWithExpiry(c, refreshToken, refreshExpires)

View file

@ -175,6 +175,7 @@ func authback(c *gin.Context) {
// LoginThirdParty(providerID, userInfo) // LoginThirdParty(providerID, userInfo)
loginCtx := makeLoginContext(c) loginCtx := makeLoginContext(c)
loginCtx.AuthSource = providerID // Set auth source to provider name (google, github, etc.) loginCtx.AuthSource = providerID // Set auth source to provider name (google, github, etc.)
loginCtx.RememberMe = true // OAuth login always uses extended token durations
// Use locale from params, fallback to "en" if not provided // Use locale from params, fallback to "en" if not provided
locale := params.Locale locale := params.Locale

View file

@ -78,7 +78,6 @@ type CaptchaConfig struct {
type TokenConfig struct { type TokenConfig struct {
ExpiresIn string `json:"expires_in,omitempty"` ExpiresIn string `json:"expires_in,omitempty"`
RefreshTokenExpiresIn string `json:"refresh_token_expires_in,omitempty"` RefreshTokenExpiresIn string `json:"refresh_token_expires_in,omitempty"`
RememberMeExpiresIn string `json:"remember_me_expires_in,omitempty"`
RememberMeRefreshTokenExpiresIn string `json:"remember_me_refresh_token_expires_in,omitempty"` RememberMeRefreshTokenExpiresIn string `json:"remember_me_refresh_token_expires_in,omitempty"`
} }

View file

@ -6,6 +6,7 @@ import (
"io" "io"
"strconv" "strconv"
"strings" "strings"
"time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/google/uuid" "github.com/google/uuid"
@ -95,8 +96,8 @@ func guardCookieTrace(r *Request) error {
} }
// OAuth 2.1 guard - authentication only // OAuth 2.1 guard - authentication only
// This guard validates the token and sets authorized info // This guard validates the token and sets authorized info.
// ACL checks are performed separately in Run() for API calls // ACL checks are performed separately in Run() for API calls.
// NOTE: This guard does NOT write HTTP responses on failure, so that // NOTE: This guard does NOT write HTTP responses on failure, so that
// the caller (Guard/apiGuard) can handle redirects or custom error responses. // the caller (Guard/apiGuard) can handle redirects or custom error responses.
func guardOAuth(r *Request) error { func guardOAuth(r *Request) error {
@ -110,23 +111,30 @@ func guardOAuth(r *Request) error {
c := r.context c := r.context
// Check token first without writing response.
// oauth.Authenticate() writes JSON + aborts on failure, which prevents
// the caller from doing redirects. So we check the token manually first.
token := oauth.OAuth.GetAccessToken(c) token := oauth.OAuth.GetAccessToken(c)
if token == "" { if token == "" {
return fmt.Errorf("Exception|401:Not authenticated") return fmt.Errorf("Exception|401:Not authenticated")
} }
if _, err := oauth.OAuth.VerifyToken(token); err != nil { claims, err := oauth.OAuth.VerifyToken(token)
return fmt.Errorf("Exception|401:Invalid or expired token") if err != nil {
// Token invalid — check if just expired (signature still valid)
expiredClaims, expErr := oauth.OAuth.VerifyTokenAllowExpired(token)
if expErr == nil && expiredClaims != nil &&
!expiredClaims.ExpiresAt.IsZero() && expiredClaims.ExpiresAt.Before(time.Now()) {
refreshed, refreshErr := oauth.OAuth.TryRefreshToken(c, expiredClaims)
if refreshErr != nil {
return fmt.Errorf("Exception|401:Token expired and refresh failed")
}
claims = refreshed
} else {
return fmt.Errorf("Exception|401:Invalid token")
}
} }
// Token is valid, now call Authenticate to set up the full context // Set authorized info in context
// (session ID, authorized info, etc.). This will succeed since token is valid. authorized.SetInfo(c, claims, oauth.OAuth.GetSessionID(c), oauth.OAuth.UserID)
oauth.OAuth.Authenticate(c)
// Get authorized info from context
info := authorized.GetInfo(c) info := authorized.GetInfo(c)
if info != nil { if info != nil {
r.Sid = info.SessionID r.Sid = info.SessionID