fix: reject cross-origin state-changing requests

This commit is contained in:
Elias Schneider
2026-10-09 20:42:05 +02:00
parent 226648c7b2
commit d31e3e1286
4 changed files with 148 additions and 0 deletions
@@ -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")
}
@@ -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 {
@@ -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
}
}
@@ -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())
})
}
}