package database_test import ( "context" "errors" "testing" "queryorchestration/internal/database/repository" "queryorchestration/internal/serviceconfig" "queryorchestration/internal/test" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) func TestExecuteTransaction(t *testing.T) { t.Run("create client", func(t *testing.T) { t.Parallel() cfg := &serviceconfig.BaseConfig{} test.CreateDB(t, cfg) err := cfg.ExecuteDBTransaction(t.Context(), func(ctx context.Context, q *repository.Queries) error { err := q.CreateClient(ctx, &repository.CreateClientParams{ Name: "example_client", Clientid: "ID", }) require.NoError(t, err) return nil }) require.NoError(t, err) client, err := cfg.GetDBQueries().GetClient(t.Context(), "ID") require.NoError(t, err) assert.Equal(t, "example_client", client.Name) assert.Equal(t, "ID", client.Clientid) }) t.Run("error in transaction", func(t *testing.T) { t.Parallel() cfg := &serviceconfig.BaseConfig{} test.CreateDB(t, cfg) err := cfg.ExecuteDBTransaction(t.Context(), func(ctx context.Context, q *repository.Queries) error { err := q.CreateClient(ctx, &repository.CreateClientParams{ Name: "example_client", Clientid: "ID", }) require.NoError(t, err) return errors.New("error") }) require.Error(t, err) _, err = cfg.GetDBQueries().GetClient(t.Context(), "ID") require.Error(t, err) }) }