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 } }