diff --git a/hscontrol/db/users.go b/hscontrol/db/users.go index 86ee936f..8e7b7223 100644 --- a/hscontrol/db/users.go +++ b/hscontrol/db/users.go @@ -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 +} diff --git a/hscontrol/db/users_test.go b/hscontrol/db/users_test.go index 2a755ff3..75b2ac1e 100644 --- a/hscontrol/db/users_test.go +++ b/hscontrol/db/users_test.go @@ -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