diff --git a/pkg/cmdrunner/cmdrunner.go b/pkg/cmdrunner/cmdrunner.go index 0caf8d58..2c2fd4bc 100644 --- a/pkg/cmdrunner/cmdrunner.go +++ b/pkg/cmdrunner/cmdrunner.go @@ -421,6 +421,9 @@ func RemoteConnectCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) if ids.Remote.RState.IsConnected() { return sstore.InfoMsgUpdate("remote %q already connected (no action taken)", ids.Remote.DisplayName), nil } + if ids.Remote.RState.Status == remote.StatusConnecting { + return sstore.InfoMsgUpdate("remote %q is already trying to connect (no action taken)", ids.Remote.DisplayName), nil + } go ids.Remote.MShell.Launch() return sstore.InfoMsgUpdate("remote %q reconnecting", ids.Remote.DisplayName), nil } @@ -431,7 +434,8 @@ func RemoteDisconnectCommand(ctx context.Context, pk *scpacket.FeCommandPacketTy return nil, err } force := resolveBool(pk.Kwargs["force"], false) - if !ids.Remote.RState.IsConnected() && !force { + status := ids.Remote.MShell.GetStatus() + if status != remote.StatusConnected && status != remote.StatusConnecting { return sstore.InfoMsgUpdate("remote %q already disconnected (no action taken)", ids.Remote.DisplayName), nil } numCommands := ids.Remote.MShell.GetNumRunningCommands() diff --git a/pkg/remote/remote.go b/pkg/remote/remote.go index 3cee7ad3..2bb9b0a8 100644 --- a/pkg/remote/remote.go +++ b/pkg/remote/remote.go @@ -66,12 +66,13 @@ type MShellProc struct { Remote *sstore.RemoteType // runtime - Status string - ServerProc *shexec.ClientProc - UName string - Err error - ControllingPty *os.File - PtyBuffer *circbuf.Buffer + Status string + ServerProc *shexec.ClientProc + UName string + Err error + ControllingPty *os.File + PtyBuffer *circbuf.Buffer + MakeClientCancelFn context.CancelFunc RunningCmds []base.CommandKey } @@ -95,6 +96,12 @@ func (state RemoteRuntimeState) IsConnected() bool { return state.Status == StatusConnected } +func (msh *MShellProc) GetStatus() string { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return msh.Status +} + func (state RemoteRuntimeState) GetBaseDisplayName() string { if state.RemoteAlias != "" { return state.RemoteAlias @@ -522,6 +529,10 @@ func (msh *MShellProc) Disconnect() { if msh.ServerProc != nil { msh.ServerProc.Close() } + if msh.MakeClientCancelFn != nil { + msh.MakeClientCancelFn() + msh.MakeClientCancelFn = nil + } } func (msh *MShellProc) GetRemoteName() string { @@ -598,6 +609,11 @@ func (msh *MShellProc) Launch() { msh.WriteToPtyBuffer("cannot launch archived remote\n") return } + curStatus := msh.GetStatus() + if curStatus == StatusConnecting { + msh.WriteToPtyBuffer("remote is already connecting, disconnect before trying to connect again\n") + return + } msh.WriteToPtyBuffer("connecting to %s...\n", remoteCopy.RemoteCanonicalName) sshOpts := convertSSHOpts(remoteCopy.SSHOpts) sshOpts.SSHErrorsToTty = true @@ -615,15 +631,22 @@ func (msh *MShellProc) Launch() { } }() go msh.RunPtyReadLoop(cmdPty) + makeClientCtx, makeClientCancelFn := context.WithCancel(context.Background()) + defer makeClientCancelFn() msh.WithLock(func() { msh.Status = StatusConnecting + msh.MakeClientCancelFn = makeClientCancelFn go msh.NotifyRemoteUpdate() }) - cproc, uname, err := shexec.MakeClientProc(ecmd) + cproc, uname, err := shexec.MakeClientProc(makeClientCtx, ecmd) msh.WithLock(func() { msh.UName = uname + msh.MakeClientCancelFn = nil // no notify here, because we'll call notify in either case below }) + if err == context.Canceled { + err = fmt.Errorf("forced disconnection") + } if err != nil { msh.setErrorStatus(err) msh.WriteToPtyBuffer("*error connecting to remote (uname=%q): %v\n", msh.UName, err)