feat(user): implement registration for the first user with concurrency handling

This commit is contained in:
2026-09-08 22:53:02 +08:00
parent b63556624f
commit 2a1aa3d4ab
6 changed files with 176 additions and 5 deletions
+31 -3
View File
@@ -4,6 +4,7 @@ import (
"crypto/sha256"
"database/sql"
"encoding/hex"
"errors"
"fmt"
"os"
"path/filepath"
@@ -81,12 +82,36 @@ func HashPassword(password string) string {
return hex.EncodeToString(hash[:])
}
// CreateUser inserts a new user into the database.
func CreateUser(username, password, level string) (int64, error) {
// 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")
result, err := UserDB.Exec(
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,
)
@@ -94,6 +119,9 @@ func CreateUser(username, password, level string) (int64, error) {
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()
}
+80
View File
@@ -0,0 +1,80 @@
package database
import (
"errors"
"path/filepath"
"sync"
"testing"
)
// initTempDB opens a fresh user database in a temp directory and registers a
// cleanup that closes it when the test finishes.
func initTempDB(t *testing.T) {
t.Helper()
path := filepath.Join(t.TempDir(), "user.db")
if err := InitUserDB(path); err != nil {
t.Fatalf("InitUserDB: %v", err)
}
t.Cleanup(CloseUserDB)
}
func TestRegisterFirstUserClosedAfterFirst(t *testing.T) {
initTempDB(t)
if _, err := RegisterFirstUser("alice", "password1", "admin"); err != nil {
t.Fatalf("first registration should succeed: %v", err)
}
if _, err := RegisterFirstUser("bob", "password2", "admin"); !errors.Is(err, ErrUsersExist) {
t.Fatalf("second registration error = %v, want ErrUsersExist", err)
}
count, err := GetUserCount()
if err != nil {
t.Fatalf("GetUserCount: %v", err)
}
if count != 1 {
t.Fatalf("user count = %d, want 1", count)
}
}
func TestRegisterFirstUserConcurrentOnlyOneWins(t *testing.T) {
initTempDB(t)
const n = 8
var wg sync.WaitGroup
results := make(chan error, n)
for i := range n {
wg.Add(1)
go func(i int) {
defer wg.Done()
_, err := RegisterFirstUser("user"+string(rune('a'+i)), "password", "admin")
results <- err
}(i)
}
wg.Wait()
close(results)
successes := 0
for err := range results {
switch {
case err == nil:
successes++
case errors.Is(err, ErrUsersExist):
// Expected for every registration that lost the race.
default:
t.Fatalf("unexpected error: %v", err)
}
}
if successes != 1 {
t.Fatalf("concurrent registrations: %d succeeded, want exactly 1", successes)
}
count, err := GetUserCount()
if err != nil {
t.Fatalf("GetUserCount: %v", err)
}
if count != 1 {
t.Fatalf("user count = %d, want 1", count)
}
}