From 8d697eda5ba6157c4b7b2fc84b76f80c2b3c3b2d Mon Sep 17 00:00:00 2001 From: Elijah Duffy Date: Mon, 29 Jun 2026 17:49:57 -0700 Subject: [PATCH] feat: add delete, context, and tx helpers Add Delete/DeleteAffected, context-aware CRUD helpers, WithTx, ContainsCol, and queryReturning deduplication for returning statements. Co-authored-by: Cursor --- doc.go | 7 ++- exec.go | 82 +++++++++++++++++++++++++++++----- exec_test.go | 99 +++++++++++++++++++++++++++++++++++++++++ go.mod | 9 +++- go.sum | 16 +++++++ jet_columns.go | 11 ++++- jet_columns_test.go | 105 ++++++++++++++++++++++++++++++++++++++++++++ query.go | 33 +++++++++++--- query_test.go | 80 +++++++++++++++++++++++++++++++++ testhelpers_test.go | 72 ++++++++++++++++++++++++++++++ tx.go | 26 +++++++++++ 11 files changed, 517 insertions(+), 23 deletions(-) create mode 100644 exec_test.go create mode 100644 jet_columns_test.go create mode 100644 query_test.go create mode 100644 testhelpers_test.go create mode 100644 tx.go diff --git a/doc.go b/doc.go index ff93c15..75d4cbb 100644 --- a/doc.go +++ b/doc.go @@ -7,14 +7,13 @@ // // # Query and mutation helpers // -// [Fetch], [FetchOne], [Insert], and [Update] wrap Jet [Statement] execution -// with consistent error semantics. Must* variants return a caller-provided error -// when no rows are found. +// [Fetch], [FetchOne], [Insert], [Update], [Delete], and their Must* and Context +// variants wrap Jet [Statement] execution with consistent error semantics. // // # Jet column utilities // // Dialect-neutral type aliases ([Column], [ColumnList]) and helpers for column -// lists ([NormalCols]) and expression building ([ExprValues]). +// lists ([NormalCols], [ContainsCol]) and expression building ([ExprValues]). // // Partial-update helpers ([ApplyPtr], [ApplyVal]) track changed fields for // repository patch logic. diff --git a/exec.go b/exec.go index d9e8a64..c93ee57 100644 --- a/exec.go +++ b/exec.go @@ -1,11 +1,19 @@ package dbx -import "errors" +import ( + "context" + "errors" +) // Insert executes an insert statement, returning the last inserted ID or an // error if the insert fails. func Insert(sqlo Executable, stmt Statement) (uint64, error) { - res, err := stmt.Exec(sqlo) + return InsertContext(context.Background(), sqlo, stmt) +} + +// InsertContext is the context-aware variant of [Insert]. +func InsertContext(ctx context.Context, sqlo Executable, stmt Statement) (uint64, error) { + res, err := stmt.ExecContext(ctx, sqlo) if err != nil { return 0, err } @@ -26,24 +34,34 @@ func Insert(sqlo Executable, stmt Statement) (uint64, error) { // The statement MUST be a Jet InsertStatement with a RETURNING clause. Returns // the inserted row object T or an error if the insert fails or no rows are returned. func InsertReturning[T any](sqlo Queryable, stmt Statement) (*T, error) { - var result T - err := stmt.Query(sqlo, &result) - if err != nil { - return nil, err - } - return &result, nil + return InsertReturningContext[T](context.Background(), sqlo, stmt) +} + +// InsertReturningContext is the context-aware variant of [InsertReturning]. +func InsertReturningContext[T any](ctx context.Context, sqlo Queryable, stmt Statement) (*T, error) { + return queryReturningContext[T](ctx, sqlo, stmt) } // Update executes an update statement, returning an error if the update fails. func Update(sqlo Executable, stmt Statement) error { - _, err := stmt.Exec(sqlo) + return UpdateContext(context.Background(), sqlo, stmt) +} + +// UpdateContext is the context-aware variant of [Update]. +func UpdateContext(ctx context.Context, sqlo Executable, stmt Statement) error { + _, err := stmt.ExecContext(ctx, sqlo) return err } // UpdateAffected executes an update statement and returns the number of rows // affected and an error if any. func UpdateAffected(sqlo Executable, stmt Statement) (int64, error) { - res, err := stmt.Exec(sqlo) + return UpdateAffectedContext(context.Background(), sqlo, stmt) +} + +// UpdateAffectedContext is the context-aware variant of [UpdateAffected]. +func UpdateAffectedContext(ctx context.Context, sqlo Executable, stmt Statement) (int64, error) { + res, err := stmt.ExecContext(ctx, sqlo) if err != nil { return 0, err } @@ -60,9 +78,49 @@ func UpdateAffected(sqlo Executable, stmt Statement) (int64, error) { // The statement MUST be a Jet UpdateStatement with a RETURNING clause. Returns // the updated row object T or an error if the update fails or no rows are returned. func UpdateReturning[T any](sqlo Queryable, stmt Statement) (*T, error) { - var result T - err := stmt.Query(sqlo, &result) + return UpdateReturningContext[T](context.Background(), sqlo, stmt) +} + +// UpdateReturningContext is the context-aware variant of [UpdateReturning]. +func UpdateReturningContext[T any](ctx context.Context, sqlo Queryable, stmt Statement) (*T, error) { + return queryReturningContext[T](ctx, sqlo, stmt) +} + +// Delete executes a delete statement, returning an error if the delete fails. +func Delete(sqlo Executable, stmt Statement) error { + return DeleteContext(context.Background(), sqlo, stmt) +} + +// DeleteContext is the context-aware variant of [Delete]. +func DeleteContext(ctx context.Context, sqlo Executable, stmt Statement) error { + _, err := stmt.ExecContext(ctx, sqlo) + return err +} + +// DeleteAffected executes a delete statement and returns the number of rows +// affected and an error if any. +func DeleteAffected(sqlo Executable, stmt Statement) (int64, error) { + return DeleteAffectedContext(context.Background(), sqlo, stmt) +} + +// DeleteAffectedContext is the context-aware variant of [DeleteAffected]. +func DeleteAffectedContext(ctx context.Context, sqlo Executable, stmt Statement) (int64, error) { + res, err := stmt.ExecContext(ctx, sqlo) if err != nil { + return 0, err + } + + rowsAffected, err := res.RowsAffected() + if err != nil { + return 0, err + } + + return rowsAffected, nil +} + +func queryReturningContext[T any](ctx context.Context, sqlo Queryable, stmt Statement) (*T, error) { + var result T + if err := stmt.QueryContext(ctx, sqlo, &result); err != nil { return nil, err } return &result, nil diff --git a/exec_test.go b/exec_test.go new file mode 100644 index 0000000..abd9343 --- /dev/null +++ b/exec_test.go @@ -0,0 +1,99 @@ +package dbx + +import ( + "context" + "database/sql" + "errors" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestInsert_ReturnsLastInsertID(t *testing.T) { + stmt := mockStatement{ + execContextFn: func(_ context.Context) (sql.Result, error) { + return mockResult{lastInsertID: 42}, nil + }, + } + + id, err := Insert(mockExecutable{}, stmt) + require.NoError(t, err) + assert.Equal(t, uint64(42), id) +} + +func TestInsert_RejectsZeroID(t *testing.T) { + stmt := mockStatement{ + execContextFn: func(_ context.Context) (sql.Result, error) { + return mockResult{lastInsertID: 0}, nil + }, + } + + _, err := Insert(mockExecutable{}, stmt) + assert.Error(t, err) +} + +func TestInsertReturning_ReturnsRow(t *testing.T) { + stmt := mockStatement{ + queryContextFn: func(_ context.Context, dest any) error { + ptr := dest.(*row) + ptr.ID = 7 + return nil + }, + } + + result, err := InsertReturning[row](mockQueryable{}, stmt) + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, 7, result.ID) +} + +func TestUpdateAffected_ReturnsCount(t *testing.T) { + stmt := mockStatement{ + execContextFn: func(_ context.Context) (sql.Result, error) { + return mockResult{rowsAffected: 3}, nil + }, + } + + n, err := UpdateAffected(mockExecutable{}, stmt) + require.NoError(t, err) + assert.Equal(t, int64(3), n) +} + +func TestDelete_Succeeds(t *testing.T) { + called := false + stmt := mockStatement{ + execContextFn: func(_ context.Context) (sql.Result, error) { + called = true + return mockResult{}, nil + }, + } + + err := Delete(mockExecutable{}, stmt) + require.NoError(t, err) + assert.True(t, called) +} + +func TestDeleteAffected_ReturnsCount(t *testing.T) { + stmt := mockStatement{ + execContextFn: func(_ context.Context) (sql.Result, error) { + return mockResult{rowsAffected: 2}, nil + }, + } + + n, err := DeleteAffected(mockExecutable{}, stmt) + require.NoError(t, err) + assert.Equal(t, int64(2), n) +} + +func TestExecContext_PropagatesError(t *testing.T) { + wantErr := errors.New("exec failed") + stmt := mockStatement{ + execContextFn: func(_ context.Context) (sql.Result, error) { + return nil, wantErr + }, + } + + err := DeleteContext(context.Background(), mockExecutable{}, stmt) + assert.ErrorIs(t, err, wantErr) +} diff --git a/go.mod b/go.mod index f285f5a..24d1dc3 100644 --- a/go.mod +++ b/go.mod @@ -13,13 +13,20 @@ require ( require ( github.com/davecgh/go-spew v1.1.1 // indirect + github.com/dustin/go-humanize v1.0.1 // indirect github.com/google/uuid v1.6.0 // indirect github.com/kr/text v0.2.0 // indirect github.com/mattn/go-colorable v0.1.13 // indirect github.com/mattn/go-isatty v0.0.20 // indirect + github.com/ncruces/go-strftime v1.0.0 // indirect github.com/niemeyer/pretty v0.0.0-20200227124842-a10e7caefd8e // indirect github.com/pmezard/go-difflib v1.0.0 // indirect - golang.org/x/sys v0.33.0 // indirect + github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect + golang.org/x/sys v0.44.0 // indirect gopkg.in/check.v1 v1.0.0-20200227125254-8fa46927fb4f // indirect gopkg.in/yaml.v3 v3.0.1 // indirect + modernc.org/libc v1.73.4 // indirect + modernc.org/mathutil v1.7.1 // indirect + modernc.org/memory v1.11.0 // indirect + modernc.org/sqlite v1.53.0 // indirect ) diff --git a/go.sum b/go.sum index 39d7fa1..2397edd 100644 --- a/go.sum +++ b/go.sum @@ -3,6 +3,8 @@ gitea.auvem.com/go-toolkit/app v0.0.0-20250530181559-231561c92698/go.mod h1:a7EN github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= +github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= github.com/fatih/color v1.18.0 h1:S8gINlzdQ840/4pfAwic/ZE0djQEH3wM94VfqLTZcOM= github.com/fatih/color v1.18.0/go.mod h1:4FelSpRwEGDpQ12mAdzqdOukCy4u8WUtOY6lkT/6HfU= github.com/go-jet/jet/v2 v2.13.0 h1:DcD2IJRGos+4X40IQRV6S6q9onoOfZY/GPdvU6ImZcQ= @@ -20,10 +22,14 @@ github.com/mattn/go-colorable v0.1.13/go.mod h1:7S9/ev0klgBDR4GtXTXX8a3vIGJpMovk github.com/mattn/go-isatty v0.0.16/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/yFXSvRLM= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= +github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= github.com/niemeyer/pretty v0.0.0-20200227124842-a10e7caefd8e h1:fD57ERR4JtEqsWbfPhv4DMiApHyliiK5xCTNVSPiaAs= github.com/niemeyer/pretty v0.0.0-20200227124842-a10e7caefd8e/go.mod h1:zD1mROLANZcx1PVRCS0qkT7pwLkGfwJo4zjcN/Tysno= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= github.com/segmentio/ksuid v1.0.4 h1:sBo2BdShXjmcugAMwjugoGUdUV0pcxY5mW4xKRn3v4c= github.com/segmentio/ksuid v1.0.4/go.mod h1:/XUiZBD3kVx5SmUOl55voK5yeAbBNNIed+2O73XgrPE= github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOfJA= @@ -34,8 +40,18 @@ golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBc golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.33.0 h1:q3i8TbbEz+JRD9ywIRlyRAQbM0qF7hu24q3teo2hbuw= golang.org/x/sys v0.33.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= +golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ= +golang.org/x/sys v0.44.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/check.v1 v1.0.0-20200227125254-8fa46927fb4f h1:BLraFXnmrev5lT+xlilqcH8XK9/i0At2xKjWk4p6zsU= gopkg.in/check.v1 v1.0.0-20200227125254-8fa46927fb4f/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +modernc.org/libc v1.73.4 h1:+ra4Ui8ngyt8HDcO1FTDPWlkAh6yOdaO2yAoh8MddQA= +modernc.org/libc v1.73.4/go.mod h1:DXZ3eO8qMCNn2SnmTNCiC71nJ9Rcq3PsnpU6Vc4rWK8= +modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU= +modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg= +modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI= +modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw= +modernc.org/sqlite v1.53.0 h1:20WG8N9q4ji/dEqGk4uiI0c6OPjSeLTNYGFCc3+7c1M= +modernc.org/sqlite v1.53.0/go.mod h1:xoEpOIpGrgT48H5iiyt/YXPCZPEzlfmfFwtk8Lklw8s= diff --git a/jet_columns.go b/jet_columns.go index 73f89ec..6a5b080 100644 --- a/jet_columns.go +++ b/jet_columns.go @@ -1,6 +1,15 @@ package dbx -import "github.com/go-jet/jet/v2/mysql" +import ( + "slices" + + "github.com/go-jet/jet/v2/mysql" +) + +// ContainsCol reports whether cols contains col. +func ContainsCol(cols ColumnList, col Column) bool { + return slices.Contains(cols, col) +} // NormalCols processes a list of columns and strips out any that implement any of // ColumnTimestamp, ColumnTime, or ColumnDate. diff --git a/jet_columns_test.go b/jet_columns_test.go new file mode 100644 index 0000000..b91ad49 --- /dev/null +++ b/jet_columns_test.go @@ -0,0 +1,105 @@ +package dbx + +import ( + "context" + "database/sql" + "errors" + "testing" + + "github.com/go-jet/jet/v2/mysql" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + _ "modernc.org/sqlite" +) + +func TestContainsCol(t *testing.T) { + colA := mysql.StringColumn("a") + colB := mysql.StringColumn("b") + cols := ColumnList{colA, colB} + + assert.True(t, ContainsCol(cols, colA)) + assert.False(t, ContainsCol(cols, mysql.StringColumn("c"))) +} + +func TestNormalCols_ExcludesTimestamp(t *testing.T) { + name := mysql.StringColumn("name") + ts := mysql.TimestampColumn("updated_at") + cols := NormalCols(name, ts) + + require.Len(t, cols, 1) + assert.Equal(t, name, cols[0]) +} + +type mockTxDB struct { + beginErr error +} + +func (m *mockTxDB) Query(string, ...any) (*sql.Rows, error) { return nil, nil } + +func (m *mockTxDB) QueryContext(context.Context, string, ...any) (*sql.Rows, error) { + return nil, nil +} + +func (m *mockTxDB) Exec(string, ...any) (sql.Result, error) { return nil, nil } + +func (m *mockTxDB) ExecContext(context.Context, string, ...any) (sql.Result, error) { + return nil, nil +} + +func (m *mockTxDB) Begin() (*sql.Tx, error) { + if m.beginErr != nil { + return nil, m.beginErr + } + return &sql.Tx{}, nil +} + +func TestWithTx_ReturnsFnError(t *testing.T) { + db := openTestDB(t) + + wantErr := errors.New("fn failed") + err := WithTx(context.Background(), db, func(_ *sql.Tx) error { + return wantErr + }) + assert.ErrorIs(t, err, wantErr) +} + +func TestWithTx_CommitsOnSuccess(t *testing.T) { + db := openTestDB(t) + called := false + + err := WithTx(context.Background(), db, func(_ *sql.Tx) error { + called = true + return nil + }) + require.NoError(t, err) + assert.True(t, called) +} + +func openTestDB(t *testing.T) *sql.DB { + t.Helper() + db, err := sql.Open("sqlite", ":memory:") + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + return db +} + +func TestWithTx_CancelledContext(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + err := WithTx(ctx, &mockTxDB{}, func(_ *sql.Tx) error { + return nil + }) + assert.ErrorIs(t, err, context.Canceled) +} + +func TestWithTx_BeginError(t *testing.T) { + wantErr := errors.New("begin failed") + err := WithTx(context.Background(), &mockTxDB{beginErr: wantErr}, func(_ *sql.Tx) error { + return nil + }) + assert.ErrorIs(t, err, wantErr) +} + +var _ QueryExecTx = (*mockTxDB)(nil) diff --git a/query.go b/query.go index a682852..73c076b 100644 --- a/query.go +++ b/query.go @@ -1,12 +1,20 @@ package dbx -import "errors" +import ( + "context" + "errors" +) // Fetch queries the database and returns the result as a slice. If the query // returns no rows, it returns an empty slice and no error. func Fetch[T any](sqlo Queryable, stmt Statement) ([]*T, error) { + return FetchContext[T](context.Background(), sqlo, stmt) +} + +// FetchContext is the context-aware variant of [Fetch]. +func FetchContext[T any](ctx context.Context, sqlo Queryable, stmt Statement) ([]*T, error) { var result []*T - if err := stmt.Query(sqlo, &result); err != nil && !errors.Is(err, ErrNoRows) { + if err := stmt.QueryContext(ctx, sqlo, &result); err != nil && !errors.Is(err, ErrNoRows) { return nil, err } return result, nil @@ -15,7 +23,12 @@ func Fetch[T any](sqlo Queryable, stmt Statement) ([]*T, error) { // MustFetch queries the database and returns the result as a slice. If the query // returns no rows, it returns an empty slice and the desired error. func MustFetch[T any](sqlo Queryable, stmt Statement, notFoundErr error) ([]*T, error) { - result, err := Fetch[T](sqlo, stmt) + return MustFetchContext[T](context.Background(), sqlo, stmt, notFoundErr) +} + +// MustFetchContext is the context-aware variant of [MustFetch]. +func MustFetchContext[T any](ctx context.Context, sqlo Queryable, stmt Statement, notFoundErr error) ([]*T, error) { + result, err := FetchContext[T](ctx, sqlo, stmt) if err != nil { return nil, err } @@ -28,7 +41,12 @@ func MustFetch[T any](sqlo Queryable, stmt Statement, notFoundErr error) ([]*T, // FetchOne queries the database and returns a single result. If the query // returns no rows, it returns nil and no error. func FetchOne[T any](sqlo Queryable, stmt Statement) (*T, error) { - result, err := Fetch[T](sqlo, stmt) + return FetchOneContext[T](context.Background(), sqlo, stmt) +} + +// FetchOneContext is the context-aware variant of [FetchOne]. +func FetchOneContext[T any](ctx context.Context, sqlo Queryable, stmt Statement) (*T, error) { + result, err := FetchContext[T](ctx, sqlo, stmt) if err != nil { return nil, err } @@ -41,7 +59,12 @@ func FetchOne[T any](sqlo Queryable, stmt Statement) (*T, error) { // MustFetchOne queries the database and returns a single result. If the query // returns no rows, it returns nil and the desired error. func MustFetchOne[T any](sqlo Queryable, stmt Statement, notFoundErr error) (*T, error) { - result, err := MustFetch[T](sqlo, stmt, notFoundErr) + return MustFetchOneContext[T](context.Background(), sqlo, stmt, notFoundErr) +} + +// MustFetchOneContext is the context-aware variant of [MustFetchOne]. +func MustFetchOneContext[T any](ctx context.Context, sqlo Queryable, stmt Statement, notFoundErr error) (*T, error) { + result, err := MustFetchContext[T](ctx, sqlo, stmt, notFoundErr) if err != nil { return nil, err } diff --git a/query_test.go b/query_test.go new file mode 100644 index 0000000..4dc7810 --- /dev/null +++ b/query_test.go @@ -0,0 +1,80 @@ +package dbx + +import ( + "context" + "errors" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestFetch_EmptyOnNoRows(t *testing.T) { + stmt := mockStatement{ + queryContextFn: func(_ context.Context, dest any) error { + return ErrNoRows + }, + } + + result, err := Fetch[row](mockQueryable{}, stmt) + require.NoError(t, err) + assert.Empty(t, result) +} + +func TestFetch_ReturnsRows(t *testing.T) { + stmt := mockStatement{ + queryContextFn: func(_ context.Context, dest any) error { + ptr := dest.(*[]*row) + *ptr = []*row{{ID: 1}, {ID: 2}} + return nil + }, + } + + result, err := Fetch[row](mockQueryable{}, stmt) + require.NoError(t, err) + assert.Len(t, result, 2) +} + +func TestMustFetch_NotFound(t *testing.T) { + errNotFound := errors.New("not found") + stmt := mockStatement{ + queryContextFn: func(_ context.Context, _ any) error { + return ErrNoRows + }, + } + + result, err := MustFetch[row](mockQueryable{}, stmt, errNotFound) + assert.Nil(t, result) + assert.ErrorIs(t, err, errNotFound) +} + +func TestFetchContext_PropagatesContext(t *testing.T) { + ctx := context.WithValue(context.Background(), testContextKey{}, "ok") + var gotCtx context.Context + stmt := mockStatement{ + queryContextFn: func(c context.Context, dest any) error { + gotCtx = c + ptr := dest.(*[]*row) + *ptr = []*row{{ID: 3}} + return nil + }, + } + + _, err := FetchContext[row](ctx, mockQueryable{}, stmt) + require.NoError(t, err) + assert.Equal(t, ctx, gotCtx) +} + +type testContextKey struct{} + +func TestFetchOne_NilWhenEmpty(t *testing.T) { + stmt := mockStatement{ + queryContextFn: func(_ context.Context, _ any) error { + return ErrNoRows + }, + } + + result, err := FetchOne[row](mockQueryable{}, stmt) + require.NoError(t, err) + assert.Nil(t, result) +} diff --git a/testhelpers_test.go b/testhelpers_test.go new file mode 100644 index 0000000..90c1c33 --- /dev/null +++ b/testhelpers_test.go @@ -0,0 +1,72 @@ +package dbx + +import ( + "context" + "database/sql" + + "github.com/go-jet/jet/v2/qrm" +) + +type mockStatement struct { + queryContextFn func(ctx context.Context, dest any) error + execContextFn func(ctx context.Context) (sql.Result, error) +} + +func (m mockStatement) Query(db qrm.Queryable, dest any) error { + return m.QueryContext(context.Background(), db, dest) +} + +func (m mockStatement) QueryContext(ctx context.Context, db qrm.Queryable, dest any) error { + if m.queryContextFn != nil { + return m.queryContextFn(ctx, dest) + } + return nil +} + +func (m mockStatement) Exec(db qrm.Executable) (sql.Result, error) { + return m.ExecContext(context.Background(), db) +} + +func (m mockStatement) ExecContext(ctx context.Context, db qrm.Executable) (sql.Result, error) { + if m.execContextFn != nil { + return m.execContextFn(ctx) + } + return nil, nil +} + +type mockQueryable struct{} + +func (mockQueryable) Query(string, ...any) (*sql.Rows, error) { return nil, nil } + +func (mockQueryable) QueryContext(context.Context, string, ...any) (*sql.Rows, error) { + return nil, nil +} + +type mockExecutable struct{} + +func (mockExecutable) Exec(string, ...any) (sql.Result, error) { return nil, nil } + +func (mockExecutable) ExecContext(context.Context, string, ...any) (sql.Result, error) { + return nil, nil +} + +type mockResult struct { + lastInsertID int64 + rowsAffected int64 + lastInsertErr error + rowsAffectedErr error +} + +func (r mockResult) LastInsertId() (int64, error) { + return r.lastInsertID, r.lastInsertErr +} + +func (r mockResult) RowsAffected() (int64, error) { + return r.rowsAffected, r.rowsAffectedErr +} + +type row struct { + ID int +} + +var _ Statement = mockStatement{} diff --git a/tx.go b/tx.go new file mode 100644 index 0000000..ec5dbc1 --- /dev/null +++ b/tx.go @@ -0,0 +1,26 @@ +package dbx + +import ( + "context" + "database/sql" +) + +// WithTx begins a transaction on db, runs fn, and commits on success. +// The transaction is rolled back if fn returns an error or commit fails. +func WithTx(ctx context.Context, db QueryExecTx, fn func(tx *sql.Tx) error) error { + if err := ctx.Err(); err != nil { + return err + } + + tx, err := db.Begin() + if err != nil { + return err + } + + if err := fn(tx); err != nil { + _ = tx.Rollback() + return err + } + + return tx.Commit() +}