diff options
| author | Kyren223 <Kyren223@proton.me> | 2025-07-15 21:05:31 +0300 |
|---|---|---|
| committer | Kyren223 <Kyren223@proton.me> | 2025-07-19 18:32:12 +0300 |
| commit | 52eb471e8901cc75525c3b5b7640fd4703ab699b (patch) | |
| tree | e5f31f42f00369d17103012724922f834069375c | |
| parent | cb3c86624821a3f5e7faf2d6a80cdd760ba837ad (diff) | |
Implemented connection-level rate limiting (to avoid reconnection
abuse), added a file for testing that, still needs to add observability
for this
| -rw-r--r-- | internal/server/server.go | 88 | ||||
| -rw-r--r-- | tools/test_rate_limit.go | 60 |
2 files changed, 147 insertions, 1 deletions
diff --git a/internal/server/server.go b/internal/server/server.go index 0620793..90d24a3 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -8,6 +8,7 @@ import ( "crypto/tls" "crypto/x509" "crypto/x509/pkix" + "encoding/binary" "encoding/pem" "errors" "fmt" @@ -36,7 +37,12 @@ var nodeId int64 = 0 const CertFile = "EKO_SERVER_CERT_FILE" -const ReadCheckCancelledInterval = 1 * time.Second +const ( + ReadCheckCancelledInterval = 1 * time.Second + RateLimitWindowSize = 1 * time.Second + RateLimitCountThresholdSus = 3 + RateLimitCountThresholdMalicious = 10 +) func getTlsConfig() *tls.Config { path, ok := os.LookupEnv(CertFile) @@ -105,6 +111,10 @@ type server struct { sessions map[snowflake.ID]*session.Session sessMu sync.RWMutex Port uint16 + ipConns map[uint32]struct { + start time.Time + count uint8 + } } // Creates a new server on the given port. @@ -119,6 +129,11 @@ func NewServer(ctx context.Context, port uint16) server { node: node, sessions: map[snowflake.ID]*session.Session{}, Port: port, + sessMu: sync.RWMutex{}, + ipConns: map[uint32]struct { + start time.Time + count uint8 + }{}, } } @@ -214,8 +229,16 @@ func (s *server) Run() { slog.Info("server context done", "error", s.ctx.Err()) break } + continue // Ignore and skip (don't connect) } + + ip := binary.BigEndian.Uint32(conn.RemoteAddr().(*net.TCPAddr).IP.To4()) + if s.isRateLimited(ip) { + _ = conn.Close() + continue + } + wg.Add(1) go func() { metrics.ConnectionsEstablished.Inc() @@ -592,3 +615,66 @@ func TokensPerRequest(requestType packet.PacketType) float64 { return 1 } + +func (s *server) isRateLimited(ip uint32) bool { + if entry, ok := s.ipConns[ip]; ok { + outsideWindow := time.Since(entry.start) > RateLimitWindowSize + notMalicious := entry.count < RateLimitCountThresholdMalicious + if outsideWindow && notMalicious { + entry.start = time.Now().UTC() + entry.count = 1 + + s.ipConns[ip] = entry + return false + } + + ipStr := formatIPv4(ip) + + if entry.count < RateLimitCountThresholdSus { + slog.Info("connection activity", "ip", ipStr, "count", entry.count) + entry.count++ + s.ipConns[ip] = entry + return false + } else if entry.count < RateLimitCountThresholdMalicious { + if entry.count == RateLimitCountThresholdSus { + slog.Warn("suspicious connection activity", "ip", ipStr, "count", entry.count) + // Only log the first one + } + entry.count++ + s.ipConns[ip] = entry + return true + } else { + if entry.count == RateLimitCountThresholdMalicious { + slog.Warn("potential malicious connection behavior", "ip", ipStr, "count", entry.count) + // Only log the first one + entry.count++ + s.ipConns[ip] = entry + // Update so it doesn't spam + } + // Don't bother to update counts, save resources + return true + } + } + + s.ipConns[ip] = struct { + start time.Time + count uint8 + }{ + start: time.Now().UTC(), + count: 1, + } + + return false +} + +func formatIPv4(ip uint32) string { + var b [15]byte // max len for "255.255.255.255" + n := strconv.AppendUint(b[:0], uint64(ip>>24), 10) + n = append(n, '.') + n = strconv.AppendUint(n, uint64((ip>>16)&0xFF), 10) + n = append(n, '.') + n = strconv.AppendUint(n, uint64((ip>>8)&0xFF), 10) + n = append(n, '.') + n = strconv.AppendUint(n, uint64(ip&0xFF), 10) + return string(n) +} diff --git a/tools/test_rate_limit.go b/tools/test_rate_limit.go new file mode 100644 index 0000000..7c25d20 --- /dev/null +++ b/tools/test_rate_limit.go @@ -0,0 +1,60 @@ +package main + +import ( + "crypto/tls" + "fmt" + "log" + "net" + "time" +) + +const addr = "localhost:7223" + +var tlsConfig = &tls.Config{ + InsecureSkipVerify: true, // skip cert verification +} + +func connect(n int, label string) { + fmt.Println("----", label) + conns := make([]net.Conn, 0, n) + for i := 0; i < n; i++ { + conn, err := tls.Dial("tcp", addr, tlsConfig) + if err != nil { + log.Printf("connect %d failed: %v", i, err) + continue + } + conns = append(conns, conn) + } + time.Sleep(300 * time.Millisecond) + for _, c := range conns { + c.Close() + } + time.Sleep(200 * time.Millisecond) +} + +func wait() { + time.Sleep(1100 * time.Millisecond) // ensure we roll over fixed 1s window +} + +func main() { + // 1. Single connection, then disconnect, should not rate limit + connect(1, "single connection") + wait() + + // 2. Two connections (under threshold), should show info + connect(2, "two connections") + wait() + connect(2, "two connections again") + wait() + + // 3. Hit 5 times to trigger suspicious threshold once + connect(5, "5 suspicious connections") + wait() + connect(5, "5 more suspicious") + wait() + + // 4. Hit 15 times to trigger malicious + connect(15, "15 malicious connections") + wait() + connect(15, "15 more malicious connections") +} |
