diff --git a/Makefile b/Makefile index f47f20c73d2..af7447dd47c 100644 --- a/Makefile +++ b/Makefile @@ -291,7 +291,6 @@ endif .PHONY: clean clean: @echo "---> Cleaning up OSS build artifacts." - rm -rf build.assets/build rm -rf $(BUILDDIR) rm -rf $(ER_BPF_BUILDDIR) rm -rf $(RS_BPF_BUILDDIR) diff --git a/build.assets/render-tests/main.go b/build.assets/render-tests/main.go index 03afbfd8e2b..8f1745fa250 100644 --- a/build.assets/render-tests/main.go +++ b/build.assets/render-tests/main.go @@ -273,4 +273,6 @@ func main() { } } } + + os.Exit(1) } diff --git a/integration/integration_test.go b/integration/integration_test.go index ab624cd709c..b2d0d2c4659 100644 --- a/integration/integration_test.go +++ b/integration/integration_test.go @@ -1022,14 +1022,35 @@ func testShutdown(t *testing.T, suite *integrationTestSuite) { } } +// errorVerifier is a function type for functions that check that a given +// error is what was expected. Implementations are expected top return nil +// if the supplied error is as expected, or an descriptive error if is is +// not +type errorVerifier func(error) error + +func errorContains(text string) errorVerifier { + return func(err error) error { + if err == nil || !strings.Contains(err.Error(), text) { + return fmt.Errorf("Expected error to contain %q, got: %v", text, err) + } + return nil + } +} + type disconnectTestCase struct { recordingMode string options types.RoleOptions disconnectTimeout time.Duration concurrentConns int sessCtlTimeout time.Duration - assertExpected func(*testing.T, error) postFunc func(context.Context, *testing.T, *TeleInstance) + + // verifyError checks if `err` reflects the error expected by the test scenario. + // It returns nil if yes, non-nil otherwise. + // It is important for verifyError to not do assertions using `*testing.T` + // itself, as those assertions must run in the main test goroutine, but + // verifyError runs in a different goroutine. + verifyError errorVerifier } // TestDisconnectScenarios tests multiple scenarios with client disconnects @@ -1074,11 +1095,7 @@ func testDisconnectScenarios(t *testing.T, suite *integrationTestSuite) { }, disconnectTimeout: 1 * time.Second, concurrentConns: 2, - assertExpected: func(t *testing.T, err error) { - if err == nil || !strings.Contains(err.Error(), "administratively prohibited") { - require.Failf(t, "Invalid error", "Expected 'administratively prohibited', got: %v", err) - } - }, + verifyError: errorContains("administratively prohibited"), }, { // "verify that concurrent connection limits are applied when recording at proxy", recordingMode: types.RecordAtProxy, @@ -1088,11 +1105,7 @@ func testDisconnectScenarios(t *testing.T, suite *integrationTestSuite) { }, disconnectTimeout: 1 * time.Second, concurrentConns: 2, - assertExpected: func(t *testing.T, err error) { - if err == nil || !strings.Contains(err.Error(), "administratively prohibited") { - require.FailNowf(t, "Invalid error", "Expected 'administratively prohibited', got: %v", err) - } - }, + verifyError: errorContains("administratively prohibited"), }, { // "verify that lost connections to auth server terminate controlled conns", recordingMode: types.RecordAtNode, @@ -1189,6 +1202,8 @@ func runDisconnectTest(t *testing.T, suite *integrationTestSuite, tc disconnectT tc.concurrentConns = 1 } + asyncErrors := make(chan error, 1) + for i := 0; i < tc.concurrentConns; i++ { person := NewTerminal(250) @@ -1208,16 +1223,24 @@ func runDisconnectTest(t *testing.T, suite *integrationTestSuite, tc disconnectT default: } - if tc.assertExpected != nil { - tc.assertExpected(t, err) + if tc.verifyError != nil { + if badErrorErr := tc.verifyError(err); badErrorErr != nil { + asyncErrors <- badErrorErr + } } else if err != nil && !trace.IsEOF(err) && !isSSHError(err) { - require.FailNowf(t, "Missing EOF", "expected EOF, ExitError, or nil, got %v instead", err) + asyncErrors <- fmt.Errorf("expected EOF, ExitError, or nil, got %v instead", err) + return } } go openSession() - go enterInput(ctx, t, person, "echo start \r\n", ".*start.*") + go func() { + err := enterInput(ctx, person, "echo start \r\n", ".*start.*") + if err != nil { + asyncErrors <- err + } + }() } if tc.postFunc != nil { @@ -1229,6 +1252,10 @@ func runDisconnectTest(t *testing.T, suite *integrationTestSuite, tc disconnectT case <-time.After(tc.disconnectTimeout + time.Second): dumpGoroutineProfile() require.FailNowf(t, "timeout", "%s timeout waiting for session to exit: %+v", timeNow(), tc) + + case ae := <-asyncErrors: + require.FailNow(t, "Async error", ae.Error()) + case <-ctx.Done(): // session closed. a test case is successful if the first // session to close encountered the expected error variant. @@ -1248,7 +1275,10 @@ func timeNow() string { return time.Now().Format(time.StampMilli) } -func enterInput(ctx context.Context, t *testing.T, person *Terminal, command, pattern string) { +// enterInput simulates entering user input into a terminal and awaiting a +// response. Returns an error if the given response text doesn't match +// the supplied regexp string. +func enterInput(ctx context.Context, person *Terminal, command, pattern string) error { person.Type(command) abortTime := time.Now().Add(10 * time.Second) var matched bool @@ -1257,17 +1287,17 @@ func enterInput(ctx context.Context, t *testing.T, person *Terminal, command, pa output = replaceNewlines(person.Output(1000)) matched, _ = regexp.MatchString(pattern, output) if matched { - return + return nil } select { case <-time.After(time.Millisecond * 50): case <-ctx.Done(): // cancellation means that we don't care about the input being // confirmed anymore; not equivalent to a timeout. - return + return nil } if time.Now().After(abortTime) { - require.FailNowf(t, "timeout", "failed to capture pattern %q in %q", pattern, output) + return fmt.Errorf("failed to capture pattern %q in %q", pattern, output) } } } diff --git a/integration/kube_integration_test.go b/integration/kube_integration_test.go index 07170a594ca..94889bd91de 100644 --- a/integration/kube_integration_test.go +++ b/integration/kube_integration_test.go @@ -1084,7 +1084,7 @@ func runKubeDisconnectTest(t *testing.T, suite *KubeSuite, tc disconnectTestCase }() // lets type something followed by "enter" and then hang the session - enterInput(sessionCtx, t, term, "echo boring platapus\r\n", ".*boring platapus.*") + require.NoError(t, enterInput(sessionCtx, term, "echo boring platapus\r\n", ".*boring platapus.*")) time.Sleep(tc.disconnectTimeout) select { case <-time.After(tc.disconnectTimeout): diff --git a/lib/utils/pipenetconn.go b/lib/utils/pipenetconn.go index 77a5d042fc4..390ac1e3934 100644 --- a/lib/utils/pipenetconn.go +++ b/lib/utils/pipenetconn.go @@ -19,11 +19,20 @@ package utils import ( "io" "net" + "sync" "time" ) -// PipeNetConn implemetns net.Conn from io.Reader,io.Writer and io.Closer +// PipeNetConn implements net.Conn from a provided io.Reader,io.Writer and +// io.Closer type PipeNetConn struct { + // Locks writing and closing the connection. If both writer & closer refer + // to the same underlying object, simultaneous write and close operations + // introduce a data race (*especially* if that object is a + // `x/crypto/ssh.channel`), so we will use this mutex to serialize write + // and close operations. + mu sync.Mutex + reader io.Reader writer io.Writer closer io.Closer @@ -31,8 +40,9 @@ type PipeNetConn struct { remoteAddr net.Addr } -// NewPipeNetConn returns a net.Conn like object -// using Pipe as an underlying implementation over reader, writer and closer +// NewPipeNetConn constructs a new PipeNetConn, providing a net.Conn +// implementation synthesized from the supplied io.Reader, io.Writer & +// io.Closer. func NewPipeNetConn(reader io.Reader, writer io.Writer, closer io.Closer, @@ -53,10 +63,16 @@ func (nc *PipeNetConn) Read(buf []byte) (n int, e error) { } func (nc *PipeNetConn) Write(buf []byte) (n int, e error) { + nc.mu.Lock() + defer nc.mu.Unlock() + return nc.writer.Write(buf) } func (nc *PipeNetConn) Close() error { + nc.mu.Lock() + defer nc.mu.Unlock() + if nc.closer != nil { return nc.closer.Close() }