fix(port,session): harden port registry and fix concurrent allocation
Docker Build and Push / Run Tests (push) Successful in 2m35s
Docker Build and Push / Build and Push Docker Image (push) Successful in 14m55s

This commit is contained in:
2026-07-18 17:10:35 +07:00
parent 34000962ef
commit 4fcc41eb8a
4 changed files with 96 additions and 36 deletions
+15 -13
View File
@@ -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)
+15 -9
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) {
@@ -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") {