258 lines
5.7 KiB
Go
258 lines
5.7 KiB
Go
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
|
|
}
|
|
}
|