mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
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:
@@ -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)
|
||||
|
||||
@@ -273,4 +273,6 @@ func main() {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,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()
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user