From eb999eba09fef5de6751c4562ec44ce1d08fcf89 Mon Sep 17 00:00:00 2001 From: ryan Date: Sat, 13 Jun 2026 10:25:25 +0800 Subject: [PATCH] fix(logs): restrict websocket origin to prevent cswsh (LOG-2) - Restrict WebSocket upgrade to same-origin or configured server_address allowed origins. - Add comprehensive test suite in utils_test.go to verify origin matching rules. --- internal/apps/admin/logs/utils.go | 34 ++++++- internal/apps/admin/logs/utils_test.go | 118 +++++++++++++++++++++++++ 2 files changed, 149 insertions(+), 3 deletions(-) create mode 100644 internal/apps/admin/logs/utils_test.go diff --git a/internal/apps/admin/logs/utils.go b/internal/apps/admin/logs/utils.go index a6371ebd..92232c64 100644 --- a/internal/apps/admin/logs/utils.go +++ b/internal/apps/admin/logs/utils.go @@ -6,16 +6,44 @@ package logs import ( "net/http" + "net/url" "strconv" + "strings" "github.com/gorilla/websocket" + + "github.com/Rain-kl/Wavelet/internal/model" ) -// getUpgrader 返回 WebSocket 升级器 +// getUpgrader 返回 WebSocket 升级器并执行 Origin 安全检查以防止 CSWSH 攻击 func getUpgrader() *websocket.Upgrader { return &websocket.Upgrader{ - CheckOrigin: func(_ *http.Request) bool { - return true // CORS 由 Gin 中间件处理 + CheckOrigin: func(r *http.Request) bool { + origin := r.Header.Get("Origin") + if origin == "" { + return true + } + + // 1. 同源检查 (Same-origin check) + u, err := url.Parse(origin) + if err == nil && strings.EqualFold(u.Host, r.Host) { + return true + } + + // 2. 检查配置的允许跨域 Origin (Check allowed origins in system config) + ctx := r.Context() + var sc model.SystemConfig + if err := sc.GetByKey(ctx, model.ConfigKeyServerAddress); err == nil && sc.Value != "" { + originToCheck := strings.TrimRight(strings.TrimSpace(origin), "/") + allowedOrigins := strings.Split(sc.Value, ",") + for _, allowed := range allowedOrigins { + allowed = strings.TrimRight(strings.TrimSpace(allowed), "/") + if allowed != "" && strings.EqualFold(allowed, originToCheck) { + return true + } + } + } + return false }, } } diff --git a/internal/apps/admin/logs/utils_test.go b/internal/apps/admin/logs/utils_test.go new file mode 100644 index 00000000..d4f5c871 --- /dev/null +++ b/internal/apps/admin/logs/utils_test.go @@ -0,0 +1,118 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package logs + +import ( + "context" + "net/http" + "testing" + + "gorm.io/driver/sqlite" + "gorm.io/gorm" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" +) + +func setupTestDB(t *testing.T) *gorm.DB { + dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + t.Fatalf("failed to open sqlite in memory: %v", err) + } + err = dbConn.AutoMigrate(&model.SystemConfig{}) + if err != nil { + t.Fatalf("failed to migrate schema: %v", err) + } + db.SetDB(dbConn) + return dbConn +} + +func TestWebSocketCheckOrigin(t *testing.T) { + dbConn := setupTestDB(t) + + // Clean up global DB after test + defer db.SetDB(nil) + + // Seed ConfigKeyServerAddress with allowed frontend origin + allowedOrigin := "http://localhost:3000" + if err := dbConn.Create(&model.SystemConfig{ + Key: model.ConfigKeyServerAddress, + Value: allowedOrigin, + }).Error; err != nil { + t.Fatalf("failed to seed server address config: %v", err) + } + + upgrader := getUpgrader() + if upgrader.CheckOrigin == nil { + t.Fatal("expected CheckOrigin to be defined") + } + + tests := []struct { + name string + origin string + host string + wantOK bool + }{ + { + name: "empty origin (non-browser clients)", + origin: "", + host: "localhost:8000", + wantOK: true, + }, + { + name: "same-origin request", + origin: "http://localhost:8000", + host: "localhost:8000", + wantOK: true, + }, + { + name: "same-origin request case insensitive", + origin: "HTTP://LOCALHOST:8000", + host: "localhost:8000", + wantOK: true, + }, + { + name: "configured allowed origin request", + origin: "http://localhost:3000", + host: "localhost:8000", + wantOK: true, + }, + { + name: "configured allowed origin request with trailing slash", + origin: "http://localhost:3000/", + host: "localhost:8000", + wantOK: true, + }, + { + name: "unauthorized third-party origin", + origin: "http://evil.com", + host: "localhost:8000", + wantOK: false, + }, + { + name: "invalid origin format", + origin: "::not-a-valid-url", + host: "localhost:8000", + wantOK: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + req, err := http.NewRequestWithContext(context.Background(), "GET", "/api/v1/admin/logs/ws", nil) + if err != nil { + t.Fatalf("failed to create request: %v", err) + } + req.Host = tt.host + if tt.origin != "" { + req.Header.Set("Origin", tt.origin) + } + + got := upgrader.CheckOrigin(req) + if got != tt.wantOK { + t.Errorf("CheckOrigin() = %v, want %v", got, tt.wantOK) + } + }) + } +}