Fix race condition in PipeNetCon (#8643)

The race condition detector is being tripped by a concurrent `Write` and
`Close` in the `PipeNetCon` in several integration tests. This is a naive
fix to serialize the write and close operations to resolve the race
condition.

The affected tests were also not handling asynchronous error reporting
correctly (i.e. it's not legal to call `require.XYZ()` from a goroutine
other than the one executing the test function.). This patch introduces
some plumbing to marshal asynchronous errors back into the main test
routine before failing the test.
This commit is contained in:
Trent Clarke
2021-10-28 09:38:51 +11:00
committed by GitHub
parent 8101a3d2aa
commit 5463c799ea
5 changed files with 71 additions and 24 deletions
-1
View File
@@ -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)
+2
View File
@@ -273,4 +273,6 @@ func main() {
}
}
}
os.Exit(1)
}
+49 -19
View File
@@ -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)
}
}
}
+1 -1
View File
@@ -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):
+19 -3
View File
@@ -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()
}