go-zero 是如何利用装饰器模式实现 HTTP 中间件?

一:自定义函数类型

Middleware 是 go-zero 定义的装饰器自定义函数类型,参数next是一个函数,返回值也是一个函数,函数类型均是http.HandlerFunc

type Middleware func(next http.HandlerFunc) http.HandlerFunc

作用就是接收下层请求处理函数,包装前置 / 后置逻辑后,返回增强后的新处理函数,对应装饰器模式。

其中http.HandlerFunc,也是标准库 net/http的一个自定义类型函数。这是用来被装饰器包装的业务处理函数。

package http
type HandlerFunc func(ResponseWriter, *Request)

如果你自定义一个请求拦截函数(类型为 Middleware),调用 use 方法后,该函数会被存入 middlewares 切片。

func (ng *engine) use(middleware Middleware) {
	ng.middlewares = append(ng.middlewares, middleware)
}
type engine struct {
	conf   RestConf
	routes []featuredRoutes
	timeout              time.Duration
	unauthorizedCallback handler.UnauthorizedCallback
	unsignedCallback     handler.UnsignedCallback
	chain                chain.Chain
	middlewares          []Middleware
	shedder              load.Shedder
	priorityShedder      load.Shedder
	tlsConfig            *tls.Config
}

上述缓存逻辑实现在引擎私有方法 use 中,Server结构体对此封装为对外暴露的公共方法,供开发者使用实现注册全局中间件。

func (s *Server) Use(middleware Middleware) {
	s.ngin.use(middleware)
}
type Server struct {
	ngin   *engine
	router httpx.Router
}

二:创建并注册请求拦截中间件

自定义创建一个请求拦截中间件,即创建请求拦截函数。注意,这个函数类型是Middleware

func IpInjectMiddleware(next http.HandlerFunc) http.HandlerFunc {
	return func(w http.ResponseWriter, r *http.Request) {
		// IP解析方法
		clientIP := utils.GetClientIPFromRequest(r)
		// 塞入上下文透传给logic
		newCtx := context.WithValue(r.Context(), ctxkey.IpKey{}, clientIP)
		next(w, r.WithContext(newCtx))
	}
}

定义完成后,调用服务实例的 Use 方法完成全局注册:

