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 } }