package middleware import ( "io" "net/http" "net/http/httptest" "os" "path/filepath" "strings" "testing" "github.com/alicebob/miniredis/v2" "github.com/gin-gonic/gin" "github.com/redis/go-redis/v9" "github.com/stretchr/testify/require" "github.com/gochat/gochat/internal/config" applogger "github.com/gochat/gochat/pkg/logger" ) func TestSecurityLimitsAreIndependentAndIgnoreForgedXFF(t *testing.T) { gin.SetMode(gin.TestMode) mr := miniredis.RunT(t) rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) t.Cleanup(func() { _ = rdb.Close() }) cfg := config.RateLimitConfig{ Enabled: true, Login: config.RouteLimitConfig{Requests: 1, WindowSeconds: 60}, Webhook: config.RouteLimitConfig{Requests: 1, WindowSeconds: 60}, } router := gin.New() require.NoError(t, router.SetTrustedProxies([]string{})) router.Use(SecurityRateLimit(rdb, cfg)) router.POST("/api/v1/auth/login", func(c *gin.Context) { c.Status(http.StatusNoContent) }) router.POST("/webhooks/provider", func(c *gin.Context) { c.Status(http.StatusNoContent) }) request := func(path, forwardedFor string) int { req := httptest.NewRequest(http.MethodPost, path, nil) req.RemoteAddr = "192.0.2.10:4321" req.Header.Set("X-Forwarded-For", forwardedFor) res := httptest.NewRecorder() router.ServeHTTP(res, req) return res.Code } require.Equal(t, http.StatusNoContent, request("/api/v1/auth/login", "198.51.100.1")) require.Equal(t, http.StatusTooManyRequests, request("/api/v1/auth/login", "198.51.100.2")) require.Equal(t, http.StatusNoContent, request("/webhooks/provider", "198.51.100.2")) } func TestSecurityRateLimitFailsClosedWhenRedisIsUnavailable(t *testing.T) { mr := miniredis.RunT(t) rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) mr.Close() t.Cleanup(func() { _ = rdb.Close() }) router := gin.New() router.Use(SecurityRateLimit(rdb, config.RateLimitConfig{Enabled: true, Login: config.RouteLimitConfig{Requests: 1, WindowSeconds: 60}})) router.POST("/api/v1/auth/login", func(c *gin.Context) { c.Status(http.StatusNoContent) }) res := httptest.NewRecorder() router.ServeHTTP(res, httptest.NewRequest(http.MethodPost, "/api/v1/auth/login", nil)) require.Equal(t, http.StatusServiceUnavailable, res.Code) } func TestRequestLoggerDoesNotWriteQueryCredentials(t *testing.T) { logPath := filepath.Join(t.TempDir(), "requests.log") require.NoError(t, applogger.Init(applogger.Config{Level: "info", Format: "json", Output: logPath, ErrorOutput: logPath})) router := gin.New() router.Use(RequestLogger()) router.GET("/cable", func(c *gin.Context) { c.Status(http.StatusNoContent) }) router.POST("/webhooks/telegram/:bot_token", func(c *gin.Context) { c.Status(http.StatusNoContent) }) res := httptest.NewRecorder() router.ServeHTTP(res, httptest.NewRequest(http.MethodGet, "/cable?ticket=do-not-log&token=also-secret", nil)) res = httptest.NewRecorder() router.ServeHTTP(res, httptest.NewRequest(http.MethodPost, "/webhooks/telegram/path-secret", nil)) applogger.Sync() file, err := os.Open(logPath) require.NoError(t, err) t.Cleanup(func() { _ = file.Close() }) contents, err := io.ReadAll(file) require.NoError(t, err) require.Contains(t, string(contents), "/cable") require.False(t, strings.Contains(string(contents), "do-not-log")) require.False(t, strings.Contains(string(contents), "also-secret")) require.False(t, strings.Contains(string(contents), "path-secret")) }