feat: 修复 WebSocket 越权与 AI 流 SSRF

This commit is contained in:
weijuesen
2026-08-14 00:45:22 +08:00
parent 1a29def480
commit 839ba91354
17 changed files with 614 additions and 136 deletions
+86 -58
View File
@@ -4,82 +4,103 @@ import (
"encoding/json"
"log/slog"
"net/http"
"strings"
"sync"
"time"
"github.com/gin-gonic/gin"
"github.com/golang-jwt/jwt/v5"
"github.com/gorilla/websocket"
"silk-server-go/internal/service"
)
var upgrader = websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool { return true },
ReadBufferSize: 1024,
WriteBufferSize: 1024,
}
// client WebSocket 客户端
type client struct {
conn *websocket.Conn
rooms map[string]bool // 订阅的房间(device:<deviceKey>
send chan []byte
conn *websocket.Conn
userID string
username string
role string
rooms map[string]bool // 订阅的房间(device:<deviceKey>
send chan []byte
}
// Hub WebSocket 中心,管理客户端和房间
type Hub struct {
jwtSecret string
clients map[*client]bool
mu sync.RWMutex
authorizer DeviceAuthorizer
allowedOrigins []string
clients map[*client]bool
mu sync.RWMutex
tickets map[string]*wsTicket
ticketsMu sync.Mutex
}
// NewHub 创建 Hub
func NewHub(jwtSecret string) *Hub {
func NewHub(jwtSecret string, authorizer DeviceAuthorizer, allowedOrigins []string) *Hub {
return &Hub{
jwtSecret: jwtSecret,
clients: make(map[*client]bool),
authorizer: authorizer,
allowedOrigins: allowedOrigins,
clients: make(map[*client]bool),
tickets: make(map[string]*wsTicket),
}
}
func (h *Hub) originAllowed(origin string) bool {
if origin == "" {
return true
}
for _, allowed := range h.allowedOrigins {
if origin == allowed {
return true
}
}
return false
}
// HandleTicket 签发一次性 WebSocket ticket,需在 JWT 鉴权路由后注册。
func (h *Hub) HandleTicket(c *gin.Context) {
userVal, ok := c.Get("user")
if !ok {
c.JSON(http.StatusUnauthorized, gin.H{"error": "未认证"})
return
}
userMap, ok := userVal.(map[string]interface{})
if !ok {
c.JSON(http.StatusUnauthorized, gin.H{"error": "用户信息无效"})
return
}
userID, _ := userMap["sub"].(string)
username, _ := userMap["username"].(string)
role, _ := userMap["role"].(string)
if userID == "" {
c.JSON(http.StatusUnauthorized, gin.H{"error": "用户信息无效"})
return
}
c.JSON(http.StatusOK, gin.H{
"ticket": h.IssueTicket(userID, username, role),
"expiresIn": int(wsTicketTTL.Seconds()),
})
}
// HandleWebSocket 处理 WebSocket 连接(Gin handler
func (h *Hub) HandleWebSocket(c *gin.Context) {
// 验证 JWT(从 query.auth.token / query.token / header.Authorization 获取)
tokenStr := ""
if t := c.Query("auth.token"); t != "" {
tokenStr = t
} else if t := c.Query("token"); t != "" {
tokenStr = t
} else if auth := c.GetHeader("Authorization"); strings.HasPrefix(auth, "Bearer ") {
tokenStr = auth[7:]
}
if tokenStr == "" {
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err == nil {
conn.WriteJSON(map[string]interface{}{"event": "auth.fail", "data": map[string]bool{"ok": false}})
conn.Close()
}
if !h.originAllowed(c.GetHeader("Origin")) {
c.JSON(http.StatusForbidden, gin.H{"error": "origin not allowed"})
return
}
// 验证 token
claims := jwt.MapClaims{}
token, err := jwt.ParseWithClaims(tokenStr, claims, func(t *jwt.Token) (interface{}, error) {
return []byte(h.jwtSecret), nil
}, jwt.WithValidMethods([]string{"HS256"}))
if err != nil || !token.Valid {
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err == nil {
conn.WriteJSON(map[string]interface{}{"event": "auth.fail", "data": map[string]bool{"ok": false}})
conn.Close()
}
ticket := h.ticketFromQuery(c.Request.URL.RawQuery)
info, ok := h.consumeTicket(ticket)
if !ok {
c.JSON(http.StatusUnauthorized, gin.H{"error": "invalid or expired ws ticket"})
return
}
// 升级为 WebSocket
upgrader := websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool { return true },
ReadBufferSize: 1024,
WriteBufferSize: 1024,
}
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
slog.Warn("WebSocket 升级失败", "err", err)
@@ -87,9 +108,12 @@ func (h *Hub) HandleWebSocket(c *gin.Context) {
}
cl := &client{
conn: conn,
rooms: make(map[string]bool),
send: make(chan []byte, 256),
conn: conn,
userID: info.userID,
username: info.username,
role: info.role,
rooms: make(map[string]bool),
send: make(chan []byte, 256),
}
h.mu.Lock()
@@ -98,8 +122,12 @@ func (h *Hub) HandleWebSocket(c *gin.Context) {
// 发送 auth.ok
h.sendJSON(cl, "auth.ok", map[string]interface{}{
"ok": true,
"user": claims,
"ok": true,
"user": map[string]interface{}{
"sub": info.userID,
"username": info.username,
"role": info.role,
},
})
go h.readPump(cl)
@@ -145,6 +173,14 @@ func (h *Hub) readPump(cl *client) {
switch data.Event {
case "subscribe.device":
if err := h.authorizeSubscription(cl.userID, data.DeviceKey); err != nil {
h.sendJSON(cl, "subscribe.denied", map[string]interface{}{
"ok": false,
"deviceKey": data.DeviceKey,
"error": err.Error(),
})
continue
}
room := "device:" + data.DeviceKey
h.mu.Lock()
cl.rooms[room] = true
@@ -204,12 +240,9 @@ func (h *Hub) BroadcastTelemetry(deviceKey string, data interface{}) {
room := "device:" + deviceKey
for cl := range h.clients {
// 推送到设备房间
if cl.rooms[room] {
h.sendJSON(cl, "telemetry", data)
}
// 全局推送
h.sendJSON(cl, "telemetry.all", data)
}
}
@@ -223,14 +256,10 @@ func (h *Hub) BroadcastAlarm(event service.AlarmEvent) {
for cl := range h.clients {
if isRecovery {
// 恢复通知
h.sendJSON(cl, "alarm.recovery", event)
if event.DeviceKey != "" && cl.rooms[room] {
h.sendJSON(cl, "alarm.device.recovery", event)
}
} else {
// 告警触发
h.sendJSON(cl, "alarm", event)
if event.DeviceKey != "" && cl.rooms[room] {
h.sendJSON(cl, "alarm.device", event)
}
@@ -246,7 +275,6 @@ func (h *Hub) BroadcastDeviceStatus(deviceKey string, status string) {
room := "device:" + deviceKey
data := map[string]interface{}{"deviceKey": deviceKey, "status": status}
for cl := range h.clients {
h.sendJSON(cl, "device.status", data)
if cl.rooms[room] {
h.sendJSON(cl, "device.status.device", data)
}