summaryrefslogtreecommitdiff
path: root/internal
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 /internal
parentcb3c86624821a3f5e7faf2d6a80cdd760ba837ad (diff)
Implemented connection-level rate limiting (to avoid reconnection
abuse), added a file for testing that, still needs to add observability for this
Diffstat (limited to 'internal')
-rw-r--r--internal/server/server.go88
1 files changed, 87 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)
+}