From fca23d77e902be201c13a44598df9af45f6bae75 Mon Sep 17 00:00:00 2001 From: Silo Server Migration Date: Mon, 25 May 2026 00:12:45 -0400 Subject: [PATCH] fix(requests): serialize quota check and clamp list limits CreateRequest previously read the user's request count outside the insert transaction, so two concurrent submissions at MaxRequests-1 could both pass the quota gate and end up at MaxRequests+1. Move the count inside the same transaction as the insert and acquire a per-user advisory lock so concurrent inserts serialize. The store reports ErrQuotaExceeded when the racing path catches the user at the limit and the service maps it back to QuotaError. normalizeListFilter previously reset limit to 50 when callers asked for more than 100, which is surprising. Clamp to the cap instead so a request for 150 returns 100 and a request for 1_000_000 still cannot hit the database with an unbounded scan. Co-Authored-By: Claude Opus 4.7 (1M context) --- internal/requests/repository.go | 24 ++++++ internal/requests/service.go | 31 ++++++- internal/requests/service_test.go | 137 +++++++++++++++++++++++++++++- internal/requests/store.go | 11 +++ 4 files changed, 197 insertions(+), 6 deletions(-) diff --git a/internal/requests/repository.go b/internal/requests/repository.go index 04a06eff..c8be6a10 100644 --- a/internal/requests/repository.go +++ b/internal/requests/repository.go @@ -182,6 +182,11 @@ func (r *Repository) ListActiveByTMDB(ctx context.Context, mediaType MediaType, return out, nil } +// quotaLockNamespace partitions advisory locks so request-quota locks do not +// collide with advisory locks held elsewhere in the database. The value is +// arbitrary; what matters is that it is stable. +const quotaLockNamespace = 139 + func (r *Repository) CreateRequest(ctx context.Context, input CreateRequestRecord) (*Request, error) { tx, err := r.pool.Begin(ctx) if err != nil { @@ -189,6 +194,25 @@ func (r *Repository) CreateRequest(ctx context.Context, input CreateRequestRecor } defer tx.Rollback(ctx) + if input.Quota != nil { + if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock($1::int4, $2::int4)`, + quotaLockNamespace, input.Quota.UserID); err != nil { + return nil, fmt.Errorf("acquire request quota lock: %w", err) + } + var count int + if err := tx.QueryRow(ctx, ` + SELECT COUNT(*) + FROM media_requests + WHERE requested_by_user_id = $1 + AND created_at >= $2 + `, input.Quota.UserID, input.Quota.WindowStart).Scan(&count); err != nil { + return nil, fmt.Errorf("count requests for quota: %w", err) + } + if count >= input.Quota.MaxRequests { + return nil, ErrQuotaExceeded + } + } + now := input.Now if now.IsZero() { now = time.Now().UTC() diff --git a/internal/requests/service.go b/internal/requests/service.go index 1ca639b2..974d9e35 100644 --- a/internal/requests/service.go +++ b/internal/requests/service.go @@ -287,18 +287,33 @@ func (s *Service) CreateRequest(ctx context.Context, viewer Viewer, input Create status = StatusApproved } } - req, err := s.store.CreateRequest(ctx, CreateRequestRecord{ + record := CreateRequestRecord{ ID: id, Input: normalized, Status: status, Outcome: OutcomeActive, Requester: viewer, Now: s.now(), - }) + } + if !policy.Unlimited { + record.Quota = &QuotaCheck{ + UserID: viewer.UserID, + WindowStart: policy.WindowStart, + MaxRequests: policy.MaxRequests, + } + } + req, err := s.store.CreateRequest(ctx, record) if err != nil { if errors.Is(err, ErrAlreadyRequested) { return nil, ErrAlreadyRequested } + if errors.Is(err, ErrQuotaExceeded) { + return nil, QuotaError{ + Used: policy.MaxRequests, + Limit: policy.MaxRequests, + WindowDays: policy.WindowDays, + } + } return nil, err } if req.Status == StatusApproved { @@ -1250,9 +1265,17 @@ func normalizeSearchMediaType(mediaType MediaType) (MediaType, error) { } } +const ( + defaultRequestListLimit = 50 + maxRequestListLimit = 100 +) + func normalizeListFilter(filter ListFilter) ListFilter { - if filter.Limit <= 0 || filter.Limit > 100 { - filter.Limit = 50 + if filter.Limit <= 0 { + filter.Limit = defaultRequestListLimit + } + if filter.Limit > maxRequestListLimit { + filter.Limit = maxRequestListLimit } if filter.Offset < 0 { filter.Offset = 0 diff --git a/internal/requests/service_test.go b/internal/requests/service_test.go index 87c2824d..7efe6fb8 100644 --- a/internal/requests/service_test.go +++ b/internal/requests/service_test.go @@ -35,6 +35,82 @@ func TestCreateRequestQuotaExceeded(t *testing.T) { } } +func TestNormalizeListFilterCapsLimit(t *testing.T) { + cases := []struct { + name string + in ListFilter + wantLim int + wantOff int + }{ + {"zero defaults", ListFilter{}, defaultRequestListLimit, 0}, + {"negative defaults", ListFilter{Limit: -10, Offset: -5}, defaultRequestListLimit, 0}, + {"under cap preserved", ListFilter{Limit: 75, Offset: 10}, 75, 10}, + {"at cap preserved", ListFilter{Limit: maxRequestListLimit, Offset: 0}, maxRequestListLimit, 0}, + {"over cap clamped", ListFilter{Limit: 1_000_000}, maxRequestListLimit, 0}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got := normalizeListFilter(tc.in) + if got.Limit != tc.wantLim { + t.Errorf("limit = %d, want %d", got.Limit, tc.wantLim) + } + if got.Offset != tc.wantOff { + t.Errorf("offset = %d, want %d", got.Offset, tc.wantOff) + } + }) + } +} + +func TestCreateRequestConcurrentSubmissionsRespectQuota(t *testing.T) { + const ( + maxRequests = 5 + goroutines = 20 + ) + store := newFakeStore() + store.settings.RequestsEnabled = true + store.settings.GlobalMaxRequests = maxRequests + service := newTestService(store) + + var ( + wg sync.WaitGroup + successMu sync.Mutex + successes int + quotaFails int + ) + for i := 0; i < goroutines; i++ { + wg.Add(1) + go func(tmdbID int) { + defer wg.Done() + _, err := service.CreateRequest(context.Background(), testViewer(1), CreateRequestInput{ + MediaType: MediaTypeMovie, + TMDBID: tmdbID, + Title: "Title", + }) + successMu.Lock() + defer successMu.Unlock() + if err == nil { + successes++ + return + } + var quota QuotaError + if errors.As(err, "a) { + quotaFails++ + } + }(1000 + i) + } + wg.Wait() + + if successes != maxRequests { + t.Fatalf("successful creations = %d, want %d", successes, maxRequests) + } + if successes+quotaFails != goroutines { + t.Fatalf("non-quota errors: successes=%d quotaFails=%d total=%d", successes, quotaFails, goroutines) + } + if len(store.created) != maxRequests { + t.Fatalf("stored creations = %d, want %d", len(store.created), maxRequests) + } +} + func TestCreateRequestActiveDuplicateBlocks(t *testing.T) { store := newFakeStore() store.settings.RequestsEnabled = true @@ -648,6 +724,7 @@ func testViewer(userID int) Viewer { } type fakeStore struct { + mu sync.Mutex settings Settings limit *UserLimit count int @@ -675,28 +752,50 @@ func newFakeStore() *fakeStore { } func (f *fakeStore) GetSettings(context.Context) (Settings, error) { + f.mu.Lock() + defer f.mu.Unlock() return f.settings, nil } func (f *fakeStore) UpdateSettings(_ context.Context, settings Settings) (Settings, error) { + f.mu.Lock() + defer f.mu.Unlock() f.settings = settings return settings, nil } func (f *fakeStore) GetUserLimit(context.Context, int) (*UserLimit, error) { + f.mu.Lock() + defer f.mu.Unlock() return f.limit, nil } func (f *fakeStore) UpsertUserLimit(_ context.Context, limit UserLimit) (*UserLimit, error) { + f.mu.Lock() + defer f.mu.Unlock() f.limit = &limit return &limit, nil } -func (f *fakeStore) CountUserRequestsSince(context.Context, int, time.Time) (int, error) { - return f.count, nil +func (f *fakeStore) CountUserRequestsSince(_ context.Context, userID int, since time.Time) (int, error) { + f.mu.Lock() + defer f.mu.Unlock() + used := f.count + for _, prior := range f.created { + if prior.Requester.UserID != userID { + continue + } + if prior.Now.Before(since) { + continue + } + used++ + } + return used, nil } func (f *fakeStore) ListActiveByTMDB(_ context.Context, mediaType MediaType, ids []int) (map[int]*Request, error) { + f.mu.Lock() + defer f.mu.Unlock() out := map[int]*Request{} for _, id := range ids { if req := f.active[mediaType][id]; req != nil { @@ -707,6 +806,23 @@ func (f *fakeStore) ListActiveByTMDB(_ context.Context, mediaType MediaType, ids } func (f *fakeStore) CreateRequest(_ context.Context, input CreateRequestRecord) (*Request, error) { + f.mu.Lock() + defer f.mu.Unlock() + if input.Quota != nil { + used := f.count + for _, prior := range f.created { + if prior.Requester.UserID != input.Quota.UserID { + continue + } + if prior.Now.Before(input.Quota.WindowStart) { + continue + } + used++ + } + if used >= input.Quota.MaxRequests { + return nil, ErrQuotaExceeded + } + } f.created = append(f.created, input) return &Request{ ID: input.ID, @@ -726,6 +842,8 @@ func (f *fakeStore) CreateRequest(_ context.Context, input CreateRequestRecord) } func (f *fakeStore) GetRequest(_ context.Context, id string) (*Request, error) { + f.mu.Lock() + defer f.mu.Unlock() req := f.requests[strings.TrimSpace(id)] if req == nil { return nil, ErrNotFound @@ -735,6 +853,8 @@ func (f *fakeStore) GetRequest(_ context.Context, id string) (*Request, error) { } func (f *fakeStore) ListReconciliationCandidates(context.Context, int) ([]*Request, error) { + f.mu.Lock() + defer f.mu.Unlock() return f.candidates, nil } @@ -747,6 +867,8 @@ func (f *fakeStore) ListAdmin(context.Context, ListFilter) ([]*Request, error) { } func (f *fakeStore) SetStatus(_ context.Context, id string, status Status, _ Viewer) (*Request, error) { + f.mu.Lock() + defer f.mu.Unlock() f.statusUpdates = append(f.statusUpdates, status) req := f.requests[id] if req == nil { @@ -759,6 +881,8 @@ func (f *fakeStore) SetStatus(_ context.Context, id string, status Status, _ Vie } func (f *fakeStore) MarkQueued(_ context.Context, id string, update QueueUpdate, _ Viewer) (*Request, error) { + f.mu.Lock() + defer f.mu.Unlock() f.queued = append(f.queued, update) req := f.requests[id] if req == nil { @@ -775,6 +899,8 @@ func (f *fakeStore) MarkQueued(_ context.Context, id string, update QueueUpdate, } func (f *fakeStore) SetOutcome(_ context.Context, id string, outcome Outcome, _ Viewer, message string) (*Request, error) { + f.mu.Lock() + defer f.mu.Unlock() req := f.requests[id] if req == nil { req = &Request{ID: id} @@ -787,6 +913,8 @@ func (f *fakeStore) SetOutcome(_ context.Context, id string, outcome Outcome, _ } func (f *fakeStore) ListIntegrations(context.Context) ([]Integration, error) { + f.mu.Lock() + defer f.mu.Unlock() return f.integrations, nil } @@ -795,6 +923,8 @@ func (f *fakeStore) UpsertIntegration(context.Context, Integration) (*Integratio } func (f *fakeStore) UpsertIntegrations(_ context.Context, integrations []Integration) ([]Integration, error) { + f.mu.Lock() + defer f.mu.Unlock() f.integrations = append([]Integration(nil), integrations...) return append([]Integration(nil), integrations...), nil } @@ -980,12 +1110,15 @@ func TestBrowseGenreMovieReturnsResults(t *testing.T) { } type fakePresence struct { + mu sync.Mutex available map[MediaType]map[int]bool byTVDB map[MediaType]map[int]int got []PresenceCandidate } func (f *fakePresence) Lookup(_ context.Context, mediaType MediaType, candidates []PresenceCandidate) (map[int]PresenceMatch, error) { + f.mu.Lock() + defer f.mu.Unlock() out := map[int]PresenceMatch{} f.got = append(f.got, candidates...) for _, candidate := range candidates { diff --git a/internal/requests/store.go b/internal/requests/store.go index 3dc218dc..2241514f 100644 --- a/internal/requests/store.go +++ b/internal/requests/store.go @@ -32,4 +32,15 @@ type CreateRequestRecord struct { Outcome Outcome Requester Viewer Now time.Time + // Quota, when non-nil, instructs the store to atomically verify the + // requester is below their per-user limit before inserting. The check + // runs inside the same transaction as the insert with a per-user + // advisory lock so concurrent submissions cannot both exceed the limit. + Quota *QuotaCheck +} + +type QuotaCheck struct { + UserID int + WindowStart time.Time + MaxRequests int }