mirror of
https://github.com/henrygd/beszel.git
synced 2026-09-21 17:07:47 +02:00
force oauth created users to start with user role
This commit is contained in:
@@ -106,6 +106,7 @@ func (h *Hub) StartHub() error {
|
|||||||
|
|
||||||
// TODO: move to users package
|
// TODO: move to users package
|
||||||
// handle default values for user / user_settings creation
|
// handle default values for user / user_settings creation
|
||||||
|
h.App.OnRecordAuthWithOAuth2Request("users").BindFunc(h.um.InitializeOAuthUserRole)
|
||||||
h.App.OnRecordCreate("users").BindFunc(h.um.InitializeUserRole)
|
h.App.OnRecordCreate("users").BindFunc(h.um.InitializeUserRole)
|
||||||
h.App.OnRecordCreate("user_settings").BindFunc(h.um.InitializeUserSettings)
|
h.App.OnRecordCreate("user_settings").BindFunc(h.um.InitializeUserSettings)
|
||||||
|
|
||||||
|
|||||||
100
internal/users/oauth_test.go
Normal file
100
internal/users/oauth_test.go
Normal file
@@ -0,0 +1,100 @@
|
|||||||
|
//go:build testing
|
||||||
|
|
||||||
|
package users_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
beszelTests "github.com/henrygd/beszel/internal/tests"
|
||||||
|
"github.com/pocketbase/pocketbase/apis"
|
||||||
|
"github.com/pocketbase/pocketbase/core"
|
||||||
|
"github.com/pocketbase/pocketbase/tools/auth"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"golang.org/x/oauth2"
|
||||||
|
)
|
||||||
|
|
||||||
|
type roleTestProvider struct {
|
||||||
|
auth.BaseProvider
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *roleTestProvider) FetchToken(string, ...oauth2.AuthCodeOption) (*oauth2.Token, error) {
|
||||||
|
return &oauth2.Token{AccessToken: "test-token"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *roleTestProvider) FetchAuthUser(*oauth2.Token) (*auth.AuthUser, error) {
|
||||||
|
return &auth.AuthUser{Id: "role-test-user", Email: "oauth@example.com"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOAuthUserRole(t *testing.T) {
|
||||||
|
t.Setenv("USER_CREATION", "true")
|
||||||
|
const provider = "beszel-role-test"
|
||||||
|
auth.Providers[provider] = func() auth.Provider { return &roleTestProvider{} }
|
||||||
|
t.Cleanup(func() { delete(auth.Providers, provider) })
|
||||||
|
|
||||||
|
for _, createData := range []string{`{}`, `{"role":"admin"}`, `{"role":"readonly"}`} {
|
||||||
|
t.Run(createData, func(t *testing.T) {
|
||||||
|
h, err := beszelTests.NewTestHub(t.TempDir())
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer h.Cleanup()
|
||||||
|
h.StartHub()
|
||||||
|
|
||||||
|
collection, err := h.FindCollectionByNameOrId("users")
|
||||||
|
require.NoError(t, err)
|
||||||
|
collection.OAuth2.Enabled = true
|
||||||
|
collection.OAuth2.Providers = []core.OAuth2ProviderConfig{{
|
||||||
|
Name: provider, ClientId: "test-client", ClientSecret: "test-secret",
|
||||||
|
}}
|
||||||
|
require.NoError(t, h.Save(collection))
|
||||||
|
r, err := apis.NewRouter(h.App)
|
||||||
|
require.NoError(t, err)
|
||||||
|
mux, err := r.BuildMux()
|
||||||
|
require.NoError(t, err)
|
||||||
|
login := func() {
|
||||||
|
body := `{"provider":"` + provider + `","code":"test-code","codeVerifier":"test-verifier","redirectUrl":"http://localhost/callback","createData":` + createData + `}`
|
||||||
|
req := httptest.NewRequest("POST", "/api/collections/users/auth-with-oauth2", strings.NewReader(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
res := httptest.NewRecorder()
|
||||||
|
mux.ServeHTTP(res, req)
|
||||||
|
require.Equal(t, 200, res.Code, res.Body.String())
|
||||||
|
}
|
||||||
|
login()
|
||||||
|
user, err := h.FindAuthRecordByEmail("users", "oauth@example.com")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, "user", user.GetString("role"))
|
||||||
|
|
||||||
|
// A later OAuth login must preserve a role assigned by an administrator.
|
||||||
|
user.Set("role", "admin")
|
||||||
|
require.NoError(t, h.Save(user))
|
||||||
|
login()
|
||||||
|
user, err = h.FindRecordById("users", user.Id)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, "admin", user.GetString("role"))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInternalUserRole(t *testing.T) {
|
||||||
|
for _, role := range []string{"", "user", "admin", "readonly"} {
|
||||||
|
t.Run("role="+role, func(t *testing.T) {
|
||||||
|
h, err := beszelTests.NewTestHub(t.TempDir())
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer h.Cleanup()
|
||||||
|
h.StartHub()
|
||||||
|
collection, err := h.FindCollectionByNameOrId("users")
|
||||||
|
require.NoError(t, err)
|
||||||
|
user := core.NewRecord(collection)
|
||||||
|
user.SetEmail("internal@example.com")
|
||||||
|
user.SetPassword("password12345")
|
||||||
|
user.Set("role", role)
|
||||||
|
require.NoError(t, h.Save(user))
|
||||||
|
user, err = h.FindRecordById("users", user.Id)
|
||||||
|
require.NoError(t, err)
|
||||||
|
if role == "" {
|
||||||
|
role = "user"
|
||||||
|
}
|
||||||
|
require.Equal(t, role, user.GetString("role"))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -24,6 +24,17 @@ func NewUserManager(app core.App) *UserManager {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// InitializeOAuthUserRole prevents self-registration from assigning a privileged role.
|
||||||
|
func (um *UserManager) InitializeOAuthUserRole(e *core.RecordAuthWithOAuth2RequestEvent) error {
|
||||||
|
if e.IsNewRecord {
|
||||||
|
if e.CreateData == nil {
|
||||||
|
e.CreateData = make(map[string]any)
|
||||||
|
}
|
||||||
|
e.CreateData["role"] = "user"
|
||||||
|
}
|
||||||
|
return e.Next()
|
||||||
|
}
|
||||||
|
|
||||||
// Initialize user role if not set
|
// Initialize user role if not set
|
||||||
func (um *UserManager) InitializeUserRole(e *core.RecordEvent) error {
|
func (um *UserManager) InitializeUserRole(e *core.RecordEvent) error {
|
||||||
if e.Record.GetString("role") == "" {
|
if e.Record.GetString("role") == "" {
|
||||||
|
|||||||
@@ -108,6 +108,9 @@ func TestCreateFirstUserAtomic(t *testing.T) {
|
|||||||
count, err := h.CountRecords("users")
|
count, err := h.CountRecords("users")
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.EqualValues(t, 1, count)
|
require.EqualValues(t, 1, count)
|
||||||
|
bootstrapUsers, err := h.FindAllRecords("users")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, "admin", bootstrapUsers[0].GetString("role"))
|
||||||
admins, err := h.FindAllRecords(core.CollectionNameSuperusers)
|
admins, err := h.FindAllRecords(core.CollectionNameSuperusers)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Len(t, admins, 1)
|
require.Len(t, admins, 1)
|
||||||
|
|||||||
Reference in New Issue
Block a user