feat: 修复 WebSocket 越权与 AI 流 SSRF
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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 从环境变量加载配置
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
Reference in New Issue
Block a user