//go:build test package db import ( "context" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) // ============================================================================= // ContextWithTx / TxFromContext round-trip // ============================================================================= func TestContextWithTx_RoundTrip(t *testing.T) { closePool() resetEnv() err := Connect() require.NoError(t, err) defer closePool() ctx := context.Background() pool := Conn.Pool() tx, err := pool.Begin(ctx) require.NoError(t, err) defer tx.Rollback(ctx) //nolint:errcheck txCtx := ContextWithTx(ctx, tx) extracted := TxFromContext(txCtx) assert.NotNil(t, extracted, "TxFromContext should return a non-nil tx") assert.Equal(t, tx, extracted, "extracted tx should be the same as the one stored") } // ============================================================================= // TxFromContext with no transaction in context // ============================================================================= func TestTxFromContext_NoTx(t *testing.T) { ctx := context.Background() tx := TxFromContext(ctx) assert.Nil(t, tx, "TxFromContext should return nil when no tx in context") } // ============================================================================= // PoolProxy.Exec routes through context transaction // ============================================================================= func TestPoolProxy_Exec_RoutesThroughTx(t *testing.T) { closePool() resetEnv() err := Connect() require.NoError(t, err) defer closePool() ctx := context.Background() pool := Conn.Pool() tx, err := pool.Begin(ctx) require.NoError(t, err) defer tx.Rollback(ctx) //nolint:errcheck txCtx := ContextWithTx(ctx, tx) _, err = Conn.Exec(txCtx, "CREATE TEMP TABLE test_exec_routing (id INT PRIMARY KEY)") require.NoError(t, err) _, err = Conn.Exec(txCtx, "INSERT INTO test_exec_routing VALUES (1)") require.NoError(t, err) var count int err = Conn.QueryRow(txCtx, "SELECT COUNT(*) FROM test_exec_routing").Scan(&count) require.NoError(t, err) assert.Equal(t, 1, count) err = tx.Rollback(ctx) require.NoError(t, err) _, err = Conn.Exec(ctx, "SELECT COUNT(*) FROM test_exec_routing") assert.Error(t, err, "expected error querying temp table after rollback") } // ============================================================================= // PoolProxy.QueryRow routes through context transaction // ============================================================================= func TestPoolProxy_QueryRow_RoutesThroughTx(t *testing.T) { closePool() resetEnv() err := Connect() require.NoError(t, err) defer closePool() ctx := context.Background() pool := Conn.Pool() tx, err := pool.Begin(ctx) require.NoError(t, err) defer tx.Rollback(ctx) //nolint:errcheck txCtx := ContextWithTx(ctx, tx) _, err = Conn.Exec(txCtx, "CREATE TEMP TABLE test_queryrow_routing (id INT PRIMARY KEY, val TEXT)") require.NoError(t, err) _, err = Conn.Exec(txCtx, "INSERT INTO test_queryrow_routing VALUES (42, 'routed')") require.NoError(t, err) var id int var val string err = Conn.QueryRow(txCtx, "SELECT id, val FROM test_queryrow_routing WHERE id = 42").Scan(&id, &val) require.NoError(t, err) assert.Equal(t, 42, id) assert.Equal(t, "routed", val) err = tx.Rollback(ctx) require.NoError(t, err) err = Conn.QueryRow(ctx, "SELECT id, val FROM test_queryrow_routing WHERE id = 42").Scan(&id, &val) assert.Error(t, err, "expected error querying temp table via pool after rollback") } // ============================================================================= // PoolProxy.Begin returns a real transaction // ============================================================================= func TestPoolProxy_Begin_ReturnsTx(t *testing.T) { closePool() resetEnv() err := Connect() require.NoError(t, err) defer closePool() ctx := context.Background() tx, err := Conn.Begin(ctx) require.NoError(t, err) assert.NotNil(t, tx) err = tx.Rollback(ctx) require.NoError(t, err) } // ============================================================================= // PoolProxy.Ping works // ============================================================================= func TestPoolProxy_Ping_Works(t *testing.T) { closePool() resetEnv() err := Connect() require.NoError(t, err) defer closePool() ctx := context.Background() err = Conn.Ping(ctx) assert.NoError(t, err) }