Skip to content

Commit d486c1b

Browse files
enhance: serialize user update methods
1 parent c21a7de commit d486c1b

6 files changed

Lines changed: 86 additions & 66 deletions

File tree

‎api/user.go‎

Lines changed: 51 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -9,19 +9,11 @@ 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

15-
// The UserDatabase interface for encapsulating database access.
16-
type UserDatabase interface {
17-
GetUsers() ([]*model.User, error)
18-
GetUserByID(id uint) (*model.User, error)
19-
GetUserByName(name string) (*model.User, error)
20-
DeleteUserByID(id uint) error
21-
UpdateUser(user *model.User) error
22-
CreateUser(user *model.User) error
23-
CountUser(condition ...any) (int64, error)
24-
}
16+
var errCannotDeleteLastAdmin = errors.New("cannot delete last admin")
2517

2618
// UserChangeNotifier notifies listeners for user changes.
2719
type UserChangeNotifier struct {
@@ -59,7 +51,7 @@ func (c *UserChangeNotifier) fireUserAdded(uid uint) error {
5951

6052
// The UserAPI provides handlers for managing users.
6153
type UserAPI struct {
62-
DB UserDatabase
54+
DB *database.GormDatabase
6355
PasswordStrength int
6456
UserChangeNotifier *UserChangeNotifier
6557
Registration bool
@@ -343,19 +335,26 @@ func (a *UserAPI) DeleteUserByID(ctx *gin.Context) {
343335
return
344336
}
345337
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
338+
for range 3 {
339+
err = a.DB.Txn(func(txdb *database.GormDatabase) error {
340+
if err := txdb.DeleteUserByID(id); err != nil {
341+
return err
342+
}
343+
anotherAdmin, err := txdb.GetUsers(&model.User{Admin: true})
344+
if err != nil {
345+
return err
346+
}
347+
if user.Admin && len(anotherAdmin) == 0 {
348+
ctx.AbortWithError(400, errCannotDeleteLastAdmin)
349+
return errCannotDeleteLastAdmin
350+
}
351+
return a.UserChangeNotifier.fireUserDeleted(id)
352+
})
353+
if err == nil || ctx.IsAborted() {
354+
return
355+
}
357356
}
358-
successOrAbort(ctx, 500, a.DB.DeleteUserByID(id))
357+
successOrAbort(ctx, 500, err)
359358
} else {
360359
ctx.AbortWithError(404, errors.New("user does not exist"))
361360
}
@@ -470,15 +469,7 @@ func (a *UserAPI) UpdateUserByID(ctx *gin.Context) {
470469
return
471470
}
472471
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-
472+
dbUserWasAdmin := dbUser.Admin
482473
dbUser.Name = updatedUser.Name
483474
dbUser.Admin = updatedUser.Admin
484475

@@ -494,10 +485,35 @@ func (a *UserAPI) UpdateUserByID(ctx *gin.Context) {
494485
}
495486
dbUser.Pass = pw
496487
}
497-
if success := successOrAbort(ctx, 500, a.DB.UpdateUser(dbUser)); !success {
498-
return
488+
489+
for range 3 {
490+
err = a.DB.Txn(func(txdb *database.GormDatabase) error {
491+
if err := txdb.UpdateUser(dbUser); err != nil {
492+
return err
493+
}
494+
495+
anotherAdmin, err := txdb.GetUsers(&model.User{Admin: true})
496+
if err != nil {
497+
return err
498+
}
499+
if !updatedUser.Admin && dbUserWasAdmin && len(anotherAdmin) == 0 {
500+
ctx.AbortWithError(400, errCannotDeleteLastAdmin)
501+
return errCannotDeleteLastAdmin
502+
}
503+
504+
return nil
505+
})
506+
507+
if ctx.IsAborted() {
508+
return
509+
}
510+
511+
if err == nil {
512+
ctx.JSON(200, toExternalUser(dbUser))
513+
return
514+
}
499515
}
500-
ctx.JSON(200, toExternalUser(dbUser))
516+
ctx.AbortWithError(500, err)
501517
} else {
502518
ctx.AbortWithError(404, errors.New("user does not exist"))
503519
}

‎api/user_test.go‎

Lines changed: 9 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,9 @@
11
package api
22

33
import (
4+
"encoding/json"
45
"errors"
6+
"io"
57
"net/http/httptest"
68
"strings"
79
"testing"
@@ -49,7 +51,7 @@ func (s *UserSuite) BeforeTest(suiteName, testName string) {
4951
s.notifiedAdd = true
5052
return nil
5153
})
52-
s.a = &UserAPI{DB: s.db, UserChangeNotifier: s.notifier}
54+
s.a = &UserAPI{DB: s.db.GormDatabase, UserChangeNotifier: s.notifier}
5355
}
5456

5557
func (s *UserSuite) AfterTest(suiteName, testName string) {
@@ -350,17 +352,6 @@ func (s *UserSuite) Test_UpdateUserByID_InvalidID() {
350352
assert.Equal(s.T(), 400, s.recorder.Code)
351353
}
352354

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-
364355
func (s *UserSuite) Test_UpdateUserByID_TooLongPassword_Expect400() {
365356
s.loginAdmin()
366357

@@ -412,6 +403,12 @@ func (s *UserSuite) Test_UpdateUserByID_UpdateNotPassword() {
412403
s.a.UpdateUserByID(s.ctx)
413404

414405
assert.Equal(s.T(), 200, s.recorder.Code)
406+
body, err := io.ReadAll(s.recorder.Body)
407+
require.NoError(s.T(), err)
408+
var retUser model.UserExternal
409+
require.NoError(s.T(), json.Unmarshal(body, &retUser))
410+
assert.Equal(s.T(), "tom", retUser.Name)
411+
assert.Equal(s.T(), true, retUser.Admin)
415412
user, err := s.db.GetUserByID(2)
416413
assert.NoError(s.T(), err)
417414
assert.NotNil(s.T(), user)

‎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)

0 commit comments

Comments
 (0)