package router import ( "base-go-v2/internal/utils/mapx" "base-go-v2/internal/utils/validate" "encoding/json" "github.com/valyala/fasthttp" "sync" ) // BodyCache 缓存 fasthttp 请求体 type BodyCache struct { once sync.Once Data []byte Err error } // Load 读取并缓存请求体 func (b *BodyCache) Load(ctx *fasthttp.RequestCtx) ([]byte, error) { b.once.Do(func() { b.Data = ctx.PostBody() }) return b.Data, b.Err } // Context fasthttp 上下文封装 type Context struct { RequestCtx *fasthttp.RequestCtx params map[string]string index int handlers []HandlerFunc bodyCache BodyCache keys map[string]interface{} aborted bool } // Next 执行下一个中间件 func (c *Context) Next() error { c.index++ if c.index < len(c.handlers) { return c.handlers[c.index](c) } return nil } // Abort 中止执行 func (c *Context) Abort() { c.aborted = true } // Param 获取路由参数 func (c *Context) Param(key string) string { return c.params[key] } // Header 获取请求头 func (c *Context) Header(key string) string { return string(c.RequestCtx.Request.Header.Peek(key)) } // Body 获取请求体 func (c *Context) Body() ([]byte, error) { return c.bodyCache.Load(c.RequestCtx) } // BindJSON 反序列化 JSON 请求体 func (c *Context) BindJSON(obj interface{}) error { body, err := c.Body() if err != nil { return err } return json.Unmarshal(body, obj) } // BodyMap 将请求体解析为 mapx.M func (c *Context) BodyMap() (mapx.M, error) { body, err := c.Body() // 使用 BodyCache 读取请求体 if err != nil { return nil, err } m := mapx.New() if err := json.Unmarshal(body, &m); err != nil { return nil, err } return m, nil } func (c *Context) GetBodyWithRequired(keys ...string) (mapx.M, error) { // 解析请求体 req, err := c.BodyMap() if err != nil { return nil, err } // 校验必填字段 if err := validate.NotEmpty(req, keys...); err != nil { return nil, err } // 返回整个请求体(包含所有字段) return req, nil } // PostForm 获取表单参数 func (c *Context) PostForm(key string) string { return string(c.RequestCtx.FormValue(key)) } // Set 存储键值 func (c *Context) Set(key string, value interface{}) { if c.keys == nil { c.keys = make(map[string]interface{}) } c.keys[key] = value } // Get 获取键值 func (c *Context) Get(key string) (interface{}, bool) { if c.keys == nil { return nil, false } val, ok := c.keys[key] return val, ok } // SetUserID 存储 userId func (c *Context) SetUserID(userID int64) { c.Set("user_id", userID) } // GetUserID 获取 userId func (c *Context) GetUserID() int64 { val, ok := c.Get("user_id") if !ok || val == nil { return 0 } switch v := val.(type) { case int64: return v case int: return int64(v) case uint64: return int64(v) case uint: return int64(v) default: return 0 } } // JSON 返回 JSON 响应,支持返回 error func (c *Context) JSON(statusCode int, data interface{}) error { c.RequestCtx.SetStatusCode(statusCode) c.RequestCtx.SetContentType("application/json; charset=utf-8") if err := json.NewEncoder(c.RequestCtx).Encode(data); err != nil { c.Abort() return err } c.Abort() return nil } // Text 返回纯文本,支持返回 error func (c *Context) Text(statusCode int, msg string) error { c.RequestCtx.SetStatusCode(statusCode) c.RequestCtx.SetContentType("text/plain; charset=utf-8") if _, err := c.RequestCtx.WriteString(msg); err != nil { c.Abort() return err } c.Abort() return nil } // AddError 记录错误,自动包装 StackError func (c *Context) AddError(err error) error { if err == nil { return nil } if c.keys == nil { c.keys = make(map[string]interface{}) } if _, exists := c.keys["errors"]; !exists { c.keys["errors"] = []*StackError{} } // 统一生成堆栈 se := WrapWithStack(err) c.keys["errors"] = append(c.keys["errors"].([]*StackError), se) return err } // Errors 返回 []*utils.StackError func (c *Context) Errors() []*StackError { if c.keys == nil { return nil } if errs, exists := c.keys["errors"]; exists { return errs.([]*StackError) } return nil }