diff --git a/cmd/setting/fix.go b/cmd/setting/fix.go index 2d57eb30..ae39aa92 100644 --- a/cmd/setting/fix.go +++ b/cmd/setting/fix.go @@ -27,7 +27,7 @@ var FixCmd = &cobra.Command{ _, err := s.Interface() if err != nil { fmt.Printf("setting %s, interface error: %v\n", k, err) - err = s.SetString(s.DefaultString()) + err = s.SetRaw(s.DefaultRaw()) if err != nil { errorCount++ fmt.Printf("setting %s fix error: %v\n", k, err) diff --git a/cmd/setting/set.go b/cmd/setting/set.go index 83e7a919..fff8198c 100644 --- a/cmd/setting/set.go +++ b/cmd/setting/set.go @@ -29,14 +29,14 @@ var SetCmd = &cobra.Command{ if !ok { return errors.New("setting not found") } - current := s.String() - err := s.SetString(args[1]) + current := s.Raw() + err := s.SetRaw(args[1]) if err != nil { - s.SetString(current) + s.SetRaw(current) fmt.Printf("set setting %s error: %v\n", args[0], err) } if v, err := s.Interface(); err != nil { - s.SetString(current) + s.SetRaw(current) fmt.Printf("set setting %s error: %v\n", args[0], err) } else { fmt.Printf("set setting success:\n%s: %v\n", args[0], v) diff --git a/internal/db/user.go b/internal/db/user.go index fef9abaa..a9dd85c4 100644 --- a/internal/db/user.go +++ b/internal/db/user.go @@ -61,13 +61,16 @@ func CreateOrLoadUser(username string, p provider.OAuth2Provider, puid uint, con return &user, nil } -func GetUserByProvider(p provider.OAuth2Provider, puid uint) (*model.User, error) { - u := &model.User{} - err := db.Preload("Providers", "provider = ? AND provider_user_id = ?", p, puid).First(u).Error - if errors.Is(err, gorm.ErrRecordNotFound) { - return u, errors.New("user not found") +func GetProviderUserID(p provider.OAuth2Provider, puid uint) (uint, error) { + var userProvider model.UserProvider + if err := db.Where("provider = ? AND provider_user_id = ?", p, puid).First(&userProvider).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return 0, errors.New("user not found") + } else { + return 0, err + } } - return u, err + return userProvider.UserID, nil } func AddUserToRoom(userID uint, roomID uint, role model.RoomRole, permission model.Permission) error { diff --git a/internal/op/users.go b/internal/op/users.go index 2b8002c4..447bf74a 100644 --- a/internal/op/users.go +++ b/internal/op/users.go @@ -76,6 +76,15 @@ func CreateOrLoadUser(username string, p provider.OAuth2Provider, pid uint, conf return u2, userCache.SetWithExpire(u.ID, u2, time.Hour) } +func GetUserByProvider(p provider.OAuth2Provider, pid uint) (*User, error) { + uid, err := db.GetProviderUserID(p, pid) + if err != nil { + return nil, err + } + + return GetUserById(uid) +} + func DeleteUserByID(userID uint) error { err := db.DeleteUserByID(userID) if err != nil { diff --git a/internal/settings/bool.go b/internal/settings/bool.go index f2cbfd22..12d47a84 100644 --- a/internal/settings/bool.go +++ b/internal/settings/bool.go @@ -49,7 +49,7 @@ func (b *Bool) Default() bool { return b.defaultValue } -func (b *Bool) DefaultString() string { +func (b *Bool) DefaultRaw() string { if b.defaultValue { return "1" } else { @@ -61,7 +61,7 @@ func (b *Bool) DefaultInterface() any { return b.Default() } -func (b *Bool) SetString(value string) error { +func (b *Bool) SetRaw(value string) error { if b.value == value { return nil } @@ -71,9 +71,9 @@ func (b *Bool) SetString(value string) error { func (b *Bool) Set(value bool) error { if value { - return b.SetString("1") + return b.SetRaw("1") } else { - return b.SetString("0") + return b.SetRaw("0") } } @@ -81,7 +81,7 @@ func (b *Bool) Get() (bool, error) { return b.value == "1", nil } -func (b *Bool) String() string { +func (b *Bool) Raw() string { return b.value } diff --git a/internal/settings/setting.go b/internal/settings/setting.go index 94530d33..673b16ed 100644 --- a/internal/settings/setting.go +++ b/internal/settings/setting.go @@ -18,9 +18,9 @@ type Setting interface { Type() model.SettingType Group() model.SettingGroup Init(string) - String() string - SetString(string) error - DefaultString() string + Raw() string + SetRaw(string) error + DefaultRaw() string DefaultInterface() any Interface() (any, error) } @@ -115,7 +115,7 @@ func initSettings(i ...Setting) error { for _, b := range i { s := &model.Setting{ Name: b.Name(), - Value: b.String(), + Value: b.Raw(), Type: b.Type(), Group: b.Group(), } diff --git a/internal/settings/var.go b/internal/settings/var.go index 7a63a6b5..374626fb 100644 --- a/internal/settings/var.go +++ b/internal/settings/var.go @@ -5,3 +5,7 @@ import "github.com/synctv-org/synctv/internal/model" var ( DisableCreateRoom = newBoolSetting("disable_create_room", false, model.SettingGroupRoom) ) + +var ( + DisableUserSignup = newBoolSetting("disable_user_signup", false, model.SettingGroupUser) +) diff --git a/server/oauth2/auth.go b/server/oauth2/auth.go index 2dda4b8a..a2b33708 100644 --- a/server/oauth2/auth.go +++ b/server/oauth2/auth.go @@ -8,6 +8,7 @@ import ( "github.com/synctv-org/synctv/internal/op" "github.com/synctv-org/synctv/internal/provider" "github.com/synctv-org/synctv/internal/provider/providers" + "github.com/synctv-org/synctv/internal/settings" "github.com/synctv-org/synctv/server/middlewares" "github.com/synctv-org/synctv/server/model" "github.com/synctv-org/synctv/utils" @@ -83,7 +84,17 @@ func OAuth2Callback(ctx *gin.Context) { return } - user, err := op.CreateOrLoadUser(ui.Username, p, ui.ProviderUserID) + disable, err := settings.DisableUserSignup.Get() + if err != nil { + ctx.AbortWithStatusJSON(http.StatusInternalServerError, model.NewApiErrorResp(err)) + return + } + var user *op.User + if disable { + user, err = op.GetUserByProvider(p, ui.ProviderUserID) + } else { + user, err = op.CreateOrLoadUser(ui.Username, p, ui.ProviderUserID) + } if err != nil { ctx.AbortWithStatusJSON(http.StatusInternalServerError, model.NewApiErrorResp(err)) return @@ -130,7 +141,17 @@ func OAuth2CallbackApi(ctx *gin.Context) { return } - user, err := op.CreateOrLoadUser(ui.Username, p, ui.ProviderUserID) + disable, err := settings.DisableUserSignup.Get() + if err != nil { + ctx.AbortWithStatusJSON(http.StatusInternalServerError, model.NewApiErrorResp(err)) + return + } + var user *op.User + if disable { + user, err = op.GetUserByProvider(p, ui.ProviderUserID) + } else { + user, err = op.CreateOrLoadUser(ui.Username, p, ui.ProviderUserID) + } if err != nil { ctx.AbortWithStatusJSON(http.StatusInternalServerError, model.NewApiErrorResp(err)) return