Files
workspace/code/base-project/base-go-v2/internal/router/context.go
T
2025-11-11 16:35:32 +08:00

214 lines
4.1 KiB
Go

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
}