From ee681f80e900a874e118f7b26ba5a1a4ddbb375e Mon Sep 17 00:00:00 2001 From: Sean Prouty Date: Tue, 21 Jul 2026 08:57:44 -0700 Subject: [PATCH] pkg/bindings: add dial-stdio fallback for SSH connections Fixes: #27814 Signed-off-by: Sean Prouty --- pkg/bindings/connection.go | 5 +- pkg/bindings/connection_stdio.go | 161 ++++++++ pkg/bindings/connection_stdio_test.go | 551 ++++++++++++++++++++++++++ test/e2e/system_dial_stdio_test.go | 3 - 4 files changed, 713 insertions(+), 7 deletions(-) create mode 100644 pkg/bindings/connection_stdio.go create mode 100644 pkg/bindings/connection_stdio_test.go diff --git a/pkg/bindings/connection.go b/pkg/bindings/connection.go index ce2e82ff0f..13d29d1ac0 100644 --- a/pkg/bindings/connection.go +++ b/pkg/bindings/connection.go @@ -297,12 +297,9 @@ func sshClient(_url *url.URL, uri string, identity string, machine bool) (Connec val := strings.TrimSuffix(b.String(), "\n") _url.Path = val } - dialContext := func(ctx context.Context, _, _ string) (net.Conn, error) { - return ssh.DialNetContext(ctx, conn, "unix", _url) - } connection.Client = &http.Client{ Transport: &http.Transport{ - DialContext: dialContext, + DialContext: newSSHDialContext(conn, _url), }, } return connection, nil diff --git a/pkg/bindings/connection_stdio.go b/pkg/bindings/connection_stdio.go new file mode 100644 index 0000000000..6f255878c1 --- /dev/null +++ b/pkg/bindings/connection_stdio.go @@ -0,0 +1,161 @@ +package bindings + +import ( + "bytes" + "context" + "errors" + "io" + "net" + "net/url" + "strings" + "sync/atomic" + "time" + + "github.com/sirupsen/logrus" + commonssh "go.podman.io/common/pkg/ssh" + "go.podman.io/storage/pkg/stringutils" + "golang.org/x/crypto/ssh" +) + +func isFallbackableChannelErr(err error) bool { + var openChannelErr *ssh.OpenChannelError + if !errors.As(err, &openChannelErr) { + return false + } + + return openChannelErr.Reason == ssh.UnknownChannelType || openChannelErr.Reason == ssh.ConnectionFailed +} + +func newSSHDialContext(client *ssh.Client, _url *url.URL) func(context.Context, string, string) (net.Conn, error) { + var useFallback atomic.Bool + return func(ctx context.Context, _, _ string) (net.Conn, error) { + if useFallback.Load() { + return dialSSHStdio(client, _url.Path) + } + + conn, err := commonssh.DialNetContext(ctx, client, "unix", _url) + if err == nil || !isFallbackableChannelErr(err) { + return conn, err + } + + logrus.Debugf("direct-streamlocal channel refused (%v), trying dial-stdio fallback", err) + fallbackConn, fallbackErr := dialSSHStdio(client, _url.Path) + if fallbackErr != nil { + return nil, errors.Join(err, fallbackErr) + } + + useFallback.Store(true) + return fallbackConn, nil + } +} + +func dialSSHStdio(client *ssh.Client, path string) (net.Conn, error) { + session, err := client.NewSession() + if err != nil { + return nil, err + } + + stdin, err := session.StdinPipe() + if err != nil { + session.Close() + return nil, err + } + + stdout, err := session.StdoutPipe() + if err != nil { + session.Close() + return nil, err + } + + var stderr bytes.Buffer + session.Stderr = &stderr + + cmd := stringutils.ShellQuoteArguments([]string{"podman", "--url", "unix://" + path, "system", "dial-stdio"}) + if err := session.Start(cmd); err != nil { + session.Close() + return nil, err + } + + conn := &sshStdioConn{ + path: path, + session: session, + writer: stdin, + reader: stdout, + sessionDone: make(chan struct{}), + closeTimeout: 5 * time.Second, + } + + go func() { + defer close(conn.sessionDone) + waitErr := session.Wait() + stderrOut := strings.TrimRight(stderr.String(), "\n") + switch { + case waitErr != nil && stderrOut != "": + logrus.Errorf("ssh session error: %v: %s", waitErr, stderrOut) + case waitErr != nil: + logrus.Errorf("ssh session error: %v", waitErr) + case stderrOut != "": + logrus.Debugf("dial-stdio stderr: %s", stderrOut) + } + }() + + return conn, nil +} + +type sshStdioConn struct { + path string + session io.Closer + writer io.WriteCloser + reader io.Reader + sessionDone chan struct{} + closeTimeout time.Duration +} + +func (c *sshStdioConn) Close() error { + err := c.writer.Close() + + select { + case <-c.sessionDone: + case <-time.After(c.closeTimeout): + logrus.Debugf("timed out waiting for dial-stdio session to exit") + } + + if sessionErr := c.session.Close(); sessionErr != nil && !errors.Is(sessionErr, io.EOF) { + err = errors.Join(err, sessionErr) + } + + return err +} + +func (c *sshStdioConn) LocalAddr() net.Addr { + return &net.UnixAddr{Name: "@", Net: "unix"} +} + +func (c *sshStdioConn) RemoteAddr() net.Addr { + return &net.UnixAddr{Name: c.path, Net: "unix"} +} + +func (c *sshStdioConn) Read(b []byte) (n int, err error) { + return c.reader.Read(b) +} + +func (c *sshStdioConn) Write(b []byte) (n int, err error) { + return c.writer.Write(b) +} + +// http.Transport wraps this connection and manages its own timeouts, so the deadline methods are no-ops. + +func (c *sshStdioConn) SetDeadline(_ time.Time) error { + logrus.Debugf("SetDeadline not implemented for sshStdioConn") + return nil +} + +func (c *sshStdioConn) SetReadDeadline(_ time.Time) error { + logrus.Debugf("SetReadDeadline not implemented for sshStdioConn") + return nil +} + +func (c *sshStdioConn) SetWriteDeadline(_ time.Time) error { + logrus.Debugf("SetWriteDeadline not implemented for sshStdioConn") + return nil +} diff --git a/pkg/bindings/connection_stdio_test.go b/pkg/bindings/connection_stdio_test.go new file mode 100644 index 0000000000..bbd3ee8599 --- /dev/null +++ b/pkg/bindings/connection_stdio_test.go @@ -0,0 +1,551 @@ +package bindings + +import ( + "crypto/ed25519" + "crypto/rand" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "net/url" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/crypto/ssh" +) + +func TestIsFallbackableChannelErr(t *testing.T) { + tests := []struct { + name string + err error + want bool + }{ + { + name: "nil error", + err: nil, + want: false, + }, + { + name: "UnknownChannelType", + err: &ssh.OpenChannelError{Reason: ssh.UnknownChannelType, Message: "unknown channel type (unsupported channel type)"}, + want: true, + }, + { + name: "wrapped UnknownChannelType", + err: fmt.Errorf("dial failed: %w", &ssh.OpenChannelError{Reason: ssh.UnknownChannelType}), + want: true, + }, + { + name: "ConnectionFailed", + err: &ssh.OpenChannelError{Reason: ssh.ConnectionFailed, Message: "open failed"}, + want: true, + }, + { + name: "wrapped ConnectionFailed", + err: fmt.Errorf("dial failed: %w", &ssh.OpenChannelError{Reason: ssh.ConnectionFailed}), + want: true, + }, + { + // An explicit administrative denial, routing around it is not our call. + name: "prohibited", + err: &ssh.OpenChannelError{Reason: ssh.Prohibited, Message: "prohibited"}, + want: false, + }, + { + name: "resource shortage", + err: &ssh.OpenChannelError{Reason: ssh.ResourceShortage, Message: "resource shortage"}, + want: false, + }, + { + name: "connection refused", + err: errors.New("connection refused"), + want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, isFallbackableChannelErr(tt.err)) + }) + } +} + +const allowStreamlocal = ssh.RejectionReason(0) + +type stdioTestServerOptions struct { + streamlocalReason ssh.RejectionReason + rejectSessions bool + backend string +} + +type stdioTestServer struct { + opts stdioTestServerOptions + addr string + + mu sync.Mutex + streamlocalAttempts int + execCommands []string +} + +func (s *stdioTestServer) StreamlocalAttempts() int { + s.mu.Lock() + defer s.mu.Unlock() + return s.streamlocalAttempts +} + +func (s *stdioTestServer) ExecCommands() []string { + s.mu.Lock() + defer s.mu.Unlock() + return append([]string(nil), s.execCommands...) +} + +func newStdioTestServer(t *testing.T, opts stdioTestServerOptions) (*stdioTestServer, *ssh.Client) { + t.Helper() + + backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + fmt.Fprintf(w, "pong %s", r.URL.Path) + })) + t.Cleanup(backend.Close) + + opts.backend = backend.Listener.Addr().String() + s := &stdioTestServer{opts: opts} + + _, priv, err := ed25519.GenerateKey(rand.Reader) + require.NoError(t, err) + signer, err := ssh.NewSignerFromKey(priv) + require.NoError(t, err) + + serverConfig := &ssh.ServerConfig{NoClientAuth: true} + serverConfig.AddHostKey(signer) + + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + t.Cleanup(func() { listener.Close() }) + s.addr = listener.Addr().String() + + go func() { + for { + nConn, err := listener.Accept() + if err != nil { + return + } + go s.handleConn(nConn, serverConfig) + } + }() + + client, err := ssh.Dial("tcp", s.addr, &ssh.ClientConfig{ + User: "test", + HostKeyCallback: ssh.FixedHostKey(signer.PublicKey()), + Timeout: 10 * time.Second, + }) + require.NoError(t, err) + t.Cleanup(func() { client.Close() }) + + return s, client +} + +func (s *stdioTestServer) handleConn(nConn net.Conn, config *ssh.ServerConfig) { + conn, chans, reqs, err := ssh.NewServerConn(nConn, config) + if err != nil { + nConn.Close() + return + } + defer conn.Close() + + go ssh.DiscardRequests(reqs) + + for newChannel := range chans { + switch newChannel.ChannelType() { + case "direct-streamlocal@openssh.com": + s.mu.Lock() + s.streamlocalAttempts++ + s.mu.Unlock() + + if s.opts.streamlocalReason != allowStreamlocal { + _ = newChannel.Reject(s.opts.streamlocalReason, "streamlocal refused") + continue + } + + channel, requests, err := newChannel.Accept() + if err != nil { + continue + } + go ssh.DiscardRequests(requests) + go s.bridge(channel, false) + case "session": + if s.opts.rejectSessions { + _ = newChannel.Reject(ssh.Prohibited, "sessions refused") + continue + } + + channel, requests, err := newChannel.Accept() + if err != nil { + continue + } + go s.handleSession(channel, requests) + default: + _ = newChannel.Reject(ssh.UnknownChannelType, "unsupported channel type") + } + } +} + +func (s *stdioTestServer) handleSession(channel ssh.Channel, requests <-chan *ssh.Request) { + for req := range requests { + if req.Type != "exec" { + _ = req.Reply(false, nil) + continue + } + + var payload struct{ Command string } + if err := ssh.Unmarshal(req.Payload, &payload); err != nil { + _ = req.Reply(false, nil) + continue + } + + s.mu.Lock() + s.execCommands = append(s.execCommands, payload.Command) + s.mu.Unlock() + + _ = req.Reply(true, nil) + go s.bridge(channel, true) + } +} + +func (s *stdioTestServer) bridge(channel ssh.Channel, exec bool) { + defer channel.Close() + + backend, err := net.Dial("tcp", s.opts.backend) + if err != nil { + return + } + defer backend.Close() + + go func() { + _, _ = io.Copy(backend, channel) + backend.Close() + }() + _, _ = io.Copy(channel, backend) + + if exec { + _, _ = channel.SendRequest("exit-status", false, ssh.Marshal(struct{ Status uint32 }{Status: 0})) + } +} + +func newTestSocketURL(path string) *url.URL { + return &url.URL{ + Scheme: "ssh", + User: url.User("test"), + Host: "localhost", + Path: path, + } +} + +const testSocketPath = "/run/podman/podman.sock" + +func newSSHTestHTTPClient(client *ssh.Client) *http.Client { + return &http.Client{ + Transport: &http.Transport{ + DialContext: newSSHDialContext(client, newTestSocketURL(testSocketPath)), + DisableKeepAlives: true, + }, + Timeout: 10 * time.Second, + } +} + +func getPing(client *http.Client) (string, error) { + resp, err := client.Get("http://d/ping") + if err != nil { + return "", err + } + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + if err != nil { + return "", err + } + return string(body), nil +} + +func TestSSHDialContextFallsBackOnChannelRejection(t *testing.T) { + tests := []struct { + name string + reason ssh.RejectionReason + }{ + {name: "server does not implement streamlocal", reason: ssh.UnknownChannelType}, + {name: "server refuses streamlocal forwarding", reason: ssh.ConnectionFailed}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + server, client := newStdioTestServer(t, stdioTestServerOptions{streamlocalReason: tt.reason}) + httpClient := newSSHTestHTTPClient(client) + + body, err := getPing(httpClient) + require.NoError(t, err) + assert.Equal(t, "pong /ping", body) + + assert.Equal(t, 1, server.StreamlocalAttempts(), "the direct channel should be tried first") + assert.Equal(t, + []string{"podman --url unix://" + testSocketPath + " system dial-stdio"}, + server.ExecCommands(), + ) + }) + } +} + +func TestSSHDialContextPrefersDirectChannel(t *testing.T) { + server, client := newStdioTestServer(t, stdioTestServerOptions{streamlocalReason: allowStreamlocal}) + httpClient := newSSHTestHTTPClient(client) + + body, err := getPing(httpClient) + require.NoError(t, err) + assert.Equal(t, "pong /ping", body) + + assert.Equal(t, 1, server.StreamlocalAttempts()) + assert.Empty(t, server.ExecCommands(), "the fallback should not run when the direct channel works") +} + +func TestSSHDialContextDoesNotFallBackOnProhibited(t *testing.T) { + server, client := newStdioTestServer(t, stdioTestServerOptions{streamlocalReason: ssh.Prohibited}) + httpClient := newSSHTestHTTPClient(client) + + _, err := getPing(httpClient) + require.Error(t, err) + assert.ErrorContains(t, err, "streamlocal refused") + assert.Empty(t, server.ExecCommands(), "an administrative denial should not be routed around") +} + +func TestSSHDialContextLatchesAfterSuccessfulFallback(t *testing.T) { + server, client := newStdioTestServer(t, stdioTestServerOptions{streamlocalReason: ssh.ConnectionFailed}) + httpClient := newSSHTestHTTPClient(client) + + for range 3 { + body, err := getPing(httpClient) + require.NoError(t, err) + assert.Equal(t, "pong /ping", body) + } + + assert.Equal(t, 1, server.StreamlocalAttempts(), "the direct channel should only be tried until it is known to fail") + assert.Len(t, server.ExecCommands(), 3) +} + +func TestSSHDialContextDoesNotLatchWhenFallbackFails(t *testing.T) { + server, client := newStdioTestServer(t, stdioTestServerOptions{ + streamlocalReason: ssh.ConnectionFailed, + rejectSessions: true, + }) + httpClient := newSSHTestHTTPClient(client) + + for range 2 { + _, err := getPing(httpClient) + require.Error(t, err) + assert.ErrorContains(t, err, "streamlocal refused") + assert.ErrorContains(t, err, "sessions refused") + } + + assert.Equal(t, 2, server.StreamlocalAttempts(), "a failed fallback should not latch") +} + +func TestDialSSHStdioQuotesSocketPath(t *testing.T) { + tests := []struct { + name string + path string + want string + }{ + { + name: "plain path", + path: "/run/user/1000/podman/podman.sock", + want: "podman --url unix:///run/user/1000/podman/podman.sock system dial-stdio", + }, + { + name: "command substitution", + path: "/run/podman/$(touch pwned).sock", + want: `podman --url 'unix:///run/podman/$(touch pwned).sock' system dial-stdio`, + }, + { + name: "embedded single quote", + path: "/run/podman/p'wn.sock", + want: `podman --url 'unix:///run/podman/p'\''wn.sock' system dial-stdio`, + }, + { + name: "command separator", + path: "/run/podman/x; touch pwned", + want: `podman --url 'unix:///run/podman/x; touch pwned' system dial-stdio`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + server, client := newStdioTestServer(t, stdioTestServerOptions{streamlocalReason: ssh.ConnectionFailed}) + + conn, err := dialSSHStdio(client, tt.path) + require.NoError(t, err) + defer conn.Close() + + assert.Eventually(t, func() bool { + return len(server.ExecCommands()) == 1 + }, 5*time.Second, 10*time.Millisecond) + assert.Equal(t, []string{tt.want}, server.ExecCommands()) + }) + } +} + +type mockSession struct { + Closed bool + CloseErr error +} + +func (m *mockSession) Close() error { + m.Closed = true + return m.CloseErr +} + +type newTestStdioConnOptions struct { + sessionDone chan struct{} + closeTimeout time.Duration + session io.Closer +} + +func newTestStdioConn(opts newTestStdioConnOptions) (*sshStdioConn, net.Conn) { + local, remote := net.Pipe() + + if opts.sessionDone == nil { + opts.sessionDone = make(chan struct{}) + close(opts.sessionDone) + } + + if opts.closeTimeout == 0 { + opts.closeTimeout = 5 * time.Second + } + + if opts.session == nil { + opts.session = &mockSession{} + } + + c := &sshStdioConn{ + writer: local, + reader: local, + path: "/run/podman/podman.sock", + sessionDone: opts.sessionDone, + session: opts.session, + closeTimeout: opts.closeTimeout, + } + return c, remote +} + +func TestSshStdioReadWrite(t *testing.T) { + conn, remote := newTestStdioConn(newTestStdioConnOptions{}) + defer conn.Close() + defer remote.Close() + + go func() { + buf := make([]byte, 32) + n, err := remote.Read(buf) + assert.NoError(t, err) + assert.Equal(t, "request", string(buf[:n])) + + _, err = remote.Write([]byte("response")) + assert.NoError(t, err) + }() + + n, err := conn.Write([]byte("request")) + assert.NoError(t, err) + assert.Equal(t, 7, n) + + buf := make([]byte, 32) + n, err = conn.Read(buf) + assert.NoError(t, err) + assert.Equal(t, "response", string(buf[:n])) +} + +func TestSshStdioConnClose(t *testing.T) { + conn, remote := newTestStdioConn(newTestStdioConnOptions{}) + defer remote.Close() + + err := conn.Close() + assert.NoError(t, err) + assert.True(t, conn.session.(*mockSession).Closed, "session should be closed") +} + +func TestSshStdioConnCloseTimeout(t *testing.T) { + mock := &mockSession{} + conn, remote := newTestStdioConn(newTestStdioConnOptions{ + sessionDone: make(chan struct{}), // left unclosed to force timeout + session: mock, + closeTimeout: 50 * time.Millisecond, + }) + + defer remote.Close() + start := time.Now() + err := conn.Close() + assert.NoError(t, err) + assert.GreaterOrEqual(t, time.Since(start), 50*time.Millisecond, "Close should wait for the full timeout") + assert.Less(t, time.Since(start), time.Second, "Close should not block indefinitely") + assert.True(t, mock.Closed, "session.Close should be called after timeout") +} + +func TestSshStdioConnCloseWaitsForSession(t *testing.T) { + session := make(chan struct{}) + conn, remote := newTestStdioConn(newTestStdioConnOptions{ + sessionDone: session, + session: &mockSession{}, + closeTimeout: 1 * time.Second, + }) + + defer remote.Close() + done := make(chan struct{}) + go func() { + conn.Close() + close(done) + }() + + select { + case <-done: + t.Fatal("Close returned before session was signaled") + case <-time.After(50 * time.Millisecond): + } + + close(session) + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("Close did not return after session was signaled") + } + + assert.True(t, conn.session.(*mockSession).Closed, "session should be closed") +} + +func TestSshStdioConnCloseJoinsSessionError(t *testing.T) { + session := &mockSession{CloseErr: errors.New("session close failed")} + conn, remote := newTestStdioConn(newTestStdioConnOptions{session: session}) + defer remote.Close() + + err := conn.Close() + assert.ErrorContains(t, err, "session close failed") + assert.True(t, session.Closed) +} + +func TestSshStdioConnCloseIgnoresEOF(t *testing.T) { + session := &mockSession{CloseErr: io.EOF} + conn, remote := newTestStdioConn(newTestStdioConnOptions{session: session}) + defer remote.Close() + + err := conn.Close() + assert.NoError(t, err) + assert.True(t, session.Closed) +} + +func TestSshStdioConnAddrs(t *testing.T) { + conn, remote := newTestStdioConn(newTestStdioConnOptions{}) + defer conn.Close() + defer remote.Close() + + assert.Equal(t, &net.UnixAddr{Name: "@", Net: "unix"}, conn.LocalAddr()) + assert.Equal(t, &net.UnixAddr{Name: "/run/podman/podman.sock", Net: "unix"}, conn.RemoteAddr()) +} diff --git a/test/e2e/system_dial_stdio_test.go b/test/e2e/system_dial_stdio_test.go index 79aad39959..b7b95588ad 100644 --- a/test/e2e/system_dial_stdio_test.go +++ b/test/e2e/system_dial_stdio_test.go @@ -15,7 +15,4 @@ var _ = Describe("podman system dial-stdio", func() { Expect(session).Should(ExitCleanly()) Expect(session.OutputToString()).To(ContainSubstring("Examples: podman system dial-stdio")) }) - - // TODO: this should have a proper connection test where we spawn a server - // and the use dial-stdio to connect to it and send data. })