From 56259e3f05118f68d5f9131960583318672a7dc2 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 27 Oct 2022 17:10:36 -0700 Subject: [PATCH] fix cmd done lock ordering (actually start the cmdwait). implement reset command to re-initialize the terminal --- pkg/cmdrunner/cmdrunner.go | 33 +++++++++++++++++++++++++- pkg/cmdrunner/shparse.go | 26 ++++++++++----------- pkg/remote/remote.go | 47 +++++++++++++++++++++++++++++++++----- pkg/remote/updatequeue.go | 6 +++++ 4 files changed, 92 insertions(+), 20 deletions(-) diff --git a/pkg/cmdrunner/cmdrunner.go b/pkg/cmdrunner/cmdrunner.go index 3bed888e..c0a57c54 100644 --- a/pkg/cmdrunner/cmdrunner.go +++ b/pkg/cmdrunner/cmdrunner.go @@ -1442,7 +1442,35 @@ func SessionCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (ssto } func ResetCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { - return nil, nil + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_Window) + if err != nil { + return nil, err + } + initPk, err := ids.Remote.MShell.ReInit(ctx) + if err != nil { + return nil, err + } + if initPk == nil || initPk.State == nil { + return nil, fmt.Errorf("invalid initpk received from remote (no remote state)") + } + remoteInst, err := sstore.UpdateRemoteState(ctx, ids.SessionId, ids.WindowId, ids.Remote.RemotePtr, *initPk.State) + if err != nil { + return nil, err + } + outputStr := "reset remote state" + cmd, err := makeStaticCmd(ctx, "reset", ids, pk.GetRawStr(), []byte(outputStr)) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + update, err := addLineForCmd(ctx, "/cd", false, ids, cmd) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + update.Interactive = pk.Interactive + update.Sessions = sstore.MakeSessionsUpdateForRemote(ids.SessionId, remoteInst) + return update, nil } func ClearCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { @@ -1622,6 +1650,9 @@ func LineShowCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sst if cmd.TermOpts != cmd.OrigTermOpts { buf.WriteString(fmt.Sprintf(" %-15s %s\n", "orig-termopts", formatTermOpts(cmd.OrigTermOpts))) } + if cmd.RtnState { + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "rtnstate", "true")) + } } update := sstore.ModelUpdate{ Info: &sstore.InfoMsgType{ diff --git a/pkg/cmdrunner/shparse.go b/pkg/cmdrunner/shparse.go index 898db34f..5a5ea021 100644 --- a/pkg/cmdrunner/shparse.go +++ b/pkg/cmdrunner/shparse.go @@ -13,6 +13,19 @@ import ( "mvdan.cc/sh/v3/syntax" ) +var ValidMetaCmdRe = regexp.MustCompile("^/([a-z][a-z0-9_-]*)(?::([a-z][a-z0-9_-]*))?$") + +type BareMetaCmdDecl struct { + CmdStr string + MetaCmd string +} + +var BareMetaCmds = []BareMetaCmdDecl{ + BareMetaCmdDecl{"cr", "cr"}, + BareMetaCmdDecl{"clear", "clear"}, + BareMetaCmdDecl{"reset", "reset"}, +} + func DumpPacket(pk *scpacket.FeCommandPacketType) { if pk == nil || pk.MetaCmd == "" { fmt.Printf("[no metacmd]\n") @@ -51,19 +64,6 @@ func getSourceStr(source string, w *syntax.Word) string { return source[offset:end] } -var ValidMetaCmdRe = regexp.MustCompile("^/([a-z][a-z0-9_-]*)(?::([a-z][a-z0-9_-]*))?$") - -type BareMetaCmdDecl struct { - CmdStr string - MetaCmd string -} - -var BareMetaCmds = []BareMetaCmdDecl{ - BareMetaCmdDecl{"cr", "cr"}, - BareMetaCmdDecl{"clear", "clear"}, - BareMetaCmdDecl{"reset", "reset"}, -} - func SubMetaCmd(cmd string) string { switch cmd { case "s": diff --git a/pkg/remote/remote.go b/pkg/remote/remote.go index 7b76ebb9..8e8e9f56 100644 --- a/pkg/remote/remote.go +++ b/pkg/remote/remote.go @@ -875,6 +875,26 @@ func (msh *MShellProc) RunInstall() { return } +func (msh *MShellProc) ReInit(ctx context.Context) (*packet.InitPacketType, error) { + reinitPk := packet.MakeReInitPacket() + reinitPk.ReqId = uuid.New().String() + resp, err := msh.PacketRpcRaw(ctx, reinitPk) + if err != nil { + return nil, err + } + if resp == nil { + return nil, fmt.Errorf("no response") + } + initPk, ok := resp.(*packet.InitPacketType) + if !ok { + return nil, fmt.Errorf("invalid reinit response (not an initpacket): %T", resp) + } + msh.WithLock(func() { + msh.Remote.InitPk = initPk + }) + return initPk, nil +} + func (msh *MShellProc) Launch() { remoteCopy := msh.GetRemoteCopy() if remoteCopy.Archived { @@ -1057,7 +1077,7 @@ func makeTermOpts(runPk *packet.RunPacketType) sstore.TermOpts { } // returns (cmdtype, allow-updates-callback, err) -func RunCommand(ctx context.Context, cmdId string, remotePtr sstore.RemotePtrType, remoteState *packet.ShellState, runPacket *packet.RunPacketType) (*sstore.CmdType, func(), error) { +func RunCommand(ctx context.Context, cmdId string, remotePtr sstore.RemotePtrType, remoteState *packet.ShellState, runPacket *packet.RunPacketType) (rtnCmd *sstore.CmdType, rtnCallback func(), rtnErr error) { if remotePtr.OwnerId != "" { return nil, nil, fmt.Errorf("cannot run command against another user's remote '%s'", remotePtr.MakeFullRemoteRef()) } @@ -1071,6 +1091,15 @@ func RunCommand(ctx context.Context, cmdId string, remotePtr sstore.RemotePtrTyp if remoteState == nil { return nil, nil, fmt.Errorf("no remote state passed to RunCommand") } + callbackFn := func() { + removeCmdWait(runPacket.CK) + } + startCmdWait(runPacket.CK) + defer func() { + if rtnErr != nil { + callbackFn() + } + }() msh.ServerProc.Output.RegisterRpc(runPacket.ReqId) err := shexec.SendRunPacketAndRunData(ctx, msh.ServerProc.Input, runPacket) if err != nil { @@ -1114,7 +1143,7 @@ func RunCommand(ctx context.Context, cmdId string, remotePtr sstore.RemotePtrTyp return nil, nil, fmt.Errorf("cannot create local ptyout file for running command: %v", err) } msh.AddRunningCmd(startPk.CK) - return cmd, func() { removeCmdWait(startPk.CK) }, nil + return cmd, callbackFn, nil } func (msh *MShellProc) AddRunningCmd(ck base.CommandKey) { @@ -1129,7 +1158,7 @@ func (msh *MShellProc) RemoveRunningCmd(ck base.CommandKey) { delete(msh.RunningCmds, ck) } -func (msh *MShellProc) PacketRpc(ctx context.Context, pk packet.RpcPacketType) (*packet.ResponsePacketType, error) { +func (msh *MShellProc) PacketRpcRaw(ctx context.Context, pk packet.RpcPacketType) (packet.RpcResponsePacketType, error) { if !msh.IsConnected() { return nil, fmt.Errorf("runner is not connected") } @@ -1147,6 +1176,14 @@ func (msh *MShellProc) PacketRpc(ctx context.Context, pk packet.RpcPacketType) ( if rtnPk == nil { return nil, ctx.Err() } + return rtnPk, nil +} + +func (msh *MShellProc) PacketRpc(ctx context.Context, pk packet.RpcPacketType) (*packet.ResponsePacketType, error) { + rtnPk, err := msh.PacketRpcRaw(ctx, pk) + if err != nil { + return nil, err + } if respPk, ok := rtnPk.(*packet.ResponsePacketType); ok { return respPk, nil } @@ -1195,9 +1232,7 @@ func (msh *MShellProc) handleCmdDonePacket(donePk *packet.CmdDonePacketType) { // fall-through (nothing to do) } update.ScreenWindows = sws - if update != nil { - sstore.MainBus.SendUpdate(donePk.CK.GetSessionId(), update) - } + sstore.MainBus.SendUpdate(donePk.CK.GetSessionId(), update) if donePk.FinalState != nil { } diff --git a/pkg/remote/updatequeue.go b/pkg/remote/updatequeue.go index 4665511f..7b978e38 100644 --- a/pkg/remote/updatequeue.go +++ b/pkg/remote/updatequeue.go @@ -4,6 +4,12 @@ import ( "github.com/scripthaus-dev/mshell/pkg/base" ) +func startCmdWait(ck base.CommandKey) { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + GlobalStore.CmdWaitMap[ck] = nil +} + func pushCmdWaitIfRequired(ck base.CommandKey, fn func()) bool { GlobalStore.Lock.Lock() defer GlobalStore.Lock.Unlock()