From f6f8c25930b7a41ca41a9dd89506e026a67af89a Mon Sep 17 00:00:00 2001 From: ryan Date: Sat, 13 Jun 2026 10:23:20 +0800 Subject: [PATCH] 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 --- internal/router/middlewares.go | 27 ++++-- internal/router/middlewares_test.go | 127 ++++++++++++++++++++++++++++ 2 files changed, 146 insertions(+), 8 deletions(-) create mode 100644 internal/router/middlewares_test.go diff --git a/internal/router/middlewares.go b/internal/router/middlewares.go index a6aea931..65db3974 100644 --- a/internal/router/middlewares.go +++ b/internal/router/middlewares.go @@ -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") diff --git a/internal/router/middlewares_test.go b/internal/router/middlewares_test.go new file mode 100644 index 00000000..6454d1bc --- /dev/null +++ b/internal/router/middlewares_test.go @@ -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) + } + }) +}