HTTP 中间件机制
一、知识点总结
1.1 什么是中间件(Middleware)
中间件是 HTTP 处理流程中的可插拔处理单元,它在请求到达最终 Handler 之前执行预处理,在 Handler 返回响应之后执行后处理。典型中间件包括:日志记录、认证鉴权、请求限流、panic 恢复、CORS 跨域、请求计时等。
中间件的核心理念可以用一个公式概括:
Middleware(Handler) → NewHandler
中间件接收一个 Handler,返回一个新的 Handler。这个新 Handler 在内部调用原始 Handler,但在调用前后插入额外的逻辑。这种模式在 Go 中称为**装饰器模式(Decorator Pattern)**的函数式实现。
1.2 洋葱模型(Onion Model)
中间件的执行顺序是理解其行为的关键。当多个中间件嵌套时,形成洋葱模型:
请求 → MiddlewareA → MiddlewareB → MiddlewareC → Handler
响应 ← MiddlewareA ← MiddlewareB ← MiddlewareC ←
即:请求阶段按注册顺序执行,响应阶段按注册顺序的逆序执行。外层中间件可以包裹内层,形成类似函数调用的栈结构。
1.3 标准库中的中间件实现方式
Go 标准库没有"中间件"这个术语,但提供了实现中间件的所有工具。最常见的两种实现方式:
方式一:Handler 包装(推荐)
func LoggerMiddleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
start := time.Now()
next.ServeHTTP(w, r)
log.Printf("%s %s %v", r.Method, r.URL.Path, time.Since(start))
})
}
方式二:HandlerFunc 包装
func LoggerMiddlewareFunc(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
start := time.Now()
next(w, r)
log.Printf("%s %s %v", r.Method, r.URL.Path, time.Since(start))
}
}
方式一更通用(接受 http.Handler),方式二更便捷(与 HandleFunc 配合)。实际项目中推荐使用方式一,因为它更标准、可复用性更强。
1.4 关键技巧:ResponseWriter 包装
很多中间件需要读取或修改响应数据(如日志记录响应状态码、Gzip 压缩响应体)。但 ResponseWriter 接口没有提供读取已写入内容的方法。解决方法是创建自定义的 ResponseWriter 包装器:
type responseRecorder struct {
http.ResponseWriter
status int
size int
}
func (rr *responseRecorder) WriteHeader(status int) {
rr.status = status
rr.ResponseWriter.WriteHeader(status)
}
func (rr *responseRecorder) Write(b []byte) (int, error) {
if rr.status == 0 {
rr.status = http.StatusOK
}
n, err := rr.ResponseWriter.Write(b)
rr.size += n
return n, err
}
这个包装器内嵌了 http.ResponseWriter,重写了 WriteHeader 和 Write 方法以捕获状态码和写入字节数。这是 Go 中间件开发中最核心的技巧之一。
1.5 Panic 恢复中间件
生产环境必须有一个 recover 中间件,防止单个请求的 panic 导致整个服务崩溃:
func RecoverMiddleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
defer func() {
if rec := recover(); rec != nil {
log.Printf("Panic recovered: %v", rec)
http.Error(w, "Internal Server Error", 500)
}
}()
next.ServeHTTP(w, r)
})
}
1.6 链式中间件组合
当有多个中间件时,可以用一个辅助函数把它们串联起来:
func Chain(h http.Handler, middlewares …func(http.Handler) http.Handler) http.Handler {
for i := len(middlewares) – 1; i >= 0; i— {
h = middlewares[i](h)
}
return h
}
// 使用
handler := Chain(myHandler, Logger, Recover, Auth)
注意循环是从后往前遍历的——这样 Middleware1 → Middleware2 → Middleware3 的注册顺序,实际执行时就是 Middleware1 在最外层包裹,Middleware3 在最内层包裹,符合直觉。
1.7 中间件 vs 框架中间件
| 接口 | Handler → Handler | c *gin.Context |
| 参数传递 | 通过 context 或闭包 | 通过 Context 对象 |
| 中断请求 | 直接 return,不调用 next | c.Abort() |
| 链式中断 | 手动控制 | 内置 Abort 机制 |
| 性能开销 | 零额外开销 | 有 Context 对象创建开销 |
二、练习代码
示例 1:日志记录中间件
package main
import (
"fmt"
"log"
"net/http"
"time"
)
// loggingResponseWriter 包装 http.ResponseWriter 以捕获状态码
type loggingResponseWriter struct {
http.ResponseWriter
status int
size int
}
func (lrw *loggingResponseWriter) WriteHeader(status int) {
lrw.status = status
lrw.ResponseWriter.WriteHeader(status)
}
func (lrw *loggingResponseWriter) Write(b []byte) (int, error) {
if lrw.status == 0 {
lrw.status = http.StatusOK
}
n, err := lrw.ResponseWriter.Write(b)
lrw.size += n
return n, err
}
// LoggerMiddleware 记录每个请求的耗时、状态码和响应大小
func LoggerMiddleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
start := time.Now()
lrw := &loggingResponseWriter{ResponseWriter: w, status: 0}
next.ServeHTTP(lrw, r)
duration := time.Since(start)
log.Printf("[%s] %s %s | Status: %d | Size: %d bytes | Duration: %v",
start.Format("2006-01-02 15:04:05"),
r.Method,
r.URL.Path,
lrw.status,
lrw.size,
duration,
)
})
}
// APIHandler 模拟业务 Handler
func APIHandler(w http.ResponseWriter, r *http.Request) {
// 模拟业务处理耗时
time.Sleep(10 * time.Millisecond)
w.Header().Set("Content-Type", "application/json")
fmt.Fprintln(w, `{"status":"ok","data":"hello"}`)
}
func ErrorHandler(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
fmt.Fprintln(w, `{"error":"something went wrong"}`)
}
func main() {
mux := http.NewServeMux()
mux.HandleFunc("/api/data", APIHandler)
mux.HandleFunc("/api/error", ErrorHandler)
mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
fmt.Fprintln(w, "Welcome! Try /api/data or /api/error")
})
// 用中间件包装整个 mux
var handler http.Handler = mux
handler = LoggerMiddleware(handler)
log.Println("Server on :8080")
log.Fatal(http.ListenAndServe(":8080", handler))
}
示例 2:Panic 恢复中间件
package main
import (
"fmt"
"log"
"net/http"
"runtime/debug"
)
// RecoverMiddleware 捕获 panic 防止服务崩溃
func RecoverMiddleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
defer func() {
if rec := recover(); rec != nil {
stack := debug.Stack()
log.Printf("[PANIC RECOVERED] %v\\n%s", rec, stack)
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
}
}()
next.ServeHTTP(w, r)
})
}
func main() {
mux := http.NewServeMux()
// 模拟一个会 panic 的 Handler
mux.HandleFunc("/panic", func(w http.ResponseWriter, r *http.Request) {
panic("Oops! Something terrible happened")
})
mux.HandleFunc("/safe", func(w http.ResponseWriter, r *http.Request) {
fmt.Fprintln(w, "This endpoint is safe")
})
// 先应用 Recover,再应用 Logger
var handler http.Handler = mux
handler = RecoverMiddleware(handler)
log.Println("Server on :8080 (try /panic and /safe)")
log.Fatal(http.ListenAndServe(":8080", handler))
}
示例 3:认证中间件 + 链式组合
package main
import (
"fmt"
"log"
"net/http"
"time"
)
// ======== 中间件定义 ========
// LoggerMiddleware 记录请求日志
func LoggerMiddleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
start := time.Now()
next.ServeHTTP(w, r)
log.Printf("[%s] %s %s – %v", r.Method, r.URL.Path, r.UserAgent(), time.Since(start))
})
}
// AuthMiddleware 简单的 Token 认证中间件
func AuthMiddleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// 公开接口不需要认证
if r.URL.Path == "/public" || r.URL.Path == "/" {
next.ServeHTTP(w, r)
return
}
// 从 Header 读取 Token
token := r.Header.Get("Authorization")
if token == "" {
w.Header().Set("WWW-Authenticate", "Bearer")
http.Error(w, "Unauthorized: missing token", http.StatusUnauthorized)
return
}
if token != "Bearer secret-token-123" {
http.Error(w, "Unauthorized: invalid token", http.StatusUnauthorized)
return
}
// 认证通过,继续执行
next.ServeHTTP(w, r)
})
}
// CORSMiddleware 简单的跨域中间件
func CORSMiddleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Access-Control-Allow-Origin", "*")
w.Header().Set("Access-Control-Allow-Methods", "GET, POST, PUT, DELETE, OPTIONS")
w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Authorization")
// 处理预检请求
if r.Method == http.MethodOptions {
w.WriteHeader(http.StatusOK)
return
}
next.ServeHTTP(w, r)
})
}
// RecoverMiddleware panic 恢复
func RecoverMiddleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
defer func() {
if rec := recover(); rec != nil {
log.Printf("[PANIC] %v", rec)
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
}
}()
next.ServeHTTP(w, r)
})
}
// ======== 链式组合工具 ========
// Chain 将多个中间件按顺序串联
// 注意:从后往前遍历,保证 middlewares[0] 在最外层
func Chain(h http.Handler, middlewares …func(http.Handler) http.Handler) http.Handler {
for i := len(middlewares) – 1; i >= 0; i— {
h = middlewares[i](h)
}
return h
}
// ======== Handlers ========
func main() {
mux := http.NewServeMux()
mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
fmt.Fprintln(w, "Public endpoint: /public, Protected: /api/data")
})
mux.HandleFunc("/public", func(w http.ResponseWriter, r *http.Request) {
fmt.Fprintln(w, `{"message":"This is a public endpoint"}`)
})
mux.HandleFunc("/api/data", func(w http.ResponseWriter, r *http.Request) {
fmt.Fprintln(w, `{"message":"Protected data","user":"admin"}`)
})
// 链式组合中间件(从外到内:CORS -> Logger -> Recover -> Auth -> mux)
// 执行顺序:
// 1. CORS 预处理 -> 2. Logger 记录开始 -> 3. Recover 捕获
// -> 4. Auth 验证 -> 5. Handler 执行
// <- 5. Handler 返回 <- 4. Auth 返回 <- 3. Recover 返回 <- 2. Logger 记录结束 <- 1. CORS 后处理
handler := Chain(mux,
CORSMiddleware,
LoggerMiddleware,
RecoverMiddleware,
AuthMiddleware,
)
log.Println("Server on :8080")
log.Println(" Public: GET http://localhost:8080/public")
log.Println(" Protected: GET http://localhost:8080/api/data")
log.Println(" Auth Header: Authorization: Bearer secret-token-123")
log.Fatal(http.ListenAndServe(":8080", handler))
}