Fix race conditions and improve lifecycle safety #150
@@ -1922,10 +1922,6 @@ func TestInteraction_Start_ProtocolSelection(t *testing.T) {
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
i := mockInteraction.(*interaction)
|
||||
if i.program != nil {
|
||||
assert.NotNil(t, i.program, "program should be initialized")
|
||||
}
|
||||
|
||||
i.Stop()
|
||||
|
||||
mockConfig.AssertExpectations(t)
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"sync"
|
||||
"testing"
|
||||
"tunnel_pls/internal/types"
|
||||
|
||||
@@ -308,3 +309,49 @@ func TestLifecycle_Close(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLifecycle_ConcurrentClose(t *testing.T) {
|
||||
mockSSHConn := &MockSSHConn{}
|
||||
mockSSHConn.On("Close").Return(nil)
|
||||
|
||||
mockForwarder := &MockForwarder{}
|
||||
mockForwarder.On("TunnelType").Return(types.TunnelTypeHTTP)
|
||||
|
||||
mockSlug := &MockSlug{}
|
||||
mockSlug.On("String").Return("test-slug")
|
||||
|
||||
mockPort := &MockPort{}
|
||||
|
||||
mockSessionRegistry := &MockSessionRegistry{}
|
||||
mockSessionRegistry.On("Remove", mock.Anything).Return()
|
||||
|
||||
mockSSHChannel := &MockSSHChannel{}
|
||||
mockSSHChannel.On("Close").Return(nil)
|
||||
|
||||
mockLifecycle := New(mockSSHConn, mockForwarder, mockSlug, mockPort, mockSessionRegistry, "mas-fuad")
|
||||
mockLifecycle.SetStatus(types.SessionStatusRUNNING)
|
||||
err := mockLifecycle.SetChannel(mockSSHChannel)
|
||||
assert.NoError(t, err)
|
||||
|
||||
const numGoroutines = 10
|
||||
var wg sync.WaitGroup
|
||||
errChan := make(chan error, numGoroutines)
|
||||
|
||||
for i := 0; i < numGoroutines; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
err := mockLifecycle.Close()
|
||||
errChan <- err
|
||||
}()
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
close(errChan)
|
||||
|
||||
for err := range errChan {
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
assert.False(t, mockLifecycle.IsActive())
|
||||
}
|
||||
|
||||
@@ -55,8 +55,7 @@ func TestHTTPServer_Serve(t *testing.T) {
|
||||
|
||||
go func() {
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
err = listener.Close()
|
||||
assert.NoError(t, err)
|
||||
_ = listener.Close()
|
||||
}()
|
||||
|
||||
err = srv.Serve(listener)
|
||||
|
||||
@@ -63,8 +63,7 @@ func TestHTTPSServer_Serve(t *testing.T) {
|
||||
|
||||
go func() {
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
err = listener.Close()
|
||||
assert.NoError(t, err)
|
||||
_ = listener.Close()
|
||||
}()
|
||||
|
||||
err = srv.Serve(listener)
|
||||
|
||||
@@ -45,8 +45,7 @@ func TestTCPServer_Serve(t *testing.T) {
|
||||
|
||||
go func() {
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
err = listener.Close()
|
||||
assert.NoError(t, err)
|
||||
_ = listener.Close()
|
||||
}()
|
||||
|
||||
err = srv.Serve(listener)
|
||||
|
||||
Reference in New Issue
Block a user