internal/sshd/ratelimit.go

v1.5.0
gitbay/internal/sshd/ratelimit.go history · blame · raw

91 lines · 1975 bytes

 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}