117 lines
3.0 KiB
Go
117 lines
3.0 KiB
Go
package router
|
|
|
|
import (
|
|
"bytes"
|
|
"encoding/json"
|
|
"io"
|
|
"net/http"
|
|
"sync"
|
|
)
|
|
|
|
// BodyCache 专门负责缓存请求体,保证只读一次
|
|
type BodyCache struct {
|
|
once sync.Once
|
|
Data []byte
|
|
Err error
|
|
}
|
|
|
|
// Load 读取请求体并缓存,只执行一次
|
|
func (b *BodyCache) Load(r *http.Request) ([]byte, error) {
|
|
b.once.Do(func() {
|
|
if r.Body == nil {
|
|
b.Err = http.ErrBodyNotAllowed
|
|
return
|
|
}
|
|
var buf bytes.Buffer
|
|
b.Err = func() error {
|
|
_, err := io.Copy(&buf, r.Body)
|
|
return err
|
|
}()
|
|
if b.Err != nil {
|
|
return
|
|
}
|
|
b.Data = buf.Bytes()
|
|
// 重新设置请求体方便后续读取
|
|
r.Body = io.NopCloser(bytes.NewReader(b.Data))
|
|
})
|
|
return b.Data, b.Err
|
|
}
|
|
|
|
// Context 自定义请求上下文,封装了请求信息、响应写入器、路由参数、中间件处理等功能
|
|
type Context struct {
|
|
writer http.ResponseWriter // HTTP 响应写入器,用于构造响应数据
|
|
request *http.Request // HTTP 请求对象,包含请求相关信息
|
|
params map[string]string // 路由参数,如动态路径中的变量值
|
|
index int // 当前执行的中间件/处理函数索引,用于控制 Next 调用流程
|
|
handlers []HandlerFunc // 本次请求的中间件和最终处理函数链
|
|
|
|
bodyCache BodyCache // 请求体缓存,确保请求体只读一次且可多次读取
|
|
keys map[string]interface{} // 用于存储请求生命周期内的自定义数据(如用户信息、orgId等)
|
|
}
|
|
|
|
// Next 执行下一个中间件或处理函数
|
|
func (c *Context) Next() {
|
|
c.index++
|
|
if c.index < len(c.handlers) {
|
|
c.handlers[c.index](c)
|
|
}
|
|
}
|
|
|
|
// Param 获取路由参数
|
|
func (c *Context) Param(key string) string {
|
|
return c.params[key]
|
|
}
|
|
|
|
// Header 获取请求头
|
|
func (c *Context) Header(key string) string {
|
|
return c.request.Header.Get(key)
|
|
}
|
|
|
|
// JSON 返回 JSON 格式响应
|
|
func (c *Context) JSON(statusCode int, data interface{}) {
|
|
c.writer.Header().Set("Content-Type", "application/json; charset=utf-8")
|
|
c.writer.WriteHeader(statusCode)
|
|
if err := json.NewEncoder(c.writer).Encode(data); err != nil {
|
|
http.Error(c.writer, err.Error(), http.StatusInternalServerError)
|
|
}
|
|
}
|
|
|
|
// Body 方便读取请求体,实际调用 BodyCache 的 Load 方法
|
|
func (c *Context) Body() ([]byte, error) {
|
|
return c.bodyCache.Load(c.request)
|
|
}
|
|
|
|
// BindJSON 反序列化 JSON 请求体到 obj
|
|
func (c *Context) BindJSON(obj interface{}) error {
|
|
body, err := c.Body()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return json.Unmarshal(body, obj)
|
|
}
|
|
|
|
// PostForm 获取表单参数
|
|
func (c *Context) PostForm(key string) string {
|
|
if err := c.request.ParseForm(); err != nil {
|
|
return ""
|
|
}
|
|
return c.request.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
|
|
}
|