diff --git a/backend/internal/api/module.go b/backend/internal/api/module.go index 4f727ee6..5dcc6153 100644 --- a/backend/internal/api/module.go +++ b/backend/internal/api/module.go @@ -3,9 +3,9 @@ package api import ( "context" - "github.com/gin-gonic/gin" "gorm.io/gorm" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/httpserver" "github.com/pocket-id/pocket-id/backend/internal/oidc" @@ -60,26 +60,23 @@ func (m *Module) DescribePermissions(ctx context.Context, audience string, keys } // RegisterRoutes mounts the admin CRUD endpoints -// adminAuth is passed in as a gin handler so the module does not import internal/middleware -func (m *Module) RegisterRoutes(apiGroup *gin.RouterGroup, adminAuth gin.HandlerFunc) { - apis := apiGroup.Group("/apis") - apis.Use(adminAuth) - apis.GET("", httpserver.Handle(m.handler.list)) - apis.POST("", httpserver.Handle(m.handler.create)) - apis.GET("/:id", httpserver.Handle(m.handler.get)) - apis.PUT("/:id", httpserver.Handle(m.handler.update)) - apis.DELETE("/:id", httpserver.Handle(m.handler.delete)) - apis.PUT("/:id/permissions", httpserver.Handle(m.handler.updatePermissions)) - apis.PUT("/:id/cimd-access", httpserver.Handle(m.handler.updateCimdAccess)) +func (m *Module) RegisterRoutes(r *authz.Router) { + apis := r.Group("/apis") + apis.GET("", authz.APIsRead, httpserver.Handle(m.handler.list)) + apis.POST("", authz.APIsWrite, httpserver.Handle(m.handler.create)) + apis.GET("/:id", authz.APIsRead, httpserver.Handle(m.handler.get)) + apis.PUT("/:id", authz.APIsWrite, httpserver.Handle(m.handler.update)) + apis.DELETE("/:id", authz.APIsWrite, httpserver.Handle(m.handler.delete)) + apis.PUT("/:id/permissions", authz.APIsWrite, httpserver.Handle(m.handler.updatePermissions)) + apis.PUT("/:id/cimd-access", authz.APIsWrite, httpserver.Handle(m.handler.updateCimdAccess)) // The same client grants are editable from either side of the relation, so the API can list and manage its clients too - apis.GET("/:id/clients", httpserver.Handle(m.handler.listClients)) - apis.GET("/:id/assignable-clients", httpserver.Handle(m.handler.listAssignableClients)) - apis.PUT("/:id/clients/:clientId", httpserver.Handle(m.handler.updateClientAccessForApi)) - apis.DELETE("/:id/clients/:clientId", httpserver.Handle(m.handler.removeClientAccessForApi)) + apis.GET("/:id/clients", authz.APIsRead, httpserver.Handle(m.handler.listClients)) + apis.GET("/:id/assignable-clients", authz.APIsRead, httpserver.Handle(m.handler.listAssignableClients)) + apis.PUT("/:id/clients/:clientId", authz.APIsWrite, httpserver.Handle(m.handler.updateClientAccessForApi)) + apis.DELETE("/:id/clients/:clientId", authz.APIsWrite, httpserver.Handle(m.handler.removeClientAccessForApi)) - access := apiGroup.Group("/api-access") - access.Use(adminAuth) - access.GET("/:clientId/apis", httpserver.Handle(m.handler.listClientApis)) - access.GET("/:clientId/assignable-apis", httpserver.Handle(m.handler.listAssignableApis)) + access := r.Group("/api-access") + access.GET("/:clientId/apis", authz.APIsRead, httpserver.Handle(m.handler.listClientApis)) + access.GET("/:clientId/assignable-apis", authz.APIsRead, httpserver.Handle(m.handler.listAssignableApis)) } diff --git a/backend/internal/apikey/handler.go b/backend/internal/apikey/handler.go index 6272a1f2..50ef78c0 100644 --- a/backend/internal/apikey/handler.go +++ b/backend/internal/apikey/handler.go @@ -5,6 +5,7 @@ import ( "github.com/gin-gonic/gin" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/httpserver" "github.com/pocket-id/pocket-id/backend/internal/utils" @@ -31,7 +32,7 @@ func newHandler(service *Service) *handler { func (h *handler) list(c *gin.Context) error { listRequestOptions := utils.ParseListRequestOptions(c) - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID apiKeys, pagination, err := h.service.ListApiKeys(c.Request.Context(), userID, listRequestOptions) if err != nil { @@ -59,7 +60,7 @@ func (h *handler) list(c *gin.Context) error { // @Success 201 {object} apiKeyResponseDto "Created API key with token" // @Router /api/api-keys [post] func (h *handler) create(c *gin.Context) error { - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID var input apiKeyCreateDto err := httpserver.BindJSON(c, &input) @@ -93,7 +94,7 @@ func (h *handler) create(c *gin.Context) error { // @Success 200 {object} apiKeyResponseDto "Renewed API key with new token" // @Router /api/api-keys/{id}/renew [post] func (h *handler) renew(c *gin.Context) error { - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID apiKeyID := c.Param("id") var input apiKeyRenewDto @@ -128,7 +129,7 @@ func (h *handler) renew(c *gin.Context) error { // @Success 204 "No Content" // @Router /api/api-keys/{id} [delete] func (h *handler) revoke(c *gin.Context) error { - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID apiKeyID := c.Param("id") err := h.service.RevokeApiKey(c.Request.Context(), userID, apiKeyID) diff --git a/backend/internal/apikey/module.go b/backend/internal/apikey/module.go index cef50da0..a8d111f0 100644 --- a/backend/internal/apikey/module.go +++ b/backend/internal/apikey/module.go @@ -5,11 +5,11 @@ import ( "errors" "fmt" - "github.com/gin-gonic/gin" francishost "github.com/italypaleale/francis/host" "gorm.io/gorm" "github.com/pocket-id/pocket-id/backend/internal/appconfig" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/httpserver" "github.com/pocket-id/pocket-id/backend/internal/model" ) @@ -63,13 +63,13 @@ func New(ctx context.Context, deps Dependencies) (*Module, error) { } // RegisterRoutes mounts the API key management endpoints -// authWithoutApiKey disables API key authentication so an API key cannot be used to mint or renew further API keys -func (m *Module) RegisterRoutes(apiGroup *gin.RouterGroup, auth, authWithoutApiKey gin.HandlerFunc) { - group := apiGroup.Group("/api-keys") - group.GET("", auth, httpserver.Handle(m.handler.list)) - group.POST("", authWithoutApiKey, httpserver.Handle(m.handler.create)) - group.POST("/:id/renew", authWithoutApiKey, httpserver.Handle(m.handler.renew)) - group.DELETE("/:id", auth, httpserver.Handle(m.handler.revoke)) +// Creating and renewing keys requires a browser session so an API key cannot mint further API keys +func (m *Module) RegisterRoutes(r *authz.Router) { + group := r.Group("/api-keys") + group.GET("", authz.AccountAPIKeys, httpserver.Handle(m.handler.list)) + group.POST("", authz.AccountAPIKeysCreate, httpserver.Handle(m.handler.create)) + group.POST("/:id/renew", authz.AccountAPIKeysCreate, httpserver.Handle(m.handler.renew)) + group.DELETE("/:id", authz.AccountAPIKeys, httpserver.Handle(m.handler.revoke)) } // ValidateApiKey resolves the user that owns the given raw API key diff --git a/backend/internal/apperror/constructors.go b/backend/internal/apperror/constructors.go index b17a8135..e9d691b7 100644 --- a/backend/internal/apperror/constructors.go +++ b/backend/internal/apperror/constructors.go @@ -90,6 +90,11 @@ func MissingPermission() *Error { return New(CodeForbidden, http.StatusForbidden, "You don't have permission to perform this action") } +// MissingScope keeps the generic forbidden code and names the scope the caller lacks so API clients can tell what to request +func MissingScope(scope string) *Error { + return MissingPermission().WithDetail("required_scope", scope) +} + func CrossOriginRequestForbidden(cause error) *Error { return Wrap(cause, CodeForbidden, http.StatusForbidden, "Cross-origin requests are not allowed") } diff --git a/backend/internal/auditlogs/handler.go b/backend/internal/auditlogs/handler.go index f95b7b76..61d982f2 100644 --- a/backend/internal/auditlogs/handler.go +++ b/backend/internal/auditlogs/handler.go @@ -5,6 +5,7 @@ import ( "github.com/gin-gonic/gin" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/utils" ) @@ -31,7 +32,7 @@ func newHandler(service *service) *handler { func (h *handler) listAuditLogsForUserHandler(c *gin.Context) error { listRequestOptions := utils.ParseListRequestOptions(c) - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID // Fetch audit logs for the user logs, pagination, err := h.service.ListAuditLogsForUser(c.Request.Context(), userID, listRequestOptions) diff --git a/backend/internal/auditlogs/handler_test.go b/backend/internal/auditlogs/handler_test.go index e99b86ea..6fa6a0ba 100644 --- a/backend/internal/auditlogs/handler_test.go +++ b/backend/internal/auditlogs/handler_test.go @@ -2,6 +2,7 @@ package auditlogs import ( "encoding/json" + "errors" "net/http" "net/http/httptest" "testing" @@ -9,6 +10,8 @@ import ( "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" + "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/model" testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing" ) @@ -28,15 +31,20 @@ func TestAuditLogRoutesPreservePermissionsAndResponses(t *testing.T) { }).Error) } - // Keep authentication lightweight while exercising which middleware each route receives + // Keep authentication lightweight while exercising which scope each route requires router := gin.New() router.Use(func(c *gin.Context) { c.Next() - if len(c.Errors) > 0 { - c.JSON(http.StatusInternalServerError, gin.H{"error": c.Errors.String()}) + if len(c.Errors) == 0 { + return } + status := http.StatusInternalServerError + if appErr, ok := errors.AsType[*apperror.Error](c.Errors.Last().Err); ok { + status = appErr.HTTPStatus() + } + c.JSON(status, gin.H{"error": c.Errors.String()}) }) - module.RegisterRoutes(router.Group("/api"), auditLogTestAuth(true), auditLogTestAuth(false)) + module.RegisterRoutes(authz.NewMiddleware(auditLogTestAuthenticator{}).Router(router.Group("/api"))) for _, path := range []string{"/audit-logs", "/audit-logs/all", "/audit-logs/filters/client-names", "/audit-logs/filters/users"} { for _, role := range []string{"", "user", "admin"} { @@ -62,19 +70,24 @@ func TestAuditLogRoutesPreservePermissionsAndResponses(t *testing.T) { } } -func auditLogTestAuth(adminRequired bool) gin.HandlerFunc { - return func(c *gin.Context) { - role := c.GetHeader("X-Test-Role") - if role == "" { - c.AbortWithStatus(http.StatusUnauthorized) - return - } - if adminRequired && role != "admin" { - c.AbortWithStatus(http.StatusForbidden) - return - } - c.Set("userID", "alice") - } +// auditLogTestAuthenticator signs in as alice with the role named in the X-Test-Role header +type auditLogTestAuthenticator struct{} + +func (auditLogTestAuthenticator) Kind() authz.PrincipalKind { + return authz.KindSession +} + +func (auditLogTestAuthenticator) Present(c *gin.Context) bool { + return c.GetHeader("X-Test-Role") != "" +} + +func (auditLogTestAuthenticator) Authenticate(c *gin.Context) (*authz.Principal, error) { + isAdmin := c.GetHeader("X-Test-Role") == "admin" + return &authz.Principal{ + Kind: authz.KindSession, + UserID: "alice", + Scopes: authz.UserScopes(isAdmin, authz.KindSession), + }, nil } func assertAuditLogRouteResponse(t *testing.T, path string, response *httptest.ResponseRecorder) { diff --git a/backend/internal/auditlogs/module.go b/backend/internal/auditlogs/module.go index 111e17c7..89606992 100644 --- a/backend/internal/auditlogs/module.go +++ b/backend/internal/auditlogs/module.go @@ -7,12 +7,12 @@ import ( "fmt" "time" - "github.com/gin-gonic/gin" francishost "github.com/italypaleale/francis/host" "github.com/lestrrat-go/jwx/v4/jwt" "gorm.io/gorm" "github.com/pocket-id/pocket-id/backend/internal/appconfig" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/httpserver" "github.com/pocket-id/pocket-id/backend/internal/iplocation" ) @@ -71,12 +71,12 @@ func New(deps Dependencies) (*Module, error) { return &Module{service: service, handler: newHandler(service)}, nil } -// RegisterRoutes mounts audit-log queries with the existing admin and current-user permissions -func (m *Module) RegisterRoutes(group *gin.RouterGroup, adminAuth, userAuth gin.HandlerFunc) { - group.GET("/audit-logs/all", adminAuth, httpserver.Handle(m.handler.listAllAuditLogsHandler)) - group.GET("/audit-logs", userAuth, httpserver.Handle(m.handler.listAuditLogsForUserHandler)) - group.GET("/audit-logs/filters/client-names", adminAuth, httpserver.Handle(m.handler.listClientNamesHandler)) - group.GET("/audit-logs/filters/users", adminAuth, httpserver.Handle(m.handler.listUserNamesWithIdsHandler)) +// RegisterRoutes mounts the audit-log queries for the caller's own events and for all events +func (m *Module) RegisterRoutes(r *authz.Router) { + r.GET("/audit-logs/all", authz.AuditLogsRead, httpserver.Handle(m.handler.listAllAuditLogsHandler)) + r.GET("/audit-logs", authz.AccountAuditLogs, httpserver.Handle(m.handler.listAuditLogsForUserHandler)) + r.GET("/audit-logs/filters/client-names", authz.AuditLogsRead, httpserver.Handle(m.handler.listClientNamesHandler)) + r.GET("/audit-logs/filters/users", authz.AuditLogsRead, httpserver.Handle(m.handler.listUserNamesWithIdsHandler)) } // Create records an event within the caller's transaction diff --git a/backend/internal/authz/middleware.go b/backend/internal/authz/middleware.go new file mode 100644 index 00000000..6bcacc13 --- /dev/null +++ b/backend/internal/authz/middleware.go @@ -0,0 +1,123 @@ +package authz + +import ( + "github.com/gin-gonic/gin" + + "github.com/pocket-id/pocket-id/backend/internal/apperror" +) + +// Authenticator resolves a principal from one kind of credential +type Authenticator interface { + // Kind reports the kind of principal this authenticator produces + Kind() PrincipalKind + + // Present reports whether the request carries this authenticator's credential at all, without validating it + Present(c *gin.Context) bool + + // Authenticate validates the credential and resolves the principal + // It returns an error with code not_signed_in when the credential is invalid, so the next authenticator gets a chance + // Any other error, such as a disabled user, rejects the request + Authenticate(c *gin.Context) (*Principal, error) +} + +// Middleware authenticates requests and enforces the scope each route declares +type Middleware struct { + authenticators []Authenticator + declared map[string]struct{} +} + +// NewMiddleware creates the authorization middleware +// Authenticators are tried in order and the first one that resolves a principal wins +func NewMiddleware(authenticators ...Authenticator) *Middleware { + return &Middleware{ + authenticators: authenticators, + declared: make(map[string]struct{}), + } +} + +// Router wraps a gin router group so every route registered through it declares its access +func (m *Middleware) Router(group *gin.RouterGroup) *Router { + return &Router{group: group, auth: m} +} + +// IsDeclared reports whether the route was registered through a Router or PublicRouter, so a test can compare gin's route table against the declarations +func (m *Middleware) IsDeclared(method, path string) bool { + _, ok := m.declared[routeKey(method, path)] + return ok +} + +func (m *Middleware) declare(method, path string) { + m.declared[routeKey(method, path)] = struct{}{} +} + +func routeKey(method, path string) string { + return method + " " + path +} + +// require returns the handler that enforces the scope on a route +// An optional route lets requests without a usable credential through as anonymous instead of rejecting them +func (m *Middleware) require(scope Scope, optional bool) gin.HandlerFunc { + return func(c *gin.Context) { + principal, kindRejected, err := m.authenticate(c, scope) + if err != nil { + c.Abort() + _ = c.Error(err) + return + } + + // Requests without a usable credential are anonymous + if principal == nil { + if optional { + c.Next() + return + } + + c.Abort() + if kindRejected { + // Only API keys can be rejected by kind today, so the error tells the caller to use a browser session instead + _ = c.Error(apperror.APIKeyAuthNotAllowed()) + return + } + _ = c.Error(apperror.NotSignedIn()) + return + } + + // A valid credential without the scope is forbidden even on optional routes, so a caller is never silently downgraded to anonymous + if !principal.Scopes.Has(scope) { + c.Abort() + _ = c.Error(apperror.MissingScope(string(scope))) + return + } + + SetPrincipal(c, principal) + c.Next() + } +} + +// authenticate resolves the principal from the first credential that can hold the scope and validates +// Credentials whose kind can never hold the scope are not validated at all, and kindRejected reports that one was present +func (m *Middleware) authenticate(c *gin.Context, scope Scope) (principal *Principal, kindRejected bool, err error) { + for _, authenticator := range m.authenticators { + if !authenticator.Present(c) { + continue + } + + // Skip credentials that could never satisfy the route so they are not validated or marked as used + if !scope.GrantableTo(authenticator.Kind()) { + kindRejected = true + continue + } + + principal, err = authenticator.Authenticate(c) + if err == nil { + return principal, false, nil + } + + // An invalid credential falls through to the next authenticator, while a valid but rejected one ends the request + if !apperror.IsCode(err, apperror.CodeNotSignedIn) { + return nil, false, err + } + } + + return nil, kindRejected, nil +} diff --git a/backend/internal/authz/middleware_test.go b/backend/internal/authz/middleware_test.go new file mode 100644 index 00000000..92655328 --- /dev/null +++ b/backend/internal/authz/middleware_test.go @@ -0,0 +1,211 @@ +package authz + +import ( + "errors" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + + "github.com/pocket-id/pocket-id/backend/internal/apperror" +) + +// fakeAuthenticator accepts any request carrying its header and resolves to the configured outcome +type fakeAuthenticator struct { + kind PrincipalKind + header string + user string + admin bool + err error + calls int +} + +func (a *fakeAuthenticator) Kind() PrincipalKind { + return a.kind +} + +func (a *fakeAuthenticator) Present(c *gin.Context) bool { + return c.GetHeader(a.header) != "" +} + +func (a *fakeAuthenticator) Authenticate(*gin.Context) (*Principal, error) { + a.calls++ + if a.err != nil { + return nil, a.err + } + + principal := &Principal{Kind: a.kind, UserID: a.user, Scopes: UserScopes(a.admin, a.kind)} + if a.kind == KindSession { + principal.AuthenticationMethod = "passkey" + principal.AuthenticationTime = time.Unix(1700000000, 0) + } + return principal, nil +} + +type middlewareResult struct { + status int + err error + principal Principal +} + +// serve runs one request through a route that requires the scope and reports what the middleware decided +func serve(t *testing.T, m *Middleware, scope Scope, optional bool, headers map[string]string) middlewareResult { + t.Helper() + gin.SetMode(gin.TestMode) + + var result middlewareResult + router := gin.New() + router.Use(func(c *gin.Context) { + c.Next() + if len(c.Errors) > 0 { + result.err = c.Errors.Last().Err + } + }) + + r := m.Router(router.Group("/api")) + if optional { + r = r.Optional() + } + r.GET("/route", scope, func(c *gin.Context) { + result.principal = PrincipalFrom(c) + c.Status(http.StatusNoContent) + }) + + req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/api/route", nil) + for key, value := range headers { + req.Header.Set(key, value) + } + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, req) + + result.status = recorder.Code + return result +} + +func TestMiddlewareAuthorizes(t *testing.T) { + session := &fakeAuthenticator{kind: KindSession, header: "X-Session", user: "session-user"} + apiKey := &fakeAuthenticator{kind: KindAPIKey, header: "X-Key", user: "key-user", admin: true} + m := NewMiddleware(session, apiKey) + + t.Run("attaches the principal", func(t *testing.T) { + result := serve(t, m, AccountRead, false, map[string]string{"X-Session": "1"}) + + require.Equal(t, http.StatusNoContent, result.status) + require.Equal(t, "session-user", result.principal.UserID) + require.Equal(t, KindSession, result.principal.Kind) + require.Equal(t, "passkey", result.principal.AuthenticationMethod) + }) + + t.Run("rejects missing credentials", func(t *testing.T) { + result := serve(t, m, AccountRead, false, nil) + + require.True(t, apperror.IsCode(result.err, apperror.CodeNotSignedIn)) + }) + + t.Run("rejects a principal without the scope and names the scope", func(t *testing.T) { + result := serve(t, m, UsersRead, false, map[string]string{"X-Session": "1"}) + + var appErr *apperror.Error + require.ErrorAs(t, result.err, &appErr) + require.Equal(t, apperror.CodeForbidden, appErr.Code()) + require.Equal(t, string(UsersRead), appErr.Details()["required_scope"]) + }) + + t.Run("the first authenticator that resolves wins", func(t *testing.T) { + result := serve(t, m, AccountRead, false, map[string]string{"X-Session": "1", "X-Key": "1"}) + + require.Equal(t, "session-user", result.principal.UserID) + }) + + t.Run("rejects a credential kind that can never hold the scope without validating it", func(t *testing.T) { + calls := apiKey.calls + result := serve(t, m, AccountSession, false, map[string]string{"X-Key": "1"}) + + require.True(t, apperror.IsCode(result.err, apperror.CodeAPIKeyAuthNotAllowed)) + require.Equal(t, calls, apiKey.calls, "the API key must not be validated or marked as used") + }) + + t.Run("a valid credential of an allowed kind wins over a rejected kind", func(t *testing.T) { + result := serve(t, m, AccountSession, false, map[string]string{"X-Session": "1", "X-Key": "1"}) + + require.Equal(t, http.StatusNoContent, result.status) + require.Equal(t, "session-user", result.principal.UserID) + }) +} + +func TestMiddlewareFallsThroughInvalidCredentials(t *testing.T) { + invalidSession := &fakeAuthenticator{kind: KindSession, header: "X-Session", err: apperror.NotSignedIn()} + apiKey := &fakeAuthenticator{kind: KindAPIKey, header: "X-Key", user: "key-user"} + m := NewMiddleware(invalidSession, apiKey) + + result := serve(t, m, AccountRead, false, map[string]string{"X-Session": "1", "X-Key": "1"}) + + require.Equal(t, http.StatusNoContent, result.status) + require.Equal(t, "key-user", result.principal.UserID) +} + +func TestMiddlewareStopsOnRejectedCredentials(t *testing.T) { + disabledSession := &fakeAuthenticator{kind: KindSession, header: "X-Session", err: apperror.UserDisabled()} + apiKey := &fakeAuthenticator{kind: KindAPIKey, header: "X-Key", user: "key-user"} + m := NewMiddleware(disabledSession, apiKey) + + for _, optional := range []bool{false, true} { + result := serve(t, m, AccountRead, optional, map[string]string{"X-Session": "1", "X-Key": "1"}) + + require.True(t, apperror.IsCode(result.err, apperror.CodeUserDisabled), "optional=%v", optional) + require.Zero(t, apiKey.calls) + } +} + +func TestMiddlewareOptional(t *testing.T) { + session := &fakeAuthenticator{kind: KindSession, header: "X-Session", user: "session-user"} + invalidSession := &fakeAuthenticator{kind: KindSession, header: "X-Expired", err: apperror.NotSignedIn()} + apiKey := &fakeAuthenticator{kind: KindAPIKey, header: "X-Key", user: "key-user"} + m := NewMiddleware(session, invalidSession, apiKey) + + t.Run("continues anonymously without credentials", func(t *testing.T) { + result := serve(t, m, AccountSession, true, nil) + + require.Equal(t, http.StatusNoContent, result.status) + require.Equal(t, Principal{}, result.principal) + }) + + t.Run("continues anonymously with an invalid credential", func(t *testing.T) { + result := serve(t, m, AccountSession, true, map[string]string{"X-Expired": "1"}) + + require.Equal(t, http.StatusNoContent, result.status) + require.Equal(t, Principal{}, result.principal) + }) + + t.Run("ignores a credential kind that can never hold the scope", func(t *testing.T) { + result := serve(t, m, AccountSession, true, map[string]string{"X-Key": "1"}) + + require.Equal(t, http.StatusNoContent, result.status) + require.Equal(t, Principal{}, result.principal) + require.Zero(t, apiKey.calls) + }) + + t.Run("attaches the principal when signed in", func(t *testing.T) { + result := serve(t, m, AccountSession, true, map[string]string{"X-Session": "1"}) + + require.Equal(t, "session-user", result.principal.UserID) + }) + + t.Run("still rejects a signed-in principal without the scope", func(t *testing.T) { + result := serve(t, m, UsersRead, true, map[string]string{"X-Session": "1"}) + + require.True(t, apperror.IsCode(result.err, apperror.CodeForbidden)) + }) +} + +func TestMiddlewarePassesThroughUnexpectedErrors(t *testing.T) { + failure := errors.New("database unavailable") + m := NewMiddleware(&fakeAuthenticator{kind: KindSession, header: "X-Session", err: failure}) + + result := serve(t, m, AccountRead, false, map[string]string{"X-Session": "1"}) + + require.ErrorIs(t, result.err, failure) +} diff --git a/backend/internal/authz/principal.go b/backend/internal/authz/principal.go new file mode 100644 index 00000000..28dd2529 --- /dev/null +++ b/backend/internal/authz/principal.go @@ -0,0 +1,36 @@ +package authz + +import ( + "time" + + "github.com/gin-gonic/gin" +) + +const principalContextKey = "authz.principal" + +// Principal is the authenticated caller of a request together with the scopes it holds +type Principal struct { + Kind PrincipalKind + UserID string + Scopes ScopeSet + + // AuthenticationMethod and AuthenticationTime describe how the session was established and are only set for KindSession + AuthenticationMethod string + AuthenticationTime time.Time +} + +// PrincipalFrom returns the principal the authorization middleware attached to the request +// Anonymous requests on optional and public routes get the zero Principal, whose UserID is empty +func PrincipalFrom(c *gin.Context) Principal { + value, _ := c.Get(principalContextKey) + if principal, ok := value.(*Principal); ok && principal != nil { + return *principal + } + return Principal{} +} + +// SetPrincipal attaches the principal to the request +// The middleware calls it after authorizing a request, and tests use it to call handlers directly +func SetPrincipal(c *gin.Context, principal *Principal) { + c.Set(principalContextKey, principal) +} diff --git a/backend/internal/authz/router.go b/backend/internal/authz/router.go new file mode 100644 index 00000000..0442e677 --- /dev/null +++ b/backend/internal/authz/router.go @@ -0,0 +1,102 @@ +package authz + +import ( + "fmt" + "net/http" + "path" + "strings" + + "github.com/gin-gonic/gin" +) + +// Router registers routes together with the scope each one requires +// Every route under /api must be registered through a Router or PublicRouter, which a test over the complete route table checks with IsDeclared +type Router struct { + group *gin.RouterGroup + auth *Middleware + optional bool +} + +// Group returns a router for routes below the relative path +func (r *Router) Group(relativePath string) *Router { + return &Router{group: r.group.Group(relativePath), auth: r.auth, optional: r.optional} +} + +// Optional returns a router whose routes let requests without a usable credential through as anonymous +// Handlers on these routes must check PrincipalFrom before relying on a signed-in user +func (r *Router) Optional() *Router { + return &Router{group: r.group, auth: r.auth, optional: true} +} + +// Public returns a router for routes that require no authentication at all +func (r *Router) Public() *PublicRouter { + return &PublicRouter{group: r.group, auth: r.auth} +} + +func (r *Router) GET(relativePath string, scope Scope, handlers ...gin.HandlerFunc) { + r.Handle(http.MethodGet, relativePath, scope, handlers...) +} + +func (r *Router) POST(relativePath string, scope Scope, handlers ...gin.HandlerFunc) { + r.Handle(http.MethodPost, relativePath, scope, handlers...) +} + +func (r *Router) PUT(relativePath string, scope Scope, handlers ...gin.HandlerFunc) { + r.Handle(http.MethodPut, relativePath, scope, handlers...) +} + +func (r *Router) PATCH(relativePath string, scope Scope, handlers ...gin.HandlerFunc) { + r.Handle(http.MethodPatch, relativePath, scope, handlers...) +} + +func (r *Router) DELETE(relativePath string, scope Scope, handlers ...gin.HandlerFunc) { + r.Handle(http.MethodDelete, relativePath, scope, handlers...) +} + +// Handle registers a route that requires the scope, running authorization before every other handler of the route +func (r *Router) Handle(method, relativePath string, scope Scope, handlers ...gin.HandlerFunc) { + // An unknown scope can never be granted, so it is a programming error just like a duplicate route in gin + if !scope.Known() { + panic(fmt.Sprintf("route %s %s requires unknown scope %q", method, joinPaths(r.group.BasePath(), relativePath), scope)) + } + + chain := make([]gin.HandlerFunc, 0, len(handlers)+1) + chain = append(chain, r.auth.require(scope, r.optional)) + chain = append(chain, handlers...) + r.group.Handle(method, relativePath, chain...) + r.auth.declare(method, joinPaths(r.group.BasePath(), relativePath)) +} + +// PublicRouter registers routes that require no authentication +// Routing them through here keeps public access an explicit decision instead of a missing middleware +type PublicRouter struct { + group *gin.RouterGroup + auth *Middleware +} + +func (r *PublicRouter) GET(relativePath string, handlers ...gin.HandlerFunc) { + r.Handle(http.MethodGet, relativePath, handlers...) +} + +func (r *PublicRouter) POST(relativePath string, handlers ...gin.HandlerFunc) { + r.Handle(http.MethodPost, relativePath, handlers...) +} + +// Handle registers a public route +func (r *PublicRouter) Handle(method, relativePath string, handlers ...gin.HandlerFunc) { + r.group.Handle(method, relativePath, handlers...) + r.auth.declare(method, joinPaths(r.group.BasePath(), relativePath)) +} + +// joinPaths mirrors how gin builds a route's absolute path so declared routes match gin's route table +func joinPaths(absolutePath, relativePath string) string { + if relativePath == "" { + return absolutePath + } + + finalPath := path.Join(absolutePath, relativePath) + if strings.HasSuffix(relativePath, "/") && !strings.HasSuffix(finalPath, "/") { + return finalPath + "/" + } + return finalPath +} diff --git a/backend/internal/authz/router_test.go b/backend/internal/authz/router_test.go new file mode 100644 index 00000000..1c2afae3 --- /dev/null +++ b/backend/internal/authz/router_test.go @@ -0,0 +1,75 @@ +package authz + +import ( + "net/http" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func noop(c *gin.Context) { + c.Status(http.StatusNoContent) +} + +func TestIsDeclared(t *testing.T) { + gin.SetMode(gin.TestMode) + engine := gin.New() + m := NewMiddleware() + apiGroup := engine.Group("/api") + api := m.Router(apiGroup) + + // Declare routes through every router variant + api.GET("/users", UsersRead, noop) + api.Group("/api-keys").POST("", AccountAPIKeysCreate, noop) + api.Group("/nested").Group("/deeper").DELETE("/:id", UsersWrite, noop) + api.Optional().GET("/optional", AccountSession, noop) + api.Public().POST("/signup", noop) + m.Router(engine.Group("/")).Optional().GET("/authorize", AccountSession, noop) + + // Register routes past the routers + apiGroup.GET("/forgotten", noop) + apiGroup.POST("/users", noop) + + declared := map[string]bool{} + for _, route := range engine.Routes() { + declared[route.Method+" "+route.Path] = m.IsDeclared(route.Method, route.Path) + } + require.Equal(t, map[string]bool{ + "GET /api/users": true, + "POST /api/api-keys": true, + "DELETE /api/nested/deeper/:id": true, + "GET /api/optional": true, + "POST /api/signup": true, + "GET /authorize": true, + "GET /api/forgotten": false, + "POST /api/users": false, + }, declared) +} + +func TestRouterRejectsUnknownScopes(t *testing.T) { + gin.SetMode(gin.TestMode) + r := NewMiddleware().Router(gin.New().Group("/api")) + + require.PanicsWithValue(t, `route GET /api/users requires unknown scope "users:everything"`, func() { + r.GET("/users", Scope("users:everything"), noop) + }) +} + +func TestJoinPathsMatchesGin(t *testing.T) { + gin.SetMode(gin.TestMode) + + for _, test := range []struct{ base, relative string }{ + {"/api", ""}, + {"/api", "/users"}, + {"/api/", "users"}, + {"/api", "/users/"}, + {"/", "/authorize"}, + {"/api/api-keys", ""}, + } { + engine := gin.New() + engine.Group(test.base).GET(test.relative, noop) + + require.Equal(t, engine.Routes()[0].Path, joinPaths(engine.Group(test.base).BasePath(), test.relative), "%+v", test) + } +} diff --git a/backend/internal/authz/scope.go b/backend/internal/authz/scope.go new file mode 100644 index 00000000..edb39cd5 --- /dev/null +++ b/backend/internal/authz/scope.go @@ -0,0 +1,149 @@ +package authz + +// Scope is a permission that a principal must hold to call a route +// Keys follow the resource:action pattern and are valid RFC 6749 scope tokens so they can later appear in API key records and OAuth access tokens unchanged +type Scope string + +// Account scopes cover the caller's own account and are held by every signed-in user +const ( + AccountRead Scope = "account:read" + AccountWrite Scope = "account:write" + AccountPasskeys Scope = "account:passkeys" + AccountAPIKeys Scope = "account:api-keys" + AccountApps Scope = "account:apps" + AccountAuditLogs Scope = "account:audit-logs" + AccountSession Scope = "account:session" + AccountPasskeysEnroll Scope = "account:passkeys:enroll" + AccountAPIKeysCreate Scope = "account:api-keys:create" +) + +// Admin scopes cover other users' data and the instance configuration +const ( + UsersRead Scope = "users:read" + UsersWrite Scope = "users:write" + GroupsRead Scope = "groups:read" + GroupsWrite Scope = "groups:write" + OidcClientsRead Scope = "oidc-clients:read" + OidcClientsWrite Scope = "oidc-clients:write" + APIsRead Scope = "apis:read" + APIsWrite Scope = "apis:write" + ConfigRead Scope = "config:read" + ConfigWrite Scope = "config:write" + AuditLogsRead Scope = "audit-logs:read" +) + +// Category groups scopes by whose data they reach +type Category int + +const ( + // CategoryAccount scopes act on the caller's own account + CategoryAccount Category = iota + 1 + // CategoryAdmin scopes act on other users or on the instance + CategoryAdmin +) + +// PrincipalKind identifies the kind of credential a principal authenticated with +// Kinds are bit flags so a scope can list every kind that may hold it +type PrincipalKind uint8 + +const ( + // KindSession is a browser session established by signing in to Pocket ID + KindSession PrincipalKind = 1 << iota + // KindAPIKey is a personal API key sent in the X-API-Key header + KindAPIKey + // KindOAuthUser is an OAuth access token issued to a client acting on behalf of a user + KindOAuthUser + // KindOAuthClient is an OAuth access token issued to a client acting as itself through the client credentials grant + KindOAuthClient +) + +// delegated lists the kinds that act for a user, which is every kind except a client acting as itself +const delegated = KindSession | KindAPIKey | KindOAuthUser + +type definition struct { + scope Scope + category Category + grantableTo PrincipalKind +} + +// catalog is the complete list of scopes +// grantableTo restricts which credential kinds can ever hold a scope, independent of the user's role +// Session-only scopes guard actions that must never be reachable with a long-lived or third-party credential, such as enrolling passkeys or minting API keys +var catalog = []definition{ + {AccountRead, CategoryAccount, delegated}, + {AccountWrite, CategoryAccount, delegated}, + {AccountPasskeys, CategoryAccount, delegated}, + {AccountAPIKeys, CategoryAccount, delegated}, + {AccountApps, CategoryAccount, delegated}, + {AccountAuditLogs, CategoryAccount, delegated}, + {AccountSession, CategoryAccount, KindSession}, + {AccountPasskeysEnroll, CategoryAccount, KindSession}, + {AccountAPIKeysCreate, CategoryAccount, KindSession}, + + {UsersRead, CategoryAdmin, delegated}, + {UsersWrite, CategoryAdmin, delegated}, + {GroupsRead, CategoryAdmin, delegated}, + {GroupsWrite, CategoryAdmin, delegated}, + {OidcClientsRead, CategoryAdmin, delegated}, + {OidcClientsWrite, CategoryAdmin, delegated}, + {APIsRead, CategoryAdmin, delegated}, + {APIsWrite, CategoryAdmin, delegated}, + {ConfigRead, CategoryAdmin, delegated}, + {ConfigWrite, CategoryAdmin, delegated}, + {AuditLogsRead, CategoryAdmin, delegated}, +} + +var definitions = indexCatalog(catalog) + +func indexCatalog(entries []definition) map[Scope]definition { + index := make(map[Scope]definition, len(entries)) + for _, entry := range entries { + index[entry.scope] = entry + } + return index +} + +// Known reports whether the scope is part of the catalog +func (s Scope) Known() bool { + _, ok := definitions[s] + return ok +} + +// GrantableTo reports whether a principal of the given kind can ever hold the scope +func (s Scope) GrantableTo(kind PrincipalKind) bool { + return definitions[s].grantableTo&kind != 0 +} + +// ScopeSet is an unordered set of scopes +type ScopeSet map[Scope]struct{} + +// NewScopeSet creates a set containing the given scopes +func NewScopeSet(scopes ...Scope) ScopeSet { + set := make(ScopeSet, len(scopes)) + for _, scope := range scopes { + set[scope] = struct{}{} + } + return set +} + +// Has reports whether the set contains the scope +func (s ScopeSet) Has(scope Scope) bool { + _, ok := s[scope] + return ok +} + +// UserScopes returns the scopes a user holds when authenticated with a credential of the given kind +// The admin flag stands in for roles: admins hold every scope and other users hold the account scopes +func UserScopes(isAdmin bool, kind PrincipalKind) ScopeSet { + set := make(ScopeSet, len(catalog)) + for _, entry := range catalog { + if entry.grantableTo&kind == 0 { + continue + } + if entry.category == CategoryAdmin && !isAdmin { + continue + } + set[entry.scope] = struct{}{} + } + return set +} diff --git a/backend/internal/authz/scope_test.go b/backend/internal/authz/scope_test.go new file mode 100644 index 00000000..370d8934 --- /dev/null +++ b/backend/internal/authz/scope_test.go @@ -0,0 +1,70 @@ +package authz + +import ( + "strings" + "testing" + + "github.com/ory/fosite" + "github.com/stretchr/testify/require" +) + +func TestCatalogInvariants(t *testing.T) { + // Scope keys end up in API key records and OAuth tokens, so they must be valid scope tokens that never collide with the identity scopes + reserved := []string{"openid", "profile", "email", "email_verified", "groups", "offline_access"} + + seen := make(map[Scope]struct{}, len(catalog)) + for _, entry := range catalog { + t.Run(string(entry.scope), func(t *testing.T) { + require.True(t, fosite.IsValidScopeToken(string(entry.scope)), "scope must be a valid RFC 6749 scope token") + require.NotContains(t, reserved, strings.ToLower(string(entry.scope))) + + _, duplicate := seen[entry.scope] + require.False(t, duplicate, "scope is listed twice") + seen[entry.scope] = struct{}{} + + require.Contains(t, []Category{CategoryAccount, CategoryAdmin}, entry.category) + require.NotZero(t, entry.grantableTo, "a scope nobody can hold can never pass a route") + + // Account scopes are named after the account and admin scopes after a resource, so the prefix alone tells callers what they reach + require.Equal(t, entry.category == CategoryAccount, strings.HasPrefix(string(entry.scope), "account:")) + }) + } + require.Len(t, definitions, len(catalog)) +} + +func TestClientCredentialsCannotHoldAnyScope(t *testing.T) { + // Service identities are not supported yet, see the scope-authorization plan + for _, entry := range catalog { + require.False(t, entry.scope.GrantableTo(KindOAuthClient), entry.scope) + } +} + +func TestUserScopes(t *testing.T) { + t.Run("admins hold every scope their credential kind allows", func(t *testing.T) { + scopes := UserScopes(true, KindSession) + require.Len(t, scopes, len(catalog)) + }) + + t.Run("regular users hold only account scopes", func(t *testing.T) { + scopes := UserScopes(false, KindSession) + for _, entry := range catalog { + require.Equal(t, entry.category == CategoryAccount, scopes.Has(entry.scope), entry.scope) + } + }) + + t.Run("API keys never hold session-only scopes, even for admins", func(t *testing.T) { + scopes := UserScopes(true, KindAPIKey) + require.True(t, scopes.Has(UsersWrite)) + require.True(t, scopes.Has(AccountAPIKeys)) + require.False(t, scopes.Has(AccountSession)) + require.False(t, scopes.Has(AccountPasskeysEnroll)) + require.False(t, scopes.Has(AccountAPIKeysCreate)) + }) +} + +func TestUnknownScope(t *testing.T) { + unknown := Scope("unknown:scope") + require.False(t, unknown.Known()) + require.False(t, unknown.GrantableTo(KindSession)) + require.True(t, UsersRead.Known()) +} diff --git a/backend/internal/bootstrap/e2etest_router_bootstrap.go b/backend/internal/bootstrap/e2etest_router_bootstrap.go index df6d3558..f27a10ec 100644 --- a/backend/internal/bootstrap/e2etest_router_bootstrap.go +++ b/backend/internal/bootstrap/e2etest_router_bootstrap.go @@ -6,17 +6,17 @@ import ( "log/slog" "os" - "github.com/gin-gonic/gin" "gorm.io/gorm" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/controller" "github.com/pocket-id/pocket-id/backend/internal/service" ) // When building for E2E tests, add the e2etest controller func init() { - registerTestControllers = []func(apiGroup *gin.RouterGroup, db *gorm.DB, svc *services){ - func(apiGroup *gin.RouterGroup, db *gorm.DB, svc *services) { + registerTestControllers = []func(apiRouter *authz.Router, db *gorm.DB, svc *services){ + func(apiRouter *authz.Router, db *gorm.DB, svc *services) { testService, err := service.NewTestService(db, svc.actors, svc.appConfigService, svc.jwtService, svc.ldapSyncModule, svc.fileStorage) if err != nil { slog.Error("Failed to initialize test service", slog.Any("error", err)) @@ -24,7 +24,7 @@ func init() { return } - controller.NewTestController(apiGroup, testService) + controller.NewTestController(apiRouter.Public(), testService) }, } } diff --git a/backend/internal/bootstrap/route_coverage_test.go b/backend/internal/bootstrap/route_coverage_test.go new file mode 100644 index 00000000..7b2062c6 --- /dev/null +++ b/backend/internal/bootstrap/route_coverage_test.go @@ -0,0 +1,68 @@ +package bootstrap + +import ( + "strings" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + + "github.com/pocket-id/pocket-id/backend/internal/api" + "github.com/pocket-id/pocket-id/backend/internal/apikey" + "github.com/pocket-id/pocket-id/backend/internal/auditlogs" + "github.com/pocket-id/pocket-id/backend/internal/devicelogin" + "github.com/pocket-id/pocket-id/backend/internal/emailverification" + "github.com/pocket-id/pocket-id/backend/internal/environment" + "github.com/pocket-id/pocket-id/backend/internal/ldapsync" + "github.com/pocket-id/pocket-id/backend/internal/logopreset" + "github.com/pocket-id/pocket-id/backend/internal/middleware" + "github.com/pocket-id/pocket-id/backend/internal/oidc" + "github.com/pocket-id/pocket-id/backend/internal/onetimeaccess" + "github.com/pocket-id/pocket-id/backend/internal/scimsync" + "github.com/pocket-id/pocket-id/backend/internal/usersignup" + "github.com/pocket-id/pocket-id/backend/internal/webauthn" +) + +// TestEveryAPIRouteDeclaresItsAccess builds the complete route table and fails when an API route was registered without declaring a scope or public access +func TestEveryAPIRouteDeclaresItsAccess(t *testing.T) { + gin.SetMode(gin.TestMode) + + // Registration only stores handler references and never calls them, so modules without dependencies are enough + svc := &services{ + apiKeyModule: &apikey.Module{}, + auditLogsModule: &auditlogs.Module{}, + deviceLoginModule: &devicelogin.Module{}, + ldapSyncModule: &ldapsync.Module{}, + scimSyncModule: &scimsync.Module{}, + oidcModule: &oidc.Module{}, + webauthnModule: &webauthn.Module{}, + userSignUpModule: &usersignup.Module{}, + oneTimeAccessModule: &onetimeaccess.Module{}, + emailVerificationModule: &emailverification.Module{}, + apiModule: &api.Module{}, + environmentModule: &environment.Module{}, + logoPresetModule: &logopreset.Module{}, + } + + engine := gin.New() + auth := middleware.NewAuthorization(nil, nil, nil) + require.NoError(t, registerRoutes(engine, nil, svc, auth, nil)) + + // Collect every API route that bypassed the authz routers + apiRoutes := 0 + var undeclared []string + for _, route := range engine.Routes() { + if !strings.HasPrefix(route.Path, "/api/") { + continue + } + apiRoutes++ + if !auth.IsDeclared(route.Method, route.Path) { + undeclared = append(undeclared, route.Method+" "+route.Path) + } + } + + require.Empty(t, undeclared, "register these routes through authz.Router with a scope, or through Public() when they need no authentication") + + // Guard against the table silently shrinking, which would make the coverage check vacuous + require.Greater(t, apiRoutes, 100) +} diff --git a/backend/internal/bootstrap/router_bootstrap.go b/backend/internal/bootstrap/router_bootstrap.go index 3863f248..e5fb8df0 100644 --- a/backend/internal/bootstrap/router_bootstrap.go +++ b/backend/internal/bootstrap/router_bootstrap.go @@ -24,6 +24,7 @@ import ( "gorm.io/gorm" "github.com/pocket-id/pocket-id/backend/frontend" + "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/controller" "github.com/pocket-id/pocket-id/backend/internal/middleware" @@ -32,14 +33,15 @@ import ( ) // This is used to register additional controllers for tests -var registerTestControllers []func(apiGroup *gin.RouterGroup, db *gorm.DB, svc *services) +var registerTestControllers []func(apiRouter *authz.Router, db *gorm.DB, svc *services) func initRouter(db *gorm.DB, svc *services, rateLimitServices map[string]*ratelimit.RateLimitService) (servicerunner.Service, error) { r, err := initEngine() if err != nil { return nil, err } - err = registerRoutes(r, db, svc, rateLimitServices) + auth := middleware.NewAuthorization(svc.apiKeyModule, svc.userService, svc.jwtService) + err = registerRoutes(r, db, svc, auth, rateLimitServices) if err != nil { return nil, err } @@ -135,7 +137,7 @@ func registerGlobalMiddleware(r *gin.Engine) { r.Use(middleware.NewCrossOriginProtectionMiddleware(common.EnvConfig.AppURL).Add()) } -func registerRoutes(r *gin.Engine, db *gorm.DB, svc *services, rateLimitServices map[string]*ratelimit.RateLimitService) error { +func registerRoutes(r *gin.Engine, db *gorm.DB, svc *services, auth *authz.Middleware, rateLimitServices map[string]*ratelimit.RateLimitService) error { err := frontend.RegisterFrontend(r) if errors.Is(err, frontend.ErrFrontendNotIncluded) { @@ -145,7 +147,6 @@ func registerRoutes(r *gin.Engine, db *gorm.DB, svc *services, rateLimitServices } // Initialize middleware for specific routes - authMiddleware := middleware.NewAuthMiddleware(svc.apiKeyModule, svc.userService, svc.jwtService) fileSizeLimitMiddleware := middleware.NewFileSizeLimitMiddleware() rateLimitMiddleware := middleware.NewRateLimitMiddleware(rateLimitServices) apiRateLimitMiddleware := rateLimitMiddleware.Add(middleware.RateLimitAPI) @@ -155,55 +156,44 @@ func registerRoutes(r *gin.Engine, db *gorm.DB, svc *services, rateLimitServices apiGroup.Use(middleware.NewClientIDParamMiddleware().Add()) baseGroup := r.Group("/", apiRateLimitMiddleware) - svc.apiKeyModule.RegisterRoutes(apiGroup, - authMiddleware.WithAdminNotRequired().Add(), - authMiddleware.WithAdminNotRequired().WithApiKeyAuthDisabled().Add(), - ) - svc.webauthnModule.RegisterRoutes(apiGroup, - authMiddleware.WithAdminNotRequired().Add(), - authMiddleware.WithAdminNotRequired().WithApiKeyAuthDisabled().Add(), + // Every route below declares the scope it requires, or that it is public, through these routers + apiRouter := auth.Router(apiGroup) + baseRouter := auth.Router(baseGroup) + + svc.apiKeyModule.RegisterRoutes(apiRouter) + svc.webauthnModule.RegisterRoutes(apiRouter, rateLimitMiddleware.Add(middleware.RateLimitWebauthnLogin), rateLimitMiddleware.Add(middleware.RateLimitWebauthnReauthenticate), ) - svc.deviceLoginModule.RegisterRoutes(apiGroup, - authMiddleware.WithAdminNotRequired().WithApiKeyAuthDisabled().Add(), + svc.deviceLoginModule.RegisterRoutes(apiRouter, rateLimitMiddleware.Add(middleware.RateLimitDeviceLoginCreate), rateLimitMiddleware.Add(middleware.RateLimitDeviceLoginExchange), rateLimitMiddleware.Add(middleware.RateLimitDeviceLoginVerification), ) - controller.NewOidcController(apiGroup, authMiddleware, fileSizeLimitMiddleware, svc.oidcService, svc.appConfigService) - controller.NewUserController(apiGroup, authMiddleware, fileSizeLimitMiddleware, svc.appConfigService, svc.userService, svc.webauthnModule, rateLimitMiddleware.Add(middleware.RateLimitUpdateOwnAccount)) - controller.NewAppConfigController(apiGroup, authMiddleware, svc.appConfigService, svc.emailModule) - svc.ldapSyncModule.RegisterRoutes(apiGroup, authMiddleware.Add()) - controller.NewAppImagesController(apiGroup, authMiddleware, fileSizeLimitMiddleware, svc.appImagesService) - svc.auditLogsModule.RegisterRoutes(apiGroup, authMiddleware.Add(), authMiddleware.WithAdminNotRequired().Add()) - controller.NewUserGroupController(apiGroup, authMiddleware, svc.appConfigService, svc.userGroupService) - svc.apiModule.RegisterRoutes(apiGroup, authMiddleware.Add()) - controller.NewCustomClaimController(apiGroup, authMiddleware, svc.customClaimService) - svc.environmentModule.RegisterRoutes(apiGroup, authMiddleware.WithAdminNotRequired().Add()) - svc.logoPresetModule.RegisterRoutes(apiGroup, authMiddleware.Add()) - svc.scimSyncModule.RegisterRoutes(apiGroup, authMiddleware.Add()) - svc.userSignUpModule.RegisterRoutes(apiGroup, - authMiddleware.Add(), - rateLimitMiddleware.Add(middleware.RateLimitSignup), - ) - svc.oneTimeAccessModule.RegisterRoutes(apiGroup, - authMiddleware.Add(), + controller.NewOidcController(apiRouter, fileSizeLimitMiddleware, svc.oidcService, svc.appConfigService) + controller.NewUserController(apiRouter, fileSizeLimitMiddleware, svc.appConfigService, svc.userService, svc.webauthnModule, rateLimitMiddleware.Add(middleware.RateLimitUpdateOwnAccount)) + controller.NewAppConfigController(apiRouter, svc.appConfigService, svc.emailModule) + svc.ldapSyncModule.RegisterRoutes(apiRouter) + controller.NewAppImagesController(apiRouter, fileSizeLimitMiddleware, svc.appImagesService) + svc.auditLogsModule.RegisterRoutes(apiRouter) + controller.NewUserGroupController(apiRouter, svc.appConfigService, svc.userGroupService) + svc.apiModule.RegisterRoutes(apiRouter) + controller.NewCustomClaimController(apiRouter, svc.customClaimService) + svc.environmentModule.RegisterRoutes(apiRouter) + svc.logoPresetModule.RegisterRoutes(apiRouter) + svc.scimSyncModule.RegisterRoutes(apiRouter) + svc.userSignUpModule.RegisterRoutes(apiRouter, rateLimitMiddleware.Add(middleware.RateLimitSignup)) + svc.oneTimeAccessModule.RegisterRoutes(apiRouter, rateLimitMiddleware.Add(middleware.RateLimitOneTimeAccessToken), rateLimitMiddleware.Add(middleware.RateLimitOneTimeAccessEmail), ) - svc.emailVerificationModule.RegisterRoutes( - apiGroup, - authMiddleware.WithAdminNotRequired().Add(), + svc.emailVerificationModule.RegisterRoutes(apiRouter, rateLimitMiddleware.Add(middleware.RateLimitSendEmailVerification), rateLimitMiddleware.Add(middleware.RateLimitVerifyEmail), ) + svc.oidcModule.RegisterRoutes(baseRouter, apiRouter) - optionalBrowserAuth := authMiddleware.WithAdminNotRequired().WithSuccessOptional().WithApiKeyAuthDisabled().Add() - browserAuth := authMiddleware.WithAdminNotRequired().WithApiKeyAuthDisabled().Add() - svc.oidcModule.RegisterRoutes(baseGroup, apiGroup, optionalBrowserAuth, browserAuth) - - registerTestRoutes(apiGroup, db, svc) + registerTestRoutes(apiRouter, db, svc) controller.NewWellKnownController(baseGroup, svc.jwtService, svc.appConfigService.GetCIMDURLAllowlist) @@ -217,13 +207,13 @@ func registerRoutes(r *gin.Engine, db *gorm.DB, svc *services, rateLimitServices return nil } -func registerTestRoutes(apiGroup *gin.RouterGroup, db *gorm.DB, svc *services) { +func registerTestRoutes(apiRouter *authz.Router, db *gorm.DB, svc *services) { if common.EnvConfig.AppEnv.IsProduction() { return } for _, f := range registerTestControllers { - f(apiGroup, db, svc) + f(apiRouter, db, svc) } } diff --git a/backend/internal/controller/app_config_controller.go b/backend/internal/controller/app_config_controller.go index 441a86be..0f814ef2 100644 --- a/backend/internal/controller/app_config_controller.go +++ b/backend/internal/controller/app_config_controller.go @@ -7,10 +7,10 @@ import ( "github.com/gin-gonic/gin" "github.com/pocket-id/pocket-id/backend/internal/appconfig" + "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/dto" "github.com/pocket-id/pocket-id/backend/internal/httpserver" - "github.com/pocket-id/pocket-id/backend/internal/middleware" "github.com/pocket-id/pocket-id/backend/internal/tracing" ) @@ -23,8 +23,7 @@ type TestEmailSender interface { // @Description Initialize routes for application configuration // @Tags Application Configuration func NewAppConfigController( - group *gin.RouterGroup, - authMiddleware *middleware.AuthMiddleware, + r *authz.Router, appConfigService *appconfig.AppConfigService, emailSender TestEmailSender, ) { @@ -33,11 +32,11 @@ func NewAppConfigController( appConfigService: appConfigService, emailSender: emailSender, } - group.GET("/application-configuration", httpserver.Handle(acc.listAppConfigHandler)) - group.GET("/application-configuration/all", authMiddleware.Add(), httpserver.Handle(acc.listAllAppConfigHandler)) - group.PUT("/application-configuration", authMiddleware.Add(), httpserver.Handle(acc.updateAppConfigHandler)) + r.Public().GET("/application-configuration", httpserver.Handle(acc.listAppConfigHandler)) + r.GET("/application-configuration/all", authz.ConfigRead, httpserver.Handle(acc.listAllAppConfigHandler)) + r.PUT("/application-configuration", authz.ConfigWrite, httpserver.Handle(acc.updateAppConfigHandler)) - group.POST("/application-configuration/test-email", authMiddleware.Add(), httpserver.Handle(acc.testEmailHandler)) + r.POST("/application-configuration/test-email", authz.ConfigWrite, httpserver.Handle(acc.testEmailHandler)) } type AppConfigController struct { @@ -169,7 +168,7 @@ func (acc *AppConfigController) testEmailHandler(c *gin.Context) error { return err } - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID err = acc.emailSender.SendTestEmail(c.Request.Context(), dbConfig, userID) if err != nil { diff --git a/backend/internal/controller/app_images_controller.go b/backend/internal/controller/app_images_controller.go index 63d3d39e..3f999490 100644 --- a/backend/internal/controller/app_images_controller.go +++ b/backend/internal/controller/app_images_controller.go @@ -10,6 +10,7 @@ import ( kitutils "github.com/italypaleale/go-kit/utils" "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/dto" "github.com/pocket-id/pocket-id/backend/internal/httpserver" "github.com/pocket-id/pocket-id/backend/internal/middleware" @@ -18,8 +19,7 @@ import ( ) func NewAppImagesController( - group *gin.RouterGroup, - authMiddleware *middleware.AuthMiddleware, + r *authz.Router, fileSizeLimitMiddleware *middleware.FileSizeLimitMiddleware, appImagesService *service.AppImagesService, ) { @@ -27,21 +27,21 @@ func NewAppImagesController( appImagesService: appImagesService, } - group.GET("/application-images/logo", httpserver.Handle(controller.getLogoHandler)) - group.GET("/application-images/email", httpserver.Handle(controller.getEmailLogoHandler)) - group.GET("/application-images/background", httpserver.Handle(controller.getBackgroundImageHandler)) - group.GET("/application-images/favicon", httpserver.Handle(controller.getFaviconHandler)) - group.GET("/application-images/default-profile-picture", authMiddleware.Add(), httpserver.Handle(controller.getDefaultProfilePicture)) + r.Public().GET("/application-images/logo", httpserver.Handle(controller.getLogoHandler)) + r.Public().GET("/application-images/email", httpserver.Handle(controller.getEmailLogoHandler)) + r.Public().GET("/application-images/background", httpserver.Handle(controller.getBackgroundImageHandler)) + r.Public().GET("/application-images/favicon", httpserver.Handle(controller.getFaviconHandler)) + r.GET("/application-images/default-profile-picture", authz.ConfigRead, httpserver.Handle(controller.getDefaultProfilePicture)) - group.PUT("/application-images/logo", authMiddleware.Add(), fileSizeLimitMiddleware.Add(10<<20), httpserver.Handle(controller.updateLogoHandler)) - group.PUT("/application-images/email", authMiddleware.Add(), fileSizeLimitMiddleware.Add(10<<20), httpserver.Handle(controller.updateEmailLogoHandler)) - group.PUT("/application-images/background", authMiddleware.Add(), fileSizeLimitMiddleware.Add(10<<20), httpserver.Handle(controller.updateBackgroundImageHandler)) - group.PUT("/application-images/favicon", authMiddleware.Add(), fileSizeLimitMiddleware.Add(10<<20), httpserver.Handle(controller.updateFaviconHandler)) - group.PUT("/application-images/default-profile-picture", authMiddleware.Add(), fileSizeLimitMiddleware.Add(10<<20), httpserver.Handle(controller.updateDefaultProfilePicture)) + r.PUT("/application-images/logo", authz.ConfigWrite, fileSizeLimitMiddleware.Add(10<<20), httpserver.Handle(controller.updateLogoHandler)) + r.PUT("/application-images/email", authz.ConfigWrite, fileSizeLimitMiddleware.Add(10<<20), httpserver.Handle(controller.updateEmailLogoHandler)) + r.PUT("/application-images/background", authz.ConfigWrite, fileSizeLimitMiddleware.Add(10<<20), httpserver.Handle(controller.updateBackgroundImageHandler)) + r.PUT("/application-images/favicon", authz.ConfigWrite, fileSizeLimitMiddleware.Add(10<<20), httpserver.Handle(controller.updateFaviconHandler)) + r.PUT("/application-images/default-profile-picture", authz.ConfigWrite, fileSizeLimitMiddleware.Add(10<<20), httpserver.Handle(controller.updateDefaultProfilePicture)) - group.DELETE("/application-images/logo", authMiddleware.Add(), httpserver.Handle(controller.deleteLogoHandler)) - group.DELETE("/application-images/background", authMiddleware.Add(), httpserver.Handle(controller.deleteBackgroundImageHandler)) - group.DELETE("/application-images/default-profile-picture", authMiddleware.Add(), httpserver.Handle(controller.deleteDefaultProfilePicture)) + r.DELETE("/application-images/logo", authz.ConfigWrite, httpserver.Handle(controller.deleteLogoHandler)) + r.DELETE("/application-images/background", authz.ConfigWrite, httpserver.Handle(controller.deleteBackgroundImageHandler)) + r.DELETE("/application-images/default-profile-picture", authz.ConfigWrite, httpserver.Handle(controller.deleteDefaultProfilePicture)) } type AppImagesController struct { diff --git a/backend/internal/controller/custom_claim_controller.go b/backend/internal/controller/custom_claim_controller.go index 831e7183..e7dc6bb7 100644 --- a/backend/internal/controller/custom_claim_controller.go +++ b/backend/internal/controller/custom_claim_controller.go @@ -4,9 +4,9 @@ import ( "net/http" "github.com/gin-gonic/gin" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/httpserver" - "github.com/pocket-id/pocket-id/backend/internal/middleware" "github.com/pocket-id/pocket-id/backend/internal/service" ) @@ -14,16 +14,13 @@ import ( // @Summary Custom claim management controller // @Description Initializes all custom claim-related API endpoints // @Tags Custom Claims -func NewCustomClaimController(group *gin.RouterGroup, authMiddleware *middleware.AuthMiddleware, customClaimService *service.CustomClaimService) { +func NewCustomClaimController(r *authz.Router, customClaimService *service.CustomClaimService) { wkc := &CustomClaimController{customClaimService: customClaimService} - customClaimsGroup := group.Group("/custom-claims") - customClaimsGroup.Use(authMiddleware.Add()) - { - customClaimsGroup.GET("/suggestions", httpserver.Handle(wkc.getSuggestionsHandler)) - customClaimsGroup.PUT("/user/:userId", httpserver.Handle(wkc.UpdateCustomClaimsForUserHandler)) - customClaimsGroup.PUT("/user-group/:userGroupId", httpserver.Handle(wkc.UpdateCustomClaimsForUserGroupHandler)) - } + customClaimsGroup := r.Group("/custom-claims") + customClaimsGroup.GET("/suggestions", authz.UsersRead, httpserver.Handle(wkc.getSuggestionsHandler)) + customClaimsGroup.PUT("/user/:userId", authz.UsersWrite, httpserver.Handle(wkc.UpdateCustomClaimsForUserHandler)) + customClaimsGroup.PUT("/user-group/:userGroupId", authz.GroupsWrite, httpserver.Handle(wkc.UpdateCustomClaimsForUserGroupHandler)) } type CustomClaimController struct { diff --git a/backend/internal/controller/e2etest_controller.go b/backend/internal/controller/e2etest_controller.go index 37a2687d..e7815a0d 100644 --- a/backend/internal/controller/e2etest_controller.go +++ b/backend/internal/controller/e2etest_controller.go @@ -7,19 +7,20 @@ import ( "github.com/gin-gonic/gin" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/httpserver" "github.com/pocket-id/pocket-id/backend/internal/service" ) -func NewTestController(group *gin.RouterGroup, testService *service.TestService) { +func NewTestController(r *authz.PublicRouter, testService *service.TestService) { testController := &TestController{TestService: testService} - group.POST("/test/reset", httpserver.Handle(testController.resetAndSeedHandler)) - group.POST("/test/accesstoken", httpserver.Handle(testController.signAccessToken)) - group.POST("/test/refreshtoken", httpserver.Handle(testController.signRefreshToken)) + r.POST("/test/reset", httpserver.Handle(testController.resetAndSeedHandler)) + r.POST("/test/accesstoken", httpserver.Handle(testController.signAccessToken)) + r.POST("/test/refreshtoken", httpserver.Handle(testController.signRefreshToken)) - group.GET("/externalidp/jwks.json", httpserver.Handle(testController.externalIdPJWKS)) - group.POST("/externalidp/sign", httpserver.Handle(testController.externalIdPSignToken)) + r.GET("/externalidp/jwks.json", httpserver.Handle(testController.externalIdPJWKS)) + r.POST("/externalidp/sign", httpserver.Handle(testController.externalIdPSignToken)) } type TestController struct { diff --git a/backend/internal/controller/oidc_controller.go b/backend/internal/controller/oidc_controller.go index fbb2a177..c1b348c3 100644 --- a/backend/internal/controller/oidc_controller.go +++ b/backend/internal/controller/oidc_controller.go @@ -10,6 +10,7 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/appconfig" "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/dto" "github.com/pocket-id/pocket-id/backend/internal/httpserver" "github.com/pocket-id/pocket-id/backend/internal/middleware" @@ -21,38 +22,38 @@ import ( // @Summary OIDC controller // @Description Initializes all OIDC-related API endpoints for authentication and client management // @Tags OIDC -func NewOidcController(group *gin.RouterGroup, authMiddleware *middleware.AuthMiddleware, fileSizeLimitMiddleware *middleware.FileSizeLimitMiddleware, oidcService *service.OidcService, appConfigService appconfig.AppConfigResolver) { +func NewOidcController(r *authz.Router, fileSizeLimitMiddleware *middleware.FileSizeLimitMiddleware, oidcService *service.OidcService, appConfigService appconfig.AppConfigResolver) { oc := &OidcController{ oidcService: oidcService, appConfigService: appConfigService, } - group.GET("/oidc/clients", authMiddleware.Add(), httpserver.Handle(oc.listClientsHandler)) - group.POST("/oidc/clients", authMiddleware.Add(), httpserver.Handle(oc.createClientHandler)) - group.GET("/oidc/clients/:id", authMiddleware.Add(), httpserver.Handle(oc.getClientHandler)) - group.GET("/oidc/clients/:id/meta", httpserver.Handle(oc.getClientMetaDataHandler)) - group.PUT("/oidc/clients/:id", authMiddleware.Add(), httpserver.Handle(oc.updateClientHandler)) - group.POST("/oidc/clients/:id/refresh", authMiddleware.Add(), httpserver.Handle(oc.refreshClientMetadataHandler)) - group.DELETE("/oidc/clients/:id", authMiddleware.Add(), httpserver.Handle(oc.deleteClientHandler)) + r.GET("/oidc/clients", authz.OidcClientsRead, httpserver.Handle(oc.listClientsHandler)) + r.POST("/oidc/clients", authz.OidcClientsWrite, httpserver.Handle(oc.createClientHandler)) + r.GET("/oidc/clients/:id", authz.OidcClientsRead, httpserver.Handle(oc.getClientHandler)) + r.Public().GET("/oidc/clients/:id/meta", httpserver.Handle(oc.getClientMetaDataHandler)) + r.PUT("/oidc/clients/:id", authz.OidcClientsWrite, httpserver.Handle(oc.updateClientHandler)) + r.POST("/oidc/clients/:id/refresh", authz.OidcClientsWrite, httpserver.Handle(oc.refreshClientMetadataHandler)) + r.DELETE("/oidc/clients/:id", authz.OidcClientsWrite, httpserver.Handle(oc.deleteClientHandler)) - group.PUT("/oidc/clients/:id/allowed-user-groups", authMiddleware.Add(), httpserver.Handle(oc.updateAllowedUserGroupsHandler)) - group.GET("/oidc/clients/:id/secrets", authMiddleware.Add(), httpserver.Handle(oc.listClientSecretsHandler)) - group.POST("/oidc/clients/:id/secrets", authMiddleware.Add(), httpserver.Handle(oc.createClientSecretHandler)) - group.DELETE("/oidc/clients/:id/secrets/:secretId", authMiddleware.Add(), httpserver.Handle(oc.deleteClientSecretHandler)) + r.PUT("/oidc/clients/:id/allowed-user-groups", authz.OidcClientsWrite, httpserver.Handle(oc.updateAllowedUserGroupsHandler)) + r.GET("/oidc/clients/:id/secrets", authz.OidcClientsRead, httpserver.Handle(oc.listClientSecretsHandler)) + r.POST("/oidc/clients/:id/secrets", authz.OidcClientsWrite, httpserver.Handle(oc.createClientSecretHandler)) + r.DELETE("/oidc/clients/:id/secrets/:secretId", authz.OidcClientsWrite, httpserver.Handle(oc.deleteClientSecretHandler)) - group.GET("/oidc/clients/:id/logo", httpserver.Handle(oc.getClientLogoHandler)) - group.DELETE("/oidc/clients/:id/logo", authMiddleware.Add(), httpserver.Handle(oc.deleteClientLogoHandler)) - group.POST("/oidc/clients/:id/logo", authMiddleware.Add(), fileSizeLimitMiddleware.Add(2<<20), httpserver.Handle(oc.updateClientLogoHandler)) + r.Public().GET("/oidc/clients/:id/logo", httpserver.Handle(oc.getClientLogoHandler)) + r.DELETE("/oidc/clients/:id/logo", authz.OidcClientsWrite, httpserver.Handle(oc.deleteClientLogoHandler)) + r.POST("/oidc/clients/:id/logo", authz.OidcClientsWrite, fileSizeLimitMiddleware.Add(2<<20), httpserver.Handle(oc.updateClientLogoHandler)) - group.GET("/oidc/clients/:id/preview/:userId", authMiddleware.Add(), httpserver.Handle(oc.getClientPreviewHandler)) + // The preview renders a user's claims, so it is guarded by the user scope rather than the client scope + r.GET("/oidc/clients/:id/preview/:userId", authz.UsersRead, httpserver.Handle(oc.getClientPreviewHandler)) - group.GET("/oidc/users/me/authorized-clients", authMiddleware.WithAdminNotRequired().Add(), httpserver.Handle(oc.listOwnAuthorizedClientsHandler)) - group.GET("/oidc/users/:id/authorized-clients", authMiddleware.Add(), httpserver.Handle(oc.listAuthorizedClientsHandler)) + r.GET("/oidc/users/me/authorized-clients", authz.AccountApps, httpserver.Handle(oc.listOwnAuthorizedClientsHandler)) + r.GET("/oidc/users/:id/authorized-clients", authz.UsersRead, httpserver.Handle(oc.listAuthorizedClientsHandler)) - group.DELETE("/oidc/users/me/authorized-clients/:clientId", authMiddleware.WithAdminNotRequired().Add(), httpserver.Handle(oc.revokeOwnClientAuthorizationHandler)) - - group.GET("/oidc/users/me/clients", authMiddleware.WithAdminNotRequired().Add(), httpserver.Handle(oc.listOwnAccessibleClientsHandler)) + r.DELETE("/oidc/users/me/authorized-clients/:clientId", authz.AccountApps, httpserver.Handle(oc.revokeOwnClientAuthorizationHandler)) + r.GET("/oidc/users/me/clients", authz.AccountApps, httpserver.Handle(oc.listOwnAccessibleClientsHandler)) } type OidcController struct { @@ -173,7 +174,7 @@ func (oc *OidcController) createClientHandler(c *gin.Context) error { return err } - client, createdSecret, err := oc.oidcService.CreateClient(c.Request.Context(), input, c.GetString("userID"), config.AutoCreateOIDCClientSecret.IsTrue()) + client, createdSecret, err := oc.oidcService.CreateClient(c.Request.Context(), input, authz.PrincipalFrom(c).UserID, config.AutoCreateOIDCClientSecret.IsTrue()) if err != nil { return err } @@ -480,7 +481,7 @@ func (oc *OidcController) updateAllowedUserGroupsHandler(c *gin.Context) error { // @Failure default {object} dto.ErrorDto "Error" // @Router /api/oidc/users/me/authorized-clients [get] func (oc *OidcController) listOwnAuthorizedClientsHandler(c *gin.Context) error { - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID return oc.listAuthorizedClients(c, userID) } @@ -536,7 +537,7 @@ func (oc *OidcController) listAuthorizedClients(c *gin.Context, userID string) e func (oc *OidcController) revokeOwnClientAuthorizationHandler(c *gin.Context) error { clientID := c.Param("clientId") - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID err := oc.oidcService.RevokeAuthorizedClient(c.Request.Context(), userID, clientID) if err != nil { @@ -564,7 +565,7 @@ func (oc *OidcController) listOwnAccessibleClientsHandler(c *gin.Context) error searchTerm := c.Query("search") listRequestOptions := utils.ParseListRequestOptions(c) - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID clients, pagination, err := oc.oidcService.ListAccessibleOidcClients(c.Request.Context(), userID, searchTerm, listRequestOptions) if err != nil { @@ -612,7 +613,7 @@ func (oc *OidcController) getClientPreviewHandler(c *gin.Context) error { clientID, userID, strings.Split(scopes, " "), - c.GetString("authenticationMethod")) + authz.PrincipalFrom(c).AuthenticationMethod) if err != nil { return err diff --git a/backend/internal/controller/upload_size_limit_test.go b/backend/internal/controller/upload_size_limit_test.go index 2378847f..1a29d1a5 100644 --- a/backend/internal/controller/upload_size_limit_test.go +++ b/backend/internal/controller/upload_size_limit_test.go @@ -43,7 +43,7 @@ func TestImageUploadRoutesLimitRequestSize(t *testing.T) { apiKeyModule, err := apikey.New(t.Context(), apikey.Dependencies{DB: db, CleanupDisabled: true}) require.NoError(t, err) - authMiddleware := middleware.NewAuthMiddleware(apiKeyModule, userService, jwtService) + auth := middleware.NewAuthorization(apiKeyModule, userService, jwtService) fileSizeLimitMiddleware := middleware.NewFileSizeLimitMiddleware() user := model.User{Username: "upload-admin", IsAdmin: true} @@ -54,9 +54,9 @@ func TestImageUploadRoutesLimitRequestSize(t *testing.T) { router := gin.New() router.Use(middleware.NewErrorHandlerMiddleware().Add()) - apiGroup := router.Group("/api") - NewUserController(apiGroup, authMiddleware, fileSizeLimitMiddleware, nil, userService, nil, func(c *gin.Context) { c.Next() }) - NewAppImagesController(apiGroup, authMiddleware, fileSizeLimitMiddleware, nil) + apiRouter := auth.Router(router.Group("/api")) + NewUserController(apiRouter, fileSizeLimitMiddleware, nil, userService, nil, func(c *gin.Context) { c.Next() }) + NewAppImagesController(apiRouter, fileSizeLimitMiddleware, nil) routes := []string{ "/api/users/user-id/profile-picture", diff --git a/backend/internal/controller/user_controller.go b/backend/internal/controller/user_controller.go index 03b426ce..dae21f5a 100644 --- a/backend/internal/controller/user_controller.go +++ b/backend/internal/controller/user_controller.go @@ -6,6 +6,7 @@ import ( "time" "github.com/pocket-id/pocket-id/backend/internal/appconfig" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/gin-gonic/gin" "github.com/pocket-id/pocket-id/backend/internal/dto" @@ -20,34 +21,33 @@ import ( // @Summary User management controller // @Description Initializes all user-related API endpoints // @Tags Users -func NewUserController(group *gin.RouterGroup, authMiddleware *middleware.AuthMiddleware, fileSizeLimitMiddleware *middleware.FileSizeLimitMiddleware, appConfigService *appconfig.AppConfigService, userService *service.UserService, webAuthnService *webauthn.Module, updateOwnAccountRateLimit gin.HandlerFunc) { +func NewUserController(r *authz.Router, fileSizeLimitMiddleware *middleware.FileSizeLimitMiddleware, appConfigService *appconfig.AppConfigService, userService *service.UserService, webAuthnService *webauthn.Module, updateOwnAccountRateLimit gin.HandlerFunc) { uc := UserController{ appConfigService: appConfigService, userService: userService, webAuthnService: webAuthnService, } - group.GET("/users", authMiddleware.Add(), httpserver.Handle(uc.listUsersHandler)) - group.GET("/users/me", authMiddleware.WithAdminNotRequired().Add(), httpserver.Handle(uc.getCurrentUserHandler)) - group.GET("/users/:id", authMiddleware.Add(), httpserver.Handle(uc.getUserHandler)) - group.POST("/users", authMiddleware.Add(), httpserver.Handle(uc.createUserHandler)) - group.PUT("/users/:id", authMiddleware.Add(), httpserver.Handle(uc.updateUserHandler)) - group.GET("/users/:id/groups", authMiddleware.Add(), httpserver.Handle(uc.getUserGroupsHandler)) - group.GET("/users/:id/webauthn-credentials", authMiddleware.Add(), httpserver.Handle(uc.listUserWebauthnCredentialsHandler)) - // Updating the own account reports whether an email or username is already taken, so it is rate limited to slow down probing for existing users - group.PUT("/users/me", authMiddleware.WithAdminNotRequired().Add(), updateOwnAccountRateLimit, httpserver.Handle(uc.updateCurrentUserHandler)) - group.DELETE("/users/:id", authMiddleware.Add(), httpserver.Handle(uc.deleteUserHandler)) - group.DELETE("/users/:id/webauthn-credentials/:credentialId", authMiddleware.Add(), httpserver.Handle(uc.deleteUserWebauthnCredentialHandler)) + r.GET("/users", authz.UsersRead, httpserver.Handle(uc.listUsersHandler)) + r.GET("/users/me", authz.AccountRead, httpserver.Handle(uc.getCurrentUserHandler)) + r.GET("/users/:id", authz.UsersRead, httpserver.Handle(uc.getUserHandler)) + r.POST("/users", authz.UsersWrite, httpserver.Handle(uc.createUserHandler)) + r.PUT("/users/:id", authz.UsersWrite, httpserver.Handle(uc.updateUserHandler)) + r.GET("/users/:id/groups", authz.UsersRead, httpserver.Handle(uc.getUserGroupsHandler)) + r.GET("/users/:id/webauthn-credentials", authz.UsersRead, httpserver.Handle(uc.listUserWebauthnCredentialsHandler)) + r.PUT("/users/me", authz.AccountWrite, updateOwnAccountRateLimit, httpserver.Handle(uc.updateCurrentUserHandler)) + r.DELETE("/users/:id", authz.UsersWrite, httpserver.Handle(uc.deleteUserHandler)) + r.DELETE("/users/:id/webauthn-credentials/:credentialId", authz.UsersWrite, httpserver.Handle(uc.deleteUserWebauthnCredentialHandler)) - group.PUT("/users/:id/user-groups", authMiddleware.Add(), httpserver.Handle(uc.updateUserGroups)) + r.PUT("/users/:id/user-groups", authz.UsersWrite, httpserver.Handle(uc.updateUserGroups)) - group.GET("/users/:id/profile-picture.png", httpserver.Handle(uc.getUserProfilePictureHandler)) + r.Public().GET("/users/:id/profile-picture.png", httpserver.Handle(uc.getUserProfilePictureHandler)) - group.PUT("/users/:id/profile-picture", authMiddleware.Add(), fileSizeLimitMiddleware.Add(10<<20), httpserver.Handle(uc.updateUserProfilePictureHandler)) - group.PUT("/users/me/profile-picture", authMiddleware.WithAdminNotRequired().Add(), fileSizeLimitMiddleware.Add(10<<20), httpserver.Handle(uc.updateCurrentUserProfilePictureHandler)) + r.PUT("/users/:id/profile-picture", authz.UsersWrite, fileSizeLimitMiddleware.Add(10<<20), httpserver.Handle(uc.updateUserProfilePictureHandler)) + r.PUT("/users/me/profile-picture", authz.AccountWrite, fileSizeLimitMiddleware.Add(10<<20), httpserver.Handle(uc.updateCurrentUserProfilePictureHandler)) - group.DELETE("/users/:id/profile-picture", authMiddleware.Add(), httpserver.Handle(uc.resetUserProfilePictureHandler)) - group.DELETE("/users/me/profile-picture", authMiddleware.WithAdminNotRequired().Add(), httpserver.Handle(uc.resetCurrentUserProfilePictureHandler)) + r.DELETE("/users/:id/profile-picture", authz.UsersWrite, httpserver.Handle(uc.resetUserProfilePictureHandler)) + r.DELETE("/users/me/profile-picture", authz.AccountWrite, httpserver.Handle(uc.resetCurrentUserProfilePictureHandler)) } type UserController struct { @@ -173,7 +173,7 @@ func (uc *UserController) getUserHandler(c *gin.Context) error { // @Failure default {object} dto.ErrorDto "Error" // @Router /api/users/me [get] func (uc *UserController) getCurrentUserHandler(c *gin.Context) error { - user, err := uc.userService.GetUser(c.Request.Context(), c.GetString("userID")) + user, err := uc.userService.GetUser(c.Request.Context(), authz.PrincipalFrom(c).UserID) if err != nil { return err } @@ -225,7 +225,7 @@ func (uc *UserController) deleteUserWebauthnCredentialHandler(c *gin.Context) er c.Param("credentialId"), c.ClientIP(), c.Request.UserAgent(), - c.GetString("userID"), + authz.PrincipalFrom(c).UserID, ) if err != nil { return err @@ -361,7 +361,7 @@ func (uc *UserController) updateUserProfilePictureHandler(c *gin.Context) error // @Failure default {object} dto.ErrorDto "Error" // @Router /api/users/me/profile-picture [put] func (uc *UserController) updateCurrentUserProfilePictureHandler(c *gin.Context) error { - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID fileHeader, err := httpserver.FormFile(c, "file") if err != nil { return err @@ -423,7 +423,7 @@ func (uc *UserController) updateUser(c *gin.Context, updateOwnUser bool) error { var userID string if updateOwnUser { - userID = c.GetString("userID") + userID = authz.PrincipalFrom(c).UserID } else { userID = c.Param("id") } @@ -471,7 +471,7 @@ func (uc *UserController) resetUserProfilePictureHandler(c *gin.Context) error { // @Failure default {object} dto.ErrorDto "Error" // @Router /api/users/me/profile-picture [delete] func (uc *UserController) resetCurrentUserProfilePictureHandler(c *gin.Context) error { - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID if err := uc.userService.ResetProfilePicture(c.Request.Context(), userID); err != nil { return err diff --git a/backend/internal/controller/user_group_controller.go b/backend/internal/controller/user_group_controller.go index 85f68abe..8c44b867 100644 --- a/backend/internal/controller/user_group_controller.go +++ b/backend/internal/controller/user_group_controller.go @@ -6,9 +6,9 @@ import ( "github.com/gin-gonic/gin" "github.com/pocket-id/pocket-id/backend/internal/appconfig" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/httpserver" - "github.com/pocket-id/pocket-id/backend/internal/middleware" "github.com/pocket-id/pocket-id/backend/internal/service" "github.com/pocket-id/pocket-id/backend/internal/utils" ) @@ -17,23 +17,20 @@ import ( // @Summary User group management controller // @Description Initializes all user group-related API endpoints // @Tags User Groups -func NewUserGroupController(group *gin.RouterGroup, authMiddleware *middleware.AuthMiddleware, appConfigService *appconfig.AppConfigService, userGroupService *service.UserGroupService) { +func NewUserGroupController(r *authz.Router, appConfigService *appconfig.AppConfigService, userGroupService *service.UserGroupService) { ugc := UserGroupController{ appConfigService: appConfigService, UserGroupService: userGroupService, } - userGroupsGroup := group.Group("/user-groups") - userGroupsGroup.Use(authMiddleware.Add()) - { - userGroupsGroup.GET("", httpserver.Handle(ugc.list)) - userGroupsGroup.GET("/:id", httpserver.Handle(ugc.get)) - userGroupsGroup.POST("", httpserver.Handle(ugc.create)) - userGroupsGroup.PUT("/:id", httpserver.Handle(ugc.update)) - userGroupsGroup.DELETE("/:id", httpserver.Handle(ugc.delete)) - userGroupsGroup.PUT("/:id/users", httpserver.Handle(ugc.updateUsers)) - userGroupsGroup.PUT("/:id/allowed-oidc-clients", httpserver.Handle(ugc.updateAllowedOidcClients)) - } + userGroupsGroup := r.Group("/user-groups") + userGroupsGroup.GET("", authz.GroupsRead, httpserver.Handle(ugc.list)) + userGroupsGroup.GET("/:id", authz.GroupsRead, httpserver.Handle(ugc.get)) + userGroupsGroup.POST("", authz.GroupsWrite, httpserver.Handle(ugc.create)) + userGroupsGroup.PUT("/:id", authz.GroupsWrite, httpserver.Handle(ugc.update)) + userGroupsGroup.DELETE("/:id", authz.GroupsWrite, httpserver.Handle(ugc.delete)) + userGroupsGroup.PUT("/:id/users", authz.GroupsWrite, httpserver.Handle(ugc.updateUsers)) + userGroupsGroup.PUT("/:id/allowed-oidc-clients", authz.GroupsWrite, httpserver.Handle(ugc.updateAllowedOidcClients)) } type UserGroupController struct { diff --git a/backend/internal/devicelogin/handler.go b/backend/internal/devicelogin/handler.go index 4350139b..a55200f2 100644 --- a/backend/internal/devicelogin/handler.go +++ b/backend/internal/devicelogin/handler.go @@ -9,6 +9,7 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/appconfig" "github.com/pocket-id/pocket-id/backend/internal/auditlogs" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/httpserver" "github.com/pocket-id/pocket-id/backend/internal/utils/cookie" @@ -139,7 +140,7 @@ func (h *handler) decideRequest(c *gin.Context) error { } reauthenticationToken, _ := c.Cookie(cookie.ReauthenticationTokenCookieName) - err = h.service.Decide(c.Request.Context(), input.Code, input.Decision, c.GetString("userID"), reauthenticationToken) + err = h.service.Decide(c.Request.Context(), input.Code, input.Decision, authz.PrincipalFrom(c).UserID, reauthenticationToken) if err != nil { return err } diff --git a/backend/internal/devicelogin/module.go b/backend/internal/devicelogin/module.go index c322c83b..bc95ccb4 100644 --- a/backend/internal/devicelogin/module.go +++ b/backend/internal/devicelogin/module.go @@ -11,6 +11,7 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/appconfig" "github.com/pocket-id/pocket-id/backend/internal/auditlogs" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/httpserver" "github.com/pocket-id/pocket-id/backend/internal/iplocation" "github.com/pocket-id/pocket-id/backend/internal/model" @@ -68,9 +69,9 @@ func New(deps Dependencies) (*Module, error) { } // RegisterRoutes mounts the public exchange and authenticated verification endpoints -func (m *Module) RegisterRoutes(apiGroup *gin.RouterGroup, browserAuth, createRateLimit, exchangeRateLimit, verificationRateLimit gin.HandlerFunc) { - apiGroup.POST("/device-login/requests", createRateLimit, httpserver.Handle(m.handler.createRequest)) - apiGroup.POST("/device-login/requests/:id/exchange", exchangeRateLimit, httpserver.Handle(m.handler.exchangeRequest)) - apiGroup.POST("/device-login/verification", verificationRateLimit, browserAuth, httpserver.Handle(m.handler.inspectRequest)) - apiGroup.POST("/device-login/verification/decision", verificationRateLimit, browserAuth, httpserver.Handle(m.handler.decideRequest)) +func (m *Module) RegisterRoutes(r *authz.Router, createRateLimit, exchangeRateLimit, verificationRateLimit gin.HandlerFunc) { + r.Public().POST("/device-login/requests", createRateLimit, httpserver.Handle(m.handler.createRequest)) + r.Public().POST("/device-login/requests/:id/exchange", exchangeRateLimit, httpserver.Handle(m.handler.exchangeRequest)) + r.POST("/device-login/verification", authz.AccountSession, verificationRateLimit, httpserver.Handle(m.handler.inspectRequest)) + r.POST("/device-login/verification/decision", authz.AccountSession, verificationRateLimit, httpserver.Handle(m.handler.decideRequest)) } diff --git a/backend/internal/emailverification/handler.go b/backend/internal/emailverification/handler.go index d30de55e..c1406312 100644 --- a/backend/internal/emailverification/handler.go +++ b/backend/internal/emailverification/handler.go @@ -7,6 +7,7 @@ import ( "github.com/gin-gonic/gin" "github.com/pocket-id/pocket-id/backend/internal/appconfig" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/httpserver" ) @@ -33,7 +34,7 @@ func (h *handler) send(c *gin.Context) error { return fmt.Errorf("error loading app configuration: %w", err) } - err = h.service.Send(c.Request.Context(), dbConfig, c.GetString("userID")) + err = h.service.Send(c.Request.Context(), dbConfig, authz.PrincipalFrom(c).UserID) if err != nil { return err } @@ -55,7 +56,7 @@ func (h *handler) verify(c *gin.Context) error { return err } - err := h.service.Verify(c.Request.Context(), c.GetString("userID"), input.Token) + err := h.service.Verify(c.Request.Context(), authz.PrincipalFrom(c).UserID, input.Token) if err != nil { return err } diff --git a/backend/internal/emailverification/module.go b/backend/internal/emailverification/module.go index 40289426..74caf2d5 100644 --- a/backend/internal/emailverification/module.go +++ b/backend/internal/emailverification/module.go @@ -8,6 +8,7 @@ import ( "gorm.io/gorm" "github.com/pocket-id/pocket-id/backend/internal/appconfig" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/httpserver" ) @@ -40,7 +41,7 @@ func New(deps Dependencies) (*Module, error) { } // RegisterRoutes mounts the email verification endpoints -func (m *Module) RegisterRoutes(apiGroup *gin.RouterGroup, userAuth, sendRateLimit, verifyRateLimit gin.HandlerFunc) { - apiGroup.POST("/users/me/send-email-verification", sendRateLimit, userAuth, httpserver.Handle(m.handler.send)) - apiGroup.POST("/users/me/verify-email", verifyRateLimit, userAuth, httpserver.Handle(m.handler.verify)) +func (m *Module) RegisterRoutes(r *authz.Router, sendRateLimit, verifyRateLimit gin.HandlerFunc) { + r.POST("/users/me/send-email-verification", authz.AccountWrite, sendRateLimit, httpserver.Handle(m.handler.send)) + r.POST("/users/me/verify-email", authz.AccountWrite, verifyRateLimit, httpserver.Handle(m.handler.verify)) } diff --git a/backend/internal/environment/module.go b/backend/internal/environment/module.go index 363d0b6c..89724811 100644 --- a/backend/internal/environment/module.go +++ b/backend/internal/environment/module.go @@ -3,8 +3,7 @@ package environment import ( "net/http" - "github.com/gin-gonic/gin" - + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/httpserver" ) @@ -31,8 +30,8 @@ func New(deps Dependencies) *Module { } // RegisterRoutes mounts the environment endpoints -func (m *Module) RegisterRoutes(apiGroup *gin.RouterGroup, auth gin.HandlerFunc) { - apiGroup.GET("/version/latest", httpserver.Handle(m.handler.getLatestVersion)) - apiGroup.GET("/version/current", auth, httpserver.Handle(m.handler.getCurrentVersion)) - apiGroup.GET("/storage/sqlite-warning", auth, httpserver.Handle(m.handler.getSqliteStorageWarning)) +func (m *Module) RegisterRoutes(r *authz.Router) { + r.Public().GET("/version/latest", httpserver.Handle(m.handler.getLatestVersion)) + r.GET("/version/current", authz.AccountRead, httpserver.Handle(m.handler.getCurrentVersion)) + r.GET("/storage/sqlite-warning", authz.AccountRead, httpserver.Handle(m.handler.getSqliteStorageWarning)) } diff --git a/backend/internal/ldapsync/module.go b/backend/internal/ldapsync/module.go index 422f7b21..bdfd549d 100644 --- a/backend/internal/ldapsync/module.go +++ b/backend/internal/ldapsync/module.go @@ -6,11 +6,11 @@ import ( "io" "net/http" - "github.com/gin-gonic/gin" francishost "github.com/italypaleale/francis/host" "gorm.io/gorm" "github.com/pocket-id/pocket-id/backend/internal/appconfig" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/httpserver" "github.com/pocket-id/pocket-id/backend/internal/model" @@ -90,9 +90,8 @@ func New(deps Dependencies) (*Module, error) { } // RegisterRoutes mounts the manual LDAP synchronization endpoint -// auth guards it, as it's an admin-only operation -func (m *Module) RegisterRoutes(apiGroup *gin.RouterGroup, auth gin.HandlerFunc) { - apiGroup.POST("/application-configuration/sync-ldap", auth, httpserver.Handle(m.handler.syncLdap)) +func (m *Module) RegisterRoutes(r *authz.Router) { + r.POST("/application-configuration/sync-ldap", authz.ConfigWrite, httpserver.Handle(m.handler.syncLdap)) } // SyncAll runs a full LDAP synchronization with the provided application configuration diff --git a/backend/internal/logopreset/module.go b/backend/internal/logopreset/module.go index 9a7e7fdb..85d6ae76 100644 --- a/backend/internal/logopreset/module.go +++ b/backend/internal/logopreset/module.go @@ -3,8 +3,7 @@ package logopreset import ( "net/http" - "github.com/gin-gonic/gin" - + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/httpserver" ) @@ -31,6 +30,6 @@ func New(deps Dependencies) *Module { } // RegisterRoutes mounts the logo preset endpoints -func (m *Module) RegisterRoutes(apiGroup *gin.RouterGroup, auth gin.HandlerFunc) { - apiGroup.GET("/oidc/logo-presets", auth, httpserver.Handle(m.handler.search)) +func (m *Module) RegisterRoutes(r *authz.Router) { + r.GET("/oidc/logo-presets", authz.OidcClientsRead, httpserver.Handle(m.handler.search)) } diff --git a/backend/internal/middleware/api_key_auth.go b/backend/internal/middleware/api_key_auth.go deleted file mode 100644 index 1ac11919..00000000 --- a/backend/internal/middleware/api_key_auth.go +++ /dev/null @@ -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 -} diff --git a/backend/internal/middleware/auth_middleware.go b/backend/internal/middleware/auth_middleware.go deleted file mode 100644 index 8957af47..00000000 --- a/backend/internal/middleware/auth_middleware.go +++ /dev/null @@ -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) - } -} diff --git a/backend/internal/middleware/auth_middleware_test.go b/backend/internal/middleware/auth_middleware_test.go deleted file mode 100644 index 16b34035..00000000 --- a/backend/internal/middleware/auth_middleware_test.go +++ /dev/null @@ -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 -} diff --git a/backend/internal/middleware/authenticator.go b/backend/internal/middleware/authenticator.go new file mode 100644 index 00000000..74683096 --- /dev/null +++ b/backend/internal/middleware/authenticator.go @@ -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 +} diff --git a/backend/internal/middleware/authenticator_test.go b/backend/internal/middleware/authenticator_test.go new file mode 100644 index 00000000..919739a0 --- /dev/null +++ b/backend/internal/middleware/authenticator_test.go @@ -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 +} diff --git a/backend/internal/middleware/jwt_auth.go b/backend/internal/middleware/jwt_auth.go deleted file mode 100644 index 25b3640c..00000000 --- a/backend/internal/middleware/jwt_auth.go +++ /dev/null @@ -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 -} diff --git a/backend/internal/oidc/authorization_handler.go b/backend/internal/oidc/authorization_handler.go index 089ccb39..469c2989 100644 --- a/backend/internal/oidc/authorization_handler.go +++ b/backend/internal/oidc/authorization_handler.go @@ -7,10 +7,10 @@ import ( "net/http" "net/url" "strings" - "time" "github.com/gin-gonic/gin" "github.com/ory/fosite" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/httpserver" "github.com/pocket-id/pocket-id/backend/internal/utils" "github.com/pocket-id/pocket-id/backend/internal/utils/cookie" @@ -35,10 +35,7 @@ func newAuthorizationHandler( func (h *authorizationHandler) authorize(c *gin.Context) { ctx := c.Request.Context() - userID := c.GetString("userID") - authenticationMethod := c.GetString("authenticationMethod") - authenticationTime, _ := c.Get("authenticationTime") - typedAuthenticationTime, _ := authenticationTime.(time.Time) + principal := authz.PrincipalFrom(c) reauthenticationToken, _ := c.Cookie(cookie.ReauthenticationTokenCookieName) // A request that resumes an interaction only carries the interaction ID; the original @@ -73,9 +70,9 @@ func (h *authorizationHandler) authorize(c *gin.Context) { } authorization, err := h.authorizationService.authorize(ctx, authorizeInput{ - userID: userID, - authenticationMethod: authenticationMethod, - authenticationTime: typedAuthenticationTime, + userID: principal.UserID, + authenticationMethod: principal.AuthenticationMethod, + authenticationTime: principal.AuthenticationTime, requester: ar, hasPushedAuthorizationRequest: hasPushedAuthorizationRequest, reauthenticationToken: reauthenticationToken, @@ -121,8 +118,7 @@ func (h *authorizationHandler) getInteractionSession(c *gin.Context) { func (h *authorizationHandler) completeInteraction(c *gin.Context) { interactionID := c.Param("id") - authenticationTime, _ := c.Get("authenticationTime") - typedAuthenticationTime, _ := authenticationTime.(time.Time) + principal := authz.PrincipalFrom(c) var request completeInteractionRequest if err := httpserver.BindJSON(c, &request); err != nil { @@ -131,7 +127,7 @@ func (h *authorizationHandler) completeInteraction(c *gin.Context) { } reauthenticationToken, _ := c.Cookie(cookie.ReauthenticationTokenCookieName) - response, err := h.authorizationService.completeInteractionStep(c.Request.Context(), interactionID, c.GetString("userID"), request.Step, reauthenticationToken, typedAuthenticationTime, requestMetaFromGin(c)) + response, err := h.authorizationService.completeInteractionStep(c.Request.Context(), interactionID, principal.UserID, request.Step, reauthenticationToken, principal.AuthenticationTime, requestMetaFromGin(c)) if err != nil { _ = c.Error(err) return diff --git a/backend/internal/oidc/authorization_handler_test.go b/backend/internal/oidc/authorization_handler_test.go index 92bc854b..3fed7dbc 100644 --- a/backend/internal/oidc/authorization_handler_test.go +++ b/backend/internal/oidc/authorization_handler_test.go @@ -15,6 +15,7 @@ import ( "time" "github.com/gin-gonic/gin" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/model" datatype "github.com/pocket-id/pocket-id/backend/internal/model/types" testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing" @@ -167,8 +168,12 @@ func testAuthorizationHandlerPAR(t *testing.T, clientType string, tt authorizati rec := httptest.NewRecorder() router := gin.New() router.Handle(tt.method, "/authorize", func(c *gin.Context) { - c.Set("userID", userID) - c.Set("authenticationTime", time.Now().UTC().Add(-time.Minute)) + authz.SetPrincipal(c, &authz.Principal{ + Kind: authz.KindSession, + UserID: userID, + Scopes: authz.UserScopes(false, authz.KindSession), + AuthenticationTime: time.Now().UTC().Add(-time.Minute), + }) handler.authorize(c) }) router.ServeHTTP(rec, req) diff --git a/backend/internal/oidc/device_handler.go b/backend/internal/oidc/device_handler.go index 9c56ba05..fae00add 100644 --- a/backend/internal/oidc/device_handler.go +++ b/backend/internal/oidc/device_handler.go @@ -4,11 +4,11 @@ import ( "errors" "log/slog" "net/http" - "time" "github.com/gin-gonic/gin" "github.com/ory/fosite" "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/utils/cookie" ) @@ -39,8 +39,7 @@ func (h *deviceHandler) authorizeDevice(c *gin.Context) { } func (h *deviceHandler) verifyDeviceCode(c *gin.Context) { - authenticationTime, _ := c.Get("authenticationTime") - typedAuthenticationTime, _ := authenticationTime.(time.Time) + principal := authz.PrincipalFrom(c) reauthenticationToken, _ := c.Cookie(cookie.ReauthenticationTokenCookieName) userCode := c.Query("code") @@ -52,9 +51,9 @@ func (h *deviceHandler) verifyDeviceCode(c *gin.Context) { err := h.deviceService.acceptDeviceCode( c.Request.Context(), userCode, - c.GetString("userID"), - c.GetString("authenticationMethod"), - typedAuthenticationTime, + principal.UserID, + principal.AuthenticationMethod, + principal.AuthenticationTime, reauthenticationToken, requestMetaFromGin(c), ) @@ -77,7 +76,7 @@ func (h *deviceHandler) deviceCodeInfo(c *gin.Context) { return } - deviceCodeInfo, err := h.deviceService.getDeviceCodeInfo(c.Request.Context(), userCode, c.GetString("userID")) + deviceCodeInfo, err := h.deviceService.getDeviceCodeInfo(c.Request.Context(), userCode, authz.PrincipalFrom(c).UserID) if err != nil { _ = c.Error(err) return diff --git a/backend/internal/oidc/end_session_handler.go b/backend/internal/oidc/end_session_handler.go index 185378a7..3097d155 100644 --- a/backend/internal/oidc/end_session_handler.go +++ b/backend/internal/oidc/end_session_handler.go @@ -5,6 +5,7 @@ import ( "net/http" "github.com/gin-gonic/gin" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/utils/cookie" ) @@ -28,7 +29,7 @@ func (h *endSessionHandler) endSession(c *gin.Context) { return } - callbackURL, err := h.endSessionService.endSession(c.Request.Context(), input, c.GetString("userID")) + callbackURL, err := h.endSessionService.endSession(c.Request.Context(), input, authz.PrincipalFrom(c).UserID) if err != nil { slog.WarnContext(c.Request.Context(), "Error getting logout callback URL, the user has to confirm the logout manually", "error", err) c.Redirect(http.StatusFound, h.baseURL+"/logout") diff --git a/backend/internal/oidc/module.go b/backend/internal/oidc/module.go index 8bf890ec..eca32843 100644 --- a/backend/internal/oidc/module.go +++ b/backend/internal/oidc/module.go @@ -7,13 +7,13 @@ import ( "net/http" "time" - "github.com/gin-gonic/gin" francishost "github.com/italypaleale/francis/host" "github.com/lestrrat-go/jwx/v4/jwa" "go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp" "gorm.io/gorm" "github.com/pocket-id/pocket-id/backend/internal/auditlogs" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/model" ) @@ -143,26 +143,26 @@ func (m *Module) RefreshClientMetadata(ctx context.Context, clientID string) (mo return m.cimdResolver.RefreshMetadataClient(ctx, clientID) } -func (m *Module) RegisterRoutes(rootGroup *gin.RouterGroup, apiGroup *gin.RouterGroup, optionalBrowserAuth gin.HandlerFunc, browserAuth gin.HandlerFunc) { - rootGroup.GET("/authorize", optionalBrowserAuth, m.authorizationHandler.authorize) - rootGroup.POST("/authorize", optionalBrowserAuth, m.authorizationHandler.authorize) +func (m *Module) RegisterRoutes(root, api *authz.Router) { + root.Optional().GET("/authorize", authz.AccountSession, m.authorizationHandler.authorize) + root.Optional().POST("/authorize", authz.AccountSession, m.authorizationHandler.authorize) - apiGroup.GET("/oidc/interactions/:id", m.authorizationHandler.getInteractionSession) - apiGroup.POST("/oidc/interactions/:id/complete", browserAuth, m.authorizationHandler.completeInteraction) + api.Public().GET("/oidc/interactions/:id", m.authorizationHandler.getInteractionSession) + api.POST("/oidc/interactions/:id/complete", authz.AccountSession, m.authorizationHandler.completeInteraction) - apiGroup.POST("/oidc/par", m.parHandler.pushedAuthorizationRequest) + api.Public().POST("/oidc/par", m.parHandler.pushedAuthorizationRequest) - apiGroup.POST("/oidc/token", m.tokenHandler.token) + api.Public().POST("/oidc/token", m.tokenHandler.token) - apiGroup.GET("/oidc/userinfo", m.userInfoHandler.userInfo) - apiGroup.POST("/oidc/userinfo", m.userInfoHandler.userInfo) + api.Public().GET("/oidc/userinfo", m.userInfoHandler.userInfo) + api.Public().POST("/oidc/userinfo", m.userInfoHandler.userInfo) - apiGroup.POST("/oidc/introspect", m.introspectionHandler.introspectToken) + api.Public().POST("/oidc/introspect", m.introspectionHandler.introspectToken) - apiGroup.GET("/oidc/end-session", optionalBrowserAuth, m.endSessionHandler.endSession) - apiGroup.POST("/oidc/end-session", optionalBrowserAuth, m.endSessionHandler.endSession) + api.Optional().GET("/oidc/end-session", authz.AccountSession, m.endSessionHandler.endSession) + api.Optional().POST("/oidc/end-session", authz.AccountSession, m.endSessionHandler.endSession) - apiGroup.POST("/oidc/device/authorize", m.deviceHandler.authorizeDevice) - apiGroup.POST("/oidc/device/verify", browserAuth, m.deviceHandler.verifyDeviceCode) - apiGroup.GET("/oidc/device/info", browserAuth, m.deviceHandler.deviceCodeInfo) + api.Public().POST("/oidc/device/authorize", m.deviceHandler.authorizeDevice) + api.POST("/oidc/device/verify", authz.AccountSession, m.deviceHandler.verifyDeviceCode) + api.GET("/oidc/device/info", authz.AccountSession, m.deviceHandler.deviceCodeInfo) } diff --git a/backend/internal/onetimeaccess/module.go b/backend/internal/onetimeaccess/module.go index d7c401bf..f94ef327 100644 --- a/backend/internal/onetimeaccess/module.go +++ b/backend/internal/onetimeaccess/module.go @@ -11,6 +11,7 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/appconfig" "github.com/pocket-id/pocket-id/backend/internal/auditlogs" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/httpserver" "github.com/pocket-id/pocket-id/backend/internal/model" ) @@ -66,10 +67,10 @@ func New(deps Dependencies) (*Module, error) { } // RegisterRoutes mounts the one-time access token endpoints -// auth guards the admin routes, while the rate limiters throttle the public exchange and email endpoints -func (m *Module) RegisterRoutes(apiGroup *gin.RouterGroup, auth, exchangeRateLimit, emailRateLimit gin.HandlerFunc) { - apiGroup.POST("/users/:id/one-time-access-token", auth, httpserver.Handle(m.handler.createTokenForUser)) - apiGroup.POST("/users/:id/one-time-access-email", auth, httpserver.Handle(m.handler.requestEmailAsAdmin)) - apiGroup.POST("/one-time-access-token/:token", exchangeRateLimit, httpserver.Handle(m.handler.exchangeToken)) - apiGroup.POST("/one-time-access-email", emailRateLimit, httpserver.Handle(m.handler.requestEmailAsUnauthenticatedUser)) +// The rate limiters throttle the public exchange and email endpoints +func (m *Module) RegisterRoutes(r *authz.Router, exchangeRateLimit, emailRateLimit gin.HandlerFunc) { + r.POST("/users/:id/one-time-access-token", authz.UsersWrite, httpserver.Handle(m.handler.createTokenForUser)) + r.POST("/users/:id/one-time-access-email", authz.UsersWrite, httpserver.Handle(m.handler.requestEmailAsAdmin)) + r.Public().POST("/one-time-access-token/:token", exchangeRateLimit, httpserver.Handle(m.handler.exchangeToken)) + r.Public().POST("/one-time-access-email", emailRateLimit, httpserver.Handle(m.handler.requestEmailAsUnauthenticatedUser)) } diff --git a/backend/internal/scimsync/module.go b/backend/internal/scimsync/module.go index 3304626d..2b34b080 100644 --- a/backend/internal/scimsync/module.go +++ b/backend/internal/scimsync/module.go @@ -6,11 +6,11 @@ import ( "log/slog" "net/http" - "github.com/gin-gonic/gin" "github.com/italypaleale/francis/actor" francishost "github.com/italypaleale/francis/host" "gorm.io/gorm" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/httpserver" ) @@ -48,13 +48,13 @@ func New(deps Dependencies) (*Module, error) { }, nil } -// RegisterRoutes mounts the SCIM service provider endpoints -func (m *Module) RegisterRoutes(apiGroup *gin.RouterGroup, auth gin.HandlerFunc) { - apiGroup.GET("/oidc/clients/:id/scim-service-provider", auth, httpserver.Handle(m.handler.getServiceProviderByClient)) - apiGroup.POST("/scim/service-provider", auth, httpserver.Handle(m.handler.createServiceProvider)) - apiGroup.POST("/scim/service-provider/:id/sync", auth, httpserver.Handle(m.handler.syncServiceProvider)) - apiGroup.PUT("/scim/service-provider/:id", auth, httpserver.Handle(m.handler.updateServiceProvider)) - apiGroup.DELETE("/scim/service-provider/:id", auth, httpserver.Handle(m.handler.deleteServiceProvider)) +// RegisterRoutes mounts the SCIM service provider endpoints, which belong to an OIDC client +func (m *Module) RegisterRoutes(r *authz.Router) { + r.GET("/oidc/clients/:id/scim-service-provider", authz.OidcClientsRead, httpserver.Handle(m.handler.getServiceProviderByClient)) + r.POST("/scim/service-provider", authz.OidcClientsWrite, httpserver.Handle(m.handler.createServiceProvider)) + r.POST("/scim/service-provider/:id/sync", authz.OidcClientsWrite, httpserver.Handle(m.handler.syncServiceProvider)) + r.PUT("/scim/service-provider/:id", authz.OidcClientsWrite, httpserver.Handle(m.handler.updateServiceProvider)) + r.DELETE("/scim/service-provider/:id", authz.OidcClientsWrite, httpserver.Handle(m.handler.deleteServiceProvider)) } // ScheduleSync schedules a debounced cluster-wide synchronization after SCIM-relevant data changes diff --git a/backend/internal/usersignup/module.go b/backend/internal/usersignup/module.go index 7875229d..102968f7 100644 --- a/backend/internal/usersignup/module.go +++ b/backend/internal/usersignup/module.go @@ -11,6 +11,7 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/appconfig" "github.com/pocket-id/pocket-id/backend/internal/auditlogs" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/httpserver" "github.com/pocket-id/pocket-id/backend/internal/model" @@ -78,12 +79,12 @@ func (m *Module) RunSignupTokenMigration(ctx context.Context) error { } // RegisterRoutes mounts the signup and signup-token management endpoints -// adminAuth guards the admin token-management routes; signupRateLimit throttles public self-signup -func (m *Module) RegisterRoutes(apiGroup *gin.RouterGroup, adminAuth, signupRateLimit gin.HandlerFunc) { - apiGroup.POST("/signup-tokens", adminAuth, httpserver.Handle(m.handler.createSignupToken)) - apiGroup.GET("/signup-tokens", adminAuth, httpserver.Handle(m.handler.listSignupTokens)) - apiGroup.DELETE("/signup-tokens/:id", adminAuth, httpserver.Handle(m.handler.deleteSignupToken)) - apiGroup.POST("/signup", signupRateLimit, httpserver.Handle(m.handler.signup)) - apiGroup.GET("/signup/setup", httpserver.Handle(m.handler.checkInitialAdminSetupAvailable)) - apiGroup.POST("/signup/setup", httpserver.Handle(m.handler.signUpInitialAdmin)) +// signupRateLimit throttles public self-signup +func (m *Module) RegisterRoutes(r *authz.Router, signupRateLimit gin.HandlerFunc) { + r.POST("/signup-tokens", authz.UsersWrite, httpserver.Handle(m.handler.createSignupToken)) + r.GET("/signup-tokens", authz.UsersRead, httpserver.Handle(m.handler.listSignupTokens)) + r.DELETE("/signup-tokens/:id", authz.UsersWrite, httpserver.Handle(m.handler.deleteSignupToken)) + r.Public().POST("/signup", signupRateLimit, httpserver.Handle(m.handler.signup)) + r.Public().GET("/signup/setup", httpserver.Handle(m.handler.checkInitialAdminSetupAvailable)) + r.Public().POST("/signup/setup", httpserver.Handle(m.handler.signUpInitialAdmin)) } diff --git a/backend/internal/webauthn/handler.go b/backend/internal/webauthn/handler.go index 401af4c3..48ecc822 100644 --- a/backend/internal/webauthn/handler.go +++ b/backend/internal/webauthn/handler.go @@ -14,6 +14,7 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/appconfig" "github.com/pocket-id/pocket-id/backend/internal/apperror" "github.com/pocket-id/pocket-id/backend/internal/auditlogs" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/httpserver" "github.com/pocket-id/pocket-id/backend/internal/utils" @@ -38,7 +39,7 @@ func (h *handler) beginRegistration(c *gin.Context) error { return fmt.Errorf("error loading app configuration: %w", err) } - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID options, err := h.service.BeginRegistration(c.Request.Context(), dbConfig, userID) if err != nil { return err @@ -60,7 +61,7 @@ func (h *handler) verifyRegistration(c *gin.Context) error { return apperror.MissingSessionID() } - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID credential, err := h.service.VerifyRegistration(c.Request.Context(), dbConfig, sessionID, userID, c.Request, c.ClientIP()) if err != nil { return err @@ -129,7 +130,7 @@ func (h *handler) verifyLogin(c *gin.Context) error { } func (h *handler) listCredentials(c *gin.Context) error { - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID credentials, err := h.service.ListCredentials(c.Request.Context(), userID) if err != nil { return err @@ -145,7 +146,7 @@ func (h *handler) listCredentials(c *gin.Context) error { } func (h *handler) deleteCredential(c *gin.Context) error { - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID credentialID := c.Param("id") clientIP := c.ClientIP() userAgent := c.Request.UserAgent() @@ -160,7 +161,7 @@ func (h *handler) deleteCredential(c *gin.Context) error { } func (h *handler) updateCredential(c *gin.Context) error { - userID := c.GetString("userID") + userID := authz.PrincipalFrom(c).UserID credentialID := c.Param("id") var input dto.WebauthnCredentialUpdateDto diff --git a/backend/internal/webauthn/module.go b/backend/internal/webauthn/module.go index 99f9ef3a..1115cdd4 100644 --- a/backend/internal/webauthn/module.go +++ b/backend/internal/webauthn/module.go @@ -13,6 +13,7 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/appconfig" "github.com/pocket-id/pocket-id/backend/internal/auditlogs" + "github.com/pocket-id/pocket-id/backend/internal/authz" "github.com/pocket-id/pocket-id/backend/internal/httpserver" "github.com/pocket-id/pocket-id/backend/internal/model" ) @@ -79,22 +80,22 @@ func New(deps Dependencies) (*Module, error) { } // RegisterRoutes mounts the WebAuthn registration, login and reauthentication endpoints -func (m *Module) RegisterRoutes(apiGroup *gin.RouterGroup, userAuth, browserAuth, loginRateLimit, reauthRateLimit gin.HandlerFunc) { - apiGroup.GET("/webauthn/register/start", browserAuth, httpserver.Handle(m.handler.beginRegistration)) - apiGroup.POST("/webauthn/register/finish", browserAuth, httpserver.Handle(m.handler.verifyRegistration)) +func (m *Module) RegisterRoutes(r *authz.Router, loginRateLimit, reauthRateLimit gin.HandlerFunc) { + r.GET("/webauthn/register/start", authz.AccountPasskeysEnroll, httpserver.Handle(m.handler.beginRegistration)) + r.POST("/webauthn/register/finish", authz.AccountPasskeysEnroll, httpserver.Handle(m.handler.verifyRegistration)) - apiGroup.GET("/webauthn/login/start", httpserver.Handle(m.handler.beginLogin)) - apiGroup.POST("/webauthn/login/finish", loginRateLimit, httpserver.Handle(m.handler.verifyLogin)) + r.Public().GET("/webauthn/login/start", httpserver.Handle(m.handler.beginLogin)) + r.Public().POST("/webauthn/login/finish", loginRateLimit, httpserver.Handle(m.handler.verifyLogin)) - apiGroup.POST("/webauthn/logout", userAuth, httpserver.Handle(m.handler.logout)) + r.POST("/webauthn/logout", authz.AccountSession, httpserver.Handle(m.handler.logout)) - apiGroup.POST("/webauthn/reauthenticate", browserAuth, reauthRateLimit, httpserver.Handle(m.handler.reauthenticate)) + r.POST("/webauthn/reauthenticate", authz.AccountSession, reauthRateLimit, httpserver.Handle(m.handler.reauthenticate)) - apiGroup.GET("/webauthn/credentials", userAuth, httpserver.Handle(m.handler.listCredentials)) - apiGroup.PATCH("/webauthn/credentials/:id", userAuth, httpserver.Handle(m.handler.updateCredential)) - apiGroup.DELETE("/webauthn/credentials/:id", userAuth, httpserver.Handle(m.handler.deleteCredential)) + r.GET("/webauthn/credentials", authz.AccountPasskeys, httpserver.Handle(m.handler.listCredentials)) + r.PATCH("/webauthn/credentials/:id", authz.AccountPasskeys, httpserver.Handle(m.handler.updateCredential)) + r.DELETE("/webauthn/credentials/:id", authz.AccountPasskeys, httpserver.Handle(m.handler.deleteCredential)) - apiGroup.GET("/webauthn/authenticator-icons/:aaguid", httpserver.Handle(m.handler.getThemedAuthenticatorIcon)) + r.Public().GET("/webauthn/authenticator-icons/:aaguid", httpserver.Handle(m.handler.getThemedAuthenticatorIcon)) } // ConsumeReauthenticationToken implements the OIDC module's ReauthenticationTokenConsumer interface