diff --git a/internal/cli/check.go b/internal/cli/check.go index fb19e63..520c043 100644 --- a/internal/cli/check.go +++ b/internal/cli/check.go @@ -11,6 +11,7 @@ import ( "path/filepath" "strconv" "strings" + "sync" "time" "github.com/dustin/go-humanize" @@ -163,7 +164,9 @@ func verifyRequiredSigner( } // reportCheckProgress renders progress updates until the channel closes. -func reportCheckProgress(progress <-chan mfer.CheckStatus) { +func reportCheckProgress(progress <-chan mfer.CheckStatus, wg *sync.WaitGroup) { + defer wg.Done() + for status := range progress { if status.ETA > 0 { log.Progressf("Checking: %d/%d files, %s/s, ETA %s, %d failures", @@ -235,11 +238,17 @@ func runCheck(ctx *cli.Context, chk *mfer.Checker, showProgress bool) (int64, er results := make(chan mfer.Result, 1) // Set up progress channel - var progress chan mfer.CheckStatus + var ( + progress chan mfer.CheckStatus + progressWg sync.WaitGroup + ) + if showProgress { progress = make(chan mfer.CheckStatus, 1) - go reportCheckProgress(progress) + progressWg.Add(1) + + go reportCheckProgress(progress, &progressWg) } // Process results in a goroutine @@ -251,6 +260,9 @@ func runCheck(ctx *cli.Context, chk *mfer.Checker, showProgress bool) (int64, er // Run check err := chk.Check(ctx.Context, results, progress) + + progressWg.Wait() + if err != nil { return 0, fmt.Errorf("check failed: %w", err) } diff --git a/internal/cli/entry_test.go b/internal/cli/entry_test.go index 359c66f..8fd7472 100644 --- a/internal/cli/entry_test.go +++ b/internal/cli/entry_test.go @@ -13,6 +13,7 @@ import ( "strings" "sync" "testing" + "time" "github.com/spf13/afero" "github.com/stretchr/testify/assert" @@ -393,6 +394,67 @@ func TestGenerateAndCheckCommand(t *testing.T) { assert.Equal(t, 0, exitCode, "check failed: %s", testStderr(t, opts)) } +// sharedWriter appends to a buffer shared with other sharedWriters, so +// output written to stdout and stderr is kept in the order it was written. +// Each write first waits for delay. +type sharedWriter struct { + mu *sync.Mutex + buf *bytes.Buffer + delay time.Duration +} + +func (w sharedWriter) Write(p []byte) (int, error) { + time.Sleep(w.delay) + + w.mu.Lock() + defer w.mu.Unlock() + + return w.buf.Write(p) +} + +// TestCheckClearsProgressBeforeSummary asserts that check --progress writes +// its last progress line and clears it before it logs the summary. Progress +// writes are slowed down, so a progress goroutine that check did not wait for +// would write after the summary, or after the run has returned. +func TestCheckClearsProgressBeforeSummary(t *testing.T) { + t.Parallel() + + fs := afero.NewMemMapFs() + require.NoError(t, fs.MkdirAll(testDir, 0o755)) + writeTestFile(t, fs, testFile1, "hello world") + + opts := testOpts([]string{testApp, cmdGenerate, "-q", "-o", testMF, testDir}, fs) + require.Equal(t, 0, runCLI(opts), "generate failed: %s", testStderr(t, opts)) + + var ( + mu sync.Mutex + output bytes.Buffer + ) + + opts = testOpts([]string{ + testApp, cmdCheck, "--progress", testFlagBase, testDir, testMF, + }, fs) + opts.Stdout = sharedWriter{mu: &mu, buf: &output, delay: 100 * time.Millisecond} + opts.Stderr = sharedWriter{mu: &mu, buf: &output} + require.Equal(t, 0, runCLI(opts)) + + mu.Lock() + got := output.String() + mu.Unlock() + + lastProgress := strings.Index(got, "Checking: 1/1 files") + progressDone := strings.Index(got, "\r\033[K") + summary := strings.Index(got, "checked 1 files") + + require.NotEqual(t, -1, lastProgress, "no last progress line in %q", got) + require.NotEqual(t, -1, progressDone, "progress line never cleared in %q", got) + require.NotEqual(t, -1, summary, "no summary in %q", got) + assert.Less(t, lastProgress, progressDone, + "progress cleared before its last line in %q", got) + assert.Less(t, progressDone, summary, + "summary logged before progress was cleared in %q", got) +} + func TestCheckCommandWithMissingFile(t *testing.T) { t.Parallel()