init: 通用库抽离——httpx/logger/wsc/wscclient/jwtx(合并分叉副本); 顶层包布局; module path git.zeroonesoft.cn/golib/zogo

This commit is contained in:
w11
2026-09-19 19:04:01 +08:00
commit f0f5421262
17 changed files with 2336 additions and 0 deletions
+187
View File
@@ -0,0 +1,187 @@
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
}
+103
View File
@@ -0,0 +1,103 @@
package wsc
import (
"encoding/json"
"sync"
"github.com/sirupsen/logrus"
)
// Room 消息房间,用于业务隔离和群组广播。
type Room struct {
name string
members map[string]*Session
mu sync.RWMutex
onEmpty func(name string) // 房间空时回调
}
func newRoom(name string) *Room {
return &Room{
name: name,
members: make(map[string]*Session),
}
}
func (r *Room) Name() string { return r.name }
func (r *Room) Len() int {
r.mu.RLock()
defer r.mu.RUnlock()
return len(r.members)
}
func (r *Room) Join(ctx *Context) {
r.mu.Lock()
r.members[ctx.SessionID] = ctx.session
r.mu.Unlock()
}
func (r *Room) Leave(ctx *Context) {
r.mu.Lock()
delete(r.members, ctx.SessionID)
empty := len(r.members) == 0
r.mu.Unlock()
if empty && r.onEmpty != nil {
r.onEmpty(r.name)
}
}
// Broadcast 向房间成员广播 Message(json.Marshal 后发原始字节)。
func (r *Room) Broadcast(excludeSessionID string, msg Message) error {
m := map[string]any{"action": msg.Action}
if len(msg.Payload) > 0 {
// Payload 已是 json.RawMessage,直接复用
m["payload"] = msg.Payload
}
raw, err := json.Marshal(m)
if err != nil {
return err
}
return r.broadcastRaw(excludeSessionID, raw)
}
// BroadcastRaw 向房间成员广播已序列化的消息(性能优化)。
func (r *Room) BroadcastRaw(excludeSessionID string, data map[string]any) error {
raw, err := json.Marshal(data)
if err != nil {
return err
}
return r.broadcastRaw(excludeSessionID, raw)
}
// broadcastRaw 向房间成员发送原始字节(锁外已序列化)。
func (r *Room) broadcastRaw(excludeSessionID string, raw []byte) error {
r.mu.RLock()
defer r.mu.RUnlock()
for id, s := range r.members {
if id == excludeSessionID {
continue
}
if err := s.sendRaw(raw); err != nil {
logrus.Errorf("[wsc] broadcast send failed: session=%s, err=%v", id[:8], err)
}
}
return nil
}
// BroadcastAll 向房间所有成员广播。
func (r *Room) BroadcastAll(msg Message) error {
return r.Broadcast("", msg)
}
// SendTo 向房间内指定成员发送消息。
func (r *Room) SendTo(sessionID string, msg Message) error {
r.mu.RLock()
s, ok := r.members[sessionID]
r.mu.RUnlock()
if !ok {
return nil
}
return s.writeJSON(msg)
}
+357
View File
@@ -0,0 +1,357 @@
package wsc
import (
"encoding/json"
"reflect"
"github.com/sirupsen/logrus"
)
// ============================================================
// 框架级消息 action 常量
// 业务消息(ping/pong/register 等)由各模块自行定义,不放在框架层。
// ============================================================
const (
// WsActionError 错误消息(框架统一回包 action)
WsActionError = "error"
)
// ============================================================
// MessageHandler — 消息处理器签名
// ============================================================
// MessageHandler 业务消息处理函数。
// ctx: WS 上下文 data: 已解析的 JSON data 字段。
// 返回的 any 会自动 JSON 序列化后通过 ctx.Write() 发回。
type MessageHandler func(ctx *Context, data json.RawMessage) (any, error)
// ============================================================
// MiddlewareFunc — 中间件签名
// ============================================================
// MiddlewareFunc 消息级中间件。
// 返回 error 时中断链路,错误消息自动发送给客户端。
type MiddlewareFunc func(ctx *Context, msg *Message) error
// ============================================================
// Router — 消息路由器
// ============================================================
// Router 消息路由器,按 type 字段分发到不同的 MessageHandler。
// 支持中间件链(类似 gin)。
type Router struct {
routes map[string]MessageHandler
middlewares []MiddlewareFunc
}
// NewRouter 创建路由器。
func NewRouter() *Router {
return &Router{
routes: make(map[string]MessageHandler),
}
}
// On 注册指定 type 的消息处理器。
// handler 签名:func(ctx *Context, data json.RawMessage) (any, error)
func (r *Router) On(msgType string, handler MessageHandler) {
r.routes[msgType] = handler
}
// Use 添加消息级中间件。按添加顺序执行。
func (r *Router) Use(mw ...MiddlewareFunc) {
r.middlewares = append(r.middlewares, mw...)
}
// dispatch 内部消息分发。先走中间件链,再走路由。
func (r *Router) dispatch(ctx *Context, msg *Message) {
// 中间件链
for _, mw := range r.middlewares {
if err := mw(ctx, msg); err != nil {
// 中间件阻断
if wErr, ok := err.(*Error); ok {
ctx.WriteError(wErr.Code, wErr.Message)
} else {
ctx.WriteError(403, err.Error())
}
return
}
}
// 路由分发
handler, ok := r.routes[msg.Action]
if !ok {
logrus.Warnf("[wsc] unknown message type: %s, session=%s, ip=%s", msg.Action, ctx.SessionID, ctx.ClientIP)
ctx.WriteError(404, "unknown message type: "+msg.Action)
return
}
resp, err := handler(ctx, msg.Payload)
if err != nil {
if wErr, ok := err.(*Error); ok {
ctx.WriteError(wErr.Code, wErr.Message)
} else {
ctx.WriteError(500, err.Error())
}
return
}
// 有返回值才写回
if resp != nil {
if rwa, ok := resp.(*responseWithAction); ok {
ctx.Write(messageOut(rwa.action, rwa.data))
} else {
ctx.Write(messageOut(msg.Action+".resp", resp))
}
}
}
// ============================================================
// Bind — 泛型辅助,自动 JSON 反序列化
// ============================================================
// Bind 将带类型的业务函数包装为 MessageHandler。
// 自动处理 JSON 反序列化和序列化。
//
// 用法:
//
// router.On("ping", wsc.Bind(func(ctx *wsc.Context, req *PingReq) (*PingResp, error) {
// return &PingResp{Pong: true}, nil
// }))
//
// responseWithAction 携带自定义响应 action 的返回包装,由 BindAs 使用。
type responseWithAction struct {
action string
data any
}
// errType error 接口的 reflect.Type,用于处理器签名校验。
var errType = reflect.TypeOf((*error)(nil)).Elem()
// isNilValue 判断 reflect.Value 是否为 nil。
// 仅对可为 nil 的类型执行 IsNil,其余类型(值类型响应)一律视为非 nil,
// 避免 reflect.Value.IsNil 对值类型 panic。
func isNilValue(v reflect.Value) bool {
switch v.Kind() {
case reflect.Chan, reflect.Func, reflect.Interface, reflect.Map, reflect.Ptr, reflect.Slice:
return v.IsNil()
default:
return false
}
}
// bindImpl 自动 JSON 反序列化并调用业务函数。
// respAction 为空时回包 action 为 msg.Action+".resp"(Bind 行为);
// 非空时回包 action 为指定值(BindAs 行为,用于兼容客户端约定的固定响应类型)。
func bindImpl(respAction string, fn any) MessageHandler {
fnVal := reflect.ValueOf(fn)
fnType := fnVal.Type()
// 验证签名:func(ctx *Context, req T) (R, error)
if fnType.Kind() != reflect.Func {
panic("wsc.Bind: argument must be a function")
}
if fnType.NumIn() != 2 {
panic("wsc.Bind: function must have 2 parameters (ctx, req)")
}
if fnType.NumOut() != 2 {
panic("wsc.Bind: function must return 2 values (resp, error)")
}
reqType := fnType.In(1) // 请求参数类型
// 构造可寻址的零值实例(始终为指针 *T)。
// 不能用 reflect.New(t).Elem().Interface():经 Interface() 拷贝后丢失可寻址性,
// 值类型请求参数时后续 .Addr() 会 panic。
newReq := func() reflect.Value {
t := reqType
if t.Kind() == reflect.Ptr {
return reflect.New(t.Elem())
}
return reflect.New(t)
}
return func(ctx *Context, data json.RawMessage) (any, error) {
req := newReq()
if len(data) > 0 {
if err := json.Unmarshal(data, req.Interface()); err != nil {
return nil, NewError(400, "invalid request data: "+err.Error())
}
}
// reqType 为指针类型时传 *T,为值类型时解引用传 T。
var reqArg reflect.Value
if reqType.Kind() == reflect.Ptr {
reqArg = req
} else {
reqArg = req.Elem()
}
results := fnVal.Call([]reflect.Value{
reflect.ValueOf(ctx),
reqArg,
})
var resp any
if !isNilValue(results[0]) {
resp = results[0].Interface()
}
var err error
if !results[1].IsNil() {
err = results[1].Interface().(error)
}
if err != nil {
return nil, err
}
if resp == nil {
return nil, nil
}
if respAction == "" {
return resp, nil
}
return &responseWithAction{action: respAction, data: resp}, nil
}
}
// Bind 将带类型的业务函数包装为 MessageHandler,自动 JSON 反序列化。
// 回包 action 为 msg.Action + ".resp"。
func Bind(fn any) MessageHandler {
return bindImpl("", fn)
}
// BindAs 与 Bind 相同,但允许指定响应 action(覆盖默认的 action+".resp")。
// 适用于客户端约定了固定响应类型(如 "pong"、"register_resp")的场景。
//
// 用法:
//
// router.On("ping", wsc.BindAs("pong", func(ctx *wsc.Context, req *PingReq) (*PingResp, error) {
// return &PingResp{Pong: true}, nil
// }))
func BindAs(respAction string, fn any) MessageHandler {
return bindImpl(respAction, fn)
}
// ============================================================
// BindNoResp — 无返回值的处理器
// ============================================================
// BindNoResp 包装无返回值的业务函数。
// 签名:func(ctx *Context, req T) error,唯一返回值须为 error。
func BindNoResp(fn any) MessageHandler {
fnVal := reflect.ValueOf(fn)
fnType := fnVal.Type()
// 验证签名:func(ctx *Context, req T) error
if fnType.Kind() != reflect.Func {
panic("wsc.BindNoResp: argument must be a function")
}
if fnType.NumIn() != 2 {
panic("wsc.BindNoResp: function must have 2 parameters (ctx, req)")
}
if fnType.NumOut() != 1 || fnType.Out(0) != errType {
panic("wsc.BindNoResp: function must return 1 value (error)")
}
reqType := fnType.In(1)
// 构造可寻址的零值实例,见 bindImpl 同名逻辑。
newReq := func() reflect.Value {
t := reqType
if t.Kind() == reflect.Ptr {
return reflect.New(t.Elem())
}
return reflect.New(t)
}
return func(ctx *Context, data json.RawMessage) (any, error) {
req := newReq()
if len(data) > 0 {
if err := json.Unmarshal(data, req.Interface()); err != nil {
return nil, NewError(400, "invalid request data: "+err.Error())
}
}
// reqType 为指针类型时传 *T,为值类型时解引用传 T。
var reqArg reflect.Value
if reqType.Kind() == reflect.Ptr {
reqArg = req
} else {
reqArg = req.Elem()
}
results := fnVal.Call([]reflect.Value{
reflect.ValueOf(ctx),
reqArg,
})
if !results[0].IsNil() {
return nil, results[0].Interface().(error)
}
return nil, nil
}
}
// ============================================================
// Error — 业务错误
// ============================================================
// Error 业务错误,自动序列化发送给客户端。
type Error struct {
Code int `json:"code"`
Message string `json:"message"`
}
func (e *Error) Error() string { return e.Message }
// NewError 创建业务错误。
func NewError(code int, msg string) *Error {
return &Error{Code: code, Message: msg}
}
// ============================================================
// 内置中间件
// ============================================================
// AuthMiddleware 认证中间件工厂。
// tokenExtractor 从消息中提取 token;validator 验证 token 并返回 uid。
// 验证通过后 uid 写入 ctx.UID。
func AuthMiddleware(
tokenExtractor func(msg *Message) string,
validator func(token string) (uid int64, err error),
) MiddlewareFunc {
return func(ctx *Context, msg *Message) error {
token := tokenExtractor(msg)
if token == "" {
return NewError(401, "token required")
}
uid, err := validator(token)
if err != nil {
return NewError(401, "invalid token: "+err.Error())
}
ctx.UID = uid
ctx.Set("uid", uid)
return nil
}
}
// RecoveryMiddleware 恢复中间件,捕获 panic。
func RecoveryMiddleware() MiddlewareFunc {
return func(ctx *Context, msg *Message) (err error) {
defer func() {
if r := recover(); r != nil {
logrus.Errorf("[wsc] panic recovered: %v, session=%s, type=%s", r, ctx.SessionID, msg.Action)
err = NewError(500, "internal server error")
}
}()
return nil
}
}
// LoggerMiddleware 日志中间件。
func LoggerMiddleware() MiddlewareFunc {
return func(ctx *Context, msg *Message) error {
logrus.Debugf("[wsc] %s | %s | %s | %s", ctx.ClientIP, ctx.SessionID, ctx.Path, msg.Action)
return nil
}
}
+123
View File
@@ -0,0 +1,123 @@
package wsc
import (
"net/http"
"sync"
"sync/atomic"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
)
// ============================================================
// 默认配置
// ============================================================
var defaultUpgrader = websocket.Upgrader{
ReadBufferSize: 4096,
WriteBufferSize: 4096,
CheckOrigin: func(r *http.Request) bool {
return true // 生产环境请限制 Origin
},
}
// ============================================================
// Server — 连接管理器
// ============================================================
// Server 管理 WebSocket 连接的全局状态。
type Server struct {
upgrader websocket.Upgrader
// 连接生命周期回调
OnConnect func(ctx *Context)
OnDisconnect func(ctx *Context)
// OnUpgrade 升级前钩子:可用于鉴权(如校验 URL token),返回 error 即拒绝连接(401)。
// 钩子中通过 c.Set(k, v) 写入的值会自动透传到 *Context(见 newContext)。
OnUpgrade func(c *gin.Context) error
// 房间管理
mu sync.RWMutex
rooms map[string]*Room
// 连接统计
connCount atomic.Int64
}
// NewServer 创建 Server。
func NewServer() *Server {
return &Server{
upgrader: defaultUpgrader,
rooms: make(map[string]*Room),
}
}
// SetUpgrader 自定义升级器。
func (s *Server) SetUpgrader(u websocket.Upgrader) {
s.upgrader = u
}
// ============================================================
// 连接统计
// ============================================================
// ActiveCount 返回当前活跃连接数。
func (s *Server) ActiveCount() int64 {
return s.connCount.Load()
}
func (s *Server) onSessionOpen() {
s.connCount.Add(1)
}
func (s *Server) onSessionClosed() {
s.connCount.Add(-1)
}
// ============================================================
// 房间管理
// ============================================================
// GetOrCreateRoom 获取或创建房间。
func (s *Server) GetOrCreateRoom(name string) *Room {
s.mu.Lock()
defer s.mu.Unlock()
if room, ok := s.rooms[name]; ok {
return room
}
room := newRoom(name)
room.onEmpty = func(name string) {
s.mu.Lock()
defer s.mu.Unlock()
// 二次判空:防止 leave→onEmpty 之间新成员加入
if r, exists := s.rooms[name]; exists && r.Len() == 0 {
delete(s.rooms, name)
}
}
s.rooms[name] = room
return room
}
// GetRoom 获取房间(不存在返回 nil)。
func (s *Server) GetRoom(name string) *Room {
s.mu.RLock()
defer s.mu.RUnlock()
return s.rooms[name]
}
// RemoveRoom 移除房间。
func (s *Server) RemoveRoom(name string) {
s.mu.Lock()
defer s.mu.Unlock()
delete(s.rooms, name)
}
// RoomCount 返回当前房间数。
func (s *Server) RoomCount() int {
s.mu.RLock()
defer s.mu.RUnlock()
return len(s.rooms)
}
+172
View File
@@ -0,0 +1,172 @@
package wsc
import (
"encoding/json"
"errors"
"io"
"sync"
"time"
"github.com/gorilla/websocket"
"github.com/sirupsen/logrus"
)
// ============================================================
// 消息帧格式
// ============================================================
// Message 统一消息结构(客户端 <-> 服务端)。
type Message struct {
Action string `json:"action"`
Payload json.RawMessage `json:"payload,omitempty"`
}
// ============================================================
// 错误定义
// ============================================================
var (
ErrConnectionClosed = errors.New("wsc: connection closed")
ErrSendBufferFull = errors.New("wsc: send buffer full")
)
// ============================================================
// Session — 单连接会话
// ============================================================
const (
writeWait = 10 * time.Second
pongWait = 60 * time.Second
pingPeriod = (pongWait * 9) / 10
maxMessageSize = 65536
)
// Session 管理单个 WebSocket 连接。
type Session struct {
wsCtx *Context
conn *websocket.Conn
server *Server
send chan []byte
closeCh chan struct{}
once sync.Once
}
// writeJSON 序列化并发送 JSON 到写通道(单次 marshal)。
func (s *Session) writeJSON(v any) error {
data, err := json.Marshal(v)
if err != nil {
return err
}
return s.sendRaw(data)
}
// sendRaw 将已序列化的字节写入发送通道。
func (s *Session) sendRaw(data []byte) error {
select {
case s.send <- data:
return nil
case <-s.closeCh:
return ErrConnectionClosed
default:
logrus.Warnf("[wsc] send buffer full, dropping message for session=%s", s.wsCtx.SessionID[:8])
return ErrSendBufferFull
}
}
// writePump 写协程,从 send 通道取数据写入 WebSocket。
func (s *Session) writePump() {
ticker := time.NewTicker(pingPeriod)
defer func() {
ticker.Stop()
s.close()
}()
for {
select {
case message := <-s.send:
s.conn.SetWriteDeadline(time.Now().Add(writeWait))
if err := s.conn.WriteMessage(websocket.TextMessage, message); err != nil {
return
}
case <-ticker.C:
s.conn.SetWriteDeadline(time.Now().Add(writeWait))
if err := s.conn.WriteMessage(websocket.PingMessage, nil); err != nil {
return
}
case <-s.closeCh:
return
}
}
}
// readPump 读协程,读取消息并路由。
func (s *Session) readPump(router *Router) {
defer func() {
if r := recover(); r != nil {
logrus.Errorf("[wsc] readPump panic recovered: %v, session=%s", r, s.wsCtx.SessionID[:8])
}
s.close()
}()
s.conn.SetReadLimit(maxMessageSize)
s.conn.SetReadDeadline(time.Now().Add(pongWait))
s.conn.SetPongHandler(func(string) error {
s.conn.SetReadDeadline(time.Now().Add(pongWait))
return nil
})
for {
_, data, err := s.conn.ReadMessage()
if err != nil {
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) &&
!errors.Is(err, io.EOF) {
logrus.Errorf("[wsc] read error: %v", err)
}
return
}
var msg Message
if err := json.Unmarshal(data, &msg); err != nil {
s.wsCtx.WriteError(400, "invalid message format")
continue
}
router.dispatch(s.wsCtx, &msg)
}
}
// close 安全关闭连接。
func (s *Session) close() {
s.once.Do(func() {
close(s.closeCh)
s.wsCtx.cancelCtx()
// 离开所有房间
// 先快照房间列表再释放读锁,避免 leave 触发 onEmpty → RemoveRoom 抢写锁死锁
s.server.mu.RLock()
rooms := make([]*Room, 0, len(s.server.rooms))
for _, room := range s.server.rooms {
rooms = append(rooms, room)
}
s.server.mu.RUnlock()
for _, room := range rooms {
room.Leave(s.wsCtx)
}
// 触发断开回调
if s.server.OnDisconnect != nil {
s.server.OnDisconnect(s.wsCtx)
}
s.server.onSessionClosed()
s.conn.Close()
// s.send 不关闭:readPump 的 dispatch 可能在 close 生效后仍并发回包,
// 向已关闭 channel 发送会 panic。writePump 退出由 closeCh 驱动,
// 残留消息随 Session 对象一起被 GC 回收。
})
}
+98
View File
@@ -0,0 +1,98 @@
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
}
}