fail2ban_state_test.go 3.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109
  1. package database
  2. import (
  3. "encoding/json"
  4. "os"
  5. "path/filepath"
  6. "runtime"
  7. "testing"
  8. "github.com/mhsanaei/3x-ui/v3/internal/config"
  9. "github.com/mhsanaei/3x-ui/v3/internal/database/model"
  10. )
  11. // stubFail2banClient puts a fail2ban-client on PATH whose exit code the test picks.
  12. func stubFail2banClient(t *testing.T, exitCode int) {
  13. t.Helper()
  14. dir := t.TempDir()
  15. script := filepath.Join(dir, "fail2ban-client")
  16. body := "#!/bin/sh\nexit " + string(rune('0'+exitCode)) + "\n"
  17. if err := os.WriteFile(script, []byte(body), 0o755); err != nil {
  18. t.Fatalf("write stub: %v", err)
  19. }
  20. t.Setenv("PATH", dir)
  21. }
  22. func TestFail2banEnforcementStateSeparatesAbsentFromUnrunnable(t *testing.T) {
  23. if runtime.GOOS == "windows" {
  24. t.Skip("fail2ban shell fixtures are Unix-only")
  25. }
  26. t.Run("absent", func(t *testing.T) {
  27. t.Setenv("PATH", t.TempDir())
  28. if got, _ := fail2banEnforcementState(); got != fail2banAbsent {
  29. t.Fatalf("state = %v, want fail2banAbsent", got)
  30. }
  31. })
  32. t.Run("present and runnable", func(t *testing.T) {
  33. stubFail2banClient(t, 0)
  34. if got, _ := fail2banEnforcementState(); got != fail2banEnforcing {
  35. t.Fatalf("state = %v, want fail2banEnforcing", got)
  36. }
  37. })
  38. t.Run("present but failing", func(t *testing.T) {
  39. stubFail2banClient(t, 1)
  40. got, err := fail2banEnforcementState()
  41. if got != fail2banUnknown {
  42. t.Fatalf("state = %v, want fail2banUnknown", got)
  43. }
  44. if err == nil {
  45. t.Fatal("want the probe error, got nil")
  46. }
  47. })
  48. }
  49. func TestResetIpLimitsKeepsConfiguredLimitsWhenProbeFails(t *testing.T) {
  50. if runtime.GOOS == "windows" {
  51. t.Skip("fail2ban shell fixtures are Unix-only")
  52. }
  53. t.Setenv("XUI_DB_FOLDER", t.TempDir())
  54. if err := InitDB(config.GetDBPath()); err != nil {
  55. t.Fatalf("init db: %v", err)
  56. }
  57. t.Cleanup(func() { _ = CloseDB() })
  58. if err := db.Where("seeder_name = ?", "ResetIpLimitNoFail2ban").Delete(&model.HistoryOfSeeders{}).Error; err != nil {
  59. t.Fatalf("clear seeder history: %v", err)
  60. }
  61. settings, err := json.Marshal(map[string]any{"clients": []any{map[string]any{"email": "[email protected]", "limitIp": 6}}})
  62. if err != nil {
  63. t.Fatalf("marshal settings: %v", err)
  64. }
  65. inbound := model.Inbound{Remark: "kept", Settings: string(settings)}
  66. if err := db.Create(&inbound).Error; err != nil {
  67. t.Fatalf("create inbound: %v", err)
  68. }
  69. record := model.ClientRecord{Email: "[email protected]", LimitIP: 2}
  70. if err := db.Create(&record).Error; err != nil {
  71. t.Fatalf("create client record: %v", err)
  72. }
  73. stubFail2banClient(t, 1)
  74. if err := resetIpLimitsWithoutFail2ban(); err != nil {
  75. t.Fatalf("reset: %v", err)
  76. }
  77. var gotInbound model.Inbound
  78. if err := db.First(&gotInbound, inbound.Id).Error; err != nil {
  79. t.Fatalf("reload inbound: %v", err)
  80. }
  81. var got map[string]any
  82. if err := json.Unmarshal([]byte(gotInbound.Settings), &got); err != nil {
  83. t.Fatalf("decode settings: %v", err)
  84. }
  85. clients := got["clients"].([]any)
  86. if limit := clients[0].(map[string]any)["limitIp"]; limit != float64(6) {
  87. t.Fatalf("inbound limitIp = %v, want 6", limit)
  88. }
  89. if err := db.First(&record, record.Id).Error; err != nil {
  90. t.Fatalf("reload client record: %v", err)
  91. }
  92. if record.LimitIP != 2 {
  93. t.Fatalf("client record limitIp = %d, want 2", record.LimitIP)
  94. }
  95. var count int64
  96. if err := db.Model(&model.HistoryOfSeeders{}).Where("seeder_name = ?", "ResetIpLimitNoFail2ban").Count(&count).Error; err != nil {
  97. t.Fatalf("count seeder history: %v", err)
  98. }
  99. if count != 0 {
  100. t.Fatalf("seeder history rows = %d, want 0", count)
  101. }
  102. }