mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-29 05:56:38 +08:00
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.
This commit is contained in:
@@ -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
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user