diff --git a/internal/store/postgres/user_repository_test.go b/internal/store/postgres/user_repository_test.go index f00aa0940..6b23c9ddf 100644 --- a/internal/store/postgres/user_repository_test.go +++ b/internal/store/postgres/user_repository_test.go @@ -366,8 +366,8 @@ func (s *UserRepositoryTestSuite) TestUpdateByEmail() { s.Run(tc.Description, func() { got, err := s.repository.UpdateByEmail(s.ctx, tc.UserToUpdate) if tc.Err != nil && tc.Err.Error() != "" { - if errors.Unwrap(err) == tc.Err { - s.T().Fatalf("got error %s, expected was %s", err.Error(), tc.Err) + if !errors.Is(err, tc.Err) { + s.T().Fatalf("got error %v, expected was %v", err, tc.Err) } } @@ -451,8 +451,8 @@ func (s *UserRepositoryTestSuite) TestUpdateByID() { s.Run(tc.Description, func() { got, err := s.repository.UpdateByID(s.ctx, tc.UserToUpdate) if tc.Err != nil && tc.Err.Error() != "" { - if errors.Unwrap(err) == tc.Err { - s.T().Fatalf("got error %s, expected was %s", err.Error(), tc.Err) + if !errors.Is(err, tc.Err) { + s.T().Fatalf("got error %v, expected was %v", err, tc.Err) } } // TODO(kushsharma): remove metadata field from ignore once metadata is refactored @@ -559,8 +559,8 @@ func (s *UserRepositoryTestSuite) TestGetByName() { s.Run(tc.Description, func() { got, err := s.repository.GetByName(s.ctx, tc.Name) if tc.Err != nil && tc.Err.Error() != "" { - if errors.Unwrap(err) == tc.Err { - s.T().Fatalf("got error %s, expected was %s", err.Error(), tc.Err) + if !errors.Is(err, tc.Err) { + s.T().Fatalf("got error %v, expected was %v", err, tc.Err) } } if !cmp.Equal(got, tc.ExpectedUser) { @@ -610,8 +610,8 @@ func (s *UserRepositoryTestSuite) TestUpdateByName() { s.Run(tc.Description, func() { got, err := s.repository.UpdateByName(s.ctx, tc.UserToUpdate) if tc.Err != nil && tc.Err.Error() != "" { - if errors.Unwrap(err) == tc.Err { - s.T().Fatalf("got error %s, expected was %s", err.Error(), tc.Err) + if !errors.Is(err, tc.Err) { + s.T().Fatalf("got error %v, expected was %v", err, tc.Err) } } diff --git a/pkg/db/db.go b/pkg/db/db.go index 92f00e97c..3d179019b 100644 --- a/pkg/db/db.go +++ b/pkg/db/db.go @@ -76,21 +76,14 @@ func (c Client) WithTxn(ctx context.Context, txnOptions sql.TxOptions, txFunc fu defer func() { if p := recover(); p != nil { - switch p := p.(type) { - case error: - err = p - default: - err = fmt.Errorf("%s", p) - } - err = txn.Rollback() + _ = txn.Rollback() panic(p) } else if err != nil { - if rlbErr := txn.Rollback(); err != nil { - err = fmt.Errorf("rollback error: %s while executing: %w", rlbErr, err) + if rlbErr := txn.Rollback(); rlbErr != nil { + err = fmt.Errorf("rollback error: %w while executing: %w", rlbErr, err) } else { err = fmt.Errorf("rollback: %w", err) } - err = fmt.Errorf("rollback: %w", err) } else { err = txn.Commit() } diff --git a/pkg/db/db_test.go b/pkg/db/db_test.go new file mode 100644 index 000000000..2f3a11128 --- /dev/null +++ b/pkg/db/db_test.go @@ -0,0 +1,141 @@ +package db + +import ( + "context" + "database/sql" + "database/sql/driver" + "errors" + "testing" + + "github.com/jmoiron/sqlx" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// fakeConn is a minimal database/sql/driver connection, in the same style as +// failConnector in internal/store/postgres/fakes_test.go. It hands out fake +// transactions, records commits and rollbacks, and can be told to fail any of +// them, so WithTxn can be tested without a real database. +type fakeConn struct { + beginErr error + commitErr error + rollbackErr error + commits int + rollbacks int +} + +func (c *fakeConn) Prepare(string) (driver.Stmt, error) { return nil, errors.New("not implemented") } +func (c *fakeConn) Close() error { return nil } +func (c *fakeConn) Begin() (driver.Tx, error) { + if c.beginErr != nil { + return nil, c.beginErr + } + return &fakeTx{conn: c}, nil +} + +type fakeTx struct{ conn *fakeConn } + +func (t *fakeTx) Commit() error { t.conn.commits++; return t.conn.commitErr } +func (t *fakeTx) Rollback() error { t.conn.rollbacks++; return t.conn.rollbackErr } + +type fakeConnector struct{ conn *fakeConn } + +func (f fakeConnector) Connect(context.Context) (driver.Conn, error) { return f.conn, nil } +func (f fakeConnector) Driver() driver.Driver { return nil } + +func newFakeClient(t *testing.T, conn *fakeConn) Client { + t.Helper() + client := Client{DB: sqlx.NewDb(sql.OpenDB(fakeConnector{conn: conn}), "postgres")} + t.Cleanup(func() { _ = client.Close() }) + return client +} + +func TestWithTxn(t *testing.T) { + t.Run("commits when the callback succeeds", func(t *testing.T) { + conn := &fakeConn{} + client := newFakeClient(t, conn) + + err := client.WithTxn(context.Background(), sql.TxOptions{}, func(*sqlx.Tx) error { + return nil + }) + + require.NoError(t, err) + assert.Equal(t, 1, conn.commits) + assert.Equal(t, 0, conn.rollbacks) + }) + + t.Run("returns the begin error without running the callback", func(t *testing.T) { + beginErr := errors.New("begin failed") + client := newFakeClient(t, &fakeConn{beginErr: beginErr}) + + called := false + err := client.WithTxn(context.Background(), sql.TxOptions{}, func(*sqlx.Tx) error { + called = true + return nil + }) + + assert.ErrorIs(t, err, beginErr) + assert.False(t, called) + }) + + t.Run("wraps the callback error once when rollback succeeds", func(t *testing.T) { + callbackErr := errors.New("insert failed") + conn := &fakeConn{} + client := newFakeClient(t, conn) + + err := client.WithTxn(context.Background(), sql.TxOptions{}, func(*sqlx.Tx) error { + return callbackErr + }) + + assert.EqualError(t, err, "rollback: insert failed") + assert.ErrorIs(t, err, callbackErr) + assert.Equal(t, 1, conn.rollbacks) + assert.Equal(t, 0, conn.commits) + }) + + t.Run("reports both errors when rollback fails", func(t *testing.T) { + callbackErr := errors.New("insert failed") + rollbackErr := errors.New("connection lost") + conn := &fakeConn{rollbackErr: rollbackErr} + client := newFakeClient(t, conn) + + err := client.WithTxn(context.Background(), sql.TxOptions{}, func(*sqlx.Tx) error { + return callbackErr + }) + + assert.EqualError(t, err, "rollback error: connection lost while executing: insert failed") + assert.ErrorIs(t, err, callbackErr) + assert.ErrorIs(t, err, rollbackErr) + assert.Equal(t, 1, conn.rollbacks) + }) + + t.Run("returns the commit error", func(t *testing.T) { + commitErr := errors.New("commit failed") + conn := &fakeConn{commitErr: commitErr} + client := newFakeClient(t, conn) + + err := client.WithTxn(context.Background(), sql.TxOptions{}, func(*sqlx.Tx) error { + return nil + }) + + assert.ErrorIs(t, err, commitErr) + assert.Equal(t, 1, conn.commits) + }) + + t.Run("rolls back and repanics when the callback panics", func(t *testing.T) { + conn := &fakeConn{} + client := newFakeClient(t, conn) + + defer func() { + p := recover() + require.NotNil(t, p) + assert.Equal(t, "boom", p) + assert.Equal(t, 1, conn.rollbacks) + assert.Equal(t, 0, conn.commits) + }() + + _ = client.WithTxn(context.Background(), sql.TxOptions{}, func(*sqlx.Tx) error { + panic("boom") + }) + }) +}