This commit is contained in:
oneao committed 2025-08-25 17:22:54 +08:00
1 parent 76a6c81af2
commit b19d083031
68 files changed
+3889 -343

No files matched your search

@@ -0,0 +1,158 @@
package router
import (
error2 "base-framework/pkg/error"
"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"] = []*error2.StackError{}
}
// 统一生成堆栈
se := error2.WrapWithStack(err)
c.keys["errors"] = append(c.keys["errors"].([]*error2.StackError), se)
}
// Errors 返回 []*utils.StackError
func (c *Context) Errors() []*error2.StackError {
if c.keys == nil {
return nil
}
if errs, exists := c.keys["errors"]; exists {
return errs.([]*error2.StackError)
}
return nil
}
@@ -0,0 +1,257 @@
package router
import (
"net/http"
"strings"
)
// HandlerFunc 路由和中间件处理函数签名
type HandlerFunc func(*Context)
// node 路由树节点
type node struct {
pattern string
part string
children []*node
isParam bool
handler HandlerFunc
}
// 匹配子节点,返回第一个匹配的
func (n *node) matchChild(part string) *node {
for _, child := range n.children {
if child.part == part || child.isParam {
return child
}
}
return nil
}
// 匹配所有匹配的子节点
func (n *node) matchChildren(part string) []*node {
nodes := make([]*node, 0)
for _, child := range n.children {
if child.part == part || child.isParam {
nodes = append(nodes, child)
}
}
return nodes
}
// 插入路由节点
func (n *node) insert(pattern string, parts []string, height int, handler HandlerFunc) {
if len(parts) == height {
n.pattern = pattern
n.handler = handler
return
}
part := parts[height]
child := n.matchChild(part)
if child == nil {
child = &node{
part: part,
isParam: len(part) > 0 && part[0] == ':',
}
n.children = append(n.children, child)
}
child.insert(pattern, parts, height+1, handler)
}
// 搜索路由节点,并收集参数
func (n *node) search(parts []string, height int, params map[string]string) *node {
if len(parts) == height || n.part == "*" {
if n.pattern == "" {
return nil
}
return n
}
part := parts[height]
children := n.matchChildren(part)
for _, child := range children {
if child.isParam {
params[child.part[1:]] = part
}
res := child.search(parts, height+1, params)
if res != nil {
return res
}
}
return nil
}
// 中间件条目,包含排除路径集合
type middlewareEntry struct {
handler HandlerFunc
excludePaths map[string]struct{}
}
// Router 路由器结构
type Router struct {
roots map[string]*node
middleware []middlewareEntry
basePath string
}
// NewRouter 创建路由器实例
func NewRouter() *Router {
return &Router{
roots: make(map[string]*node),
}
}
// MiddlewareHandle 用于链式配置中间件排除路径
type MiddlewareHandle struct {
router *Router
entryIdx int
}
// Use 注册中间件,返回链式配置句柄
func (r *Router) Use(handler HandlerFunc) *MiddlewareHandle {
entry := middlewareEntry{
handler: handler,
excludePaths: nil,
}
r.middleware = append(r.middleware, entry)
return &MiddlewareHandle{
router: r,
entryIdx: len(r.middleware) - 1,
}
}
// ExcludePaths 设置排除路径
func (mh *MiddlewareHandle) ExcludePaths(paths ...string) *Router {
excludeMap := make(map[string]struct{}, len(paths))
for _, p := range paths {
excludeMap[p] = struct{}{}
}
mh.router.middleware[mh.entryIdx].excludePaths = excludeMap
return mh.router
}
// Group 路由分组,返回新的 Router,继承根节点和中间件,basePath拼接
func (r *Router) Group(prefix string, m ...HandlerFunc) *Router {
// 复制原middleware
newMiddleware := make([]middlewareEntry, len(r.middleware))
copy(newMiddleware, r.middleware)
// 将新增的 HandlerFunc 转成 middlewareEntry
for _, handler := range m {
entry := middlewareEntry{
handler: handler,
excludePaths: nil,
}
newMiddleware = append(newMiddleware, entry)
}
return &Router{
roots: r.roots,
middleware: newMiddleware,
basePath: joinPaths(r.basePath, prefix),
}
}
// 辅助函数,提取middlewareEntry中的handler为HandlerFunc slice
func (r *Router) middlewareHandlers() []HandlerFunc {
handlers := make([]HandlerFunc, 0, len(r.middleware))
for _, m := range r.middleware {
handlers = append(handlers, m.handler)
}
return handlers
}
// GET 注册GET请求路由
func (r *Router) GET(path string, handler HandlerFunc) *Router {
return r.handle("GET", path, handler)
}
// POST 注册POST请求路由
func (r *Router) POST(path string, handler HandlerFunc) *Router {
return r.handle("POST", path, handler)
}
// handle 注册具体方法和路径的处理函数
func (r *Router) handle(method, path string, handler HandlerFunc) *Router {
fullPath := joinPaths(r.basePath, path)
parts := parsePattern(fullPath)
if r.roots[method] == nil {
r.roots[method] = &node{}
}
r.roots[method].insert(fullPath, parts, 0, handler)
return r
}
// ServeHTTP 实现 http.Handler,执行路由匹配和中间件
func (r *Router) ServeHTTP(w http.ResponseWriter, req *http.Request) {
root := r.roots[req.Method]
if root == nil {
http.NotFound(w, req)
return
}
searchParts := parsePattern(req.URL.Path)
params := make(map[string]string)
n := root.search(searchParts, 0, params)
if n == nil || n.handler == nil {
http.NotFound(w, req)
return
}
// 按排除路径过滤中间件
handlers := make([]HandlerFunc, 0)
for _, m := range r.middleware {
if m.excludePaths != nil {
if _, excluded := m.excludePaths[req.URL.Path]; excluded {
continue
}
}
handlers = append(handlers, m.handler)
}
handlers = append(handlers, n.handler)
c := &Context{
Writer: w,
Request: req,
params: params,
index: -1,
handlers: handlers,
}
c.Next()
}
// 解析路径,去除空字符串
func parsePattern(pattern string) []string {
vs := strings.Split(strings.Trim(pattern, "/"), "/")
parts := make([]string, 0, len(vs))
for _, item := range vs {
if item != "" {
parts = append(parts, item)
}
}
return parts
}
// 拼接两个路径字符串,保证中间只有一个 '/'
func joinPaths(a, b string) string {
if a == "" {
if b == "" {
return "/"
}
if !strings.HasPrefix(b, "/") {
return "/" + b
}
return b
}
if b == "" {
return a
}
aslash := strings.HasSuffix(a, "/")
bslash := strings.HasPrefix(b, "/")
switch {
case aslash && bslash:
return a + b[1:]
case !aslash && !bslash:
return a + "/" + b
default:
return a + b
}
}