diff --git a/cmd/internal/agentcontainer/credentials_server.go b/cmd/internal/agentcontainer/credentials_server.go index b69536114..8184c2151 100644 --- a/cmd/internal/agentcontainer/credentials_server.go +++ b/cmd/internal/agentcontainer/credentials_server.go @@ -27,8 +27,6 @@ import ( "github.com/spf13/cobra" ) -const ExitCodeIO int = 64 - // CredentialsServerCmd holds the cmd flags. type CredentialsServerCmd struct { *flags.GlobalFlags @@ -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) } diff --git a/cmd/internal/agentcontainer/setup.go b/cmd/internal/agentcontainer/setup.go index 22eb15fbb..7ef936283 100644 --- a/cmd/internal/agentcontainer/setup.go +++ b/cmd/internal/agentcontainer/setup.go @@ -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) } diff --git a/cmd/internal/agentworkspace/up.go b/cmd/internal/agentworkspace/up.go index d368978e4..aacde7186 100644 --- a/cmd/internal/agentworkspace/up.go +++ b/cmd/internal/agentworkspace/up.go @@ -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) } diff --git a/cmd/internal/ssh_server.go b/cmd/internal/ssh_server.go index e58278e3b..174fe3e38 100644 --- a/cmd/internal/ssh_server.go +++ b/cmd/internal/ssh_server.go @@ -5,6 +5,7 @@ import ( "encoding/base64" "errors" "fmt" + "net" "os" "time" @@ -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)) } @@ -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 diff --git a/cmd/internal/ssh_server_test.go b/cmd/internal/ssh_server_test.go index 0dd4b5a94..12f7e165d 100644 --- a/cmd/internal/ssh_server_test.go +++ b/cmd/internal/ssh_server_test.go @@ -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) diff --git a/cmd/machine/ssh.go b/cmd/machine/ssh.go index b9d927aad..4fb8ccfe6 100644 --- a/cmd/machine/ssh.go +++ b/cmd/machine/ssh.go @@ -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 } diff --git a/cmd/snapshot/create.go b/cmd/snapshot/create.go index 833cf1bc7..240b74a77 100644 --- a/cmd/snapshot/create.go +++ b/cmd/snapshot/create.go @@ -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() diff --git a/cmd/workspace/logs.go b/cmd/workspace/logs.go index e5fd3c658..6943e2102 100644 --- a/cmd/workspace/logs.go +++ b/cmd/workspace/logs.go @@ -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 } diff --git a/go.mod b/go.mod index 1c7a7b2e5..d483dcf09 100644 --- a/go.mod +++ b/go.mod @@ -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 @@ -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 @@ -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 diff --git a/go.sum b/go.sum index 0d4ad052f..58343d83e 100644 --- a/go.sum +++ b/go.sum @@ -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= @@ -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= @@ -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= diff --git a/pkg/agent/tunnelserver/client.go b/pkg/agent/tunnelserver/client.go index a8213d665..9a2549b6c 100644 --- a/pkg/agent/tunnelserver/client.go +++ b/pkg/agent/tunnelserver/client.go @@ -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 diff --git a/pkg/agent/tunnelserver/tunnelserver.go b/pkg/agent/tunnelserver/tunnelserver.go index 7386280bd..b31826878 100644 --- a/pkg/agent/tunnelserver/tunnelserver.go +++ b/pkg/agent/tunnelserver/tunnelserver.go @@ -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) diff --git a/pkg/devcontainer/sshtunnel/sshtunnel.go b/pkg/devcontainer/sshtunnel/sshtunnel.go index 21b0af4c8..32e37a2b9 100644 --- a/pkg/devcontainer/sshtunnel/sshtunnel.go +++ b/pkg/devcontainer/sshtunnel/sshtunnel.go @@ -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) } diff --git a/pkg/ssh/helper.go b/pkg/ssh/helper.go index 3c6de53d9..568c7de60 100644 --- a/pkg/ssh/helper.go +++ b/pkg/ssh/helper.go @@ -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{}, @@ -48,17 +67,16 @@ 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( @@ -66,9 +84,8 @@ func StdioClientFromKeyBytesWithUser( 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 @@ -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) { diff --git a/pkg/ssh/keepalive_test.go b/pkg/ssh/keepalive_test.go new file mode 100644 index 000000000..e015903b6 --- /dev/null +++ b/pkg/ssh/keepalive_test.go @@ -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) +} diff --git a/pkg/ssh/server/ssh.go b/pkg/ssh/server/ssh.go index ee26906ee..4d248869f 100644 --- a/pkg/ssh/server/ssh.go +++ b/pkg/ssh/server/ssh.go @@ -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 } diff --git a/pkg/stdio/conn.go b/pkg/stdio/conn.go index 38c48a073..832962b5d 100644 --- a/pkg/stdio/conn.go +++ b/pkg/stdio/conn.go @@ -3,7 +3,7 @@ package stdio import ( "io" "net" - "os" + "sync" "time" ) @@ -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, } } @@ -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. diff --git a/pkg/stdio/conn_test.go b/pkg/stdio/conn_test.go new file mode 100644 index 000000000..aa481904c --- /dev/null +++ b/pkg/stdio/conn_test.go @@ -0,0 +1,62 @@ +package stdio + +import ( + "io" + "strings" + "sync/atomic" + "testing" +) + +type trackingWriteCloser struct { + closeCalls atomic.Int32 + closeErr error +} + +func (w *trackingWriteCloser) Write(p []byte) (int, error) { + return len(p), nil +} + +func (w *trackingWriteCloser) Close() error { + w.closeCalls.Add(1) + return w.closeErr +} + +func TestStdioStreamCloseDoesNotExitAndIsIdempotent(t *testing.T) { + writer := &trackingWriteCloser{} + var callbackCalls atomic.Int32 + stream := newStdioStream(strings.NewReader(""), writer, func() { + callbackCalls.Add(1) + }) + + if err := stream.Close(); err != nil { + t.Fatalf("first Close() error = %v", err) + } + if err := stream.Close(); err != nil { + t.Fatalf("second Close() error = %v", err) + } + + if got := writer.closeCalls.Load(); got != 1 { + t.Fatalf("writer Close calls = %d, want 1", got) + } + if got := callbackCalls.Load(); got != 1 { + t.Fatalf("close callback calls = %d, want 1", got) + } + + // Reaching this assertion proves Close returned control to the process. + if _, err := io.ReadAll(stream); err != nil { + t.Fatalf("read after Close() error = %v", err) + } +} + +func TestStdioStreamCloseReturnsUnderlyingErrorOnce(t *testing.T) { + wantErr := io.ErrClosedPipe + writer := &trackingWriteCloser{closeErr: wantErr} + stream := NewStdioStream(strings.NewReader(""), writer) + + if err := stream.Close(); err != wantErr { + t.Fatalf("first Close() error = %v, want %v", err, wantErr) + } + if err := stream.Close(); err != wantErr { + t.Fatalf("second Close() error = %v, want %v", err, wantErr) + } +} diff --git a/pkg/stdio/listener.go b/pkg/stdio/listener.go index bfa1d4963..6b9542471 100644 --- a/pkg/stdio/listener.go +++ b/pkg/stdio/listener.go @@ -3,24 +3,27 @@ package stdio import ( "io" "net" + "sync" ) -// StdioListener implements the listener interface. +// StdioListener implements the listener interface for one stdio connection. type StdioListener struct { - connChan chan net.Conn -} + conn net.Conn -// NewStdioListener creates a new stdio listener. -func NewStdioListener(reader io.Reader, writer io.WriteCloser, exitOnClose bool) *StdioListener { - conn := NewStdioStream(reader, writer, exitOnClose, 0) - connChan := make(chan net.Conn) - go func() { - connChan <- conn - }() + mu sync.Mutex + accepted bool + closed bool + closedCh chan struct{} + once sync.Once +} - return &StdioListener{ - connChan: connChan, +// NewStdioListener creates a new one-shot stdio listener. +func NewStdioListener(reader io.Reader, writer io.WriteCloser) *StdioListener { + lis := &StdioListener{ + closedCh: make(chan struct{}), } + lis.conn = newStdioStream(reader, writer, lis.markClosed) + return lis } // Ready implements interface. @@ -29,11 +32,35 @@ func (lis *StdioListener) Ready(conn net.Conn) { // Accept implements interface. func (lis *StdioListener) Accept() (net.Conn, error) { - return <-lis.connChan, nil + lis.mu.Lock() + if lis.closed { + lis.mu.Unlock() + return nil, net.ErrClosed + } + if !lis.accepted { + lis.accepted = true + conn := lis.conn + lis.mu.Unlock() + return conn, nil + } + lis.mu.Unlock() + + <-lis.closedCh + return nil, net.ErrClosed } -// Close implements interface. +// Close closes the listener and its stdio connection. func (lis *StdioListener) Close() error { + lis.once.Do(func() { + lis.mu.Lock() + lis.closed = true + close(lis.closedCh) + lis.mu.Unlock() + }) + + if lis.conn != nil { + _ = lis.conn.Close() + } return nil } @@ -41,3 +68,13 @@ func (lis *StdioListener) Close() error { func (lis *StdioListener) Addr() net.Addr { return NewStdinAddr("listener") } + +// markClosed transitions the listener to the closed state. +func (lis *StdioListener) markClosed() { + lis.once.Do(func() { + lis.mu.Lock() + lis.closed = true + close(lis.closedCh) + lis.mu.Unlock() + }) +} diff --git a/pkg/stdio/listener_test.go b/pkg/stdio/listener_test.go new file mode 100644 index 000000000..252cee550 --- /dev/null +++ b/pkg/stdio/listener_test.go @@ -0,0 +1,118 @@ +package stdio + +import ( + "errors" + "net" + "strings" + "testing" + "time" +) + +func acceptWithTimeout(t *testing.T, listener net.Listener) (net.Conn, error) { + t.Helper() + result := make(chan struct { + conn net.Conn + err error + }, 1) + go func() { + conn, err := listener.Accept() + result <- struct { + conn net.Conn + err error + }{conn, err} + }() + select { + case result := <-result: + return result.conn, result.err + case <-time.After(2 * time.Second): + t.Fatal("listener Accept timed out") + return nil, nil + } +} + +func TestStdioListenerAcceptsOneConnection(t *testing.T) { + listener := NewStdioListener(strings.NewReader(""), &trackingWriteCloser{}) + conn, err := acceptWithTimeout(t, listener) + if err != nil { + t.Fatalf("first Accept() error = %v", err) + } + if conn == nil { + t.Fatal("first Accept() returned nil connection") + } + defer func() { _ = listener.Close() }() + + secondDone := make(chan error, 1) + go func() { + _, err := listener.Accept() + secondDone <- err + }() + if err := conn.Close(); err != nil { + t.Fatalf("connection Close() error = %v", err) + } + + select { + case err := <-secondDone: + if !errors.Is(err, net.ErrClosed) { + t.Fatalf("second Accept() error = %v, want net.ErrClosed", err) + } + case <-time.After(2 * time.Second): + t.Fatal("second Accept() remained blocked after connection close") + } +} + +func TestStdioListenerCloseUnblocksAccept(t *testing.T) { + listener := NewStdioListener(strings.NewReader(""), &trackingWriteCloser{}) + if _, err := acceptWithTimeout(t, listener); err != nil { + t.Fatalf("first Accept() error = %v", err) + } + + secondDone := make(chan error, 1) + go func() { + _, err := listener.Accept() + secondDone <- err + }() + if err := listener.Close(); err != nil { + t.Fatalf("listener Close() error = %v", err) + } + + select { + case err := <-secondDone: + if !errors.Is(err, net.ErrClosed) { + t.Fatalf("second Accept() error = %v, want net.ErrClosed", err) + } + case <-time.After(2 * time.Second): + t.Fatal("second Accept() remained blocked after listener close") + } +} + +func TestStdioListenerClosedBeforeAccept(t *testing.T) { + listener := NewStdioListener(strings.NewReader(""), &trackingWriteCloser{}) + if err := listener.Close(); err != nil { + t.Fatalf("listener Close() error = %v", err) + } + + if _, err := listener.Accept(); !errors.Is(err, net.ErrClosed) { + t.Fatalf("Accept() error = %v, want net.ErrClosed", err) + } +} + +func TestStdioListenerRepeatedCloseOrders(t *testing.T) { + listener := NewStdioListener(strings.NewReader(""), &trackingWriteCloser{}) + conn, err := acceptWithTimeout(t, listener) + if err != nil { + t.Fatalf("first Accept() error = %v", err) + } + + if err := conn.Close(); err != nil { + t.Fatalf("connection Close() error = %v", err) + } + if err := listener.Close(); err != nil { + t.Fatalf("listener Close() after connection close error = %v", err) + } + if err := listener.Close(); err != nil { + t.Fatalf("repeated listener Close() error = %v", err) + } + if err := conn.Close(); err != nil { + t.Fatalf("repeated connection Close() error = %v", err) + } +} diff --git a/pkg/tunnel/container.go b/pkg/tunnel/container.go index db351efcf..a0ca3c437 100644 --- a/pkg/tunnel/container.go +++ b/pkg/tunnel/container.go @@ -64,7 +64,7 @@ func (c *ContainerTunnel) Run( return c.runHostTunnel(ctx, stdin, stdout, timeout) }, func(ctx context.Context, stdout, stdin *os.File) error { - sshClient, err := devssh.StdioClient(stdout, stdin, false) + sshClient, err := devssh.StdioClient(stdout, stdin) if err != nil { return fmt.Errorf("create ssh client: %w", err) } @@ -202,7 +202,7 @@ func (c *ContainerTunnel) runInContainer( }) }() - containerClient, err := devssh.StdioClient(pb.StdoutReader, pb.StdinWriter, false) + containerClient, err := devssh.StdioClient(pb.StdoutReader, pb.StdinWriter) if err != nil { select { // check if the tunnel goroutine has already returned an error case tunnelErr := <-tunnelDone: diff --git a/pkg/tunnel/direct.go b/pkg/tunnel/direct.go index 9ad60546a..e0c8091a5 100644 --- a/pkg/tunnel/direct.go +++ b/pkg/tunnel/direct.go @@ -28,7 +28,7 @@ func NewTunnel(ctx context.Context, tunnel Tunnel, handler Handler) error { return tunnel(ctx, stdin, stdout) }, func(ctx context.Context, stdout, stdin *os.File) error { - sshClient, err := devssh.StdioClient(stdout, stdin, false) + sshClient, err := devssh.StdioClient(stdout, stdin) if err != nil { return err } diff --git a/pkg/tunnel/pipebridge.go b/pkg/tunnel/pipebridge.go index de5a52692..e0b501edf 100644 --- a/pkg/tunnel/pipebridge.go +++ b/pkg/tunnel/pipebridge.go @@ -6,6 +6,7 @@ import ( "sync" "time" + "github.com/devsy-org/devsy/pkg/log" "github.com/devsy-org/devsy/pkg/util/iojoin" ) @@ -102,5 +103,12 @@ func (pb *PipeBridge) RunPair( // Run the two sides concurrently and wait for both to finish (or the slower // side to be abandoned after joinTimeout). tunnelErr, handlerErr := iojoin.Join(tunnelSide, handlerSide, joinTimeout, stop) + log.Debugf( + "pipe bridge completed: tunnel_err=%v handler_err=%v parent_err=%v pair_err=%v", + tunnelErr, + handlerErr, + ctx.Err(), + pairCtx.Err(), + ) return ClassifyTunnelErrors(tunnelErr, handlerErr) }