mengstack-api/internal/app/middleware/middleware_test.go
MengStack Dev 9148d2f6da
Some checks failed
CI / Build & Test (push) Failing after 33s
feat: Swagger 全量注解 + M8 质量加固
- 37 个端点全部添加 godoc Swagger 注解(RBAC 11 / Org 5 / Audit 1 / Settings 7 / Notification 7)
- 重新生成 docs/(swagger.json/yaml/docs.go)
- M8 质量横切:IP 限流、安全头、Body 限制、结构化日志增强、优雅关闭
- 新增 Dockerfile 多阶段构建 + .dockerignore
- 新增 testutil 测试工具包
- 修复 testutil.go 编译错误
2026-10-03 01:57:39 +08:00

188 lines
4.1 KiB
Go

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)
}
}