mirror of
https://github.com/juanfont/headscale.git
synced 2026-09-26 10:14:52 +09:00
db: return non-NotFound errors from user lookups (#3480)
This commit is contained in:
+17
-18
@@ -122,15 +122,7 @@ func (hsdb *HSDatabase) GetUserByID(uid types.UserID) (*types.User, error) {
|
||||
}
|
||||
|
||||
func GetUserByID(tx *gorm.DB, uid types.UserID) (*types.User, error) {
|
||||
user := types.User{}
|
||||
if result := tx.First(&user, "id = ?", uid); errors.Is(
|
||||
result.Error,
|
||||
gorm.ErrRecordNotFound,
|
||||
) {
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
|
||||
return &user, nil
|
||||
return firstUser(tx, "id = ?", uid)
|
||||
}
|
||||
|
||||
func (hsdb *HSDatabase) GetUserByOIDCIdentifier(id string) (*types.User, error) {
|
||||
@@ -140,15 +132,7 @@ func (hsdb *HSDatabase) GetUserByOIDCIdentifier(id string) (*types.User, error)
|
||||
}
|
||||
|
||||
func GetUserByOIDCIdentifier(tx *gorm.DB, id string) (*types.User, error) {
|
||||
user := types.User{}
|
||||
if result := tx.First(&user, "provider_identifier = ?", id); errors.Is(
|
||||
result.Error,
|
||||
gorm.ErrRecordNotFound,
|
||||
) {
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
|
||||
return &user, nil
|
||||
return firstUser(tx, "provider_identifier = ?", id)
|
||||
}
|
||||
|
||||
func (hsdb *HSDatabase) ListUsers(filter *types.User) ([]types.User, error) {
|
||||
@@ -230,3 +214,18 @@ func (hsdb *HSDatabase) CreateUsersForTest(count int, namePrefix ...string) []*t
|
||||
|
||||
return users
|
||||
}
|
||||
|
||||
func firstUser(tx *gorm.DB, query string, arg any) (*types.User, error) {
|
||||
user := types.User{}
|
||||
|
||||
err := tx.First(&user, query, arg).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
@@ -214,6 +214,58 @@ func TestDestroyUserErrors(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetUserErrorPropagation(t *testing.T) {
|
||||
lookups := []struct {
|
||||
name string
|
||||
get func(*HSDatabase) (*types.User, error)
|
||||
}{
|
||||
{
|
||||
name: "by_id",
|
||||
get: func(db *HSDatabase) (*types.User, error) { return db.GetUserByID(1) },
|
||||
},
|
||||
{
|
||||
name: "by_oidc_identifier",
|
||||
get: func(db *HSDatabase) (*types.User, error) { return db.GetUserByOIDCIdentifier("oidc-id") },
|
||||
},
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
closeDB bool
|
||||
wantErr error
|
||||
}{
|
||||
{name: "missing_row_is_not_found", wantErr: ErrUserNotFound},
|
||||
{name: "query_failure_is_returned", closeDB: true},
|
||||
}
|
||||
|
||||
for _, lookup := range lookups {
|
||||
for _, tt := range tests {
|
||||
t.Run(lookup.name+"/"+tt.name, func(t *testing.T) {
|
||||
db, err := newSQLiteTestDB()
|
||||
require.NoError(t, err)
|
||||
|
||||
if tt.closeDB {
|
||||
sqlDB, err := db.DB.DB()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, sqlDB.Close())
|
||||
}
|
||||
|
||||
user, err := lookup.get(db)
|
||||
|
||||
// A swallowed error surfaces as a zero user and a nil error.
|
||||
require.Error(t, err)
|
||||
assert.Nil(t, user)
|
||||
|
||||
if tt.wantErr != nil {
|
||||
assert.ErrorIs(t, err, tt.wantErr)
|
||||
} else {
|
||||
assert.NotErrorIs(t, err, ErrUserNotFound)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenameUser(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
|
||||
Reference in New Issue
Block a user