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 <cursoragent@cursor.com>
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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=
|
||||
|
||||
+10
-1
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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{}
|
||||
@@ -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()
|
||||
}
|
||||
Reference in New Issue
Block a user