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