diff --git a/api/user.go b/api/user.go index 204b5e34..afd42e78 100644 --- a/api/user.go +++ b/api/user.go @@ -9,19 +9,11 @@ import ( "github.com/gin-gonic/gin" "github.com/gotify/server/v3/auth" "github.com/gotify/server/v3/auth/password" + "github.com/gotify/server/v3/database" "github.com/gotify/server/v3/model" ) -// The UserDatabase interface for encapsulating database access. -type UserDatabase interface { - GetUsers() ([]*model.User, error) - GetUserByID(id uint) (*model.User, error) - GetUserByName(name string) (*model.User, error) - DeleteUserByID(id uint) error - UpdateUser(user *model.User) error - CreateUser(user *model.User) error - CountUser(condition ...any) (int64, error) -} +var errCannotDeleteLastAdmin = errors.New("cannot delete last admin") // UserChangeNotifier notifies listeners for user changes. type UserChangeNotifier struct { @@ -59,7 +51,7 @@ func (c *UserChangeNotifier) fireUserAdded(uid uint) error { // The UserAPI provides handlers for managing users. type UserAPI struct { - DB UserDatabase + DB *database.GormDatabase PasswordStrength int UserChangeNotifier *UserChangeNotifier Registration bool @@ -343,19 +335,26 @@ func (a *UserAPI) DeleteUserByID(ctx *gin.Context) { return } if user != nil { - adminCount, err := a.DB.CountUser(&model.User{Admin: true}) - if success := successOrAbort(ctx, 500, err); !success { - return - } - if user.Admin && adminCount == 1 { - ctx.AbortWithError(400, errors.New("cannot delete last admin")) - return - } - if err := a.UserChangeNotifier.fireUserDeleted(id); err != nil { - ctx.AbortWithError(500, err) - return + for range 3 { + err = a.DB.Txn(func(txdb *database.GormDatabase) error { + if err := txdb.DeleteUserByID(id); err != nil { + return err + } + anotherAdmin, err := txdb.GetUsers(&model.User{Admin: true}) + if err != nil { + return err + } + if user.Admin && len(anotherAdmin) == 0 { + ctx.AbortWithError(400, errCannotDeleteLastAdmin) + return errCannotDeleteLastAdmin + } + return a.UserChangeNotifier.fireUserDeleted(id) + }) + if err == nil || ctx.IsAborted() { + return + } } - successOrAbort(ctx, 500, a.DB.DeleteUserByID(id)) + successOrAbort(ctx, 500, err) } else { ctx.AbortWithError(404, errors.New("user does not exist")) } @@ -470,15 +469,7 @@ func (a *UserAPI) UpdateUserByID(ctx *gin.Context) { return } if dbUser != nil { - adminCount, err := a.DB.CountUser(&model.User{Admin: true}) - if success := successOrAbort(ctx, 500, err); !success { - return - } - if !updatedUser.Admin && dbUser.Admin && adminCount == 1 { - ctx.AbortWithError(400, errors.New("cannot delete last admin")) - return - } - + dbUserWasAdmin := dbUser.Admin dbUser.Name = updatedUser.Name dbUser.Admin = updatedUser.Admin @@ -494,10 +485,35 @@ func (a *UserAPI) UpdateUserByID(ctx *gin.Context) { } dbUser.Pass = pw } - if success := successOrAbort(ctx, 500, a.DB.UpdateUser(dbUser)); !success { - return + + for range 3 { + err = a.DB.Txn(func(txdb *database.GormDatabase) error { + if err := txdb.UpdateUser(dbUser); err != nil { + return err + } + + anotherAdmin, err := txdb.GetUsers(&model.User{Admin: true}) + if err != nil { + return err + } + if !updatedUser.Admin && dbUserWasAdmin && len(anotherAdmin) == 0 { + ctx.AbortWithError(400, errCannotDeleteLastAdmin) + return errCannotDeleteLastAdmin + } + + return nil + }) + + if ctx.IsAborted() { + return + } + + if err == nil { + ctx.JSON(200, toExternalUser(dbUser)) + return + } } - ctx.JSON(200, toExternalUser(dbUser)) + ctx.AbortWithError(500, err) } else { ctx.AbortWithError(404, errors.New("user does not exist")) } diff --git a/api/user_test.go b/api/user_test.go index 583bb64e..d0f6f8ca 100644 --- a/api/user_test.go +++ b/api/user_test.go @@ -1,7 +1,9 @@ package api import ( + "encoding/json" "errors" + "io" "net/http/httptest" "strings" "testing" @@ -49,7 +51,7 @@ func (s *UserSuite) BeforeTest(suiteName, testName string) { s.notifiedAdd = true return nil }) - s.a = &UserAPI{DB: s.db, UserChangeNotifier: s.notifier} + s.a = &UserAPI{DB: s.db.GormDatabase, UserChangeNotifier: s.notifier} } func (s *UserSuite) AfterTest(suiteName, testName string) { @@ -350,17 +352,6 @@ func (s *UserSuite) Test_UpdateUserByID_InvalidID() { assert.Equal(s.T(), 400, s.recorder.Code) } -func (s *UserSuite) Test_UpdateUserByID_EmptyPassword_Expect400() { - s.loginAdmin() - - s.ctx.Params = gin.Params{{Key: "id", Value: "1"}} - - s.ctx.Request = httptest.NewRequest("POST", "/user/1", strings.NewReader(`{"name": "admin", "pass": "", "admin": false}`)) - s.ctx.Request.Header.Set("Content-Type", "application/json") - s.a.UpdateUserByID(s.ctx) - assert.Equal(s.T(), 400, s.recorder.Code) -} - func (s *UserSuite) Test_UpdateUserByID_TooLongPassword_Expect400() { s.loginAdmin() @@ -412,6 +403,12 @@ func (s *UserSuite) Test_UpdateUserByID_UpdateNotPassword() { s.a.UpdateUserByID(s.ctx) assert.Equal(s.T(), 200, s.recorder.Code) + body, err := io.ReadAll(s.recorder.Body) + require.NoError(s.T(), err) + var retUser model.UserExternal + require.NoError(s.T(), json.Unmarshal(body, &retUser)) + assert.Equal(s.T(), "tom", retUser.Name) + assert.Equal(s.T(), true, retUser.Admin) user, err := s.db.GetUserByID(2) assert.NoError(s.T(), err) assert.NotNil(s.T(), user) diff --git a/database/database.go b/database/database.go index 8da80906..ec322177 100644 --- a/database/database.go +++ b/database/database.go @@ -172,14 +172,24 @@ func createDirectoryIfSqlite(dialect, connection string) { // GormDatabase is a wrapper for the gorm framework. type GormDatabase struct { - DB *gorm.DB + DB *gorm.DB + Nested bool } // Close closes the gorm database connection. func (d *GormDatabase) Close() { + if d.Nested { + return + } sqldb, err := d.DB.DB() if err != nil { return } sqldb.Close() } + +func (d *GormDatabase) Txn(fn func(txdb *GormDatabase) error) error { + return d.DB.Transaction(func(tx *gorm.DB) error { + return fn(&GormDatabase{DB: tx, Nested: true}) + }, &sql.TxOptions{Isolation: sql.LevelSerializable}) +} diff --git a/database/user.go b/database/user.go index a2bdcde1..d1acfc6e 100644 --- a/database/user.go +++ b/database/user.go @@ -44,23 +44,19 @@ func (d *GormDatabase) GetUserByID(id uint) (*model.User, error) { return nil, err } -// CountUser returns the user count which satisfies the given condition. -func (d *GormDatabase) CountUser(condition ...any) (int64, error) { - c := int64(-1) +// GetUsers returns the users which satisfy the given condition. +func (d *GormDatabase) GetUsers(condition ...any) ([]*model.User, error) { + users := make([]*model.User, 0) handle := d.DB.Model(new(model.User)) if len(condition) == 1 { handle = handle.Where(condition[0]) } else if len(condition) > 1 { handle = handle.Where(condition[0], condition[1:]...) } - err := handle.Count(&c).Error - return c, err -} - -// GetUsers returns all users. -func (d *GormDatabase) GetUsers() ([]*model.User, error) { - var users []*model.User - err := d.DB.Find(&users).Error + err := handle.Find(&users).Error + if err == gorm.ErrRecordNotFound { + return nil, nil + } return users, err } diff --git a/database/user_test.go b/database/user_test.go index c10d97a8..aefe182c 100644 --- a/database/user_test.go +++ b/database/user_test.go @@ -19,9 +19,10 @@ func (s *DatabaseSuite) TestUser() { require.NoError(s.T(), err) assert.NotNil(s.T(), jmattheis, "on bootup the first user should be automatically created") - adminCount, err := s.db.CountUser("admin = ?", true) + admins, err := s.db.GetUsers("admin = ?", true) require.NoError(s.T(), err) - assert.Equal(s.T(), int64(1), adminCount, "there is initially one admin") + assert.Len(s.T(), admins, 1) + assert.True(s.T(), admins[0].Admin, "the admin user should be an admin") users, err := s.db.GetUsers() require.NoError(s.T(), err) @@ -31,9 +32,9 @@ func (s *DatabaseSuite) TestUser() { nicories := &model.User{Name: "nicories", Pass: []byte{1, 2, 3, 4}, Admin: false} s.db.CreateUser(nicories) assert.NotEqual(s.T(), 0, nicories.ID, "on create user a new id should be assigned") - userCount, err := s.db.CountUser() + users, err = s.db.GetUsers() require.NoError(s.T(), err) - assert.Equal(s.T(), int64(2), userCount, "two users should exist") + assert.Len(s.T(), users, 2, "two users should exist") user, err = s.db.GetUserByName("nicories") require.NoError(s.T(), err) @@ -58,9 +59,9 @@ func (s *DatabaseSuite) TestUser() { require.NoError(s.T(), err) assert.Len(s.T(), users, 2) - adminCount, err = s.db.CountUser(&model.User{Admin: true}) + admins, err = s.db.GetUsers(&model.User{Admin: true}) require.NoError(s.T(), err) - assert.Equal(s.T(), int64(2), adminCount, "two admins exist") + assert.Len(s.T(), admins, 2, "two admins exist") require.NoError(s.T(), s.db.DeleteUserByID(tom.ID)) users, err = s.db.GetUsers() diff --git a/plugin/manager.go b/plugin/manager.go index feac8efa..f23de4c4 100644 --- a/plugin/manager.go +++ b/plugin/manager.go @@ -23,7 +23,7 @@ import ( // The Database interface for encapsulating database access. type Database interface { - GetUsers() ([]*model.User, error) + GetUsers(condition ...any) ([]*model.User, error) GetPluginConfByUserAndPath(userid uint, path string) (*model.PluginConf, error) CreatePluginConf(p *model.PluginConf) error GetPluginConfByApplicationID(appid uint) (*model.PluginConf, error)