package store import ( "context" "database/sql" "testing" "github.com/DATA-DOG/go-sqlmock" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) func TestDB_WithTx_Commit(t *testing.T) { db, mock, err := sqlmock.New() require.NoError(t, err) defer db.Close() sdb := New(db) mock.ExpectBegin() mock.ExpectExec("INSERT INTO test").WillReturnResult(sqlmock.NewResult(1, 1)) mock.ExpectCommit() err = sdb.WithTx(context.Background(), func(tx *sql.Tx) error { _, err := tx.Exec("INSERT INTO test VALUES (1)") return err }) assert.NoError(t, err) assert.NoError(t, mock.ExpectationsWereMet()) } func TestDB_WithTx_Rollback(t *testing.T) { db, mock, err := sqlmock.New() require.NoError(t, err) defer db.Close() sdb := New(db) mock.ExpectBegin() mock.ExpectRollback() testErr := assert.AnError err = sdb.WithTx(context.Background(), func(tx *sql.Tx) error { return testErr }) assert.Error(t, err) assert.NoError(t, mock.ExpectationsWereMet()) } // mockEntity for generic store tests type mockEntity struct { ID string Name string } func (m *mockEntity) ScanRow(rows *sql.Rows) error { return rows.Scan(&m.ID, &m.Name) } func scanRows(rows *sql.Rows) (*mockEntity, error) { m := &mockEntity{} err := m.ScanRow(rows) return m, err } func scanRow(row *sql.Row) (*mockEntity, error) { m := &mockEntity{} err := row.Scan(&m.ID, &m.Name) return m, err } func TestStore_List(t *testing.T) { db, mock, err := sqlmock.New() require.NoError(t, err) defer db.Close() sdb := New(db) store := NewStore(sdb, "test_table", []string{"id", "name"}, scanRows, scanRow) mock.ExpectQuery("SELECT id, name FROM test_table"). WillReturnRows(sqlmock.NewRows([]string{"id", "name"}). AddRow("1", "Alice"). AddRow("2", "Bob")) results, err := store.List(context.Background(), "") require.NoError(t, err) assert.Len(t, results, 2) assert.Equal(t, "Alice", results[0].Name) assert.Equal(t, "Bob", results[1].Name) assert.NoError(t, mock.ExpectationsWereMet()) } func TestStore_List_WithWhere(t *testing.T) { db, mock, err := sqlmock.New() require.NoError(t, err) defer db.Close() sdb := New(db) store := NewStore(sdb, "test_table", []string{"id", "name"}, scanRows, scanRow) mock.ExpectQuery("SELECT id, name FROM test_table WHERE status = \\$1 ORDER BY created_at DESC LIMIT 100"). WithArgs("active"). WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow("1", "Alice")) results, err := store.List(context.Background(), "status = $1", "active") require.NoError(t, err) assert.Len(t, results, 1) assert.NoError(t, mock.ExpectationsWereMet()) } func TestStore_Get(t *testing.T) { db, mock, err := sqlmock.New() require.NoError(t, err) defer db.Close() sdb := New(db) store := NewStore(sdb, "test_table", []string{"id", "name"}, scanRows, scanRow) mock.ExpectQuery("SELECT id, name FROM test_table WHERE id = \\$1"). WithArgs("1"). WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow("1", "Alice")) result, err := store.Get(context.Background(), "1") require.NoError(t, err) assert.Equal(t, "Alice", result.Name) assert.NoError(t, mock.ExpectationsWereMet()) } func TestStore_Get_NotFound(t *testing.T) { db, mock, err := sqlmock.New() require.NoError(t, err) defer db.Close() sdb := New(db) store := NewStore(sdb, "test_table", []string{"id", "name"}, func(rows *sql.Rows) (*mockEntity, error) { return nil, nil }, scanRow, ) mock.ExpectQuery("SELECT id, name FROM test_table WHERE id = \\$1"). WithArgs("999"). WillReturnError(sql.ErrNoRows) _, err = store.Get(context.Background(), "999") assert.Error(t, err) assert.Contains(t, err.Error(), "not found") assert.NoError(t, mock.ExpectationsWereMet()) } func TestStore_Delete(t *testing.T) { db, mock, err := sqlmock.New() require.NoError(t, err) defer db.Close() sdb := New(db) store := NewStore(sdb, "test_table", []string{"id", "name"}, func(rows *sql.Rows) (*mockEntity, error) { return nil, nil }, func(row *sql.Row) (*mockEntity, error) { return nil, nil }, ) mock.ExpectExec("DELETE FROM test_table WHERE id = \\$1"). WithArgs("1"). WillReturnResult(sqlmock.NewResult(0, 1)) err = store.Delete(context.Background(), "1") assert.NoError(t, err) assert.NoError(t, mock.ExpectationsWereMet()) }