欢迎光临
我们一直在努力

Golang HTTP 中间件机制

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 框架中间件

特性标准库手写Gin/Echo 框架
接口 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))
}


三、今日思考题

  • 为什么 Chain 函数要从后往前遍历中间件数组?如果从前向后遍历会有什么问题?
  • loggingResponseWriter 为什么要内嵌 http.ResponseWriter 而不是直接实现完整的 http.ResponseWriter 接口?(提示:http.ResponseWriter 还包含哪些方法?)
  • 中间件中的 panic recovery 为什么一定要放在 defer 中?如果直接写 recover() 会怎样?

  • 赞(0)
    未经允许不得转载:171主机测评 » Golang HTTP 中间件机制
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址