186 lines
5.5 KiB
Go
186 lines
5.5 KiB
Go
package database
|
|
|
|
import (
|
|
"crypto/sha256"
|
|
"database/sql"
|
|
"encoding/hex"
|
|
"errors"
|
|
"fmt"
|
|
"os"
|
|
"path/filepath"
|
|
"time"
|
|
|
|
_ "modernc.org/sqlite"
|
|
)
|
|
|
|
// UserDB is the global database handle for user data.
|
|
var UserDB *sql.DB
|
|
|
|
// InitUserDB opens (or creates) the user database and initializes the users table.
|
|
func InitUserDB(dbPath string) error {
|
|
dir := filepath.Dir(dbPath)
|
|
if dir != "." {
|
|
if err := os.MkdirAll(dir, 0755); err != nil {
|
|
return fmt.Errorf("failed to create database directory: %w", err)
|
|
}
|
|
}
|
|
|
|
var err error
|
|
UserDB, err = sql.Open("sqlite", dbPath)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to open user database: %w", err)
|
|
}
|
|
|
|
UserDB.SetMaxOpenConns(1) // SQLite serializes writes.
|
|
UserDB.SetMaxIdleConns(1)
|
|
UserDB.SetConnMaxLifetime(5 * time.Minute)
|
|
|
|
if err := UserDB.Ping(); err != nil {
|
|
return fmt.Errorf("failed to ping user database: %w", err)
|
|
}
|
|
|
|
// Use the rollback journal (DELETE) instead of WAL.
|
|
//
|
|
// WAL memory-maps a -shm index file. modernc.org/sqlite (a C-to-Go
|
|
// transpile) is built without SEH on Windows, so an I/O page fault on that
|
|
// mapping (common when the database lives on an SMB/network share)
|
|
// terminates the process with EXCEPTION_IN_PAGE_ERROR instead of being
|
|
// caught and retried. With SetMaxOpenConns(1) above, WAL provides no
|
|
// concurrency benefit either, so the rollback journal is strictly better.
|
|
if _, err := UserDB.Exec("PRAGMA journal_mode=DELETE"); err != nil {
|
|
return fmt.Errorf("failed to set journal mode: %w", err)
|
|
}
|
|
|
|
// Wait up to 5 seconds for a busy database instead of immediately failing with SQLITE_BUSY.
|
|
if _, err := UserDB.Exec("PRAGMA busy_timeout=5000"); err != nil {
|
|
return fmt.Errorf("failed to set busy timeout: %w", err)
|
|
}
|
|
|
|
if err := initUserTable(); err != nil {
|
|
return fmt.Errorf("failed to initialize user table: %w", err)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func initUserTable() error {
|
|
createTableSQL := `
|
|
CREATE TABLE IF NOT EXISTS users (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
username TEXT NOT NULL UNIQUE,
|
|
password TEXT NOT NULL,
|
|
level TEXT NOT NULL DEFAULT 'admin',
|
|
register_date TEXT NOT NULL
|
|
);`
|
|
_, err := UserDB.Exec(createTableSQL)
|
|
return err
|
|
}
|
|
|
|
// HashPassword returns the SHA-256 hex hash of a password.
|
|
func HashPassword(password string) string {
|
|
hash := sha256.Sum256([]byte(password))
|
|
return hex.EncodeToString(hash[:])
|
|
}
|
|
|
|
// ErrUsersExist is returned by RegisterFirstUser when the users table is not
|
|
// empty. Registration is only ever allowed for the very first user.
|
|
var ErrUsersExist = errors.New("registration rejected: only the first user can be registered via this endpoint")
|
|
|
|
// RegisterFirstUser atomically creates the first user, but only while the users
|
|
// table is empty. The emptiness check and the insert run inside a single
|
|
// transaction. Because InitUserDB caps the pool at one connection, a concurrent
|
|
// registration blocks at Begin until the in-flight transaction commits, so it
|
|
// cannot observe the table as empty between the check and the insert. This
|
|
// closes the check-then-insert race two simultaneous first-run registrations
|
|
// would otherwise hit.
|
|
func RegisterFirstUser(username, password, level string) (int64, error) {
|
|
hashedPassword := HashPassword(password)
|
|
registerDate := time.Now().Format("2006-01-02 15:04:05")
|
|
|
|
tx, err := UserDB.Begin()
|
|
if err != nil {
|
|
return 0, fmt.Errorf("failed to begin transaction: %w", err)
|
|
}
|
|
defer tx.Rollback()
|
|
|
|
var count int
|
|
if err := tx.QueryRow("SELECT COUNT(*) FROM users").Scan(&count); err != nil {
|
|
return 0, fmt.Errorf("failed to check existing users: %w", err)
|
|
}
|
|
if count > 0 {
|
|
return 0, errors.New("registration closed: users already exist")
|
|
}
|
|
|
|
result, err := tx.Exec(
|
|
"INSERT INTO users (username, password, level, register_date) VALUES (?, ?, ?, ?)",
|
|
username, hashedPassword, level, registerDate,
|
|
)
|
|
if err != nil {
|
|
return 0, fmt.Errorf("failed to create user: %w", err)
|
|
}
|
|
|
|
if err := tx.Commit(); err != nil {
|
|
return 0, fmt.Errorf("failed to commit user creation: %w", err)
|
|
}
|
|
return result.LastInsertId()
|
|
}
|
|
|
|
// GetUserByUsername retrieves a user by username and validates the password.
|
|
func GetUserByUsername(username, password string) (int64, string, string, string, error) {
|
|
hashedPassword := HashPassword(password)
|
|
|
|
var id int64
|
|
var dbUsername string
|
|
var level string
|
|
var registerDate string
|
|
|
|
err := UserDB.QueryRow(
|
|
"SELECT id, username, level, register_date FROM users WHERE username = ? AND password = ?",
|
|
username, hashedPassword,
|
|
).Scan(&id, &dbUsername, &level, ®isterDate)
|
|
|
|
if err == sql.ErrNoRows {
|
|
return 0, "", "", "", fmt.Errorf("invalid username or password")
|
|
}
|
|
if err != nil {
|
|
return 0, "", "", "", fmt.Errorf("database error: %w", err)
|
|
}
|
|
|
|
return id, dbUsername, level, registerDate, nil
|
|
}
|
|
|
|
// GetUserCount returns the total number of users in the database.
|
|
func GetUserCount() (int, error) {
|
|
var count int
|
|
err := UserDB.QueryRow("SELECT COUNT(*) FROM users").Scan(&count)
|
|
return count, err
|
|
}
|
|
|
|
// GetUserByID retrieves a user by their ID.
|
|
func GetUserByID(id int64) (string, string, string, error) {
|
|
var username string
|
|
var level string
|
|
var registerDate string
|
|
|
|
err := UserDB.QueryRow(
|
|
"SELECT username, level, register_date FROM users WHERE id = ?",
|
|
id,
|
|
).Scan(&username, &level, ®isterDate)
|
|
|
|
if err == sql.ErrNoRows {
|
|
return "", "", "", fmt.Errorf("user not found")
|
|
}
|
|
if err != nil {
|
|
return "", "", "", fmt.Errorf("database error: %w", err)
|
|
}
|
|
|
|
return username, level, registerDate, nil
|
|
}
|
|
|
|
// CloseUserDB closes the user database connection.
|
|
func CloseUserDB() {
|
|
if UserDB != nil {
|
|
UserDB.Close()
|
|
}
|
|
}
|