diff --git a/build/opt.go b/build/opt.go index fb8b5fe393d2..0121f1da4ad5 100644 --- a/build/opt.go +++ b/build/opt.go @@ -501,6 +501,9 @@ func toSolveOpt(ctx context.Context, np *noderesolver.ResolvedNode, multiDriver } defers = append(defers, func(error) { cancel() + // Close waits for the "importing to docker" progress + // goroutine so it cannot Write after printer.Wait. + _ = w.Close() }) so.Exports[i].Output = func(_ map[string]string) (io.WriteCloser, error) { return w, nil diff --git a/util/dockerutil/client.go b/util/dockerutil/client.go index 68ff9a29c281..513a91d1dc5f 100644 --- a/util/dockerutil/client.go +++ b/util/dockerutil/client.go @@ -110,6 +110,11 @@ func (w *waitingWriter) Write(dt []byte) (int, error) { func (w *waitingWriter) Close() error { err := w.PipeWriter.Close() + // If Write never ran, the loader goroutine was never started and + // done would otherwise block forever. Unblock that path here. + w.once.Do(func() { + close(w.done) + }) <-w.done if err == nil { w.mu.Lock() diff --git a/util/dockerutil/client_test.go b/util/dockerutil/client_test.go new file mode 100644 index 000000000000..c9fd827025a6 --- /dev/null +++ b/util/dockerutil/client_test.go @@ -0,0 +1,63 @@ +package dockerutil + +import ( + "io" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestWaitingWriterCloseWithoutWrite(t *testing.T) { + t.Parallel() + + pr, pw := io.Pipe() + defer pr.Close() + done := make(chan struct{}) + w := &waitingWriter{ + PipeWriter: pw, + f: func() { + t.Error("loader should not start when Close runs before Write") + }, + done: done, + } + + errCh := make(chan error, 1) + go func() { + errCh <- w.Close() + }() + + select { + case err := <-errCh: + require.NoError(t, err) + case <-time.After(2 * time.Second): + t.Fatal("Close hung when Write was never called") + } +} + +func TestWaitingWriterCloseWaitsForLoader(t *testing.T) { + t.Parallel() + + pr, pw := io.Pipe() + done := make(chan struct{}) + started := make(chan struct{}) + w := &waitingWriter{ + PipeWriter: pw, + f: func() { + close(started) + _, _ = io.Copy(io.Discard, pr) + time.Sleep(80 * time.Millisecond) + close(done) + }, + done: done, + } + + n, err := w.Write([]byte("layer")) + require.NoError(t, err) + require.Equal(t, 5, n) + <-started + + start := time.Now() + require.NoError(t, w.Close()) + require.GreaterOrEqual(t, time.Since(start), 60*time.Millisecond) +} diff --git a/util/progress/printer.go b/util/progress/printer.go index dcc2294569c1..66fc13db8d0b 100644 --- a/util/progress/printer.go +++ b/util/progress/printer.go @@ -86,12 +86,25 @@ func (p *Printer) Resume() { } func (p *Printer) Write(s *client.SolveStatus) { - p.status <- s + if !p.sendStatus(s) { + return + } if p.metrics != nil { p.metrics.Write(s) } } +// sendStatus is a no-op if Wait already closed the status channel. +func (p *Printer) sendStatus(s *client.SolveStatus) (sent bool) { + defer func() { + if recover() != nil { + sent = false + } + }() + p.status <- s + return true +} + func (p *Printer) Warnings() []client.VertexWarning { return dedupWarnings(p.warnings) } diff --git a/util/progress/printer_test.go b/util/progress/printer_test.go new file mode 100644 index 000000000000..909fc03c4a22 --- /dev/null +++ b/util/progress/printer_test.go @@ -0,0 +1,52 @@ +package progress + +import ( + "sync" + "testing" + + "github.com/moby/buildkit/client" + "github.com/stretchr/testify/require" +) + +func TestPrinterWriteAfterWait(t *testing.T) { + p := &Printer{ + status: make(chan *client.SolveStatus), + done: make(chan struct{}), + } + go func() { + for range p.status { + } + close(p.done) + }() + + p.Write(&client.SolveStatus{}) + require.NoError(t, p.Wait()) + require.NotPanics(t, func() { + p.Write(&client.SolveStatus{}) + }) +} + +func TestPrinterWriteRaceWait(t *testing.T) { + p := &Printer{ + status: make(chan *client.SolveStatus), + done: make(chan struct{}), + } + go func() { + for range p.status { + } + close(p.done) + }() + + var wg sync.WaitGroup + for i := 0; i < 32; i++ { + wg.Add(1) + go func() { + defer wg.Done() + require.NotPanics(t, func() { + p.Write(&client.SolveStatus{}) + }) + }() + } + require.NoError(t, p.Wait()) + wg.Wait() +}