Files

188 lines
4.5 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}