diff --git a/database/user.go b/database/user.go index 2013609..bdaeec5 100644 --- a/database/user.go +++ b/database/user.go @@ -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() } diff --git a/database/user_test.go b/database/user_test.go new file mode 100644 index 0000000..62fa8e8 --- /dev/null +++ b/database/user_test.go @@ -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) + } +} diff --git a/handler/user.go b/handler/user.go index 7792de3..c42766b 100644 --- a/handler/user.go +++ b/handler/user.go @@ -2,6 +2,7 @@ package handler import ( "encoding/json" + "errors" "net/http" db "nukumizu-backend/database" @@ -98,10 +99,16 @@ func UserRegisterHandler(w http.ResponseWriter, r *http.Request) { return } - userID, err := db.CreateUser(req.Username, req.Password, "admin") + userID, err := db.RegisterFirstUser(req.Username, req.Password, "admin") if err != nil { + // A concurrent registration may have won between the count check above + // and this insert; both map to the same "registration is closed" answer. + if errors.Is(err, db.ErrUsersExist) { + utils.SendErrorResponse(w, http.StatusForbidden, "registration is closed: users already exist") + return + } postLog.Error("Failed to register user: " + err.Error()) - utils.SendErrorResponse(w, http.StatusInternalServerError, "failed to register user, username may already exist") + utils.SendErrorResponse(w, http.StatusInternalServerError, "failed to register user") return } diff --git a/handler/user_test.go b/handler/user_test.go new file mode 100644 index 0000000..608f941 --- /dev/null +++ b/handler/user_test.go @@ -0,0 +1,54 @@ +package handler + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "path/filepath" + "strconv" + "testing" + "time" + + db "nukumizu-backend/database" +) + +// initHandlerUserDB opens a fresh user database in a temp directory for the +// register handler tests. +func initHandlerUserDB(t *testing.T) { + t.Helper() + path := filepath.Join(t.TempDir(), "user.db") + if err := db.InitUserDB(path); err != nil { + t.Fatalf("InitUserDB: %v", err) + } + t.Cleanup(db.CloseUserDB) +} + +func registerRequest(t *testing.T, username, password string) *http.Request { + t.Helper() + body, err := json.Marshal(map[string]string{"username": username, "password": password}) + if err != nil { + t.Fatalf("marshal body: %v", err) + } + req := httptest.NewRequest(http.MethodPost, "/api/user/register", bytes.NewReader(body)) + req.Header.Set("X-Timestamp", strconv.FormatInt(time.Now().Unix(), 10)) + return req +} + +// TestRegisterOnlyFirstUser exercises the API rule: the database accepts +// exactly the first registration and rejects every later one. +func TestRegisterOnlyFirstUser(t *testing.T) { + initHandlerUserDB(t) + + w := httptest.NewRecorder() + UserRegisterHandler(w, registerRequest(t, "alice", "password1")) + if w.Code != http.StatusOK { + t.Fatalf("first register status = %d, body = %s", w.Code, w.Body.String()) + } + + w2 := httptest.NewRecorder() + UserRegisterHandler(w2, registerRequest(t, "bob", "password2")) + if w2.Code != http.StatusForbidden { + t.Fatalf("second register status = %d, body = %s", w2.Code, w2.Body.String()) + } +} diff --git a/~/.agents/skills/avoid-ai-design b/~/.agents/skills/avoid-ai-design new file mode 160000 index 0000000..8337060 --- /dev/null +++ b/~/.agents/skills/avoid-ai-design @@ -0,0 +1 @@ +Subproject commit 8337060636a8cf12e32e883eb367becd702aa526 diff --git a/~/.claude/skills/avoid-ai-design b/~/.claude/skills/avoid-ai-design new file mode 160000 index 0000000..8337060 --- /dev/null +++ b/~/.claude/skills/avoid-ai-design @@ -0,0 +1 @@ +Subproject commit 8337060636a8cf12e32e883eb367becd702aa526