diff --git a/go.mod b/go.mod --- a/go.mod +++ b/go.mod @@ -34,6 +34,7 @@ github.com/gorilla/feeds v1.2.0 github.com/gorilla/sessions v1.4.0 github.com/gorilla/websocket v1.5.4-0.20250319132907-e064f32e3674 + github.com/hashicorp/golang-lru/v2 v2.0.7 github.com/hiddeco/sshsig v0.2.0 github.com/hpcloud/tail v1.0.0 github.com/ipfs/go-cid v0.6.0 @@ -159,7 +160,6 @@ github.com/hashicorp/go-secure-stdlib/strutil v0.1.2 // indirect github.com/hashicorp/go-sockaddr v1.0.7 // indirect github.com/hashicorp/golang-lru v1.0.2 // indirect - github.com/hashicorp/golang-lru/v2 v2.0.7 // indirect github.com/hashicorp/hcl v1.0.1-vault-7 // indirect github.com/hexops/gotextdiff v1.0.3 // indirect github.com/ipfs/bbloom v0.0.4 // indirect diff --git a/appview/db/migration.go b/appview/db/migration.go --- a/appview/db/migration.go +++ b/appview/db/migration.go @@ -79,8 +79,20 @@ func EnqueuePdsRecordMigration(ctx context.Context, e Execer, name string, did syntax.DID, collection syntax.NSID, rkey syntax.RecordKey) error { _, err := e.ExecContext(ctx, `insert into pds_migration (name, did, collection, rkey) - values (?, ?, ?, ?)`, + values (?, ?, ?, ?) + on conflict(name, did, collection, rkey) do update set + status = case when pds_migration.status = 'failed' then 'pending' else pds_migration.status end, + retry_count = case when pds_migration.status = 'failed' then 0 else pds_migration.retry_count end, + retry_after = case when pds_migration.status = 'failed' then 0 else pds_migration.retry_after end, + error_msg = case when pds_migration.status = 'failed' then null else pds_migration.error_msg end`, name, did, collection, rkey, + ) + return err +} + +func ReapStaleRunningMigrations(ctx context.Context, e Execer) error { + _, err := e.ExecContext(ctx, + `update pds_migration set status = 'pending' where status = 'running'`, ) return err } diff --git a/appview/migration/migration.go b/appview/migration/migration.go --- a/appview/migration/migration.go +++ b/appview/migration/migration.go @@ -20,23 +20,34 @@ const maxConcurrentMigrations = 8 +type migrator func(ctx context.Context, client *atclient.APIClient, did syntax.DID, aturi syntax.ATURI) error + +type permAuthErrHandler func(ctx context.Context, did syntax.DID, sessId string, err error) bool + type Migration struct { - db *db.DB - oauth *oauth.OAuth - dir identity.Directory - logger *slog.Logger - inflight sync.Map - sem chan struct{} + db *db.DB + oauth *oauth.OAuth + dir identity.Directory + logger *slog.Logger + inflight sync.Map + sem chan struct{} + migrators map[string]migrator + onPermAuthErr permAuthErrHandler } func NewMigration(db *db.DB, oauth *oauth.OAuth, dir identity.Directory, logger *slog.Logger) *Migration { - return &Migration{ - db: db, - oauth: oauth, - dir: dir, - logger: logger, - sem: make(chan struct{}, maxConcurrentMigrations), + m := &Migration{ + db: db, + oauth: oauth, + dir: dir, + logger: logger, + sem: make(chan struct{}, maxConcurrentMigrations), + onPermAuthErr: oauth.HandlePermanentAuthErr, } + m.migrators = map[string]migrator{ + "add-repo-did": m.migrateAddRepoDid, + } + return m } func (s *Migration) BackgroundMigrationMiddleware(next http.Handler) http.Handler { @@ -64,6 +75,7 @@ return } + sessId := s.oauth.GetSessIdFromCookie(r) client, err := s.oauth.AuthorizedClient(r) if err != nil || client.AccountDID == nil { <-s.sem @@ -74,12 +86,12 @@ go func() { defer s.inflight.Delete(did) defer func() { <-s.sem }() - s.runPendingMigrations(context.Background(), *client.AccountDID, client) + s.runPendingMigrations(context.Background(), *client.AccountDID, sessId, client) }() }) } -func (s *Migration) runPendingMigrations(ctx context.Context, did syntax.DID, client *atclient.APIClient) { +func (s *Migration) runPendingMigrations(ctx context.Context, did syntax.DID, sessId string, client *atclient.APIClient) { l := s.logger.With("did", did) migrations, err := db.ListPendingPdsRecordMigrations(ctx, s.db, did) if err != nil { @@ -88,25 +100,23 @@ } for _, migration := range migrations { - if err := s.migrate(ctx, client, migration); err != nil { + if err := s.migrate(ctx, client, sessId, migration); err != nil { l.Error("migration failed", "err", err) } } } -func (s *Migration) migrate(ctx context.Context, client *atclient.APIClient, migration *models.PDSMigration) error { +func (s *Migration) migrate(ctx context.Context, client *atclient.APIClient, sessId string, migration *models.PDSMigration) error { l := s.logger.With( "name", migration.Name, "aturi", migration.RecordAtUri(), ) - var err error - switch migration.Name { - case "add-repo-did": - err = s.migrateAddRepoDid(ctx, client, migration.Did, migration.RecordAtUri()) - default: + mig, ok := s.migrators[migration.Name] + if !ok { return fmt.Errorf("unexpected migration name %s", migration.Name) } + err := mig(ctx, client, migration.Did, migration.RecordAtUri()) if err == nil { l.Info("migrated") @@ -114,20 +124,24 @@ } else { l.Warn("failed to migrate", "err", err) - errMsg := err.Error() - var retryCount = migration.RetryCount + 1 - var retryAfter = time.Now().Add(3 * time.Second).Unix() - - // remove null bytes - errMsg = strings.ReplaceAll(errMsg, "\x00", "") - - migration.Status = models.PDSMigrationStatusPending + errMsg := strings.ReplaceAll(err.Error(), "\x00", "") migration.ErrorMsg = &errMsg - migration.RetryCount = retryCount - migration.RetryAfter = retryAfter + migration.RetryCount++ + + if s.onPermAuthErr(ctx, migration.Did, sessId, err) { + migration.Status = models.PDSMigrationStatusFailed + migration.RetryAfter = 0 + } else { + migration.Status = models.PDSMigrationStatusPending + migration.RetryAfter = time.Now().Add(retryBackoff(migration.RetryCount)).Unix() + } } if err := db.UpdatePdsRecordMigration(ctx, s.db, migration); err != nil { return fmt.Errorf("failed to update migration status: %w", err) } return nil +} + +func retryBackoff(retries int) time.Duration { + return min(time.Duration(retries)*5*time.Second, time.Hour) } diff --git a/appview/migration/migration_test.go b/appview/migration/migration_test.go new file mode 100644 --- /dev/null +++ b/appview/migration/migration_test.go @@ -0,0 +1,205 @@ +package migration + +import ( + "context" + "errors" + "io" + "log/slog" + "path/filepath" + "sync/atomic" + "testing" + + "github.com/bluesky-social/indigo/atproto/atclient" + "github.com/bluesky-social/indigo/atproto/syntax" + + "tangled.org/core/appview/db" + "tangled.org/core/appview/models" + "tangled.org/core/appview/oauth" +) + +func newTestDB(t *testing.T) *db.DB { + t.Helper() + d, err := db.Make(context.Background(), filepath.Join(t.TempDir(), "test.db")) + if err != nil { + t.Fatalf("db.Make: %v", err) + } + t.Cleanup(func() { d.Close() }) + return d +} + +func seedMigration(t *testing.T, d *db.DB, mig *models.PDSMigration) { + t.Helper() + if err := db.EnqueuePdsRecordMigration(context.Background(), d, mig.Name, mig.Did, mig.Collection, mig.Rkey); err != nil { + t.Fatalf("EnqueuePdsRecordMigration: %v", err) + } +} + +func fetch(t *testing.T, d *db.DB, did syntax.DID) *models.PDSMigration { + t.Helper() + rows, err := d.QueryContext(context.Background(), + `select name, did, collection, rkey, status, error_msg, retry_count, retry_after from pds_migration where did = ?`, did) + if err != nil { + t.Fatalf("query: %v", err) + } + defer rows.Close() + if !rows.Next() { + t.Fatalf("no row for did %s", did) + } + var m models.PDSMigration + if err := rows.Scan(&m.Name, &m.Did, &m.Collection, &m.Rkey, &m.Status, &m.ErrorMsg, &m.RetryCount, &m.RetryAfter); err != nil { + t.Fatalf("scan: %v", err) + } + return &m +} + +func newTestMigration(t *testing.T, mig migrator, onPerm permAuthErrHandler) *Migration { + t.Helper() + return &Migration{ + db: newTestDB(t), + logger: slog.New(slog.NewTextHandler(io.Discard, nil)), + sem: make(chan struct{}, maxConcurrentMigrations), + migrators: map[string]migrator{"add-repo-did": mig}, + onPermAuthErr: onPerm, + } +} + +func newPDSMigration(did syntax.DID) *models.PDSMigration { + return &models.PDSMigration{ + Name: "add-repo-did", + Did: did, + Collection: "sh.tangled.repo", + Rkey: "abc", + Status: models.PDSMigrationStatusRunning, + } +} + +func TestMigrateInvalidGrantMarksFailed(t *testing.T) { + did := syntax.DID("did:plc:boltless") + var permCalled atomic.Int32 + m := newTestMigration(t, + func(context.Context, *atclient.APIClient, syntax.DID, syntax.ATURI) error { + return errors.New("put record: failed to refresh OAuth tokens: token refresh failed: auth server request failed (HTTP 400): invalid_grant") + }, + func(_ context.Context, _ syntax.DID, _ string, err error) bool { + permCalled.Add(1) + return oauth.IsPermanentAuthErr(err) + }, + ) + seedMigration(t, m.db, newPDSMigration(did)) + pm := newPDSMigration(did) + + if err := m.migrate(context.Background(), &atclient.APIClient{}, "sess1", pm); err != nil { + t.Fatalf("migrate: %v", err) + } + + got := fetch(t, m.db, did) + if got.Status != models.PDSMigrationStatusFailed { + t.Fatalf("status = %s, want failed", got.Status) + } + if got.RetryAfter != 0 { + t.Fatalf("RetryAfter = %d, want 0", got.RetryAfter) + } + if permCalled.Load() != 1 { + t.Fatalf("onPermAuthErr called %d times, want 1", permCalled.Load()) + } +} + +func TestMigrateTransientErrorStaysPending(t *testing.T) { + did := syntax.DID("did:plc:akshay") + m := newTestMigration(t, + func(context.Context, *atclient.APIClient, syntax.DID, syntax.ATURI) error { + return errors.New("put record: failed to refresh OAuth tokens: token refresh failed (HTTP 429): rate_limited") + }, + func(_ context.Context, _ syntax.DID, _ string, err error) bool { + return oauth.IsPermanentAuthErr(err) + }, + ) + seedMigration(t, m.db, newPDSMigration(did)) + pm := newPDSMigration(did) + + if err := m.migrate(context.Background(), &atclient.APIClient{}, "sess1", pm); err != nil { + t.Fatalf("migrate: %v", err) + } + + got := fetch(t, m.db, did) + if got.Status != models.PDSMigrationStatusPending { + t.Fatalf("status = %s, want pending", got.Status) + } + if got.RetryCount != 1 { + t.Fatalf("RetryCount = %d, want 1", got.RetryCount) + } + if got.RetryAfter == 0 { + t.Fatalf("RetryAfter not scheduled") + } +} + +func TestMigrateSuccess(t *testing.T) { + did := syntax.DID("did:plc:boltless") + m := newTestMigration(t, + func(context.Context, *atclient.APIClient, syntax.DID, syntax.ATURI) error { return nil }, + func(context.Context, syntax.DID, string, error) bool { return false }, + ) + seedMigration(t, m.db, newPDSMigration(did)) + pm := newPDSMigration(did) + + if err := m.migrate(context.Background(), &atclient.APIClient{}, "sess1", pm); err != nil { + t.Fatalf("migrate: %v", err) + } + + got := fetch(t, m.db, did) + if got.Status != models.PDSMigrationStatusDone { + t.Fatalf("status = %s, want done", got.Status) + } +} + +func TestEnqueueResetsFailedToPending(t *testing.T) { + d := newTestDB(t) + did := syntax.DID("did:plc:boltless") + seed := newPDSMigration(did) + seedMigration(t, d, seed) + + errMsg := "some prior failure" + failed := *seed + failed.Status = models.PDSMigrationStatusFailed + failed.RetryCount = 7 + failed.ErrorMsg = &errMsg + if err := db.UpdatePdsRecordMigration(context.Background(), d, &failed); err != nil { + t.Fatalf("UpdatePdsRecordMigration: %v", err) + } + + if err := db.EnqueuePdsRecordMigration(context.Background(), d, seed.Name, seed.Did, seed.Collection, seed.Rkey); err != nil { + t.Fatalf("re-enqueue: %v", err) + } + + got := fetch(t, d, did) + if got.Status != models.PDSMigrationStatusPending { + t.Fatalf("status = %s, want pending", got.Status) + } + if got.RetryCount != 0 { + t.Fatalf("RetryCount = %d, want 0", got.RetryCount) + } + if got.ErrorMsg != nil { + t.Fatalf("ErrorMsg = %v, want nil", got.ErrorMsg) + } +} + +func TestReapStaleRunning(t *testing.T) { + d := newTestDB(t) + did := syntax.DID("did:plc:akshay") + seed := newPDSMigration(did) + seedMigration(t, d, seed) + running := *seed + running.Status = models.PDSMigrationStatusRunning + if err := db.UpdatePdsRecordMigration(context.Background(), d, &running); err != nil { + t.Fatalf("update: %v", err) + } + + if err := db.ReapStaleRunningMigrations(context.Background(), d); err != nil { + t.Fatalf("reap: %v", err) + } + + got := fetch(t, d, did) + if got.Status != models.PDSMigrationStatusPending { + t.Fatalf("status = %s, want pending", got.Status) + } +} diff --git a/appview/models/migration.go b/appview/models/migration.go --- a/appview/models/migration.go +++ b/appview/models/migration.go @@ -30,6 +30,7 @@ PDSMigrationStatusPending PDSMigrationStatus = "pending" PDSMigrationStatusRunning PDSMigrationStatus = "running" PDSMigrationStatusDone PDSMigrationStatus = "done" + PDSMigrationStatusFailed PDSMigrationStatus = "failed" ) func (m *PDSMigration) RecordAtUri() syntax.ATURI { diff --git a/appview/oauth/cache_test.go b/appview/oauth/cache_test.go new file mode 100644 --- /dev/null +++ b/appview/oauth/cache_test.go @@ -0,0 +1,198 @@ +package oauth + +import ( + "context" + "errors" + "io" + "log/slog" + "sync" + "sync/atomic" + "testing" + + "github.com/bluesky-social/indigo/atproto/atcrypto" + "github.com/bluesky-social/indigo/atproto/auth/oauth" + "github.com/bluesky-social/indigo/atproto/syntax" + "github.com/hashicorp/golang-lru/v2/expirable" +) + +func discardLogger(t *testing.T) *slog.Logger { + t.Helper() + return slog.New(slog.NewTextHandler(io.Discard, nil)) +} + +type stubStore struct { + mu sync.Mutex + data map[string]oauth.ClientSessionData + getSessionCalls atomic.Int32 + deleteCalls atomic.Int32 +} + +func (s *stubStore) key(did syntax.DID, sessId string) string { + return string(did) + ":" + sessId +} + +func (s *stubStore) GetSession(_ context.Context, did syntax.DID, sessId string) (*oauth.ClientSessionData, error) { + s.getSessionCalls.Add(1) + s.mu.Lock() + defer s.mu.Unlock() + v, ok := s.data[s.key(did, sessId)] + if !ok { + return nil, errors.New("no such session") + } + clone := v + return &clone, nil +} + +func (s *stubStore) SaveSession(_ context.Context, sess oauth.ClientSessionData) error { + s.mu.Lock() + defer s.mu.Unlock() + s.data[s.key(sess.AccountDID, sess.SessionID)] = sess + return nil +} + +func (s *stubStore) DeleteSession(_ context.Context, did syntax.DID, sessId string) error { + s.deleteCalls.Add(1) + s.mu.Lock() + defer s.mu.Unlock() + delete(s.data, s.key(did, sessId)) + return nil +} + +func (s *stubStore) GetAuthRequestInfo(context.Context, string) (*oauth.AuthRequestData, error) { + return nil, errors.New("not used") +} +func (s *stubStore) SaveAuthRequestInfo(context.Context, oauth.AuthRequestData) error { + return nil +} +func (s *stubStore) DeleteAuthRequestInfo(context.Context, string) error { return nil } + +func newTestOAuth(t *testing.T) (*OAuth, *stubStore) { + t.Helper() + priv, err := atcrypto.GeneratePrivateKeyP256() + if err != nil { + t.Fatalf("generate key: %v", err) + } + store := &stubStore{data: map[string]oauth.ClientSessionData{}} + store.data[store.key("did:plc:boltless", "sess1")] = oauth.ClientSessionData{ + AccountDID: "did:plc:boltless", + SessionID: "sess1", + HostURL: "https://pds.example", + AuthServerURL: "https://pds.example", + AuthServerTokenEndpoint: "https://pds.example/oauth/token", + DPoPPrivateKeyMultibase: priv.Multibase(), + } + + cfg := oauth.NewLocalhostConfig("http://127.0.0.1/cb", []string{"atproto"}) + app := oauth.NewClientApp(&cfg, store) + o := &OAuth{ + ClientApp: app, + Logger: discardLogger(t), + sessionCache: expirable.NewLRU[string, *oauth.ClientSession](sessionCacheSize, nil, sessionCacheTTL), + } + return o, store +} + +func TestResumeSessionSingleflightDedupes(t *testing.T) { + o, store := newTestOAuth(t) + + const n = 32 + var wg sync.WaitGroup + results := make([]*oauth.ClientSession, n) + errs := make([]error, n) + wg.Add(n) + for i := range n { + go func() { + defer wg.Done() + sess, err := o.resumeSession(context.Background(), "did:plc:boltless", "sess1") + results[i] = sess + errs[i] = err + }() + } + wg.Wait() + + for i, err := range errs { + if err != nil { + t.Fatalf("goroutine %d: %v", i, err) + } + } + first := results[0] + if first == nil { + t.Fatal("first session is nil") + } + for i, s := range results { + if s != first { + t.Fatalf("goroutine %d got different *ClientSession (%p vs %p)", i, s, first) + } + } + calls := store.getSessionCalls.Load() + if calls > 1 { + t.Fatalf("GetSession called %d times, want 1", calls) + } +} + +func TestResumeSessionReuseAfterCache(t *testing.T) { + o, store := newTestOAuth(t) + + a, err := o.resumeSession(context.Background(), "did:plc:boltless", "sess1") + if err != nil { + t.Fatalf("first: %v", err) + } + b, err := o.resumeSession(context.Background(), "did:plc:boltless", "sess1") + if err != nil { + t.Fatalf("second: %v", err) + } + if a != b { + t.Fatalf("expected same pointer across cache hit") + } + if got := store.getSessionCalls.Load(); got != 1 { + t.Fatalf("GetSession called %d times, want 1", got) + } +} + +func TestHandlePermanentAuthErrEvictsAndLogsOut(t *testing.T) { + o, store := newTestOAuth(t) + + if _, err := o.resumeSession(context.Background(), "did:plc:boltless", "sess1"); err != nil { + t.Fatalf("seed: %v", err) + } + if _, ok := o.sessionCache.Get(sessionCacheKey("did:plc:boltless", "sess1")); !ok { + t.Fatal("cache missing after resume") + } + + handled := o.HandlePermanentAuthErr( + context.Background(), "did:plc:boltless", "sess1", + errors.New("auth server request failed (HTTP 400): invalid_grant"), + ) + if !handled { + t.Fatal("HandlePermanentAuthErr returned false") + } + if _, ok := o.sessionCache.Get(sessionCacheKey("did:plc:boltless", "sess1")); ok { + t.Fatal("cache still holds entry after HandlePermanentAuthErr") + } + if got := store.deleteCalls.Load(); got != 1 { + t.Fatalf("store.DeleteSession called %d times, want 1", got) + } + if _, ok := store.data[store.key("did:plc:boltless", "sess1")]; ok { + t.Fatal("store still holds session after Logout") + } +} + +func TestHandlePermanentAuthErrIgnoresTransient(t *testing.T) { + o, store := newTestOAuth(t) + if _, err := o.resumeSession(context.Background(), "did:plc:boltless", "sess1"); err != nil { + t.Fatalf("seed: %v", err) + } + handled := o.HandlePermanentAuthErr( + context.Background(), "did:plc:boltless", "sess1", + errors.New("token refresh failed (HTTP 429): rate_limited"), + ) + if handled { + t.Fatal("HandlePermanentAuthErr matched a transient error") + } + if _, ok := o.sessionCache.Get(sessionCacheKey("did:plc:boltless", "sess1")); !ok { + t.Fatal("transient error evicted cache") + } + if got := store.deleteCalls.Load(); got != 0 { + t.Fatalf("store.DeleteSession called %d times, want 0", got) + } +} diff --git a/appview/oauth/errors.go b/appview/oauth/errors.go new file mode 100644 --- /dev/null +++ b/appview/oauth/errors.go @@ -0,0 +1,22 @@ +package oauth + +import "regexp" + +var ( + permanentAuthErrorRe = regexp.MustCompile(`\b(invalid_grant|invalid_client|unauthorized_client)\b`) + staleAccessTokenErrRe = regexp.MustCompile(`HTTP 401\b.*(AuthenticationRequired|Invalid OAuth access token|invalid_token)`) +) + +func IsPermanentAuthErr(err error) bool { + if err == nil { + return false + } + return permanentAuthErrorRe.MatchString(err.Error()) +} + +func IsStaleAccessTokenErr(err error) bool { + if err == nil { + return false + } + return staleAccessTokenErrRe.MatchString(err.Error()) +} diff --git a/appview/oauth/errors_test.go b/appview/oauth/errors_test.go new file mode 100644 --- /dev/null +++ b/appview/oauth/errors_test.go @@ -0,0 +1,58 @@ +package oauth + +import ( + "errors" + "fmt" + "testing" +) + +func TestIsPermanentAuthErr(t *testing.T) { + cases := []struct { + name string + err error + want bool + }{ + {"nil", nil, false}, + {"empty", errors.New(""), false}, + {"random", errors.New("network unreachable"), false}, + {"rate limited", errors.New("token refresh failed (HTTP 429): rate_limited"), false}, + {"invalid grant direct", errors.New("token refresh failed (HTTP 400): invalid_grant"), true}, + {"invalid grant wrapped", fmt.Errorf("put record: %w", errors.New("failed to refresh OAuth tokens: token refresh failed: auth server request failed (HTTP 400): invalid_grant")), true}, + {"invalid client", errors.New("auth server request failed (HTTP 401): invalid_client"), true}, + {"unauthorized client", errors.New("token refresh failed (HTTP 400): unauthorized_client"), true}, + {"substring trap", errors.New("our invalid_grant_alternative ran out"), false}, + {"case-sensitive", errors.New("INVALID_GRANT"), false}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + got := IsPermanentAuthErr(c.err) + if got != c.want { + t.Fatalf("got %v want %v", got, c.want) + } + }) + } +} + +func TestIsStaleAccessTokenErr(t *testing.T) { + cases := []struct { + name string + err error + want bool + }{ + {"nil", nil, false}, + {"random", errors.New("hello"), false}, + {"500", errors.New("API request failed (HTTP 500): InternalError"), false}, + {"401 auth required", errors.New("API request failed (HTTP 401): AuthenticationRequired: Invalid OAuth access token"), true}, + {"401 invalid token", errors.New("API request failed (HTTP 401): invalid_token"), true}, + {"401 wrapped", fmt.Errorf("put record: %w", errors.New("API request failed (HTTP 401): AuthenticationRequired")), true}, + {"403 forbidden", errors.New("API request failed (HTTP 403): Forbidden"), false}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + got := IsStaleAccessTokenErr(c.err) + if got != c.want { + t.Fatalf("got %v want %v", got, c.want) + } + }) + } +} diff --git a/appview/oauth/handler.go b/appview/oauth/handler.go --- a/appview/oauth/handler.go +++ b/appview/oauth/handler.go @@ -238,7 +238,7 @@ l.Debug("creating empty Tangled profile") - sess, err := o.ClientApp.ResumeSession(ctx, sessData.AccountDID, sessData.SessionID) + sess, err := o.resumeSession(ctx, sessData.AccountDID, sessData.SessionID) if err != nil { l.Error("failed to resume session for profile creation", "err", err) return diff --git a/appview/oauth/oauth.go b/appview/oauth/oauth.go --- a/appview/oauth/oauth.go +++ b/appview/oauth/oauth.go @@ -17,11 +17,18 @@ "github.com/bluesky-social/indigo/atproto/syntax" xrpc "github.com/bluesky-social/indigo/xrpc" "github.com/gorilla/sessions" + "github.com/hashicorp/golang-lru/v2/expirable" "github.com/posthog/posthog-go" + "golang.org/x/sync/singleflight" "tangled.org/core/appview/config" "tangled.org/core/appview/db" "tangled.org/core/idresolver" "tangled.org/core/rbac" +) + +const ( + sessionCacheSize = 10000 + sessionCacheTTL = time.Hour ) type OAuth struct { @@ -39,6 +46,50 @@ appPasswordSession *AppPasswordSession appPasswordSessionMu sync.Mutex + + sessionCache *expirable.LRU[string, *oauth.ClientSession] + sessionSF singleflight.Group +} + +func sessionCacheKey(did syntax.DID, sessionId string) string { + return string(did) + ":" + sessionId +} + +func (o *OAuth) resumeSession(ctx context.Context, did syntax.DID, sessionId string) (*oauth.ClientSession, error) { + key := sessionCacheKey(did, sessionId) + if v, ok := o.sessionCache.Get(key); ok { + return v, nil + } + v, err, _ := o.sessionSF.Do(key, func() (any, error) { + if v, ok := o.sessionCache.Get(key); ok { + return v, nil + } + sess, err := o.ClientApp.ResumeSession(ctx, did, sessionId) + if err != nil { + return nil, err + } + o.sessionCache.Add(key, sess) + return sess, nil + }) + if err != nil { + return nil, err + } + return v.(*oauth.ClientSession), nil +} + +func (o *OAuth) EvictSession(did syntax.DID, sessionId string) { + o.sessionCache.Remove(sessionCacheKey(did, sessionId)) +} + +func (o *OAuth) HandlePermanentAuthErr(ctx context.Context, did syntax.DID, sessionId string, err error) bool { + if !IsPermanentAuthErr(err) { + return false + } + o.EvictSession(did, sessionId) + if logoutErr := o.ClientApp.Logout(ctx, did, sessionId); logoutErr != nil { + o.Logger.Warn("store logout after permanent auth error failed", "did", did, "err", logoutErr) + } + return true } func New(config *config.Config, ph posthog.Client, db *db.DB, enforcer *rbac.Enforcer, res *idresolver.Resolver, logger *slog.Logger) (*OAuth, error) { @@ -89,17 +140,18 @@ logger.Info("oauth setup successfully", "IsConfidential", clientApp.Config.IsConfidential()) return &OAuth{ - ClientApp: clientApp, - Config: config, - SessStore: sessStore, - JwksUri: jwksUri, - ClientName: clientName, - ClientUri: clientUri, - Posthog: ph, - Db: db, - Enforcer: enforcer, - IdResolver: res, - Logger: logger, + ClientApp: clientApp, + Config: config, + SessStore: sessStore, + JwksUri: jwksUri, + ClientName: clientName, + ClientUri: clientUri, + Posthog: ph, + Db: db, + Enforcer: enforcer, + IdResolver: res, + Logger: logger, + sessionCache: expirable.NewLRU[string, *oauth.ClientSession](sessionCacheSize, nil, sessionCacheTTL), }, nil } @@ -148,7 +200,7 @@ sessId := userSession.Values[SessionId].(string) - clientSess, err := o.ClientApp.ResumeSession(r.Context(), sessDid, sessId) + clientSess, err := o.resumeSession(r.Context(), sessDid, sessId) if err != nil { return nil, fmt.Errorf("failed to resume session: %w", err) } @@ -173,11 +225,14 @@ sessId := userSession.Values[SessionId].(string) + o.EvictSession(sessDid, sessId) + // delete the session err1 := o.ClientApp.Logout(r.Context(), sessDid, sessId) if err1 != nil { err1 = fmt.Errorf("failed to logout: %w", err1) } + o.EvictSession(sessDid, sessId) // remove the cookie userSession.Options.MaxAge = -1 @@ -201,7 +256,7 @@ return fmt.Errorf("invalid DID: %w", err) } - sess, err := o.ClientApp.ResumeSession(r.Context(), did, account.SessionId) + sess, err := o.resumeSession(r.Context(), did, account.SessionId) if err != nil { registry.RemoveAccount(targetDid) _ = o.saveAccounts(w, r, registry) @@ -230,7 +285,9 @@ did, err := syntax.ParseDID(targetDid) if err == nil { + o.EvictSession(did, account.SessionId) _ = o.ClientApp.Logout(r.Context(), did, account.SessionId) + o.EvictSession(did, account.SessionId) } registry.RemoveAccount(targetDid) @@ -259,6 +316,18 @@ return "" } return parsed +} + +func (o *OAuth) GetSessIdFromCookie(r *http.Request) string { + userSession, err := o.SessStore.Get(r, SessionName) + if err != nil || userSession.IsNew { + return "" + } + s, ok := userSession.Values[SessionId].(string) + if !ok { + return "" + } + return s } func (o *OAuth) AuthorizedClient(r *http.Request) (*atclient.APIClient, error) { diff --git a/appview/state/router.go b/appview/state/router.go --- a/appview/state/router.go +++ b/appview/state/router.go @@ -1,6 +1,7 @@ package state import ( + "context" "database/sql" "errors" "net/http" @@ -13,7 +14,7 @@ "tangled.org/core/appview/labels" "tangled.org/core/appview/metrics" "tangled.org/core/appview/middleware" - "tangled.org/core/appview/migration" + // "tangled.org/core/appview/migration" "tangled.org/core/appview/notifications" "tangled.org/core/appview/pipelines" "tangled.org/core/appview/pulls" @@ -41,8 +42,12 @@ router.Use(metrics.Middleware) - m := migration.NewMigration(s.db, s.oauth, s.idResolver.Directory(), s.logger) - router.Use(m.BackgroundMigrationMiddleware) + if err := db.ReapStaleRunningMigrations(context.Background(), s.db); err != nil { + s.logger.Warn("failed to reap stale running migrations", "err", err) + } + // PDS record migrator disabled while we isolate OAuth refresh behaviour. + // m := migration.NewMigration(s.db, s.oauth, s.idResolver.Directory(), s.logger) + // router.Use(m.BackgroundMigrationMiddleware) router.Get("/pwa-manifest.json", s.WebAppManifest) router.Get("/robots.txt", s.RobotsTxt)