253 lines
5.3 KiB
Go
253 lines
5.3 KiB
Go
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 {
|
||
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 支持 /** 前缀匹配
|
||
func (mh *MiddlewareHandle) ExcludePaths(paths ...string) *Router {
|
||
excludeMap := make(map[string]struct{}, 0)
|
||
excludePrefix := make([]string, 0)
|
||
|
||
for _, p := range paths {
|
||
if strings.HasSuffix(p, "/**") {
|
||
base := strings.TrimSuffix(p, "/**")
|
||
excludePrefix = append(excludePrefix, base)
|
||
} else {
|
||
excludeMap[p] = struct{}{}
|
||
}
|
||
}
|
||
|
||
mh.router.middleware[mh.entryIdx].excludePaths = excludeMap
|
||
mh.router.middleware[mh.entryIdx].excludePrefix = excludePrefix
|
||
|
||
return mh.router
|
||
}
|
||
|
||
// 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("404 Not Found", 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
|
||
}
|
||
}
|