This commit is contained in:
oneao committed 2025-11-01 17:29:25 +08:00
1 parent b116690769
commit 1d6bb2b7f5
23 files changed
+922 -116

No files matched your search

@@ -0,0 +1,222 @@
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{}
}
// Router 主结构
type Router struct {
roots map[string]*node
middleware []middlewareEntry
basePath string
}
// NewRouter 创建
func NewRouter(basePath ...string) *Router {
path := "/"
if len(basePath) > 0 && basePath[0] != "" {
path = basePath[0]
}
return &Router{
roots: make(map[string]*node),
basePath: path,
}
}
type MiddlewareHandle struct {
router *Router
entryIdx int
}
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,
}
}
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
}
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),
}
}
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.excludePaths != nil {
if _, excluded := m.excludePaths[path]; excluded {
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
}
}