chore: update ws security
This commit is contained in:
@@ -1,5 +1,7 @@
|
|||||||
package constant
|
package constant
|
||||||
|
|
||||||
|
import "time"
|
||||||
|
|
||||||
const (
|
const (
|
||||||
|
|
||||||
// ConfigPath 配置文件路径
|
// ConfigPath 配置文件路径
|
||||||
@@ -159,6 +161,14 @@ const (
|
|||||||
// Env Type
|
// Env Type
|
||||||
EnvTypeNormal = "normal"
|
EnvTypeNormal = "normal"
|
||||||
EnvTypeSecret = "secret"
|
EnvTypeSecret = "secret"
|
||||||
|
|
||||||
|
// WebSocket 安全常量
|
||||||
|
// PongWait 收到 pong 的超时时间
|
||||||
|
PongWait = 60 * time.Second
|
||||||
|
// PingPeriod 发送 ping 的周期
|
||||||
|
PingPeriod = (PongWait * 9) / 10
|
||||||
|
// MaxMessageSize 允许的最大消息大小
|
||||||
|
MaxMessageSize = 1024 * 1024 // 1MB
|
||||||
)
|
)
|
||||||
|
|
||||||
// TablePrefix 表前缀,从配置文件读取
|
// TablePrefix 表前缀,从配置文件读取
|
||||||
|
|||||||
@@ -20,9 +20,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
var agentUpgrader = websocket.Upgrader{
|
var agentUpgrader = websocket.Upgrader{
|
||||||
CheckOrigin: func(r *http.Request) bool {
|
CheckOrigin: utils.CheckWSOrigin,
|
||||||
return true
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// AgentController Agent 控制器
|
// AgentController Agent 控制器
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package controllers
|
|||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/engigu/baihu-panel/internal/constant"
|
||||||
"github.com/engigu/baihu-panel/internal/database"
|
"github.com/engigu/baihu-panel/internal/database"
|
||||||
"github.com/engigu/baihu-panel/internal/models"
|
"github.com/engigu/baihu-panel/internal/models"
|
||||||
"github.com/engigu/baihu-panel/internal/services/tasks"
|
"github.com/engigu/baihu-panel/internal/services/tasks"
|
||||||
@@ -10,6 +11,7 @@ import (
|
|||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/gorilla/websocket"
|
"github.com/gorilla/websocket"
|
||||||
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
type LogWSController struct{}
|
type LogWSController struct{}
|
||||||
@@ -32,6 +34,23 @@ func (lc *LogWSController) StreamLog(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
defer conn.Close()
|
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. 检查数据库中是否已结束
|
// 1. 检查数据库中是否已结束
|
||||||
var taskLog models.TaskLog
|
var taskLog models.TaskLog
|
||||||
res := database.DB.Where("id = ?", logID).Limit(1).Find(&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()
|
sub := tl.Subscribe()
|
||||||
defer tl.Unsubscribe(sub)
|
defer tl.Unsubscribe(sub)
|
||||||
|
|
||||||
|
ticker := time.NewTicker(constant.PingPeriod)
|
||||||
|
defer ticker.Stop()
|
||||||
|
|
||||||
// 推送更新
|
// 推送更新
|
||||||
for {
|
for {
|
||||||
select {
|
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:
|
case data, ok := <-sub:
|
||||||
if !ok {
|
if !ok {
|
||||||
// 任务结束,尝试刷新最后一次库内完整内容
|
// 任务结束,尝试刷新最后一次库内完整内容
|
||||||
|
|||||||
@@ -4,11 +4,11 @@ import (
|
|||||||
"bufio"
|
"bufio"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"runtime"
|
"runtime"
|
||||||
"sync"
|
"sync"
|
||||||
|
"time"
|
||||||
"unicode/utf8"
|
"unicode/utf8"
|
||||||
|
|
||||||
"github.com/engigu/baihu-panel/internal/constant"
|
"github.com/engigu/baihu-panel/internal/constant"
|
||||||
@@ -22,6 +22,7 @@ import (
|
|||||||
"golang.org/x/text/transform"
|
"golang.org/x/text/transform"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
type TerminalController struct {
|
type TerminalController struct {
|
||||||
envService *services.EnvService
|
envService *services.EnvService
|
||||||
}
|
}
|
||||||
@@ -33,9 +34,7 @@ func NewTerminalController(envService *services.EnvService) *TerminalController
|
|||||||
}
|
}
|
||||||
|
|
||||||
var upgrader = websocket.Upgrader{
|
var upgrader = websocket.Upgrader{
|
||||||
CheckOrigin: func(r *http.Request) bool {
|
CheckOrigin: utils.CheckWSOrigin,
|
||||||
return true
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// toUTF8 将可能是 GBK 编码的字节转换为 UTF-8
|
// toUTF8 将可能是 GBK 编码的字节转换为 UTF-8
|
||||||
@@ -98,6 +97,13 @@ func (tc *TerminalController) HandleWebSocket(c *gin.Context) {
|
|||||||
|
|
||||||
// handlePtyMode 使用 PTY 处理终端(Unix/macOS)
|
// handlePtyMode 使用 PTY 处理终端(Unix/macOS)
|
||||||
func (tc *TerminalController) handlePtyMode(conn *websocket.Conn, userID string) {
|
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 模式标识
|
// 发送 PTY 模式标识
|
||||||
conn.WriteMessage(websocket.TextMessage, []byte("__PTY_MODE__"))
|
conn.WriteMessage(websocket.TextMessage, []byte("__PTY_MODE__"))
|
||||||
|
|
||||||
@@ -124,6 +130,7 @@ func (tc *TerminalController) handlePtyMode(conn *websocket.Conn, userID string)
|
|||||||
writeMessage := func(data []byte) {
|
writeMessage := func(data []byte) {
|
||||||
connMu.Lock()
|
connMu.Lock()
|
||||||
defer connMu.Unlock()
|
defer connMu.Unlock()
|
||||||
|
conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||||
conn.WriteMessage(websocket.TextMessage, data)
|
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 {
|
for {
|
||||||
_, message, err := conn.ReadMessage()
|
_, message, err := conn.ReadMessage()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -167,6 +195,7 @@ func (tc *TerminalController) handlePtyMode(conn *websocket.Conn, userID string)
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
close(pingDone)
|
||||||
cmd.Process.Kill()
|
cmd.Process.Kill()
|
||||||
cmd.Wait()
|
cmd.Wait()
|
||||||
wg.Wait()
|
wg.Wait()
|
||||||
@@ -174,6 +203,13 @@ func (tc *TerminalController) handlePtyMode(conn *websocket.Conn, userID string)
|
|||||||
|
|
||||||
// handlePipeMode 使用 pipe 处理终端(Windows)
|
// handlePipeMode 使用 pipe 处理终端(Windows)
|
||||||
func (tc *TerminalController) handlePipeMode(conn *websocket.Conn, userID string) {
|
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 模式标识
|
// 发送 pipe 模式标识
|
||||||
conn.WriteMessage(websocket.TextMessage, []byte("__PIPE_MODE__"))
|
conn.WriteMessage(websocket.TextMessage, []byte("__PIPE_MODE__"))
|
||||||
|
|
||||||
@@ -215,6 +251,7 @@ func (tc *TerminalController) handlePipeMode(conn *websocket.Conn, userID string
|
|||||||
writeMessage := func(data []byte) {
|
writeMessage := func(data []byte) {
|
||||||
connMu.Lock()
|
connMu.Lock()
|
||||||
defer connMu.Unlock()
|
defer connMu.Unlock()
|
||||||
|
conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||||
conn.WriteMessage(websocket.TextMessage, data)
|
conn.WriteMessage(websocket.TextMessage, data)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -238,6 +275,27 @@ func (tc *TerminalController) handlePipeMode(conn *websocket.Conn, userID string
|
|||||||
go readOutput(stdout)
|
go readOutput(stdout)
|
||||||
go readOutput(stderr)
|
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 {
|
for {
|
||||||
_, message, err := conn.ReadMessage()
|
_, message, err := conn.ReadMessage()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -258,6 +316,7 @@ func (tc *TerminalController) handlePipeMode(conn *websocket.Conn, userID string
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
close(pingDone)
|
||||||
stdin.Close()
|
stdin.Close()
|
||||||
cmd.Process.Kill()
|
cmd.Process.Kill()
|
||||||
cmd.Wait()
|
cmd.Wait()
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package router
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"strings"
|
"strings"
|
||||||
|
"os"
|
||||||
|
|
||||||
"github.com/engigu/baihu-panel/internal/controllers"
|
"github.com/engigu/baihu-panel/internal/controllers"
|
||||||
"github.com/engigu/baihu-panel/internal/middleware"
|
"github.com/engigu/baihu-panel/internal/middleware"
|
||||||
@@ -31,7 +32,9 @@ type Controllers struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func Setup(c *Controllers) *gin.Engine {
|
func Setup(c *Controllers) *gin.Engine {
|
||||||
gin.SetMode(gin.ReleaseMode)
|
if os.Getenv("GIN_MODE") == "" {
|
||||||
|
gin.SetMode(gin.ReleaseMode)
|
||||||
|
}
|
||||||
router := gin.New()
|
router := gin.New()
|
||||||
router.Use(middleware.GinLogger(), middleware.GinRecovery())
|
router.Use(middleware.GinLogger(), middleware.GinRecovery())
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user