u
This commit is contained in:
1 parent
b116690769
commit
1d6bb2b7f5
23 files changed
+922
-116
No files matched your search
@@ -0,0 +1,123 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"sync"
|
||||
|
||||
"github.com/valyala/fasthttp"
|
||||
)
|
||||
|
||||
// BodyCache 缓存 fasthttp 请求体
|
||||
type BodyCache struct {
|
||||
once sync.Once
|
||||
Data []byte
|
||||
Err error
|
||||
}
|
||||
|
||||
// Load 读取并缓存请求体
|
||||
func (b *BodyCache) Load(ctx *fasthttp.RequestCtx) ([]byte, error) {
|
||||
b.once.Do(func() {
|
||||
b.Data = ctx.PostBody()
|
||||
})
|
||||
return b.Data, b.Err
|
||||
}
|
||||
|
||||
// Context fasthttp 上下文封装
|
||||
type Context struct {
|
||||
RequestCtx *fasthttp.RequestCtx
|
||||
params map[string]string
|
||||
index int
|
||||
handlers []HandlerFunc
|
||||
|
||||
bodyCache BodyCache
|
||||
keys map[string]interface{}
|
||||
aborted bool
|
||||
}
|
||||
|
||||
// Next 执行下一个中间件
|
||||
func (c *Context) Next() error {
|
||||
c.index++
|
||||
if c.index < len(c.handlers) {
|
||||
return c.handlers[c.index](c)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Abort 中止执行
|
||||
func (c *Context) Abort() {
|
||||
c.aborted = true
|
||||
}
|
||||
|
||||
// Param 获取路由参数
|
||||
func (c *Context) Param(key string) string {
|
||||
return c.params[key]
|
||||
}
|
||||
|
||||
// Header 获取请求头
|
||||
func (c *Context) Header(key string) string {
|
||||
return string(c.RequestCtx.Request.Header.Peek(key))
|
||||
}
|
||||
|
||||
// Body 获取请求体
|
||||
func (c *Context) Body() ([]byte, error) {
|
||||
return c.bodyCache.Load(c.RequestCtx)
|
||||
}
|
||||
|
||||
// BindJSON 反序列化 JSON 请求体
|
||||
func (c *Context) BindJSON(obj interface{}) error {
|
||||
body, err := c.Body()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return json.Unmarshal(body, obj)
|
||||
}
|
||||
|
||||
// PostForm 获取表单参数
|
||||
func (c *Context) PostForm(key string) string {
|
||||
return string(c.RequestCtx.FormValue(key))
|
||||
}
|
||||
|
||||
// Set 存储键值
|
||||
func (c *Context) Set(key string, value interface{}) {
|
||||
if c.keys == nil {
|
||||
c.keys = make(map[string]interface{})
|
||||
}
|
||||
c.keys[key] = value
|
||||
}
|
||||
|
||||
// Get 获取键值
|
||||
func (c *Context) Get(key string) (interface{}, bool) {
|
||||
if c.keys == nil {
|
||||
return nil, false
|
||||
}
|
||||
val, ok := c.keys[key]
|
||||
return val, ok
|
||||
}
|
||||
|
||||
// JSON 返回 JSON 响应,支持返回 error
|
||||
func (c *Context) JSON(statusCode int, data interface{}) error {
|
||||
c.RequestCtx.SetStatusCode(statusCode)
|
||||
c.RequestCtx.SetContentType("application/json; charset=utf-8")
|
||||
|
||||
if err := json.NewEncoder(c.RequestCtx).Encode(data); err != nil {
|
||||
c.Abort()
|
||||
return err
|
||||
}
|
||||
|
||||
c.Abort()
|
||||
return nil
|
||||
}
|
||||
|
||||
// Text 返回纯文本,支持返回 error
|
||||
func (c *Context) Text(statusCode int, msg string) error {
|
||||
c.RequestCtx.SetStatusCode(statusCode)
|
||||
c.RequestCtx.SetContentType("text/plain; charset=utf-8")
|
||||
|
||||
if _, err := c.RequestCtx.WriteString(msg); err != nil {
|
||||
c.Abort()
|
||||
return err
|
||||
}
|
||||
|
||||
c.Abort()
|
||||
return nil
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
package router
|
||||
|
||||
type HandlerFunc func(c *Context) error
|
||||
|
||||
// node 路由树节点
|
||||
type node struct {
|
||||
pattern string
|
||||
part string
|
||||
children []*node
|
||||
isParam bool
|
||||
handler HandlerFunc
|
||||
}
|
||||
Reference in new issue
Block a user