Something went wrong. Try again.
Monorepo for Tangled
Something went wrong. Try again.
Go
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120package oauth
import ( "context" "database/sql" "encoding/json" "errors" "fmt" "time"
indigooauth "github.com/bluesky-social/indigo/atproto/auth/oauth" "github.com/bluesky-social/indigo/atproto/syntax" "tangled.org/core/migrator/crypto" "tangled.org/core/migrator/db")
type Store struct { db *db.DB key crypto.MasterKey}
func NewStore(database *db.DB, key crypto.MasterKey) *Store { return &Store{db: database, key: key}}
func sessionAAD(did syntax.DID) []byte { return []byte("oauth-session|" + did.String()) }func requestAAD(state string) []byte { return []byte("oauth-request|" + state) }
func sealJSON[T any](key crypto.MasterKey, value T, aad []byte) (string, error) { raw, err := json.Marshal(value) if err != nil { return "", err } return crypto.Encrypt(key, string(raw), aad)}
func openJSON[T any](key crypto.MasterKey, envelope string, aad []byte) (*T, error) { raw, err := crypto.Decrypt(key, envelope, aad) if err != nil { return nil, err } var value T if err := json.Unmarshal([]byte(raw), &value); err != nil { return nil, err } return &value, nil}
func (s *Store) GetSession(ctx context.Context, did syntax.DID, _ string) (*indigooauth.ClientSessionData, error) { var encrypted string err := s.db.QueryRowContext(ctx, "select encrypted_data from oauth_sessions where owner_did = ?", did.String()).Scan(&encrypted) if errors.Is(err, sql.ErrNoRows) { return nil, fmt.Errorf("oauth session not found for %s", did) } if err != nil { return nil, err } return openJSON[indigooauth.ClientSessionData](s.key, encrypted, sessionAAD(did))}
func (s *Store) SaveSession(ctx context.Context, session indigooauth.ClientSessionData) error { encrypted, err := sealJSON(s.key, session, sessionAAD(session.AccountDID)) if err != nil { return err } _, err = s.db.ExecContext(ctx, ` insert into oauth_sessions (owner_did, session_id, encrypted_data, updated_at) values (?, ?, ?, ?) on conflict(owner_did) do update set session_id = excluded.session_id, encrypted_data = excluded.encrypted_data, updated_at = excluded.updated_at `, session.AccountDID.String(), session.SessionID, encrypted, time.Now().UTC().Format(time.RFC3339)) return err}
func (s *Store) DeleteSession(ctx context.Context, did syntax.DID, _ string) error { _, err := s.db.ExecContext(ctx, "delete from oauth_sessions where owner_did = ?", did.String()) return err}
func (s *Store) GetAuthRequestInfo(ctx context.Context, state string) (*indigooauth.AuthRequestData, error) { var encrypted string cutoff := time.Now().UTC().Add(-30 * time.Minute).Format(time.RFC3339) err := s.db.QueryRowContext(ctx, "select encrypted_data from oauth_requests where state = ? and created_at >= ?", state, cutoff).Scan(&encrypted) if errors.Is(err, sql.ErrNoRows) { return nil, fmt.Errorf("oauth request not found") } if err != nil { return nil, err } return openJSON[indigooauth.AuthRequestData](s.key, encrypted, requestAAD(state))}
func (s *Store) SaveAuthRequestInfo(ctx context.Context, info indigooauth.AuthRequestData) error { encrypted, err := sealJSON(s.key, info, requestAAD(info.State)) if err != nil { return err } if _, err := s.db.ExecContext(ctx, "delete from oauth_requests where created_at < ?", time.Now().UTC().Add(-30*time.Minute).Format(time.RFC3339)); err != nil { return err } _, err = s.db.ExecContext(ctx, "insert into oauth_requests (state, encrypted_data) values (?, ?)", info.State, encrypted) return err}
func (s *Store) DeleteAuthRequestInfo(ctx context.Context, state string) error { _, err := s.db.ExecContext(ctx, "delete from oauth_requests where state = ?", state) return err}
func (s *Store) SessionID(ctx context.Context, did string) (string, error) { var sessionID string err := s.db.QueryRowContext(ctx, "select session_id from oauth_sessions where owner_did = ?", did).Scan(&sessionID) if errors.Is(err, sql.ErrNoRows) { return "", db.ErrNotFound } return sessionID, err}