File size: 3,730 Bytes
bf9e111 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 | // Package middleware provides HTTP middleware components for the CLI Proxy API server.
// This file contains unit tests for the middleware components.
package middleware
import (
"errors"
"net/http"
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
)
func setupTestRouter() *gin.Engine {
gin.SetMode(gin.TestMode)
return gin.New()
}
func TestCorrelationIDMiddleware(t *testing.T) {
router := setupTestRouter()
router.Use(CorrelationIDMiddleware())
router.GET("/test", func(c *gin.Context) {
id := GetCorrelationID(c)
c.JSON(200, gin.H{"correlation_id": id})
})
t.Run("generates correlation ID when not provided", func(t *testing.T) {
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/test", nil)
router.ServeHTTP(w, req)
assert.Equal(t, 200, w.Code)
// Check that response contains a correlation ID
assert.Contains(t, w.Body.String(), "correlation_id")
// Check that response header contains the correlation ID
assert.NotEmpty(t, w.Header().Get(CorrelationIDHeader))
})
t.Run("uses existing correlation ID from header", func(t *testing.T) {
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/test", nil)
req.Header.Set(CorrelationIDHeader, "test-correlation-id-123")
router.ServeHTTP(w, req)
assert.Equal(t, 200, w.Code)
assert.Contains(t, w.Body.String(), "test-correlation-id-123")
assert.Equal(t, "test-correlation-id-123", w.Header().Get(CorrelationIDHeader))
})
}
func TestErrorHandlerMiddleware(t *testing.T) {
router := setupTestRouter()
router.Use(CorrelationIDMiddleware())
router.Use(ErrorHandlerMiddleware())
router.GET("/error", func(c *gin.Context) {
c.Error(errors.New("test error"))
c.Abort()
})
router.GET("/success", func(c *gin.Context) {
c.JSON(200, gin.H{"status": "ok"})
})
t.Run("handles errors gracefully", func(t *testing.T) {
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/error", nil)
router.ServeHTTP(w, req)
assert.Equal(t, 500, w.Code)
assert.Contains(t, w.Body.String(), "INTERNAL_ERROR")
})
t.Run("passes through successful requests", func(t *testing.T) {
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/success", nil)
router.ServeHTTP(w, req)
assert.Equal(t, 200, w.Code)
assert.Contains(t, w.Body.String(), "ok")
})
}
func TestRecoveryMiddleware(t *testing.T) {
router := setupTestRouter()
router.Use(RecoveryMiddleware(nil))
router.Use(CorrelationIDMiddleware())
router.Use(ErrorHandlerMiddleware())
router.GET("/panic", func(c *gin.Context) {
panic("test panic")
})
router.GET("/normal", func(c *gin.Context) {
c.JSON(200, gin.H{"status": "ok"})
})
t.Run("recovers from panic", func(t *testing.T) {
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/panic", nil)
router.ServeHTTP(w, req)
assert.Equal(t, 500, w.Code)
assert.Contains(t, w.Body.String(), "INTERNAL_ERROR")
})
t.Run("normal requests work after panic recovery", func(t *testing.T) {
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/normal", nil)
router.ServeHTTP(w, req)
assert.Equal(t, 200, w.Code)
assert.Contains(t, w.Body.String(), "ok")
})
}
func TestSafeHandler(t *testing.T) {
router := setupTestRouter()
router.Use(CorrelationIDMiddleware())
router.Use(ErrorHandlerMiddleware())
router.GET("/safe-panic", SafeHandler(func(c *gin.Context) {
panic("safe handler panic")
}))
t.Run("safe handler recovers from panic", func(t *testing.T) {
w := httptest.NewRecorder()
req, _ := http.NewRequest("GET", "/safe-panic", nil)
router.ServeHTTP(w, req)
assert.Equal(t, 500, w.Code)
assert.Contains(t, w.Body.String(), "INTERNAL_ERROR")
})
}
|