From 353605f815916cf5ff1fd51c649101218f2b8305 Mon Sep 17 00:00:00 2001 From: sawka Date: Wed, 6 Jul 2022 18:59:46 -0700 Subject: [PATCH] bug fixes and updates for running server with scripthaus --- pkg/server/server.go | 14 +++++++++----- pkg/shexec/client.go | 18 ++++++------------ pkg/shexec/shexec.go | 32 ++++++++++++++++++++++++++------ 3 files changed, 41 insertions(+), 23 deletions(-) diff --git a/pkg/server/server.go b/pkg/server/server.go index 91637778..40ab29cc 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -7,6 +7,7 @@ package server import ( + "context" "fmt" "os" "sync" @@ -52,12 +53,16 @@ func (m *MServer) runCommand(runPacket *packet.RunPacketType) { m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("server run packets require valid ck: %s", err)) return } - cproc, err := shexec.MakeClientProc(runPacket.CK) + ecmd, err := shexec.SSHOpts{}.MakeMShellSingleCmd() + if err != nil { + m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("server run packets require valid ck: %s", err)) + return + } + cproc, err := shexec.MakeClientProc(ecmd) if err != nil { m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("starting mshell client: %s", err)) return } - fmt.Printf("client start: %v\n", runPacket.CK) m.Lock.Lock() m.ClientMap[runPacket.CK] = cproc m.Lock.Unlock() @@ -67,10 +72,9 @@ func (m *MServer) runCommand(runPacket *packet.RunPacketType) { delete(m.ClientMap, runPacket.CK) m.Lock.Unlock() cproc.Close() - fmt.Printf("client done: %v\n", runPacket.CK) }() - shexec.SendRunPacketAndRunData(cproc.Input, runPacket) - cproc.ProxyOutput(m.Sender) + shexec.SendRunPacketAndRunData(context.Background(), cproc.Input, runPacket) + cproc.ProxySingleOutput(runPacket.CK, m.Sender) }() } diff --git a/pkg/shexec/client.go b/pkg/shexec/client.go index f322ed38..f95b252f 100644 --- a/pkg/shexec/client.go +++ b/pkg/shexec/client.go @@ -12,7 +12,7 @@ import ( type ClientProc struct { Cmd *exec.Cmd - CK base.CommandKey + InitPk *packet.InitPacketType StartTs time.Time StdinWriter io.WriteCloser StdoutReader io.ReadCloser @@ -21,11 +21,7 @@ type ClientProc struct { Output *packet.PacketParser } -func MakeClientProc(ck base.CommandKey) (*ClientProc, error) { - ecmd, err := SSHOpts{}.MakeMShellSingleCmd() - if err != nil { - return nil, err - } +func MakeClientProc(ecmd *exec.Cmd) (*ClientProc, error) { inputWriter, err := ecmd.StdinPipe() if err != nil { return nil, fmt.Errorf("creating stdin pipe: %v", err) @@ -49,7 +45,6 @@ func MakeClientProc(ck base.CommandKey) (*ClientProc, error) { packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser) cproc := &ClientProc{ Cmd: ecmd, - CK: ck, StartTs: startTs, StdinWriter: inputWriter, StdoutReader: stdoutReader, @@ -57,7 +52,6 @@ func MakeClientProc(ck base.CommandKey) (*ClientProc, error) { Input: sender, Output: packetParser, } - versionOk := false for pk := range packetParser.MainCh { if pk.GetType() != packet.InitPacketStr { cproc.Close() @@ -72,10 +66,10 @@ func MakeClientProc(ck base.CommandKey) (*ClientProc, error) { cproc.Close() return nil, fmt.Errorf("invalid remote mshell version 'v%s', must be v%s", initPk.Version, base.MShellVersion) } - versionOk = true + cproc.InitPk = initPk break } - if !versionOk { + if cproc.InitPk == nil { cproc.Close() return nil, fmt.Errorf("no init packet received from mshell client") } @@ -100,7 +94,7 @@ func (cproc *ClientProc) Close() { } } -func (cproc *ClientProc) ProxyOutput(sender *packet.PacketSender) { +func (cproc *ClientProc) ProxySingleOutput(ck base.CommandKey, sender *packet.PacketSender) { sentDonePk := false for pk := range cproc.Output.MainCh { if pk.GetType() == packet.CmdDonePacketStr { @@ -112,7 +106,7 @@ func (cproc *ClientProc) ProxyOutput(sender *packet.PacketSender) { if !sentDonePk { endTs := time.Now() cmdDuration := endTs.Sub(cproc.StartTs) - donePacket := packet.MakeCmdDonePacket(cproc.CK) + donePacket := packet.MakeCmdDonePacket(ck) donePacket.Ts = endTs.UnixMilli() donePacket.ExitCode = GetExitCode(exitErr) donePacket.DurationMs = int64(cmdDuration / time.Millisecond) diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 6b549f56..03d0ba6e 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -7,6 +7,7 @@ package shexec import ( + "context" "encoding/base64" "fmt" "io" @@ -359,6 +360,15 @@ func (opts SSHOpts) MakeSSHInstallCmd() (*exec.Cmd, error) { return opts.MakeSSHExecCmd(InstallCommand), nil } +func (opts SSHOpts) MakeMShellServerCmd() (*exec.Cmd, error) { + msPath, err := base.GetMShellPath() + if err != nil { + return nil, err + } + ecmd := exec.Command(msPath, "--server") + return ecmd, nil +} + func (opts SSHOpts) MakeMShellSingleCmd() (*exec.Cmd, error) { if opts.SSHHost == "" { execFile, err := os.Executable() @@ -716,7 +726,7 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdCon if !versionOk { return nil, fmt.Errorf("did not receive version from remote mshell") } - SendRunPacketAndRunData(sender, runPacket) + SendRunPacketAndRunData(context.Background(), sender, runPacket) if debug { cmd.Multiplexer.Debug = true } @@ -735,10 +745,13 @@ func min(v1 int, v2 int) int { return v2 } -func SendRunPacketAndRunData(sender *packet.PacketSender, runPacket *packet.RunPacketType) { - sender.SendPacket(runPacket) +func SendRunPacketAndRunData(ctx context.Context, sender *packet.PacketSender, runPacket *packet.RunPacketType) error { + err := sender.SendPacketCtx(ctx, runPacket) + if err != nil { + return err + } if len(runPacket.RunData) == 0 { - return + return nil } for _, runData := range runPacket.RunData { sendBuf := runData.Data @@ -751,10 +764,17 @@ func SendRunPacketAndRunData(sender *packet.PacketSender, runPacket *packet.RunP dataPk.Data64 = base64.StdEncoding.EncodeToString(chunk) dataPk.Eof = (len(chunk) == len(sendBuf)) sendBuf = sendBuf[chunkSize:] - sender.SendPacket(dataPk) + err = sender.SendPacketCtx(ctx, dataPk) + if err != nil { + return err + } } } - sender.SendPacket(packet.MakeDataEndPacket(runPacket.CK)) + err = sender.SendPacketCtx(ctx, packet.MakeDataEndPacket(runPacket.CK)) + if err != nil { + return err + } + return nil } func DetectGoArch(uname string) (string, string, error) {