diff --git a/internal/session/session.go b/internal/session/session.go index 785be25..a2c3d81 100644 --- a/internal/session/session.go +++ b/internal/session/session.go @@ -363,19 +363,28 @@ func (s *session) HandleTCPForward(req *ssh.Request, addr string, portToBind uin } } + releasePort := func() { + if err := s.lifecycle.PortRegistry().SetStatus(portToBind, false); err != nil { + log.Printf("failed to release port %d: %v", portToBind, err) + } + } + tcpServer := transport.NewTCPServer(portToBind, s.forwarder) listener, err := tcpServer.Listen() if err != nil { + releasePort() return s.denyForwardingRequest(req, nil, listener, fmt.Sprintf("Port %d is already in use or restricted", portToBind)) } key := types.SessionKey{Id: fmt.Sprintf("%d", portToBind), Type: types.TunnelTypeTCP} if !s.registry.Register(key, s) { + releasePort() return s.denyForwardingRequest(req, nil, listener, fmt.Sprintf("Failed to register TunnelTypeTCP client with id: %s", key.Id)) } err = s.finalizeForwarding(req, portToBind, listener, types.TunnelTypeTCP, key.Id) if err != nil { + releasePort() return s.denyForwardingRequest(req, &key, listener, fmt.Sprintf("Failed to finalize forwarding: %s", err)) } diff --git a/internal/session/session_test.go b/internal/session/session_test.go index 07d6024..dc6251d 100644 --- a/internal/session/session_test.go +++ b/internal/session/session_test.go @@ -705,6 +705,7 @@ func TestForwardingFailures(t *testing.T) { s, mRegistry, mPort, _, _, sReqs, cConn, cleanup := setup(t) defer cleanup() mPort.On("Claim", mock.Anything).Return(true) + mPort.On("SetStatus", uint16(1234), false).Return(nil) mRegistry.On("Register", mock.Anything, mock.Anything).Return(false) payload := make([]byte, 4+9+4) @@ -723,7 +724,7 @@ func TestForwardingFailures(t *testing.T) { }) t.Run("Finalize Forwarding Failure", func(t *testing.T) { - s, mRegistry, _, mRandom, _, sReqs, cConn, cleanup := setup(t) + s, mRegistry, _, mRandom, sConn, sReqs, cConn, cleanup := setup(t) defer cleanup() mRandom.On("String", 20).Return("test-slug", nil) mRegistry.On("Register", mock.Anything, mock.Anything).Return(true) @@ -742,7 +743,7 @@ func TestForwardingFailures(t *testing.T) { err := cConn.Close() assert.NoError(t, err) - time.Sleep(50 * time.Millisecond) + _ = sConn.Wait() err = s.HandleTCPIPForward(req) assert.Error(t, err) @@ -764,6 +765,7 @@ func TestForwardingFailures(t *testing.T) { }(l) _, portStr, _ := net.SplitHostPort(l.Addr().String()) port, _ := strconv.Atoi(portStr) + mPort.On("SetStatus", uint16(port), false).Return(nil) payload := make([]byte, 4+9+4) binary.BigEndian.PutUint32(payload[0:4], 9) @@ -1247,6 +1249,7 @@ func TestHandleTCPForward_Failures(t *testing.T) { assert.NoError(t, err) }(l) port := uint16(l.Addr().(*net.TCPAddr).Port) + mPort.On("SetStatus", port, false).Return(nil) err = s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", port, false) if err == nil { @@ -1254,12 +1257,14 @@ func TestHandleTCPForward_Failures(t *testing.T) { } else if !strings.Contains(err.Error(), "already in use") { t.Errorf("expected error to contain %q, got %q", "already in use", err.Error()) } + mPort.AssertExpectations(t) }) t.Run("Registry Register fail", func(t *testing.T) { s, mRegistry, mPort, _, sReqs, cConn, cleanup := setup(t) defer cleanup() mPort.On("Claim", mock.Anything).Return(true) + mPort.On("SetStatus", mock.AnythingOfType("uint16"), false).Return(nil) mRegistry.On("Register", mock.Anything, mock.Anything).Return(false) err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 0, false) if err == nil { @@ -1267,12 +1272,14 @@ func TestHandleTCPForward_Failures(t *testing.T) { } else if !strings.Contains(err.Error(), "Failed to register") { t.Errorf("expected error to contain %q, got %q", "Failed to register", err.Error()) } + mPort.AssertExpectations(t) }) t.Run("Finalize fail (Reply fail)", func(t *testing.T) { s, mRegistry, mPort, sConn, sReqs, cConn, cleanup := setup(t) defer cleanup() mPort.On("Claim", mock.Anything).Return(true) + mPort.On("SetStatus", mock.AnythingOfType("uint16"), false).Return(nil) mRegistry.On("Register", mock.Anything, mock.Anything).Return(true) req := getReq(t, cConn, sReqs) err := cConn.Close() @@ -1285,6 +1292,7 @@ func TestHandleTCPForward_Failures(t *testing.T) { } else if !strings.Contains(err.Error(), "Failed to finalize forwarding") { t.Errorf("expected error to contain %q, got %q", "Failed to finalize forwarding", err.Error()) } + mPort.AssertExpectations(t) }) }