fix(port,session): harden port registry and fix concurrent allocation
This commit is contained in:
@@ -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) {
|
||||
@@ -1037,7 +1043,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 +1051,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 +1062,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 +1226,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") {
|
||||
@@ -1242,7 +1248,7 @@ func TestHandleTCPForward_Failures(t *testing.T) {
|
||||
}(l)
|
||||
port := uint16(l.Addr().(*net.TCPAddr).Port)
|
||||
|
||||
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") {
|
||||
@@ -1255,7 +1261,7 @@ func TestHandleTCPForward_Failures(t *testing.T) {
|
||||
defer cleanup()
|
||||
mPort.On("Claim", mock.Anything).Return(true)
|
||||
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") {
|
||||
@@ -1264,16 +1270,16 @@ func TestHandleTCPForward_Failures(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()
|
||||
mPort.On("Claim", mock.Anything).Return(true)
|
||||
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") {
|
||||
|
||||
Reference in New Issue
Block a user