accounts and email

This commit is contained in:
2026-06-04 08:50:34 -04:00
parent 73d505ac7e
commit ddc7d5cbf5
46 changed files with 2685 additions and 309 deletions

View File

@@ -49,16 +49,18 @@ func (r *Repository) UpsertUser(ctx context.Context, user User) (User, error) {
func (r *Repository) GetUserByEmail(ctx context.Context, email string) (User, error) {
var user User
var created, lastSeen, passwordHash, role string
var disabled int
err := r.db.QueryRowContext(ctx, `
SELECT id, email, COALESCE(display_name, ''), COALESCE(password_hash, ''), COALESCE(role, 'viewer'), created_at, COALESCE(last_seen_at, '')
SELECT id, email, COALESCE(display_name, ''), COALESCE(password_hash, ''), COALESCE(role, 'viewer'), disabled, created_at, COALESCE(last_seen_at, '')
FROM users
WHERE email = ?
`, email).Scan(&user.ID, &user.Email, &user.DisplayName, &passwordHash, &role, &created, &lastSeen)
`, email).Scan(&user.ID, &user.Email, &user.DisplayName, &passwordHash, &role, &disabled, &created, &lastSeen)
if err != nil {
return User{}, err
}
user.PasswordHash = passwordHash
user.Role = Role(role)
user.Disabled = disabled == 1
user.CreatedAt, err = time.Parse(time.RFC3339, created)
if err != nil {
return User{}, fmt.Errorf("parse user created_at: %w", err)
@@ -94,7 +96,7 @@ func (r *Repository) CountUsers(ctx context.Context) (int, error) {
func (r *Repository) CountAdmins(ctx context.Context) (int, error) {
var count int
if err := r.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM users WHERE role = 'admin'`).Scan(&count); err != nil {
if err := r.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM users WHERE role = 'admin' AND disabled = 0`).Scan(&count); err != nil {
return 0, err
}
return count, nil
@@ -189,7 +191,7 @@ func (r *Repository) UpdateSignupsEnabled(ctx context.Context, enabled bool) err
func (r *Repository) ListUsers(ctx context.Context) ([]User, error) {
rows, err := r.db.QueryContext(ctx, `
SELECT id, email, COALESCE(display_name, ''), COALESCE(password_hash, ''), COALESCE(role, 'viewer'), created_at, COALESCE(last_seen_at, '')
SELECT id, email, COALESCE(display_name, ''), COALESCE(password_hash, ''), COALESCE(role, 'viewer'), disabled, created_at, COALESCE(last_seen_at, '')
FROM users
ORDER BY created_at DESC
`)
@@ -202,10 +204,12 @@ func (r *Repository) ListUsers(ctx context.Context) ([]User, error) {
for rows.Next() {
var user User
var created, lastSeen, role string
if err := rows.Scan(&user.ID, &user.Email, &user.DisplayName, &user.PasswordHash, &role, &created, &lastSeen); err != nil {
var disabled int
if err := rows.Scan(&user.ID, &user.Email, &user.DisplayName, &user.PasswordHash, &role, &disabled, &created, &lastSeen); err != nil {
return nil, fmt.Errorf("scan user: %w", err)
}
user.Role = Role(role)
user.Disabled = disabled == 1
user.CreatedAt, _ = time.Parse(time.RFC3339, created)
if lastSeen != "" {
user.LastSeenAt, _ = time.Parse(time.RFC3339, lastSeen)
@@ -245,6 +249,41 @@ func (r *Repository) UpdateUserRole(ctx context.Context, userID string, role Rol
return nil
}
func (r *Repository) UpdateUserAccess(ctx context.Context, updates []UserAccessUpdate) error {
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return fmt.Errorf("begin user access update: %w", err)
}
defer func() { _ = tx.Rollback() }()
for _, update := range updates {
result, err := tx.ExecContext(ctx, `
UPDATE users
SET role = ?, disabled = ?
WHERE id = ?
`, string(update.Role), boolInt(update.Disabled), update.ID)
if err != nil {
return fmt.Errorf("update user access: %w", err)
}
if rows, _ := result.RowsAffected(); rows == 0 {
return sql.ErrNoRows
}
}
var admins int
if err := tx.QueryRowContext(ctx, `SELECT COUNT(*) FROM users WHERE role = 'admin' AND disabled = 0`).Scan(&admins); err != nil {
return fmt.Errorf("count enabled admins: %w", err)
}
if admins == 0 {
return fmt.Errorf("cannot disable or demote the last admin account")
}
if err := tx.Commit(); err != nil {
return fmt.Errorf("commit user access update: %w", err)
}
return nil
}
func (r *Repository) DeleteUser(ctx context.Context, userID string) error {
result, err := r.db.ExecContext(ctx, `DELETE FROM users WHERE id = ?`, userID)
if err != nil {
@@ -411,6 +450,87 @@ func (r *Repository) RevokeSession(ctx context.Context, sessionID string) error
return err
}
func (r *Repository) CreatePasswordResetToken(ctx context.Context, token PasswordResetToken) error {
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return fmt.Errorf("begin password reset token: %w", err)
}
defer func() { _ = tx.Rollback() }()
now := time.Now().UTC().Format(time.RFC3339)
if _, err := tx.ExecContext(ctx, `
UPDATE password_reset_tokens
SET used_at = ?
WHERE user_id = ? AND used_at IS NULL
`, now, token.UserID); err != nil {
return fmt.Errorf("expire prior password reset tokens: %w", err)
}
if _, err := tx.ExecContext(ctx, `
INSERT INTO password_reset_tokens (id, user_id, token_hash, created_at, expires_at, initiated_by)
VALUES (?, ?, ?, ?, ?, ?)
`, token.ID, token.UserID, token.TokenHash, token.CreatedAt.Format(time.RFC3339), token.ExpiresAt.Format(time.RFC3339), nullString(token.InitiatedBy)); err != nil {
return fmt.Errorf("create password reset token: %w", err)
}
if err := tx.Commit(); err != nil {
return fmt.Errorf("commit password reset token: %w", err)
}
return nil
}
func (r *Repository) ExpirePasswordResetToken(ctx context.Context, tokenHash string) {
_, _ = r.db.ExecContext(ctx, `
UPDATE password_reset_tokens
SET used_at = ?
WHERE token_hash = ? AND used_at IS NULL
`, time.Now().UTC().Format(time.RFC3339), tokenHash)
}
func (r *Repository) CompletePasswordReset(ctx context.Context, tokenHash, passwordHash string) (string, error) {
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return "", fmt.Errorf("begin password reset: %w", err)
}
defer func() { _ = tx.Rollback() }()
var id, userID, expires, used string
if err := tx.QueryRowContext(ctx, `
SELECT id, user_id, expires_at, COALESCE(used_at, '')
FROM password_reset_tokens
WHERE token_hash = ?
`, tokenHash).Scan(&id, &userID, &expires, &used); err != nil {
return "", err
}
if used != "" {
return "", fmt.Errorf("password reset link has already been used")
}
expiresAt, err := time.Parse(time.RFC3339, expires)
if err != nil {
return "", fmt.Errorf("parse password reset expiry: %w", err)
}
if time.Now().UTC().After(expiresAt) {
return "", fmt.Errorf("password reset link has expired")
}
now := time.Now().UTC().Format(time.RFC3339)
result, err := tx.ExecContext(ctx, `UPDATE password_reset_tokens SET used_at = ? WHERE id = ? AND used_at IS NULL`, now, id)
if err != nil {
return "", fmt.Errorf("consume password reset token: %w", err)
}
if rows, _ := result.RowsAffected(); rows == 0 {
return "", fmt.Errorf("password reset link has already been used")
}
if _, err := tx.ExecContext(ctx, `UPDATE users SET password_hash = ? WHERE id = ?`, passwordHash, userID); err != nil {
return "", fmt.Errorf("update reset password: %w", err)
}
if _, err := tx.ExecContext(ctx, `UPDATE sessions SET revoked_at = ? WHERE user_id = ? AND revoked_at IS NULL`, now, userID); err != nil {
return "", fmt.Errorf("revoke password reset sessions: %w", err)
}
if err := tx.Commit(); err != nil {
return "", fmt.Errorf("commit password reset: %w", err)
}
return userID, nil
}
func (r *Repository) CreateAPIKey(ctx context.Context, key APIKey) error {
if _, err := r.db.ExecContext(ctx, `
INSERT INTO api_keys (id, user_id, name, key_hash, scopes, created_at, expires_at, last_used_at, revoked_at)

View File

@@ -18,13 +18,19 @@ const (
sessionTTL = 24 * time.Hour
challengeTTL = 5 * time.Minute
deviceCodeTTL = 15 * time.Minute
passwordResetTTL = time.Hour
devicePollSeconds = 5
)
type EmailSender interface {
Send(ctx context.Context, to []string, subject, body string) error
}
type Service struct {
repo *Repository
webauthn *webauthn.WebAuthn
publicOrigin string
email EmailSender
}
func NewService(repo *Repository, publicOrigin string) (*Service, error) {
@@ -55,6 +61,10 @@ func NewService(repo *Repository, publicOrigin string) (*Service, error) {
return &Service{repo: repo, webauthn: web, publicOrigin: publicOrigin}, nil
}
func (s *Service) SetEmailSender(sender EmailSender) {
s.email = sender
}
func (s *Service) GetInstanceSettings(ctx context.Context) (InstanceSettings, error) {
return s.repo.GetInstanceSettings(ctx)
}
@@ -119,7 +129,7 @@ func (s *Service) LoginPassword(ctx context.Context, email, password, ip, userAg
return Principal{}, "", fmt.Errorf("invalid credentials")
}
ok, err := verifyPassword(password, user.PasswordHash)
if err != nil || !ok {
if err != nil || !ok || user.Disabled {
return Principal{}, "", fmt.Errorf("invalid credentials")
}
principal, token, err := s.createSession(ctx, user, ip, userAgent)
@@ -185,6 +195,9 @@ func (s *Service) BeginPasskeyLogin(ctx context.Context, email string) (any, str
if err != nil {
return nil, "", err
}
if user.Disabled {
return nil, "", fmt.Errorf("invalid credentials")
}
assertion, session, err := s.webauthn.BeginLogin(user)
if err != nil {
return nil, "", err
@@ -205,6 +218,9 @@ func (s *Service) FinishPasskeyLogin(ctx context.Context, challengeID string, r
if err != nil {
return Principal{}, "", err
}
if user.Disabled {
return Principal{}, "", fmt.Errorf("invalid credentials")
}
credential, err := s.webauthn.FinishLogin(user, session, r)
if err != nil {
return Principal{}, "", err
@@ -224,6 +240,9 @@ func (s *Service) ValidateSessionToken(ctx context.Context, token string) (Princ
if err != nil {
return Principal{}, err
}
if user.Disabled {
return Principal{}, fmt.Errorf("account disabled")
}
return principalFromUser(user, "session", session.ID, "", nil, session.ExpiresAt), nil
}
@@ -324,37 +343,129 @@ func (s *Service) PrincipalForUser(ctx context.Context, userID string) (Principa
if err != nil {
return Principal{}, err
}
if user.Disabled {
return Principal{}, fmt.Errorf("account disabled")
}
return principalFromUser(user, "dev", "", "", nil, time.Time{}), nil
}
func (s *Service) UpdateUserRole(ctx context.Context, actor Principal, userID, role string) (User, error) {
if !Allows(actor, ScopeAdmin) {
return User{}, fmt.Errorf("admin role required")
}
normalizedRole := Role(strings.ToLower(strings.TrimSpace(role)))
switch normalizedRole {
case RoleViewer, RoleEditor, RoleAdmin:
default:
return User{}, fmt.Errorf("invalid role")
}
user, err := s.repo.GetUserByID(ctx, userID)
if err != nil {
return User{}, err
}
if user.Role == RoleAdmin && normalizedRole != RoleAdmin {
admins, err := s.repo.CountAdmins(ctx)
if err != nil {
return User{}, err
}
if admins <= 1 {
return User{}, fmt.Errorf("cannot demote the last admin account")
}
}
if err := s.repo.UpdateUserRole(ctx, user.ID, normalizedRole); err != nil {
users, err := s.UpdateUserAccess(ctx, actor, []UserAccessUpdate{{
ID: userID,
Role: Role(role),
Disabled: user.Disabled,
}})
if err != nil {
return User{}, err
}
s.repo.Audit(ctx, actor.UserID, "auth.user.role.update", "user", user.ID, "", "", map[string]any{"role": normalizedRole})
return s.repo.GetUserByID(ctx, user.ID)
return users[0], nil
}
func (s *Service) UpdateUserAccess(ctx context.Context, actor Principal, updates []UserAccessUpdate) ([]User, error) {
if !Allows(actor, ScopeAdmin) {
return nil, fmt.Errorf("admin role required")
}
if len(updates) == 0 {
return nil, fmt.Errorf("at least one user update is required")
}
normalized := make([]UserAccessUpdate, 0, len(updates))
seen := make(map[string]struct{}, len(updates))
for _, update := range updates {
update.ID = strings.TrimSpace(update.ID)
if update.ID == "" {
return nil, fmt.Errorf("user id is required")
}
if _, ok := seen[update.ID]; ok {
return nil, fmt.Errorf("duplicate user update")
}
seen[update.ID] = struct{}{}
update.Role = Role(strings.ToLower(strings.TrimSpace(string(update.Role))))
switch update.Role {
case RoleViewer, RoleEditor, RoleAdmin:
default:
return nil, fmt.Errorf("invalid role")
}
normalized = append(normalized, update)
}
if err := s.repo.UpdateUserAccess(ctx, normalized); err != nil {
return nil, err
}
users := make([]User, 0, len(normalized))
for _, update := range normalized {
user, err := s.repo.GetUserByID(ctx, update.ID)
if err != nil {
return nil, err
}
s.repo.Audit(ctx, actor.UserID, "auth.user.access.update", "user", user.ID, "", "", map[string]any{
"role": user.Role,
"disabled": user.Disabled,
})
users = append(users, user)
}
return users, nil
}
func (s *Service) SendPasswordReset(ctx context.Context, actor Principal, userID string) error {
if !Allows(actor, ScopeAdmin) {
return fmt.Errorf("admin role required")
}
if s.email == nil {
return fmt.Errorf("email sender is not configured")
}
user, err := s.repo.GetUserByID(ctx, userID)
if err != nil {
return err
}
secret, err := randomToken(32)
if err != nil {
return err
}
id, err := randomToken(18)
if err != nil {
return err
}
now := time.Now().UTC()
if err := s.repo.CreatePasswordResetToken(ctx, PasswordResetToken{
ID: "reset:" + id,
UserID: user.ID,
TokenHash: hashSecret(secret),
CreatedAt: now,
ExpiresAt: now.Add(passwordResetTTL),
InitiatedBy: actor.UserID,
}); err != nil {
return err
}
link := s.publicOrigin + "/password-reset?token=" + url.QueryEscape(secret)
body := fmt.Sprintf("A Cairnquire administrator requested a password reset for your account.\n\nSet a new password within one hour:\n%s\n\nIf you did not expect this message, contact your administrator.", link)
if err := s.email.Send(ctx, []string{user.Email}, "Reset your Cairnquire password", body); err != nil {
s.repo.ExpirePasswordResetToken(ctx, hashSecret(secret))
return err
}
s.repo.Audit(ctx, actor.UserID, "auth.password.reset.request", "user", user.ID, "", "", nil)
return nil
}
func (s *Service) ResetPassword(ctx context.Context, token, newPassword string) error {
if len(newPassword) < 12 {
return fmt.Errorf("password must be at least 12 characters")
}
if strings.TrimSpace(token) == "" {
return fmt.Errorf("password reset token is required")
}
passwordHash, err := hashPassword(newPassword)
if err != nil {
return err
}
userID, err := s.repo.CompletePasswordReset(ctx, hashSecret(token), passwordHash)
if err != nil {
return err
}
s.repo.Audit(ctx, userID, "auth.password.reset.complete", "user", userID, "", "", nil)
return nil
}
func (s *Service) CreateAPIKey(ctx context.Context, userID, name string, scopes []Scope, expiresAt *time.Time) (CreatedAPIKey, error) {
@@ -400,6 +511,9 @@ func (s *Service) ValidateBearerToken(ctx context.Context, token string) (Princi
if err != nil {
return Principal{}, err
}
if user.Disabled {
return Principal{}, fmt.Errorf("account disabled")
}
expiresAt := time.Time{}
if key.ExpiresAt != nil {
expiresAt = *key.ExpiresAt
@@ -491,6 +605,9 @@ func (s *Service) PollDeviceFlow(ctx context.Context, deviceCode string) (Create
}
func (s *Service) createSession(ctx context.Context, user User, ip, userAgent string) (Principal, string, error) {
if user.Disabled {
return Principal{}, "", fmt.Errorf("account disabled")
}
token, err := randomToken(32)
if err != nil {
return Principal{}, "", err

View File

@@ -3,6 +3,8 @@ package auth
import (
"context"
"database/sql"
"errors"
"net/url"
"strings"
"testing"
"time"
@@ -10,6 +12,20 @@ import (
"github.com/tim/cairnquire/apps/server/internal/database"
)
type testEmailSender struct {
to []string
subject string
body string
err error
}
func (s *testEmailSender) Send(ctx context.Context, to []string, subject, body string) error {
s.to = append([]string(nil), to...)
s.subject = subject
s.body = body
return s.err
}
func setupAuthTestService(t *testing.T) *Service {
t.Helper()
@@ -158,11 +174,131 @@ func TestCannotDemoteOrDeleteLastAdmin(t *testing.T) {
if _, err := service.UpdateUserRole(ctx, principal, user.ID, string(RoleEditor)); err == nil {
t.Fatal("expected demoting last admin to fail")
}
if _, err := service.UpdateUserAccess(ctx, principal, []UserAccessUpdate{{ID: user.ID, Role: RoleAdmin, Disabled: true}}); err == nil {
t.Fatal("expected disabling last admin to fail")
}
if err := service.DeleteAccount(ctx, principal, "correct horse battery staple"); err == nil {
t.Fatal("expected deleting last admin to fail")
}
}
func TestDisabledUserCannotAuthenticateWithPasswordSessionOrAPIKey(t *testing.T) {
service := setupAuthTestService(t)
ctx := context.Background()
admin := setupInitialAdmin(t, service, true)
user, err := service.RegisterPasswordUser(ctx, "viewer@example.com", "Viewer", "correct horse battery staple", "viewer")
if err != nil {
t.Fatalf("RegisterPasswordUser() error = %v", err)
}
_, sessionToken, err := service.LoginPassword(ctx, user.Email, "correct horse battery staple", "127.0.0.1", "test")
if err != nil {
t.Fatalf("LoginPassword() before disable error = %v", err)
}
apiKey, err := service.CreateAPIKey(ctx, user.ID, "CLI", []Scope{ScopeDocsRead}, nil)
if err != nil {
t.Fatalf("CreateAPIKey() error = %v", err)
}
adminPrincipal := principalFromUser(admin, "session", "sess:admin", "", nil, time.Now().Add(time.Hour))
if _, err := service.UpdateUserAccess(ctx, adminPrincipal, []UserAccessUpdate{{ID: user.ID, Role: RoleViewer, Disabled: true}}); err != nil {
t.Fatalf("UpdateUserAccess() error = %v", err)
}
if _, _, err := service.LoginPassword(ctx, user.Email, "correct horse battery staple", "127.0.0.1", "test"); err == nil {
t.Fatal("disabled user password login succeeded")
}
if _, err := service.ValidateSessionToken(ctx, sessionToken); err == nil {
t.Fatal("disabled user session remained valid")
}
if _, err := service.ValidateBearerToken(ctx, apiKey.Token); err == nil {
t.Fatal("disabled user api token remained valid")
}
}
func TestPasswordResetEmailReplacesPriorLinkAndRevokesSessions(t *testing.T) {
service := setupAuthTestService(t)
ctx := context.Background()
admin := setupInitialAdmin(t, service, true)
user, err := service.RegisterPasswordUser(ctx, "viewer@example.com", "Viewer", "correct horse battery staple", "viewer")
if err != nil {
t.Fatalf("RegisterPasswordUser() error = %v", err)
}
_, sessionToken, err := service.LoginPassword(ctx, user.Email, "correct horse battery staple", "127.0.0.1", "test")
if err != nil {
t.Fatalf("LoginPassword() before reset error = %v", err)
}
sender := &testEmailSender{}
service.SetEmailSender(sender)
adminPrincipal := principalFromUser(admin, "session", "sess:admin", "", nil, time.Now().Add(time.Hour))
if err := service.SendPasswordReset(ctx, adminPrincipal, user.ID); err != nil {
t.Fatalf("SendPasswordReset() first error = %v", err)
}
firstToken := passwordResetTokenFromBody(t, sender.body)
if err := service.SendPasswordReset(ctx, adminPrincipal, user.ID); err != nil {
t.Fatalf("SendPasswordReset() second error = %v", err)
}
secondToken := passwordResetTokenFromBody(t, sender.body)
if firstToken == secondToken {
t.Fatal("expected each password reset email to contain a fresh token")
}
if err := service.ResetPassword(ctx, firstToken, "new correct horse battery staple"); err == nil {
t.Fatal("first password reset link remained valid after requesting a new one")
}
if err := service.ResetPassword(ctx, secondToken, "new correct horse battery staple"); err != nil {
t.Fatalf("ResetPassword() error = %v", err)
}
if err := service.ResetPassword(ctx, secondToken, "another correct horse battery staple"); err == nil {
t.Fatal("password reset link was reusable")
}
if _, err := service.ValidateSessionToken(ctx, sessionToken); err == nil {
t.Fatal("password reset did not revoke existing sessions")
}
if _, _, err := service.LoginPassword(ctx, user.Email, "correct horse battery staple", "127.0.0.1", "test"); err == nil {
t.Fatal("old password still works after reset")
}
if _, _, err := service.LoginPassword(ctx, user.Email, "new correct horse battery staple", "127.0.0.1", "test"); err != nil {
t.Fatalf("new password login error = %v", err)
}
}
func TestPasswordResetEmailFailureInvalidatesLink(t *testing.T) {
service := setupAuthTestService(t)
ctx := context.Background()
admin := setupInitialAdmin(t, service, true)
user, err := service.RegisterPasswordUser(ctx, "viewer@example.com", "Viewer", "correct horse battery staple", "viewer")
if err != nil {
t.Fatalf("RegisterPasswordUser() error = %v", err)
}
sender := &testEmailSender{err: errors.New("smtp unavailable")}
service.SetEmailSender(sender)
adminPrincipal := principalFromUser(admin, "session", "sess:admin", "", nil, time.Now().Add(time.Hour))
if err := service.SendPasswordReset(ctx, adminPrincipal, user.ID); err == nil {
t.Fatal("expected password reset email delivery failure")
}
token := passwordResetTokenFromBody(t, sender.body)
if err := service.ResetPassword(ctx, token, "new correct horse battery staple"); err == nil {
t.Fatal("password reset link remained valid after email delivery failed")
}
}
func passwordResetTokenFromBody(t *testing.T, body string) string {
t.Helper()
for _, line := range strings.Split(body, "\n") {
if !strings.HasPrefix(line, "http") {
continue
}
parsed, err := url.Parse(line)
if err != nil {
t.Fatalf("parse password reset URL: %v", err)
}
if token := parsed.Query().Get("token"); token != "" {
return token
}
}
t.Fatalf("password reset body does not contain a link: %q", body)
return ""
}
func TestPublicRegistrationCannotAttachCredentialsToExistingUser(t *testing.T) {
service := setupAuthTestService(t)
ctx := context.Background()

View File

@@ -41,12 +41,29 @@ type User struct {
Email string
DisplayName string
Role Role
Disabled bool
PasswordHash string
CreatedAt time.Time
LastSeenAt time.Time
Credentials []webauthn.Credential
}
type UserAccessUpdate struct {
ID string `json:"id"`
Role Role `json:"role"`
Disabled bool `json:"disabled"`
}
type PasswordResetToken struct {
ID string
UserID string
TokenHash string
CreatedAt time.Time
ExpiresAt time.Time
UsedAt *time.Time
InitiatedBy string
}
func (u User) WebAuthnID() []byte {
return []byte(u.ID)
}