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..785be25 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,9 +356,11 @@ 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)) + } } tcpServer := transport.NewTCPServer(portToBind, s.forwarder) diff --git a/internal/session/session_test.go b/internal/session/session_test.go index c217947..07d6024 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) { @@ -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") {