diff --git a/run.go b/run.go index 61997a4..bcf606c 100644 --- a/run.go +++ b/run.go @@ -8,6 +8,13 @@ import ( "time" ) +// RunOptions controls how the non-go-test runner executes registered tests. +type RunOptions struct { + // PackageConcurrency is the maximum number of packages executed at once. + // Values <= 1 preserve the historical serial package execution behavior. + PackageConcurrency int +} + // RunAsTest runs all registered tests under Go's testing framework. // // To run tests on a per-package basis, put a test file in each package containing a single test that calls this function. @@ -89,182 +96,231 @@ func RunAsTest(t *testing.T) { // // TODO: channel for results to support progressive result loading? func Run() TestResult { + return RunWithOptions(RunOptions{PackageConcurrency: 1}) +} + +// RunWithOptions runs all registered tests and returns result information about them. +func RunWithOptions(opts RunOptions) TestResult { start := time.Now() results := TestResult{ Name: "Test Suite", Started: start, } - anyFailures := false - // TODO run packages in parallel like go test does + pkgs := collectPackages() + results.Subtests = make([]TestResult, len(pkgs)) + + packageConcurrency := opts.PackageConcurrency + if packageConcurrency < 1 { + packageConcurrency = 1 + } + + if packageConcurrency == 1 { + for i := range pkgs { + results.Subtests[i] = runPackage(pkgs[i].name, pkgs[i].tests) + } + } else { + runPackagesConcurrently(pkgs, results.Subtests, packageConcurrency) + } + + r := ResultPassed + for _, pkgResult := range results.Subtests { + if pkgResult.Result == ResultFailed { + r = ResultFailed + break + } + } + results.Result = r + dur := time.Since(start).Round(time.Millisecond) + results.Dur = dur + results.DurHuman = dur.String() + return results +} + +func runPackagesConcurrently(pkgs []runPackageInput, results []TestResult, packageConcurrency int) { + sem := make(chan struct{}, packageConcurrency) + wg := sync.WaitGroup{} + for i := range pkgs { + wg.Add(1) + sem <- struct{}{} + go func(i int) { + defer wg.Done() + defer func() { <-sem }() + results[i] = runPackage(pkgs[i].name, pkgs[i].tests) + }(i) + } + wg.Wait() +} + +type runPackageInput struct { + name string + tests *testPkg +} + +func collectPackages() []runPackageInput { + var pkgs []runPackageInput instance.tests.Iterate(func(pkg string, pkgTests *testPkg) bool { - pkgStart := time.Now() - results.Subtests = append(results.Subtests, TestResult{ - Package: pkg, - Name: "Package", - Started: pkgStart, - }) - pkgResults := &results.Subtests[len(results.Subtests)-1] + pkgs = append(pkgs, runPackageInput{name: pkg, tests: pkgTests}) + return true + }) + return pkgs +} - pkgHelperT := &t{} - pkgAnyFailures := false +func runPackage(pkg string, pkgTests *testPkg) TestResult { + pkgStart := time.Now() + pkgResults := TestResult{ + Package: pkg, + Name: "Package", + Started: pkgStart, + } - // we have to hold onto any panics here to be able to run AfterPackage - var beforePkgErr any - if pkgTests.BeforePackage != nil { - func() { - defer func() { - if beforePkgErr = recover(); beforePkgErr != nil { - beforePkgErr = fmt.Sprintf("before package: %v\n\n%s", beforePkgErr, debug.Stack()) - } - }() - pkgTests.BeforePackage(pkgHelperT) - }() + pkgHelperT := &t{} + pkgAnyFailures := false - if beforePkgErr != nil { - pkgAnyFailures = true - pkgResults.Msgs = []Msg{ - { - Msg: fmt.Sprintf("%v", beforePkgErr), - Level: LevelError, - }, + // we have to hold onto any panics here to be able to run AfterPackage + var beforePkgErr any + if pkgTests.BeforePackage != nil { + func() { + defer func() { + if beforePkgErr = recover(); beforePkgErr != nil { + beforePkgErr = fmt.Sprintf("before package: %v\n\n%s", beforePkgErr, debug.Stack()) } + }() + pkgTests.BeforePackage(pkgHelperT) + }() + + if beforePkgErr != nil { + pkgAnyFailures = true + pkgResults.Msgs = []Msg{ + { + Msg: fmt.Sprintf("%v", beforePkgErr), + Level: LevelError, + }, } } + } - // we still have to iterate even if there was a BeforePackage panic to be able to fail all the tests - pkgTests.tests.Iterate(func(name string, test testCase) bool { - // only run the tests if BeforePackage didn't panic - if beforePkgErr == nil { - testHelperT := &t{} - - // we have to hold onto any panics here to be able to run AfterTest - var beforeTestErr any - if pkgTests.BeforeTest != nil { - func() { - defer func() { - if beforeTestErr = recover(); beforeTestErr != nil { - beforeTestErr = fmt.Sprintf("before test: %v\n\n%s", beforeTestErr, debug.Stack()) - } - }() - pkgTests.BeforeTest(testHelperT) + // we still have to iterate even if there was a BeforePackage panic to be able to fail all the tests + pkgTests.tests.Iterate(func(name string, test testCase) bool { + // only run the tests if BeforePackage didn't panic + if beforePkgErr == nil { + testHelperT := &t{} + + // we have to hold onto any panics here to be able to run AfterTest + var beforeTestErr any + if pkgTests.BeforeTest != nil { + func() { + defer func() { + if beforeTestErr = recover(); beforeTestErr != nil { + beforeTestErr = fmt.Sprintf("before test: %v\n\n%s", beforeTestErr, debug.Stack()) + } }() - } + pkgTests.BeforeTest(testHelperT) + }() + } - // only run the tests if any BeforeTest didn't panic - if beforeTestErr == nil { - res := runTest(pkg, test.Name, test.tester) - if res.Result == ResultFailed { - pkgAnyFailures = true - } - pkgResults.Subtests = append(pkgResults.Subtests, res) - } else { + // only run the tests if any BeforeTest didn't panic + if beforeTestErr == nil { + res := runTest(pkg, test.Name, test.tester) + if res.Result == ResultFailed { pkgAnyFailures = true - pkgResults.Subtests = append(pkgResults.Subtests, TestResult{ - Package: pkg, - Name: name, - Started: time.Now(), - Result: ResultFailed, - Dur: 0, - DurHuman: "0s", - Msgs: append(testHelperT.msgs, Msg{ - Msg: fmt.Sprintf("%v", beforeTestErr), - Level: LevelError, - }), - }) - } - - if pkgTests.AfterTest != nil { - var afterTestErr any - func() { - defer func() { - if afterTestErr = recover(); afterTestErr != nil { - afterTestErr = fmt.Sprintf("after test: %v\n\n%s", afterTestErr, debug.Stack()) - } - }() - pkgTests.AfterTest(testHelperT) - }() - - if afterTestErr != nil { - pkgAnyFailures = true - // update test results marking it failed and with this panic message. - r := &pkgResults.Subtests[len(pkgResults.Subtests)-1] - r.Result = ResultFailed - r.Msgs = append(r.Msgs, append(testHelperT.msgs, Msg{ - Msg: fmt.Sprintf("%v", afterTestErr), - Level: LevelError, - })...) - } } + pkgResults.Subtests = append(pkgResults.Subtests, res) } else { pkgAnyFailures = true - // BeforePackage panicked, so simply mark the test as failed with its message pkgResults.Subtests = append(pkgResults.Subtests, TestResult{ Package: pkg, Name: name, - Started: pkgStart, + Started: time.Now(), Result: ResultFailed, Dur: 0, DurHuman: "0s", - Msgs: append(pkgHelperT.msgs, Msg{ - Msg: fmt.Sprintf("%v", beforePkgErr), + Msgs: append(testHelperT.msgs, Msg{ + Msg: fmt.Sprintf("%v", beforeTestErr), Level: LevelError, }), }) } - return true - }) - - var afterPkgErr any - if pkgTests.AfterPackage != nil { - func() { - defer func() { - if afterPkgErr = recover(); afterPkgErr != nil { - afterPkgErr = fmt.Sprintf("after package: %v\n\n%s", afterPkgErr, debug.Stack()) - } + if pkgTests.AfterTest != nil { + var afterTestErr any + func() { + defer func() { + if afterTestErr = recover(); afterTestErr != nil { + afterTestErr = fmt.Sprintf("after test: %v\n\n%s", afterTestErr, debug.Stack()) + } + }() + pkgTests.AfterTest(testHelperT) }() - pkgTests.AfterPackage(pkgHelperT) - }() - } - // update test results if AfterPackage panicked - if afterPkgErr != nil { - pkgAnyFailures = true - m := Msg{ - Msg: fmt.Sprintf("%v", afterPkgErr), - Level: LevelError, - } - for i := range pkgResults.Subtests { - r := &pkgResults.Subtests[i] - r.Result = ResultFailed - r.Msgs = append(r.Msgs, append(pkgHelperT.msgs, m)...) + if afterTestErr != nil { + pkgAnyFailures = true + // update test results marking it failed and with this panic message. + r := &pkgResults.Subtests[len(pkgResults.Subtests)-1] + r.Result = ResultFailed + r.Msgs = append(r.Msgs, append(testHelperT.msgs, Msg{ + Msg: fmt.Sprintf("%v", afterTestErr), + Level: LevelError, + })...) + } } - pkgResults.Msgs = append(pkgResults.Msgs, m) - } - - r := ResultPassed - if pkgAnyFailures { - r = ResultFailed - anyFailures = true + } else { + pkgAnyFailures = true + // BeforePackage panicked, so simply mark the test as failed with its message + pkgResults.Subtests = append(pkgResults.Subtests, TestResult{ + Package: pkg, + Name: name, + Started: pkgStart, + Result: ResultFailed, + Dur: 0, + DurHuman: "0s", + Msgs: append(pkgHelperT.msgs, Msg{ + Msg: fmt.Sprintf("%v", beforePkgErr), + Level: LevelError, + }), + }) } - pkgResults.Result = r - dur := time.Since(pkgStart).Round(time.Millisecond) - pkgResults.Dur = dur - pkgResults.DurHuman = dur.String() return true }) + var afterPkgErr any + if pkgTests.AfterPackage != nil { + func() { + defer func() { + if afterPkgErr = recover(); afterPkgErr != nil { + afterPkgErr = fmt.Sprintf("after package: %v\n\n%s", afterPkgErr, debug.Stack()) + } + }() + pkgTests.AfterPackage(pkgHelperT) + }() + } + + // update test results if AfterPackage panicked + if afterPkgErr != nil { + pkgAnyFailures = true + m := Msg{ + Msg: fmt.Sprintf("%v", afterPkgErr), + Level: LevelError, + } + for i := range pkgResults.Subtests { + r := &pkgResults.Subtests[i] + r.Result = ResultFailed + r.Msgs = append(r.Msgs, append(pkgHelperT.msgs, m)...) + } + pkgResults.Msgs = append(pkgResults.Msgs, m) + } + r := ResultPassed - if anyFailures { + if pkgAnyFailures { r = ResultFailed } - results.Result = r - dur := time.Since(start).Round(time.Millisecond) - results.Dur = dur - results.DurHuman = dur.String() - return results + pkgResults.Result = r + dur := time.Since(pkgStart).Round(time.Millisecond) + pkgResults.Dur = dur + pkgResults.DurHuman = dur.String() + + return pkgResults } func runTest(pkg, baseName string, tester Tester) TestResult { diff --git a/run_test.go b/run_test.go index 6fc289b..75c9f56 100644 --- a/run_test.go +++ b/run_test.go @@ -1,11 +1,15 @@ package testy import ( + "fmt" + "sync/atomic" "testing" "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + "github.com/gametimesf/testy/internal/orderedmap" ) // know that before/after package/test and the test itself have run and when they were run @@ -443,3 +447,79 @@ func TestRun(t *testing.T) { }) } } + +func TestRunWithOptionsRunsPackagesConcurrentlyAndPreservesOrder(t *testing.T) { + instance = testy{ + tests: orderedmap.OrderedMap[string, *testPkg]{}, + } + defer func() { + instance = testy{} + }() + + started := make(chan string, 3) + release := make(chan struct{}) + var running int32 + var maxRunning int32 + + for i := 0; i < 3; i++ { + pkg := fmt.Sprintf("github.com/gametimesf/testy/concurrency/pkg%d", i) + testName := fmt.Sprintf("test%d", i) + instance.tests[pkg] = &testPkg{ + name: pkg, + tests: orderedmap.OrderedMap[string, testCase]{}, + } + instance.tests[pkg].tests[testName] = testCase{ + Package: pkg, + Name: testName, + tester: func(t TestingT) { + nowRunning := atomic.AddInt32(&running, 1) + for { + max := atomic.LoadInt32(&maxRunning) + if nowRunning <= max || atomic.CompareAndSwapInt32(&maxRunning, max, nowRunning) { + break + } + } + started <- pkg + <-release + atomic.AddInt32(&running, -1) + }, + } + } + + done := make(chan TestResult, 1) + go func() { + done <- RunWithOptions(RunOptions{PackageConcurrency: 2}) + }() + + for i := 0; i < 2; i++ { + select { + case <-started: + case <-time.After(time.Second): + t.Fatalf("timed out waiting for package %d to start", i+1) + } + } + + select { + case pkg := <-started: + t.Fatalf("package %s started before a concurrency slot was released", pkg) + case <-time.After(50 * time.Millisecond): + } + + close(release) + + var res TestResult + select { + case res = <-done: + case <-time.After(time.Second): + t.Fatal("timed out waiting for concurrent run to finish") + } + + assert.LessOrEqual(t, atomic.LoadInt32(&maxRunning), int32(2)) + assert.Equal(t, ResultPassed, res.Result) + require.Len(t, res.Subtests, 3) + for i, pkgResult := range res.Subtests { + assert.Equal(t, fmt.Sprintf("github.com/gametimesf/testy/concurrency/pkg%d", i), pkgResult.Package) + require.Len(t, pkgResult.Subtests, 1) + assert.Equal(t, fmt.Sprintf("test%d", i), pkgResult.Subtests[0].Name) + } +}