chore: update ws security
This commit is contained in:
@@ -20,9 +20,7 @@ import (
|
||||
)
|
||||
|
||||
var agentUpgrader = websocket.Upgrader{
|
||||
CheckOrigin: func(r *http.Request) bool {
|
||||
return true
|
||||
},
|
||||
CheckOrigin: utils.CheckWSOrigin,
|
||||
}
|
||||
|
||||
// AgentController Agent 控制器
|
||||
|
||||
@@ -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 {
|
||||
// 任务结束,尝试刷新最后一次库内完整内容
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user