mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat(cli): implement ssh remote forward (#8515)
This commit is contained in:
+41
-57
@@ -6,7 +6,6 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/url"
|
||||
"os"
|
||||
"os/exec"
|
||||
@@ -27,7 +26,6 @@ import (
|
||||
"cdr.dev/slog"
|
||||
"cdr.dev/slog/sloggers/sloghuman"
|
||||
|
||||
"github.com/coder/coder/agent/agentssh"
|
||||
"github.com/coder/coder/cli/clibase"
|
||||
"github.com/coder/coder/cli/cliui"
|
||||
"github.com/coder/coder/coderd/autobuild/notify"
|
||||
@@ -53,6 +51,7 @@ func (r *RootCmd) ssh() *clibase.Cmd {
|
||||
waitEnum string
|
||||
noWait bool
|
||||
logDirPath string
|
||||
remoteForward string
|
||||
)
|
||||
client := new(codersdk.Client)
|
||||
cmd := &clibase.Cmd{
|
||||
@@ -122,6 +121,16 @@ func (r *RootCmd) ssh() *clibase.Cmd {
|
||||
client.SetLogger(logger)
|
||||
}
|
||||
|
||||
if remoteForward != "" {
|
||||
isValid := validateRemoteForward(remoteForward)
|
||||
if !isValid {
|
||||
return xerrors.Errorf(`invalid format of remote-forward, expected: remote_port:local_address:local_port`)
|
||||
}
|
||||
if isValid && stdio {
|
||||
return xerrors.Errorf(`remote-forward can't be enabled in the stdio mode`)
|
||||
}
|
||||
}
|
||||
|
||||
workspace, workspaceAgent, err := getWorkspaceAndAgent(ctx, inv, client, codersdk.Me, inv.Args[0])
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -198,6 +207,7 @@ func (r *RootCmd) ssh() *clibase.Cmd {
|
||||
}
|
||||
defer conn.Close()
|
||||
conn.AwaitReachable(ctx)
|
||||
|
||||
stopPolling := tryPollWorkspaceAutostop(ctx, client, workspace)
|
||||
defer stopPolling()
|
||||
|
||||
@@ -300,6 +310,19 @@ func (r *RootCmd) ssh() *clibase.Cmd {
|
||||
defer closer.Close()
|
||||
}
|
||||
|
||||
if remoteForward != "" {
|
||||
localAddr, remoteAddr, err := parseRemoteForward(remoteForward)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
closer, err := sshRemoteForward(ctx, inv.Stderr, sshClient, localAddr, remoteAddr)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("ssh remote forward: %w", err)
|
||||
}
|
||||
defer closer.Close()
|
||||
}
|
||||
|
||||
stdoutFile, validOut := inv.Stdout.(*os.File)
|
||||
stdinFile, validIn := inv.Stdin.(*os.File)
|
||||
if validOut && validIn && isatty.IsTerminal(stdoutFile.Fd()) {
|
||||
@@ -424,6 +447,13 @@ func (r *RootCmd) ssh() *clibase.Cmd {
|
||||
FlagShorthand: "l",
|
||||
Value: clibase.StringOf(&logDirPath),
|
||||
},
|
||||
{
|
||||
Flag: "remote-forward",
|
||||
Description: "Enable remote port forwarding (remote_port:local_address:local_port).",
|
||||
Env: "CODER_SSH_REMOTE_FORWARD",
|
||||
FlagShorthand: "R",
|
||||
Value: clibase.StringOf(&remoteForward),
|
||||
},
|
||||
}
|
||||
return cmd
|
||||
}
|
||||
@@ -568,8 +598,15 @@ func getWorkspaceAndAgent(ctx context.Context, inv *clibase.Invocation, client *
|
||||
// of the CLI running simultaneously.
|
||||
func tryPollWorkspaceAutostop(ctx context.Context, client *codersdk.Client, workspace codersdk.Workspace) (stop func()) {
|
||||
lock := flock.New(filepath.Join(os.TempDir(), "coder-autostop-notify-"+workspace.ID.String()))
|
||||
condition := notifyCondition(ctx, client, workspace.ID, lock)
|
||||
return notify.Notify(condition, workspacePollInterval, autostopNotifyCountdown...)
|
||||
conditionCtx, cancelCondition := context.WithCancel(ctx)
|
||||
condition := notifyCondition(conditionCtx, client, workspace.ID, lock)
|
||||
stopFunc := notify.Notify(condition, workspacePollInterval, autostopNotifyCountdown...)
|
||||
return func() {
|
||||
// With many "ssh" processes running, `lock.TryLockContext` can be hanging until the context canceled.
|
||||
// Without this cancellation, a CLI process with failed remote-forward could be hanging indefinitely.
|
||||
cancelCondition()
|
||||
stopFunc()
|
||||
}
|
||||
}
|
||||
|
||||
// Notify the user if the workspace is due to shutdown.
|
||||
@@ -752,56 +789,3 @@ func remoteGPGAgentSocket(sshClient *gossh.Client) (string, error) {
|
||||
|
||||
return string(bytes.TrimSpace(remoteSocket)), nil
|
||||
}
|
||||
|
||||
// cookieAddr is a special net.Addr accepted by sshForward() which includes a
|
||||
// cookie which is written to the connection before forwarding.
|
||||
type cookieAddr struct {
|
||||
net.Addr
|
||||
cookie []byte
|
||||
}
|
||||
|
||||
// sshForwardRemote starts forwarding connections from a remote listener to a
|
||||
// local address via SSH in a goroutine.
|
||||
//
|
||||
// Accepts a `cookieAddr` as the local address.
|
||||
func sshForwardRemote(ctx context.Context, stderr io.Writer, sshClient *gossh.Client, localAddr, remoteAddr net.Addr) (io.Closer, error) {
|
||||
listener, err := sshClient.Listen(remoteAddr.Network(), remoteAddr.String())
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("listen on remote SSH address %s: %w", remoteAddr.String(), err)
|
||||
}
|
||||
|
||||
go func() {
|
||||
for {
|
||||
remoteConn, err := listener.Accept()
|
||||
if err != nil {
|
||||
if ctx.Err() == nil {
|
||||
_, _ = fmt.Fprintf(stderr, "Accept SSH listener connection: %+v\n", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
go func() {
|
||||
defer remoteConn.Close()
|
||||
|
||||
localConn, err := net.Dial(localAddr.Network(), localAddr.String())
|
||||
if err != nil {
|
||||
_, _ = fmt.Fprintf(stderr, "Dial local address %s: %+v\n", localAddr.String(), err)
|
||||
return
|
||||
}
|
||||
defer localConn.Close()
|
||||
|
||||
if c, ok := localAddr.(cookieAddr); ok {
|
||||
_, err = localConn.Write(c.cookie)
|
||||
if err != nil {
|
||||
_, _ = fmt.Fprintf(stderr, "Write cookie to local connection: %+v\n", err)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
agentssh.Bicopy(ctx, localConn, remoteConn)
|
||||
}()
|
||||
}
|
||||
}()
|
||||
|
||||
return listener, nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user