From 61a4791bcbdb5e055a5cbe0299d9d866d561e10e Mon Sep 17 00:00:00 2001 From: Bagas Date: Sun, 19 Jul 2026 10:59:22 +0700 Subject: [PATCH] fix(port, session): harden port registry and fix concurent allocation (#155) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Critical fixes to the port registry that prevented correct tunnel port allocation under concurrency and allowed out-of-range ports to be claimed. Bug Fixes: - Fixed AddRange infinite loop when endPort == 65535 (uint16 wraparound) -Fixed Claim bypassing the allowed range — now rejects out-of-range ports - Fixed Unassigned handing out the same port to concurrent callers — now reserves atomically - Fixed SetStatus silently creating entries for unknown ports — now returns an error Reviewed-on: https://git.fossy.my.id/bagas/tunnel-please/pulls/155 Co-authored-by: Bagas Co-committed-by: Bagas --- internal/port/port.go | 25 ++++++++------- internal/port/port_test.go | 55 ++++++++++++++++++++++++++++++-- internal/session/session.go | 37 +++++++++++++-------- internal/session/session_test.go | 36 ++++++++++++++------- 4 files changed, 115 insertions(+), 38 deletions(-) diff --git a/internal/port/port.go b/internal/port/port.go index 6c60fbb..15f08a8 100644 --- a/internal/port/port.go +++ b/internal/port/port.go @@ -33,10 +33,15 @@ func (pm *port) AddRange(startPort, endPort uint16) error { if startPort > endPort { return fmt.Errorf("start port cannot be greater than end port") } - for index := startPort; index <= endPort; index++ { - if _, exists := pm.ports[index]; !exists { - pm.ports[index] = false - pm.sortedPorts = append(pm.sortedPorts, index) + for index := startPort; ; index++ { + if index != 0 { + if _, exists := pm.ports[index]; !exists { + pm.ports[index] = false + pm.sortedPorts = append(pm.sortedPorts, index) + } + } + if index == endPort { + break } } sort.Slice(pm.sortedPorts, func(i, j int) bool { @@ -51,6 +56,7 @@ func (pm *port) Unassigned() (uint16, bool) { for _, index := range pm.sortedPorts { if !pm.ports[index] { + pm.ports[index] = true return index, true } } @@ -61,6 +67,9 @@ func (pm *port) SetStatus(port uint16, assigned bool) error { pm.mu.Lock() defer pm.mu.Unlock() + if _, exists := pm.ports[port]; !exists { + return fmt.Errorf("port %d is not in the allowed range", port) + } pm.ports[port] = assigned return nil } @@ -70,16 +79,10 @@ func (pm *port) Claim(port uint16) (claimed bool) { defer pm.mu.Unlock() status, exists := pm.ports[port] - - if exists && status { + if !exists || status { return false } - if !exists { - pm.ports[port] = true - return true - } - pm.ports[port] = true return true } diff --git a/internal/port/port_test.go b/internal/port/port_test.go index fcc64d3..9f9a3d9 100644 --- a/internal/port/port_test.go +++ b/internal/port/port_test.go @@ -16,6 +16,8 @@ func TestAddRange(t *testing.T) { {"normal range", 1000, 1002, false}, {"invalid range", 2000, 1999, true}, {"single port range", 3000, 3000, false}, + {"range ending at max uint16", 65533, 65535, false}, + {"range including port zero", 0, 2, false}, } for _, tt := range tests { @@ -31,6 +33,22 @@ func TestAddRange(t *testing.T) { } } +func TestAddRangeBoundaries(t *testing.T) { + pm := New() + err := pm.AddRange(0, 3) + assert.NoError(t, err) + _, hasZero := pm.(*port).ports[0] + assert.False(t, hasZero, "port 0 must be skipped") + assert.Len(t, pm.(*port).ports, 3) + + pm2 := New() + err = pm2.AddRange(65533, 65535) + assert.NoError(t, err) + assert.Len(t, pm2.(*port).ports, 3) + _, hasMax := pm2.(*port).ports[65535] + assert.True(t, hasMax) +} + func TestUnassigned(t *testing.T) { pm := New() _ = pm.AddRange(1000, 1002) @@ -58,6 +76,21 @@ func TestUnassigned(t *testing.T) { } } +func TestUnassignedReservesPort(t *testing.T) { + pm := New() + _ = pm.AddRange(1000, 1002) + + p1, ok1 := pm.Unassigned() + assert.True(t, ok1) + assert.Equal(t, uint16(1000), p1) + + p2, ok2 := pm.Unassigned() + assert.True(t, ok2) + assert.Equal(t, uint16(1001), p2) + + assert.True(t, pm.(*port).ports[1000], "Unassigned must reserve the port") +} + func TestSetStatus(t *testing.T) { pm := New() _ = pm.AddRange(1000, 1002) @@ -83,6 +116,17 @@ func TestSetStatus(t *testing.T) { } } +func TestSetStatusUnknownPort(t *testing.T) { + pm := New() + _ = pm.AddRange(1000, 1002) + + err := pm.SetStatus(5000, true) + assert.Error(t, err) + + _, exists := pm.(*port).ports[5000] + assert.False(t, exists, "SetStatus must not create entries for unknown ports") +} + func TestClaim(t *testing.T) { pm := New() _ = pm.AddRange(1000, 1002) @@ -95,7 +139,7 @@ func TestClaim(t *testing.T) { }{ {"claim unassigned port", 1000, false, true}, {"claim already assigned port", 1001, true, false}, - {"claim non-existent port", 5000, false, true}, + {"claim non-existent port", 5000, false, false}, } for _, tt := range tests { @@ -107,8 +151,13 @@ func TestClaim(t *testing.T) { got := pm.Claim(tt.port) assert.Equal(t, tt.want, got) - finalState := pm.(*port).ports[tt.port] - assert.True(t, finalState) + finalState, exists := pm.(*port).ports[tt.port] + if !tt.want && tt.port == 5000 { + assert.False(t, exists, "out-of-range port must not be added to the registry") + } else { + assert.True(t, exists) + assert.True(t, finalState) + } }) } } diff --git a/internal/session/session.go b/internal/session/session.go index 80b3db6..a2c3d81 100644 --- a/internal/session/session.go +++ b/internal/session/session.go @@ -26,7 +26,7 @@ type Session interface { HandleGlobalRequest(ch <-chan *ssh.Request) error HandleTCPIPForward(req *ssh.Request) error HandleHTTPForward(req *ssh.Request, port uint16) error - HandleTCPForward(req *ssh.Request, addr string, port uint16) error + HandleTCPForward(req *ssh.Request, addr string, port uint16, reserved bool) error Lifecycle() lifecycle.Lifecycle Interaction() interaction.Interaction Forwarder() forwarder.Forwarder @@ -254,35 +254,35 @@ func (s *session) HandleGlobalRequest(GlobalRequest <-chan *ssh.Request) error { return nil } -func (s *session) parseForwardPayload(payload []byte) (address string, port uint16, err error) { +func (s *session) parseForwardPayload(payload []byte) (address string, port uint16, reserved bool, err error) { var forwardPayload struct { BindAddr string BindPort uint32 } if err = ssh.Unmarshal(payload, &forwardPayload); err != nil { - return "", 0, fmt.Errorf("failed to unmarshal forward payload: %w", err) + return "", 0, false, fmt.Errorf("failed to unmarshal forward payload: %w", err) } if forwardPayload.BindPort > 65535 { - return "", 0, fmt.Errorf("port is larger than allowed port of 65535") + return "", 0, false, fmt.Errorf("port is larger than allowed port of 65535") } port = uint16(forwardPayload.BindPort) if isBlockedPort(port) { - return "", 0, fmt.Errorf("port is blocked") + return "", 0, false, fmt.Errorf("port is blocked") } if port == 0 { unassigned, ok := s.lifecycle.PortRegistry().Unassigned() if !ok { - return "", 0, fmt.Errorf("no available port") + return "", 0, false, fmt.Errorf("no available port") } - return forwardPayload.BindAddr, unassigned, nil + return forwardPayload.BindAddr, unassigned, true, nil } - return forwardPayload.BindAddr, port, nil + return forwardPayload.BindAddr, port, false, nil } func (s *session) denyForwardingRequest(req *ssh.Request, key *types.SessionKey, listener io.Closer, msg string) error { @@ -326,7 +326,7 @@ func (s *session) finalizeForwarding(req *ssh.Request, portToBind uint16, listen } func (s *session) HandleTCPIPForward(req *ssh.Request) error { - address, port, err := s.parseForwardPayload(req.Payload) + address, port, reserved, err := s.parseForwardPayload(req.Payload) if err != nil { return s.denyForwardingRequest(req, nil, nil, fmt.Sprintf("cannot parse forwarded payload: %s", err.Error())) } @@ -335,7 +335,7 @@ func (s *session) HandleTCPIPForward(req *ssh.Request) error { case 80, 443: return s.HandleHTTPForward(req, port) default: - return s.HandleTCPForward(req, address, port) + return s.HandleTCPForward(req, address, port, reserved) } } @@ -356,24 +356,35 @@ func (s *session) HandleHTTPForward(req *ssh.Request, portToBind uint16) error { return nil } -func (s *session) HandleTCPForward(req *ssh.Request, addr string, portToBind uint16) error { - if claimed := s.lifecycle.PortRegistry().Claim(portToBind); !claimed { - return s.denyForwardingRequest(req, nil, nil, fmt.Sprintf("Port %d is already in use or restricted", portToBind)) +func (s *session) HandleTCPForward(req *ssh.Request, addr string, portToBind uint16, reserved bool) error { + if !reserved { + if claimed := s.lifecycle.PortRegistry().Claim(portToBind); !claimed { + return s.denyForwardingRequest(req, nil, nil, fmt.Sprintf("Port %d is already in use or restricted", portToBind)) + } + } + + 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 c217947..dc6251d 100644 --- a/internal/session/session_test.go +++ b/internal/session/session_test.go @@ -396,6 +396,12 @@ func TestHandleTCPIPForward_Table(t *testing.T) { err := s.HandleTCPIPForward(req) assert.NoError(t, err) assert.Equal(t, uint16(12345), s.forwarder.ForwardedPort()) + + defer func() { + if l := s.forwarder.Listener(); l != nil { + _ = l.Close() + } + }() }) t.Run("Invalid Payload", func(t *testing.T) { @@ -699,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) @@ -717,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) @@ -736,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) @@ -758,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) @@ -1037,7 +1045,7 @@ func TestParseForwardPayload_Errors(t *testing.T) { s := &session{} t.Run("Short Address", func(t *testing.T) { - _, _, err := s.parseForwardPayload([]byte{0, 0, 0, 4}) + _, _, _, err := s.parseForwardPayload([]byte{0, 0, 0, 4}) if err == nil { t.Error("expected error, got nil") } @@ -1045,7 +1053,7 @@ func TestParseForwardPayload_Errors(t *testing.T) { t.Run("Short Port", func(t *testing.T) { payload := append([]byte{0, 0, 0, 4}, []byte("addr")...) - _, _, err := s.parseForwardPayload(payload) + _, _, _, err := s.parseForwardPayload(payload) if err == nil { t.Error("expected error, got nil") } @@ -1056,7 +1064,7 @@ func TestParseForwardPayload_Errors(t *testing.T) { portBuf := make([]byte, 4) binary.BigEndian.PutUint32(portBuf, 22) payload = append(payload, portBuf...) - _, _, err := s.parseForwardPayload(payload) + _, _, _, err := s.parseForwardPayload(payload) if err == nil { t.Error("expected error, got nil") } else if !strings.Contains(err.Error(), "port is block") { @@ -1220,7 +1228,7 @@ func TestHandleTCPForward_Failures(t *testing.T) { s, _, mPort, _, sReqs, cConn, cleanup := setup(t) defer cleanup() mPort.On("Claim", mock.Anything).Return(false) - err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 1234) + err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 1234, false) if err == nil { t.Error("expected error, got nil") } else if !strings.Contains(err.Error(), "already in use") { @@ -1241,44 +1249,50 @@ 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) + err = s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", port, false) if err == nil { t.Error("expected error, got nil") } 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) + err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 0, false) if err == nil { t.Error("expected error, got nil") } 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, _, sReqs, cConn, cleanup := setup(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() assert.NoError(t, err) - time.Sleep(100 * time.Millisecond) + _ = sConn.Wait() - err = s.HandleTCPForward(req, "localhost", 0) + err = s.HandleTCPForward(req, "localhost", 0, false) if err == nil { t.Error("expected error, got nil") } 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) }) }