diff --git a/cmd/aperture/serve.go b/cmd/aperture/serve.go index d9813c9..ae4bafa 100644 --- a/cmd/aperture/serve.go +++ b/cmd/aperture/serve.go @@ -7,14 +7,12 @@ import ( "net/http" "os" "os/signal" - "strings" "syscall" "time" "github.com/mayvqt/aperture/internal/config" "github.com/mayvqt/aperture/internal/db" "github.com/mayvqt/aperture/internal/httpserver" - "github.com/mayvqt/aperture/internal/mediaserver" "github.com/mayvqt/aperture/internal/mediaserver/router" ) @@ -47,59 +45,51 @@ func serve(args []string) error { cfg.APIKey, cfg.EncryptionKey, cfg.SessionSecret, cfg.InviteSecret, settings.APIKey, settings.SessionSecret, settings.InviteSecret, ) - if cfg.ProviderManaged { - if err := store.ValidateMediaProvider(context.Background(), cfg.MediaProvider); err != nil { - return err - } - } else if provider, ok := mediaserver.ParseProvider(settings.Provider); ok { - cfg.MediaProvider = string(provider) - } - if !cfg.PublicURLManaged && settings.PublicURL != "" { - cfg.PublicURL = settings.PublicURL - } - if !cfg.CookieManaged { - cfg.CookieSecure = strings.HasPrefix(cfg.PublicURL, "https://") - } - provider, _ := mediaserver.ParseProvider(cfg.MediaProvider) - media, err := router.New(provider) - if err != nil { - return err - } ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) defer stop() + handler := httpserver.NewServer(cfg, store, router.New) + if err := handler.Initialize(ctx); err != nil { + return err + } workerDone := make(chan struct{}) go func() { defer close(workerDone) - httpserver.RunMaintenanceWorker(ctx, cfg, store, media) + handler.RunMaintenance(ctx) }() - handler, waitForWebhooks := httpserver.NewWithShutdown(cfg, store, media) srv := &http.Server{ Addr: cfg.HTTPAddr, Handler: handler, ReadHeaderTimeout: 10 * time.Second, ReadTimeout: 30 * time.Second, - WriteTimeout: 30 * time.Second, + WriteTimeout: 4 * time.Minute, IdleTimeout: 120 * time.Second, MaxHeaderBytes: 64 << 10, } + shutdownDone := make(chan struct{}) go func() { + defer close(shutdownDone) <-ctx.Done() - shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + handler.CloseAdmission() + shutdownCtx, cancel := context.WithTimeout(context.Background(), 4*time.Minute) defer cancel() - _ = srv.Shutdown(shutdownCtx) + if err := srv.Shutdown(shutdownCtx); err != nil { + slog.Warn("HTTP shutdown deadline reached", "error", err) + _ = srv.Close() + } }() slog.Info("starting aperture", "addr", cfg.HTTPAddr, "db", cfg.DBPath) err = srv.ListenAndServe() stop() + <-shutdownDone <-workerDone - deliveryCtx, cancelDeliveries := context.WithTimeout(context.Background(), 10*time.Second) - defer cancelDeliveries() - if waitErr := waitForWebhooks(deliveryCtx); waitErr != nil { - slog.Warn("webhook deliveries did not finish before shutdown", "error", waitErr) + // Each accepted operation and webhook already has its own deadline. Closing + // SQLite after a separate drain timeout could race their final state writes. + if waitErr := handler.Drain(context.Background()); waitErr != nil { + return waitErr } if errors.Is(err, http.ErrServerClosed) { return nil diff --git a/docker-compose.yml b/docker-compose.yml index b7646a7..94b5ab6 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -6,6 +6,7 @@ services: image: aperture:local container_name: aperture restart: unless-stopped + stop_grace_period: 270s init: true security_opt: - no-new-privileges:true diff --git a/docs/SECURITY.md b/docs/SECURITY.md index ef7d28c..0633a90 100644 --- a/docs/SECURITY.md +++ b/docs/SECURITY.md @@ -21,12 +21,19 @@ Initial setup is serialized within the application process so a delayed request cannot overwrite completed setup. Run one Aperture process per state directory. After an invite use is reserved, account provisioning continues for a bounded period even if the browser disconnects. Incomplete accounts are disabled when -the media server is reachable; failed cleanup requires administrator intervention. +the media server is reachable. Cleanup is persisted before access changes and +retried after failures or restarts, independently of the template retry limit. +The account ID is saved before Emby password setup begins. Accounts with incomplete password setup retain their external ID for review and cannot be enabled through automatic or manual template-only retries. +Template policies always disable administrator access, regardless of property +casing. Conflicting property names are rejected. Applying a template preserves +the target account's authentication and password-reset providers and merges its +remaining defaults, so imported authentication settings do not cross accounts. + Before upgrading an existing installation, review older `needs_attention` registrations in the media server, especially password-setup failures. Earlier -records are not reclassified by this release. Keep any incomplete accounts +records may not identify the interrupted password step. Keep any incomplete accounts disabled and finish password setup before allowing a template retry. Do this before restarting Aperture, since maintenance can retry eligible records. diff --git a/docs/capabilities.md b/docs/capabilities.md index d90ea9f..c6a3744 100644 --- a/docs/capabilities.md +++ b/docs/capabilities.md @@ -4,10 +4,11 @@ Aperture provides: - browser-based first-run configuration for Jellyfin or Emby; - administrator authentication through the configured media server; -- reusable non-administrator policy templates, including import from an existing user; +- reusable non-administrator policy templates that preserve target authentication defaults, including import from an existing user; - bounded, expiring invite links with usage limits and optional account expiry; -- account creation, policy application, retry, disable, and recovery workflows; -- managed-user and registration history views; +- account creation, policy application, retry, disable, and recovery workflows with durable incomplete-account cleanup; +- paginated invite and registration history, including older records that need review; +- server-scoped account tracking with explicit review when changing servers; - Discord and generic JSON webhooks with selected events; and - administrative audit history with bounded retention. diff --git a/docs/configuration.md b/docs/configuration.md index c6bef53..ce15200 100644 --- a/docs/configuration.md +++ b/docs/configuration.md @@ -20,3 +20,32 @@ and use a compatible Aperture version; never delete the database as a normal upg Admins can configure Discord or generic JSON webhooks, events, and optional Discord role IDs in the web UI. URLs are encrypted. Failed template application is retried automatically up to six times. + +## Account recovery + +Registrations show access retries and pending disables separately. After six +automatic access retries, review the account and use **Retry access** when ready. +Failed disables keep retrying with backoff until the account is disabled or no +longer exists. Password-incomplete accounts require administrator review and +cannot be enabled by retrying a template. + +An account's expiry is fixed when its invite use is reserved. Recovery never +extends it, and an expired account can only be disabled. Accounts undergoing +creation or recovery cannot be removed until that operation finishes. + +## Changing servers and reviewing history + +Ownership includes the provider, normalized server URL and authenticated server ID. +Changing any of these signs out all administrators. Existing invites and accounts +remain associated with their original server; automatic work pauses for records +that belong elsewhere. A cloned server ID at a different URL is a separate server. + +After an upgrade, older records show **Unverified server**. Use **Review server** on +an invite or registration to inspect the destination and confirm its assignment. +Review tracked-only users under **Users → Tracked user history**. Assignment keeps +invite links, usage counts and account deadlines. It never assigns all accounts +from an invite together. The account ID must exist on the destination server with +a verified non-administrator policy; Aperture does not guess ownership from names. + +Use **Needs review** to find older records requiring attention. Invite and +registration history pages show 50 records at a time, with links to older pages. diff --git a/docs/development/architecture.md b/docs/development/architecture.md index 22831cd..0ad401d 100644 --- a/docs/development/architecture.md +++ b/docs/development/architecture.md @@ -1,8 +1,8 @@ # Architecture `cmd/aperture` loads configuration, opens the store, initializes schema and -runtime secrets, selects a media-server adapter, starts the maintenance worker, -and serves `internal/httpserver`. Handlers depend on narrow store and media-server +runtime secrets, selects a media-server adapter, and starts one HTTP server that +owns maintenance and notification delivery. Handlers depend on narrow store and media-server interfaces; persistence stays in `internal/db`, while outbound Jellyfin/Emby protocol details stay in `internal/mediaserver`. @@ -20,6 +20,9 @@ protocol details stay in `internal/mediaserver`. - Jellyfin and Emby request shapes and authorization are provider contracts. Verify changes against the supported upstream documentation and cover them with deterministic HTTP fixtures before any opt-in live smoke test. +- Policy application reads the target user's complete policy before merging + template overrides. Shared normalization rejects ambiguous case aliases, + preserves target authentication providers, and forces non-administrator access. ## Security and concurrency boundaries @@ -30,8 +33,20 @@ and no-redirect outbound clients. SQLite is deliberately limited to one connection and uses WAL, foreign keys, and a busy timeout. Setup is serialized in-process. Registration reserves invite -capacity before provisioning; ambiguous external failures retain evidence rather -than silently releasing capacity. The maintenance worker reconciles stale work, -retries templates, disables expired users, and prunes audit events. Webhook -deliveries are tracked and receive a bounded graceful-shutdown window. Changes to +capacity and an immutable expiry before provisioning; ambiguous external failures +retain evidence rather than silently releasing capacity. Shared account recovery +coordinates manual and automatic retries with per-registration operation guards, +transactional claims, and durable cleanup. The maintenance worker reconciles +stale work, disables incomplete and expired accounts before retrying templates, +and prunes audit events. Graceful shutdown drains accepted HTTP and maintenance +work before notifications and before closing SQLite. Changes to these flows must preserve idempotency, bounded work, and cancellation behavior. + +`internal/connection` resolves deployment overrides and publishes immutable +connection snapshots through a media adapter factory. Each accepted operation +retains its adapter, URL, credentials and verified ownership. Administrative +sessions must match both binding and generation before their tokens are sent; +authenticated system information detects a replacement at the same URL. +Background account queries filter ownership before limits and grouping. History +remains readable across origins; explicit assignment uses account operation guards +and verifies an existing non-administrator account at the destination. diff --git a/docs/development/codebase-map.md b/docs/development/codebase-map.md index 2af5d5d..3cc57db 100644 --- a/docs/development/codebase-map.md +++ b/docs/development/codebase-map.md @@ -4,10 +4,11 @@ | --- | --- | | `cmd/aperture` | CLI dispatch, configuration startup, logging, process lifecycle, and version output. | | `internal/config` | Environment/flag parsing, defaults, URL validation, and generated encryption-key bootstrap. | +| `internal/connection` | Effective configuration, immutable operation snapshots, server identity verification and atomic connection publication. | | `internal/db` | SQLite schema, migrations, settings, invites, templates, sessions, registrations, managed users, and audit records. | -| `internal/httpserver` | Routes, middleware, browser workflows, embedded templates/assets, webhooks, and maintenance coordination. | +| `internal/httpserver` | Routes, middleware, browser workflows, embedded templates/assets, shared account recovery, webhooks, and maintenance coordination. | | `internal/mediaserver` | Provider-neutral contracts and URL rules. | -| `internal/mediaserver/jellyfin`, `emby`, `protocol`, `router` | Provider adapters, HTTP protocol, and runtime provider selection. | +| `internal/mediaserver/jellyfin`, `emby`, `protocol`, `router` | Provider adapters, HTTP protocol, and immutable adapter factory. | | `internal/security` | Encryption, token helpers, and diagnostic redaction. | | `scripts`, `docker-entrypoint.sh`, `Dockerfile`, `docker-compose.yml` | Container build, startup, ownership, and regression checks. | | `templates/unraid` | Unraid application template. | diff --git a/docs/development/data.md b/docs/development/data.md index 6598c9c..d23b7ab 100644 --- a/docs/development/data.md +++ b/docs/development/data.md @@ -31,3 +31,16 @@ Every schema change requires: Never repair an upgrade by deleting the database. Back up and restore the whole state set as described in [Operations](operations.md). + +Revision 5 adds durable account cleanup state. Expiry is captured in the invite +reservation transaction and retained through failures and recovery. Migration +preserves existing deadlines and reconstructs missing finite deadlines from the +original registration time and retained invite duration. Incomplete accounts +with known upstream IDs enter cleanup independently of access retry attempts. + +Revision 6 adds server ownership using provider, normalized URL and the authenticated +server ID. Existing invite, registration and tracked-user rows retain their data +with unknown ownership until an administrator reviews each record. Connection +publication atomically saves settings and revokes sessions on an origin change. +Reassigning an account preserves its deadline and resets the old server's disable +acknowledgement so expired access is checked again. diff --git a/docs/development/operations.md b/docs/development/operations.md index 52ddef2..4597143 100644 --- a/docs/development/operations.md +++ b/docs/development/operations.md @@ -31,3 +31,18 @@ the pre-upgrade state with the previous immutable image. Stop on failed health o data verification; preserve failed state for diagnosis. Maintenance or repair actions that mutate users, registrations, or SQLite state are explicit, operator-approved procedures, never automatic troubleshooting steps. + +Revision 5 resumes incomplete-account disables from durable state. Review older +password-setup failures before upgrading as described in [Security](../SECURITY.md). +Allow four minutes for accepted account operations to finish before SQLite +closes. The supplied Compose and Unraid configurations allow 270 seconds, +including notification delivery, before forcing the process to stop. + +After upgrading to revision 6, sign in again and review saved invites and accounts. +Legacy ownership is unknown, so automatic account changes pause until each account +is assigned. On **Invites** and **Registrations**, use **Review server**; previously +imported users are under **Users → Tracked user history**. Confirm the displayed +server and account ID before assigning. Existing deadlines are preserved, so an +expired or incomplete assigned account may be disabled on the next maintenance +run. Reconnecting the same provider, normalized URL and server ID restores its +existing ownership; an outage does not silently assign records elsewhere. diff --git a/docs/development/validation.md b/docs/development/validation.md index 9a5f63b..5a561e1 100644 --- a/docs/development/validation.md +++ b/docs/development/validation.md @@ -16,6 +16,15 @@ Format touched Go files with `gofmt -w`. Performance or refactor claims require representative before/after benchmark or trace and a regression threshold; a clean test run alone is not performance evidence. +Media-server policy fixtures cover complete target defaults, case-insensitive +overrides, duplicate access flags, and disabling an account after an ambiguous +failure. Keep these checks when changing template import or application. + +Account lifecycle regressions cover interrupted password setup, lost policy and +completion responses, exhausted template retries with pending cleanup, fixed +expiry, operation collisions, and fresh/legacy/rollback migration behavior with +encrypted-value preservation. + ## UI evidence Rebuild before visual checks because templates, CSS, and JavaScript are embedded. @@ -37,3 +46,8 @@ The one comprehensive gate is the complete GitHub Actions `CI` workflow in whitespace, `go test ./...`, the entrypoint, vet, race detection, pinned Staticcheck and govulncheck versions, a release-style build, and the Docker image. Do not claim release readiness until both jobs pass on that revision. + +Origin regressions cover upgrades from revisions 1–5, encrypted invite preservation, +rollback, unknown ownership, cloned IDs at different URLs, replacement servers, +API-key rotation, session revocation, lost settings acknowledgements, optional-key +login, scoped work queues and paginated review of old records. diff --git a/docs/setup.md b/docs/setup.md index 4810142..8bb928e 100644 --- a/docs/setup.md +++ b/docs/setup.md @@ -9,6 +9,10 @@ Open `http://localhost:8099` from a trusted network. Choose Jellyfin or Emby, en and API key, then sign in as a media-server administrator. Create or import a non-admin template before creating an invite. +The default template uses the media server's account defaults with administrator +access disabled. Imported policies keep the new account's authentication and +password-reset providers. JSON property names must be unique regardless of casing. + Standalone installs use `aperture serve`. Persist and back up the config directory. Do not expose Aperture publicly until the unauthenticated first-run setup is complete. diff --git a/internal/connection/manager.go b/internal/connection/manager.go new file mode 100644 index 0000000..a03f172 --- /dev/null +++ b/internal/connection/manager.go @@ -0,0 +1,235 @@ +// Package connection owns effective configuration and immutable media-server +// snapshots. Publishing settings never mutates an adapter already in use. +package connection + +import ( + "context" + "errors" + "strings" + "sync" + "time" + + "github.com/mayvqt/aperture/internal/config" + "github.com/mayvqt/aperture/internal/db" + "github.com/mayvqt/aperture/internal/mediaserver" +) + +var ErrUnavailable = errors.New("media-server identity is unavailable") + +type Factory func(mediaserver.Provider) (mediaserver.Server, error) + +type Store interface { + Settings(context.Context) (db.Settings, error) + MediaConnection(context.Context) (db.MediaConnection, error) + PublishMediaConnection(context.Context, db.ConnectionUpdate) (db.MediaConnection, error) +} + +type Snapshot struct { + Settings db.Settings + Identity db.MediaConnection + Media mediaserver.Server + CookieSecure bool + Verified bool + version uint64 +} + +type Manager struct { + cfg config.Config + store Store + factory Factory + updateMu sync.Mutex + mu sync.RWMutex + current Snapshot + revision uint64 +} + +func New(cfg config.Config, store Store, factory Factory) *Manager { + return &Manager{cfg: cfg, store: store, factory: factory} +} + +// Current resolves environment overrides before deciding whether saved session +// ownership still applies. It makes no upstream requests. +func (m *Manager) Current(ctx context.Context) (Snapshot, error) { + m.mu.RLock() + current := m.current + m.mu.RUnlock() + if current.version != 0 { + return current, nil + } + m.updateMu.Lock() + defer m.updateMu.Unlock() + m.mu.RLock() + current = m.current + m.mu.RUnlock() + if current.version != 0 { + return current, nil + } + settings, err := m.store.Settings(ctx) + if err != nil { + return Snapshot{}, err + } + settings = m.effective(settings) + provider, ok := mediaserver.ParseProvider(settings.Provider) + if !ok { + return Snapshot{}, errors.New("unsupported media provider") + } + if settings.ServerURL != "" { + settings.ServerURL, err = mediaserver.NormalizeBaseURL(provider, settings.ServerURL) + if err != nil { + return Snapshot{}, err + } + } + media, err := m.factory(provider) + if err != nil { + return Snapshot{}, err + } + identity, err := m.store.MediaConnection(ctx) + if err != nil { + return Snapshot{}, err + } + if identity.Provider != settings.Provider || identity.BaseURL != settings.ServerURL { + identity, err = m.store.PublishMediaConnection(ctx, db.ConnectionUpdate{Origin: db.MediaBinding{Provider: settings.Provider, BaseURL: settings.ServerURL}, ExpectedGeneration: identity.Generation}) + if err != nil { + return Snapshot{}, err + } + } + m.revision++ + current = Snapshot{Settings: settings, Identity: identity, Media: media, CookieSecure: m.cookieSecure(settings.PublicURL), version: m.revision} + m.replace(current) + return current, nil +} + +func (m *Manager) effective(s db.Settings) db.Settings { + if m.cfg.ProviderManaged || s.Provider == "" { + s.Provider = m.cfg.MediaProvider + } + if s.Provider == "" { + s.Provider = string(mediaserver.ProviderJellyfin) + } + if m.cfg.PublicURLManaged || s.PublicURL == "" { + s.PublicURL = m.cfg.PublicURL + } + if m.cfg.ServerURLManaged { + s.ServerURL = m.cfg.ServerURL + } + if m.cfg.APIKeyManaged || m.cfg.APIKey != "" { + s.APIKey = m.cfg.APIKey + } + if m.cfg.SessionSecret != "" { + s.SessionSecret = m.cfg.SessionSecret + } + if m.cfg.InviteSecret != "" { + s.InviteSecret = m.cfg.InviteSecret + } + return s +} + +func (m *Manager) cookieSecure(publicURL string) bool { + if m.cfg.CookieManaged { + return m.cfg.CookieSecure + } + return strings.HasPrefix(publicURL, "https://") +} + +func (m *Manager) replace(s Snapshot) { m.mu.Lock(); m.current = s; m.mu.Unlock() } +func (m *Manager) Peek() (Snapshot, bool) { + m.mu.RLock() + defer m.mu.RUnlock() + return m.current, m.current.version != 0 +} + +// Verify checks the configured server for this operation. A previously saved +// binding alone cannot authorize changes after a server replacement. +func (m *Manager) Verify(ctx context.Context, s Snapshot, token, deviceID string) (Snapshot, error) { + if s.Settings.ServerURL == "" || token == "" { + return Snapshot{}, ErrUnavailable + } + checkCtx, cancel := context.WithTimeout(ctx, 5*time.Second) + defer cancel() + info, err := s.Media.Inspect(checkCtx, s.Settings.ServerURL, token, deviceID) + if err != nil { + return Snapshot{}, err + } + m.updateMu.Lock() + defer m.updateMu.Unlock() + m.mu.RLock() + current := m.current + m.mu.RUnlock() + if current.version != s.version { + return Snapshot{}, db.ErrConnectionChanged + } + if current.Identity.Binding.ID == 0 || current.Identity.Binding.ServerID != info.ID { + identity, err := m.store.PublishMediaConnection(ctx, db.ConnectionUpdate{Origin: db.MediaBinding{Provider: s.Settings.Provider, BaseURL: s.Settings.ServerURL, ServerID: info.ID, Name: info.Name}, ExpectedGeneration: s.Identity.Generation}) + if err != nil { + m.replace(Snapshot{}) + return Snapshot{}, err + } + current.Identity = identity + m.revision++ + current.version = m.revision + m.replace(current) + } + current.Verified = true + return current, nil +} + +// Publish validates a candidate independently before atomically making its +// settings and ownership current. Operations already accepted retain their copy. +func (m *Manager) Publish(ctx context.Context, expected Snapshot, target db.Settings, u db.ConnectionUpdate) (Snapshot, error) { + m.updateMu.Lock() + defer m.updateMu.Unlock() + m.mu.RLock() + current := m.current + m.mu.RUnlock() + if current.version != expected.version { + return Snapshot{}, db.ErrConnectionChanged + } + target = m.effective(target) + provider, ok := mediaserver.ParseProvider(target.Provider) + if !ok { + return Snapshot{}, errors.New("unsupported media provider") + } + baseURL, err := mediaserver.NormalizeBaseURL(provider, target.ServerURL) + if err != nil { + return Snapshot{}, err + } + target.ServerURL = baseURL + media, err := m.factory(provider) + if err != nil { + return Snapshot{}, err + } + origin := db.MediaBinding{Provider: target.Provider, BaseURL: baseURL} + if target.APIKey != "" { + checkCtx, cancel := context.WithTimeout(ctx, 5*time.Second) + info, err := media.Inspect(checkCtx, baseURL, target.APIKey, "aperture") + cancel() + if err != nil { + return Snapshot{}, err + } + origin.ServerID, origin.Name = info.ID, info.Name + } else if current.Settings.Provider == target.Provider && current.Settings.ServerURL == baseURL { + origin = current.Identity.Binding + origin.Provider, origin.BaseURL = target.Provider, baseURL + } + u.Origin = origin + u.ExpectedGeneration = current.Identity.Generation + identity, err := m.store.PublishMediaConnection(ctx, u) + if err != nil { + m.replace(Snapshot{}) + return Snapshot{}, err + } + m.revision++ + updated := Snapshot{Settings: target, Identity: identity, Media: media, CookieSecure: m.cookieSecure(target.PublicURL), version: m.revision} + m.replace(updated) + return updated, nil +} + +type snapshotKey struct{} + +func WithSnapshot(ctx context.Context, s Snapshot) context.Context { + return context.WithValue(ctx, snapshotKey{}, s) +} +func FromContext(ctx context.Context) (Snapshot, bool) { + s, ok := ctx.Value(snapshotKey{}).(Snapshot) + return s, ok +} diff --git a/internal/connection/manager_test.go b/internal/connection/manager_test.go new file mode 100644 index 0000000..603b962 --- /dev/null +++ b/internal/connection/manager_test.go @@ -0,0 +1,185 @@ +package connection + +import ( + "context" + "errors" + "github.com/mayvqt/aperture/internal/config" + "github.com/mayvqt/aperture/internal/db" + "github.com/mayvqt/aperture/internal/mediaserver" + "path/filepath" + "testing" + "time" +) + +type identityMedia struct { + mediaserver.Server + id string + calls int +} + +func (m *identityMedia) Inspect(context.Context, string, string, string) (mediaserver.ServerInfo, error) { + m.calls++ + return mediaserver.ServerInfo{ID: m.id, Name: "Synthetic"}, nil +} + +type acknowledgementStore struct { + *db.Store + fail bool +} + +func (s *acknowledgementStore) PublishMediaConnection(ctx context.Context, u db.ConnectionUpdate) (db.MediaConnection, error) { + v, err := s.Store.PublishMediaConnection(ctx, u) + if err == nil && s.fail { + s.fail = false + return db.MediaConnection{}, errors.New("acknowledgement lost") + } + return v, err +} +func fixture(t *testing.T) (*db.Store, *Manager, *identityMedia) { + t.Helper() + store, err := db.Open(filepath.Join(t.TempDir(), "aperture.db"), "synthetic-encryption-key-at-least-32") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { store.Close() }) + if err := store.InitSchema(t.Context()); err != nil { + t.Fatal(err) + } + for key, value := range map[string]string{"media_provider": "jellyfin", "public_url": "https://join.test", "server_url": "http://media.test", "api_key": "synthetic-key"} { + if err := store.SetSetting(t.Context(), key, value, key == "api_key"); err != nil { + t.Fatal(err) + } + } + media := &identityMedia{id: "server-A"} + manager := New(config.Config{}, store, func(mediaserver.Provider) (mediaserver.Server, error) { return media, nil }) + return store, manager, media +} +func verified(t *testing.T, m *Manager) Snapshot { + t.Helper() + s, err := m.Current(t.Context()) + if err != nil { + t.Fatal(err) + } + s, err = m.Verify(t.Context(), s, "synthetic-token", "device") + if err != nil { + t.Fatal(err) + } + return s +} +func session(t *testing.T, store *db.Store, s Snapshot) string { + t.Helper() + id, _, err := store.CreateSession(t.Context(), db.SessionInput{UserID: "admin", Username: "admin", AccessToken: "secret", DeviceID: "device", TTL: time.Hour, BindingID: s.Identity.Binding.ID, Generation: s.Identity.Generation}) + if err != nil { + t.Fatal(err) + } + return id +} + +func TestOriginChangesRevokeSessionsButKeyRotationPreservesThem(t *testing.T) { + store, m, media := fixture(t) + old := verified(t, m) + id := session(t, store, old) + target := old.Settings + target.APIKey = "rotated-key" + rotated, err := m.Publish(t.Context(), old, target, db.ConnectionUpdate{APIKey: &target.APIKey}) + if err != nil { + t.Fatal(err) + } + if rotated.Identity != old.Identity { + t.Fatal("credential rotation changed ownership") + } + if _, err := store.Session(t.Context(), id); err != nil { + t.Fatal("same-origin session revoked") + } + target.ServerURL = "http://clone.test" + clone, err := m.Publish(t.Context(), rotated, target, db.ConnectionUpdate{ServerURL: &target.ServerURL}) + if err != nil { + t.Fatal(err) + } + if clone.Identity.Binding.ID == old.Identity.Binding.ID || clone.Identity.Binding.ServerID != old.Identity.Binding.ServerID { + t.Fatal("cloned server ID aliased another URL") + } + if _, err := store.Session(t.Context(), id); !errors.Is(err, db.ErrNotFound) { + t.Fatal("old session survived URL change") + } + if old.Settings.ServerURL != "http://media.test" || old.Settings.APIKey != "synthetic-key" { + t.Fatal("accepted operation snapshot changed") + } + if _, err := m.Publish(t.Context(), old, old.Settings, db.ConnectionUpdate{}); !errors.Is(err, db.ErrConnectionChanged) { + t.Fatal("stale settings overwrote new origin") + } + id = session(t, store, clone) + media.id = "replacement-server" + replacement := verified(t, m) + if replacement.Identity.Binding.ID == clone.Identity.Binding.ID { + t.Fatal("replacement reused ownership") + } + if _, err := store.Session(t.Context(), id); !errors.Is(err, db.ErrNotFound) { + t.Fatal("replacement kept old session") + } + if _, _, err := store.CreateSession(t.Context(), db.SessionInput{BindingID: clone.Identity.Binding.ID, Generation: clone.Identity.Generation, TTL: time.Hour}); !errors.Is(err, db.ErrConnectionChanged) { + t.Fatal("in-flight login created old-origin session") + } +} + +func TestEnvironmentOriginIsInvalidatedBeforeNetworkUse(t *testing.T) { + store, m, media := fixture(t) + old := verified(t, m) + id := session(t, store, old) + calls := media.calls + restarted := New(config.Config{ServerURLManaged: true, ServerURL: "http://new-media.test"}, store, func(mediaserver.Provider) (mediaserver.Server, error) { return media, nil }) + s, err := restarted.Current(t.Context()) + if err != nil { + t.Fatal(err) + } + if s.Identity.Binding.ID != 0 || s.Verified || media.calls != calls { + t.Fatal("startup trusted or contacted an unverified origin") + } + if _, err := store.Session(t.Context(), id); !errors.Is(err, db.ErrNotFound) { + t.Fatal("environment move did not revoke session") + } +} + +func TestPublicationReloadsAfterLostAcknowledgement(t *testing.T) { + store, m, _ := fixture(t) + old := verified(t, m) + wrapper := &acknowledgementStore{Store: store, fail: true} + m.store = wrapper + target := old.Settings + target.ServerURL = "http://new.test" + if _, err := m.Publish(t.Context(), old, target, db.ConnectionUpdate{ServerURL: &target.ServerURL}); err == nil { + t.Fatal("missing injected failure") + } + current, err := m.Current(t.Context()) + if err != nil { + t.Fatal(err) + } + if current.Settings.ServerURL != target.ServerURL || current.Identity.Generation <= old.Identity.Generation { + t.Fatal("cached old state survived ambiguous publication") + } + verified(t, m) + if _, err := m.Publish(t.Context(), old, old.Settings, db.ConnectionUpdate{}); !errors.Is(err, db.ErrConnectionChanged) { + t.Fatal("reload reused a stale snapshot version") + } +} + +func TestOptionalAPIKeyCanBindUsingLoginToken(t *testing.T) { + store, m, media := fixture(t) + old := verified(t, m) + target := old.Settings + target.ServerURL = "http://login-only.test" + target.APIKey = "" + calls := media.calls + s, err := m.Publish(t.Context(), old, target, db.ConnectionUpdate{ServerURL: &target.ServerURL, APIKey: &target.APIKey}) + if err != nil { + t.Fatal(err) + } + if s.Identity.Binding.ID != 0 || media.calls != calls { + t.Fatal("keyless setup was rejected or guessed ownership") + } + s, err = m.Verify(t.Context(), s, "administrator-login-token", "login-device") + if err != nil { + t.Fatal(err) + } + session(t, store, s) +} diff --git a/internal/db/account_cleanup.go b/internal/db/account_cleanup.go new file mode 100644 index 0000000..1b706e5 --- /dev/null +++ b/internal/db/account_cleanup.go @@ -0,0 +1,55 @@ +package db + +import ( + "context" + "database/sql" + "fmt" +) + +func migrateAccountCleanup(ctx context.Context, tx *sql.Tx) error { + for _, column := range []struct{ name, definition string }{ + {"cleanup_pending", "INTEGER NOT NULL DEFAULT 0"}, {"cleanup_error", "TEXT"}, + } { + exists, err := schemaColumnExists(ctx, tx, "registrations", column.name) + if err != nil { + return err + } + if !exists { + if _, err := tx.ExecContext(ctx, "ALTER TABLE registrations ADD COLUMN "+column.name+" "+column.definition); err != nil { + return fmt.Errorf("add account cleanup state: %w", err) + } + } + } + // Existing successful deadlines are immutable. Failed legacy registrations + // lacked one; use the original creation time and the retained invite policy. + _, err := tx.ExecContext(ctx, ` + UPDATE registrations SET user_disable_at = datetime(created_at, '+' || + (SELECT user_expiry_days FROM invites WHERE id = registrations.invite_id) || ' days') + WHERE user_disable_at IS NULL AND EXISTS + (SELECT 1 FROM invites WHERE id = registrations.invite_id AND user_expiry_days > 0); + UPDATE registrations SET cleanup_pending = 1, next_disable_attempt_at = NULL + WHERE external_user_id IS NOT NULL AND status IN + ('creating_user','applying_template','retrying_template','needs_attention','failed_create_user','failed_apply_template','pending'); + CREATE INDEX IF NOT EXISTS idx_registrations_cleanup ON registrations(next_disable_attempt_at) WHERE cleanup_pending = 1; + `) + return err +} + +// RequireAccountCleanup preserves the known account ID even when the preceding +// result write failed or timed out. Password-incomplete creation cannot recover +// through a template-only retry. +func (s *Store) RequireAccountCleanup(ctx context.Context, id int64, userID string) error { + result, err := s.db.ExecContext(ctx, ` + UPDATE registrations SET external_user_id = COALESCE(external_user_id, NULLIF(?, '')), + cleanup_pending = 1, next_disable_attempt_at = NULL, + status = CASE WHEN status = 'creating_user' THEN 'failed_create_user' + WHEN status IN ('applying_template','retrying_template') THEN 'needs_attention' ELSE status END, + updated_at = CURRENT_TIMESTAMP + WHERE id = ? AND status NOT IN ('complete','disabled_expired') + AND (external_user_id IS NULL OR external_user_id = ?) + `, userID, id, userID) + if err != nil { + return err + } + return requireSingleTransition(result) +} diff --git a/internal/db/account_cleanup_test.go b/internal/db/account_cleanup_test.go new file mode 100644 index 0000000..22d4cc7 --- /dev/null +++ b/internal/db/account_cleanup_test.go @@ -0,0 +1,277 @@ +package db + +import ( + "context" + "database/sql" + "errors" + "fmt" + "path/filepath" + "testing" + "time" +) + +func TestAccountCleanupSurvivesExhaustedTemplateRetries(t *testing.T) { + ctx, store := testStore(t) + inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "cleanup", TemplateID: 1, MaxUses: 1, UserExpiryDays: 7, BindingID: 1}) + if err != nil { + t.Fatal(err) + } + id, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "", "", "alice") + if err != nil { + t.Fatal(err) + } + reserved, err := store.Registration(ctx, id) + if err != nil { + t.Fatal(err) + } + if !reserved.UserDisableAt.Valid || reserved.UserDisableAt.Time.Sub(reserved.CreatedAt) != 7*24*time.Hour { + t.Fatalf("reservation expiry=%+v", reserved) + } + if err := store.BeginUserCreation(ctx, id); err != nil { + t.Fatal(err) + } + if err := store.RecordProvisioningUser(ctx, id, "alice-id"); err != nil { + t.Fatal(err) + } + if err := store.RecordCreatedUser(ctx, id, "alice-id"); err != nil { + t.Fatal(err) + } + if err := store.CompleteRegistration(ctx, id, RegistrationNeedsAttention, "policy response lost"); err != nil { + t.Fatal(err) + } + for attempt := 0; attempt < 6; attempt++ { + if _, err := store.ClaimTemplateRecovery(ctx, id, 1, false); err != nil { + t.Fatal(err) + } + if err := store.RecordTemplateRetryFailure(ctx, id, "response lost"); err != nil { + t.Fatal(err) + } + } + if ids, err := store.ListDueTemplateRecoveryIDs(ctx, 1, 10); err != nil || len(ids) != 0 { + t.Fatalf("exhausted retries=%v %v", ids, err) + } + if _, err := store.ClaimTemplateRecovery(ctx, id, 1, true); !errors.Is(err, ErrRegistrationTransition) { + t.Fatalf("automatic claim bypassed cap: %v", err) + } + due, err := store.DueUserDisables(ctx, 1, 10) + if err != nil || len(due) != 1 { + t.Fatalf("missing durable cleanup=%+v %v", due, err) + } + if err := store.MarkUserDisabled(ctx, id); err != nil { + t.Fatal(err) + } + clean, err := store.Registration(ctx, id) + if err != nil { + t.Fatal(err) + } + if clean.CleanupPending || clean.UserDisabledAt.Valid || !clean.UserDisableAt.Time.Equal(reserved.UserDisableAt.Time) || clean.Status != RegistrationNeedsAttention { + t.Fatalf("cleanup changed access window: %+v", clean) + } + if _, err := store.ClaimTemplateRecovery(ctx, id, 1, false); err != nil { + t.Fatal(err) + } + if err := store.CompleteTemplateRecovery(ctx, id); err != nil { + t.Fatal(err) + } + complete, err := store.Registration(ctx, id) + if err != nil { + t.Fatal(err) + } + if complete.CleanupPending || !complete.UserDisableAt.Time.Equal(reserved.UserDisableAt.Time) { + t.Fatalf("recovery extended expiry: %+v", complete) + } +} + +func TestInterruptedPasswordSetupHasCleanupButCannotActivate(t *testing.T) { + ctx, store := testStore(t) + inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "password-crash", TemplateID: 1, MaxUses: 1, BindingID: 1}) + if err != nil { + t.Fatal(err) + } + id, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "", "", "alice") + if err != nil { + t.Fatal(err) + } + if err := store.BeginUserCreation(ctx, id); err != nil { + t.Fatal(err) + } + if err := store.RecordProvisioningUser(ctx, id, "partial-id"); err != nil { + t.Fatal(err) + } + if _, err := store.db.ExecContext(ctx, `UPDATE registrations SET updated_at = datetime('now','-1 hour') WHERE id = ?`, id); err != nil { + t.Fatal(err) + } + if _, err := store.ReconcileStaleRegistrations(ctx, time.Now().Add(-15*time.Minute), 10); err != nil { + t.Fatal(err) + } + if _, err := store.ClaimTemplateRecovery(ctx, id, 1, false); !errors.Is(err, ErrRegistrationTransition) { + t.Fatalf("password-incomplete activation: %v", err) + } + due, err := store.DueUserDisables(ctx, 1, 10) + if err != nil || len(due) != 1 || due[0].ExternalUserID.String != "partial-id" { + t.Fatalf("interrupted password cleanup=%+v %v", due, err) + } +} + +func TestCleanupBackoffCannotPostponeAccountExpiry(t *testing.T) { + ctx, store := testStore(t) + inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "cleanup-expiry", TemplateID: 1, MaxUses: 1, BindingID: 1}) + if err != nil { + t.Fatal(err) + } + id, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "", "", "alice") + if err != nil { + t.Fatal(err) + } + if _, err := store.db.ExecContext(ctx, `UPDATE registrations SET external_user_id='alice', status='needs_attention',cleanup_pending=1,disable_attempts=9,user_disable_at=datetime('now','+5 minutes') WHERE id=?`, id); err != nil { + t.Fatal(err) + } + if err := store.MarkUserDisableFailed(ctx, id, "offline"); err != nil { + t.Fatal(err) + } + r, err := store.Registration(ctx, id) + if err != nil { + t.Fatal(err) + } + if !r.NextDisableAttemptAt.Time.Equal(r.UserDisableAt.Time) { + t.Fatalf("cleanup backoff passed expiry: %+v", r) + } +} + +func TestCleanupMigrationPreservesHistoryAndSecrets(t *testing.T) { + for revision := 1; revision <= 5; revision++ { + t.Run(fmt.Sprint(revision), func(t *testing.T) { + store := legacyCleanupStore(t, revision) + ctx := t.Context() + var before string + if err := store.db.QueryRowContext(ctx, `SELECT value FROM settings WHERE key='api_key'`).Scan(&before); err != nil { + t.Fatal(err) + } + if err := store.InitSchema(ctx); err != nil { + t.Fatal(err) + } + if err := store.InitSchema(ctx); err != nil { + t.Fatal(err) + } + var after string + if err := store.db.QueryRowContext(ctx, `SELECT value FROM settings WHERE key='api_key'`).Scan(&after); err != nil { + t.Fatal(err) + } + if after != before { + t.Fatal("migration rewrote encrypted secret") + } + settings, err := store.Settings(ctx) + if err != nil || settings.APIKey != "synthetic-api-key" { + t.Fatalf("secret preservation: %v", err) + } + inv, err := store.Invite(ctx, 1) + if err != nil || inv.Token != "synthetic-invite-token" { + t.Fatalf("invite preservation: %v", err) + } + rows, err := store.RecentRegistrations(ctx, 10) + if err != nil || len(rows) != 2 { + t.Fatalf("history=%+v %v", rows, err) + } + for _, r := range rows { + want := time.Date(2020, 1, 8, 0, 0, 0, 0, time.UTC) + if r.Username == "existing-deadline" { + want = time.Date(2020, 1, 3, 0, 0, 0, 0, time.UTC) + } + if !r.UserDisableAt.Valid || !r.UserDisableAt.Time.Equal(want) || !r.CleanupPending || r.NextDisableAttemptAt.Valid { + t.Fatalf("migration extended or deferred cleanup: %+v", r) + } + } + }) + } +} + +func TestCleanupMigrationRollsBackAsOneTransaction(t *testing.T) { + store := legacyCleanupStore(t, 4) + if _, err := store.db.Exec(`CREATE TRIGGER reject_cleanup BEFORE UPDATE ON registrations BEGIN SELECT RAISE(ABORT,'synthetic migration failure'); END`); err != nil { + t.Fatal(err) + } + if err := store.InitSchema(t.Context()); err == nil { + t.Fatal("expected migration failure") + } + var revision int + if err := store.db.QueryRow(`PRAGMA user_version`).Scan(&revision); err != nil { + t.Fatal(err) + } + exists, err := schemaColumnExists(t.Context(), store.db, "registrations", "cleanup_pending") + if err != nil || exists || revision != 4 { + t.Fatalf("partial migration persisted: column=%v revision=%d err=%v", exists, revision, err) + } + var deadline sql.NullTime + if err := store.db.QueryRow(`SELECT user_disable_at FROM registrations WHERE username='missing-deadline'`).Scan(&deadline); err != nil || deadline.Valid { + t.Fatalf("failed migration altered deadline: %+v %v", deadline, err) + } +} + +func legacyCleanupStore(t *testing.T, revision int) *Store { + t.Helper() + store, err := Open(filepath.Join(t.TempDir(), "aperture.db"), "test-encryption-key-32-characters") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { store.Close() }) + if _, err := store.db.Exec(schema); err != nil { + t.Fatal(err) + } + if revision == 1 { + if _, err := store.db.Exec(`DROP INDEX idx_registrations_template_retry; ALTER TABLE registrations DROP COLUMN template_attempts; ALTER TABLE registrations DROP COLUMN next_template_attempt_at`); err != nil { + t.Fatal(err) + } + } + if revision <= 2 { + if _, err := store.db.Exec(`ALTER TABLE webhooks DROP COLUMN kind; ALTER TABLE webhooks DROP COLUMN role_ids`); err != nil { + t.Fatal(err) + } + } + if revision < 4 { + if _, err := store.db.Exec(`DROP TABLE managed_users`); err != nil { + t.Fatal(err) + } + } + if _, err := store.db.Exec(fmt.Sprintf(`PRAGMA user_version=%d`, revision)); err != nil { + t.Fatal(err) + } + ctx := context.Background() + if err := store.SetSetting(ctx, "api_key", "synthetic-api-key", true); err != nil { + t.Fatal(err) + } + if _, err := store.db.Exec(`INSERT INTO templates(id,name,policy_json,created_at,updated_at) VALUES(1,'Legacy','{}',CURRENT_TIMESTAMP,CURRENT_TIMESTAMP)`); err != nil { + t.Fatal(err) + } + + encrypted, err := store.encryptor.EncryptString("synthetic-invite-token") + if err != nil { + t.Fatal(err) + } + result, err := store.db.Exec(`INSERT INTO invites(token_hash,token_encrypted,template_id,max_uses,user_expiry_days,created_at,updated_at) VALUES('legacy-token-hash',?,1,2,7,CURRENT_TIMESTAMP,CURRENT_TIMESTAMP)`, encrypted) + if err != nil { + t.Fatal(err) + } + id, err := result.LastInsertId() + if err != nil { + t.Fatal(err) + } + if _, err := store.db.Exec(`INSERT INTO registrations(invite_id,username,external_user_id,status,user_disable_at,next_disable_attempt_at,created_at,updated_at) VALUES + (?, 'missing-deadline','first-user','needs_attention',NULL,datetime('now','+30 days'),'2020-01-01 00:00:00','2020-01-01 00:00:00'), + (?, 'existing-deadline','second-user','failed_apply_template','2020-01-03 00:00:00',NULL,'2020-01-01 00:00:00','2020-01-01 00:00:00')`, id, id); err != nil { + t.Fatal(err) + } + if revision == 5 { + tx, err := store.db.BeginTx(ctx, nil) + if err != nil { + t.Fatal(err) + } + if err := migrateAccountCleanup(ctx, tx); err != nil { + tx.Rollback() + t.Fatal(err) + } + if err := tx.Commit(); err != nil { + t.Fatal(err) + } + } + return store +} diff --git a/internal/db/history.go b/internal/db/history.go new file mode 100644 index 0000000..3c400eb --- /dev/null +++ b/internal/db/history.go @@ -0,0 +1,156 @@ +package db + +import ( + "context" + "database/sql" + "errors" + "strings" +) + +func (s *Store) Invite(ctx context.Context, id int64) (Invite, error) { + v, err := s.scanInvite(s.db.QueryRowContext(ctx, `SELECT `+inviteColumns+` FROM invites i JOIN templates t ON t.id=i.template_id WHERE i.id=? AND i.deleted_at IS NULL`, id)) + if errors.Is(err, sql.ErrNoRows) { + return Invite{}, ErrNotFound + } + return v, err +} + +// History pages use id cursors so new arrivals do not shift older pages. +func (s *Store) InvitePage(ctx context.Context, before int64, limit int) ([]Invite, error) { + query := `SELECT ` + inviteColumns + ` FROM invites i JOIN templates t ON t.id=i.template_id WHERE i.deleted_at IS NULL` + args := []any{} + if before > 0 { + query += ` AND i.id < ?` + args = append(args, before) + } + query += ` ORDER BY i.id DESC LIMIT ?` + args = append(args, limit) + rows, err := s.db.QueryContext(ctx, query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + var out []Invite + for rows.Next() { + v, err := s.scanInvite(rows) + if err != nil { + return nil, err + } + out = append(out, v) + } + return out, rows.Err() +} + +func (s *Store) RegistrationPage(ctx context.Context, bindingID, before int64, review bool, limit int) ([]Registration, error) { + query := `SELECT ` + registrationColumns + ` FROM registrations WHERE 1=1` + args := []any{} + if before > 0 { + query += ` AND id < ?` + args = append(args, before) + } + if review { + query += ` AND (COALESCE(binding_id,0)<>? OR cleanup_pending=1 OR status IN ('needs_attention','failed_create_user','failed_apply_template','disable_failed'))` + args = append(args, bindingID) + } + query += ` ORDER BY id DESC LIMIT ?` + args = append(args, limit) + rows, err := s.db.QueryContext(ctx, query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + var out []Registration + for rows.Next() { + v, err := scanRegistration(rows) + if err != nil { + return nil, err + } + out = append(out, v) + } + return out, rows.Err() +} + +func (s *Store) InvitePageActivity(ctx context.Context, ids []int64) (map[int64]InviteActivity, error) { + out := map[int64]InviteActivity{} + if len(ids) == 0 { + return out, nil + } + args := make([]any, len(ids)) + for i, id := range ids { + args[i] = id + } + rows, err := s.db.QueryContext(ctx, `SELECT r.invite_id,r.username,r.status,r.created_at FROM registrations r WHERE r.id IN (SELECT MAX(id) FROM registrations WHERE invite_id IN (`+strings.TrimRight(strings.Repeat("?,", len(ids)), ",")+`) GROUP BY invite_id)`, args...) + if err != nil { + return nil, err + } + defer rows.Close() + for rows.Next() { + var v InviteActivity + if err := rows.Scan(&v.InviteID, &v.Username, &v.Status, &v.CreatedAt); err != nil { + return nil, err + } + out[v.InviteID] = v + } + return out, rows.Err() +} + +func (s *Store) ManagedUserReviewPage(ctx context.Context, bindingID, before int64, limit int) ([]ManagedUser, error) { + query := `SELECT id,external_user_id,username,created_at,updated_at,COALESCE(binding_id,0) FROM managed_users WHERE COALESCE(binding_id,0)<>?` + args := []any{bindingID} + if before > 0 { + query += ` AND id < ?` + args = append(args, before) + } + query += ` ORDER BY id DESC LIMIT ?` + args = append(args, limit) + rows, err := s.db.QueryContext(ctx, query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + var out []ManagedUser + for rows.Next() { + var v ManagedUser + if err := rows.Scan(&v.ID, &v.ExternalUserID, &v.Username, &v.CreatedAt, &v.UpdatedAt, &v.BindingID); err != nil { + return nil, err + } + out = append(out, v) + } + return out, rows.Err() +} +func (s *Store) ManagedUser(ctx context.Context, id int64) (ManagedUser, error) { + var v ManagedUser + err := s.db.QueryRowContext(ctx, `SELECT id,external_user_id,username,created_at,updated_at,COALESCE(binding_id,0) FROM managed_users WHERE id=?`, id).Scan(&v.ID, &v.ExternalUserID, &v.Username, &v.CreatedAt, &v.UpdatedAt, &v.BindingID) + if errors.Is(err, sql.ErrNoRows) { + err = ErrNotFound + } + return v, err +} +func (s *Store) AdoptManagedUser(ctx context.Context, id, bindingID int64) error { + if bindingID <= 0 { + return ErrConnectionChanged + } + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return err + } + defer tx.Rollback() + var externalID string + if err := tx.QueryRowContext(ctx, `SELECT external_user_id FROM managed_users WHERE id=? AND COALESCE(binding_id,0)<>?`, id, bindingID).Scan(&externalID); errors.Is(err, sql.ErrNoRows) { + return ErrRegistrationTransition + } else if err != nil { + return err + } + // Tracking a live account before reviewing its legacy row must not strand + // that history. Keep the reviewed record and the earliest tracking date. + if _, err := tx.ExecContext(ctx, `UPDATE managed_users SET created_at=MIN(created_at,COALESCE((SELECT created_at FROM managed_users WHERE binding_id=? AND external_user_id=?),created_at)) WHERE id=?`, bindingID, externalID, id); err != nil { + return err + } + if _, err := tx.ExecContext(ctx, `DELETE FROM managed_users WHERE binding_id=? AND external_user_id=?`, bindingID, externalID); err != nil { + return err + } + if _, err := tx.ExecContext(ctx, `UPDATE managed_users SET binding_id=?,updated_at=CURRENT_TIMESTAMP WHERE id=?`, bindingID, id); err != nil { + return err + } + return tx.Commit() +} diff --git a/internal/db/invites.go b/internal/db/invites.go index 87d9f64..75207a0 100644 --- a/internal/db/invites.go +++ b/internal/db/invites.go @@ -7,9 +7,12 @@ import ( ) const inviteColumns = `i.id, i.token_hash, COALESCE(i.token_prefix, ''), COALESCE(i.token_encrypted, ''), COALESCE(i.label, ''), i.template_id, t.name, - i.expires_at, i.max_uses, i.uses, i.enabled, i.user_expiry_days, i.last_used_at, i.deleted_at, i.created_at, i.updated_at` + i.expires_at, i.max_uses, i.uses, i.enabled, i.user_expiry_days, i.last_used_at, i.deleted_at, i.created_at, i.updated_at, COALESCE(i.binding_id,0)` func (s *Store) CreateInvite(ctx context.Context, invite Invite) (int64, error) { + if invite.BindingID <= 0 { + return 0, ErrConnectionChanged + } encryptedToken := "" if invite.Token != "" { var err error @@ -19,40 +22,17 @@ func (s *Store) CreateInvite(ctx context.Context, invite Invite) (int64, error) } } result, err := s.db.ExecContext(ctx, ` - INSERT INTO invites (token_hash, token_prefix, token_encrypted, label, template_id, expires_at, max_uses, uses, enabled, user_expiry_days, created_by_user_id, created_at, updated_at) - VALUES (?, ?, NULLIF(?, ''), ?, ?, ?, ?, 0, 1, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) - `, invite.TokenHash, invite.TokenPrefix, encryptedToken, invite.Label, invite.TemplateID, invite.ExpiresAt, invite.MaxUses, invite.UserExpiryDays, invite.CreatedByUserID) + INSERT INTO invites (token_hash, token_prefix, token_encrypted, label, template_id, expires_at, max_uses, uses, enabled, user_expiry_days, created_by_user_id, binding_id, created_at, updated_at) + VALUES (?, ?, NULLIF(?, ''), ?, ?, ?, ?, 0, 1, ?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) + `, invite.TokenHash, invite.TokenPrefix, encryptedToken, invite.Label, invite.TemplateID, invite.ExpiresAt, invite.MaxUses, invite.UserExpiryDays, invite.CreatedByUserID, invite.BindingID) if err != nil { return 0, err } return result.LastInsertId() } -func (s *Store) ListInvites(ctx context.Context) ([]Invite, error) { - rows, err := s.db.QueryContext(ctx, ` - SELECT `+inviteColumns+` - FROM invites i - JOIN templates t ON t.id = i.template_id - WHERE i.deleted_at IS NULL - ORDER BY i.created_at DESC - `) - if err != nil { - return nil, err - } - defer rows.Close() - var invites []Invite - for rows.Next() { - invite, err := s.scanInvite(rows) - if err != nil { - return nil, err - } - invites = append(invites, invite) - } - return invites, rows.Err() -} - func (s *Store) InvitePreview(ctx context.Context, limit int) ([]Invite, error) { rows, err := s.db.QueryContext(ctx, ` - SELECT id, COALESCE(label, ''), expires_at, max_uses, uses, enabled, user_expiry_days, deleted_at + SELECT id, COALESCE(label, ''), expires_at, max_uses, uses, enabled, user_expiry_days, deleted_at, COALESCE(binding_id,0) FROM invites WHERE deleted_at IS NULL ORDER BY created_at DESC @@ -65,7 +45,7 @@ func (s *Store) InvitePreview(ctx context.Context, limit int) ([]Invite, error) var invites []Invite for rows.Next() { var invite Invite - if err := rows.Scan(&invite.ID, &invite.Label, &invite.ExpiresAt, &invite.MaxUses, &invite.Uses, &invite.Enabled, &invite.UserExpiryDays, &invite.DeletedAt); err != nil { + if err := rows.Scan(&invite.ID, &invite.Label, &invite.ExpiresAt, &invite.MaxUses, &invite.Uses, &invite.Enabled, &invite.UserExpiryDays, &invite.DeletedAt, &invite.BindingID); err != nil { return nil, err } invites = append(invites, invite) @@ -90,13 +70,13 @@ func (s *Store) InvitePreset(ctx context.Context, id int64) (Invite, error) { return invite, nil } -func (s *Store) InviteByHash(ctx context.Context, hash string) (Invite, error) { +func (s *Store) InviteByHash(ctx context.Context, hash string, bindingID int64) (Invite, error) { row := s.db.QueryRowContext(ctx, ` SELECT `+inviteColumns+` FROM invites i JOIN templates t ON t.id = i.template_id - WHERE i.token_hash = ? - `, hash) + WHERE i.token_hash = ? AND i.binding_id = ? + `, hash, bindingID) invite, err := s.scanInvite(row) if errors.Is(err, sql.ErrNoRows) { return Invite{}, ErrNotFound @@ -106,7 +86,7 @@ func (s *Store) InviteByHash(ctx context.Context, hash string) (Invite, error) { } return invite, nil } -func (s *Store) ReserveInviteUse(ctx context.Context, inviteID int64, ip, ua, username string) (int64, Template, error) { +func (s *Store) ReserveInviteUse(ctx context.Context, inviteID, bindingID int64, ip, ua, username string) (int64, Template, error) { tx, err := s.db.BeginTx(ctx, nil) if err != nil { return 0, Template{}, err @@ -118,8 +98,8 @@ func (s *Store) ReserveInviteUse(ctx context.Context, inviteID int64, ip, ua, us SELECT t.id, t.name, COALESCE(t.description, ''), t.policy_json, t.is_default, t.created_at, t.updated_at FROM invites i JOIN templates t ON t.id = i.template_id - WHERE i.id = ? - `, inviteID).Scan(&template.ID, &template.Name, &template.Description, &template.PolicyJSON, &template.IsDefault, &template.CreatedAt, &template.UpdatedAt) + WHERE i.id = ? AND i.binding_id = ? + `, inviteID, bindingID).Scan(&template.ID, &template.Name, &template.Description, &template.PolicyJSON, &template.IsDefault, &template.CreatedAt, &template.UpdatedAt) if errors.Is(err, sql.ErrNoRows) { return 0, Template{}, ErrInviteUnavailable } @@ -130,12 +110,12 @@ func (s *Store) ReserveInviteUse(ctx context.Context, inviteID int64, ip, ua, us result, err := tx.ExecContext(ctx, ` UPDATE invites SET uses = uses + 1, last_used_at = CURRENT_TIMESTAMP, updated_at = CURRENT_TIMESTAMP - WHERE id = ? + WHERE id = ? AND binding_id = ? AND enabled = 1 AND deleted_at IS NULL AND uses < max_uses AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP) - `, inviteID) + `, inviteID, bindingID) if err != nil { return 0, Template{}, err } @@ -149,10 +129,11 @@ func (s *Store) ReserveInviteUse(ctx context.Context, inviteID int64, ip, ua, us result, err = tx.ExecContext(ctx, ` INSERT INTO registrations ( invite_id, username, status, template_name, template_policy_json, - ip_address, user_agent, created_at, updated_at + ip_address, user_agent, user_disable_at, binding_id, created_at, updated_at ) - VALUES (?, ?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) - `, inviteID, username, RegistrationReserved, template.Name, template.PolicyJSON, ip, ua) + SELECT ?, ?, ?, ?, ?, ?, ?, CASE WHEN user_expiry_days > 0 THEN datetime(CURRENT_TIMESTAMP, '+' || user_expiry_days || ' days') END, binding_id, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP + FROM invites WHERE id = ? + `, inviteID, username, RegistrationReserved, template.Name, template.PolicyJSON, ip, ua, inviteID) if err != nil { return 0, Template{}, err } @@ -185,12 +166,12 @@ func (s *Store) BeginUserCreation(ctx context.Context, registrationID int64) err return requireSingleTransition(result) } -func (s *Store) SetInviteEnabled(ctx context.Context, id int64, enabled bool) error { +func (s *Store) SetInviteEnabled(ctx context.Context, id, bindingID int64, enabled bool) error { value := 0 if enabled { value = 1 } - result, err := s.db.ExecContext(ctx, `UPDATE invites SET enabled = ?, updated_at = CURRENT_TIMESTAMP WHERE id = ? AND deleted_at IS NULL`, value, id) + result, err := s.db.ExecContext(ctx, `UPDATE invites SET enabled = ?, updated_at = CURRENT_TIMESTAMP WHERE id = ? AND deleted_at IS NULL AND (? = 0 OR binding_id = ?)`, value, id, value, bindingID) if err != nil { return err } @@ -221,7 +202,7 @@ func (s *Store) DeleteInvite(ctx context.Context, id int64) error { func (s *Store) scanInvite(scanner rowScanner) (Invite, error) { var invite Invite var encryptedToken string - err := scanner.Scan(&invite.ID, &invite.TokenHash, &invite.TokenPrefix, &encryptedToken, &invite.Label, &invite.TemplateID, &invite.Template, &invite.ExpiresAt, &invite.MaxUses, &invite.Uses, &invite.Enabled, &invite.UserExpiryDays, &invite.LastUsedAt, &invite.DeletedAt, &invite.CreatedAt, &invite.UpdatedAt) + err := scanner.Scan(&invite.ID, &invite.TokenHash, &invite.TokenPrefix, &encryptedToken, &invite.Label, &invite.TemplateID, &invite.Template, &invite.ExpiresAt, &invite.MaxUses, &invite.Uses, &invite.Enabled, &invite.UserExpiryDays, &invite.LastUsedAt, &invite.DeletedAt, &invite.CreatedAt, &invite.UpdatedAt, &invite.BindingID) if err != nil { return Invite{}, err } diff --git a/internal/db/invites_test.go b/internal/db/invites_test.go index 17c2444..c85e50e 100644 --- a/internal/db/invites_test.go +++ b/internal/db/invites_test.go @@ -13,19 +13,19 @@ func TestReserveInviteUseHonorsMaxUses(t *testing.T) { TokenPrefix: "prefix", Label: "test", TemplateID: 1, - MaxUses: 1, + MaxUses: 1, BindingID: 1, }) if err != nil { t.Fatal(err) } - regID, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "alice") + regID, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "alice") if err != nil { t.Fatal(err) } if regID == 0 { t.Fatal("expected registration id") } - _, _, err = store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "bob") + _, _, err = store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "bob") if !errors.Is(err, ErrInviteUnavailable) { t.Fatalf("expected ErrInviteUnavailable, got %v", err) } @@ -40,7 +40,7 @@ func TestCreateInviteEncryptsRetainedTokenForCopyLinks(t *testing.T) { Label: "test", TemplateID: 1, MaxUses: 1, - CreatedByUserID: "admin-id", + CreatedByUserID: "admin-id", BindingID: 1, }) if err != nil { t.Fatal(err) @@ -62,7 +62,7 @@ func TestCreateInviteEncryptsRetainedTokenForCopyLinks(t *testing.T) { if createdBy != "admin-id" { t.Fatalf("created_by_user_id = %q", createdBy) } - invites, err := store.ListInvites(ctx) + invites, err := store.InvitePage(ctx, 0, 50) if err != nil { t.Fatal(err) } @@ -93,7 +93,7 @@ func TestCreateInviteEncryptsRetainedTokenForCopyLinks(t *testing.T) { func TestInviteStateMutationsReportMissingRows(t *testing.T) { ctx, store := testStore(t) - if err := store.SetInviteEnabled(ctx, 404, false); !errors.Is(err, ErrNotFound) { + if err := store.SetInviteEnabled(ctx, 404, 1, false); !errors.Is(err, ErrNotFound) { t.Fatalf("SetInviteEnabled missing error = %v, want ErrNotFound", err) } if err := store.DeleteInvite(ctx, 404); !errors.Is(err, ErrNotFound) { @@ -108,7 +108,7 @@ func TestDeletedInviteCannotBeReenabledOrDeletedAgain(t *testing.T) { TokenPrefix: "prefix", Label: "test", TemplateID: 1, - MaxUses: 1, + MaxUses: 1, BindingID: 1, }) if err != nil { t.Fatal(err) @@ -119,7 +119,7 @@ func TestDeletedInviteCannotBeReenabledOrDeletedAgain(t *testing.T) { if _, err := store.InvitePreset(ctx, inviteID); !errors.Is(err, ErrNotFound) { t.Fatalf("InvitePreset deleted error = %v, want ErrNotFound", err) } - if err := store.SetInviteEnabled(ctx, inviteID, true); !errors.Is(err, ErrNotFound) { + if err := store.SetInviteEnabled(ctx, inviteID, 1, true); !errors.Is(err, ErrNotFound) { t.Fatalf("SetInviteEnabled deleted error = %v, want ErrNotFound", err) } if err := store.DeleteInvite(ctx, inviteID); !errors.Is(err, ErrNotFound) { diff --git a/internal/db/managed_users.go b/internal/db/managed_users.go index 8f21f96..2f86ff5 100644 --- a/internal/db/managed_users.go +++ b/internal/db/managed_users.go @@ -4,8 +4,8 @@ import ( "context" ) -func (s *Store) ListManagedUsers(ctx context.Context) ([]ManagedUser, error) { - rows, err := s.db.QueryContext(ctx, `SELECT external_user_id, username, created_at, updated_at FROM managed_users ORDER BY username COLLATE NOCASE`) +func (s *Store) ListManagedUsers(ctx context.Context, bindingID int64) ([]ManagedUser, error) { + rows, err := s.db.QueryContext(ctx, `SELECT id,external_user_id, username, created_at, updated_at,COALESCE(binding_id,0) FROM managed_users WHERE binding_id = ? ORDER BY username COLLATE NOCASE`, bindingID) if err != nil { return nil, err } @@ -13,7 +13,7 @@ func (s *Store) ListManagedUsers(ctx context.Context) ([]ManagedUser, error) { var users []ManagedUser for rows.Next() { var user ManagedUser - if err := rows.Scan(&user.ExternalUserID, &user.Username, &user.CreatedAt, &user.UpdatedAt); err != nil { + if err := rows.Scan(&user.ID, &user.ExternalUserID, &user.Username, &user.CreatedAt, &user.UpdatedAt, &user.BindingID); err != nil { return nil, err } users = append(users, user) @@ -22,25 +22,63 @@ func (s *Store) ListManagedUsers(ctx context.Context) ([]ManagedUser, error) { } func (s *Store) SaveManagedUser(ctx context.Context, user ManagedUser) error { + if user.BindingID <= 0 { + return ErrConnectionChanged + } _, err := s.db.ExecContext(ctx, ` - INSERT INTO managed_users (external_user_id, username, created_at, updated_at) - VALUES (?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) - ON CONFLICT(external_user_id) DO UPDATE SET username = excluded.username, updated_at = CURRENT_TIMESTAMP - `, user.ExternalUserID, user.Username) + INSERT INTO managed_users (external_user_id, username, binding_id, created_at, updated_at) + VALUES (?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) + ON CONFLICT(binding_id,external_user_id) DO UPDATE SET username = excluded.username, updated_at = CURRENT_TIMESTAMP + `, user.ExternalUserID, user.Username, user.BindingID) return err } -func (s *Store) DeleteUserRecords(ctx context.Context, externalUserID string) error { +func (s *Store) DeleteUserRecords(ctx context.Context, bindingID int64, externalUserID string) error { tx, err := s.db.BeginTx(ctx, nil) if err != nil { return err } defer tx.Rollback() - if _, err := tx.ExecContext(ctx, `DELETE FROM registrations WHERE external_user_id = ?`, externalUserID); err != nil { + if _, err := tx.ExecContext(ctx, `DELETE FROM registrations WHERE external_user_id = ? AND binding_id = ?`, externalUserID, bindingID); err != nil { return err } - if _, err := tx.ExecContext(ctx, `DELETE FROM managed_users WHERE external_user_id = ?`, externalUserID); err != nil { + if _, err := tx.ExecContext(ctx, `DELETE FROM managed_users WHERE external_user_id = ? AND binding_id = ?`, externalUserID, bindingID); err != nil { return err } return tx.Commit() } + +func (s *Store) UserDeletionRegistrations(ctx context.Context, bindingID int64, id string) ([]Registration, error) { + rows, err := s.db.QueryContext(ctx, `SELECT `+registrationColumns+` FROM registrations WHERE external_user_id = ? AND binding_id = ? ORDER BY id`, id, bindingID) + if err != nil { + return nil, err + } + defer rows.Close() + var regs []Registration + for rows.Next() { + reg, err := scanRegistration(rows) + if err != nil { + return nil, err + } + if IsRegistrationActive(reg.Status) { + return nil, ErrRegistrationTransition + } + regs = append(regs, reg) + } + if err := rows.Err(); err != nil { + return nil, err + } + if err := rows.Close(); err != nil { + return nil, err + } + if len(regs) == 0 { + var tracked bool + if err := s.db.QueryRowContext(ctx, `SELECT EXISTS(SELECT 1 FROM managed_users WHERE external_user_id = ? AND binding_id = ?)`, id, bindingID).Scan(&tracked); err != nil { + return nil, err + } + if !tracked { + return nil, ErrNotFound + } + } + return regs, nil +} diff --git a/internal/db/managed_users_test.go b/internal/db/managed_users_test.go index 3386d16..46a4c40 100644 --- a/internal/db/managed_users_test.go +++ b/internal/db/managed_users_test.go @@ -4,7 +4,7 @@ import "testing" func TestManagedUserRoundTrip(t *testing.T) { ctx, store := testStore(t) - user := ManagedUser{ExternalUserID: "media-1", Username: "Alice"} + user := ManagedUser{ExternalUserID: "media-1", Username: "Alice", BindingID: 1} if err := store.SaveManagedUser(ctx, user); err != nil { t.Fatal(err) } @@ -12,17 +12,17 @@ func TestManagedUserRoundTrip(t *testing.T) { if err := store.SaveManagedUser(ctx, user); err != nil { t.Fatal(err) } - users, err := store.ListManagedUsers(ctx) + users, err := store.ListManagedUsers(ctx, 1) if err != nil { t.Fatal(err) } if len(users) != 1 || users[0].Username != user.Username { t.Fatalf("users = %#v", users) } - if err := store.DeleteUserRecords(ctx, user.ExternalUserID); err != nil { + if err := store.DeleteUserRecords(ctx, 1, user.ExternalUserID); err != nil { t.Fatal(err) } - users, err = store.ListManagedUsers(ctx) + users, err = store.ListManagedUsers(ctx, 1) if err != nil { t.Fatal(err) } diff --git a/internal/db/media_bindings.go b/internal/db/media_bindings.go new file mode 100644 index 0000000..1eaadae --- /dev/null +++ b/internal/db/media_bindings.go @@ -0,0 +1,202 @@ +package db + +import ( + "context" + "database/sql" + "errors" + "fmt" +) + +var ErrConnectionChanged = errors.New("media-server connection changed") + +type MediaBinding struct { + ID int64 + Provider, BaseURL, ServerID, Name string +} + +type MediaConnection struct { + Binding MediaBinding + Provider, BaseURL string + Generation int64 +} + +type ConnectionUpdate struct { + Origin MediaBinding + ExpectedGeneration int64 + Provider, PublicURL, ServerURL, APIKey *string +} + +const bindingSchema = ` +CREATE TABLE IF NOT EXISTS media_bindings ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + provider TEXT NOT NULL, base_url TEXT NOT NULL, server_id TEXT NOT NULL, + name TEXT NOT NULL DEFAULT '', + UNIQUE(provider, base_url, server_id) +); +CREATE TABLE IF NOT EXISTS media_connection ( + id INTEGER PRIMARY KEY CHECK(id = 1), + provider TEXT NOT NULL DEFAULT '', base_url TEXT NOT NULL DEFAULT '', + binding_id INTEGER REFERENCES media_bindings(id), + generation INTEGER NOT NULL DEFAULT 0 +); +INSERT OR IGNORE INTO media_connection(id) VALUES(1); +` + +func migrateMediaBindings(ctx context.Context, tx *sql.Tx) error { + if _, err := tx.ExecContext(ctx, bindingSchema); err != nil { + return err + } + for _, table := range []string{"invites", "registrations", "sessions"} { + if _, err := tx.ExecContext(ctx, "ALTER TABLE "+table+" ADD COLUMN binding_id INTEGER REFERENCES media_bindings(id)"); err != nil { + return fmt.Errorf("bind %s to media origin: %w", table, err) + } + } + _, err := tx.ExecContext(ctx, ` + ALTER TABLE sessions ADD COLUMN connection_generation INTEGER NOT NULL DEFAULT 0; + CREATE TABLE managed_users_bound ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + binding_id INTEGER REFERENCES media_bindings(id), + external_user_id TEXT NOT NULL, username TEXT NOT NULL, + created_at DATETIME NOT NULL, updated_at DATETIME NOT NULL, + UNIQUE(binding_id,external_user_id) + ); + INSERT INTO managed_users_bound(external_user_id,username,created_at,updated_at) + SELECT external_user_id,username,created_at,updated_at FROM managed_users; + DROP TABLE managed_users; + ALTER TABLE managed_users_bound RENAME TO managed_users; + CREATE INDEX idx_registrations_binding_user ON registrations(binding_id,external_user_id,id); + CREATE INDEX idx_registrations_binding_disable ON registrations(binding_id,next_disable_attempt_at); + CREATE INDEX idx_registrations_binding_retry ON registrations(binding_id,status,next_template_attempt_at); + CREATE INDEX idx_invites_binding ON invites(binding_id,id); + `) + return err +} + +type connectionQuerier interface { + QueryRowContext(context.Context, string, ...any) *sql.Row +} + +func readMediaConnection(ctx context.Context, q connectionQuerier) (MediaConnection, error) { + var c MediaConnection + err := q.QueryRowContext(ctx, `SELECT c.provider,c.base_url,c.generation,COALESCE(b.id,0),COALESCE(b.provider,''),COALESCE(b.base_url,''),COALESCE(b.server_id,''),COALESCE(b.name,'') FROM media_connection c LEFT JOIN media_bindings b ON b.id=c.binding_id WHERE c.id=1`).Scan(&c.Provider, &c.BaseURL, &c.Generation, &c.Binding.ID, &c.Binding.Provider, &c.Binding.BaseURL, &c.Binding.ServerID, &c.Binding.Name) + return c, err +} + +func (s *Store) MediaConnection(ctx context.Context) (MediaConnection, error) { + return readMediaConnection(ctx, s.db) +} + +func (s *Store) MediaBindings(ctx context.Context) ([]MediaBinding, error) { + rows, err := s.db.QueryContext(ctx, `SELECT id,provider,base_url,server_id,name FROM media_bindings ORDER BY id`) + if err != nil { + return nil, err + } + defer rows.Close() + var bindings []MediaBinding + for rows.Next() { + var b MediaBinding + if err := rows.Scan(&b.ID, &b.Provider, &b.BaseURL, &b.ServerID, &b.Name); err != nil { + return nil, err + } + bindings = append(bindings, b) + } + return bindings, rows.Err() +} + +// Adoption is an explicit administrator action after reviewing the destination. +// Invites and registrations are independent: one old invite can span servers. +func (s *Store) AdoptInvite(ctx context.Context, id, bindingID int64) error { + if bindingID <= 0 { + return ErrConnectionChanged + } + result, err := s.db.ExecContext(ctx, `UPDATE invites SET binding_id=?,updated_at=CURRENT_TIMESTAMP WHERE id=? AND deleted_at IS NULL AND COALESCE(binding_id,0)<>?`, bindingID, id, bindingID) + if err != nil { + return err + } + return requireSingleTransition(result) +} + +func (s *Store) AdoptRegistration(ctx context.Context, id, bindingID int64) error { + if bindingID <= 0 { + return ErrConnectionChanged + } + result, err := s.db.ExecContext(ctx, `UPDATE registrations SET binding_id=?, user_disabled_at=NULL, cleanup_pending=CASE WHEN external_user_id IS NOT NULL AND (status<>'complete' OR user_disable_at<=CURRENT_TIMESTAMP) THEN 1 ELSE 0 END, cleanup_error=NULL, next_disable_attempt_at=NULL, updated_at=CURRENT_TIMESTAMP WHERE id=? AND COALESCE(binding_id,0)<>? AND status NOT IN ('reserved','creating_user','applying_template','retrying_template','pending')`, bindingID, id, bindingID) + if err != nil { + return err + } + return requireSingleTransition(result) +} + +// PublishMediaConnection commits browser-managed settings, verified origin and +// session generation together. Legacy ownership is deliberately left unknown. +func (s *Store) PublishMediaConnection(ctx context.Context, u ConnectionUpdate) (MediaConnection, error) { + var encryptedKey string + if u.APIKey != nil { + var err error + encryptedKey, _, err = s.settingValue(*u.APIKey, true) + if err != nil { + return MediaConnection{}, err + } + } + s.settingsMu.Lock() + defer s.settingsMu.Unlock() + defer func() { s.settingsCached = false }() + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return MediaConnection{}, err + } + defer tx.Rollback() + current, err := readMediaConnection(ctx, tx) + if err != nil { + return MediaConnection{}, err + } + if current.Generation != u.ExpectedGeneration { + return MediaConnection{}, ErrConnectionChanged + } + for _, setting := range []struct { + key string + value *string + secret int + }{ + {"media_provider", u.Provider, 0}, {"public_url", u.PublicURL, 0}, {"server_url", u.ServerURL, 0}, + } { + if setting.value != nil { + if err := upsertSetting(ctx, tx, setting.key, *setting.value, setting.secret); err != nil { + return MediaConnection{}, err + } + } + } + if u.APIKey != nil { + if err := upsertSetting(ctx, tx, "api_key", encryptedKey, 1); err != nil { + return MediaConnection{}, err + } + } + binding := u.Origin + binding.ID = 0 + if binding.ServerID != "" { + if binding.Provider == "" || binding.BaseURL == "" { + return MediaConnection{}, errors.New("verified media origin is incomplete") + } + if _, err := tx.ExecContext(ctx, `INSERT INTO media_bindings(provider,base_url,server_id,name) VALUES(?,?,?,?) ON CONFLICT(provider,base_url,server_id) DO UPDATE SET name=excluded.name`, binding.Provider, binding.BaseURL, binding.ServerID, binding.Name); err != nil { + return MediaConnection{}, err + } + if err := tx.QueryRowContext(ctx, `SELECT id FROM media_bindings WHERE provider=? AND base_url=? AND server_id=?`, binding.Provider, binding.BaseURL, binding.ServerID).Scan(&binding.ID); err != nil { + return MediaConnection{}, err + } + } + generation := current.Generation + if current.Provider != binding.Provider || current.BaseURL != binding.BaseURL || current.Binding.ID != binding.ID { + generation++ + if _, err := tx.ExecContext(ctx, `DELETE FROM sessions`); err != nil { + return MediaConnection{}, err + } + } + if _, err := tx.ExecContext(ctx, `UPDATE media_connection SET provider=?,base_url=?,binding_id=NULLIF(?,0),generation=? WHERE id=1`, binding.Provider, binding.BaseURL, binding.ID, generation); err != nil { + return MediaConnection{}, err + } + if err := tx.Commit(); err != nil { + return MediaConnection{}, err + } + s.settingsCached = false + return MediaConnection{Binding: binding, Provider: binding.Provider, BaseURL: binding.BaseURL, Generation: generation}, nil +} diff --git a/internal/db/media_bindings_test.go b/internal/db/media_bindings_test.go new file mode 100644 index 0000000..89205c2 --- /dev/null +++ b/internal/db/media_bindings_test.go @@ -0,0 +1,255 @@ +package db + +import ( + "errors" + "fmt" + "testing" + "time" +) + +func TestBindingMigrationQuarantinesAndPreservesLegacyHistory(t *testing.T) { + for revision := 1; revision <= 5; revision++ { + t.Run(fmt.Sprint(revision), func(t *testing.T) { + store := legacyCleanupStore(t, revision) + ctx := t.Context() + var encrypted string + if err := store.db.QueryRow(`SELECT token_encrypted FROM invites WHERE id=1`).Scan(&encrypted); err != nil { + t.Fatal(err) + } + if revision >= 4 { + if _, err := store.db.Exec(`INSERT INTO managed_users(external_user_id,username,created_at,updated_at) VALUES('legacy-user','Legacy','2020-01-01','2020-01-01')`); err != nil { + t.Fatal(err) + } + } + if err := store.InitSchema(ctx); err != nil { + t.Fatal(err) + } + if err := store.InitSchema(ctx); err != nil { + t.Fatal(err) + } + seedTestBinding(t, ctx, store) + v, err := store.Invite(ctx, 1) + if err != nil || v.BindingID != 0 || v.Token != "synthetic-invite-token" { + t.Fatalf("legacy invite assigned/changed: %+v %v", v, err) + } + var after string + if err := store.db.QueryRow(`SELECT token_encrypted FROM invites WHERE id=1`).Scan(&after); err != nil { + t.Fatal(err) + } + if after != encrypted { + t.Fatal("encrypted invite rewritten") + } + if _, err := store.InviteByHash(ctx, "legacy-token-hash", 1); !errors.Is(err, ErrNotFound) { + t.Fatal("unverified invite usable") + } + if due, err := store.DueUserDisables(ctx, 1, 10); err != nil || len(due) != 0 { + t.Fatal("legacy cleanup targeted current server") + } + before, err := store.Registration(ctx, 1) + if err != nil { + t.Fatal(err) + } + if err := store.AdoptRegistration(ctx, 1, 1); err != nil { + t.Fatal(err) + } + adopted, err := store.Registration(ctx, 1) + if err != nil || !adopted.UserDisableAt.Time.Equal(before.UserDisableAt.Time) || adopted.BindingID != 1 { + t.Fatal("adoption changed access deadline") + } + if err := store.AdoptInvite(ctx, 1, 1); err != nil { + t.Fatal(err) + } + if _, err := store.InviteByHash(ctx, "legacy-token-hash", 1); err != nil { + t.Fatal(err) + } + other, err := store.Registration(ctx, 2) + if err != nil || other.BindingID != 0 { + t.Fatal("invite adoption assigned unrelated account") + } + if revision >= 4 { + users, err := store.ManagedUserReviewPage(ctx, 1, 0, 50) + if err != nil || len(users) != 1 { + t.Fatal("lost tracked-only history") + } + before := users[0] + if err := store.AdoptManagedUser(ctx, before.ID, 1); err != nil { + t.Fatal(err) + } + users, err = store.ListManagedUsers(ctx, 1) + if err != nil || len(users) != 1 || !users[0].CreatedAt.Equal(before.CreatedAt) { + t.Fatal("tracked user adoption lost creation history") + } + } + }) + } +} + +func TestBindingMigrationRollbackPreservesLegacySchema(t *testing.T) { + store := legacyCleanupStore(t, 5) + if _, err := store.db.Exec(`CREATE TABLE managed_users_bound(collision INTEGER)`); err != nil { + t.Fatal(err) + } + if err := store.InitSchema(t.Context()); err == nil { + t.Fatal("expected migration failure") + } + exists, err := schemaColumnExists(t.Context(), store.db, "registrations", "binding_id") + if err != nil || exists { + t.Fatal("partial binding columns persisted") + } + var revision int + if err := store.db.QueryRow(`PRAGMA user_version`).Scan(&revision); err != nil || revision != 5 { + t.Fatal("failed migration advanced version") + } + if _, err := store.db.Exec(`DROP TABLE managed_users_bound`); err != nil { + t.Fatal(err) + } + if err := store.InitSchema(t.Context()); err != nil { + t.Fatal(err) + } +} + +func TestBindingsScopeQueriesBeforeLimitsAndGroupings(t *testing.T) { + ctx, store := testStore(t) + first, err := store.CreateInvite(ctx, Invite{BindingID: 1, TokenHash: "first", TemplateID: 1, MaxUses: 40}) + if err != nil { + t.Fatal(err) + } + for i := 0; i < 30; i++ { + id, _, err := store.ReserveInviteUse(ctx, first, 1, "", "", fmt.Sprint(i)) + if err != nil { + t.Fatal(err) + } + if _, err := store.db.Exec(`UPDATE registrations SET status='needs_attention',external_user_id='same-id',cleanup_pending=1 WHERE id=?`, id); err != nil { + t.Fatal(err) + } + } + c, err := store.MediaConnection(ctx) + if err != nil { + t.Fatal(err) + } + c, err = store.PublishMediaConnection(ctx, ConnectionUpdate{ExpectedGeneration: c.Generation, Origin: MediaBinding{Provider: "jellyfin", BaseURL: "http://other.test", ServerID: "synthetic-server"}}) + if err != nil { + t.Fatal(err) + } + second, err := store.CreateInvite(ctx, Invite{BindingID: c.Binding.ID, TokenHash: "second", TemplateID: 1, MaxUses: 2}) + if err != nil { + t.Fatal(err) + } + id, _, err := store.ReserveInviteUse(ctx, second, c.Binding.ID, "", "", "current") + if err != nil { + t.Fatal(err) + } + if _, err := store.db.Exec(`UPDATE registrations SET status='needs_attention',external_user_id='same-id',cleanup_pending=1 WHERE id=?`, id); err != nil { + t.Fatal(err) + } + due, err := store.DueUserDisables(ctx, c.Binding.ID, 1) + if err != nil || len(due) != 1 || due[0].ID != id { + t.Fatal("foreign cleanup starved current work") + } + ids, err := store.ListDueTemplateRecoveryIDs(ctx, c.Binding.ID, 1) + if err != nil || len(ids) != 1 || ids[0] != id { + t.Fatal("foreign recovery starved current work") + } + if _, err := store.ClaimTemplateRecovery(ctx, id, 1, false); !errors.Is(err, ErrNotFound) { + t.Fatal("cross-origin claim accepted") + } + users, err := store.RegistrationUsers(ctx, 1) + if err != nil || len(users) != 1 || users[0].BindingID != 1 { + t.Fatal("foreign maximum hid original account") + } + for _, bindingID := range []int64{1, c.Binding.ID} { + if err := store.SaveManagedUser(ctx, ManagedUser{BindingID: bindingID, ExternalUserID: "same-id", Username: "same"}); err != nil { + t.Fatal(err) + } + } + if err := store.DeleteUserRecords(ctx, c.Binding.ID, "same-id"); err != nil { + t.Fatal(err) + } + tracked, err := store.ListManagedUsers(ctx, 1) + if err != nil || len(tracked) != 1 { + t.Fatal("foreign user collision deleted original tracking") + } + users, err = store.RegistrationUsers(ctx, 1) + if err != nil || len(users) != 1 { + t.Fatal("foreign deletion removed history") + } +} + +func TestAdoptionRechecksExpiredAccessOnNewServer(t *testing.T) { + ctx, store := testStore(t) + id, err := store.CreateInvite(ctx, Invite{BindingID: 1, TokenHash: "expired", TemplateID: 1, MaxUses: 1}) + if err != nil { + t.Fatal(err) + } + regID, _, err := store.ReserveInviteUse(ctx, id, 1, "", "", "alice") + if err != nil { + t.Fatal(err) + } + if _, err := store.db.Exec(`UPDATE registrations SET binding_id=NULL,status='disabled_expired',external_user_id='alice',user_disable_at='2020-01-01',user_disabled_at='2020-01-02',next_disable_attempt_at='2099-01-01' WHERE id=?`, regID); err != nil { + t.Fatal(err) + } + if err := store.AdoptRegistration(ctx, regID, 1); err != nil { + t.Fatal(err) + } + r, err := store.Registration(ctx, regID) + if err != nil || r.UserDisabledAt.Valid || !r.NeedsDisable(time.Now()) { + t.Fatalf("old-server disable acknowledgement reused: %+v %v", r, err) + } +} + +func TestHistoryPagingFindsOldUnverifiedWork(t *testing.T) { + ctx, store := testStore(t) + invite, err := store.CreateInvite(ctx, Invite{BindingID: 1, TokenHash: "history", TemplateID: 1, MaxUses: 110}) + if err != nil { + t.Fatal(err) + } + for i := 0; i < 105; i++ { + id, _, err := store.ReserveInviteUse(ctx, invite, 1, "", "", fmt.Sprint(i)) + if err != nil { + t.Fatal(err) + } + if _, err := store.db.Exec(`UPDATE registrations SET status='complete' WHERE id=?`, id); err != nil { + t.Fatal(err) + } + } + if _, err := store.db.Exec(`UPDATE registrations SET binding_id=NULL WHERE id=1`); err != nil { + t.Fatal(err) + } + review, err := store.RegistrationPage(ctx, 1, 0, true, 50) + if err != nil || len(review) != 1 || review[0].ID != 1 { + t.Fatal("review filter hid old work") + } + first, err := store.RegistrationPage(ctx, 1, 0, false, 50) + if err != nil || len(first) != 50 { + t.Fatal("invalid first page") + } + second, err := store.RegistrationPage(ctx, 1, first[49].ID, false, 50) + if err != nil || len(second) != 50 || second[0].ID >= first[49].ID { + t.Fatal("history cursor repeated or omitted page") + } +} + +func TestTrackedUserAdoptionMergesDuplicateWithoutDeletingAccountHistory(t *testing.T) { + ctx, store := testStore(t) + result, err := store.db.Exec(`INSERT INTO managed_users(external_user_id,username,created_at,updated_at) VALUES('same','Legacy','2020-01-01','2020-01-01')`) + if err != nil { + t.Fatal(err) + } + id, err := result.LastInsertId() + if err != nil { + t.Fatal(err) + } + if err := store.SaveManagedUser(ctx, ManagedUser{BindingID: 1, ExternalUserID: "same", Username: "Current"}); err != nil { + t.Fatal(err) + } + if err := store.AdoptManagedUser(ctx, id, 1); err != nil { + t.Fatal(err) + } + users, err := store.ListManagedUsers(ctx, 1) + if err != nil || len(users) != 1 || users[0].ID != id || users[0].CreatedAt.Year() != 2020 { + t.Fatalf("duplicate adoption lost history: %+v %v", users, err) + } + if pending, err := store.ManagedUserReviewPage(ctx, 1, 0, 50); err != nil || len(pending) != 0 { + t.Fatal("review row stranded") + } +} diff --git a/internal/db/models.go b/internal/db/models.go index a003045..9877f67 100644 --- a/internal/db/models.go +++ b/internal/db/models.go @@ -43,6 +43,7 @@ type Template struct { type Invite struct { ID int64 + BindingID int64 TokenHash string TokenPrefix string Token string @@ -99,8 +100,23 @@ func CanRetryRegistrationTemplate(status string) bool { return status == RegistrationNeedsAttention || status == RegistrationFailedApplyTemplate } +func IsRegistrationActive(status string) bool { + switch status { + case RegistrationReserved, RegistrationCreatingUser, RegistrationApplyingTemplate, RegistrationRetryingTemplate, RegistrationLegacyPending: + return true + } + return false +} + +func (r Registration) NeedsDisable(now time.Time) bool { + return r.ExternalUserID.Valid && !IsRegistrationActive(r.Status) && + (r.CleanupPending || (r.UserDisableAt.Valid && !r.UserDisableAt.Time.After(now) && !r.UserDisabledAt.Valid)) && + (!r.NextDisableAttemptAt.Valid || !r.NextDisableAttemptAt.Time.After(now)) +} + type Registration struct { ID int64 + BindingID int64 InviteID int64 ExternalUserID sql.NullString Username string @@ -112,6 +128,8 @@ type Registration struct { NextDisableAttemptAt sql.NullTime TemplateAttempts int NextTemplateAttemptAt sql.NullTime + CleanupPending bool + CleanupError sql.NullString CreatedAt time.Time UpdatedAt time.Time } @@ -141,6 +159,8 @@ type Webhook struct { } type ManagedUser struct { + ID int64 + BindingID int64 ExternalUserID string Username string CreatedAt time.Time diff --git a/internal/db/operations.go b/internal/db/operations.go index 9fef8de..66f3e70 100644 --- a/internal/db/operations.go +++ b/internal/db/operations.go @@ -2,8 +2,6 @@ package db import ( "context" - "database/sql" - "errors" "time" ) @@ -83,8 +81,8 @@ func (s *Store) DeleteWebhook(ctx context.Context, id int64) error { } return err } -func (s *Store) DueTemplateRecoveries(ctx context.Context, limit int) ([]RegistrationRecovery, error) { - rows, err := s.db.QueryContext(ctx, `SELECT r.id FROM registrations r WHERE r.external_user_id IS NOT NULL AND r.status IN (?,?) AND r.template_attempts < 6 AND (r.next_template_attempt_at IS NULL OR r.next_template_attempt_at <= CURRENT_TIMESTAMP) ORDER BY COALESCE(r.next_template_attempt_at,r.updated_at),r.id LIMIT ?`, RegistrationNeedsAttention, RegistrationFailedApplyTemplate, limit) +func (s *Store) ListDueTemplateRecoveryIDs(ctx context.Context, bindingID int64, limit int) ([]int64, error) { + rows, err := s.db.QueryContext(ctx, `SELECT r.id FROM registrations r WHERE r.binding_id = ? AND r.external_user_id IS NOT NULL AND r.status IN (?,?) AND r.template_attempts < 6 AND (r.next_template_attempt_at IS NULL OR r.next_template_attempt_at <= CURRENT_TIMESTAMP) ORDER BY COALESCE(r.next_template_attempt_at,r.updated_at),r.id LIMIT ?`, bindingID, RegistrationNeedsAttention, RegistrationFailedApplyTemplate, limit) if err != nil { return nil, err } @@ -100,16 +98,5 @@ func (s *Store) DueTemplateRecoveries(ctx context.Context, limit int) ([]Registr if err := rows.Err(); err != nil { return nil, err } - var out []RegistrationRecovery - for _, id := range ids { - v, err := s.ClaimTemplateRecovery(ctx, id) - if errors.Is(err, ErrRegistrationTransition) || errors.Is(err, sql.ErrNoRows) { - continue - } - if err != nil { - return nil, err - } - out = append(out, v) - } - return out, nil + return ids, nil } diff --git a/internal/db/registration_recovery.go b/internal/db/registration_recovery.go index c73b613..0f7f090 100644 --- a/internal/db/registration_recovery.go +++ b/internal/db/registration_recovery.go @@ -78,6 +78,10 @@ func (s *Store) ReconcileStaleRegistrations(ctx context.Context, staleBefore tim } message := "Registration was interrupted after media-server user creation may have begun." + status := RegistrationNeedsAttention + if reg.status == RegistrationCreatingUser || reg.status == RegistrationLegacyPending { + status = RegistrationFailedCreateUser + } if reg.status == RegistrationApplyingTemplate || reg.status == RegistrationRetryingTemplate { message = "Registration was interrupted after media-server user creation and before template application completed." } @@ -87,7 +91,7 @@ func (s *Store) ReconcileStaleRegistrations(ctx context.Context, staleBefore tim error_message = ?, updated_at = CURRENT_TIMESTAMP WHERE id = ? AND status = ? - `, RegistrationNeedsAttention, message, reg.id, reg.status) + `, status, message, reg.id, reg.status) if err != nil { return ReconciliationResult{}, err } @@ -99,7 +103,7 @@ func (s *Store) ReconcileStaleRegistrations(ctx context.Context, staleBefore tim return result, tx.Commit() } -func (s *Store) ClaimTemplateRecovery(ctx context.Context, registrationID int64) (RegistrationRecovery, error) { +func (s *Store) ClaimTemplateRecovery(ctx context.Context, registrationID, bindingID int64, automatic bool) (RegistrationRecovery, error) { tx, err := s.db.BeginTx(ctx, nil) if err != nil { return RegistrationRecovery{}, err @@ -117,13 +121,14 @@ func (s *Store) ClaimTemplateRecovery(ctx context.Context, registrationID int64) SELECT r.id, r.invite_id, r.external_user_id, r.username, r.status, r.error_message, r.user_disable_at, r.user_disabled_at, r.disable_attempts, r.next_disable_attempt_at, r.template_attempts, r.next_template_attempt_at, r.created_at, r.updated_at, + r.cleanup_pending, r.cleanup_error, COALESCE(r.binding_id,0), COALESCE(r.template_name, ''), COALESCE(r.template_policy_json, ''), i.user_expiry_days FROM registrations r JOIN invites i ON i.id = r.invite_id - WHERE r.id = ? - `, registrationID).Scan(destinations...) + WHERE r.id = ? AND r.binding_id = ? + `, registrationID, bindingID).Scan(destinations...) if errors.Is(err, sql.ErrNoRows) { return RegistrationRecovery{}, ErrNotFound } @@ -134,9 +139,12 @@ func (s *Store) ClaimTemplateRecovery(ctx context.Context, registrationID int64) !CanRetryRegistrationTemplate(recovery.Registration.Status) { return RegistrationRecovery{}, ErrRegistrationTransition } + if automatic && (recovery.Registration.TemplateAttempts >= 6 || (recovery.Registration.NextTemplateAttemptAt.Valid && recovery.Registration.NextTemplateAttemptAt.Time.After(time.Now()))) { + return RegistrationRecovery{}, ErrRegistrationTransition + } claimed, err := tx.ExecContext(ctx, ` UPDATE registrations - SET status = ?, updated_at = CURRENT_TIMESTAMP + SET status = ?, cleanup_pending = 1, cleanup_error = NULL, next_disable_attempt_at = NULL, updated_at = CURRENT_TIMESTAMP WHERE id = ? AND external_user_id IS NOT NULL AND status IN (?, ?) `, RegistrationRetryingTemplate, registrationID, RegistrationNeedsAttention, RegistrationFailedApplyTemplate) if err != nil { @@ -157,7 +165,7 @@ func (s *Store) RecordTemplateRetryFailure(ctx context.Context, registrationID i UPDATE registrations SET status = ?, error_message = NULLIF(?, ''), template_attempts = template_attempts + 1, - next_template_attempt_at = datetime(CURRENT_TIMESTAMP, '+' || CASE WHEN template_attempts >= 5 THEN 360 ELSE (1 << template_attempts) END || ' minutes'), + next_template_attempt_at = CASE WHEN template_attempts < 5 THEN datetime(CURRENT_TIMESTAMP, '+' || (1 << template_attempts) || ' minutes') END, updated_at = CURRENT_TIMESTAMP WHERE id = ? AND external_user_id IS NOT NULL @@ -169,16 +177,17 @@ func (s *Store) RecordTemplateRetryFailure(ctx context.Context, registrationID i return requireSingleTransition(result) } -func (s *Store) CompleteTemplateRecovery(ctx context.Context, registrationID int64, disableAt sql.NullTime) error { +func (s *Store) CompleteTemplateRecovery(ctx context.Context, registrationID int64) error { result, err := s.db.ExecContext(ctx, ` UPDATE registrations SET status = ?, error_message = NULL, next_template_attempt_at = NULL, - user_disable_at = ?, next_disable_attempt_at = ?, + cleanup_pending = 0, cleanup_error = NULL, user_disabled_at = NULL, next_disable_attempt_at = user_disable_at, updated_at = CURRENT_TIMESTAMP WHERE id = ? AND external_user_id IS NOT NULL AND status = ? - `, RegistrationComplete, disableAt, disableAt, registrationID, RegistrationRetryingTemplate) + AND (user_disable_at IS NULL OR user_disable_at > CURRENT_TIMESTAMP) + `, RegistrationComplete, registrationID, RegistrationRetryingTemplate) if err != nil { return err } diff --git a/internal/db/registration_recovery_test.go b/internal/db/registration_recovery_test.go index 000c20b..054bda2 100644 --- a/internal/db/registration_recovery_test.go +++ b/internal/db/registration_recovery_test.go @@ -17,11 +17,11 @@ func TestTemplateRecoveryUsesReservationSnapshot(t *testing.T) { }); err != nil { t.Fatal(err) } - inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "hash-recovery", TemplateID: 1, MaxUses: 1, UserExpiryDays: 7}) + inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "hash-recovery", TemplateID: 1, MaxUses: 1, UserExpiryDays: 7, BindingID: 1}) if err != nil { t.Fatal(err) } - registrationID, snapshot, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "alice") + registrationID, snapshot, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "alice") if err != nil { t.Fatal(err) } @@ -34,7 +34,7 @@ func TestTemplateRecoveryUsesReservationSnapshot(t *testing.T) { if err := store.RecordCreatedUser(ctx, registrationID, "jf-alice"); err != nil { t.Fatal(err) } - if err := store.CompleteRegistration(ctx, registrationID, RegistrationNeedsAttention, "policy failed", sql.NullTime{}); err != nil { + if err := store.CompleteRegistration(ctx, registrationID, RegistrationNeedsAttention, "policy failed"); err != nil { t.Fatal(err) } if err := store.UpdateTemplate(ctx, Template{ @@ -46,7 +46,7 @@ func TestTemplateRecoveryUsesReservationSnapshot(t *testing.T) { t.Fatal(err) } - recovery, err := store.ClaimTemplateRecovery(ctx, registrationID) + recovery, err := store.ClaimTemplateRecovery(ctx, registrationID, 1, false) if err != nil { t.Fatal(err) } @@ -56,7 +56,7 @@ func TestTemplateRecoveryUsesReservationSnapshot(t *testing.T) { if recovery.UserExpiryDays != 7 || recovery.Registration.ExternalUserID.String != "jf-alice" { t.Fatalf("recovery metadata = %#v", recovery) } - if _, err := store.ClaimTemplateRecovery(ctx, registrationID); !errors.Is(err, ErrRegistrationTransition) { + if _, err := store.ClaimTemplateRecovery(ctx, registrationID, 1, false); !errors.Is(err, ErrRegistrationTransition) { t.Fatalf("concurrent ClaimTemplateRecovery error = %v", err) } if _, err := store.db.ExecContext(ctx, `UPDATE registrations SET updated_at = datetime('now', '-1 hour') WHERE id = ?`, registrationID); err != nil { @@ -69,37 +69,36 @@ func TestTemplateRecoveryUsesReservationSnapshot(t *testing.T) { if reconciled.FlaggedAmbiguous != 1 { t.Fatalf("retry reconciliation = %#v", reconciled) } - recovery, err = store.ClaimTemplateRecovery(ctx, registrationID) + recovery, err = store.ClaimTemplateRecovery(ctx, registrationID, 1, false) if err != nil { t.Fatal(err) } - disableAt := sql.NullTime{Time: recovery.Registration.CreatedAt.AddDate(0, 0, recovery.UserExpiryDays), Valid: true} - if err := store.CompleteTemplateRecovery(ctx, registrationID, disableAt); err != nil { + if err := store.CompleteTemplateRecovery(ctx, registrationID); err != nil { t.Fatal(err) } - if err := store.CompleteTemplateRecovery(ctx, registrationID, disableAt); !errors.Is(err, ErrRegistrationTransition) { + if err := store.CompleteTemplateRecovery(ctx, registrationID); !errors.Is(err, ErrRegistrationTransition) { t.Fatalf("second CompleteTemplateRecovery error = %v", err) } } func TestReconcileStaleRegistrationsReleasesOnlySafeReservations(t *testing.T) { ctx, store := testStore(t) - inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "hash-reconcile", TemplateID: 1, MaxUses: 4}) + inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "hash-reconcile", TemplateID: 1, MaxUses: 4, BindingID: 1}) if err != nil { t.Fatal(err) } - reservedID, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "reserved") + reservedID, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "reserved") if err != nil { t.Fatal(err) } - creatingID, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "creating") + creatingID, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "creating") if err != nil { t.Fatal(err) } if err := store.BeginUserCreation(ctx, creatingID); err != nil { t.Fatal(err) } - applyingID, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "applying") + applyingID, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "applying") if err != nil { t.Fatal(err) } @@ -109,7 +108,7 @@ func TestReconcileStaleRegistrationsReleasesOnlySafeReservations(t *testing.T) { if err := store.RecordCreatedUser(ctx, applyingID, "jf-applying"); err != nil { t.Fatal(err) } - legacyID, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "legacy") + legacyID, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "legacy") if err != nil { t.Fatal(err) } @@ -156,7 +155,11 @@ func TestReconcileStaleRegistrationsReleasesOnlySafeReservations(t *testing.T) { t.Fatalf("reserved status = %q", statuses[reservedID]) } for _, id := range []int64{creatingID, applyingID, legacyID} { - if statuses[id] != RegistrationNeedsAttention { + want := RegistrationFailedCreateUser + if id == applyingID { + want = RegistrationNeedsAttention + } + if statuses[id] != want { t.Fatalf("ambiguous registration %d status = %q", id, statuses[id]) } } diff --git a/internal/db/registrations.go b/internal/db/registrations.go index 0a0e524..2b9ca40 100644 --- a/internal/db/registrations.go +++ b/internal/db/registrations.go @@ -6,7 +6,20 @@ import ( "errors" ) -const registrationColumns = `id, invite_id, external_user_id, username, status, error_message, user_disable_at, user_disabled_at, disable_attempts, next_disable_attempt_at, template_attempts, next_template_attempt_at, created_at, updated_at` +const registrationColumns = `id, invite_id, external_user_id, username, status, error_message, user_disable_at, user_disabled_at, disable_attempts, next_disable_attempt_at, template_attempts, next_template_attempt_at, created_at, updated_at, cleanup_pending, cleanup_error, COALESCE(binding_id,0)` + +// RecordProvisioningUser is called before a second provider request can begin. +// Password setup is still incomplete and template recovery remains forbidden. +func (s *Store) RecordProvisioningUser(ctx context.Context, id int64, userID string) error { + if userID == "" { + return ErrRegistrationTransition + } + result, err := s.db.ExecContext(ctx, `UPDATE registrations SET external_user_id = ?, cleanup_pending = 1, updated_at = CURRENT_TIMESTAMP WHERE id = ? AND status = ?`, userID, id, RegistrationCreatingUser) + if err != nil { + return err + } + return requireSingleTransition(result) +} func (s *Store) FailUserCreation(ctx context.Context, registrationID int64, message string) error { result, err := s.db.ExecContext(ctx, ` @@ -28,7 +41,7 @@ func (s *Store) RecordFailedUserCreation(ctx context.Context, registrationID int } result, err := s.db.ExecContext(ctx, ` UPDATE registrations - SET external_user_id = ?, status = ?, error_message = NULLIF(?, ''), updated_at = CURRENT_TIMESTAMP + SET external_user_id = ?, status = ?, error_message = NULLIF(?, ''), cleanup_pending = 1, updated_at = CURRENT_TIMESTAMP WHERE id = ? AND status = ? `, externalUserID, RegistrationFailedCreateUser, message, registrationID, RegistrationCreatingUser) if err != nil { @@ -43,7 +56,7 @@ func (s *Store) RecordCreatedUser(ctx context.Context, registrationID int64, ext } result, err := s.db.ExecContext(ctx, ` UPDATE registrations - SET external_user_id = ?, status = ?, error_message = NULL, updated_at = CURRENT_TIMESTAMP + SET external_user_id = ?, status = ?, error_message = NULL, cleanup_pending = 1, updated_at = CURRENT_TIMESTAMP WHERE id = ? AND status = ? `, externalUserID, RegistrationApplyingTemplate, registrationID, RegistrationCreatingUser) if err != nil { @@ -52,34 +65,34 @@ func (s *Store) RecordCreatedUser(ctx context.Context, registrationID int64, ext return requireSingleTransition(result) } -func (s *Store) CompleteRegistration(ctx context.Context, registrationID int64, status, message string, disableAt sql.NullTime) error { +func (s *Store) CompleteRegistration(ctx context.Context, registrationID int64, status, message string) error { if status != RegistrationComplete && status != RegistrationNeedsAttention { return ErrRegistrationTransition } result, err := s.db.ExecContext(ctx, ` UPDATE registrations SET status = ?, error_message = NULLIF(?, ''), - user_disable_at = ?, next_disable_attempt_at = ?, updated_at = CURRENT_TIMESTAMP + cleanup_pending = CASE WHEN ? = 'complete' THEN 0 ELSE 1 END, + cleanup_error = NULL, next_disable_attempt_at = CASE WHEN ? = 'complete' THEN user_disable_at END, updated_at = CURRENT_TIMESTAMP WHERE id = ? AND status = ? - `, status, message, disableAt, disableAt, registrationID, RegistrationApplyingTemplate) + `, status, message, status, status, registrationID, RegistrationApplyingTemplate) if err != nil { return err } return requireSingleTransition(result) } -func (s *Store) DueUserDisables(ctx context.Context, limit int) ([]Registration, error) { +func (s *Store) DueUserDisables(ctx context.Context, bindingID int64, limit int) ([]Registration, error) { rows, err := s.db.QueryContext(ctx, ` SELECT `+registrationColumns+` FROM registrations - WHERE external_user_id IS NOT NULL - AND user_disable_at IS NOT NULL - AND user_disable_at <= CURRENT_TIMESTAMP - AND user_disabled_at IS NULL + WHERE binding_id = ? AND external_user_id IS NOT NULL + AND (cleanup_pending = 1 OR (user_disable_at <= CURRENT_TIMESTAMP AND user_disabled_at IS NULL)) + AND status NOT IN ('reserved','creating_user','applying_template','retrying_template','pending') AND (next_disable_attempt_at IS NULL OR next_disable_attempt_at <= CURRENT_TIMESTAMP) ORDER BY next_disable_attempt_at ASC, user_disable_at ASC LIMIT ? - `, limit) + `, bindingID, limit) if err != nil { return nil, err } @@ -97,12 +110,13 @@ func (s *Store) DueUserDisables(ctx context.Context, limit int) ([]Registration, func (s *Store) MarkUserDisabled(ctx context.Context, registrationID int64) error { result, err := s.db.ExecContext(ctx, ` UPDATE registrations - SET user_disabled_at = CURRENT_TIMESTAMP, status = ?, - error_message = NULL, next_disable_attempt_at = NULL, updated_at = CURRENT_TIMESTAMP + SET user_disabled_at = CASE WHEN user_disable_at <= CURRENT_TIMESTAMP THEN CURRENT_TIMESTAMP ELSE user_disabled_at END, + status = CASE WHEN user_disable_at <= CURRENT_TIMESTAMP THEN ? ELSE status END, + cleanup_pending = 0, cleanup_error = NULL, + next_disable_attempt_at = user_disable_at, updated_at = CURRENT_TIMESTAMP WHERE id = ? AND external_user_id IS NOT NULL - AND user_disable_at IS NOT NULL - AND user_disabled_at IS NULL + AND (cleanup_pending = 1 OR (user_disable_at <= CURRENT_TIMESTAMP AND user_disabled_at IS NULL)) `, RegistrationDisabledExpired, registrationID) if err != nil { return err @@ -112,21 +126,20 @@ func (s *Store) MarkUserDisabled(ctx context.Context, registrationID int64) erro func (s *Store) MarkUserDisableFailed(ctx context.Context, registrationID int64, message string) error { result, err := s.db.ExecContext(ctx, ` UPDATE registrations - SET status = ?, - error_message = NULLIF(?, ''), + SET status = CASE WHEN user_disable_at <= CURRENT_TIMESTAMP THEN ? ELSE status END, + cleanup_error = NULLIF(?, ''), disable_attempts = disable_attempts + 1, - next_disable_attempt_at = datetime( + next_disable_attempt_at = MIN(datetime( CURRENT_TIMESTAMP, '+' || CASE WHEN disable_attempts >= 9 THEN 360 ELSE (1 << disable_attempts) END || ' minutes' - ), + ), COALESCE(CASE WHEN user_disable_at > CURRENT_TIMESTAMP THEN user_disable_at END, '9999-12-31 23:59:59')), updated_at = CURRENT_TIMESTAMP WHERE id = ? AND external_user_id IS NOT NULL - AND user_disable_at IS NOT NULL - AND user_disabled_at IS NULL + AND (cleanup_pending = 1 OR (user_disable_at <= CURRENT_TIMESTAMP AND user_disabled_at IS NULL)) `, RegistrationDisableFailed, message, registrationID) if err != nil { return err @@ -162,14 +175,14 @@ func (s *Store) RecentRegistrations(ctx context.Context, limit int) ([]Registrat return regs, rows.Err() } -func (s *Store) RegistrationUsers(ctx context.Context) ([]Registration, error) { +func (s *Store) RegistrationUsers(ctx context.Context, bindingID int64) ([]Registration, error) { rows, err := s.db.QueryContext(ctx, ` SELECT `+registrationColumns+` FROM registrations WHERE external_user_id IS NOT NULL - AND id IN (SELECT MAX(id) FROM registrations WHERE external_user_id IS NOT NULL GROUP BY external_user_id) + AND binding_id = ? AND id IN (SELECT MAX(id) FROM registrations WHERE binding_id = ? AND external_user_id IS NOT NULL GROUP BY external_user_id) ORDER BY username COLLATE NOCASE - `) + `, bindingID, bindingID) if err != nil { return nil, err } @@ -195,7 +208,7 @@ func (s *Store) Registration(ctx context.Context, id int64) (Registration, error } func (s *Store) DeleteRegistration(ctx context.Context, id int64) error { - result, err := s.db.ExecContext(ctx, `DELETE FROM registrations WHERE id = ?`, id) + result, err := s.db.ExecContext(ctx, `DELETE FROM registrations WHERE id = ? AND external_user_id IS NULL AND status NOT IN ('reserved','creating_user','applying_template','retrying_template','pending')`, id) if err != nil { return err } @@ -204,64 +217,40 @@ func (s *Store) DeleteRegistration(ctx context.Context, id int64) error { return err } if changed != 1 { - return ErrNotFound + return ErrRegistrationTransition } return nil } -func (s *Store) DashboardCounts(ctx context.Context) (DashboardCounts, error) { +func (s *Store) DashboardCounts(ctx context.Context, bindingID int64) (DashboardCounts, error) { var counts DashboardCounts err := s.db.QueryRowContext(ctx, ` SELECT (SELECT COUNT(*) FROM invites - WHERE enabled = 1 + WHERE binding_id = ? AND enabled = 1 AND deleted_at IS NULL AND uses < max_uses AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP)), (SELECT COUNT(*) FROM templates), (SELECT COUNT(*) FROM registrations - WHERE status IN (?, ?, ?, ?)), + WHERE COALESCE(binding_id,0)<>? OR cleanup_pending=1 OR status IN (?, ?, ?, ?)), (SELECT COUNT(*) FROM registrations - WHERE external_user_id IS NOT NULL + WHERE binding_id = ? AND external_user_id IS NOT NULL AND user_disable_at IS NOT NULL AND user_disabled_at IS NULL) `, + bindingID, bindingID, RegistrationNeedsAttention, RegistrationFailedCreateUser, RegistrationFailedApplyTemplate, - RegistrationDisableFailed, + RegistrationDisableFailed, bindingID, ).Scan(&counts.ActiveInvites, &counts.Templates, &counts.NeedsAttention, &counts.ScheduledUserDisables) return counts, err } -func (s *Store) LatestInviteActivity(ctx context.Context) (map[int64]InviteActivity, error) { - rows, err := s.db.QueryContext(ctx, ` - SELECT r.invite_id, r.username, r.status, r.created_at - FROM registrations r - JOIN ( - SELECT invite_id, MAX(id) AS id - FROM registrations - GROUP BY invite_id - ) latest ON latest.id = r.id - `) - if err != nil { - return nil, err - } - defer rows.Close() - activity := map[int64]InviteActivity{} - for rows.Next() { - var item InviteActivity - if err := rows.Scan(&item.InviteID, &item.Username, &item.Status, &item.CreatedAt); err != nil { - return nil, err - } - activity[item.InviteID] = item - } - return activity, rows.Err() -} - func scanRegistration(scanner rowScanner) (Registration, error) { var reg Registration err := scanner.Scan(registrationScanDestinations(®)...) @@ -284,6 +273,9 @@ func registrationScanDestinations(reg *Registration) []any { ®.NextTemplateAttemptAt, ®.CreatedAt, ®.UpdatedAt, + ®.CleanupPending, + ®.CleanupError, + ®.BindingID, } } diff --git a/internal/db/registrations_test.go b/internal/db/registrations_test.go index f45a772..f0713d9 100644 --- a/internal/db/registrations_test.go +++ b/internal/db/registrations_test.go @@ -1,7 +1,6 @@ package db import ( - "database/sql" "errors" "testing" "time" @@ -15,16 +14,16 @@ func TestDueUserDisablesOnlyReturnsDueEnabledUsers(t *testing.T) { Label: "test", TemplateID: 1, MaxUses: 3, - UserExpiryDays: 7, + UserExpiryDays: 7, BindingID: 1, }) if err != nil { t.Fatal(err) } - dueRegID, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "alice") + dueRegID, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "alice") if err != nil { t.Fatal(err) } - futureRegID, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "bob") + futureRegID, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "bob") if err != nil { t.Fatal(err) } @@ -39,14 +38,20 @@ func TestDueUserDisablesOnlyReturnsDueEnabledUsers(t *testing.T) { if err := store.RecordCreatedUser(ctx, futureRegID, "jellyfin-bob"); err != nil { t.Fatal(err) } - if err := store.CompleteRegistration(ctx, dueRegID, RegistrationComplete, "", sql.NullTime{Time: time.Now().Add(-time.Hour).UTC(), Valid: true}); err != nil { + if _, err := store.db.ExecContext(ctx, `UPDATE registrations SET user_disable_at = ?, next_disable_attempt_at = ? WHERE id = ?`, time.Now().Add(-time.Hour).UTC(), time.Now().Add(-time.Hour).UTC(), dueRegID); err != nil { t.Fatal(err) } - if err := store.CompleteRegistration(ctx, futureRegID, RegistrationComplete, "", sql.NullTime{Time: time.Now().Add(time.Hour).UTC(), Valid: true}); err != nil { + if err := store.CompleteRegistration(ctx, dueRegID, RegistrationComplete, ""); err != nil { + t.Fatal(err) + } + if _, err := store.db.ExecContext(ctx, `UPDATE registrations SET user_disable_at = ?, next_disable_attempt_at = ? WHERE id = ?`, time.Now().Add(time.Hour).UTC(), time.Now().Add(time.Hour).UTC(), futureRegID); err != nil { + t.Fatal(err) + } + if err := store.CompleteRegistration(ctx, futureRegID, RegistrationComplete, ""); err != nil { t.Fatal(err) } - due, err := store.DueUserDisables(ctx, 10) + due, err := store.DueUserDisables(ctx, 1, 10) if err != nil { t.Fatal(err) } @@ -56,7 +61,7 @@ func TestDueUserDisablesOnlyReturnsDueEnabledUsers(t *testing.T) { if err := store.MarkUserDisabled(ctx, dueRegID); err != nil { t.Fatal(err) } - due, err = store.DueUserDisables(ctx, 10) + due, err = store.DueUserDisables(ctx, 1, 10) if err != nil { t.Fatal(err) } @@ -70,21 +75,30 @@ func TestDueUserDisablesOnlyReturnsDueEnabledUsers(t *testing.T) { func TestDeleteRegistrationRemovesOnlyTheHistoryRecord(t *testing.T) { ctx, store := testStore(t) - inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "delete-history", TemplateID: 1, MaxUses: 1}) + inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "delete-history", TemplateID: 1, MaxUses: 1, BindingID: 1}) if err != nil { t.Fatal(err) } - registrationID, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "alice") + registrationID, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "alice") if err != nil { t.Fatal(err) } + if err := store.DeleteRegistration(ctx, registrationID); !errors.Is(err, ErrRegistrationTransition) { + t.Fatalf("deleted active reservation: %v", err) + } + if err := store.BeginUserCreation(ctx, registrationID); err != nil { + t.Fatal(err) + } + if err := store.FailUserCreation(ctx, registrationID, "creation rejected"); err != nil { + t.Fatal(err) + } if err := store.DeleteRegistration(ctx, registrationID); err != nil { t.Fatal(err) } if _, err := store.Registration(ctx, registrationID); !errors.Is(err, ErrNotFound) { t.Fatalf("lookup after delete error = %v, want ErrNotFound", err) } - invite, err := store.InviteByHash(ctx, "delete-history") + invite, err := store.InviteByHash(ctx, "delete-history", 1) if err != nil { t.Fatal(err) } @@ -95,11 +109,11 @@ func TestDeleteRegistrationRemovesOnlyTheHistoryRecord(t *testing.T) { func TestDisableFailureUsesDurableExponentialBackoff(t *testing.T) { ctx, store := testStore(t) - inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "hash-backoff", TemplateID: 1, MaxUses: 1}) + inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "hash-backoff", TemplateID: 1, MaxUses: 1, BindingID: 1}) if err != nil { t.Fatal(err) } - registrationID, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "alice") + registrationID, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "alice") if err != nil { t.Fatal(err) } @@ -109,7 +123,10 @@ func TestDisableFailureUsesDurableExponentialBackoff(t *testing.T) { if err := store.RecordCreatedUser(ctx, registrationID, "jf-alice"); err != nil { t.Fatal(err) } - if err := store.CompleteRegistration(ctx, registrationID, RegistrationComplete, "", sql.NullTime{Time: time.Now().Add(-time.Hour), Valid: true}); err != nil { + if _, err := store.db.ExecContext(ctx, `UPDATE registrations SET user_disable_at = ?, next_disable_attempt_at = ? WHERE id = ?`, time.Now().Add(-time.Hour).UTC(), time.Now().Add(-time.Hour).UTC(), registrationID); err != nil { + t.Fatal(err) + } + if err := store.CompleteRegistration(ctx, registrationID, RegistrationComplete, ""); err != nil { t.Fatal(err) } before := time.Now().UTC() @@ -127,7 +144,7 @@ func TestDisableFailureUsesDurableExponentialBackoff(t *testing.T) { if reg.NextDisableAttemptAt.Time.Before(before.Add(30*time.Second)) || reg.NextDisableAttemptAt.Time.After(before.Add(2*time.Minute)) { t.Fatalf("next attempt = %s, want about one minute after failure", reg.NextDisableAttemptAt.Time) } - due, err := store.DueUserDisables(ctx, 10) + due, err := store.DueUserDisables(ctx, 1, 10) if err != nil { t.Fatal(err) } @@ -138,11 +155,11 @@ func TestDisableFailureUsesDurableExponentialBackoff(t *testing.T) { func TestCompleteRegistrationRejectsRepeatedTransition(t *testing.T) { ctx, store := testStore(t) - inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "hash-transition", TemplateID: 1, MaxUses: 1}) + inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "hash-transition", TemplateID: 1, MaxUses: 1, BindingID: 1}) if err != nil { t.Fatal(err) } - registrationID, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "alice") + registrationID, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "alice") if err != nil { t.Fatal(err) } @@ -152,21 +169,21 @@ func TestCompleteRegistrationRejectsRepeatedTransition(t *testing.T) { if err := store.RecordCreatedUser(ctx, registrationID, "jf-alice"); err != nil { t.Fatal(err) } - if err := store.CompleteRegistration(ctx, registrationID, RegistrationComplete, "", sql.NullTime{}); err != nil { + if err := store.CompleteRegistration(ctx, registrationID, RegistrationComplete, ""); err != nil { t.Fatal(err) } - if err := store.CompleteRegistration(ctx, registrationID, RegistrationComplete, "", sql.NullTime{}); !errors.Is(err, ErrRegistrationTransition) { + if err := store.CompleteRegistration(ctx, registrationID, RegistrationComplete, ""); !errors.Is(err, ErrRegistrationTransition) { t.Fatalf("second CompleteRegistration error = %v, want ErrRegistrationTransition", err) } } func TestUncertainUserCreationRetainsInviteUseAndNullUserID(t *testing.T) { ctx, store := testStore(t) - inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "hash-uncertain", TemplateID: 1, MaxUses: 1}) + inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "hash-uncertain", TemplateID: 1, MaxUses: 1, BindingID: 1}) if err != nil { t.Fatal(err) } - registrationID, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "alice") + registrationID, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "alice") if err != nil { t.Fatal(err) } @@ -183,7 +200,7 @@ func TestUncertainUserCreationRetainsInviteUseAndNullUserID(t *testing.T) { if regs[0].ExternalUserID.Valid || regs[0].Status != RegistrationFailedCreateUser { t.Fatalf("uncertain registration = %#v", regs[0]) } - if _, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "bob"); !errors.Is(err, ErrInviteUnavailable) { + if _, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "bob"); !errors.Is(err, ErrInviteUnavailable) { t.Fatalf("second reservation error = %v, want ErrInviteUnavailable", err) } } @@ -195,19 +212,19 @@ func TestLatestInviteActivityReturnsNewestRegistrationPerInvite(t *testing.T) { TokenPrefix: "prefix", Label: "Family", TemplateID: 1, - MaxUses: 5, + MaxUses: 5, BindingID: 1, }) if err != nil { t.Fatal(err) } - if _, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "alice"); err != nil { + if _, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "alice"); err != nil { t.Fatal(err) } - if _, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "bob"); err != nil { + if _, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "bob"); err != nil { t.Fatal(err) } - activity, err := store.LatestInviteActivity(ctx) + activity, err := store.InvitePageActivity(ctx, []int64{inviteID}) if err != nil { t.Fatal(err) } @@ -222,16 +239,16 @@ func TestDashboardCountsCoverFullHistory(t *testing.T) { TokenHash: "hash-dashboard", Label: "Dashboard", TemplateID: 1, - MaxUses: 3, + MaxUses: 3, BindingID: 1, }) if err != nil { t.Fatal(err) } - attentionID, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "alice") + attentionID, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "alice") if err != nil { t.Fatal(err) } - scheduledID, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "bob") + scheduledID, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "bob") if err != nil { t.Fatal(err) } @@ -246,14 +263,17 @@ func TestDashboardCountsCoverFullHistory(t *testing.T) { if err := store.RecordCreatedUser(ctx, scheduledID, "jf-bob"); err != nil { t.Fatal(err) } - if err := store.CompleteRegistration(ctx, attentionID, RegistrationNeedsAttention, "policy failed", sql.NullTime{}); err != nil { + if err := store.CompleteRegistration(ctx, attentionID, RegistrationNeedsAttention, "policy failed"); err != nil { + t.Fatal(err) + } + if _, err := store.db.ExecContext(ctx, `UPDATE registrations SET user_disable_at = ?, next_disable_attempt_at = ? WHERE id = ?`, time.Now().Add(time.Hour).UTC(), time.Now().Add(time.Hour).UTC(), scheduledID); err != nil { t.Fatal(err) } - if err := store.CompleteRegistration(ctx, scheduledID, RegistrationComplete, "", sql.NullTime{Time: time.Now().Add(time.Hour), Valid: true}); err != nil { + if err := store.CompleteRegistration(ctx, scheduledID, RegistrationComplete, ""); err != nil { t.Fatal(err) } - counts, err := store.DashboardCounts(ctx) + counts, err := store.DashboardCounts(ctx, 1) if err != nil { t.Fatal(err) } @@ -264,11 +284,11 @@ func TestDashboardCountsCoverFullHistory(t *testing.T) { func TestPartialUserCreationCannotRetryTemplate(t *testing.T) { ctx, store := testStore(t) - inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "partial-account", TemplateID: 1, MaxUses: 1}) + inviteID, err := store.CreateInvite(ctx, Invite{TokenHash: "partial-account", TemplateID: 1, MaxUses: 1, BindingID: 1}) if err != nil { t.Fatal(err) } - id, _, err := store.ReserveInviteUse(ctx, inviteID, "127.0.0.1", "test", "alice") + id, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "127.0.0.1", "test", "alice") if err != nil { t.Fatal(err) } @@ -291,16 +311,16 @@ func TestPartialUserCreationCannotRetryTemplate(t *testing.T) { if registration.Status != RegistrationFailedCreateUser || registration.ExternalUserID.String != "partial-user" { t.Fatalf("partial account was not durably isolated: status=%q id=%q", registration.Status, registration.ExternalUserID.String) } - if _, err := store.ClaimTemplateRecovery(ctx, id); !errors.Is(err, ErrRegistrationTransition) { + if _, err := store.ClaimTemplateRecovery(ctx, id, 1, false); !errors.Is(err, ErrRegistrationTransition) { t.Fatalf("partial account allowed manual template retry: %v", err) } - if due, err := store.DueTemplateRecoveries(ctx, 10); err != nil || len(due) != 0 { + if due, err := store.ListDueTemplateRecoveryIDs(ctx, 1, 10); err != nil || len(due) != 0 { t.Fatalf("partial account allowed automatic template retry: count=%d err=%v", len(due), err) } if err := store.RecordFailedUserCreation(ctx, id, "other-user", "replacement"); !errors.Is(err, ErrRegistrationTransition) { t.Fatalf("partial account allowed repeated transition: %v", err) } - invite, err := store.InviteByHash(ctx, "partial-account") + invite, err := store.InviteByHash(ctx, "partial-account", 1) if err != nil || invite.Uses != 1 { t.Fatalf("partial creation released invite use: uses=%d err=%v", invite.Uses, err) } diff --git a/internal/db/schema.go b/internal/db/schema.go index dd757a5..612fc79 100644 --- a/internal/db/schema.go +++ b/internal/db/schema.go @@ -6,7 +6,7 @@ import ( "fmt" ) -const schemaRevision = 4 +const schemaRevision = 6 const schema = ` CREATE TABLE IF NOT EXISTS settings ( @@ -143,7 +143,7 @@ func (s *Store) InitSchema(ctx context.Context) error { return fmt.Errorf("inspect database schema: %w", err) } if tables != 0 { - return fmt.Errorf("unsupported database schema; remove the database and restart Aperture") + return fmt.Errorf("unsupported database schema; preserve the database and restore a supported backup before restarting Aperture") } } @@ -193,6 +193,16 @@ func (s *Store) InitSchema(ctx context.Context) error { if _, err := tx.ExecContext(ctx, schema); err != nil { return fmt.Errorf("create schema: %w", err) } + if revision < 5 { + if err := migrateAccountCleanup(ctx, tx); err != nil { + return err + } + } + if revision < 6 { + if err := migrateMediaBindings(ctx, tx); err != nil { + return fmt.Errorf("migrate media ownership: %w", err) + } + } if _, err := tx.ExecContext(ctx, ` INSERT INTO templates (name, description, policy_json, is_default, created_at, updated_at) SELECT 'Default', 'Restricted default template', '{"IsAdministrator":false}', 1, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP @@ -200,7 +210,7 @@ func (s *Store) InitSchema(ctx context.Context) error { `); err != nil { return fmt.Errorf("seed default template: %w", err) } - if _, err := tx.ExecContext(ctx, `PRAGMA user_version = 4`); err != nil { + if _, err := tx.ExecContext(ctx, `PRAGMA user_version = 6`); err != nil { return fmt.Errorf("record schema revision: %w", err) } if err := tx.Commit(); err != nil { diff --git a/internal/db/schema_test.go b/internal/db/schema_test.go index e16e88e..52e3c9a 100644 --- a/internal/db/schema_test.go +++ b/internal/db/schema_test.go @@ -45,23 +45,19 @@ func TestInitSchemaRejectsPreV1Database(t *testing.T) { } err = store.InitSchema(context.Background()) - if err == nil || !strings.Contains(err.Error(), "remove the database") { + if err == nil || !strings.Contains(err.Error(), "preserve the database") { t.Fatalf("legacy database error = %v", err) } } func TestInitSchemaMigratesRevisionThreeWithManagedUsers(t *testing.T) { - ctx, store := testStore(t) - if _, err := store.db.ExecContext(ctx, `DROP TABLE managed_users`); err != nil { - t.Fatal(err) - } - if _, err := store.db.ExecContext(ctx, `PRAGMA user_version = 3`); err != nil { - t.Fatal(err) - } + store := legacyCleanupStore(t, 3) + ctx := t.Context() if err := store.InitSchema(ctx); err != nil { t.Fatal(err) } - if err := store.SaveManagedUser(ctx, ManagedUser{ExternalUserID: "user-1", Username: "Alice"}); err != nil { + seedTestBinding(t, ctx, store) + if err := store.SaveManagedUser(ctx, ManagedUser{ExternalUserID: "user-1", Username: "Alice", BindingID: 1}); err != nil { t.Fatalf("managed_users was not created during migration: %v", err) } } diff --git a/internal/db/sessions.go b/internal/db/sessions.go index d66696b..667a96a 100644 --- a/internal/db/sessions.go +++ b/internal/db/sessions.go @@ -19,9 +19,20 @@ type Session struct { DeviceID string CSRFSecret string ExpiresAt time.Time + BindingID int64 + Generation int64 } -func (s *Store) CreateSession(ctx context.Context, userID, username, accessToken, deviceID string, ttl time.Duration) (string, string, error) { +type SessionInput struct { + UserID, Username, AccessToken, DeviceID string + TTL time.Duration + BindingID, Generation int64 +} + +func (s *Store) CreateSession(ctx context.Context, input SessionInput) (string, string, error) { + if input.BindingID <= 0 { + return "", "", ErrConnectionChanged + } now := time.Now().UTC() sessionID, err := security.RandomToken(32) if err != nil { @@ -32,7 +43,7 @@ func (s *Store) CreateSession(ctx context.Context, userID, username, accessToken return "", "", err } sessionHash := hashSessionID(sessionID) - encryptedAccessToken, err := s.encryptor.EncryptString(accessToken) + encryptedAccessToken, err := s.encryptor.EncryptString(input.AccessToken) if err != nil { return "", "", err } @@ -40,13 +51,19 @@ func (s *Store) CreateSession(ctx context.Context, userID, username, accessToken if err != nil { return "", "", err } - _, err = s.db.ExecContext(ctx, ` - INSERT INTO sessions (id, external_user_id, external_username, access_token, device_id, csrf_secret, expires_at, created_at) - VALUES (?, ?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP) - `, sessionHash, userID, username, encryptedAccessToken, deviceID, encryptedCSRF, now.Add(ttl)) + result, err := s.db.ExecContext(ctx, ` + INSERT INTO sessions (id, external_user_id, external_username, access_token, device_id, csrf_secret, expires_at, binding_id, connection_generation, created_at) + SELECT ?, ?, ?, ?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP + FROM media_connection WHERE id=1 AND binding_id=? AND generation=? + `, sessionHash, input.UserID, input.Username, encryptedAccessToken, input.DeviceID, encryptedCSRF, now.Add(input.TTL), input.BindingID, input.Generation, input.BindingID, input.Generation) if err != nil { return "", "", err } + if changed, err := result.RowsAffected(); err != nil { + return "", "", err + } else if changed != 1 { + return "", "", ErrConnectionChanged + } _, _ = s.db.ExecContext(ctx, ` DELETE FROM sessions WHERE id IN ( @@ -62,10 +79,10 @@ func (s *Store) Session(ctx context.Context, id string) (Session, error) { sessionHash := hashSessionID(id) var session Session err := s.db.QueryRowContext(ctx, ` - SELECT id, external_user_id, external_username, access_token, device_id, csrf_secret, expires_at + SELECT id, external_user_id, external_username, access_token, device_id, csrf_secret, expires_at, COALESCE(binding_id,0),connection_generation FROM sessions WHERE id = ? AND expires_at > CURRENT_TIMESTAMP - `, sessionHash).Scan(&session.ID, &session.UserID, &session.Username, &session.AccessToken, &session.DeviceID, &session.CSRFSecret, &session.ExpiresAt) + `, sessionHash).Scan(&session.ID, &session.UserID, &session.Username, &session.AccessToken, &session.DeviceID, &session.CSRFSecret, &session.ExpiresAt, &session.BindingID, &session.Generation) if errors.Is(err, sql.ErrNoRows) { return Session{}, ErrNotFound } diff --git a/internal/db/sessions_test.go b/internal/db/sessions_test.go index 1be4603..3a4acea 100644 --- a/internal/db/sessions_test.go +++ b/internal/db/sessions_test.go @@ -9,7 +9,7 @@ import ( func TestSessionLifecycleEncryptsTokenAndExpires(t *testing.T) { ctx, store := testStore(t) - sessionID, csrf, err := store.CreateSession(ctx, "jf-user", "admin", "access-token", "device-id", time.Hour) + sessionID, csrf, err := store.CreateSession(ctx, SessionInput{UserID: "jf-user", Username: "admin", AccessToken: "access-token", DeviceID: "device-id", TTL: time.Hour, BindingID: 1, Generation: 1}) if err != nil { t.Fatal(err) } @@ -47,7 +47,7 @@ func TestSessionLifecycleEncryptsTokenAndExpires(t *testing.T) { func TestCreateSessionRemovesExpiredSessions(t *testing.T) { ctx, store := testStore(t) - expiredID, _, err := store.CreateSession(ctx, "old-user", "old-admin", "old-token", "old-device", time.Hour) + expiredID, _, err := store.CreateSession(ctx, SessionInput{UserID: "old-user", Username: "old-admin", AccessToken: "old-token", DeviceID: "old-device", TTL: time.Hour, BindingID: 1, Generation: 1}) if err != nil { t.Fatal(err) } @@ -58,7 +58,7 @@ func TestCreateSessionRemovesExpiredSessions(t *testing.T) { t.Fatalf("expired Session error = %v, want ErrNotFound", err) } - if _, _, err := store.CreateSession(ctx, "new-user", "new-admin", "new-token", "new-device", time.Hour); err != nil { + if _, _, err := store.CreateSession(ctx, SessionInput{UserID: "new-user", Username: "new-admin", AccessToken: "new-token", DeviceID: "new-device", TTL: time.Hour, BindingID: 1, Generation: 1}); err != nil { t.Fatal(err) } var count int @@ -114,13 +114,13 @@ func TestExpiredSessionCleanupIsBoundedAndContinues(t *testing.T) { } return count } - if _, _, err := store.CreateSession(ctx, "new-user", "new-admin", "new-token", "new-device-1", time.Hour); err != nil { + if _, _, err := store.CreateSession(ctx, SessionInput{UserID: "new-user", Username: "new-admin", AccessToken: "new-token", DeviceID: "new-device-1", TTL: time.Hour, BindingID: 1, Generation: 1}); err != nil { t.Fatal(err) } if got := remaining(); got != 1 { t.Fatalf("expired sessions after first bounded cleanup = %d, want 1", got) } - if _, _, err := store.CreateSession(ctx, "new-user", "new-admin", "new-token", "new-device-2", time.Hour); err != nil { + if _, _, err := store.CreateSession(ctx, SessionInput{UserID: "new-user", Username: "new-admin", AccessToken: "new-token", DeviceID: "new-device-2", TTL: time.Hour, BindingID: 1, Generation: 1}); err != nil { t.Fatal(err) } if got := remaining(); got != 0 { diff --git a/internal/db/settings.go b/internal/db/settings.go index bba8927..1d77670 100644 --- a/internal/db/settings.go +++ b/internal/db/settings.go @@ -4,8 +4,6 @@ import ( "context" "database/sql" "errors" - "fmt" - "strings" "github.com/mayvqt/aperture/internal/security" ) @@ -33,17 +31,6 @@ func (s *Store) ensureRuntimeSecret(ctx context.Context, key, configured string) return s.SetSetting(ctx, key, generated, true) } -func (s *Store) ValidateMediaProvider(ctx context.Context, configured string) error { - settings, err := s.Settings(ctx) - if err != nil { - return err - } - saved := strings.TrimSpace(settings.Provider) - if saved != "" && saved != configured { - return fmt.Errorf("configured media provider %q does not match database provider %q; reset the database or restore APERTURE_MEDIA_PROVIDER=%s", configured, saved, saved) - } - return nil -} func (s *Store) Settings(ctx context.Context) (Settings, error) { s.settingsMu.RLock() if s.settingsCached { @@ -133,82 +120,6 @@ func (s *Store) SetSetting(ctx context.Context, key, value string, secret bool) return err } -func (s *Store) UpdateApplicationSettings(ctx context.Context, provider, publicURL, serverURL, apiKey *string) error { - var encryptedAPIKey string - if apiKey != nil { - var err error - encryptedAPIKey, _, err = s.settingValue(*apiKey, true) - if err != nil { - return err - } - } - - s.settingsMu.Lock() - defer s.settingsMu.Unlock() - tx, err := s.db.BeginTx(ctx, nil) - if err != nil { - return err - } - defer tx.Rollback() - if provider != nil { - if err := upsertSetting(ctx, tx, "media_provider", *provider, 0); err != nil { - return err - } - } - if publicURL != nil { - if err := upsertSetting(ctx, tx, "public_url", *publicURL, 0); err != nil { - return err - } - } - if serverURL != nil { - if err := upsertSetting(ctx, tx, "server_url", *serverURL, 0); err != nil { - return err - } - } - if apiKey != nil { - if err := upsertSetting(ctx, tx, "api_key", encryptedAPIKey, 1); err != nil { - return err - } - } - if err := tx.Commit(); err != nil { - return err - } - s.settingsCached = false - return nil -} - -func (s *Store) UpdateSetupSettings(ctx context.Context, provider, publicURL, serverURL, apiKey string) error { - encryptedAPIKey, _, err := s.settingValue(apiKey, true) - if err != nil { - return err - } - s.settingsMu.Lock() - defer s.settingsMu.Unlock() - tx, err := s.db.BeginTx(ctx, nil) - if err != nil { - return err - } - defer tx.Rollback() - values := []struct{ key, value string }{ - {"media_provider", provider}, - {"public_url", publicURL}, - {"server_url", serverURL}, - } - for _, setting := range values { - if err := upsertSetting(ctx, tx, setting.key, setting.value, 0); err != nil { - return err - } - } - if err := upsertSetting(ctx, tx, "api_key", encryptedAPIKey, 1); err != nil { - return err - } - if err := tx.Commit(); err != nil { - return err - } - s.settingsCached = false - return nil -} - func (s *Store) settingValue(value string, secret bool) (string, int, error) { if !secret { return value, 0, nil diff --git a/internal/db/settings_test.go b/internal/db/settings_test.go index 34d2e16..1ac30f5 100644 --- a/internal/db/settings_test.go +++ b/internal/db/settings_test.go @@ -119,7 +119,7 @@ func TestUpdateApplicationSettingsWritesAtomically(t *testing.T) { publicURL := "https://join.example.com" url := "http://media:8096/emby" apiKey := "api-key" - if err := store.UpdateApplicationSettings(ctx, &provider, &publicURL, &url, &apiKey); err != nil { + if _, err := store.PublishMediaConnection(ctx, ConnectionUpdate{ExpectedGeneration: 1, Origin: MediaBinding{Provider: provider, BaseURL: url, ServerID: "emby-server"}, Provider: &provider, PublicURL: &publicURL, ServerURL: &url, APIKey: &apiKey}); err != nil { t.Fatal(err) } settings, err := store.Settings(ctx) @@ -152,7 +152,7 @@ func TestUpdateApplicationSettingsRollsBackPartialWrite(t *testing.T) { } url := "http://must-not-persist:8096" apiKey := "api-key" - if err := store.UpdateApplicationSettings(ctx, nil, nil, &url, &apiKey); err == nil { + if _, err := store.PublishMediaConnection(ctx, ConnectionUpdate{ExpectedGeneration: 1, Origin: MediaBinding{Provider: "jellyfin", BaseURL: url, ServerID: "other-server"}, ServerURL: &url, APIKey: &apiKey}); err == nil { t.Fatal("atomic settings update unexpectedly succeeded") } if _, err := store.Setting(ctx, "server_url"); !errors.Is(err, ErrNotFound) { @@ -160,20 +160,6 @@ func TestUpdateApplicationSettingsRollsBackPartialWrite(t *testing.T) { } } -func TestValidateMediaProviderRejectsProviderSwitch(t *testing.T) { - ctx, store := testStore(t) - provider := "emby" - if err := store.UpdateApplicationSettings(ctx, &provider, nil, nil, nil); err != nil { - t.Fatal(err) - } - if err := store.ValidateMediaProvider(ctx, "jellyfin"); err == nil { - t.Fatal("provider switch was accepted") - } - if err := store.ValidateMediaProvider(ctx, "emby"); err != nil { - t.Fatalf("matching provider rejected: %v", err) - } -} - func TestEnsureRuntimeSecretsGeneratesAndReusesSecrets(t *testing.T) { ctx, store := testStore(t) if err := store.EnsureRuntimeSecrets(ctx, "", ""); err != nil { diff --git a/internal/db/templates.go b/internal/db/templates.go index dba9cf2..4d54b2c 100644 --- a/internal/db/templates.go +++ b/internal/db/templates.go @@ -90,8 +90,13 @@ func (s *Store) SetDefaultTemplate(ctx context.Context, id int64) error { } func (s *Store) DeleteTemplate(ctx context.Context, id int64) error { + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return err + } + defer tx.Rollback() var isDefault bool - if err := s.db.QueryRowContext(ctx, `SELECT is_default FROM templates WHERE id = ?`, id).Scan(&isDefault); errors.Is(err, sql.ErrNoRows) { + if err := tx.QueryRowContext(ctx, `SELECT is_default FROM templates WHERE id = ?`, id).Scan(&isDefault); errors.Is(err, sql.ErrNoRows) { return ErrNotFound } else if err != nil { return err @@ -100,14 +105,17 @@ func (s *Store) DeleteTemplate(ctx context.Context, id int64) error { return ErrTemplateIsDefault } var uses int - if err := s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM invites WHERE template_id = ?`, id).Scan(&uses); err != nil { + if err := tx.QueryRowContext(ctx, `SELECT COUNT(*) FROM invites WHERE template_id = ?`, id).Scan(&uses); err != nil { return err } if uses > 0 { return ErrTemplateInUse } - _, err := s.db.ExecContext(ctx, `DELETE FROM templates WHERE id = ?`, id) - return err + _, err = tx.ExecContext(ctx, `DELETE FROM templates WHERE id = ?`, id) + if err != nil { + return err + } + return tx.Commit() } func scanTemplate(scanner rowScanner) (Template, error) { diff --git a/internal/db/templates_test.go b/internal/db/templates_test.go index cbe5ae6..c775f5b 100644 --- a/internal/db/templates_test.go +++ b/internal/db/templates_test.go @@ -107,7 +107,7 @@ func TestTemplateDefaultAndDeleteRules(t *testing.T) { TokenHash: "hash-second", Label: "invite", TemplateID: secondID, - MaxUses: 1, + MaxUses: 1, BindingID: 1, }); err != nil { t.Fatal(err) } diff --git a/internal/db/test_support_test.go b/internal/db/test_support_test.go index a951c77..3ff9c45 100644 --- a/internal/db/test_support_test.go +++ b/internal/db/test_support_test.go @@ -21,5 +21,22 @@ func testStore(t *testing.T) (context.Context, *Store) { if err := store.InitSchema(ctx); err != nil { t.Fatal(err) } + + seedTestBinding(t, ctx, store) + // Ordinary lifecycle fixtures represent one already verified server. + if _, err := store.db.ExecContext(ctx, `CREATE TRIGGER fixture_registration_binding AFTER INSERT ON registrations WHEN NEW.binding_id IS NULL BEGIN UPDATE registrations SET binding_id=1 WHERE id=NEW.id; END`); err != nil { + t.Fatal(err) + } return ctx, store } + +func seedTestBinding(t *testing.T, ctx context.Context, store *Store) { + t.Helper() + c, err := store.MediaConnection(ctx) + if err != nil { + t.Fatal(err) + } + if _, err := store.PublishMediaConnection(ctx, ConnectionUpdate{ExpectedGeneration: c.Generation, Origin: MediaBinding{Provider: "jellyfin", BaseURL: "http://media:8096", ServerID: "synthetic-server"}}); err != nil { + t.Fatal(err) + } +} diff --git a/internal/httpserver/account_recovery.go b/internal/httpserver/account_recovery.go new file mode 100644 index 0000000..6ae4525 --- /dev/null +++ b/internal/httpserver/account_recovery.go @@ -0,0 +1,210 @@ +package httpserver + +import ( + "context" + "errors" + "log/slog" + "strconv" + "strings" + "time" + + "github.com/mayvqt/aperture/internal/db" + "github.com/mayvqt/aperture/internal/mediaserver" +) + +const accountCleanupTimeout = 45 * time.Second + +// recoverAccount is shared by the administrator action and maintenance. The +// local gate serializes recovery, disable and deletion of each registration; the store claim also +// rejects stale actions. One Aperture process owns each state directory. +func (s *Server) recoverAccount(parent context.Context, id int64, automatic bool) (err error) { + op, err := operationSnapshot(parent) + if err != nil { + return err + } + settings := op.Settings + release, err := s.claimAccountOperations(id) + if err != nil { + return err + } + defer release() + if err := parent.Err(); err != nil { + return err + } + saved, err := s.store.Registration(parent, id) + if err != nil { + return err + } + if saved.BindingID != op.Identity.Binding.ID { + return db.ErrRegistrationTransition + } + releaseUser, err := s.claimMediaUser(op.Identity.Binding.ID, saved.ExternalUserID.String) + if err != nil { + return err + } + defer releaseUser() + recovery, err := s.store.ClaimTemplateRecovery(parent, id, op.Identity.Binding.ID, automatic) + if err != nil { + return err + } + ctx, cancel := context.WithTimeout(context.WithoutCancel(parent), registrationProvisioningTimeout) + defer cancel() + reg := recovery.Registration + if reg.UserDisableAt.Valid && !reg.UserDisableAt.Time.After(time.Now()) { + return s.disableAccount(ctx, reg) + } + succeeded := false + defer func() { + if succeeded { + return + } + persistCtx, persistCancel := context.WithTimeout(context.WithoutCancel(parent), 5*time.Second) + defer persistCancel() + if recordErr := s.store.RecordTemplateRetryFailure(persistCtx, id, safeError(err)); recordErr != nil { + slog.Error("could not record template retry failure", "registration_id", id, "error", safeError(recordErr)) + } + s.secureIncompleteAccount(parent, id, reg.ExternalUserID.String) + if automatic { + description := "Aperture will retry access with backoff." + if reg.TemplateAttempts+1 >= 6 { + description = "Automatic access retries are finished. Review this registration before retrying access." + } + s.notify(webhookNotice{Event: "template.failed", Title: "Template retry failed", Description: description, Color: 0xe67e22, Fields: map[string]string{"Username": reg.Username, "Registration": strconv.FormatInt(id, 10), "Attempt": strconv.Itoa(reg.TemplateAttempts + 1), "Error": safeError(err)}}) + } + }() + if strings.TrimSpace(recovery.Template.PolicyJSON) == "" { + return errors.New("saved registration template is unavailable") + } + if err = op.Media.ApplyTemplate(ctx, settings.ServerURL, settings.APIKey, reg.ExternalUserID.String, recovery.Template); err != nil { + return err + } + if err = s.store.CompleteTemplateRecovery(ctx, id); err != nil && !s.accountCompletionConfirmed(parent, id, reg.ExternalUserID.String) { + return err + } + err = nil + succeeded = true + s.notify(webhookNotice{Event: "template.recovered", Title: "Access template recovered", Color: 0x2ecc71, Fields: map[string]string{"Username": reg.Username, "Registration": strconv.FormatInt(id, 10), "Template": recovery.Template.Name}}) + return nil +} + +func (s *Server) accountCompletionConfirmed(parent context.Context, id int64, userID string) bool { + ctx, cancel := context.WithTimeout(context.WithoutCancel(parent), 5*time.Second) + defer cancel() + reg, err := s.store.Registration(ctx, id) + return err == nil && reg.Status == db.RegistrationComplete && reg.ExternalUserID.Valid && reg.ExternalUserID.String == userID +} + +func (s *Server) claimAccountOperations(ids ...int64) (func(), error) { + s.accountMu.Lock() + defer s.accountMu.Unlock() + if s.closing { + return nil, db.ErrRegistrationTransition + } + if s.activeAccounts == nil { + s.activeAccounts = make(map[int64]bool) + } + for _, id := range ids { + if s.activeAccounts[id] { + return nil, db.ErrRegistrationTransition + } + } + for _, id := range ids { + s.activeAccounts[id] = true + } + s.accountWG.Add(1) + return func() { + s.accountMu.Lock() + for _, id := range ids { + delete(s.activeAccounts, id) + } + s.accountMu.Unlock() + s.accountWG.Done() + }, nil +} + +func (s *Server) secureIncompleteAccount(parent context.Context, id int64, userID string) { + op, err := operationSnapshot(parent) + if err != nil { + return + } + settings := op.Settings + ctx, cancel := context.WithTimeout(context.WithoutCancel(parent), accountCleanupTimeout) + defer cancel() + if err := s.store.RequireAccountCleanup(ctx, id, userID); err != nil { + slog.Error("could not persist incomplete account cleanup", "registration_id", id, "error", safeError(err)) + if errors.Is(err, db.ErrRegistrationTransition) { + return + } + if current, readErr := s.store.Registration(ctx, id); readErr == nil && current.Status == db.RegistrationComplete { + return + } + } + err = op.Media.DisableUser(ctx, settings.ServerURL, settings.APIKey, userID) + persistCtx, persistCancel := context.WithTimeout(context.WithoutCancel(parent), 5*time.Second) + defer persistCancel() + var recordErr error + if err != nil && !userAbsent(err) { + recordErr = s.store.MarkUserDisableFailed(persistCtx, id, safeError(err)) + slog.Warn("incomplete account disable will retry", "registration_id", id, "error", safeError(err)) + s.notify(webhookNotice{Event: "user.disable_failed", Title: "Incomplete account disable failed", Description: "Aperture will retry disabling this account.", Color: 0xe67e22, Fields: map[string]string{"Registration": strconv.FormatInt(id, 10), "Error": safeError(err)}}) + } else { + recordErr = s.store.MarkUserDisabled(persistCtx, id) + } + if recordErr != nil { + slog.Error("could not record account cleanup result", "registration_id", id, "error", safeError(recordErr)) + } +} + +func userAbsent(err error) bool { + var httpErr *mediaserver.HTTPError + return errors.Is(err, mediaserver.ErrUserNotFound) || (errors.As(err, &httpErr) && httpErr.StatusCode == 404) +} + +func (s *Server) disableAccount(parent context.Context, reg db.Registration) error { + op, err := operationSnapshot(parent) + if err != nil { + return err + } + if reg.BindingID != op.Identity.Binding.ID { + return db.ErrConnectionChanged + } + settings := op.Settings + ctx, cancel := context.WithTimeout(parent, accountCleanupTimeout) + defer cancel() + err = op.Media.DisableUser(ctx, settings.ServerURL, settings.APIKey, reg.ExternalUserID.String) + persistCtx, persistCancel := context.WithTimeout(context.WithoutCancel(parent), 5*time.Second) + defer persistCancel() + if err != nil && !userAbsent(err) { + if recordErr := s.store.MarkUserDisableFailed(persistCtx, reg.ID, safeError(err)); recordErr != nil { + slog.Error("could not record account disable failure", "registration_id", reg.ID, "error", safeError(recordErr)) + } + s.notify(webhookNotice{Event: "user.disable_failed", Title: "Account disable failed", Description: "Aperture will retry disabling this account.", Color: 0xe67e22, Fields: map[string]string{"Username": reg.Username, "Registration": strconv.FormatInt(reg.ID, 10), "Error": safeError(err)}}) + return err + } + if err := s.store.MarkUserDisabled(persistCtx, reg.ID); err != nil { + return err + } + title := "Incomplete account disabled" + if reg.UserDisableAt.Valid && !reg.UserDisableAt.Time.After(time.Now()) { + title = "Expired user disabled" + } + s.notify(webhookNotice{Event: "user.disabled", Title: title, Color: 0x2ecc71, Fields: map[string]string{"Username": reg.Username, "Registration": strconv.FormatInt(reg.ID, 10)}}) + return nil +} + +// The server-scoped account gate also covers separately imported/history rows +// that refer to the same upstream account. +func (s *Server) claimMediaUser(bindingID int64, userID string) (func(), error) { + key := strconv.FormatInt(bindingID, 10) + ":" + userID + s.accountMu.Lock() + defer s.accountMu.Unlock() + if s.closing || s.activeUsers[key] { + return nil, db.ErrRegistrationTransition + } + if s.activeUsers == nil { + s.activeUsers = map[string]bool{} + } + s.activeUsers[key] = true + s.accountWG.Add(1) + return func() { s.accountMu.Lock(); delete(s.activeUsers, key); s.accountMu.Unlock(); s.accountWG.Done() }, nil +} diff --git a/internal/httpserver/account_recovery_test.go b/internal/httpserver/account_recovery_test.go new file mode 100644 index 0000000..25fb37f --- /dev/null +++ b/internal/httpserver/account_recovery_test.go @@ -0,0 +1,256 @@ +package httpserver + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "net/url" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/mayvqt/aperture/internal/db" + "github.com/mayvqt/aperture/internal/security" +) + +type acknowledgementFailureStore struct{ Store } + +func (s acknowledgementFailureStore) CompleteTemplateRecovery(ctx context.Context, id int64) error { + if err := s.Store.CompleteTemplateRecovery(ctx, id); err != nil { + return err + } + return errors.New("completion acknowledgement lost") +} + +type pausedReservationStore struct { + *fakeStore + entered, release chan struct{} +} + +func (s *pausedReservationStore) ReserveInviteUse(ctx context.Context, id, bindingID int64, ip, ua, username string) (int64, db.Template, error) { + reg, tmpl, err := s.fakeStore.ReserveInviteUse(ctx, id, bindingID, ip, ua, username) + close(s.entered) + <-s.release + return reg, tmpl, err +} + +func TestShutdownWaitsForReservationAndRejectsLateProvisioning(t *testing.T) { + store := &pausedReservationStore{fakeStore: newFakeStore(), entered: make(chan struct{}), release: make(chan struct{})} + media := &fakeMediaServer{} + s := NewServer(testConfig(), store, testMediaFactory(media)) + form := url.Values{"csrf": {"csrf"}, "username": {"new_user"}, "password": {"correct horse"}, "confirm_password": {"correct horse"}} + req := httptest.NewRequest(http.MethodPost, "/i/token/register", strings.NewReader(form.Encode())) + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + req.AddCookie(&http.Cookie{Name: publicCSRFCookie, Value: security.HashToken(store.settings.InviteSecret, "token:csrf")}) + finished := make(chan struct{}) + go func() { defer close(finished); s.ServeHTTP(httptest.NewRecorder(), req) }() + <-store.entered + s.CloseAdmission() + drained := make(chan error, 1) + go func() { drained <- s.Drain(t.Context()) }() + select { + case <-drained: + t.Fatal("drain returned while a reservation handler still used the store") + case <-time.After(20 * time.Millisecond): + } + close(store.release) + <-finished + if err := <-drained; err != nil { + t.Fatal(err) + } + if media.createdUser || store.beganUserCreation { + t.Fatal("provisioning started after shutdown admission closed") + } + rr := httptest.NewRecorder() + s.ServeHTTP(rr, httptest.NewRequest(http.MethodGet, "/guide", nil)) + if rr.Code != http.StatusServiceUnavailable { + t.Fatalf("new request accepted during shutdown: %d", rr.Code) + } +} + +type recoveryRoundTrip func(*http.Request) (*http.Response, error) + +func (f recoveryRoundTrip) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) } + +func TestRecoveryPreservesFailureAlertsAndCleanupMeaning(t *testing.T) { + notices := make(chan webhookNotice, 4) + previous := http.DefaultTransport + http.DefaultTransport = recoveryRoundTrip(func(r *http.Request) (*http.Response, error) { + var n struct{ Event, Title, Description string } + if err := json.NewDecoder(r.Body).Decode(&n); err != nil { + return nil, err + } + notices <- webhookNotice{Event: n.Event, Title: n.Title, Description: n.Description} + return &http.Response{StatusCode: 204, Body: io.NopCloser(strings.NewReader("")), Header: make(http.Header)}, nil + }) + defer func() { http.DefaultTransport = previous }() + store := newFakeStore() + store.webhooks = []db.Webhook{{URL: "http://webhook.test/events", Kind: "generic", Enabled: true, Events: "template.failed,user.disable_failed,user.disabled"}} + store.recovery.Registration.TemplateAttempts = 5 + media := &fakeMediaServer{applyErr: errors.New("response lost"), disableErr: errors.New("offline")} + s := NewServer(testConfig(), store, testMediaFactory(media)) + if err := s.recoverAccount(verifiedTestContext(t, s), 1, true); err == nil { + t.Fatal("expected template failure") + } + if err := s.waitForWebhooks(t.Context()); err != nil { + t.Fatal(err) + } + seen := map[string]webhookNotice{} + for len(notices) > 0 { + n := <-notices + seen[n.Event] = n + } + if _, ok := seen["user.disable_failed"]; !ok { + t.Fatal("disable failure alert dropped") + } + if n, ok := seen["template.failed"]; !ok || !strings.Contains(n.Description, "finished") { + t.Fatalf("final retry alert=%+v", n) + } + media.disableErr = nil + r := store.recovery.Registration + r.UserDisableAt = sql.NullTime{Time: time.Now().Add(time.Hour), Valid: true} + if err := s.disableAccount(verifiedTestContext(t, s), r); err != nil { + t.Fatal(err) + } + if err := s.waitForWebhooks(t.Context()); err != nil { + t.Fatal(err) + } + if n := <-notices; n.Event != "user.disabled" || n.Title != "Incomplete account disabled" { + t.Fatalf("early cleanup misreported: %+v", n) + } +} + +func TestCommittedRecoveryIsNotDisabledWhenAcknowledgementFails(t *testing.T) { + store, _, id := recoveryFixture(t) + media := &fakeMediaServer{} + s := NewServer(testConfig(), acknowledgementFailureStore{store}, testMediaFactory(media)) + if err := s.recoverAccount(verifiedTestContext(t, s), id, false); err != nil { + t.Fatalf("confirmed completion reported failure: %v", err) + } + r, err := store.Registration(t.Context(), id) + if err != nil { + t.Fatal(err) + } + if r.Status != db.RegistrationComplete || r.CleanupPending || media.disabledUserID != "" { + t.Fatalf("committed recovery was disabled: %+v disabled=%q", r, media.disabledUserID) + } +} + +func TestFailedPolicyResponseLeavesDurableCleanupAfterRetryLimit(t *testing.T) { + store, _, id := recoveryFixture(t) + media := &fakeMediaServer{applyErr: errors.New("policy response lost"), disableErr: errors.New("offline")} + s := NewServer(testConfig(), store, testMediaFactory(media)) + for i := 0; i < 6; i++ { + if err := s.recoverAccount(verifiedTestContext(t, s), id, false); err == nil { + t.Fatal("expected policy failure") + } + } + r, err := store.Registration(t.Context(), id) + if err != nil { + t.Fatal(err) + } + if r.TemplateAttempts != 6 || !r.CleanupPending || r.DisableAttempts != 6 || !r.NextDisableAttemptAt.Valid { + t.Fatalf("cleanup lost: %+v", r) + } + if ids, err := store.ListDueTemplateRecoveryIDs(t.Context(), 1, 10); err != nil || len(ids) > 0 { + t.Fatalf("exhausted retries=%v %v", ids, err) + } + // Retry only the security cleanup using the same persisted registration. + media.disableErr = nil + if err := s.disableAccount(verifiedTestContext(t, s), r); err != nil { + t.Fatal(err) + } + r, err = store.Registration(t.Context(), id) + if err != nil { + t.Fatal(err) + } + if r.CleanupPending || r.Status != db.RegistrationNeedsAttention { + t.Fatalf("cleanup result=%+v", r) + } +} + +func TestExpiredRecoveryOnlyAttemptsDisable(t *testing.T) { + store := newFakeStore() + store.recovery.Registration.UserDisableAt = sql.NullTime{Time: time.Now().Add(-time.Hour), Valid: true} + media := &fakeMediaServer{disableErr: errors.New("offline")} + s := NewServer(testConfig(), store, testMediaFactory(media)) + if err := s.recoverAccount(verifiedTestContext(t, s), 1, false); err == nil { + t.Fatal("expected disable failure") + } + if media.appliedTemplate || store.recoveryFailure != "" || store.failedDisableID != 1 || store.recordedUserID != "" { + t.Fatal("expired recovery ran template failure/compensation path") + } +} + +func TestAccountOperationGuardsAreScopedAndReleased(t *testing.T) { + s := NewServer(testConfig(), newFakeStore(), testMediaFactory(&fakeMediaServer{})) + release, err := s.claimAccountOperations(1, 2) + if err != nil { + t.Fatal(err) + } + if _, err := s.claimAccountOperations(2, 3); !errors.Is(err, db.ErrRegistrationTransition) { + t.Fatalf("overlapping operation accepted: %v", err) + } + other, err := s.claimAccountOperations(3) + if err != nil { + t.Fatalf("independent account blocked: %v", err) + } + other() + release() + release, err = s.claimAccountOperations(1) + if err != nil { + t.Fatal(err) + } + release() + if err := s.Drain(t.Context()); err != nil { + t.Fatal(err) + } +} + +func recoveryFixture(t *testing.T) (*db.Store, db.Settings, int64) { + t.Helper() + store, err := db.Open(filepath.Join(t.TempDir(), "aperture.db"), "test-encryption-key-32-characters") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { store.Close() }) + ctx := t.Context() + if err := store.InitSchema(ctx); err != nil { + t.Fatal(err) + } + for key, value := range map[string]string{"media_provider": "jellyfin", "public_url": "https://join.test", "server_url": "http://media.test", "api_key": "synthetic-key"} { + if err := store.SetSetting(ctx, key, value, key == "api_key"); err != nil { + t.Fatal(err) + } + } + settings, err := store.Settings(ctx) + if err != nil { + t.Fatal(err) + } + if _, err := store.PublishMediaConnection(ctx, db.ConnectionUpdate{Origin: db.MediaBinding{Provider: "jellyfin", BaseURL: "http://media:8096", ServerID: "synthetic-server"}}); err != nil { + t.Fatal(err) + } + inviteID, err := store.CreateInvite(ctx, db.Invite{TokenHash: "retry-fixture", TemplateID: 1, MaxUses: 1, UserExpiryDays: 7, BindingID: 1}) + if err != nil { + t.Fatal(err) + } + id, _, err := store.ReserveInviteUse(ctx, inviteID, 1, "", "", "alice") + if err != nil { + t.Fatal(err) + } + if err := store.BeginUserCreation(ctx, id); err != nil { + t.Fatal(err) + } + if err := store.RecordCreatedUser(ctx, id, "alice-id"); err != nil { + t.Fatal(err) + } + if err := store.CompleteRegistration(ctx, id, db.RegistrationNeedsAttention, "response lost"); err != nil { + t.Fatal(err) + } + return store, settings, id +} diff --git a/internal/httpserver/assets/aperture.css b/internal/httpserver/assets/aperture.css index 73bfb40..d78f162 100644 --- a/internal/httpserver/assets/aperture.css +++ b/internal/httpserver/assets/aperture.css @@ -424,3 +424,31 @@ select { padding: 22px 16px; } } + +.nav-toggle, .nav-toggle[hidden], nav[hidden] { + display: none; +} + +@media (max-width: 1079px) { + header { + display: grid; + grid-template-columns: minmax(0, 1fr) auto; + align-items: center; + } + + .nav-toggle:not([hidden]) { + display: block; + } + + header nav.admin-nav:not([hidden]) { + grid-column: 1 / -1; + display: grid; + grid-template-columns: repeat(2, minmax(0, 1fr)); + overflow: visible; + } + + header nav a, header nav button { + text-align: left; + min-height: 40px; + } +} diff --git a/internal/httpserver/assets/app.js b/internal/httpserver/assets/app.js index 85257a4..2fb8da1 100644 --- a/internal/httpserver/assets/app.js +++ b/internal/httpserver/assets/app.js @@ -54,16 +54,25 @@ document.addEventListener("submit", function (event) { }); var adminNav = document.querySelector("nav.admin-nav"); -if (adminNav) { - var currentNavItem = adminNav.querySelector('a[aria-current="page"]'); - var reducedMotion = window.matchMedia("(prefers-reduced-motion: reduce)"); - if (currentNavItem && adminNav.scrollWidth > adminNav.clientWidth) { - currentNavItem.scrollIntoView({ - block: "nearest", - inline: "nearest", - behavior: reducedMotion.matches ? "auto" : "smooth" - }); +var navToggle = document.querySelector(".nav-toggle"); +if (adminNav && navToggle) { + var compactNav = window.matchMedia("(max-width: 1079px)"); + function setNavOpen(open) { + adminNav.hidden = compactNav.matches && !open; + navToggle.setAttribute("aria-expanded", String(!adminNav.hidden)); } + navToggle.hidden = false; + setNavOpen(false); + navToggle.addEventListener("click", function () { + setNavOpen(adminNav.hidden); + }); + compactNav.addEventListener("change", function () { setNavOpen(false); }); + adminNav.addEventListener("keydown", function (event) { + if (event.key === "Escape" && compactNav.matches) { + setNavOpen(false); + navToggle.focus(); + } + }); } var animatingWorkflows = new WeakSet(); diff --git a/internal/httpserver/assets/components.css b/internal/httpserver/assets/components.css index 1ef6dbd..ddca6f8 100644 --- a/internal/httpserver/assets/components.css +++ b/internal/httpserver/assets/components.css @@ -750,3 +750,7 @@ pre { max-height: 240px; overflow: auto; } + +.server-review p { + overflow-wrap: anywhere; +} diff --git a/internal/httpserver/contracts.go b/internal/httpserver/contracts.go index 3eee98d..0d341c2 100644 --- a/internal/httpserver/contracts.go +++ b/internal/httpserver/contracts.go @@ -2,18 +2,15 @@ package httpserver import ( "context" - "database/sql" "time" + "github.com/mayvqt/aperture/internal/connection" "github.com/mayvqt/aperture/internal/db" - "github.com/mayvqt/aperture/internal/mediaserver" ) type Store interface { - Settings(context.Context) (db.Settings, error) - UpdateApplicationSettings(context.Context, *string, *string, *string, *string) error - UpdateSetupSettings(context.Context, string, string, string, string) error - CreateSession(context.Context, string, string, string, string, time.Duration) (string, string, error) + connection.Store + CreateSession(context.Context, db.SessionInput) (string, string, error) Session(context.Context, string) (db.Session, error) DeleteSession(context.Context, string) error ListTemplates(context.Context) ([]db.Template, error) @@ -23,46 +20,52 @@ type Store interface { SetDefaultTemplate(context.Context, int64) error DeleteTemplate(context.Context, int64) error CreateInvite(context.Context, db.Invite) (int64, error) - ListInvites(context.Context) ([]db.Invite, error) InvitePreview(context.Context, int) ([]db.Invite, error) InvitePreset(context.Context, int64) (db.Invite, error) - InviteByHash(context.Context, string) (db.Invite, error) - ReserveInviteUse(context.Context, int64, string, string, string) (int64, db.Template, error) + InviteByHash(context.Context, string, int64) (db.Invite, error) + ReserveInviteUse(context.Context, int64, int64, string, string, string) (int64, db.Template, error) BeginUserCreation(context.Context, int64) error FailUserCreation(context.Context, int64, string) error RecordFailedUserCreation(context.Context, int64, string, string) error RecordCreatedUser(context.Context, int64, string) error - CompleteRegistration(context.Context, int64, string, string, sql.NullTime) error + RecordProvisioningUser(context.Context, int64, string) error + CompleteRegistration(context.Context, int64, string, string) error + RequireAccountCleanup(context.Context, int64, string) error ReconcileStaleRegistrations(context.Context, time.Time, int) (db.ReconciliationResult, error) - DueUserDisables(context.Context, int) ([]db.Registration, error) + DueUserDisables(context.Context, int64, int) ([]db.Registration, error) MarkUserDisabled(context.Context, int64) error MarkUserDisableFailed(context.Context, int64, string) error - SetInviteEnabled(context.Context, int64, bool) error + SetInviteEnabled(context.Context, int64, int64, bool) error DeleteInvite(context.Context, int64) error Audit(context.Context, string, string, string, string, string, string, string) error RecentRegistrations(context.Context, int) ([]db.Registration, error) - RegistrationUsers(context.Context) ([]db.Registration, error) + RegistrationUsers(context.Context, int64) ([]db.Registration, error) Registration(context.Context, int64) (db.Registration, error) DeleteRegistration(context.Context, int64) error - ListManagedUsers(context.Context) ([]db.ManagedUser, error) + ListManagedUsers(context.Context, int64) ([]db.ManagedUser, error) SaveManagedUser(context.Context, db.ManagedUser) error - DeleteUserRecords(context.Context, string) error - ClaimTemplateRecovery(context.Context, int64) (db.RegistrationRecovery, error) + DeleteUserRecords(context.Context, int64, string) error + ClaimTemplateRecovery(context.Context, int64, int64, bool) (db.RegistrationRecovery, error) + UserDeletionRegistrations(context.Context, int64, string) ([]db.Registration, error) RecordTemplateRetryFailure(context.Context, int64, string) error - CompleteTemplateRecovery(context.Context, int64, sql.NullTime) error - DashboardCounts(context.Context) (db.DashboardCounts, error) - LatestInviteActivity(context.Context) (map[int64]db.InviteActivity, error) + CompleteTemplateRecovery(context.Context, int64) error + DashboardCounts(context.Context, int64) (db.DashboardCounts, error) ListAuditEvents(context.Context, int) ([]db.AuditEvent, error) PruneAuditEvents(context.Context, time.Time, int) (int64, error) ListWebhooks(context.Context) ([]db.Webhook, error) CreateWebhook(context.Context, db.Webhook) (int64, error) DeleteWebhook(context.Context, int64) error - DueTemplateRecoveries(context.Context, int) ([]db.RegistrationRecovery, error) -} - -type MediaServer interface { - mediaserver.Server - SetProvider(mediaserver.Provider) error + ListDueTemplateRecoveryIDs(context.Context, int64, int) ([]int64, error) + Invite(context.Context, int64) (db.Invite, error) + InvitePage(context.Context, int64, int) ([]db.Invite, error) + InvitePageActivity(context.Context, []int64) (map[int64]db.InviteActivity, error) + RegistrationPage(context.Context, int64, int64, bool, int) ([]db.Registration, error) + ManagedUserReviewPage(context.Context, int64, int64, int) ([]db.ManagedUser, error) + ManagedUser(context.Context, int64) (db.ManagedUser, error) + AdoptManagedUser(context.Context, int64, int64) error + MediaBindings(context.Context) ([]db.MediaBinding, error) + AdoptInvite(context.Context, int64, int64) error + AdoptRegistration(context.Context, int64, int64) error } var _ Store = (*db.Store)(nil) diff --git a/internal/httpserver/dashboard.go b/internal/httpserver/dashboard.go index 109e110..f8fd816 100644 --- a/internal/httpserver/dashboard.go +++ b/internal/httpserver/dashboard.go @@ -2,9 +2,7 @@ package httpserver import ( "context" - "database/sql" "log/slog" - "strconv" "strings" "time" @@ -20,7 +18,11 @@ func (s *Server) dashboardData(ctx context.Context) ([]db.Invite, []db.Registrat if err != nil { return nil, nil, db.DashboardCounts{}, err } - counts, err := s.store.DashboardCounts(ctx) + op, err := operationSnapshot(ctx) + if err != nil { + return nil, nil, db.DashboardCounts{}, err + } + counts, err := s.store.DashboardCounts(ctx, op.Identity.Binding.ID) if err != nil { return nil, nil, db.DashboardCounts{}, err } @@ -86,6 +88,9 @@ func (s *Server) dashboardHealth(ctx context.Context, settings db.Settings, coun }) } } + if counts.NeedsAttention > 0 { + checks = append(checks, healthCheck{Level: "warn", Title: "Registrations need review", Detail: "Review incomplete accounts and records saved for an unverified or previous server.", URL: "/admin/registrations?review=1", Action: "Review accounts"}) + } if counts.Templates > 0 && counts.ActiveInvites == 0 { checks = append(checks, healthCheck{ Level: "warn", @@ -98,45 +103,51 @@ func (s *Server) dashboardHealth(ctx context.Context, settings db.Settings, coun return checks } -func disableAtFor(days int) sql.NullTime { - return disableAtFrom(time.Now(), days) -} - -func disableAtFrom(createdAt time.Time, days int) sql.NullTime { - if days <= 0 { - return sql.NullTime{} - } - return sql.NullTime{Time: createdAt.AddDate(0, 0, days).UTC(), Valid: true} -} func (s *Server) processDueUserDisables(ctx context.Context) int { settings, err := s.settings(ctx) if err != nil || settings.ServerURL == "" || settings.APIKey == "" { return 0 } - regs, err := s.store.DueUserDisables(ctx, 25) + op, err := operationSnapshot(ctx) + if err != nil { + return 0 + } + regs, err := s.store.DueUserDisables(ctx, op.Identity.Binding.ID, 25) if err != nil { slog.Warn("could not list due media-server user disables", "error", safeError(err)) return 0 } disabled := 0 for _, reg := range regs { - if !reg.ExternalUserID.Valid { + if ctx.Err() != nil { + break + } + operationCtx, err := s.refreshAccountContext(ctx) + if err != nil { + break + } + release, err := s.claimAccountOperations(reg.ID) + if err != nil { continue } - if err := s.media.DisableUser(ctx, settings.ServerURL, settings.APIKey, reg.ExternalUserID.String); err != nil { - if recordErr := s.store.MarkUserDisableFailed(ctx, reg.ID, safeError(err)); recordErr != nil { - slog.Error("could not record media-server user disable failure", "registration_id", reg.ID, "error", safeError(recordErr)) - } - slog.Warn("media-server user disable failed", "registration_id", reg.ID, "error", safeError(err)) - s.notify(webhookNotice{Event: "user.disable_failed", Title: "Expired user disable failed", Description: "Aperture will retry automatically.", Color: 0xe67e22, Fields: map[string]string{"Username": reg.Username, "Registration": strconv.FormatInt(reg.ID, 10), "Error": safeError(err)}}) + current, err := s.store.Registration(ctx, reg.ID) + if err != nil || !current.NeedsDisable(time.Now()) { + release() continue } - if err := s.store.MarkUserDisabled(ctx, reg.ID); err != nil { - slog.Warn("could not mark media-server user disabled", "registration_id", reg.ID, "error", safeError(err)) + releaseUser, claimErr := s.claimMediaUser(op.Identity.Binding.ID, current.ExternalUserID.String) + if claimErr != nil { + release() + continue + } + err = s.disableAccount(operationCtx, current) + releaseUser() + release() + if err != nil { + slog.Warn("media-server user disable will retry", "registration_id", reg.ID, "error", safeError(err)) continue } disabled++ - s.notify(webhookNotice{Event: "user.disabled", Title: "Expired user disabled", Color: 0x2ecc71, Fields: map[string]string{"Username": reg.Username, "Registration": strconv.FormatInt(reg.ID, 10)}}) } return disabled } @@ -146,22 +157,25 @@ func (s *Server) processDueTemplateRetries(ctx context.Context) { if err != nil || settings.ServerURL == "" || settings.APIKey == "" { return } - recoveries, err := s.store.DueTemplateRecoveries(ctx, 10) + op, err := operationSnapshot(ctx) if err != nil { - slog.Warn("could not claim template retries", "error", safeError(err)) return } - for _, recovery := range recoveries { - reg := recovery.Registration - if err := s.media.ApplyTemplate(ctx, settings.ServerURL, settings.APIKey, reg.ExternalUserID.String, recovery.Template); err != nil { - _ = s.store.RecordTemplateRetryFailure(ctx, reg.ID, safeError(err)) - s.notify(webhookNotice{Event: "template.failed", Title: "Template retry failed", Description: "Aperture will retry with bounded backoff.", Color: 0xe67e22, Fields: map[string]string{"Username": reg.Username, "Registration": strconv.FormatInt(reg.ID, 10), "Attempt": strconv.Itoa(reg.TemplateAttempts + 1), "Error": safeError(err)}}) - continue + ids, err := s.store.ListDueTemplateRecoveryIDs(ctx, op.Identity.Binding.ID, 10) + if err != nil { + slog.Warn("could not list template retries", "error", safeError(err)) + return + } + for _, id := range ids { + if ctx.Err() != nil { + return } - if err := s.store.CompleteTemplateRecovery(ctx, reg.ID, disableAtFrom(reg.CreatedAt, recovery.UserExpiryDays)); err != nil { - slog.Warn("could not complete automatic template recovery", "registration_id", reg.ID, "error", safeError(err)) - continue + operationCtx, err := s.refreshAccountContext(ctx) + if err != nil { + return + } + if err := s.recoverAccount(operationCtx, id, true); err != nil { + slog.Warn("template recovery did not complete", "registration_id", id, "error", safeError(err)) } - s.notify(webhookNotice{Event: "template.recovered", Title: "Access template recovered", Color: 0x2ecc71, Fields: map[string]string{"Username": reg.Username, "Registration": strconv.FormatInt(reg.ID, 10), "Template": recovery.Template.Name}}) } } diff --git a/internal/httpserver/fake_media_server_test.go b/internal/httpserver/fake_media_server_test.go index ac8773b..555168a 100644 --- a/internal/httpserver/fake_media_server_test.go +++ b/internal/httpserver/fake_media_server_test.go @@ -23,6 +23,8 @@ type fakeMediaServer struct { applyErr error importErr error pingErr error + inspectErr error + inspectAPIKey atomic.Value pingCalls atomic.Int32 pingStarted chan struct{} pingRelease chan struct{} @@ -34,9 +36,12 @@ type fakeMediaServer struct { disableErr error } -func (f *fakeMediaServer) SetProvider(provider mediaserver.Provider) error { - f.provider = provider - return nil +func (f *fakeMediaServer) Inspect(ctx context.Context, baseURL, token, deviceID string) (mediaserver.ServerInfo, error) { + f.inspectAPIKey.Store(token) + if err := f.inspectErr; err != nil && deviceID == "aperture" { + return mediaserver.ServerInfo{}, err + } + return mediaserver.ServerInfo{ID: "synthetic-server", Name: "Media"}, nil } func (f *fakeMediaServer) Authenticate(context.Context, string, string, string) (mediaserver.AuthResult, error) { @@ -62,12 +67,12 @@ func (f *fakeMediaServer) Ping(_ context.Context, _, apiKey string) error { } return nil } -func (f *fakeMediaServer) CreateUser(context.Context, string, string, string, string) (mediaserver.User, error) { +func (f *fakeMediaServer) CreateUser(_ context.Context, _, _, _, _ string, created func(mediaserver.User) error) (mediaserver.User, error) { f.createdUser = true if f.createErr != nil { return mediaserver.User{}, f.createErr } - return mediaserver.User{ID: "new-media-user", Name: "new_user"}, nil + return mediaserver.User{ID: "new-media-user", Name: "new_user"}, created(mediaserver.User{ID: "new-media-user", Name: "new_user"}) } func (f *fakeMediaServer) ApplyTemplate(_ context.Context, _, _, userID string, template db.Template) error { f.appliedTemplate = true @@ -101,6 +106,6 @@ func (f *fakeMediaServer) ImportTemplate(context.Context, string, string, string return mediaserver.TemplateData{PolicyJSON: db.TemplatePolicyDefaultJSON}, nil } -var _ MediaServer = (*fakeMediaServer)(nil) +var _ mediaserver.Server = (*fakeMediaServer)(nil) var errFakePing = errors.New("ping failed") diff --git a/internal/httpserver/fake_store_test.go b/internal/httpserver/fake_store_test.go index f226564..325bf07 100644 --- a/internal/httpserver/fake_store_test.go +++ b/internal/httpserver/fake_store_test.go @@ -10,6 +10,7 @@ import ( type fakeStore struct { settings db.Settings + connection db.MediaConnection session db.Session template db.Template invite db.Invite @@ -64,19 +65,20 @@ func newFakeStore() *fakeStore { InviteSecret: "invite-secret-with-at-least-32-characters", } tmpl := db.Template{ID: 1, Name: "Default", PolicyJSON: `{"IsAdministrator":false}`, IsDefault: true} - invite := db.Invite{ID: 1, TokenHash: "hash", TokenPrefix: "prefix", Token: "saved-token", Label: "Family", TemplateID: 1, Template: "Default", MaxUses: 3, Enabled: true, UserExpiryDays: 0} + invite := db.Invite{ID: 1, TokenHash: "hash", TokenPrefix: "prefix", Token: "saved-token", Label: "Family", TemplateID: 1, Template: "Default", MaxUses: 3, Enabled: true, UserExpiryDays: 0, BindingID: 1} return &fakeStore{ settings: settings, - session: db.Session{ID: "session-id", UserID: "admin-id", Username: "admin", AccessToken: "access-token", DeviceID: "device-id", CSRFSecret: "csrf-secret", ExpiresAt: time.Now().Add(time.Hour)}, + connection: db.MediaConnection{Provider: settings.Provider, BaseURL: settings.ServerURL, Generation: 1, Binding: db.MediaBinding{ID: 1, Provider: settings.Provider, BaseURL: settings.ServerURL, ServerID: "synthetic-server", Name: "Media"}}, + session: db.Session{ID: "session-id", UserID: "admin-id", Username: "admin", AccessToken: "access-token", DeviceID: "device-id", CSRFSecret: "csrf-secret", ExpiresAt: time.Now().Add(time.Hour), BindingID: 1, Generation: 1}, template: tmpl, invite: invite, invites: []db.Invite{invite}, - registrations: []db.Registration{{ID: 1, Username: "alice", Status: db.RegistrationComplete, ExternalUserID: sql.NullString{String: "media-alice", Valid: true}}}, + registrations: []db.Registration{{BindingID: 1, ID: 1, Username: "alice", Status: db.RegistrationComplete, ExternalUserID: sql.NullString{String: "media-alice", Valid: true}}}, inviteActivity: map[int64]db.InviteActivity{ 1: {InviteID: 1, Username: "alice", Status: db.RegistrationComplete, CreatedAt: time.Date(2030, 1, 2, 3, 4, 5, 0, time.UTC)}, }, recovery: db.RegistrationRecovery{ - Registration: db.Registration{ID: 1, InviteID: 1, ExternalUserID: sql.NullString{String: "media-alice", Valid: true}, Status: db.RegistrationNeedsAttention, CreatedAt: time.Now().Add(-time.Hour)}, + Registration: db.Registration{ID: 1, InviteID: 1, ExternalUserID: sql.NullString{String: "media-alice", Valid: true}, Status: db.RegistrationNeedsAttention, CreatedAt: time.Now().Add(-time.Hour), BindingID: 1}, Template: tmpl, UserExpiryDays: 7, }, @@ -118,29 +120,9 @@ func (f *fakeStore) UpdateApplicationSettings(ctx context.Context, provider, pub } return nil } -func (f *fakeStore) UpdateSetupSettings(ctx context.Context, provider, publicURL, serverURL, apiKey string) error { - if f.setupSettingsErr != nil { - return f.setupSettingsErr - } - for _, setting := range []struct { - key, value string - secret bool - }{ - {"media_provider", provider, false}, - {"public_url", publicURL, false}, - {"server_url", serverURL, false}, - {"api_key", apiKey, true}, - } { - if err := f.SetSetting(ctx, setting.key, setting.value, setting.secret); err != nil { - return err - } - } - f.settings.PublicURL = publicURL - return nil -} -func (f *fakeStore) CreateSession(_ context.Context, _, _, _, deviceID string, ttl time.Duration) (string, string, error) { - f.createdDeviceID = deviceID - f.createdSessionTTL = ttl +func (f *fakeStore) CreateSession(_ context.Context, input db.SessionInput) (string, string, error) { + f.createdDeviceID = input.DeviceID + f.createdSessionTTL = input.TTL return f.session.ID, f.session.CSRFSecret, nil } func (f *fakeStore) Session(_ context.Context, id string) (db.Session, error) { @@ -183,7 +165,6 @@ func (f *fakeStore) CreateInvite(_ context.Context, invite db.Invite) (int64, er f.createdInvite.ID = 99 return 99, nil } -func (f *fakeStore) ListInvites(context.Context) ([]db.Invite, error) { return f.invites, nil } func (f *fakeStore) InvitePreview(_ context.Context, limit int) ([]db.Invite, error) { if limit > len(f.invites) { limit = len(f.invites) @@ -208,17 +189,20 @@ func (f *fakeStore) InvitePreset(_ context.Context, id int64) (db.Invite, error) } return db.Invite{}, db.ErrNotFound } -func (f *fakeStore) InviteByHash(context.Context, string) (db.Invite, error) { +func (f *fakeStore) InviteByHash(context.Context, string, int64) (db.Invite, error) { if f.inviteErr != nil { return db.Invite{}, f.inviteErr } return f.invite, nil } -func (f *fakeStore) ReserveInviteUse(context.Context, int64, string, string, string) (int64, db.Template, error) { +func (f *fakeStore) ReserveInviteUse(context.Context, int64, int64, string, string, string) (int64, db.Template, error) { if f.templateErr != nil { return 0, db.Template{}, f.templateErr } f.reservedInviteUse = true + if f.invite.UserExpiryDays > 0 { + f.completedDisableAt = sql.NullTime{Time: time.Now().AddDate(0, 0, f.invite.UserExpiryDays), Valid: true} + } return 123, f.template, nil } func (f *fakeStore) BeginUserCreation(context.Context, int64) error { @@ -241,16 +225,15 @@ func (f *fakeStore) RecordCreatedUser(_ context.Context, _ int64, userID string) f.recordedUserID = userID return nil } -func (f *fakeStore) CompleteRegistration(_ context.Context, _ int64, status string, _ string, disableAt sql.NullTime) error { +func (f *fakeStore) CompleteRegistration(_ context.Context, _ int64, status string, _ string) error { f.completedStatus = status - f.completedDisableAt = disableAt return nil } func (f *fakeStore) ReconcileStaleRegistrations(context.Context, time.Time, int) (db.ReconciliationResult, error) { f.reconciledStale = true return f.reconciliation, nil } -func (f *fakeStore) DueUserDisables(context.Context, int) ([]db.Registration, error) { +func (f *fakeStore) DueUserDisables(context.Context, int64, int) ([]db.Registration, error) { return f.dueDisables, nil } func (f *fakeStore) MarkUserDisabled(_ context.Context, registrationID int64) error { @@ -261,19 +244,19 @@ func (f *fakeStore) MarkUserDisableFailed(_ context.Context, id int64, _ string) f.failedDisableID = id return nil } -func (f *fakeStore) SetInviteEnabled(context.Context, int64, bool) error { return nil } -func (f *fakeStore) DeleteInvite(context.Context, int64) error { return nil } +func (f *fakeStore) SetInviteEnabled(context.Context, int64, int64, bool) error { return nil } +func (f *fakeStore) DeleteInvite(context.Context, int64) error { return nil } func (f *fakeStore) Audit(context.Context, string, string, string, string, string, string, string) error { return nil } func (f *fakeStore) RecentRegistrations(context.Context, int) ([]db.Registration, error) { return f.registrations, nil } -func (f *fakeStore) RegistrationUsers(context.Context) ([]db.Registration, error) { +func (f *fakeStore) RegistrationUsers(context.Context, int64) ([]db.Registration, error) { return f.registrations, nil } func (f *fakeStore) Registration(_ context.Context, id int64) (db.Registration, error) { - for _, registration := range f.registrations { + for _, registration := range append(append([]db.Registration{}, f.registrations...), f.dueDisables...) { if registration.ID == id { return registration, nil } @@ -284,7 +267,7 @@ func (f *fakeStore) DeleteRegistration(_ context.Context, id int64) error { f.deletedRegistrationID = id return nil } -func (f *fakeStore) ListManagedUsers(context.Context) ([]db.ManagedUser, error) { +func (f *fakeStore) ListManagedUsers(context.Context, int64) ([]db.ManagedUser, error) { return f.managedUsers, nil } func (f *fakeStore) SaveManagedUser(_ context.Context, user db.ManagedUser) error { @@ -297,7 +280,7 @@ func (f *fakeStore) SaveManagedUser(_ context.Context, user db.ManagedUser) erro f.managedUsers = append(f.managedUsers, user) return nil } -func (f *fakeStore) DeleteUserRecords(_ context.Context, id string) error { +func (f *fakeStore) DeleteUserRecords(_ context.Context, _ int64, id string) error { for i := len(f.registrations) - 1; i >= 0; i-- { if f.registrations[i].ExternalUserID.String == id { f.registrations = append(f.registrations[:i], f.registrations[i+1:]...) @@ -310,7 +293,7 @@ func (f *fakeStore) DeleteUserRecords(_ context.Context, id string) error { } return nil } -func (f *fakeStore) ClaimTemplateRecovery(context.Context, int64) (db.RegistrationRecovery, error) { +func (f *fakeStore) ClaimTemplateRecovery(context.Context, int64, int64, bool) (db.RegistrationRecovery, error) { recovery := f.recovery recovery.Registration.Status = db.RegistrationRetryingTemplate return recovery, nil @@ -319,11 +302,11 @@ func (f *fakeStore) RecordTemplateRetryFailure(_ context.Context, _ int64, messa f.recoveryFailure = message return nil } -func (f *fakeStore) CompleteTemplateRecovery(context.Context, int64, sql.NullTime) error { +func (f *fakeStore) CompleteTemplateRecovery(context.Context, int64) error { f.recoveryCompleted = true return nil } -func (f *fakeStore) DashboardCounts(context.Context) (db.DashboardCounts, error) { +func (f *fakeStore) DashboardCounts(context.Context, int64) (db.DashboardCounts, error) { stats := dashboardStatsFrom(f.invites, f.registrations) return db.DashboardCounts{ ActiveInvites: stats.ActiveInvites, @@ -332,9 +315,6 @@ func (f *fakeStore) DashboardCounts(context.Context) (db.DashboardCounts, error) ScheduledUserDisables: stats.ScheduledUserDisables, }, nil } -func (f *fakeStore) LatestInviteActivity(context.Context) (map[int64]db.InviteActivity, error) { - return f.inviteActivity, nil -} func (f *fakeStore) ListAuditEvents(context.Context, int) ([]db.AuditEvent, error) { return f.auditEvents, nil } @@ -357,8 +337,102 @@ func (f *fakeStore) DeleteWebhook(_ context.Context, id int64) error { } return db.ErrNotFound } -func (f *fakeStore) DueTemplateRecoveries(context.Context, int) ([]db.RegistrationRecovery, error) { - return f.dueTemplateRetries, nil +func (f *fakeStore) ListDueTemplateRecoveryIDs(context.Context, int64, int) ([]int64, error) { + var ids []int64 + for _, r := range f.dueTemplateRetries { + ids = append(ids, r.Registration.ID) + } + return ids, nil } var _ Store = (*fakeStore)(nil) + +func (f *fakeStore) RecordProvisioningUser(_ context.Context, _ int64, id string) error { + f.recordedUserID = id + return nil +} +func (f *fakeStore) RequireAccountCleanup(_ context.Context, _ int64, id string) error { + f.recordedUserID = id + return nil +} +func (f *fakeStore) UserDeletionRegistrations(_ context.Context, _ int64, id string) ([]db.Registration, error) { + var regs []db.Registration + for _, r := range f.registrations { + if r.ExternalUserID.String == id { + if db.IsRegistrationActive(r.Status) { + return nil, db.ErrRegistrationTransition + } + regs = append(regs, r) + } + } + if len(regs) > 0 { + return regs, nil + } + for _, u := range f.managedUsers { + if u.ExternalUserID == id { + return nil, nil + } + } + return nil, db.ErrNotFound +} + +func (f *fakeStore) MediaConnection(context.Context) (db.MediaConnection, error) { + return f.connection, nil +} +func (f *fakeStore) PublishMediaConnection(ctx context.Context, u db.ConnectionUpdate) (db.MediaConnection, error) { + if f.setupSettingsErr != nil { + return db.MediaConnection{}, f.setupSettingsErr + } + if u.ExpectedGeneration != f.connection.Generation { + return db.MediaConnection{}, db.ErrConnectionChanged + } + b := u.Origin + b.ID = f.connection.Binding.ID + if b.Provider != f.connection.Provider || b.BaseURL != f.connection.BaseURL || b.ServerID != f.connection.Binding.ServerID { + f.connection.Generation++ + b.ID++ + } + if b.ServerID == "" { + b.ID = 0 + } + if b.ServerID != "" && b.ID == 0 { + b.ID = 1 + } + if err := f.UpdateApplicationSettings(ctx, u.Provider, u.PublicURL, u.ServerURL, u.APIKey); err != nil { + return db.MediaConnection{}, err + } + f.connection.Binding = b + f.connection.Provider = b.Provider + f.connection.BaseURL = b.BaseURL + return f.connection, nil +} +func (f *fakeStore) MediaBindings(context.Context) ([]db.MediaBinding, error) { + return []db.MediaBinding{f.connection.Binding}, nil +} +func (f *fakeStore) AdoptInvite(context.Context, int64, int64) error { return nil } +func (f *fakeStore) AdoptRegistration(context.Context, int64, int64) error { return nil } + +func (f *fakeStore) Invite(_ context.Context, id int64) (db.Invite, error) { + for _, v := range f.invites { + if v.ID == id { + return v, nil + } + } + return db.Invite{}, db.ErrNotFound +} +func (f *fakeStore) InvitePage(context.Context, int64, int) ([]db.Invite, error) { + return f.invites, nil +} +func (f *fakeStore) InvitePageActivity(context.Context, []int64) (map[int64]db.InviteActivity, error) { + return f.inviteActivity, nil +} +func (f *fakeStore) RegistrationPage(context.Context, int64, int64, bool, int) ([]db.Registration, error) { + return f.registrations, nil +} +func (f *fakeStore) ManagedUserReviewPage(context.Context, int64, int64, int) ([]db.ManagedUser, error) { + return nil, nil +} +func (f *fakeStore) ManagedUser(context.Context, int64) (db.ManagedUser, error) { + return db.ManagedUser{}, db.ErrNotFound +} +func (f *fakeStore) AdoptManagedUser(context.Context, int64, int64) error { return nil } diff --git a/internal/httpserver/handlers_admin.go b/internal/httpserver/handlers_admin.go index 2909915..16299f4 100644 --- a/internal/httpserver/handlers_admin.go +++ b/internal/httpserver/handlers_admin.go @@ -45,6 +45,8 @@ func (s *Server) settingsForm(w http.ResponseWriter, r *http.Request, session db render(w, "settings", data) } func (s *Server) settingsPost(w http.ResponseWriter, r *http.Request, session db.Session) { + s.setupMu.Lock() + defer s.setupMu.Unlock() if err := r.ParseForm(); err != nil { s.message(w, "Invalid request", "The settings form could not be read.", http.StatusBadRequest) return @@ -99,51 +101,21 @@ func (s *Server) settingsPost(w http.ResponseWriter, r *http.Request, session db } else if apiKey == "" { effectiveAPIKey = current.APIKey } - previousProvider, _, _ := s.runtimeSettings() - restoreProvider := func() { - if previous, ok := mediaserver.ParseProvider(previousProvider); ok { - _ = s.media.SetProvider(previous) - } - } - if err := s.media.SetProvider(provider); err != nil { - s.renderSettingsError(w, r, session, string(provider), publicURL, serverURL, "Could not select that media server.") + expected, err := s.snapshot(r.Context()) + if err != nil { + s.error(w, err) return } - if effectiveAPIKey != "" { - if err := s.media.Ping(r.Context(), serverURL, effectiveAPIKey); err != nil { - restoreProvider() - s.renderSettingsError(w, r, session, string(provider), publicURL, serverURL, "Could not reach the media server with those connection details.") - return - } - } - var providerUpdate, publicURLUpdate, serverURLUpdate, apiKeyUpdate *string - if !s.cfg.ProviderManaged { - value := string(provider) - providerUpdate = &value - } - if !s.cfg.PublicURLManaged { - publicURLUpdate = &publicURL - } - if !s.cfg.ServerURLManaged { - serverURLUpdate = &serverURL - } - if !s.cfg.APIKeyManaged && (apiKey != "" || removeAPIKey) { - apiKeyUpdate = &apiKey - } - if providerUpdate != nil || publicURLUpdate != nil || serverURLUpdate != nil || apiKeyUpdate != nil { - if err := s.store.UpdateApplicationSettings(r.Context(), providerUpdate, publicURLUpdate, serverURLUpdate, apiKeyUpdate); err != nil { - restoreProvider() - s.error(w, err) - return - } - } - cookieSecure := strings.HasPrefix(publicURL, "https://") - if s.cfg.CookieManaged { - cookieSecure = s.cfg.CookieSecure + target := current + target.Provider, target.PublicURL, target.ServerURL, target.APIKey = string(provider), publicURL, serverURL, effectiveAPIKey + update := s.settingsUpdate(target, apiKey != "" || removeAPIKey) + published, err := s.connections.Publish(r.Context(), expected, target, update) + if err != nil { + s.renderSettingsError(w, r, session, string(provider), publicURL, serverURL, "Could not verify or save this connection. Refresh the page, check the details, and try again.") + return } - s.setRuntime(string(provider), publicURL, cookieSecure) - s.audit(r, session, "settings.update", "settings", "application", map[string]any{"provider": string(provider), "public_url": publicURL, "server_url": serverURL, "api_key_changed": apiKeyUpdate != nil}) - if current.Provider != string(provider) || current.ServerURL != serverURL { + s.audit(r, session, "settings.update", "settings", "application", map[string]any{"provider": string(provider), "public_url": publicURL, "server_url": serverURL, "api_key_changed": update.APIKey != nil}) + if published.Identity.Generation != session.Generation { s.clearSession(w, r, session.ID) http.Redirect(w, r, "/login", http.StatusSeeOther) return @@ -169,3 +141,20 @@ func (s *Server) renderSettingsError(w http.ResponseWriter, r *http.Request, ses s.setMediaViewData(&data) render(w, "settings", data) } + +func (s *Server) settingsUpdate(target db.Settings, replaceKey bool) db.ConnectionUpdate { + var update db.ConnectionUpdate + if !s.cfg.ProviderManaged { + update.Provider = &target.Provider + } + if !s.cfg.PublicURLManaged { + update.PublicURL = &target.PublicURL + } + if !s.cfg.ServerURLManaged { + update.ServerURL = &target.ServerURL + } + if !s.cfg.APIKeyManaged && replaceKey { + update.APIKey = &target.APIKey + } + return update +} diff --git a/internal/httpserver/handlers_admin_test.go b/internal/httpserver/handlers_admin_test.go index 78f0df0..49272d5 100644 --- a/internal/httpserver/handlers_admin_test.go +++ b/internal/httpserver/handlers_admin_test.go @@ -20,7 +20,7 @@ func TestAdminDashboardLeavesDueUserDisablesToBackgroundWorker(t *testing.T) { ExternalUserID: sql.NullString{String: "media-expired", Valid: true}, }} media := &fakeMediaServer{} - handler := New(testConfig(), store, media) + handler := New(testConfig(), store, testMediaFactory(media)) req := adminRequest(t, http.MethodGet, "/admin", nil) rr := httptest.NewRecorder() @@ -36,7 +36,7 @@ func TestAdminDashboardLeavesDueUserDisablesToBackgroundWorker(t *testing.T) { func TestSettingsEnvironmentValuesAreManagedAndNotWritable(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) getReq := adminRequest(t, http.MethodGet, "/admin/settings", nil) getRR := httptest.NewRecorder() @@ -100,7 +100,7 @@ func TestSettingsUpdatesAllBrowserManagedApplicationSettings(t *testing.T) { cfg.ServerURLManaged = false cfg.APIKeyManaged = false media := &fakeMediaServer{provider: mediaserver.ProviderJellyfin} - handler := New(cfg, store, media) + handler := New(cfg, store, testMediaFactory(media)) form := url.Values{ "csrf": {store.session.CSRFSecret}, @@ -121,7 +121,7 @@ func TestSettingsUpdatesAllBrowserManagedApplicationSettings(t *testing.T) { if rr.Header().Get("Location") != "/login" || store.deletedSessionID != store.session.ID { t.Fatalf("provider switch did not clear the old media-server session") } - if media.provider != mediaserver.ProviderEmby { + if handler.(*Server).connectionsStateProvider() != "emby" { t.Fatalf("active provider = %q, want Emby", media.provider) } if store.settings.Provider != "emby" || store.settings.PublicURL != "https://join.example.com" || @@ -138,7 +138,7 @@ func TestSettingsCanRemoveSavedAPIKey(t *testing.T) { cfg := testConfig() cfg.APIKey = "" cfg.APIKeyManaged = false - handler := New(cfg, store, &fakeMediaServer{}) + handler := New(cfg, store, testMediaFactory(&fakeMediaServer{})) form := url.Values{ "csrf": {store.session.CSRFSecret}, "remove_api_key": {"on"}, @@ -167,7 +167,7 @@ func TestSettingsNeverRendersAPIKey(t *testing.T) { cfg := testConfig() cfg.APIKey = "environment-api-key-never-render" cfg.APIKeyManaged = false - handler := New(cfg, store, &fakeMediaServer{}) + handler := New(cfg, store, testMediaFactory(&fakeMediaServer{})) req := adminRequest(t, http.MethodGet, "/admin/settings", nil) rr := httptest.NewRecorder() @@ -239,7 +239,7 @@ func TestSettingsPostWritesOnlyUIManagedValues(t *testing.T) { store := newFakeStore() cfg := testConfig() tt.configure(&cfg) - handler := New(cfg, store, &fakeMediaServer{}) + handler := New(cfg, store, testMediaFactory(&fakeMediaServer{})) form := url.Values{ "csrf": {store.session.CSRFSecret}, "server_url": {"http://new-media:8096"}, @@ -270,7 +270,7 @@ func TestSettingsPostRejectsUnreachableConnectionBeforeWriting(t *testing.T) { cfg.APIKey = "" cfg.ServerURLManaged = false cfg.APIKeyManaged = false - handler := New(cfg, store, &fakeMediaServer{pingErr: errFakePing}) + handler := New(cfg, store, testMediaFactory(&fakeMediaServer{inspectErr: errFakePing})) form := url.Values{ "csrf": {store.session.CSRFSecret}, "server_url": {"http://new-media:8096"}, @@ -282,7 +282,7 @@ func TestSettingsPostRejectsUnreachableConnectionBeforeWriting(t *testing.T) { handler.ServeHTTP(rr, req) - if rr.Code != http.StatusOK || !strings.Contains(rr.Body.String(), "Could not reach the media server") { + if rr.Code != http.StatusOK || !strings.Contains(rr.Body.String(), "Could not verify or save this connection") { t.Fatalf("status = %d; body %s", rr.Code, rr.Body.String()) } if len(store.settingWrites) != 0 { @@ -294,7 +294,7 @@ func TestAdminDashboardShowsHealthChecks(t *testing.T) { store := newFakeStore() store.invites = nil media := &fakeMediaServer{pingErr: errFakePing} - handler := New(testConfig(), store, media) + handler := New(testConfig(), store, testMediaFactory(media)) req := adminRequest(t, http.MethodGet, "/admin", nil) rr := httptest.NewRecorder() @@ -313,7 +313,7 @@ func TestAdminDashboardShowsHealthChecks(t *testing.T) { func TestAdminDashboardDoesNotRenderRetainedInviteTokens(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) req := adminRequest(t, http.MethodGet, "/admin", nil) rr := httptest.NewRecorder() @@ -337,7 +337,7 @@ func TestAdminDashboardCachesMediaServerHealthCheck(t *testing.T) { cfg.APIKey = "" cfg.ServerURLManaged = false cfg.APIKeyManaged = false - handler := New(cfg, store, media) + handler := New(cfg, store, testMediaFactory(media)) for range 2 { req := adminRequest(t, http.MethodGet, "/admin", nil) @@ -361,7 +361,7 @@ func TestAdminDashboardHealthCacheChangesWithCredentials(t *testing.T) { cfg.APIKey = "" cfg.ServerURLManaged = false cfg.APIKeyManaged = false - handler := New(cfg, store, media) + handler := New(cfg, store, testMediaFactory(media)) request := func() { t.Helper() @@ -374,7 +374,16 @@ func TestAdminDashboardHealthCacheChangesWithCredentials(t *testing.T) { } request() - store.settings.APIKey = "replacement-api-key" + server := handler.(*Server) + current, err := server.connections.Current(t.Context()) + if err != nil { + t.Fatal(err) + } + target := current.Settings + target.APIKey = "replacement-api-key" + if _, err := server.connections.Publish(t.Context(), current, target, server.settingsUpdate(target, true)); err != nil { + t.Fatal(err) + } request() if media.pingCalls.Load() != 2 { t.Fatalf("media-server ping calls after credential change = %d, want 2", media.pingCalls.Load()) diff --git a/internal/httpserver/handlers_auth.go b/internal/httpserver/handlers_auth.go index 20211b3..fab0c87 100644 --- a/internal/httpserver/handlers_auth.go +++ b/internal/httpserver/handlers_auth.go @@ -8,6 +8,8 @@ import ( "strings" "time" + "github.com/mayvqt/aperture/internal/connection" + "github.com/mayvqt/aperture/internal/db" "github.com/mayvqt/aperture/internal/mediaserver" ) @@ -54,6 +56,12 @@ func (s *Server) setupPost(w http.ResponseWriter, r *http.Request) { // A waiting setup request must recheck completion before changing providers. s.setupMu.Lock() defer s.setupMu.Unlock() + current, err := s.connections.Current(r.Context()) + if err != nil { + s.error(w, err) + return + } + r = r.WithContext(connection.WithSnapshot(r.Context(), current)) complete, err := s.setupComplete(r.Context()) if err != nil { s.error(w, err) @@ -106,33 +114,13 @@ func (s *Server) setupPost(w http.ResponseWriter, r *http.Request) { s.renderSetupError(w, r, string(provider), publicURL, serverURL, err.Error()) return } - previousProvider, _, _ := s.runtimeSettings() - restoreProvider := func() { - if previous, ok := mediaserver.ParseProvider(previousProvider); ok { - _ = s.media.SetProvider(previous) - } - } - if err := s.media.SetProvider(provider); err != nil { - s.renderSetupError(w, r, string(provider), publicURL, serverURL, "Could not select that media server.") + target := current.Settings + target.Provider, target.PublicURL, target.ServerURL, target.APIKey = string(provider), publicURL, serverURL, apiKey + update := s.settingsUpdate(target, true) + if _, err := s.connections.Publish(r.Context(), current, target, update); err != nil { + s.renderSetupError(w, r, string(provider), publicURL, serverURL, "Could not verify or save that media-server connection. Check the connection details and try again.") return } - if apiKey != "" { - if err := s.media.Ping(r.Context(), serverURL, apiKey); err != nil { - restoreProvider() - s.renderSetupError(w, r, string(provider), publicURL, serverURL, "Could not reach the media server with that API key.") - return - } - } - if err := s.store.UpdateSetupSettings(r.Context(), string(provider), publicURL, serverURL, apiKey); err != nil { - restoreProvider() - s.error(w, err) - return - } - cookieSecure := strings.HasPrefix(publicURL, "https://") - if s.cfg.CookieManaged { - cookieSecure = s.cfg.CookieSecure - } - s.setRuntime(string(provider), publicURL, cookieSecure) http.Redirect(w, r, "/login", http.StatusSeeOther) } func (s *Server) loginForm(w http.ResponseWriter, r *http.Request) { @@ -169,17 +157,27 @@ func (s *Server) loginPost(w http.ResponseWriter, r *http.Request) { s.renderLoginFailure(w, r) return } - auth, err := s.media.Authenticate(r.Context(), settings.ServerURL, username, password) + snapshot, err := s.snapshot(r.Context()) + if err != nil { + s.error(w, err) + return + } + auth, err := snapshot.Media.Authenticate(r.Context(), settings.ServerURL, username, password) if err != nil || !auth.IsAdmin { s.renderLoginFailure(w, r) return } + verified, err := s.connections.Verify(r.Context(), snapshot, auth.AccessToken, auth.DeviceID) + if err != nil { + s.message(w, "Could not verify server", "Sign-in could not be completed because the media server changed or its identity could not be verified. Try again shortly.", http.StatusBadGateway) + return + } s.limiter.Reset(rateKey) sessionTTL := defaultSessionTTL if r.FormValue("remember_me") == "on" { sessionTTL = rememberMeTTL } - sessionID, _, err := s.store.CreateSession(r.Context(), auth.UserID, auth.Username, auth.AccessToken, auth.DeviceID, sessionTTL) + sessionID, _, err := s.store.CreateSession(r.Context(), db.SessionInput{UserID: auth.UserID, Username: auth.Username, AccessToken: auth.AccessToken, DeviceID: auth.DeviceID, TTL: sessionTTL, BindingID: verified.Identity.Binding.ID, Generation: verified.Identity.Generation}) if err != nil { s.error(w, err) return diff --git a/internal/httpserver/handlers_auth_test.go b/internal/httpserver/handlers_auth_test.go index f8d6574..c27d226 100644 --- a/internal/httpserver/handlers_auth_test.go +++ b/internal/httpserver/handlers_auth_test.go @@ -18,7 +18,7 @@ import ( func TestAdminLoginBrowserFlow(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) getRequest := httptest.NewRequest(http.MethodGet, "/login", nil) getResponse := httptest.NewRecorder() @@ -77,7 +77,7 @@ func TestLoginRememberMeSessionDuration(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) form := url.Values{ "csrf": {"csrf-value"}, "username": {"admin"}, @@ -110,7 +110,7 @@ func TestLoginRememberMeSessionDuration(t *testing.T) { func TestLoginFormIncludesRememberMeControl(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) rr := httptest.NewRecorder() handler.ServeHTTP(rr, httptest.NewRequest(http.MethodGet, "/login", nil)) @@ -125,7 +125,7 @@ func TestLoginFormIncludesRememberMeControl(t *testing.T) { func TestSuccessfulLoginsDoNotExhaustClientIPLimit(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) csrf := "csrf-value" for i := range 60 { @@ -155,7 +155,7 @@ func TestSetupPostRejectsInvalidMediaServerURL(t *testing.T) { CookieSecure: false, SessionSecret: store.settings.SessionSecret, InviteSecret: store.settings.InviteSecret, - }, store, &fakeMediaServer{}) + }, store, testMediaFactory(&fakeMediaServer{})) csrf := "csrf-value" form := url.Values{ "csrf": {csrf}, @@ -183,7 +183,7 @@ func TestSetupPostRejectsInvalidMediaServerURL(t *testing.T) { func TestSetupRedirectsWhenMediaServerURLIsEnvironmentManaged(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) for _, method := range []string{http.MethodGet, http.MethodPost} { req := httptest.NewRequest(method, "/setup", nil) @@ -208,7 +208,7 @@ func TestSetupUsesEnvironmentManagedAPIKeyWithoutStoringIt(t *testing.T) { cfg.ServerURLManaged = false cfg.APIKeyManaged = true media := &fakeMediaServer{} - handler := New(cfg, store, media) + handler := New(cfg, store, testMediaFactory(media)) getReq := httptest.NewRequest(http.MethodGet, "/setup", nil) getRR := httptest.NewRecorder() @@ -241,10 +241,10 @@ func TestSetupUsesEnvironmentManagedAPIKeyWithoutStoringIt(t *testing.T) { if postRR.Code != http.StatusSeeOther { t.Fatalf("POST status = %d, want redirect; body %s", postRR.Code, postRR.Body.String()) } - if len(store.settingWrites) != 4 || store.settingWrites[0].key != "media_provider" || store.settingWrites[1].key != "public_url" { - t.Fatalf("setup writes = %#v, want complete browser settings", store.settingWrites) + if len(store.settingWrites) != 1 || store.settingWrites[0].key != "server_url" { + t.Fatalf("setup writes = %#v, want only browser-managed server URL", store.settingWrites) } - if got, _ := media.pingAPIKey.Load().(string); got != cfg.APIKey { + if got, _ := media.inspectAPIKey.Load().(string); got != cfg.APIKey { t.Fatalf("setup ping API key = %q, want environment-managed key", got) } } @@ -259,7 +259,7 @@ func TestSetupConfiguresEmbyEntirelyFromBrowser(t *testing.T) { CookieSecure: false, } media := &fakeMediaServer{} - handler := New(cfg, store, media) + handler := New(cfg, store, testMediaFactory(media)) getReq := httptest.NewRequest(http.MethodGet, "/setup", nil) getRR := httptest.NewRecorder() @@ -297,7 +297,7 @@ func TestSetupConfiguresEmbyEntirelyFromBrowser(t *testing.T) { if postRR.Code != http.StatusSeeOther || postRR.Header().Get("Location") != "/login" { t.Fatalf("status = %d, location = %q; body %s", postRR.Code, postRR.Header().Get("Location"), postRR.Body.String()) } - if media.provider != mediaserver.ProviderEmby { + if handler.(*Server).connectionsStateProvider() != "emby" { t.Fatalf("active provider = %q", media.provider) } if store.settings.Provider != "emby" || store.settings.PublicURL != "https://join.example.com" || @@ -315,7 +315,7 @@ func TestSetupRestoresProviderWhenSettingsCannotBeSaved(t *testing.T) { store.setupSettingsErr = errors.New("save failed") cfg := config.Config{MediaProvider: "jellyfin"} media := &fakeMediaServer{provider: mediaserver.ProviderJellyfin} - handler := New(cfg, store, media) + handler := New(cfg, store, testMediaFactory(media)) csrf := "csrf-value" form := url.Values{ @@ -345,7 +345,7 @@ func TestSetupRestoresProviderWhenSettingsCannotBeSaved(t *testing.T) { func TestLoginInvalidCSRFDoesNotConsumeRateLimit(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) for range 10 { form := url.Values{"username": {"admin"}, "password": {"password"}} @@ -377,7 +377,7 @@ func TestLoginInvalidCSRFDoesNotConsumeRateLimit(t *testing.T) { func TestLoginRateLimitsRotatingUsernamesByClientIP(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{authErr: errors.New("bad credentials")}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{authErr: errors.New("bad credentials")})) csrf := "csrf-value" for i := range 51 { @@ -410,10 +410,10 @@ type setupPingMedia struct { pingError error } -func (m *setupPingMedia) Ping(context.Context, string, string) error { +func (m *setupPingMedia) Inspect(context.Context, string, string, string) (mediaserver.ServerInfo, error) { close(m.started) <-m.release - return m.pingError + return mediaserver.ServerInfo{ID: "synthetic-server"}, m.pingError } func TestConcurrentSetupCannotOverwriteCompletedConfiguration(t *testing.T) { @@ -427,7 +427,7 @@ func TestConcurrentSetupCannotOverwriteCompletedConfiguration(t *testing.T) { if failFirst { media.pingError = errors.New("upstream unavailable") } - handler := New(config.Config{MediaProvider: "jellyfin"}, store, media) + handler := New(config.Config{MediaProvider: "jellyfin"}, store, testMediaFactory(media)) secret := store.settings.SessionSecret post := func(provider, serverURL, apiKey string) *httptest.ResponseRecorder { form := url.Values{"csrf": {"csrf-value"}, "provider": {provider}, "public_url": {"https://join.example.test"}, "server_url": {serverURL}, "api_key": {apiKey}} @@ -473,7 +473,7 @@ func TestConcurrentSetupCannotOverwriteCompletedConfiguration(t *testing.T) { if failFirst { wantProvider, wantURL = "jellyfin", "http://second-media:8096" } - if store.settings.Provider != wantProvider || store.settings.ServerURL != wantURL || string(media.provider) != wantProvider || len(store.settingWrites) != 4 { + if store.settings.Provider != wantProvider || store.settings.ServerURL != wantURL || handler.(*Server).connectionsStateProvider() != wantProvider || len(store.settingWrites) != 4 { t.Fatalf("setup did not preserve the single successful configuration: stored provider=%q runtime provider=%q writes=%d", store.settings.Provider, media.provider, len(store.settingWrites)) } }) diff --git a/internal/httpserver/handlers_invites.go b/internal/httpserver/handlers_invites.go index 0835a35..b07d8eb 100644 --- a/internal/httpserver/handlers_invites.go +++ b/internal/httpserver/handlers_invites.go @@ -11,22 +11,39 @@ import ( ) func (s *Server) invitesList(w http.ResponseWriter, r *http.Request, session db.Session) { - invites, err := s.store.ListInvites(r.Context()) + before := historyCursor(r) + invites, err := s.store.InvitePage(r.Context(), before, 51) if err != nil { s.error(w, err) return } - activity, err := s.store.LatestInviteActivity(r.Context()) + next := int64(0) + if len(invites) > 50 { + invites = invites[:50] + next = invites[49].ID + } + ids := make([]int64, len(invites)) + for i, invite := range invites { + ids[i] = invite.ID + } + activity, err := s.store.InvitePageActivity(r.Context(), ids) if err != nil { s.error(w, err) return } data := s.data(r, session) + data.HistoryBefore = before + data.NextBefore = next data.Invites = invites data.InviteRows = inviteRows(invites, activity) _, publicURL, _ := s.runtimeSettings() data.PublicURL = strings.TrimRight(publicURL, "/") - data.Stats = dashboardStatsFrom(invites, nil) + counts, err := s.store.DashboardCounts(r.Context(), session.BindingID) + if err != nil { + s.error(w, err) + return + } + data.Stats.ActiveInvites = counts.ActiveInvites render(w, "invites", data) } func (s *Server) invitesNew(w http.ResponseWriter, r *http.Request, session db.Session) { @@ -96,6 +113,7 @@ func (s *Server) invitesCreate(w http.ResponseWriter, r *http.Request, session d } hash := security.HashToken(settings.InviteSecret, token) inviteID, err := s.store.CreateInvite(r.Context(), db.Invite{ + BindingID: session.BindingID, TokenHash: hash, TokenPrefix: security.Prefix(token), Token: token, @@ -139,7 +157,7 @@ func (s *Server) setInviteState(w http.ResponseWriter, r *http.Request, session s.message(w, "Invalid invite", "That invite does not exist.", http.StatusBadRequest) return } - if err := s.store.SetInviteEnabled(r.Context(), id, enabled); err != nil { + if err := s.store.SetInviteEnabled(r.Context(), id, session.BindingID, enabled); err != nil { s.inviteError(w, err) return } diff --git a/internal/httpserver/handlers_invites_test.go b/internal/httpserver/handlers_invites_test.go index 8d14304..622cfbf 100644 --- a/internal/httpserver/handlers_invites_test.go +++ b/internal/httpserver/handlers_invites_test.go @@ -10,7 +10,7 @@ import ( func TestInvitesListShowsSavedCopyButtons(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) req := adminRequest(t, http.MethodGet, "/admin/invites", nil) rr := httptest.NewRecorder() @@ -32,7 +32,7 @@ func TestInvitesListShowsSavedCopyButtons(t *testing.T) { func TestInvitesNewCustomExpiryIsProgressivelyEnhanced(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) req := adminRequest(t, http.MethodGet, "/admin/invites/new", nil) rr := httptest.NewRecorder() @@ -57,7 +57,7 @@ func TestInvitesNewCustomExpiryIsProgressivelyEnhanced(t *testing.T) { func TestInvitesNewCanReuseExistingInviteAsPreset(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) req := adminRequest(t, http.MethodGet, "/admin/invites/new?preset_id=1", nil) rr := httptest.NewRecorder() @@ -76,7 +76,7 @@ func TestInvitesNewCanReuseExistingInviteAsPreset(t *testing.T) { func TestInvitesNewIgnoresMissingPreset(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) req := adminRequest(t, http.MethodGet, "/admin/invites/new?preset_id=404", nil) rr := httptest.NewRecorder() @@ -99,7 +99,7 @@ func TestInvitesNewIgnoresMissingPreset(t *testing.T) { func TestInvitesCreateStoresUserExpiryDaysWithoutPuttingTokenInRedirect(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) form := url.Values{ "csrf": {store.session.CSRFSecret}, "label": {"Family"}, @@ -138,7 +138,7 @@ func TestInvitesCreateRequiresAPIKey(t *testing.T) { cfg := testConfig() cfg.APIKey = "" cfg.APIKeyManaged = false - handler := New(cfg, store, &fakeMediaServer{}) + handler := New(cfg, store, testMediaFactory(&fakeMediaServer{})) form := url.Values{ "csrf": {store.session.CSRFSecret}, "label": {"Keep this label"}, @@ -168,7 +168,7 @@ func TestInvitesCreateRequiresAPIKey(t *testing.T) { func TestInvitesCreateValidationRendersFormAndPreservesSafeFields(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) form := url.Values{ "csrf": {store.session.CSRFSecret}, "label": {"Keep this label"}, @@ -204,7 +204,7 @@ func TestInvitesCreateValidationPreservesQuickExpiryChoice(t *testing.T) { for _, choice := range []string{"1", "7", "30"} { t.Run(choice, func(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) form := url.Values{ "csrf": {store.session.CSRFSecret}, "label": {"Family"}, @@ -236,7 +236,7 @@ func TestInvitesCreateValidationPreservesQuickExpiryChoice(t *testing.T) { func TestInvitesCreateSupportsQuickExpiryChoices(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) form := url.Values{ "csrf": {store.session.CSRFSecret}, "label": {"Family"}, diff --git a/internal/httpserver/handlers_layout_test.go b/internal/httpserver/handlers_layout_test.go index 1825d00..ce7cc25 100644 --- a/internal/httpserver/handlers_layout_test.go +++ b/internal/httpserver/handlers_layout_test.go @@ -9,7 +9,7 @@ import ( func TestAdminNavShowsDashboardWithoutBrandIcon(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) req := adminRequest(t, http.MethodGet, "/admin", nil) rr := httptest.NewRecorder() @@ -40,7 +40,7 @@ func TestAdminNavShowsDashboardWithoutBrandIcon(t *testing.T) { func TestAuthPagesUseFullViewportShell(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) req := httptest.NewRequest(http.MethodGet, "/login", nil) rr := httptest.NewRecorder() @@ -62,7 +62,7 @@ func TestAuthPagesUseFullViewportShell(t *testing.T) { } func TestAuthTitlesAndAdminNavScript(t *testing.T) { - handler := New(testConfig(), newFakeStore(), &fakeMediaServer{}) + handler := New(testConfig(), newFakeStore(), testMediaFactory(&fakeMediaServer{})) req := httptest.NewRequest(http.MethodGet, "/guide", nil) rr := httptest.NewRecorder() handler.ServeHTTP(rr, req) @@ -74,7 +74,7 @@ func TestAuthTitlesAndAdminNavScript(t *testing.T) { assetRR := httptest.NewRecorder() handler.ServeHTTP(assetRR, assetReq) asset := assetRR.Body.String() - for _, want := range []string{`nav.admin-nav`, `a[aria-current="page"]`, `scrollIntoView`, `prefers-reduced-motion`} { + for _, want := range []string{`nav.admin-nav`, `.nav-toggle`, `aria-expanded`, `Escape`, `prefers-reduced-motion`} { if !strings.Contains(asset, want) { t.Fatalf("app.js missing %q:\n%s", want, asset) } @@ -83,7 +83,7 @@ func TestAuthTitlesAndAdminNavScript(t *testing.T) { func TestStaticAssetsAndCSPUseExternalCSSAndJS(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) assetReq := httptest.NewRequest(http.MethodGet, "/assets/app.css", nil) assetRR := httptest.NewRecorder() @@ -122,7 +122,7 @@ func TestStaticAssetsAndCSPUseExternalCSSAndJS(t *testing.T) { } func TestHSTSOnlyAppliesToConfiguredHTTPSHost(t *testing.T) { - handler := New(testConfig(), newFakeStore(), &fakeMediaServer{}) + handler := New(testConfig(), newFakeStore(), testMediaFactory(&fakeMediaServer{})) req := httptest.NewRequest(http.MethodGet, "/login", nil) req.Host = "internal-proxy:8099" rr := httptest.NewRecorder() diff --git a/internal/httpserver/handlers_public.go b/internal/httpserver/handlers_public.go index 781dde4..d4d9d19 100644 --- a/internal/httpserver/handlers_public.go +++ b/internal/httpserver/handlers_public.go @@ -2,7 +2,6 @@ package httpserver import ( "context" - "database/sql" "errors" "log/slog" "net/http" @@ -11,6 +10,7 @@ import ( "time" "github.com/mayvqt/aperture/internal/db" + "github.com/mayvqt/aperture/internal/mediaserver" "github.com/mayvqt/aperture/internal/security" ) @@ -18,7 +18,7 @@ const registrationProvisioningTimeout = 2 * time.Minute func (s *Server) publicInvite(w http.ResponseWriter, r *http.Request) { token := r.PathValue("token") - invite, err := s.validInvite(r.Context(), token) + invite, err := s.lookupInvite(r.Context(), token) if err != nil { s.message(w, "Invite unavailable", "This invite is no longer available.", http.StatusNotFound) return @@ -56,7 +56,7 @@ func (s *Server) publicRegister(w http.ResponseWriter, r *http.Request) { } username := strings.TrimSpace(r.FormValue("username")) password := r.FormValue("password") - invite, err := s.validInvite(r.Context(), token) + invite, err := s.lookupInvite(r.Context(), token) if err != nil { s.message(w, "Invite unavailable", "This invite is no longer available.", http.StatusNotFound) return @@ -70,6 +70,17 @@ func (s *Server) publicRegister(w http.ResponseWriter, r *http.Request) { render(w, "public-invite", viewData{AuthTitle: "Create account · Aperture", Token: token, CSRF: csrf, Invite: invite, FormUsername: username, Error: validationError}) return } + ctx, err := s.verifiedAPIContext(r.Context()) + if err != nil { + s.message(w, "Registration unavailable", "Aperture could not verify the media server. Ask the server admin to check the connection.", http.StatusServiceUnavailable) + return + } + r = r.WithContext(ctx) + invite, err = s.validInvite(r.Context(), token) + if err != nil { + s.message(w, "Invite unavailable", "This invite is no longer available.", http.StatusNotFound) + return + } settings, err := s.settings(r.Context()) if err != nil { s.error(w, err) @@ -79,7 +90,12 @@ func (s *Server) publicRegister(w http.ResponseWriter, r *http.Request) { s.message(w, "Registration unavailable", "Ask the server admin to finish configuring account registration.", http.StatusServiceUnavailable) return } - regID, template, err := s.store.ReserveInviteUse(r.Context(), invite.ID, remoteIP, requestUserAgent(r), username) + op, err := operationSnapshot(r.Context()) + if err != nil { + s.error(w, err) + return + } + regID, template, err := s.store.ReserveInviteUse(r.Context(), invite.ID, invite.BindingID, remoteIP, requestUserAgent(r), username) if err != nil { if errors.Is(err, db.ErrInviteUnavailable) { s.message(w, "Invite unavailable", "This invite is no longer available.", http.StatusConflict) @@ -92,23 +108,29 @@ func (s *Server) publicRegister(w http.ResponseWriter, r *http.Request) { // even if the browser disconnects. The operation remains time-bounded. provisionCtx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), registrationProvisioningTimeout) defer cancel() + release, err := s.claimAccountOperations(regID) + if err != nil { + s.error(w, err) + return + } + defer release() r = r.WithContext(provisionCtx) if err := s.store.BeginUserCreation(r.Context(), regID); err != nil { s.error(w, err) return } - user, err := s.media.CreateUser(r.Context(), settings.ServerURL, settings.APIKey, username, password) + user, err := op.Media.CreateUser(r.Context(), settings.ServerURL, settings.APIKey, username, password, func(user mediaserver.User) error { + persistCtx, persistCancel := context.WithTimeout(context.WithoutCancel(r.Context()), 5*time.Second) + defer persistCancel() + return s.store.RecordProvisioningUser(persistCtx, regID, user.ID) + }) provisioned := false if user.ID != "" { defer func() { if provisioned { return } - cleanupCtx, cleanupCancel := context.WithTimeout(context.WithoutCancel(r.Context()), 10*time.Second) - defer cleanupCancel() - if err := s.media.DisableUser(cleanupCtx, settings.ServerURL, settings.APIKey, user.ID); err != nil { - slog.Error("could not disable incomplete account", "registration_id", regID, "error", safeError(err)) - } + s.secureIncompleteAccount(r.Context(), regID, user.ID) }() } if err != nil { @@ -127,15 +149,15 @@ func (s *Server) publicRegister(w http.ResponseWriter, r *http.Request) { } slog.Warn("media-server user creation failed", "error", safeError(err)) s.notify(webhookNotice{Event: "registration.failed", Title: "Registration failed", Color: 0xe74c3c, Fields: map[string]string{"Username": username, "Registration": strconv.FormatInt(regID, 10), "Error": safeError(err)}}) - s.message(w, "Registration failed", "The account could not be created. Ask the server admin to check this invite.", http.StatusInternalServerError) + s.message(w, "Registration failed", "Account setup could not be confirmed. Ask the server admin to check this invite before trying again.", http.StatusInternalServerError) return } if err := s.store.RecordCreatedUser(r.Context(), regID, user.ID); err != nil { s.error(w, err) return } - if err := s.media.ApplyTemplate(r.Context(), settings.ServerURL, settings.APIKey, user.ID, template); err != nil { - if recordErr := s.store.CompleteRegistration(r.Context(), regID, db.RegistrationNeedsAttention, safeError(err), sql.NullTime{}); recordErr != nil { + if err := op.Media.ApplyTemplate(r.Context(), settings.ServerURL, settings.APIKey, user.ID, template); err != nil { + if recordErr := s.store.CompleteRegistration(r.Context(), regID, db.RegistrationNeedsAttention, safeError(err)); recordErr != nil { slog.Error("could not record partial media-server registration", "registration_id", regID, "error", safeError(recordErr)) } slog.Warn("media-server template application failed", "error", safeError(err)) @@ -143,7 +165,7 @@ func (s *Server) publicRegister(w http.ResponseWriter, r *http.Request) { s.message(w, "Account needs review", "The account was created, but an admin needs to finish applying access.", http.StatusAccepted) return } - if err := s.store.CompleteRegistration(r.Context(), regID, db.RegistrationComplete, "", disableAtFor(invite.UserExpiryDays)); err != nil { + if err := s.store.CompleteRegistration(r.Context(), regID, db.RegistrationComplete, ""); err != nil && !s.accountCompletionConfirmed(r.Context(), regID, user.ID) { s.error(w, err) return } diff --git a/internal/httpserver/handlers_public_test.go b/internal/httpserver/handlers_public_test.go index 3aeb448..c59e1aa 100644 --- a/internal/httpserver/handlers_public_test.go +++ b/internal/httpserver/handlers_public_test.go @@ -2,7 +2,6 @@ package httpserver import ( "context" - "database/sql" "errors" "net/http" "net/http/httptest" @@ -20,7 +19,7 @@ func TestPublicRegisterSchedulesUserDisable(t *testing.T) { store := newFakeStore() store.invite.UserExpiryDays = 7 media := &fakeMediaServer{} - handler := New(testConfig(), store, media) + handler := New(testConfig(), store, testMediaFactory(media)) token := "public-invite-token" csrf := "csrf-value" form := url.Values{ @@ -66,7 +65,7 @@ func TestPublicRegisterSchedulesUserDisable(t *testing.T) { func TestPublicRegisterChecksTemplateBeforeReservingInvite(t *testing.T) { store := newFakeStore() store.templateErr = errors.New("template read failed") - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) token := "public-invite-token" csrf := "csrf-value" form := url.Values{ @@ -96,7 +95,7 @@ func TestPublicRegisterChecksTemplateBeforeReservingInvite(t *testing.T) { func TestPublicRegisterRetainsInviteUseAfterUncertainUserCreationFailure(t *testing.T) { store := newFakeStore() media := &fakeMediaServer{createErr: errors.New("create failed")} - handler := New(testConfig(), store, media) + handler := New(testConfig(), store, testMediaFactory(media)) token := "public-invite-token" csrf := "csrf-value" form := url.Values{ @@ -132,7 +131,7 @@ func TestPublicRegisterWithoutAPIKeyDoesNotConsumeInvite(t *testing.T) { cfg := testConfig() cfg.APIKey = "" cfg.APIKeyManaged = false - handler := New(cfg, store, &fakeMediaServer{}) + handler := New(cfg, store, testMediaFactory(&fakeMediaServer{})) token := "public-invite-token" csrf := "csrf-value" form := url.Values{ @@ -161,7 +160,7 @@ func TestPublicInviteShowsOnlyAccountCreationDetails(t *testing.T) { store.invite.Label = "Admin label" store.invite.Template = "Internal template" store.invite.UserExpiryDays = 14 - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) req := httptest.NewRequest(http.MethodGet, "/i/public-token", nil) rr := httptest.NewRecorder() @@ -189,7 +188,7 @@ func TestPublicInviteFailsEarlyWithoutAPIKey(t *testing.T) { cfg := testConfig() cfg.APIKey = "" cfg.APIKeyManaged = false - handler := New(cfg, store, &fakeMediaServer{}) + handler := New(cfg, store, testMediaFactory(&fakeMediaServer{})) req := httptest.NewRequest(http.MethodGet, "/i/public-token", nil) rr := httptest.NewRecorder() @@ -206,7 +205,7 @@ func TestPublicInviteFailsEarlyWithoutAPIKey(t *testing.T) { func TestPublicInviteUnavailableDoesNotRevealDetails(t *testing.T) { store := newFakeStore() store.inviteErr = db.ErrNotFound - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) req := httptest.NewRequest(http.MethodGet, "/i/missing-token", nil) rr := httptest.NewRecorder() @@ -222,7 +221,7 @@ func TestPublicInviteUnavailableDoesNotRevealDetails(t *testing.T) { } func TestGuideRendersAccountCreatedMessage(t *testing.T) { - handler := New(testConfig(), newFakeStore(), &fakeMediaServer{}) + handler := New(testConfig(), newFakeStore(), testMediaFactory(&fakeMediaServer{})) req := httptest.NewRequest(http.MethodGet, "/guide", nil) rr := httptest.NewRecorder() @@ -255,14 +254,14 @@ func (s *provisioningStore) RecordCreatedUser(ctx context.Context, id int64, use return s.fakeStore.RecordCreatedUser(ctx, id, userID) } -func (s *provisioningStore) CompleteRegistration(ctx context.Context, id int64, status, message string, expiry sql.NullTime) error { +func (s *provisioningStore) CompleteRegistration(ctx context.Context, id int64, status, message string) error { if err := ctx.Err(); err != nil { return err } if s.completeError != nil { return s.completeError } - return s.fakeStore.CompleteRegistration(ctx, id, status, message, expiry) + return s.fakeStore.CompleteRegistration(ctx, id, status, message) } type provisioningMedia struct { @@ -271,8 +270,8 @@ type provisioningMedia struct { partialError error } -func (m *provisioningMedia) CreateUser(ctx context.Context, baseURL, key, username, password string) (mediaserver.User, error) { - user, err := m.fakeMediaServer.CreateUser(ctx, baseURL, key, username, password) +func (m *provisioningMedia) CreateUser(ctx context.Context, baseURL, key, username, password string, created func(mediaserver.User) error) (mediaserver.User, error) { + user, err := m.fakeMediaServer.CreateUser(ctx, baseURL, key, username, password, created) if m.cancel != nil { m.cancel() } @@ -316,7 +315,7 @@ func TestPublicRegisterProtectsIncompleteProvisioning(t *testing.T) { case "completion failure": store.completeError = errors.New("database write failed") } - handler := New(testConfig(), store, media) + handler := New(testConfig(), store, testMediaFactory(media)) token, csrf := "public-invite-token", "csrf-value" form := url.Values{"csrf": {csrf}, "username": {"new_user"}, "password": {"correct horse"}, "confirm_password": {"correct horse"}} req := httptest.NewRequest(http.MethodPost, "/i/"+token+"/register", strings.NewReader(form.Encode())).WithContext(ctx) diff --git a/internal/httpserver/handlers_registrations.go b/internal/httpserver/handlers_registrations.go index 7d06a1e..f278f54 100644 --- a/internal/httpserver/handlers_registrations.go +++ b/internal/httpserver/handlers_registrations.go @@ -2,7 +2,6 @@ package httpserver import ( "errors" - "log/slog" "net/http" "strconv" "strings" @@ -11,12 +10,20 @@ import ( ) func (s *Server) registrationsList(w http.ResponseWriter, r *http.Request, session db.Session) { - regs, err := s.store.RecentRegistrations(r.Context(), 100) + before := historyCursor(r) + review := r.URL.Query().Get("review") == "1" + regs, err := s.store.RegistrationPage(r.Context(), session.BindingID, before, review, 51) if err != nil { s.error(w, err) return } data := s.data(r, session) + data.HistoryBefore = before + data.ReviewOnly = review + if len(regs) > 50 { + regs = regs[:50] + data.NextBefore = regs[49].ID + } data.Registrations = regs render(w, "registrations", data) } @@ -36,26 +43,14 @@ func (s *Server) registrationsRetryTemplate(w http.ResponseWriter, r *http.Reque s.message(w, "API key required", "Save a media-server API key before retrying template application.", http.StatusConflict) return } - recovery, err := s.store.ClaimTemplateRecovery(r.Context(), id) - if err != nil { - s.registrationRecoveryError(w, err) - return - } - if strings.TrimSpace(recovery.Template.PolicyJSON) == "" { - s.recordTemplateRetryFailure(r, id, errors.New("saved registration template is unavailable")) - s.message(w, "Recovery unavailable", "The saved registration template is unavailable.", http.StatusConflict) - return - } - if err := s.media.ApplyTemplate(r.Context(), settings.ServerURL, settings.APIKey, recovery.Registration.ExternalUserID.String, recovery.Template); err != nil { - s.recordTemplateRetryFailure(r, id, err) - s.message(w, "Template retry failed", "The media server did not accept the template. Review the saved details and try again.", http.StatusBadGateway) - return - } - if err := s.store.CompleteTemplateRecovery(r.Context(), id, disableAtFrom(recovery.Registration.CreatedAt, recovery.UserExpiryDays)); err != nil { - s.registrationRecoveryError(w, err) + if err := s.recoverAccount(r.Context(), id, false); err != nil { + if errors.Is(err, db.ErrNotFound) || errors.Is(err, db.ErrRegistrationTransition) { + s.registrationRecoveryError(w, err) + } else { + s.message(w, "Template retry failed", "Access could not be confirmed. Aperture will keep trying to disable the incomplete account; review the registration details before trying again.", http.StatusBadGateway) + } return } - s.notify(webhookNotice{Event: "template.recovered", Title: "Access template recovered", Color: 0x2ecc71, Fields: map[string]string{"Username": recovery.Registration.Username, "Registration": strconv.FormatInt(id, 10), "Template": recovery.Template.Name}}) s.audit(r, session, "registration.retry_template", "registration", strconv.FormatInt(id, 10), nil) http.Redirect(w, r, "/admin/registrations", http.StatusSeeOther) } @@ -71,6 +66,15 @@ func (s *Server) registrationsDelete(w http.ResponseWriter, r *http.Request, ses s.registrationRecoveryError(w, err) return } + op, err := operationSnapshot(r.Context()) + if err != nil { + s.error(w, err) + return + } + if registration.BindingID != op.Identity.Binding.ID { + s.message(w, "Server review needed", "Assign this registration to the correct server before managing its account.", http.StatusConflict) + return + } deletedUpstream := false if registration.ExternalUserID.Valid && strings.TrimSpace(registration.ExternalUserID.String) != "" { settings, err := s.settings(r.Context()) @@ -82,14 +86,22 @@ func (s *Server) registrationsDelete(w http.ResponseWriter, r *http.Request, ses s.message(w, "API key required", "Aperture must verify that the media-server user has been deleted first.", http.StatusConflict) return } - deletedUpstream, err = s.deleteUserAndRecords(r.Context(), settings, registration.ExternalUserID.String) + deletedUpstream, err = s.deleteUserAndRecords(r.Context(), registration.ExternalUserID.String) if err != nil { s.userDeleteError(w, err, deletedUpstream) return } - } else if err := s.store.DeleteRegistration(r.Context(), id); err != nil { - s.registrationRecoveryError(w, err) - return + } else { + release, err := s.claimAccountOperations(id) + if err != nil { + s.registrationRecoveryError(w, err) + return + } + defer release() + if err := s.store.DeleteRegistration(r.Context(), id); err != nil { + s.registrationRecoveryError(w, err) + return + } } s.audit(r, session, "registration.delete", "registration", strconv.FormatInt(id, 10), map[string]any{"username": registration.Username, "deleted_from_media_server": deletedUpstream}) destination := "/admin/registrations" @@ -99,12 +111,6 @@ func (s *Server) registrationsDelete(w http.ResponseWriter, r *http.Request, ses http.Redirect(w, r, destination, http.StatusSeeOther) } -func (s *Server) recordTemplateRetryFailure(r *http.Request, registrationID int64, err error) { - if recordErr := s.store.RecordTemplateRetryFailure(r.Context(), registrationID, safeError(err)); recordErr != nil { - slog.Error("could not record template retry failure", "registration_id", registrationID, "error", safeError(recordErr)) - } -} - func (s *Server) registrationRecoveryError(w http.ResponseWriter, err error) { switch { case errors.Is(err, db.ErrNotFound): diff --git a/internal/httpserver/handlers_registrations_test.go b/internal/httpserver/handlers_registrations_test.go index 3e213f8..1cd3f6f 100644 --- a/internal/httpserver/handlers_registrations_test.go +++ b/internal/httpserver/handlers_registrations_test.go @@ -17,7 +17,7 @@ func TestRegistrationTemplateRetryUsesSavedSnapshot(t *testing.T) { store := newFakeStore() store.recovery.Template.PolicyJSON = `{"EnableAllFolders":false}` media := &fakeMediaServer{} - handler := New(testConfig(), store, media) + handler := New(testConfig(), store, testMediaFactory(media)) form := url.Values{"csrf": {store.session.CSRFSecret}} req := adminRequest(t, http.MethodPost, "/admin/registrations/1/retry-template", strings.NewReader(form.Encode())) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") @@ -39,12 +39,13 @@ func TestRegistrationTemplateRetryUsesSavedSnapshot(t *testing.T) { func TestRegistrationsPageShowsTemplateRecoveryAction(t *testing.T) { store := newFakeStore() store.registrations = []db.Registration{{ + BindingID: 1, ID: 7, Status: db.RegistrationNeedsAttention, ExternalUserID: sql.NullString{String: "media-alice", Valid: true}, ErrorMessage: sql.NullString{String: "policy failed", Valid: true}, }} - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) req := adminRequest(t, http.MethodGet, "/admin/registrations", nil) rr := httptest.NewRecorder() @@ -80,7 +81,7 @@ func TestRegistrationsPageShowsTemplateRecoveryAction(t *testing.T) { func TestRegistrationTemplateRetryRecordsFailure(t *testing.T) { store := newFakeStore() media := &fakeMediaServer{applyErr: errors.New("apply failed")} - handler := New(testConfig(), store, media) + handler := New(testConfig(), store, testMediaFactory(media)) form := url.Values{"csrf": {store.session.CSRFSecret}} req := adminRequest(t, http.MethodPost, "/admin/registrations/1/retry-template", strings.NewReader(form.Encode())) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") @@ -103,12 +104,12 @@ func TestRegistrationDeleteRemovesUpstreamUserWhenPresent(t *testing.T) { wantUpstreamDelete bool }{ {name: "user already deleted"}, - {name: "user still exists", liveUsers: []mediaserver.User{{ID: "media-alice", Name: "alice"}}, wantUpstreamDelete: true}, + {name: "user still exists", liveUsers: []mediaserver.User{{ID: "media-alice", Name: "alice", Policy: []byte(`{"IsAdministrator":false}`)}}, wantUpstreamDelete: true}, } { t.Run(test.name, func(t *testing.T) { store := newFakeStore() media := &fakeMediaServer{users: test.liveUsers} - handler := New(testConfig(), store, media) + handler := New(testConfig(), store, testMediaFactory(media)) form := url.Values{"csrf": {store.session.CSRFSecret}} req := adminRequest(t, http.MethodPost, "/admin/registrations/1/delete", strings.NewReader(form.Encode())) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") diff --git a/internal/httpserver/handlers_templates.go b/internal/httpserver/handlers_templates.go index 74d7d0c..74b0015 100644 --- a/internal/httpserver/handlers_templates.go +++ b/internal/httpserver/handlers_templates.go @@ -130,7 +130,12 @@ func (s *Server) templatesImport(w http.ResponseWriter, r *http.Request, session s.message(w, "Import needs media-server access", "Save an API key in settings or log in again with a media-server administrator account.", http.StatusBadRequest) return } - imported, err := s.media.ImportTemplate(r.Context(), settings.ServerURL, token, deviceID, userRef) + op, err := operationSnapshot(r.Context()) + if err != nil { + s.error(w, err) + return + } + imported, err := op.Media.ImportTemplate(r.Context(), settings.ServerURL, token, deviceID, userRef) if err != nil { slog.Warn("template import failed", "error", safeError(err)) status := http.StatusBadGateway diff --git a/internal/httpserver/handlers_templates_test.go b/internal/httpserver/handlers_templates_test.go index 9383ae6..fa67bb7 100644 --- a/internal/httpserver/handlers_templates_test.go +++ b/internal/httpserver/handlers_templates_test.go @@ -13,7 +13,7 @@ import ( func TestTemplatesCreateStripsAdminPrivilege(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) form := url.Values{ "csrf": {store.session.CSRFSecret}, "name": {"Imported"}, @@ -43,7 +43,7 @@ func TestTemplatesCreateStripsAdminPrivilege(t *testing.T) { func TestTemplatesPageUsesExpandableWorkflows(t *testing.T) { store := newFakeStore() store.template.IsDefault = false - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) rr := httptest.NewRecorder() handler.ServeHTTP(rr, adminRequest(t, http.MethodGet, "/admin/templates", nil)) @@ -69,7 +69,7 @@ func TestTemplateDetailShowsPreviewAndUpdatesTemplate(t *testing.T) { Description: "limited", PolicyJSON: `{"IsAdministrator":false,"EnabledFolders":["Movies"]}`, } - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) req := adminRequest(t, http.MethodGet, "/admin/templates/7", nil) rr := httptest.NewRecorder() @@ -130,7 +130,7 @@ func TestTemplateDetailShowsPreviewAndUpdatesTemplate(t *testing.T) { func TestTemplateDefaultAndDeleteHandlers(t *testing.T) { store := newFakeStore() store.template = db.Template{ID: 7, Name: "Kids", PolicyJSON: `{"IsAdministrator":false}`} - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) form := url.Values{"csrf": {store.session.CSRFSecret}} req := adminRequest(t, http.MethodPost, "/admin/templates/7/default", strings.NewReader(form.Encode())) @@ -160,7 +160,7 @@ func TestTemplateDefaultAndDeleteHandlers(t *testing.T) { func TestTemplateImportReturnsNotFoundForUnknownMediaUser(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{importErr: mediaserver.ErrUserNotFound}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{importErr: mediaserver.ErrUserNotFound})) form := url.Values{ "csrf": {store.session.CSRFSecret}, "name": {"Imported"}, diff --git a/internal/httpserver/health.go b/internal/httpserver/health.go index 587608b..14e282f 100644 --- a/internal/httpserver/health.go +++ b/internal/httpserver/health.go @@ -17,7 +17,11 @@ type mediaHealthCache struct { } func (s *Server) mediaHealthy(ctx context.Context, baseURL, apiKey string) bool { - provider, _, _ := s.runtimeSettings() + snapshot, err := s.snapshot(ctx) + if err != nil { + return false + } + provider := snapshot.Settings.Provider key := sha256.Sum256([]byte(provider + "\x00" + baseURL + "\x00" + apiKey)) now := time.Now() @@ -29,7 +33,7 @@ func (s *Server) mediaHealthy(ctx context.Context, baseURL, apiKey string) bool pingCtx, cancel := context.WithTimeout(ctx, 3*time.Second) defer cancel() - healthy := s.media.Ping(pingCtx, baseURL, apiKey) == nil + healthy := snapshot.Media.Ping(pingCtx, baseURL, apiKey) == nil if ctx.Err() == nil { s.healthCache.key = key s.healthCache.checkedAt = time.Now() diff --git a/internal/httpserver/health_test.go b/internal/httpserver/health_test.go index b232f69..b47dd16 100644 --- a/internal/httpserver/health_test.go +++ b/internal/httpserver/health_test.go @@ -2,6 +2,7 @@ package httpserver import ( "context" + "github.com/mayvqt/aperture/internal/connection" "sync" "testing" ) @@ -11,7 +12,7 @@ func TestMediaHealthCacheCoalescesConcurrentChecks(t *testing.T) { pingStarted: make(chan struct{}), pingRelease: make(chan struct{}), } - s := &Server{media: media} + s := NewServer(testConfig(), newFakeStore(), testMediaFactory(media)) const requests = 20 results := make(chan bool, requests) @@ -40,14 +41,19 @@ func TestMediaHealthCacheCoalescesConcurrentChecks(t *testing.T) { func TestMediaHealthCacheSeparatesProviders(t *testing.T) { media := &fakeMediaServer{} - s := &Server{media: media, provider: "jellyfin"} + s := NewServer(testConfig(), newFakeStore(), testMediaFactory(media)) if !s.mediaHealthy(t.Context(), "http://media:8096", "api-key") { t.Fatal("initial health check failed") } media.pingErr = errFakePing - s.setRuntime("emby", "", false) - if s.mediaHealthy(t.Context(), "http://media:8096", "api-key") { + snap, err := s.snapshot(t.Context()) + if err != nil { + t.Fatal(err) + } + snap.Settings.Provider = "emby" + ctx := connection.WithSnapshot(t.Context(), snap) + if s.mediaHealthy(ctx, "http://media:8096", "api-key") { t.Fatal("Emby health check reused Jellyfin cache entry") } if media.pingCalls.Load() != 2 { diff --git a/internal/httpserver/helpers.go b/internal/httpserver/helpers.go index 4807782..281a21a 100644 --- a/internal/httpserver/helpers.go +++ b/internal/httpserver/helpers.go @@ -11,17 +11,22 @@ import ( "time" "github.com/mayvqt/aperture/internal/config" + "github.com/mayvqt/aperture/internal/connection" "github.com/mayvqt/aperture/internal/db" "github.com/mayvqt/aperture/internal/mediaserver" "github.com/mayvqt/aperture/internal/security" ) -func (s *Server) validInvite(ctx context.Context, token string) (db.Invite, error) { - settings, err := s.settings(ctx) +func (s *Server) lookupInvite(ctx context.Context, token string) (db.Invite, error) { + snapshot, err := s.snapshot(ctx) if err != nil { return db.Invite{}, err } - invite, err := s.store.InviteByHash(ctx, security.HashToken(settings.InviteSecret, token)) + if snapshot.Identity.Binding.ID <= 0 { + return db.Invite{}, db.ErrInviteUnavailable + } + settings := snapshot.Settings + invite, err := s.store.InviteByHash(ctx, security.HashToken(settings.InviteSecret, token), snapshot.Identity.Binding.ID) if err != nil { return db.Invite{}, err } @@ -31,43 +36,55 @@ func (s *Server) validInvite(ctx context.Context, token string) (db.Invite, erro } return invite, nil } +func (s *Server) validInvite(ctx context.Context, token string) (db.Invite, error) { + if _, err := operationSnapshot(ctx); err != nil { + return db.Invite{}, err + } + return s.lookupInvite(ctx, token) +} func (s *Server) serverName() string { value, _, _ := s.runtimeSettings() provider, _ := mediaserver.ParseProvider(value) return provider.Name() } -func (s *Server) settings(ctx context.Context) (db.Settings, error) { - settings, err := s.store.Settings(ctx) - if err != nil { - return db.Settings{}, err - } - if s.cfg.ProviderManaged { - settings.Provider = s.cfg.MediaProvider - } - if s.cfg.PublicURLManaged { - settings.PublicURL = s.cfg.PublicURL +func (s *Server) snapshot(ctx context.Context) (connection.Snapshot, error) { + if snapshot, ok := connection.FromContext(ctx); ok { + return snapshot, nil } - if s.cfg.ServerURLManaged { - settings.ServerURL = s.cfg.ServerURL - } - if s.cfg.APIKey != "" { - settings.APIKey = s.cfg.APIKey + return s.connections.Current(ctx) +} + +func (s *Server) settings(ctx context.Context) (db.Settings, error) { + snapshot, err := s.snapshot(ctx) + return snapshot.Settings, err +} + +func operationSnapshot(ctx context.Context) (connection.Snapshot, error) { + snapshot, ok := connection.FromContext(ctx) + if !ok || !snapshot.Verified || snapshot.Identity.Binding.ID <= 0 { + return connection.Snapshot{}, connection.ErrUnavailable } - if s.cfg.SessionSecret != "" { - settings.SessionSecret = s.cfg.SessionSecret + return snapshot, nil +} + +func (s *Server) verifiedAPIContext(ctx context.Context) (context.Context, error) { + snapshot, err := s.snapshot(ctx) + if err != nil { + return nil, err } - if s.cfg.InviteSecret != "" { - settings.InviteSecret = s.cfg.InviteSecret + verified, err := s.connections.Verify(ctx, snapshot, snapshot.Settings.APIKey, "aperture") + if err != nil { + return nil, err } - return settings, nil + return connection.WithSnapshot(ctx, verified), nil } + func (s *Server) setupComplete(ctx context.Context) (bool, error) { settings, err := s.settings(ctx) if err != nil { return false, err } - _, publicURL, _ := s.runtimeSettings() - return settings.Provider != "" && settings.ServerURL != "" && publicURL != "", nil + return settings.Provider != "" && settings.ServerURL != "" && settings.PublicURL != "", nil } func (s *Server) message(w http.ResponseWriter, title, message string, status int) { renderStatus(w, "message", viewData{AuthTitle: "Message · Aperture", Title: title, Message: message}, status) @@ -208,3 +225,24 @@ func safeAdminImportError(err error) string { return "Aperture could not import that media-server user's template: " + msg } } + +// A maintenance batch may span minutes. Verify each new account operation, +// stopping the old batch if a replacement or settings change occurred. +func (s *Server) refreshAccountContext(ctx context.Context) (context.Context, error) { + previous, err := operationSnapshot(ctx) + if err != nil { + return nil, err + } + fresh, err := s.verifiedAPIContext(ctx) + if err != nil { + return nil, err + } + current, err := operationSnapshot(fresh) + if err != nil { + return nil, err + } + if current.Identity.Binding.ID != previous.Identity.Binding.ID || current.Identity.Generation != previous.Identity.Generation { + return nil, db.ErrConnectionChanged + } + return fresh, nil +} diff --git a/internal/httpserver/helpers_test.go b/internal/httpserver/helpers_test.go index 4ee20dc..3b1ff34 100644 --- a/internal/httpserver/helpers_test.go +++ b/internal/httpserver/helpers_test.go @@ -46,7 +46,7 @@ func TestClientIPRejectsSpoofedForwardedPrefix(t *testing.T) { func TestValidateBaseURL(t *testing.T) { cfg := testConfig() - s := &Server{cfg: cfg, provider: cfg.MediaProvider} + s := &Server{cfg: cfg} for _, value := range []string{"http://jellyfin:8096", "https://jellyfin.example"} { if _, err := s.validateServerURL(value); err != nil { t.Fatalf("validateBaseURL(%q) unexpected error: %v", value, err) @@ -97,19 +97,6 @@ type assertErr string func (e assertErr) Error() string { return string(e) } -func TestDisableAtFor(t *testing.T) { - if got := disableAtFor(0); got.Valid { - t.Fatalf("disableAtFor(0) = %#v, want invalid", got) - } - got := disableAtFor(2) - if !got.Valid { - t.Fatal("disableAtFor(2) should be valid") - } - if days := int(time.Until(got.Time).Hours() / 24); days < 1 || days > 2 { - t.Fatalf("disableAtFor(2) is about %d days away, want about 2", days) - } -} - func TestDashboardStatsFrom(t *testing.T) { now := time.Now() stats := dashboardStatsFrom([]db.Invite{ diff --git a/internal/httpserver/maintenance_worker_test.go b/internal/httpserver/maintenance_worker_test.go index d8ae792..37f57b0 100644 --- a/internal/httpserver/maintenance_worker_test.go +++ b/internal/httpserver/maintenance_worker_test.go @@ -5,20 +5,17 @@ import ( "database/sql" "errors" "testing" - "time" - "github.com/mayvqt/aperture/internal/config" "github.com/mayvqt/aperture/internal/db" ) func TestMaintenanceWorkerReconcilesAndProcessesExpiryImmediately(t *testing.T) { store := newFakeStore() - store.dueDisables = []db.Registration{{ID: 42, ExternalUserID: sql.NullString{String: "expired-user", Valid: true}}} + store.dueDisables = []db.Registration{{BindingID: 1, ID: 42, CleanupPending: true, ExternalUserID: sql.NullString{String: "expired-user", Valid: true}}} media := &fakeMediaServer{} - ctx, cancel := context.WithCancel(context.Background()) - cancel() - - runMaintenanceWorker(ctx, config.Config{}, store, media, time.Hour) + ctx := context.Background() + s := NewServer(testConfig(), store, testMediaFactory(media)) + s.runMaintenance(ctx) if !store.reconciledStale { t.Fatal("maintenance worker did not reconcile stale registrations") } @@ -35,12 +32,11 @@ func TestMaintenanceWorkerReconcilesAndProcessesExpiryImmediately(t *testing.T) func TestMaintenanceWorkerRecordsExpiredUserDisableFailure(t *testing.T) { store := newFakeStore() - store.dueDisables = []db.Registration{{ID: 42, ExternalUserID: sql.NullString{String: "expired-user", Valid: true}}} + store.dueDisables = []db.Registration{{BindingID: 1, ID: 42, CleanupPending: true, ExternalUserID: sql.NullString{String: "expired-user", Valid: true}}} media := &fakeMediaServer{disableErr: errors.New("media unavailable")} - ctx, cancel := context.WithCancel(context.Background()) - cancel() - - runMaintenanceWorker(ctx, config.Config{}, store, media, time.Hour) + ctx := context.Background() + s := NewServer(testConfig(), store, testMediaFactory(media)) + s.runMaintenance(ctx) if store.markedDisabledID != 0 || store.failedDisableID != 42 { t.Fatalf("disabled/failed IDs = %d/%d", store.markedDisabledID, store.failedDisableID) } diff --git a/internal/httpserver/middleware.go b/internal/httpserver/middleware.go index 21ce5f4..299554b 100644 --- a/internal/httpserver/middleware.go +++ b/internal/httpserver/middleware.go @@ -22,9 +22,8 @@ func (s *Server) securityHeaders(next http.Handler) http.Handler { w.Header().Set("Content-Security-Policy", "default-src 'self'; style-src 'self'; script-src 'self'; form-action 'self'; frame-ancestors 'none'") w.Header().Set("Cross-Origin-Opener-Policy", "same-origin") w.Header().Set("X-Frame-Options", "DENY") - s.runtimeMu.RLock() - hstsHost := s.hstsHost - s.runtimeMu.RUnlock() + _, publicURL, _ := s.runtimeSettings() + hstsHost := securePublicHost(publicURL) if hstsHost != "" && strings.EqualFold(hstsHost, r.Host) { w.Header().Set("Strict-Transport-Security", "max-age=31536000") } diff --git a/internal/httpserver/routes.go b/internal/httpserver/routes.go index 551bc09..97656d7 100644 --- a/internal/httpserver/routes.go +++ b/internal/httpserver/routes.go @@ -31,6 +31,9 @@ func (s *Server) routes() { s.mux.HandleFunc("POST /admin/templates/{id}/delete", s.adminPost(s.templatesDelete)) s.mux.HandleFunc("POST /admin/templates/import-from-user", s.adminPost(s.templatesImport)) s.mux.HandleFunc("GET /admin/registrations", s.admin(s.registrationsList)) + s.mux.HandleFunc("GET /admin/server-review/{kind}/{id}", s.admin(s.serverReview)) + s.mux.HandleFunc("POST /admin/server-review/{kind}/{id}", s.adminPost(s.serverReview)) + s.mux.HandleFunc("GET /admin/users/review", s.admin(s.managedUserReview)) s.mux.HandleFunc("GET /admin/users", s.admin(s.usersList)) s.mux.HandleFunc("POST /admin/users/{id}/track", s.adminPost(s.usersTrack)) s.mux.HandleFunc("POST /admin/users/{id}/delete", s.adminPost(s.usersDelete)) diff --git a/internal/httpserver/server.go b/internal/httpserver/server.go index e63046d..8c60b65 100644 --- a/internal/httpserver/server.go +++ b/internal/httpserver/server.go @@ -9,6 +9,7 @@ import ( "time" "github.com/mayvqt/aperture/internal/config" + "github.com/mayvqt/aperture/internal/connection" ) const ( @@ -20,46 +21,88 @@ const ( ) type Server struct { - cfg config.Config - store Store - media MediaServer - mux *http.ServeMux - limiter *rateLimiter - loginIPLimiter *rateLimiter - healthCache mediaHealthCache - trustedProxies []*net.IPNet - hstsHost string - runtimeMu sync.RWMutex - setupMu sync.Mutex - provider string - runtimePublicURL string - cookieSecure bool - webhookWG sync.WaitGroup + cfg config.Config + store Store + connections *connection.Manager + mux *http.ServeMux + limiter *rateLimiter + loginIPLimiter *rateLimiter + healthCache mediaHealthCache + trustedProxies []*net.IPNet + setupMu sync.Mutex + webhookWG sync.WaitGroup + accountMu sync.Mutex + activeAccounts map[int64]bool + activeUsers map[string]bool + accountWG sync.WaitGroup + httpWG sync.WaitGroup + closing bool + handler http.Handler } -func New(cfg config.Config, store Store, media MediaServer) http.Handler { - handler, _ := NewWithShutdown(cfg, store, media) +func New(cfg config.Config, store Store, factory connection.Factory) http.Handler { + handler, _ := NewWithShutdown(cfg, store, factory) return handler } // NewWithShutdown returns the HTTP handler and a function that waits for // accepted webhook deliveries to finish during graceful shutdown. -func NewWithShutdown(cfg config.Config, store Store, media MediaServer) (http.Handler, func(context.Context) error) { +func NewWithShutdown(cfg config.Config, store Store, factory connection.Factory) (http.Handler, func(context.Context) error) { + s := NewServer(cfg, store, factory) + return s, s.Drain +} + +// NewServer owns HTTP, maintenance and accepted background deliveries together. +func NewServer(cfg config.Config, store Store, factory connection.Factory) *Server { s := &Server{ - cfg: cfg, - store: store, - media: media, - mux: http.NewServeMux(), - limiter: newRateLimiter(10, 10*time.Minute), - loginIPLimiter: newRateLimiter(50, 10*time.Minute), - trustedProxies: parseTrustedProxies(cfg.TrustedProxyCIDRs), - hstsHost: securePublicHost(cfg.PublicURL), - provider: cfg.MediaProvider, - runtimePublicURL: cfg.PublicURL, - cookieSecure: cfg.CookieSecure, + cfg: cfg, + store: store, + connections: connection.New(cfg, store, factory), + mux: http.NewServeMux(), + limiter: newRateLimiter(10, 10*time.Minute), + loginIPLimiter: newRateLimiter(50, 10*time.Minute), + trustedProxies: parseTrustedProxies(cfg.TrustedProxyCIDRs), } s.routes() - return s.securityHeaders(s.mux), s.waitForWebhooks + s.handler = s.securityHeaders(s.mux) + return s +} + +func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { + s.accountMu.Lock() + if s.closing { + s.accountMu.Unlock() + http.Error(w, "Aperture is restarting. Try again shortly.", http.StatusServiceUnavailable) + return + } + s.httpWG.Add(1) + s.accountMu.Unlock() + defer s.httpWG.Done() + snapshot, err := s.connections.Current(r.Context()) + if err != nil { + s.error(w, err) + return + } + s.handler.ServeHTTP(w, r.WithContext(connection.WithSnapshot(r.Context(), snapshot))) +} + +func (s *Server) CloseAdmission() { + s.accountMu.Lock() + s.closing = true + s.accountMu.Unlock() +} + +// Drain is called after the HTTP server and maintenance stop accepting work. +func (s *Server) Drain(ctx context.Context) error { + s.CloseAdmission() + done := make(chan struct{}) + go func() { s.httpWG.Wait(); s.accountWG.Wait(); close(done) }() + select { + case <-done: + return s.waitForWebhooks(ctx) + case <-ctx.Done(): + return ctx.Err() + } } func (s *Server) waitForWebhooks(ctx context.Context) error { @@ -76,29 +119,26 @@ func (s *Server) waitForWebhooks(ctx context.Context) error { } } -func (s *Server) setRuntime(provider, publicURL string, cookieSecure bool) { - s.runtimeMu.Lock() - s.provider = provider - s.runtimePublicURL = publicURL - s.cookieSecure = cookieSecure - s.hstsHost = securePublicHost(publicURL) - s.runtimeMu.Unlock() +func (s *Server) Initialize(ctx context.Context) error { + _, err := s.connections.Current(ctx) + return err } func (s *Server) runtimeSettings() (provider, publicURL string, cookieSecure bool) { - s.runtimeMu.RLock() - defer s.runtimeMu.RUnlock() - return s.provider, s.runtimePublicURL, s.cookieSecure -} - -func RunMaintenanceWorker(ctx context.Context, cfg config.Config, store Store, media MediaServer) { - runMaintenanceWorker(ctx, cfg, store, media, maintenanceInterval) + if s.connections != nil { + if current, ok := s.connections.Peek(); ok { + return current.Settings.Provider, current.Settings.PublicURL, current.CookieSecure + } + } + return s.cfg.MediaProvider, s.cfg.PublicURL, s.cfg.CookieSecure } -func runMaintenanceWorker(ctx context.Context, cfg config.Config, store Store, media MediaServer, interval time.Duration) { - s := &Server{cfg: cfg, store: store, media: media} +func (s *Server) RunMaintenance(ctx context.Context) { + if ctx.Err() != nil { + return + } s.runMaintenance(ctx) - ticker := time.NewTicker(interval) + ticker := time.NewTicker(maintenanceInterval) defer ticker.Stop() for { select { @@ -112,8 +152,10 @@ func runMaintenanceWorker(ctx context.Context, cfg config.Config, store Store, m func (s *Server) runMaintenance(ctx context.Context) { s.reconcileAndLogStaleRegistrations(ctx) - s.processDueTemplateRetries(ctx) - s.processAndLogDueUserDisables(ctx) + if verified, err := s.verifiedAPIContext(ctx); err == nil { + s.processAndLogDueUserDisables(verified) + s.processDueTemplateRetries(verified) + } s.pruneAuditLog(ctx) } diff --git a/internal/httpserver/server_review.go b/internal/httpserver/server_review.go new file mode 100644 index 0000000..de5ca66 --- /dev/null +++ b/internal/httpserver/server_review.go @@ -0,0 +1,178 @@ +package httpserver + +import ( + "context" + "github.com/mayvqt/aperture/internal/db" + "net/http" + "strconv" + "time" +) + +// Review is deliberately per record: sharing an invite never proves that its +// historical accounts belong to the server currently configured. +type serverReview struct { + Kind string + ID, BindingID int64 + Name, UserID, Source, Destination, AccountName string + Expiry string +} + +func (s *Server) reviewRecord(ctx context.Context, kind string, id int64) (serverReview, error) { + v := serverReview{Kind: kind, ID: id} + switch kind { + case "invite": + r, err := s.store.Invite(ctx, id) + if err != nil { + return v, err + } + v.Name = r.Label + v.BindingID = r.BindingID + case "registration": + r, err := s.store.Registration(ctx, id) + if err != nil { + return v, err + } + if db.IsRegistrationActive(r.Status) { + return v, db.ErrRegistrationTransition + } + v.Name = r.Username + v.UserID = r.ExternalUserID.String + v.BindingID = r.BindingID + v.Expiry = "No expiry" + if r.UserDisableAt.Valid { + v.Expiry = r.UserDisableAt.Time.UTC().Format("2006-01-02 15:04 UTC") + } + case "user": + r, err := s.store.ManagedUser(ctx, id) + if err != nil { + return v, err + } + v.Name = r.Username + v.UserID = r.ExternalUserID + v.BindingID = r.BindingID + default: + return v, db.ErrNotFound + } + return v, nil +} + +func (s *Server) serverReview(w http.ResponseWriter, r *http.Request, session db.Session) { + id, err := idFromPath(r, "id") + if err != nil { + s.registrationRecoveryError(w, db.ErrNotFound) + return + } + kind := r.PathValue("kind") + release, err := s.claimAccountOperations(id) + if err != nil { + s.registrationRecoveryError(w, err) + return + } + defer release() + v, err := s.reviewRecord(r.Context(), kind, id) + if err != nil { + s.registrationRecoveryError(w, err) + return + } + op, err := operationSnapshot(r.Context()) + if err != nil { + s.error(w, err) + return + } + if v.BindingID == op.Identity.Binding.ID { + s.message(w, "Already assigned", "This record already belongs to the current server.", http.StatusConflict) + return + } + bindings, err := s.store.MediaBindings(r.Context()) + if err != nil { + s.error(w, err) + return + } + v.Source = "Unverified (record saved before server tracking)" + for _, b := range bindings { + if b.ID == v.BindingID { + v.Source = b.Provider + " · " + b.BaseURL + " · " + b.ServerID + } + } + b := op.Identity.Binding + v.Destination = b.Provider + " · " + b.BaseURL + " · " + b.ServerID + if v.UserID != "" { + releaseUser, err := s.claimMediaUser(op.Identity.Binding.ID, v.UserID) + if err != nil { + s.registrationRecoveryError(w, err) + return + } + defer releaseUser() + if op.Settings.APIKey == "" { + s.message(w, "API key required", "Save an API key to verify the account before assigning it.", http.StatusConflict) + return + } + ctx, cancel := context.WithTimeout(r.Context(), 10*time.Second) + defer cancel() + user, found, err := op.Media.GetUser(ctx, op.Settings.ServerURL, op.Settings.APIKey, v.UserID) + if err != nil { + s.message(w, "Could not verify account", "Try again when the media server is available.", http.StatusBadGateway) + return + } + if !found || user.ID != v.UserID || userIsAdministrator(user) { + s.message(w, "Account cannot be assigned", "The account ID must exist on this server and have a verified non-administrator policy.", http.StatusConflict) + return + } + v.AccountName = user.Name + } + if r.Method == http.MethodPost { + if r.FormValue("confirmed") != "yes" || r.FormValue("binding_id") != strconv.FormatInt(b.ID, 10) || r.FormValue("source_id") != strconv.FormatInt(v.BindingID, 10) { + s.message(w, "Review needed", "Review the current destination and confirm this assignment.", http.StatusConflict) + return + } + switch kind { + case "invite": + err = s.store.AdoptInvite(r.Context(), id, b.ID) + case "registration": + err = s.store.AdoptRegistration(r.Context(), id, b.ID) + case "user": + err = s.store.AdoptManagedUser(r.Context(), id, b.ID) + } + if err != nil { + s.registrationRecoveryError(w, err) + return + } + s.audit(r, session, "server.assign", kind, strconv.FormatInt(id, 10), map[string]any{"previous_binding": v.BindingID, "binding": b.ID}) + destination := "/admin/registrations?review=1" + if kind == "invite" { + destination = "/admin/invites" + } + if kind == "user" { + destination = "/admin/users/review" + } + http.Redirect(w, r, destination, http.StatusSeeOther) + return + } + data := s.data(r, session) + data.Review = v + render(w, "server-review", data) +} + +func (s *Server) managedUserReview(w http.ResponseWriter, r *http.Request, session db.Session) { + before := historyCursor(r) + users, err := s.store.ManagedUserReviewPage(r.Context(), session.BindingID, before, 51) + if err != nil { + s.error(w, err) + return + } + data := s.data(r, session) + data.HistoryBefore = before + if len(users) > 50 { + users = users[:50] + data.NextBefore = users[49].ID + } + data.ManagedHistory = users + render(w, "users-review", data) +} +func historyCursor(r *http.Request) int64 { + v, _ := strconv.ParseInt(r.URL.Query().Get("before"), 10, 64) + if v < 0 { + return 0 + } + return v +} diff --git a/internal/httpserver/server_review_test.go b/internal/httpserver/server_review_test.go new file mode 100644 index 0000000..6f77ed1 --- /dev/null +++ b/internal/httpserver/server_review_test.go @@ -0,0 +1,155 @@ +package httpserver + +import ( + "context" + "database/sql" + "github.com/mayvqt/aperture/internal/db" + "github.com/mayvqt/aperture/internal/mediaserver" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" + "time" +) + +type observedMedia struct { + *fakeMediaServer + inspections, adminChecks int + serverID string +} + +func (m *observedMedia) Inspect(context.Context, string, string, string) (mediaserver.ServerInfo, error) { + m.inspections++ + return mediaserver.ServerInfo{ID: m.serverID}, nil +} +func (m *observedMedia) IsAdmin(context.Context, string, string, string, string) (bool, error) { + m.adminChecks++ + return true, nil +} + +func TestAdminSessionIsRejectedBeforeTokenReachesDifferentOrigin(t *testing.T) { + store := newFakeStore() + cfg := testConfig() + cfg.ServerURL = "http://different-media:8096" + media := &observedMedia{fakeMediaServer: &fakeMediaServer{}, serverID: "synthetic-server"} + s := NewServer(cfg, store, testMediaFactory(media)) + rr := httptest.NewRecorder() + s.ServeHTTP(rr, adminRequest(t, http.MethodGet, "/admin", nil)) + if rr.Code != http.StatusSeeOther || media.inspections != 0 || media.adminChecks != 0 { + t.Fatalf("old token forwarded: status=%d identity=%d admin=%d", rr.Code, media.inspections, media.adminChecks) + } +} +func TestReplacementAtSameURLDoesNotRunOldSessionHandler(t *testing.T) { + store := newFakeStore() + media := &observedMedia{fakeMediaServer: &fakeMediaServer{}, serverID: "replacement"} + s := NewServer(testConfig(), store, testMediaFactory(media)) + rr := httptest.NewRecorder() + s.ServeHTTP(rr, adminRequest(t, http.MethodGet, "/admin", nil)) + if rr.Code != http.StatusSeeOther || media.inspections != 1 || media.adminChecks != 0 { + t.Fatal("replacement server accepted prior session") + } +} +func TestUnverifiedRegistrationCannotDeleteCurrentAccount(t *testing.T) { + store := newFakeStore() + store.registrations[0].BindingID = 0 + media := &fakeMediaServer{users: []mediaserver.User{{ID: "media-alice", Policy: []byte(`{"IsAdministrator":false}`)}}} + s := NewServer(testConfig(), store, testMediaFactory(media)) + form := url.Values{"csrf": {"csrf-secret"}} + req := adminRequest(t, http.MethodPost, "/admin/registrations/1/delete", strings.NewReader(form.Encode())) + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + rr := httptest.NewRecorder() + s.ServeHTTP(rr, req) + if rr.Code != http.StatusConflict || media.deletedUserID != "" || len(store.registrations) != 1 { + t.Fatal("unverified registration deleted a current-server account") + } +} +func TestServerReviewRequiresCurrentAccountAndConfirmation(t *testing.T) { + store := newFakeStore() + store.registrations[0].BindingID = 0 + media := &fakeMediaServer{users: []mediaserver.User{{ID: "media-alice", Name: "Current Alice", Policy: []byte(`{"IsAdministrator":false}`)}}} + s := NewServer(testConfig(), store, testMediaFactory(media)) + rr := httptest.NewRecorder() + s.ServeHTTP(rr, adminRequest(t, http.MethodGet, "/admin/server-review/registration/1", nil)) + if rr.Code != 200 || !strings.Contains(rr.Body.String(), "Current Alice") || !strings.Contains(rr.Body.String(), "Unverified") { + t.Fatalf("review omitted account identity: %d", rr.Code) + } + req := adminRequest(t, http.MethodPost, "/admin/server-review/registration/1", strings.NewReader("csrf=csrf-secret&binding_id=1&source_id=0")) + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + rr = httptest.NewRecorder() + s.ServeHTTP(rr, req) + if rr.Code != http.StatusConflict { + t.Fatal("assignment accepted without confirmation") + } + release, err := s.claimAccountOperations(1) + if err != nil { + t.Fatal(err) + } + defer release() + rr = httptest.NewRecorder() + s.ServeHTTP(rr, adminRequest(t, http.MethodGet, "/admin/server-review/registration/1", nil)) + if rr.Code != http.StatusConflict { + t.Fatal("review bypassed ongoing account operation") + } +} +func TestUnsafeUserPoliciesCannotPermitDestructiveActions(t *testing.T) { + for _, raw := range []string{"", `null`, `{}`, `{"IsAdministrator":null}`, `{"IsAdministrator":"false"}`, `{"IsAdministrator":false,"isadministrator":true}`} { + if !userIsAdministrator(mediaserver.User{Policy: []byte(raw)}) { + t.Fatalf("unsafe policy allowed account management: %s", raw) + } + } + if userIsAdministrator(mediaserver.User{Policy: []byte(`{"IsAdministrator":false}`)}) { + t.Fatal("verified regular user rejected") + } +} +func TestInviteStatusIncludesExpiredAndUnverifiedLinks(t *testing.T) { + invite := db.Invite{BindingID: 1, Enabled: true, MaxUses: 1, ExpiresAt: sql.NullTime{Time: time.Now().Add(-time.Hour), Valid: true}} + if inviteState(invite, 1) != "Expired" { + t.Fatal("expired link shown active") + } + invite.BindingID = 0 + if inviteState(invite, 1) != "Unverified server" { + t.Fatal("unknown link shown active") + } +} + +func TestInvalidInviteDoesNotContactMediaServer(t *testing.T) { + store := newFakeStore() + store.inviteErr = db.ErrNotFound + media := &observedMedia{fakeMediaServer: &fakeMediaServer{}, serverID: "synthetic-server"} + s := NewServer(testConfig(), store, testMediaFactory(media)) + rr := httptest.NewRecorder() + s.ServeHTTP(rr, httptest.NewRequest(http.MethodGet, "/i/invalid", nil)) + if rr.Code != http.StatusNotFound || media.inspections != 0 { + t.Fatal("invalid public token triggered an upstream request") + } +} + +type replacingMaintenanceMedia struct { + *fakeMediaServer + inspections int + disabled []string +} + +func (m *replacingMaintenanceMedia) Inspect(context.Context, string, string, string) (mediaserver.ServerInfo, error) { + m.inspections++ + id := "synthetic-server" + if m.inspections >= 3 { + id = "replacement" + } + return mediaserver.ServerInfo{ID: id}, nil +} +func (m *replacingMaintenanceMedia) DisableUser(_ context.Context, _, _, id string) error { + m.disabled = append(m.disabled, id) + return nil +} +func TestMaintenanceVerifiesIdentityBeforeEachAccount(t *testing.T) { + store := newFakeStore() + store.dueDisables = []db.Registration{{ID: 41, BindingID: 1, Status: db.RegistrationNeedsAttention, CleanupPending: true, ExternalUserID: sql.NullString{String: "first", Valid: true}}, {ID: 42, BindingID: 1, Status: db.RegistrationNeedsAttention, CleanupPending: true, ExternalUserID: sql.NullString{String: "second", Valid: true}}} + media := &replacingMaintenanceMedia{fakeMediaServer: &fakeMediaServer{}} + s := NewServer(testConfig(), store, testMediaFactory(media)) + s.runMaintenance(t.Context()) + if len(media.disabled) != 1 || media.disabled[0] != "first" { + t.Fatalf("batch mutated replacement server: %v", media.disabled) + } +} diff --git a/internal/httpserver/session.go b/internal/httpserver/session.go index c64c37a..f8fa40d 100644 --- a/internal/httpserver/session.go +++ b/internal/httpserver/session.go @@ -7,6 +7,7 @@ import ( "strings" "time" + "github.com/mayvqt/aperture/internal/connection" "github.com/mayvqt/aperture/internal/db" "github.com/mayvqt/aperture/internal/mediaserver" "github.com/mayvqt/aperture/internal/security" @@ -41,13 +42,37 @@ func (s *Server) admin(next func(http.ResponseWriter, *http.Request, db.Session) http.Redirect(w, r, "/login", http.StatusSeeOther) return } - settings, err := s.settings(r.Context()) + snapshot, err := s.snapshot(r.Context()) if err != nil { s.error(w, err) return } + matches := func(candidate connection.Snapshot) bool { + return session.BindingID > 0 && session.BindingID == candidate.Identity.Binding.ID && session.Generation == candidate.Identity.Generation + } + if !matches(snapshot) { + s.clearSession(w, r, session.ID) + http.Redirect(w, r, "/login", http.StatusSeeOther) + return + } + verified, err := s.connections.Verify(r.Context(), snapshot, session.AccessToken, session.DeviceID) + if err != nil { + if invalidMediaSession(err) || errors.Is(err, db.ErrConnectionChanged) { + s.clearSession(w, r, session.ID) + http.Redirect(w, r, "/login", http.StatusSeeOther) + return + } + s.message(w, "Could not verify server", "Aperture could not confirm the media server's identity. Try again shortly.", http.StatusBadGateway) + return + } + if !matches(verified) { + s.clearSession(w, r, session.ID) + http.Redirect(w, r, "/login", http.StatusSeeOther) + return + } + r = r.WithContext(connection.WithSnapshot(r.Context(), verified)) checkCtx, cancel := context.WithTimeout(r.Context(), adminCheckTimeout) - isAdmin, err := s.media.IsAdmin(checkCtx, settings.ServerURL, session.AccessToken, session.DeviceID, session.UserID) + isAdmin, err := verified.Media.IsAdmin(checkCtx, verified.Settings.ServerURL, session.AccessToken, session.DeviceID, session.UserID) cancel() if err != nil { if invalidMediaSession(err) { @@ -113,7 +138,7 @@ func (s *Server) data(r *http.Request, session db.Session) viewData { case strings.HasPrefix(path, "/admin/settings"): page, title = "/admin/settings", "Settings" } - return viewData{Admin: true, Username: session.Username, CSRF: session.CSRFSecret, Title: title, CurrentPage: page} + return viewData{BindingID: session.BindingID, Admin: true, Username: session.Username, CSRF: session.CSRFSecret, Title: title, CurrentPage: page} } func (s *Server) validCSRF(r *http.Request, expected string) bool { actual, ok := formCSRF(r) diff --git a/internal/httpserver/session_test.go b/internal/httpserver/session_test.go index 8602104..7a4a273 100644 --- a/internal/httpserver/session_test.go +++ b/internal/httpserver/session_test.go @@ -21,7 +21,7 @@ func TestValidCSRFDeniesEmptyValues(t *testing.T) { func TestAdminRevocationInvalidatesApertureSession(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{adminRevoked: true}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{adminRevoked: true})) req := adminRequest(t, http.MethodGet, "/admin", nil) rr := httptest.NewRecorder() @@ -40,7 +40,7 @@ func TestAdminRevocationInvalidatesApertureSession(t *testing.T) { } func TestAdminVerificationFailsClosed(t *testing.T) { - handler := New(testConfig(), newFakeStore(), &fakeMediaServer{isAdminErr: errors.New("jellyfin unavailable")}) + handler := New(testConfig(), newFakeStore(), testMediaFactory(&fakeMediaServer{isAdminErr: errors.New("jellyfin unavailable")})) req := adminRequest(t, http.MethodGet, "/admin", nil) rr := httptest.NewRecorder() @@ -53,9 +53,9 @@ func TestAdminVerificationFailsClosed(t *testing.T) { func TestAdminInvalidMediaServerTokenClearsSession(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{ + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{ isAdminErr: &mediaserver.HTTPError{StatusCode: http.StatusUnauthorized, Status: "401 Unauthorized"}, - }) + })) req := adminRequest(t, http.MethodGet, "/admin", nil) rr := httptest.NewRecorder() @@ -75,7 +75,7 @@ func TestAdminInvalidMediaServerTokenClearsSession(t *testing.T) { func TestLogoutRequiresValidCSRFBeforeClearingCookie(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) req := adminRequest(t, http.MethodPost, "/logout", nil) rr := httptest.NewRecorder() @@ -94,7 +94,7 @@ func TestLogoutRequiresValidCSRFBeforeClearingCookie(t *testing.T) { func TestLogoutWithValidCSRFDeletesSessionAndClearsCookie(t *testing.T) { store := newFakeStore() - handler := New(testConfig(), store, &fakeMediaServer{}) + handler := New(testConfig(), store, testMediaFactory(&fakeMediaServer{})) form := url.Values{"csrf": {store.session.CSRFSecret}} req := adminRequest(t, http.MethodPost, "/logout", strings.NewReader(form.Encode())) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") diff --git a/internal/httpserver/template_funcs.go b/internal/httpserver/template_funcs.go index a59a9f6..be87c62 100644 --- a/internal/httpserver/template_funcs.go +++ b/internal/httpserver/template_funcs.go @@ -18,6 +18,25 @@ type templateSummaryItem struct { } var templateFuncs = template.FuncMap{ + "registrationActive": db.IsRegistrationActive, + "registrationNextRetry": func(reg db.Registration) string { + if reg.CleanupPending || reg.Status == db.RegistrationDisableFailed { + if reg.NextDisableAttemptAt.Valid { + return reg.NextDisableAttemptAt.Time.UTC().Format("2006-01-02 15:04 UTC") + } + return "Disable pending" + } + if db.CanRetryRegistrationTemplate(reg.Status) { + if reg.TemplateAttempts >= 6 { + return "Review needed" + } + if reg.NextTemplateAttemptAt.Valid { + return reg.NextTemplateAttemptAt.Time.UTC().Format("2006-01-02 15:04 UTC") + } + return "Access retry pending" + } + return "—" + }, "stylesheetHash": func() string { return stylesheetHash }, "scriptHash": func() string { return scriptHash }, "boolText": func(v bool) string { @@ -48,7 +67,7 @@ var templateFuncs = template.FuncMap{ case db.RegistrationNeedsAttention, db.RegistrationFailedApplyTemplate: return "Review needed" case db.RegistrationFailedCreateUser: - return "Account not created" + return "Setup incomplete" case db.RegistrationDisableFailed: return "Could not disable" case db.RegistrationDisabledExpired: @@ -59,23 +78,16 @@ var templateFuncs = template.FuncMap{ return "In progress" } }, - "inviteState": func(enabled bool, uses, maxUses int) string { - if !enabled { - return "Disabled" - } - if uses >= maxUses { - return "Used" - } - return "Active" - }, - "inviteStateClass": func(enabled bool, uses, maxUses int) string { - if !enabled { + "inviteState": inviteState, + "inviteStateClass": func(invite db.Invite, bindingID int64) string { + switch inviteState(invite, bindingID) { + case "Active": + return "good" + case "Disabled": return "bad" - } - if uses >= maxUses { + default: return "warn" } - return "good" }, "inviteURL": func(publicURL, token string) string { if token == "" { @@ -183,3 +195,22 @@ func presenceText(count int) string { } return fmt.Sprintf("%d settings", count) } + +func inviteState(invite db.Invite, bindingID int64) string { + if invite.BindingID != bindingID || bindingID <= 0 { + if invite.BindingID == 0 { + return "Unverified server" + } + return "Previous server" + } + if !invite.Enabled { + return "Disabled" + } + if invite.ExpiresAt.Valid && !invite.ExpiresAt.Time.After(time.Now()) { + return "Expired" + } + if invite.Uses >= invite.MaxUses { + return "Used" + } + return "Active" +} diff --git a/internal/httpserver/templates/dashboard.html b/internal/httpserver/templates/dashboard.html index fe4b345..f36af7f 100644 --- a/internal/httpserver/templates/dashboard.html +++ b/internal/httpserver/templates/dashboard.html @@ -32,7 +32,7 @@

Dashboard

{{.Stats.NeedsAttention}} - review needed + review needed
{{.Stats.ScheduledUserDisables}} @@ -60,7 +60,7 @@

Invites

{{.Label}} {{.Uses}} / {{.MaxUses}} - {{inviteState .Enabled .Uses .MaxUses}} + {{inviteState . $.BindingID}} {{userExpiryText .UserExpiryDays}} @@ -91,16 +91,16 @@

Recent signups

{{range .Registrations}} - {{.Username}} + {{.Username}}{{if ne .BindingID $.BindingID}} · {{if .BindingID}}Previous server{{else}}Unverified server{{end}}{{end}} {{registrationStatus .Status}} {{nullDate .UserDisableAt}} - {{nullDateTime .NextDisableAttemptAt}} -
+ {{registrationNextRetry .}} + {{if and (eq .BindingID $.BindingID) (not (registrationActive .Status))}} -
+ {{end}} {{else}} @@ -132,7 +132,7 @@

Application setup

- {{if not .ProviderManaged}}Changing the provider or server signs you out. Existing account history still refers to the previous server.{{end}} + {{if not .ProviderManaged}}Changing the provider or server signs you out. Saved invites and accounts stay with their original server. Review records individually before assigning them to another server.{{end}}
@@ -185,25 +186,25 @@

Registrations

{{range .Registrations}} - {{.Username}} + {{.Username}}{{if ne .BindingID $.BindingID}} · {{if .BindingID}}Previous server{{else}}Unverified server{{end}}{{end}} {{registrationStatus .Status}} {{.ExternalUserID.String}} - Ends {{nullDate .UserDisableAt}}Disabled {{nullDate .UserDisabledAt}} - {{.DisableAttempts}} attemptsNext {{nullDateTime .NextDisableAttemptAt}} - {{.TemplateAttempts}} attemptsNext {{nullDateTime .NextTemplateAttemptAt}} -
{{if canRetryTemplate .Status .ExternalUserID}} + {{if .UserDisableAt.Valid}}Ends {{nullDate .UserDisableAt}}{{else}}No expiry{{end}}{{if .UserDisabledAt.Valid}}Disabled {{nullDate .UserDisabledAt}}{{end}} + {{if .CleanupPending}}Disable pending{{else}}{{.DisableAttempts}} attempts{{end}}{{if .CleanupError.Valid}}Last attempt failed{{end}}{{if and .NextDisableAttemptAt.Valid (not .UserDisabledAt.Valid)}}Next {{nullDateTime .NextDisableAttemptAt}}{{end}} + {{.TemplateAttempts}} attempts{{if canRetryTemplate .Status .ExternalUserID}}{{if ge .TemplateAttempts 6}}Automatic retries finished; review access{{else if .NextTemplateAttemptAt.Valid}}Next {{nullDateTime .NextTemplateAttemptAt}}{{else}}Retry pending{{end}}{{end}} +
{{if ne .BindingID $.BindingID}}Review server{{end}}{{if and (eq .BindingID $.BindingID) (canRetryTemplate .Status .ExternalUserID)}}
{{end}} -
+ {{if and (eq .BindingID $.BindingID) (not (registrationActive .Status))}} -
- {{.CreatedAt}} + {{end}}
+ {{dateTime .CreatedAt}} {{else}} @@ -212,4 +213,4 @@

Registrations

{{end}} -{{template "base-end" .}}{{end}} +{{template "base-end" .}}{{end}} diff --git a/internal/httpserver/templates/invites.html b/internal/httpserver/templates/invites.html index dbdf7c6..e27e304 100644 --- a/internal/httpserver/templates/invites.html +++ b/internal/httpserver/templates/invites.html @@ -15,7 +15,7 @@

Invites

{{.Invite.Label}}

{{.Invite.Template}} template

- {{inviteState .Invite.Enabled .Invite.Uses .Invite.MaxUses}} + {{inviteState .Invite $.BindingID}}
{{.Invite.Uses}} of {{.Invite.MaxUses}} signups used @@ -23,20 +23,20 @@

{{.Invite.Label}}

{{userExpiryText .Invite.UserExpiryDays}} account access {{if .Activity.Username}}{{.Activity.Username}} last signed up {{dateOnly .Activity.CreatedAt}}{{else}}No signups yet{{end}}
- {{if .Invite.Token}} + {{if and .Invite.Token (eq .Invite.BindingID $.BindingID)}} - {{else}}

This link is no longer available to copy.

{{end}} -
{{if .Invite.Enabled}} + {{else}}

{{if ne .Invite.BindingID $.BindingID}}Review the server before sharing this invite.{{else}}This link is no longer available to copy.{{end}}

{{end}} +
{{if ne .Invite.BindingID $.BindingID}}Review server{{end}}{{if .Invite.Enabled}}
- {{else}} + {{else if eq .Invite.BindingID $.BindingID}}
@@ -51,7 +51,7 @@

{{.Invite.Label}}

{{else}}

No invites yet.

{{end}} -
{{template "base-end" .}}{{end}} +{{template "base-end" .}}{{end}} {{define "invite-new"}}{{template "base-start" .}}
@@ -65,7 +65,7 @@

Create invite

-