From ddfa0cc4c36417555cdf9dac54892492d4fdab9f Mon Sep 17 00:00:00 2001 From: Yumechi Date: Wed, 2 Sep 2026 16:51:08 +0800 Subject: [PATCH] enhance: serialize user update methods --- api/user.go | 108 ++++++++++++++++++++++++++-------------- api/user_test.go | 15 +----- database/database.go | 12 ++++- database/user.go | 18 +++---- database/user_test.go | 13 ++--- plugin/manager.go | 2 +- router/router.go | 2 +- test/testdb/database.go | 6 +++ 8 files changed, 107 insertions(+), 69 deletions(-) diff --git a/api/user.go b/api/user.go index 204b5e34..c1ab343b 100644 --- a/api/user.go +++ b/api/user.go @@ -12,15 +12,17 @@ import ( "github.com/gotify/server/v3/model" ) +var errCannotDeleteLastAdmin = errors.New("cannot delete last admin") + // The UserDatabase interface for encapsulating database access. -type UserDatabase interface { - GetUsers() ([]*model.User, error) +type UserDatabase[T UserDatabase[T]] interface { + Txn(fn func(txdb T) error) error + GetUsers(condition ...any) ([]*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) } // UserChangeNotifier notifies listeners for user changes. @@ -58,8 +60,8 @@ func (c *UserChangeNotifier) fireUserAdded(uid uint) error { } // The UserAPI provides handlers for managing users. -type UserAPI struct { - DB UserDatabase +type UserAPI[T UserDatabase[T]] struct { + DB T PasswordStrength int UserChangeNotifier *UserChangeNotifier Registration bool @@ -90,7 +92,7 @@ type UserAPI struct { // description: Forbidden // schema: // $ref: "#/definitions/Error" -func (a *UserAPI) GetUsers(ctx *gin.Context) { +func (a *UserAPI[T]) GetUsers(ctx *gin.Context) { users, err := a.DB.GetUsers() if success := successOrAbort(ctx, 500, err); !success { return @@ -126,7 +128,7 @@ func (a *UserAPI) GetUsers(ctx *gin.Context) { // description: Forbidden // schema: // $ref: "#/definitions/Error" -func (a *UserAPI) GetCurrentUser(ctx *gin.Context) { +func (a *UserAPI[T]) GetCurrentUser(ctx *gin.Context) { user, err := a.DB.GetUserByID(auth.GetUserID(ctx)) if success := successOrAbort(ctx, 500, err); !success { return @@ -185,7 +187,7 @@ func (a *UserAPI) GetCurrentUser(ctx *gin.Context) { // description: Forbidden // schema: // $ref: "#/definitions/Error" -func (a *UserAPI) CreateUser(ctx *gin.Context) { +func (a *UserAPI[T]) CreateUser(ctx *gin.Context) { user := model.CreateUserExternal{} if err := ctx.Bind(&user); err == nil { if err := password.ValidateNewPassword(user.Pass); err != nil { @@ -286,7 +288,7 @@ func (a *UserAPI) CreateUser(ctx *gin.Context) { // description: Not Found // schema: // $ref: "#/definitions/Error" -func (a *UserAPI) GetUserByID(ctx *gin.Context) { +func (a *UserAPI[T]) GetUserByID(ctx *gin.Context) { withID(ctx, "id", func(id uint) { user, err := a.DB.GetUserByID(id) if success := successOrAbort(ctx, 500, err); !success { @@ -336,26 +338,41 @@ func (a *UserAPI) GetUserByID(ctx *gin.Context) { // description: Not Found // schema: // $ref: "#/definitions/Error" -func (a *UserAPI) DeleteUserByID(ctx *gin.Context) { +func (a *UserAPI[T]) DeleteUserByID(ctx *gin.Context) { withID(ctx, "id", func(id uint) { user, err := a.DB.GetUserByID(id) if success := successOrAbort(ctx, 500, err); !success { 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 { + commitError := false + err = a.DB.Txn(func(txdb T) error { + if success := successOrAbort(ctx, 500, txdb.DeleteUserByID(id)); !success { + return err + } + anotherAdmin, err := txdb.GetUsers(&model.User{Admin: true}) + if success := successOrAbort(ctx, 500, err); !success { + return err + } + if user.Admin && len(anotherAdmin) == 0 { + ctx.AbortWithError(400, errCannotDeleteLastAdmin) + return errCannotDeleteLastAdmin + } + if success := successOrAbort(ctx, 500, a.UserChangeNotifier.fireUserDeleted(id)); !success { + return err + } + commitError = true + return nil + }) + if !commitError || err == nil { + break + } + if err != nil { + ctx.AbortWithError(500, err) + return + } } - successOrAbort(ctx, 500, a.DB.DeleteUserByID(id)) } else { ctx.AbortWithError(404, errors.New("user does not exist")) } @@ -395,7 +412,7 @@ func (a *UserAPI) DeleteUserByID(ctx *gin.Context) { // description: Forbidden // schema: // $ref: "#/definitions/Error" -func (a *UserAPI) ChangePassword(ctx *gin.Context) { +func (a *UserAPI[T]) ChangePassword(ctx *gin.Context) { pw := model.UserExternalPass{} if err := ctx.Bind(&pw); err == nil { if err := password.ValidateNewPassword(pw.Pass); err != nil { @@ -461,7 +478,7 @@ func (a *UserAPI) ChangePassword(ctx *gin.Context) { // description: Not Found // schema: // $ref: "#/definitions/Error" -func (a *UserAPI) UpdateUserByID(ctx *gin.Context) { +func (a *UserAPI[T]) UpdateUserByID(ctx *gin.Context) { withID(ctx, "id", func(id uint) { var updatedUser *model.UpdateUserExternal if err := ctx.Bind(&updatedUser); err == nil { @@ -470,15 +487,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 +503,37 @@ func (a *UserAPI) UpdateUserByID(ctx *gin.Context) { } dbUser.Pass = pw } - if success := successOrAbort(ctx, 500, a.DB.UpdateUser(dbUser)); !success { - return + + for range 3 { + commitError := false + + err = a.DB.Txn(func(txdb T) error { + if success := successOrAbort(ctx, 500, txdb.UpdateUser(dbUser)); !success { + return err + } + + anotherAdmin, err := txdb.GetUsers(&model.User{Admin: true}) + if success := successOrAbort(ctx, 500, err); !success { + return err + } + if !updatedUser.Admin && dbUserWasAdmin && len(anotherAdmin) == 0 { + ctx.AbortWithError(400, errCannotDeleteLastAdmin) + return errCannotDeleteLastAdmin + } + + commitError = true + + return nil + }) + + if !commitError || err == nil { + break + } + } + + if err == nil { + ctx.JSON(200, toExternalUser(dbUser)) } - ctx.JSON(200, toExternalUser(dbUser)) } else { ctx.AbortWithError(404, errors.New("user does not exist")) } diff --git a/api/user_test.go b/api/user_test.go index 583bb64e..e7111ee9 100644 --- a/api/user_test.go +++ b/api/user_test.go @@ -25,7 +25,7 @@ func TestUserSuite(t *testing.T) { type UserSuite struct { suite.Suite db *testdb.Database - a *UserAPI + a *UserAPI[*testdb.Database] ctx *gin.Context recorder *httptest.ResponseRecorder notifiedAdd bool @@ -49,7 +49,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[*testdb.Database]{DB: s.db, UserChangeNotifier: s.notifier} } func (s *UserSuite) AfterTest(suiteName, testName string) { @@ -350,17 +350,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() 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) diff --git a/router/router.go b/router/router.go index 74e770d3..771650c5 100644 --- a/router/router.go +++ b/router/router.go @@ -104,7 +104,7 @@ func Create(db *database.GormDatabase, vInfo *model.VersionInfo, conf *config.Co } sessionHandler := api.SessionAPI{DB: db, NotifyDeleted: streamHandler.NotifyDeletedClient, SecureCookie: conf.Server.SecureCookie, LocalAuthEnabled: conf.LocalAuthEnabled} userChangeNotifier := new(api.UserChangeNotifier) - userHandler := api.UserAPI{DB: db, PasswordStrength: conf.PassStrength, UserChangeNotifier: userChangeNotifier, Registration: conf.Registration} + userHandler := api.UserAPI[*database.GormDatabase]{DB: db, PasswordStrength: conf.PassStrength, UserChangeNotifier: userChangeNotifier, Registration: conf.Registration} pluginManager, err := plugin.NewManager(db, conf.PluginsDir, g.Group("/plugin/:id/custom/"), streamHandler) if err != nil { diff --git a/test/testdb/database.go b/test/testdb/database.go index d7e60f38..c0cf4f5d 100644 --- a/test/testdb/database.go +++ b/test/testdb/database.go @@ -20,6 +20,12 @@ type Database struct { t *testing.T } +func (d *Database) Txn(fn func(txdb *Database) error) error { + return d.GormDatabase.Txn(func(txdb *database.GormDatabase) error { + return fn(&Database{GormDatabase: txdb, t: d.t}) + }) +} + // AppClientBuilder has helper methods to create applications and clients. type AppClientBuilder struct { userID uint