diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 2ff19ca..5eb2e25 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -2,50 +2,104 @@ name: CI on: push: - branches: [ master, main ] + branches: [master, main] pull_request: - branches: [ master, main ] + branches: [master, main] + +permissions: + contents: read + +env: + GOIMPORTS_VERSION: v0.49.0 + GOLANGCI_LINT_VERSION: v2.13.2 + GOVULNCHECK_VERSION: v1.7.0 jobs: test: + name: Test (Go ${{ matrix.go-version }}) runs-on: ubuntu-latest - strategy: + fail-fast: false matrix: - go-version: ['1.21', '1.22', '1.23'] - + go-version: ["1.23.x", stable] + steps: + - uses: actions/checkout@v7 + - uses: actions/setup-go@v7 + with: + go-version: ${{ matrix.go-version }} + cache-dependency-path: go.sum + - run: go test ./... -count=1 + - run: go build ./... + + quality: + name: Format and lint + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v7 + - uses: actions/setup-go@v7 + with: + go-version: stable + cache-dependency-path: go.sum + - name: Install pinned tools + run: | + go install "golang.org/x/tools/cmd/goimports@${GOIMPORTS_VERSION}" + go install "github.com/golangci/golangci-lint/v2/cmd/golangci-lint@${GOLANGCI_LINT_VERSION}" + - run: make fmt-check + - run: make lint + - name: Check module metadata + run: go mod tidy -diff + + race: + name: Race detector + runs-on: ubuntu-latest steps: - - uses: actions/checkout@v4 - - - name: Set up Go - uses: actions/setup-go@v4 - with: - go-version: ${{ matrix.go-version }} - - - name: Cache Go modules - uses: actions/cache@v3 - with: - path: | - ~/.cache/go-build - ~/go/pkg/mod - key: ${{ runner.os }}-go-${{ matrix.go-version }}-${{ hashFiles('**/go.sum') }} - restore-keys: | - ${{ runner.os }}-go-${{ matrix.go-version }}- - - - name: Install tools - run: | - go install golang.org/x/tools/cmd/goimports@latest - go install honnef.co/go/tools/cmd/staticcheck@latest - echo "$(go env GOPATH)/bin" >> $GITHUB_PATH - - - name: Format - run: make fmt - - - name: Lint - run: make lint - - - name: Test - run: make test - - - name: Build - run: go build ./... \ No newline at end of file + - uses: actions/checkout@v7 + - uses: actions/setup-go@v7 + with: + go-version: stable + cache-dependency-path: go.sum + - run: make test-race + + vulnerability: + name: Vulnerability scan + runs-on: ubuntu-latest + timeout-minutes: 10 + steps: + - uses: actions/checkout@v7 + - uses: actions/setup-go@v7 + with: + go-version: stable + cache-dependency-path: go.sum + - name: Install govulncheck + run: go install "golang.org/x/vuln/cmd/govulncheck@${GOVULNCHECK_VERSION}" + - name: Run govulncheck + shell: bash + run: | + for attempt in 1 2 3; do + if govulncheck ./...; then + exit 0 + else + status=$? + fi + + if [[ ${status} -eq 3 ]]; then + exit "${status}" + fi + + if [[ ${status} -ne 1 ]]; then + exit "${status}" + fi + + if [[ ${attempt} -eq 3 ]]; then + exit "${status}" + fi + + if [[ ${attempt} -eq 1 ]]; then + delay=10 + else + delay=30 + fi + + echo "::warning::govulncheck failed with an operational error (exit ${status}); retrying in ${delay}s" + sleep "${delay}" + done diff --git a/.github/workflows/static-analysis.yml b/.github/workflows/static-analysis.yml deleted file mode 100644 index 9652f5b..0000000 --- a/.github/workflows/static-analysis.yml +++ /dev/null @@ -1,43 +0,0 @@ -name: Static Code Analysis - -on: - push: - branches: [ master, main ] - pull_request: - branches: [ master, main ] - -jobs: - static-analysis: - name: Security & Code Quality Analysis - runs-on: ubuntu-latest - - steps: - - uses: actions/checkout@v4 - - - name: Set up Go - uses: actions/setup-go@v4 - with: - go-version: '1.23' - - - name: Cache Go modules - uses: actions/cache@v3 - with: - path: | - ~/.cache/go-build - ~/go/pkg/mod - key: ${{ runner.os }}-go-static-analysis-${{ hashFiles('**/go.sum') }} - restore-keys: | - ${{ runner.os }}-go-static-analysis- - - - name: Install analysis tools - run: | - go install github.com/securego/gosec/v2/cmd/gosec@latest - go install github.com/golangci/golangci-lint/cmd/golangci-lint@latest - - - name: Run gosec security scanner - run: gosec ./... - continue-on-error: true - - - name: Run comprehensive static analysis - run: golangci-lint run - continue-on-error: true \ No newline at end of file diff --git a/.golangci.yml b/.golangci.yml index 41b6cca..3f60c9a 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -1,56 +1,18 @@ -linters-settings: - errcheck: - check-type-assertions: true - check-blank: true - govet: - enable: - - shadow - - fieldalignment - gocyclo: - min-complexity: 15 - dupl: - threshold: 100 - goconst: - min-len: 2 - min-occurrences: 2 - misspell: - locale: US - lll: - line-length: 140 - goimports: - local-prefixes: github.com/ziflex/dbx - gocritic: - enabled-tags: - - diagnostic - - performance - - style - disabled-checks: - - dupImport - - ifElseChain - - octalLiteral - - whyNoLint - - wrapperFunc - depguard: - rules: - main: - deny: - - pkg: "unsafe" - desc: "unsafe package should not be used" - +version: "2" +run: + issues-exit-code: 1 linters: - disable-all: true + default: none enable: - bodyclose + - depguard - dupl - errcheck - gochecknoinits - goconst - gocritic - gocyclo - - gofmt - - goimports - gosec - - gosimple - govet - ineffassign - lll @@ -61,28 +23,79 @@ linters: - rowserrcheck - sqlclosecheck - staticcheck - - stylecheck - - typecheck - unconvert - unparam - unused - whitespace - -issues: - exclude-rules: - - path: _test\.go - linters: - - gocyclo - - errcheck - - dupl - - gosec - - lll - - govet - - gocritic - exclude: - # Exclude some staticcheck messages - - "SA1029:" # should not use built-in type string as key for value - -run: - timeout: 5m - issues-exit-code: 1 \ No newline at end of file + settings: + depguard: + rules: + main: + deny: + - pkg: unsafe + desc: unsafe package should not be used + dupl: + threshold: 100 + errcheck: + check-type-assertions: true + check-blank: true + goconst: + min-len: 2 + min-occurrences: 2 + gocritic: + disabled-checks: + - dupImport + - ifElseChain + - octalLiteral + - whyNoLint + - wrapperFunc + enabled-tags: + - diagnostic + - performance + - style + gocyclo: + min-complexity: 15 + govet: + enable: + - shadow + - fieldalignment + lll: + line-length: 140 + misspell: + locale: US + exclusions: + generated: lax + presets: + - comments + - common-false-positives + - legacy + - std-error-handling + rules: + - linters: + - dupl + - gocritic + - gocyclo + - gosec + - govet + - lll + path: _test\.go + - path: (.+)\.go$ + text: 'SA1029:' + paths: + - third_party$ + - builtin$ + - examples$ +formatters: + enable: + - gofmt + - goimports + settings: + goimports: + local-prefixes: + - github.com/ziflex/dbx + exclusions: + generated: lax + paths: + - third_party$ + - builtin$ + - examples$ diff --git a/Makefile b/Makefile index 4839e5c..cfdc613 100644 --- a/Makefile +++ b/Makefile @@ -1,12 +1,18 @@ -default: fmt lint test +default: fmt-check lint test test: go test ./... lint: - go vet ./... && \ - staticcheck -tests=false ./... + golangci-lint run fmt: go fmt ./... && \ - goimports -w . \ No newline at end of file + goimports -w -local github.com/ziflex/dbx . + +fmt-check: + @files="$$(gofmt -l .)"; if [ -n "$$files" ]; then echo "gofmt required for:"; echo "$$files"; exit 1; fi + @files="$$(goimports -l -local github.com/ziflex/dbx .)"; if [ -n "$$files" ]; then echo "goimports required for:"; echo "$$files"; exit 1; fi + +test-race: + go test -race ./... diff --git a/README.md b/README.md index 056cb4c..bf4559f 100644 --- a/README.md +++ b/README.md @@ -2,7 +2,7 @@ A lightweight, context-aware abstraction layer for Go's `database/sql` package that simplifies database operations and transaction management. -[![API Documentation](https://godoc.org/github.com/ziflex/dbx?status.svg)](https://godoc.org/github.com/ziflex/dbx) +[![Go Reference](https://pkg.go.dev/badge/github.com/ziflex/dbx.svg)](https://pkg.go.dev/github.com/ziflex/dbx) ## Table of Contents - [Why dbx?](#why-dbx) @@ -21,7 +21,7 @@ A lightweight, context-aware abstraction layer for Go's `database/sql` package t The standard `database/sql` package is powerful but requires boilerplate code for common patterns. `dbx` addresses several pain points: - **Context Management**: Eliminates the need to pass both `context.Context` and database connections separately -- **Transaction Handling**: Automatic transaction lifecycle management with support for nested transactions +- **Transaction Handling**: Automatic transaction lifecycle management with transaction reuse for nested operations - **Unified Interface**: Same API for both direct database operations and transactions - **Testing**: Easier to mock and test database operations - **Clean Architecture**: Promotes separation of concerns between business logic and data access @@ -44,14 +44,18 @@ go get github.com/ziflex/dbx@latest ## Key Concepts -### Database Interface -The `Database` interface wraps a `*sql.DB` and provides context creation: +### Database Interfaces +The `Database` interface provides connection management, transaction creation, and query execution. Context creation is exposed separately by `DatabaseWithContext`: ```go type Database interface { io.Closer - ContextCreator // Creates dbx.Context - Beginner // Begins transactions - Executor // Executes queries directly + Beginner // Begins transactions + Executor // Executes queries directly +} + +type DatabaseWithContext interface { + Database + ContextCreator // Creates dbx.Context } ``` @@ -89,8 +93,8 @@ import ( "database/sql" "fmt" "log" - - _ "github.com/lib/pq" + + _ "github.com/lib/pq" "github.com/ziflex/dbx" ) @@ -104,7 +108,7 @@ type User struct { func getUserNames(ctx dbx.Context) ([]User, error) { executor := ctx.Executor() - rows, err := executor.Query("SELECT id, name FROM users ORDER BY name") + rows, err := executor.QueryContext(ctx, "SELECT id, name FROM users ORDER BY name") if err != nil { return nil, fmt.Errorf("failed to query users: %w", err) } @@ -169,7 +173,7 @@ func directExample() { func getUserCount(ctx dbx.Context) (int, error) { var count int - err := ctx.Executor().QueryRow("SELECT COUNT(*) FROM users").Scan(&count) + err := ctx.Executor().QueryRowContext(ctx, "SELECT COUNT(*) FROM users").Scan(&count) return count, err } ``` @@ -201,14 +205,14 @@ func main() { ``` ### Context Helper Functions -- `dbx.Is(ctx)` - Check if context contains dbx context -- `dbx.As(ctx)` - Extract dbx context with ok flag -- `dbx.FromContext(ctx)` - Extract dbx context (returns nil if not found) +- `dbx.Is(ctx)` - Check whether the context directly implements `dbx.Context` +- `dbx.As(ctx)` - Type-assert a direct `dbx.Context` with an ok flag +- `dbx.FromContext(ctx)` - Extract a direct or embedded dbx context (returns nil if not found) - `dbx.WithContext(ctx, dbxCtx)` - Embed dbx context into regular context ## Transaction Management -`dbx` provides powerful transaction management with automatic lifecycle handling and support for nested operations. +`dbx` provides transaction management with automatic lifecycle handling and transaction reuse for nested operations. ### Basic Transactions @@ -216,20 +220,18 @@ func main() { func createUserWithProfile(ctx context.Context, db dbx.Database, userName, email string) error { return dbx.Transaction(ctx, db, func(txCtx dbx.Context) error { // Insert user - result, err := txCtx.Executor().Exec( - "INSERT INTO users (name) VALUES ($1) RETURNING id", userName) - if err != nil { - return fmt.Errorf("failed to insert user: %w", err) - } - var userID int64 - userID, err = result.LastInsertId() + err := txCtx.Executor().QueryRowContext( + txCtx, + "INSERT INTO users (name) VALUES ($1) RETURNING id", + userName, + ).Scan(&userID) if err != nil { return fmt.Errorf("failed to get user ID: %w", err) } // Insert profile - _, err = txCtx.Executor().Exec( + _, err = txCtx.Executor().ExecContext(txCtx, "INSERT INTO profiles (user_id, email) VALUES ($1, $2)", userID, email) if err != nil { return fmt.Errorf("failed to insert profile: %w", err) @@ -272,13 +274,13 @@ Use `TransactionWithResult` when you need to return values from transactions: ```go func createUserAndGetID(ctx context.Context, db dbx.Database, name string) (int64, error) { return dbx.TransactionWithResult(ctx, db, func(txCtx dbx.Context) (int64, error) { - result, err := txCtx.Executor().Exec( - "INSERT INTO users (name) VALUES ($1)", name) - if err != nil { - return 0, err - } - - return result.LastInsertId() + var userID int64 + err := txCtx.Executor().QueryRowContext( + txCtx, + "INSERT INTO users (name) VALUES ($1) RETURNING id", + name, + ).Scan(&userID) + return userID, err }) } ``` @@ -287,7 +289,7 @@ func createUserAndGetID(ctx context.Context, db dbx.Database, name string) (int6 ### Transaction Options -Control transaction behavior with options: +Control transaction behavior with options. Isolation and read-only options apply only when `dbx` creates a transaction; a reused transaction retains the options selected by its owner: ```go // Read-only transaction @@ -301,21 +303,23 @@ err := dbx.Transaction(ctx, db, func(txCtx dbx.Context) error { return performCriticalOperation(txCtx) }, dbx.WithIsolationLevel(sql.LevelSerializable)) -// Force new transaction (disable reuse) +// Force an independent transaction (disable reuse; this is not a savepoint) err := dbx.Transaction(ctx, db, func(txCtx dbx.Context) error { return independentOperation(txCtx) }, dbx.WithNewTransaction()) ``` +An independent transaction may use another pooled connection and commits separately from the outer transaction. `WithNewTransaction` does not create a database savepoint. + ### Error Handling Patterns -`dbx` automatically handles transaction rollback on errors: +`dbx` automatically handles transaction rollback on errors and panics. An operation error is returned unchanged when rollback succeeds; if rollback also fails, the returned error contains both failures: ```go func transferFunds(ctx context.Context, db dbx.Database, fromID, toID int, amount decimal.Decimal) error { return dbx.Transaction(ctx, db, func(txCtx dbx.Context) error { // Debit source account - result, err := txCtx.Executor().Exec( + result, err := txCtx.Executor().ExecContext(txCtx, "UPDATE accounts SET balance = balance - $1 WHERE id = $2 AND balance >= $1", amount, fromID) if err != nil { @@ -331,7 +335,7 @@ func transferFunds(ctx context.Context, db dbx.Database, fromID, toID int, amoun } // Credit destination account - _, err = txCtx.Executor().Exec( + _, err = txCtx.Executor().ExecContext(txCtx, "UPDATE accounts SET balance = balance + $1 WHERE id = $2", amount, toID) if err != nil { @@ -346,21 +350,30 @@ func transferFunds(ctx context.Context, db dbx.Database, fromID, toID int, amoun ### Working with Prepared Statements -Since `dbx.Context.Executor()` returns the underlying `sql.DB` or `sql.Tx`, you can use prepared statements: +`Executor` intentionally exposes only the query methods common to `sql.DB` and `sql.Tx`. Both standard implementations also support prepared statements, which can be accessed through a narrow local capability interface without expanding `dbx.Executor`: ```go +type statementPreparer interface { + PrepareContext(context.Context, string) (*sql.Stmt, error) +} + func batchInsertUsers(ctx dbx.Context, users []User) error { - executor := ctx.Executor() - - // Prepare statement (works with both DB and Tx) - stmt, err := executor.Prepare("INSERT INTO users (name, email) VALUES ($1, $2)") + preparer, ok := ctx.Executor().(statementPreparer) + if !ok { + return fmt.Errorf("executor does not support prepared statements") + } + + stmt, err := preparer.PrepareContext( + ctx, + "INSERT INTO users (name, email) VALUES ($1, $2)", + ) if err != nil { return err } defer stmt.Close() for _, user := range users { - if _, err := stmt.Exec(user.Name, user.Email); err != nil { + if _, err := stmt.ExecContext(ctx, user.Name, user.Email); err != nil { return fmt.Errorf("failed to insert user %s: %w", user.Name, err) } } @@ -413,10 +426,10 @@ func TestTransferFunds(t *testing.T) { // Setup transaction expectations mock.ExpectBegin() mock.ExpectExec("UPDATE accounts SET balance"). - WithArgs(100, 1, 100). + WithArgs(sqlmock.AnyArg(), 1). WillReturnResult(sqlmock.NewResult(0, 1)) mock.ExpectExec("UPDATE accounts SET balance"). - WithArgs(100, 2). + WithArgs(sqlmock.AnyArg(), 2). WillReturnResult(sqlmock.NewResult(0, 1)) mock.ExpectCommit() @@ -432,21 +445,21 @@ func TestTransferFunds(t *testing.T) { ### Core Functions -- `dbx.New(db *sql.DB) Database` - Creates a new dbx Database wrapper -- `dbx.Transaction(ctx context.Context, db Database, op Operation, opts ...Option) error` - Executes operation in transaction -- `dbx.TransactionWithResult[T](ctx context.Context, db Database, op OperationWithResult[T], opts ...Option) (T, error)` - Executes operation in transaction with return value +- `dbx.New(db *sql.DB) DatabaseWithContext` - Creates a new dbx database wrapper with context creation +- `dbx.Transaction(ctx context.Context, beginner Beginner, op Operation, opts ...Option) error` - Executes an operation in a transaction +- `dbx.TransactionWithResult[T](ctx context.Context, beginner Beginner, op OperationWithResult[T], opts ...Option) (T, error)` - Executes a transaction and returns a typed result ### Context Functions - `dbx.FromContext(ctx context.Context) Context` - Extract dbx context from context - `dbx.WithContext(ctx context.Context, dbxCtx Context) context.Context` - Embed dbx context -- `dbx.Is(ctx context.Context) bool` - Check if context contains dbx context -- `dbx.As(ctx context.Context) (Context, bool)` - Extract dbx context with ok flag +- `dbx.Is(ctx context.Context) bool` - Check whether the context directly implements `dbx.Context` +- `dbx.As(ctx context.Context) (Context, bool)` - Type-assert a direct `dbx.Context` ### Transaction Options - `dbx.WithIsolationLevel(level sql.IsolationLevel)` - Set transaction isolation level - `dbx.WithReadOnly(readOnly bool)` - Set read-only flag -- `dbx.WithNewTransaction()` - Force creation of new transaction (disable reuse) +- `dbx.WithNewTransaction()` - Force creation of an independent transaction (disable reuse) -For complete API documentation, see [GoDoc](https://godoc.org/github.com/ziflex/dbx). \ No newline at end of file +For complete API documentation, see [Go Reference](https://pkg.go.dev/github.com/ziflex/dbx). diff --git a/context.go b/context.go index 7da88ef..9986e43 100644 --- a/context.go +++ b/context.go @@ -82,7 +82,7 @@ func As(ctx context.Context) (Context, bool) { // // sqlDB, _ := sql.Open("postgres", connectionString) // dbCtx := dbx.NewContext(context.Background(), sqlDB) -// result, err := dbCtx.Executor().Exec("INSERT INTO users (name) VALUES (?)", "John") +// result, err := dbCtx.Executor().ExecContext(dbCtx, "INSERT INTO users (name) VALUES (?)", "John") func NewContext(parent context.Context, exec Executor) Context { return &defaultContext{ parent: parent, @@ -105,7 +105,7 @@ func NewContext(parent context.Context, exec Executor) Context { // // db := dbx.New(sqlDB) // dbCtx := dbx.NewDatabaseContext(context.Background(), db) -// result, err := dbCtx.Executor().Exec("INSERT INTO users (name) VALUES (?)", "John") +// result, err := dbCtx.Executor().ExecContext(dbCtx, "INSERT INTO users (name) VALUES (?)", "John") func NewDatabaseContext(parent context.Context, db Database) Context { return &defaultContext{ parent: parent, diff --git a/context_test.go b/context_test.go index 8f32de7..626860d 100644 --- a/context_test.go +++ b/context_test.go @@ -5,17 +5,15 @@ import ( "testing" "time" - "github.com/DATA-DOG/go-sqlmock" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/ziflex/dbx" ) func TestContextHelpers(t *testing.T) { t.Run("Is should return true for dbx context", func(t *testing.T) { - mockDB, _, err := sqlmock.New() - require.NoError(t, err) - defer mockDB.Close() + mockDB, _ := newSQLMock(t) db := dbx.New(mockDB) ctx := db.Context(context.Background()) @@ -29,9 +27,7 @@ func TestContextHelpers(t *testing.T) { }) t.Run("As should return dbx context when present", func(t *testing.T) { - mockDB, _, err := sqlmock.New() - require.NoError(t, err) - defer mockDB.Close() + mockDB, _ := newSQLMock(t) db := dbx.New(mockDB) originalCtx := db.Context(context.Background()) @@ -50,9 +46,7 @@ func TestContextHelpers(t *testing.T) { }) t.Run("WithContext and FromContext should work together", func(t *testing.T) { - mockDB, _, err := sqlmock.New() - require.NoError(t, err) - defer mockDB.Close() + mockDB, _ := newSQLMock(t) db := dbx.New(mockDB) dbCtx := db.Context(context.Background()) @@ -64,6 +58,10 @@ func TestContextHelpers(t *testing.T) { // Extract it back extractedCtx := dbx.FromContext(embeddedCtx) assert.Equal(t, dbCtx, extractedCtx) + assert.False(t, dbx.Is(embeddedCtx)) + assert.NotEqual(t, dbCtx, embeddedCtx) + _, ok := dbx.As(embeddedCtx) + assert.False(t, ok) }) t.Run("FromContext should return nil for regular context", func(t *testing.T) { @@ -73,9 +71,7 @@ func TestContextHelpers(t *testing.T) { }) t.Run("FromContext should work with direct dbx context", func(t *testing.T) { - mockDB, _, err := sqlmock.New() - require.NoError(t, err) - defer mockDB.Close() + mockDB, _ := newSQLMock(t) db := dbx.New(mockDB) dbCtx := db.Context(context.Background()) @@ -85,11 +81,57 @@ func TestContextHelpers(t *testing.T) { }) } +func TestNewContextFrom(t *testing.T) { + t.Run("returns an existing context", func(t *testing.T) { + database, _ := newSQLMock(t) + existing := dbx.NewContext(context.Background(), dbx.New(database)) + + actual := dbx.NewContextFrom(existing, struct{}{}) + + assert.Same(t, existing, actual) + }) + + t.Run("creates a context with a ContextCreator", func(t *testing.T) { + database, _ := newSQLMock(t) + creator := dbx.New(database) + + actual := dbx.NewContextFrom(context.Background(), creator) + + assert.Equal(t, creator, actual.Executor()) + }) + + t.Run("creates a context with a Database", func(t *testing.T) { + database, _ := newSQLMock(t) + input := struct{ dbx.Database }{Database: dbx.New(database)} + + actual := dbx.NewContextFrom(context.Background(), input) + + assert.Equal(t, input, actual.Executor()) + }) + + t.Run("creates a context with a Transactor", func(t *testing.T) { + database, mock := newSQLMock(t) + mock.ExpectBegin() + mock.ExpectRollback() + + tx, err := database.BeginTx(context.Background(), nil) + require.NoError(t, err) + + actual := dbx.NewContextFrom(context.Background(), tx) + assert.Equal(t, tx, actual.Executor()) + require.NoError(t, tx.Rollback()) + }) + + t.Run("panics for unsupported input", func(t *testing.T) { + assert.PanicsWithValue(t, "input must implement ContextCreator, Database, or Transactor", func() { + dbx.NewContextFrom(context.Background(), struct{}{}) + }) + }) +} + func TestDefaultContext(t *testing.T) { t.Run("context methods should delegate to parent", func(t *testing.T) { - mockDB, _, err := sqlmock.New() - require.NoError(t, err) - defer mockDB.Close() + mockDB, _ := newSQLMock(t) // Create a context with deadline deadline := time.Now().Add(5 * time.Second) @@ -112,7 +154,7 @@ func TestDefaultContext(t *testing.T) { // Expected } - // Test Err (should be nil since not cancelled/timed out) + // Test Err (should be nil since not canceled/timed out) assert.NoError(t, dbCtx.Err()) // Test Value @@ -125,10 +167,8 @@ func TestDefaultContext(t *testing.T) { assert.Nil(t, dbCtxWithValue.Value("nonexistent-key")) }) - t.Run("cancelled context should propagate cancellation", func(t *testing.T) { - mockDB, _, err := sqlmock.New() - require.NoError(t, err) - defer mockDB.Close() + t.Run("canceled context should propagate cancellation", func(t *testing.T) { + mockDB, _ := newSQLMock(t) parentCtx, cancel := context.WithCancel(context.Background()) db := dbx.New(mockDB) @@ -151,9 +191,7 @@ func TestDefaultContext(t *testing.T) { }) t.Run("Executor should return the provided executor", func(t *testing.T) { - mockDB, _, err := sqlmock.New() - require.NoError(t, err) - defer mockDB.Close() + mockDB, _ := newSQLMock(t) db := dbx.New(mockDB) ctx := dbx.NewContext(context.Background(), db) diff --git a/database.go b/database.go index 73b276d..77b216b 100644 --- a/database.go +++ b/database.go @@ -21,8 +21,8 @@ type defaultDatabase struct { // context creation via the Context method. // // Parameters: -// - db: A properly initialized sql.DB instance. The caller retains ownership -// and responsibility for the sql.DB's configuration and driver setup. +// - db: A properly initialized sql.DB instance. The caller remains responsible +// for its configuration and driver setup. Closing the returned wrapper closes db. // // Returns: // - DatabaseWithContext: A dbx Database that can create contexts and manage transactions. @@ -33,13 +33,11 @@ type defaultDatabase struct { // if err != nil { // return err // } -// defer sqlDB.Close() -// // dbxDB := dbx.New(sqlDB) // defer dbxDB.Close() // // ctx := dbxDB.Context(context.Background()) -// rows, err := ctx.Executor().Query("SELECT * FROM users") +// rows, err := ctx.Executor().QueryContext(ctx, "SELECT * FROM users") func New(db *sql.DB) DatabaseWithContext { return &defaultDatabase{db} } @@ -66,6 +64,8 @@ func (d *defaultDatabase) Context(ctx context.Context) Context { // Begin starts a transaction with default options. // This method delegates to the underlying sql.DB's Begin method. +// +//nolint:noctx // Beginner preserves database/sql's legacy Begin API for compatibility. func (d *defaultDatabase) Begin() (*sql.Tx, error) { return d.db.Begin() } @@ -91,6 +91,8 @@ func (d *defaultDatabase) BeginTx(ctx context.Context, opts *sql.TxOptions) (*sq // Returns: // - sql.Result: Contains information about the query execution (rows affected, last insert ID) // - error: Any error that occurred during query execution +// +//nolint:noctx // Executor intentionally preserves database/sql's legacy Exec API. func (d *defaultDatabase) Exec(query string, args ...interface{}) (sql.Result, error) { return d.db.Exec(query, args...) } @@ -105,6 +107,8 @@ func (d *defaultDatabase) Exec(query string, args ...interface{}) (sql.Result, e // Returns: // - *sql.Rows: Rows returned by the query. Must be closed after use. // - error: Any error that occurred during query execution +// +//nolint:noctx // Executor intentionally preserves database/sql's legacy Query API. func (d *defaultDatabase) Query(query string, args ...interface{}) (*sql.Rows, error) { return d.db.Query(query, args...) } @@ -119,6 +123,8 @@ func (d *defaultDatabase) Query(query string, args ...interface{}) (*sql.Rows, e // // Returns: // - *sql.Row: Single row result. Use Scan() to extract values and check for errors. +// +//nolint:noctx // Executor intentionally preserves database/sql's legacy QueryRow API. func (d *defaultDatabase) QueryRow(query string, args ...interface{}) *sql.Row { return d.db.QueryRow(query, args...) } diff --git a/database_test.go b/database_test.go index 5a4f9ad..9bbd1a6 100644 --- a/database_test.go +++ b/database_test.go @@ -7,6 +7,7 @@ import ( "github.com/DATA-DOG/go-sqlmock" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/ziflex/dbx" ) @@ -91,7 +92,7 @@ func TestDatabase(t *testing.T) { require.NoError(t, err) defer mockDB.Close() - rows := sqlmock.NewRows([]string{"id", "name"}). + rows := sqlmock.NewRows([]string{testIDColumn, testNameColumn}). AddRow(1, "Alice"). AddRow(2, "Bob") mock.ExpectQuery("SELECT id, name FROM users").WillReturnRows(rows) @@ -117,6 +118,7 @@ func TestDatabase(t *testing.T) { assert.NoError(t, err) users = append(users, user) } + assert.NoError(t, result.Err()) assert.Len(t, users, 2) assert.Equal(t, "Alice", users[0].Name) @@ -129,7 +131,7 @@ func TestDatabase(t *testing.T) { require.NoError(t, err) defer mockDB.Close() - rows := sqlmock.NewRows([]string{"count"}).AddRow(42) + rows := sqlmock.NewRows([]string{testCountColumn}).AddRow(42) mock.ExpectQuery("SELECT COUNT\\(\\*\\) FROM users").WillReturnRows(rows) db := dbx.New(mockDB) @@ -170,7 +172,7 @@ func TestDatabase(t *testing.T) { require.NoError(t, err) defer mockDB.Close() - rows := sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "Alice") + rows := sqlmock.NewRows([]string{testIDColumn, testNameColumn}).AddRow(1, "Alice") mock.ExpectQuery("SELECT id, name FROM users WHERE id"). WithArgs(1). WillReturnRows(rows) @@ -188,6 +190,7 @@ func TestDatabase(t *testing.T) { var name string err = result.Scan(&id, &name) assert.NoError(t, err) + assert.NoError(t, result.Err()) assert.Equal(t, 1, id) assert.Equal(t, "Alice", name) @@ -199,7 +202,7 @@ func TestDatabase(t *testing.T) { require.NoError(t, err) defer mockDB.Close() - rows := sqlmock.NewRows([]string{"name"}).AddRow("Alice") + rows := sqlmock.NewRows([]string{testNameColumn}).AddRow("Alice") mock.ExpectQuery("SELECT name FROM users WHERE id"). WithArgs(1). WillReturnRows(rows) diff --git a/go.mod b/go.mod index 7648630..cc3c652 100644 --- a/go.mod +++ b/go.mod @@ -2,8 +2,6 @@ module github.com/ziflex/dbx go 1.23 -toolchain go1.24.6 - require ( github.com/DATA-DOG/go-sqlmock v1.5.2 github.com/stretchr/testify v1.11.1 @@ -11,9 +9,7 @@ require ( require ( github.com/davecgh/go-spew v1.1.1 // indirect - github.com/kisielk/sqlstruct v0.0.0-20210630145711-dae28ed37023 // indirect github.com/kr/pretty v0.3.1 // indirect - github.com/kr/text v0.2.0 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect github.com/rogpeppe/go-internal v1.14.1 // indirect gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c // indirect diff --git a/go.sum b/go.sum index 11224ca..896d6fb 100644 --- a/go.sum +++ b/go.sum @@ -4,14 +4,10 @@ github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ3 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/kisielk/sqlstruct v0.0.0-20201105191214-5f3e10d3ab46/go.mod h1:yyMNCyc/Ib3bDTKd379tNMpB/7/H5TjM2Y9QJ5THLbE= -github.com/kisielk/sqlstruct v0.0.0-20210630145711-dae28ed37023 h1:/pb3UJ+3ZtSEUKWnufwsoVF7f0AX5ytPULbTwHMgbq4= -github.com/kisielk/sqlstruct v0.0.0-20210630145711-dae28ed37023/go.mod h1:yyMNCyc/Ib3bDTKd379tNMpB/7/H5TjM2Y9QJ5THLbE= -github.com/kr/pretty v0.2.1 h1:Fmg33tUaq4/8ym9TJN1x7sLJnHVwhP33CNkpYV/7rwI= github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI= github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ= -github.com/kr/text v0.1.0 h1:45sCR5RtlFHMR4UwH9sdQ5TC8v0qDQCHnXt+kaKSTVE= github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI= github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= diff --git a/lib.go b/lib.go index f74ad75..3dc2caa 100644 --- a/lib.go +++ b/lib.go @@ -7,7 +7,7 @@ // // Key Features: // - Context-driven design that embeds database connections within Go contexts -// - Automatic transaction lifecycle management with support for nested transactions +// - Automatic transaction lifecycle management with reuse for nested operations // - Unified interface that works the same for both direct DB operations and transactions // - Interface-based design for maximum flexibility and testability // - Zero magic - predictable behavior with no hidden surprises @@ -19,8 +19,6 @@ // if err != nil { // return err // } -// defer db.Close() -// // // Create a dbx database instance // dbxDB := dbx.New(db) // defer dbxDB.Close() @@ -30,12 +28,12 @@ // dbCtx := dbxDB.Context(ctx) // // // Execute queries -// rows, err := dbCtx.Executor().Query("SELECT * FROM users") +// rows, err := dbCtx.Executor().QueryContext(dbCtx, "SELECT * FROM users") // // Transaction Example: // // err := dbx.Transaction(ctx, dbxDB, func(txCtx dbx.Context) error { -// _, err := txCtx.Executor().Exec("INSERT INTO users (name) VALUES (?)", "John") +// _, err := txCtx.Executor().ExecContext(txCtx, "INSERT INTO users (name) VALUES (?)", "John") // return err // }) package dbx @@ -135,7 +133,7 @@ type ( // // Example: // op := func(ctx dbx.Context) error { - // _, err := ctx.Executor().Exec("INSERT INTO users (name) VALUES (?)", "John") + // _, err := ctx.Executor().ExecContext(ctx, "INSERT INTO users (name) VALUES (?)", "John") // return err // } Operation func(ctx Context) error @@ -146,7 +144,7 @@ type ( // // Example: // op := func(ctx dbx.Context) (int64, error) { - // result, err := ctx.Executor().Exec("INSERT INTO users (name) VALUES (?)", "John") + // result, err := ctx.Executor().ExecContext(ctx, "INSERT INTO users (name) VALUES (?)", "John") // if err != nil { // return 0, err // } diff --git a/options.go b/options.go index b693c22..94b3f0e 100644 --- a/options.go +++ b/options.go @@ -46,7 +46,8 @@ func newOptions(setters []Option) *options { // WithIsolationLevel sets the isolation level for the transaction. // This option configures how the transaction isolates its operations -// from other concurrent transactions. +// from other concurrent transactions. The option is applied only when a new +// transaction is created; it cannot change an existing reused transaction. // // Parameters: // - level: The SQL isolation level (e.g., sql.LevelReadCommitted, sql.LevelSerializable) @@ -67,7 +68,8 @@ func WithIsolationLevel(level sql.IsolationLevel) Option { // WithReadOnly sets the read-only flag for the transaction. // Read-only transactions can provide performance benefits and prevent -// accidental data modifications. +// accidental data modifications. The option is applied only when a new +// transaction is created; it cannot change an existing reused transaction. // // Parameters: // - readOnly: true to make the transaction read-only, false for read-write @@ -87,13 +89,13 @@ func WithReadOnly(readOnly bool) Option { } } -// WithNewTransaction forces the creation of a new transaction even if there -// is an existing transaction in the context. This is useful when you need -// a separate transaction scope that can be committed or rolled back -// independently of the outer transaction. +// WithNewTransaction forces the creation of an independent transaction even if +// there is an existing transaction in the context. The new transaction is not a +// nested transaction or savepoint: it can be committed or rolled back independently +// and may use a different connection from the database pool. // // By default, dbx reuses existing transactions found in the context to avoid -// nested transaction issues. Use this option when you explicitly need a new +// nested transaction issues. Use this option when you explicitly need an independent // transaction boundary. // // Returns: diff --git a/options_test.go b/options_test.go index 446899b..8bfb43d 100644 --- a/options_test.go +++ b/options_test.go @@ -3,232 +3,179 @@ package dbx_test import ( "context" "database/sql" + "errors" "testing" "github.com/DATA-DOG/go-sqlmock" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/ziflex/dbx" ) func TestTransactionOptions(t *testing.T) { - t.Run("WithIsolationLevel should set isolation level", func(t *testing.T) { - mockDB, mock, err := sqlmock.New() - require.NoError(t, err) - defer mockDB.Close() - - mock.ExpectBegin().WillReturnError(nil) // Default expectation - mock.ExpectExec("SELECT 1").WillReturnResult(sqlmock.NewResult(1, 1)) - mock.ExpectCommit() - - db := dbx.New(mockDB) - ctx := context.Background() - - err = dbx.Transaction(ctx, db, func(c dbx.Context) error { - _, e := c.Executor().Exec("SELECT 1") - return e - }, dbx.WithIsolationLevel(sql.LevelReadCommitted)) - - assert.NoError(t, err) - assert.NoError(t, mock.ExpectationsWereMet()) - }) - - t.Run("WithReadOnly should set read-only flag", func(t *testing.T) { - mockDB, mock, err := sqlmock.New() - require.NoError(t, err) - defer mockDB.Close() + t.Run("passes isolation and read-only options to BeginTx", func(t *testing.T) { + database, mock := newSQLMock(t) + beginner := &beginnerOnly{db: database} mock.ExpectBegin() - mock.ExpectQuery("SELECT \\* FROM users").WillReturnRows(sqlmock.NewRows([]string{"id", "name"})) mock.ExpectCommit() - db := dbx.New(mockDB) - ctx := context.Background() - - err = dbx.Transaction(ctx, db, func(c dbx.Context) error { - _, e := c.Executor().Query("SELECT * FROM users") - return e - }, dbx.WithReadOnly(true)) + err := dbx.Transaction( + context.Background(), + beginner, + func(dbx.Context) error { return nil }, + dbx.WithIsolationLevel(sql.LevelSerializable), + dbx.WithReadOnly(true), + ) - assert.NoError(t, err) - assert.NoError(t, mock.ExpectationsWereMet()) + require.NoError(t, err) + assert.Equal(t, 1, beginner.beginTxCalls) + assert.Equal(t, sql.LevelSerializable, beginner.lastOptions.Isolation) + assert.True(t, beginner.lastOptions.ReadOnly) }) - t.Run("WithNewTransaction should create new transaction even if one exists", func(t *testing.T) { - mockDB, mock, err := sqlmock.New() - require.NoError(t, err) - defer mockDB.Close() + t.Run("creates an independent transaction when requested", func(t *testing.T) { + database, mock := newSQLMock(t) + db := dbx.New(database) - // Expect two separate transactions - mock.ExpectBegin() // First transaction + mock.ExpectBegin() mock.ExpectExec("SELECT 1").WillReturnResult(sqlmock.NewResult(1, 1)) - mock.ExpectBegin() // Second transaction (new one) + mock.ExpectBegin() mock.ExpectExec("SELECT 2").WillReturnResult(sqlmock.NewResult(1, 1)) - mock.ExpectCommit() // Second transaction commit - mock.ExpectCommit() // First transaction commit + mock.ExpectCommit() + mock.ExpectCommit() + + err := dbx.Transaction(context.Background(), db, func(outer dbx.Context) error { + if _, err := outer.Executor().Exec("SELECT 1"); err != nil { + return err + } - db := dbx.New(mockDB) - ctx := context.Background() + return dbx.Transaction(outer, db, func(inner dbx.Context) error { + assert.NotEqual(t, outer.Executor(), inner.Executor()) - err = dbx.Transaction(ctx, db, func(c1 dbx.Context) error { - c1.Executor().Exec("SELECT 1") + _, err := inner.Executor().Exec("SELECT 2") - // This should create a NEW transaction instead of reusing - return dbx.Transaction(c1, db, func(c2 dbx.Context) error { - c2.Executor().Exec("SELECT 2") - return nil + return err }, dbx.WithNewTransaction()) }) - assert.NoError(t, err) - assert.NoError(t, mock.ExpectationsWereMet()) - }) - - t.Run("multiple options should work together", func(t *testing.T) { - mockDB, mock, err := sqlmock.New() require.NoError(t, err) - defer mockDB.Close() - - mock.ExpectBegin() - mock.ExpectQuery("SELECT \\* FROM users").WillReturnRows(sqlmock.NewRows([]string{"id"})) - mock.ExpectCommit() - - db := dbx.New(mockDB) - ctx := context.Background() - - err = dbx.Transaction(ctx, db, func(c dbx.Context) error { - _, e := c.Executor().Query("SELECT * FROM users") - return e - }, dbx.WithIsolationLevel(sql.LevelSerializable), dbx.WithReadOnly(true)) - - assert.NoError(t, err) - assert.NoError(t, mock.ExpectationsWereMet()) }) } func TestTransactionWithResult(t *testing.T) { - t.Run("should return result on success", func(t *testing.T) { - mockDB, mock, err := sqlmock.New() - require.NoError(t, err) - defer mockDB.Close() + t.Run("returns the operation result on success", func(t *testing.T) { + database, mock := newSQLMock(t) + db := dbx.New(database) + rows := sqlmock.NewRows([]string{testCountColumn}).AddRow(5) - rows := sqlmock.NewRows([]string{"count"}).AddRow(5) mock.ExpectBegin() mock.ExpectQuery("SELECT COUNT\\(\\*\\) FROM users").WillReturnRows(rows) mock.ExpectCommit() - db := dbx.New(mockDB) - ctx := context.Background() - - count, err := dbx.TransactionWithResult(ctx, db, func(c dbx.Context) (int, error) { + count, err := dbx.TransactionWithResult(context.Background(), db, func(ctx dbx.Context) (int, error) { var count int - err := c.Executor().QueryRow("SELECT COUNT(*) FROM users").Scan(&count) + err := ctx.Executor().QueryRow("SELECT COUNT(*) FROM users").Scan(&count) + return count, err }) - assert.NoError(t, err) + require.NoError(t, err) assert.Equal(t, 5, count) - assert.NoError(t, mock.ExpectationsWereMet()) }) - t.Run("should return zero value on error", func(t *testing.T) { - mockDB, mock, err := sqlmock.New() - require.NoError(t, err) - defer mockDB.Close() + t.Run("returns the zero value on an operation error", func(t *testing.T) { + database, mock := newSQLMock(t) + db := dbx.New(database) + operationErr := errors.New("query users") - testErr := assert.AnError mock.ExpectBegin() - mock.ExpectQuery("SELECT COUNT\\(\\*\\) FROM users").WillReturnError(testErr) + mock.ExpectQuery("SELECT COUNT\\(\\*\\) FROM users").WillReturnError(operationErr) mock.ExpectRollback() - db := dbx.New(mockDB) - ctx := context.Background() - - count, err := dbx.TransactionWithResult(ctx, db, func(c dbx.Context) (int, error) { + count, err := dbx.TransactionWithResult(context.Background(), db, func(ctx dbx.Context) (int, error) { var count int - err := c.Executor().QueryRow("SELECT COUNT(*) FROM users").Scan(&count) + err := ctx.Executor().QueryRow("SELECT COUNT(*) FROM users").Scan(&count) + return count, err }) - assert.Error(t, err) - assert.Equal(t, 0, count) // zero value for int - assert.NoError(t, mock.ExpectationsWereMet()) + require.ErrorIs(t, err, operationErr) + assert.Zero(t, count) }) - t.Run("should work with custom types", func(t *testing.T) { - type User struct { + t.Run("returns zero and leaves a reused transaction usable after an operation error", func(t *testing.T) { + database, mock := newSQLMock(t) + db := dbx.New(database) + unusedBeginner := &beginnerOnly{db: database} + operationErr := errors.New("nested operation error") + + mock.ExpectBegin() + mock.ExpectExec("SELECT 1").WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() + + err := dbx.Transaction(context.Background(), db, func(outer dbx.Context) error { + result, err := dbx.TransactionWithResult(outer, unusedBeginner, func(inner dbx.Context) (string, error) { + assert.Same(t, outer, inner) + + return "discarded", operationErr + }) + + require.ErrorIs(t, err, operationErr) + assert.Equal(t, operationErr, err) + assert.Empty(t, result) + + _, err = outer.Executor().ExecContext(outer, "SELECT 1") + + return err + }) + + require.NoError(t, err) + assert.Zero(t, unusedBeginner.beginTxCalls) + }) + + t.Run("supports custom result types", func(t *testing.T) { + type user struct { ID int Name string } - mockDB, mock, err := sqlmock.New() - require.NoError(t, err) - defer mockDB.Close() + database, mock := newSQLMock(t) + db := dbx.New(database) + rows := sqlmock.NewRows([]string{testIDColumn, testNameColumn}).AddRow(1, "Alice") - rows := sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "Alice") mock.ExpectBegin() mock.ExpectQuery("SELECT id, name FROM users WHERE id"). WithArgs(1). WillReturnRows(rows) mock.ExpectCommit() - db := dbx.New(mockDB) - ctx := context.Background() + result, err := dbx.TransactionWithResult(context.Background(), db, func(ctx dbx.Context) (user, error) { + var result user + err := ctx.Executor().QueryRow("SELECT id, name FROM users WHERE id = ?", 1). + Scan(&result.ID, &result.Name) - user, err := dbx.TransactionWithResult(ctx, db, func(c dbx.Context) (User, error) { - var user User - err := c.Executor().QueryRow("SELECT id, name FROM users WHERE id = ?", 1). - Scan(&user.ID, &user.Name) - return user, err + return result, err }) - assert.NoError(t, err) - assert.Equal(t, User{ID: 1, Name: "Alice"}, user) - assert.NoError(t, mock.ExpectationsWereMet()) - }) - - t.Run("should handle begin error", func(t *testing.T) { - mockDB, mock, err := sqlmock.New() require.NoError(t, err) - defer mockDB.Close() - - testErr := assert.AnError - mock.ExpectBegin().WillReturnError(testErr) - - db := dbx.New(mockDB) - ctx := context.Background() - - result, err := dbx.TransactionWithResult(ctx, db, func(c dbx.Context) (string, error) { - return "should not reach here", nil - }) - - assert.Error(t, err) - assert.Equal(t, testErr, err) - assert.Equal(t, "", result) // zero value for string - assert.NoError(t, mock.ExpectationsWereMet()) + assert.Equal(t, user{ID: 1, Name: "Alice"}, result) }) - t.Run("should handle commit error", func(t *testing.T) { - mockDB, mock, err := sqlmock.New() - require.NoError(t, err) - defer mockDB.Close() + t.Run("returns the zero value on a commit error", func(t *testing.T) { + database, mock := newSQLMock(t) + db := dbx.New(database) + commitErr := errors.New("commit transaction") - testErr := assert.AnError mock.ExpectBegin() - mock.ExpectExec("SELECT 1").WillReturnResult(sqlmock.NewResult(1, 1)) - mock.ExpectCommit().WillReturnError(testErr) - - db := dbx.New(mockDB) - ctx := context.Background() + mock.ExpectCommit().WillReturnError(commitErr) - result, err := dbx.TransactionWithResult(ctx, db, func(c dbx.Context) (string, error) { - c.Executor().Exec("SELECT 1") - return "success", nil + result, err := dbx.TransactionWithResult(context.Background(), db, func(dbx.Context) (string, error) { + return "result", nil }) - assert.Error(t, err) - assert.Equal(t, testErr, err) - assert.Equal(t, "", result) // zero value for string - assert.NoError(t, mock.ExpectationsWereMet()) + require.ErrorIs(t, err, commitErr) + assert.Empty(t, result) }) } diff --git a/test_helpers_test.go b/test_helpers_test.go new file mode 100644 index 0000000..5772de4 --- /dev/null +++ b/test_helpers_test.go @@ -0,0 +1,53 @@ +package dbx_test + +import ( + "context" + "database/sql" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +const ( + testCountColumn = "count" + testIDColumn = "id" + testNameColumn = "name" +) + +type beginnerOnly struct { + db *sql.DB + beginCalls int + beginTxCalls int + lastOptions sql.TxOptions +} + +func (b *beginnerOnly) Begin() (*sql.Tx, error) { + b.beginCalls++ + + return b.db.Begin() //nolint:noctx // Beginner includes the legacy Begin method for compatibility. +} + +func (b *beginnerOnly) BeginTx(ctx context.Context, opts *sql.TxOptions) (*sql.Tx, error) { + b.beginTxCalls++ + if opts != nil { + b.lastOptions = *opts + } + + return b.db.BeginTx(ctx, opts) +} + +func newSQLMock(t *testing.T) (*sql.DB, sqlmock.Sqlmock) { + t.Helper() + + database, mock, err := sqlmock.New() + require.NoError(t, err) + + t.Cleanup(func() { + assert.NoError(t, mock.ExpectationsWereMet()) + _ = database.Close() // SQL behavior is asserted before best-effort test cleanup. + }) + + return database, mock +} diff --git a/transaction.go b/transaction.go index 93574b0..45450a9 100644 --- a/transaction.go +++ b/transaction.go @@ -2,6 +2,9 @@ package dbx import ( "context" + "database/sql" + "errors" + "fmt" ) // Transaction begins or reuses a transaction, executes the provided operation, @@ -12,6 +15,8 @@ import ( // - If a new transaction is created, it's committed on successful operation or rolled back on error. // - If an existing transaction is reused, commit/rollback is left to the outer transaction. // - Any panic during operation execution triggers rollback if a new transaction was created. +// - If an operation and its rollback both fail, the returned error contains both failures. +// - Transaction options apply only when a new transaction is created. // // Parameters: // - ctx: Parent Go context. @@ -25,9 +30,9 @@ import ( // Example: // // err := dbx.Transaction(ctx, db, func(txCtx dbx.Context) error { -// _, err := txCtx.Executor().Exec("INSERT INTO users (name) VALUES (?)", "John") +// _, err := txCtx.Executor().ExecContext(txCtx, "INSERT INTO users (name) VALUES (?)", "John") // if err != nil { return err } // triggers automatic rollback -// _, err = txCtx.Executor().Exec("INSERT INTO profiles (user_id) VALUES (?)", userID) +// _, err = txCtx.Executor().ExecContext(txCtx, "INSERT INTO profiles (user_id) VALUES (?)", userID) // return err // }) func Transaction(ctx context.Context, beginner Beginner, op Operation, opts ...Option) error { @@ -54,9 +59,13 @@ func Transaction(ctx context.Context, beginner Beginner, op Operation, opts ...O // Example: // // userID, err := dbx.TransactionWithResult(ctx, db, func(txCtx dbx.Context) (int64, error) { -// result, err := txCtx.Executor().Exec("INSERT INTO users (name) VALUES (?)", "John") -// if err != nil { return 0, err } -// return result.LastInsertId() +// var userID int64 +// err := txCtx.Executor().QueryRowContext( +// txCtx, +// "INSERT INTO users (name) VALUES (?) RETURNING id", +// "John", +// ).Scan(&userID) +// return userID, err // }) func TransactionWithResult[T any](ctx context.Context, beginner Beginner, op OperationWithResult[T], setters ...Option) (T, error) { return transactionWithInternal(ctx, beginner, op, setters) @@ -83,54 +92,54 @@ func TransactionWithResult[T any](ctx context.Context, beginner Beginner, op Ope // - T: Operation result (zero value if error). // - error: Any error from transaction handling or op execution. func transactionWithInternal[T any](ctx context.Context, beginner Beginner, op OperationWithResult[T], setters []Option) (T, error) { - var tx Transactor - var createdTx bool - var dbCtx Context + var zero T opts := newOptions(setters) if !opts.AlwaysCreate { - // retrieve existing or create a new context - dbCtx = NewContextFrom(ctx, beginner) - executor := dbCtx.Executor() - - // check if the executor is a transaction - transactor, ok := executor.(Transactor) - - // if the executor is a transaction, use it - if ok { - tx = transactor + // Reuse an existing transaction without requiring beginner to provide + // unrelated context or execution capabilities. + if dbCtx := FromContext(ctx); dbCtx != nil { + if _, ok := dbCtx.Executor().(Transactor); ok { + out, err := op(dbCtx) + if err != nil { + return zero, err + } + + return out, nil + } } } - if tx == nil { - var err error - createdTx = true + tx, err := beginner.BeginTx(ctx, opts.TxOptions) + if err != nil { + return zero, err + } - // create a new transaction - tx, err = beginner.BeginTx(ctx, opts.TxOptions) + dbCtx := NewContext(ctx, tx) + var operationReturned bool - if err != nil { - return *new(T), err + // Keep rollback armed until op returns so panics are cleaned up without + // issuing another rollback after normal lifecycle handling. + defer func() { + if !operationReturned { + _ = tx.Rollback() //nolint:errcheck // A panic must retain its original value. } - - // create a new context with the transaction - dbCtx = NewContext(ctx, tx) - } + }() out, err := op(dbCtx) + operationReturned = true if err != nil { - if createdTx { - tx.Rollback() + rollbackErr := tx.Rollback() + if rollbackErr != nil && !errors.Is(rollbackErr, sql.ErrTxDone) { + return zero, errors.Join(err, fmt.Errorf("rollback transaction: %w", rollbackErr)) } - return *new(T), err + return zero, err } - if createdTx { - if e := tx.Commit(); e != nil { - return *new(T), e - } + if e := tx.Commit(); e != nil { + return zero, e } return out, nil diff --git a/transaction_test.go b/transaction_test.go index cb00c4b..2a09bcc 100644 --- a/transaction_test.go +++ b/transaction_test.go @@ -2,169 +2,207 @@ package dbx_test import ( "context" + "database/sql" "errors" "testing" "github.com/DATA-DOG/go-sqlmock" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/ziflex/dbx" ) -func TestTransaction(test *testing.T) { - test.Run("should handle single transaction", func(t *testing.T) { - dbMock, dmock, _ := sqlmock.New() - defer dbMock.Close() +func TestTransaction(t *testing.T) { + t.Run("commits a successful transaction", func(t *testing.T) { + database, mock := newSQLMock(t) + db := dbx.New(database) - ctx := context.Background() + mock.ExpectBegin() + mock.ExpectExec("SELECT 1").WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() - db := dbx.New(dbMock) - dmock.ExpectBegin() - dmock.ExpectExec("SELECT 1").WillReturnResult(sqlmock.NewResult(1, 1)) - dmock.ExpectExec("SELECT 2").WillReturnResult(sqlmock.NewResult(1, 1)) - dmock.ExpectExec("SELECT 3").WillReturnResult(sqlmock.NewResult(1, 1)) - dmock.ExpectCommit() + err := dbx.Transaction(context.Background(), db, func(ctx dbx.Context) error { + _, err := ctx.Executor().Exec("SELECT 1") - err := dbx.Transaction(ctx, db, func(c dbx.Context) error { - executor := c.Executor() + return err + }) - if _, e := executor.Exec("SELECT 1"); e != nil { - return e - } + require.NoError(t, err) + }) - if _, e := executor.Exec("SELECT 2"); e != nil { - return e - } + t.Run("accepts a minimal Beginner implementation", func(t *testing.T) { + database, mock := newSQLMock(t) + beginner := &beginnerOnly{db: database} - if _, e := executor.Exec("SELECT 3"); e != nil { - return e - } + mock.ExpectBegin() + mock.ExpectExec("SELECT 1").WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() - return nil + err := dbx.Transaction(context.Background(), beginner, func(ctx dbx.Context) error { + _, err := ctx.Executor().Exec("SELECT 1") + + return err }) - assert.NoError(t, err) + require.NoError(t, err) + assert.Zero(t, beginner.beginCalls) + assert.Equal(t, 1, beginner.beginTxCalls) }) - test.Run("should handle tx begin errors", func(t *testing.T) { - dbMock, dmock, _ := sqlmock.New() - defer dbMock.Close() - - ctx := context.Background() + t.Run("returns a begin error without running the operation", func(t *testing.T) { + database, mock := newSQLMock(t) + db := dbx.New(database) + beginErr := errors.New("begin transaction") + operationCalled := false - testErr := errors.New("test error") - db := dbx.New(dbMock) - dmock.ExpectBegin().WillReturnError(testErr) + mock.ExpectBegin().WillReturnError(beginErr) - err := dbx.Transaction(ctx, db, func(c dbx.Context) error { - executor := c.Executor() - executor.Exec("SELECT 1") + err := dbx.Transaction(context.Background(), db, func(dbx.Context) error { + operationCalled = true return nil }) - assert.Error(t, err) - assert.Equal(t, testErr, err) + require.ErrorIs(t, err, beginErr) + assert.False(t, operationCalled) }) - test.Run("should handle tx commit errors", func(t *testing.T) { - dbMock, dmock, _ := sqlmock.New() - defer dbMock.Close() + t.Run("returns a commit error", func(t *testing.T) { + database, mock := newSQLMock(t) + db := dbx.New(database) + commitErr := errors.New("commit transaction") - ctx := context.Background() + mock.ExpectBegin() + mock.ExpectExec("SELECT 1").WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit().WillReturnError(commitErr) - testErr := errors.New("test error") - db := dbx.New(dbMock) - dmock.ExpectBegin() - dmock.ExpectExec("SELECT 1").WillReturnResult(sqlmock.NewResult(1, 1)) - dmock.ExpectCommit().WillReturnError(testErr) + err := dbx.Transaction(context.Background(), db, func(ctx dbx.Context) error { + _, err := ctx.Executor().Exec("SELECT 1") - err := dbx.Transaction(ctx, db, func(c dbx.Context) error { - executor := c.Executor() - executor.Exec("SELECT 1") + return err + }) - return nil + require.ErrorIs(t, err, commitErr) + }) + + t.Run("rolls back on an operation error", func(t *testing.T) { + database, mock := newSQLMock(t) + db := dbx.New(database) + operationErr := errors.New("run operation") + + mock.ExpectBegin() + mock.ExpectRollback() + + err := dbx.Transaction(context.Background(), db, func(dbx.Context) error { + return operationErr }) - assert.Error(t, err) - assert.Equal(t, testErr, err) + require.ErrorIs(t, err, operationErr) + assert.Equal(t, operationErr, err) }) - test.Run("should handle single transaction and rollback on errors", func(t *testing.T) { - dbMock, dmock, _ := sqlmock.New() - defer dbMock.Close() + t.Run("joins operation and rollback errors", func(t *testing.T) { + database, mock := newSQLMock(t) + db := dbx.New(database) + operationErr := errors.New("run operation") + rollbackErr := errors.New("rollback transaction") + + mock.ExpectBegin() + mock.ExpectRollback().WillReturnError(rollbackErr) - ctx := context.Background() + err := dbx.Transaction(context.Background(), db, func(dbx.Context) error { + return operationErr + }) - testErr := errors.New("test error") - db := dbx.New(dbMock) - dmock.ExpectBegin() - dmock.ExpectExec("SELECT 1").WillReturnResult(sqlmock.NewResult(1, 1)) - dmock.ExpectExec("SELECT 2").WillReturnResult(sqlmock.NewResult(1, 1)) - dmock.ExpectExec("SELECT 3").WillReturnResult(sqlmock.NewResult(1, 1)) - dmock.ExpectRollback() + require.ErrorIs(t, err, operationErr) + require.ErrorIs(t, err, rollbackErr) + }) - err := dbx.Transaction(ctx, db, func(c dbx.Context) error { - executor := c.Executor() - executor.Exec("SELECT 1") - executor.Exec("SELECT 2") - executor.Exec("SELECT 3") + t.Run("ignores ErrTxDone while returning the operation error", func(t *testing.T) { + database, mock := newSQLMock(t) + db := dbx.New(database) + operationErr := errors.New("run operation") - return testErr + mock.ExpectBegin() + mock.ExpectRollback().WillReturnError(sql.ErrTxDone) + + err := dbx.Transaction(context.Background(), db, func(dbx.Context) error { + return operationErr }) - assert.Error(t, err) - assert.Equal(t, testErr, err) + require.ErrorIs(t, err, operationErr) + assert.Equal(t, operationErr, err) + assert.NotErrorIs(t, err, sql.ErrTxDone) }) - test.Run("should reuse nested transaction", func(t *testing.T) { - dbMock, dmock, _ := sqlmock.New() - defer dbMock.Close() + t.Run("rolls back and preserves a panic", func(t *testing.T) { + database, mock := newSQLMock(t) + db := dbx.New(database) - ctx := context.Background() + mock.ExpectBegin() + mock.ExpectRollback() - db := dbx.New(dbMock) - dmock.ExpectBegin() - dmock.ExpectExec("SELECT 1").WillReturnResult(sqlmock.NewResult(1, 1)) - dmock.ExpectExec("SELECT 2").WillReturnResult(sqlmock.NewResult(1, 1)) - dmock.ExpectExec("SELECT 3").WillReturnResult(sqlmock.NewResult(1, 1)) - dmock.ExpectExec("SELECT 4").WillReturnResult(sqlmock.NewResult(1, 1)) - dmock.ExpectExec("SELECT 5").WillReturnResult(sqlmock.NewResult(1, 1)) - dmock.ExpectExec("SELECT 6").WillReturnResult(sqlmock.NewResult(1, 1)) - dmock.ExpectCommit() + assert.PanicsWithValue(t, "operation panic", func() { + _ = dbx.Transaction(context.Background(), db, func(dbx.Context) error { //nolint:errcheck // The operation must panic before returning. + panic("operation panic") + }) + }) + }) - err := dbx.Transaction(ctx, db, func(c1 dbx.Context) error { - executor := c1.Executor() - executor.Exec("SELECT 1") - executor.Exec("SELECT 2") - executor.Exec("SELECT 3") + t.Run("reuses an existing transaction", func(t *testing.T) { + database, mock := newSQLMock(t) + db := dbx.New(database) + unusedBeginner := &beginnerOnly{db: database} - return dbx.Transaction(c1, db, func(c2 dbx.Context) error { - executor2 := c2.Executor() - executor2.Exec("SELECT 4") + mock.ExpectBegin() + mock.ExpectExec("SELECT 1").WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectExec("SELECT 2").WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() - assert.Equal(t, c1, c2) - assert.Equal(t, executor, executor2) + err := dbx.Transaction(context.Background(), db, func(outer dbx.Context) error { + if _, err := outer.Executor().Exec("SELECT 1"); err != nil { + return err + } - return dbx.Transaction(c2, db, func(c3 dbx.Context) error { - executor3 := c3.Executor() - executor3.Exec("SELECT 5") + return dbx.Transaction(outer, unusedBeginner, func(inner dbx.Context) error { + assert.Same(t, outer, inner) + assert.Equal(t, outer.Executor(), inner.Executor()) - assert.Equal(t, c2, c3) - assert.Equal(t, executor2, executor3) + _, err := inner.Executor().Exec("SELECT 2") - return dbx.Transaction(c3, db, func(c4 dbx.Context) error { - executor4 := c4.Executor() - executor4.Exec("SELECT 6") + return err + }, dbx.WithReadOnly(true)) + }) - assert.Equal(t, c3, c4) - assert.Equal(t, executor3, executor4) + require.NoError(t, err) + assert.Zero(t, unusedBeginner.beginTxCalls) + }) - return nil - }) + t.Run("leaves a reused transaction usable after a panic", func(t *testing.T) { + database, mock := newSQLMock(t) + db := dbx.New(database) + unusedBeginner := &beginnerOnly{db: database} + panicValue := errors.New("nested operation panic") + + mock.ExpectBegin() + mock.ExpectExec("SELECT 1").WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() + + err := dbx.Transaction(context.Background(), db, func(outer dbx.Context) error { + assert.PanicsWithValue(t, panicValue, func() { + _ = dbx.Transaction(outer, unusedBeginner, func(inner dbx.Context) error { //nolint:errcheck // The operation must panic before returning. + assert.Same(t, outer, inner) + panic(panicValue) }) }) + + _, err := outer.Executor().ExecContext(outer, "SELECT 1") + + return err }) - assert.NoError(t, err) + require.NoError(t, err) + assert.Zero(t, unusedBeginner.beginTxCalls) }) }