Add nasctl: Go NAS control plane with React frontend
This commit is contained in:
@@ -0,0 +1,62 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"embed"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
//go:embed migrations/*.sql
|
||||
var migrationsFS embed.FS
|
||||
|
||||
type DB struct {
|
||||
conn *sql.DB
|
||||
}
|
||||
|
||||
func Open(path string) (*DB, error) {
|
||||
conn, err := sql.Open("sqlite", path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open sqlite: %w", err)
|
||||
}
|
||||
if _, err := conn.Exec(`PRAGMA foreign_keys = ON`); err != nil {
|
||||
_ = conn.Close()
|
||||
return nil, fmt.Errorf("enable foreign keys: %w", err)
|
||||
}
|
||||
return &DB{conn: conn}, nil
|
||||
}
|
||||
|
||||
func (d *DB) Close() error {
|
||||
return d.conn.Close()
|
||||
}
|
||||
|
||||
func (d *DB) Conn() *sql.DB {
|
||||
return d.conn
|
||||
}
|
||||
|
||||
func (d *DB) Migrate() error {
|
||||
entries, err := fs.ReadDir(migrationsFS, "migrations")
|
||||
if err != nil {
|
||||
return fmt.Errorf("read migrations: %w", err)
|
||||
}
|
||||
sort.Slice(entries, func(i, j int) bool {
|
||||
return entries[i].Name() < entries[j].Name()
|
||||
})
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() || !strings.HasSuffix(entry.Name(), ".sql") {
|
||||
continue
|
||||
}
|
||||
content, err := migrationsFS.ReadFile("migrations/" + entry.Name())
|
||||
if err != nil {
|
||||
return fmt.Errorf("read migration %s: %w", entry.Name(), err)
|
||||
}
|
||||
if _, err := d.conn.Exec(string(content)); err != nil {
|
||||
return fmt.Errorf("apply migration %s: %w", entry.Name(), err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
CREATE TABLE IF NOT EXISTS system_users (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
username TEXT NOT NULL UNIQUE,
|
||||
uid INTEGER,
|
||||
gid INTEGER,
|
||||
groups TEXT NOT NULL DEFAULT '[]',
|
||||
smb_enabled INTEGER NOT NULL DEFAULT 0,
|
||||
disabled INTEGER NOT NULL DEFAULT 0,
|
||||
pending_password TEXT NOT NULL DEFAULT '',
|
||||
created_at TEXT NOT NULL DEFAULT (datetime('now')),
|
||||
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS admins (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
username TEXT NOT NULL UNIQUE,
|
||||
password_hash TEXT NOT NULL,
|
||||
created_at TEXT NOT NULL DEFAULT (datetime('now')),
|
||||
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS settings (
|
||||
key TEXT PRIMARY KEY,
|
||||
value TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS samba_shares (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name TEXT NOT NULL UNIQUE,
|
||||
path TEXT NOT NULL,
|
||||
comment TEXT NOT NULL DEFAULT '',
|
||||
read_only INTEGER NOT NULL DEFAULT 0,
|
||||
guest_ok INTEGER NOT NULL DEFAULT 0,
|
||||
valid_users TEXT NOT NULL DEFAULT '[]',
|
||||
valid_groups TEXT NOT NULL DEFAULT '[]',
|
||||
created_at TEXT NOT NULL DEFAULT (datetime('now')),
|
||||
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS nfs_exports (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
path TEXT NOT NULL,
|
||||
clients TEXT NOT NULL DEFAULT '[]',
|
||||
options TEXT NOT NULL DEFAULT 'rw,sync,no_root_squash',
|
||||
created_at TEXT NOT NULL DEFAULT (datetime('now')),
|
||||
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS deleted_system_users (
|
||||
username TEXT PRIMARY KEY,
|
||||
marked_at TEXT NOT NULL DEFAULT (datetime('now'))
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS dirty_modules (
|
||||
module TEXT PRIMARY KEY,
|
||||
marked_at TEXT NOT NULL DEFAULT (datetime('now'))
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS apply_log (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
module TEXT NOT NULL,
|
||||
message TEXT NOT NULL,
|
||||
success INTEGER NOT NULL,
|
||||
created_at TEXT NOT NULL DEFAULT (datetime('now'))
|
||||
);
|
||||
@@ -0,0 +1,60 @@
|
||||
package db
|
||||
|
||||
import "time"
|
||||
|
||||
type User struct {
|
||||
ID int64 `json:"id"`
|
||||
Username string `json:"username"`
|
||||
UID *int64 `json:"uid,omitempty"`
|
||||
GID *int64 `json:"gid,omitempty"`
|
||||
Groups []string `json:"groups"`
|
||||
SMBEnabled bool `json:"smb_enabled"`
|
||||
Disabled bool `json:"disabled"`
|
||||
// PendingPassword holds a not-yet-applied password. Never serialized.
|
||||
PendingPassword string `json:"-"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
type Admin struct {
|
||||
ID int64 `json:"id"`
|
||||
Username string `json:"username"`
|
||||
PasswordHash string `json:"-"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
type SambaShare struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Path string `json:"path"`
|
||||
Comment string `json:"comment"`
|
||||
ReadOnly bool `json:"read_only"`
|
||||
GuestOK bool `json:"guest_ok"`
|
||||
ValidUsers []string `json:"valid_users"`
|
||||
ValidGroups []string `json:"valid_groups"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
type NFSExport struct {
|
||||
ID int64 `json:"id"`
|
||||
Path string `json:"path"`
|
||||
Clients []string `json:"clients"`
|
||||
Options string `json:"options"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
type DirtyModule struct {
|
||||
Module string `json:"module"`
|
||||
MarkedAt time.Time `json:"marked_at"`
|
||||
}
|
||||
|
||||
type ApplyLogEntry struct {
|
||||
ID int64 `json:"id"`
|
||||
Module string `json:"module"`
|
||||
Message string `json:"message"`
|
||||
Success bool `json:"success"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
@@ -0,0 +1,286 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
func parseTime(value string) time.Time {
|
||||
t, err := time.Parse(time.RFC3339, value)
|
||||
if err != nil {
|
||||
t, _ = time.Parse("2006-01-02 15:04:05", value)
|
||||
}
|
||||
return t
|
||||
}
|
||||
|
||||
func encodeJSONStrings(values []string) (string, error) {
|
||||
if values == nil {
|
||||
values = []string{}
|
||||
}
|
||||
data, err := json.Marshal(values)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(data), nil
|
||||
}
|
||||
|
||||
func decodeJSONStrings(raw string) ([]string, error) {
|
||||
if raw == "" {
|
||||
return []string{}, nil
|
||||
}
|
||||
var values []string
|
||||
if err := json.Unmarshal([]byte(raw), &values); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return values, nil
|
||||
}
|
||||
|
||||
func scanSambaShare(row interface {
|
||||
Scan(dest ...any) error
|
||||
}) (SambaShare, error) {
|
||||
var share SambaShare
|
||||
var readOnly, guestOK int
|
||||
var validUsers, validGroups, createdAt, updatedAt string
|
||||
if err := row.Scan(
|
||||
&share.ID,
|
||||
&share.Name,
|
||||
&share.Path,
|
||||
&share.Comment,
|
||||
&readOnly,
|
||||
&guestOK,
|
||||
&validUsers,
|
||||
&validGroups,
|
||||
&createdAt,
|
||||
&updatedAt,
|
||||
); err != nil {
|
||||
return SambaShare{}, err
|
||||
}
|
||||
share.ReadOnly = readOnly != 0
|
||||
share.GuestOK = guestOK != 0
|
||||
var err error
|
||||
share.ValidUsers, err = decodeJSONStrings(validUsers)
|
||||
if err != nil {
|
||||
return SambaShare{}, err
|
||||
}
|
||||
share.ValidGroups, err = decodeJSONStrings(validGroups)
|
||||
if err != nil {
|
||||
return SambaShare{}, err
|
||||
}
|
||||
share.CreatedAt = parseTime(createdAt)
|
||||
share.UpdatedAt = parseTime(updatedAt)
|
||||
return share, nil
|
||||
}
|
||||
|
||||
func (d *DB) MarkDirty(module string) error {
|
||||
_, err := d.conn.Exec(
|
||||
`INSERT INTO dirty_modules (module, marked_at) VALUES (?, datetime('now'))
|
||||
ON CONFLICT(module) DO UPDATE SET marked_at = datetime('now')`,
|
||||
module,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("mark dirty %q: %w", module, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (d *DB) ClearDirty(module string) error {
|
||||
_, err := d.conn.Exec(`DELETE FROM dirty_modules WHERE module = ?`, module)
|
||||
if err != nil {
|
||||
return fmt.Errorf("clear dirty %q: %w", module, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (d *DB) IsDirty(module string) (bool, error) {
|
||||
var count int
|
||||
err := d.conn.QueryRow(`SELECT COUNT(1) FROM dirty_modules WHERE module = ?`, module).Scan(&count)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("is dirty %q: %w", module, err)
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func (d *DB) ListDirty() ([]DirtyModule, error) {
|
||||
rows, err := d.conn.Query(`SELECT module, marked_at FROM dirty_modules ORDER BY marked_at ASC`)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list dirty: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var modules []DirtyModule
|
||||
for rows.Next() {
|
||||
var module DirtyModule
|
||||
var markedAt string
|
||||
if err := rows.Scan(&module.Module, &markedAt); err != nil {
|
||||
return nil, fmt.Errorf("scan dirty: %w", err)
|
||||
}
|
||||
module.MarkedAt = parseTime(markedAt)
|
||||
modules = append(modules, module)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return modules, nil
|
||||
}
|
||||
|
||||
func (d *DB) AppendApplyLog(module, message string, success bool) error {
|
||||
successInt := 0
|
||||
if success {
|
||||
successInt = 1
|
||||
}
|
||||
_, err := d.conn.Exec(
|
||||
`INSERT INTO apply_log (module, message, success, created_at) VALUES (?, ?, ?, datetime('now'))`,
|
||||
module, message, successInt,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("append apply log: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (d *DB) ListApplyLog(limit int) ([]ApplyLogEntry, error) {
|
||||
if limit <= 0 {
|
||||
limit = 100
|
||||
}
|
||||
rows, err := d.conn.Query(
|
||||
`SELECT id, module, message, success, created_at FROM apply_log ORDER BY id DESC LIMIT ?`,
|
||||
limit,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list apply log: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var entries []ApplyLogEntry
|
||||
for rows.Next() {
|
||||
var entry ApplyLogEntry
|
||||
var success int
|
||||
var createdAt string
|
||||
if err := rows.Scan(&entry.ID, &entry.Module, &entry.Message, &success, &createdAt); err != nil {
|
||||
return nil, fmt.Errorf("scan apply log: %w", err)
|
||||
}
|
||||
entry.Success = success != 0
|
||||
entry.CreatedAt = parseTime(createdAt)
|
||||
entries = append(entries, entry)
|
||||
}
|
||||
return entries, rows.Err()
|
||||
}
|
||||
|
||||
func (d *DB) ListSambaShares() ([]SambaShare, error) {
|
||||
rows, err := d.conn.Query(`
|
||||
SELECT id, name, path, comment, read_only, guest_ok, valid_users, valid_groups, created_at, updated_at
|
||||
FROM samba_shares ORDER BY name ASC`)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list samba shares: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var shares []SambaShare
|
||||
for rows.Next() {
|
||||
share, err := scanSambaShare(rows)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("scan samba share: %w", err)
|
||||
}
|
||||
shares = append(shares, share)
|
||||
}
|
||||
return shares, rows.Err()
|
||||
}
|
||||
|
||||
func (d *DB) GetSambaShare(id int64) (SambaShare, error) {
|
||||
row := d.conn.QueryRow(`
|
||||
SELECT id, name, path, comment, read_only, guest_ok, valid_users, valid_groups, created_at, updated_at
|
||||
FROM samba_shares WHERE id = ?`, id)
|
||||
share, err := scanSambaShare(row)
|
||||
if err == sql.ErrNoRows {
|
||||
return SambaShare{}, fmt.Errorf("samba share not found")
|
||||
}
|
||||
if err != nil {
|
||||
return SambaShare{}, fmt.Errorf("get samba share: %w", err)
|
||||
}
|
||||
return share, nil
|
||||
}
|
||||
|
||||
func (d *DB) CreateSambaShare(share SambaShare) (SambaShare, error) {
|
||||
validUsers, err := encodeJSONStrings(share.ValidUsers)
|
||||
if err != nil {
|
||||
return SambaShare{}, err
|
||||
}
|
||||
validGroups, err := encodeJSONStrings(share.ValidGroups)
|
||||
if err != nil {
|
||||
return SambaShare{}, err
|
||||
}
|
||||
readOnly := 0
|
||||
if share.ReadOnly {
|
||||
readOnly = 1
|
||||
}
|
||||
guestOK := 0
|
||||
if share.GuestOK {
|
||||
guestOK = 1
|
||||
}
|
||||
result, err := d.conn.Exec(`
|
||||
INSERT INTO samba_shares (name, path, comment, read_only, guest_ok, valid_users, valid_groups)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?)`,
|
||||
share.Name, share.Path, share.Comment, readOnly, guestOK, validUsers, validGroups,
|
||||
)
|
||||
if err != nil {
|
||||
return SambaShare{}, fmt.Errorf("create samba share: %w", err)
|
||||
}
|
||||
id, err := result.LastInsertId()
|
||||
if err != nil {
|
||||
return SambaShare{}, fmt.Errorf("last insert id: %w", err)
|
||||
}
|
||||
return d.GetSambaShare(id)
|
||||
}
|
||||
|
||||
func (d *DB) UpdateSambaShare(id int64, share SambaShare) (SambaShare, error) {
|
||||
validUsers, err := encodeJSONStrings(share.ValidUsers)
|
||||
if err != nil {
|
||||
return SambaShare{}, err
|
||||
}
|
||||
validGroups, err := encodeJSONStrings(share.ValidGroups)
|
||||
if err != nil {
|
||||
return SambaShare{}, err
|
||||
}
|
||||
readOnly := 0
|
||||
if share.ReadOnly {
|
||||
readOnly = 1
|
||||
}
|
||||
guestOK := 0
|
||||
if share.GuestOK {
|
||||
guestOK = 1
|
||||
}
|
||||
result, err := d.conn.Exec(`
|
||||
UPDATE samba_shares
|
||||
SET name = ?, path = ?, comment = ?, read_only = ?, guest_ok = ?, valid_users = ?, valid_groups = ?, updated_at = datetime('now')
|
||||
WHERE id = ?`,
|
||||
share.Name, share.Path, share.Comment, readOnly, guestOK, validUsers, validGroups, id,
|
||||
)
|
||||
if err != nil {
|
||||
return SambaShare{}, fmt.Errorf("update samba share: %w", err)
|
||||
}
|
||||
rows, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return SambaShare{}, err
|
||||
}
|
||||
if rows == 0 {
|
||||
return SambaShare{}, fmt.Errorf("samba share not found")
|
||||
}
|
||||
return d.GetSambaShare(id)
|
||||
}
|
||||
|
||||
func (d *DB) DeleteSambaShare(id int64) error {
|
||||
result, err := d.conn.Exec(`DELETE FROM samba_shares WHERE id = ?`, id)
|
||||
if err != nil {
|
||||
return fmt.Errorf("delete samba share: %w", err)
|
||||
}
|
||||
rows, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if rows == 0 {
|
||||
return fmt.Errorf("samba share not found")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
func scanNFSExport(row interface {
|
||||
Scan(dest ...any) error
|
||||
}) (NFSExport, error) {
|
||||
var export NFSExport
|
||||
var clients, createdAt, updatedAt string
|
||||
if err := row.Scan(
|
||||
&export.ID,
|
||||
&export.Path,
|
||||
&clients,
|
||||
&export.Options,
|
||||
&createdAt,
|
||||
&updatedAt,
|
||||
); err != nil {
|
||||
return NFSExport{}, err
|
||||
}
|
||||
var err error
|
||||
export.Clients, err = decodeJSONStrings(clients)
|
||||
if err != nil {
|
||||
return NFSExport{}, err
|
||||
}
|
||||
export.CreatedAt = parseTime(createdAt)
|
||||
export.UpdatedAt = parseTime(updatedAt)
|
||||
return export, nil
|
||||
}
|
||||
|
||||
func (d *DB) ListNFSExports() ([]NFSExport, error) {
|
||||
rows, err := d.conn.Query(`
|
||||
SELECT id, path, clients, options, created_at, updated_at
|
||||
FROM nfs_exports ORDER BY path ASC`)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list nfs exports: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var exports []NFSExport
|
||||
for rows.Next() {
|
||||
export, err := scanNFSExport(rows)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("scan nfs export: %w", err)
|
||||
}
|
||||
exports = append(exports, export)
|
||||
}
|
||||
return exports, rows.Err()
|
||||
}
|
||||
|
||||
func (d *DB) GetNFSExport(id int64) (NFSExport, error) {
|
||||
row := d.conn.QueryRow(`
|
||||
SELECT id, path, clients, options, created_at, updated_at
|
||||
FROM nfs_exports WHERE id = ?`, id)
|
||||
export, err := scanNFSExport(row)
|
||||
if err == sql.ErrNoRows {
|
||||
return NFSExport{}, fmt.Errorf("nfs export not found")
|
||||
}
|
||||
if err != nil {
|
||||
return NFSExport{}, fmt.Errorf("get nfs export: %w", err)
|
||||
}
|
||||
return export, nil
|
||||
}
|
||||
|
||||
func (d *DB) CreateNFSExport(export NFSExport) (NFSExport, error) {
|
||||
clients, err := encodeJSONStrings(export.Clients)
|
||||
if err != nil {
|
||||
return NFSExport{}, err
|
||||
}
|
||||
result, err := d.conn.Exec(`
|
||||
INSERT INTO nfs_exports (path, clients, options)
|
||||
VALUES (?, ?, ?)`,
|
||||
export.Path, clients, export.Options,
|
||||
)
|
||||
if err != nil {
|
||||
return NFSExport{}, fmt.Errorf("create nfs export: %w", err)
|
||||
}
|
||||
id, err := result.LastInsertId()
|
||||
if err != nil {
|
||||
return NFSExport{}, fmt.Errorf("last insert id: %w", err)
|
||||
}
|
||||
return d.GetNFSExport(id)
|
||||
}
|
||||
|
||||
func (d *DB) UpdateNFSExport(id int64, export NFSExport) (NFSExport, error) {
|
||||
clients, err := encodeJSONStrings(export.Clients)
|
||||
if err != nil {
|
||||
return NFSExport{}, err
|
||||
}
|
||||
result, err := d.conn.Exec(`
|
||||
UPDATE nfs_exports
|
||||
SET path = ?, clients = ?, options = ?, updated_at = datetime('now')
|
||||
WHERE id = ?`,
|
||||
export.Path, clients, export.Options, id,
|
||||
)
|
||||
if err != nil {
|
||||
return NFSExport{}, fmt.Errorf("update nfs export: %w", err)
|
||||
}
|
||||
rows, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return NFSExport{}, err
|
||||
}
|
||||
if rows == 0 {
|
||||
return NFSExport{}, fmt.Errorf("nfs export not found")
|
||||
}
|
||||
return d.GetNFSExport(id)
|
||||
}
|
||||
|
||||
func (d *DB) DeleteNFSExport(id int64) error {
|
||||
result, err := d.conn.Exec(`DELETE FROM nfs_exports WHERE id = ?`, id)
|
||||
if err != nil {
|
||||
return fmt.Errorf("delete nfs export: %w", err)
|
||||
}
|
||||
rows, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if rows == 0 {
|
||||
return fmt.Errorf("nfs export not found")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,263 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user