summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorKyren223 <Kyren223@proton.me>2025-07-15 21:05:31 +0300
committerKyren223 <Kyren223@proton.me>2025-07-19 18:32:12 +0300
commit52eb471e8901cc75525c3b5b7640fd4703ab699b (patch)
treee5f31f42f00369d17103012724922f834069375c
parentcb3c86624821a3f5e7faf2d6a80cdd760ba837ad (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.go88
-rw-r--r--tools/test_rate_limit.go60
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")
+}