diff --git a/route/v1/user.go b/route/v1/user.go index 587227d..90ab2d3 100644 --- a/route/v1/user.go +++ b/route/v1/user.go @@ -5,6 +5,7 @@ import ( "crypto/ecdsa" "encoding/base64" json2 "encoding/json" + "errors" "image" "image/png" "io" @@ -31,7 +32,6 @@ import ( "github.com/IceWhaleTech/CasaOS-UserService/pkg/utils/file" model2 "github.com/IceWhaleTech/CasaOS-UserService/service/model" "github.com/labstack/echo/v4" - uuid "github.com/satori/go.uuid" "github.com/tidwall/gjson" "go.uber.org/zap" "golang.org/x/time/rate" @@ -48,11 +48,6 @@ func PostUserRegister(ctx echo.Context) error { username := json["username"] pwd := json["password"] key := json["key"] - if _, ok := service.UserRegisterHash[key]; !ok { - return ctx.JSON(common_err.CLIENT_ERROR, - model.Result{Success: common_err.KEY_NOT_EXIST, Message: common_err.GetMsg(common_err.KEY_NOT_EXIST)}) - } - if len(username) == 0 || len(pwd) == 0 { return ctx.JSON(common_err.CLIENT_ERROR, model.Result{Success: common_err.INVALID_PARAMS, Message: common_err.GetMsg(common_err.INVALID_PARAMS)}) @@ -72,12 +67,21 @@ func PostUserRegister(ctx echo.Context) error { user.Password = encryption.GetMD5ByStr(pwd) user.Role = "admin" - user = service.MyService.User().CreateUser(user) - if user.Id == 0 { + var err error + user, err = service.MyService.User().RegisterInitialUser(key, user) + if errors.Is(err, service.ErrRegistrationKeyInvalid) { + return ctx.JSON(common_err.CLIENT_ERROR, + model.Result{Success: common_err.KEY_NOT_EXIST, Message: common_err.GetMsg(common_err.KEY_NOT_EXIST)}) + } + if errors.Is(err, service.ErrAlreadyInitialized) { + return ctx.JSON(http.StatusForbidden, + model.Result{Success: common_err.USER_EXIST, Message: common_err.GetMsg(common_err.USER_EXIST)}) + } + if err != nil { + logger.Error("register initial user error", zap.Error(err)) return ctx.JSON(common_err.SERVICE_ERROR, model.Result{Success: common_err.SERVICE_ERROR, Message: common_err.GetMsg(common_err.SERVICE_ERROR)}) } file.MkDir(config.AppInfo.UserDataPath + "/" + strconv.Itoa(user.Id)) - delete(service.UserRegisterHash, key) return ctx.JSON(common_err.SUCCESS, model.Result{Success: common_err.SUCCESS, Message: common_err.GetMsg(common_err.SUCCESS)}) } @@ -753,12 +757,16 @@ func DeleteUserAll(ctx echo.Context) error { func GetUserStatus(ctx echo.Context) error { data := make(map[string]interface{}, 2) - if service.MyService.User().GetUserCount() > 0 { + initialized, key, err := service.MyService.User().RegistrationStatus() + if err != nil { + logger.Error("get registration status error", zap.Error(err)) + return ctx.JSON(http.StatusInternalServerError, + model.Result{Success: common_err.SERVICE_ERROR, Message: common_err.GetMsg(common_err.SERVICE_ERROR)}) + } + if initialized { data["initialized"] = true data["key"] = "" } else { - key := uuid.NewV4().String() - service.UserRegisterHash[key] = key data["key"] = key data["initialized"] = false } diff --git a/service/registration_test.go b/service/registration_test.go new file mode 100644 index 0000000..db5547c --- /dev/null +++ b/service/registration_test.go @@ -0,0 +1,74 @@ +package service + +import ( + "errors" + "path/filepath" + "sync" + "sync/atomic" + "testing" + + "github.com/IceWhaleTech/CasaOS-UserService/service/model" + "github.com/glebarez/sqlite" + "gorm.io/gorm" +) + +func registrationTestService(t *testing.T) *userService { + t.Helper() + registrationMu.Lock() + registrationKey = "" + registrationMu.Unlock() + db, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "users.db")), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + if err := db.AutoMigrate(&model.UserDBModel{}); err != nil { + t.Fatal(err) + } + return &userService{db: db} +} + +func TestRegistrationKeyIsUniqueAndInvalidated(t *testing.T) { + u := registrationTestService(t) + first, key, err := u.RegistrationStatus() + if err != nil || first || key == "" { + t.Fatalf("first status = (%v, %q, %v)", first, key, err) + } + initialized, again, err := u.RegistrationStatus() + if err != nil || initialized || again != key { + t.Fatalf("second status did not reuse the single key") + } + user, err := u.RegisterInitialUser(key, model.UserDBModel{Username: "owner", Role: "admin"}) + if err != nil || user.Id == 0 { + t.Fatalf("initial registration failed: %v", err) + } + if _, err := u.RegisterInitialUser(key, model.UserDBModel{Username: "second", Role: "admin"}); !errors.Is(err, ErrRegistrationKeyInvalid) { + t.Fatalf("reused key error = %v", err) + } + initialized, again, err = u.RegistrationStatus() + if err != nil || !initialized || again != "" || u.GetUserCount() != 1 { + t.Fatalf("post-registration status = (%v, %q, %v), count = %d", initialized, again, err, u.GetUserCount()) + } +} + +func TestConcurrentInitialRegistrationCreatesOneUser(t *testing.T) { + u := registrationTestService(t) + _, key, err := u.RegistrationStatus() + if err != nil { + t.Fatal(err) + } + var successes atomic.Int32 + var wg sync.WaitGroup + for i := 0; i < 16; i++ { + wg.Add(1) + go func() { + defer wg.Done() + if _, err := u.RegisterInitialUser(key, model.UserDBModel{Username: "owner", Role: "admin"}); err == nil { + successes.Add(1) + } + }() + } + wg.Wait() + if successes.Load() != 1 || u.GetUserCount() != 1 { + t.Fatalf("registrations = %d, users = %d; want one", successes.Load(), u.GetUserCount()) + } +} diff --git a/service/user.go b/service/user.go index 54bed10..ed505d7 100644 --- a/service/user.go +++ b/service/user.go @@ -11,13 +11,16 @@ package service import ( "crypto/ecdsa" + "errors" "io" "mime/multipart" "os" + "sync" "github.com/IceWhaleTech/CasaOS-Common/utils/jwt" "github.com/IceWhaleTech/CasaOS-Common/utils/logger" "github.com/IceWhaleTech/CasaOS-UserService/service/model" + uuid "github.com/satori/go.uuid" "go.uber.org/zap" "gorm.io/gorm" ) @@ -25,6 +28,8 @@ import ( type UserService interface { UpLoadFile(file multipart.File, name string) error CreateUser(m model.UserDBModel) model.UserDBModel + RegistrationStatus() (bool, string, error) + RegisterInitialUser(key string, user model.UserDBModel) (model.UserDBModel, error) GetUserCount() (userCount int64) UpdateUser(m model.UserDBModel) UpdateUserPassword(m model.UserDBModel) @@ -39,7 +44,12 @@ type UserService interface { GetKeyPair() (*ecdsa.PrivateKey, *ecdsa.PublicKey) } -var UserRegisterHash = make(map[string]string) +var ( + registrationMu sync.Mutex + registrationKey string + ErrRegistrationKeyInvalid = errors.New("invalid registration key") + ErrAlreadyInitialized = errors.New("user service already initialized") +) type userService struct { privateKey *ecdsa.PrivateKey // keep this private - NEVER expose it!!! @@ -66,6 +76,51 @@ func (u *userService) CreateUser(m model.UserDBModel) model.UserDBModel { return m } +func (u *userService) RegistrationStatus() (bool, string, error) { + registrationMu.Lock() + defer registrationMu.Unlock() + + var count int64 + if err := u.db.Model(&model.UserDBModel{}).Count(&count).Error; err != nil { + return false, "", err + } + if count > 0 { + registrationKey = "" + return true, "", nil + } + if registrationKey == "" { + registrationKey = uuid.NewV4().String() + } + return false, registrationKey, nil +} + +func (u *userService) RegisterInitialUser(key string, user model.UserDBModel) (model.UserDBModel, error) { + registrationMu.Lock() + defer registrationMu.Unlock() + + if key == "" || key != registrationKey { + return model.UserDBModel{}, ErrRegistrationKeyInvalid + } + err := u.db.Transaction(func(tx *gorm.DB) error { + var count int64 + if err := tx.Model(&model.UserDBModel{}).Count(&count).Error; err != nil { + return err + } + if count != 0 { + return ErrAlreadyInitialized + } + return tx.Create(&user).Error + }) + if err != nil { + if errors.Is(err, ErrAlreadyInitialized) { + registrationKey = "" + } + return model.UserDBModel{}, err + } + registrationKey = "" + return user, nil +} + func (u *userService) GetUserCount() (userCount int64) { u.db.Find(&model.UserDBModel{}).Count(&userCount) return