package router import ( "strings" "github.com/valyala/fasthttp" ) 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{} excludePrefix []string } // 判断路径是否被排除 func (m *middlewareEntry) IsExcluded(path string) bool { path = strings.TrimSuffix(path, "/") // 去掉尾部斜杠 if m.excludePaths != nil { if _, ok := m.excludePaths[path]; ok { return true } } for _, prefix := range m.excludePrefix { if strings.HasPrefix(path, prefix) { return true } } return false } // Router 主结构 type Router struct { roots map[string]*node middleware []middlewareEntry basePath string } // 创建 Router func NewRouter(basePath ...string) *Router { path := "/" if len(basePath) > 0 && basePath[0] != "" { path = basePath[0] } return &Router{ roots: make(map[string]*node), basePath: path, } } // MiddlewareHandle 用于链式排除路径 type MiddlewareHandle struct { router *Router entryIdx int } // Use 注册中间件 func (r *Router) Use(handler HandlerFunc) *MiddlewareHandle { entry := middlewareEntry{handler: handler} r.middleware = append(r.middleware, entry) return &MiddlewareHandle{ router: r, entryIdx: len(r.middleware) - 1, } } // ExcludePaths 支持 /** 前缀匹配 // ExcludePaths 支持 /** 前缀匹配,并自动带上 basePath func (mh *MiddlewareHandle) ExcludePaths(paths ...string) *MiddlewareHandle { excludeMap := make(map[string]struct{}) excludePrefix := make([]string, 0) basePath := strings.TrimSuffix(mh.router.basePath, "/") // 去掉尾部斜杠 for _, p := range paths { fullPath := p if !strings.HasPrefix(p, "/") { fullPath = "/" + p } // 自动加上 basePath if basePath != "/" { fullPath = basePath + fullPath } if strings.HasSuffix(fullPath, "/**") { base := strings.TrimSuffix(fullPath, "/**") excludePrefix = append(excludePrefix, base) } else { excludeMap[fullPath] = struct{}{} } } mh.router.middleware[mh.entryIdx].excludePaths = excludeMap mh.router.middleware[mh.entryIdx].excludePrefix = excludePrefix return mh } // Group 创建子路由组 func (r *Router) Group(prefix string, m ...HandlerFunc) *Router { newMiddleware := make([]middlewareEntry, len(r.middleware)) copy(newMiddleware, r.middleware) for _, handler := range m { entry := middlewareEntry{handler: handler} newMiddleware = append(newMiddleware, entry) } return &Router{ roots: r.roots, middleware: newMiddleware, basePath: joinPaths(r.basePath, prefix), } } // GET/POST func (r *Router) GET(path string, handler HandlerFunc) *Router { return r.handle("GET", path, handler) } func (r *Router) POST(path string, handler HandlerFunc) *Router { return r.handle("POST", path, handler) } 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 } // Handler 入口(fasthttp) func (r *Router) Handler(ctx *fasthttp.RequestCtx) { method := string(ctx.Method()) root := r.roots[method] if root == nil { ctx.Error("未找到该页面", fasthttp.StatusNotFound) return } searchParts := parsePattern(string(ctx.Path())) params := make(map[string]string) n := root.search(searchParts, 0, params) if n == nil || n.handler == nil { ctx.Error("404 Not Found", fasthttp.StatusNotFound) return } // 组装中间件链 handlers := make([]HandlerFunc, 0) path := string(ctx.Path()) for _, m := range r.middleware { if m.IsExcluded(path) { continue } handlers = append(handlers, m.handler) } handlers = append(handlers, n.handler) // 最终 handler c := &Context{ RequestCtx: ctx, 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 } }