This commit is contained in:
ShortArrow 2026-04-19 15:25:30 -05:00 committed by GitHub
commit 7b8fb0a0f8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 195 additions and 25 deletions

View file

@ -9,6 +9,7 @@ import (
"os"
"os/exec"
"path"
"runtime"
"time"
)
@ -21,11 +22,12 @@ type CmdKiller interface {
type SSHHandler struct {
oSCommand CmdKiller
dialContext func(ctx context.Context, network, addr string) (io.Closer, error)
startCmd func(*exec.Cmd) error
tempDir func(dir string, pattern string) (name string, err error)
getenv func(key string) string
setenv func(key, value string) error
dialContext func(ctx context.Context, network, addr string) (io.Closer, error)
startCmd func(*exec.Cmd) error
tempDir func(dir string, pattern string) (name string, err error)
findFreePort func() (int, error)
getenv func(key string) string
setenv func(key, value string) error
}
func NewSSHHandler(oSCommand CmdKiller) *SSHHandler {
@ -37,8 +39,17 @@ func NewSSHHandler(oSCommand CmdKiller) *SSHHandler {
},
startCmd: func(cmd *exec.Cmd) error { return cmd.Start() },
tempDir: os.MkdirTemp,
getenv: os.Getenv,
setenv: os.Setenv,
findFreePort: func() (int, error) {
listener, err := net.Listen("tcp", "localhost:0")
if err != nil {
return 0, err
}
port := listener.Addr().(*net.TCPAddr).Port
listener.Close()
return port, nil
},
getenv: os.Getenv,
setenv: os.Setenv,
}
}
@ -85,8 +96,17 @@ func (t *tunneledDockerHost) Close() error {
return t.oSCommand.Kill(t.cmd)
}
const socketTunnelTimeout = 8 * time.Second
func (self *SSHHandler) createDockerHostTunnel(ctx context.Context, remoteHost string) (*tunneledDockerHost, error) {
socketDir, err := self.tempDir("/tmp", "lazydocker-sshtunnel-")
if runtime.GOOS == "windows" {
return self.createDockerHostTunnelTCP(ctx, remoteHost)
}
return self.createDockerHostTunnelUnix(ctx, remoteHost)
}
func (self *SSHHandler) createDockerHostTunnelUnix(ctx context.Context, remoteHost string) (*tunneledDockerHost, error) {
socketDir, err := self.tempDir(os.TempDir(), "lazydocker-sshtunnel-")
if err != nil {
return nil, fmt.Errorf("create ssh tunnel tmp file: %w", err)
}
@ -99,11 +119,10 @@ func (self *SSHHandler) createDockerHostTunnel(ctx context.Context, remoteHost s
// set a reasonable timeout, then wait for the socket to dial successfully
// before attempting to create a new docker client
const socketTunnelTimeout = 8 * time.Second
ctx, cancel := context.WithTimeout(ctx, socketTunnelTimeout)
defer cancel()
err = self.retrySocketDial(ctx, localSocket)
err = self.retrySocketDial(ctx, "unix", localSocket)
if err != nil {
return nil, fmt.Errorf("ssh tunneled socket never became available: %w", err)
}
@ -119,7 +138,7 @@ func (self *SSHHandler) createDockerHostTunnel(ctx context.Context, remoteHost s
// Attempt to dial the socket until it becomes available.
// The retry loop will continue until the parent context is canceled.
func (self *SSHHandler) retrySocketDial(ctx context.Context, socketPath string) error {
func (self *SSHHandler) retrySocketDial(ctx context.Context, network, address string) error {
t := time.NewTicker(1 * time.Second)
defer t.Stop()
@ -130,7 +149,7 @@ func (self *SSHHandler) retrySocketDial(ctx context.Context, socketPath string)
case <-t.C:
}
// attempt to dial the socket, exit on success
err := self.tryDial(ctx, socketPath)
err := self.tryDial(ctx, network, address)
if err != nil {
continue
}
@ -138,9 +157,9 @@ func (self *SSHHandler) retrySocketDial(ctx context.Context, socketPath string)
}
}
// Try to dial the specified unix socket, immediately close the connection if successfully created.
func (self *SSHHandler) tryDial(ctx context.Context, socketPath string) error {
conn, err := self.dialContext(ctx, "unix", socketPath)
// Try to dial the specified socket, immediately close the connection if successfully created.
func (self *SSHHandler) tryDial(ctx context.Context, network, address string) error {
conn, err := self.dialContext(ctx, network, address)
if err != nil {
return err
}
@ -157,3 +176,33 @@ func (self *SSHHandler) tunnelSSH(ctx context.Context, host, localSocket string)
}
return cmd, nil
}
func (self *SSHHandler) createDockerHostTunnelTCP(ctx context.Context, remoteHost string) (*tunneledDockerHost, error) {
port, err := self.findFreePort()
if err != nil {
return nil, fmt.Errorf("find free port for ssh tunnel: %w", err)
}
localAddr := fmt.Sprintf("localhost:%d", port)
cmd, err := self.tunnelSSH(ctx, remoteHost, localAddr)
if err != nil {
return nil, fmt.Errorf("tunnel docker host over ssh: %w", err)
}
ctx, cancel := context.WithTimeout(ctx, socketTunnelTimeout)
defer cancel()
err = self.retrySocketDial(ctx, "tcp", localAddr)
if err != nil {
self.oSCommand.Kill(cmd)
return nil, fmt.Errorf("ssh tunneled socket never became available: %w", err)
}
newDockerHostURL := url.URL{Scheme: "tcp", Host: localAddr}
return &tunneledDockerHost{
socketPath: newDockerHostURL.String(),
cmd: cmd,
oSCommand: self.oSCommand,
}, nil
}

View file

@ -2,8 +2,11 @@ package ssh
import (
"context"
"fmt"
"io"
"os"
"os/exec"
"runtime"
"testing"
"github.com/stretchr/testify/assert"
@ -50,21 +53,41 @@ func TestSSHHandlerHandleSSHDockerHost(t *testing.T) {
}
tempDir := func(dir string, pattern string) (string, error) {
assert.Equal(t, "/tmp", dir)
assert.Equal(t, os.TempDir(), dir)
assert.Equal(t, "lazydocker-sshtunnel-", pattern)
return "/tmp/lazydocker-ssh-tunnel-12345", nil
}
findFreePort := func() (int, error) {
return 12345, nil
}
var expectedDockerHost string
var expectedNetwork string
var expectedAddress string
var expectedCmdArgs []string
if runtime.GOOS == "windows" {
expectedDockerHost = "tcp://localhost:12345"
expectedNetwork = "tcp"
expectedAddress = "localhost:12345"
expectedCmdArgs = []string{"ssh", "-L", "localhost:12345:/var/run/docker.sock", "192.168.5.178", "-N"}
} else {
expectedDockerHost = "unix:///tmp/lazydocker-ssh-tunnel-12345/dockerhost.sock"
expectedNetwork = "unix"
expectedAddress = "/tmp/lazydocker-ssh-tunnel-12345/dockerhost.sock"
expectedCmdArgs = []string{"ssh", "-L", "/tmp/lazydocker-ssh-tunnel-12345/dockerhost.sock:/var/run/docker.sock", "192.168.5.178", "-N"}
}
setenv := func(key, value string) error {
assert.Equal(t, "DOCKER_HOST", key)
assert.Equal(t, "unix:///tmp/lazydocker-ssh-tunnel-12345/dockerhost.sock", value)
assert.Equal(t, expectedDockerHost, value)
return nil
}
startCmdCount := 0
startCmd := func(cmd *exec.Cmd) error {
assert.EqualValues(t, []string{"ssh", "-L", "/tmp/lazydocker-ssh-tunnel-12345/dockerhost.sock:/var/run/docker.sock", "192.168.5.178", "-N"}, cmd.Args)
assert.EqualValues(t, expectedCmdArgs, cmd.Args)
startCmdCount++
@ -73,8 +96,8 @@ func TestSSHHandlerHandleSSHDockerHost(t *testing.T) {
dialContextCount := 0
dialContext := func(ctx context.Context, network string, address string) (io.Closer, error) {
assert.Equal(t, "unix", network)
assert.Equal(t, "/tmp/lazydocker-ssh-tunnel-12345/dockerhost.sock", address)
assert.Equal(t, expectedNetwork, network)
assert.Equal(t, expectedAddress, address)
dialContextCount++
@ -84,11 +107,12 @@ func TestSSHHandlerHandleSSHDockerHost(t *testing.T) {
handler := &SSHHandler{
oSCommand: &fakeCmdKiller{},
dialContext: dialContext,
startCmd: startCmd,
tempDir: tempDir,
getenv: getenv,
setenv: setenv,
dialContext: dialContext,
startCmd: startCmd,
tempDir: tempDir,
findFreePort: findFreePort,
getenv: getenv,
setenv: setenv,
}
_, err := handler.HandleSSHDockerHost()
@ -100,6 +124,103 @@ func TestSSHHandlerHandleSSHDockerHost(t *testing.T) {
}
}
func TestCreateDockerHostTunnelUnix(t *testing.T) {
remoteHost := "192.168.5.178"
socketDir := "/tmp/lazydocker-ssh-tunnel-12345"
expectedNetwork := "unix"
expectedAddress := socketDir + "/dockerhost.sock"
expectedDockerHost := "unix://" + expectedAddress
expectedCmdArgs := []string{"ssh", "-L", expectedAddress + ":/var/run/docker.sock", remoteHost, "-N"}
tempDir := func(dir string, pattern string) (string, error) {
return socketDir, nil
}
startCmd := func(cmd *exec.Cmd) error {
assert.EqualValues(t, expectedCmdArgs, cmd.Args)
return nil
}
dialContext := func(ctx context.Context, network string, address string) (io.Closer, error) {
assert.Equal(t, expectedNetwork, network)
assert.Equal(t, expectedAddress, address)
return noopCloser{}, nil
}
handler := &SSHHandler{
oSCommand: &fakeCmdKiller{},
dialContext: dialContext,
startCmd: startCmd,
tempDir: tempDir,
findFreePort: func() (int, error) { return 0, nil },
getenv: os.Getenv,
setenv: func(k, v string) error { return nil },
}
tunnel, err := handler.createDockerHostTunnelUnix(context.Background(), remoteHost)
assert.NoError(t, err)
assert.Equal(t, expectedDockerHost, tunnel.socketPath)
}
func TestCreateDockerHostTunnelTCP(t *testing.T) {
remoteHost := "192.168.5.178"
port := 54321
expectedNetwork := "tcp"
expectedAddress := fmt.Sprintf("localhost:%d", port)
expectedDockerHost := "tcp://" + expectedAddress
expectedCmdArgs := []string{"ssh", "-L", expectedAddress + ":/var/run/docker.sock", remoteHost, "-N"}
startCmd := func(cmd *exec.Cmd) error {
assert.EqualValues(t, expectedCmdArgs, cmd.Args)
return nil
}
dialContext := func(ctx context.Context, network string, address string) (io.Closer, error) {
assert.Equal(t, expectedNetwork, network)
assert.Equal(t, expectedAddress, address)
return noopCloser{}, nil
}
handler := &SSHHandler{
oSCommand: &fakeCmdKiller{},
dialContext: dialContext,
startCmd: startCmd,
tempDir: func(string, string) (string, error) { return "", nil },
findFreePort: func() (int, error) { return port, nil },
getenv: os.Getenv,
setenv: func(k, v string) error { return nil },
}
tunnel, err := handler.createDockerHostTunnelTCP(context.Background(), remoteHost)
assert.NoError(t, err)
assert.Equal(t, expectedDockerHost, tunnel.socketPath)
}
func TestCreateDockerHostTunnelTCP_FindFreePortError(t *testing.T) {
remoteHost := "192.168.5.178"
handler := &SSHHandler{
oSCommand: &fakeCmdKiller{},
findFreePort: func() (int, error) { return 0, fmt.Errorf("no ports available") },
}
_, err := handler.createDockerHostTunnelTCP(context.Background(), remoteHost)
assert.Error(t, err)
assert.Contains(t, err.Error(), "find free port for ssh tunnel")
}
func TestCreateDockerHostTunnelTCP_TunnelSSHError(t *testing.T) {
remoteHost := "192.168.5.178"
handler := &SSHHandler{
oSCommand: &fakeCmdKiller{},
findFreePort: func() (int, error) { return 54321, nil },
startCmd: func(cmd *exec.Cmd) error { return fmt.Errorf("ssh not found") },
}
_, err := handler.createDockerHostTunnelTCP(context.Background(), remoteHost)
assert.Error(t, err)
assert.Contains(t, err.Error(), "tunnel docker host over ssh")
}
type fakeCmdKiller struct{}
func (self *fakeCmdKiller) Kill(cmd *exec.Cmd) error {