Files
silo-server/internal/api/handlers/rate_limits_test.go

356 lines
11 KiB
Go

package handlers
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/Silo-Server/silo-server/internal/ratelimit"
)
type fakeRateLimitStore struct {
values map[string]string
setCalls int
setManyCalls int
atomicCalls int
}
func newFakeRateLimitStore() *fakeRateLimitStore {
return &fakeRateLimitStore{values: make(map[string]string)}
}
func (s *fakeRateLimitStore) Get(_ context.Context, key string) (string, error) {
return s.values[key], nil
}
func (s *fakeRateLimitStore) Set(_ context.Context, key, value string) error {
s.setCalls++
s.values[key] = value
return nil
}
func (s *fakeRateLimitStore) SetMany(_ context.Context, values map[string]string) error {
s.setManyCalls++
for key, value := range values {
s.values[key] = value
}
return nil
}
func (s *fakeRateLimitStore) GetAll(_ context.Context) (map[string]string, error) {
out := make(map[string]string, len(s.values))
for k, v := range s.values {
out[k] = v
}
return out, nil
}
func (s *fakeRateLimitStore) UpdateAtomic(
ctx context.Context,
update func(current map[string]string) (map[string]string, error),
) error {
s.atomicCalls++
current, err := s.GetAll(ctx)
if err != nil {
return err
}
writes, err := update(current)
if err != nil || len(writes) == 0 {
return err
}
return s.SetMany(ctx, writes)
}
type postCommitOverwriteRateLimitStore struct {
*fakeRateLimitStore
}
func (s *postCommitOverwriteRateLimitStore) UpdateAtomic(
ctx context.Context,
update func(current map[string]string) (map[string]string, error),
) error {
s.atomicCalls++
current, err := s.GetAll(ctx)
if err != nil {
return err
}
writes, err := update(current)
if err != nil || len(writes) == 0 {
return err
}
if err := s.SetMany(ctx, writes); err != nil {
return err
}
// Model a later request committing before this handler performs its
// post-commit reload and restart decision.
s.values["ratelimit.enabled"] = "false"
return nil
}
func newRunningRateLimitMiddleware(t *testing.T, store ratelimit.SettingsStore) *ratelimit.Middleware {
t.Helper()
perKey := ratelimit.NewMemoryLimiter()
global := ratelimit.NewMemoryLimiter()
t.Cleanup(func() {
perKey.Close()
global.Close()
})
mw := ratelimit.NewMiddleware(perKey, global, store, true)
if err := mw.Init(context.Background()); err != nil {
t.Fatalf("init middleware: %v", err)
}
return mw
}
func putRateLimitConfig(t *testing.T, h *RateLimitHandler, body string) map[string]any {
t.Helper()
req := httptest.NewRequest(http.MethodPut, "/admin/rate-limits/config", strings.NewReader(body))
rec := httptest.NewRecorder()
h.HandleUpdateConfig(rec, req)
if rec.Code != http.StatusOK {
t.Fatalf("PUT status = %d, body = %s", rec.Code, rec.Body.String())
}
var resp map[string]any
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
t.Fatalf("decode PUT response: %v", err)
}
return resp
}
func getRateLimitConfig(t *testing.T, h *RateLimitHandler) rateLimitConfigResponse {
t.Helper()
req := httptest.NewRequest(http.MethodGet, "/admin/rate-limits/config", nil)
rec := httptest.NewRecorder()
h.HandleGetConfig(rec, req)
if rec.Code != http.StatusOK {
t.Fatalf("GET status = %d, body = %s", rec.Code, rec.Body.String())
}
var resp rateLimitConfigResponse
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
t.Fatalf("decode GET response: %v", err)
}
return resp
}
func TestRateLimitHandlerWithoutRunningLimiter(t *testing.T) {
// The limiter is nil when rate limiting was disabled at boot. The config
// endpoints must still work so an admin can re-enable it from the UI.
store := newFakeRateLimitStore()
store.values["ratelimit.enabled"] = "false"
tracker := NewServerRestartStatusTracker()
h := NewRateLimitHandler(store, nil, nil, tracker, true)
got := getRateLimitConfig(t, h)
if got.Enabled {
t.Error("GET enabled = true, want false")
}
if got.Active {
t.Error("GET active = true, want false when no limiter is running")
}
resp := putRateLimitConfig(t, h, `{"enabled":true,"backend":"redis"}`)
if resp["restart_required"] != true {
t.Errorf("enabling without a running limiter: restart_required = %v, want true", resp["restart_required"])
}
if store.values["ratelimit.enabled"] != "true" {
t.Errorf("ratelimit.enabled = %q, want \"true\"", store.values["ratelimit.enabled"])
}
if store.values["ratelimit.backend"] != "redis" {
t.Errorf("ratelimit.backend = %q, want \"redis\"", store.values["ratelimit.backend"])
}
if store.setManyCalls != 1 {
t.Fatalf("atomic writes = %d, want 1", store.setManyCalls)
}
if !tracker.Snapshot().RestartRequired {
t.Fatal("enabling a stopped limiter did not update server restart status")
}
resp = putRateLimitConfig(t, h, `{"enabled":false,"backend":"redis"}`)
if resp["restart_required"] != false {
t.Errorf("saving disabled config without a running limiter: restart_required = %v, want false", resp["restart_required"])
}
}
func TestRateLimitHandlerUsesLatestCommittedSettingsAfterAtomicWrite(t *testing.T) {
baseStore := newFakeRateLimitStore()
baseStore.values["ratelimit.enabled"] = "false"
store := &postCommitOverwriteRateLimitStore{fakeRateLimitStore: baseStore}
h := NewRateLimitHandler(store, nil, nil, NewServerRestartStatusTracker())
resp := putRateLimitConfig(t, h, `{"enabled":true}`)
if resp["restart_required"] != false {
t.Fatalf("restart_required = %v, want false from latest committed disabled state", resp["restart_required"])
}
if got := store.values["ratelimit.enabled"]; got != "false" {
t.Fatalf("ratelimit.enabled = %q, want later committed value false", got)
}
}
func TestRateLimitHandlerWithRunningLimiter(t *testing.T) {
store := newFakeRateLimitStore()
mw := newRunningRateLimitMiddleware(t, store)
h := NewRateLimitHandler(store, mw, nil, NewServerRestartStatusTracker(), true)
got := getRateLimitConfig(t, h)
if !got.Active {
t.Error("GET active = false, want true with a running limiter")
}
if got.ActiveBackend != "memory" {
t.Errorf("GET active_backend = %q, want \"memory\"", got.ActiveBackend)
}
// Toggling enabled hot-reloads; no restart needed when the backend matches.
resp := putRateLimitConfig(t, h, `{"enabled":false,"backend":"memory"}`)
if resp["restart_required"] != false {
t.Errorf("toggle on running limiter: restart_required = %v, want false", resp["restart_required"])
}
// Switching backend only takes effect at boot.
resp = putRateLimitConfig(t, h, `{"enabled":true,"backend":"redis"}`)
if resp["restart_required"] != true {
t.Errorf("backend change: restart_required = %v, want true", resp["restart_required"])
}
}
func TestRateLimitHandlerRejectsExplicitZeroWithoutWriting(t *testing.T) {
store := newFakeRateLimitStore()
h := NewRateLimitHandler(store, nil, nil, NewServerRestartStatusTracker())
req := httptest.NewRequest(http.MethodPut, "/admin/rate-limits/config", strings.NewReader(
`{"tiers":{"standard":{"requests_per_second":0}}}`,
))
rec := httptest.NewRecorder()
h.HandleUpdateConfig(rec, req)
if rec.Code != http.StatusBadRequest {
t.Fatalf("PUT status = %d, want 400; body = %s", rec.Code, rec.Body.String())
}
if store.setManyCalls != 0 || store.setCalls != 0 {
t.Fatalf("invalid request wrote settings: SetMany=%d Set=%d", store.setManyCalls, store.setCalls)
}
}
func TestRateLimitHandlerRejectsOverflowingGlobalRateWithoutWriting(t *testing.T) {
store := newFakeRateLimitStore()
h := NewRateLimitHandler(store, nil, nil, NewServerRestartStatusTracker())
req := httptest.NewRequest(http.MethodPut, "/admin/rate-limits/config", strings.NewReader(
`{"global_requests_per_second":1e308}`,
))
rec := httptest.NewRecorder()
h.HandleUpdateConfig(rec, req)
if rec.Code != http.StatusBadRequest {
t.Fatalf("PUT status = %d, want 400; body = %s", rec.Code, rec.Body.String())
}
if store.setManyCalls != 0 || store.setCalls != 0 {
t.Fatalf("invalid request wrote settings: SetMany=%d Set=%d", store.setManyCalls, store.setCalls)
}
}
func TestRateLimitHandlerRejectsOverflowingBurstWithoutWriting(t *testing.T) {
store := newFakeRateLimitStore()
h := NewRateLimitHandler(store, nil, nil, NewServerRestartStatusTracker())
req := httptest.NewRequest(http.MethodPut, "/admin/rate-limits/config", strings.NewReader(
`{"ip_burst":9223372036854775807}`,
))
rec := httptest.NewRecorder()
h.HandleUpdateConfig(rec, req)
if rec.Code != http.StatusBadRequest {
t.Fatalf("PUT status = %d, want 400; body = %s", rec.Code, rec.Body.String())
}
if store.setManyCalls != 0 || store.setCalls != 0 {
t.Fatalf("invalid request wrote settings: SetMany=%d Set=%d", store.setManyCalls, store.setCalls)
}
}
func TestRateLimitHandlerRejectsRedisWithoutRedisConfiguration(t *testing.T) {
store := newFakeRateLimitStore()
h := NewRateLimitHandler(store, nil, nil, NewServerRestartStatusTracker())
req := httptest.NewRequest(http.MethodPut, "/admin/rate-limits/config", strings.NewReader(
`{"backend":"redis"}`,
))
rec := httptest.NewRecorder()
h.HandleUpdateConfig(rec, req)
if rec.Code != http.StatusBadRequest {
t.Fatalf("PUT status = %d, want 400; body = %s", rec.Code, rec.Body.String())
}
if store.setManyCalls != 0 {
t.Fatalf("invalid Redis selection wrote settings: SetMany=%d", store.setManyCalls)
}
}
func TestRateLimitHandlerRejectsMalformedPersistedRedisURL(t *testing.T) {
store := newFakeRateLimitStore()
store.values["redis.url"] = "not-a-url"
h := NewRateLimitHandler(store, nil, nil, NewServerRestartStatusTracker())
req := httptest.NewRequest(http.MethodPut, "/admin/rate-limits/config", strings.NewReader(
`{"backend":"redis"}`,
))
rec := httptest.NewRecorder()
h.HandleUpdateConfig(rec, req)
if rec.Code != http.StatusBadRequest {
t.Fatalf("PUT status = %d, want 400; body = %s", rec.Code, rec.Body.String())
}
if store.setManyCalls != 0 {
t.Fatalf("invalid Redis selection wrote settings: SetMany=%d", store.setManyCalls)
}
}
func TestRateLimitHandlerAcceptsBootstrapRedisDespiteStalePersistedURL(t *testing.T) {
store := newFakeRateLimitStore()
store.values["redis.url"] = " redis://cache.example.invalid:6379 "
h := NewRateLimitHandler(store, nil, nil, NewServerRestartStatusTracker(), true)
req := httptest.NewRequest(http.MethodPut, "/admin/rate-limits/config", strings.NewReader(
`{"backend":"redis"}`,
))
rec := httptest.NewRecorder()
h.HandleUpdateConfig(rec, req)
if rec.Code != http.StatusOK {
t.Fatalf("PUT status = %d, want 200; body = %s", rec.Code, rec.Body.String())
}
if store.values["ratelimit.backend"] != "redis" {
t.Fatalf("ratelimit.backend = %q, want redis", store.values["ratelimit.backend"])
}
}
func TestRateLimitHandlerDoesNotTreatActiveRedisAsDurableConfiguration(t *testing.T) {
store := newFakeRateLimitStore()
perKey := ratelimit.NewMemoryLimiter()
global := ratelimit.NewMemoryLimiter()
t.Cleanup(func() {
perKey.Close()
global.Close()
})
// isMemory=false models a process currently using Redis. With no stored or
// bootstrap transport, that active state cannot survive the next restart.
mw := ratelimit.NewMiddleware(perKey, global, store, false)
h := NewRateLimitHandler(store, mw, nil, NewServerRestartStatusTracker())
req := httptest.NewRequest(http.MethodPut, "/admin/rate-limits/config", strings.NewReader(
`{"backend":"redis"}`,
))
rec := httptest.NewRecorder()
h.HandleUpdateConfig(rec, req)
if rec.Code != http.StatusBadRequest {
t.Fatalf("PUT status = %d, want 400; body = %s", rec.Code, rec.Body.String())
}
if store.setManyCalls != 0 {
t.Fatalf("invalid Redis selection wrote settings: SetMany=%d", store.setManyCalls)
}
}