Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 1 addition & 3 deletions cmd/internal/agentcontainer/credentials_server.go
Original file line number Diff line number Diff line change
Expand Up @@ -27,8 +27,6 @@ import (
"github.com/spf13/cobra"
)

const ExitCodeIO int = 64

// CredentialsServerCmd holds the cmd flags.
type CredentialsServerCmd struct {
*flags.GlobalFlags
Expand Down Expand Up @@ -103,7 +101,7 @@ func (cmd *CredentialsServerCmd) Run(ctx context.Context, port int) error {
}
defer func() { _ = ln.Close() }()

tunnelClient, err := tunnelserver.NewTunnelClient(os.Stdin, os.Stdout, true, ExitCodeIO)
tunnelClient, err := tunnelserver.NewTunnelClient(os.Stdin, os.Stdout)
if err != nil {
return fmt.Errorf("error creating tunnel client: %w", err)
}
Expand Down
2 changes: 1 addition & 1 deletion cmd/internal/agentcontainer/setup.go
Original file line number Diff line number Diff line change
Expand Up @@ -401,7 +401,7 @@ func buildDeferredHooksCmd(
func (cmd *SetupContainerCmd) initializeTunnelClient(
ctx context.Context,
) (tunnel.TunnelClient, error) {
tunnelClient, err := tunnelserver.NewTunnelClient(os.Stdin, os.Stdout, true, 0)
tunnelClient, err := tunnelserver.NewTunnelClient(os.Stdin, os.Stdout)
if err != nil {
return nil, fmt.Errorf("initializing tunnel client: %w", err)
}
Expand Down
2 changes: 1 addition & 1 deletion cmd/internal/agentworkspace/up.go
Original file line number Diff line number Diff line change
Expand Up @@ -370,7 +370,7 @@ func (w *workspaceInitializer) tryConfigureDockerDaemon(ctx context.Context) {
}

func (w *workspaceInitializer) initializeTunnel(ctx context.Context) error {
client, err := tunnelserver.NewTunnelClient(os.Stdin, os.Stdout, true, 0)
client, err := tunnelserver.NewTunnelClient(os.Stdin, os.Stdout)
if err != nil {
return fmt.Errorf("error creating tunnel client: %w", err)
}
Expand Down
9 changes: 5 additions & 4 deletions cmd/internal/ssh_server.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import (
"encoding/base64"
"errors"
"fmt"
"net"
"os"
"time"

Expand Down Expand Up @@ -106,7 +107,7 @@ func (cmd *sshServerCmd) serveStdio(ctx context.Context, server sshserver.Server
go runActivityHeartbeat(ctx, config.ContainerActivityFile)
}
go shutdownOnCancel(ctx, server) // #nosec G118 -- see shutdownOnCancel.
lis := stdio.NewStdioListener(os.Stdin, os.Stdout, true)
lis := stdio.NewStdioListener(os.Stdin, os.Stdout)
return ignoreServerClosed(server.Serve(lis))
}

Expand Down Expand Up @@ -135,10 +136,10 @@ func shutdownOnCancel(ctx context.Context, server sshserver.Server) {
}
}

// ignoreServerClosed turns the expected post-Shutdown error into a clean
// return so cobra doesn't surface "ssh: Server closed" as a failure.
// ignoreServerClosed turns expected listener shutdown errors into a clean
// return so cobra doesn't surface them as a failure.
func ignoreServerClosed(err error) error {
if err == nil || errors.Is(err, ssh.ErrServerClosed) {
if err == nil || errors.Is(err, ssh.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
return nil
}
return err
Expand Down
3 changes: 3 additions & 0 deletions cmd/internal/ssh_server_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -165,6 +165,9 @@ func TestIgnoreServerClosed(t *testing.T) {
if err := ignoreServerClosed(ssh.ErrServerClosed); err != nil {
t.Errorf("ErrServerClosed should map to nil, got %v", err)
}
if err := ignoreServerClosed(net.ErrClosed); err != nil {
t.Errorf("net.ErrClosed should map to nil, got %v", err)
}
other := errors.New("boom")
if err := ignoreServerClosed(other); err != other {
t.Errorf("unrelated error should pass through, got %v", err)
Expand Down
2 changes: 1 addition & 1 deletion cmd/machine/ssh.go
Original file line number Diff line number Diff line change
Expand Up @@ -188,7 +188,7 @@ func StartSSHSession(ctx context.Context, options StartSSHSessionOptions) error
return options.Exec(ctx, stdin, stdout, options.Stderr)
},
func(ctx context.Context, stdout, stdin *os.File) error {
sshClient, err := devssh.StdioClientWithUser(stdout, stdin, options.User, false)
sshClient, err := devssh.StdioClientWithUser(stdout, stdin, options.User)
if err != nil {
return err
}
Expand Down
2 changes: 1 addition & 1 deletion cmd/snapshot/create.go
Original file line number Diff line number Diff line change
Expand Up @@ -406,7 +406,7 @@ func newLocalTunnelClient(
serverDone <- tunnelServ.Run(serverCtx, clientToServerR, serverToClientW)
}()

tunnelClient, err := tunnelserver.NewTunnelClient(serverToClientR, clientToServerW, false, 0)
tunnelClient, err := tunnelserver.NewTunnelClient(serverToClientR, clientToServerW)
if err != nil {
cancel()
_ = clientToServerW.Close()
Expand Down
2 changes: 1 addition & 1 deletion cmd/workspace/logs.go
Original file line number Diff line number Diff line change
Expand Up @@ -150,7 +150,7 @@ type injectLogsAgentParams struct {
}

func runLogsSession(stdout, stdin *os.File, client clientpkg.WorkspaceClient) error {
sshClient, err := ssh.StdioClientWithUser(stdout, stdin, "", false)
sshClient, err := ssh.StdioClientWithUser(stdout, stdin, "")
if err != nil {
return err
}
Expand Down
6 changes: 3 additions & 3 deletions go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@ require (
github.com/devsy-org/agentapi v1.0.1
github.com/devsy-org/api v1.1.0
github.com/devsy-org/apiserver v1.5.3
github.com/devsy-org/ssh v1.2.2
github.com/devsy-org/ssh v1.2.5
github.com/distribution/reference v0.6.0
github.com/docker/cli v29.7.1+incompatible
github.com/docker/docker v28.5.2+incompatible
Expand Down Expand Up @@ -63,7 +63,7 @@ require (
go.uber.org/atomic v1.11.0
go.uber.org/goleak v1.3.0
go.uber.org/zap v1.28.0
golang.org/x/crypto v0.54.0
golang.org/x/crypto v0.55.0
golang.org/x/mod v0.38.0
golang.org/x/sync v0.22.0
golang.org/x/sys v0.47.0
Expand Down Expand Up @@ -489,7 +489,7 @@ require (
golang.org/x/net v0.57.1-0.20260729233039-99c3b0a8f463 // indirect
golang.org/x/oauth2 v0.36.0 // indirect
golang.org/x/telemetry v0.0.0-20260708182218-49f421fb7959 // indirect
golang.org/x/text v0.40.0 // indirect
golang.org/x/text v0.41.0 // indirect
golang.org/x/time v0.15.0 // indirect
golang.org/x/tools v0.48.0 // indirect
golang.org/x/xerrors v0.0.0-20240903120638-7835f813f4da // indirect
Expand Down
12 changes: 6 additions & 6 deletions go.sum
Original file line number Diff line number Diff line change
Expand Up @@ -407,8 +407,8 @@ github.com/devsy-org/api v1.1.0 h1:l7T9k7RVwatwN4lxeDTF3iN6EYmfGZgR3ZTJMDGha1M=
github.com/devsy-org/api v1.1.0/go.mod h1:mAZklKdnywJYiXDReBLte/H+3m69z6G7RHB3n1lI53Q=
github.com/devsy-org/apiserver v1.5.3 h1:tFKMgPxxfvojJ+C+wo0oqSmbE9PTJ2ur1E2yFj0eISI=
github.com/devsy-org/apiserver v1.5.3/go.mod h1:m7gpbrh++Hp8iEM5jaP7vjjODPbs9Bmh3l4+iPFg1jE=
github.com/devsy-org/ssh v1.2.2 h1:Ylr62ag6nc1U5T6soj0OecPzoLJqPJODUbnJhpaDYSY=
github.com/devsy-org/ssh v1.2.2/go.mod h1:x5NsT8LXW/SBykm115/l5j2jynZAFwIALDcbpi4MlU0=
github.com/devsy-org/ssh v1.2.5 h1:Z7gTanYs2ZslT1swTw4leoVVuDEmuNNhQi48W2kNqMU=
github.com/devsy-org/ssh v1.2.5/go.mod h1:6r5tZ+H9JFoMl6NrxXRgltg9HhEryTd69avY831Lio8=
github.com/devsy-org/tailscale v1.102.2 h1:9SB6htvO+HmG8alal8WGCshHapfHR7dUFRSMYiFIzIM=
github.com/devsy-org/tailscale v1.102.2/go.mod h1:kQUA0lYb/bqCJJZzShx+gqLSxggEFgXB4VDVZDzxWc8=
github.com/dghubble/go-twitter v0.0.0-20211115160449-93a8679adecb h1:7ENzkH+O3juL+yj2undESLTaAeRllHwCs/b8z6aWSfc=
Expand Down Expand Up @@ -1397,8 +1397,8 @@ golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5y
golang.org/x/crypto v0.0.0-20220622213112-05595931fe9d/go.mod h1:IxCIyHEi3zRg3s0A5j5BB6A9Jmi73HwBIUl50j+osU4=
golang.org/x/crypto v0.0.0-20220722155217-630584e8d5aa/go.mod h1:IxCIyHEi3zRg3s0A5j5BB6A9Jmi73HwBIUl50j+osU4=
golang.org/x/crypto v0.17.0/go.mod h1:gCAAfMLgwOJRpTjQ2zCCt2OcSfYMTeZVSRtQlPC7Nq4=
golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw=
golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk=
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
golang.org/x/exp v0.0.0-20260603202125-055de637280b h1:v1uXiEBHo8QA0LiGCo7UgHMzHT4Kdfpl2zmtH5vaP1Q=
golang.org/x/exp v0.0.0-20260603202125-055de637280b/go.mod h1:d2fgXJLVs4dYDHUk5lwMIfzRzSrWCfGZb0ZqeLa/Vcw=
golang.org/x/exp/typeparams v0.0.0-20240314144324-c7f7c6466f7f h1:phY1HzDcf18Aq9A8KkmRtY9WvOFIxN8wgfvy6Zm1DV8=
Expand Down Expand Up @@ -1479,8 +1479,8 @@ golang.org/x/text v0.4.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8=
golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8=
golang.org/x/text v0.9.0/go.mod h1:e1OnstbJyHTd6l/uOt8jFFHp6TRDWZR/bV3emEE/zU8=
golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs=
golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY=
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U=
golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
Expand Down
4 changes: 1 addition & 3 deletions pkg/agent/tunnelserver/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -15,10 +15,8 @@ import (
func NewTunnelClient(
reader io.Reader,
writer io.WriteCloser,
exitOnClose bool,
exitCode int,
) (tunnel.TunnelClient, error) {
pipe := stdio.NewStdioStream(reader, writer, exitOnClose, exitCode)
pipe := stdio.NewStdioStream(reader, writer)

// After moving from deprecated grpc.Dial to grpc.NewClient we need to setup resolver first
// https://github.com/grpc/grpc-go/issues/1786#issuecomment-2119088770
Expand Down
2 changes: 1 addition & 1 deletion pkg/agent/tunnelserver/tunnelserver.go
Original file line number Diff line number Diff line change
Expand Up @@ -131,7 +131,7 @@ func (t *tunnelServer) RunWithResult(
reader io.Reader,
writer io.WriteCloser,
) (*config.Result, error) {
lis := stdio.NewStdioListener(reader, writer, false)
lis := stdio.NewStdioListener(reader, writer)
s := grpc.NewServer()
tunnel.RegisterTunnelServer(s, t)
reflection.Register(s)
Expand Down
2 changes: 1 addition & 1 deletion pkg/devcontainer/sshtunnel/sshtunnel.go
Original file line number Diff line number Diff line change
Expand Up @@ -136,7 +136,7 @@ func runSSHTunnel(ctx context.Context, p sshTunnelParams) (*config2.Result, erro
defer func() { log.Infof("tunnel: setup complete elapsed=%s", time.Since(start)) }()

log.Debug("creating SSH client")
sshClient, err := devssh.StdioClient(p.stdout, p.stdin, false)
sshClient, err := devssh.StdioClient(p.stdout, p.stdin)
if err != nil {
return nil, fmt.Errorf("failed to create SSH client: %w", err)
}
Expand Down
31 changes: 24 additions & 7 deletions pkg/ssh/helper.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,25 @@ import (
"golang.org/x/crypto/ssh"
)

const keepAliveRequestType = "keepalive@openssh.com"

func handleKeepAliveRequests(in <-chan *ssh.Request) <-chan *ssh.Request {
out := make(chan *ssh.Request)
go func() {
defer close(out)
for req := range in {
if req.Type == keepAliveRequestType {
if req.WantReply {
_ = req.Reply(true, nil)
}
continue
}
out <- req
}
}()
return out
}

func NewSSHPassClient(user, addr, password string) (*ssh.Client, error) {
clientConfig := &ssh.ClientConfig{
Auth: []ssh.AuthMethod{},
Expand Down Expand Up @@ -48,27 +67,25 @@ func NewSSHClient(user, addr string, keyBytes []byte) (*ssh.Client, error) {
return client, nil
}

func StdioClient(reader io.Reader, writer io.WriteCloser, exitOnClose bool) (*ssh.Client, error) {
return StdioClientFromKeyBytesWithUser(nil, reader, writer, "", exitOnClose)
func StdioClient(reader io.Reader, writer io.WriteCloser) (*ssh.Client, error) {
return StdioClientFromKeyBytesWithUser(nil, reader, writer, "")
}

func StdioClientWithUser(
reader io.Reader,
writer io.WriteCloser,
user string,
exitOnClose bool,
) (*ssh.Client, error) {
return StdioClientFromKeyBytesWithUser(nil, reader, writer, user, exitOnClose)
return StdioClientFromKeyBytesWithUser(nil, reader, writer, user)
}

func StdioClientFromKeyBytesWithUser(
keyBytes []byte,
reader io.Reader,
writer io.WriteCloser,
user string,
exitOnClose bool,
) (*ssh.Client, error) {
conn := stdio.NewStdioStream(reader, writer, exitOnClose, 0)
conn := stdio.NewStdioStream(reader, writer)
clientConfig, err := ConfigFromKeyBytes(keyBytes)
if err != nil {
return nil, err
Expand All @@ -80,7 +97,7 @@ func StdioClientFromKeyBytesWithUser(
return nil, err
}

return ssh.NewClient(c, chans, req), nil
return ssh.NewClient(c, chans, handleKeepAliveRequests(req)), nil
}

func ConfigFromKeyBytes(keyBytes []byte) (*ssh.ClientConfig, error) {
Expand Down
21 changes: 21 additions & 0 deletions pkg/ssh/keepalive_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,21 @@
package ssh

import (
"testing"

"github.com/stretchr/testify/assert"
gossh "golang.org/x/crypto/ssh"
)

func TestHandleKeepAliveRequests(t *testing.T) {
requests := make(chan *gossh.Request, 2)
requests <- &gossh.Request{Type: keepAliveRequestType}
other := &gossh.Request{Type: "other"}
requests <- other
close(requests)

forwarded := handleKeepAliveRequests(requests)
assert.Same(t, other, <-forwarded)
_, open := <-forwarded
assert.False(t, open)
}
3 changes: 3 additions & 0 deletions pkg/ssh/server/ssh.go
Original file line number Diff line number Diff line change
Expand Up @@ -192,6 +192,9 @@ func NewServer(

server.sshServer.Handler = server.handler
server.sshServer.ConnCallback = server.connCallback
server.sshServer.ConnectionCompleteCallback = func(conn *gossh.ServerConn, err error) {
log.Debugf("ssh transport completed remote=%v err=%v", conn.RemoteAddr(), err)
}
server.sshServer.ConnectionClosingCallback = cleanupAgentOnConnClosing
return server, nil
}
Expand Down
37 changes: 21 additions & 16 deletions pkg/stdio/conn.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@ package stdio
import (
"io"
"net"
"os"
"sync"
"time"
)

Expand All @@ -14,19 +14,23 @@ type StdioStream struct {
local *StdinAddr
remote *StdinAddr

exitOnClose bool
exitCode int
closeOnce sync.Once
closeErr error
onClose func()
}

// NewStdioStream is used to implement the connection interface.
func NewStdioStream(in io.Reader, out io.WriteCloser, exitOnClose bool, exitCode int) *StdioStream {
func NewStdioStream(in io.Reader, out io.WriteCloser) *StdioStream {
return newStdioStream(in, out, nil)
}

func newStdioStream(in io.Reader, out io.WriteCloser, onClose func()) *StdioStream {
return &StdioStream{
local: NewStdinAddr("local"),
remote: NewStdinAddr("remote"),
in: in,
out: out,
exitOnClose: exitOnClose,
exitCode: exitCode,
local: NewStdinAddr("local"),
remote: NewStdinAddr("remote"),
in: in,
out: out,
onClose: onClose,
}
}

Expand All @@ -52,12 +56,13 @@ func (s *StdioStream) Write(b []byte) (n int, err error) {

// Close implements interface.
func (s *StdioStream) Close() error {
if s.exitOnClose {
// We kill ourself here because the streams are closed
os.Exit(s.exitCode)
}

return s.out.Close()
s.closeOnce.Do(func() {
s.closeErr = s.out.Close()
if s.onClose != nil {
s.onClose()
}
})
return s.closeErr
}

// SetDeadline implements interface.
Expand Down
Loading
Loading