mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-07 08:06:37 +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 (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"net/url"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
|
||||||
"github.com/gorilla/websocket"
|
"github.com/gorilla/websocket"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
)
|
)
|
||||||
|
|
||||||
// getUpgrader 返回 WebSocket 升级器
|
// getUpgrader 返回 WebSocket 升级器并执行 Origin 安全检查以防止 CSWSH 攻击
|
||||||
func getUpgrader() *websocket.Upgrader {
|
func getUpgrader() *websocket.Upgrader {
|
||||||
return &websocket.Upgrader{
|
return &websocket.Upgrader{
|
||||||
CheckOrigin: func(_ *http.Request) bool {
|
CheckOrigin: func(r *http.Request) bool {
|
||||||
return true // CORS 由 Gin 中间件处理
|
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