Middleware в Go
Middleware -- функция, которая оборачивает http.Handler, добавляя поведение до и/или после обработки запроса. Это основной механизм расширения HTTP-сервера в Go.
Базовый паттерн
// Middleware type: takes a handler, returns a wrapped handler
type Middleware func(http.Handler) http.Handler
// Example: simple middleware structure
func myMiddleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// Before handler
fmt.Println("before")
next.ServeHTTP(w, r) // call the next handler
// After handler
fmt.Println("after")
})
}
Цепочка Middleware
// Chain applies middleware in order: first middleware wraps outermost
func Chain(h http.Handler, middlewares ...Middleware) http.Handler {
// Apply in reverse so first middleware is outermost
for i := len(middlewares) - 1; i >= 0; i-- {
h = middlewares[i](h)
}
return h
}
// Usage
mux := http.NewServeMux()
mux.HandleFunc("GET /api/users", listUsers)
handler := Chain(mux,
recoveryMiddleware(), // 1st: outermost, catches panics
requestIDMiddleware(), // 2nd: assigns request ID
loggingMiddleware(), // 3rd: logs request
corsMiddleware(), // 4th: handles CORS
authMiddleware(), // 5th: checks authentication
)
http.ListenAndServe(":8080", handler)
Logging Middleware
// responseRecorder captures the status code written by the handler.
type responseRecorder struct {
http.ResponseWriter
statusCode int
written int64
}
func newResponseRecorder(w http.ResponseWriter) *responseRecorder {
return &responseRecorder{ResponseWriter: w, statusCode: http.StatusOK}
}
func (r *responseRecorder) WriteHeader(code int) {
r.statusCode = code
r.ResponseWriter.WriteHeader(code)
}
func (r *responseRecorder) Write(b []byte) (int, error) {
n, err := r.ResponseWriter.Write(b)
r.written += int64(n)
return n, err
}
// LoggingMiddleware logs every request with duration, status, and size.
func LoggingMiddleware(logger *slog.Logger) Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
start := time.Now()
rec := newResponseRecorder(w)
// Get request ID from context (set by RequestID middleware)
requestID, _ := r.Context().Value(requestIDKey).(string)
next.ServeHTTP(rec, r)
duration := time.Since(start)
logger.Info("http request",
"method", r.Method,
"path", r.URL.Path,
"status", rec.statusCode,
"duration", duration,
"bytes", rec.written,
"remote_addr", r.RemoteAddr,
"user_agent", r.UserAgent(),
"request_id", requestID,
)
})
}
}
Request ID Middleware
type contextKey string
const requestIDKey contextKey = "request_id"
// RequestIDMiddleware assigns a unique ID to each request.
func RequestIDMiddleware() Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// Check for existing request ID (from load balancer, API gateway)
id := r.Header.Get("X-Request-ID")
if id == "" {
id = uuid.New().String()
}
// Set response header
w.Header().Set("X-Request-ID", id)
// Add to context
ctx := context.WithValue(r.Context(), requestIDKey, id)
next.ServeHTTP(w, r.WithContext(ctx))
})
}
}
// GetRequestID extracts request ID from context.
func GetRequestID(ctx context.Context) string {
id, _ := ctx.Value(requestIDKey).(string)
return id
}
Authentication Middleware
JWT Authentication
// AuthMiddleware validates JWT tokens.
func AuthMiddleware(secretKey []byte) Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// Extract token from Authorization header
authHeader := r.Header.Get("Authorization")
if authHeader == "" {
writeError(w, http.StatusUnauthorized, "missing authorization header")
return
}
// Expect "Bearer <token>"
parts := strings.SplitN(authHeader, " ", 2)
if len(parts) != 2 || parts[0] != "Bearer" {
writeError(w, http.StatusUnauthorized, "invalid authorization format")
return
}
token := parts[1]
// Validate and parse JWT
claims, err := validateJWT(token, secretKey)
if err != nil {
writeError(w, http.StatusUnauthorized, "invalid token")
return
}
// Add claims to context
ctx := context.WithValue(r.Context(), userClaimsKey, claims)
next.ServeHTTP(w, r.WithContext(ctx))
})
}
}
// SkipAuth returns a middleware that skips auth for specified paths.
func SkipAuth(authMW Middleware, skipPaths ...string) Middleware {
skip := make(map[string]bool)
for _, p := range skipPaths {
skip[p] = true
}
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if skip[r.URL.Path] {
next.ServeHTTP(w, r)
return
}
authMW(next).ServeHTTP(w, r)
})
}
}
API Key Authentication
// APIKeyMiddleware validates API key from header or query param.
func APIKeyMiddleware(validKeys map[string]bool) Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
key := r.Header.Get("X-API-Key")
if key == "" {
key = r.URL.Query().Get("api_key")
}
if key == "" || !validKeys[key] {
writeError(w, http.StatusUnauthorized, "invalid API key")
return
}
next.ServeHTTP(w, r)
})
}
}
CORS Middleware
// CORSConfig holds CORS middleware configuration.
type CORSConfig struct {
AllowedOrigins []string
AllowedMethods []string
AllowedHeaders []string
ExposedHeaders []string
AllowCredentials bool
MaxAge int // seconds
}
// DefaultCORSConfig returns permissive CORS config for development.
func DefaultCORSConfig() CORSConfig {
return CORSConfig{
AllowedOrigins: []string{"*"},
AllowedMethods: []string{"GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"},
AllowedHeaders: []string{"Accept", "Authorization", "Content-Type", "X-Request-ID"},
MaxAge: 86400,
}
}
// CORSMiddleware handles Cross-Origin Resource Sharing.
func CORSMiddleware(cfg CORSConfig) Middleware {
allowedOrigins := make(map[string]bool)
for _, o := range cfg.AllowedOrigins {
allowedOrigins[o] = true
}
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
origin := r.Header.Get("Origin")
// Check if origin is allowed
if allowedOrigins["*"] || allowedOrigins[origin] {
w.Header().Set("Access-Control-Allow-Origin", origin)
}
if cfg.AllowCredentials {
w.Header().Set("Access-Control-Allow-Credentials", "true")
}
// Handle preflight
if r.Method == http.MethodOptions {
w.Header().Set("Access-Control-Allow-Methods",
strings.Join(cfg.AllowedMethods, ", "))
w.Header().Set("Access-Control-Allow-Headers",
strings.Join(cfg.AllowedHeaders, ", "))
if cfg.MaxAge > 0 {
w.Header().Set("Access-Control-Max-Age",
strconv.Itoa(cfg.MaxAge))
}
w.WriteHeader(http.StatusNoContent)
return
}
if len(cfg.ExposedHeaders) > 0 {
w.Header().Set("Access-Control-Expose-Headers",
strings.Join(cfg.ExposedHeaders, ", "))
}
next.ServeHTTP(w, r)
})
}
}
Rate Limiting Middleware
// RateLimiter provides per-IP rate limiting using token bucket.
type RateLimiter struct {
mu sync.Mutex
visitors map[string]*rate.Limiter
limit rate.Limit
burst int
}
// NewRateLimiter creates a rate limiter: limit requests per second, burst size.
func NewRateLimiter(rps float64, burst int) *RateLimiter {
rl := &RateLimiter{
visitors: make(map[string]*rate.Limiter),
limit: rate.Limit(rps),
burst: burst,
}
// Clean up old entries periodically
go rl.cleanup()
return rl
}
func (rl *RateLimiter) getLimiter(ip string) *rate.Limiter {
rl.mu.Lock()
defer rl.mu.Unlock()
limiter, ok := rl.visitors[ip]
if !ok {
limiter = rate.NewLimiter(rl.limit, rl.burst)
rl.visitors[ip] = limiter
}
return limiter
}
func (rl *RateLimiter) cleanup() {
ticker := time.NewTicker(time.Minute)
defer ticker.Stop()
for range ticker.C {
rl.mu.Lock()
// Simple cleanup: just reset the map periodically
rl.visitors = make(map[string]*rate.Limiter)
rl.mu.Unlock()
}
}
// RateLimitMiddleware limits requests per IP.
func RateLimitMiddleware(rps float64, burst int) Middleware {
limiter := NewRateLimiter(rps, burst)
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
ip := r.RemoteAddr
// Handle X-Forwarded-For from reverse proxy
if forwarded := r.Header.Get("X-Forwarded-For"); forwarded != "" {
ip = strings.Split(forwarded, ",")[0]
}
if !limiter.getLimiter(ip).Allow() {
w.Header().Set("Retry-After", "1")
writeError(w, http.StatusTooManyRequests, "rate limit exceeded")
return
}
next.ServeHTTP(w, r)
})
}
}
Recovery Middleware
// RecoveryMiddleware recovers from panics and returns 500.
func RecoveryMiddleware(logger *slog.Logger) Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
defer func() {
if rec := recover(); rec != nil {
// Get stack trace
stack := make([]byte, 4096)
n := runtime.Stack(stack, false)
logger.Error("panic recovered",
"panic", rec,
"method", r.Method,
"path", r.URL.Path,
"stack", string(stack[:n]),
)
writeError(w, http.StatusInternalServerError, "internal server error")
}
}()
next.ServeHTTP(w, r)
})
}
}
Timeout Middleware
// TimeoutMiddleware cancels requests exceeding the given duration.
func TimeoutMiddleware(timeout time.Duration) Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
ctx, cancel := context.WithTimeout(r.Context(), timeout)
defer cancel()
// Create a channel for completion
done := make(chan struct{})
go func() {
next.ServeHTTP(w, r.WithContext(ctx))
close(done)
}()
select {
case <-done:
// Handler completed normally
case <-ctx.Done():
// Timeout exceeded
writeError(w, http.StatusGatewayTimeout, "request timeout")
}
})
}
}
// Note: for simpler cases, use http.TimeoutHandler from stdlib
handler := http.TimeoutHandler(mux, 30*time.Second, "request timeout")
Compression Middleware
// GzipMiddleware compresses responses for clients that support it.
func GzipMiddleware() Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// Check if client accepts gzip
if !strings.Contains(r.Header.Get("Accept-Encoding"), "gzip") {
next.ServeHTTP(w, r)
return
}
gz := gzip.NewWriter(w)
defer gz.Close()
w.Header().Set("Content-Encoding", "gzip")
w.Header().Del("Content-Length") // length changes after compression
gzw := &gzipResponseWriter{Writer: gz, ResponseWriter: w}
next.ServeHTTP(gzw, r)
})
}
}
type gzipResponseWriter struct {
io.Writer
http.ResponseWriter
}
func (g *gzipResponseWriter) Write(b []byte) (int, error) {
return g.Writer.Write(b)
}
Порядок Middleware -- важно!
Порядок применения middleware критически важен. Первый middleware в цепочке -- самый внешний:
handler := Chain(mux,
RecoveryMiddleware(logger), // 1. Catches panics (must be outermost!)
RequestIDMiddleware(), // 2. Assigns ID before logging
LoggingMiddleware(logger), // 3. Logs with request ID
CORSMiddleware(corsConfig), // 4. Handles CORS before auth rejects
RateLimitMiddleware(100, 10), // 5. Rate limit before expensive auth
AuthMiddleware(secretKey), // 6. Authenticates request
TimeoutMiddleware(30*time.Second), // 7. Timeout for handler execution
)
Порядок обработки запроса:
Request → Recovery → RequestID → Logging → CORS → RateLimit → Auth → Timeout → Handler
Response ← Recovery ← RequestID ← Logging ← CORS ← RateLimit ← Auth ← Timeout ← Handler