package db import ( "database/sql" "fmt" ) func scanUser(row interface { Scan(dest ...any) error }) (User, error) { var user User var uid, gid sql.NullInt64 var groups, createdAt, updatedAt string var smbEnabled, disabled int if err := row.Scan( &user.ID, &user.Username, &uid, &gid, &groups, &smbEnabled, &disabled, &user.PendingPassword, &createdAt, &updatedAt, ); err != nil { return User{}, err } if uid.Valid { user.UID = &uid.Int64 } if gid.Valid { user.GID = &gid.Int64 } user.SMBEnabled = smbEnabled != 0 user.Disabled = disabled != 0 var err error user.Groups, err = decodeJSONStrings(groups) if err != nil { return User{}, err } user.CreatedAt = parseTime(createdAt) user.UpdatedAt = parseTime(updatedAt) return user, nil } const userColumns = `id, username, uid, gid, groups, smb_enabled, disabled, pending_password, created_at, updated_at` func (d *DB) ListUsers() ([]User, error) { rows, err := d.conn.Query(`SELECT ` + userColumns + ` FROM system_users ORDER BY username ASC`) if err != nil { return nil, fmt.Errorf("list users: %w", err) } defer rows.Close() var users []User for rows.Next() { user, err := scanUser(rows) if err != nil { return nil, fmt.Errorf("scan user: %w", err) } users = append(users, user) } return users, rows.Err() } func (d *DB) GetUser(id int64) (User, error) { row := d.conn.QueryRow(`SELECT `+userColumns+` FROM system_users WHERE id = ?`, id) user, err := scanUser(row) if err == sql.ErrNoRows { return User{}, fmt.Errorf("user not found") } if err != nil { return User{}, fmt.Errorf("get user: %w", err) } return user, nil } func boolToInt(b bool) int { if b { return 1 } return 0 } func (d *DB) CreateUser(user User) (User, error) { groups, err := encodeJSONStrings(user.Groups) if err != nil { return User{}, err } result, err := d.conn.Exec(` INSERT INTO system_users (username, groups, smb_enabled, disabled, pending_password) VALUES (?, ?, ?, ?, ?)`, user.Username, groups, boolToInt(user.SMBEnabled), boolToInt(user.Disabled), user.PendingPassword, ) if err != nil { return User{}, fmt.Errorf("create user: %w", err) } id, err := result.LastInsertId() if err != nil { return User{}, fmt.Errorf("last insert id: %w", err) } return d.GetUser(id) } func (d *DB) UpdateUser(id int64, user User) (User, error) { groups, err := encodeJSONStrings(user.Groups) if err != nil { return User{}, err } result, err := d.conn.Exec(` UPDATE system_users SET groups = ?, smb_enabled = ?, disabled = ?, pending_password = ?, updated_at = datetime('now') WHERE id = ?`, groups, boolToInt(user.SMBEnabled), boolToInt(user.Disabled), user.PendingPassword, id, ) if err != nil { return User{}, fmt.Errorf("update user: %w", err) } rows, err := result.RowsAffected() if err != nil { return User{}, err } if rows == 0 { return User{}, fmt.Errorf("user not found") } return d.GetUser(id) } func (d *DB) SetUserPendingPassword(id int64, password string) error { _, err := d.conn.Exec( `UPDATE system_users SET pending_password = ?, updated_at = datetime('now') WHERE id = ?`, password, id, ) if err != nil { return fmt.Errorf("set pending password: %w", err) } return nil } // ClearUserPendingPassword removes the stored password once it has been applied. func (d *DB) ClearUserPendingPassword(id int64) error { _, err := d.conn.Exec( `UPDATE system_users SET pending_password = '' WHERE id = ?`, id, ) if err != nil { return fmt.Errorf("clear pending password: %w", err) } return nil } // DeleteUser removes the user row and records a tombstone so the users module // can run userdel on the next apply. func (d *DB) DeleteUser(id int64) error { user, err := d.GetUser(id) if err != nil { return err } tx, err := d.conn.Begin() if err != nil { return fmt.Errorf("begin tx: %w", err) } defer func() { _ = tx.Rollback() }() if _, err := tx.Exec(`DELETE FROM system_users WHERE id = ?`, id); err != nil { return fmt.Errorf("delete user: %w", err) } if _, err := tx.Exec( `INSERT INTO deleted_system_users (username, marked_at) VALUES (?, datetime('now')) ON CONFLICT(username) DO UPDATE SET marked_at = datetime('now')`, user.Username, ); err != nil { return fmt.Errorf("record deleted user: %w", err) } return tx.Commit() } func (d *DB) ListDeletedUsers() ([]string, error) { rows, err := d.conn.Query(`SELECT username FROM deleted_system_users ORDER BY username ASC`) if err != nil { return nil, fmt.Errorf("list deleted users: %w", err) } defer rows.Close() var usernames []string for rows.Next() { var username string if err := rows.Scan(&username); err != nil { return nil, err } usernames = append(usernames, username) } return usernames, rows.Err() } func (d *DB) ClearDeletedUser(username string) error { _, err := d.conn.Exec(`DELETE FROM deleted_system_users WHERE username = ?`, username) if err != nil { return fmt.Errorf("clear deleted user: %w", err) } return nil } func (d *DB) GetAdminByUsername(username string) (Admin, error) { row := d.conn.QueryRow( `SELECT id, username, password_hash, created_at, updated_at FROM admins WHERE username = ?`, username, ) var admin Admin var createdAt, updatedAt string err := row.Scan(&admin.ID, &admin.Username, &admin.PasswordHash, &createdAt, &updatedAt) if err == sql.ErrNoRows { return Admin{}, fmt.Errorf("admin not found") } if err != nil { return Admin{}, fmt.Errorf("get admin: %w", err) } admin.CreatedAt = parseTime(createdAt) admin.UpdatedAt = parseTime(updatedAt) return admin, nil } func (d *DB) CountAdmins() (int, error) { var count int if err := d.conn.QueryRow(`SELECT COUNT(1) FROM admins`).Scan(&count); err != nil { return 0, fmt.Errorf("count admins: %w", err) } return count, nil } func (d *DB) CreateAdmin(username, passwordHash string) (Admin, error) { if _, err := d.conn.Exec( `INSERT INTO admins (username, password_hash) VALUES (?, ?)`, username, passwordHash, ); err != nil { return Admin{}, fmt.Errorf("create admin: %w", err) } return d.GetAdminByUsername(username) } func (d *DB) GetSetting(key string) (string, bool, error) { var value string err := d.conn.QueryRow(`SELECT value FROM settings WHERE key = ?`, key).Scan(&value) if err == sql.ErrNoRows { return "", false, nil } if err != nil { return "", false, fmt.Errorf("get setting %q: %w", key, err) } return value, true, nil } func (d *DB) SetSetting(key, value string) error { _, err := d.conn.Exec( `INSERT INTO settings (key, value) VALUES (?, ?) ON CONFLICT(key) DO UPDATE SET value = excluded.value`, key, value, ) if err != nil { return fmt.Errorf("set setting %q: %w", key, err) } return nil }