fix(session): release port when error happen

This commit is contained in:
2026-07-18 17:43:17 +07:00
parent 4fcc41eb8a
commit 1ec4518a9e
2 changed files with 19 additions and 2 deletions
+9
View File
@@ -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) tcpServer := transport.NewTCPServer(portToBind, s.forwarder)
listener, err := tcpServer.Listen() listener, err := tcpServer.Listen()
if err != nil { if err != nil {
releasePort()
return s.denyForwardingRequest(req, nil, listener, fmt.Sprintf("Port %d is already in use or restricted", portToBind)) 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} key := types.SessionKey{Id: fmt.Sprintf("%d", portToBind), Type: types.TunnelTypeTCP}
if !s.registry.Register(key, s) { 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)) 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) err = s.finalizeForwarding(req, portToBind, listener, types.TunnelTypeTCP, key.Id)
if err != nil { if err != nil {
releasePort()
return s.denyForwardingRequest(req, &key, listener, fmt.Sprintf("Failed to finalize forwarding: %s", err)) return s.denyForwardingRequest(req, &key, listener, fmt.Sprintf("Failed to finalize forwarding: %s", err))
} }
+10 -2
View File
@@ -705,6 +705,7 @@ func TestForwardingFailures(t *testing.T) {
s, mRegistry, mPort, _, _, sReqs, cConn, cleanup := setup(t) s, mRegistry, mPort, _, _, sReqs, cConn, cleanup := setup(t)
defer cleanup() defer cleanup()
mPort.On("Claim", mock.Anything).Return(true) mPort.On("Claim", mock.Anything).Return(true)
mPort.On("SetStatus", uint16(1234), false).Return(nil)
mRegistry.On("Register", mock.Anything, mock.Anything).Return(false) mRegistry.On("Register", mock.Anything, mock.Anything).Return(false)
payload := make([]byte, 4+9+4) payload := make([]byte, 4+9+4)
@@ -723,7 +724,7 @@ func TestForwardingFailures(t *testing.T) {
}) })
t.Run("Finalize Forwarding Failure", func(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() defer cleanup()
mRandom.On("String", 20).Return("test-slug", nil) mRandom.On("String", 20).Return("test-slug", nil)
mRegistry.On("Register", mock.Anything, mock.Anything).Return(true) mRegistry.On("Register", mock.Anything, mock.Anything).Return(true)
@@ -742,7 +743,7 @@ func TestForwardingFailures(t *testing.T) {
err := cConn.Close() err := cConn.Close()
assert.NoError(t, err) assert.NoError(t, err)
time.Sleep(50 * time.Millisecond) _ = sConn.Wait()
err = s.HandleTCPIPForward(req) err = s.HandleTCPIPForward(req)
assert.Error(t, err) assert.Error(t, err)
@@ -764,6 +765,7 @@ func TestForwardingFailures(t *testing.T) {
}(l) }(l)
_, portStr, _ := net.SplitHostPort(l.Addr().String()) _, portStr, _ := net.SplitHostPort(l.Addr().String())
port, _ := strconv.Atoi(portStr) port, _ := strconv.Atoi(portStr)
mPort.On("SetStatus", uint16(port), false).Return(nil)
payload := make([]byte, 4+9+4) payload := make([]byte, 4+9+4)
binary.BigEndian.PutUint32(payload[0:4], 9) binary.BigEndian.PutUint32(payload[0:4], 9)
@@ -1247,6 +1249,7 @@ func TestHandleTCPForward_Failures(t *testing.T) {
assert.NoError(t, err) assert.NoError(t, err)
}(l) }(l)
port := uint16(l.Addr().(*net.TCPAddr).Port) port := uint16(l.Addr().(*net.TCPAddr).Port)
mPort.On("SetStatus", port, false).Return(nil)
err = s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", port, false) err = s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", port, false)
if err == nil { if err == nil {
@@ -1254,12 +1257,14 @@ func TestHandleTCPForward_Failures(t *testing.T) {
} else if !strings.Contains(err.Error(), "already in use") { } else if !strings.Contains(err.Error(), "already in use") {
t.Errorf("expected error to contain %q, got %q", "already in use", err.Error()) 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) { t.Run("Registry Register fail", func(t *testing.T) {
s, mRegistry, mPort, _, sReqs, cConn, cleanup := setup(t) s, mRegistry, mPort, _, sReqs, cConn, cleanup := setup(t)
defer cleanup() defer cleanup()
mPort.On("Claim", mock.Anything).Return(true) 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) mRegistry.On("Register", mock.Anything, mock.Anything).Return(false)
err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 0, false) err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 0, false)
if err == nil { if err == nil {
@@ -1267,12 +1272,14 @@ func TestHandleTCPForward_Failures(t *testing.T) {
} else if !strings.Contains(err.Error(), "Failed to register") { } else if !strings.Contains(err.Error(), "Failed to register") {
t.Errorf("expected error to contain %q, got %q", "Failed to register", err.Error()) 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) { t.Run("Finalize fail (Reply fail)", func(t *testing.T) {
s, mRegistry, mPort, sConn, sReqs, cConn, cleanup := setup(t) s, mRegistry, mPort, sConn, sReqs, cConn, cleanup := setup(t)
defer cleanup() defer cleanup()
mPort.On("Claim", mock.Anything).Return(true) 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) mRegistry.On("Register", mock.Anything, mock.Anything).Return(true)
req := getReq(t, cConn, sReqs) req := getReq(t, cConn, sReqs)
err := cConn.Close() err := cConn.Close()
@@ -1285,6 +1292,7 @@ func TestHandleTCPForward_Failures(t *testing.T) {
} else if !strings.Contains(err.Error(), "Failed to finalize forwarding") { } 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()) t.Errorf("expected error to contain %q, got %q", "Failed to finalize forwarding", err.Error())
} }
mPort.AssertExpectations(t)
}) })
} }