init: 通用库抽离——httpx/logger/wsc/wscclient/jwtx(合并分叉副本); 顶层包布局; module path git.zeroonesoft.cn/golib/zogo
This commit is contained in:
+187
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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 回收。
|
||||
})
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user