package middleware import ( "net/http" "net/http/httptest" "strings" "testing" "time" "github.com/gin-gonic/gin" "go.uber.org/zap/zaptest" ) func init() { gin.SetMode(gin.TestMode) } func TestRequestID_GeneratesNew(t *testing.T) { r := gin.New() r.Use(RequestID()) r.GET("/test", func(c *gin.Context) { c.String(200, c.GetString("trace_id")) }) w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) r.ServeHTTP(w, req) if w.Header().Get("X-Request-ID") == "" { t.Error("expected X-Request-ID header to be set") } if w.Body.String() == "" { t.Error("expected trace_id in context") } } func TestRequestID_PreservesExisting(t *testing.T) { r := gin.New() r.Use(RequestID()) r.GET("/test", func(c *gin.Context) { c.String(200, c.GetString("trace_id")) }) w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) req.Header.Set("X-Request-ID", "existing-id") r.ServeHTTP(w, req) if w.Header().Get("X-Request-ID") != "existing-id" { t.Errorf("expected existing-id, got %s", w.Header().Get("X-Request-ID")) } } func TestSecurityHeaders(t *testing.T) { r := gin.New() r.Use(SecurityHeaders()) r.GET("/test", func(c *gin.Context) { c.String(200, "ok") }) w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) r.ServeHTTP(w, req) headers := map[string]string{ "X-Content-Type-Options": "nosniff", "X-Frame-Options": "DENY", "X-XSS-Protection": "1; mode=block", "Referrer-Policy": "strict-origin-when-cross-origin", } for key, expected := range headers { if got := w.Header().Get(key); got != expected { t.Errorf("header %s: expected %q, got %q", key, expected, got) } } } func TestRateLimit_AllowsWithinLimit(t *testing.T) { r := gin.New() r.Use(RateLimit(5, time.Minute)) r.GET("/test", func(c *gin.Context) { c.String(200, "ok") }) for i := 0; i < 5; i++ { w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) r.ServeHTTP(w, req) if w.Code != 200 { t.Errorf("request %d: expected 200, got %d", i+1, w.Code) } } } func TestRateLimit_BlocksOverLimit(t *testing.T) { r := gin.New() r.Use(RateLimit(2, time.Minute)) r.GET("/test", func(c *gin.Context) { c.String(200, "ok") }) for i := 0; i < 2; i++ { w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) r.ServeHTTP(w, req) } w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) r.ServeHTTP(w, req) if w.Code != http.StatusTooManyRequests { t.Errorf("expected 429, got %d", w.Code) } } func TestRequestLogger_DoesNotPanic(t *testing.T) { log := zaptest.NewLogger(t) r := gin.New() r.Use(RequestID()) r.Use(RequestLogger(log)) r.GET("/test", func(c *gin.Context) { c.String(200, "ok") }) w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) r.ServeHTTP(w, req) if w.Code != 200 { t.Errorf("expected 200, got %d", w.Code) } } func TestRecovery_HandlesPanic(t *testing.T) { log := zaptest.NewLogger(t) r := gin.New() r.Use(RequestID()) r.Use(Recovery(log)) r.GET("/panic", func(c *gin.Context) { panic("test panic") }) w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/panic", nil) r.ServeHTTP(w, req) if w.Code != 500 { t.Errorf("expected 500, got %d", w.Code) } } func TestBodyLimit_RejectsLargeBody(t *testing.T) { r := gin.New() r.Use(BodyLimit(10)) r.POST("/test", func(c *gin.Context) { c.String(200, "ok") }) w := httptest.NewRecorder() body := strings.NewReader(strings.Repeat("x", 100)) req, _ := http.NewRequest("POST", "/test", body) req.Header.Set("Content-Length", "100") r.ServeHTTP(w, req) if w.Code != http.StatusRequestEntityTooLarge { t.Errorf("expected 413, got %d", w.Code) } } func TestBodyLimit_AllowsSmallBody(t *testing.T) { r := gin.New() r.Use(BodyLimit(1024)) r.POST("/test", func(c *gin.Context) { c.String(200, "ok") }) w := httptest.NewRecorder() body := strings.NewReader("hello") req, _ := http.NewRequest("POST", "/test", body) r.ServeHTTP(w, req) if w.Code != 200 { t.Errorf("expected 200, got %d", w.Code) } }