Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 33 additions & 0 deletions coderd/database/dbtestutil/tx.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import (
"testing"

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/xerrors"

"github.com/coder/coder/v2/coderd/database"
Expand Down Expand Up @@ -71,3 +72,35 @@ func (tx *DBTx) Done() error {
close(tx.done)
return <-tx.finalErr
}

var errRollbackTestTx = xerrors.New("roll back test transaction")

// StartRolledBackTx returns a Store bound to a transaction that is rolled back
// when the test finishes. It lets parallel subtests that only need query
// isolation share one database instead of creating one each.
func StartRolledBackTx(t testing.TB, db database.Store) database.Store {
t.Helper()
done := make(chan struct{})
txC := make(chan database.Store, 1)
errC := make(chan error, 1)

go func() {
errC <- db.InTx(func(tx database.Store) error {
txC <- tx
<-done
return errRollbackTestTx
}, nil)
}()

var tx database.Store
select {
case tx = <-txC:
case err := <-errC:
require.NoError(t, err, "start transaction")
}
t.Cleanup(func() {
close(done)
assert.ErrorIs(t, <-errC, errRollbackTestTx)
})
return tx
}
38 changes: 25 additions & 13 deletions coderd/database/querier_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -6213,6 +6213,10 @@ func TestGroupRemovalTrigger(t *testing.T) {
func TestGetUserStatusCounts(t *testing.T) {
t.Parallel()

// Every leaf subtest runs in its own rolled-back transaction, so one
// database serves the whole timezone x date matrix.
store, _ := dbtestutil.NewDB(t)

type testCase struct {
timezone string
location *time.Location
Expand Down Expand Up @@ -6269,7 +6273,7 @@ func TestGetUserStatusCounts(t *testing.T) {

t.Run("No Users", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
db := dbtestutil.StartRolledBackTx(t, store)
ctx := testutil.Context(t, testutil.WaitShort)

counts, err := db.GetUserStatusCounts(ctx, database.GetUserStatusCountsParams{
Expand Down Expand Up @@ -6305,7 +6309,7 @@ func TestGetUserStatusCounts(t *testing.T) {
for _, stc := range subTestCases {
t.Run(stc.name, func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
db := dbtestutil.StartRolledBackTx(t, store)
ctx := testutil.Context(t, testutil.WaitShort)

dbgen.User(t, db, database.User{
Expand Down Expand Up @@ -6486,7 +6490,7 @@ func TestGetUserStatusCounts(t *testing.T) {
for _, stc := range subTestCases {
t.Run(stc.name, func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
db := dbtestutil.StartRolledBackTx(t, store)
ctx := testutil.Context(t, testutil.WaitShort)

user := dbgen.User(t, db, database.User{
Expand Down Expand Up @@ -6621,7 +6625,7 @@ func TestGetUserStatusCounts(t *testing.T) {
t.Run(stc.name, func(t *testing.T) {
t.Parallel()

db, _ := dbtestutil.NewDB(t)
db := dbtestutil.StartRolledBackTx(t, store)
ctx := testutil.Context(t, testutil.WaitShort)

user1 := dbgen.User(t, db, database.User{
Expand Down Expand Up @@ -6700,7 +6704,7 @@ func TestGetUserStatusCounts(t *testing.T) {

t.Run("User precedes and survives query range", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
db := dbtestutil.StartRolledBackTx(t, store)
ctx := testutil.Context(t, testutil.WaitShort)

_ = dbgen.User(t, db, database.User{
Expand Down Expand Up @@ -6732,7 +6736,7 @@ func TestGetUserStatusCounts(t *testing.T) {

t.Run("User deleted before query range", func(t *testing.T) {
t.Parallel()
db, _, sqlDB := dbtestutil.NewDBWithSQLDB(t)
db := dbtestutil.StartRolledBackTx(t, store)
ctx := testutil.Context(t, testutil.WaitShort)

user := dbgen.User(t, db, database.User{
Expand All @@ -6741,10 +6745,14 @@ func TestGetUserStatusCounts(t *testing.T) {
UpdatedAt: userCreatedAt,
})

err := db.UpdateUserDeletedByID(ctx, user.ID)
// The deletion trigger records users.updated_at as deleted_at.
_, err := db.UpdateUserStatus(ctx, database.UpdateUserStatusParams{
ID: user.ID,
Status: user.Status,
UpdatedAt: tc.reportUntil,
})
require.NoError(t, err)

_, err = sqlDB.ExecContext(ctx, "UPDATE user_deleted SET deleted_at = $1 WHERE user_id = $2", tc.reportUntil, user.ID)
err = db.UpdateUserDeletedByID(ctx, user.ID)
require.NoError(t, err)

userStatusChanges, err := db.GetUserStatusCounts(ctx, database.GetUserStatusCountsParams{
Expand All @@ -6759,7 +6767,7 @@ func TestGetUserStatusCounts(t *testing.T) {
t.Run("User deleted during query range", func(t *testing.T) {
t.Parallel()

db, _, sqlDB := dbtestutil.NewDBWithSQLDB(t)
db := dbtestutil.StartRolledBackTx(t, store)
ctx := testutil.Context(t, testutil.WaitShort)

user := dbgen.User(t, db, database.User{
Expand All @@ -6768,10 +6776,14 @@ func TestGetUserStatusCounts(t *testing.T) {
UpdatedAt: userCreatedAt,
})

err := db.UpdateUserDeletedByID(ctx, user.ID)
// The deletion trigger records users.updated_at as deleted_at.
_, err := db.UpdateUserStatus(ctx, database.UpdateUserStatusParams{
ID: user.ID,
Status: user.Status,
UpdatedAt: tc.reportUntil,
})
require.NoError(t, err)

_, err = sqlDB.ExecContext(ctx, "UPDATE user_deleted SET deleted_at = $1 WHERE user_id = $2", tc.reportUntil, user.ID)
err = db.UpdateUserDeletedByID(ctx, user.ID)
require.NoError(t, err)

userStatusChanges, err := db.GetUserStatusCounts(ctx, database.GetUserStatusCountsParams{
Expand Down
Loading