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)