diff --git a/core/api/jwt.go b/core/api/jwt.go index 21d64c5..48cb0eb 100644 --- a/core/api/jwt.go +++ b/core/api/jwt.go @@ -1,11 +1,13 @@ package api import ( + "errors" "net/http" "github.com/gaucho-racing/sentinel/core/config" "github.com/gaucho-racing/sentinel/core/service" "github.com/gin-gonic/gin" + "gorm.io/gorm" ) func JWKS(c *gin.Context) { @@ -67,6 +69,10 @@ func RevokeToken(c *gin.Context) { id := c.Param("id") if err := service.RevokeToken(id); err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + c.JSON(http.StatusNotFound, gin.H{"error": "token not found"}) + return + } c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return } diff --git a/core/service/jwt.go b/core/service/jwt.go index 2bfbe5a..b20aec2 100644 --- a/core/service/jwt.go +++ b/core/service/jwt.go @@ -215,5 +215,8 @@ func RevokeToken(id string) error { logger.SugarLogger.Errorf("Failed to revoke token: %v", result.Error) return result.Error } + if result.RowsAffected == 0 { + return gorm.ErrRecordNotFound + } return nil } diff --git a/oauth/api/login.go b/oauth/api/login.go index 19c39ca..138dc39 100644 --- a/oauth/api/login.go +++ b/oauth/api/login.go @@ -71,18 +71,19 @@ func RefreshSession(c *gin.Context) { return } - entityID, _ := claims["sub"].(string) - scope, _ := claims["scope"].(string) - if entityID == "" || !service.ScopesContain(scope, "refresh_token") { + refreshClaims, err := parseRefreshTokenClaims(claims, config.SentinelClientID) + if err != nil { c.JSON(http.StatusUnauthorized, gin.H{"error": "not a refresh token"}) return } - if tokenID, ok := claims["jti"].(string); ok { - sentinel.Delete("/api/core/token/"+tokenID, nil) + if err := sentinel.Delete("/api/core/token/"+refreshClaims.TokenID, nil); err != nil { + logger.SugarLogger.Errorf("Failed to revoke first-party refresh token %s: %v", refreshClaims.TokenID, err) + c.JSON(http.StatusBadGateway, gin.H{"error": "failed to rotate refresh token"}) + return } - resp, err := mintFirstPartySession(c, entityID) + resp, err := mintFirstPartySession(c, refreshClaims.EntityID) if err != nil { c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return diff --git a/oauth/api/refresh_claims.go b/oauth/api/refresh_claims.go new file mode 100644 index 0000000..5e6caa6 --- /dev/null +++ b/oauth/api/refresh_claims.go @@ -0,0 +1,48 @@ +package api + +import ( + "errors" + + "github.com/gaucho-racing/sentinel/oauth/service" +) + +var errInvalidRefreshTokenClaims = errors.New("invalid refresh token claims") + +type refreshTokenClaims struct { + EntityID string + Scope string + TokenID string +} + +func parseRefreshTokenClaims(claims map[string]interface{}, expectedAudience string) (refreshTokenClaims, error) { + entityID, _ := claims["sub"].(string) + scope, _ := claims["scope"].(string) + tokenID, _ := claims["jti"].(string) + if entityID == "" || tokenID == "" || !service.ScopesContain(scope, "refresh_token") { + return refreshTokenClaims{}, errInvalidRefreshTokenClaims + } + if !audienceMatches(claims["aud"], expectedAudience) { + return refreshTokenClaims{}, errInvalidRefreshTokenClaims + } + return refreshTokenClaims{EntityID: entityID, Scope: scope, TokenID: tokenID}, nil +} + +func audienceMatches(raw interface{}, expected string) bool { + if expected == "" { + return false + } + switch audience := raw.(type) { + case string: + return audience == expected + case []interface{}: + if len(audience) != 1 { + return false + } + value, ok := audience[0].(string) + return ok && value == expected + case []string: + return len(audience) == 1 && audience[0] == expected + default: + return false + } +}