diff --git a/docker-compose.yml b/docker-compose.yml index 4b9116d..d381de5 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -5,6 +5,7 @@ services: container_name: augur restart: unless-stopped init: true + stop_grace_period: 45s security_opt: - no-new-privileges:true tmpfs: diff --git a/docs/capabilities.md b/docs/capabilities.md index 35dc3ab..f814a67 100644 --- a/docs/capabilities.md +++ b/docs/capabilities.md @@ -1,14 +1,15 @@ # Capabilities Augur lets Discord members search Seerr, preview a movie or show, choose TV -seasons, submit a request, and review their recent requests without leaving +seasons across multiple pages, submit a request once, and review their recent requests without leaving Discord. Members can opt in or out of direct messages for approvals, declines, and availability. Server administrators can send pending Seerr requests to a Discord channel as approval cards. Those cards also cover requests created outside Discord. Augur polls Seerr for request decisions and availability, remembers delivery state -across restarts, and can expose private health, readiness, and metrics endpoints. +and pending cleanup across restarts. Decision notifications and card repair run +independently of new-request polling. Augur can expose private health, readiness, and metrics endpoints. For exact commands and permissions, see [Commands and approvals](features.md). For deployment settings, see [Configuration](configuration.md). diff --git a/docs/development/architecture.md b/docs/development/architecture.md index fbc6c2d..4369307 100644 --- a/docs/development/architecture.md +++ b/docs/development/architecture.md @@ -2,7 +2,7 @@ `cmd/augur` loads configuration and creates an `app.Runner`. The runner owns the Discord connection, Seerr client, SQLite store, optional health server, and -background monitor. Discord handlers ask the runner to perform application +background monitor and delivery worker. Discord handlers ask the runner to perform application operations; protocol details stay in `internal/discordbot` and `internal/seer`, and durable state stays in `internal/storage`. @@ -30,7 +30,28 @@ so the Manage Server or Administrator permission check is a security boundary. Account-name matching must never replace explicit Seerr linking. The runner serializes lifecycle changes and approval reconciliation separately. -Background polling, Discord callbacks, and shutdown can overlap. Keep operations +Background polling, decision delivery, Discord callbacks, and shutdown can overlap. Keep operations idempotent, make state changes durable before sending notifications where practical, bound retries and waits, and honor context cancellation. SQLite uses one connection and WAL mode; do not add parallel database owners. + +## Approval recovery + +The application serializes Seerr decisions and persists the intended action +before the external write. SQLite owns canonical decision records, recipient +receipts, send leases, backoff and orphan-message cleanup. Discord owns how cards +and notifications are rendered. Approval commands and delivery use the same +claim/finalization path. + +The delivery worker repairs cards and reconciles uncertain decisions even when +the pending-request poll fails. A complete, recent pending-ID snapshot avoids +reloading covered requests; other status lookups are shared across guild cards. +The same worker sends decision DMs and performs retention cleanup. There is no +per-card background goroutine. Shutdown and failed startup close interaction +admission, cancel external calls and join accepted work before storage closes. + +Request confirmation is claimed atomically per search. An interaction lock also +orders updates to that search's Discord response. Other searches can proceed +independently. Explicit Seerr rejection, an already-requested no-op, and an +uncertain submission have different user responses; uncertain POSTs are never +automatically retried. diff --git a/docs/development/codebase-map.md b/docs/development/codebase-map.md index a33fee4..8774abb 100644 --- a/docs/development/codebase-map.md +++ b/docs/development/codebase-map.md @@ -3,11 +3,11 @@ | Path | What lives here | | --- | --- | | `cmd/augur` | Process entry point, signal handling, configuration loading, and logging setup. | -| `internal/app` | Application lifecycle, request/approval coordination, polling, health endpoints, and metrics. | +| `internal/app` | Application lifecycle, request/approval coordination, polling, durable-decision delivery, health endpoints, and metrics. | | `internal/config` | JSON and environment configuration, defaults, normalization, and validation. | -| `internal/discordbot` | Discord session lifecycle, slash commands, buttons, previews, approval cards, formatting, and delivery cache. | +| `internal/discordbot` | Discord session lifecycle, slash commands, buttons, previews, approval controls, card presentation/delivery, formatting, and selection cache. | | `internal/seer` | Seerr HTTP client, account links, requests, status, and user lookup. | -| `internal/storage` | SQLite schema, migrations, subscriptions, approvals, notification preferences, and scans. | +| `internal/storage` | SQLite schema, migrations, subscriptions, delivery leases, decision intents/jobs, card cleanup, notification preferences, and scans. | | `internal/safelog` | Log redaction and safe error output. | | `Dockerfile`, `docker-entrypoint.sh`, `docker-compose.yml` | Container build, startup, ownership, and example deployment. | | `unraid` | Unraid application template. | diff --git a/docs/development/data.md b/docs/development/data.md index ab70fde..5e98248 100644 --- a/docs/development/data.md +++ b/docs/development/data.md @@ -5,7 +5,9 @@ - request subscriptions and availability completion; - approval settings per Discord server; -- approval-message IDs and decisions; +- approval-message IDs, delivery leases and retry deadlines; +- saved decision intents and confirmed decision notification jobs; +- cleanup of sent cards that could not be tracked; - notification preferences per Discord user; and - decision-notification deduplication. @@ -25,3 +27,23 @@ supported version, preservation of existing rows, repeat startup, and failure behavior. Keep SQL and migration compatibility in `internal/storage`. Restore the full data directory for rollback; an older binary may not understand a newer schema. + +## Delivery state in revision 6 + +A claim token and two-minute lease fence each approval send. Retries honor the +current enabled channel, and acknowledgement/deletion of a card requires its +physical channel and message ID. Backoff grows from 30 seconds to 30 minutes. +Known untracked messages have separate durable cleanup records. + +A decision intent preserves the moderator, reason and card presentation before +Seerr is called. An uncertain response is observed without replaying the write. +Confirmed decisions retain their first attribution and queue notification work +before Discord rendering. A single worker resolves current recipient links, +checks preferences, sends DMs and records successful or deliberately suppressed +recipients. Card removal does not remove notification work. + +Discord delivery is at least once: a process crash after Discord accepts a message +but before SQLite records the receipt can cause a duplicate. Existing revision 5 +receipts are retained; migration does not replay historical decisions. Legacy +blank card claims become recoverable. See [Operations](operations.md) before an +upgrade or rollback. diff --git a/docs/development/operations.md b/docs/development/operations.md index 0b9dc95..75b3cc9 100644 --- a/docs/development/operations.md +++ b/docs/development/operations.md @@ -1,6 +1,6 @@ # Operations -Run Augur with the smallest practical permissions: one writable `/data` +Run one Augur process per SQLite data directory, with the smallest practical permissions: one writable `/data` mount, outbound access to Discord and Seerr, and no public inbound port unless the optional health server is deliberately exposed to a private monitoring network. Keep the Discord token, Seerr API key, `.env`, `config.json`, data @@ -32,3 +32,23 @@ The release sequence is: update [release notes](../releases.md), run the final C gate, create an annotated `vX.Y.Z` tag, publish the matching GitHub release, verify the versioned image, and announce any operator action. Repairs that change Discord messages, Seerr requests, or SQLite state must be explicit and opt-in. + +## Revision 6 delivery recovery + +Back up the complete data directory before upgrading. Revision 6 preserves +subscriptions, settings and existing notification receipts, recovers interrupted +blank approval claims, and adds durable decision and card-cleanup work. It does +not send historical decision notifications as part of migration. An older binary +requires the matching pre-upgrade data for rollback. + +The Compose and Unraid examples allow 45 seconds for a graceful stop. Keep that +allowance when customizing deployment so accepted interactions can finish +persisting state. Retry logs distinguish decision checks, recipient delivery, +card updates and cleanup. Resolve expired credentials or Discord permissions; +queued work resumes with bounded backoff. Do not delete the database to clear a +failed delivery. + +A lost Discord acknowledgement can result in a duplicate message because Discord +and SQLite cannot commit atomically. Known extra cards are queued for cleanup. +Keep the data directory private: saved moderator names, decline reasons and +request relationships are part of recovery state. diff --git a/docs/development/validation.md b/docs/development/validation.md index 8a5c24e..74939d9 100644 --- a/docs/development/validation.md +++ b/docs/development/validation.md @@ -31,3 +31,24 @@ The final gate is the complete `CI` workflow in `.github/workflows/ci.yml` on the exact revision. It checks formatting, `go mod tidy -diff`, whitespace, tests, vet, the race detector, pinned Staticcheck and govulncheck versions, the release build, Docker build, entrypoint behavior, and runtime ownership. + +## Recovery and interaction regressions + +The deterministic suites cover upgrades from revisions 1–5, repeat startup and +migration rollback; expired claims, changed/disabled destinations, physical-card +fences and orphan cleanup; accepted decisions with lost responses or failed +local persistence; failed DMs, preference suppression and independent card repair; +modal components decoded from Discord JSON; and startup/shutdown draining. + +Request-flow tests cover finite but non-exhausted quotas, clearing and merging +season pages, duplicate confirmations, stale response ordering, missing account +link data, ambiguous links across pages, and terminal versus uncertain POST +outcomes. Compare bounded account lookup to serial lookup with: + +```sh +go test ./internal/seer -run '^$' -bench BenchmarkLinkedUserNotificationScan -benchtime=5x +``` + +The benchmark models 20 users with one millisecond of notification-endpoint +latency. The concurrency regression requires four workers and joins all of them; +benchmark results describe that fixture, not production Seerr latency. diff --git a/docs/features.md b/docs/features.md index 24022cc..bbb83b6 100644 --- a/docs/features.md +++ b/docs/features.md @@ -13,7 +13,12 @@ | `/approvals disable` | Disables approval cards for the server. | Requests use your linked Seerr account and its permissions and limits by default. -The all-seasons option requires an unlimited TV quota. Augur checks availability +The all-seasons option requires an unlimited TV quota. Shows with more than 25 +seasons have Previous/Next controls; selections stay selected across pages and +share the same remaining quota. Clear a page's selection to free space. + +A confirmation submits once. If Seerr's response cannot be confirmed, check +`/requests` or Seerr before starting another request. Augur checks availability in the background and sends completion DMs for requests it tracks. ## Approval cards @@ -27,7 +32,18 @@ Cards cover pending requests from all Seerr sources, including requests made outside Discord. Choose a channel whose members may see those requests. The bot needs View Channel, Send Messages, Embed Links, and Manage Messages there. -Declines can include an optional reason, stored by Augur. Linked requesters can -receive decision DMs according to their notification preferences. Cards are -removed two minutes after a decision. Discord privacy settings must allow DMs -from the bot for notifications to arrive. +Declines can include an optional reason, saved before the decision is sent to +Seerr. A failed update keeps the card available for retry. If Seerr accepted the +decision but its response was lost, Augur checks its status and preserves the +saved reason. A stale button shows the recorded decision and its original author. + +Linked requesters can receive decision DMs according to their notification +preferences. Failed deliveries retry after restarts, independently of card +updates. Disabled notifications are skipped; enabling them later does not replay +those past decisions. Discord privacy settings must allow DMs from the bot. + +Cards are removed about two minutes after their decided state is displayed. +Failed sends, updates and cleanup retry with backoff. Changing the approval +channel applies to new deliveries; existing cards remain in their original +channel until decided and removed. Disabling approvals prevents further card +sends for that server. diff --git a/internal/app/approval_monitor.go b/internal/app/approval_monitor.go index 52da4a1..57e0c58 100644 --- a/internal/app/approval_monitor.go +++ b/internal/app/approval_monitor.go @@ -18,16 +18,21 @@ func (r *Runner) reconcileApprovals(ctx context.Context) { return } approvals := make([]seer.ApprovalRequest, 0, len(requests)) + pendingIDs := make(map[int]bool, len(requests)) for _, request := range requests { if ctx.Err() != nil { return } + if !seer.IsPendingRequest(request.Status) || request.ID <= 0 { + continue + } + pendingIDs[request.ID] = true approval, ok := r.pendingApproval(ctx, request) if ok { approvals = append(approvals, approval) } } - if err := r.bot.ReconcileApprovals(ctx, approvals); err != nil && ctx.Err() == nil { + if err := r.bot.ReconcileApprovals(ctx, approvals, pendingIDs); err != nil && ctx.Err() == nil { r.logger.Error("reconcile Discord approval messages", "error", err) } } diff --git a/internal/app/approvals.go b/internal/app/approvals.go index 099c196..29f94fa 100644 --- a/internal/app/approvals.go +++ b/internal/app/approvals.go @@ -2,6 +2,8 @@ package app import ( "context" + "errors" + "fmt" "strings" "time" @@ -25,29 +27,21 @@ func (r *Runner) ApprovalDestinations(ctx context.Context) ([]storage.ApprovalSe return r.store.EnabledApprovalSettings(ctx) } -func (r *Runner) ClaimApproval(ctx context.Context, requestID int, guildID, channelID string) (bool, error) { - return r.store.ClaimApprovalMessage(ctx, storage.ApprovalMessage{RequestID: requestID, GuildID: guildID, ChannelID: channelID}) +func (r *Runner) ClaimApproval(ctx context.Context, requestID int, guildID, channelID string) (storage.ApprovalMessage, bool, error) { + return r.store.ClaimApprovalMessage(ctx, storage.ApprovalMessage{RequestID: requestID, GuildID: guildID, ChannelID: channelID}, time.Now().UTC()) } - -func (r *Runner) FinishApproval(ctx context.Context, requestID int, guildID, channelID, messageID string) error { - return r.store.FinishApprovalMessage(ctx, storage.ApprovalMessage{RequestID: requestID, GuildID: guildID, ChannelID: channelID, MessageID: messageID}) -} - -func (r *Runner) ReleaseApproval(ctx context.Context, requestID int, guildID string) error { - return r.store.ReleaseApprovalMessage(ctx, requestID, guildID) +func (r *Runner) FinishApproval(ctx context.Context, message storage.ApprovalMessage) error { + return r.store.FinishApprovalMessage(ctx, message) } - -func (r *Runner) MarkApprovalDecided(ctx context.Context, requestID int, guildID string, decidedAt time.Time) error { - return r.store.MarkApprovalMessageDecided(ctx, requestID, guildID, decidedAt) -} -func (r *Runner) SetApprovalDecision(ctx context.Context, requestID int, guildID, status, reason string) error { - return r.store.SetApprovalDecision(ctx, requestID, guildID, status, reason) -} -func (r *Runner) ClaimDecisionNotification(ctx context.Context, requestID int, discordID, status string) (bool, error) { - return r.store.ClaimDecisionNotification(ctx, requestID, discordID, status) +func (r *Runner) RetryApproval(ctx context.Context, message storage.ApprovalMessage) error { + attempt := message.Attempts + if message.MessageID != "" { + attempt++ + } + return r.store.RetryApprovalMessage(ctx, message, time.Now().Add(deliveryRetryDelay(attempt))) } -func (r *Runner) ReleaseDecisionNotification(ctx context.Context, requestID int, discordID, status string) error { - return r.store.ReleaseDecisionNotification(ctx, requestID, discordID, status) +func (r *Runner) MarkApprovalDecided(ctx context.Context, message storage.ApprovalMessage, decidedAt time.Time) error { + return r.store.MarkApprovalMessageDecided(ctx, message, decidedAt) } func (r *Runner) NotificationPreferences(ctx context.Context, discordID string) (storage.NotificationPreferences, error) { return r.store.NotificationPreferences(ctx, discordID) @@ -76,8 +70,8 @@ func (r *Runner) ApprovalMessages(ctx context.Context) ([]storage.ApprovalMessag return r.store.ApprovalMessages(ctx) } -func (r *Runner) DeleteApprovalRecord(ctx context.Context, requestID int, guildID string) error { - return r.store.DeleteApprovalMessage(ctx, requestID, guildID) +func (r *Runner) DeleteApprovalRecord(ctx context.Context, message storage.ApprovalMessage) error { + return r.store.DeleteApprovalMessage(ctx, message) } func (r *Runner) RequesterDiscordIDs(ctx context.Context, userID int) ([]string, error) { @@ -88,19 +82,130 @@ func (r *Runner) RequesterDiscordIDs(ctx context.Context, userID int) ([]string, return settings.DiscordIDs, nil } -func (r *Runner) DecideRequest(ctx context.Context, requestID int, action string) (seer.Request, error) { +func (r *Runner) DecideRequest(ctx context.Context, requestID int, action string, presentation storage.ApprovalDecision) (storage.ApprovalDecision, bool, error) { r.approvalMu.Lock() defer r.approvalMu.Unlock() + if err := r.store.Ping(ctx); err != nil { + return storage.ApprovalDecision{}, false, err + } current, err := r.seer.Request(ctx, requestID) if err != nil { - return seer.Request{}, err + return storage.ApprovalDecision{}, false, err + } + changed := seer.IsPendingRequest(current.Status) + if changed { + action = strings.ToLower(strings.TrimSpace(action)) + status := "Approved" + if action == "decline" { + status = "Declined" + } else if action != "approve" { + return storage.ApprovalDecision{}, false, errors.New("invalid decision action") + } + intent, exists, err := r.store.DecisionIntent(ctx, requestID) + if err != nil { + return storage.ApprovalDecision{}, false, err + } + if exists && time.Since(intent.CreatedAt) < 2*time.Minute { + return storage.ApprovalDecision{}, false, &userFacingError{message: "The previous decision is still being checked. Try again shortly."} + } + proposal := storage.DecisionIntent{RequestID: requestID, Status: status, Actor: presentation.Actor, Reason: presentation.Reason, Title: presentation.Title, URL: presentation.URL, PosterURL: presentation.PosterURL, CreatedAt: time.Now().UTC()} + if status != "Declined" { + proposal.Reason = "" + } + if err := r.store.SaveDecisionIntent(ctx, proposal); err != nil { + return storage.ApprovalDecision{}, false, err + } + updated, err := r.seer.UpdateRequestStatus(ctx, requestID, strings.ToLower(strings.TrimSpace(action))) + if err != nil { + if seer.IsDefiniteRejection(err) { + persistCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) + clearErr := r.store.ClearDecisionIntent(persistCtx, requestID) + cancel() + if clearErr != nil { + r.logger.Error("clear rejected decision intent", "request_id", requestID, "error", clearErr) + } + } + return storage.ApprovalDecision{}, false, err + } + if updated.RequestedBy == nil { + updated.RequestedBy = current.RequestedBy + } + if updated.Media == nil { + updated.Media = current.Media + } + if updated.Type == "" { + updated.Type = current.Type + } + if seer.IsPendingRequest(updated.Status) { + return storage.ApprovalDecision{}, false, errors.New("seerr did not confirm the decision") + } + current = updated + if seer.RequestStatusLabel(current.Status) != status { + changed = false + presentation.Actor = "Seerr" + presentation.Reason = "" + } + } else { + // A second administrator must not replace the first decision's attribution. + presentation.Actor = "Seerr" + presentation.Reason = "" + } + persistCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) + defer cancel() + decision, err := r.recordDecision(persistCtx, current, presentation) + return decision, changed, err +} + +func (r *Runner) ObserveApprovalDecision(ctx context.Context, request seer.Request, presentation storage.ApprovalDecision) (storage.ApprovalDecision, error) { + r.approvalMu.Lock() + defer r.approvalMu.Unlock() + presentation.Actor = "Seerr" + presentation.Reason = "" + return r.recordDecision(ctx, request, presentation) +} + +func (r *Runner) recordDecision(ctx context.Context, request seer.Request, d storage.ApprovalDecision) (storage.ApprovalDecision, error) { + d.RequestID = request.ID + d.Status = seer.RequestStatusLabel(request.Status) + d.DecidedAt = time.Now().UTC() + if request.RequestedBy != nil { + d.RequesterID = request.RequestedBy.ID + } + if request.Media != nil { + d.MediaID = request.Media.TMDBID + d.MediaType = request.Media.MediaType } - if !seer.IsPendingRequest(current.Status) { - return current, nil + if request.Type != "" { + d.MediaType = request.Type } - updated, err := r.seer.UpdateRequestStatus(ctx, requestID, strings.ToLower(strings.TrimSpace(action))) - if err == nil && updated.RequestedBy == nil { - updated.RequestedBy = current.RequestedBy + if d.Status != "Declined" { + d.Reason = "" } - return updated, err + if d.Title == "" { + d.Title = fmt.Sprintf("Request #%d", request.ID) + } + switch d.Status { + case "Unknown", "Pending approval": + return storage.ApprovalDecision{}, errors.New("seerr did not confirm a final request status") + case "Failed", "Completed": + return d, r.store.ClearDecisionIntent(ctx, d.RequestID) + } + return r.store.RecordApprovalDecision(ctx, d) +} + +func (r *Runner) ApprovalDecision(ctx context.Context, requestID int, status string) (storage.ApprovalDecision, bool, error) { + return r.store.ApprovalDecision(ctx, requestID, status) +} + +func (r *Runner) QueueUntrackedApprovalCleanup(ctx context.Context, message storage.ApprovalMessage) (bool, error) { + return r.store.QueueUntrackedApprovalCleanup(ctx, message) +} +func (r *Runner) DueApprovalCleanup(ctx context.Context) ([]storage.ApprovalMessage, error) { + return r.store.DueApprovalCleanup(ctx, time.Now()) +} +func (r *Runner) CompleteApprovalCleanup(ctx context.Context, message storage.ApprovalMessage) error { + return r.store.CompleteApprovalCleanup(ctx, message) +} +func (r *Runner) RetryApprovalCleanup(ctx context.Context, message storage.ApprovalMessage) error { + return r.store.RetryApprovalCleanup(ctx, message, time.Now().Add(deliveryRetryDelay(message.Attempts+1))) } diff --git a/internal/app/decision_worker.go b/internal/app/decision_worker.go new file mode 100644 index 0000000..1581f70 --- /dev/null +++ b/internal/app/decision_worker.go @@ -0,0 +1,155 @@ +package app + +import ( + "context" + "errors" + "strings" + "time" + + "github.com/mayvqt/Augur/internal/config" + "github.com/mayvqt/Augur/internal/seer" + "github.com/mayvqt/Augur/internal/storage" +) + +func deliveryRetryDelay(attempt int) time.Duration { + return min(30*time.Second*time.Duration(1< quota.TV.Remaining { + if quota != nil && quota.TV.Limited() && len(unique) > quota.TV.Remaining { return seer.SeasonSelection{}, &userFacingError{ message: fmt.Sprintf("You can request %d more TV season(s) in the current quota window.", quota.TV.Remaining), } diff --git a/internal/app/runner.go b/internal/app/runner.go index 76711c4..9d2a240 100644 --- a/internal/app/runner.go +++ b/internal/app/runner.go @@ -129,8 +129,9 @@ func (r *Runner) Run(ctx context.Context) error { r.health.SetReady(true) } - r.wg.Add(1) + r.wg.Add(2) go r.runMonitor(runCtx) + go r.runDecisionWorker(runCtx) <-runCtx.Done() r.logger.Info("shutdown requested") @@ -150,7 +151,7 @@ func (r *Runner) Close() error { if running { select { case <-runDone: - case <-time.After(10 * time.Second): + case <-time.After(30 * time.Second): return errors.New("runner shutdown timed out") } } diff --git a/internal/app/runner_test.go b/internal/app/runner_test.go index 20a6053..4025499 100644 --- a/internal/app/runner_test.go +++ b/internal/app/runner_test.go @@ -149,8 +149,8 @@ func TestRunnerRejectsRequestWithoutSeerrID(t *testing.T) { if _, err := runner.Request(context.Background(), "123456789012345678", seer.SearchResult{ ID: 9, MediaType: "movie", Title: "Arrival", - }, seer.SeasonSelection{}); err == nil || !strings.Contains(err.Error(), "valid ID") { - t.Fatalf("Request() error = %v, want missing Seerr request ID", err) + }, seer.SeasonSelection{}); !errors.Is(err, seer.ErrSubmissionUnknown) { + t.Fatalf("Request() error = %v, want uncertain submission outcome", err) } } @@ -181,7 +181,7 @@ func TestRunnerTVRequestEnforcesRemainingSeasonQuota(t *testing.T) { seerClient := &fakeSeer{ user: seer.User{ID: 7}, found: true, - quota: seer.Quota{TV: seer.QuotaUsage{Restricted: true, Remaining: 3}}, + quota: seer.Quota{TV: seer.QuotaUsage{Limit: 5, Used: 2, Restricted: false, Remaining: 3}}, tvDetails: seer.TVDetails{Seasons: []seer.Season{ {SeasonNumber: 1}, {SeasonNumber: 2}, {SeasonNumber: 3}, {SeasonNumber: 4}, }}, @@ -213,7 +213,7 @@ func TestRunnerTVRequestSubmitsSelectedSeasons(t *testing.T) { seerClient := &fakeSeer{ user: seer.User{ID: 7}, found: true, - quota: seer.Quota{TV: seer.QuotaUsage{Restricted: true, Remaining: 3}}, + quota: seer.Quota{TV: seer.QuotaUsage{Limit: 5, Used: 2, Restricted: false, Remaining: 3}}, tvDetails: seer.TVDetails{Seasons: []seer.Season{ {SeasonNumber: 1}, {SeasonNumber: 2}, {SeasonNumber: 3}, }}, diff --git a/internal/app/test_helpers_test.go b/internal/app/test_helpers_test.go index 83de9ff..b8feef9 100644 --- a/internal/app/test_helpers_test.go +++ b/internal/app/test_helpers_test.go @@ -152,6 +152,7 @@ func (f *fakeSeer) UpdateRequestStatus(ctx context.Context, id int, action strin } type fakeStore struct { + stateStore pending []storage.Subscription pendingByID map[int]storage.Subscription added []storage.Subscription @@ -231,19 +232,19 @@ func (f *fakeStore) NeedsApprovalMessage(ctx context.Context, requestID int) (bo return f.approval.Enabled, ctx.Err() } -func (f *fakeStore) ClaimApprovalMessage(ctx context.Context, message storage.ApprovalMessage) (bool, error) { - return true, ctx.Err() +func (f *fakeStore) ClaimApprovalMessage(ctx context.Context, message storage.ApprovalMessage, now time.Time) (storage.ApprovalMessage, bool, error) { + return message, true, ctx.Err() } func (f *fakeStore) FinishApprovalMessage(ctx context.Context, message storage.ApprovalMessage) error { return ctx.Err() } -func (f *fakeStore) ReleaseApprovalMessage(ctx context.Context, requestID int, guildID string) error { +func (f *fakeStore) RetryApprovalMessage(ctx context.Context, message storage.ApprovalMessage, retryAt time.Time) error { return ctx.Err() } -func (f *fakeStore) MarkApprovalMessageDecided(ctx context.Context, requestID int, guildID string, decidedAt time.Time) error { +func (f *fakeStore) MarkApprovalMessageDecided(ctx context.Context, message storage.ApprovalMessage, decidedAt time.Time) error { return ctx.Err() } @@ -251,20 +252,14 @@ func (f *fakeStore) DueApprovalMessages(ctx context.Context, before time.Time) ( return nil, ctx.Err() } -func (f *fakeStore) DeleteApprovalMessage(ctx context.Context, requestID int, guildID string) error { +func (f *fakeStore) DeleteApprovalMessage(ctx context.Context, message storage.ApprovalMessage) error { return ctx.Err() } func (f *fakeStore) ApprovalMessages(ctx context.Context) ([]storage.ApprovalMessage, error) { return nil, ctx.Err() } -func (f *fakeStore) SetApprovalDecision(ctx context.Context, requestID int, guildID, status, reason string) error { - return ctx.Err() -} -func (f *fakeStore) ClaimDecisionNotification(ctx context.Context, requestID int, discordID, status string) (bool, error) { - return true, ctx.Err() -} -func (f *fakeStore) ReleaseDecisionNotification(ctx context.Context, requestID int, discordID, status string) error { - return ctx.Err() +func (f *fakeStore) DueDecisionJobs(ctx context.Context, now time.Time, limit int) ([]storage.DecisionJob, error) { + return nil, ctx.Err() } func (f *fakeStore) NotificationPreferences(ctx context.Context, discordID string) (storage.NotificationPreferences, error) { if f.notificationPrefs != nil { @@ -293,7 +288,7 @@ type fakeNotifier struct { notifyErr error } -func (f *fakeNotifier) ReconcileApprovals(ctx context.Context, approvals []seer.ApprovalRequest) error { +func (f *fakeNotifier) ReconcileApprovals(ctx context.Context, approvals []seer.ApprovalRequest, pendingIDs map[int]bool) error { if err := ctx.Err(); err != nil { return err } @@ -322,3 +317,12 @@ func (f *fakeNotifier) NotifyComplete(ctx context.Context, discordID string, med func (f *fakeNotifier) Close() error { return nil } + +func (f *fakeNotifier) NotifyDecision(ctx context.Context, id string, d storage.ApprovalDecision) error { + return errors.Join(ctx.Err(), f.notifyErr) +} +func (f *fakeNotifier) MaintainApprovals(ctx context.Context) error { return ctx.Err() } + +func (f *fakeStore) DueDecisionIntents(ctx context.Context, now time.Time) ([]storage.DecisionIntent, error) { + return nil, ctx.Err() +} diff --git a/internal/discordbot/approval_delivery.go b/internal/discordbot/approval_delivery.go new file mode 100644 index 0000000..1856a1b --- /dev/null +++ b/internal/discordbot/approval_delivery.go @@ -0,0 +1,265 @@ +package discordbot + +import ( + "context" + "errors" + "time" + + "github.com/bwmarrin/discordgo" + "github.com/mayvqt/Augur/internal/seer" + "github.com/mayvqt/Augur/internal/storage" +) + +const approvalMessageRetention = 2 * time.Minute + +func (b *Bot) postApproval(ctx context.Context, guildID string, requestID int, requesterID string, result seer.SearchResult, seasons seer.SeasonSelection) { + if guildID == "" || requestID <= 0 { + return + } + channelID, enabled, err := b.handler.ApprovalChannel(ctx, guildID) + if err != nil { + b.logger.Error("load approval channel", "error", err) + return + } + if !enabled { + return + } + b.sendApproval(ctx, storage.ApprovalSettings{GuildID: guildID, ChannelID: channelID, Enabled: true}, seer.ApprovalRequest{RequestID: requestID, RequesterID: requesterID, Media: result, Seasons: seasons}) +} + +func (b *Bot) sendApproval(ctx context.Context, destination storage.ApprovalSettings, approval seer.ApprovalRequest) { + claim, claimed, err := b.handler.ClaimApproval(ctx, approval.RequestID, destination.GuildID, destination.ChannelID) + if err != nil { + b.logger.Error("claim approval message", "error", err) + return + } + if !claimed { + return + } + message, sendErr := b.sendApprovalMessage(ctx, destination.ChannelID, &discordgo.MessageSend{Embeds: []*discordgo.MessageEmbed{b.approvalEmbed(approval)}, Components: approvalComponents(approval.RequestID), AllowedMentions: noMentions()}) + persistCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) + defer cancel() + if sendErr == nil && message != nil && message.ID != "" { + claim.MessageID = message.ID + if err := b.handler.FinishApproval(persistCtx, claim); err == nil { + return + } else { + b.logger.Error("save approval message", "request_id", approval.RequestID, "error", err) + } + // A separate deadline allows cleanup even when saving the card timed out. + cleanupCtx, cleanupCancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) + queued, queueErr := b.handler.QueueUntrackedApprovalCleanup(cleanupCtx, claim) + if queueErr != nil { + b.logger.Error("retain untracked approval cleanup", "request_id", claim.RequestID, "error", queueErr) + } + if queued { + b.deleteUntrackedApproval(cleanupCtx, claim) + } + cleanupCancel() + } else { + b.logger.Warn("send approval message", "request_id", claim.RequestID, "error", sendErr) + } + retryCtx, retryCancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) + defer retryCancel() + if err := b.handler.RetryApproval(retryCtx, claim); err != nil { + b.logger.Error("schedule approval retry", "request_id", claim.RequestID, "error", err) + } +} + +func (b *Bot) ReconcileApprovals(ctx context.Context, approvals []seer.ApprovalRequest, pendingIDs map[int]bool) error { + b.pendingMu.Lock() + b.pendingIDs = make(map[int]bool, len(pendingIDs)) + for id, pending := range pendingIDs { + b.pendingIDs[id] = pending + } + b.pendingAt = time.Now() + b.pendingMu.Unlock() + destinations, err := b.handler.ApprovalDestinations(ctx) + if err != nil { + return err + } + for _, approval := range approvals { + for _, destination := range destinations { + if ctx.Err() != nil { + return ctx.Err() + } + b.sendApproval(ctx, destination, approval) + } + } + return nil +} + +func (b *Bot) recentlyPending(requestID int) bool { + b.pendingMu.Lock() + defer b.pendingMu.Unlock() + return time.Since(b.pendingAt) < time.Minute && b.pendingIDs[requestID] +} + +func (b *Bot) MaintainApprovals(ctx context.Context) error { + if err := b.cleanupApprovals(ctx); err != nil { + b.logger.Warn("clean up decided cards", "error", err) + } + records, err := b.handler.ApprovalMessages(ctx) + if err != nil { + return err + } + type observed struct { + decision storage.ApprovalDecision + err error + pending bool + } + outcomes := make(map[int]observed) + for _, record := range records { + if ctx.Err() != nil { + return ctx.Err() + } + if !record.DecidedAt.IsZero() { + continue + } + outcome, seen := outcomes[record.RequestID] + if !seen { + decision, found, lookupErr := b.handler.ApprovalDecision(ctx, record.RequestID, record.Status) + if lookupErr != nil { + outcome.err = lookupErr + } else if found { + outcome.decision = decision + } else if b.recentlyPending(record.RequestID) { + outcome.pending = true + } else { + request, err := b.handler.RequestStatus(ctx, record.RequestID) + outcome.err = err + if err == nil { + outcome.pending = seer.IsPendingRequest(request.Status) + if !outcome.pending { + presentation := storage.ApprovalDecision{} + if request.Media != nil && request.Media.TMDBID > 0 { + mediaType := request.Type + if mediaType == "" { + mediaType = request.Media.MediaType + } + if media, err := b.handler.MediaDetails(ctx, mediaType, request.Media.TMDBID); err == nil { + presentation = approvalPresentation(&discordgo.Message{Embeds: []*discordgo.MessageEmbed{b.mediaPreview(media, nil)}}) + } + } + outcome.decision, outcome.err = b.handler.ObserveApprovalDecision(ctx, request, presentation) + } + } + } + outcomes[record.RequestID] = outcome + } + if outcome.err != nil || outcome.pending { + if outcome.err != nil { + b.retryApprovalCard(ctx, record) + } + continue + } + if record.MessageID == "" { + if err := b.saveApprovalOutcome(ctx, func(saveCtx context.Context) error { return b.handler.DeleteApprovalRecord(saveCtx, record) }); err != nil { + b.logger.Warn("remove finished delivery claim", "error", err) + } + continue + } + // The decision job already exists even if this card was deleted or damaged. + message, err := b.fetchApprovalMessage(ctx, record.ChannelID, record.MessageID) + if err != nil { + if discordNotFound(err) { + _ = b.saveApprovalOutcome(ctx, func(saveCtx context.Context) error { return b.handler.DeleteApprovalRecord(saveCtx, record) }) + } else { + b.retryApprovalCard(ctx, record) + } + continue + } + var source *discordgo.MessageEmbed + if message != nil && len(message.Embeds) > 0 { + source = message.Embeds[0] + } + embed := approvalDecisionEmbed(source, outcome.decision) + empty := "" + _, err = b.editApprovalMessage(ctx, &discordgo.MessageEdit{ID: record.MessageID, Channel: record.ChannelID, Content: &empty, Embeds: &[]*discordgo.MessageEmbed{embed}, Components: &[]discordgo.MessageComponent{}, AllowedMentions: noMentions()}) + if err != nil { + if discordNotFound(err) { + _ = b.saveApprovalOutcome(ctx, func(saveCtx context.Context) error { return b.handler.DeleteApprovalRecord(saveCtx, record) }) + } else { + b.retryApprovalCard(ctx, record) + } + continue + } + record.Status = outcome.decision.Status + if err := b.saveApprovalOutcome(ctx, func(saveCtx context.Context) error { + return b.handler.MarkApprovalDecided(saveCtx, record, time.Now().UTC()) + }); err != nil { + b.logger.Error("record approval card update", "request_id", record.RequestID, "error", err) + } + } + return nil +} + +func (b *Bot) cleanupApprovals(ctx context.Context) error { + orphans, err := b.handler.DueApprovalCleanup(ctx) + if err != nil { + return err + } + for _, message := range orphans { + if ctx.Err() != nil { + return ctx.Err() + } + b.deleteUntrackedApproval(ctx, message) + } + + messages, err := b.handler.DueApprovalMessages(ctx, time.Now().UTC().Add(-approvalMessageRetention)) + if err != nil { + return err + } + for _, message := range messages { + if ctx.Err() != nil { + return ctx.Err() + } + if err := b.cleanupApprovalMessage(ctx, message.ChannelID, message.MessageID); err != nil && !discordNotFound(err) { + b.logger.Warn("delete decided approval message", "request_id", message.RequestID, "error", err) + b.retryApprovalCard(ctx, message) + continue + } + if err := b.saveApprovalOutcome(ctx, func(saveCtx context.Context) error { return b.handler.DeleteApprovalRecord(saveCtx, message) }); err != nil { + b.logger.Error("delete approval record", "request_id", message.RequestID, "error", err) + } + } + return nil +} + +func (b *Bot) NotifyDecision(ctx context.Context, discordID string, d storage.ApprovalDecision) error { + channel, err := b.session.UserChannelCreate(discordID, discordgo.WithContext(ctx)) + if err != nil { + return err + } + _, err = b.session.ChannelMessageSendComplex(channel.ID, &discordgo.MessageSend{Embeds: []*discordgo.MessageEmbed{decisionEmbed(decisionSource(d), d.Status, d.Reason)}, AllowedMentions: noMentions()}, discordgo.WithContext(ctx)) + return err +} +func discordNotFound(err error) bool { + var rest *discordgo.RESTError + return errors.As(err, &rest) && rest.Response != nil && rest.Response.StatusCode == 404 +} + +func (b *Bot) retryApprovalCard(ctx context.Context, message storage.ApprovalMessage) { + if err := b.saveApprovalOutcome(ctx, func(saveCtx context.Context) error { return b.handler.RetryApproval(saveCtx, message) }); err != nil { + b.logger.Warn("schedule approval card retry", "request_id", message.RequestID, "error", err) + } +} + +func (b *Bot) deleteUntrackedApproval(ctx context.Context, message storage.ApprovalMessage) { + if err := b.cleanupApprovalMessage(ctx, message.ChannelID, message.MessageID); err != nil && !discordNotFound(err) { + b.logger.Warn("delete untracked approval card", "error", err) + if retryErr := b.saveApprovalOutcome(ctx, func(saveCtx context.Context) error { return b.handler.RetryApprovalCleanup(saveCtx, message) }); retryErr != nil { + b.logger.Warn("retain cleanup retry", "error", retryErr) + } + return + } + if err := b.saveApprovalOutcome(ctx, func(saveCtx context.Context) error { return b.handler.CompleteApprovalCleanup(saveCtx, message) }); err != nil { + b.logger.Warn("complete untracked card cleanup", "error", err) + } +} + +func (b *Bot) saveApprovalOutcome(ctx context.Context, save func(context.Context) error) error { + persistCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) + defer cancel() + return save(persistCtx) +} diff --git a/internal/discordbot/approval_flow_test.go b/internal/discordbot/approval_flow_test.go new file mode 100644 index 0000000..871240e --- /dev/null +++ b/internal/discordbot/approval_flow_test.go @@ -0,0 +1,320 @@ +package discordbot + +import ( + "context" + "encoding/json" + "errors" + "io" + "log/slog" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/bwmarrin/discordgo" + "github.com/mayvqt/Augur/internal/config" + "github.com/mayvqt/Augur/internal/seer" + "github.com/mayvqt/Augur/internal/storage" +) + +type decisionHandler struct { + Handler + err error + presentation storage.ApprovalDecision + decision storage.ApprovalDecision + changed bool + marked storage.ApprovalMessage +} + +func (h *decisionHandler) DecideRequest(_ context.Context, _ int, _ string, d storage.ApprovalDecision) (storage.ApprovalDecision, bool, error) { + h.presentation = d + return h.decision, h.changed, h.err +} +func (h *decisionHandler) MarkApprovalDecided(_ context.Context, m storage.ApprovalMessage, _ time.Time) error { + h.marked = m + return nil +} +func approvalInteraction() *discordgo.InteractionCreate { + return &discordgo.InteractionCreate{Interaction: &discordgo.Interaction{Type: discordgo.InteractionMessageComponent, GuildID: "123", ChannelID: "456", Member: &discordgo.Member{Permissions: discordgo.PermissionManageGuild, User: &discordgo.User{ID: "123456789012345678", Username: "Second moderator"}}, Message: &discordgo.Message{ID: "789", ChannelID: "456", Embeds: []*discordgo.MessageEmbed{{Title: "Arrival", Fields: []*discordgo.MessageEmbedField{{Name: "Status", Value: "Pending approval"}}}}}}} +} +func quietLogger() *slog.Logger { return slog.New(slog.NewTextHandler(io.Discard, nil)) } + +func TestApprovalErrorPreservesCardAndEmptyLegacyCardCanRecover(t *testing.T) { + h := &decisionHandler{err: errors.New("temporary error"), decision: storage.ApprovalDecision{RequestID: 42, Status: "Approved", Title: "Arrival", Actor: "First moderator"}} + b := &Bot{handler: h, ctx: context.Background(), logger: quietLogger()} + s := &fakeInteractionSession{} + i := approvalInteraction() + b.handleApproval(s, i, discordgo.MessageComponentInteractionData{CustomID: componentApprove + "42"}, "approve") + if s.edit == nil || s.edit.Embeds != nil || len(*s.edit.Components) == 0 { + t.Fatal("failed approval erased the source card") + } + h.err = nil + i.Message.Embeds = nil + b.handleApproval(s, i, discordgo.MessageComponentInteractionData{CustomID: componentApprove + "42"}, "approve") + if s.edit == nil || s.edit.Embeds == nil || len(*s.edit.Embeds) != 1 || (*s.edit.Embeds)[0].Title != "Arrival" { + t.Fatal("legacy empty card did not recover") + } + if !strings.Contains((*s.edit.Embeds)[0].Footer.Text, "First moderator") || !strings.Contains(*s.edit.Content, "already approved") { + t.Fatal("stale click changed decision attribution") + } + if h.marked.MessageID != "789" || h.marked.ChannelID != "456" { + t.Fatal("card acknowledgement was not fenced", h.marked) + } +} + +func TestDeclineModalUsesDecodedComponentsAndAcknowledgesBeforeLookup(t *testing.T) { + for _, fail := range []bool{false, true} { + t.Run(map[bool]string{false: "success", true: "failure"}[fail], func(t *testing.T) { + var data discordgo.ModalSubmitInteractionData + if err := json.Unmarshal([]byte(`{"custom_id":"augur:decline-reason:42:456:789","components":[{"type":1,"components":[{"type":4,"custom_id":"reason","value":" Not this week "}]}]}`), &data); err != nil { + t.Fatal(err) + } + h := &decisionHandler{changed: true, decision: storage.ApprovalDecision{RequestID: 42, Status: "Declined", Title: "Arrival", Actor: "Moderator", Reason: "Not this week"}} + if fail { + h.err = errors.New("offline") + } + s := &fakeInteractionSession{} + i := approvalInteraction() + i.Type = discordgo.InteractionModalSubmit + i.Data = data + b := &Bot{handler: h, ctx: context.Background(), logger: quietLogger()} + b.fetchApprovalMessage = func(context.Context, string, string) (*discordgo.Message, error) { + if s.response == nil || s.response.Type != discordgo.InteractionResponseDeferredChannelMessageWithSource { + t.Fatal("modal was not acknowledged before I/O") + } + return &discordgo.Message{ID: "789", ChannelID: "456"}, nil + } + edited := false + b.editApprovalMessage = func(_ context.Context, edit *discordgo.MessageEdit) (*discordgo.Message, error) { + edited = true + return &discordgo.Message{}, nil + } + b.handleDeclineModal(s, i) + if h.presentation.Reason != "Not this week" { + t.Fatal("decoded reason was lost", h.presentation) + } + if s.edit == nil || s.edit.Components == nil || len(*s.edit.Components) != 0 { + t.Fatal("orphan approval controls were attached to modal response") + } + if !fail && (!edited || *s.edit.Content != "Request declined.") { + t.Fatal("missing modal completion", s.edit) + } + }) + } +} + +func TestDecidedEmbedDoesNotMutatePendingMessage(t *testing.T) { + source := &discordgo.MessageEmbed{Title: "Arrival", Fields: []*discordgo.MessageEmbedField{{Name: "Status", Value: "Pending approval"}}} + _ = decidedApprovalEmbed(source, 42, "Approved", "Moderator") + if source.Fields[0].Value != "Pending approval" { + t.Fatal("source embed was mutated") + } +} + +type maintenanceHandler struct { + Handler + store *storage.Store + request seer.Request + requestCalls int +} + +func (h *maintenanceHandler) DueApprovalCleanup(ctx context.Context) ([]storage.ApprovalMessage, error) { + return h.store.DueApprovalCleanup(ctx, time.Now()) +} +func (h *maintenanceHandler) DueApprovalMessages(ctx context.Context, before time.Time) ([]storage.ApprovalMessage, error) { + return h.store.DueApprovalMessages(ctx, before) +} +func (h *maintenanceHandler) ApprovalMessages(ctx context.Context) ([]storage.ApprovalMessage, error) { + return h.store.ApprovalMessages(ctx) +} +func (h *maintenanceHandler) ApprovalDecision(ctx context.Context, id int, status string) (storage.ApprovalDecision, bool, error) { + return h.store.ApprovalDecision(ctx, id, status) +} +func (h *maintenanceHandler) RequestStatus(context.Context, int) (seer.Request, error) { + h.requestCalls++ + return h.request, nil +} +func (h *maintenanceHandler) MediaDetails(context.Context, string, int) (seer.SearchResult, error) { + return seer.SearchResult{ID: 99, Title: "Arrival", MediaType: "movie"}, nil +} +func (h *maintenanceHandler) ObserveApprovalDecision(ctx context.Context, request seer.Request, d storage.ApprovalDecision) (storage.ApprovalDecision, error) { + d.RequestID = request.ID + d.Status = seer.RequestStatusLabel(request.Status) + d.RequesterID = 7 + d.Actor = "Seerr" + return h.store.RecordApprovalDecision(ctx, d) +} +func (h *maintenanceHandler) DeleteApprovalRecord(ctx context.Context, m storage.ApprovalMessage) error { + return h.store.DeleteApprovalMessage(ctx, m) +} +func (h *maintenanceHandler) MarkApprovalDecided(ctx context.Context, m storage.ApprovalMessage, at time.Time) error { + return h.store.MarkApprovalMessageDecided(ctx, m, at) +} +func (h *maintenanceHandler) RetryApproval(ctx context.Context, m storage.ApprovalMessage) error { + return h.store.RetryApprovalMessage(ctx, m, time.Now().Add(time.Minute)) +} + +func TestBlankApprovalDeliveriesObserveExternalDecisionOnceAcrossGuilds(t *testing.T) { + ctx := context.Background() + store, err := storage.Open(filepath.Join(t.TempDir(), "state.db")) + if err != nil { + t.Fatal(err) + } + defer store.Close() + for _, guild := range []string{"one", "two"} { + store.SetApprovalSettings(ctx, storage.ApprovalSettings{GuildID: guild, ChannelID: "channel", Enabled: true}) + if _, ok, err := store.ClaimApprovalMessage(ctx, storage.ApprovalMessage{RequestID: 42, GuildID: guild, ChannelID: "channel"}, time.Now().Add(-time.Hour)); err != nil || !ok { + t.Fatal(err) + } + } + h := &maintenanceHandler{store: store, request: seer.Request{ID: 42, Status: 2, RequestedBy: &seer.User{ID: 7}, Media: &seer.Media{TMDBID: 99, MediaType: "movie"}}} + b := &Bot{handler: h, ctx: ctx, logger: quietLogger()} + b.fetchApprovalMessage = func(context.Context, string, string) (*discordgo.Message, error) { + t.Fatal("looked up an empty message ID") + return nil, nil + } + if err := b.MaintainApprovals(ctx); err != nil { + t.Fatal(err) + } + if h.requestCalls != 1 { + t.Fatal("request fetched once per guild", h.requestCalls) + } + if jobs, err := store.DueDecisionJobs(ctx, time.Now(), 25); err != nil || len(jobs) != 1 || jobs[0].Title != "Arrival" { + t.Fatal("decision lost with failed card send", jobs, err) + } + if messages, err := store.ApprovalMessages(ctx); err != nil || len(messages) != 0 { + t.Fatal("finished blank claims retained", messages, err) + } +} + +func TestCanonicalDecisionRepairsCardWithoutPendingPoll(t *testing.T) { + ctx := context.Background() + store, err := storage.Open(filepath.Join(t.TempDir(), "state.db")) + if err != nil { + t.Fatal(err) + } + defer store.Close() + store.SetApprovalSettings(ctx, storage.ApprovalSettings{GuildID: "g", ChannelID: "c", Enabled: true}) + claim, _, err := store.ClaimApprovalMessage(ctx, storage.ApprovalMessage{RequestID: 42, GuildID: "g", ChannelID: "c"}, time.Now()) + if err != nil { + t.Fatal(err) + } + claim.MessageID = "m" + if err := store.FinishApprovalMessage(ctx, claim); err != nil { + t.Fatal(err) + } + store.RecordApprovalDecision(ctx, storage.ApprovalDecision{RequestID: 42, RequesterID: 7, Status: "Declined", Title: "Arrival", Reason: "Not this week", Actor: "First moderator"}) + h := &maintenanceHandler{store: store} + b := &Bot{handler: h, ctx: ctx, logger: quietLogger(), pendingIDs: map[int]bool{42: true}, pendingAt: time.Now()} + b.fetchApprovalMessage = func(context.Context, string, string) (*discordgo.Message, error) { + return &discordgo.Message{ID: "m", ChannelID: "c"}, nil + } + attempts := 0 + b.editApprovalMessage = func(_ context.Context, edit *discordgo.MessageEdit) (*discordgo.Message, error) { + attempts++ + if attempts == 1 { + return nil, errors.New("Discord offline") + } + embed := (*edit.Embeds)[0] + if embed.Title != "Arrival" || !strings.Contains(embed.Footer.Text, "First moderator") { + t.Fatal("lost canonical card content") + } + return &discordgo.Message{}, nil + } + b.MaintainApprovals(ctx) + store.RetryApprovalMessage(ctx, claim, time.Now().Add(-time.Second)) + if err := b.MaintainApprovals(ctx); err != nil { + t.Fatal(err) + } + if attempts != 2 || h.requestCalls != 0 { + t.Fatal("card repair depended on Seerr polling", attempts, h.requestCalls) + } + if due, err := store.DueApprovalMessages(ctx, time.Now().Add(time.Hour)); err != nil || len(due) != 1 { + t.Fatal("retention started before successful render", due, err) + } +} + +func TestStartupFailureCancelsAndDrainsAcceptedInteraction(t *testing.T) { + b, err := New(config.DiscordConfig{Token: "synthetic"}, config.LinkConfig{}, &decisionHandler{}, quietLogger()) + if err != nil { + t.Fatal(err) + } + cancelled := make(chan struct{}) + release := make(chan struct{}) + b.openSession = func() error { + b.lifecycleMu.Lock() + b.interactions.Add(1) + b.lifecycleMu.Unlock() + go func() { defer b.interactions.Done(); <-b.ctx.Done(); close(cancelled); <-release }() + return nil + } + b.closeSession = func() error { return nil } + b.registerApplicationCommands = func() error { return errors.New("synthetic command registration failure") } + done := make(chan error, 1) + go func() { done <- b.Start(context.Background()) }() + select { + case <-cancelled: + case <-time.After(time.Second): + t.Fatal("startup failure did not cancel interaction") + } + select { + case <-done: + t.Fatal("startup returned before interaction persistence finished") + default: + } + close(release) + select { + case err := <-done: + if err == nil || !strings.Contains(err.Error(), "registration failure") { + t.Fatal(err) + } + case <-time.After(time.Second): + t.Fatal("startup failed to drain") + } + b.onInteraction(nil, approvalInteraction()) // closed admission must not invoke handlers +} + +func (h *maintenanceHandler) CompleteApprovalCleanup(ctx context.Context, m storage.ApprovalMessage) error { + return h.store.CompleteApprovalCleanup(ctx, m) +} +func (h *maintenanceHandler) RetryApprovalCleanup(ctx context.Context, m storage.ApprovalMessage) error { + return h.store.RetryApprovalCleanup(ctx, m, time.Now().Add(time.Minute)) +} + +func TestMaintenanceTimeoutDefersSlowCleanupAndAllowsNextCard(t *testing.T) { + ctx := context.Background() + store, err := storage.Open(filepath.Join(t.TempDir(), "state.db")) + if err != nil { + t.Fatal(err) + } + defer store.Close() + for _, id := range []string{"1-slow", "2-healthy"} { + if _, err := store.QueueUntrackedApprovalCleanup(ctx, storage.ApprovalMessage{ChannelID: "c", MessageID: id}); err != nil { + t.Fatal(err) + } + } + h := &maintenanceHandler{store: store} + b := &Bot{handler: h, ctx: ctx, logger: quietLogger()} + healthyDeleted := false + b.cleanupApprovalMessage = func(ctx context.Context, channel, id string) error { + if id == "1-slow" { + <-ctx.Done() + return ctx.Err() + } + healthyDeleted = true + return nil + } + timed, cancel := context.WithTimeout(ctx, 10*time.Millisecond) + _ = b.MaintainApprovals(timed) + cancel() + due, err := store.DueApprovalCleanup(ctx, time.Now()) + if err != nil || len(due) != 1 || due[0].MessageID != "2-healthy" { + t.Fatal("timed out row did not retain backoff", due, err) + } + if err := b.MaintainApprovals(ctx); err != nil { + t.Fatal(err) + } + if !healthyDeleted { + t.Fatal("slow first row starved healthy cleanup") + } +} diff --git a/internal/discordbot/approval_interactions.go b/internal/discordbot/approval_interactions.go new file mode 100644 index 0000000..6a618dc --- /dev/null +++ b/internal/discordbot/approval_interactions.go @@ -0,0 +1,188 @@ +package discordbot + +import ( + "context" + "fmt" + "strconv" + "strings" + "time" + + "github.com/bwmarrin/discordgo" + "github.com/mayvqt/Augur/internal/storage" +) + +const componentDeclineModal = "augur:decline-reason:" + +func (b *Bot) handleApproval(s interactionSession, i *discordgo.InteractionCreate, data discordgo.MessageComponentInteractionData, action string) { + if i.GuildID == "" || !canManageServer(i) { + b.ephemeral(s, i, "You need Manage Server permission to approve or decline requests.") + return + } + requestID, ok := approvalRequestID(data.CustomID, action) + if !ok { + b.ephemeral(s, i, "That approval button is invalid.") + return + } + if action == "decline" { + if i.Message == nil || i.Message.ID == "" || i.ChannelID == "" { + b.ephemeral(s, i, "That approval message is no longer available. Please retry from the approval channel.") + return + } + if err := s.InteractionRespond(i.Interaction, &discordgo.InteractionResponse{Type: discordgo.InteractionResponseModal, Data: &discordgo.InteractionResponseData{CustomID: componentDeclineModal + strconv.Itoa(requestID) + ":" + i.ChannelID + ":" + i.Message.ID, Title: "Decline request", Components: []discordgo.MessageComponent{discordgo.ActionsRow{Components: []discordgo.MessageComponent{discordgo.TextInput{CustomID: "reason", Label: "Reason (optional)", Style: discordgo.TextInputParagraph, Required: false, MaxLength: 500}}}}}}, discordgo.WithContext(b.ctx)); err != nil { + b.logger.Warn("open decline form", "error", err) + } + return + } + if !b.deferComponentUpdate(s, i) { + return + } + ctx, cancel := context.WithTimeout(b.ctx, 30*time.Second) + defer cancel() + presentation := approvalPresentation(i.Message) + presentation.Actor = interactionDisplayName(i) + decision, changed, err := b.handler.DecideRequest(ctx, requestID, action, presentation) + if err != nil { + b.logger.Error("update Seerr request", "request_id", requestID, "action", action, "error", err) + // Omitted embeds preserve the card through a transient failure and its retry. + content := "Could not finish updating this request. Try again to check its current status." + components := approvalComponents(requestID) + if _, err := s.InteractionResponseEdit(i.Interaction, &discordgo.WebhookEdit{Content: &content, Components: &components, AllowedMentions: noMentions()}, discordgo.WithContext(b.ctx)); err != nil { + b.logger.Warn("restore approval controls", "error", err) + } + return + } + b.finishApprovalInteraction(ctx, s, i, decision, changed) +} + +func modalReason(components []discordgo.MessageComponent) string { + for _, component := range components { + var children []discordgo.MessageComponent + switch row := component.(type) { + case *discordgo.ActionsRow: + if row != nil { + children = row.Components + } + case discordgo.ActionsRow: + children = row.Components + } + for _, child := range children { + var input *discordgo.TextInput + switch value := child.(type) { + case *discordgo.TextInput: + input = value + case discordgo.TextInput: + input = &value + } + if input != nil && input.CustomID == "reason" { + return truncate(strings.TrimSpace(input.Value), 500) + } + } + } + return "" +} + +func (b *Bot) handleDeclineModal(s interactionSession, i *discordgo.InteractionCreate) { + if i.GuildID == "" || !canManageServer(i) { + b.ephemeral(s, i, "You need Manage Server permission to decline requests.") + return + } + value, found := strings.CutPrefix(i.ModalSubmitData().CustomID, componentDeclineModal) + parts := strings.Split(value, ":") + if !found || len(parts) != 3 || !decimalID(parts[1]) || !decimalID(parts[2]) || parts[1] != i.ChannelID { + b.ephemeral(s, i, "That decline form is invalid.") + return + } + id, err := strconv.Atoi(parts[0]) + if err != nil || id <= 0 { + b.ephemeral(s, i, "That decline form is invalid.") + return + } + // Acknowledge before Discord/Seerr I/O, including the message lookup. + if !b.deferInteraction(s, i) { + return + } + ctx, cancel := context.WithTimeout(b.ctx, 30*time.Second) + defer cancel() + message, err := b.fetchApprovalMessage(ctx, parts[1], parts[2]) + if err != nil || message == nil || message.ID != parts[2] || message.ChannelID != parts[1] { + b.edit(s, i, "That approval message is no longer available. Please retry from the approval channel.") + return + } + i.Message = message + presentation := approvalPresentation(message) + presentation.Actor = interactionDisplayName(i) + presentation.Reason = modalReason(i.ModalSubmitData().Components) + decision, changed, err := b.handler.DecideRequest(ctx, id, "decline", presentation) + if err != nil { + b.logger.Warn("decline request", "request_id", id, "error", err) + b.edit(s, i, "Could not finish updating this request. Try again from the approval card to check its current status.") + return + } + b.finishApprovalInteraction(ctx, s, i, decision, changed) +} + +func (b *Bot) finishApprovalInteraction(ctx context.Context, s interactionSession, i *discordgo.InteractionCreate, d storage.ApprovalDecision, changed bool) { + var source *discordgo.MessageEmbed + if i.Message != nil && len(i.Message.Embeds) > 0 { + source = i.Message.Embeds[0] + } + embed := approvalDecisionEmbed(source, d) + content := "" + if !changed { + content = fmt.Sprintf("This request is already %s.", strings.ToLower(d.Status)) + } + var err error + if i.Type == discordgo.InteractionModalSubmit { + _, err = b.editApprovalMessage(ctx, &discordgo.MessageEdit{ID: i.Message.ID, Channel: i.ChannelID, Content: &content, Embeds: &[]*discordgo.MessageEmbed{embed}, Components: &[]discordgo.MessageComponent{}, AllowedMentions: noMentions()}) + response := fmt.Sprintf("Request %s.", strings.ToLower(d.Status)) + if !changed { + response = content + } + if err != nil { + response += " The approval card will update when Discord is available." + } + b.edit(s, i, response) + } else { + _, err = s.InteractionResponseEdit(i.Interaction, &discordgo.WebhookEdit{Content: &content, Embeds: &[]*discordgo.MessageEmbed{embed}, Components: &[]discordgo.MessageComponent{}, AllowedMentions: noMentions()}) + } + if err != nil { + b.logger.Warn("update approval card", "request_id", d.RequestID, "error", err) + return + } + if err := b.saveApprovalOutcome(ctx, func(saveCtx context.Context) error { + return b.handler.MarkApprovalDecided(saveCtx, storage.ApprovalMessage{RequestID: d.RequestID, GuildID: i.GuildID, ChannelID: i.ChannelID, MessageID: interactionMessageID(i), Status: d.Status}, time.Now().UTC()) + }); err != nil { + b.logger.Error("record approval card update", "request_id", d.RequestID, "error", err) + } +} + +func decimalID(value string) bool { + if value == "" { + return false + } + for _, r := range value { + if r < '0' || r > '9' { + return false + } + } + return true +} +func approvalRequestID(customID, action string) (int, bool) { + prefix := componentApprove + if action == "decline" { + prefix = componentDecline + } + value, found := strings.CutPrefix(customID, prefix) + if !found || !decimalID(value) { + return 0, false + } + id, err := strconv.Atoi(value) + return id, err == nil && id > 0 +} + +func interactionMessageID(i *discordgo.InteractionCreate) string { + if i.Message == nil { + return "" + } + return i.Message.ID +} diff --git a/internal/discordbot/approval_presentation.go b/internal/discordbot/approval_presentation.go new file mode 100644 index 0000000..7bfd8e2 --- /dev/null +++ b/internal/discordbot/approval_presentation.go @@ -0,0 +1,110 @@ +package discordbot + +import ( + "fmt" + "github.com/bwmarrin/discordgo" + "github.com/mayvqt/Augur/internal/seer" + "github.com/mayvqt/Augur/internal/storage" + "strconv" + "strings" +) + +func (b *Bot) approvalEmbed(approval seer.ApprovalRequest) *discordgo.MessageEmbed { + embed := b.mediaPreview(approval.Media, nil) + embed.Color = 0xFEE75C + requester := normalizeInlineText(approval.Requester) + if approval.RequesterID != "" { + requester = "<@" + approval.RequesterID + ">" + } else if requester == "" { + requester = "Unknown Seerr user" + } + embed.Fields = []*discordgo.MessageEmbedField{ + {Name: "Requested by", Value: requester, Inline: true}, + {Name: "Status", Value: "Pending approval", Inline: true}, + } + if approval.Media.MediaType == "tv" && (approval.Seasons.All || len(approval.Seasons.Numbers) > 0) { + embed.Fields = append(embed.Fields, &discordgo.MessageEmbedField{Name: "Seasons", Value: truncate(seasonSelectionLabel(approval.Seasons), 1024), Inline: true}) + } + embed.Footer = &discordgo.MessageEmbedFooter{Text: fmt.Sprintf("Seerr request #%d", approval.RequestID)} + return embed +} + +func approvalComponents(requestID int) []discordgo.MessageComponent { + id := strconv.Itoa(requestID) + return []discordgo.MessageComponent{discordgo.ActionsRow{Components: []discordgo.MessageComponent{ + discordgo.Button{CustomID: componentApprove + id, Label: "Approve", Style: discordgo.SuccessButton}, + discordgo.Button{CustomID: componentDecline + id, Label: "Decline", Style: discordgo.DangerButton}, + }}} +} + +func decisionSource(d storage.ApprovalDecision) *discordgo.MessageEmbed { + embed := &discordgo.MessageEmbed{Title: truncate(d.Title, 256), URL: d.URL} + if d.PosterURL != "" { + embed.Thumbnail = &discordgo.MessageEmbedThumbnail{URL: d.PosterURL} + } + return embed +} +func approvalPresentation(message *discordgo.Message) storage.ApprovalDecision { + var d storage.ApprovalDecision + if message == nil || len(message.Embeds) == 0 || message.Embeds[0] == nil { + return d + } + source := message.Embeds[0] + d.Title = source.Title + d.URL = source.URL + if source.Thumbnail != nil { + d.PosterURL = source.Thumbnail.URL + } + return d +} +func decidedApprovalEmbed(source *discordgo.MessageEmbed, requestID int, status, actor string) *discordgo.MessageEmbed { + if source == nil { + source = &discordgo.MessageEmbed{Title: fmt.Sprintf("Request #%d", requestID)} + } + embed := *source + embed.Color = 0x57F287 + if status == "Declined" { + embed.Color = 0xED4245 + } + embed.Fields = nil + for _, field := range source.Fields { + if field == nil || field.Name == "Status" || field.Name == "Decline reason" { + continue + } + copy := *field + embed.Fields = append(embed.Fields, ©) + } + embed.Fields = append(embed.Fields, &discordgo.MessageEmbedField{Name: "Status", Value: status, Inline: true}) + embed.Footer = &discordgo.MessageEmbedFooter{Text: truncate(fmt.Sprintf("Seerr request #%d · %s by %s", requestID, status, actor), 2048)} + return &embed +} +func approvalDecisionEmbed(source *discordgo.MessageEmbed, d storage.ApprovalDecision) *discordgo.MessageEmbed { + if source == nil { + source = decisionSource(d) + } + embed := decidedApprovalEmbed(source, d.RequestID, d.Status, d.Actor) + if d.Reason != "" { + embed.Fields = append(embed.Fields, &discordgo.MessageEmbedField{Name: "Decline reason", Value: truncate(d.Reason, 1000)}) + } + return embed +} +func decisionEmbed(source *discordgo.MessageEmbed, status string, reasons ...string) *discordgo.MessageEmbed { + embed := &discordgo.MessageEmbed{Title: "Your request was " + strings.ToLower(status), Color: 0x57F287} + if status == "Declined" { + embed.Color = 0xED4245 + } + if source != nil { + embed.Description = source.Title + embed.URL = source.URL + embed.Thumbnail = source.Thumbnail + } + embed.Footer = &discordgo.MessageEmbedFooter{Text: "Status: " + status} + reason := "" + if len(reasons) > 0 { + reason = reasons[0] + } + if strings.TrimSpace(reason) != "" { + embed.Fields = append(embed.Fields, &discordgo.MessageEmbedField{Name: "Decline reason", Value: truncate(reason, 1000)}) + } + return embed +} diff --git a/internal/discordbot/approvals.go b/internal/discordbot/approvals.go index 84969ec..48984ae 100644 --- a/internal/discordbot/approvals.go +++ b/internal/discordbot/approvals.go @@ -2,21 +2,12 @@ package discordbot import ( "context" - "errors" - "fmt" - "strconv" "strings" "time" "github.com/bwmarrin/discordgo" - "github.com/mayvqt/Augur/internal/config" - "github.com/mayvqt/Augur/internal/seer" - "github.com/mayvqt/Augur/internal/storage" ) -const approvalMessageRetention = 2 * time.Minute -const componentDeclineModal = "augur:decline-reason:" - func (b *Bot) handleApprovals(s interactionSession, i *discordgo.InteractionCreate) { if i.GuildID == "" || !canManageServer(i) { b.ephemeral(s, i, "You need Manage Server permission to configure approval messages.") @@ -43,7 +34,7 @@ func (b *Bot) handleApprovals(s interactionSession, i *discordgo.InteractionCrea return } if enabled { - permissions, err := b.session.UserChannelPermissions(b.session.State.User.ID, channelID) + permissions, err := b.session.UserChannelPermissions(b.session.State.User.ID, channelID, discordgo.WithContext(b.ctx)) if err != nil { b.logger.Warn("check approval channel permissions", "guild_id", i.GuildID, "channel_id", channelID, "error", err) b.ephemeral(s, i, "I could not inspect that channel. Make sure it belongs to this server and try again.") @@ -104,462 +95,6 @@ func missingApprovalPermissions(permissions int64) []string { return missing } -func (b *Bot) postApproval(ctx context.Context, guildID string, requestID int, requesterID string, result seer.SearchResult, seasons seer.SeasonSelection) { - if guildID == "" || requestID <= 0 { - return - } - channelID, enabled, err := b.handler.ApprovalChannel(ctx, guildID) - if err != nil || !enabled { - if err != nil { - b.logger.Error("load approval channel", "guild_id", guildID, "error", err) - } - return - } - claimed, err := b.handler.ClaimApproval(ctx, requestID, guildID, channelID) - if err != nil || !claimed { - if err != nil { - b.logger.Error("claim approval message", "request_id", requestID, "error", err) - } - return - } - b.sendApproval(ctx, storage.ApprovalSettings{GuildID: guildID, ChannelID: channelID, Enabled: true}, seer.ApprovalRequest{ - RequestID: requestID, RequesterID: requesterID, Media: result, Seasons: seasons, - }, true) -} - -func (b *Bot) ReconcileApprovals(ctx context.Context, approvals []seer.ApprovalRequest) error { - if err := b.cleanupDueApprovals(ctx); err != nil { - b.logger.Error("clean up decided approval messages", "error", err) - } - destinations, err := b.handler.ApprovalDestinations(ctx) - if err != nil { - return err - } - // Requests can be decided directly in Seerr; bring persisted cards up to date. - known := make(map[int]struct{}, len(approvals)) - for _, approval := range approvals { - known[approval.RequestID] = struct{}{} - } - if records, err := b.handler.ApprovalMessages(ctx); err == nil { - for _, record := range records { - if _, pending := known[record.RequestID]; pending || !record.DecidedAt.IsZero() { - continue - } - request, err := b.handler.RequestStatus(ctx, record.RequestID) - if err != nil || seer.IsPendingRequest(request.Status) { - continue - } - status := seer.RequestStatusLabel(request.Status) - message, err := b.session.ChannelMessage(record.ChannelID, record.MessageID) - if err != nil { - if discordNotFound(err) { - _ = b.handler.DeleteApprovalRecord(ctx, record.RequestID, record.GuildID) - } - continue - } - if len(message.Embeds) == 0 { - continue - } - embed := decidedApprovalEmbed(message.Embeds[0], record.RequestID, status, "Seerr") - if record.Reason != "" { - embed.Fields = append(embed.Fields, &discordgo.MessageEmbedField{Name: "Decline reason", Value: truncate(record.Reason, 1000)}) - } - if _, err := b.session.ChannelMessageEditComplex(&discordgo.MessageEdit{ID: record.MessageID, Channel: record.ChannelID, Embeds: &[]*discordgo.MessageEmbed{embed}, Components: &[]discordgo.MessageComponent{}}); err != nil { - if discordNotFound(err) { - _ = b.handler.DeleteApprovalRecord(ctx, record.RequestID, record.GuildID) - } - continue - } - if request.RequestedBy != nil { - b.notifyRequestDecision(ctx, request, embed, status, record.Reason) - } - _ = b.handler.MarkApprovalDecided(ctx, record.RequestID, record.GuildID, time.Now().UTC()) - } - } - for _, approval := range approvals { - for _, destination := range destinations { - if ctx.Err() != nil { - return ctx.Err() - } - b.sendApproval(ctx, destination, approval, false) - } - } - return nil -} - -func (b *Bot) sendApproval(ctx context.Context, destination storage.ApprovalSettings, approval seer.ApprovalRequest, alreadyClaimed bool) { - if !alreadyClaimed { - claimed, err := b.handler.ClaimApproval(ctx, approval.RequestID, destination.GuildID, destination.ChannelID) - if err != nil || !claimed { - if err != nil { - b.logger.Error("claim approval message", "request_id", approval.RequestID, "error", err) - } - return - } - } - embed := b.approvalEmbed(approval) - sendMessage := b.sendApprovalMessage - if sendMessage == nil { - sendMessage = func(channelID string, data *discordgo.MessageSend) (*discordgo.Message, error) { - return b.session.ChannelMessageSendComplex(channelID, data) - } - } - message, err := sendMessage(destination.ChannelID, &discordgo.MessageSend{ - Embeds: []*discordgo.MessageEmbed{embed}, Components: approvalComponents(approval.RequestID), AllowedMentions: noMentions(), - }) - if err != nil { - _ = b.handler.ReleaseApproval(ctx, approval.RequestID, destination.GuildID) - b.logger.Error("send approval message", "request_id", approval.RequestID, "channel_id", destination.ChannelID, "error", err) - return - } - if err := b.handler.FinishApproval(ctx, approval.RequestID, destination.GuildID, destination.ChannelID, message.ID); err != nil { - b.logger.Error("save approval message", "request_id", approval.RequestID, "error", err) - deleteMessage := b.cleanupApprovalMessage - if deleteMessage == nil { - deleteMessage = func(channelID, messageID string) error { - return b.session.ChannelMessageDelete(channelID, messageID) - } - } - if cleanupErr := deleteMessage(destination.ChannelID, message.ID); cleanupErr != nil { - b.logger.Error("delete untracked approval message", "request_id", approval.RequestID, "channel_id", destination.ChannelID, "message_id", message.ID, "error", cleanupErr) - } - if releaseErr := b.handler.ReleaseApproval(ctx, approval.RequestID, destination.GuildID); releaseErr != nil { - b.logger.Error("release approval claim", "request_id", approval.RequestID, "guild_id", destination.GuildID, "error", releaseErr) - } - } -} - -func (b *Bot) approvalEmbed(approval seer.ApprovalRequest) *discordgo.MessageEmbed { - embed := b.mediaPreview(approval.Media, nil) - embed.Color = 0xFEE75C - requester := normalizeInlineText(approval.Requester) - if approval.RequesterID != "" { - requester = "<@" + approval.RequesterID + ">" - } else if requester == "" { - requester = "Unknown Seerr user" - } - embed.Fields = []*discordgo.MessageEmbedField{ - {Name: "Requested by", Value: requester, Inline: true}, - {Name: "Status", Value: "Pending approval", Inline: true}, - } - if approval.Media.MediaType == "tv" && (approval.Seasons.All || len(approval.Seasons.Numbers) > 0) { - embed.Fields = append(embed.Fields, &discordgo.MessageEmbedField{Name: "Seasons", Value: seasonSelectionLabel(approval.Seasons), Inline: true}) - } - embed.Footer = &discordgo.MessageEmbedFooter{Text: fmt.Sprintf("Seerr request #%d", approval.RequestID)} - return embed -} - -func approvalComponents(requestID int) []discordgo.MessageComponent { - id := strconv.Itoa(requestID) - return []discordgo.MessageComponent{discordgo.ActionsRow{Components: []discordgo.MessageComponent{ - discordgo.Button{CustomID: componentApprove + id, Label: "Approve", Style: discordgo.SuccessButton}, - discordgo.Button{CustomID: componentDecline + id, Label: "Decline", Style: discordgo.DangerButton}, - }}} -} - -func (b *Bot) handleApproval(s interactionSession, i *discordgo.InteractionCreate, data discordgo.MessageComponentInteractionData, action string) { - if i.GuildID == "" || !canManageServer(i) { - b.ephemeral(s, i, "You need Manage Server permission to approve or decline requests.") - return - } - requestID, ok := approvalRequestID(data.CustomID, action) - if !ok { - b.ephemeral(s, i, "That approval button is invalid.") - return - } - if action == "decline" { - if i.Message == nil || i.Message.ID == "" || i.ChannelID == "" { - b.ephemeral(s, i, "That approval message is no longer available. Please retry from the approval channel.") - return - } - _ = s.InteractionRespond(i.Interaction, &discordgo.InteractionResponse{Type: discordgo.InteractionResponseModal, Data: &discordgo.InteractionResponseData{CustomID: componentDeclineModal + strconv.Itoa(requestID) + ":" + i.ChannelID + ":" + i.Message.ID, Title: "Decline request", Components: []discordgo.MessageComponent{discordgo.ActionsRow{Components: []discordgo.MessageComponent{discordgo.TextInput{CustomID: "reason", Label: "Reason (optional)", Style: discordgo.TextInputParagraph, Required: false, MaxLength: 500}}}}}}) - return - } - if !b.deferComponentUpdate(s, i) { - return - } - ctx, cancel := context.WithTimeout(b.ctx, 30*time.Second) - defer cancel() - request, err := b.handler.DecideRequest(ctx, requestID, action) - if err != nil { - b.logger.Error("update Seerr request", "request_id", requestID, "action", action, "error", err) - b.editContent(s, i, "Seerr could not update this request. Try again.", approvalComponents(requestID), "restore failed approval") - return - } - b.finishApprovalInteraction(ctx, s, i, request) -} - -func (b *Bot) handleDeclineModal(s interactionSession, i *discordgo.InteractionCreate) { - if i.GuildID == "" || !canManageServer(i) { - b.ephemeral(s, i, "You need Manage Server permission to decline requests.") - return - } - customID := i.ModalSubmitData().CustomID - value, found := strings.CutPrefix(customID, componentDeclineModal) - parts := strings.Split(value, ":") - if !found || len(parts) != 3 || !decimalID(parts[1]) || !decimalID(parts[2]) { - b.ephemeral(s, i, "That decline form is invalid.") - return - } - id, err := strconv.Atoi(parts[0]) - if err != nil || id <= 0 { - b.ephemeral(s, i, "That decline form is invalid.") - return - } - reason := "" - for _, row := range i.ModalSubmitData().Components { - if r, ok := row.(discordgo.ActionsRow); ok { - for _, c := range r.Components { - if input, ok := c.(discordgo.TextInput); ok && input.CustomID == "reason" { - reason = strings.TrimSpace(input.Value) - } - } - } - } - message, fetchErr := b.session.ChannelMessage(parts[1], parts[2]) - if fetchErr != nil || message == nil || message.ID != parts[2] || message.ChannelID != parts[1] || len(message.Embeds) == 0 { - b.ephemeral(s, i, "That approval message is no longer available. Please retry from the approval channel.") - return - } - i.Message = message - i.ChannelID = parts[1] - if !b.deferInteraction(s, i) { - return - } - ctx, cancel := context.WithTimeout(b.ctx, 30*time.Second) - defer cancel() - request, err := b.handler.DecideRequest(ctx, id, "decline") - if err != nil { - b.editContent(s, i, "Seerr could not update this request. Try again.", approvalComponents(id), "restore failed approval") - return - } - b.finishApprovalInteractionWithReason(ctx, s, i, request, reason) -} - -func decimalID(value string) bool { - if value == "" { - return false - } - for _, r := range value { - if r < '0' || r > '9' { - return false - } - } - return true -} - -func approvalRequestID(customID, action string) (int, bool) { - prefix := componentApprove - if action == "decline" { - prefix = componentDecline - } - value, found := strings.CutPrefix(customID, prefix) - if !found { - return 0, false - } - if value == "" { - return 0, false - } - if !decimalID(value) { - return 0, false - } - id, err := strconv.Atoi(value) - return id, err == nil && id > 0 -} - -func (b *Bot) finishApprovalInteraction(ctx context.Context, s interactionSession, i *discordgo.InteractionCreate, request seer.Request) { - b.finishApprovalInteractionWithReason(ctx, s, i, request, "") -} -func (b *Bot) finishApprovalInteractionWithReason(ctx context.Context, s interactionSession, i *discordgo.InteractionCreate, request seer.Request, reason string) { - if i.Message == nil || len(i.Message.Embeds) == 0 { - b.logger.Error("approval message has no embed", "request_id", request.ID) - return - } - status := seer.RequestStatusLabel(request.Status) - embed := decidedApprovalEmbed(i.Message.Embeds[0], request.ID, status, interactionDisplayName(i)) - if reason != "" { - embed.Fields = append(embed.Fields, &discordgo.MessageEmbedField{Name: "Decline reason", Value: truncate(reason, 1000)}) - } - empty := "" - components := []discordgo.MessageComponent{} - var editErr error - if i.Type == discordgo.InteractionModalSubmit && i.Message != nil { - _, editErr = b.session.ChannelMessageEditComplex(&discordgo.MessageEdit{ID: i.Message.ID, Channel: i.ChannelID, Content: &empty, Embeds: &[]*discordgo.MessageEmbed{embed}, Components: &components, AllowedMentions: noMentions()}) - } else if _, err := s.InteractionResponseEdit(i.Interaction, &discordgo.WebhookEdit{Content: &empty, Embeds: &[]*discordgo.MessageEmbed{embed}, Components: &components, AllowedMentions: noMentions()}); err != nil { - editErr = err - } - if editErr != nil { - b.logger.Error("update approval message", "request_id", request.ID, "error", editErr) - if i.Type == discordgo.InteractionModalSubmit { - b.editContent(s, i, "Could not update the approval message. Try again.", approvalComponents(request.ID), "restore failed approval") - } - return - } - if i.Type == discordgo.InteractionModalSubmit { - content := "" - if _, err := s.InteractionResponseEdit(i.Interaction, &discordgo.WebhookEdit{Content: &content, AllowedMentions: noMentions()}); err != nil { - b.logger.Warn("complete decline modal response", "request_id", request.ID, "error", err) - } - } - decidedAt := time.Now().UTC() - if err := b.handler.MarkApprovalDecided(ctx, request.ID, i.GuildID, decidedAt); err != nil { - b.logger.Error("persist approval decision", "request_id", request.ID, "error", err) - } - _ = b.handler.SetApprovalDecision(ctx, request.ID, i.GuildID, status, reason) - b.notifyRequestDecision(ctx, request, embed, status, reason) - channelID := i.ChannelID - if i.Message.ChannelID != "" { - channelID = i.Message.ChannelID - } - b.scheduleApprovalCleanup(storage.ApprovalMessage{ - RequestID: request.ID, GuildID: i.GuildID, ChannelID: channelID, MessageID: i.Message.ID, DecidedAt: decidedAt, - }) -} - -func decidedApprovalEmbed(source *discordgo.MessageEmbed, requestID int, status, actor string) *discordgo.MessageEmbed { - embed := *source - embed.Color = 0x57F287 - if status == "Declined" { - embed.Color = 0xED4245 - } - for _, field := range embed.Fields { - if field.Name == "Status" { - field.Value = status - } - } - embed.Footer = &discordgo.MessageEmbedFooter{Text: fmt.Sprintf("Seerr request #%d · %s by %s", requestID, status, actor)} - return &embed -} - -func (b *Bot) notifyRequestDecision(ctx context.Context, request seer.Request, source *discordgo.MessageEmbed, status, reason string) { - if request.RequestedBy == nil || request.RequestedBy.ID <= 0 { - return - } - ids, err := b.handler.RequesterDiscordIDs(ctx, request.RequestedBy.ID) - if err != nil { - b.logger.Error("load requester Discord IDs", "request_id", request.ID, "error", err) - return - } - seen := make(map[string]struct{}, len(ids)) - for _, discordID := range ids { - discordID = strings.TrimSpace(discordID) - if !config.IsDiscordID(discordID) { - continue - } - if _, duplicate := seen[discordID]; duplicate { - continue - } - seen[discordID] = struct{}{} - prefs, err := b.handler.NotificationPreferences(ctx, discordID) - if err != nil { - continue - } - if !decisionNotificationEnabled(status, prefs) { - continue - } - claimed, err := b.handler.ClaimDecisionNotification(ctx, request.ID, discordID, status) - if err != nil || !claimed { - continue - } - if err := b.sendDecisionDM(discordID, request.ID, decisionEmbed(source, status, reason)); err != nil { - if releaseErr := b.handler.ReleaseDecisionNotification(ctx, request.ID, discordID, status); releaseErr != nil { - b.logger.Warn("release failed decision notification claim", "request_id", request.ID, "error", releaseErr) - } - } - } -} - -func decisionNotificationEnabled(status string, prefs storage.NotificationPreferences) bool { - switch status { - case "Approved": - return prefs.Approved - case "Declined": - return prefs.Declined - default: - return false - } -} - -func (b *Bot) sendDecisionDM(discordID string, requestID int, embed *discordgo.MessageEmbed) error { - channel, err := b.session.UserChannelCreate(discordID) - if err != nil { - b.logger.Warn("open requester DM", "request_id", requestID, "error", err) - return err - } - if _, err := b.session.ChannelMessageSendComplex(channel.ID, &discordgo.MessageSend{Embeds: []*discordgo.MessageEmbed{embed}, AllowedMentions: noMentions()}); err != nil { - b.logger.Warn("send request decision DM", "request_id", requestID, "error", err) - return err - } - return nil -} - -func decisionEmbed(source *discordgo.MessageEmbed, status string, reasons ...string) *discordgo.MessageEmbed { - embed := &discordgo.MessageEmbed{Title: "Your request was " + strings.ToLower(status), Color: 0x57F287} - if status == "Declined" { - embed.Color = 0xED4245 - } - if source != nil { - embed.Description = source.Title - embed.URL = source.URL - embed.Thumbnail = source.Thumbnail - } - embed.Footer = &discordgo.MessageEmbedFooter{Text: "Status: " + status} - reason := "" - if len(reasons) > 0 { - reason = reasons[0] - } - if strings.TrimSpace(reason) != "" { - embed.Fields = append(embed.Fields, &discordgo.MessageEmbedField{Name: "Decline reason", Value: truncate(reason, 1000)}) - } - return embed -} - -func (b *Bot) scheduleApprovalCleanup(message storage.ApprovalMessage) { - delay := max(time.Until(message.DecidedAt.Add(approvalMessageRetention)), 0) - go func() { - timer := time.NewTimer(delay) - defer timer.Stop() - select { - case <-b.ctx.Done(): - return - case <-timer.C: - } - ctx, cancel := context.WithTimeout(b.ctx, 15*time.Second) - defer cancel() - b.deleteApprovalMessage(ctx, message) - }() -} - -func (b *Bot) cleanupDueApprovals(ctx context.Context) error { - messages, err := b.handler.DueApprovalMessages(ctx, time.Now().UTC().Add(-approvalMessageRetention)) - if err != nil { - return err - } - for _, message := range messages { - if ctx.Err() != nil { - return ctx.Err() - } - b.deleteApprovalMessage(ctx, message) - } - return nil -} - -func (b *Bot) deleteApprovalMessage(ctx context.Context, message storage.ApprovalMessage) { - err := b.session.ChannelMessageDelete(message.ChannelID, message.MessageID) - if err != nil && !discordNotFound(err) { - b.logger.Warn("delete decided approval message", "request_id", message.RequestID, "error", err) - return - } - if err := b.handler.DeleteApprovalRecord(ctx, message.RequestID, message.GuildID); err != nil { - b.logger.Error("delete approval record", "request_id", message.RequestID, "error", err) - } -} - -func discordNotFound(err error) bool { - var restErr *discordgo.RESTError - return errors.As(err, &restErr) && restErr.Response != nil && restErr.Response.StatusCode == 404 -} - func interactionDisplayName(i *discordgo.InteractionCreate) string { if i.Member != nil { if i.Member.Nick != "" { diff --git a/internal/discordbot/bot.go b/internal/discordbot/bot.go index 83249df..1d6f8be 100644 --- a/internal/discordbot/bot.go +++ b/internal/discordbot/bot.go @@ -4,6 +4,7 @@ import ( "context" "errors" "log/slog" + "sync" "time" "github.com/mayvqt/Augur/internal/config" @@ -21,20 +22,22 @@ type Handler interface { ConfigureApprovals(ctx context.Context, guildID, channelID string, enabled bool) error ApprovalChannel(ctx context.Context, guildID string) (string, bool, error) ApprovalDestinations(ctx context.Context) ([]storage.ApprovalSettings, error) - ClaimApproval(ctx context.Context, requestID int, guildID, channelID string) (bool, error) - FinishApproval(ctx context.Context, requestID int, guildID, channelID, messageID string) error - ReleaseApproval(ctx context.Context, requestID int, guildID string) error - MarkApprovalDecided(ctx context.Context, requestID int, guildID string, decidedAt time.Time) error + QueueUntrackedApprovalCleanup(ctx context.Context, message storage.ApprovalMessage) (bool, error) + DueApprovalCleanup(ctx context.Context) ([]storage.ApprovalMessage, error) + CompleteApprovalCleanup(ctx context.Context, message storage.ApprovalMessage) error + RetryApprovalCleanup(ctx context.Context, message storage.ApprovalMessage) error + ClaimApproval(ctx context.Context, requestID int, guildID, channelID string) (storage.ApprovalMessage, bool, error) + FinishApproval(ctx context.Context, message storage.ApprovalMessage) error + RetryApproval(ctx context.Context, message storage.ApprovalMessage) error + MarkApprovalDecided(ctx context.Context, message storage.ApprovalMessage, decidedAt time.Time) error DueApprovalMessages(ctx context.Context, before time.Time) ([]storage.ApprovalMessage, error) - DeleteApprovalRecord(ctx context.Context, requestID int, guildID string) error - DecideRequest(ctx context.Context, requestID int, action string) (seer.Request, error) - RequesterDiscordIDs(ctx context.Context, userID int) ([]string, error) + DeleteApprovalRecord(ctx context.Context, message storage.ApprovalMessage) error + DecideRequest(ctx context.Context, requestID int, action string, presentation storage.ApprovalDecision) (storage.ApprovalDecision, bool, error) + ApprovalDecision(ctx context.Context, requestID int, status string) (storage.ApprovalDecision, bool, error) + ObserveApprovalDecision(ctx context.Context, request seer.Request, presentation storage.ApprovalDecision) (storage.ApprovalDecision, error) RequestsForUser(ctx context.Context, discordID string, limit int) ([]seer.Request, error) NotificationPreferences(ctx context.Context, discordID string) (storage.NotificationPreferences, error) SetNotificationPreferences(ctx context.Context, preferences storage.NotificationPreferences) error - SetApprovalDecision(ctx context.Context, requestID int, guildID, status, reason string) error - ClaimDecisionNotification(ctx context.Context, requestID int, discordID, status string) (bool, error) - ReleaseDecisionNotification(ctx context.Context, requestID int, discordID, status string) error ApprovalMessages(ctx context.Context) ([]storage.ApprovalMessage, error) RequestStatus(ctx context.Context, requestID int) (seer.Request, error) MediaDetails(ctx context.Context, mediaType string, mediaID int) (seer.SearchResult, error) @@ -46,15 +49,29 @@ type interactionSession interface { } type Bot struct { - session *discordgo.Session - cfg config.DiscordConfig - link config.LinkConfig - handler Handler - logger *slog.Logger - ctx context.Context - cache selectionCache - sendApprovalMessage func(string, *discordgo.MessageSend) (*discordgo.Message, error) - cleanupApprovalMessage func(string, string) error + session *discordgo.Session + cfg config.DiscordConfig + link config.LinkConfig + handler Handler + logger *slog.Logger + ctx context.Context + cache selectionCache + pendingMu sync.Mutex + pendingIDs map[int]bool + pendingAt time.Time + cancel context.CancelFunc + closeOnce sync.Once + closeErr error + openSession func() error + closeSession func() error + registerApplicationCommands func() error + lifecycleMu sync.Mutex + closing bool + interactions sync.WaitGroup + sendApprovalMessage func(context.Context, string, *discordgo.MessageSend) (*discordgo.Message, error) + cleanupApprovalMessage func(context.Context, string, string) error + fetchApprovalMessage func(context.Context, string, string) (*discordgo.Message, error) + editApprovalMessage func(context.Context, *discordgo.MessageEdit) (*discordgo.Message, error) } func New(cfg config.DiscordConfig, link config.LinkConfig, handler Handler, logger *slog.Logger) (*Bot, error) { @@ -69,12 +86,21 @@ func New(cfg config.DiscordConfig, link config.LinkConfig, handler Handler, logg return nil, err } bot := &Bot{session: session, cfg: cfg, link: link, handler: handler, logger: logger, ctx: context.Background()} - bot.sendApprovalMessage = func(channelID string, data *discordgo.MessageSend) (*discordgo.Message, error) { - return session.ChannelMessageSendComplex(channelID, data) + bot.sendApprovalMessage = func(ctx context.Context, channelID string, data *discordgo.MessageSend) (*discordgo.Message, error) { + return session.ChannelMessageSendComplex(channelID, data, discordgo.WithContext(ctx)) } - bot.cleanupApprovalMessage = func(channelID, messageID string) error { - return session.ChannelMessageDelete(channelID, messageID) + bot.cleanupApprovalMessage = func(ctx context.Context, channelID, messageID string) error { + return session.ChannelMessageDelete(channelID, messageID, discordgo.WithContext(ctx)) } + bot.fetchApprovalMessage = func(ctx context.Context, channelID, messageID string) (*discordgo.Message, error) { + return session.ChannelMessage(channelID, messageID, discordgo.WithContext(ctx)) + } + bot.editApprovalMessage = func(ctx context.Context, edit *discordgo.MessageEdit) (*discordgo.Message, error) { + return session.ChannelMessageEditComplex(edit, discordgo.WithContext(ctx)) + } + bot.openSession = session.Open + bot.closeSession = session.Close + bot.registerApplicationCommands = bot.registerCommands session.AddHandler(bot.onReady) session.AddHandler(bot.onInteraction) return bot, nil diff --git a/internal/discordbot/bot_test.go b/internal/discordbot/bot_test.go index 5e1fef8..4ed0064 100644 --- a/internal/discordbot/bot_test.go +++ b/internal/discordbot/bot_test.go @@ -21,21 +21,27 @@ type fakeApprovalHandler struct { released bool } -func (f *fakeApprovalHandler) ClaimApproval(context.Context, int, string, string) (bool, error) { +func (f *fakeApprovalHandler) ClaimApproval(context.Context, int, string, string) (storage.ApprovalMessage, bool, error) { f.claimed = true - return true, nil + return storage.ApprovalMessage{RequestID: 42, GuildID: "guild-1", ChannelID: "channel-1", ClaimToken: "claim"}, true, nil } -func (f *fakeApprovalHandler) FinishApproval(context.Context, int, string, string, string) error { +func (f *fakeApprovalHandler) FinishApproval(context.Context, storage.ApprovalMessage) error { f.finished = true return errors.New("finish failed") } -func (f *fakeApprovalHandler) ReleaseApproval(context.Context, int, string) error { +func (f *fakeApprovalHandler) RetryApproval(context.Context, storage.ApprovalMessage) error { f.released = true return nil } +func (f *fakeApprovalHandler) QueueUntrackedApprovalCleanup(context.Context, storage.ApprovalMessage) (bool, error) { + return true, nil +} +func (f *fakeApprovalHandler) CompleteApprovalCleanup(context.Context, storage.ApprovalMessage) error { + return nil +} func TestSendApprovalCleansUpWhenFinishFails(t *testing.T) { t.Parallel() handler := &fakeApprovalHandler{} @@ -44,16 +50,16 @@ func TestSendApprovalCleansUpWhenFinishFails(t *testing.T) { handler: handler, logger: slog.Default(), ctx: context.Background(), - sendApprovalMessage: func(string, *discordgo.MessageSend) (*discordgo.Message, error) { + sendApprovalMessage: func(context.Context, string, *discordgo.MessageSend) (*discordgo.Message, error) { return &discordgo.Message{ID: "message-1"}, nil }, - cleanupApprovalMessage: func(channelID, messageID string) error { + cleanupApprovalMessage: func(_ context.Context, channelID, messageID string) error { deleted = channelID == "channel-1" && messageID == "message-1" return nil }, } - bot.sendApproval(context.Background(), storage.ApprovalSettings{GuildID: "guild-1", ChannelID: "channel-1"}, seer.ApprovalRequest{RequestID: 42, Media: seer.SearchResult{Title: "Arrival"}}, false) + bot.sendApproval(context.Background(), storage.ApprovalSettings{GuildID: "guild-1", ChannelID: "channel-1"}, seer.ApprovalRequest{RequestID: 42, Media: seer.SearchResult{Title: "Arrival"}}) if !handler.claimed || !handler.finished || !deleted || !handler.released { t.Fatalf("claim=%t finish=%t deleted=%t released=%t, want all true", handler.claimed, handler.finished, deleted, handler.released) @@ -145,7 +151,7 @@ func TestAvailableTitleOffersSeerrLinkAndBack(t *testing.T) { func TestRelevantQuotaLabelOnlyShowsSelectedMediaType(t *testing.T) { t.Parallel() quota := &seer.Quota{ - Movie: seer.QuotaUsage{Days: 7, Limit: 10, Used: 6, Remaining: 4, Restricted: true}, + Movie: seer.QuotaUsage{Days: 7, Limit: 10, Used: 6, Remaining: 4, Restricted: false}, TV: seer.QuotaUsage{Used: 2}, } got := relevantQuotaLabel("movie", quota) @@ -163,8 +169,8 @@ func TestSeasonPickerEnforcesLimitedQuotaAndHidesAllSeasons(t *testing.T) { {SeasonNumber: 3, Name: "Season 3", EpisodeCount: 6}, {SeasonNumber: 4, Name: "Season 4", EpisodeCount: 4}, } - quota := &seer.Quota{TV: seer.QuotaUsage{Restricted: true, Remaining: 3}} - components := seasonPickerComponents("cache", "result", seasons, quota, seer.SeasonSelection{}) + quota := &seer.Quota{TV: seer.QuotaUsage{Limit: 5, Used: 2, Restricted: false, Remaining: 3}} + components := seasonPickerComponents("cache", "result", seasons, quota, seer.SeasonSelection{}, 0) menu := components[0].(discordgo.ActionsRow).Components[0].(discordgo.SelectMenu) if menu.MaxValues != 3 { @@ -180,7 +186,7 @@ func TestSeasonPickerShowsAllSeasonsOnlyForUnlimitedQuota(t *testing.T) { t.Parallel() seasons := []seer.Season{{SeasonNumber: 1}, {SeasonNumber: 2}} quota := &seer.Quota{TV: seer.QuotaUsage{Restricted: false}} - components := seasonPickerComponents("cache", "result", seasons, quota, seer.SeasonSelection{}) + components := seasonPickerComponents("cache", "result", seasons, quota, seer.SeasonSelection{}, 0) buttons := components[1].(discordgo.ActionsRow).Components if len(buttons) != 2 { @@ -195,14 +201,14 @@ func TestSeasonPickerShowsAllSeasonsOnlyForUnlimitedQuota(t *testing.T) { func TestSeasonPickerKeepsSelectionAndShowsRequestAction(t *testing.T) { t.Parallel() seasons := []seer.Season{{SeasonNumber: 1}, {SeasonNumber: 2}} - quota := &seer.Quota{TV: seer.QuotaUsage{Restricted: true, Remaining: 2}} - components := seasonPickerComponents("cache", "result", seasons, quota, seer.SeasonSelection{Numbers: []int{2}}) + quota := &seer.Quota{TV: seer.QuotaUsage{Limit: 5, Used: 3, Restricted: false, Remaining: 2}} + components := seasonPickerComponents("cache", "result", seasons, quota, seer.SeasonSelection{Numbers: []int{2}}, 0) menu := components[0].(discordgo.ActionsRow).Components[0].(discordgo.SelectMenu) if menu.Options[0].Default || !menu.Options[1].Default { t.Fatalf("season defaults = %#v", menu.Options) } button := components[1].(discordgo.ActionsRow).Components[0].(discordgo.Button) - if button.CustomID != componentConfirm+"cache:result" || button.Label != "Request selected seasons" { + if button.CustomID != componentConfirm+"cache:result" || button.Label != "Request 1 selected season(s)" { t.Fatalf("request button = %#v", button) } } @@ -317,10 +323,10 @@ func TestFormatRequestLinesIncludesTerminalStatus(t *testing.T) { } func TestDecisionNotificationPreferencesSuppressMatchingStatus(t *testing.T) { - if decisionNotificationEnabled("Approved", storage.NotificationPreferences{Approved: false, Declined: true}) { + if (storage.NotificationPreferences{Approved: false, Declined: true}).DecisionEnabled("Approved") { t.Fatal("approved notification was enabled") } - if !decisionNotificationEnabled("Declined", storage.NotificationPreferences{Declined: true}) { + if !(storage.NotificationPreferences{Declined: true}).DecisionEnabled("Declined") { t.Fatal("declined notification was suppressed") } } diff --git a/internal/discordbot/cache.go b/internal/discordbot/cache.go index 84b5b7a..c9065b3 100644 --- a/internal/discordbot/cache.go +++ b/internal/discordbot/cache.go @@ -3,7 +3,9 @@ package discordbot import ( "crypto/rand" "encoding/hex" + "errors" "fmt" + "sort" "strconv" "sync" "time" @@ -22,12 +24,15 @@ type selectionCache struct { } type cachedSearch struct { - ownerID string - query string - expiresAt time.Time - results []cachedSelection - quota *seer.Quota - quotaKnown bool + interactionMu sync.Mutex + ownerID string + query string + expiresAt time.Time + results []cachedSelection + quota *seer.Quota + quotaKnown bool + submitting bool + submitted bool } type cachedSelection struct { @@ -64,7 +69,7 @@ func (c *selectionCache) setResults(cacheID, ownerID string, results []seer.Sear c.mu.Lock() defer c.mu.Unlock() search, ok := c.searchLocked(cacheID, ownerID) - if !ok { + if !ok || search.submitting || search.submitted { return false } search.results = make([]cachedSelection, len(results)) @@ -147,21 +152,6 @@ func (c *selectionCache) getQuota(cacheID, ownerID string) (*seer.Quota, bool) { return "aCopy, true } -func (c *selectionCache) setSeasons(cacheID, key, ownerID string, seasons seer.SeasonSelection) bool { - c.mu.Lock() - defer c.mu.Unlock() - - selection, ok := c.selectionLocked(cacheID, key, ownerID) - if !ok { - return false - } - selection.selectedSeasons = seer.SeasonSelection{ - Numbers: append([]int(nil), seasons.Numbers...), - All: seasons.All, - } - return true -} - func (c *selectionCache) setAvailableSeasons(cacheID, key, ownerID string, seasons []seer.Season) bool { c.mu.Lock() defer c.mu.Unlock() @@ -186,18 +176,115 @@ func (c *selectionCache) selection(cacheID, key, ownerID string) (seer.SearchRes }, true } -func (c *selectionCache) discard(cacheID, ownerID string) { +// A search owns its response ordering. Other searches remain independent; a +// concurrent click gets an immediate ephemeral acknowledgement instead of +// waiting behind a network call and later overwriting the completed response. +func (c *selectionCache) beginInteraction(cacheID, ownerID string) (func(), bool) { c.mu.Lock() defer c.mu.Unlock() + search, ok := c.searchLocked(cacheID, ownerID) + if !ok || !search.interactionMu.TryLock() { + return nil, false + } + search.expiresAt = time.Now().Add(selectionTTL) + return search.interactionMu.Unlock, true +} - if _, ok := c.searchLocked(cacheID, ownerID); ok { - delete(c.searches, cacheID) +// submissionState is checked before acknowledging components so stale clicks do +// not replace the response belonging to the request already in progress. +func (c *selectionCache) submissionState(cacheID, ownerID string) string { + c.mu.Lock() + defer c.mu.Unlock() + search, ok := c.searchLocked(cacheID, ownerID) + if !ok { + return "expired" + } + if search.submitted { + return "submitted" } + if search.submitting { + return "submitting" + } + return "" } -func (c *selectionCache) selectionLocked(cacheID, key, ownerID string) (*cachedSelection, bool) { +func (c *selectionCache) beginSubmission(cacheID, key, ownerID string, all bool) (seer.SearchResult, seer.SeasonSelection, string) { + c.mu.Lock() + defer c.mu.Unlock() search, ok := c.searchLocked(cacheID, ownerID) if !ok { + return seer.SearchResult{}, seer.SeasonSelection{}, "expired" + } + if search.submitted { + return seer.SearchResult{}, seer.SeasonSelection{}, "submitted" + } + if search.submitting { + return seer.SearchResult{}, seer.SeasonSelection{}, "submitting" + } + selection, ok := c.selectionLocked(cacheID, key, ownerID) + if !ok { + return seer.SearchResult{}, seer.SeasonSelection{}, "expired" + } + if all { + selection.selectedSeasons = seer.SeasonSelection{All: true} + } + search.submitting = true + return selection.result, seer.SeasonSelection{All: selection.selectedSeasons.All, Numbers: append([]int(nil), selection.selectedSeasons.Numbers...)}, "" +} + +func (c *selectionCache) finishSubmission(cacheID, ownerID string, consumed bool) { + c.mu.Lock() + defer c.mu.Unlock() + // Retain a tombstone until expiry. Delayed duplicate clicks must leave the + // successful response intact, including a request accepted before a local error. + if search, ok := c.searches[cacheID]; ok && search.ownerID == ownerID { + search.submitting = false + search.submitted = consumed + } +} + +func (c *selectionCache) selectSeasonPage(cacheID, key, ownerID string, page int, numbers []int) error { + c.mu.Lock() + defer c.mu.Unlock() + selection, ok := c.selectionLocked(cacheID, key, ownerID) + if !ok { + return errors.New("that season picker expired; run `/request` again") + } + visible, ok := seasonPage(selection.availableSeasons, page) + if !ok { + return errors.New("that season page is no longer available") + } + allowed := make(map[int]bool, len(visible)) + for _, season := range visible { + allowed[season.SeasonNumber] = true + } + merged := make(map[int]bool) + for _, number := range selection.selectedSeasons.Numbers { + if !allowed[number] { + merged[number] = true + } + } + for _, number := range numbers { + if !allowed[number] { + return errors.New("choose seasons from the displayed page") + } + merged[number] = true + } + next := seer.SeasonSelection{} + for number := range merged { + next.Numbers = append(next.Numbers, number) + } + sort.Ints(next.Numbers) + if err := validatePickerSelection(next, c.searches[cacheID].quota); err != nil { + return err + } + selection.selectedSeasons = next + return nil +} + +func (c *selectionCache) selectionLocked(cacheID, key, ownerID string) (*cachedSelection, bool) { + search, ok := c.searchLocked(cacheID, ownerID) + if !ok || search.submitting || search.submitted { return nil, false } index, err := strconv.Atoi(key) diff --git a/internal/discordbot/cache_test.go b/internal/discordbot/cache_test.go index b78df27..76249dd 100644 --- a/internal/discordbot/cache_test.go +++ b/internal/discordbot/cache_test.go @@ -48,7 +48,7 @@ func TestSelectionCacheSession(t *testing.T) { if _, ok := cache.getQuota("search", "owner"); ok { t.Fatal("quota was reported as cached before it was loaded") } - quota := &seer.Quota{TV: seer.QuotaUsage{Restricted: true, Remaining: 3}} + quota := &seer.Quota{TV: seer.QuotaUsage{Limit: 5, Used: 2, Restricted: false, Remaining: 3}} if !cache.setQuota("search", "owner", quota) { t.Fatal("owner could not store quota") } @@ -61,13 +61,13 @@ func TestSelectionCacheSession(t *testing.T) { if !ok || gotQuota.TV.Remaining != 3 { t.Fatal("caller mutated the quota stored in the cache") } - if !cache.setSeasons("search", "1", "owner", seer.SeasonSelection{Numbers: []int{1, 3}}) { - t.Fatal("owner could not store selected seasons") - } available := []seer.Season{{SeasonNumber: 1}, {SeasonNumber: 3}} if !cache.setAvailableSeasons("search", "1", "owner", available) { t.Fatal("owner could not store available seasons") } + if err := cache.selectSeasonPage("search", "1", "owner", 0, []int{1, 3}); err != nil { + t.Fatal(err) + } _, gotAvailable, seasons, ok := cache.selection("search", "1", "owner") if !ok { t.Fatal("owner could not read the selection") @@ -78,9 +78,9 @@ func TestSelectionCacheSession(t *testing.T) { if len(seasons.Numbers) != 2 || seasons.Numbers[0] != 1 || seasons.Numbers[1] != 3 { t.Fatalf("selected seasons = %#v", seasons) } - cache.discard("search", "owner") + cache.finishSubmission("search", "owner", true) if _, _, _, ok := cache.selection("search", "1", "owner"); ok { - t.Fatal("discard did not invalidate the search") + t.Fatal("submission did not invalidate the selection") } } diff --git a/internal/discordbot/commands.go b/internal/discordbot/commands.go index 6a34639..be2a616 100644 --- a/internal/discordbot/commands.go +++ b/internal/discordbot/commands.go @@ -9,6 +9,7 @@ const ( commandRequests = "requests" commandNotifications = "notifications" componentPick = "augur:pick:" + componentSeasonPage = "augur:season-page:" componentSeasons = "augur:seasons:" componentAll = "augur:all:" componentConfirm = "augur:confirm:" diff --git a/internal/discordbot/handlers.go b/internal/discordbot/handlers.go index d07169b..b6983e5 100644 --- a/internal/discordbot/handlers.go +++ b/internal/discordbot/handlers.go @@ -232,9 +232,27 @@ func (b *Bot) searchAndShow(ctx context.Context, s interactionSession, i *discor func (b *Bot) handleComponent(s interactionSession, i *discordgo.InteractionCreate) { data := i.MessageComponentData() + for _, prefix := range []string{componentPick, componentSeasons, componentSeasonPage, componentAll, componentConfirm, componentBack, componentRetry, componentSearch} { + if value, ok := strings.CutPrefix(data.CustomID, prefix); ok { + cacheID, _, _ := strings.Cut(value, ":") + unlock, ok := b.cache.beginInteraction(cacheID, interactionUserID(i)) + if !ok { + b.ephemeral(s, i, "This search is busy or expired. Wait for the current action, or run `/request` again.") + return + } + defer unlock() + if state := b.cache.submissionState(cacheID, interactionUserID(i)); state == "submitting" || state == "submitted" { + b.submissionNotice(s, i, state) + return + } + break + } + } switch { case strings.HasPrefix(data.CustomID, componentPick): b.handlePick(s, i, data) + case strings.HasPrefix(data.CustomID, componentSeasonPage): + b.handleSeasonPage(s, i, data) case strings.HasPrefix(data.CustomID, componentSeasons): b.handleSeasons(s, i, data) case strings.HasPrefix(data.CustomID, componentAll): @@ -296,22 +314,27 @@ func (b *Bot) handlePick(s interactionSession, i *discordgo.InteractionCreate, d b.edit(s, i, expiredPickerMessage) return } - b.editPreview( - s, - i, - b.mediaPreview(result, quota), - seasonPickerComponents(cacheID, key, seasons, quota, seer.SeasonSelection{}), - ) + b.showSeasonSelection(s, i, cacheID, key, 0, "") } -func (b *Bot) handleSeasons(s interactionSession, i *discordgo.InteractionCreate, data discordgo.MessageComponentInteractionData) { - if len(data.Values) == 0 { - return +func seasonComponentSelection(customID, prefix string) (cacheID, key string, page int, ok bool) { + cacheID, tail, valid := componentSelection(customID, prefix) + if !valid { + return "", "", 0, false } + key, pageText, hasPage := strings.Cut(tail, ":") + if !hasPage { + return cacheID, key, 0, key != "" + } + page, err := strconv.Atoi(pageText) + return cacheID, key, page, err == nil && page >= 0 && key != "" +} + +func (b *Bot) handleSeasons(s interactionSession, i *discordgo.InteractionCreate, data discordgo.MessageComponentInteractionData) { if !b.deferComponentUpdate(s, i) { return } - cacheID, key, ok := componentSelection(data.CustomID, componentSeasons) + cacheID, key, page, ok := seasonComponentSelection(data.CustomID, componentSeasons) if !ok { b.edit(s, i, "That season picker is invalid. Run `/request` again.") return @@ -321,35 +344,37 @@ func (b *Bot) handleSeasons(s interactionSession, i *discordgo.InteractionCreate b.edit(s, i, err.Error()) return } - b.showSeasonSelection(s, i, cacheID, key, selection) + notice := "" + if err := b.cache.selectSeasonPage(cacheID, key, interactionUserID(i), page, selection.Numbers); err != nil { + notice = err.Error() + } + b.showSeasonSelection(s, i, cacheID, key, page, notice) } -func (b *Bot) handleAllSeasons(s interactionSession, i *discordgo.InteractionCreate, data discordgo.MessageComponentInteractionData) { +func (b *Bot) handleSeasonPage(s interactionSession, i *discordgo.InteractionCreate, data discordgo.MessageComponentInteractionData) { if !b.deferComponentUpdate(s, i) { return } - cacheID, key, ok := componentSelection(data.CustomID, componentAll) + cacheID, key, page, ok := seasonComponentSelection(data.CustomID, componentSeasonPage) if !ok { - b.edit(s, i, "That season selection is invalid. Run `/request` again.") + b.edit(s, i, "That season page is invalid. Run `/request` again.") return } - ownerID := interactionUserID(i) - if !b.cache.setSeasons(cacheID, key, ownerID, seer.SeasonSelection{All: true}) { - b.edit(s, i, expiredSeasonPickerMessage) + b.showSeasonSelection(s, i, cacheID, key, page, "") +} + +func (b *Bot) handleAllSeasons(s interactionSession, i *discordgo.InteractionCreate, data discordgo.MessageComponentInteractionData) { + cacheID, key, ok := componentSelection(data.CustomID, componentAll) + if !ok { + b.ephemeral(s, i, "That season selection is invalid. Run `/request` again.") return } - b.submitRequest(s, i, cacheID, key) + b.submitRequest(s, i, cacheID, key, true) } -func (b *Bot) showSeasonSelection( - s interactionSession, - i *discordgo.InteractionCreate, - cacheID string, - key string, - selection seer.SeasonSelection, -) { +func (b *Bot) showSeasonSelection(s interactionSession, i *discordgo.InteractionCreate, cacheID, key string, page int, notice string) { ownerID := interactionUserID(i) - result, seasons, _, ok := b.cache.selection(cacheID, key, ownerID) + result, seasons, selection, ok := b.cache.selection(cacheID, key, ownerID) if !ok || result.MediaType != "tv" { b.edit(s, i, expiredSeasonPickerMessage) return @@ -359,40 +384,52 @@ func (b *Bot) showSeasonSelection( b.edit(s, i, expiredSeasonPickerMessage) return } - if err := validatePickerSelection(selection, quota); err != nil { - b.edit(s, i, err.Error()) - return - } - if !b.cache.setSeasons(cacheID, key, ownerID, selection) { - b.edit(s, i, expiredSeasonPickerMessage) - return + if _, valid := seasonPage(seasons, page); !valid { + page = 0 } embed := b.mediaPreview(result, quota) - b.editPreview( - s, - i, - embed, - seasonPickerComponents(cacheID, key, seasons, quota, selection), - ) + if len(selection.Numbers) > 0 { + embed.Fields = append(embed.Fields, &discordgo.MessageEmbedField{Name: "Selected seasons", Value: truncate(seasonSelectionLabel(selection), 1024)}) + } + if notice != "" { + embed.Footer.Text = truncate(notice, 2048) + } + if len(seasons) == 0 { + embed.Footer.Text = "No seasons are available to select for this show." + } + b.editPreview(s, i, embed, seasonPickerComponents(cacheID, key, seasons, quota, selection, page)) } func (b *Bot) handleConfirm(s interactionSession, i *discordgo.InteractionCreate, data discordgo.MessageComponentInteractionData) { - if !b.deferComponentUpdate(s, i) { - return - } cacheID, key, ok := componentSelection(data.CustomID, componentConfirm) if !ok { - b.edit(s, i, "That confirmation is invalid. Run `/request` again.") + b.ephemeral(s, i, "That confirmation is invalid. Run `/request` again.") return } - b.submitRequest(s, i, cacheID, key) + b.submitRequest(s, i, cacheID, key, false) } -func (b *Bot) submitRequest(s interactionSession, i *discordgo.InteractionCreate, cacheID, key string) { +func (b *Bot) submissionNotice(s interactionSession, i *discordgo.InteractionCreate, state string) { + message := "That confirmation expired. Run `/request` again." + if state == "submitting" { + message = "Your request is being submitted. Please wait." + } + if state == "submitted" { + message = "This confirmation has already been used. Check `/requests` or Seerr for its status." + } + b.ephemeral(s, i, message) +} + +func (b *Bot) submitRequest(s interactionSession, i *discordgo.InteractionCreate, cacheID, key string, all bool) { ownerID := interactionUserID(i) - result, _, seasons, ok := b.cache.selection(cacheID, key, ownerID) - if !ok { - b.edit(s, i, "That confirmation expired or was already used. Run `/request` again.") + result, seasons, state := b.cache.beginSubmission(cacheID, key, ownerID, all) + if state != "" { + b.submissionNotice(s, i, state) + return + } + consumed := false + defer func() { b.cache.finishSubmission(cacheID, ownerID, consumed) }() + if !b.deferComponentUpdate(s, i) { return } ctx, cancel := context.WithTimeout(b.ctx, 30*time.Second) @@ -405,26 +442,26 @@ func (b *Bot) submitRequest(s interactionSession, i *discordgo.InteractionCreate message = userErr.UserMessage() } var submitted interface{ RequestSubmitted() bool } - if errors.As(err, &submitted) && submitted.RequestSubmitted() { - b.cache.discard(cacheID, ownerID) - b.edit(s, i, truncate(message, 300)) - b.logger.Error("track submitted request", "request_id", req.ID, "media_type", result.MediaType, "media_id", result.ID, "error", err) - return + var retry interface{ SafeToRetry() bool } + consumed = (errors.As(err, &submitted) && submitted.RequestSubmitted()) || (errors.As(err, &retry) && !retry.SafeToRetry()) + if consumed { + b.edit(s, i, truncate(message, 500)) + } else { + b.editRetry(s, i, truncate(message, 300), cacheID, key) } - b.editRetry(s, i, truncate(message, 180), cacheID, key) - b.logger.Error("request media failed", "media_type", result.MediaType, "media_id", result.ID, "error", err) + b.logger.Error("request media failed", "request_id", req.ID, "media_type", result.MediaType, "media_id", result.ID, "error", err) return } - b.cache.discard(cacheID, ownerID) - if seer.IsPendingRequest(req.Status) { - b.postApproval(ctx, i.GuildID, req.ID, ownerID, result, seasons) - } + consumed = true title := escapeMarkdown(optionLabel(result)) selection := "" if result.MediaType == "tv" { - selection = "\nSelected: **" + escapeMarkdown(seasonSelectionLabel(seasons)) + "**" + selection = "\nSelected: **" + truncate(escapeMarkdown(seasonSelectionLabel(seasons)), 900) + "**" + } + b.edit(s, i, fmt.Sprintf("Your request for **%s** has been submitted.%s\nCheck `/requests` for its status. DM preferences are available in `/notifications`.", title, selection)) + if seer.IsPendingRequest(req.Status) { + b.postApproval(ctx, i.GuildID, req.ID, ownerID, result, seasons) } - b.edit(s, i, fmt.Sprintf("Your request for **%s** has been submitted successfully.%s\nYou will receive a direct message when it is fully available.", title, selection)) } func (b *Bot) handleBack(s interactionSession, i *discordgo.InteractionCreate, data discordgo.MessageComponentInteractionData) { @@ -448,15 +485,12 @@ func (b *Bot) handleBack(s interactionSession, i *discordgo.InteractionCreate, d } func (b *Bot) handleRetry(s interactionSession, i *discordgo.InteractionCreate, data discordgo.MessageComponentInteractionData) { - if !b.deferComponentUpdate(s, i) { - return - } cacheID, key, ok := componentSelection(data.CustomID, componentRetry) if !ok { - b.edit(s, i, "That retry expired. Run `/request` again.") + b.ephemeral(s, i, "That retry expired. Run `/request` again.") return } - b.submitRequest(s, i, cacheID, key) + b.submitRequest(s, i, cacheID, key, false) } func (b *Bot) handleSearchRetry(s interactionSession, i *discordgo.InteractionCreate, data discordgo.MessageComponentInteractionData) { @@ -489,9 +523,6 @@ func parseSeasonValues(values []string) (seer.SeasonSelection, error) { seen[number] = struct{}{} selection.Numbers = append(selection.Numbers, number) } - if len(selection.Numbers) == 0 { - return seer.SeasonSelection{}, errors.New("select at least one season") - } sort.Ints(selection.Numbers) return selection, nil } @@ -507,12 +538,12 @@ func componentSelection(customID, prefix string) (cacheID, key string, ok bool) func validatePickerSelection(selection seer.SeasonSelection, quota *seer.Quota) error { if selection.All { - if quota == nil || quota.TV.Restricted { + if quota == nil || quota.TV.Limited() { return errors.New("all seasons are only available with an unlimited TV request limit") } return nil } - if quota != nil && quota.TV.Restricted && len(selection.Numbers) > quota.TV.Remaining { + if quota != nil && quota.TV.Limited() && len(selection.Numbers) > quota.TV.Remaining { return fmt.Errorf("you can select up to %d more TV season(s)", quota.TV.Remaining) } return nil diff --git a/internal/discordbot/lifecycle.go b/internal/discordbot/lifecycle.go index 582f532..a20a404 100644 --- a/internal/discordbot/lifecycle.go +++ b/internal/discordbot/lifecycle.go @@ -10,41 +10,58 @@ import ( "github.com/mayvqt/Augur/internal/seer" ) -func (b *Bot) Start(ctx context.Context) error { +func (b *Bot) Start(ctx context.Context) (startErr error) { + // Existing cards can receive interactions as soon as the gateway opens. + // Every failed startup must close admission and join those interactions. + defer func() { + if startErr != nil { + startErr = errors.Join(startErr, b.Close()) + } + }() if err := ctx.Err(); err != nil { return err } - b.ctx = ctx + b.ctx, b.cancel = context.WithCancel(ctx) b.session.Identify.Intents = discordgo.IntentsGuilds | discordgo.IntentsDirectMessages - if err := b.session.Open(); err != nil { + if err := b.openSession(); err != nil { return err } if err := ctx.Err(); err != nil { - return errors.Join(err, b.session.Close()) + return err } if err := b.applyPresence(); err != nil { b.logger.Error("discord presence update failed", "error", err) } if err := ctx.Err(); err != nil { - return errors.Join(err, b.session.Close()) + return err } - if err := b.registerCommands(); err != nil { - return errors.Join(err, b.session.Close()) + if err := b.registerApplicationCommands(); err != nil { + return err } b.logger.Info("discord slash commands registered", "guild_id", b.cfg.GuildID) return nil } func (b *Bot) Close() error { - b.logger.Info("closing discord session") - return b.session.Close() + b.closeOnce.Do(func() { + b.lifecycleMu.Lock() + b.closing = true + if b.cancel != nil { + b.cancel() + } + b.lifecycleMu.Unlock() + b.logger.Info("closing discord session") + b.closeErr = b.closeSession() + b.interactions.Wait() + }) + return b.closeErr } func (b *Bot) NotifyComplete(ctx context.Context, discordID string, media seer.SearchResult) error { if err := ctx.Err(); err != nil { return err } - channel, err := b.session.UserChannelCreate(discordID) + channel, err := b.session.UserChannelCreate(discordID, discordgo.WithContext(ctx)) if err != nil { return err } @@ -55,17 +72,17 @@ func (b *Bot) NotifyComplete(ctx context.Context, discordID string, media seer.S _, err = b.session.ChannelMessageSendComplex(channel.ID, &discordgo.MessageSend{ Embeds: []*discordgo.MessageEmbed{embed}, AllowedMentions: noMentions(), - }) + }, discordgo.WithContext(ctx)) return err } func (b *Bot) registerCommands() error { appID := b.session.State.User.ID - if _, err := b.session.ApplicationCommandBulkOverwrite(appID, b.cfg.GuildID, slashCommands()); err != nil { + if _, err := b.session.ApplicationCommandBulkOverwrite(appID, b.cfg.GuildID, slashCommands(), discordgo.WithContext(b.ctx)); err != nil { return fmt.Errorf("register slash commands: %w", err) } if b.cfg.GuildID != "" { - if _, err := b.session.ApplicationCommandBulkOverwrite(appID, "", []*discordgo.ApplicationCommand{}); err != nil { + if _, err := b.session.ApplicationCommandBulkOverwrite(appID, "", []*discordgo.ApplicationCommand{}, discordgo.WithContext(b.ctx)); err != nil { return fmt.Errorf("remove duplicate global slash commands: %w", err) } return nil @@ -74,7 +91,7 @@ func (b *Bot) registerCommands() error { if guild == nil { continue } - if _, err := b.session.ApplicationCommandBulkOverwrite(appID, guild.ID, []*discordgo.ApplicationCommand{}); err != nil { + if _, err := b.session.ApplicationCommandBulkOverwrite(appID, guild.ID, []*discordgo.ApplicationCommand{}, discordgo.WithContext(b.ctx)); err != nil { return fmt.Errorf("remove duplicate slash commands from guild %s: %w", guild.ID, err) } } @@ -92,6 +109,14 @@ func (b *Bot) onInteraction(s *discordgo.Session, interaction *discordgo.Interac if interaction == nil || interaction.Interaction == nil { return } + b.lifecycleMu.Lock() + if b.closing || b.ctx.Err() != nil { + b.lifecycleMu.Unlock() + return + } + b.interactions.Add(1) + b.lifecycleMu.Unlock() + defer b.interactions.Done() switch interaction.Type { case discordgo.InteractionApplicationCommand: b.handleCommand(s, interaction) diff --git a/internal/discordbot/preview.go b/internal/discordbot/preview.go index 0200df5..408f563 100644 --- a/internal/discordbot/preview.go +++ b/internal/discordbot/preview.go @@ -90,65 +90,67 @@ func resultPickerRow(cacheID, placeholder string, options []discordgo.SelectMenu }} } -func seasonPickerComponents(cacheID, key string, seasons []seer.Season, quota *seer.Quota, selected seer.SeasonSelection) []discordgo.MessageComponent { - maxSelections := len(seasons) - if quota != nil && quota.TV.Restricted && quota.TV.Remaining < maxSelections { - maxSelections = quota.TV.Remaining +const seasonsPerPage = 25 + +func seasonPage(seasons []seer.Season, page int) ([]seer.Season, bool) { + if page < 0 || page > (len(seasons)-1)/seasonsPerPage || len(seasons) == 0 { + return nil, false } - if maxSelections <= 0 { + start := page * seasonsPerPage + return seasons[start:min(start+seasonsPerPage, len(seasons))], true +} + +func seasonPickerComponents(cacheID, key string, seasons []seer.Season, quota *seer.Quota, selected seer.SeasonSelection, page int) []discordgo.MessageComponent { + visible, ok := seasonPage(seasons, page) + if !ok { return backComponents(cacheID) } - - options := make([]discordgo.SelectMenuOption, 0, min(len(seasons), 25)) - for _, season := range seasons { - if len(options) == 25 { - break + onPage := 0 + options := make([]discordgo.SelectMenuOption, 0, len(visible)) + for _, season := range visible { + chosen := containsSeason(selected.Numbers, season.SeasonNumber) + if chosen { + onPage++ } - label := seasonLabel(season) description := "" if season.EpisodeCount > 0 { description = fmt.Sprintf("%d episodes", season.EpisodeCount) } options = append(options, discordgo.SelectMenuOption{ - Label: truncate(label, 100), - Description: description, - Value: strconv.Itoa(season.SeasonNumber), - Default: containsSeason(selected.Numbers, season.SeasonNumber), + Label: truncate(seasonLabel(season), 100), Description: description, + Value: strconv.Itoa(season.SeasonNumber), Default: chosen, }) } - if len(options) == 0 { - return backComponents(cacheID) + maxSelections := len(options) + if quota != nil && quota.TV.Limited() { + maxSelections = min(maxSelections, max(0, quota.TV.Remaining-(len(selected.Numbers)-onPage))) } - maxSelections = min(maxSelections, len(options)) placeholder := fmt.Sprintf("Choose up to %d season(s)", maxSelections) - buttons := []discordgo.MessageComponent{} - if quota != nil && !quota.TV.Restricted { - buttons = append(buttons, discordgo.Button{ - CustomID: componentAll + cacheID + ":" + key, - Label: "Request all seasons", - Style: discordgo.PrimaryButton, - }) + if maxSelections == 0 { + placeholder = "No quota left; clear other selections first" + } + rows := []discordgo.MessageComponent{discordgo.ActionsRow{Components: []discordgo.MessageComponent{ + discordgo.SelectMenu{ + CustomID: componentSeasons + cacheID + ":" + key + ":" + strconv.Itoa(page), Placeholder: placeholder, + MinValues: intPtr(0), MaxValues: max(1, maxSelections), Disabled: maxSelections == 0, Options: options, + }, + }}} + pages := (len(seasons) + seasonsPerPage - 1) / seasonsPerPage + if pages > 1 { + rows = append(rows, discordgo.ActionsRow{Components: []discordgo.MessageComponent{ + discordgo.Button{CustomID: componentSeasonPage + cacheID + ":" + key + ":" + strconv.Itoa(max(0, page-1)), Label: "Previous seasons", Style: discordgo.SecondaryButton, Disabled: page == 0}, + discordgo.Button{CustomID: componentSeasonPage + cacheID + ":" + key + ":" + strconv.Itoa(page+1), Label: fmt.Sprintf("Next seasons (%d/%d)", page+1, pages), Style: discordgo.SecondaryButton, Disabled: page == pages-1}, + }}) } + buttons := []discordgo.MessageComponent{} if len(selected.Numbers) > 0 { - buttons = append([]discordgo.MessageComponent{discordgo.Button{ - CustomID: componentConfirm + cacheID + ":" + key, - Label: "Request selected seasons", - Style: discordgo.SuccessButton, - }}, buttons...) + buttons = append(buttons, discordgo.Button{CustomID: componentConfirm + cacheID + ":" + key, Label: fmt.Sprintf("Request %d selected season(s)", len(selected.Numbers)), Style: discordgo.SuccessButton}) } - buttons = append(buttons, backButton(cacheID)) - return []discordgo.MessageComponent{ - discordgo.ActionsRow{Components: []discordgo.MessageComponent{ - discordgo.SelectMenu{ - CustomID: componentSeasons + cacheID + ":" + key, - Placeholder: placeholder, - MinValues: intPtr(1), - MaxValues: maxSelections, - Options: options, - }, - }}, - discordgo.ActionsRow{Components: buttons}, + if quota != nil && !quota.TV.Limited() { + buttons = append(buttons, discordgo.Button{CustomID: componentAll + cacheID + ":" + key, Label: "Request all seasons", Style: discordgo.PrimaryButton}) } + buttons = append(buttons, backButton(cacheID)) + return append(rows, discordgo.ActionsRow{Components: buttons}) } func backComponents(cacheID string) []discordgo.MessageComponent { @@ -230,7 +232,7 @@ func relevantQuotaLabel(mediaType string, quota *seer.Quota) string { } func quotaUsageLabel(usage seer.QuotaUsage) string { - if !usage.Restricted { + if !usage.Limited() { return fmt.Sprintf("%d used · Unlimited", usage.Used) } window := "rolling window" diff --git a/internal/discordbot/responses.go b/internal/discordbot/responses.go index 7c26fe7..330aa8a 100644 --- a/internal/discordbot/responses.go +++ b/internal/discordbot/responses.go @@ -7,7 +7,7 @@ import ( ) func (b *Bot) respond(s interactionSession, i *discordgo.InteractionCreate, data *discordgo.InteractionResponseData) { - if err := s.InteractionRespond(i.Interaction, &discordgo.InteractionResponse{Type: discordgo.InteractionResponseChannelMessageWithSource, Data: data}); err != nil { + if err := s.InteractionRespond(i.Interaction, &discordgo.InteractionResponse{Type: discordgo.InteractionResponseChannelMessageWithSource, Data: data}, discordgo.WithContext(b.ctx)); err != nil { b.logger.Error("respond interaction", "error", err) } } @@ -20,7 +20,7 @@ func (b *Bot) deferInteraction(s interactionSession, i *discordgo.InteractionCre err := s.InteractionRespond(i.Interaction, &discordgo.InteractionResponse{ Type: discordgo.InteractionResponseDeferredChannelMessageWithSource, Data: &discordgo.InteractionResponseData{Flags: discordgo.MessageFlagsEphemeral, AllowedMentions: noMentions()}, - }) + }, discordgo.WithContext(b.ctx)) if err != nil { b.logger.Error("defer interaction", "error", err) return false @@ -31,7 +31,7 @@ func (b *Bot) deferInteraction(s interactionSession, i *discordgo.InteractionCre func (b *Bot) deferComponentUpdate(s interactionSession, i *discordgo.InteractionCreate) bool { err := s.InteractionRespond(i.Interaction, &discordgo.InteractionResponse{ Type: discordgo.InteractionResponseDeferredMessageUpdate, - }) + }, discordgo.WithContext(b.ctx)) if err != nil { b.logger.Error("defer component interaction", "error", err) return false @@ -50,7 +50,7 @@ func (b *Bot) editPreview(s interactionSession, i *discordgo.InteractionCreate, Embeds: &[]*discordgo.MessageEmbed{embed}, AllowedMentions: noMentions(), Components: &components, - }) + }, discordgo.WithContext(b.ctx)) if err != nil { b.logger.Error("edit request preview", "error", err) } @@ -88,7 +88,7 @@ func (b *Bot) editContent(s interactionSession, i *discordgo.InteractionCreate, Embeds: &[]*discordgo.MessageEmbed{}, Components: &components, AllowedMentions: noMentions(), - }) + }, discordgo.WithContext(b.ctx)) if err != nil { b.logger.Error(logMessage, "error", err) } diff --git a/internal/discordbot/selection_flow_test.go b/internal/discordbot/selection_flow_test.go new file mode 100644 index 0000000..c7ff8af --- /dev/null +++ b/internal/discordbot/selection_flow_test.go @@ -0,0 +1,182 @@ +package discordbot + +import ( + "context" + "log/slog" + "reflect" + "strconv" + "strings" + "sync" + "sync/atomic" + "testing" + + "github.com/bwmarrin/discordgo" + "github.com/mayvqt/Augur/internal/seer" +) + +func testSeasons(n int) []seer.Season { + seasons := make([]seer.Season, n) + for i := range seasons { + seasons[i] = seer.Season{SeasonNumber: i + 1} + } + return seasons +} + +func TestSeasonPagesPreserveSelectionsAndQuota(t *testing.T) { + var cache selectionCache + cache.set("s", "owner", "show", []seer.SearchResult{{ID: 42, MediaType: "tv"}}) + quota := &seer.Quota{TV: seer.QuotaUsage{Limit: 5, Used: 2, Remaining: 3}} + cache.setQuota("s", "owner", quota) + seasons := testSeasons(51) + cache.setAvailableSeasons("s", "0", "owner", seasons) + selectPage := func(page int, numbers ...int) { + t.Helper() + if err := cache.selectSeasonPage("s", "0", "owner", page, numbers); err != nil { + t.Fatal(err) + } + } + menu := func(page int) discordgo.SelectMenu { + t.Helper() + _, _, selected, _ := cache.selection("s", "0", "owner") + return seasonPickerComponents("s", "0", seasons, quota, selected, page)[0].(discordgo.ActionsRow).Components[0].(discordgo.SelectMenu) + } + selectPage(0, 1, 2, 3) + if m := menu(1); !m.Disabled || m.MaxValues != 1 || *m.MinValues != 0 || m.Options[0].Value != "26" { + t.Fatalf("full quota menu: %#v", m) + } + if menu(0).Disabled { + t.Fatal("cannot clear previous page") + } + selectPage(0) + if menu(1).Disabled { + t.Fatal("quota not released by clearing selections") + } + selectPage(1, 26) + selectPage(2, 51) + selectPage(0, 2) + _, _, selected, _ := cache.selection("s", "0", "owner") + if !reflect.DeepEqual(selected.Numbers, []int{2, 26, 51}) { + t.Fatal(selected) + } + if err := cache.selectSeasonPage("s", "0", "owner", 0, []int{26}); err == nil { + t.Fatal("accepted a forged page selection") + } + if err := cache.selectSeasonPage("s", "0", "owner", 0, []int{2, 3}); err == nil { + t.Fatal("exceeded finite quota with restricted=false") + } + selectPage(1) + _, _, selected, _ = cache.selection("s", "0", "owner") + if !reflect.DeepEqual(selected.Numbers, []int{2, 51}) { + t.Fatal(selected) + } +} + +func TestSeasonPagePayloadLimits(t *testing.T) { + for _, n := range []int{0, 1, 25, 26, 50, 51} { + t.Run(strconv.Itoa(n), func(t *testing.T) { + seasons := testSeasons(n) + for page := 0; page < max(1, (n+24)/25); page++ { + rows := seasonPickerComponents("s", "0", seasons, &seer.Quota{}, seer.SeasonSelection{}, page) + if len(rows) > 5 { + t.Fatal("too many rows") + } + for _, row := range rows { + for _, component := range row.(discordgo.ActionsRow).Components { + if menu, ok := component.(discordgo.SelectMenu); ok { + if len(menu.Options) < 1 || len(menu.Options) > 25 || menu.MaxValues < 1 || menu.MaxValues > 25 || *menu.MinValues != 0 { + t.Fatalf("invalid Discord select: %#v", menu) + } + } + } + } + } + }) + } +} + +type requestFlowHandler struct { + Handler + calls atomic.Int32 + entered chan struct{} + release chan struct{} + err error +} + +func (f *requestFlowHandler) Request(context.Context, string, seer.SearchResult, seer.SeasonSelection) (seer.Request, error) { + f.calls.Add(1) + if f.entered != nil { + close(f.entered) + <-f.release + } + return seer.Request{ID: 42, Status: 2}, f.err +} +func componentInteraction(id string) *discordgo.InteractionCreate { + return &discordgo.InteractionCreate{Interaction: &discordgo.Interaction{Type: discordgo.InteractionMessageComponent, User: &discordgo.User{ID: "owner"}, Data: discordgo.MessageComponentInteractionData{CustomID: id}}} +} + +func TestConcurrentConfirmationPostsOnceAndPreservesSuccess(t *testing.T) { + handler := &requestFlowHandler{entered: make(chan struct{}), release: make(chan struct{})} + bot := &Bot{handler: handler, ctx: context.Background(), logger: slog.Default()} + bot.cache.set("s", "owner", "movie", []seer.SearchResult{{ID: 1, MediaType: "movie", Title: "Movie"}}) + first, duplicate := &fakeInteractionSession{}, &fakeInteractionSession{} + var wg sync.WaitGroup + wg.Go(func() { bot.handleComponent(first, componentInteraction(componentConfirm+"s:0")) }) + <-handler.entered + bot.handleComponent(duplicate, componentInteraction(componentConfirm+"s:0")) + if duplicate.edit != nil || duplicate.response == nil || duplicate.response.Data.Flags != discordgo.MessageFlagsEphemeral { + t.Fatal("duplicate touched original response") + } + close(handler.release) + wg.Wait() + bot.handleComponent(duplicate, componentInteraction(componentBack+"s")) + if handler.calls.Load() != 1 { + t.Fatal("duplicate POST") + } + if first.edit == nil || !strings.Contains(*first.edit.Content, "has been submitted") { + t.Fatalf("lost success: %#v", first.edit) + } + if duplicate.edit != nil { + t.Fatal("stale click overwrote success") + } +} + +func TestUnknownSubmissionConsumesConfirmationWithoutRetry(t *testing.T) { + handler := &requestFlowHandler{err: seer.ErrSubmissionUnknown} + bot := &Bot{handler: handler, ctx: context.Background(), logger: slog.Default()} + bot.cache.set("s", "owner", "movie", []seer.SearchResult{{ID: 1, MediaType: "movie"}}) + session := &fakeInteractionSession{} + bot.handleComponent(session, componentInteraction(componentConfirm+"s:0")) + if session.edit == nil || len(*session.edit.Components) != 0 || !strings.Contains(*session.edit.Content, "may have accepted") { + t.Fatalf("unknown response: %#v", session.edit) + } + bot.handleComponent(&fakeInteractionSession{}, componentInteraction(componentRetry+"s:0")) + if handler.calls.Load() != 1 { + t.Fatal("retried unknown submission") + } +} + +func TestSubmissionClaimIsAtomicAndCanRetryDefiniteFailure(t *testing.T) { + var cache selectionCache + cache.set("s", "owner", "movie", []seer.SearchResult{{ID: 1, MediaType: "movie"}}) + var accepted atomic.Int32 + var wg sync.WaitGroup + for range 30 { + wg.Go(func() { + if _, _, state := cache.beginSubmission("s", "0", "owner", false); state == "" { + accepted.Add(1) + } + }) + } + wg.Wait() + if accepted.Load() != 1 { + t.Fatal(accepted.Load()) + } + cache.finishSubmission("s", "owner", false) + if _, _, state := cache.beginSubmission("s", "0", "owner", false); state != "" { + t.Fatal(state) + } + cache.finishSubmission("s", "owner", true) + if state := cache.submissionState("s", "owner"); state != "submitted" { + t.Fatal(state) + } +} diff --git a/internal/seer/client.go b/internal/seer/client.go index a9e8058..3b42062 100644 --- a/internal/seer/client.go +++ b/internal/seer/client.go @@ -89,6 +89,9 @@ type QuotaUsage struct { Restricted bool `json:"restricted"` } +// Limited distinguishes a finite quota from Seerr's exhausted-quota flag. +func (q QuotaUsage) Limited() bool { return q.Limit > 0 } + type TVDetails struct { Seasons []Season `json:"seasons"` } @@ -200,54 +203,6 @@ func (c *Client) Search(ctx context.Context, query string) ([]SearchResult, erro return filtered, nil } -func (c *Client) FindUserByDiscordID(ctx context.Context, discordID string) (User, bool, error) { - discordID = strings.TrimSpace(discordID) - if discordID == "" { - return User{}, false, errors.New("discord ID is required") - } - var matched User - seenUsers := make(map[int]struct{}) - for skip := 0; ; skip += 100 { - values := url.Values{} - values.Set("take", "100") - values.Set("skip", strconv.Itoa(skip)) - var page struct { - Results []User `json:"results"` - } - if err := c.do(ctx, http.MethodGet, "/api/v1/user?"+values.Encode(), nil, &page); err != nil { - return User{}, false, err - } - for _, user := range page.Results { - if user.ID <= 0 { - continue - } - if _, duplicate := seenUsers[user.ID]; duplicate { - continue - } - seenUsers[user.ID] = struct{}{} - settings, err := c.NotificationSettings(ctx, user.ID) - if err != nil { - return User{}, false, err - } - if settings.HasDiscordID(discordID) { - if matched.ID != 0 && matched.ID != user.ID { - return User{}, false, fmt.Errorf("discord ID is linked to multiple seerr users (%d and %d)", matched.ID, user.ID) - } - matched = user - } - } - if len(page.Results) < 100 { - return matched, matched.ID != 0, nil - } - if skip > math.MaxInt-100 { - return User{}, false, errors.New("seerr user pagination overflowed") - } - if len(seenUsers) <= skip { - return User{}, false, errors.New("seerr user pagination did not advance") - } - } -} - func (c *Client) RequestMedia(ctx context.Context, userID int, mediaType string, mediaID int, seasons SeasonSelection) (Request, error) { if userID < 0 { return Request{}, errors.New("user ID must not be negative") @@ -290,12 +245,20 @@ func (c *Client) RequestMedia(ctx context.Context, userID int, mediaType string, body["seasons"] = seasons.Numbers } } - var out Request - if err := c.doAsUser(ctx, http.MethodPost, "/api/v1/request", body, &out, userID); err != nil { + if err := ctx.Err(); err != nil { return Request{}, err } - if out.ID <= 0 { - return Request{}, errors.New("seerr create-request response is missing a valid request ID") + var out Request + status, err := c.doResponse(ctx, http.MethodPost, "/api/v1/request", body, &out, userID) + switch { + case status == http.StatusAccepted || status == http.StatusConflict: + return Request{}, &submissionError{message: "No new request was created. The title or selected seasons are already requested or available. Check Seerr or choose different seasons."} + case status == http.StatusTooManyRequests: + return Request{}, err + case status >= 400 && status < 500 && status != http.StatusRequestTimeout: + return Request{}, &submissionError{message: "Seerr rejected this request. Check your permissions and quota, then start a new search.", cause: err} + case err != nil || out.ID <= 0: + return Request{}, &submissionError{message: ErrSubmissionUnknown.message, cause: err} } return out, nil } @@ -433,6 +396,9 @@ func (c *Client) NotificationSettings(ctx context.Context, userID int) (Notifica } var out NotificationSettings err := c.do(ctx, http.MethodGet, "/api/v1/user/"+strconv.Itoa(userID)+"/settings/notifications", nil, &out) + if err == nil && out.DiscordIDs == nil { + err = errors.New("seerr notification settings are missing discordIds") + } return out, err } @@ -440,27 +406,49 @@ func (c *Client) UserQuota(ctx context.Context, userID int) (Quota, error) { if userID <= 0 { return Quota{}, errors.New("user ID must be positive") } - var out Quota - err := c.do(ctx, http.MethodGet, "/api/v1/user/"+strconv.Itoa(userID)+"/quota", nil, &out) - return out, err + type usage struct { + Days int `json:"days"` + Limit *int `json:"limit"` + Used *int `json:"used"` + Remaining *int `json:"remaining"` + Restricted bool `json:"restricted"` + } + var out struct { + Movie *usage `json:"movie"` + TV *usage `json:"tv"` + } + if err := c.do(ctx, http.MethodGet, "/api/v1/user/"+strconv.Itoa(userID)+"/quota", nil, &out); err != nil { + return Quota{}, err + } + valid := func(u *usage) bool { + return u != nil && u.Limit != nil && u.Used != nil && u.Remaining != nil && *u.Limit >= 0 && *u.Used >= 0 && *u.Remaining >= 0 + } + if !valid(out.Movie) || !valid(out.TV) { + return Quota{}, errors.New("seerr quota response is incomplete") + } + convert := func(u *usage) QuotaUsage { + return QuotaUsage{Days: u.Days, Limit: *u.Limit, Used: *u.Used, Remaining: *u.Remaining, Restricted: u.Restricted} + } + return Quota{Movie: convert(out.Movie), TV: convert(out.TV)}, nil } -func (c *Client) do(ctx context.Context, method, path string, body any, out any) error { - return c.doAsUser(ctx, method, path, body, out, 0) +func (c *Client) do(ctx context.Context, method, path string, body, out any) error { + _, err := c.doResponse(ctx, method, path, body, out, 0) + return err } -func (c *Client) doAsUser(ctx context.Context, method, path string, body any, out any, userID int) error { +func (c *Client) doResponse(ctx context.Context, method, path string, body any, out any, userID int) (int, error) { endpoint := safeEndpointPath(path) req, err := c.newRequest(ctx, method, path, body) if err != nil { - return err + return 0, err } if userID > 0 { req.Header.Set("X-Api-User", strconv.Itoa(userID)) } resp, err := c.httpClient.Do(req) if err != nil { - return err + return 0, err } defer resp.Body.Close() if resp.StatusCode < 200 || resp.StatusCode >= 300 { @@ -468,24 +456,24 @@ func (c *Client) doAsUser(ctx context.Context, method, path string, body any, ou // Response bodies are deliberately excluded from errors. Upstream errors can // contain reflected request headers, credentials, or other sensitive data, // and these errors are subsequently written to application logs. - return &responseError{method: method, path: endpoint, statusCode: resp.StatusCode} + return resp.StatusCode, &responseError{method: method, path: endpoint, statusCode: resp.StatusCode} } data, err := readResponseBody(resp.Body) if err != nil { - return err + return resp.StatusCode, err } if out == nil || len(data) == 0 { - return nil + return resp.StatusCode, nil } decoder := json.NewDecoder(bytes.NewReader(data)) decoder.UseNumber() if err := decoder.Decode(out); err != nil { - return fmt.Errorf("decode seerr %s %s response: %w", method, endpoint, err) + return resp.StatusCode, fmt.Errorf("decode seerr %s %s response: %w", method, endpoint, err) } if err := ensureJSONEOF(decoder); err != nil { - return fmt.Errorf("decode seerr %s %s response: %w", method, endpoint, err) + return resp.StatusCode, fmt.Errorf("decode seerr %s %s response: %w", method, endpoint, err) } - return nil + return resp.StatusCode, nil } func drainResponseBody(body io.Reader) { diff --git a/internal/seer/reliability_test.go b/internal/seer/reliability_test.go new file mode 100644 index 0000000..3cf7866 --- /dev/null +++ b/internal/seer/reliability_test.go @@ -0,0 +1,222 @@ +package seer + +import ( + "context" + "errors" + "fmt" + "io" + "net/http" + "strings" + "sync/atomic" + "testing" + "time" +) + +func TestCreateRequestOutcomeClassification(t *testing.T) { + for _, tt := range []struct { + name string + status int + body string + transportErr error + success bool + }{ + {name: "created", status: 201, body: `{"id":42}`, success: true}, + {name: "no eligible seasons", status: 202, body: `{}`}, + {name: "duplicate", status: 409, body: `{}`}, + {name: "forbidden", status: 403, body: `{}`}, + {name: "server failure", status: 500, body: `{}`}, + {name: "empty created", status: 201}, + {name: "missing ID", status: 201, body: `{}`}, + {name: "malformed created", status: 201, body: `{"id":`}, + {name: "oversized", status: 201, body: strings.Repeat("x", maxResponseBodyBytes+1)}, + {name: "lost acknowledgement", transportErr: io.ErrUnexpectedEOF}, + } { + t.Run(tt.name, func(t *testing.T) { + var calls int + client := newTestClient(t) + client.httpClient.Transport = roundTripFunc(func(r *http.Request) (*http.Response, error) { + calls++ + io.Copy(io.Discard, r.Body) + if tt.transportErr != nil { + return nil, tt.transportErr + } + return &http.Response{StatusCode: tt.status, Body: io.NopCloser(strings.NewReader(tt.body))}, nil + }) + req, err := client.RequestMedia(context.Background(), 7, "tv", 42, SeasonSelection{Numbers: []int{1}}) + if calls != 1 { + t.Fatal(calls) + } + if tt.success { + if err != nil || req.ID != 42 { + t.Fatalf("%#v %v", req, err) + } + return + } + var outcome interface{ SafeToRetry() bool } + if err == nil || !errors.As(err, &outcome) || outcome.SafeToRetry() { + t.Fatalf("error must suppress POST replay: %v", err) + } + if req.ID != 0 { + t.Fatal("unconfirmed ID exposed") + } + if tt.status == 202 && strings.Contains(err.Error(), "may have accepted") { + t.Fatal("known no-op reported as uncertain") + } + }) + } +} + +func TestRequestCancellationBeforeAndAfterDispatch(t *testing.T) { + client := newTestClient(t) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + client.httpClient.Transport = roundTripFunc(func(r *http.Request) (*http.Response, error) { + t.Fatal("pre-cancelled request dispatched") + return nil, nil + }) + if _, err := client.RequestMedia(ctx, 7, "movie", 1, SeasonSelection{}); !errors.Is(err, context.Canceled) { + t.Fatal(err) + } + ctx, cancel = context.WithCancel(context.Background()) + client.httpClient.Transport = roundTripFunc(func(r *http.Request) (*http.Response, error) { cancel(); return nil, r.Context().Err() }) + _, err := client.RequestMedia(ctx, 7, "movie", 1, SeasonSelection{}) + var outcome interface{ SafeToRetry() bool } + if !errors.As(err, &outcome) || outcome.SafeToRetry() { + t.Fatal(err) + } +} + +func TestLinkedUserScanRejectsIncompleteResponses(t *testing.T) { + for _, body := range []string{"", `null`, `{}`, `{"discordIds":null}`} { + t.Run(body, func(t *testing.T) { + client := newTestClient(t) + client.httpClient.Transport = roundTripFunc(func(r *http.Request) (*http.Response, error) { + if r.URL.Path == "/api/v1/user" { + return jsonResponse(t, map[string]any{"results": []User{{ID: 1}, {ID: 2}}}), nil + } + if strings.Contains(r.URL.Path, "/1/") { + return jsonResponse(t, map[string]any{"discordIds": []string{"match"}}), nil + } + return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(body))}, nil + }) + if user, ok, err := client.FindUserByDiscordID(context.Background(), "match"); err == nil || ok || user.ID != 0 { + t.Fatalf("partial lookup authorized: %#v %t %v", user, ok, err) + } + }) + } +} + +func TestUserLookupConcurrencyIsBoundedAndJoined(t *testing.T) { + client := newTestClient(t) + var active, peak atomic.Int32 + entered := make(chan struct{}, userLookupConcurrency) + release := make(chan struct{}) + client.httpClient.Transport = roundTripFunc(func(r *http.Request) (*http.Response, error) { + if r.URL.Path == "/api/v1/user" { + return jsonResponse(t, map[string]any{"results": []User{{ID: 1}, {ID: 2}, {ID: 3}, {ID: 4}, {ID: 5}, {ID: 6}}}), nil + } + current := active.Add(1) + defer active.Add(-1) + for old := peak.Load(); current > old && !peak.CompareAndSwap(old, current); old = peak.Load() { + } + select { + case entered <- struct{}{}: + default: + } + select { + case <-release: + case <-r.Context().Done(): + return nil, r.Context().Err() + } + return jsonResponse(t, map[string]any{"discordIds": []string{}}), nil + }) + done := make(chan error, 1) + go func() { _, _, err := client.FindUserByDiscordID(context.Background(), "match"); done <- err }() + for range userLookupConcurrency { + select { + case <-entered: + case <-time.After(time.Second): + t.Fatal("lookup was serialized") + } + } + close(release) + if err := <-done; err != nil { + t.Fatal(err) + } + if peak.Load() != userLookupConcurrency || active.Load() != 0 { + t.Fatalf("peak=%d active=%d", peak.Load(), active.Load()) + } +} + +func TestLookupDetectsAmbiguityAcrossPages(t *testing.T) { + client := newTestClient(t) + client.httpClient.Transport = roundTripFunc(func(r *http.Request) (*http.Response, error) { + if r.URL.Path == "/api/v1/user" { + users := []User{{ID: 101}} + if r.URL.Query().Get("skip") == "0" { + users = make([]User, 100) + for i := range users { + users[i].ID = i + 1 + } + } + return jsonResponse(t, map[string]any{"results": users}), nil + } + ids := []string{} + if r.URL.Path == "/api/v1/user/1/settings/notifications" || r.URL.Path == "/api/v1/user/101/settings/notifications" { + ids = append(ids, "match") + } + return jsonResponse(t, map[string]any{"discordIds": ids}), nil + }) + if _, ok, err := client.FindUserByDiscordID(context.Background(), "match"); err == nil || ok { + t.Fatalf("ambiguous scan: %t %v", ok, err) + } +} + +func TestQuotaRequiresCompleteUsage(t *testing.T) { + for _, body := range []string{"", `{}`, `{"movie":{},"tv":{}}`} { + t.Run(fmt.Sprint(len(body)), func(t *testing.T) { + client := newTestClient(t) + client.httpClient.Transport = roundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(body))}, nil + }) + if _, err := client.UserQuota(context.Background(), 7); err == nil { + t.Fatal("incomplete quota appeared unlimited") + } + }) + } +} + +func BenchmarkLinkedUserNotificationScan(b *testing.B) { + users := make([]User, 20) + for i := range users { + users[i].ID = i + 1 + } + for _, parallel := range []bool{false, true} { + b.Run(fmt.Sprintf("parallel=%t", parallel), func(b *testing.B) { + client, _ := New(Config{BaseURL: "http://seerr.test", APIKey: "synthetic", Timeout: time.Second}) + client.httpClient.Transport = roundTripFunc(func(r *http.Request) (*http.Response, error) { + timer := time.NewTimer(time.Millisecond) + defer timer.Stop() + select { + case <-timer.C: + case <-r.Context().Done(): + return nil, r.Context().Err() + } + return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"discordIds":[]}`))}, nil + }) + for b.Loop() { + if parallel { + if _, err := client.matchDiscordUsers(context.Background(), users, "match"); err != nil { + b.Fatal(err) + } + } else { + for _, user := range users { + if _, err := client.NotificationSettings(context.Background(), user.ID); err != nil { + b.Fatal(err) + } + } + } + } + }) + } +} diff --git a/internal/seer/submission.go b/internal/seer/submission.go new file mode 100644 index 0000000..5c339e2 --- /dev/null +++ b/internal/seer/submission.go @@ -0,0 +1,27 @@ +package seer + +import ( + "errors" + "net/http" +) + +// A dispatched POST has no idempotency key. Losing its acknowledgement does not +// prove rejection; the caller must not offer an automatic or one-click retry. +type submissionError struct { + message string + cause error +} + +func (e *submissionError) Error() string { return e.message } +func (e *submissionError) UserMessage() string { return e.message } +func (e *submissionError) SafeToRetry() bool { return false } +func (e *submissionError) Unwrap() error { return e.cause } + +var ErrSubmissionUnknown = &submissionError{message: "Seerr may have accepted your request, but its response could not be confirmed. Check `/requests` or Seerr before starting another request."} + +// IsDefiniteRejection is only used to discard a saved mutation intent when +// Seerr explicitly rejects it. Transport and server failures remain uncertain. +func IsDefiniteRejection(err error) bool { + var response *responseError + return errors.As(err, &response) && response.statusCode >= 400 && response.statusCode < 500 && response.statusCode != http.StatusRequestTimeout +} diff --git a/internal/seer/user_lookup.go b/internal/seer/user_lookup.go new file mode 100644 index 0000000..65f9172 --- /dev/null +++ b/internal/seer/user_lookup.go @@ -0,0 +1,112 @@ +package seer + +import ( + "context" + "errors" + "fmt" + "math" + "net/http" + "net/url" + "strconv" + "strings" + "sync" +) + +const userLookupConcurrency = 4 + +func (c *Client) FindUserByDiscordID(ctx context.Context, discordID string) (User, bool, error) { + discordID = strings.TrimSpace(discordID) + if discordID == "" { + return User{}, false, errors.New("discord ID is required") + } + var matched User + seen := make(map[int]bool) + for skip := 0; ; skip += 100 { + values := url.Values{"take": {"100"}, "skip": {strconv.Itoa(skip)}} + var page struct { + Results []User `json:"results"` + } + if err := c.do(ctx, http.MethodGet, "/api/v1/user?"+values.Encode(), nil, &page); err != nil { + return User{}, false, err + } + if page.Results == nil { + return User{}, false, errors.New("seerr user page is missing results") + } + users := make([]User, 0, len(page.Results)) + for _, user := range page.Results { + if user.ID <= 0 { + return User{}, false, errors.New("seerr user page contains an invalid user ID") + } + if !seen[user.ID] { + seen[user.ID] = true + users = append(users, user) + } + } + matches, err := c.matchDiscordUsers(ctx, users, discordID) + if err != nil { + return User{}, false, err + } + for _, user := range matches { + if matched.ID != 0 && matched.ID != user.ID { + return User{}, false, errors.New("discord ID is linked to multiple seerr users") + } + matched = user + } + if len(page.Results) < 100 { + return matched, matched.ID != 0, nil + } + if skip > math.MaxInt-100 { + return User{}, false, errors.New("seerr user pagination overflowed") + } + if len(users) == 0 { + return User{}, false, errors.New("seerr user pagination did not advance") + } + } +} + +// Each worker owns one result slot at a time. Join every worker before returning, +// including errors, so an incomplete scan can never authorize a partial match. +func (c *Client) matchDiscordUsers(ctx context.Context, users []User, discordID string) ([]User, error) { + lookupCtx, cancel := context.WithCancel(ctx) + defer cancel() + jobs := make(chan int, len(users)) + for i := range users { + jobs <- i + } + close(jobs) + matched := make([]bool, len(users)) + failures := make([]error, len(users)) + var wg sync.WaitGroup + for range min(userLookupConcurrency, len(users)) { + wg.Go(func() { + for i := range jobs { + if lookupCtx.Err() != nil { + return + } + settings, err := c.NotificationSettings(lookupCtx, users[i].ID) + if err != nil { + failures[i] = fmt.Errorf("verify seerr account link: %w", err) + cancel() + return + } + matched[i] = settings.HasDiscordID(discordID) + } + }) + } + wg.Wait() + for _, err := range failures { + if err != nil && !errors.Is(err, context.Canceled) { + return nil, err + } + } + if err := lookupCtx.Err(); err != nil { + return nil, err + } + var out []User + for i, yes := range matched { + if yes { + out = append(out, users[i]) + } + } + return out, nil +} diff --git a/internal/storage/approval_cleanup.go b/internal/storage/approval_cleanup.go new file mode 100644 index 0000000..d2e8ceb --- /dev/null +++ b/internal/storage/approval_cleanup.go @@ -0,0 +1,47 @@ +package storage + +import ( + "context" + "errors" + "time" +) + +// QueueUntrackedApprovalCleanup rechecks tracking after an uncertain database +// acknowledgement. A card that was saved successfully must never be deleted. +func (s *Store) QueueUntrackedApprovalCleanup(ctx context.Context, message ApprovalMessage) (bool, error) { + if message.ChannelID == "" || message.MessageID == "" { + return false, errors.New("channel and message are required") + } + result, err := s.db.ExecContext(ctx, `INSERT INTO approval_cleanup(channel_id,message_id) + SELECT ?,? WHERE NOT EXISTS(SELECT 1 FROM approval_messages WHERE channel_id=? AND message_id=?) + ON CONFLICT(channel_id,message_id) DO UPDATE SET message_id=excluded.message_id`, message.ChannelID, message.MessageID, message.ChannelID, message.MessageID) + if err != nil { + return false, err + } + n, err := result.RowsAffected() + return n == 1, err +} +func (s *Store) DueApprovalCleanup(ctx context.Context, now time.Time) ([]ApprovalMessage, error) { + rows, err := s.db.QueryContext(ctx, `SELECT channel_id,message_id,attempts FROM approval_cleanup WHERE next_attempt_at<=? ORDER BY next_attempt_at,channel_id,message_id LIMIT 100`, now.UnixMilli()) + if err != nil { + return nil, err + } + defer rows.Close() + var messages []ApprovalMessage + for rows.Next() { + var m ApprovalMessage + if err := rows.Scan(&m.ChannelID, &m.MessageID, &m.Attempts); err != nil { + return nil, err + } + messages = append(messages, m) + } + return messages, rows.Err() +} +func (s *Store) CompleteApprovalCleanup(ctx context.Context, message ApprovalMessage) error { + _, err := s.db.ExecContext(ctx, `DELETE FROM approval_cleanup WHERE channel_id=? AND message_id=?`, message.ChannelID, message.MessageID) + return err +} +func (s *Store) RetryApprovalCleanup(ctx context.Context, message ApprovalMessage, retryAt time.Time) error { + _, err := s.db.ExecContext(ctx, `UPDATE approval_cleanup SET attempts=attempts+1,next_attempt_at=? WHERE channel_id=? AND message_id=?`, retryAt.UnixMilli(), message.ChannelID, message.MessageID) + return err +} diff --git a/internal/storage/approval_delivery.go b/internal/storage/approval_delivery.go new file mode 100644 index 0000000..95904fe --- /dev/null +++ b/internal/storage/approval_delivery.go @@ -0,0 +1,77 @@ +package storage + +import ( + "context" + "crypto/rand" + "database/sql" + "errors" + "strings" + "time" +) + +func (s *Store) NeedsApprovalMessage(ctx context.Context, requestID int) (bool, error) { + if requestID <= 0 { + return false, errors.New("request_id must be positive") + } + var needed bool + err := s.db.QueryRowContext(ctx, `SELECT EXISTS ( + SELECT 1 FROM approval_settings settings LEFT JOIN approval_messages messages + ON messages.guild_id=settings.guild_id AND messages.request_id=? + WHERE settings.enabled=1 AND (messages.request_id IS NULL OR + (messages.message_id='' AND (messages.channel_id!=settings.channel_id OR + (messages.lease_until<=? AND messages.next_attempt_at<=?)))))`, requestID, time.Now().UnixMilli(), time.Now().UnixMilli()).Scan(&needed) + return needed, err +} + +// The token fences late senders after a crash, lease expiry or channel change. +func (s *Store) ClaimApprovalMessage(ctx context.Context, message ApprovalMessage, now time.Time) (ApprovalMessage, bool, error) { + message.GuildID = strings.TrimSpace(message.GuildID) + message.ChannelID = strings.TrimSpace(message.ChannelID) + if message.RequestID <= 0 || message.GuildID == "" || message.ChannelID == "" { + return ApprovalMessage{}, false, errors.New("request, guild and channel are required") + } + message.ClaimToken = rand.Text() + message.LeaseUntil = now.Add(2 * time.Minute) + err := s.db.QueryRowContext(ctx, `INSERT INTO approval_messages(request_id,guild_id,channel_id,claim_token,lease_until,attempts,next_attempt_at) + SELECT ?,?,?,?, ?,1,? FROM approval_settings WHERE guild_id=? AND channel_id=? AND enabled=1 + ON CONFLICT(request_id,guild_id) DO UPDATE SET channel_id=excluded.channel_id,claim_token=excluded.claim_token, + lease_until=excluded.lease_until,attempts=approval_messages.attempts+1,next_attempt_at=excluded.next_attempt_at + WHERE approval_messages.message_id='' AND (approval_messages.channel_id!=excluded.channel_id OR + (approval_messages.lease_until<=? AND approval_messages.next_attempt_at<=?)) + RETURNING attempts`, message.RequestID, message.GuildID, message.ChannelID, message.ClaimToken, message.LeaseUntil.UnixMilli(), message.LeaseUntil.UnixMilli(), message.GuildID, message.ChannelID, now.UnixMilli(), now.UnixMilli()).Scan(&message.Attempts) + if errors.Is(err, sql.ErrNoRows) { + return ApprovalMessage{}, false, nil + } + return message, err == nil, err +} + +func (s *Store) FinishApprovalMessage(ctx context.Context, message ApprovalMessage) error { + if message.MessageID == "" || message.ClaimToken == "" { + return errors.New("message ID and claim token are required") + } + result, err := s.db.ExecContext(ctx, `UPDATE approval_messages SET message_id=?,lease_until=0,next_attempt_at=0,attempts=0 + WHERE request_id=? AND guild_id=? AND channel_id=? AND claim_token=? AND message_id='' + AND EXISTS(SELECT 1 FROM approval_settings WHERE guild_id=? AND channel_id=? AND enabled=1)`, message.MessageID, message.RequestID, message.GuildID, message.ChannelID, message.ClaimToken, message.GuildID, message.ChannelID) + if err == nil { + n, e := result.RowsAffected() + if e == nil && n == 1 { + return nil + } + err = errors.New("approval message claim is no longer current") + } + // A successful write can lose its acknowledgement. Confirm it before the + // caller considers deleting the newly sent message. + var saved string + if readErr := s.db.QueryRowContext(ctx, `SELECT message_id FROM approval_messages WHERE request_id=? AND guild_id=? AND claim_token=?`, message.RequestID, message.GuildID, message.ClaimToken).Scan(&saved); readErr == nil && saved == message.MessageID { + return nil + } + return err +} + +func (s *Store) RetryApprovalMessage(ctx context.Context, message ApprovalMessage, retryAt time.Time) error { + _, err := s.db.ExecContext(ctx, `UPDATE approval_messages SET claim_token='',lease_until=0,next_attempt_at=?, + attempts=attempts+CASE WHEN message_id!='' THEN 1 ELSE 0 END + WHERE request_id=? AND guild_id=? AND channel_id=? AND + ((claim_token=? AND message_id='') OR (message_id=? AND message_id!=''))`, retryAt.UnixMilli(), message.RequestID, message.GuildID, message.ChannelID, message.ClaimToken, message.MessageID) + return err +} diff --git a/internal/storage/decision_intents.go b/internal/storage/decision_intents.go new file mode 100644 index 0000000..2651dae --- /dev/null +++ b/internal/storage/decision_intents.go @@ -0,0 +1,65 @@ +package storage + +import ( + "context" + "database/sql" + "errors" + "time" +) + +type DecisionIntent struct { + RequestID int + Status, Actor, Reason string + Title, URL, PosterURL string + CreatedAt time.Time +} + +func (s *Store) SaveDecisionIntent(ctx context.Context, intent DecisionIntent) error { + if intent.RequestID <= 0 || (intent.Status != "Approved" && intent.Status != "Declined") { + return errors.New("a valid decision intent is required") + } + _, err := s.db.ExecContext(ctx, `INSERT INTO decision_intents(request_id,status,actor,reason,title,url,poster_url,created_at) VALUES(?,?,?,?,?,?,?,?) + ON CONFLICT(request_id) DO UPDATE SET status=excluded.status,actor=excluded.actor,reason=excluded.reason,title=excluded.title,url=excluded.url,poster_url=excluded.poster_url,created_at=excluded.created_at,next_attempt_at=0`, intent.RequestID, intent.Status, intent.Actor, intent.Reason, intent.Title, intent.URL, intent.PosterURL, formatTime(intent.CreatedAt)) + return err +} +func scanIntent(row interface{ Scan(...any) error }) (DecisionIntent, error) { + var intent DecisionIntent + var at string + err := row.Scan(&intent.RequestID, &intent.Status, &intent.Actor, &intent.Reason, &intent.Title, &intent.URL, &intent.PosterURL, &at) + if err != nil { + return intent, err + } + intent.CreatedAt, err = parseTime(at) + return intent, err +} +func (s *Store) DecisionIntent(ctx context.Context, requestID int) (DecisionIntent, bool, error) { + intent, err := scanIntent(s.db.QueryRowContext(ctx, `SELECT request_id,status,actor,reason,title,url,poster_url,created_at FROM decision_intents WHERE request_id=?`, requestID)) + if errors.Is(err, sql.ErrNoRows) { + return DecisionIntent{}, false, nil + } + return intent, err == nil, err +} +func (s *Store) DueDecisionIntents(ctx context.Context, now time.Time) ([]DecisionIntent, error) { + rows, err := s.db.QueryContext(ctx, `SELECT request_id,status,actor,reason,title,url,poster_url,created_at FROM decision_intents WHERE next_attempt_at<=? ORDER BY next_attempt_at,request_id LIMIT 25`, now.UnixMilli()) + if err != nil { + return nil, err + } + defer rows.Close() + var intents []DecisionIntent + for rows.Next() { + intent, err := scanIntent(rows) + if err != nil { + return nil, err + } + intents = append(intents, intent) + } + return intents, rows.Err() +} +func (s *Store) ClearDecisionIntent(ctx context.Context, requestID int) error { + _, err := s.db.ExecContext(ctx, `DELETE FROM decision_intents WHERE request_id=?`, requestID) + return err +} +func (s *Store) RetryDecisionIntent(ctx context.Context, intent DecisionIntent, retryAt time.Time) error { + _, err := s.db.ExecContext(ctx, `UPDATE decision_intents SET next_attempt_at=? WHERE request_id=? AND created_at=?`, retryAt.UnixMilli(), intent.RequestID, formatTime(intent.CreatedAt)) + return err +} diff --git a/internal/storage/decision_jobs.go b/internal/storage/decision_jobs.go new file mode 100644 index 0000000..452906a --- /dev/null +++ b/internal/storage/decision_jobs.go @@ -0,0 +1,147 @@ +package storage + +import ( + "context" + "database/sql" + "errors" + "time" +) + +// ApprovalDecision is the durable result and presentation of one Seerr decision. +// Discord cards and notification delivery may fail independently of this row. +type ApprovalDecision struct { + RequestID int + Status string + RequesterID int + MediaID int + MediaType string + Title string + URL string + PosterURL string + Actor string + Reason string + DecidedAt time.Time +} + +type DecisionJob struct { + ApprovalDecision + Attempts int +} + +const decisionColumns = `request_id,status,requester_id,media_id,media_type,title,url,poster_url,actor,reason,decided_at,attempts` + +func scanDecision(row interface{ Scan(...any) error }) (DecisionJob, error) { + var job DecisionJob + var at string + err := row.Scan(&job.RequestID, &job.Status, &job.RequesterID, &job.MediaID, &job.MediaType, &job.Title, &job.URL, &job.PosterURL, &job.Actor, &job.Reason, &at, &job.Attempts) + if err != nil { + return job, err + } + job.DecidedAt, err = parseTime(at) + return job, err +} + +func (s *Store) RecordApprovalDecision(ctx context.Context, d ApprovalDecision) (ApprovalDecision, error) { + if d.RequestID <= 0 || (d.Status != "Approved" && d.Status != "Declined") { + return ApprovalDecision{}, errors.New("a valid approval decision is required") + } + if d.DecidedAt.IsZero() { + d.DecidedAt = time.Now().UTC() + } + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return ApprovalDecision{}, err + } + defer tx.Rollback() + // Recover attribution saved before an acknowledged or uncertain Seerr write. + var status, actor, reason, title, url, poster string + intentErr := tx.QueryRowContext(ctx, `SELECT status,actor,reason,title,url,poster_url FROM decision_intents WHERE request_id=?`, d.RequestID).Scan(&status, &actor, &reason, &title, &url, &poster) + if intentErr != nil && !errors.Is(intentErr, sql.ErrNoRows) { + return ApprovalDecision{}, intentErr + } + if intentErr == nil && status == d.Status { + d.Actor = actor + d.Reason = reason + if title != "" { + d.Title = title + } + if url != "" { + d.URL = url + } + if poster != "" { + d.PosterURL = poster + } + } + job, err := scanDecision(tx.QueryRowContext(ctx, `INSERT INTO decision_jobs(request_id,status,requester_id,media_id,media_type,title,url,poster_url,actor,reason,decided_at) + VALUES(?,?,?,?,?,?,?,?,?,?,?) ON CONFLICT(request_id,status) DO UPDATE SET + requester_id=CASE WHEN decision_jobs.requester_id=0 THEN excluded.requester_id ELSE decision_jobs.requester_id END + RETURNING `+decisionColumns, d.RequestID, d.Status, d.RequesterID, d.MediaID, d.MediaType, d.Title, d.URL, d.PosterURL, d.Actor, d.Reason, formatTime(d.DecidedAt))) + if err != nil { + return ApprovalDecision{}, err + } + canonical := job.ApprovalDecision + if _, err = tx.ExecContext(ctx, `UPDATE approval_messages SET decided_at=CASE WHEN status!=? THEN NULL ELSE decided_at END,status=?,reason=? WHERE request_id=?`, canonical.Status, canonical.Status, canonical.Reason, canonical.RequestID); err != nil { + return ApprovalDecision{}, err + } + if _, err := tx.ExecContext(ctx, `DELETE FROM decision_intents WHERE request_id=?`, d.RequestID); err != nil { + return ApprovalDecision{}, err + } + return canonical, tx.Commit() +} + +func (s *Store) DueDecisionJobs(ctx context.Context, now time.Time, limit int) ([]DecisionJob, error) { + if limit <= 0 || limit > 100 { + limit = 25 + } + rows, err := s.db.QueryContext(ctx, `SELECT `+decisionColumns+` FROM decision_jobs WHERE completed_at IS NULL AND next_attempt_at<=? ORDER BY next_attempt_at,request_id,status LIMIT ?`, now.UnixMilli(), limit) + if err != nil { + return nil, err + } + defer rows.Close() + var jobs []DecisionJob + for rows.Next() { + job, err := scanDecision(rows) + if err != nil { + return nil, err + } + jobs = append(jobs, job) + } + return jobs, rows.Err() +} + +func (s *Store) RetryDecisionJob(ctx context.Context, job DecisionJob, retryAt time.Time) error { + _, err := s.db.ExecContext(ctx, `UPDATE decision_jobs SET attempts=attempts+1,next_attempt_at=? WHERE request_id=? AND status=? AND completed_at IS NULL`, retryAt.UnixMilli(), job.RequestID, job.Status) + return err +} +func (s *Store) CompleteDecisionJob(ctx context.Context, d ApprovalDecision, at time.Time) error { + _, err := s.db.ExecContext(ctx, `UPDATE decision_jobs SET completed_at=? WHERE request_id=? AND status=? AND completed_at IS NULL`, formatTime(at), d.RequestID, d.Status) + return err +} +func (s *Store) DecisionNotificationHandled(ctx context.Context, requestID int, discordID, status string) (bool, error) { + var handled bool + err := s.db.QueryRowContext(ctx, `SELECT EXISTS(SELECT 1 FROM decision_notifications WHERE request_id=? AND discord_id=? AND status=?)`, requestID, discordID, status).Scan(&handled) + return handled, err +} +func (s *Store) RecordDecisionNotification(ctx context.Context, requestID int, discordID, status string) error { + _, err := s.db.ExecContext(ctx, `INSERT INTO decision_notifications(request_id,discord_id,status) VALUES(?,?,?) ON CONFLICT DO NOTHING`, requestID, discordID, status) + return err +} + +func (p NotificationPreferences) DecisionEnabled(status string) bool { + switch status { + case "Approved": + return p.Approved + case "Declined": + return p.Declined + default: + return false + } +} + +func (s *Store) ApprovalDecision(ctx context.Context, requestID int, status string) (ApprovalDecision, bool, error) { + job, err := scanDecision(s.db.QueryRowContext(ctx, `SELECT `+decisionColumns+` FROM decision_jobs WHERE request_id=? AND status=?`, requestID, status)) + if errors.Is(err, sql.ErrNoRows) { + return ApprovalDecision{}, false, nil + } + return job.ApprovalDecision, err == nil, err +} diff --git a/internal/storage/delivery_test.go b/internal/storage/delivery_test.go new file mode 100644 index 0000000..147dff8 --- /dev/null +++ b/internal/storage/delivery_test.go @@ -0,0 +1,335 @@ +package storage + +import ( + "context" + "database/sql" + "fmt" + "path/filepath" + "testing" + "time" +) + +func TestApprovalLeaseRecoveryAndDestinationFencing(t *testing.T) { + ctx := context.Background() + path := filepath.Join(t.TempDir(), "state.db") + store, err := Open(path) + if err != nil { + t.Fatal(err) + } + settings := ApprovalSettings{GuildID: "guild", ChannelID: "old", Enabled: true} + if err := store.SetApprovalSettings(ctx, settings); err != nil { + t.Fatal(err) + } + candidate := ApprovalMessage{RequestID: 42, GuildID: "guild", ChannelID: "old"} + now := time.Now() + first, ok, err := store.ClaimApprovalMessage(ctx, candidate, now) + if err != nil || !ok { + t.Fatalf("first claim: %t %v", ok, err) + } + store.Close() + store, err = Open(path) + if err != nil { + t.Fatal(err) + } + defer store.Close() + if _, ok, err := store.ClaimApprovalMessage(ctx, candidate, now.Add(time.Minute)); err != nil || ok { + t.Fatalf("active lease reclaimed: %t %v", ok, err) + } + second, ok, err := store.ClaimApprovalMessage(ctx, candidate, now.Add(3*time.Minute)) + if err != nil || !ok || second.ClaimToken == first.ClaimToken || second.Attempts != 2 { + t.Fatalf("recovery: %#v %t %v", second, ok, err) + } + first.MessageID = "late" + if err := store.FinishApprovalMessage(ctx, first); err == nil { + t.Fatal("late sender took over claim") + } + if err := store.RetryApprovalMessage(ctx, first, now); err != nil { + t.Fatal(err) + } + second.MessageID = "current" + if err := store.FinishApprovalMessage(ctx, second); err != nil { + t.Fatal(err) + } + if err := store.FinishApprovalMessage(ctx, second); err != nil { + t.Fatal("lost-ack confirmation was not idempotent", err) + } + candidate.RequestID = 43 + old, ok, err := store.ClaimApprovalMessage(ctx, candidate, now) + if err != nil || !ok { + t.Fatal(err) + } + settings.ChannelID = "new" + if err := store.SetApprovalSettings(ctx, settings); err != nil { + t.Fatal(err) + } + old.MessageID = "old-destination" + if err := store.FinishApprovalMessage(ctx, old); err == nil { + t.Fatal("finished after destination changed") + } + if _, ok, err := store.ClaimApprovalMessage(ctx, candidate, now); err != nil || ok { + t.Fatal("retried old destination", err) + } + candidate.ChannelID = "new" + claim, ok, err := store.ClaimApprovalMessage(ctx, candidate, now) + if err != nil || !ok { + t.Fatal("current destination not claimable", err) + } + settings.Enabled = false + if err := store.SetApprovalSettings(ctx, settings); err != nil { + t.Fatal(err) + } + claim.MessageID = "disabled" + if err := store.FinishApprovalMessage(ctx, claim); err == nil { + t.Fatal("finished disabled delivery") + } + if _, ok, err := store.ClaimApprovalMessage(ctx, candidate, now.Add(time.Hour)); err != nil || ok { + t.Fatal("disabled destination claimed", err) + } +} + +func TestApprovalRetryPersistsAndDoesNotBlockOtherCleanup(t *testing.T) { + ctx := context.Background() + store, err := Open(filepath.Join(t.TempDir(), "state.db")) + if err != nil { + t.Fatal(err) + } + defer store.Close() + store.SetApprovalSettings(ctx, ApprovalSettings{GuildID: "g", ChannelID: "c", Enabled: true}) + claim, _, err := store.ClaimApprovalMessage(ctx, ApprovalMessage{RequestID: 1, GuildID: "g", ChannelID: "c"}, time.Now()) + if err != nil { + t.Fatal(err) + } + future := time.Now().Add(time.Hour) + if err := store.RetryApprovalMessage(ctx, claim, future); err != nil { + t.Fatal(err) + } + if needed, err := store.NeedsApprovalMessage(ctx, 1); err != nil || needed { + t.Fatal("backoff ignored", err) + } + if _, ok, err := store.ClaimApprovalMessage(ctx, claim, time.Now()); err != nil || ok { + t.Fatal("claim ignored backoff", err) + } + claim, ok, err := store.ClaimApprovalMessage(ctx, claim, future.Add(time.Second)) + if err != nil || !ok { + t.Fatal("due claim unavailable", err) + } + claim.MessageID = "message" + if err := store.FinishApprovalMessage(ctx, claim); err != nil { + t.Fatal(err) + } + claim.Status = "Approved" + store.MarkApprovalMessageDecided(ctx, claim, time.Now().Add(-time.Hour)) + if err := store.RetryApprovalMessage(ctx, claim, future); err != nil { + t.Fatal(err) + } + if due, err := store.DueApprovalMessages(ctx, time.Now()); err != nil || len(due) != 0 { + t.Fatal("cleanup backoff ignored", due, err) + } +} + +func TestDecisionJobsSurviveCardDeletionAndKeepFirstDecision(t *testing.T) { + ctx := context.Background() + path := filepath.Join(t.TempDir(), "state.db") + store, err := Open(path) + if err != nil { + t.Fatal(err) + } + original := ApprovalDecision{RequestID: 42, Status: "Declined", RequesterID: 7, Title: "Arrival", Actor: "First moderator", Reason: "Not this week", DecidedAt: time.Now().UTC()} + canonical, err := store.RecordApprovalDecision(ctx, original) + if err != nil { + t.Fatal(err) + } + later := original + later.Actor = "Second moderator" + later.Reason = "Different reason" + got, err := store.RecordApprovalDecision(ctx, later) + if err != nil || got != canonical { + t.Fatalf("decision overwritten: %#v %v", got, err) + } + store.DeleteApprovalMessage(ctx, ApprovalMessage{RequestID: 42, GuildID: "guild", ChannelID: "c", MessageID: "old"}) + store.Close() + store, err = Open(path) + if err != nil { + t.Fatal(err) + } + defer store.Close() + jobs, err := store.DueDecisionJobs(ctx, time.Now(), 25) + if err != nil || len(jobs) != 1 || jobs[0].Actor != original.Actor { + t.Fatalf("lost decision job: %#v %v", jobs, err) + } + if err := store.RecordDecisionNotification(ctx, 42, "user", "Declined"); err != nil { + t.Fatal(err) + } + if err := store.CompleteDecisionJob(ctx, original, time.Now()); err != nil { + t.Fatal(err) + } + store.RecordApprovalDecision(ctx, later) + if jobs, err := store.DueDecisionJobs(ctx, time.Now(), 25); err != nil || len(jobs) != 0 { + t.Fatal("replayed completed job", jobs, err) + } + if handled, err := store.DecisionNotificationHandled(ctx, 42, "user", "Declined"); err != nil || !handled { + t.Fatal("lost receipt", err) + } +} + +func legacyDeliveryDB(t *testing.T, version int) string { + t.Helper() + path := filepath.Join(t.TempDir(), "state.db") + db, err := sql.Open("sqlite", path) + if err != nil { + t.Fatal(err) + } + store := &Store{db: db} + db.SetMaxOpenConns(1) + if _, err := db.Exec(`CREATE TABLE schema_migrations(version INTEGER PRIMARY KEY,applied_at TEXT NOT NULL)`); err != nil { + t.Fatal(err) + } + for v := 1; v <= version; v++ { + if err := store.applyMigration(context.Background(), v); err != nil { + t.Fatal(err) + } + } + if _, err := db.Exec(`INSERT INTO approval_settings(guild_id,channel_id,enabled) VALUES('g','c',1); + INSERT INTO approval_messages(request_id,guild_id,channel_id,message_id) VALUES(42,'g','c',''),(43,'g','c','existing'); + INSERT INTO subscriptions(request_id,discord_id,title,media_type,created_at) VALUES(42,'user','Arrival','movie','2026-01-01T00:00:00Z')`); err != nil { + t.Fatal(err) + } + if version >= 5 { + if _, err := db.Exec(`INSERT INTO decision_notifications VALUES(41,'user','Approved')`); err != nil { + t.Fatal(err) + } + } + db.Close() + return path +} + +func TestDeliveryUpgradePreservesEveryPriorRevision(t *testing.T) { + for version := 1; version <= 5; version++ { + t.Run(fmt.Sprint(version), func(t *testing.T) { + path := legacyDeliveryDB(t, version) + for range 2 { + store, err := Open(path) + if err != nil { + t.Fatal(err) + } + var count int + if err := store.db.QueryRow(`SELECT count(*) FROM subscriptions WHERE title='Arrival'`).Scan(&count); err != nil || count != 1 { + t.Fatal("lost subscription", err) + } + if err := store.db.QueryRow(`SELECT count(*) FROM approval_messages WHERE message_id='existing'`).Scan(&count); err != nil || count != 1 { + t.Fatal("lost card", err) + } + if err := store.db.QueryRow(`SELECT count(*) FROM decision_jobs`).Scan(&count); err != nil || count != 0 { + t.Fatal("historical notifications were replayed", err) + } + if needed, err := store.NeedsApprovalMessage(context.Background(), 42); err != nil || !needed { + t.Fatal("legacy blank claim is stranded", err) + } + if version == 5 { + if handled, err := store.DecisionNotificationHandled(context.Background(), 41, "user", "Approved"); err != nil || !handled { + t.Fatal("legacy receipt lost", err) + } + } + store.Close() + } + }) + } +} + +func TestDeliveryMigrationRollsBackOnFailure(t *testing.T) { + path := legacyDeliveryDB(t, 5) + db, err := sql.Open("sqlite", path) + if err != nil { + t.Fatal(err) + } + if _, err := db.Exec(`CREATE TRIGGER reject_delivery_version BEFORE INSERT ON schema_migrations WHEN NEW.version=6 BEGIN SELECT RAISE(ABORT,'synthetic migration failure'); END`); err != nil { + t.Fatal(err) + } + db.Close() + if store, err := Open(path); err == nil { + store.Close() + t.Fatal("failed migration was accepted") + } + db, err = sql.Open("sqlite", path) + if err != nil { + t.Fatal(err) + } + defer db.Close() + var count int + if err := db.QueryRow(`SELECT count(*) FROM pragma_table_info('approval_messages') WHERE name='claim_token'`).Scan(&count); err != nil || count != 0 { + t.Fatal("partial migration remained", err) + } + if _, err := db.Exec(`DROP TRIGGER reject_delivery_version`); err != nil { + t.Fatal(err) + } + store, err := Open(path) + if err != nil { + t.Fatal(err) + } + store.Close() +} + +func TestCardAcknowledgementAndDeletionRequirePhysicalIdentity(t *testing.T) { + ctx := context.Background() + store, err := Open(filepath.Join(t.TempDir(), "state.db")) + if err != nil { + t.Fatal(err) + } + defer store.Close() + store.SetApprovalSettings(ctx, ApprovalSettings{GuildID: "g", ChannelID: "c", Enabled: true}) + claim, _, err := store.ClaimApprovalMessage(ctx, ApprovalMessage{RequestID: 42, GuildID: "g", ChannelID: "c"}, time.Now()) + if err != nil { + t.Fatal(err) + } + claim.MessageID = "new" + if err := store.FinishApprovalMessage(ctx, claim); err != nil { + t.Fatal(err) + } + stale := claim + stale.Status = "Approved" + stale.MessageID = "old" + store.MarkApprovalMessageDecided(ctx, stale, time.Now().Add(-time.Hour)) + store.DeleteApprovalMessage(ctx, stale) + messages, err := store.ApprovalMessages(ctx) + if err != nil || len(messages) != 1 || messages[0].MessageID != "new" || !messages[0].DecidedAt.IsZero() { + t.Fatal("stale card affected replacement", messages, err) + } +} + +func TestOrphanCleanupSurvivesRestartAndProtectsTrackedCard(t *testing.T) { + ctx := context.Background() + path := filepath.Join(t.TempDir(), "state.db") + store, err := Open(path) + if err != nil { + t.Fatal(err) + } + orphan := ApprovalMessage{ChannelID: "c", MessageID: "orphan"} + if queued, err := store.QueueUntrackedApprovalCleanup(ctx, orphan); err != nil || !queued { + t.Fatal("orphan not retained", err) + } + store.Close() + store, err = Open(path) + if err != nil { + t.Fatal(err) + } + defer store.Close() + if due, err := store.DueApprovalCleanup(ctx, time.Now()); err != nil || len(due) != 1 || due[0].MessageID != "orphan" { + t.Fatal("orphan lost on restart", due, err) + } + store.RetryApprovalCleanup(ctx, orphan, time.Now().Add(time.Hour)) + if due, err := store.DueApprovalCleanup(ctx, time.Now()); err != nil || len(due) != 0 { + t.Fatal("orphan retry ignored backoff", due, err) + } + store.SetApprovalSettings(ctx, ApprovalSettings{GuildID: "g", ChannelID: "c", Enabled: true}) + claim, _, err := store.ClaimApprovalMessage(ctx, ApprovalMessage{RequestID: 42, GuildID: "g", ChannelID: "c"}, time.Now()) + if err != nil { + t.Fatal(err) + } + claim.MessageID = "tracked" + if err := store.FinishApprovalMessage(ctx, claim); err != nil { + t.Fatal(err) + } + if queued, err := store.QueueUntrackedApprovalCleanup(ctx, claim); err != nil || queued { + t.Fatal("tracked card queued for orphan deletion", err) + } +} diff --git a/internal/storage/store.go b/internal/storage/store.go index 573a90d..b2428ef 100644 --- a/internal/storage/store.go +++ b/internal/storage/store.go @@ -37,13 +37,16 @@ type ApprovalSettings struct { } type ApprovalMessage struct { - RequestID int - GuildID string - ChannelID string - MessageID string - DecidedAt time.Time - Status string - Reason string + RequestID int + GuildID string + ChannelID string + MessageID string + DecidedAt time.Time + Status string + Reason string + ClaimToken string + Attempts int + LeaseUntil time.Time } type NotificationPreferences struct { @@ -55,7 +58,7 @@ type NotificationPreferences struct { const subscriptionColumnList = "request_id, discord_id, title, media_type, overview, poster_path, release_year, language, rating, created_at, completed_at" -const currentMigrationVersion = 5 +const currentMigrationVersion = 6 func Open(path string) (*Store, error) { dbPath, err := storagePath(path) @@ -219,14 +222,22 @@ func (s *Store) SetApprovalSettings(ctx context.Context, settings ApprovalSettin if settings.Enabled && settings.ChannelID == "" { return errors.New("channel_id is required when approvals are enabled") } - _, err := s.db.ExecContext(ctx, ` - INSERT INTO approval_settings (guild_id, channel_id, enabled) VALUES (?, ?, ?) - ON CONFLICT(guild_id) DO UPDATE SET channel_id = excluded.channel_id, enabled = excluded.enabled - `, settings.GuildID, settings.ChannelID, settings.Enabled) + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return err + } + defer tx.Rollback() + _, err = tx.ExecContext(ctx, ` + INSERT INTO approval_settings (guild_id, channel_id, enabled) VALUES (?, ?, ?) + ON CONFLICT(guild_id) DO UPDATE SET channel_id = excluded.channel_id, enabled = excluded.enabled + `, settings.GuildID, settings.ChannelID, settings.Enabled) if err != nil { return fmt.Errorf("save approval settings: %w", err) } - return nil + if _, err = tx.ExecContext(ctx, `UPDATE approval_messages SET claim_token='',lease_until=0,next_attempt_at=0 WHERE guild_id=? AND message_id=''`, settings.GuildID); err != nil { + return err + } + return tx.Commit() } func (s *Store) ApprovalSettings(ctx context.Context, guildID string) (ApprovalSettings, bool, error) { @@ -262,94 +273,27 @@ func (s *Store) EnabledApprovalSettings(ctx context.Context) ([]ApprovalSettings return settings, rows.Err() } -func (s *Store) NeedsApprovalMessage(ctx context.Context, requestID int) (bool, error) { - if requestID <= 0 { - return false, errors.New("request_id must be positive") - } - var needed int - err := s.db.QueryRowContext(ctx, ` - SELECT EXISTS ( - SELECT 1 FROM approval_settings settings - LEFT JOIN approval_messages messages - ON messages.guild_id = settings.guild_id AND messages.request_id = ? - WHERE settings.enabled = 1 AND messages.request_id IS NULL - ) - `, requestID).Scan(&needed) - if err != nil { - return false, fmt.Errorf("check approval message coverage: %w", err) - } - return needed != 0, nil -} - -func (s *Store) ClaimApprovalMessage(ctx context.Context, message ApprovalMessage) (bool, error) { - result, err := s.db.ExecContext(ctx, ` - INSERT INTO approval_messages (request_id, guild_id, channel_id, message_id) - VALUES (?, ?, ?, '') ON CONFLICT(request_id, guild_id) DO NOTHING - `, message.RequestID, strings.TrimSpace(message.GuildID), strings.TrimSpace(message.ChannelID)) - if err != nil { - return false, fmt.Errorf("claim approval message: %w", err) - } - n, err := result.RowsAffected() - return n == 1, err -} - -func (s *Store) FinishApprovalMessage(ctx context.Context, message ApprovalMessage) error { - result, err := s.db.ExecContext(ctx, `UPDATE approval_messages SET message_id = ? WHERE request_id = ? AND guild_id = ? AND message_id = ''`, message.MessageID, message.RequestID, message.GuildID) - if err != nil { - return fmt.Errorf("finish approval message: %w", err) - } - n, err := result.RowsAffected() - if err != nil || n != 1 { - return errors.New("approval message claim was not found") - } - return nil -} - -func (s *Store) ReleaseApprovalMessage(ctx context.Context, requestID int, guildID string) error { - _, err := s.db.ExecContext(ctx, `DELETE FROM approval_messages WHERE request_id = ? AND guild_id = ? AND message_id = ''`, requestID, guildID) - return err -} - -func (s *Store) MarkApprovalMessageDecided(ctx context.Context, requestID int, guildID string, decidedAt time.Time) error { - if requestID <= 0 || strings.TrimSpace(guildID) == "" { +func (s *Store) MarkApprovalMessageDecided(ctx context.Context, message ApprovalMessage, decidedAt time.Time) error { + if message.RequestID <= 0 || message.GuildID == "" || message.ChannelID == "" || message.MessageID == "" || message.Status == "" { return errors.New("request_id and guild_id are required") } if decidedAt.IsZero() { decidedAt = time.Now().UTC() } - _, err := s.db.ExecContext(ctx, `UPDATE approval_messages SET decided_at = ? WHERE request_id = ? AND guild_id = ?`, formatTime(decidedAt), requestID, guildID) + _, err := s.db.ExecContext(ctx, `UPDATE approval_messages SET decided_at = ?,status=? WHERE request_id = ? AND guild_id = ? AND channel_id=? AND message_id=? AND (status='' OR status=?)`, formatTime(decidedAt), message.Status, message.RequestID, message.GuildID, message.ChannelID, message.MessageID, message.Status) if err != nil { return fmt.Errorf("mark approval message decided: %w", err) } return nil } -func (s *Store) SetApprovalDecision(ctx context.Context, requestID int, guildID, status, reason string) error { - _, err := s.db.ExecContext(ctx, `UPDATE approval_messages SET status = ?, reason = ? WHERE request_id = ? AND guild_id = ?`, strings.TrimSpace(status), strings.TrimSpace(reason), requestID, strings.TrimSpace(guildID)) - return err -} - -func (s *Store) ClaimDecisionNotification(ctx context.Context, requestID int, discordID, status string) (bool, error) { - result, err := s.db.ExecContext(ctx, `INSERT INTO decision_notifications(request_id, discord_id, status) VALUES(?,?,?) ON CONFLICT(request_id, discord_id, status) DO NOTHING`, requestID, strings.TrimSpace(discordID), strings.TrimSpace(status)) - if err != nil { - return false, err - } - n, err := result.RowsAffected() - return n == 1, err -} - -func (s *Store) ReleaseDecisionNotification(ctx context.Context, requestID int, discordID, status string) error { - _, err := s.db.ExecContext(ctx, `DELETE FROM decision_notifications WHERE request_id = ? AND discord_id = ? AND status = ?`, requestID, strings.TrimSpace(discordID), strings.TrimSpace(status)) - return err -} - func (s *Store) DueApprovalMessages(ctx context.Context, before time.Time) ([]ApprovalMessage, error) { rows, err := s.db.QueryContext(ctx, ` - SELECT request_id, guild_id, channel_id, message_id, decided_at, status, reason + SELECT request_id, guild_id, channel_id, message_id, decided_at, status, reason, attempts FROM approval_messages - WHERE decided_at IS NOT NULL AND decided_at <= ? AND message_id != '' - ORDER BY decided_at, request_id - `, formatTime(before.UTC())) + WHERE decided_at IS NOT NULL AND decided_at <= ? AND message_id != '' AND next_attempt_at <= ? + ORDER BY decided_at, request_id LIMIT 100 + `, formatTime(before.UTC()), time.Now().UnixMilli()) if err != nil { return nil, fmt.Errorf("list due approval messages: %w", err) } @@ -358,7 +302,7 @@ func (s *Store) DueApprovalMessages(ctx context.Context, before time.Time) ([]Ap for rows.Next() { var message ApprovalMessage var decidedAt string - if err := rows.Scan(&message.RequestID, &message.GuildID, &message.ChannelID, &message.MessageID, &decidedAt, &message.Status, &message.Reason); err != nil { + if err := rows.Scan(&message.RequestID, &message.GuildID, &message.ChannelID, &message.MessageID, &decidedAt, &message.Status, &message.Reason, &message.Attempts); err != nil { return nil, fmt.Errorf("scan due approval message: %w", err) } parsed, err := parseTime(decidedAt) @@ -371,16 +315,13 @@ func (s *Store) DueApprovalMessages(ctx context.Context, before time.Time) ([]Ap return messages, rows.Err() } -func (s *Store) DeleteApprovalMessage(ctx context.Context, requestID int, guildID string) error { - _, err := s.db.ExecContext(ctx, `DELETE FROM approval_messages WHERE request_id = ? AND guild_id = ?`, requestID, strings.TrimSpace(guildID)) - if err != nil { - return fmt.Errorf("delete approval message: %w", err) - } - return nil +func (s *Store) DeleteApprovalMessage(ctx context.Context, message ApprovalMessage) error { + _, err := s.db.ExecContext(ctx, `DELETE FROM approval_messages WHERE request_id=? AND guild_id=? AND channel_id=? AND message_id=? AND (message_id!='' OR claim_token=?)`, message.RequestID, message.GuildID, message.ChannelID, message.MessageID, message.ClaimToken) + return err } func (s *Store) ApprovalMessages(ctx context.Context) ([]ApprovalMessage, error) { - rows, err := s.db.QueryContext(ctx, `SELECT request_id, guild_id, channel_id, message_id, decided_at, status, reason FROM approval_messages WHERE message_id != ''`) + rows, err := s.db.QueryContext(ctx, `SELECT request_id, guild_id, channel_id, message_id, decided_at, status, reason, attempts,claim_token FROM approval_messages WHERE decided_at IS NULL AND next_attempt_at<=? ORDER BY next_attempt_at,request_id,guild_id`, time.Now().UnixMilli()) if err != nil { return nil, err } @@ -389,7 +330,7 @@ func (s *Store) ApprovalMessages(ctx context.Context) ([]ApprovalMessage, error) for rows.Next() { var m ApprovalMessage var decided sql.NullString - if err := rows.Scan(&m.RequestID, &m.GuildID, &m.ChannelID, &m.MessageID, &decided, &m.Status, &m.Reason); err != nil { + if err := rows.Scan(&m.RequestID, &m.GuildID, &m.ChannelID, &m.MessageID, &decided, &m.Status, &m.Reason, &m.Attempts, &m.ClaimToken); err != nil { return nil, err } if decided.Valid { @@ -499,6 +440,26 @@ func (s *Store) applyMigration(ctx context.Context, version int) error { 3: {`ALTER TABLE approval_messages ADD COLUMN status TEXT NOT NULL DEFAULT ''`, `ALTER TABLE approval_messages ADD COLUMN reason TEXT NOT NULL DEFAULT ''`}, 4: {`CREATE TABLE IF NOT EXISTS notification_preferences (discord_id TEXT PRIMARY KEY, approved INTEGER NOT NULL DEFAULT 1 CHECK (approved IN (0,1)), declined INTEGER NOT NULL DEFAULT 1 CHECK (declined IN (0,1)), available INTEGER NOT NULL DEFAULT 1 CHECK (available IN (0,1)))`}, 5: {`CREATE TABLE IF NOT EXISTS decision_notifications (request_id INTEGER NOT NULL, discord_id TEXT NOT NULL, status TEXT NOT NULL, PRIMARY KEY(request_id, discord_id, status))`}, + 6: { + `CREATE TABLE approval_cleanup(channel_id TEXT NOT NULL,message_id TEXT NOT NULL,attempts INTEGER NOT NULL DEFAULT 0,next_attempt_at INTEGER NOT NULL DEFAULT 0,PRIMARY KEY(channel_id,message_id))`, + `CREATE INDEX idx_approval_cleanup_due ON approval_cleanup(next_attempt_at,channel_id,message_id)`, + `CREATE TABLE decision_intents(request_id INTEGER PRIMARY KEY CHECK(request_id>0),status TEXT NOT NULL CHECK(status IN ('Approved','Declined')),actor TEXT NOT NULL,reason TEXT NOT NULL,title TEXT NOT NULL DEFAULT '',url TEXT NOT NULL DEFAULT '',poster_url TEXT NOT NULL DEFAULT '',created_at TEXT NOT NULL,next_attempt_at INTEGER NOT NULL DEFAULT 0)`, + `CREATE INDEX idx_decision_intents_due ON decision_intents(next_attempt_at,request_id)`, + `ALTER TABLE approval_messages ADD COLUMN claim_token TEXT NOT NULL DEFAULT ''`, + `ALTER TABLE approval_messages ADD COLUMN lease_until INTEGER NOT NULL DEFAULT 0`, + `ALTER TABLE approval_messages ADD COLUMN attempts INTEGER NOT NULL DEFAULT 0`, + `ALTER TABLE approval_messages ADD COLUMN next_attempt_at INTEGER NOT NULL DEFAULT 0`, + `CREATE TABLE decision_jobs ( + request_id INTEGER NOT NULL CHECK(request_id>0), status TEXT NOT NULL CHECK(status IN ('Approved','Declined')), + requester_id INTEGER NOT NULL DEFAULT 0, media_id INTEGER NOT NULL DEFAULT 0, media_type TEXT NOT NULL DEFAULT '', + title TEXT NOT NULL DEFAULT '', url TEXT NOT NULL DEFAULT '', poster_url TEXT NOT NULL DEFAULT '', + actor TEXT NOT NULL DEFAULT '', reason TEXT NOT NULL DEFAULT '', decided_at TEXT NOT NULL, + attempts INTEGER NOT NULL DEFAULT 0, next_attempt_at INTEGER NOT NULL DEFAULT 0, completed_at TEXT, + PRIMARY KEY(request_id,status))`, + `CREATE INDEX idx_decision_jobs_due ON decision_jobs(next_attempt_at,request_id) WHERE completed_at IS NULL`, + `CREATE INDEX idx_approval_cleanup ON approval_messages(decided_at,request_id) WHERE decided_at IS NOT NULL`, + `CREATE INDEX idx_approval_refresh ON approval_messages(next_attempt_at,request_id) WHERE decided_at IS NULL AND message_id!=''`, + }, } columns := map[string]map[string]struct{}{} if version == 2 || version == 3 { diff --git a/internal/storage/store_test.go b/internal/storage/store_test.go index 22d83c0..2ba2ebe 100644 --- a/internal/storage/store_test.go +++ b/internal/storage/store_test.go @@ -113,7 +113,7 @@ func TestExistingMainDatabaseUpgradesWithoutLosingRows(t *testing.T) { if err != nil || !ok || settings.ChannelID != "channel" || !settings.Enabled { t.Fatalf("upgraded approval settings = %#v, %t, err %v", settings, ok, err) } - messages, err := store.ApprovalMessages(context.Background()) + messages, err := store.DueApprovalMessages(context.Background(), time.Now().UTC()) if err != nil || len(messages) != 1 || messages[0].MessageID != "message" || messages[0].DecidedAt.IsZero() { t.Fatalf("upgraded approval messages = %#v, err %v", messages, err) } @@ -323,14 +323,15 @@ func TestApprovalSettingsAndMessageDedupe(t *testing.T) { if err != nil || !needed { t.Fatalf("NeedsApprovalMessage before claim = %t, %v", needed, err) } - claimed, err := store.ClaimApprovalMessage(ctx, message) + claim, claimed, err := store.ClaimApprovalMessage(ctx, message, time.Now()) if err != nil || !claimed { t.Fatalf("first claim = %t, %v", claimed, err) } - claimed, err = store.ClaimApprovalMessage(ctx, message) + _, claimed, err = store.ClaimApprovalMessage(ctx, message, time.Now()) if err != nil || claimed { t.Fatalf("duplicate claim = %t, %v", claimed, err) } + message = claim message.MessageID = "789" if err := store.FinishApprovalMessage(ctx, message); err != nil { t.Fatal(err) @@ -339,15 +340,16 @@ func TestApprovalSettingsAndMessageDedupe(t *testing.T) { if err != nil || needed { t.Fatalf("NeedsApprovalMessage after claim = %t, %v", needed, err) } + message.Status = "Approved" decidedAt := time.Now().UTC().Add(-3 * time.Minute) - if err := store.MarkApprovalMessageDecided(ctx, 42, "123", decidedAt); err != nil { + if err := store.MarkApprovalMessageDecided(ctx, message, decidedAt); err != nil { t.Fatal(err) } due, err := store.DueApprovalMessages(ctx, time.Now().UTC().Add(-2*time.Minute)) if err != nil || len(due) != 1 || due[0].MessageID != "789" { t.Fatalf("DueApprovalMessages() = %#v, %v", due, err) } - if err := store.DeleteApprovalMessage(ctx, 42, "123"); err != nil { + if err := store.DeleteApprovalMessage(ctx, message); err != nil { t.Fatal(err) } due, err = store.DueApprovalMessages(ctx, time.Now().UTC()) diff --git a/unraid/augur.xml b/unraid/augur.xml index 739b620..92de695 100644 --- a/unraid/augur.xml +++ b/unraid/augur.xml @@ -12,7 +12,7 @@ MediaApp:Other Network:Other https://raw.githubusercontent.com/mayvqt/Augur/main/unraid/augur.xml - --restart=unless-stopped --init --tmpfs /tmp --security-opt=no-new-privileges:true + --restart=unless-stopped --stop-timeout=45 --init --tmpfs /tmp --security-opt=no-new-privileges:true