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
+3 -1
View File
@@ -5,6 +5,7 @@ import (
"log/slog"
"os"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
@@ -52,7 +53,7 @@ func main() {
}
// 5. 创建 WebSocket Hub
hub := ws.NewHub(cfg.JWTSecret)
hub := ws.NewHub(cfg.JWTSecret, ws.NewDBDeviceAuthorizer(db), strings.Split(cfg.WSAllowedOrigins, ","))
// 6. 创建并启动 MQTT 服务
mqttSvc := service.NewMQTTService(cfg.MQTT, db, iotdb, hub)
@@ -89,6 +90,7 @@ func main() {
// 11. API 路由组(经过 JWT auth 中间件,白名单路径自动跳过)
api := r.Group("/api/v1")
api.Use(middleware.Auth(cfg))
api.GET("/ws/ticket", hub.HandleTicket)
// 公开路由(auth 白名单中跳过鉴权)
handler.RegisterVideoStreamRoutes(api, transcodeSvc, db, mediaSvc, cfg)
+30 -29
View File
@@ -6,39 +6,40 @@ import (
// Config 全局配置,从环境变量加载
type Config struct {
AppEnv string `env:"APP_ENV" envDefault:"development"`
AllowDevAutoMigrate bool `env:"ALLOW_DEV_AUTOMIGRATE" envDefault:"false"`
PG string `env:"PG" envDefault:"postgresql://postgres:pan@localhost:5432/silk"`
Redis string `env:"REDIS" envDefault:"redis://:pan@localhost:6379"`
JWTSecret string `env:"JWT_SECRET" envDefault:"silk-secret-please-change-me"`
JWTExpiresIn string `env:"JWT_EXPIRES_IN" envDefault:"2h"`
MQTT string `env:"MQTT" envDefault:"mqtt://pan:pan@localhost:1883"`
IoTDBURL string `env:"IOTDB_URL" envDefault:"http://127.0.0.1:18081"`
S3Endpoint string `env:"S3_ENDPOINT" envDefault:"http://100.83.103.1:7480"`
S3AccessKey string `env:"S3_ACCESS_KEY" envDefault:"silk-app"`
S3SecretKey string `env:"S3_SECRET_KEY" envDefault:"Silk-App-Secret-2026!"`
S3Bucket string `env:"S3_BUCKET" envDefault:"silk-video-events"`
S3BucketImages string `env:"S3_BUCKET_IMAGES" envDefault:"silk-images"`
S3Region string `env:"S3_REGION" envDefault:"us-east-1"`
WVPAPIBase string `env:"WVP_API_BASE" envDefault:"http://localhost:18978"`
WVPUsername string `env:"WVP_USERNAME" envDefault:"admin"`
WVPPassword string `env:"WVP_PASSWORD" envDefault:"admin"`
ZLMAPIBase string `env:"ZLM_API_BASE" envDefault:"http://100.83.103.1:8081"`
ZLMSecret string `env:"ZLM_SECRET" envDefault:"su6TiedN2rVAmBbIDX0aa0QTiBJLBdcf"`
RecorderAPIBase string `env:"RECORDER_API_BASE" envDefault:"http://localhost:9090"`
AIServiceBase string `env:"AI_SERVICE_BASE" envDefault:"http://localhost:8000"`
WechatAppID string `env:"WECHAT_APPID" envDefault:""`
WechatSecret string `env:"WECHAT_SECRET" envDefault:""`
WechatTemplateAlarm string `env:"WECHAT_TEMPLATE_ALARM" envDefault:""`
AppEnv string `env:"APP_ENV" envDefault:"development"`
AllowDevAutoMigrate bool `env:"ALLOW_DEV_AUTOMIGRATE" envDefault:"false"`
WSAllowedOrigins string `env:"WS_ALLOWED_ORIGINS" envDefault:"http://localhost:5174,http://localhost:3000,http://127.0.0.1:5174,http://127.0.0.1:3000,http://100.83.103.1:5174,http://100.83.103.1:3000"`
PG string `env:"PG" envDefault:"postgresql://postgres:pan@localhost:5432/silk"`
Redis string `env:"REDIS" envDefault:"redis://:pan@localhost:6379"`
JWTSecret string `env:"JWT_SECRET" envDefault:"silk-secret-please-change-me"`
JWTExpiresIn string `env:"JWT_EXPIRES_IN" envDefault:"2h"`
MQTT string `env:"MQTT" envDefault:"mqtt://pan:pan@localhost:1883"`
IoTDBURL string `env:"IOTDB_URL" envDefault:"http://127.0.0.1:18081"`
S3Endpoint string `env:"S3_ENDPOINT" envDefault:"http://100.83.103.1:7480"`
S3AccessKey string `env:"S3_ACCESS_KEY" envDefault:"silk-app"`
S3SecretKey string `env:"S3_SECRET_KEY" envDefault:"Silk-App-Secret-2026!"`
S3Bucket string `env:"S3_BUCKET" envDefault:"silk-video-events"`
S3BucketImages string `env:"S3_BUCKET_IMAGES" envDefault:"silk-images"`
S3Region string `env:"S3_REGION" envDefault:"us-east-1"`
WVPAPIBase string `env:"WVP_API_BASE" envDefault:"http://localhost:18978"`
WVPUsername string `env:"WVP_USERNAME" envDefault:"admin"`
WVPPassword string `env:"WVP_PASSWORD" envDefault:"admin"`
ZLMAPIBase string `env:"ZLM_API_BASE" envDefault:"http://100.83.103.1:8081"`
ZLMSecret string `env:"ZLM_SECRET" envDefault:"su6TiedN2rVAmBbIDX0aa0QTiBJLBdcf"`
RecorderAPIBase string `env:"RECORDER_API_BASE" envDefault:"http://localhost:9090"`
AIServiceBase string `env:"AI_SERVICE_BASE" envDefault:"http://localhost:8000"`
WechatAppID string `env:"WECHAT_APPID" envDefault:""`
WechatSecret string `env:"WECHAT_SECRET" envDefault:""`
WechatTemplateAlarm string `env:"WECHAT_TEMPLATE_ALARM" envDefault:""`
WechatTemplateInspection string `env:"WECHAT_TEMPLATE_INSPECTION" envDefault:""`
QWeatherAPIKey string `env:"QWEATHER_API_KEY" envDefault:""`
QWeatherLocation string `env:"QWEATHER_LOCATION" envDefault:""`
QWeatherIntervalMin int `env:"QWEATHER_INTERVAL_MIN" envDefault:"30"`
InternalAPIKey string `env:"INTERNAL_API_KEY" envDefault:"silk-internal-2026"`
Port int `env:"PORT" envDefault:"3000"`
DefaultAdminUsername string `env:"DEFAULT_ADMIN_USERNAME" envDefault:"admin"`
DefaultAdminPassword string `env:"DEFAULT_ADMIN_PASSWORD" envDefault:"silk@123"`
DefaultAdminEmail string `env:"DEFAULT_ADMIN_EMAIL" envDefault:"admin@silk.local"`
InternalAPIKey string `env:"INTERNAL_API_KEY" envDefault:"silk-internal-2026"`
Port int `env:"PORT" envDefault:"3000"`
DefaultAdminUsername string `env:"DEFAULT_ADMIN_USERNAME" envDefault:"admin"`
DefaultAdminPassword string `env:"DEFAULT_ADMIN_PASSWORD" envDefault:"silk@123"`
DefaultAdminEmail string `env:"DEFAULT_ADMIN_EMAIL" envDefault:"admin@silk.local"`
}
// Load 从环境变量加载配置
+59
View File
@@ -0,0 +1,59 @@
package ws
import (
"errors"
"fmt"
"silk-server-go/internal/model"
"gorm.io/gorm"
)
// DeviceAuthorizer 校验用户是否有权读取指定设备。
type DeviceAuthorizer interface {
CanReadDevice(userID, deviceKey string) (bool, error)
}
// DBDeviceAuthorizer 使用现有 RBAC 权限判断设备读取权;当前未做用户级资源 ACL。
type DBDeviceAuthorizer struct {
db *gorm.DB
}
func NewDBDeviceAuthorizer(db *gorm.DB) *DBDeviceAuthorizer {
return &DBDeviceAuthorizer{db: db}
}
func (a *DBDeviceAuthorizer) CanReadDevice(userID, deviceKey string) (bool, error) {
if deviceKey == "" {
return false, nil
}
var user model.User
if err := a.db.Where("id = ?", userID).First(&user).Error; err != nil {
return false, err
}
if user.Role == model.RoleAdmin {
return true, nil
}
var count int64
err := a.db.Table("role_permissions").
Joins("JOIN permissions ON permissions.id = role_permissions.permission_id").
Where("role_permissions.role = ? AND permissions.code = ?", user.Role, "device:read").
Count(&count).Error
return count > 0, err
}
func (h *Hub) authorizeSubscription(userID, deviceKey string) error {
if h.authorizer == nil {
return errors.New("device authorizer not configured")
}
ok, err := h.authorizer.CanReadDevice(userID, deviceKey)
if err != nil {
return err
}
if !ok {
return fmt.Errorf("device subscription denied: %s", deviceKey)
}
return nil
}
+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)
}
+120
View File
@@ -0,0 +1,120 @@
package ws
import (
"errors"
"testing"
"silk-server-go/internal/service"
)
type fakeAuthorizer struct {
allowed map[string]bool
}
func (f fakeAuthorizer) CanReadDevice(userID, deviceKey string) (bool, error) {
if f.allowed == nil {
return false, nil
}
return f.allowed[userID+"|"+deviceKey], nil
}
func newTestHub() *Hub {
return NewHub(
"test-secret",
fakeAuthorizer{allowed: map[string]bool{"user-1|device-a": true}},
[]string{"http://localhost:5174"},
)
}
func TestWebSocketOriginRejectsUnknownOrigin(t *testing.T) {
hub := newTestHub()
if hub.originAllowed("http://evil.example") {
t.Fatal("unknown origin should be rejected")
}
if !hub.originAllowed("http://localhost:5174") {
t.Fatal("configured origin should be allowed")
}
}
func TestWebSocketTicketIsOneTimeAndExpires(t *testing.T) {
hub := newTestHub()
ticket := hub.IssueTicket("user-1", "admin", "User One")
if ticket == "" {
t.Fatal("expected non-empty ticket")
}
info, ok := hub.consumeTicket(ticket)
if !ok || info.userID != "user-1" {
t.Fatalf("first ticket consume failed: ok=%v info=%+v", ok, info)
}
if _, ok := hub.consumeTicket(ticket); ok {
t.Fatal("ticket should be single-use")
}
}
func TestWebSocketRejectsLongLivedJWTQuery(t *testing.T) {
hub := newTestHub()
if hub.ticketFromQuery("") != "" {
t.Fatal("missing ticket should not be accepted")
}
if hub.ticketFromQuery("token=eyJhbGciOiJIUzI1NiJ9.abc") != "" {
t.Fatal("long-lived token query must not be accepted")
}
}
func TestWebSocketDeviceAuthorizerRejectsUnauthorized(t *testing.T) {
hub := newTestHub()
if err := hub.authorizeSubscription("user-1", "device-a"); err != nil {
t.Fatalf("authorized device rejected: %v", err)
}
if err := hub.authorizeSubscription("user-1", "device-b"); err == nil {
t.Fatal("unauthorized device should be rejected")
}
}
func TestBroadcastTelemetryDoesNotLeakGlobal(t *testing.T) {
hub := newTestHub()
subscribed := &client{rooms: map[string]bool{"device:device-a": true}, send: make(chan []byte, 1)}
unsubscribed := &client{rooms: map[string]bool{}, send: make(chan []byte, 1)}
hub.clients[subscribed] = true
hub.clients[unsubscribed] = true
hub.BroadcastTelemetry("device-a", map[string]interface{}{"value": 1})
if len(subscribed.send) == 0 {
t.Fatal("subscribed client should receive telemetry")
}
if len(unsubscribed.send) != 0 {
t.Fatal("unsubscribed client must not receive global telemetry")
}
}
func TestBroadcastAlarmDoesNotLeakGlobal(t *testing.T) {
hub := newTestHub()
subscribed := &client{rooms: map[string]bool{"device:device-a": true}, send: make(chan []byte, 1)}
unsubscribed := &client{rooms: map[string]bool{}, send: make(chan []byte, 1)}
hub.clients[subscribed] = true
hub.clients[unsubscribed] = true
hub.BroadcastAlarm(service.AlarmEvent{DeviceKey: "device-a", Code: "high"})
if len(unsubscribed.send) != 0 {
t.Fatal("unsubscribed client must not receive global alarm")
}
}
type brokenAuthorizer struct{}
func (brokenAuthorizer) CanReadDevice(userID, deviceKey string) (bool, error) {
return false, errors.New("authorizer unavailable")
}
func TestWebSocketAuthorizerErrorIsDenied(t *testing.T) {
hub := NewHub("test-secret", brokenAuthorizer{}, nil)
if err := hub.authorizeSubscription("user-1", "device-a"); err == nil {
t.Fatal("authorizer error should deny subscription")
}
if _, ok := hub.consumeTicket("missing"); ok {
t.Fatal("missing ticket should not be consumed")
}
}
+57
View File
@@ -0,0 +1,57 @@
package ws
import (
"crypto/rand"
"encoding/hex"
"net/url"
"time"
)
const wsTicketTTL = 60 * time.Second
type wsTicket struct {
userID string
username string
role string
expiresAt time.Time
used bool
}
// IssueTicket 签发一次性 WebSocket ticket,避免长期 JWT 进入 URL。
func (h *Hub) IssueTicket(userID, username, role string) string {
buf := make([]byte, 32)
if _, err := rand.Read(buf); err != nil {
return ""
}
ticket := hex.EncodeToString(buf)
h.ticketsMu.Lock()
h.tickets[ticket] = &wsTicket{
userID: userID,
username: username,
role: role,
expiresAt: time.Now().Add(wsTicketTTL),
}
h.ticketsMu.Unlock()
return ticket
}
func (h *Hub) consumeTicket(ticket string) (wsTicket, bool) {
h.ticketsMu.Lock()
defer h.ticketsMu.Unlock()
info, ok := h.tickets[ticket]
if !ok || info.used || time.Now().After(info.expiresAt) {
return wsTicket{}, false
}
info.used = true
return *info, true
}
func (h *Hub) ticketFromQuery(rawQuery string) string {
values, err := url.ParseQuery(rawQuery)
if err != nil {
return ""
}
return values.Get("ticket")
}