package wsc import ( "net/http" "github.com/gin-gonic/gin" ) // HandleWS 一站式 WebSocket 注册。 // // 用法: // // wsc.HandleWS(r, "/ws/chat", svcCtx, func(router *wsc.Router) { // router.On("ping", wsc.Bind(func(ctx *wsc.Context, req *PingReq) (*PingResp, error) { // return &PingResp{Pong: true}, nil // })) // }) func HandleWS(r *gin.RouterGroup, path string, server *Server, register func(router *Router)) { // 预处理:创建路由表(只执行一次) router := NewRouter() if register != nil { register(router) } r.GET(path, func(c *gin.Context) { // 升级前钩子:鉴权失败直接 401 拒绝连接 if server.OnUpgrade != nil { if err := server.OnUpgrade(c); err != nil { c.AbortWithStatusJSON(401, gin.H{"error": err.Error()}) return } } conn, err := server.upgrader.Upgrade(c.Writer, c.Request, nil) if err != nil { return } // 创建会话 session := &Session{ conn: conn, server: server, send: make(chan []byte, 256), closeCh: make(chan struct{}), } // 创建上下文(从 gin 提取元数据) wsCtx := newContext(c, session) session.wsCtx = wsCtx wsCtx.session = session // 连接回调 server.onSessionOpen() if server.OnConnect != nil { server.OnConnect(wsCtx) } // 启动读写协程 go session.writePump() go session.readPump(router) }) } // HandleWSDefault 内部自动创建默认 Server 的注册,等价于 HandleWS(r, path, NewServer(), register)。 func HandleWSDefault(r *gin.RouterGroup, path string, register func(router *Router)) { HandleWS(r, path, NewServer(), register) } // ============================================================ // QuickServer — 快捷方式,省去手动创建 Server // ============================================================ // QuickServer 创建带常用配置的 Server。 func QuickServer(onConnect, onDisconnect func(ctx *Context)) *Server { return &Server{ upgrader: defaultUpgrader, OnConnect: onConnect, OnDisconnect: onDisconnect, rooms: make(map[string]*Room), } } // ============================================================ // CORS 辅助 // ============================================================ // SetCheckOrigin 设置允许的 Origin。 func (s *Server) SetCheckOrigin(origins ...string) { s.upgrader.CheckOrigin = func(r *http.Request) bool { origin := r.Header.Get("Origin") for _, o := range origins { if o == "*" || o == origin { return true } } return false } }