got basic mshell client working -- still need detectfds and extra files support

This commit is contained in:
sawka
2022-06-24 13:25:09 -07:00
parent 0267836376
commit 5223760a76
6 changed files with 226 additions and 93 deletions
+6 -35
View File
@@ -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() {
+18 -14
View File
@@ -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)
+26 -16
View File
@@ -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
}
+80 -15
View File
@@ -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
}
+7 -1
View File
@@ -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
+89 -12
View File
@@ -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 {