diff --git a/backend/internal/apperror/constructors.go b/backend/internal/apperror/constructors.go index 636b4c14..b17a8135 100644 --- a/backend/internal/apperror/constructors.go +++ b/backend/internal/apperror/constructors.go @@ -90,6 +90,10 @@ func MissingPermission() *Error { return New(CodeForbidden, http.StatusForbidden, "You don't have permission to perform this action") } +func CrossOriginRequestForbidden(cause error) *Error { + return Wrap(cause, CodeForbidden, http.StatusForbidden, "Cross-origin requests are not allowed") +} + func TooManyRequests() *Error { return New(CodeRateLimited, http.StatusTooManyRequests, "Too many requests") } diff --git a/backend/internal/bootstrap/router_bootstrap.go b/backend/internal/bootstrap/router_bootstrap.go index 22ca0bb9..3863f248 100644 --- a/backend/internal/bootstrap/router_bootstrap.go +++ b/backend/internal/bootstrap/router_bootstrap.go @@ -132,6 +132,7 @@ func registerGlobalMiddleware(r *gin.Engine) { r.Use(middleware.NewCorsMiddleware().Add()) r.Use(middleware.NewCspMiddleware().Add()) r.Use(middleware.NewErrorHandlerMiddleware().Add()) + r.Use(middleware.NewCrossOriginProtectionMiddleware(common.EnvConfig.AppURL).Add()) } func registerRoutes(r *gin.Engine, db *gorm.DB, svc *services, rateLimitServices map[string]*ratelimit.RateLimitService) error { diff --git a/backend/internal/middleware/cross_origin_protection.go b/backend/internal/middleware/cross_origin_protection.go new file mode 100644 index 00000000..526f8e7a --- /dev/null +++ b/backend/internal/middleware/cross_origin_protection.go @@ -0,0 +1,62 @@ +package middleware + +import ( + "log/slog" + "net/http" + + "github.com/gin-gonic/gin" + + "github.com/pocket-id/pocket-id/backend/internal/apperror" +) + +// CrossOriginProtectionMiddleware rejects state-changing browser requests that come from another origin to prevent CSRF +type CrossOriginProtectionMiddleware struct { + protection *http.CrossOriginProtection +} + +func NewCrossOriginProtectionMiddleware(appURL string) *CrossOriginProtectionMiddleware { + protection := http.NewCrossOriginProtection() + + // Trust APP_URL so browsers without Sec-Fetch-Site still pass when a reverse proxy rewrites the Host header + if appURL != "" { + if err := protection.AddTrustedOrigin(appURL); err != nil { + slog.Warn("Failed to add APP_URL as a trusted origin for cross-origin protection", slog.Any("error", err)) + } + } + + return &CrossOriginProtectionMiddleware{protection: protection} +} + +func (m *CrossOriginProtectionMiddleware) Add() gin.HandlerFunc { + return func(c *gin.Context) { + if isCrossOriginAllowedPath(c.FullPath()) { + c.Next() + return + } + + if err := m.protection.Check(c.Request); err != nil { + _ = c.Error(apperror.CrossOriginRequestForbidden(err)) + c.Abort() + return + } + + c.Next() + } +} + +// isCrossOriginAllowedPath reports whether a route must accept requests from other origins +// These are OAuth/OIDC protocol endpoints that relying parties call from their own origin, either via CORS or a cross-site form post +func isCrossOriginAllowedPath(path string) bool { + if isCorsPath(path) { + return true + } + + switch path { + case "/authorize", + "/api/oidc/par", + "/api/oidc/device/authorize": + return true + default: + return false + } +} diff --git a/backend/internal/middleware/cross_origin_protection_test.go b/backend/internal/middleware/cross_origin_protection_test.go new file mode 100644 index 00000000..0e10701e --- /dev/null +++ b/backend/internal/middleware/cross_origin_protection_test.go @@ -0,0 +1,81 @@ +package middleware + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func newCrossOriginProtectionTestRouter() *gin.Engine { + gin.SetMode(gin.TestMode) + router := gin.New() + router.Use(NewErrorHandlerMiddleware().Add()) + router.Use(NewCrossOriginProtectionMiddleware("https://id.example.com").Add()) + + ok := func(c *gin.Context) { c.Status(http.StatusOK) } + router.GET("/api/oidc/device/verify", ok) + router.POST("/api/oidc/device/verify", ok) + router.DELETE("/api/users/:id", ok) + router.POST("/authorize", ok) + router.POST("/api/oidc/token", ok) + router.POST("/api/oidc/userinfo", ok) + router.POST("/api/oidc/introspect", ok) + router.POST("/api/oidc/end-session", ok) + router.POST("/api/oidc/par", ok) + router.POST("/api/oidc/device/authorize", ok) + + return router +} + +func TestCrossOriginProtectionMiddleware(t *testing.T) { + router := newCrossOriginProtectionTestRouter() + + type testCase struct { + name string + method string + path string + headers map[string]string + want int + } + + tests := []testCase{ + // Browser requests from a sibling origin on the same site carry SameSite=Lax cookies, so they must be rejected + {"same-site sibling origin", http.MethodPost, "/api/oidc/device/verify?code=ABCD", map[string]string{"Sec-Fetch-Site": "same-site", "Origin": "https://app.example.com"}, http.StatusForbidden}, + {"cross-site origin", http.MethodPost, "/api/oidc/device/verify?code=ABCD", map[string]string{"Sec-Fetch-Site": "cross-site", "Origin": "https://evil.test"}, http.StatusForbidden}, + {"cross-site delete", http.MethodDelete, "/api/users/1", map[string]string{"Sec-Fetch-Site": "cross-site"}, http.StatusForbidden}, + {"old browser with foreign origin", http.MethodPost, "/api/oidc/device/verify", map[string]string{"Origin": "https://app.example.com"}, http.StatusForbidden}, + + // Requests from the SPA itself and direct navigations are allowed + {"same origin", http.MethodPost, "/api/oidc/device/verify", map[string]string{"Sec-Fetch-Site": "same-origin", "Origin": "https://id.example.com"}, http.StatusOK}, + {"user-initiated navigation", http.MethodPost, "/api/oidc/device/verify", map[string]string{"Sec-Fetch-Site": "none"}, http.StatusOK}, + {"old browser with matching host", http.MethodPost, "/api/oidc/device/verify", map[string]string{"Origin": "http://example.com"}, http.StatusOK}, + {"old browser behind host-rewriting proxy", http.MethodPost, "/api/oidc/device/verify", map[string]string{"Origin": "https://id.example.com"}, http.StatusOK}, + + // Non-browser clients such as API key scripts send neither header + {"non-browser client", http.MethodPost, "/api/oidc/device/verify", nil, http.StatusOK}, + + // Safe methods are never checked + {"cross-site get", http.MethodGet, "/api/oidc/device/verify", map[string]string{"Sec-Fetch-Site": "cross-site"}, http.StatusOK}, + } + + // Protocol endpoints are called by relying parties from their own origin + for _, path := range []string{"/authorize", "/api/oidc/token", "/api/oidc/userinfo", "/api/oidc/introspect", "/api/oidc/end-session", "/api/oidc/par", "/api/oidc/device/authorize"} { + tests = append(tests, testCase{"cross-site protocol endpoint " + path, http.MethodPost, path, map[string]string{"Sec-Fetch-Site": "cross-site", "Origin": "https://rp.test"}, http.StatusOK}) + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + req := httptest.NewRequestWithContext(t.Context(), tt.method, tt.path, http.NoBody) + for k, v := range tt.headers { + req.Header.Set(k, v) + } + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + + require.Equal(t, tt.want, w.Code, w.Body.String()) + }) + } +}