diff options
Diffstat (limited to '')
| -rw-r--r-- | internal/controlplane/identity.go | 291 |
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 +} |
