d3d59b4e45
Backend: - ListStorageJobs, ListSambaShares, ListNFSExports, ListUsers, ListDirty, ListApplyLog, ListWatchedMounts: initialize with make([]T, 0) instead of var x []T to avoid JSON null on empty. Fixes "Cannot read properties of null" crash on /storage. Frontend: - Storage.tsx: defensive setJobs(res.jobs ?? []) to guard against API returning null. Tests: - Add TestListStorageJobsEmptyReturnsSlice. Version: 0.7.1
311 lines
7.8 KiB
Go
311 lines
7.8 KiB
Go
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()
|
|
|
|
users := make([]User, 0)
|
|
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) ReplaceUsers(users []User) error {
|
|
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`); err != nil {
|
|
return fmt.Errorf("clear system_users: %w", err)
|
|
}
|
|
|
|
for _, u := range users {
|
|
groups, err := encodeJSONStrings(u.Groups)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
_, err = tx.Exec(`
|
|
INSERT INTO system_users (username, groups, smb_enabled, disabled, pending_password)
|
|
VALUES (?, ?, ?, ?, '')`,
|
|
u.Username, groups, boolToInt(u.SMBEnabled), boolToInt(u.Disabled),
|
|
)
|
|
if err != nil {
|
|
return fmt.Errorf("insert user %s: %w", u.Username, err)
|
|
}
|
|
}
|
|
|
|
return tx.Commit()
|
|
}
|
|
|
|
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) UpdateAdminPasswordHash(id int64, hash string) error {
|
|
result, err := d.conn.Exec(
|
|
`UPDATE admins SET password_hash = ?, updated_at = datetime('now') WHERE id = ?`,
|
|
hash, id,
|
|
)
|
|
if err != nil {
|
|
return fmt.Errorf("update admin password: %w", err)
|
|
}
|
|
rows, err := result.RowsAffected()
|
|
if err != nil {
|
|
return fmt.Errorf("rows affected: %w", err)
|
|
}
|
|
if rows == 0 {
|
|
return fmt.Errorf("admin not found")
|
|
}
|
|
return nil
|
|
}
|
|
|
|
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
|
|
}
|