Fix race conditions and improve lifecycle safety #150
@@ -80,6 +80,12 @@ func (l *lifecycle) User() string {
|
|||||||
func (l *lifecycle) SetChannel(channel ssh.Channel) error {
|
func (l *lifecycle) SetChannel(channel ssh.Channel) error {
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
defer l.mu.Unlock()
|
defer l.mu.Unlock()
|
||||||
|
if l.status == types.SessionStatusCLOSED {
|
||||||
|
return fmt.Errorf("lifecycle is closed")
|
||||||
|
}
|
||||||
|
if channel == nil {
|
||||||
|
return fmt.Errorf("channel cannot be nil")
|
||||||
|
}
|
||||||
if l.channel != nil {
|
if l.channel != nil {
|
||||||
return fmt.Errorf("channel already set")
|
return fmt.Errorf("channel already set")
|
||||||
}
|
}
|
||||||
@@ -100,6 +106,9 @@ func (l *lifecycle) Connection() ssh.Conn {
|
|||||||
func (l *lifecycle) SetStatus(status types.SessionStatus) {
|
func (l *lifecycle) SetStatus(status types.SessionStatus) {
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
defer l.mu.Unlock()
|
defer l.mu.Unlock()
|
||||||
|
if l.status == types.SessionStatusCLOSED {
|
||||||
|
return
|
||||||
|
}
|
||||||
l.status = status
|
l.status = status
|
||||||
if status == types.SessionStatusRUNNING && l.startedAt.IsZero() {
|
if status == types.SessionStatusRUNNING && l.startedAt.IsZero() {
|
||||||
l.startedAt = time.Now()
|
l.startedAt = time.Now()
|
||||||
@@ -114,41 +123,41 @@ func (l *lifecycle) IsActive() bool {
|
|||||||
|
|
||||||
func (l *lifecycle) Close() error {
|
func (l *lifecycle) Close() error {
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
if l.status == types.SessionStatusCLOSED {
|
if l.status == types.SessionStatusCLOSED {
|
||||||
return l.closeErr
|
closeErr := l.closeErr
|
||||||
|
l.mu.Unlock()
|
||||||
|
return closeErr
|
||||||
}
|
}
|
||||||
l.status = types.SessionStatusCLOSED
|
l.status = types.SessionStatusCLOSED
|
||||||
|
|
||||||
|
channel := l.channel
|
||||||
|
conn := l.conn
|
||||||
|
l.mu.Unlock()
|
||||||
|
|
||||||
var errs []error
|
var errs []error
|
||||||
errs = append(errs, l.closeChannel())
|
if channel != nil {
|
||||||
errs = append(errs, l.closeConnection())
|
if err := channel.Close(); err != nil && !isClosedError(err) {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if conn != nil {
|
||||||
|
if err := conn.Close(); err != nil && !isClosedError(err) {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
l.cleanupRegistry()
|
l.cleanupRegistry()
|
||||||
errs = append(errs, l.cleanupForwarder())
|
if err := l.cleanupForwarder(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
|
||||||
l.closeErr = errors.Join(errs...)
|
closeErr := errors.Join(errs...)
|
||||||
return l.closeErr
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *lifecycle) closeChannel() error {
|
l.mu.Lock()
|
||||||
if l.channel == nil {
|
l.closeErr = closeErr
|
||||||
return nil
|
l.mu.Unlock()
|
||||||
}
|
|
||||||
if err := l.channel.Close(); err != nil && !isClosedError(err) {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *lifecycle) closeConnection() error {
|
return closeErr
|
||||||
if l.conn == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
if err := l.conn.Close(); err != nil && !isClosedError(err) {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (l *lifecycle) cleanupRegistry() {
|
func (l *lifecycle) cleanupRegistry() {
|
||||||
|
|||||||
@@ -355,3 +355,70 @@ func TestLifecycle_ConcurrentClose(t *testing.T) {
|
|||||||
|
|
||||||
assert.False(t, mockLifecycle.IsActive())
|
assert.False(t, mockLifecycle.IsActive())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestLifecycle_SetChannel_AfterClose(t *testing.T) {
|
||||||
|
mockSSHConn := new(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)
|
||||||
|
|
||||||
|
err = mockLifecycle.Close()
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
anotherChannel := &MockSSHChannel{}
|
||||||
|
err = mockLifecycle.SetChannel(anotherChannel)
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "lifecycle is closed")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLifecycle_SetChannel_Nil(t *testing.T) {
|
||||||
|
mockSSHConn := new(MockSSHConn)
|
||||||
|
mockForwarder := &MockForwarder{}
|
||||||
|
mockSlug := &MockSlug{}
|
||||||
|
mockPort := &MockPort{}
|
||||||
|
mockSessionRegistry := &MockSessionRegistry{}
|
||||||
|
|
||||||
|
mockLifecycle := New(mockSSHConn, mockForwarder, mockSlug, mockPort, mockSessionRegistry, "mas-fuad")
|
||||||
|
|
||||||
|
err := mockLifecycle.SetChannel(nil)
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "channel cannot be nil")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLifecycle_SetStatus_AfterClose(t *testing.T) {
|
||||||
|
mockSSHConn := new(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)
|
||||||
|
|
||||||
|
err = mockLifecycle.Close()
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.False(t, mockLifecycle.IsActive())
|
||||||
|
|
||||||
|
mockLifecycle.SetStatus(types.SessionStatusRUNNING)
|
||||||
|
assert.False(t, mockLifecycle.IsActive(), "SetStatus should be ignored after Close")
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user