summaryrefslogtreecommitdiff
path: root/internal/controlplane/identity.go
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-05 14:48:00 +1200
committerChia <Chia@93.nz>2026-08-05 14:48:00 +1200
commitcd0dd91ab93653631904f2ea0e574ccde6d60339 (patch)
treec65417b880a3f4a35c504c44edae821bc2122f70 /internal/controlplane/identity.go
parent86b1f42e3c5601ff10621a9779cf0076590797a1 (diff)
add passkey, totp.
Diffstat (limited to 'internal/controlplane/identity.go')
-rw-r--r--internal/controlplane/identity.go291
1 files changed, 291 insertions, 0 deletions
diff --git a/internal/controlplane/identity.go b/internal/controlplane/identity.go
new file mode 100644
index 0000000..07fe0d7
--- /dev/null
+++ b/internal/controlplane/identity.go
@@ -0,0 +1,291 @@
+package controlplane
+
+import (
+ "context"
+ "crypto/sha256"
+ "errors"
+ "fmt"
+ "net/mail"
+ "strings"
+ "time"
+
+ "aigw/internal/security"
+
+ "github.com/jackc/pgx/v5"
+)
+
+var (
+ ErrActionTokenInvalid = errors.New("account action link is invalid or expired")
+ ErrMFARequired = errors.New("multi-factor authentication is required")
+)
+
+func (s *Store) RegisterTenantPending(ctx context.Context, input RegisterInput, currency, publicURL, remoteIP string) (ConsoleActor, int64, error) {
+ input.Organization = strings.TrimSpace(input.Organization)
+ input.TenantSlug = strings.ToLower(strings.TrimSpace(input.TenantSlug))
+ input.DisplayName = strings.TrimSpace(input.DisplayName)
+ input.Email = strings.ToLower(strings.TrimSpace(input.Email))
+ if input.Organization == "" || input.DisplayName == "" || !slugPattern.MatchString(input.TenantSlug) || !validEmail(input.Email) {
+ return ConsoleActor{}, 0, errors.New("registration requires organization, a valid tenant_slug, display_name, and email")
+ }
+ if len(currency) != 3 {
+ return ConsoleActor{}, 0, errors.New("registration currency is invalid")
+ }
+ hash, salt, iterations, err := security.HashPassword(input.Password)
+ if err != nil {
+ return ConsoleActor{}, 0, err
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return ConsoleActor{}, 0, err
+ }
+ defer tx.Rollback(ctx)
+ var tenantID string
+ if err := tx.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,$2) RETURNING id::text`, input.TenantSlug, input.Organization).Scan(&tenantID); err != nil {
+ return ConsoleActor{}, 0, fmt.Errorf("create registered tenant: %w", err)
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO projects (tenant_id,slug,name) VALUES ($1,'default','Default project')`, tenantID); err != nil {
+ return ConsoleActor{}, 0, fmt.Errorf("create default project: %w", err)
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO tenant_wallets (tenant_id,currency) VALUES ($1,$2)`, tenantID, strings.ToLower(currency)); err != nil {
+ return ConsoleActor{}, 0, fmt.Errorf("create tenant wallet: %w", err)
+ }
+ actor := ConsoleActor{TenantID: tenantID, Email: input.Email, DisplayName: input.DisplayName, Role: RoleTenantAdmin}
+ if err := tx.QueryRow(ctx, `INSERT INTO console_users
+ (tenant_id,email,display_name,role,password_hash,password_salt,password_iterations,password_changed_at,status)
+ VALUES ($1,$2,$3,$4,$5,$6,$7,now(),'pending_verification') RETURNING id::text`, tenantID, input.Email,
+ input.DisplayName, RoleTenantAdmin, hash, salt, iterations).Scan(&actor.ID); err != nil {
+ return ConsoleActor{}, 0, fmt.Errorf("create pending tenant administrator: %w", err)
+ }
+ if _, err := s.issueActionToken(ctx, tx, actor.ID, actor.Email, actor.DisplayName, actionVerifyEmail, publicURL, remoteIP, 24*time.Hour); err != nil {
+ return ConsoleActor{}, 0, err
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return ConsoleActor{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return ConsoleActor{}, 0, err
+ }
+ return actor, generation, nil
+}
+
+func (s *Store) InviteConsoleUser(ctx context.Context, input CreateConsoleUserInput, publicURL, remoteIP string) (CreatedConsoleUser, error) {
+ input.Email = strings.ToLower(strings.TrimSpace(input.Email))
+ input.DisplayName = strings.TrimSpace(input.DisplayName)
+ input.TenantID = strings.TrimSpace(input.TenantID)
+ address, addressErr := mail.ParseAddress(input.Email)
+ if addressErr != nil || address.Address != input.Email || input.DisplayName == "" {
+ return CreatedConsoleUser{}, errors.New("invitation requires a valid email and display_name")
+ }
+ if err := validateConsoleRole(input.Role, input.TenantID); err != nil {
+ return CreatedConsoleUser{}, err
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return CreatedConsoleUser{}, err
+ }
+ defer tx.Rollback(ctx)
+ var result CreatedConsoleUser
+ err = tx.QueryRow(ctx, `INSERT INTO console_users
+ (tenant_id,email,display_name,role,status,invited_at)
+ VALUES (NULLIF($1,'')::uuid,$2,$3,$4,'invited',now())
+ RETURNING id::text,COALESCE(tenant_id::text,''),email,display_name,role,COALESCE(token_prefix,''),
+ password_hash IS NOT NULL,status,FALSE,'{}'::text[],last_used_at,created_at`,
+ input.TenantID, input.Email, input.DisplayName, input.Role,
+ ).Scan(&result.ID, &result.TenantID, &result.Email, &result.DisplayName, &result.Role, &result.TokenPrefix,
+ &result.HasPassword, &result.Status, &result.EmailVerified, &result.MFAMethods, &result.LastUsedAt, &result.CreatedAt)
+ if err != nil {
+ return CreatedConsoleUser{}, fmt.Errorf("create invited console user: %w", err)
+ }
+ if _, err := s.issueActionToken(ctx, tx, result.ID, result.Email, result.DisplayName, actionInvite, publicURL, remoteIP, 7*24*time.Hour); err != nil {
+ return CreatedConsoleUser{}, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return CreatedConsoleUser{}, err
+ }
+ return result, nil
+}
+
+func validateConsoleRole(role, tenantID string) error {
+ platform := role == RolePlatformAdmin || role == RolePlatformViewer
+ tenant := role == RoleTenantAdmin || role == RoleTenantBilling || role == RoleTenantDeveloper || role == RoleTenantViewer
+ if (!platform && !tenant) || (platform && tenantID != "") || (tenant && tenantID == "") {
+ return errors.New("console user role and tenant_id are inconsistent")
+ }
+ return nil
+}
+
+func (s *Store) VerifyEmailToken(ctx context.Context, rawToken string) (ConsoleActor, error) {
+ tokenHash := sha256.Sum256([]byte(strings.TrimSpace(rawToken)))
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return ConsoleActor{}, err
+ }
+ defer tx.Rollback(ctx)
+ var tokenID string
+ var actor ConsoleActor
+ err = tx.QueryRow(ctx, `SELECT a.id::text,u.id::text,COALESCE(u.tenant_id::text,''),u.email,u.display_name,u.role
+ FROM console_action_tokens a JOIN console_users u ON u.id=a.user_id
+ WHERE a.token_hash=$1 AND a.purpose='verify_email' AND a.consumed_at IS NULL AND a.expires_at>now()
+ AND u.status='pending_verification' FOR UPDATE OF a,u`, tokenHash[:],
+ ).Scan(&tokenID, &actor.ID, &actor.TenantID, &actor.Email, &actor.DisplayName, &actor.Role)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return ConsoleActor{}, ErrActionTokenInvalid
+ }
+ if err != nil {
+ return ConsoleActor{}, err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE console_action_tokens SET consumed_at=now() WHERE id=$1`, tokenID); err != nil {
+ return ConsoleActor{}, err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE console_users SET status='active',email_verified_at=now(),accepted_at=COALESCE(accepted_at,now()) WHERE id=$1`, actor.ID); err != nil {
+ return ConsoleActor{}, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return ConsoleActor{}, err
+ }
+ return actor, nil
+}
+
+func (s *Store) ResendVerification(ctx context.Context, email, publicURL, remoteIP string) error {
+ email = strings.ToLower(strings.TrimSpace(email))
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return err
+ }
+ defer tx.Rollback(ctx)
+ var id, displayName string
+ err = tx.QueryRow(ctx, `SELECT id::text,display_name FROM console_users
+ WHERE lower(email)=$1 AND status='pending_verification' FOR UPDATE`, email).Scan(&id, &displayName)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return nil
+ }
+ if err != nil {
+ return err
+ }
+ if _, err := s.issueActionToken(ctx, tx, id, email, displayName, actionVerifyEmail, publicURL, remoteIP, 24*time.Hour); err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+}
+
+func (s *Store) RequestPasswordReset(ctx context.Context, email, publicURL, remoteIP string) error {
+ email = strings.ToLower(strings.TrimSpace(email))
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return err
+ }
+ defer tx.Rollback(ctx)
+ var id, displayName string
+ err = tx.QueryRow(ctx, `SELECT id::text,display_name FROM console_users
+ WHERE lower(email)=$1 AND status='active' AND email_verified_at IS NOT NULL FOR UPDATE`, email).Scan(&id, &displayName)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return nil
+ }
+ if err != nil {
+ return err
+ }
+ if _, err := s.issueActionToken(ctx, tx, id, email, displayName, actionPasswordReset, publicURL, remoteIP, 30*time.Minute); err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+}
+
+func (s *Store) ResetPassword(ctx context.Context, input PasswordResetInput) error {
+ hash, salt, iterations, err := security.HashPassword(input.NewPassword)
+ if err != nil {
+ return err
+ }
+ tokenHash := sha256.Sum256([]byte(strings.TrimSpace(input.Token)))
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return err
+ }
+ defer tx.Rollback(ctx)
+ var tokenID, userID string
+ err = tx.QueryRow(ctx, `SELECT a.id::text,u.id::text FROM console_action_tokens a
+ JOIN console_users u ON u.id=a.user_id WHERE a.token_hash=$1 AND a.purpose='password_reset'
+ AND a.consumed_at IS NULL AND a.expires_at>now() AND u.status='active' FOR UPDATE OF a,u`, tokenHash[:]).Scan(&tokenID, &userID)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return ErrActionTokenInvalid
+ }
+ if err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE console_users SET password_hash=$2,password_salt=$3,password_iterations=$4,password_changed_at=now() WHERE id=$1`, userID, hash, salt, iterations); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE console_action_tokens SET consumed_at=now() WHERE id=$1`, tokenID); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE console_sessions SET revoked_at=now() WHERE user_id=$1 AND revoked_at IS NULL`, userID); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE console_auth_challenges SET consumed_at=now() WHERE user_id=$1 AND consumed_at IS NULL`, userID); err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+}
+
+func (s *Store) AcceptInvite(ctx context.Context, input InviteAcceptInput) (ConsoleActor, error) {
+ hash, salt, iterations, err := security.HashPassword(input.Password)
+ if err != nil {
+ return ConsoleActor{}, err
+ }
+ input.DisplayName = strings.TrimSpace(input.DisplayName)
+ if input.DisplayName == "" {
+ return ConsoleActor{}, errors.New("display_name is required")
+ }
+ tokenHash := sha256.Sum256([]byte(strings.TrimSpace(input.Token)))
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return ConsoleActor{}, err
+ }
+ defer tx.Rollback(ctx)
+ var tokenID string
+ var actor ConsoleActor
+ err = tx.QueryRow(ctx, `SELECT a.id::text,u.id::text,COALESCE(u.tenant_id::text,''),u.email,u.role
+ FROM console_action_tokens a JOIN console_users u ON u.id=a.user_id
+ WHERE a.token_hash=$1 AND a.purpose='invite' AND a.consumed_at IS NULL AND a.expires_at>now()
+ AND u.status='invited' FOR UPDATE OF a,u`, tokenHash[:],
+ ).Scan(&tokenID, &actor.ID, &actor.TenantID, &actor.Email, &actor.Role)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return ConsoleActor{}, ErrActionTokenInvalid
+ }
+ if err != nil {
+ return ConsoleActor{}, err
+ }
+ actor.DisplayName = input.DisplayName
+ if _, err := tx.Exec(ctx, `UPDATE console_users SET display_name=$2,password_hash=$3,password_salt=$4,
+ password_iterations=$5,password_changed_at=now(),status='active',email_verified_at=now(),accepted_at=now()
+ WHERE id=$1`, actor.ID, actor.DisplayName, hash, salt, iterations); err != nil {
+ return ConsoleActor{}, err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE console_action_tokens SET consumed_at=now() WHERE id=$1`, tokenID); err != nil {
+ return ConsoleActor{}, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return ConsoleActor{}, err
+ }
+ return actor, nil
+}
+
+func (s *Store) ConsumeRateLimit(ctx context.Context, scope, key string, limit int, window time.Duration) error {
+ if limit < 1 || window < time.Second {
+ return errors.New("invalid rate limit")
+ }
+ bucket := sha256.Sum256([]byte(scope + "\x00" + strings.ToLower(strings.TrimSpace(key))))
+ var hits int
+ err := s.db.QueryRow(ctx, `INSERT INTO console_rate_limits (bucket_hash,hits) VALUES ($1,1)
+ ON CONFLICT (bucket_hash) DO UPDATE SET
+ hits=CASE WHEN console_rate_limits.window_started_at < now()-($2::bigint * interval '1 millisecond') THEN 1 ELSE console_rate_limits.hits+1 END,
+ window_started_at=CASE WHEN console_rate_limits.window_started_at < now()-($2::bigint * interval '1 millisecond') THEN now() ELSE console_rate_limits.window_started_at END,
+ updated_at=now() RETURNING hits`, bucket[:], window.Milliseconds()).Scan(&hits)
+ if err != nil {
+ return fmt.Errorf("consume account rate limit: %w", err)
+ }
+ if hits > limit {
+ return ErrConsoleRateLimited
+ }
+ return nil
+}