147 lines
3.9 KiB
Go
147 lines
3.9 KiB
Go
package abuseguard
|
|||
|
|
|
||
|
|
import (
|
||
|
|
"net"
|
||
|
|
"path/filepath"
|
||
|
|
"strconv"
|
||
|
|
"testing"
|
||
|
|
|
||
|
|
"gopkg.in/ini.v1"
|
||
|
|
|
||
|
|
"mailgoserver/internal/db"
|
||
|
|
)
|
||
|
|
|
||
|
|
func testCfg(t *testing.T, threshold int) *ini.File {
|
||
|
|
t.Helper()
|
||
|
|
cfg := ini.Empty()
|
||
|
|
sec, _ := cfg.NewSection("Security")
|
||
|
|
sec.NewKey("abuse_detection_enabled", "true")
|
||
|
|
sec.NewKey("abuse_failure_threshold", strconv.Itoa(threshold))
|
||
|
|
sec.NewKey("abuse_detection_window_minutes", "10")
|
||
|
|
sec.NewKey("abuse_blacklist_base_hours", "12")
|
||
|
|
sec.NewKey("abuse_blacklist_max_hours", "168")
|
||
|
|
return cfg
|
||
|
|
}
|
||
|
|
|
||
|
|
func openTestDB(t *testing.T) *db.DB {
|
||
|
|
t.Helper()
|
||
|
|
database, err := db.Open(filepath.Join(t.TempDir(), "test.db"))
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
t.Cleanup(func() { database.Close() })
|
||
|
|
return database
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRecordFailureAndMaybeBlacklistTripsThreshold(t *testing.T) {
|
||
|
|
database := openTestDB(t)
|
||
|
|
cfg := testCfg(t, 3)
|
||
|
|
const ip = "203.0.113.50"
|
||
|
|
|
||
|
|
for i := 0; i < 2; i++ {
|
||
|
|
database.LogAuthAttempt("sender", "victim@example.com", ip, false, "bad password")
|
||
|
|
RecordFailureAndMaybeBlacklist(database, cfg, nil, ip)
|
||
|
|
}
|
||
|
|
if blocked, _ := database.IsIPBlacklisted(ip); blocked {
|
||
|
|
t.Fatal("should not be blacklisted before threshold")
|
||
|
|
}
|
||
|
|
|
||
|
|
database.LogAuthAttempt("sender", "victim@example.com", ip, false, "bad password")
|
||
|
|
RecordFailureAndMaybeBlacklist(database, cfg, nil, ip)
|
||
|
|
|
||
|
|
blocked, err := database.IsIPBlacklisted(ip)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if !blocked {
|
||
|
|
t.Fatal("expected IP to be blacklisted after hitting the threshold")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRecordFailureAndMaybeBlacklistSkipsWhitelisted(t *testing.T) {
|
||
|
|
database := openTestDB(t)
|
||
|
|
cfg := testCfg(t, 2)
|
||
|
|
const ip = "203.0.113.51"
|
||
|
|
if err := database.AddAbuseWhitelist(ip, "trusted"); err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
|
||
|
|
for i := 0; i < 5; i++ {
|
||
|
|
database.LogAuthAttempt("sender", "victim@example.com", ip, false, "bad password")
|
||
|
|
RecordFailureAndMaybeBlacklist(database, cfg, nil, ip)
|
||
|
|
}
|
||
|
|
|
||
|
|
if blocked, _ := database.IsIPBlacklisted(ip); blocked {
|
||
|
|
t.Fatal("whitelisted IP should never be blacklisted")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// fakeListener yields exactly one already-open in-memory connection pair, then EOF-like
|
||
|
|
// closed errors, letting GuardListener's Accept loop be tested without real sockets.
|
||
|
|
type fakeListener struct {
|
||
|
|
conns chan net.Conn
|
||
|
|
done chan struct{}
|
||
|
|
}
|
||
|
|
|
||
|
|
func newFakeListener(conns ...net.Conn) *fakeListener {
|
||
|
|
ch := make(chan net.Conn, len(conns))
|
||
|
|
for _, c := range conns {
|
||
|
|
ch <- c
|
||
|
|
}
|
||
|
|
return &fakeListener{conns: ch, done: make(chan struct{})}
|
||
|
|
}
|
||
|
|
|
||
|
|
func (f *fakeListener) Accept() (net.Conn, error) {
|
||
|
|
select {
|
||
|
|
case c := <-f.conns:
|
||
|
|
return c, nil
|
||
|
|
case <-f.done:
|
||
|
|
return nil, net.ErrClosed
|
||
|
|
}
|
||
|
|
}
|
||
|
|
func (f *fakeListener) Close() error { close(f.done); return nil }
|
||
|
|
func (f *fakeListener) Addr() net.Addr { return dummyAddr{} }
|
||
|
|
|
||
|
|
type dummyAddr struct{}
|
||
|
|
|
||
|
|
func (dummyAddr) Network() string { return "tcp" }
|
||
|
|
func (dummyAddr) String() string { return "0.0.0.0:0" }
|
||
|
|
|
||
|
|
func TestGuardListenerRejectsBlacklistedIP(t *testing.T) {
|
||
|
|
database := openTestDB(t)
|
||
|
|
const blockedIP = "198.51.100.77"
|
||
|
|
if err := database.AddManualBlacklistEntry(blockedIP, "test", 1); err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
|
||
|
|
blockedConn, blockedPeer := net.Pipe()
|
||
|
|
defer blockedPeer.Close()
|
||
|
|
|
||
|
|
inner := newFakeListener(&addrOverrideConn{Conn: blockedConn, remote: hostPortAddr(blockedIP)})
|
||
|
|
guarded := GuardListener(inner, database, nil)
|
||
|
|
|
||
|
|
go func() {
|
||
|
|
guarded.Accept()
|
||
|
|
inner.Close()
|
||
|
|
}()
|
||
|
|
|
||
|
|
// The blocked connection's peer end should observe the connection close rather
|
||
|
|
// than any protocol banner, since GuardListener closes it before returning it.
|
||
|
|
buf := make([]byte, 1)
|
||
|
|
if _, err := blockedPeer.Read(buf); err == nil {
|
||
|
|
t.Fatal("expected blocked connection to be closed by GuardListener, got readable data instead")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
type addrOverrideConn struct {
|
||
|
|
net.Conn
|
||
|
|
remote net.Addr
|
||
|
|
}
|
||
|
|
|
||
|
|
func (c *addrOverrideConn) RemoteAddr() net.Addr { return c.remote }
|
||
|
|
|
||
|
|
type hostPortAddr string
|
||
|
|
|
||
|
|
func (hostPortAddr) Network() string { return "tcp" }
|
||
|
|
func (a hostPortAddr) String() string { return string(a) + ":12345" }
|