188 lines
4.5 KiB
Go
188 lines
4.5 KiB
Go
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
|
||
}
|