From cf9acf3ea046a93290aafd23549cc5e8c6eeeb65 Mon Sep 17 00:00:00 2001 From: duorameng <2997944583@qq.com> Date: Mon, 27 Apr 2026 10:54:00 +0800 Subject: [PATCH] chore: update ws security --- internal/constant/constant.go | 10 +++ internal/controllers/agent_controller.go | 4 +- internal/controllers/log_ws_controller.go | 27 +++++++++ internal/controllers/terminal_controller.go | 67 +++++++++++++++++++-- internal/router/router.go | 5 +- internal/utils/network.go | 57 ++++++++++++++++++ 6 files changed, 162 insertions(+), 8 deletions(-) create mode 100644 internal/utils/network.go diff --git a/internal/constant/constant.go b/internal/constant/constant.go index c78b359..e91c82c 100644 --- a/internal/constant/constant.go +++ b/internal/constant/constant.go @@ -1,5 +1,7 @@ package constant +import "time" + const ( // ConfigPath 配置文件路径 @@ -159,6 +161,14 @@ const ( // Env Type EnvTypeNormal = "normal" EnvTypeSecret = "secret" + + // WebSocket 安全常量 + // PongWait 收到 pong 的超时时间 + PongWait = 60 * time.Second + // PingPeriod 发送 ping 的周期 + PingPeriod = (PongWait * 9) / 10 + // MaxMessageSize 允许的最大消息大小 + MaxMessageSize = 1024 * 1024 // 1MB ) // TablePrefix 表前缀,从配置文件读取 diff --git a/internal/controllers/agent_controller.go b/internal/controllers/agent_controller.go index 0f15160..e5f0106 100644 --- a/internal/controllers/agent_controller.go +++ b/internal/controllers/agent_controller.go @@ -20,9 +20,7 @@ import ( ) var agentUpgrader = websocket.Upgrader{ - CheckOrigin: func(r *http.Request) bool { - return true - }, + CheckOrigin: utils.CheckWSOrigin, } // AgentController Agent 控制器 diff --git a/internal/controllers/log_ws_controller.go b/internal/controllers/log_ws_controller.go index 4c77135..33e8195 100644 --- a/internal/controllers/log_ws_controller.go +++ b/internal/controllers/log_ws_controller.go @@ -3,6 +3,7 @@ package controllers import ( "fmt" + "github.com/engigu/baihu-panel/internal/constant" "github.com/engigu/baihu-panel/internal/database" "github.com/engigu/baihu-panel/internal/models" "github.com/engigu/baihu-panel/internal/services/tasks" @@ -10,6 +11,7 @@ import ( "github.com/gin-gonic/gin" "github.com/gorilla/websocket" + "time" ) type LogWSController struct{} @@ -32,6 +34,23 @@ func (lc *LogWSController) StreamLog(c *gin.Context) { } defer conn.Close() + // DoS 保护与心跳设置 + conn.SetReadLimit(constant.MaxMessageSize) // 使用与终端一致的限制 + conn.SetReadDeadline(time.Now().Add(constant.PongWait)) + conn.SetPongHandler(func(string) error { + conn.SetReadDeadline(time.Now().Add(constant.PongWait)) + return nil + }) + + // 启动一个读取循环,用于处理 pong 和检测断开 + go func() { + for { + if _, _, err := conn.ReadMessage(); err != nil { + break + } + } + }() + // 1. 检查数据库中是否已结束 var taskLog models.TaskLog res := database.DB.Where("id = ?", logID).Limit(1).Find(&taskLog) @@ -68,9 +87,17 @@ func (lc *LogWSController) StreamLog(c *gin.Context) { sub := tl.Subscribe() defer tl.Unsubscribe(sub) + ticker := time.NewTicker(constant.PingPeriod) + defer ticker.Stop() + // 推送更新 for { select { + case <-ticker.C: + conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) + if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil { + return + } case data, ok := <-sub: if !ok { // 任务结束,尝试刷新最后一次库内完整内容 diff --git a/internal/controllers/terminal_controller.go b/internal/controllers/terminal_controller.go index 5f12de4..1ad7e19 100644 --- a/internal/controllers/terminal_controller.go +++ b/internal/controllers/terminal_controller.go @@ -4,11 +4,11 @@ import ( "bufio" "encoding/json" "io" - "net/http" "os" "path/filepath" "runtime" "sync" + "time" "unicode/utf8" "github.com/engigu/baihu-panel/internal/constant" @@ -22,6 +22,7 @@ import ( "golang.org/x/text/transform" ) + type TerminalController struct { envService *services.EnvService } @@ -33,9 +34,7 @@ func NewTerminalController(envService *services.EnvService) *TerminalController } var upgrader = websocket.Upgrader{ - CheckOrigin: func(r *http.Request) bool { - return true - }, + CheckOrigin: utils.CheckWSOrigin, } // toUTF8 将可能是 GBK 编码的字节转换为 UTF-8 @@ -98,6 +97,13 @@ func (tc *TerminalController) HandleWebSocket(c *gin.Context) { // handlePtyMode 使用 PTY 处理终端(Unix/macOS) func (tc *TerminalController) handlePtyMode(conn *websocket.Conn, userID string) { + conn.SetReadLimit(constant.MaxMessageSize) + conn.SetReadDeadline(time.Now().Add(constant.PongWait)) + conn.SetPongHandler(func(string) error { + conn.SetReadDeadline(time.Now().Add(constant.PongWait)) + return nil + }) + // 发送 PTY 模式标识 conn.WriteMessage(websocket.TextMessage, []byte("__PTY_MODE__")) @@ -124,6 +130,7 @@ func (tc *TerminalController) handlePtyMode(conn *websocket.Conn, userID string) writeMessage := func(data []byte) { connMu.Lock() defer connMu.Unlock() + conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) conn.WriteMessage(websocket.TextMessage, data) } @@ -143,6 +150,27 @@ func (tc *TerminalController) handlePtyMode(conn *websocket.Conn, userID string) } }() + // 启动 ping 协程 + pingDone := make(chan struct{}) + go func() { + ticker := time.NewTicker(constant.PingPeriod) + defer ticker.Stop() + for { + select { + case <-ticker.C: + connMu.Lock() + conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) + if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil { + connMu.Unlock() + return + } + connMu.Unlock() + case <-pingDone: + return + } + } + }() + for { _, message, err := conn.ReadMessage() if err != nil { @@ -167,6 +195,7 @@ func (tc *TerminalController) handlePtyMode(conn *websocket.Conn, userID string) } } + close(pingDone) cmd.Process.Kill() cmd.Wait() wg.Wait() @@ -174,6 +203,13 @@ func (tc *TerminalController) handlePtyMode(conn *websocket.Conn, userID string) // handlePipeMode 使用 pipe 处理终端(Windows) func (tc *TerminalController) handlePipeMode(conn *websocket.Conn, userID string) { + conn.SetReadLimit(constant.MaxMessageSize) + conn.SetReadDeadline(time.Now().Add(constant.PongWait)) + conn.SetPongHandler(func(string) error { + conn.SetReadDeadline(time.Now().Add(constant.PongWait)) + return nil + }) + // 发送 pipe 模式标识 conn.WriteMessage(websocket.TextMessage, []byte("__PIPE_MODE__")) @@ -215,6 +251,7 @@ func (tc *TerminalController) handlePipeMode(conn *websocket.Conn, userID string writeMessage := func(data []byte) { connMu.Lock() defer connMu.Unlock() + conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) conn.WriteMessage(websocket.TextMessage, data) } @@ -238,6 +275,27 @@ func (tc *TerminalController) handlePipeMode(conn *websocket.Conn, userID string go readOutput(stdout) go readOutput(stderr) + // 启动 ping 协程 + pingDone := make(chan struct{}) + go func() { + ticker := time.NewTicker(constant.PingPeriod) + defer ticker.Stop() + for { + select { + case <-ticker.C: + connMu.Lock() + conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) + if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil { + connMu.Unlock() + return + } + connMu.Unlock() + case <-pingDone: + return + } + } + }() + for { _, message, err := conn.ReadMessage() if err != nil { @@ -258,6 +316,7 @@ func (tc *TerminalController) handlePipeMode(conn *websocket.Conn, userID string } } + close(pingDone) stdin.Close() cmd.Process.Kill() cmd.Wait() diff --git a/internal/router/router.go b/internal/router/router.go index e1a732b..6f1747d 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -2,6 +2,7 @@ package router import ( "strings" + "os" "github.com/engigu/baihu-panel/internal/controllers" "github.com/engigu/baihu-panel/internal/middleware" @@ -31,7 +32,9 @@ type Controllers struct { } func Setup(c *Controllers) *gin.Engine { - gin.SetMode(gin.ReleaseMode) + if os.Getenv("GIN_MODE") == "" { + gin.SetMode(gin.ReleaseMode) + } router := gin.New() router.Use(middleware.GinLogger(), middleware.GinRecovery()) diff --git a/internal/utils/network.go b/internal/utils/network.go new file mode 100644 index 0000000..9a1aa7f --- /dev/null +++ b/internal/utils/network.go @@ -0,0 +1,57 @@ +package utils + +import ( + "net/http" + "net/url" + "os" + "strings" + + "github.com/gin-gonic/gin" +) + +// CheckWSOrigin 校验 WebSocket 的 Origin 来源是否安全。 +// 默认仅允许同源请求,可通过环境变量 BH_ALLOWED_ORIGINS 配置额外的允许列表(逗号分隔)。 +func CheckWSOrigin(r *http.Request) bool { + origin := r.Header.Get("Origin") + if origin == "" { + // 非浏览器发起的请求(如直接用脚本连接)通常不带 Origin,默认放行。 + return true + } + + u, err := url.Parse(origin) + if err != nil { + return false + } + + // 0. 开发环境校验:如果是非 Release 模式,默认放行 + if gin.Mode() != gin.ReleaseMode { + return true + } + + // 1. 同源校验:Origin 的 Host 与请求头中的 Host 一致 + if strings.EqualFold(u.Host, r.Host) { + return true + } + + // 2. 环境变量配置的允许列表校验 + allowedOrigins := os.Getenv("BH_ALLOWED_ORIGINS") + if allowedOrigins != "" { + origins := strings.Split(allowedOrigins, ",") + for _, o := range origins { + o = strings.TrimSpace(o) + if o == "*" { + return true + } + // 匹配完整 Origin (如 http://localhost:5173) 或仅 Host 部分 + if strings.EqualFold(o, origin) || strings.EqualFold(o, u.Host) { + return true + } + } + } + + // 3. 兜底策略:如果是开发环境常见的 localhost/127.0.0.1,且端口不一致的情况, + // 如果用户没有配置允许列表,我们在非 Release 模式下可以考虑放行, + // 但为了安全,默认应严格限制。建议开发时通过 BH_ALLOWED_ORIGINS=localhost:5173 显式开启。 + + return false +}