Skip to content

Commit e787139

Browse files
enhance: serialize user update methods
1 parent 14bfc25 commit e787139

7 files changed

Lines changed: 78 additions & 60 deletions

File tree

‎api/user.go‎

Lines changed: 50 additions & 28 deletions
Original file line numberDiff line numberDiff line change
@@ -9,18 +9,21 @@ import (
99
"github.com/gin-gonic/gin"
1010
"github.com/gotify/server/v3/auth"
1111
"github.com/gotify/server/v3/auth/password"
12+
"github.com/gotify/server/v3/database"
1213
"github.com/gotify/server/v3/model"
1314
)
1415

16+
var errCannotDeleteLastAdmin = errors.New("cannot delete last admin")
17+
1518
// The UserDatabase interface for encapsulating database access.
16-
type UserDatabase interface {
17-
GetUsers() ([]*model.User, error)
19+
type UserDatabase[T UserDatabase[T]] interface {
20+
Txn(fn func(txdb T) error) error
21+
GetUsers(condition ...any) ([]*model.User, error)
1822
GetUserByID(id uint) (*model.User, error)
1923
GetUserByName(name string) (*model.User, error)
2024
DeleteUserByID(id uint) error
2125
UpdateUser(user *model.User) error
2226
CreateUser(user *model.User) error
23-
CountUser(condition ...any) (int64, error)
2427
}
2528

2629
// UserChangeNotifier notifies listeners for user changes.
@@ -59,7 +62,7 @@ func (c *UserChangeNotifier) fireUserAdded(uid uint) error {
5962

6063
// The UserAPI provides handlers for managing users.
6164
type UserAPI struct {
62-
DB UserDatabase
65+
DB database.GormDatabase
6366
PasswordStrength int
6467
UserChangeNotifier *UserChangeNotifier
6568
Registration bool
@@ -343,19 +346,26 @@ func (a *UserAPI) DeleteUserByID(ctx *gin.Context) {
343346
return
344347
}
345348
if user != nil {
346-
adminCount, err := a.DB.CountUser(&model.User{Admin: true})
347-
if success := successOrAbort(ctx, 500, err); !success {
348-
return
349-
}
350-
if user.Admin && adminCount == 1 {
351-
ctx.AbortWithError(400, errors.New("cannot delete last admin"))
352-
return
353-
}
354-
if err := a.UserChangeNotifier.fireUserDeleted(id); err != nil {
355-
ctx.AbortWithError(500, err)
356-
return
349+
for range 3 {
350+
err = a.DB.Txn(func(txdb *database.GormDatabase) error {
351+
if err := txdb.DeleteUserByID(id); err != nil {
352+
return err
353+
}
354+
anotherAdmin, err := txdb.GetUsers(&model.User{Admin: true})
355+
if err != nil {
356+
return err
357+
}
358+
if user.Admin && len(anotherAdmin) == 0 {
359+
ctx.AbortWithError(400, errCannotDeleteLastAdmin)
360+
return errCannotDeleteLastAdmin
361+
}
362+
return a.UserChangeNotifier.fireUserDeleted(id)
363+
})
364+
if err == nil || ctx.IsAborted() {
365+
return
366+
}
357367
}
358-
successOrAbort(ctx, 500, a.DB.DeleteUserByID(id))
368+
successOrAbort(ctx, 500, err)
359369
} else {
360370
ctx.AbortWithError(404, errors.New("user does not exist"))
361371
}
@@ -470,15 +480,7 @@ func (a *UserAPI) UpdateUserByID(ctx *gin.Context) {
470480
return
471481
}
472482
if dbUser != nil {
473-
adminCount, err := a.DB.CountUser(&model.User{Admin: true})
474-
if success := successOrAbort(ctx, 500, err); !success {
475-
return
476-
}
477-
if !updatedUser.Admin && dbUser.Admin && adminCount == 1 {
478-
ctx.AbortWithError(400, errors.New("cannot delete last admin"))
479-
return
480-
}
481-
483+
dbUserWasAdmin := dbUser.Admin
482484
dbUser.Name = updatedUser.Name
483485
dbUser.Admin = updatedUser.Admin
484486

@@ -494,10 +496,30 @@ func (a *UserAPI) UpdateUserByID(ctx *gin.Context) {
494496
}
495497
dbUser.Pass = pw
496498
}
497-
if success := successOrAbort(ctx, 500, a.DB.UpdateUser(dbUser)); !success {
498-
return
499+
500+
for range 3 {
501+
err = a.DB.Txn(func(txdb *database.GormDatabase) error {
502+
if err := txdb.UpdateUser(dbUser); err != nil {
503+
return err
504+
}
505+
506+
anotherAdmin, err := txdb.GetUsers(&model.User{Admin: true})
507+
if err != nil {
508+
return err
509+
}
510+
if !updatedUser.Admin && dbUserWasAdmin && len(anotherAdmin) == 0 {
511+
ctx.AbortWithError(400, errCannotDeleteLastAdmin)
512+
return errCannotDeleteLastAdmin
513+
}
514+
515+
return nil
516+
})
517+
518+
if err == nil || ctx.IsAborted() {
519+
return
520+
}
499521
}
500-
ctx.JSON(200, toExternalUser(dbUser))
522+
successOrAbort(ctx, 500, err)
501523
} else {
502524
ctx.AbortWithError(404, errors.New("user does not exist"))
503525
}

‎api/user_test.go‎

Lines changed: 1 addition & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -49,7 +49,7 @@ func (s *UserSuite) BeforeTest(suiteName, testName string) {
4949
s.notifiedAdd = true
5050
return nil
5151
})
52-
s.a = &UserAPI{DB: s.db, UserChangeNotifier: s.notifier}
52+
s.a = &UserAPI{DB: *s.db.GormDatabase, UserChangeNotifier: s.notifier}
5353
}
5454

5555
func (s *UserSuite) AfterTest(suiteName, testName string) {
@@ -350,17 +350,6 @@ func (s *UserSuite) Test_UpdateUserByID_InvalidID() {
350350
assert.Equal(s.T(), 400, s.recorder.Code)
351351
}
352352

353-
func (s *UserSuite) Test_UpdateUserByID_EmptyPassword_Expect400() {
354-
s.loginAdmin()
355-
356-
s.ctx.Params = gin.Params{{Key: "id", Value: "1"}}
357-
358-
s.ctx.Request = httptest.NewRequest("POST", "/user/1", strings.NewReader(`{"name": "admin", "pass": "", "admin": false}`))
359-
s.ctx.Request.Header.Set("Content-Type", "application/json")
360-
s.a.UpdateUserByID(s.ctx)
361-
assert.Equal(s.T(), 400, s.recorder.Code)
362-
}
363-
364353
func (s *UserSuite) Test_UpdateUserByID_TooLongPassword_Expect400() {
365354
s.loginAdmin()
366355

‎database/database.go‎

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -172,14 +172,24 @@ func createDirectoryIfSqlite(dialect, connection string) {
172172

173173
// GormDatabase is a wrapper for the gorm framework.
174174
type GormDatabase struct {
175-
DB *gorm.DB
175+
DB *gorm.DB
176+
Nested bool
176177
}
177178

178179
// Close closes the gorm database connection.
179180
func (d *GormDatabase) Close() {
181+
if d.Nested {
182+
return
183+
}
180184
sqldb, err := d.DB.DB()
181185
if err != nil {
182186
return
183187
}
184188
sqldb.Close()
185189
}
190+
191+
func (d *GormDatabase) Txn(fn func(txdb *GormDatabase) error) error {
192+
return d.DB.Transaction(func(tx *gorm.DB) error {
193+
return fn(&GormDatabase{DB: tx, Nested: true})
194+
}, &sql.TxOptions{Isolation: sql.LevelSerializable})
195+
}

‎database/user.go‎

Lines changed: 7 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -44,23 +44,19 @@ func (d *GormDatabase) GetUserByID(id uint) (*model.User, error) {
4444
return nil, err
4545
}
4646

47-
// CountUser returns the user count which satisfies the given condition.
48-
func (d *GormDatabase) CountUser(condition ...any) (int64, error) {
49-
c := int64(-1)
47+
// GetUsers returns the users which satisfy the given condition.
48+
func (d *GormDatabase) GetUsers(condition ...any) ([]*model.User, error) {
49+
users := make([]*model.User, 0)
5050
handle := d.DB.Model(new(model.User))
5151
if len(condition) == 1 {
5252
handle = handle.Where(condition[0])
5353
} else if len(condition) > 1 {
5454
handle = handle.Where(condition[0], condition[1:]...)
5555
}
56-
err := handle.Count(&c).Error
57-
return c, err
58-
}
59-
60-
// GetUsers returns all users.
61-
func (d *GormDatabase) GetUsers() ([]*model.User, error) {
62-
var users []*model.User
63-
err := d.DB.Find(&users).Error
56+
err := handle.Find(&users).Error
57+
if err == gorm.ErrRecordNotFound {
58+
return nil, nil
59+
}
6460
return users, err
6561
}
6662

‎database/user_test.go‎

Lines changed: 7 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -19,9 +19,10 @@ func (s *DatabaseSuite) TestUser() {
1919
require.NoError(s.T(), err)
2020
assert.NotNil(s.T(), jmattheis, "on bootup the first user should be automatically created")
2121

22-
adminCount, err := s.db.CountUser("admin = ?", true)
22+
admins, err := s.db.GetUsers("admin = ?", true)
2323
require.NoError(s.T(), err)
24-
assert.Equal(s.T(), int64(1), adminCount, "there is initially one admin")
24+
assert.Len(s.T(), admins, 1)
25+
assert.True(s.T(), admins[0].Admin, "the admin user should be an admin")
2526

2627
users, err := s.db.GetUsers()
2728
require.NoError(s.T(), err)
@@ -31,9 +32,9 @@ func (s *DatabaseSuite) TestUser() {
3132
nicories := &model.User{Name: "nicories", Pass: []byte{1, 2, 3, 4}, Admin: false}
3233
s.db.CreateUser(nicories)
3334
assert.NotEqual(s.T(), 0, nicories.ID, "on create user a new id should be assigned")
34-
userCount, err := s.db.CountUser()
35+
users, err = s.db.GetUsers()
3536
require.NoError(s.T(), err)
36-
assert.Equal(s.T(), int64(2), userCount, "two users should exist")
37+
assert.Len(s.T(), users, 2, "two users should exist")
3738

3839
user, err = s.db.GetUserByName("nicories")
3940
require.NoError(s.T(), err)
@@ -58,9 +59,9 @@ func (s *DatabaseSuite) TestUser() {
5859
require.NoError(s.T(), err)
5960
assert.Len(s.T(), users, 2)
6061

61-
adminCount, err = s.db.CountUser(&model.User{Admin: true})
62+
admins, err = s.db.GetUsers(&model.User{Admin: true})
6263
require.NoError(s.T(), err)
63-
assert.Equal(s.T(), int64(2), adminCount, "two admins exist")
64+
assert.Len(s.T(), admins, 2, "two admins exist")
6465

6566
require.NoError(s.T(), s.db.DeleteUserByID(tom.ID))
6667
users, err = s.db.GetUsers()

‎plugin/manager.go‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,7 @@ import (
2323

2424
// The Database interface for encapsulating database access.
2525
type Database interface {
26-
GetUsers() ([]*model.User, error)
26+
GetUsers(condition ...any) ([]*model.User, error)
2727
GetPluginConfByUserAndPath(userid uint, path string) (*model.PluginConf, error)
2828
CreatePluginConf(p *model.PluginConf) error
2929
GetPluginConfByApplicationID(appid uint) (*model.PluginConf, error)

‎router/router.go‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -104,7 +104,7 @@ func Create(db *database.GormDatabase, vInfo *model.VersionInfo, conf *config.Co
104104
}
105105
sessionHandler := api.SessionAPI{DB: db, NotifyDeleted: streamHandler.NotifyDeletedClient, SecureCookie: conf.Server.SecureCookie, LocalAuthEnabled: conf.LocalAuthEnabled}
106106
userChangeNotifier := new(api.UserChangeNotifier)
107-
userHandler := api.UserAPI{DB: db, PasswordStrength: conf.PassStrength, UserChangeNotifier: userChangeNotifier, Registration: conf.Registration}
107+
userHandler := api.UserAPI{DB: *db, PasswordStrength: conf.PassStrength, UserChangeNotifier: userChangeNotifier, Registration: conf.Registration}
108108

109109
pluginManager, err := plugin.NewManager(db, conf.PluginsDir, g.Group("/plugin/:id/custom/"), streamHandler)
110110
if err != nil {

0 commit comments

Comments
 (0)