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"
"os/exec" "os/exec"
"path" "path"
"runtime"
"time" "time"
) )
@ -21,11 +22,12 @@ type CmdKiller interface {
type SSHHandler struct { type SSHHandler struct {
oSCommand CmdKiller oSCommand CmdKiller
dialContext func(ctx context.Context, network, addr string) (io.Closer, error) dialContext func(ctx context.Context, network, addr string) (io.Closer, error)
startCmd func(*exec.Cmd) error startCmd func(*exec.Cmd) error
tempDir func(dir string, pattern string) (name string, err error) tempDir func(dir string, pattern string) (name string, err error)
getenv func(key string) string findFreePort func() (int, error)
setenv func(key, value string) error getenv func(key string) string
setenv func(key, value string) error
} }
func NewSSHHandler(oSCommand CmdKiller) *SSHHandler { func NewSSHHandler(oSCommand CmdKiller) *SSHHandler {
@ -37,8 +39,17 @@ func NewSSHHandler(oSCommand CmdKiller) *SSHHandler {
}, },
startCmd: func(cmd *exec.Cmd) error { return cmd.Start() }, startCmd: func(cmd *exec.Cmd) error { return cmd.Start() },
tempDir: os.MkdirTemp, tempDir: os.MkdirTemp,
getenv: os.Getenv, findFreePort: func() (int, error) {
setenv: os.Setenv, 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) return t.oSCommand.Kill(t.cmd)
} }
const socketTunnelTimeout = 8 * time.Second
func (self *SSHHandler) createDockerHostTunnel(ctx context.Context, remoteHost string) (*tunneledDockerHost, error) { 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 { if err != nil {
return nil, fmt.Errorf("create ssh tunnel tmp file: %w", err) 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 // set a reasonable timeout, then wait for the socket to dial successfully
// before attempting to create a new docker client // before attempting to create a new docker client
const socketTunnelTimeout = 8 * time.Second
ctx, cancel := context.WithTimeout(ctx, socketTunnelTimeout) ctx, cancel := context.WithTimeout(ctx, socketTunnelTimeout)
defer cancel() defer cancel()
err = self.retrySocketDial(ctx, localSocket) err = self.retrySocketDial(ctx, "unix", localSocket)
if err != nil { if err != nil {
return nil, fmt.Errorf("ssh tunneled socket never became available: %w", err) 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. // Attempt to dial the socket until it becomes available.
// The retry loop will continue until the parent context is canceled. // 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) t := time.NewTicker(1 * time.Second)
defer t.Stop() defer t.Stop()
@ -130,7 +149,7 @@ func (self *SSHHandler) retrySocketDial(ctx context.Context, socketPath string)
case <-t.C: case <-t.C:
} }
// attempt to dial the socket, exit on success // attempt to dial the socket, exit on success
err := self.tryDial(ctx, socketPath) err := self.tryDial(ctx, network, address)
if err != nil { if err != nil {
continue 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. // Try to dial the specified socket, immediately close the connection if successfully created.
func (self *SSHHandler) tryDial(ctx context.Context, socketPath string) error { func (self *SSHHandler) tryDial(ctx context.Context, network, address string) error {
conn, err := self.dialContext(ctx, "unix", socketPath) conn, err := self.dialContext(ctx, network, address)
if err != nil { if err != nil {
return err return err
} }
@ -157,3 +176,33 @@ func (self *SSHHandler) tunnelSSH(ctx context.Context, host, localSocket string)
} }
return cmd, nil 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 ( import (
"context" "context"
"fmt"
"io" "io"
"os"
"os/exec" "os/exec"
"runtime"
"testing" "testing"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
@ -50,21 +53,41 @@ func TestSSHHandlerHandleSSHDockerHost(t *testing.T) {
} }
tempDir := func(dir string, pattern string) (string, error) { 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) assert.Equal(t, "lazydocker-sshtunnel-", pattern)
return "/tmp/lazydocker-ssh-tunnel-12345", nil 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 { setenv := func(key, value string) error {
assert.Equal(t, "DOCKER_HOST", key) 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 return nil
} }
startCmdCount := 0 startCmdCount := 0
startCmd := func(cmd *exec.Cmd) error { 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++ startCmdCount++
@ -73,8 +96,8 @@ func TestSSHHandlerHandleSSHDockerHost(t *testing.T) {
dialContextCount := 0 dialContextCount := 0
dialContext := func(ctx context.Context, network string, address string) (io.Closer, error) { dialContext := func(ctx context.Context, network string, address string) (io.Closer, error) {
assert.Equal(t, "unix", network) assert.Equal(t, expectedNetwork, network)
assert.Equal(t, "/tmp/lazydocker-ssh-tunnel-12345/dockerhost.sock", address) assert.Equal(t, expectedAddress, address)
dialContextCount++ dialContextCount++
@ -84,11 +107,12 @@ func TestSSHHandlerHandleSSHDockerHost(t *testing.T) {
handler := &SSHHandler{ handler := &SSHHandler{
oSCommand: &fakeCmdKiller{}, oSCommand: &fakeCmdKiller{},
dialContext: dialContext, dialContext: dialContext,
startCmd: startCmd, startCmd: startCmd,
tempDir: tempDir, tempDir: tempDir,
getenv: getenv, findFreePort: findFreePort,
setenv: setenv, getenv: getenv,
setenv: setenv,
} }
_, err := handler.HandleSSHDockerHost() _, 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{} type fakeCmdKiller struct{}
func (self *fakeCmdKiller) Kill(cmd *exec.Cmd) error { func (self *fakeCmdKiller) Kill(cmd *exec.Cmd) error {