Merge pull request #9 from maxcask/patch-1

kata 19: The Cleanup Chain
This commit is contained in:
medunes
2026-07-09 19:15:10 +02:00
committed by GitHub
2 changed files with 693 additions and 0 deletions
@@ -0,0 +1,146 @@
package defer_cleanup_chain
import (
"bufio"
"context"
"errors"
"fmt"
"os"
)
// Tx models the transaction lifecycle we need for this kata.
// Keeping this small makes it easy to provide fake implementations in tests.
type Tx interface {
Commit() error
Rollback() error
}
// Rows models a streaming result set.
// Values returns the current row payload after Next() moved to a valid item.
type Rows interface {
Next() bool
Values() (id int, payload string, err error)
Err() error
Close() error
}
// DB is the minimal database contract needed by BackupDatabase.
// The implementation can be real SQL, in-memory, or fully mocked in tests.
type DB interface {
Begin(ctx context.Context) (Tx, error)
QueryBackupRows(ctx context.Context, tx Tx) (Rows, error)
Close() error
}
// Connector abstracts DB creation so tests can inject controlled failures.
type Connector func(ctx context.Context, dbURL string) (DB, error)
// defaultConnector is intentionally unconfigured.
// In real code this would be wired to an actual driver in init/bootstrap code.
var defaultConnector Connector = func(ctx context.Context, dbURL string) (DB, error) {
return nil, errors.New("no database connector configured")
}
// BackupDatabase performs the backup flow while preserving both operation and
// cleanup errors. It uses a named return error so deferred cleanups can append
// their own failures via errors.Join.
func BackupDatabase(ctx context.Context, dbURL, filename string) (err error) {
return backupDatabase(ctx, dbURL, filename, defaultConnector, os.Create)
}
// backupDatabase is split from BackupDatabase to keep the public API clean while
// allowing tests to inject open/connect behavior and force edge cases.
func backupDatabase(
ctx context.Context,
dbURL, filename string,
connect Connector,
openFile func(name string) (*os.File, error),
) (err error) {
// Acquire #1: output file.
// Defer is registered immediately so every return path closes the file.
file, err := openFile(filename)
if err != nil {
return fmt.Errorf("open output file: %w", err)
}
defer func() {
if closeErr := file.Close(); closeErr != nil {
err = errors.Join(err, fmt.Errorf("close output file: %w", closeErr))
}
}()
// Wrap the file with a buffered writer. This is another acquired resource:
// data can remain in memory until Flush runs, so we must defer Flush as well.
writer := bufio.NewWriter(file)
defer func() {
if flushErr := writer.Flush(); flushErr != nil {
err = errors.Join(err, fmt.Errorf("flush writer: %w", flushErr))
}
}()
// Acquire #2: DB connection.
db, err := connect(ctx, dbURL)
if err != nil {
return fmt.Errorf("connect DB: %w", err)
}
defer func() {
if closeErr := db.Close(); closeErr != nil {
err = errors.Join(err, fmt.Errorf("close DB connection: %w", closeErr))
}
}()
// Acquire #3: transaction.
// We avoid manual rollback branches by using a committed flag consumed by defer.
tx, err := db.Begin(ctx)
if err != nil {
return fmt.Errorf("begin transaction: %w", err)
}
committed := false
defer func() {
if committed {
return
}
if rollbackErr := tx.Rollback(); rollbackErr != nil {
err = errors.Join(err, fmt.Errorf("rollback transaction: %w", rollbackErr))
}
}()
// Acquire #4: row stream (logical cursor/network resource).
rows, err := db.QueryBackupRows(ctx, tx)
if err != nil {
return fmt.Errorf("query backup rows: %w", err)
}
defer func() {
if closeErr := rows.Close(); closeErr != nil {
err = errors.Join(err, fmt.Errorf("close rows: %w", closeErr))
}
}()
// Stream each row to the backup file.
// We intentionally do not defer anything inside this loop.
for rows.Next() {
if ctxErr := ctx.Err(); ctxErr != nil {
return fmt.Errorf("backup canceled: %w", ctxErr)
}
id, payload, rowErr := rows.Values()
if rowErr != nil {
return fmt.Errorf("read row values: %w", rowErr)
}
if _, writeErr := fmt.Fprintf(writer, "%d,%s\n", id, payload); writeErr != nil {
return fmt.Errorf("write backup row: %w", writeErr)
}
}
if rowsErr := rows.Err(); rowsErr != nil {
return fmt.Errorf("iterate backup rows: %w", rowsErr)
}
// Commit is the success boundary.
// If commit fails, committed remains false and deferred rollback still runs.
if commitErr := tx.Commit(); commitErr != nil {
return fmt.Errorf("commit transaction: %w", commitErr)
}
committed = true
return nil
}
@@ -0,0 +1,547 @@
package defer_cleanup_chain
import (
"context"
"errors"
"os"
"path/filepath"
"runtime"
"strconv"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
type fakeRow struct {
id int
payload string
err error
}
type fakeRows struct {
data []fakeRow
nextIdx int
iterErr error
closeErr error
closed bool
}
func (r *fakeRows) Next() bool {
if r.nextIdx >= len(r.data) {
return false
}
r.nextIdx++
return true
}
func (r *fakeRows) Values() (int, string, error) {
if r.nextIdx == 0 || r.nextIdx > len(r.data) {
return 0, "", errors.New("values called without current row")
}
row := r.data[r.nextIdx-1]
if row.err != nil {
return 0, "", row.err
}
return row.id, row.payload, nil
}
func (r *fakeRows) Err() error {
return r.iterErr
}
func (r *fakeRows) Close() error {
r.closed = true
return r.closeErr
}
type fakeTx struct {
commitErr error
rollbackErr error
commitCalled int
rollbackCalled int
}
func (tx *fakeTx) Commit() error {
tx.commitCalled++
return tx.commitErr
}
func (tx *fakeTx) Rollback() error {
tx.rollbackCalled++
return tx.rollbackErr
}
type fakeDB struct {
tx Tx
rows Rows
beginErr error
queryErr error
closeErr error
beginCalled int
queryCalled int
closeCalled int
}
func (db *fakeDB) Begin(context.Context) (Tx, error) {
db.beginCalled++
if db.beginErr != nil {
return nil, db.beginErr
}
return db.tx, nil
}
func (db *fakeDB) QueryBackupRows(context.Context, Tx) (Rows, error) {
db.queryCalled++
if db.queryErr != nil {
return nil, db.queryErr
}
return db.rows, nil
}
func (db *fakeDB) Close() error {
db.closeCalled++
return db.closeErr
}
func TestBackupDatabase_Matrix(t *testing.T) {
type tc struct {
name string
connectErr error
openErr error
beginErr error
queryErr error
rows []fakeRow
iterErr error
rowsCloseErr error
commitErr error
rollbackErr error
dbCloseErr error
cancelCtx bool
openReadOnly bool
largePayloadRow bool
wantErr bool
wantErrContains []string
wantErrIs []error
wantCommitCalls int
wantRollbackCalls int
wantDBCloseCalls int
wantRowsClosed bool
wantOutput string
}
errConnectBoom := errors.New("connect boom")
errOpenBoom := errors.New("open boom")
errBeginBoom := errors.New("begin boom")
errQueryBoom := errors.New("query boom")
errRollbackBoom := errors.New("rollback boom")
errDBCloseBoom := errors.New("db close boom")
errValuesBoom := errors.New("values boom")
errIterateBoom := errors.New("iterate boom")
errCommitBoom := errors.New("commit boom")
tests := []tc{
{
name: "success commits and writes all rows",
rows: []fakeRow{{id: 1, payload: "alpha"}, {id: 2, payload: "beta"}},
wantCommitCalls: 1,
wantRollbackCalls: 0,
wantDBCloseCalls: 1,
wantRowsClosed: true,
wantOutput: "1,alpha\n2,beta\n",
},
{
name: "connect failure returns wrapped error",
connectErr: errConnectBoom,
wantErr: true,
wantErrContains: []string{"connect DB"},
wantErrIs: []error{errConnectBoom},
wantCommitCalls: 0,
wantRollbackCalls: 0,
wantDBCloseCalls: 0,
wantRowsClosed: false,
},
{
name: "begin failure keeps cleanup error with join",
beginErr: errBeginBoom,
dbCloseErr: errDBCloseBoom,
wantErr: true,
wantErrContains: []string{"begin transaction", "close DB connection"},
wantErrIs: []error{errBeginBoom, errDBCloseBoom},
wantCommitCalls: 0,
wantRollbackCalls: 0,
wantDBCloseCalls: 1,
wantRowsClosed: false,
},
{
name: "query failure triggers rollback and preserves cleanup failures",
queryErr: errQueryBoom,
rollbackErr: errRollbackBoom,
dbCloseErr: errDBCloseBoom,
wantErr: true,
wantErrContains: []string{"query backup rows", "rollback transaction", "close DB connection"},
wantErrIs: []error{errQueryBoom, errRollbackBoom, errDBCloseBoom},
wantCommitCalls: 0,
wantRollbackCalls: 1,
wantDBCloseCalls: 1,
wantRowsClosed: false,
},
{
name: "row values failure stops stream and rolls back",
rows: []fakeRow{{id: 1, payload: "ok"}, {err: errValuesBoom}},
wantErr: true,
wantErrContains: []string{"read row values"},
wantErrIs: []error{errValuesBoom},
wantCommitCalls: 0,
wantRollbackCalls: 1,
wantDBCloseCalls: 1,
wantRowsClosed: true,
},
{
name: "rows err after iteration rolls back",
rows: []fakeRow{{id: 1, payload: "ok"}},
iterErr: errIterateBoom,
wantErr: true,
wantErrContains: []string{"iterate backup rows"},
wantErrIs: []error{errIterateBoom},
wantCommitCalls: 0,
wantRollbackCalls: 1,
wantDBCloseCalls: 1,
wantRowsClosed: true,
},
{
name: "commit failure keeps rollback and close errors",
rows: []fakeRow{{id: 1, payload: "ok"}},
commitErr: errCommitBoom,
rollbackErr: errRollbackBoom,
dbCloseErr: errDBCloseBoom,
wantErr: true,
wantErrContains: []string{"commit transaction", "rollback transaction", "close DB connection"},
wantErrIs: []error{errCommitBoom, errRollbackBoom, errDBCloseBoom},
wantCommitCalls: 1,
wantRollbackCalls: 1,
wantDBCloseCalls: 1,
wantRowsClosed: true,
},
{
name: "context cancellation in loop returns cancellation error",
rows: []fakeRow{{id: 1, payload: "ok"}},
cancelCtx: true,
wantErr: true,
wantErrContains: []string{"backup canceled", context.Canceled.Error()},
wantErrIs: []error{context.Canceled},
wantCommitCalls: 0,
wantRollbackCalls: 1,
wantDBCloseCalls: 1,
wantRowsClosed: true,
},
{
name: "flush failure after successful work is returned",
rows: []fakeRow{{id: 1, payload: "ok"}},
openReadOnly: true,
wantErr: true,
wantErrContains: []string{"flush writer"},
wantCommitCalls: 1,
wantRollbackCalls: 0,
wantDBCloseCalls: 1,
wantRowsClosed: true,
},
{
name: "write failure in loop is returned",
rows: []fakeRow{{id: 1, payload: "ok"}},
openReadOnly: true,
largePayloadRow: true,
wantErr: true,
wantErrContains: []string{"write backup row"},
wantCommitCalls: 0,
wantRollbackCalls: 1,
wantDBCloseCalls: 1,
wantRowsClosed: true,
},
{
name: "open file failure short-circuits immediately",
openErr: errOpenBoom,
wantErr: true,
wantErrContains: []string{"open output file"},
wantErrIs: []error{errOpenBoom},
wantCommitCalls: 0,
wantRollbackCalls: 0,
wantDBCloseCalls: 0,
wantRowsClosed: false,
},
{
name: "rows close failure is preserved",
rows: []fakeRow{{id: 1, payload: "ok"}},
rowsCloseErr: errors.New("rows close boom"),
wantErr: true,
wantErrContains: []string{"close rows"},
wantCommitCalls: 1,
wantRollbackCalls: 0,
wantDBCloseCalls: 1,
wantRowsClosed: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
dir := t.TempDir()
filename := filepath.Join(dir, "backup.csv")
if tt.openReadOnly {
require.NoError(t, os.WriteFile(filename, nil, 0o644))
}
rowsData := tt.rows
if tt.largePayloadRow {
rowsData = []fakeRow{{id: 1, payload: strings.Repeat("x", 9000)}}
}
tx := &fakeTx{commitErr: tt.commitErr, rollbackErr: tt.rollbackErr}
rows := &fakeRows{data: rowsData, iterErr: tt.iterErr, closeErr: tt.rowsCloseErr}
db := &fakeDB{
tx: tx,
rows: rows,
beginErr: tt.beginErr,
queryErr: tt.queryErr,
closeErr: tt.dbCloseErr,
}
connect := func(context.Context, string) (DB, error) {
if tt.connectErr != nil {
return nil, tt.connectErr
}
return db, nil
}
openFile := os.Create
if tt.openErr != nil {
openFile = func(string) (*os.File, error) {
return nil, tt.openErr
}
} else if tt.openReadOnly {
openFile = func(name string) (*os.File, error) {
return os.Open(name)
}
}
ctx := context.Background()
if tt.cancelCtx {
ctxWithCancel, cancel := context.WithCancel(context.Background())
cancel()
ctx = ctxWithCancel
}
err := backupDatabase(ctx, "db://test", filename, connect, openFile)
if tt.wantErr {
require.Error(t, err)
for _, substr := range tt.wantErrContains {
assert.ErrorContains(t, err, substr)
}
for _, targetErr := range tt.wantErrIs {
assert.ErrorIs(t, err, targetErr)
}
} else {
require.NoError(t, err)
}
assert.Equal(t, tt.wantCommitCalls, tx.commitCalled)
assert.Equal(t, tt.wantRollbackCalls, tx.rollbackCalled)
assert.Equal(t, tt.wantDBCloseCalls, db.closeCalled)
assert.Equal(t, tt.wantRowsClosed, rows.closed)
if tt.wantOutput != "" {
got, readErr := os.ReadFile(filename)
require.NoError(t, readErr)
assert.Equal(t, tt.wantOutput, string(got))
}
})
}
}
func TestBackupDatabase_SelfCorrection_BeginFails_ClosesFileAndDB(t *testing.T) {
errBegin := errors.New("begin failed")
dir := t.TempDir()
filename := filepath.Join(dir, "backup.csv")
var openedFile *os.File
openFile := func(name string) (*os.File, error) {
f, err := os.Create(name)
if err != nil {
return nil, err
}
openedFile = f
return f, nil
}
db := &fakeDB{beginErr: errBegin}
connect := func(context.Context, string) (DB, error) {
return db, nil
}
err := backupDatabase(context.Background(), "db://test", filename, connect, openFile)
require.Error(t, err)
assert.ErrorContains(t, err, "begin transaction")
assert.ErrorIs(t, err, errBegin)
assert.Equal(t, 1, db.closeCalled, "db close must run even when Begin fails")
require.NotNil(t, openedFile)
assert.ErrorIs(t, openedFile.Close(), os.ErrClosed, "file must already be closed by defer")
}
func TestBackupDatabase_SelfCorrection_CommitFailsAndFileCloseFails_BothErrorsReturned(t *testing.T) {
errCommit := errors.New("commit failed")
badFile := os.NewFile(^uintptr(0), "invalid-fd")
openFile := func(string) (*os.File, error) {
return badFile, nil
}
rows := &fakeRows{}
tx := &fakeTx{commitErr: errCommit}
db := &fakeDB{tx: tx, rows: rows}
connect := func(context.Context, string) (DB, error) {
return db, nil
}
err := backupDatabase(context.Background(), "db://test", "ignored.csv", connect, openFile)
require.Error(t, err)
assert.ErrorContains(t, err, "commit transaction")
assert.ErrorContains(t, err, "close output file")
assert.ErrorIs(t, err, errCommit)
assert.Equal(t, 1, tx.commitCalled)
assert.Equal(t, 1, tx.rollbackCalled, "rollback must run after failed commit")
}
func TestBackupDatabase_SelfCorrection_NoFDLeak_1000Runs(t *testing.T) {
before, err := countOpenFDs()
require.NoError(t, err)
for i := 0; i < 1000; i++ {
dir := t.TempDir()
filename := filepath.Join(dir, "backup.csv")
rows := &fakeRows{}
tx := &fakeTx{}
db := &fakeDB{tx: tx, rows: rows}
connect := func(context.Context, string) (DB, error) {
return db, nil
}
err := backupDatabase(context.Background(), "db://test", filename, connect, os.Create)
require.NoError(t, err)
}
runtime.GC()
after, err := countOpenFDs()
require.NoError(t, err)
// Allow a tiny delta for runtime/testing internals while asserting no growth trend.
assert.LessOrEqual(t, after, before+3, "open file descriptors should not grow after 1000 runs")
}
func TestBackupDatabase_PublicWrapper_Success(t *testing.T) {
dir := t.TempDir()
filename := filepath.Join(dir, "backup.csv")
oldConnector := defaultConnector
t.Cleanup(func() {
defaultConnector = oldConnector
})
tx := &fakeTx{}
rows := &fakeRows{data: []fakeRow{{id: 7, payload: "public"}}}
db := &fakeDB{tx: tx, rows: rows}
defaultConnector = func(context.Context, string) (DB, error) {
return db, nil
}
err := BackupDatabase(context.Background(), "db://public", filename)
require.NoError(t, err)
assert.Equal(t, 1, tx.commitCalled)
assert.Equal(t, 0, tx.rollbackCalled)
assert.Equal(t, 1, db.closeCalled)
assert.True(t, rows.closed)
got, readErr := os.ReadFile(filename)
require.NoError(t, readErr)
assert.Equal(t, "7,public\n", string(got))
}
func TestBackupDatabase_PublicWrapper_ConnectorError(t *testing.T) {
errConnect := errors.New("public connector fail")
oldConnector := defaultConnector
t.Cleanup(func() {
defaultConnector = oldConnector
})
defaultConnector = func(context.Context, string) (DB, error) {
return nil, errConnect
}
filename := filepath.Join(t.TempDir(), "backup.csv")
err := BackupDatabase(context.Background(), "db://public", filename)
require.Error(t, err)
assert.ErrorContains(t, err, "connect DB")
assert.ErrorIs(t, err, errConnect)
}
func TestBackupDatabase_PublicWrapper_DefaultConnectorUnconfigured(t *testing.T) {
oldConnector := defaultConnector
t.Cleanup(func() {
defaultConnector = oldConnector
})
defaultConnector = oldConnector
filename := filepath.Join(t.TempDir(), "backup.csv")
err := BackupDatabase(context.Background(), "db://public", filename)
require.Error(t, err)
assert.ErrorContains(t, err, "connect DB")
assert.ErrorContains(t, err, "no database connector configured")
}
func countOpenFDs() (int, error) {
var lastErr error
for _, path := range []string{"/proc/self/fd", "/dev/fd"} {
count, err := countNumericDirEntries(path)
if err == nil {
return count, nil
}
lastErr = err
}
if lastErr == nil {
lastErr = errors.New("no fd directory available")
}
return 0, lastErr
}
func countNumericDirEntries(path string) (int, error) {
dir, err := os.Open(path)
if err != nil {
return 0, err
}
defer dir.Close()
names, err := dir.Readdirnames(-1)
if err != nil {
return 0, err
}
count := 0
for _, name := range names {
if _, convErr := strconv.Atoi(name); convErr == nil {
count++
}
}
return count, nil
}