func main() {
	server := rest.MustNewServer(c.RestConf)
	defer server.Stop()
	...
	// 注册全局中间件 跨域、鉴权
	server.Use(middleware.IpInjectMiddleware)

三:请求拦截实现

1:注册路由

package rest

type Route struct {
	Method  string
	Path    string
	Handler http.HandlerFunc
}
	

结构体Route是一个单条路由规则,用来描述一条接口路由的三要素:请求方式、请求路径、接口处理函数。

所有接口组装成 []Route 路由切片,统一交给 server.AddRoutes() 注册。

func RegisterHandlers(server *rest.Server, serverCtx *svc.ServiceContext) {
	server.AddRoutes(
		[]rest.Route{
			{
				Method:  http.MethodGet,
				Path:    "/user/:id",
				Handler: user.GetUserInfoHandler(serverCtx),
			},
			{
				Method:  http.MethodPost,
				Path:    "/user/login",
				Handler: user.LoginHandler(serverCtx),
			},
			{
				Method:  http.MethodPost,
				Path:    "/user/login-send-code",
				Handler: user.LoginSendCodeHandler(serverCtx),
			},
			{
				Method:  http.MethodGet,
				Path:    "/user/user-count",
				Handler: user.GetUserCountHandler(serverCtx),
			},
		},
		rest.WithPrefix("/v1"),
	)
}
	

调用 server.AddRoutes(rs []Route, opts …RouteOption) 时,内部会初始化 featuredRoutes 对象,将传入的 []Route 存入其 routes 字段;再遍历传入的路由配置选项(如 WithPrefix、WithJwt、WithTimeout)填充分组配置,最后调用 s.ngin.addRoutes(r)。

func (s *Server) AddRoutes(rs []Route, opts ...RouteOption) {
	r := featuredRoutes{
		routes: rs,
	}
	for _, opt := range opts {
		opt(&r)
	}
	s.ngin.addRoutes(r)
}
	
package rest

type featuredRoutes struct {
	timeout   *time.Duration
	priority  bool
	jwt       jwtSetting
	signature signatureSetting
	sse       bool
	routes    []Route
	maxBytes  int64
}
	

featuredRoutes作用:go-zero 支持按路由分组差异化配置,比如部分接口需要 JWT、部分不需要;部分接口单独设置短超时,全部通过这个结构体承载分组配置。

注意:featuredRoutes 结构体本身没有存放自定义中间件的字段,不能直接在它内部定义局部自定义中间件;所以路由分组差异化只针对默认的中间件,如果自定义的中间件要实现由特定的某个路由执行,最简单的方式是在中间件内部加一层判断,匹配指定路由才走的逻辑,其他接口直接跳过。

执行 s.ngin.addRoutes(r) 后,该分组 featuredRoutes 会追加存入 engine 的 routes []featuredRoutes 切片中统一保存:

type engine struct {
	conf   RestConf
	routes []featuredRoutes // 存储全部路由分组
	timeout              time.Duration
	unauthorizedCallback handler.UnauthorizedCallback
	unsignedCallback     handler.UnsignedCallback
	chain                chain.Chain
	middlewares          []Middleware // 存放server.Use注册的全局自定义中间件
	shedder              load.Shedder
	priorityShedder      load.Shedder
	tlsConfig            *tls.Config
}
package rest

type Server struct {
	ngin   *engine
	router httpx.Router
}
	

handler.RegisterHandlers(server, ctx) 执行完后。Server 内部 engine 已经完整保存了所有路由分组、每条路由的原始业务 Handler、全局自定义中间件列表。但此时还未组装中间件执行链条,路由的请求的业务函数被自定义的中间件拦截功能,还未生效。

2:启动服务中间件绑定业务请求

现在,数据全部已经保存到位,只是没有组装、没有绑定到路由树

启动服务,经过以下步骤:

server.Start() -> s.ngin.start(s.router) -> ng.bindRoutes(router) -> ng.bindFeaturedRoutes(router, fr, metrics)

\rest\chain\chain.go

type chain struct {
	middlewares          []Middleware
}

func (ng *engine) bindRoute(fr featuredRoutes, router httpx.Router, metrics *stat.Metrics,
	route Route, verifier func(chain.Chain) chain.Chain) error {
	chn := ng.chain
	if chn == nil {
		chn = ng.buildChainWithNativeMiddlewares(fr, route, metrics)
	}
	chn = ng.appendAuthHandler(fr, chn, verifier)

	// 循环追加你自定义的全局中间件
	for _, middleware := range ng.middlewares {
		chn = chn.Append(convertMiddleware(middleware))
	}
	handle := chn.ThenFunc(route.Handler)

	return router.Handle(route.Method, route.Path, handle)
}

ng.middlewares:数组,存放上面说的所有 server.Use(xxxMiddleware) 注册的拦截函数(上面的 IpInjectMiddleware中间件)

convertMiddleware(middleware)这个转换比较费脑,我一步步梳理:

type Middleware func(next http.HandlerFunc) http.HandlerFunc
// 1. 自定义中间件
func IpInjectMiddleware(next http.HandlerFunc) http.HandlerFunc {
	return func(w http.ResponseWriter, r *http.Request) {
		clientIP := utils.GetClientIPFromRequest(r)
		newCtx := context.WithValue(r.Context(), ctxkey.IpKey{}, clientIP)
		next(w, r.WithContext(newCtx))
	}
}

// 2. 转换器
func convertMiddleware(ware Middleware) func(http.Handler) http.Handler {
	return func(next http.Handler) http.Handler {
		// 核心行:return ware(next.ServeHTTP)
		return ware(next.ServeHTTP)
	}
}

第 1 层:替换 ware


func convertMiddleware(IpInjectMiddleware) func(http.Handler) http.Handler {
	return func(next http.Handler) http.Handler {
		// ware = IpInjectMiddleware
		return IpInjectMiddleware(next.ServeHTTP)
	}
}

第二层:


func convertMiddleware(IpInjectMiddleware) func(http.Handler) http.Handler {
	return func(next http.Handler) http.Handler {
		// 完整展开 IpInjectMiddleware(next.ServeHTTP)
		return (
			// IpInjectMiddleware 定义体
			func(next http.HandlerFunc) http.HandlerFunc {
				return func(w http.ResponseWriter, r *http.Request) {
					clientIP := utils.GetClientIPFromRequest(r)
					newCtx := context.WithValue(r.Context(), ctxkey.IpKey{}, clientIP)
					// 这里的 next = 传入的实参:next.ServeHTTP
					next(w, r.WithContext(newCtx))
				}
			}
		)(next.ServeHTTP) // 实参:外层next.ServeHTTP(HandlerFunc类型)
	}
}

把内层形参改名,看得更清楚:

func convertMiddleware(IpInjectMiddleware) func(http.Handler) http.Handler {
	return func(nextHandler http.Handler) http.Handler {
		return (
			// 形参 f:接收外层传进来的 nextHandler.ServeHTTP
			func(f http.HandlerFunc) http.HandlerFunc {
				return func(w http.ResponseWriter, r *http.Request) {
					clientIP := utils.GetClientIPFromRequest(r)
					newCtx := context.WithValue(r.Context(), ctxkey.IpKey{}, clientIP)
					// f 等价原来内层的 next,就是下层处理器
					f(w, r.WithContext(newCtx))
				}
			}
		)(nextHandler.ServeHTTP)
	}
}

然后将整条中间件链条包裹当前路由的原始业务 HandlerFunc,生成最终统一处理器 handle

// handle := chn.ThenFunc(route.Handler)
func (c chain) Then(h http.Handler) http.Handler {
	if h == nil {
		h = http.DefaultServeMux
	}

	for i := range c.middlewares {
		h = c.middlewares[len(c.middlewares)-1-i](h)
	}

	return h
}
return router.Handle(route.Method, route.Path, handle)

把组装完成、包含全套拦截逻辑的 handle 绑定到路由树:请求方法、路径映射到这个复合 Handler。 后续客户端匹配该请求时,路由取出 handle 执行 ServeHTTP(w,r),触发完整中间件和业务流程。

四:客户端发起请求

Request完全是标准库原生结构体,不是go-zero框架自定义的。

import "net/http"
type Request struct {
	Method string
	URL    *url.URL
	Header Header
	Body   io.ReadCloser
	// 还有 RemoteAddr、ContentLength、ctx 等一堆字段
}

Go net/http 底层 accept 连接,读取 socket 二进制数据,解析 HTTP 协议(拆分请求行、Header、Body),接着把解析完成后的结构化数据,填充进 http.Request{} 结构体,go-zero框架拿到的 r *http.Request,是解析完成后的结构化对象

type Handler interface {
	ServeHTTP(w http.ResponseWriter, r *http.Request)
}

ServeHTTP 是处理单次 HTTP 请求的入口函数,只要一个类型实现了这个方法,它就是合法的 HTTP 处理器,标准库收到客户端连接后,会自动调用这个方法完成整套逻辑。即:接收请求载体 r 和响应输出工具 w,完成一次请求的读取请求 → 执行业务 → 组装返回数据全流程。

注意:这个函数不需要写 return,因为http.ResponseWriter的方法w.WriteHeader()、w.Write() 可以直接往响应缓冲区写入数据,函数执行完毕后,底层直接把 w 缓存的响应推送客户端。