internal/sshd/ratelimit.go
91 lines · 1975 bytes
8 symbols in this file
1package sshd
2
3import (
4 "net"
5 "sync"
6 "time"
7)
8
9// rateLimiter throttles per-IP authentication FAILURES: successful auths
10// never count (a busy CLI makes many connections per minute) and clear the
11// IP's slate. limits.ssh_auth_rate failures per window lock the IP out
12// until the window passes.
13type rateLimiter struct {
14 mu sync.Mutex
15 limit int
16 window time.Duration
17 seen map[string]*ipWindow
18}
19
20type ipWindow struct {
21 start time.Time
22 count int
23 audited bool
24}
25
26func newRateLimiter(limit int, window time.Duration) *rateLimiter {
27 return &rateLimiter{limit: limit, window: window, seen: map[string]*ipWindow{}}
28}
29
30// allow reports whether ip may attempt authentication at all.
31func (r *rateLimiter) allow(ip string) bool {
32 if r.limit <= 0 {
33 return true
34 }
35 r.mu.Lock()
36 defer r.mu.Unlock()
37 now := time.Now()
38 if len(r.seen) > 4096 {
39 for k, w := range r.seen {
40 if now.Sub(w.start) > r.window {
41 delete(r.seen, k)
42 }
43 }
44 }
45 w := r.seen[ip]
46 if w == nil || now.Sub(w.start) > r.window {
47 delete(r.seen, ip)
48 return true
49 }
50 return w.count < r.limit
51}
52
53// fail records an authentication failure for ip.
54func (r *rateLimiter) fail(ip string) {
55 r.mu.Lock()
56 defer r.mu.Unlock()
57 now := time.Now()
58 w := r.seen[ip]
59 if w == nil || now.Sub(w.start) > r.window {
60 r.seen[ip] = &ipWindow{start: now, count: 1}
61 return
62 }
63 w.count++
64}
65
66// success clears the IP's failure slate.
67func (r *rateLimiter) success(ip string) {
68 r.mu.Lock()
69 defer r.mu.Unlock()
70 delete(r.seen, ip)
71}
72
73// firstThrottle reports true exactly once per throttled window, so the
74// audit log records a burst rather than every rejected attempt.
75func (r *rateLimiter) firstThrottle(ip string) bool {
76 r.mu.Lock()
77 defer r.mu.Unlock()
78 w := r.seen[ip]
79 if w == nil || w.audited {
80 return false
81 }
82 w.audited = true
83 return true
84}
85
86func remoteIP(addr net.Addr) string {
87 if host, _, err := net.SplitHostPort(addr.String()); err == nil {
88 return host
89 }
90 return addr.String()
91}