Files
pocket-id/backend/internal/middleware/cross_origin_protection.go
T

63 lines
1.6 KiB
Go

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
}
}