mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-10-10 03:39:05 +02:00
fix: reject cross-origin state-changing requests
This commit is contained in:
@@ -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())
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user