package router import ( "base-framework/pkg/utils" "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 { Request *http.Request Writer http.ResponseWriter params map[string]string index int handlers []HandlerFunc bodyCache BodyCache keys map[string]interface{} aborted bool // 是否中止 } // Next 执行下一个中间件或处理函数 func (c *Context) Next() { c.index++ for c.index < len(c.handlers) { if c.aborted { break } c.handlers[c.index](c) c.index++ } } // 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 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 方便读取请求体 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 } // AddError 记录错误,自动包装 StackError func (c *Context) AddError(err error) { if err == nil { return } if c.keys == nil { c.keys = make(map[string]interface{}) } if _, exists := c.keys["errors"]; !exists { c.keys["errors"] = []*utils.StackError{} } // 统一生成堆栈 se := utils.WrapWithStack(err) c.keys["errors"] = append(c.keys["errors"].([]*utils.StackError), se) } // Errors 返回 []*utils.StackError func (c *Context) Errors() []*utils.StackError { if c.keys == nil { return nil } if errs, exists := c.keys["errors"]; exists { return errs.([]*utils.StackError) } return nil }