feat: 修复 WebSocket 越权与 AI 流 SSRF
This commit is contained in:
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user