Files
Alphaeus Mote 6fdf9a9d54 feat(users): assign RBAC roles to users; OIDC default-role dropdown + clearer SSO config
Roles:
- Add GET /api/v1/roles (built-in SuperAdmin/Admin/Operator/Viewer) and
  wire role assignment into user create/update (UserRepository.SetRoles /
  GetRoleNames / ListRoles). User responses now include roles.
- User dialog gains a roles multi-select; the users list shows role chips.

OIDC / SSO config UI:
- Default role is now a dropdown populated from /roles.
- Issuer URL shows real provider examples (Entra/Okta/Google).
- Replace the confusing manual "Redirect URL" field with a read-only,
  auto-derived callback URL (from the browser origin / public URL) plus a
  copy button — the exact value to register at the IdP. The backend still
  auto-derives the callback and honors an env override.

Verified: /roles lists the four roles; creating/updating a user with roles
round-trips through GET.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-09-02 21:25:59 -04:00

391 lines
12 KiB
Go

// Package repository provides database access for all entities
package repository
import (
"database/sql"
"time"
"github.com/Grace-Solutions/OrchestrAD/internal/models"
"github.com/google/uuid"
)
// UserRepository handles user database operations
type UserRepository struct {
db *sql.DB
}
// NewUserRepository creates a new UserRepository
func NewUserRepository(db *sql.DB) *UserRepository {
return &UserRepository{db: db}
}
// ListRoles returns all defined roles, ordered by their privilege tier.
func (r *UserRepository) ListRoles() ([]models.Role, error) {
rows, err := r.db.Query(`
SELECT id, name, description, is_system_role, created_utc, updated_utc
FROM roles
ORDER BY CASE name
WHEN 'SuperAdmin' THEN 0 WHEN 'Admin' THEN 1
WHEN 'Operator' THEN 2 WHEN 'Viewer' THEN 3 ELSE 4 END, name`)
if err != nil {
return nil, err
}
defer rows.Close()
var roles []models.Role
for rows.Next() {
var role models.Role
var isSystem int
var createdStr, updatedStr string
if err := rows.Scan(&role.ID, &role.Name, &role.Description, &isSystem, &createdStr, &updatedStr); err != nil {
return nil, err
}
role.IsSystemRole = isSystem != 0
role.CreatedUTC = parseTimeOrZero(createdStr)
role.UpdatedUTC = parseTimeOrZero(updatedStr)
roles = append(roles, role)
}
return roles, rows.Err()
}
// GetRoleNames returns the names of the roles assigned to a user.
func (r *UserRepository) GetRoleNames(userID string) ([]string, error) {
rows, err := r.db.Query(`
SELECT ro.name FROM user_roles ur
JOIN roles ro ON ro.id = ur.role_id
WHERE ur.user_id = ?
ORDER BY ro.name`, userID)
if err != nil {
return nil, err
}
defer rows.Close()
var names []string
for rows.Next() {
var n string
if err := rows.Scan(&n); err != nil {
return nil, err
}
names = append(names, n)
}
return names, rows.Err()
}
// SetRoles replaces a user's role assignments with the named roles (unknown
// names are ignored). Runs in a transaction so the set change is atomic.
func (r *UserRepository) SetRoles(userID string, roleNames []string) error {
tx, err := r.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
if _, err := tx.Exec(`DELETE FROM user_roles WHERE user_id = ?`, userID); err != nil {
return err
}
now := time.Now().UTC().Format(time.RFC3339)
for _, name := range roleNames {
var roleID string
err := tx.QueryRow(`SELECT id FROM roles WHERE name = ?`, name).Scan(&roleID)
if err == sql.ErrNoRows {
continue
}
if err != nil {
return err
}
if _, err := tx.Exec(`INSERT OR IGNORE INTO user_roles (user_id, role_id, created_utc) VALUES (?, ?, ?)`,
userID, roleID, now); err != nil {
return err
}
}
return tx.Commit()
}
// Create creates a new user
func (r *UserRepository) Create(user *models.User) error {
if user.ID == "" {
user.ID = uuid.New().String()
}
now := time.Now().UTC()
user.CreatedUTC = now
user.UpdatedUTC = now
_, err := r.db.Exec(`
INSERT INTO users (
id, username, email, password_hash, display_name,
is_active, is_oidc_user, oidc_provider_id, oidc_subject,
password_reset_required, created_utc, updated_utc
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`,
user.ID, user.Username, user.Email, user.PasswordHash, user.DisplayName,
boolToInt(user.IsActive), boolToInt(user.IsOIDCUser),
user.OIDCProviderID, user.OIDCSubject,
boolToInt(user.PasswordResetRequired),
formatTime(user.CreatedUTC), formatTime(user.UpdatedUTC),
)
return err
}
// GetByID retrieves a user by ID
func (r *UserRepository) GetByID(id string) (*models.User, error) {
user := &models.User{}
var isActive, isOIDCUser, passwordResetRequired int
var lastLogin, deleted sql.NullString
var createdStr, updatedStr string
err := r.db.QueryRow(`
SELECT id, username, email, password_hash, display_name,
is_active, is_oidc_user, oidc_provider_id, oidc_subject,
password_reset_required,
last_login_utc, created_utc, updated_utc, deleted_utc
FROM users WHERE id = ? AND deleted_utc IS NULL
`, id).Scan(
&user.ID, &user.Username, &user.Email, &user.PasswordHash, &user.DisplayName,
&isActive, &isOIDCUser, &user.OIDCProviderID, &user.OIDCSubject,
&passwordResetRequired,
&lastLogin, &createdStr, &updatedStr, &deleted,
)
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
user.IsActive = intToBool(isActive)
user.IsOIDCUser = intToBool(isOIDCUser)
user.PasswordResetRequired = intToBool(passwordResetRequired)
user.LastLoginUTC = parseNullTime(lastLogin)
user.CreatedUTC = parseTimeOrZero(createdStr)
user.UpdatedUTC = parseTimeOrZero(updatedStr)
user.DeletedUTC = parseNullTime(deleted)
return user, nil
}
// GetByUsername retrieves a user by username
func (r *UserRepository) GetByUsername(username string) (*models.User, error) {
user := &models.User{}
var isActive, isOIDCUser, passwordResetRequired int
var lastLogin, deleted sql.NullString
var createdStr, updatedStr string
err := r.db.QueryRow(`
SELECT id, username, email, password_hash, display_name,
is_active, is_oidc_user, oidc_provider_id, oidc_subject,
password_reset_required,
last_login_utc, created_utc, updated_utc, deleted_utc
FROM users WHERE username = ? AND deleted_utc IS NULL
`, username).Scan(
&user.ID, &user.Username, &user.Email, &user.PasswordHash, &user.DisplayName,
&isActive, &isOIDCUser, &user.OIDCProviderID, &user.OIDCSubject,
&passwordResetRequired,
&lastLogin, &createdStr, &updatedStr, &deleted,
)
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
user.IsActive = intToBool(isActive)
user.IsOIDCUser = intToBool(isOIDCUser)
user.PasswordResetRequired = intToBool(passwordResetRequired)
user.LastLoginUTC = parseNullTime(lastLogin)
user.CreatedUTC = parseTimeOrZero(createdStr)
user.UpdatedUTC = parseTimeOrZero(updatedStr)
user.DeletedUTC = parseNullTime(deleted)
return user, nil
}
// GetByOIDCSubject retrieves a non-deleted user by their OIDC provider id and
// subject, or nil when none matches. This is the stable identity link for
// federated (SSO) users.
func (r *UserRepository) GetByOIDCSubject(providerID, subject string) (*models.User, error) {
user := &models.User{}
var isActive, isOIDCUser, passwordResetRequired int
var lastLogin, deleted sql.NullString
var createdStr, updatedStr string
err := r.db.QueryRow(`
SELECT id, username, email, password_hash, display_name,
is_active, is_oidc_user, oidc_provider_id, oidc_subject,
password_reset_required,
last_login_utc, created_utc, updated_utc, deleted_utc
FROM users WHERE oidc_provider_id = ? AND oidc_subject = ? AND deleted_utc IS NULL
`, providerID, subject).Scan(
&user.ID, &user.Username, &user.Email, &user.PasswordHash, &user.DisplayName,
&isActive, &isOIDCUser, &user.OIDCProviderID, &user.OIDCSubject,
&passwordResetRequired,
&lastLogin, &createdStr, &updatedStr, &deleted,
)
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
user.IsActive = intToBool(isActive)
user.IsOIDCUser = intToBool(isOIDCUser)
user.PasswordResetRequired = intToBool(passwordResetRequired)
user.LastLoginUTC = parseNullTime(lastLogin)
user.CreatedUTC = parseTimeOrZero(createdStr)
user.UpdatedUTC = parseTimeOrZero(updatedStr)
user.DeletedUTC = parseNullTime(deleted)
return user, nil
}
// List returns a page of non-deleted users ordered by username, along with
// the total count of matching rows for pagination metadata.
func (r *UserRepository) List(offset, limit int) ([]models.User, int, error) {
var total int
if err := r.db.QueryRow(`SELECT COUNT(*) FROM users WHERE deleted_utc IS NULL`).Scan(&total); err != nil {
return nil, 0, err
}
rows, err := r.db.Query(`
SELECT id, username, email, password_hash, display_name,
is_active, is_oidc_user, oidc_provider_id, oidc_subject,
password_reset_required,
last_login_utc, created_utc, updated_utc, deleted_utc
FROM users WHERE deleted_utc IS NULL
ORDER BY username ASC
LIMIT ? OFFSET ?
`, limit, offset)
if err != nil {
return nil, 0, err
}
defer rows.Close()
users := make([]models.User, 0, limit)
for rows.Next() {
var u models.User
var isActive, isOIDCUser, passwordResetRequired int
var lastLogin, deleted sql.NullString
var createdStr, updatedStr string
if err := rows.Scan(
&u.ID, &u.Username, &u.Email, &u.PasswordHash, &u.DisplayName,
&isActive, &isOIDCUser, &u.OIDCProviderID, &u.OIDCSubject,
&passwordResetRequired,
&lastLogin, &createdStr, &updatedStr, &deleted,
); err != nil {
return nil, 0, err
}
u.IsActive = intToBool(isActive)
u.IsOIDCUser = intToBool(isOIDCUser)
u.PasswordResetRequired = intToBool(passwordResetRequired)
u.LastLoginUTC = parseNullTime(lastLogin)
u.CreatedUTC = parseTimeOrZero(createdStr)
u.UpdatedUTC = parseTimeOrZero(updatedStr)
u.DeletedUTC = parseNullTime(deleted)
users = append(users, u)
}
if err := rows.Err(); err != nil {
return nil, 0, err
}
return users, total, nil
}
// Update updates an existing user
func (r *UserRepository) Update(user *models.User) error {
user.UpdatedUTC = time.Now().UTC()
_, err := r.db.Exec(`
UPDATE users SET
username = ?, email = ?, password_hash = ?, display_name = ?,
is_active = ?, password_reset_required = ?, updated_utc = ?
WHERE id = ?
`,
user.Username, user.Email, user.PasswordHash, user.DisplayName,
boolToInt(user.IsActive), boolToInt(user.PasswordResetRequired),
formatTime(user.UpdatedUTC), user.ID,
)
return err
}
// UpdateLastLogin updates the last login timestamp
func (r *UserRepository) UpdateLastLogin(id string) error {
now := time.Now().UTC()
_, err := r.db.Exec(`
UPDATE users SET last_login_utc = ?, updated_utc = ?
WHERE id = ?
`, formatTime(now), formatTime(now), id)
return err
}
// SetPasswordResetRequired toggles the password_reset_required flag for a
// user. Used by the bootstrap path to mark the default admin as needing a
// password change on first login, and by the change-password handler to
// clear the flag once the user replaces their credentials.
func (r *UserRepository) SetPasswordResetRequired(id string, required bool) error {
now := time.Now().UTC()
_, err := r.db.Exec(`
UPDATE users SET password_reset_required = ?, updated_utc = ?
WHERE id = ?
`, boolToInt(required), formatTime(now), id)
return err
}
// UpdatePassword replaces a user's password hash and clears the reset
// flag in a single statement so the two fields cannot drift.
func (r *UserRepository) UpdatePassword(id string, passwordHash string) error {
now := time.Now().UTC()
_, err := r.db.Exec(`
UPDATE users SET password_hash = ?, password_reset_required = 0,
updated_utc = ?
WHERE id = ?
`, passwordHash, formatTime(now), id)
return err
}
// SoftDelete marks a user as deleted
func (r *UserRepository) SoftDelete(id string) error {
now := time.Now().UTC()
_, err := r.db.Exec(`
UPDATE users SET deleted_utc = ?, updated_utc = ?
WHERE id = ?
`, formatTime(now), formatTime(now), id)
return err
}
// Helper functions
func boolToInt(b bool) int {
if b {
return 1
}
return 0
}
func intToBool(i int) bool {
return i != 0
}
func formatTime(t time.Time) string {
return t.UTC().Format(time.RFC3339)
}
func parseNullTime(ns sql.NullString) *time.Time {
if !ns.Valid {
return nil
}
t, err := time.Parse(time.RFC3339, ns.String)
if err != nil {
return nil
}
return &t
}
// parseTimeOrZero parses a non-null RFC3339 timestamp, returning the zero
// value if the string is empty or unparseable. Used for NOT NULL TEXT
// timestamp columns that the sqlite3 driver can't auto-convert.
func parseTimeOrZero(s string) time.Time {
if s == "" {
return time.Time{}
}
t, err := time.Parse(time.RFC3339, s)
if err != nil {
return time.Time{}
}
return t
}