Skip to content

Commit 3421a40

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

13 files changed

Lines changed: 205 additions & 158 deletions

File tree

‎api/oidc.go‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -519,7 +519,7 @@ func (a *OIDCAPI) registerUser(username, oidcID string, hasAdminGroup bool) (*mo
519519
return nil, http.StatusInternalServerError, fmt.Errorf("failed to create user: %w", err)
520520
}
521521
log.Info().Str("oidc_id", oidcID).Str("username", user.Name).Bool("admin", user.Admin).Msg("OIDC auto registration")
522-
if err := a.UserChangeNotifier.fireUserAdded(user.ID); err != nil {
522+
if err := a.UserChangeNotifier.fireUserAdded(a.DB, user.ID); err != nil {
523523
log.Error().Err(err).Uint("user_id", user.ID).Msg("Could not notify user change")
524524
}
525525
return user, 0, nil

‎api/oidc_test.go‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@ import (
1313

1414
"github.com/gin-gonic/gin"
1515
"github.com/gotify/server/v3/auth"
16+
"github.com/gotify/server/v3/database"
1617
"github.com/gotify/server/v3/decaymap"
1718
"github.com/gotify/server/v3/mode"
1819
"github.com/gotify/server/v3/model"
@@ -49,7 +50,7 @@ func (s *OIDCSuite) BeforeTest(suiteName, testName string) {
4950
s.db = testdb.NewDB(s.T())
5051
s.notified = false
5152
notifier := new(UserChangeNotifier)
52-
notifier.OnUserAdded(func(uint) error {
53+
notifier.OnUserAdded(func(tx *database.GormDatabase, uid uint) error {
5354
s.notified = true
5455
return nil
5556
})

‎api/plugin_test.go‎

Lines changed: 34 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,7 @@ type PluginSuite struct {
2929
suite.Suite
3030
db *testdb.Database
3131
a *PluginAPI
32+
u *UserAPI
3233
ctx *gin.Context
3334
recorder *httptest.ResponseRecorder
3435
manager *plugin.Manager
@@ -39,19 +40,20 @@ func (s *PluginSuite) BeforeTest(suiteName, testName string) {
3940
mode.Set(mode.TestDev)
4041
s.db = testdb.NewDB(s.T())
4142
s.resetRecorder()
42-
manager, err := plugin.NewManager(s.db, "", nil, s)
43+
manager, err := plugin.NewManager(s.db.GormDatabase, "", nil, s)
4344
assert.Nil(s.T(), err)
4445
s.manager = manager
4546
withURL(s.ctx, "http", "example.com")
4647
s.a = &PluginAPI{DB: s.db, Manager: manager, Notifier: s}
48+
s.u = &UserAPI{DB: s.db.GormDatabase, UserChangeNotifier: &UserChangeNotifier{}}
4749

4850
mockPluginCompat := new(mock.Plugin)
4951
assert.Nil(s.T(), s.manager.LoadPlugin(mockPluginCompat))
5052

51-
s.db.User(1)
52-
assert.Nil(s.T(), s.manager.InitializeForUserID(1))
53-
s.db.User(2)
54-
assert.Nil(s.T(), s.manager.InitializeForUserID(2))
53+
s.db.NewUserWithNameAdmin(1, "user1", true)
54+
assert.Nil(s.T(), s.manager.InitializeForUserID(s.db.GormDatabase, 1))
55+
s.db.NewUserWithNameAdmin(2, "user2", true)
56+
assert.Nil(s.T(), s.manager.InitializeForUserID(s.db.GormDatabase, 2))
5557

5658
s.db.CreatePluginConf(&model.PluginConf{
5759
UserID: 1,
@@ -97,6 +99,31 @@ func (s *PluginSuite) Test_GetPlugins() {
9799
assert.False(s.T(), pluginConfs[0].Enabled, "Plugins should be disabled by default")
98100
}
99101

102+
func (s *PluginSuite) Test_DeleteUser() {
103+
test.WithUser(s.ctx, 1)
104+
105+
s.ctx.Request = httptest.NewRequest("POST", "/plugin/1/enable", nil)
106+
s.ctx.Params = gin.Params{{Key: "id", Value: "1"}}
107+
s.a.EnablePlugin(s.ctx)
108+
109+
assert.Equal(s.T(), 200, s.recorder.Code)
110+
111+
if pluginConf, err := s.db.GetPluginConfByUserAndPath(1, mock.ModulePath); assert.NoError(s.T(), err) {
112+
assert.True(s.T(), pluginConf.Enabled)
113+
}
114+
s.resetRecorder()
115+
116+
s.ctx.Request = httptest.NewRequest("DELETE", "/user/1", nil)
117+
s.ctx.Params = gin.Params{{Key: "id", Value: "1"}}
118+
s.u.DeleteUserByID(s.ctx)
119+
120+
assert.Equal(s.T(), 200, s.recorder.Code)
121+
122+
user, err := s.db.GetUserByID(1)
123+
assert.NoError(s.T(), err)
124+
assert.Nil(s.T(), user)
125+
}
126+
100127
func (s *PluginSuite) Test_EnableDisablePlugin() {
101128
{
102129
test.WithUser(s.ctx, 1)
@@ -161,7 +188,7 @@ func (s *PluginSuite) Test_EnableDisablePlugin() {
161188

162189
func (s *PluginSuite) Test_EnableDisablePlugin_EnableReturnsError_expect500() {
163190
s.db.User(16)
164-
assert.Nil(s.T(), s.manager.InitializeForUserID(16))
191+
assert.Nil(s.T(), s.manager.InitializeForUserID(s.db.GormDatabase, 16))
165192
mock.ReturnErrorOnEnableForUser(16, errors.New("test error"))
166193
conf, err := s.db.GetPluginConfByUserAndPath(16, mock.ModulePath)
167194
assert.NoError(s.T(), err)
@@ -183,7 +210,7 @@ func (s *PluginSuite) Test_EnableDisablePlugin_EnableReturnsError_expect500() {
183210

184211
func (s *PluginSuite) Test_EnableDisablePlugin_DisableReturnsError_expect500() {
185212
s.db.User(17)
186-
assert.Nil(s.T(), s.manager.InitializeForUserID(17))
213+
assert.Nil(s.T(), s.manager.InitializeForUserID(s.db.GormDatabase, 17))
187214
mock.ReturnErrorOnDisableForUser(17, errors.New("test error"))
188215
conf, err := s.db.GetPluginConfByUserAndPath(17, mock.ModulePath)
189216
assert.NoError(s.T(), err)

‎api/stream/stream.go‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@ import (
1111
"github.com/gin-gonic/gin"
1212
"github.com/gorilla/websocket"
1313
"github.com/gotify/server/v3/auth"
14+
"github.com/gotify/server/v3/database"
1415
"github.com/gotify/server/v3/model"
1516
)
1617

@@ -50,7 +51,7 @@ func (a *API) CollectConnectedClientTokens() []string {
5051
}
5152

5253
// NotifyDeletedUser closes existing connections for the given user.
53-
func (a *API) NotifyDeletedUser(userID uint) error {
54+
func (a *API) NotifyDeletedUser(tx *database.GormDatabase, userID uint) error {
5455
a.lock.Lock()
5556
defer a.lock.Unlock()
5657
if clients, ok := a.clients[userID]; ok {

‎api/stream/stream_test.go‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -321,7 +321,7 @@ func TestDeleteUser(t *testing.T) {
321321
expectNoMessage(userTwo...)
322322
expectNoMessage(userThree...)
323323

324-
api.NotifyDeletedUser(1)
324+
api.NotifyDeletedUser(nil, 1)
325325

326326
api.Notify(1, &model.MessageExternal{ID: 2, Message: "there"})
327327
expectNoMessage(userOne...)

‎api/user.go‎

Lines changed: 67 additions & 48 deletions
Original file line numberDiff line numberDiff line change
@@ -9,48 +9,40 @@ 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 {
28-
userDeletedCallbacks []func(uid uint) error
29-
userAddedCallbacks []func(uid uint) error
20+
userDeletedCallbacks []func(tx *database.GormDatabase, uid uint) error
21+
userAddedCallbacks []func(tx *database.GormDatabase, uid uint) error
3022
}
3123

3224
// OnUserDeleted is called on user deletion.
33-
func (c *UserChangeNotifier) OnUserDeleted(cb func(uid uint) error) {
25+
func (c *UserChangeNotifier) OnUserDeleted(cb func(tx *database.GormDatabase, uid uint) error) {
3426
c.userDeletedCallbacks = append(c.userDeletedCallbacks, cb)
3527
}
3628

3729
// OnUserAdded is called on user creation.
38-
func (c *UserChangeNotifier) OnUserAdded(cb func(uid uint) error) {
30+
func (c *UserChangeNotifier) OnUserAdded(cb func(tx *database.GormDatabase, uid uint) error) {
3931
c.userAddedCallbacks = append(c.userAddedCallbacks, cb)
4032
}
4133

42-
func (c *UserChangeNotifier) fireUserDeleted(uid uint) error {
34+
func (c *UserChangeNotifier) fireUserDeleted(tx *database.GormDatabase, uid uint) error {
4335
for _, cb := range c.userDeletedCallbacks {
44-
if err := cb(uid); err != nil {
36+
if err := cb(tx, uid); err != nil {
4537
return err
4638
}
4739
}
4840
return nil
4941
}
5042

51-
func (c *UserChangeNotifier) fireUserAdded(uid uint) error {
43+
func (c *UserChangeNotifier) fireUserAdded(tx *database.GormDatabase, uid uint) error {
5244
for _, cb := range c.userAddedCallbacks {
53-
if err := cb(uid); err != nil {
45+
if err := cb(tx, uid); err != nil {
5446
return err
5547
}
5648
}
@@ -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
@@ -233,11 +225,14 @@ func (a *UserAPI) CreateUser(ctx *gin.Context) {
233225
}
234226

235227
if existingUser == nil {
236-
if success := successOrAbort(ctx, 500, a.DB.CreateUser(internal)); !success {
237-
return
238-
}
239-
if err := a.UserChangeNotifier.fireUserAdded(internal.ID); err != nil {
240-
ctx.AbortWithError(500, err)
228+
// this should not cause conflicts, so no need to retry
229+
err = a.DB.Txn(func(txdb *database.GormDatabase) error {
230+
if err := txdb.CreateUser(internal); err != nil {
231+
return err
232+
}
233+
return a.UserChangeNotifier.fireUserAdded(txdb, internal.ID)
234+
})
235+
if success := successOrAbort(ctx, 500, err); !success {
241236
return
242237
}
243238
ctx.JSON(200, toExternalUser(internal))
@@ -343,19 +338,26 @@ func (a *UserAPI) DeleteUserByID(ctx *gin.Context) {
343338
return
344339
}
345340
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
341+
for range 3 {
342+
err = a.DB.Txn(func(txdb *database.GormDatabase) error {
343+
if err := txdb.DeleteUserByID(id); err != nil {
344+
return err
345+
}
346+
anotherAdmin, err := txdb.GetUsers(&model.User{Admin: true})
347+
if err != nil {
348+
return err
349+
}
350+
if user.Admin && len(anotherAdmin) == 0 {
351+
ctx.AbortWithError(400, errCannotDeleteLastAdmin)
352+
return errCannotDeleteLastAdmin
353+
}
354+
return a.UserChangeNotifier.fireUserDeleted(txdb, id)
355+
})
356+
if err == nil || ctx.IsAborted() {
357+
return
358+
}
357359
}
358-
successOrAbort(ctx, 500, a.DB.DeleteUserByID(id))
360+
successOrAbort(ctx, 500, err)
359361
} else {
360362
ctx.AbortWithError(404, errors.New("user does not exist"))
361363
}
@@ -470,15 +472,7 @@ func (a *UserAPI) UpdateUserByID(ctx *gin.Context) {
470472
return
471473
}
472474
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-
475+
dbUserWasAdmin := dbUser.Admin
482476
dbUser.Name = updatedUser.Name
483477
dbUser.Admin = updatedUser.Admin
484478

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

0 commit comments

Comments
 (0)