From 5223760a7687e694f72929b4f43c56803300a9b0 Mon Sep 17 00:00:00 2001 From: sawka Date: Fri, 24 Jun 2022 13:25:09 -0700 Subject: [PATCH] got basic mshell client working -- still need detectfds and extra files support --- main-mshell.go | 41 +++-------------- pkg/mpio/bufreader.go | 32 +++++++------ pkg/mpio/bufwriter.go | 42 +++++++++++------- pkg/mpio/mpio.go | 95 ++++++++++++++++++++++++++++++++------- pkg/packet/packet.go | 8 +++- pkg/shexec/shexec.go | 101 +++++++++++++++++++++++++++++++++++++----- 6 files changed, 226 insertions(+), 93 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index d9608ccc..418769f2 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -9,7 +9,6 @@ package main import ( "fmt" "os" - "os/exec" "os/signal" "os/user" "strings" @@ -251,7 +250,7 @@ func handleRemote() { defer cmd.Close() startPacket := cmd.MakeCmdStartPacket() sender.SendPacket(startPacket) - cmd.RunIOAndWait(packetCh, sender) + cmd.RunRemoteIOAndWait(packetCh, sender) } func handleServer() { @@ -261,17 +260,8 @@ func detectOpenFds() { } -type ClientOpts struct { - IsSSH bool - SSHOptsTerm bool - SSHOpts []string - Command string - Fds []packet.RemoteFd - Cwd string -} - -func parseClientOpts() (*ClientOpts, error) { - opts := &ClientOpts{} +func parseClientOpts() (*shexec.ClientOpts, error) { + opts := &shexec.ClientOpts{} iter := base.MakeOptsIter(os.Args[1:]) for iter.HasNext() { argStr := iter.Next() @@ -313,7 +303,6 @@ func parseClientOpts() (*ClientOpts, error) { } func handleClient() (int, error) { - fmt.Printf("mshell client\n") opts, err := parseClientOpts() if err != nil { return 1, fmt.Errorf("parsing opts: %w", err) @@ -321,29 +310,11 @@ func handleClient() (int, error) { if !opts.IsSSH { return 1, fmt.Errorf("when running in client mode '--ssh' option must be present") } - fmt.Printf("opts: %v\n", opts) - sshRemoteCommand := `PATH=$PATH:~/.mshell; mshell --remote` - sshOpts := append(opts.SSHOpts, sshRemoteCommand) - ecmd := exec.Command("ssh", sshOpts...) - inputWriter, err := ecmd.StdinPipe() + donePacket, err := shexec.RunClientSSHCommandAndWait(opts) if err != nil { - return 1, fmt.Errorf("creating stdin pipe: %v", err) + return 1, err } - outputReader, err := ecmd.StdoutPipe() - if err != nil { - return 1, fmt.Errorf("creating stdout pipe: %v", err) - } - ecmd.Stderr = ecmd.Stdout - err = ecmd.Start() - if err != nil { - return 1, fmt.Errorf("running ssh command: %w", err) - } - parser := packet.PacketParser(outputReader) - go func() { - fmt.Printf("%v %v\n", parser, inputWriter) - }() - exitErr := ecmd.Wait() - return shexec.GetExitCode(exitErr), nil + return donePacket.ExitCode, nil } func handleUsage() { diff --git a/pkg/mpio/bufreader.go b/pkg/mpio/bufreader.go index f78ba4b1..756f9e83 100644 --- a/pkg/mpio/bufreader.go +++ b/pkg/mpio/bufreader.go @@ -15,21 +15,23 @@ import ( ) type FdReader struct { - CVar *sync.Cond - M *Multiplexer - FdNum int - Fd *os.File - BufSize int - Closed bool + CVar *sync.Cond + M *Multiplexer + FdNum int + Fd *os.File + BufSize int + Closed bool + ShouldCloseFd bool } -func MakeFdReader(m *Multiplexer, fd *os.File, fdNum int) *FdReader { +func MakeFdReader(m *Multiplexer, fd *os.File, fdNum int, shouldCloseFd bool) *FdReader { fr := &FdReader{ - CVar: sync.NewCond(&sync.Mutex{}), - M: m, - FdNum: fdNum, - Fd: fd, - BufSize: 0, + CVar: sync.NewCond(&sync.Mutex{}), + M: m, + FdNum: fdNum, + Fd: fd, + BufSize: 0, + ShouldCloseFd: shouldCloseFd, } return fr } @@ -40,7 +42,7 @@ func (r *FdReader) Close() { if r.Closed { return } - if r.Fd != nil { + if r.Fd != nil && r.ShouldCloseFd { r.Fd.Close() } r.CVar.Broadcast() @@ -110,7 +112,9 @@ func (r *FdReader) isClosed() bool { func (r *FdReader) ReadLoop(wg *sync.WaitGroup) { defer r.Close() - defer wg.Done() + if wg != nil { + defer wg.Done() + } buf := make([]byte, 4096) for { nr, err := r.Fd.Read(buf) diff --git a/pkg/mpio/bufwriter.go b/pkg/mpio/bufwriter.go index 16e139ca..9b389678 100644 --- a/pkg/mpio/bufwriter.go +++ b/pkg/mpio/bufwriter.go @@ -13,21 +13,23 @@ import ( ) type FdWriter struct { - CVar *sync.Cond - M *Multiplexer - FdNum int - Buffer []byte - Fd *os.File - Eof bool - Closed bool + CVar *sync.Cond + M *Multiplexer + FdNum int + Buffer []byte + Fd *os.File + Eof bool + Closed bool + ShouldCloseFd bool } -func MakeFdWriter(m *Multiplexer, fd *os.File, fdNum int) *FdWriter { +func MakeFdWriter(m *Multiplexer, fd *os.File, fdNum int, shouldCloseFd bool) *FdWriter { fw := &FdWriter{ - CVar: sync.NewCond(&sync.Mutex{}), - Fd: fd, - M: m, - FdNum: fdNum, + CVar: sync.NewCond(&sync.Mutex{}), + Fd: fd, + M: m, + FdNum: fdNum, + ShouldCloseFd: shouldCloseFd, } return fw } @@ -39,7 +41,7 @@ func (w *FdWriter) Close() { return } w.Closed = true - if w.Fd != nil { + if w.Fd != nil && w.ShouldCloseFd { w.Fd.Close() } w.Buffer = nil @@ -65,6 +67,9 @@ func (w *FdWriter) AddData(data []byte, eof bool) error { if w.Closed { return fmt.Errorf("write to closed file") } + if w.Eof { + return fmt.Errorf("write to closed file (eof)") + } if len(data) > 0 { if len(data)+len(w.Buffer) > WriteBufSize { return fmt.Errorf("write exceeds buffer size") @@ -78,8 +83,11 @@ func (w *FdWriter) AddData(data []byte, eof bool) error { return nil } -func (w *FdWriter) WriteLoop() { +func (w *FdWriter) WriteLoop(wg *sync.WaitGroup) { defer w.Close() + if wg != nil { + defer wg.Done() + } for { data, isEof := w.WaitForData() // chunk the writes to make sure we send ample ack packets @@ -90,8 +98,10 @@ func (w *FdWriter) WriteLoop() { chunkSize := min(len(data), MaxSingleWriteSize) chunk := data[0:chunkSize] nw, err := w.Fd.Write(chunk) - ack := w.M.makeDataAckPacket(w.FdNum, nw, err) - w.M.sendPacket(ack) + if nw > 0 || err != nil { + ack := w.M.makeDataAckPacket(w.FdNum, nw, err) + w.M.sendPacket(ack) + } if err != nil { return } diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index c3f82d61..31a7c577 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -45,17 +45,32 @@ func (m *Multiplexer) Close() { m.Lock.Lock() defer m.Lock.Unlock() - for _, fd := range m.FdReaders { - fd.Close() + for _, fr := range m.FdReaders { + fr.Close() } - for _, fd := range m.FdWriters { - fd.Close() + for _, fw := range m.FdWriters { + fw.Close() } for _, fd := range m.CloseAfterStart { fd.Close() } } +func (m *Multiplexer) HandleInputDone() { + m.Lock.Lock() + defer m.Lock.Unlock() + + // close readers (obviously the done command needs no more input) + for _, fr := range m.FdReaders { + fr.Close() + } + + // ensure EOF on all writers (ignore error) + for _, fw := range m.FdWriters { + fw.AddData(nil, true) + } +} + // returns the *writer* to connect to process, reader is put in FdReaders func (m *Multiplexer) MakeReaderPipe(fdNum int) (*os.File, error) { pr, pw, err := os.Pipe() @@ -64,7 +79,7 @@ func (m *Multiplexer) MakeReaderPipe(fdNum int) (*os.File, error) { } m.Lock.Lock() defer m.Lock.Unlock() - m.FdReaders[fdNum] = MakeFdReader(m, pr, fdNum) + m.FdReaders[fdNum] = MakeFdReader(m, pr, fdNum, true) m.CloseAfterStart = append(m.CloseAfterStart, pw) return pw, nil } @@ -77,11 +92,23 @@ func (m *Multiplexer) MakeWriterPipe(fdNum int) (*os.File, error) { } m.Lock.Lock() defer m.Lock.Unlock() - m.FdWriters[fdNum] = MakeFdWriter(m, pw, fdNum) + m.FdWriters[fdNum] = MakeFdWriter(m, pw, fdNum, true) m.CloseAfterStart = append(m.CloseAfterStart, pr) return pr, nil } +func (m *Multiplexer) MakeRawFdReader(fdNum int, fd *os.File) { + m.Lock.Lock() + defer m.Lock.Unlock() + m.FdReaders[fdNum] = MakeFdReader(m, fd, fdNum, false) +} + +func (m *Multiplexer) MakeRawFdWriter(fdNum int, fd *os.File) { + m.Lock.Lock() + defer m.Lock.Unlock() + m.FdWriters[fdNum] = MakeFdWriter(m, fd, fdNum, false) +} + func (m *Multiplexer) makeDataAckPacket(fdNum int, ackLen int, err error) *packet.DataAckPacketType { ack := packet.MakeDataAckPacket() ack.SessionId = m.SessionId @@ -110,18 +137,23 @@ func (m *Multiplexer) sendPacket(p packet.PacketType) { m.Sender.SendPacket(p) } -func (m *Multiplexer) launchWriters() { +func (m *Multiplexer) launchWriters(wg *sync.WaitGroup) { m.Lock.Lock() defer m.Lock.Unlock() + if wg != nil { + wg.Add(len(m.FdWriters)) + } for _, fw := range m.FdWriters { - go fw.WriteLoop() + go fw.WriteLoop(wg) } } func (m *Multiplexer) launchReaders(wg *sync.WaitGroup) { m.Lock.Lock() defer m.Lock.Unlock() - wg.Add(len(m.FdReaders)) + if wg != nil { + wg.Add(len(m.FdReaders)) + } for _, fr := range m.FdReaders { go fr.ReadLoop(wg) } @@ -138,7 +170,8 @@ func (m *Multiplexer) startIO(packetCh chan packet.PacketType, sender *packet.Pa m.Started = true } -func (m *Multiplexer) runPacketInputLoop() { +func (m *Multiplexer) runPacketInputLoop() *packet.CmdDonePacketType { + defer m.HandleInputDone() for pk := range m.Input { if pk.GetType() == packet.DataPacketStr { dataPacket := pk.(*packet.DataPacketType) @@ -152,9 +185,15 @@ func (m *Multiplexer) runPacketInputLoop() { if pk.GetType() == packet.DataAckPacketStr { ackPacket := pk.(*packet.DataAckPacketType) m.processAckPacket(ackPacket) + continue + } + if pk.GetType() == packet.CmdDonePacketStr { + donePacket := pk.(*packet.CmdDonePacketType) + return donePacket } // other packet types are ignored } + return nil } func (m *Multiplexer) processDataPacket(dataPacket *packet.DataPacketType) error { @@ -163,7 +202,7 @@ func (m *Multiplexer) processDataPacket(dataPacket *packet.DataPacketType) error fw := m.FdWriters[dataPacket.FdNum] if fw == nil { // add a closed FdWriter as a placeholder so we only send one error - fw := MakeFdWriter(m, nil, dataPacket.FdNum) + fw := MakeFdWriter(m, nil, dataPacket.FdNum, false) fw.Close() m.FdWriters[dataPacket.FdNum] = fw return fmt.Errorf("write to closed file") @@ -195,12 +234,38 @@ func (m *Multiplexer) closeTempStartFds() { m.CloseAfterStart = nil } -func (m *Multiplexer) RunIOAndWait(packetCh chan packet.PacketType, sender *packet.PacketSender) { +func (m *Multiplexer) RunIOAndWait(packetCh chan packet.PacketType, sender *packet.PacketSender, waitOnReaders bool, waitOnWriters bool, waitForInputLoop bool) *packet.CmdDonePacketType { m.startIO(packetCh, sender) m.closeTempStartFds() var wg sync.WaitGroup - m.launchReaders(&wg) - m.launchWriters() - go m.runPacketInputLoop() + if waitOnReaders { + m.launchReaders(&wg) + } else { + m.launchReaders(nil) + } + if waitOnWriters { + m.launchWriters(&wg) + } else { + m.launchWriters(nil) + } + var donePacket *packet.CmdDonePacketType + if waitForInputLoop { + wg.Add(1) + } + go func() { + if waitForInputLoop { + defer wg.Done() + } + pkRtn := m.runPacketInputLoop() + if pkRtn != nil { + m.Lock.Lock() + donePacket = pkRtn + m.Lock.Unlock() + } + }() wg.Wait() + + m.Lock.Lock() + defer m.Lock.Unlock() + return donePacket } diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 5f6f3bd2..9217b17f 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -23,6 +23,8 @@ import ( // server: init, run, ping, cmdstart, cmddone, cd, resp, getcmd, untailcmd, cmddata, input, data, [comp] // all: error, message +var GlobalDebug = false + const ( RunPacketStr = "run" PingPacketStr = "ping" @@ -353,7 +355,7 @@ type RunPacketType struct { Command string `json:"command"` Cwd string `json:"cwd,omitempty"` Env map[string]string `json:"env,omitempty"` - TermSize TermSize `json:"termsize,omitempty"` + TermSize *TermSize `json:"termsize,omitempty"` Fds []RemoteFd `json:"fds,omitempty"` Detached bool `json:"detached,omitempty"` } @@ -430,6 +432,10 @@ func SendPacket(w io.Writer, packet PacketType) error { outBuf.WriteString(fmt.Sprintf("##%d", len(jsonBytes))) outBuf.Write(jsonBytes) outBuf.WriteByte('\n') + if GlobalDebug { + outBytes := outBuf.Bytes() + fmt.Printf("SEND>%s", string(outBytes[1:])) + } _, err = w.Write(outBuf.Bytes()) if err != nil { return err diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index e8fc52ad..cc372d92 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -33,19 +33,21 @@ const FirstExtraFilesFdNum = 3 type ShExecType struct { Lock *sync.Mutex StartTs time.Time - RunPacket *packet.RunPacketType + SessionId string + CmdId string FileNames *base.CommandFileNames Cmd *exec.Cmd CmdPty *os.File Multiplexer *mpio.Multiplexer } -func MakeShExec(pk *packet.RunPacketType) *ShExecType { +func MakeShExec(sessionId string, cmdId string) *ShExecType { return &ShExecType{ Lock: &sync.Mutex{}, StartTs: time.Now(), - RunPacket: pk, - Multiplexer: mpio.MakeMultiplexer(pk.SessionId, pk.CmdId), + SessionId: sessionId, + CmdId: cmdId, + Multiplexer: mpio.MakeMultiplexer(sessionId, cmdId), } } @@ -59,8 +61,8 @@ func (c *ShExecType) 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.SessionId = c.SessionId + startPacket.CmdId = c.CmdId startPacket.Pid = c.Cmd.Process.Pid startPacket.MShellPid = os.Getpid() return startPacket @@ -209,14 +211,89 @@ func RunCommand(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecT } } -func (cmd *ShExecType) RunIOAndWait(packetCh chan packet.PacketType, sender *packet.PacketSender) { - cmd.Multiplexer.RunIOAndWait(packetCh, sender) +type ClientOpts struct { + IsSSH bool + SSHOptsTerm bool + SSHOpts []string + Command string + Fds []packet.RemoteFd + Cwd string +} + +func (opts *ClientOpts) MakeRunPacket() *packet.RunPacketType { + runPacket := packet.MakeRunPacket() + runPacket.Command = opts.Command + runPacket.Cwd = opts.Cwd + runPacket.Fds = opts.Fds + return runPacket +} + +func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, error) { + // packet.GlobalDebug = true + cmd := MakeShExec("", "") + sshRemoteCommand := `PATH=$PATH:~/.mshell; mshell --remote` + var fullSshOpts []string + fullSshOpts = append(fullSshOpts, opts.SSHOpts...) + fullSshOpts = append(fullSshOpts, sshRemoteCommand) + ecmd := exec.Command("ssh", fullSshOpts...) + cmd.Cmd = ecmd + inputWriter, err := ecmd.StdinPipe() + if err != nil { + return nil, fmt.Errorf("creating stdin pipe: %v", err) + } + stdoutReader, err := ecmd.StdoutPipe() + if err != nil { + return nil, fmt.Errorf("creating stdout pipe: %v", err) + } + stderrReader, err := ecmd.StderrPipe() + if err != nil { + return nil, fmt.Errorf("creating stderr pipe: %v", err) + } + err = ecmd.Start() + if err != nil { + return nil, fmt.Errorf("running ssh command: %w", err) + } + defer cmd.Close() + packetCh := packet.PacketParser(stdoutReader) + go func() { + io.Copy(os.Stderr, stderrReader) + }() + sender := packet.MakePacketSender(inputWriter) + for pk := range packetCh { + if pk.GetType() == packet.RawPacketStr { + rawPk := pk.(*packet.RawPacketType) + fmt.Printf("%s\n", rawPk.Data) + continue + } + if pk.GetType() == packet.InitPacketStr { + initPk := pk.(*packet.InitPacketType) + if initPk.Version != "0.1.0" { + return nil, fmt.Errorf("invalid remote mshell version 'v%s', must be v0.1.0", initPk.Version) + } + break + } + } + runPacket := opts.MakeRunPacket() + sender.SendPacket(runPacket) + cmd.Multiplexer.MakeRawFdReader(0, os.Stdin) + cmd.Multiplexer.MakeRawFdWriter(1, os.Stdout) + cmd.Multiplexer.MakeRawFdWriter(2, os.Stderr) + remoteDonePacket := cmd.Multiplexer.RunIOAndWait(packetCh, sender, false, true, true) + donePacket := cmd.WaitForCommand() + if remoteDonePacket != nil { + donePacket = remoteDonePacket + } + return donePacket, nil +} + +func (cmd *ShExecType) RunRemoteIOAndWait(packetCh chan packet.PacketType, sender *packet.PacketSender) { + cmd.Multiplexer.RunIOAndWait(packetCh, sender, true, false, false) donePacket := cmd.WaitForCommand() sender.SendPacket(donePacket) } func runCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, error) { - cmd := MakeShExec(pk) + cmd := MakeShExec(pk.SessionId, pk.CmdId) cmd.Cmd = exec.Command("bash", "-c", pk.Command) UpdateCmdEnv(cmd.Cmd, pk.Env) if pk.Cwd != "" { @@ -316,7 +393,7 @@ func runCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) ( defer func() { cmdTty.Close() }() - rtn := MakeShExec(pk) + rtn := MakeShExec(pk.SessionId, pk.CmdId) ecmd := MakeExecCmd(pk, cmdTty) err = ecmd.Start() if err != nil { @@ -364,8 +441,8 @@ func (c *ShExecType) WaitForCommand() *packet.CmdDonePacketType { exitCode := GetExitCode(exitErr) donePacket := packet.MakeCmdDonePacket() donePacket.Ts = endTs.UnixMilli() - donePacket.SessionId = c.RunPacket.SessionId - donePacket.CmdId = c.RunPacket.CmdId + donePacket.SessionId = c.SessionId + donePacket.CmdId = c.CmdId donePacket.ExitCode = exitCode donePacket.DurationMs = int64(cmdDuration / time.Millisecond) if c.FileNames != nil {