diff --git a/e2e/tests/machineprovider/testdata/machineprovider/provider.yaml b/e2e/tests/machineprovider/testdata/machineprovider/provider.yaml index 23dcd28be..aac71c23c 100644 --- a/e2e/tests/machineprovider/testdata/machineprovider/provider.yaml +++ b/e2e/tests/machineprovider/testdata/machineprovider/provider.yaml @@ -5,10 +5,14 @@ description: |- options: LOCATION: description: The location Devsy should use + INACTIVITY_TIMEOUT: + description: The timeout until the machine will be stopped + default: 30s agent: local: true docker: install: false + inactivityTimeout: ${INACTIVITY_TIMEOUT} exec: create: |- mkdir -p ${LOCATION}/${MACHINE_ID} diff --git a/pkg/inject/inject.go b/pkg/inject/inject.go index f5517be11..f1785fdb8 100644 --- a/pkg/inject/inject.go +++ b/pkg/inject/inject.go @@ -305,26 +305,50 @@ func readLine(reader io.Reader) (string, error) { } } +const pipeSecondDirTimeout = 5 * time.Second + func pipe( toStdin io.WriteCloser, fromStdin io.Reader, toStdout io.Writer, fromStdout io.ReadCloser, ) error { - errChan := make(chan error, 2) + stdinErr := make(chan error, 1) + stdoutErr := make(chan error, 1) + go func() { _, err := io.Copy(toStdout, fromStdout) - errChan <- err + stdoutErr <- err }() go func() { _, err := io.Copy(toStdin, fromStdin) - errChan <- err + stdinErr <- err }() - first := <-errChan + // Wait for whichever direction completes first. + var firstErr error + var otherCh <-chan error + select { + case firstErr = <-stdinErr: + otherCh = stdoutErr + case firstErr = <-stdoutErr: + otherCh = stdinErr + } + + // Give the other direction time to finish naturally so we can + // capture any real error and avoid interrupting data in flight. + // If it doesn't finish in time, close pipes to force completion. + var secondErr error + timer := time.NewTimer(pipeSecondDirTimeout) + defer timer.Stop() + select { + case secondErr = <-otherCh: + case <-timer.C: + } _ = toStdin.Close() _ = fromStdout.Close() - <-errChan - - return first + if firstErr != nil { + return firstErr + } + return secondErr } diff --git a/pkg/inject/inject_test.go b/pkg/inject/inject_test.go index 845627d0f..848254f25 100644 --- a/pkg/inject/inject_test.go +++ b/pkg/inject/inject_test.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "errors" + "fmt" "io" "os" "runtime" @@ -153,6 +154,39 @@ func (s *PipeTestSuite) TestPipe_NoGoroutineLeak() { s.LessOrEqual(after, before+5, "goroutine leak detected: before=%d after=%d", before, after) } +func (s *PipeTestSuite) TestPipe_ConcurrentCopyRaceRegression() { + for i := range 100 { + func() { + fromStdinReader, fromStdinWriter := io.Pipe() + toStdoutBuf := &bytes.Buffer{} + + fromStdoutReader, fromStdoutWriter := io.Pipe() + toStdinPipeReader, toStdinPipeWriter := io.Pipe() + + errCh := make(chan error, 1) + go func() { + errCh <- pipe(toStdinPipeWriter, fromStdinReader, toStdoutBuf, fromStdoutReader) + }() + + msg := fmt.Sprintf("iteration-%d", i) + _, err := fromStdinWriter.Write([]byte(msg)) + s.Require().NoError(err) + _ = fromStdinWriter.Close() + + _, err = fromStdoutWriter.Write([]byte(msg)) + s.Require().NoError(err) + _ = fromStdoutWriter.Close() + + received, err := io.ReadAll(toStdinPipeReader) + s.Require().NoError(err) + + s.NoError(<-errCh) + s.Equal(msg, string(received), "stdin data lost at iteration %d", i) + s.Equal(msg, toStdoutBuf.String(), "stdout data lost at iteration %d", i) + }() + } +} + func TestPipeSuite(t *testing.T) { suite.Run(t, new(PipeTestSuite)) }