mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-10-10 03:39:05 +02:00
refactor: authorize API routes with per-endpoint scopes (#1823)
Co-authored-by: copilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com>
This commit is contained in:
co-authored by
copilot-swe-agent[bot]
parent
1bd6f006c8
commit
0ec6bfa191
@@ -1,54 +0,0 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apikey"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/service"
|
||||
)
|
||||
|
||||
type ApiKeyAuthMiddleware struct {
|
||||
apiKeyModule *apikey.Module
|
||||
jwtService *service.JwtService
|
||||
}
|
||||
|
||||
func NewApiKeyAuthMiddleware(apiKeyModule *apikey.Module, jwtService *service.JwtService) *ApiKeyAuthMiddleware {
|
||||
return &ApiKeyAuthMiddleware{
|
||||
apiKeyModule: apiKeyModule,
|
||||
jwtService: jwtService,
|
||||
}
|
||||
}
|
||||
|
||||
func (m *ApiKeyAuthMiddleware) Add(adminRequired bool) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
userID, isAdmin, err := m.Verify(c, adminRequired)
|
||||
if err != nil {
|
||||
c.Abort()
|
||||
_ = c.Error(err)
|
||||
return
|
||||
}
|
||||
|
||||
c.Set("userID", userID)
|
||||
c.Set("userIsAdmin", isAdmin)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
func (m *ApiKeyAuthMiddleware) Verify(c *gin.Context, adminRequired bool) (userID string, isAdmin bool, err error) {
|
||||
apiKey := c.GetHeader("X-API-Key")
|
||||
|
||||
user, err := m.apiKeyModule.ValidateApiKey(c.Request.Context(), apiKey)
|
||||
if err != nil {
|
||||
return "", false, apperror.NotSignedIn()
|
||||
}
|
||||
|
||||
if user.Disabled {
|
||||
return "", false, apperror.UserDisabled()
|
||||
}
|
||||
|
||||
if adminRequired && !user.IsAdmin {
|
||||
return "", false, apperror.MissingPermission()
|
||||
}
|
||||
|
||||
return user.ID, user.IsAdmin, nil
|
||||
}
|
||||
@@ -1,132 +0,0 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apikey"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/service"
|
||||
)
|
||||
|
||||
// AuthMiddleware is a wrapper middleware that delegates to either API key or JWT authentication
|
||||
type AuthMiddleware struct {
|
||||
apiKeyMiddleware *ApiKeyAuthMiddleware
|
||||
jwtMiddleware *JwtAuthMiddleware
|
||||
options AuthOptions
|
||||
}
|
||||
|
||||
type AuthOptions struct {
|
||||
AdminRequired bool
|
||||
SuccessOptional bool
|
||||
AllowApiKeyAuth bool
|
||||
}
|
||||
|
||||
func NewAuthMiddleware(
|
||||
apiKeyModule *apikey.Module,
|
||||
userService *service.UserService,
|
||||
jwtService *service.JwtService,
|
||||
) *AuthMiddleware {
|
||||
return &AuthMiddleware{
|
||||
apiKeyMiddleware: NewApiKeyAuthMiddleware(apiKeyModule, jwtService),
|
||||
jwtMiddleware: NewJwtAuthMiddleware(jwtService, userService),
|
||||
options: AuthOptions{
|
||||
AdminRequired: true,
|
||||
SuccessOptional: false,
|
||||
AllowApiKeyAuth: true,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// WithAdminNotRequired allows the middleware to continue with the request even if the user is not an admin
|
||||
func (m *AuthMiddleware) WithAdminNotRequired() *AuthMiddleware {
|
||||
// Create a new instance to avoid modifying the original
|
||||
clone := &AuthMiddleware{
|
||||
apiKeyMiddleware: m.apiKeyMiddleware,
|
||||
jwtMiddleware: m.jwtMiddleware,
|
||||
options: m.options,
|
||||
}
|
||||
clone.options.AdminRequired = false
|
||||
return clone
|
||||
}
|
||||
|
||||
// WithSuccessOptional allows the middleware to continue with the request even if authentication fails
|
||||
func (m *AuthMiddleware) WithSuccessOptional() *AuthMiddleware {
|
||||
// Create a new instance to avoid modifying the original
|
||||
clone := &AuthMiddleware{
|
||||
apiKeyMiddleware: m.apiKeyMiddleware,
|
||||
jwtMiddleware: m.jwtMiddleware,
|
||||
options: m.options,
|
||||
}
|
||||
clone.options.SuccessOptional = true
|
||||
return clone
|
||||
}
|
||||
|
||||
// WithApiKeyAuthDisabled disables API key authentication fallback and requires JWT auth.
|
||||
func (m *AuthMiddleware) WithApiKeyAuthDisabled() *AuthMiddleware {
|
||||
clone := &AuthMiddleware{
|
||||
apiKeyMiddleware: m.apiKeyMiddleware,
|
||||
jwtMiddleware: m.jwtMiddleware,
|
||||
options: m.options,
|
||||
}
|
||||
clone.options.AllowApiKeyAuth = false
|
||||
return clone
|
||||
}
|
||||
|
||||
func (m *AuthMiddleware) Add() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
userID, isAdmin, authenticationMethod, authenticationTime, err := m.jwtMiddleware.Verify(c, m.options.AdminRequired)
|
||||
if err == nil {
|
||||
c.Set("userID", userID)
|
||||
c.Set("userIsAdmin", isAdmin)
|
||||
c.Set("authenticationMethod", authenticationMethod)
|
||||
c.Set("authenticationTime", authenticationTime)
|
||||
if c.IsAborted() {
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
// If JWT auth failed for a reason other than missing credentials, abort the request
|
||||
if !apperror.IsCode(err, apperror.CodeNotSignedIn) {
|
||||
c.Abort()
|
||||
_ = c.Error(err)
|
||||
return
|
||||
}
|
||||
|
||||
if !m.options.AllowApiKeyAuth {
|
||||
if m.options.SuccessOptional {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
c.Abort()
|
||||
if c.GetHeader("X-API-Key") != "" {
|
||||
_ = c.Error(apperror.APIKeyAuthNotAllowed())
|
||||
return
|
||||
}
|
||||
_ = c.Error(err)
|
||||
return
|
||||
}
|
||||
|
||||
// JWT auth failed, try API key auth
|
||||
userID, isAdmin, err = m.apiKeyMiddleware.Verify(c, m.options.AdminRequired)
|
||||
if err == nil {
|
||||
c.Set("userID", userID)
|
||||
c.Set("userIsAdmin", isAdmin)
|
||||
if c.IsAborted() {
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
if m.options.SuccessOptional {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
// Both JWT and API key auth failed
|
||||
c.Abort()
|
||||
_ = c.Error(err)
|
||||
}
|
||||
}
|
||||
@@ -1,108 +0,0 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apikey"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/common"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/instanceid"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/service"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils"
|
||||
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
|
||||
)
|
||||
|
||||
func TestWithApiKeyAuthDisabled(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
originalEnvConfig := common.EnvConfig
|
||||
defer func() {
|
||||
common.EnvConfig = originalEnvConfig
|
||||
}()
|
||||
common.EnvConfig.AppURL = "https://test.example.com"
|
||||
common.EnvConfig.EncryptionKey = []byte("0123456789abcdef0123456789abcdef")
|
||||
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
|
||||
instanceID, err := instanceid.Load(t.Context(), db)
|
||||
require.NoError(t, err)
|
||||
|
||||
jwtService, err := service.NewJwtService(t.Context(), db, instanceID)
|
||||
require.NoError(t, err)
|
||||
|
||||
userService := service.NewUserService(db, jwtService, nil, nil, nil, nil, nil)
|
||||
apiKeyModule, err := apikey.New(t.Context(), apikey.Dependencies{DB: db, CleanupDisabled: true})
|
||||
require.NoError(t, err)
|
||||
|
||||
authMiddleware := NewAuthMiddleware(apiKeyModule, userService, jwtService)
|
||||
|
||||
user := createUserForAuthMiddlewareTest(t, db)
|
||||
jwtToken, err := jwtService.GenerateAccessToken(user, "", time.Hour)
|
||||
require.NoError(t, err)
|
||||
|
||||
apiKeyToken := "middleware-test-api-key-raw-token"
|
||||
apiKeyRecord := apikey.ApiKey{
|
||||
Name: "Middleware API Key",
|
||||
Key: utils.CreateSha256Hash(apiKeyToken),
|
||||
UserID: user.ID,
|
||||
ExpiresAt: datatype.DateTime(time.Now().Add(24 * time.Hour)),
|
||||
}
|
||||
require.NoError(t, db.Create(&apiKeyRecord).Error)
|
||||
|
||||
router := gin.New()
|
||||
router.Use(NewErrorHandlerMiddleware().Add())
|
||||
router.GET("/api/protected", authMiddleware.WithAdminNotRequired().WithApiKeyAuthDisabled().Add(), func(c *gin.Context) {
|
||||
c.Status(http.StatusNoContent)
|
||||
})
|
||||
|
||||
t.Run("rejects API key auth when API key auth is disabled", func(t *testing.T) {
|
||||
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/api/protected", nil)
|
||||
req.Header.Set("X-API-Key", apiKeyToken)
|
||||
recorder := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(recorder, req)
|
||||
|
||||
require.Equal(t, http.StatusForbidden, recorder.Code)
|
||||
|
||||
var body map[string]string
|
||||
err := json.Unmarshal(recorder.Body.Bytes(), &body)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "API key authentication is not allowed for this endpoint", body["error"])
|
||||
})
|
||||
|
||||
t.Run("allows JWT auth when API key auth is disabled", func(t *testing.T) {
|
||||
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/api/protected", nil)
|
||||
req.Header.Set("Authorization", "Bearer "+jwtToken)
|
||||
recorder := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(recorder, req)
|
||||
|
||||
require.Equal(t, http.StatusNoContent, recorder.Code)
|
||||
})
|
||||
}
|
||||
|
||||
func createUserForAuthMiddlewareTest(t *testing.T, db *gorm.DB) model.User {
|
||||
t.Helper()
|
||||
|
||||
user := model.User{
|
||||
Username: "auth-user",
|
||||
Email: new("auth@example.com"),
|
||||
FirstName: "Auth",
|
||||
LastName: "User",
|
||||
DisplayName: "Auth User",
|
||||
}
|
||||
|
||||
err := db.Create(&user).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
return user
|
||||
}
|
||||
@@ -0,0 +1,122 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apikey"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/authz"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/service"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils/cookie"
|
||||
)
|
||||
|
||||
// #nosec G101 -- this is the name of the header that carries the API key, not a credential
|
||||
const apiKeyHeader = "X-API-Key"
|
||||
|
||||
// NewAuthorization creates the authorization middleware with every credential Pocket ID accepts
|
||||
// A browser session is tried before an API key, so a signed-in browser is never mistaken for an API client
|
||||
func NewAuthorization(apiKeyModule *apikey.Module, userService *service.UserService, jwtService *service.JwtService) *authz.Middleware {
|
||||
return authz.NewMiddleware(
|
||||
NewSessionAuthenticator(jwtService, userService),
|
||||
NewAPIKeyAuthenticator(apiKeyModule),
|
||||
)
|
||||
}
|
||||
|
||||
// SessionAuthenticator authenticates the session access token Pocket ID issues after sign-in
|
||||
type SessionAuthenticator struct {
|
||||
jwtService *service.JwtService
|
||||
userService *service.UserService
|
||||
}
|
||||
|
||||
func NewSessionAuthenticator(jwtService *service.JwtService, userService *service.UserService) *SessionAuthenticator {
|
||||
return &SessionAuthenticator{jwtService: jwtService, userService: userService}
|
||||
}
|
||||
|
||||
func (a *SessionAuthenticator) Kind() authz.PrincipalKind {
|
||||
return authz.KindSession
|
||||
}
|
||||
|
||||
func (a *SessionAuthenticator) Present(c *gin.Context) bool {
|
||||
return sessionToken(c) != ""
|
||||
}
|
||||
|
||||
func (a *SessionAuthenticator) Authenticate(c *gin.Context) (*authz.Principal, error) {
|
||||
// Verify the token signature, audience and type
|
||||
token, err := a.jwtService.VerifyAccessToken(sessionToken(c))
|
||||
if err != nil {
|
||||
return nil, apperror.NotSignedIn()
|
||||
}
|
||||
authenticationMethod, err := a.jwtService.GetAuthenticationMethod(token)
|
||||
if err != nil {
|
||||
return nil, apperror.NotSignedIn()
|
||||
}
|
||||
authenticationTime, _ := token.IssuedAt()
|
||||
|
||||
subject, ok := token.Subject()
|
||||
if !ok {
|
||||
return nil, apperror.TokenInvalid()
|
||||
}
|
||||
|
||||
// Load the user so disabling an account or changing its admin flag takes effect before the token expires
|
||||
user, err := a.userService.GetUser(c, subject)
|
||||
if err != nil {
|
||||
return nil, apperror.NotSignedIn()
|
||||
}
|
||||
if user.Disabled {
|
||||
return nil, apperror.UserDisabled()
|
||||
}
|
||||
|
||||
return &authz.Principal{
|
||||
Kind: authz.KindSession,
|
||||
UserID: user.ID,
|
||||
Scopes: authz.UserScopes(user.IsAdmin, authz.KindSession),
|
||||
AuthenticationMethod: authenticationMethod,
|
||||
AuthenticationTime: authenticationTime,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// sessionToken reads the session access token from its cookie, or from the Authorization header when the cookie is absent
|
||||
// An invalid cookie deliberately does not fall back to the header
|
||||
func sessionToken(c *gin.Context) string {
|
||||
if accessToken, err := c.Cookie(cookie.AccessTokenCookieName); err == nil {
|
||||
return accessToken
|
||||
}
|
||||
|
||||
_, accessToken, _ := strings.Cut(c.GetHeader("Authorization"), " ")
|
||||
return accessToken
|
||||
}
|
||||
|
||||
// APIKeyAuthenticator authenticates a personal API key sent in the X-API-Key header
|
||||
type APIKeyAuthenticator struct {
|
||||
apiKeyModule *apikey.Module
|
||||
}
|
||||
|
||||
func NewAPIKeyAuthenticator(apiKeyModule *apikey.Module) *APIKeyAuthenticator {
|
||||
return &APIKeyAuthenticator{apiKeyModule: apiKeyModule}
|
||||
}
|
||||
|
||||
func (a *APIKeyAuthenticator) Kind() authz.PrincipalKind {
|
||||
return authz.KindAPIKey
|
||||
}
|
||||
|
||||
func (a *APIKeyAuthenticator) Present(c *gin.Context) bool {
|
||||
return c.GetHeader(apiKeyHeader) != ""
|
||||
}
|
||||
|
||||
func (a *APIKeyAuthenticator) Authenticate(c *gin.Context) (*authz.Principal, error) {
|
||||
user, err := a.apiKeyModule.ValidateApiKey(c.Request.Context(), c.GetHeader(apiKeyHeader))
|
||||
if err != nil {
|
||||
return nil, apperror.NotSignedIn()
|
||||
}
|
||||
if user.Disabled {
|
||||
return nil, apperror.UserDisabled()
|
||||
}
|
||||
|
||||
return &authz.Principal{
|
||||
Kind: authz.KindAPIKey,
|
||||
UserID: user.ID,
|
||||
Scopes: authz.UserScopes(user.IsAdmin, authz.KindAPIKey),
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apikey"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/authz"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/common"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/instanceid"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/service"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils/cookie"
|
||||
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
|
||||
)
|
||||
|
||||
func TestAuthorizationWithRealCredentials(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
originalEnvConfig := common.EnvConfig
|
||||
defer func() {
|
||||
common.EnvConfig = originalEnvConfig
|
||||
}()
|
||||
common.EnvConfig.AppURL = "https://test.example.com"
|
||||
common.EnvConfig.EncryptionKey = []byte("0123456789abcdef0123456789abcdef")
|
||||
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
|
||||
instanceID, err := instanceid.Load(t.Context(), db)
|
||||
require.NoError(t, err)
|
||||
|
||||
jwtService, err := service.NewJwtService(t.Context(), db, instanceID)
|
||||
require.NoError(t, err)
|
||||
|
||||
userService := service.NewUserService(db, jwtService, nil, nil, nil, nil, nil)
|
||||
apiKeyModule, err := apikey.New(t.Context(), apikey.Dependencies{DB: db, CleanupDisabled: true})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create one credential of each kind for a regular user, an admin and a disabled user
|
||||
user := createUserForAuthorizationTest(t, db, "auth-user", false, false)
|
||||
admin := createUserForAuthorizationTest(t, db, "auth-admin", true, false)
|
||||
disabled := createUserForAuthorizationTest(t, db, "auth-disabled", true, true)
|
||||
|
||||
userSession, err := jwtService.GenerateAccessToken(user, "", time.Hour)
|
||||
require.NoError(t, err)
|
||||
userAPIKey := createAPIKeyForAuthorizationTest(t, db, user, "user-raw-api-key")
|
||||
adminAPIKey := createAPIKeyForAuthorizationTest(t, db, admin, "admin-raw-api-key")
|
||||
disabledAPIKey := createAPIKeyForAuthorizationTest(t, db, disabled, "disabled-raw-api-key")
|
||||
|
||||
// Mount one route per kind of requirement
|
||||
router := gin.New()
|
||||
router.Use(NewErrorHandlerMiddleware().Add())
|
||||
apiRouter := NewAuthorization(apiKeyModule, userService, jwtService).Router(router.Group("/api"))
|
||||
ok := func(c *gin.Context) {
|
||||
c.String(http.StatusOK, authz.PrincipalFrom(c).UserID)
|
||||
}
|
||||
apiRouter.GET("/session-only", authz.AccountSession, ok)
|
||||
apiRouter.GET("/account", authz.AccountRead, ok)
|
||||
apiRouter.GET("/admin", authz.UsersRead, ok)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
path string
|
||||
authorization string
|
||||
sessionCookie string
|
||||
apiKey string
|
||||
status int
|
||||
userID string
|
||||
errorCode string
|
||||
}{
|
||||
{name: "session on session-only route", path: "/api/session-only", authorization: "Bearer " + userSession, status: http.StatusOK, userID: user.ID},
|
||||
{name: "API key on session-only route", path: "/api/session-only", apiKey: adminAPIKey, status: http.StatusForbidden, errorCode: "api_key_auth_not_allowed"},
|
||||
{name: "API key on account route", path: "/api/account", apiKey: userAPIKey, status: http.StatusOK, userID: user.ID},
|
||||
{name: "regular user session on admin route", path: "/api/admin", authorization: "Bearer " + userSession, status: http.StatusForbidden, errorCode: "forbidden"},
|
||||
{name: "regular user API key on admin route", path: "/api/admin", apiKey: userAPIKey, status: http.StatusForbidden, errorCode: "forbidden"},
|
||||
{name: "admin API key on admin route", path: "/api/admin", apiKey: adminAPIKey, status: http.StatusOK, userID: admin.ID},
|
||||
{name: "invalid session cookie falls back to API key", path: "/api/admin", sessionCookie: "not-a-jwt", apiKey: adminAPIKey, status: http.StatusOK, userID: admin.ID},
|
||||
{name: "disabled user's API key", path: "/api/account", apiKey: disabledAPIKey, status: http.StatusForbidden, errorCode: "user_disabled"},
|
||||
{name: "unknown API key", path: "/api/account", apiKey: "unknown", status: http.StatusUnauthorized, errorCode: "not_signed_in"},
|
||||
{name: "no credentials", path: "/api/account", status: http.StatusUnauthorized, errorCode: "not_signed_in"},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, test.path, nil)
|
||||
if test.authorization != "" {
|
||||
req.Header.Set("Authorization", test.authorization)
|
||||
}
|
||||
if test.sessionCookie != "" {
|
||||
req.AddCookie(&http.Cookie{Name: cookie.AccessTokenCookieName, Value: test.sessionCookie, Secure: true, HttpOnly: true, SameSite: http.SameSiteLaxMode})
|
||||
}
|
||||
if test.apiKey != "" {
|
||||
req.Header.Set("X-API-Key", test.apiKey)
|
||||
}
|
||||
recorder := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(recorder, req)
|
||||
|
||||
require.Equal(t, test.status, recorder.Code, recorder.Body.String())
|
||||
if test.status == http.StatusOK {
|
||||
require.Equal(t, test.userID, recorder.Body.String())
|
||||
return
|
||||
}
|
||||
|
||||
var body map[string]any
|
||||
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &body))
|
||||
require.Equal(t, test.errorCode, body["code"])
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func createUserForAuthorizationTest(t *testing.T, db *gorm.DB, username string, isAdmin, disabled bool) model.User {
|
||||
t.Helper()
|
||||
|
||||
user := model.User{
|
||||
Username: username,
|
||||
Email: new(username + "@example.com"),
|
||||
FirstName: "Auth",
|
||||
LastName: "User",
|
||||
DisplayName: "Auth User",
|
||||
IsAdmin: isAdmin,
|
||||
Disabled: disabled,
|
||||
}
|
||||
|
||||
err := db.Create(&user).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
return user
|
||||
}
|
||||
|
||||
func createAPIKeyForAuthorizationTest(t *testing.T, db *gorm.DB, owner model.User, rawToken string) string {
|
||||
t.Helper()
|
||||
|
||||
err := db.Create(&apikey.ApiKey{
|
||||
Name: owner.Username + " key",
|
||||
Key: utils.CreateSha256Hash(rawToken),
|
||||
UserID: owner.ID,
|
||||
ExpiresAt: datatype.DateTime(time.Now().Add(24 * time.Hour)),
|
||||
}).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
return rawToken
|
||||
}
|
||||
@@ -1,81 +0,0 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/service"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils/cookie"
|
||||
)
|
||||
|
||||
type JwtAuthMiddleware struct {
|
||||
userService *service.UserService
|
||||
jwtService *service.JwtService
|
||||
}
|
||||
|
||||
func NewJwtAuthMiddleware(jwtService *service.JwtService, userService *service.UserService) *JwtAuthMiddleware {
|
||||
return &JwtAuthMiddleware{jwtService: jwtService, userService: userService}
|
||||
}
|
||||
|
||||
func (m *JwtAuthMiddleware) Add(adminRequired bool) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
userID, isAdmin, authenticationMethod, authenticationTime, err := m.Verify(c, adminRequired)
|
||||
if err != nil {
|
||||
c.Abort()
|
||||
_ = c.Error(err)
|
||||
return
|
||||
}
|
||||
|
||||
c.Set("userID", userID)
|
||||
c.Set("userIsAdmin", isAdmin)
|
||||
c.Set("authenticationMethod", authenticationMethod)
|
||||
c.Set("authenticationTime", authenticationTime)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
func (m *JwtAuthMiddleware) Verify(c *gin.Context, adminRequired bool) (subject string, isAdmin bool, authenticationMethod string, authenticationTime time.Time, err error) {
|
||||
// Extract the token from the cookie
|
||||
accessToken, err := c.Cookie(cookie.AccessTokenCookieName)
|
||||
if err != nil {
|
||||
// Try to extract the token from the Authorization header if it's not in the cookie
|
||||
var ok bool
|
||||
_, accessToken, ok = strings.Cut(c.GetHeader("Authorization"), " ")
|
||||
if !ok || accessToken == "" {
|
||||
return "", false, "", time.Time{}, apperror.NotSignedIn()
|
||||
}
|
||||
}
|
||||
|
||||
token, err := m.jwtService.VerifyAccessToken(accessToken)
|
||||
if err != nil {
|
||||
return "", false, "", time.Time{}, apperror.NotSignedIn()
|
||||
}
|
||||
authenticationMethod, err = m.jwtService.GetAuthenticationMethod(token)
|
||||
if err != nil {
|
||||
return "", false, "", time.Time{}, apperror.NotSignedIn()
|
||||
}
|
||||
authenticationTime, _ = token.IssuedAt()
|
||||
|
||||
subject, ok := token.Subject()
|
||||
if !ok {
|
||||
_ = c.Error(apperror.TokenInvalid())
|
||||
return "", false, "", time.Time{}, apperror.TokenInvalid()
|
||||
}
|
||||
|
||||
user, err := m.userService.GetUser(c, subject)
|
||||
if err != nil {
|
||||
return "", false, "", time.Time{}, apperror.NotSignedIn()
|
||||
}
|
||||
|
||||
if user.Disabled {
|
||||
return "", false, "", time.Time{}, apperror.UserDisabled()
|
||||
}
|
||||
|
||||
if adminRequired && !user.IsAdmin {
|
||||
return "", false, "", time.Time{}, apperror.MissingPermission()
|
||||
}
|
||||
|
||||
return subject, user.IsAdmin, authenticationMethod, authenticationTime, nil
|
||||
}
|
||||
Reference in New Issue
Block a user