package wsc import ( "context" "sync" "time" "github.com/gin-gonic/gin" "github.com/google/uuid" ) // Context WebSocket 上下文,实现 context.Context 接口。 // 在连接建立时从 gin.Context 提取元数据,贯穿整个连接生命周期。 type Context struct { // === 从 gin 提取的元数据 === // 公开字段仅供读取,正常业务流程不应修改。 SessionID string `json:"sessionId"` ClientIP string `json:"clientIp"` UserAgent string `json:"userAgent"` Path string `json:"path"` TraceID string `json:"traceId"` // UID 由认证中间件注入,业务层可直接读取。 UID int64 `json:"uid,omitempty"` // === 内部字段 === ctx context.Context cancel context.CancelFunc session *Session mu sync.RWMutex values map[string]any } // newContext 从 gin.Context 创建 ws 上下文。 func newContext(c *gin.Context, session *Session) *Context { ctx, cancel := context.WithCancel(context.Background()) wsCtx := &Context{ SessionID: uuid.New().String(), ClientIP: c.ClientIP(), UserAgent: c.GetHeader("User-Agent"), Path: c.FullPath(), TraceID: c.GetHeader("X-Trace-Id"), ctx: ctx, cancel: cancel, session: session, values: make(map[string]any), } // 透传 gin 上下文中的自定义值(如 OnUpgrade 中设置的 userId/userName) // gin v1.12.0 的 Keys 为 map[any]any,需将 key 断言为 string。 if c.Keys != nil { for k, v := range c.Keys { if key, ok := k.(string); ok { wsCtx.values[key] = v } } } return wsCtx } // ============================================================ // context.Context 接口实现 // ============================================================ func (c *Context) Deadline() (time.Time, bool) { return c.ctx.Deadline() } func (c *Context) Done() <-chan struct{} { return c.ctx.Done() } func (c *Context) Err() error { return c.ctx.Err() } func (c *Context) Value(key any) any { if s, ok := key.(string); ok { c.mu.RLock() v, exists := c.values[s] c.mu.RUnlock() if exists { return v // 含 nil } } return c.ctx.Value(key) } // ============================================================ // 写回方法 // ============================================================ // Write 发送 JSON 消息(原始数据,不带 type 包装)。 func (c *Context) Write(v any) error { return c.session.writeJSON(v) } // WriteMessage 发送带 type 的结构化消息。 func (c *Context) WriteMessage(msgType string, data any) error { return c.Write(messageOut(msgType, data)) } // WriteError 发送错误消息。 func (c *Context) WriteError(code int, msg string) error { return c.WriteMessage(WsActionError, map[string]any{ "code": code, "message": msg, }) } // ============================================================ // 自定义值存取 // ============================================================ func (c *Context) Set(key string, value any) { c.mu.Lock() c.values[key] = value c.mu.Unlock() } func (c *Context) Get(key string) (any, bool) { c.mu.RLock() v, ok := c.values[key] c.mu.RUnlock() return v, ok } func (c *Context) GetString(key string) string { v, _ := c.Get(key) s, _ := v.(string) return s } func (c *Context) GetInt64(key string) int64 { v, _ := c.Get(key) n, _ := v.(int64) return n } // ============================================================ // 房间操作 // ============================================================ func (c *Context) JoinRoom(name string) { c.session.server.GetOrCreateRoom(name).Join(c) } func (c *Context) LeaveRoom(name string) { room := c.session.server.GetRoom(name) if room != nil { room.Leave(c) } } func (c *Context) BroadcastToRoom(roomName string, msgType string, data any) error { room := c.session.server.GetRoom(roomName) if room == nil { return nil } return room.BroadcastRaw(c.SessionID, messageOut(msgType, data)) } func (c *Context) GetRoom(name string) *Room { return c.session.server.GetRoom(name) } // ============================================================ // 生命周期 // ============================================================ func (c *Context) Close() { c.session.close() } func (c *Context) cancelCtx() { c.cancel() } // ============================================================ // 内部辅助 // ============================================================ // messageOut 构造写出的消息 map(单次序列化)。 func messageOut(msgType string, data any) map[string]any { m := map[string]any{"action": msgType} if data != nil { m["payload"] = data } return m }