diff --git a/internal/session/lifecycle/lifecycle.go b/internal/session/lifecycle/lifecycle.go index c04473d..0191e40 100644 --- a/internal/session/lifecycle/lifecycle.go +++ b/internal/session/lifecycle/lifecycle.go @@ -80,6 +80,12 @@ func (l *lifecycle) User() string { func (l *lifecycle) SetChannel(channel ssh.Channel) error { l.mu.Lock() 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 { return fmt.Errorf("channel already set") } @@ -100,6 +106,9 @@ func (l *lifecycle) Connection() ssh.Conn { func (l *lifecycle) SetStatus(status types.SessionStatus) { l.mu.Lock() defer l.mu.Unlock() + if l.status == types.SessionStatusCLOSED { + return + } l.status = status if status == types.SessionStatusRUNNING && l.startedAt.IsZero() { l.startedAt = time.Now() @@ -114,41 +123,41 @@ func (l *lifecycle) IsActive() bool { func (l *lifecycle) Close() error { l.mu.Lock() - defer l.mu.Unlock() - if l.status == types.SessionStatusCLOSED { - return l.closeErr + closeErr := l.closeErr + l.mu.Unlock() + return closeErr } l.status = types.SessionStatusCLOSED + channel := l.channel + conn := l.conn + l.mu.Unlock() + var errs []error - errs = append(errs, l.closeChannel()) - errs = append(errs, l.closeConnection()) + if channel != nil { + 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() - errs = append(errs, l.cleanupForwarder()) + if err := l.cleanupForwarder(); err != nil { + errs = append(errs, err) + } - l.closeErr = errors.Join(errs...) - return l.closeErr -} + closeErr := errors.Join(errs...) -func (l *lifecycle) closeChannel() error { - if l.channel == nil { - return nil - } - if err := l.channel.Close(); err != nil && !isClosedError(err) { - return err - } - return nil -} + l.mu.Lock() + l.closeErr = closeErr + l.mu.Unlock() -func (l *lifecycle) closeConnection() error { - if l.conn == nil { - return nil - } - if err := l.conn.Close(); err != nil && !isClosedError(err) { - return err - } - return nil + return closeErr } func (l *lifecycle) cleanupRegistry() { diff --git a/internal/session/lifecycle/lifecycle_test.go b/internal/session/lifecycle/lifecycle_test.go index e16fd48..b5b2ff6 100644 --- a/internal/session/lifecycle/lifecycle_test.go +++ b/internal/session/lifecycle/lifecycle_test.go @@ -355,3 +355,70 @@ func TestLifecycle_ConcurrentClose(t *testing.T) { 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") +}