From c43d3ecc85bedf08dd9c7ff8dc6cf91edf800348 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 23 Jun 2022 12:48:45 -0700 Subject: [PATCH] checkpoint got stdout/stderr data packets working with new remote handler --- main-mshell.go | 116 ++++++++++++++++++++--- pkg/base/base.go | 17 ++-- pkg/packet/packet.go | 77 ++++++++-------- pkg/shexec/shexec.go | 212 ++++++++++++++++++++++++++++++++++++------- 4 files changed, 329 insertions(+), 93 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index dd551372..0bdb4e77 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -11,6 +11,7 @@ import ( "os" "os/signal" "os/user" + "strings" "syscall" "time" @@ -21,8 +22,10 @@ import ( "github.com/scripthaus-dev/mshell/pkg/shexec" ) -// in single run mode, we don't want the runner to die from signals -// since we want the single runner to persist even if session / main runner +const MShellVersion = "0.1.0" + +// in single run mode, we don't want mshell to die from signals +// since we want the single mshell to persist even if session / main mshell // is terminated. func setupSingleSignals(cmd *shexec.ShExecType) { sigCh := make(chan os.Signal, 1) @@ -46,7 +49,7 @@ func doSingle(cmdId string) { runPacket, _ = pk.(*packet.RunPacketType) break } - sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to runner", pk.GetType())) + sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", pk.GetType())) return } if runPacket == nil { @@ -66,13 +69,9 @@ func doSingle(cmdId string) { return } setupSingleSignals(cmd) - startPacket := packet.MakeCmdStartPacket() - startPacket.Ts = time.Now().UnixMilli() - startPacket.CmdId = runPacket.CmdId - startPacket.Pid = cmd.Cmd.Process.Pid - startPacket.RunnerPid = os.Getpid() + startPacket := cmd.MakeCmdStartPacket() sender.SendPacket(startPacket) - donePacket := cmd.WaitForCommand(runPacket.CmdId) + donePacket := cmd.WaitForCommand() sender.SendPacket(donePacket) sender.CloseSendCh() sender.WaitForDone() @@ -94,7 +93,7 @@ func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { } cmd, err := shexec.MakeRunnerExec(pk.CmdId) if err != nil { - sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot make runner command: %v", err))) + sender.SendPacket(packet.MakeIdErrorPacket(pk.CmdId, fmt.Sprintf("cannot make mshell command: %v", err))) return } cmdStdin, err := cmd.StdinPipe() @@ -155,7 +154,7 @@ func doMain() { packet.SendErrorPacket(os.Stdout, fmt.Sprintf("cannot change directory to $HOME '%s': %v", homeDir, err)) return } - err = base.EnsureMShellPath() + _, err = base.GetMShellPath() if err != nil { packet.SendErrorPacket(os.Stdout, err.Error()) return @@ -168,7 +167,7 @@ func doMain() { return } go tailer.Run() - initPacket := packet.MakeRunnerInitPacket() + initPacket := packet.MakeInitPacket() initPacket.Env = os.Environ() initPacket.HomeDir = homeDir initPacket.ScHomeDir = scHomeDir @@ -207,19 +206,106 @@ func doMain() { } if pk.GetType() == packet.ErrorPacketStr { errPk := pk.(*packet.ErrorPacketType) - errPk.Error = "invalid packet sent to runner: " + errPk.Error + errPk.Error = "invalid packet sent to mshell: " + errPk.Error sender.SendPacket(errPk) continue } - sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to runner", pk.GetType())) + sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", pk.GetType())) } } +func handleRemote() { + packetCh := packet.PacketParser(os.Stdin) + sender := packet.MakePacketSender(os.Stdout) + defer func() { + // wait for sender to complete + close(sender.SendCh) + <-sender.DoneCh + }() + initPacket := packet.MakeInitPacket() + initPacket.Version = MShellVersion + sender.SendPacket(initPacket) + var runPacket *packet.RunPacketType + for pk := range packetCh { + if pk.GetType() == packet.PingPacketStr { + continue + } + if pk.GetType() == packet.RunPacketStr { + runPacket, _ = pk.(*packet.RunPacketType) + break + } + sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", pk.GetType())) + return + } + cmd, err := shexec.RunCommand(runPacket, sender) + if err != nil { + sender.SendErrorPacket(fmt.Sprintf("error running command: %v", err)) + return + } + defer cmd.Close() + startPacket := cmd.MakeCmdStartPacket() + sender.SendPacket(startPacket) + cmd.RunIOAndWait(sender) +} + +func handleServer() { +} + +func handleClient() { + fmt.Printf("mshell client\n") +} + +func handleUsage(extended bool) { + usage := ` +Client Usage: mshell [mshell-opts] [ssh-opts] user@host [command] + +mshell multiplexes input and output streams to a remote command over ssh. + +Options: + --env 'X=Y,A=B' - set remote environment variables for command, comma or newline separated + --env-file [file] - load environment variables from [file] (.env format) + --env-copy [glob] - copy local environment variables to remote using [glob] pattern + --cwd [dir] - execute remote command in [dir] + --no-auto-fds - do not auto-detect additional fds + --fds [fdspec] - open fds based off [fdspec], comma separated (implies --no-auto-fds) + <[num] opens for reading + >[num] opens for writing + <>[num] opens for read/write + e.g. --fds '<5,>6,<>7' + +mshell is licensed under the MPLv2 +Please see https://github.com/scripthaus-dev/mshell for extended usage modes, source code, bugs, and feature requests +` + fmt.Printf("%s\n\n", strings.TrimSpace(usage)) +} + func main() { + if len(os.Args) == 1 { + handleUsage(false) + return + } + firstArg := os.Args[1] + if firstArg == "--help" { + handleUsage(true) + return + } else if firstArg == "--version" { + fmt.Printf("mshell v%s\n", MShellVersion) + return + } else if firstArg == "--remote" { + handleRemote() + return + } else if firstArg == "--server" { + handleServer() + return + } else { + handleClient() + return + } + if len(os.Args) >= 2 { cmdId, err := uuid.Parse(os.Args[1]) if err != nil { - packet.SendErrorPacket(os.Stdout, fmt.Sprintf("invalid non-cmdid passed to runner", err)) + packet.SendErrorPacket(os.Stdout, fmt.Sprintf("invalid non-cmdid passed to mshell", err)) return } doSingle(cmdId.String()) diff --git a/pkg/base/base.go b/pkg/base/base.go index 91f2959b..dc13157f 100644 --- a/pkg/base/base.go +++ b/pkg/base/base.go @@ -17,6 +17,7 @@ import ( ) const DefaultMShellPath = "mshell" +const DefaultUserMShellPath = ".mshell/mshell" const MShellPathVarName = "MSHELL_PATH" const SSHCommandVarName = "SSH_COMMAND" const ScHomeVarName = "SCRIPTHAUS_HOME" @@ -128,21 +129,17 @@ func EnsureSessionDir(sessionId string) (string, error) { return sdir, nil } -func GetMShellPath() string { +func GetMShellPath() (string, error) { msPath := os.Getenv(MShellPathVarName) if msPath != "" { - return msPath + return exec.LookPath(msPath) } - return DefaultMShellPath -} - -func EnsureMShellPath() error { - msPath := GetMShellPath() - _, err := exec.LookPath(msPath) + userMShellPath := path.Join(GetHomeDir(), DefaultUserMShellPath) + msPath, err := exec.LookPath(userMShellPath) if err != nil { - return err + return msPath, nil } - return nil + return exec.LookPath(DefaultMShellPath) } func GetScSessionsDir() (string, error) { diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 863fc418..33989978 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -18,28 +18,28 @@ import ( "sync" ) -// remote: runnerinit, run, ping, data, cmdstart, cmddone -// remote(detached): runnerinit, run, cmdstart -// server: runnerinit, run, ping, cmdstart, cmddone, cd, resp, getcmd, untailcmd, cmddata, input, data, [comp] +// remote: init, run, ping, data, cmdstart, cmddone +// remote(detached): init, run, cmdstart +// server: init, run, ping, cmdstart, cmddone, cd, resp, getcmd, untailcmd, cmddata, input, data, [comp] // all: error, message const ( - RunPacketStr = "run" - PingPacketStr = "ping" - RunnerInitPacketStr = "runnerinit" - DataPacketStr = "data" - CmdStartPacketStr = "cmdstart" - CmdDonePacketStr = "cmddone" - ResponsePacketStr = "resp" - DonePacketStr = "done" - ErrorPacketStr = "error" - MessagePacketStr = "message" - GetCmdPacketStr = "getcmd" - UntailCmdPacketStr = "untailcmd" - CdPacketStr = "cd" - CmdDataPacketStr = "cmddata" - RawPacketStr = "raw" - InputPacketStr = "input" + RunPacketStr = "run" + PingPacketStr = "ping" + InitPacketStr = "init" + DataPacketStr = "data" + CmdStartPacketStr = "cmdstart" + CmdDonePacketStr = "cmddone" + ResponsePacketStr = "resp" + DonePacketStr = "done" + ErrorPacketStr = "error" + MessagePacketStr = "message" + GetCmdPacketStr = "getcmd" + UntailCmdPacketStr = "untailcmd" + CdPacketStr = "cd" + CmdDataPacketStr = "cmddata" + RawPacketStr = "raw" + InputPacketStr = "input" ) var TypeStrToFactory map[string]reflect.Type @@ -56,7 +56,7 @@ func init() { TypeStrToFactory[CmdDonePacketStr] = reflect.TypeOf(CmdDonePacketType{}) TypeStrToFactory[GetCmdPacketStr] = reflect.TypeOf(GetCmdPacketType{}) TypeStrToFactory[UntailCmdPacketStr] = reflect.TypeOf(UntailCmdPacketType{}) - TypeStrToFactory[RunnerInitPacketStr] = reflect.TypeOf(RunnerInitPacketType{}) + TypeStrToFactory[InitPacketStr] = reflect.TypeOf(InitPacketType{}) TypeStrToFactory[CdPacketStr] = reflect.TypeOf(CdPacketType{}) TypeStrToFactory[CmdDataPacketStr] = reflect.TypeOf(CmdDataPacketType{}) TypeStrToFactory[RawPacketStr] = reflect.TypeOf(RawPacketType{}) @@ -112,18 +112,20 @@ func MakePingPacket() *PingPacketType { type DataPacketType struct { Type string `json:"type"` - SessionId string `json:"sessionid"` - CmdId string `json:"cmdid"` + SessionId string `json:"sessionid,omitempty"` + CmdId string `json:"cmdid,omitempty"` FdNum int `json:"fdnum"` Data string `json:"data"` + Eof bool `json:"eof,omitempty"` + Error string `json:"error,omitempty"` } func (*DataPacketType) GetType() string { return DataPacketStr } -func MakeDataPacket(fdNum int, data string) *DataPacketType { - return &DataPacketType{Type: DataPacketStr, FdNum: fdNum, Data: data} +func MakeDataPacket() *DataPacketType { + return &DataPacketType{Type: DataPacketStr} } // InputData gets written to PTY directly @@ -249,7 +251,7 @@ func FmtMessagePacket(fmtStr string, args ...interface{}) *MessagePacketType { return &MessagePacketType{Type: MessagePacketStr, Message: message} } -type RunnerInitPacketType struct { +type InitPacketType struct { Type string `json:"type"` Version string `json:"version"` ScHomeDir string `json:"schomedir,omitempty"` @@ -258,12 +260,12 @@ type RunnerInitPacketType struct { User string `json:"user,omitempty"` } -func (*RunnerInitPacketType) GetType() string { - return RunnerInitPacketStr +func (*InitPacketType) GetType() string { + return InitPacketStr } -func MakeRunnerInitPacket() *RunnerInitPacketType { - return &RunnerInitPacketType{Type: RunnerInitPacketStr} +func MakeInitPacket() *InitPacketType { + return &InitPacketType{Type: InitPacketStr} } type DonePacketType struct { @@ -281,7 +283,8 @@ func MakeDonePacket() *DonePacketType { type CmdDonePacketType struct { Type string `json:"type"` Ts int64 `json:"ts"` - CmdId string `json:"cmdid"` + SessionId string `json:"sessionid,omitempty"` + CmdId string `json:"cmdid,omitempty"` ExitCode int `json:"exitcode"` DurationMs int64 `json:"durationms"` } @@ -297,9 +300,10 @@ func MakeCmdDonePacket() *CmdDonePacketType { type CmdStartPacketType struct { Type string `json:"type"` Ts int64 `json:"ts"` - CmdId string `json:"cmdid"` + SessionId string `json:"sessionid,omitempty"` + CmdId string `json:"cmdid,omitempty"` Pid int `json:"pid"` - RunnerPid int `json:"runnerpid"` + MShellPid int `json:"mshellpid"` } func (*CmdStartPacketType) GetType() string { @@ -323,13 +327,14 @@ type RemoteFd struct { type RunPacketType struct { Type string `json:"type"` - SessionId string `json:"sessionid"` - CmdId string `json:"cmdid"` + SessionId string `json:"sessionid,omitempty"` + CmdId string `json:"cmdid,omitempty"` Command string `json:"command"` Cwd string `json:"cwd,omitempty"` Env map[string]string `json:"env,omitempty"` - TermSize TermSize `json:"termsize"` - Fds []RemoteFd `json:"fds"` + TermSize TermSize `json:"termsize,omitempty"` + Fds []RemoteFd `json:"fds,omitempty"` + Detached bool `json:"detached,omitempty"` } func (*RunPacketType) GetType() string { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 600dd379..e2cd6a8e 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -12,6 +12,7 @@ import ( "os" "os/exec" "strings" + "sync" "syscall" "time" @@ -27,14 +28,48 @@ const MaxRows = 1024 const MaxCols = 1024 type ShExecType struct { - FileNames *base.CommandFileNames - Cmd *exec.Cmd - CmdPty *os.File - StartTs time.Time + StartTs time.Time + RunPacket *packet.RunPacketType + FileNames *base.CommandFileNames + Cmd *exec.Cmd + CmdPty *os.File + FdReaders map[int]*os.File + FdWriters map[int]*os.File + CloseAfterStart []*os.File +} + +func MakeShExec(pk *packet.RunPacketType) *ShExecType { + return &ShExecType{ + StartTs: time.Now(), + RunPacket: pk, + FdReaders: make(map[int]*os.File), + FdWriters: make(map[int]*os.File), + } } func (c *ShExecType) Close() { - c.CmdPty.Close() + if c.CmdPty != nil { + c.CmdPty.Close() + } + for _, fd := range c.FdReaders { + fd.Close() + } + for _, fd := range c.FdWriters { + fd.Close() + } + for _, fd := range c.CloseAfterStart { + fd.Close() + } +} + +func (c *ShExecType) MakeCmdStartPacket() *packet.CmdStartPacketType { + startPacket := packet.MakeCmdStartPacket() + startPacket.Ts = time.Now().UnixMilli() + startPacket.SessionId = c.RunPacket.SessionId + startPacket.CmdId = c.RunPacket.CmdId + startPacket.Pid = c.Cmd.Process.Pid + startPacket.MShellPid = os.Getpid() + return startPacket } func getEnvStrKey(envStr string) string { @@ -93,7 +128,10 @@ func MakeExecCmd(pk *packet.RunPacketType, cmdTty *os.File) *exec.Cmd { } func MakeRunnerExec(cmdId string) (*exec.Cmd, error) { - msPath := base.GetMShellPath() + msPath, err := base.GetMShellPath() + if err != nil { + return nil, err + } ecmd := exec.Command(msPath, cmdId) return ecmd, nil } @@ -124,19 +162,21 @@ func ValidateRunPacket(pk *packet.RunPacketType) error { if pk.Type != packet.RunPacketStr { return fmt.Errorf("run packet has wrong type: %s", pk.Type) } - if pk.SessionId == "" { - return fmt.Errorf("run packet does not have sessionid") - } - _, err := uuid.Parse(pk.SessionId) - if err != nil { - return fmt.Errorf("invalid sessionid '%s' for command", pk.SessionId) - } - if pk.CmdId == "" { - return fmt.Errorf("run packet does not have cmdid") - } - _, err = uuid.Parse(pk.CmdId) - if err != nil { - return fmt.Errorf("invalid cmdid '%s' for command", pk.CmdId) + if pk.Detached { + if pk.SessionId == "" { + return fmt.Errorf("run packet does not have sessionid") + } + _, err := uuid.Parse(pk.SessionId) + if err != nil { + return fmt.Errorf("invalid sessionid '%s' for command", pk.SessionId) + } + if pk.CmdId == "" { + return fmt.Errorf("run packet does not have cmdid") + } + _, err = uuid.Parse(pk.CmdId) + if err != nil { + return fmt.Errorf("invalid cmdid '%s' for command", pk.CmdId) + } } if pk.Cwd != "" { dirInfo, err := os.Stat(pk.Cwd) @@ -164,13 +204,120 @@ func GetWinsize(p *packet.RunPacketType) *pty.Winsize { // when err is nil, the command will have already been started func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { - if pk.CmdId == "" { - pk.CmdId = uuid.New().String() - } err := ValidateRunPacket(pk) if err != nil { return nil, err } + if !pk.Detached { + return runCommandSimple(pk, sender) + } else { + return runCommandDetached(pk, sender) + } +} + +// returns the *writer* to connect to process, reader is put in FdReaders +func (cmd *ShExecType) makeReaderPipe(fdNum int) (*os.File, error) { + pr, pw, err := os.Pipe() + if err != nil { + return nil, err + } + cmd.FdReaders[fdNum] = pr + cmd.CloseAfterStart = append(cmd.CloseAfterStart, pw) + return pw, nil +} + +// returns the *reader* to connect to process, writer is put in FdWriters +func (cmd *ShExecType) makeWriterPipe(fdNum int) (*os.File, error) { + pr, pw, err := os.Pipe() + if err != nil { + return nil, err + } + cmd.FdWriters[fdNum] = pw + cmd.CloseAfterStart = append(cmd.CloseAfterStart, pr) + return pr, nil +} + +func (cmd *ShExecType) MakeDataPacket(fdNum int, data []byte) *packet.DataPacketType { + pk := packet.MakeDataPacket() + pk.SessionId = cmd.RunPacket.SessionId + pk.CmdId = cmd.RunPacket.CmdId + pk.FdNum = fdNum + pk.Data = string(data) + return pk +} + +func (cmd *ShExecType) runReadLoop(wg *sync.WaitGroup, fdNum int, fd *os.File, sender *packet.PacketSender) { + go func() { + defer fd.Close() + defer wg.Done() + buf := make([]byte, 4096) + for { + nr, err := fd.Read(buf) + pk := cmd.MakeDataPacket(fdNum, buf[0:nr]) + if err == io.EOF { + pk.Eof = true + sender.SendPacket(pk) + break + } else if err != nil { + pk.Error = err.Error() + sender.SendPacket(pk) + break + } else { + sender.SendPacket(pk) + } + } + }() +} + +func (cmd *ShExecType) RunIOAndWait(sender *packet.PacketSender) { + var wg sync.WaitGroup + wg.Add(len(cmd.FdReaders)) + go func() { + for fdNum, fd := range cmd.FdReaders { + cmd.runReadLoop(&wg, fdNum, fd, sender) + } + }() + donePacket := cmd.WaitForCommand() + wg.Wait() + sender.SendPacket(donePacket) +} + +func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { + cmd := MakeShExec(pk) + cmd.Cmd = exec.Command("bash", "-c", pk.Command) + UpdateCmdEnv(cmd.Cmd, pk.Env) + if pk.Cwd != "" { + cmd.Cmd.Dir = pk.Cwd + } + var err error + cmd.Cmd.Stdin, err = cmd.makeWriterPipe(0) + if err != nil { + cmd.Close() + return nil, err + } + cmd.Cmd.Stdout, err = cmd.makeReaderPipe(1) + if err != nil { + cmd.Close() + return nil, err + } + cmd.Cmd.Stderr, err = cmd.makeReaderPipe(2) + if err != nil { + cmd.Close() + return nil, err + } + err = cmd.Cmd.Start() + if err != nil { + cmd.Close() + return nil, err + } + for _, fd := range cmd.CloseAfterStart { + fd.Close() + } + cmd.CloseAfterStart = nil + return cmd, nil +} + +func runCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { fileNames, err := base.GetCommandFileNames(pk.SessionId, pk.CmdId) if err != nil { return nil, err @@ -190,7 +337,7 @@ func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecT defer func() { cmdTty.Close() }() - startTs := time.Now() + rtn := MakeShExec(pk) ecmd := MakeExecCmd(pk, cmdTty) err = ecmd.Start() if err != nil { @@ -214,12 +361,10 @@ func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecT sender.SendErrorPacket(fmt.Sprintf("reading from stdin fifo: %v", copyFifoErr)) } }() - return &ShExecType{ - FileNames: fileNames, - Cmd: ecmd, - CmdPty: cmdPty, - StartTs: startTs, - }, nil + rtn.FileNames = fileNames + rtn.Cmd = ecmd + rtn.CmdPty = cmdPty + return rtn, nil } func GetExitCode(err error) int { @@ -233,16 +378,19 @@ func GetExitCode(err error) int { } } -func (c *ShExecType) WaitForCommand(cmdId string) *packet.CmdDonePacketType { +func (c *ShExecType) WaitForCommand() *packet.CmdDonePacketType { exitErr := c.Cmd.Wait() endTs := time.Now() cmdDuration := endTs.Sub(c.StartTs) exitCode := GetExitCode(exitErr) donePacket := packet.MakeCmdDonePacket() donePacket.Ts = endTs.UnixMilli() - donePacket.CmdId = cmdId + donePacket.SessionId = c.RunPacket.SessionId + donePacket.CmdId = c.RunPacket.CmdId donePacket.ExitCode = exitCode donePacket.DurationMs = int64(cmdDuration / time.Millisecond) - os.Remove(c.FileNames.StdinFifo) // best effort (no need to check error) + if c.FileNames != nil { + os.Remove(c.FileNames.StdinFifo) // best effort (no need to check error) + } return donePacket }