99 lines
2.5 KiB
Go
99 lines
2.5 KiB
Go
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
|
|
}
|
|
}
|