feat(user): implement registration for the first user with concurrency handling
This commit is contained in:
+31
-3
@@ -4,6 +4,7 @@ import (
|
|||||||
"crypto/sha256"
|
"crypto/sha256"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
@@ -81,12 +82,36 @@ func HashPassword(password string) string {
|
|||||||
return hex.EncodeToString(hash[:])
|
return hex.EncodeToString(hash[:])
|
||||||
}
|
}
|
||||||
|
|
||||||
// CreateUser inserts a new user into the database.
|
// ErrUsersExist is returned by RegisterFirstUser when the users table is not
|
||||||
func CreateUser(username, password, level string) (int64, error) {
|
// 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)
|
hashedPassword := HashPassword(password)
|
||||||
registerDate := time.Now().Format("2006-01-02 15:04:05")
|
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 (?, ?, ?, ?)",
|
"INSERT INTO users (username, password, level, register_date) VALUES (?, ?, ?, ?)",
|
||||||
username, hashedPassword, level, registerDate,
|
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)
|
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()
|
return result.LastInsertId()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
+9
-2
@@ -2,6 +2,7 @@ package handler
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
|
||||||
db "nukumizu-backend/database"
|
db "nukumizu-backend/database"
|
||||||
@@ -98,10 +99,16 @@ func UserRegisterHandler(w http.ResponseWriter, r *http.Request) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
userID, err := db.CreateUser(req.Username, req.Password, "admin")
|
userID, err := db.RegisterFirstUser(req.Username, req.Password, "admin")
|
||||||
if err != nil {
|
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())
|
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
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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())
|
||||||
|
}
|
||||||
|
}
|
||||||
Submodule
+1
Submodule ~/.agents/skills/avoid-ai-design added at 8337060636
Submodule
+1
Submodule ~/.claude/skills/avoid-ai-design added at 8337060636
Reference in New Issue
Block a user