checkpoint got stdout/stderr data packets working with new remote handler

This commit is contained in:
sawka
2022-06-23 12:48:45 -07:00
parent 766d19f1bc
commit c43d3ecc85
4 changed files with 329 additions and 93 deletions
+101 -15
View File
@@ -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())
+7 -10
View File
@@ -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) {
+41 -36
View File
@@ -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 {
+180 -32
View File
@@ -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
}