diff --git a/internal/web/job/check_client_ip_ban_commit_test.go b/internal/web/job/check_client_ip_ban_commit_test.go
new file mode 100644
index 000000000..0ce329632
--- /dev/null
+++ b/internal/web/job/check_client_ip_ban_commit_test.go
@@ -0,0 +1,212 @@
+package job
+
+import (
+ "errors"
+ "os"
+ "testing"
+ "time"
+
+ "github.com/mhsanaei/3x-ui/v3/internal/database"
+
+ "gorm.io/gorm"
+)
+
+// installClientIpCommitFailure fails the transaction at COMMIT, not at a
+// statement, so every write inside it succeeds before the rollback.
+func installClientIpCommitFailure(t *testing.T) {
+ t.Helper()
+ db := database.GetDB()
+ switch db.Name() {
+ case "sqlite":
+ sqlDB, err := db.DB()
+ if err != nil {
+ t.Fatalf("get sql DB: %v", err)
+ }
+ // foreign_keys is per connection, so the injection only holds while the
+ // pool cannot hand the scan a fresh one with the pragma back at OFF.
+ sqlDB.SetMaxOpenConns(1)
+ for _, statement := range []string{
+ "DROP TRIGGER IF EXISTS commitfail_on_ips",
+ "DROP TABLE IF EXISTS commitfail_child",
+ "DROP TABLE IF EXISTS commitfail_parent",
+ "PRAGMA foreign_keys = ON",
+ "CREATE TABLE commitfail_parent (id INTEGER PRIMARY KEY)",
+ "CREATE TABLE commitfail_child (parent_id INTEGER, FOREIGN KEY(parent_id) REFERENCES commitfail_parent(id) DEFERRABLE INITIALLY DEFERRED)",
+ "CREATE TRIGGER commitfail_on_ips AFTER UPDATE OF ips ON inbound_client_ips BEGIN INSERT INTO commitfail_child(parent_id) VALUES (999); END",
+ } {
+ if err := db.Exec(statement).Error; err != nil {
+ t.Fatalf("install SQLite commit-failure injection %q: %v", statement, err)
+ }
+ }
+ t.Cleanup(func() {
+ _ = db.Exec("DROP TRIGGER IF EXISTS commitfail_on_ips").Error
+ _ = db.Exec("DROP TABLE IF EXISTS commitfail_child").Error
+ _ = db.Exec("DROP TABLE IF EXISTS commitfail_parent").Error
+ _ = db.Exec("PRAGMA foreign_keys = OFF").Error
+ })
+ case "postgres":
+ for _, statement := range []string{
+ "DROP TABLE IF EXISTS commitfail_child",
+ "DROP TABLE IF EXISTS commitfail_parent",
+ "CREATE TABLE commitfail_parent (id bigint PRIMARY KEY)",
+ "CREATE TABLE commitfail_child (id bigint PRIMARY KEY, parent_id bigint REFERENCES commitfail_parent(id) DEFERRABLE INITIALLY DEFERRED)",
+ } {
+ if err := db.Exec(statement).Error; err != nil {
+ t.Fatalf("install PostgreSQL commit-failure injection %q: %v", statement, err)
+ }
+ }
+ const callbackName = "test:client_ip_commit_failure"
+ if err := db.Callback().Update().After("gorm:update").Register(callbackName, func(tx *gorm.DB) {
+ stmt := tx.Statement
+ if stmt == nil || stmt.Schema == nil || stmt.Schema.Table != "inbound_client_ips" {
+ return
+ }
+ result := tx.Session(&gorm.Session{NewDB: true}).Exec("INSERT INTO commitfail_child (id, parent_id) VALUES (1, 999)")
+ if result.Error != nil {
+ _ = tx.AddError(result.Error)
+ }
+ }); err != nil {
+ t.Fatalf("register PostgreSQL commit-failure callback: %v", err)
+ }
+ t.Cleanup(func() {
+ _ = db.Callback().Update().Remove(callbackName)
+ _ = db.Exec("DROP TABLE IF EXISTS commitfail_child").Error
+ _ = db.Exec("DROP TABLE IF EXISTS commitfail_parent").Error
+ })
+ default:
+ t.Fatalf("unsupported test database dialect %q", db.Name())
+ }
+}
+
+// A fail2ban line is not a row a rollback can take back, so nothing may be
+// appended, and bannedSeen not advanced, until the scan has committed.
+func TestProcessObserved_CommitFailureDoesNotPublishBan(t *testing.T) {
+ setupIntegrationDB(t)
+
+ const email = "rollback-must-not-ban@x"
+ seedLinkedInboundWithClient(t, "rollback-must-not-ban", email, 1)
+ now := time.Now().Unix()
+ seedClientIps(t, email, []IPWithTimestamp{{IP: "198.51.100.10", Timestamp: now - 2}})
+
+ installClientIpCommitFailure(t)
+
+ j := NewCheckClientIpJob()
+ cleaned := j.processObserved(map[string]map[string]int64{
+ email: {
+ "198.51.100.10": now - 1,
+ "198.51.100.11": now,
+ },
+ }, true, true)
+ if cleaned {
+ t.Errorf("processObserved reported a published ban after the commit failed")
+ }
+ if got := ipSet(readClientIps(t, email)); len(got) != 1 || got["198.51.100.10"] != now-2 {
+ t.Errorf("rolled-back IP row = %v, want only the original client address", got)
+ }
+ if _, err := os.Stat(readIpLimitLogPath()); !os.IsNotExist(err) {
+ body, _ := os.ReadFile(readIpLimitLogPath())
+ t.Errorf("the rollback still touched the fail2ban trigger file (stat=%v):\n%s", err, body)
+ }
+ if _, seen := j.bannedSeen[email+"|198.51.100.10"]; seen {
+ t.Errorf("the rollback advanced bannedSeen and would suppress the retry")
+ }
+}
+
+// The committed row has already dropped the address, so a bannedSeen entry
+// recorded ahead of a failed write would suppress the ban for good.
+func TestProcessObserved_PublishFailureLeavesBanRetryable(t *testing.T) {
+ setupIntegrationDB(t)
+
+ const email = "publish-failure@x"
+ seedLinkedInboundWithClient(t, "publish-failure", email, 1)
+ now := time.Now().Unix()
+
+ // A directory where the log file belongs makes every open fail.
+ if err := os.MkdirAll(readIpLimitLogPath(), 0o755); err != nil {
+ t.Fatalf("block the log path: %v", err)
+ }
+
+ j := NewCheckClientIpJob()
+ observed := map[string]map[string]int64{
+ email: {"198.51.100.20": now - 1, "198.51.100.21": now},
+ }
+ if cleaned := j.processObserved(observed, true, true); cleaned {
+ t.Errorf("processObserved reported a publication that could not happen")
+ }
+ for key := range j.bannedSeen {
+ t.Errorf("bannedSeen recorded %q although nothing was written", key)
+ }
+}
+
+// A client back under its limit produces no candidates, so pruning cannot live
+// in the selection step: a surviving entry suppresses its next real ban.
+func TestProcessObserved_ForgetsBannedSeenWhenClientReturnsUnderLimit(t *testing.T) {
+ setupIntegrationDB(t)
+
+ const email = "prune-banned-seen@x"
+ seedLinkedInboundWithClient(t, "prune-banned-seen", email, 1)
+ now := time.Now().Unix()
+ seedClientIps(t, email, []IPWithTimestamp{{IP: "203.0.113.1", Timestamp: now - 500}})
+ j := NewCheckClientIpJob()
+
+ j.processObserved(map[string]map[string]int64{
+ email: {"203.0.113.1": now - 400, "203.0.113.2": now - 300},
+ }, true, true)
+ if got := banLineCount(t, email); got != 1 {
+ t.Fatalf("ban lines after the first scan = %d, want 1", got)
+ }
+
+ // Back under the limit: no candidates, so the stale entry must be dropped here.
+ j.processObserved(map[string]map[string]int64{
+ email: {"203.0.113.1": now - 400},
+ }, true, true)
+ if len(j.bannedSeen) != 0 {
+ t.Fatalf("bannedSeen = %v, want empty once the client is under its limit", j.bannedSeen)
+ }
+
+ j.processObserved(map[string]map[string]int64{
+ email: {"203.0.113.1": now - 400, "203.0.113.3": now},
+ }, true, true)
+ if got := banLineCount(t, email); got != 2 {
+ t.Fatalf("ban lines after the client goes over again = %d, want 2", got)
+ }
+}
+
+type failingWriter struct{ err error }
+
+func (f failingWriter) Write([]byte) (int, error) { return 0, f.err }
+
+// A dropped write error would let publishBans record an address the jail never
+// sees, so it has to reach the caller.
+func TestWriteBanLinesSurfacesWriteFailure(t *testing.T) {
+ want := errors.New("no space left on device")
+ err := writeBanLines(failingWriter{err: want}, "write-failure@x", []IPWithTimestamp{
+ {IP: "203.0.113.9", Timestamp: time.Now().Unix()},
+ })
+ if !errors.Is(err, want) {
+ t.Fatalf("writeBanLines error = %v, want the writer's own error", err)
+ }
+}
+
+// A committed over-limit scan writes its line and hands the client on.
+func TestProcessObserved_PublishesBanForCommittedScan(t *testing.T) {
+ setupIntegrationDB(t)
+
+ const email = "published-ban@x"
+ seedLinkedInboundWithClient(t, "published-ban", email, 1)
+ now := time.Now().Unix()
+ seedClientIps(t, email, []IPWithTimestamp{{IP: "203.0.113.50", Timestamp: now - 500}})
+
+ j := NewCheckClientIpJob()
+ if cleaned := j.processObserved(map[string]map[string]int64{
+ email: {"203.0.113.50": now - 400, "203.0.113.51": now},
+ }, true, true); !cleaned {
+ t.Fatalf("a published ban must report the access log as worth cleaning")
+ }
+ if got := banLineCount(t, email); got != 1 {
+ t.Fatalf("ban lines = %d, want 1", got)
+ }
+ if _, seen := j.bannedSeen[email+"|203.0.113.50"]; !seen {
+ t.Fatalf("a published address must be recorded so the next scan does not repeat it")
+ }
+}
diff --git a/internal/web/job/check_client_ip_frozen_ban_test.go b/internal/web/job/check_client_ip_frozen_ban_test.go
index fd25fc7ee..79daccb1f 100644
--- a/internal/web/job/check_client_ip_frozen_ban_test.go
+++ b/internal/web/job/check_client_ip_frozen_ban_test.go
@@ -5,8 +5,6 @@ import (
"strings"
"testing"
"time"
-
- "github.com/mhsanaei/3x-ui/v3/internal/database"
)
func banLineCount(t *testing.T, email string) int {
@@ -41,14 +39,14 @@ func TestUpdateInboundClientIps_FrozenLastSeenBannedOnce(t *testing.T) {
}
row := seedClientIps(t, email, nil)
- if _, banned := j.updateInboundClientIps(database.GetDB(), row, inbound, email, 1, live, true, true); !banned {
+ if banned, _ := j.enforceIpLimitForTest(t, row, inbound, email, 1, live, true); !banned {
t.Fatalf("first scan: the over-limit stale IP must be banned")
}
if got := banLineCount(t, email); got != 1 {
t.Fatalf("ban lines after first scan = %d, want 1", got)
}
- if _, banned := j.updateInboundClientIps(database.GetDB(), row, inbound, email, 1, live, true, true); banned {
+ if banned, _ := j.enforceIpLimitForTest(t, row, inbound, email, 1, live, true); banned {
t.Fatalf("second scan with a frozen lastSeen must not re-ban a dead connection")
}
if got := banLineCount(t, email); got != 1 {
@@ -59,7 +57,7 @@ func TestUpdateInboundClientIps_FrozenLastSeenBannedOnce(t *testing.T) {
{IP: "10.2.0.1", Timestamp: now + 30},
{IP: "192.0.2.7", Timestamp: now + 60},
}
- if _, banned := j.updateInboundClientIps(database.GetDB(), row, inbound, email, 1, reconnected, true, true); !banned {
+ if banned, _ := j.enforceIpLimitForTest(t, row, inbound, email, 1, reconnected, true); !banned {
t.Fatalf("a reconnect (advanced lastSeen) must be banned again")
}
if got := banLineCount(t, email); got != 2 {
diff --git a/internal/web/job/check_client_ip_job.go b/internal/web/job/check_client_ip_job.go
index 75290cf99..471d75c8a 100644
--- a/internal/web/job/check_client_ip_job.go
+++ b/internal/web/job/check_client_ip_job.go
@@ -4,7 +4,8 @@ import (
"context"
"encoding/json"
"errors"
- "log"
+ "fmt"
+ "io"
"os"
"os/exec"
"runtime"
@@ -32,11 +33,10 @@ type IPWithTimestamp struct {
// API; no access log is involved. On a core too old to expose that API the job
// simply skips the run (the bundled core always supports it).
type CheckClientIpJob struct {
- disAllowedIps []string
- bannedSeen map[string]int64
- xrayService service.XrayService
- allowlist ipLimitAllowlist
- lastIpPrune int64
+ bannedSeen map[string]int64
+ xrayService service.XrayService
+ allowlist ipLimitAllowlist
+ lastIpPrune int64
}
var job *CheckClientIpJob
@@ -290,11 +290,7 @@ func (j *CheckClientIpJob) processObserved(observed map[string]map[string]int64,
// be recorded under this panel's own guid for cross-node IP attribution.
attribution := make(map[string][]model.ClientIpEntry, len(observed))
- type pendingDisconnect struct {
- inbound *model.Inbound
- email string
- }
- var disconnects []pendingDisconnect
+ var bans []pendingBan
db := database.GetDB()
tx := db.Begin()
@@ -364,23 +360,23 @@ func (j *CheckClientIpJob) processObserved(observed map[string]map[string]int64,
continue
}
- cleaned, banned := j.updateInboundClientIps(tx, clientIpsRecord, inbound, email, limitByEmail[email], ipsWithTime, enforce, observedAreLive)
- shouldCleanLog = cleaned || shouldCleanLog
- if banned {
- disconnects = append(disconnects, pendingDisconnect{inbound: inbound, email: email})
- }
+ candidates, keptLive := j.updateInboundClientIps(tx, clientIpsRecord, inbound, email, limitByEmail[email], ipsWithTime, enforce, observedAreLive)
+ bans = append(bans, pendingBan{inbound: inbound, email: email, candidates: candidates, keptLive: keptLive})
}
if err := tx.Commit().Error; err != nil {
j.checkError(err)
- return shouldCleanLog
+ return false
}
committed = true
+ published := j.publishBans(bans)
+
// Xray disconnects run after the commit so their network round-trips never
// extend the scan's write transaction (node syncs upsert the same table).
+ shouldCleanLog = shouldCleanLog || len(published) > 0
clientsCache := make(map[int][]model.Client)
- for _, d := range disconnects {
+ for _, d := range published {
clients, cached := clientsCache[d.inbound.Id]
if !cached {
clients, _ = service.ParseInboundSettingsClients(d.inbound.Settings)
@@ -495,13 +491,12 @@ func (j *CheckClientIpJob) delInboundClientIps(tx *gorm.DB, clientEmail string)
}
// updateInboundClientIps merges one email's observed IPs into its tracking row
-// and applies the IP limit. limitIp comes from the caller (the clients table);
-// writes go through the caller's transaction. banned=true asks the caller to
-// disconnect the client after the transaction commits.
-func (j *CheckClientIpJob) updateInboundClientIps(tx *gorm.DB, inboundClientIps *model.InboundClientIps, inbound *model.Inbound, clientEmail string, limitIp int, newIpsWithTime []IPWithTimestamp, enforce, observedAreLive bool) (shouldCleanLog, banned bool) {
+// and applies the IP limit. Ban candidates are returned, not written: the
+// fail2ban log is the point of no return and must wait for the commit.
+func (j *CheckClientIpJob) updateInboundClientIps(tx *gorm.DB, inboundClientIps *model.InboundClientIps, inbound *model.Inbound, clientEmail string, limitIp int, newIpsWithTime []IPWithTimestamp, enforce, observedAreLive bool) (banCandidates []IPWithTimestamp, keptLiveCount int) {
if inbound.Settings == "" {
logger.Debug("wrong data:", inbound)
- return false, false
+ return nil, 0
}
if !enforce || limitIp <= 0 || !inbound.Enable {
@@ -512,7 +507,7 @@ func (j *CheckClientIpJob) updateInboundClientIps(tx *gorm.DB, inboundClientIps
if err := tx.Save(inboundClientIps).Error; err != nil {
logger.Error("failed to save inboundClientIps:", err)
}
- return false, false
+ return nil, 0
}
// Parse old IPs from database
@@ -531,40 +526,15 @@ func (j *CheckClientIpJob) updateInboundClientIps(tx *gorm.DB, inboundClientIps
}
liveIps, historicalIps := partitionLiveIps(ipMap, observedThisScan)
- j.disAllowedIps = []string{}
-
// historical db-only ips are excluded from this count on purpose.
limitedIps, allowedIps := j.allowlist.split(liveIps)
keptLive, bannedLive := selectIpsToBan(limitedIps, limitIp)
// Allowlisted addresses stay connected and out of the count: charging them
// against the limit would still cut the shared network the entry protects.
keptLive = append(keptLive, allowedIps...)
- actionable := j.filterAdvancedSinceLastBan(clientEmail, bannedLive)
- if len(actionable) > 0 {
- shouldCleanLog = true
- banned = true
- logIpFile, err := os.OpenFile(xray.GetIPLimitLogPath(), os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o644)
- if err != nil {
- logger.Errorf("failed to open IP limit log file: %s", err)
- return false, false
- }
- defer logIpFile.Close()
- ipLogger := log.New(logIpFile, "", log.LstdFlags)
-
- // log format is load-bearing: x-ui.sh create_iplimit_jails builds
- // filter.d/3x-ipl.conf with
- // failregex = \[LIMIT_IP\]\s*Email\s*=\s*.+\s*\|\|\s*Disconnecting OLD IP\s*=\s*\s*\|\|\s*Timestamp\s*=\s*\d+
- // don't change the wording.
- for _, ipTime := range actionable {
- j.disAllowedIps = append(j.disAllowedIps, ipTime.IP)
- ipLogger.Printf("[LIMIT_IP] Email = %s || Disconnecting OLD IP = %s || Timestamp = %d", clientEmail, ipTime.IP, ipTime.Timestamp)
- }
- }
-
- // keep kept-live + historical in the blob so the panel keeps showing
- // recently seen ips. banned live ips are already in the fail2ban log
- // and will reappear in the next scan if they reconnect.
+ // keep kept-live + historical in the blob so the panel keeps showing recently
+ // seen ips; banned live ips reappear in the next scan if they reconnect.
dbIps := make([]IPWithTimestamp, 0, len(keptLive)+len(historicalIps))
dbIps = append(dbIps, keptLive...)
dbIps = append(dbIps, historicalIps...)
@@ -573,33 +543,97 @@ func (j *CheckClientIpJob) updateInboundClientIps(tx *gorm.DB, inboundClientIps
if err := tx.Save(inboundClientIps).Error; err != nil {
logger.Error("failed to save inboundClientIps:", err)
- return false, banned
+ return nil, 0
}
- if len(j.disAllowedIps) > 0 {
- logger.Infof("[LIMIT_IP] Client %s: Kept %d live IPs, queued %d old IPs for fail2ban", clientEmail, len(keptLive), len(j.disAllowedIps))
- }
-
- return shouldCleanLog, banned
+ return bannedLive, len(keptLive)
}
-// filterAdvancedSinceLastBan keeps only banned pairs whose lastSeen advanced since
-// the previous ban: the core refreshes lastSeen solely on a new dispatch, so a
-// frozen value is a dead connection it hasn't reaped yet, not a reconnect.
-func (j *CheckClientIpJob) filterAdvancedSinceLastBan(email string, banned []IPWithTimestamp) []IPWithTimestamp {
+// pendingBan carries one client's enforcement outcome from inside the scan's
+// transaction to the publication that may only follow a successful commit.
+type pendingBan struct {
+ inbound *model.Inbound
+ email string
+ candidates []IPWithTimestamp
+ keptLive int
+}
+
+// publishBans returns the clients whose lines reached the log. bannedSeen
+// advances only for those, so a failed write leaves the address retryable.
+func (j *CheckClientIpJob) publishBans(bans []pendingBan) []pendingBan {
+ published := make([]pendingBan, 0, len(bans))
+ var logIpFile *os.File
+ defer func() {
+ if logIpFile == nil {
+ return
+ }
+ if err := logIpFile.Close(); err != nil {
+ logger.Errorf("failed to close IP limit log file: %s", err)
+ }
+ }()
+ for _, b := range bans {
+ actionable := j.selectAdvancedSinceLastBan(b.email, b.candidates)
+ if len(actionable) == 0 {
+ j.recordBannedSeen(b.email, b.candidates, nil)
+ continue
+ }
+ if logIpFile == nil {
+ f, err := os.OpenFile(xray.GetIPLimitLogPath(), os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o644)
+ if err != nil {
+ logger.Errorf("failed to open IP limit log file: %s", err)
+ return published
+ }
+ logIpFile = f
+ }
+ if err := writeBanLines(logIpFile, b.email, actionable); err != nil {
+ logger.Errorf("failed to write IP limit bans for %s: %s", b.email, err)
+ continue
+ }
+ j.recordBannedSeen(b.email, b.candidates, actionable)
+ logger.Infof("[LIMIT_IP] Client %s: Kept %d live IPs, queued %d old IPs for fail2ban", b.email, b.keptLive, len(actionable))
+ published = append(published, b)
+ }
+ return published
+}
+
+// writeBanLines emits one line per address; the wording is load-bearing, since
+// x-ui.sh create_iplimit_jails builds filter.d/3x-ipl.conf failregex from it.
+func writeBanLines(w io.Writer, clientEmail string, actionable []IPWithTimestamp) error {
+ stamp := time.Now().Format("2006/01/02 15:04:05")
+ for _, ipTime := range actionable {
+ if _, err := fmt.Fprintf(w, "%s [LIMIT_IP] Email = %s || Disconnecting OLD IP = %s || Timestamp = %d\n",
+ stamp, clientEmail, ipTime.IP, ipTime.Timestamp); err != nil {
+ return err
+ }
+ }
+ return nil
+}
+
+// selectAdvancedSinceLastBan drops pairs with a frozen lastSeen: the core
+// refreshes it only on a new dispatch, so those are unreaped dead connections.
+func (j *CheckClientIpJob) selectAdvancedSinceLastBan(email string, banned []IPWithTimestamp) []IPWithTimestamp {
+ actionable := make([]IPWithTimestamp, 0, len(banned))
+ for _, ipTime := range banned {
+ if last, ok := j.bannedSeen[email+"|"+ipTime.IP]; ok && ipTime.Timestamp <= last {
+ continue
+ }
+ actionable = append(actionable, ipTime)
+ }
+ return actionable
+}
+
+// recordBannedSeen marks published pairs and forgets addresses this scan no
+// longer bans; it runs for every enforced client, which is what prunes the map.
+func (j *CheckClientIpJob) recordBannedSeen(email string, banned, published []IPWithTimestamp) {
if j.bannedSeen == nil {
j.bannedSeen = make(map[string]int64)
}
+ for _, ipTime := range published {
+ j.bannedSeen[email+"|"+ipTime.IP] = ipTime.Timestamp
+ }
current := make(map[string]struct{}, len(banned))
- actionable := make([]IPWithTimestamp, 0, len(banned))
for _, ipTime := range banned {
- key := email + "|" + ipTime.IP
- current[key] = struct{}{}
- if last, ok := j.bannedSeen[key]; ok && ipTime.Timestamp <= last {
- continue
- }
- j.bannedSeen[key] = ipTime.Timestamp
- actionable = append(actionable, ipTime)
+ current[email+"|"+ipTime.IP] = struct{}{}
}
prefix := email + "|"
for key := range j.bannedSeen {
@@ -609,7 +643,6 @@ func (j *CheckClientIpJob) filterAdvancedSinceLastBan(email string, banned []IPW
}
}
}
- return actionable
}
// disconnectClientTemporarily drops a client's credential for a moment, so new
diff --git a/internal/web/job/check_client_ip_job_integration_test.go b/internal/web/job/check_client_ip_job_integration_test.go
index a30c12e47..a25870f93 100644
--- a/internal/web/job/check_client_ip_job_integration_test.go
+++ b/internal/web/job/check_client_ip_job_integration_test.go
@@ -24,6 +24,7 @@ var loggerInitOnce sync.Once
// updateInboundClientIps can run end to end. closes the db before
// TempDir cleanup so windows doesn't complain about the file being in
// use.
+
func setupIntegrationDB(t *testing.T) {
t.Helper()
@@ -57,6 +58,19 @@ func setupIntegrationDB(t *testing.T) {
})
}
+// enforceIpLimitForTest runs the same two steps processObserved does: select
+// inside the transaction, publish only once it would have committed.
+func (j *CheckClientIpJob) enforceIpLimitForTest(t *testing.T, row *model.InboundClientIps, inbound *model.Inbound, email string, limit int, live []IPWithTimestamp, observedAreLive bool) (banned bool, published []IPWithTimestamp) {
+ t.Helper()
+ candidates, keptLive := j.updateInboundClientIps(database.GetDB(), row, inbound, email, limit, live, true, observedAreLive)
+ actionable := j.selectAdvancedSinceLastBan(email, candidates)
+ done := j.publishBans([]pendingBan{{inbound: inbound, email: email, candidates: candidates, keptLive: keptLive}})
+ if len(done) == 0 {
+ return false, nil
+ }
+ return true, actionable
+}
+
// seed an inbound whose settings json has a single client with the
// given email and ip limit.
func seedInboundWithClient(t *testing.T, tag, email string, limitIp int) {
@@ -206,16 +220,13 @@ func TestUpdateInboundClientIps_LiveIpNotBannedByStillFreshHistoricals(t *testin
if err != nil {
t.Fatalf("getInboundByEmail: %v", err)
}
- shouldCleanLog, banned := j.updateInboundClientIps(database.GetDB(), row, inbound, email, 3, live, true, false)
+ banned, published := j.enforceIpLimitForTest(t, row, inbound, email, 3, live, false)
- if shouldCleanLog {
- t.Fatalf("shouldCleanLog must be false, nothing should have been banned with 1 live ip under limit 3")
- }
if banned {
t.Fatalf("banned must be false with 1 live ip under limit 3")
}
- if len(j.disAllowedIps) != 0 {
- t.Fatalf("disAllowedIps must be empty, got %v", j.disAllowedIps)
+ if len(published) != 0 {
+ t.Fatalf("published bans must be empty, got %v", published)
}
persisted := ipSet(readClientIps(t, email))
@@ -262,16 +273,13 @@ func TestUpdateInboundClientIps_ExcessLiveIpIsStillBanned(t *testing.T) {
if err != nil {
t.Fatalf("getInboundByEmail: %v", err)
}
- shouldCleanLog, banned := j.updateInboundClientIps(database.GetDB(), row, inbound, email, 1, live, true, false)
+ banned, published := j.enforceIpLimitForTest(t, row, inbound, email, 1, live, false)
- if !shouldCleanLog {
- t.Fatalf("shouldCleanLog must be true when the live set exceeds the limit")
- }
if !banned {
t.Fatalf("banned must be true when the live set exceeds the limit")
}
- if len(j.disAllowedIps) != 1 || j.disAllowedIps[0] != "10.1.0.1" {
- t.Fatalf("expected 10.1.0.1 to be banned; disAllowedIps = %v", j.disAllowedIps)
+ if len(published) != 1 || published[0].IP != "10.1.0.1" {
+ t.Fatalf("expected 10.1.0.1 to be banned; published = %v", published)
}
persisted := ipSet(readClientIps(t, email))
@@ -446,13 +454,13 @@ func TestUpdateInboundClientIps_AllowlistedIpIsNeitherCountedNorBanned(t *testin
if err != nil {
t.Fatalf("getInboundByEmail: %v", err)
}
- _, banned := j.updateInboundClientIps(database.GetDB(), row, inbound, email, 1, live, true, false)
+ banned, published := j.enforceIpLimitForTest(t, row, inbound, email, 1, live, false)
if banned {
t.Fatal("an allowlisted address pushed the client over its limit and something was banned")
}
- if len(j.disAllowedIps) != 0 {
- t.Fatalf("disAllowedIps = %v, want none", j.disAllowedIps)
+ if len(published) != 0 {
+ t.Fatalf("published = %v, want none", published)
}
persisted := ipSet(readClientIps(t, email))