Files
workspace/code/base-project/base-go-v2/internal/router/router.go
T
2025-11-06 17:32:51 +08:00

253 lines
5.3 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{}
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
}
}