MidПрактика7 min

Middleware

Паттерн middleware в Go: цепочки, логирование, аутентификация, CORS, rate limiting, recovery

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

Проверь себя

Какой тип имеет middleware в Go?

Зачем responseRecorder перехватывает WriteHeader?

Почему Recovery middleware должен быть самым внешним в цепочке?