mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-29 05:56:38 +08:00
fix(router): restrict CORS origin reflection to allowed hosts (AUTH-ROUTE-5)
- Extract isOriginAllowed helper to match Origin against server_address configurations - Ensure arbitrary origins are not reflected and credentials are not allowed when server_address is unconfigured or mismatched - Trim trailing slashes from allowed origins configuration for robust matching - Add TestCORSMiddleware to cover all CORS matching and rejection scenarios
This commit is contained in:
@@ -5,8 +5,10 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
@@ -67,17 +69,26 @@ func loggerMiddleware() gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
func isOriginAllowed(ctx context.Context, origin string) bool {
|
||||
var sc model.SystemConfig
|
||||
if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err != nil || sc.Value == "" {
|
||||
return false
|
||||
}
|
||||
allowedOrigins := strings.Split(sc.Value, ",")
|
||||
for _, allowed := range allowedOrigins {
|
||||
allowed = strings.TrimRight(strings.TrimSpace(allowed), "/")
|
||||
if allowed != "" && strings.EqualFold(allowed, origin) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func corsMiddleware() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
origin := c.Request.Header.Get("Origin")
|
||||
if origin != "" {
|
||||
var sc model.SystemConfig
|
||||
// Fetch from system config. We use request context which supports trace
|
||||
if err := sc.GetByKey(c.Request.Context(), model.ConfigKeyServerAddress); err == nil && sc.Value != "" {
|
||||
c.Writer.Header().Set("Access-Control-Allow-Origin", sc.Value)
|
||||
} else {
|
||||
c.Writer.Header().Set("Access-Control-Allow-Origin", origin)
|
||||
}
|
||||
if origin != "" && isOriginAllowed(c.Request.Context(), origin) {
|
||||
c.Writer.Header().Set("Access-Control-Allow-Origin", origin)
|
||||
c.Writer.Header().Set("Access-Control-Allow-Credentials", "true")
|
||||
c.Writer.Header().Set("Access-Control-Allow-Headers", "Content-Type, Content-Length, Accept-Encoding, X-CSRF-Token, Authorization, accept, origin, Cache-Control, X-Requested-With, X-Access-Token, X-Cap-Token")
|
||||
c.Writer.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS, GET, PUT, DELETE, PATCH")
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package router
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func TestCORSMiddleware(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
// Helper to clear config cache
|
||||
clearConfigCache := func() {
|
||||
_ = db.Redis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)).Err()
|
||||
}
|
||||
|
||||
t.Run("missing server_address configuration returns no CORS headers", func(t *testing.T) {
|
||||
clearConfigCache()
|
||||
// Ensure it's empty in DB
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyServerAddress).Update("value", "").Error; err != nil {
|
||||
t.Fatalf("failed to update config: %v", err)
|
||||
}
|
||||
clearConfigCache()
|
||||
|
||||
r := gin.New()
|
||||
r.Use(corsMiddleware())
|
||||
r.GET("/test", func(c *gin.Context) {
|
||||
c.String(http.StatusOK, "ok")
|
||||
})
|
||||
|
||||
req, _ := http.NewRequest(http.MethodGet, "/test", nil)
|
||||
req.Header.Set("Origin", "http://attacker.com")
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d", w.Code)
|
||||
}
|
||||
if val := w.Header().Get("Access-Control-Allow-Origin"); val != "" {
|
||||
t.Errorf("expected empty Access-Control-Allow-Origin header, got %q", val)
|
||||
}
|
||||
if val := w.Header().Get("Access-Control-Allow-Credentials"); val != "" {
|
||||
t.Errorf("expected empty Access-Control-Allow-Credentials header, got %q", val)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("matching server_address allows origin and sets credential headers", func(t *testing.T) {
|
||||
clearConfigCache()
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyServerAddress).Update("value", "https://trusted.com, http://localhost:3000/").Error; err != nil {
|
||||
t.Fatalf("failed to update config: %v", err)
|
||||
}
|
||||
clearConfigCache()
|
||||
|
||||
r := gin.New()
|
||||
r.Use(corsMiddleware())
|
||||
r.GET("/test", func(c *gin.Context) {
|
||||
c.String(http.StatusOK, "ok")
|
||||
})
|
||||
|
||||
// Test trusted origin 1
|
||||
req1, _ := http.NewRequest(http.MethodGet, "/test", nil)
|
||||
req1.Header.Set("Origin", "https://trusted.com")
|
||||
w1 := httptest.NewRecorder()
|
||||
r.ServeHTTP(w1, req1)
|
||||
|
||||
if w1.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d", w1.Code)
|
||||
}
|
||||
if val := w1.Header().Get("Access-Control-Allow-Origin"); val != "https://trusted.com" {
|
||||
t.Errorf("expected Access-Control-Allow-Origin 'https://trusted.com', got %q", val)
|
||||
}
|
||||
if val := w1.Header().Get("Access-Control-Allow-Credentials"); val != "true" {
|
||||
t.Errorf("expected Access-Control-Allow-Credentials 'true', got %q", val)
|
||||
}
|
||||
|
||||
// Test trusted origin 2 (trimmed trailing slash)
|
||||
req2, _ := http.NewRequest(http.MethodGet, "/test", nil)
|
||||
req2.Header.Set("Origin", "http://localhost:3000")
|
||||
w2 := httptest.NewRecorder()
|
||||
r.ServeHTTP(w2, req2)
|
||||
|
||||
if w2.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d", w2.Code)
|
||||
}
|
||||
if val := w2.Header().Get("Access-Control-Allow-Origin"); val != "http://localhost:3000" {
|
||||
t.Errorf("expected Access-Control-Allow-Origin 'http://localhost:3000', got %q", val)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("non-matching origin is denied CORS headers", func(t *testing.T) {
|
||||
clearConfigCache()
|
||||
if err := dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyServerAddress).Update("value", "https://trusted.com").Error; err != nil {
|
||||
t.Fatalf("failed to update config: %v", err)
|
||||
}
|
||||
clearConfigCache()
|
||||
|
||||
r := gin.New()
|
||||
r.Use(corsMiddleware())
|
||||
r.GET("/test", func(c *gin.Context) {
|
||||
c.String(http.StatusOK, "ok")
|
||||
})
|
||||
|
||||
req, _ := http.NewRequest(http.MethodGet, "/test", nil)
|
||||
req.Header.Set("Origin", "https://attacker.com")
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d", w.Code)
|
||||
}
|
||||
if val := w.Header().Get("Access-Control-Allow-Origin"); val != "" {
|
||||
t.Errorf("expected empty Access-Control-Allow-Origin header, got %q", val)
|
||||
}
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user