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() } }