fix(port, session): harden port registry and fix concurent allocation (#155)
SonarQube Scan / SonarQube Trigger (push) Successful in 3m53s

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: #155
Co-authored-by: Bagas <bagas@fossy.my.id>
Co-committed-by: Bagas <bagas@fossy.my.id>
This commit was merged in pull request #155.
This commit is contained in:
2026-07-19 10:59:22 +07:00
committed by bagas
parent 4118872ed8
commit 61a4791bcb
4 changed files with 115 additions and 38 deletions
+25 -11
View File
@@ -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)
})
}