master
go 215 lines 3.78 KB
Raw
1 // SPDX-License-Identifier: GPL-3.0-or-later
2
3 package logger
4
5 import (
6 "fmt"
7 "log/slog"
8 "sync"
9 "time"
10 )
11
12 const (
13 rateLimiterSweepEvery = 4096
14 rateLimiterTTL = time.Hour
15 )
16
17 type limitMode uint8
18
19 const (
20 modeOnce limitMode = iota + 1
21 modeLimit
22 )
23
24 type limitKey struct {
25 mode limitMode
26 key string
27 }
28
29 type limitEntry struct {
30 limit int
31 window time.Duration
32 count int
33 windowStart time.Time
34 lastSeen time.Time
35 }
36
37 type rateLimiter struct {
38 mu sync.Mutex
39 entries map[limitKey]*limitEntry
40 sweepEvery uint64
41 ttl time.Duration
42 calls uint64
43 now func() time.Time
44 }
45
46 func newRateLimiter() *rateLimiter {
47 return &rateLimiter{
48 entries: make(map[limitKey]*limitEntry),
49 sweepEvery: rateLimiterSweepEvery,
50 ttl: rateLimiterTTL,
51 now: time.Now,
52 }
53 }
54
55 type LimitedLogger struct {
56 l *Logger
57 mode limitMode
58 key string
59 limit int
60 window time.Duration
61 }
62
63 func (l *Logger) Once(key string) LimitedLogger {
64 return LimitedLogger{
65 l: l,
66 mode: modeOnce,
67 key: key,
68 limit: 1,
69 window: 0,
70 }
71 }
72
73 func (l *Logger) Limit(key string, n int, d time.Duration) LimitedLogger {
74 if n <= 0 {
75 n = 1
76 }
77 if d < 0 {
78 d = 0
79 }
80 return LimitedLogger{
81 l: l,
82 mode: modeLimit,
83 key: key,
84 limit: n,
85 window: d,
86 }
87 }
88
89 func (l *Logger) ResetAllOnce() {
90 if l == nil || l.rl == nil {
91 return
92 }
93 l.rl.resetAllMode(modeOnce)
94 }
95
96 func (l LimitedLogger) Error(a ...any) {
97 l.logArgs(slog.LevelError, a...)
98 }
99
100 func (l LimitedLogger) Warning(a ...any) {
101 l.logArgs(slog.LevelWarn, a...)
102 }
103
104 func (l LimitedLogger) Notice(a ...any) {
105 l.logArgs(levelNotice, a...)
106 }
107
108 func (l LimitedLogger) Info(a ...any) {
109 l.logArgs(slog.LevelInfo, a...)
110 }
111
112 func (l LimitedLogger) Debug(a ...any) {
113 l.logArgs(slog.LevelDebug, a...)
114 }
115
116 func (l LimitedLogger) Errorf(format string, a ...any) {
117 l.logf(slog.LevelError, format, a...)
118 }
119
120 func (l LimitedLogger) Warningf(format string, a ...any) {
121 l.logf(slog.LevelWarn, format, a...)
122 }
123
124 func (l LimitedLogger) Noticef(format string, a ...any) {
125 l.logf(levelNotice, format, a...)
126 }
127
128 func (l LimitedLogger) Infof(format string, a ...any) {
129 l.logf(slog.LevelInfo, format, a...)
130 }
131
132 func (l LimitedLogger) Debugf(format string, a ...any) {
133 l.logf(slog.LevelDebug, format, a...)
134 }
135
136 func (l LimitedLogger) logArgs(level slog.Level, a ...any) {
137 if !l.l.canLog(level) || !l.allow() {
138 return
139 }
140 l.l.log(level, fmt.Sprint(a...))
141 }
142
143 func (l LimitedLogger) logf(level slog.Level, format string, a ...any) {
144 if !l.l.canLog(level) || !l.allow() {
145 return
146 }
147 l.l.log(level, fmt.Sprintf(format, a...))
148 }
149
150 func (l LimitedLogger) allow() bool {
151 if l.l == nil || l.l.rl == nil {
152 // Preserve nil logger behavior: no panic and no additional suppression.
153 return true
154 }
155 return l.l.rl.allow(l.mode, l.key, l.limit, l.window)
156 }
157
158 func (r *rateLimiter) allow(mode limitMode, key string, limit int, window time.Duration) bool {
159 now := r.now()
160 allow := false
161
162 r.mu.Lock()
163 defer r.mu.Unlock()
164
165 r.calls++
166 if r.sweepEvery > 0 && r.calls%r.sweepEvery == 0 {
167 r.sweep(now)
168 }
169
170 k := limitKey{mode: mode, key: key}
171 e, ok := r.entries[k]
172 if !ok {
173 e = &limitEntry{
174 limit: limit,
175 window: window,
176 windowStart: now,
177 }
178 r.entries[k] = e
179 }
180
181 e.lastSeen = now
182
183 // d==0 means infinite window, so rollover must be disabled.
184 if e.window > 0 && now.Sub(e.windowStart) >= e.window {
185 e.windowStart = now
186 e.count = 0
187 }
188
189 if e.count < e.limit {
190 e.count++
191 allow = true
192 }
193
194 return allow
195 }
196
197 func (r *rateLimiter) resetAllMode(mode limitMode) {
198 r.mu.Lock()
199 defer r.mu.Unlock()
200
201 for k := range r.entries {
202 if k.mode == mode {
203 delete(r.entries, k)
204 }
205 }
206 }
207
208 func (r *rateLimiter) sweep(now time.Time) {
209 cutoff := now.Add(-r.ttl)
210 for k, e := range r.entries {
211 if e.lastSeen.Before(cutoff) {
212 delete(r.entries, k)
213 }
214 }
215 }