fix(06-12): fail closed on invalid limiter definitions
This commit is contained in:
@@ -2,6 +2,7 @@ package surf
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"math"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"strconv"
|
"strconv"
|
||||||
@@ -68,6 +69,16 @@ func (l *FixedWindowLimiter) RegisterBucket(pluginID, name string, b Bucket) err
|
|||||||
if name == "" {
|
if name == "" {
|
||||||
return fmt.Errorf("surf: plugin %q registered empty bucket", pluginID)
|
return fmt.Errorf("surf: plugin %q registered empty bucket", pluginID)
|
||||||
}
|
}
|
||||||
|
switch {
|
||||||
|
case l.store == nil:
|
||||||
|
return fmt.Errorf("surf: plugin %q bucket %q: limiter has no store", pluginID, name)
|
||||||
|
case b.Key == nil:
|
||||||
|
return fmt.Errorf("surf: plugin %q bucket %q has nil Key", pluginID, name)
|
||||||
|
case b.Max < 1:
|
||||||
|
return fmt.Errorf("surf: plugin %q bucket %q has Max < 1", pluginID, name)
|
||||||
|
case b.Decay <= 0:
|
||||||
|
return fmt.Errorf("surf: plugin %q bucket %q has non-positive Decay", pluginID, name)
|
||||||
|
}
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
defer l.mu.Unlock()
|
defer l.mu.Unlock()
|
||||||
if existing, ok := l.owners[name]; ok {
|
if existing, ok := l.owners[name]; ok {
|
||||||
@@ -82,16 +93,20 @@ func (l *FixedWindowLimiter) RegisterBucket(pluginID, name string, b Bucket) err
|
|||||||
// Middleware builds the throttle: factory body. param is either a
|
// Middleware builds the throttle: factory body. param is either a
|
||||||
// registered bucket name or a literal "N,M" pair (inline throttle).
|
// registered bucket name or a literal "N,M" pair (inline throttle).
|
||||||
func (l *FixedWindowLimiter) Middleware(param string) pact.Middleware {
|
func (l *FixedWindowLimiter) Middleware(param string) pact.Middleware {
|
||||||
|
failClosed := func(next http.Handler) http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusInternalServerError)
|
||||||
|
})
|
||||||
|
}
|
||||||
if l == nil {
|
if l == nil {
|
||||||
return func(next http.Handler) http.Handler { return next }
|
return failClosed
|
||||||
}
|
}
|
||||||
b, err := l.resolve(param)
|
b, err := l.resolve(param)
|
||||||
|
if err != nil || b.Key == nil || l.store == nil {
|
||||||
|
return failClosed
|
||||||
|
}
|
||||||
return func(next http.Handler) http.Handler {
|
return func(next http.Handler) http.Handler {
|
||||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
if err != nil || b.Key == nil || l.store == nil {
|
|
||||||
next.ServeHTTP(w, r)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
key := b.Key(r)
|
key := b.Key(r)
|
||||||
allowed, attempts, retryAfter := l.store.Attempt(key, b.Max, b.Decay)
|
allowed, attempts, retryAfter := l.store.Attempt(key, b.Max, b.Decay)
|
||||||
if !allowed {
|
if !allowed {
|
||||||
@@ -123,6 +138,9 @@ func (l *FixedWindowLimiter) ValidateThrottle(param string) error {
|
|||||||
if l == nil {
|
if l == nil {
|
||||||
return fmt.Errorf("surf: limiter is nil")
|
return fmt.Errorf("surf: limiter is nil")
|
||||||
}
|
}
|
||||||
|
if l.store == nil {
|
||||||
|
return fmt.Errorf("surf: limiter has no store")
|
||||||
|
}
|
||||||
_, err := l.resolve(param)
|
_, err := l.resolve(param)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -142,7 +160,7 @@ func (l *FixedWindowLimiter) resolve(param string) (Bucket, error) {
|
|||||||
}
|
}
|
||||||
n, errN := strconv.Atoi(strings.TrimSpace(nStr))
|
n, errN := strconv.Atoi(strings.TrimSpace(nStr))
|
||||||
m, errM := strconv.Atoi(strings.TrimSpace(mStr))
|
m, errM := strconv.Atoi(strings.TrimSpace(mStr))
|
||||||
if errN != nil || errM != nil || n < 1 || m < 1 {
|
if errN != nil || errM != nil || n < 1 || m < 1 || int64(m) > math.MaxInt64/int64(time.Minute) {
|
||||||
return Bucket{}, fmt.Errorf("surf: malformed throttle %q", param)
|
return Bucket{}, fmt.Errorf("surf: malformed throttle %q", param)
|
||||||
}
|
}
|
||||||
trusted := l.trusted
|
trusted := l.trusted
|
||||||
|
|||||||
@@ -389,3 +389,17 @@ func stackThrottle(lim *FixedWindowLimiter, next http.Handler, names ...string)
|
|||||||
}
|
}
|
||||||
|
|
||||||
var _ pact.Router = (*Router)(nil)
|
var _ pact.Router = (*Router)(nil)
|
||||||
|
|
||||||
|
func TestRegisterBucketRejectsInvalidDefinitions(t *testing.T) {
|
||||||
|
l := NewFixedWindowLimiter(NewMemoryStore(time.Minute), nil)
|
||||||
|
key := func(*http.Request) string { return "k" }
|
||||||
|
if err := l.RegisterBucket("p", "n", Bucket{Max: 1, Decay: 0, Key: key}); err == nil {
|
||||||
|
t.Fatal("zero decay accepted")
|
||||||
|
}
|
||||||
|
if err := l.RegisterBucket("p", "n", Bucket{Max: 1, Decay: time.Minute}); err == nil {
|
||||||
|
t.Fatal("nil key accepted")
|
||||||
|
}
|
||||||
|
if err := l.ValidateThrottle("1,9223372036854775807"); err == nil {
|
||||||
|
t.Fatal("overflowing minutes accepted")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user