Files

358 lines
10 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}
}