mirror of
https://tangled.org/evan.jarrett.net/at-container-registry
synced 2026-09-21 09:44:15 +00:00
421 lines
12 KiB
Go
421 lines
12 KiB
Go
package db
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"encoding/json"
|
|
"fmt"
|
|
"log/slog"
|
|
"time"
|
|
|
|
atoauth "atcr.io/pkg/auth/oauth"
|
|
"github.com/bluesky-social/indigo/atproto/auth/oauth"
|
|
"github.com/bluesky-social/indigo/atproto/syntax"
|
|
)
|
|
|
|
// OAuthStore implements oauth.ClientAuthStore with SQLite persistence
|
|
type OAuthStore struct {
|
|
db *sql.DB
|
|
}
|
|
|
|
// NewOAuthStore creates a new SQLite-backed OAuth store
|
|
func NewOAuthStore(db *sql.DB) *OAuthStore {
|
|
return &OAuthStore{db: db}
|
|
}
|
|
|
|
// GetSession retrieves a session by DID and session ID
|
|
func (s *OAuthStore) GetSession(ctx context.Context, did syntax.DID, sessionID string) (*oauth.ClientSessionData, error) {
|
|
sessionKey := makeSessionKey(did.String(), sessionID)
|
|
|
|
var sessionDataJSON string
|
|
|
|
err := s.db.QueryRowContext(ctx, `
|
|
SELECT session_data
|
|
FROM oauth_sessions
|
|
WHERE session_key = ?
|
|
`, sessionKey).Scan(&sessionDataJSON)
|
|
|
|
if err == sql.ErrNoRows {
|
|
return nil, fmt.Errorf("session not found: %s/%s", did, sessionID)
|
|
}
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to query session: %w", err)
|
|
}
|
|
|
|
// Parse session data JSON
|
|
var sessionData oauth.ClientSessionData
|
|
if err := json.Unmarshal([]byte(sessionDataJSON), &sessionData); err != nil {
|
|
return nil, fmt.Errorf("failed to parse session data: %w", err)
|
|
}
|
|
|
|
return &sessionData, nil
|
|
}
|
|
|
|
// SaveSession saves or updates a session (upsert)
|
|
func (s *OAuthStore) SaveSession(ctx context.Context, sess oauth.ClientSessionData) error {
|
|
sessionKey := makeSessionKey(sess.AccountDID.String(), sess.SessionID)
|
|
|
|
// Marshal entire session to JSON
|
|
sessionDataJSON, err := json.Marshal(sess)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to marshal session data: %w", err)
|
|
}
|
|
|
|
_, err = s.db.ExecContext(ctx, `
|
|
INSERT INTO oauth_sessions (
|
|
session_key, account_did, session_id, session_data,
|
|
created_at, updated_at
|
|
) VALUES (?, ?, ?, ?, datetime('now'), datetime('now'))
|
|
ON CONFLICT(session_key) DO UPDATE SET
|
|
session_data = excluded.session_data,
|
|
updated_at = datetime('now')
|
|
`,
|
|
sessionKey,
|
|
sess.AccountDID.String(),
|
|
sess.SessionID,
|
|
string(sessionDataJSON),
|
|
)
|
|
|
|
if err != nil {
|
|
return fmt.Errorf("failed to save session: %w", err)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// DeleteSession removes a session
|
|
func (s *OAuthStore) DeleteSession(ctx context.Context, did syntax.DID, sessionID string) error {
|
|
sessionKey := makeSessionKey(did.String(), sessionID)
|
|
|
|
_, err := s.db.ExecContext(ctx, `
|
|
DELETE FROM oauth_sessions WHERE session_key = ?
|
|
`, sessionKey)
|
|
|
|
return err
|
|
}
|
|
|
|
// DeleteSessionsForDID removes all sessions for a given DID
|
|
// This is useful for logout flows where we want to revoke all OAuth sessions
|
|
func (s *OAuthStore) DeleteSessionsForDID(ctx context.Context, did string) error {
|
|
result, err := s.db.ExecContext(ctx, `
|
|
DELETE FROM oauth_sessions WHERE account_did = ?
|
|
`, did)
|
|
|
|
if err != nil {
|
|
return fmt.Errorf("failed to delete sessions for DID: %w", err)
|
|
}
|
|
|
|
deleted, _ := result.RowsAffected()
|
|
if deleted > 0 {
|
|
slog.Info("Deleted OAuth sessions for DID", "count", deleted, "did", did)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// DeleteOldSessionsForDID removes all sessions for a DID except the specified session to keep
|
|
// This is used during OAuth callback to clean up stale sessions with expired refresh tokens
|
|
func (s *OAuthStore) DeleteOldSessionsForDID(ctx context.Context, did string, keepSessionID string) error {
|
|
result, err := s.db.ExecContext(ctx, `
|
|
DELETE FROM oauth_sessions WHERE account_did = ? AND session_id != ?
|
|
`, did, keepSessionID)
|
|
|
|
if err != nil {
|
|
return fmt.Errorf("failed to delete old sessions for DID: %w", err)
|
|
}
|
|
|
|
deleted, _ := result.RowsAffected()
|
|
if deleted > 0 {
|
|
slog.Info("Deleted old OAuth sessions for DID", "count", deleted, "did", did, "kept", keepSessionID)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// GetAuthRequestInfo retrieves authentication request data by state
|
|
func (s *OAuthStore) GetAuthRequestInfo(ctx context.Context, state string) (*oauth.AuthRequestData, error) {
|
|
var requestDataJSON string
|
|
|
|
err := s.db.QueryRowContext(ctx, `
|
|
SELECT request_data FROM oauth_auth_requests WHERE state = ?
|
|
`, state).Scan(&requestDataJSON)
|
|
|
|
if err == sql.ErrNoRows {
|
|
return nil, fmt.Errorf("auth request not found: %s", state)
|
|
}
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to query auth request: %w", err)
|
|
}
|
|
|
|
var requestData oauth.AuthRequestData
|
|
if err := json.Unmarshal([]byte(requestDataJSON), &requestData); err != nil {
|
|
return nil, fmt.Errorf("failed to parse auth request data: %w", err)
|
|
}
|
|
|
|
return &requestData, nil
|
|
}
|
|
|
|
// SaveAuthRequestInfo saves authentication request data
|
|
func (s *OAuthStore) SaveAuthRequestInfo(ctx context.Context, info oauth.AuthRequestData) error {
|
|
requestDataJSON, err := json.Marshal(info)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to marshal auth request data: %w", err)
|
|
}
|
|
|
|
_, err = s.db.ExecContext(ctx, `
|
|
INSERT INTO oauth_auth_requests (state, request_data, created_at)
|
|
VALUES (?, ?, datetime('now'))
|
|
`, info.State, string(requestDataJSON))
|
|
|
|
if err != nil {
|
|
return fmt.Errorf("failed to save auth request: %w", err)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// DeleteAuthRequestInfo removes authentication request data
|
|
func (s *OAuthStore) DeleteAuthRequestInfo(ctx context.Context, state string) error {
|
|
_, err := s.db.ExecContext(ctx, `
|
|
DELETE FROM oauth_auth_requests WHERE state = ?
|
|
`, state)
|
|
|
|
return err
|
|
}
|
|
|
|
// GetLatestSessionForDID returns the most recently updated session for a DID
|
|
// This is the key improvement over the file-based store - we can query by timestamp
|
|
func (s *OAuthStore) GetLatestSessionForDID(ctx context.Context, did string) (*oauth.ClientSessionData, string, error) {
|
|
var sessionDataJSON string
|
|
var sessionID string
|
|
|
|
err := s.db.QueryRowContext(ctx, `
|
|
SELECT session_id, session_data
|
|
FROM oauth_sessions
|
|
WHERE account_did = ?
|
|
ORDER BY updated_at DESC
|
|
LIMIT 1
|
|
`, did).Scan(&sessionID, &sessionDataJSON)
|
|
|
|
if err == sql.ErrNoRows {
|
|
return nil, "", fmt.Errorf("no session found for DID: %s", did)
|
|
}
|
|
if err != nil {
|
|
return nil, "", fmt.Errorf("failed to query session: %w", err)
|
|
}
|
|
|
|
// Parse session data JSON
|
|
var sessionData oauth.ClientSessionData
|
|
if err := json.Unmarshal([]byte(sessionDataJSON), &sessionData); err != nil {
|
|
return nil, "", fmt.Errorf("failed to parse session data: %w", err)
|
|
}
|
|
|
|
return &sessionData, sessionID, nil
|
|
}
|
|
|
|
// CleanupOldSessions removes sessions older than the specified duration
|
|
func (s *OAuthStore) CleanupOldSessions(ctx context.Context, olderThan time.Duration) {
|
|
cutoff := time.Now().Add(-olderThan)
|
|
|
|
result, err := s.db.ExecContext(ctx, `
|
|
DELETE FROM oauth_sessions
|
|
WHERE updated_at < ?
|
|
`, cutoff)
|
|
|
|
if err != nil {
|
|
slog.Warn("Failed to cleanup old OAuth sessions", "component", "oauth_store", "error", err)
|
|
return
|
|
}
|
|
|
|
deleted, _ := result.RowsAffected()
|
|
if deleted > 0 {
|
|
slog.Info("Cleaned up old OAuth sessions", "count", deleted, "older_than", olderThan)
|
|
}
|
|
}
|
|
|
|
// CleanupExpiredAuthRequests removes auth requests older than 10 minutes
|
|
func (s *OAuthStore) CleanupExpiredAuthRequests(ctx context.Context) {
|
|
cutoff := time.Now().Add(-10 * time.Minute)
|
|
|
|
result, err := s.db.ExecContext(ctx, `
|
|
DELETE FROM oauth_auth_requests
|
|
WHERE created_at < ?
|
|
`, cutoff)
|
|
|
|
if err != nil {
|
|
slog.Warn("Failed to cleanup expired auth requests", "component", "oauth_store", "error", err)
|
|
return
|
|
}
|
|
|
|
deleted, _ := result.RowsAffected()
|
|
if deleted > 0 {
|
|
slog.Info("Cleaned up expired auth requests", "count", deleted)
|
|
}
|
|
}
|
|
|
|
// InvalidateSessionsWithMismatchedScopes removes all sessions whose scopes don't match the desired scopes
|
|
// This is called on AppView startup to ensure all sessions have current scopes
|
|
// Returns the count of invalidated sessions
|
|
func (s *OAuthStore) InvalidateSessionsWithMismatchedScopes(ctx context.Context, desiredScopes []string) (int, error) {
|
|
// Query all sessions
|
|
rows, err := s.db.QueryContext(ctx, `
|
|
SELECT session_key, account_did, session_id, session_data
|
|
FROM oauth_sessions
|
|
`)
|
|
if err != nil {
|
|
return 0, fmt.Errorf("failed to query sessions: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
|
|
var sessionsToDelete []string
|
|
for rows.Next() {
|
|
var sessionKey, accountDID, sessionID, sessionDataJSON string
|
|
if err := rows.Scan(&sessionKey, &accountDID, &sessionID, &sessionDataJSON); err != nil {
|
|
slog.Warn("Failed to scan session row", "component", "oauth/store", "error", err)
|
|
continue
|
|
}
|
|
|
|
// Parse session data
|
|
var sessionData oauth.ClientSessionData
|
|
if err := json.Unmarshal([]byte(sessionDataJSON), &sessionData); err != nil {
|
|
slog.Warn("Failed to parse session data", "component", "oauth/store", "session_key", sessionKey, "error", err)
|
|
// Delete malformed sessions
|
|
sessionsToDelete = append(sessionsToDelete, sessionKey)
|
|
continue
|
|
}
|
|
|
|
// Check if scopes match (expands include: scopes before comparing)
|
|
if !atoauth.ScopesMatch(sessionData.Scopes, desiredScopes) {
|
|
slog.Debug("Session has mismatched scopes",
|
|
"component", "oauth/store",
|
|
"session_key", sessionKey,
|
|
"account_did", accountDID,
|
|
"session_scopes", sessionData.Scopes,
|
|
"desired_scopes", desiredScopes,
|
|
)
|
|
sessionsToDelete = append(sessionsToDelete, sessionKey)
|
|
}
|
|
}
|
|
|
|
if err := rows.Err(); err != nil {
|
|
return 0, fmt.Errorf("error iterating sessions: %w", err)
|
|
}
|
|
|
|
// Delete sessions with mismatched scopes
|
|
if len(sessionsToDelete) > 0 {
|
|
for _, key := range sessionsToDelete {
|
|
_, err := s.db.ExecContext(ctx, `
|
|
DELETE FROM oauth_sessions WHERE session_key = ?
|
|
`, key)
|
|
if err != nil {
|
|
slog.Warn("Failed to delete session", "component", "oauth/store", "session_key", key, "error", err)
|
|
}
|
|
}
|
|
slog.Info("Invalidated OAuth sessions with mismatched scopes", "count", len(sessionsToDelete))
|
|
}
|
|
|
|
return len(sessionsToDelete), nil
|
|
}
|
|
|
|
// GetSessionStats returns statistics about stored OAuth sessions
|
|
// Useful for monitoring and debugging session health
|
|
func (s *OAuthStore) GetSessionStats(ctx context.Context) (map[string]any, error) {
|
|
stats := make(map[string]any)
|
|
|
|
// Total sessions
|
|
var totalSessions int
|
|
err := s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM oauth_sessions`).Scan(&totalSessions)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to count sessions: %w", err)
|
|
}
|
|
stats["total_sessions"] = totalSessions
|
|
|
|
// Sessions by age
|
|
var sessionsOlderThan1Hour, sessionsOlderThan1Day, sessionsOlderThan7Days int
|
|
|
|
err = s.db.QueryRowContext(ctx, `
|
|
SELECT COUNT(*) FROM oauth_sessions
|
|
WHERE updated_at < datetime('now', '-1 hour')
|
|
`).Scan(&sessionsOlderThan1Hour)
|
|
if err == nil {
|
|
stats["sessions_idle_1h+"] = sessionsOlderThan1Hour
|
|
}
|
|
|
|
err = s.db.QueryRowContext(ctx, `
|
|
SELECT COUNT(*) FROM oauth_sessions
|
|
WHERE updated_at < datetime('now', '-1 day')
|
|
`).Scan(&sessionsOlderThan1Day)
|
|
if err == nil {
|
|
stats["sessions_idle_1d+"] = sessionsOlderThan1Day
|
|
}
|
|
|
|
err = s.db.QueryRowContext(ctx, `
|
|
SELECT COUNT(*) FROM oauth_sessions
|
|
WHERE updated_at < datetime('now', '-7 days')
|
|
`).Scan(&sessionsOlderThan7Days)
|
|
if err == nil {
|
|
stats["sessions_idle_7d+"] = sessionsOlderThan7Days
|
|
}
|
|
|
|
// Recent sessions (updated in last 5 minutes)
|
|
var recentSessions int
|
|
err = s.db.QueryRowContext(ctx, `
|
|
SELECT COUNT(*) FROM oauth_sessions
|
|
WHERE updated_at > datetime('now', '-5 minutes')
|
|
`).Scan(&recentSessions)
|
|
if err == nil {
|
|
stats["sessions_active_5m"] = recentSessions
|
|
}
|
|
|
|
return stats, nil
|
|
}
|
|
|
|
// ListSessionsForMonitoring returns a list of all sessions with basic info for monitoring
|
|
// Returns: DID, session age (minutes), last update time
|
|
func (s *OAuthStore) ListSessionsForMonitoring(ctx context.Context) ([]map[string]any, error) {
|
|
rows, err := s.db.QueryContext(ctx, `
|
|
SELECT
|
|
account_did,
|
|
session_id,
|
|
created_at,
|
|
updated_at,
|
|
CAST((julianday('now') - julianday(updated_at)) * 24 * 60 AS INTEGER) as idle_minutes
|
|
FROM oauth_sessions
|
|
ORDER BY updated_at DESC
|
|
`)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to query sessions: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
|
|
var sessions []map[string]any
|
|
for rows.Next() {
|
|
var did, sessionID, createdAt, updatedAt string
|
|
var idleMinutes int
|
|
|
|
if err := rows.Scan(&did, &sessionID, &createdAt, &updatedAt, &idleMinutes); err != nil {
|
|
slog.Warn("Failed to scan session row", "error", err)
|
|
continue
|
|
}
|
|
|
|
sessions = append(sessions, map[string]any{
|
|
"did": did,
|
|
"session_id": sessionID,
|
|
"created_at": createdAt,
|
|
"updated_at": updatedAt,
|
|
"idle_minutes": idleMinutes,
|
|
})
|
|
}
|
|
|
|
if err := rows.Err(); err != nil {
|
|
return nil, fmt.Errorf("error iterating sessions: %w", err)
|
|
}
|
|
|
|
return sessions, nil
|
|
}
|
|
|
|
// makeSessionKey creates a composite key for session storage
|
|
func makeSessionKey(did, sessionID string) string {
|
|
return fmt.Sprintf("%s:%s", did, sessionID)
|
|
}
|