Files
workspace/code/base-project/base-go-v2/internal/router/router.go
T
2025-11-01 17:29:25 +08:00

223 lines
4.6 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}
}