Files
at-container-registry/pkg/auth/oauth/refresher.go
T

225 lines
6.0 KiB
Go

package oauth
import (
"context"
"fmt"
"net/http"
"sync"
"github.com/bluesky-social/indigo/atproto/auth/oauth"
"github.com/bluesky-social/indigo/atproto/syntax"
)
// SessionCache represents a cached OAuth session
type SessionCache struct {
Session *oauth.ClientSession
SessionID string
}
// Refresher manages OAuth sessions and token refresh for AppView
type Refresher struct {
app *App
sessions map[string]*SessionCache // Key: DID string
mu sync.RWMutex
refreshLocks map[string]*sync.Mutex // Per-DID locks for refresh operations
refreshLockMu sync.Mutex // Protects refreshLocks map
}
// NewRefresher creates a new session refresher
func NewRefresher(app *App) *Refresher {
return &Refresher{
app: app,
sessions: make(map[string]*SessionCache),
refreshLocks: make(map[string]*sync.Mutex),
}
}
// GetSession gets a fresh OAuth session for a DID
// Returns cached session if still valid, otherwise resumes from store
func (r *Refresher) GetSession(ctx context.Context, did string) (*oauth.ClientSession, error) {
// Check cache first (fast path)
r.mu.RLock()
cached, ok := r.sessions[did]
r.mu.RUnlock()
if ok && cached.Session != nil {
// Session cached, tokens will auto-refresh if needed
return cached.Session, nil
}
// Session not cached, need to resume from store
// Get or create per-DID lock to prevent concurrent resume operations
r.refreshLockMu.Lock()
didLock, ok := r.refreshLocks[did]
if !ok {
didLock = &sync.Mutex{}
r.refreshLocks[did] = didLock
}
r.refreshLockMu.Unlock()
// Acquire DID-specific lock
didLock.Lock()
defer didLock.Unlock()
// Double-check cache after acquiring lock (another goroutine might have loaded it)
r.mu.RLock()
cached, ok = r.sessions[did]
r.mu.RUnlock()
if ok && cached.Session != nil {
return cached.Session, nil
}
// Actually resume the session
return r.resumeSession(ctx, did)
}
// GetAccessToken gets a fresh access token for a DID
// This is a convenience method that extracts the access token from the session
func (r *Refresher) GetAccessToken(ctx context.Context, did string) (string, error) {
session, err := r.GetSession(ctx, did)
if err != nil {
return "", err
}
// Get access token and DPoP nonce from session
accessToken, _ := session.GetHostAccessData()
return accessToken, nil
}
// GetHTTPClient returns an HTTP client with DPoP authentication for a DID
// The client automatically adds DPoP headers and refreshes tokens as needed
func (r *Refresher) GetHTTPClient(ctx context.Context, did string) (*http.Client, error) {
session, err := r.GetSession(ctx, did)
if err != nil {
return nil, err
}
// Get API client from session
// This client automatically handles DPoP and token refresh
apiClient := session.APIClient()
return apiClient.Client, nil
}
// resumeSession loads a session from storage and caches it
func (r *Refresher) resumeSession(ctx context.Context, did string) (*oauth.ClientSession, error) {
// Parse DID
accountDID, err := syntax.ParseDID(did)
if err != nil {
return nil, fmt.Errorf("failed to parse DID: %w", err)
}
// Get all sessions for this DID from store
fileStore, ok := r.app.clientApp.Store.(*FileStore)
if !ok {
return nil, fmt.Errorf("store is not a FileStore")
}
// Find a session for this DID
sessions := fileStore.ListSessions()
var sessionID string
for _, sessionData := range sessions {
if sessionData.AccountDID.String() == did {
sessionID = sessionData.SessionID
break
}
}
if sessionID == "" {
return nil, fmt.Errorf("no session found for DID: %s", did)
}
// Resume session
session, err := r.app.ResumeSession(ctx, accountDID, sessionID)
if err != nil {
return nil, fmt.Errorf("failed to resume session: %w", err)
}
// Cache the session
r.mu.Lock()
r.sessions[did] = &SessionCache{
Session: session,
SessionID: sessionID,
}
r.mu.Unlock()
return session, nil
}
// InvalidateSession removes a cached session for a DID
// This is useful when a new OAuth flow creates a fresh session
func (r *Refresher) InvalidateSession(did string) {
r.mu.Lock()
delete(r.sessions, did)
r.mu.Unlock()
}
// RevokeSession removes a session from both cache and storage
func (r *Refresher) RevokeSession(ctx context.Context, did string) error {
// Remove from cache
r.mu.Lock()
cached, ok := r.sessions[did]
delete(r.sessions, did)
r.mu.Unlock()
if !ok {
// Not cached, still try to delete from storage
accountDID, err := syntax.ParseDID(did)
if err != nil {
return fmt.Errorf("failed to parse DID: %w", err)
}
// Find session ID from store
fileStore, ok := r.app.clientApp.Store.(*FileStore)
if !ok {
return fmt.Errorf("store is not a FileStore")
}
sessions := fileStore.ListSessions()
for _, sessionData := range sessions {
if sessionData.AccountDID.String() == did {
return r.app.clientApp.Store.DeleteSession(ctx, accountDID, sessionData.SessionID)
}
}
return fmt.Errorf("no session found for DID: %s", did)
}
// Revoke the session via OAuth
if err := cached.Session.RevokeSession(ctx); err != nil {
fmt.Printf("WARNING: failed to revoke session for %s: %v\n", did, err)
// Continue anyway to delete from storage
}
// Delete from storage
accountDID, err := syntax.ParseDID(did)
if err != nil {
return fmt.Errorf("failed to parse DID: %w", err)
}
return r.app.clientApp.Store.DeleteSession(ctx, accountDID, cached.SessionID)
}
// CleanupExpiredSessions removes expired sessions from cache
// Note: indigo handles token expiry automatically, but we clean up orphaned cache entries
func (r *Refresher) CleanupExpiredSessions(ctx context.Context) {
r.mu.Lock()
defer r.mu.Unlock()
// For each cached session, verify it still exists in storage
for did, cached := range r.sessions {
accountDID, err := syntax.ParseDID(did)
if err != nil {
delete(r.sessions, did)
continue
}
// Try to get session from store
_, err = r.app.clientApp.Store.GetSession(ctx, accountDID, cached.SessionID)
if err != nil {
// Session no longer exists, remove from cache
delete(r.sessions, did)
}
}
}