fix(port, session): harden port registry and fix concurent allocation (#155)
SonarQube Scan / SonarQube Trigger (push) Successful in 3m53s
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:
+14
-11
@@ -33,10 +33,15 @@ func (pm *port) AddRange(startPort, endPort uint16) error {
|
|||||||
if startPort > endPort {
|
if startPort > endPort {
|
||||||
return fmt.Errorf("start port cannot be greater than end port")
|
return fmt.Errorf("start port cannot be greater than end port")
|
||||||
}
|
}
|
||||||
for index := startPort; index <= endPort; index++ {
|
for index := startPort; ; index++ {
|
||||||
if _, exists := pm.ports[index]; !exists {
|
if index != 0 {
|
||||||
pm.ports[index] = false
|
if _, exists := pm.ports[index]; !exists {
|
||||||
pm.sortedPorts = append(pm.sortedPorts, index)
|
pm.ports[index] = false
|
||||||
|
pm.sortedPorts = append(pm.sortedPorts, index)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if index == endPort {
|
||||||
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
sort.Slice(pm.sortedPorts, func(i, j int) bool {
|
sort.Slice(pm.sortedPorts, func(i, j int) bool {
|
||||||
@@ -51,6 +56,7 @@ func (pm *port) Unassigned() (uint16, bool) {
|
|||||||
|
|
||||||
for _, index := range pm.sortedPorts {
|
for _, index := range pm.sortedPorts {
|
||||||
if !pm.ports[index] {
|
if !pm.ports[index] {
|
||||||
|
pm.ports[index] = true
|
||||||
return index, true
|
return index, true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -61,6 +67,9 @@ func (pm *port) SetStatus(port uint16, assigned bool) error {
|
|||||||
pm.mu.Lock()
|
pm.mu.Lock()
|
||||||
defer pm.mu.Unlock()
|
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
|
pm.ports[port] = assigned
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -70,16 +79,10 @@ func (pm *port) Claim(port uint16) (claimed bool) {
|
|||||||
defer pm.mu.Unlock()
|
defer pm.mu.Unlock()
|
||||||
|
|
||||||
status, exists := pm.ports[port]
|
status, exists := pm.ports[port]
|
||||||
|
if !exists || status {
|
||||||
if exists && status {
|
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
if !exists {
|
|
||||||
pm.ports[port] = true
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
pm.ports[port] = true
|
pm.ports[port] = true
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -16,6 +16,8 @@ func TestAddRange(t *testing.T) {
|
|||||||
{"normal range", 1000, 1002, false},
|
{"normal range", 1000, 1002, false},
|
||||||
{"invalid range", 2000, 1999, true},
|
{"invalid range", 2000, 1999, true},
|
||||||
{"single port range", 3000, 3000, false},
|
{"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 {
|
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) {
|
func TestUnassigned(t *testing.T) {
|
||||||
pm := New()
|
pm := New()
|
||||||
_ = pm.AddRange(1000, 1002)
|
_ = 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) {
|
func TestSetStatus(t *testing.T) {
|
||||||
pm := New()
|
pm := New()
|
||||||
_ = pm.AddRange(1000, 1002)
|
_ = 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) {
|
func TestClaim(t *testing.T) {
|
||||||
pm := New()
|
pm := New()
|
||||||
_ = pm.AddRange(1000, 1002)
|
_ = pm.AddRange(1000, 1002)
|
||||||
@@ -95,7 +139,7 @@ func TestClaim(t *testing.T) {
|
|||||||
}{
|
}{
|
||||||
{"claim unassigned port", 1000, false, true},
|
{"claim unassigned port", 1000, false, true},
|
||||||
{"claim already assigned port", 1001, true, false},
|
{"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 {
|
for _, tt := range tests {
|
||||||
@@ -107,8 +151,13 @@ func TestClaim(t *testing.T) {
|
|||||||
got := pm.Claim(tt.port)
|
got := pm.Claim(tt.port)
|
||||||
assert.Equal(t, tt.want, got)
|
assert.Equal(t, tt.want, got)
|
||||||
|
|
||||||
finalState := pm.(*port).ports[tt.port]
|
finalState, exists := pm.(*port).ports[tt.port]
|
||||||
assert.True(t, finalState)
|
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)
|
||||||
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+24
-13
@@ -26,7 +26,7 @@ type Session interface {
|
|||||||
HandleGlobalRequest(ch <-chan *ssh.Request) error
|
HandleGlobalRequest(ch <-chan *ssh.Request) error
|
||||||
HandleTCPIPForward(req *ssh.Request) error
|
HandleTCPIPForward(req *ssh.Request) error
|
||||||
HandleHTTPForward(req *ssh.Request, port uint16) 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
|
Lifecycle() lifecycle.Lifecycle
|
||||||
Interaction() interaction.Interaction
|
Interaction() interaction.Interaction
|
||||||
Forwarder() forwarder.Forwarder
|
Forwarder() forwarder.Forwarder
|
||||||
@@ -254,35 +254,35 @@ func (s *session) HandleGlobalRequest(GlobalRequest <-chan *ssh.Request) error {
|
|||||||
return nil
|
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 {
|
var forwardPayload struct {
|
||||||
BindAddr string
|
BindAddr string
|
||||||
BindPort uint32
|
BindPort uint32
|
||||||
}
|
}
|
||||||
|
|
||||||
if err = ssh.Unmarshal(payload, &forwardPayload); err != nil {
|
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 {
|
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)
|
port = uint16(forwardPayload.BindPort)
|
||||||
|
|
||||||
if isBlockedPort(port) {
|
if isBlockedPort(port) {
|
||||||
return "", 0, fmt.Errorf("port is blocked")
|
return "", 0, false, fmt.Errorf("port is blocked")
|
||||||
}
|
}
|
||||||
|
|
||||||
if port == 0 {
|
if port == 0 {
|
||||||
unassigned, ok := s.lifecycle.PortRegistry().Unassigned()
|
unassigned, ok := s.lifecycle.PortRegistry().Unassigned()
|
||||||
if !ok {
|
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 {
|
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 {
|
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 {
|
if err != nil {
|
||||||
return s.denyForwardingRequest(req, nil, nil, fmt.Sprintf("cannot parse forwarded payload: %s", err.Error()))
|
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:
|
case 80, 443:
|
||||||
return s.HandleHTTPForward(req, port)
|
return s.HandleHTTPForward(req, port)
|
||||||
default:
|
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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *session) HandleTCPForward(req *ssh.Request, addr string, portToBind uint16) error {
|
func (s *session) HandleTCPForward(req *ssh.Request, addr string, portToBind uint16, reserved bool) error {
|
||||||
if claimed := s.lifecycle.PortRegistry().Claim(portToBind); !claimed {
|
if !reserved {
|
||||||
return s.denyForwardingRequest(req, nil, nil, fmt.Sprintf("Port %d is already in use or restricted", portToBind))
|
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)
|
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))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -396,6 +396,12 @@ func TestHandleTCPIPForward_Table(t *testing.T) {
|
|||||||
err := s.HandleTCPIPForward(req)
|
err := s.HandleTCPIPForward(req)
|
||||||
assert.NoError(t, err)
|
assert.NoError(t, err)
|
||||||
assert.Equal(t, uint16(12345), s.forwarder.ForwardedPort())
|
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) {
|
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)
|
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)
|
||||||
@@ -717,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)
|
||||||
@@ -736,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)
|
||||||
@@ -758,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)
|
||||||
@@ -1037,7 +1045,7 @@ func TestParseForwardPayload_Errors(t *testing.T) {
|
|||||||
s := &session{}
|
s := &session{}
|
||||||
|
|
||||||
t.Run("Short Address", func(t *testing.T) {
|
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 {
|
if err == nil {
|
||||||
t.Error("expected error, got 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) {
|
t.Run("Short Port", func(t *testing.T) {
|
||||||
payload := append([]byte{0, 0, 0, 4}, []byte("addr")...)
|
payload := append([]byte{0, 0, 0, 4}, []byte("addr")...)
|
||||||
_, _, err := s.parseForwardPayload(payload)
|
_, _, _, err := s.parseForwardPayload(payload)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Error("expected error, got nil")
|
t.Error("expected error, got nil")
|
||||||
}
|
}
|
||||||
@@ -1056,7 +1064,7 @@ func TestParseForwardPayload_Errors(t *testing.T) {
|
|||||||
portBuf := make([]byte, 4)
|
portBuf := make([]byte, 4)
|
||||||
binary.BigEndian.PutUint32(portBuf, 22)
|
binary.BigEndian.PutUint32(portBuf, 22)
|
||||||
payload = append(payload, portBuf...)
|
payload = append(payload, portBuf...)
|
||||||
_, _, err := s.parseForwardPayload(payload)
|
_, _, _, err := s.parseForwardPayload(payload)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Error("expected error, got nil")
|
t.Error("expected error, got nil")
|
||||||
} else if !strings.Contains(err.Error(), "port is block") {
|
} 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)
|
s, _, mPort, _, sReqs, cConn, cleanup := setup(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
mPort.On("Claim", mock.Anything).Return(false)
|
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 {
|
if err == nil {
|
||||||
t.Error("expected error, got nil")
|
t.Error("expected error, got nil")
|
||||||
} else if !strings.Contains(err.Error(), "already in use") {
|
} else if !strings.Contains(err.Error(), "already in use") {
|
||||||
@@ -1241,44 +1249,50 @@ 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)
|
err = s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", port, false)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Error("expected error, got nil")
|
t.Error("expected error, got nil")
|
||||||
} 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)
|
err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 0, false)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Error("expected error, got nil")
|
t.Error("expected error, got nil")
|
||||||
} 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, _, 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()
|
||||||
assert.NoError(t, err)
|
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 {
|
if err == nil {
|
||||||
t.Error("expected error, got nil")
|
t.Error("expected error, got nil")
|
||||||
} 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)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user