diff --git a/.gitignore b/.gitignore index 47932004..058324df 100644 --- a/.gitignore +++ b/.gitignore @@ -3,13 +3,17 @@ dist-dev/ node_modules/ *~ *.log +*.out out/ .DS_Store bin/ +waveshell/bin/ +wavesrv/bin/ dev-bin local-server-bin *.pw build/ *.dmg webshare/dist/ -webshare/dist-dev/ \ No newline at end of file +webshare/dist-dev/ + diff --git a/waveshell/go.mod b/waveshell/go.mod new file mode 100644 index 00000000..7fd6d4cd --- /dev/null +++ b/waveshell/go.mod @@ -0,0 +1,13 @@ +module github.com/commandlinedev/apishell + +go 1.18 + +require ( + github.com/alessio/shellescape v1.4.1 + github.com/creack/pty v1.1.18 + github.com/fsnotify/fsnotify v1.6.0 + github.com/google/uuid v1.3.0 + golang.org/x/mod v0.5.1 + golang.org/x/sys v0.10.0 + mvdan.cc/sh/v3 v3.7.0 +) diff --git a/waveshell/go.sum b/waveshell/go.sum new file mode 100644 index 00000000..9897e263 --- /dev/null +++ b/waveshell/go.sum @@ -0,0 +1,20 @@ +github.com/alessio/shellescape v1.4.1 h1:V7yhSDDn8LP4lc4jS8pFkt0zCnzVJlG5JXy9BVKJUX0= +github.com/alessio/shellescape v1.4.1/go.mod h1:PZAiSCk0LJaZkiCSkPv8qIobYglO3FPpyFjDCtHLS30= +github.com/creack/pty v1.1.18 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY= +github.com/creack/pty v1.1.18/go.mod h1:MOBLtS5ELjhRRrroQr9kyvTxUAFNvYEK993ew/Vr4O4= +github.com/frankban/quicktest v1.14.5 h1:dfYrrRyLtiqT9GyKXgdh+k4inNeTvmGbuSgZ3lx3GhA= +github.com/fsnotify/fsnotify v1.6.0 h1:n+5WquG0fcWoWp6xPWfHdbskMCQaFnG6PfBrh1Ky4HY= +github.com/fsnotify/fsnotify v1.6.0/go.mod h1:sl3t1tCWJFWoRz9R8WJCbQihKKwmorjAbSClcnxKAGw= +github.com/google/go-cmp v0.5.9 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38= +github.com/google/uuid v1.3.0 h1:t6JiXgmwXMjEs8VusXIJk2BXHsn+wx8BZdTaoZ5fu7I= +github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= +github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= +github.com/rogpeppe/go-internal v1.10.1-0.20230524175051-ec119421bb97 h1:3RPlVWzZ/PDqmVuf/FKHARG5EMid/tl7cv54Sw/QRVY= +golang.org/x/mod v0.5.1 h1:OJxoQ/rynoF0dcCdI7cLPktw/hR2cueqYfjm43oqK38= +golang.org/x/mod v0.5.1/go.mod h1:5OXOZSfqPIIbmVBIIKWRFfZjPR0E5r58TLhUjH0a2Ro= +golang.org/x/sys v0.0.0-20220908164124-27713097b956/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.10.0 h1:SqMFp9UcQJZa+pmYuAKjd9xq1f0j5rLcDIk0mj4qAsA= +golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +mvdan.cc/sh/v3 v3.7.0 h1:lSTjdP/1xsddtaKfGg7Myu7DnlHItd3/M2tomOcNNBg= +mvdan.cc/sh/v3 v3.7.0/go.mod h1:K2gwkaesF/D7av7Kxl0HbF5kGOd2ArupNTX3X44+8l8= diff --git a/waveshell/main-waveshell.go b/waveshell/main-waveshell.go new file mode 100644 index 00000000..36b8ce1c --- /dev/null +++ b/waveshell/main-waveshell.go @@ -0,0 +1,586 @@ +package main + +import ( + "bytes" + "fmt" + "os" + "strconv" + "strings" + "syscall" + "time" + + "github.com/commandlinedev/apishell/pkg/base" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/apishell/pkg/server" + "github.com/commandlinedev/apishell/pkg/shexec" + "golang.org/x/sys/unix" +) + +var BuildTime = "0" + +// func doMainRun(pk *packet.RunPacketType, sender *packet.PacketSender) { +// err := shexec.ValidateRunPacket(pk) +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("invalid run packet: %v", err)) +// return +// } +// fileNames, err := base.GetCommandFileNames(pk.CK) +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot get command file names: %v", err)) +// return +// } +// cmd, err := shexec.MakeRunnerExec(pk.CK) +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot make mshell command: %v", err)) +// return +// } +// cmdStdin, err := cmd.StdinPipe() +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot pipe stdin to command: %v", err)) +// return +// } +// // touch ptyout file (should exist for tailer to work correctly) +// ptyOutFd, err := os.OpenFile(fileNames.PtyOutFile, os.O_CREATE|os.O_TRUNC|os.O_APPEND|os.O_WRONLY, 0600) +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot open pty out file '%s': %v", fileNames.PtyOutFile, err)) +// return +// } +// ptyOutFd.Close() // just opened to create the file, can close right after +// runnerOutFd, err := os.OpenFile(fileNames.RunnerOutFile, os.O_CREATE|os.O_TRUNC|os.O_APPEND|os.O_WRONLY, 0600) +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("cannot open runner out file '%s': %v", fileNames.RunnerOutFile, err)) +// return +// } +// defer runnerOutFd.Close() +// cmd.Stdout = runnerOutFd +// cmd.Stderr = runnerOutFd +// err = cmd.Start() +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("error starting command: %v", err)) +// return +// } +// go func() { +// err = packet.SendPacket(cmdStdin, pk) +// if err != nil { +// sender.SendCKErrorPacket(pk.CK, fmt.Sprintf("error sending forked runner command: %v", err)) +// return +// } +// cmdStdin.Close() + +// // clean up zombies +// cmd.Wait() +// }() +// } + +// func doGetCmd(tailer *cmdtail.Tailer, pk *packet.GetCmdPacketType, sender *packet.PacketSender) error { +// err := tailer.AddWatch(pk) +// if err != nil { +// return err +// } +// return nil +// } + +// func doMain() { +// homeDir := base.GetHomeDir() +// err := os.Chdir(homeDir) +// if err != nil { +// packet.SendErrorPacket(os.Stdout, fmt.Sprintf("cannot change directory to $HOME '%s': %v", homeDir, err)) +// return +// } +// _, err = base.GetMShellPath() +// if err != nil { +// packet.SendErrorPacket(os.Stdout, err.Error()) +// return +// } +// packetParser := packet.MakePacketParser(os.Stdin) +// sender := packet.MakePacketSender(os.Stdout) +// tailer, err := cmdtail.MakeTailer(sender) +// if err != nil { +// packet.SendErrorPacket(os.Stdout, err.Error()) +// return +// } +// go tailer.Run() +// initPacket := shexec.MakeInitPacket() +// sender.SendPacket(initPacket) +// for pk := range packetParser.MainCh { +// if pk.GetType() == packet.RunPacketStr { +// doMainRun(pk.(*packet.RunPacketType), sender) +// continue +// } +// if pk.GetType() == packet.GetCmdPacketStr { +// err = doGetCmd(tailer, pk.(*packet.GetCmdPacketType), sender) +// if err != nil { +// errPk := packet.MakeErrorPacket(err.Error()) +// sender.SendPacket(errPk) +// continue +// } +// continue +// } +// if pk.GetType() == packet.CdPacketStr { +// cdPacket := pk.(*packet.CdPacketType) +// err := os.Chdir(cdPacket.Dir) +// resp := packet.MakeResponsePacket(cdPacket.ReqId) +// if err != nil { +// resp.Error = err.Error() +// } else { +// resp.Success = true +// } +// sender.SendPacket(resp) +// continue +// } +// if pk.GetType() == packet.ErrorPacketStr { +// errPk := pk.(*packet.ErrorPacketType) +// errPk.Error = "invalid packet sent to mshell: " + errPk.Error +// sender.SendPacket(errPk) +// continue +// } +// sender.SendErrorPacket(fmt.Sprintf("invalid packet '%s' sent to mshell", pk.GetType())) +// } +// } + +func readFullRunPacket(packetParser *packet.PacketParser) (*packet.RunPacketType, error) { + rpb := packet.MakeRunPacketBuilder() + for pk := range packetParser.MainCh { + ok, runPacket := rpb.ProcessPacket(pk) + if runPacket != nil { + return runPacket, nil + } + if !ok { + return nil, fmt.Errorf("invalid packet '%s' sent to mshell", pk.GetType()) + } + } + return nil, fmt.Errorf("no run packet received") +} + +func handleSingle(fromServer bool) { + packetParser := packet.MakePacketParser(os.Stdin, false) + sender := packet.MakePacketSender(os.Stdout, nil) + defer func() { + sender.Close() + sender.WaitForDone() + }() + initPacket := shexec.MakeInitPacket() + sender.SendPacket(initPacket) + if len(os.Args) >= 3 && os.Args[2] == "--version" { + return + } + runPacket, err := readFullRunPacket(packetParser) + if err != nil { + sender.SendErrorResponse(runPacket.ReqId, err) + return + } + err = shexec.ValidateRunPacket(runPacket) + if err != nil { + sender.SendErrorResponse(runPacket.ReqId, err) + return + } + if fromServer { + err = runPacket.CK.Validate("run packet") + if err != nil { + sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("run packets from server must have a CK: %v", err)) + } + } + if runPacket.Detached { + cmd, startPk, err := shexec.RunCommandDetached(runPacket, sender) + if err != nil { + sender.SendErrorResponse(runPacket.ReqId, err) + return + } + sender.SendPacket(startPk) + sender.Close() + sender.WaitForDone() + cmd.DetachedWait(startPk) + return + } else { + shexec.IgnoreSigPipe() + ticker := time.NewTicker(1 * time.Minute) + go func() { + for range ticker.C { + // this will let the command detect when the server has gone away + // that will then trigger cmd.SendHup() to send SIGHUP to the exec'ed process + sender.SendPacket(packet.MakePingPacket()) + } + }() + defer ticker.Stop() + cmd, err := shexec.RunCommandSimple(runPacket, sender, true) + if err != nil { + sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("error running command: %w", err)) + return + } + defer cmd.Close() + startPacket := cmd.MakeCmdStartPacket(runPacket.ReqId) + sender.SendPacket(startPacket) + go func() { + exitErr := sender.WaitForDone() + if exitErr != nil { + base.Logf("I/O error talking to server, sending SIGHUP to children\n") + cmd.SendSignal(syscall.SIGHUP) + } + }() + cmd.RunRemoteIOAndWait(packetParser, sender) + return + } +} + +func detectOpenFds() ([]packet.RemoteFd, error) { + var fds []packet.RemoteFd + for fdNum := 3; fdNum <= 64; fdNum++ { + flags, err := unix.FcntlInt(uintptr(fdNum), unix.F_GETFL, 0) + if err != nil { + continue + } + flags = flags & 3 + rfd := packet.RemoteFd{FdNum: fdNum} + if flags&2 == 2 { + return nil, fmt.Errorf("invalid fd=%d, mshell does not support fds open for reading and writing", fdNum) + } + if flags&1 == 1 { + rfd.Write = true + } else { + rfd.Read = true + } + fds = append(fds, rfd) + } + return fds, nil +} + +func parseInstallOpts() (*shexec.InstallOpts, error) { + opts := &shexec.InstallOpts{} + iter := base.MakeOptsIter(os.Args[2:]) // first arg is --install + for iter.HasNext() { + argStr := iter.Next() + found, err := tryParseSSHOpt(iter, &opts.SSHOpts) + if err != nil { + return nil, err + } + if found { + continue + } + if argStr == "--detect" { + opts.Detect = true + continue + } + if base.IsOption(argStr) { + return nil, fmt.Errorf("invalid option '%s' passed to mshell --install", argStr) + } + opts.ArchStr = argStr + break + } + return opts, nil +} + +func tryParseSSHOpt(iter *base.OptsIter, sshOpts *shexec.SSHOpts) (bool, error) { + argStr := iter.Current() + if argStr == "--ssh" { + if !iter.IsNextPlain() { + return false, fmt.Errorf("'--ssh [user@host]' missing host") + } + sshOpts.SSHHost = iter.Next() + return true, nil + } + if argStr == "--ssh-opts" { + if !iter.HasNext() { + return false, fmt.Errorf("'--ssh-opts [options]' missing options") + } + sshOpts.SSHOptsStr = iter.Next() + return true, nil + } + if argStr == "-i" { + if !iter.IsNextPlain() { + return false, fmt.Errorf("-i [identity-file]' missing file") + } + sshOpts.SSHIdentity = iter.Next() + return true, nil + } + if argStr == "-l" { + if !iter.IsNextPlain() { + return false, fmt.Errorf("-l [user]' missing user") + } + sshOpts.SSHUser = iter.Next() + return true, nil + } + if argStr == "-p" { + if !iter.IsNextPlain() { + return false, fmt.Errorf("-p [port]' missing port") + } + nextArgStr := iter.Next() + portVal, err := strconv.Atoi(nextArgStr) + if err != nil { + return false, fmt.Errorf("-p [port]' invalid port: %v", err) + } + if portVal <= 0 { + return false, fmt.Errorf("-p [port]' invalid port: %d", portVal) + } + sshOpts.SSHPort = portVal + return true, nil + } + return false, nil +} + +func parseClientOpts() (*shexec.ClientOpts, error) { + opts := &shexec.ClientOpts{} + iter := base.MakeOptsIter(os.Args[1:]) + for iter.HasNext() { + argStr := iter.Next() + found, err := tryParseSSHOpt(iter, &opts.SSHOpts) + if err != nil { + return nil, err + } + if found { + continue + } + if argStr == "--cwd" { + if !iter.IsNextPlain() { + return nil, fmt.Errorf("'--cwd [dir]' missing directory") + } + opts.Cwd = iter.Next() + continue + } + if argStr == "--detach" { + opts.Detach = true + continue + } + if argStr == "--pty" { + opts.UsePty = true + continue + } + if argStr == "--debug" { + opts.Debug = true + continue + } + if argStr == "--sudo" { + opts.Sudo = true + continue + } + if argStr == "--sudo-with-password" { + if !iter.HasNext() { + return nil, fmt.Errorf("'--sudo-with-password [pw]', missing password") + } + opts.Sudo = true + opts.SudoWithPass = true + opts.SudoPw = iter.Next() + continue + } + if argStr == "--sudo-with-passfile" { + if !iter.IsNextPlain() { + return nil, fmt.Errorf("'--sudo-with-passfile [file]', missing file") + } + opts.Sudo = true + opts.SudoWithPass = true + fileName := iter.Next() + contents, err := os.ReadFile(fileName) + if err != nil { + return nil, fmt.Errorf("cannot read --sudo-with-passfile file '%s': %w", fileName, err) + } + if newlineIdx := bytes.Index(contents, []byte{'\n'}); newlineIdx != -1 { + contents = contents[0:newlineIdx] + } + opts.SudoPw = string(contents) + "\n" + continue + } + if argStr == "--" { + if !iter.HasNext() { + return nil, fmt.Errorf("'--' should be followed by command") + } + opts.Command = strings.Join(iter.Rest(), " ") + break + } + return nil, fmt.Errorf("invalid option '%s' passed to mshell", argStr) + } + return opts, nil +} + +func handleClient() (int, error) { + opts, err := parseClientOpts() + if err != nil { + return 1, fmt.Errorf("parsing opts: %w", err) + } + if opts.Debug { + packet.GlobalDebug = true + } + if opts.Command == "" { + return 1, fmt.Errorf("no [command] specified. [command] follows '--' option (see usage)") + } + fds, err := detectOpenFds() + if err != nil { + return 1, err + } + opts.Fds = fds + err = shexec.ValidateRemoteFds(opts.Fds) + if err != nil { + return 1, err + } + runPacket, err := opts.MakeRunPacket() // modifies opts + if err != nil { + return 1, err + } + if runPacket.Detached { + return 1, fmt.Errorf("cannot run detached command from command line client") + } + donePacket, err := shexec.RunClientSSHCommandAndWait(runPacket, shexec.StdContext{}, opts.SSHOpts, nil, opts.Debug) + if err != nil { + return 1, err + } + return donePacket.ExitCode, nil +} + +func handleInstall() (int, error) { + opts, err := parseInstallOpts() + if err != nil { + return 1, fmt.Errorf("parsing opts: %w", err) + } + if opts.SSHOpts.SSHHost == "" { + return 1, fmt.Errorf("cannot install without '--ssh user@host' option") + } + if opts.Detect && opts.ArchStr != "" { + return 1, fmt.Errorf("cannot supply both --detect and arch '%s'", opts.ArchStr) + } + if opts.ArchStr == "" && !opts.Detect { + return 1, fmt.Errorf("must supply an arch string or '--detect' to auto detect") + } + if opts.ArchStr != "" { + fullArch := opts.ArchStr + fields := strings.SplitN(fullArch, ".", 2) + if len(fields) != 2 { + return 1, fmt.Errorf("invalid arch format '%s' passed to mshell --install", fullArch) + } + goos, goarch := fields[0], fields[1] + if !base.ValidGoArch(goos, goarch) { + return 1, fmt.Errorf("invalid arch '%s' passed to mshell --install", fullArch) + } + optName := base.GoArchOptFile(base.MShellVersion, goos, goarch) + _, err = os.Stat(optName) + if err != nil { + return 1, fmt.Errorf("cannot install mshell to remote host, cannot read '%s': %w", optName, err) + } + opts.OptName = optName + } + err = shexec.RunInstallFromOpts(opts) + if err != nil { + return 1, err + } + return 0, nil +} + +func handleEnv() (int, error) { + cwd, err := os.Getwd() + if err != nil { + return 1, err + } + fmt.Printf("%s\x00\x00", cwd) + fullEnv := os.Environ() + var linePrinted bool + for _, envLine := range fullEnv { + if envLine != "" { + fmt.Printf("%s\x00", envLine) + linePrinted = true + } + } + if linePrinted { + fmt.Printf("\x00") + } else { + fmt.Printf("\x00\x00") + } + return 0, nil +} + +func handleUsage() { + usage := ` +Client Usage: mshell [opts] --ssh user@host -- [command] + +mshell multiplexes input and output streams to a remote command over ssh. + +Options: + -i [identity-file] - used to set '-i' option for ssh command + -l [user] - used to set '-l' option for ssh command + --cwd [dir] - execute remote command in [dir] + --ssh-opts [opts] - addition options to pass to ssh command + [command] - the remote command to execute + +Sudo Options: + --sudo - use only if sudo never requires a password + --sudo-with-password [pw] - not recommended, use --sudo-with-passfile if possible + --sudo-with-passfile [file] + +Sudo options allow you to run the given command using "sudo". The first +option only works when you can sudo without a password. Your password will be passed +securely through a high numbered fd to "sudo -S". Note that to use high numbered +file descriptors with sudo, you will need to add this line to your /etc/sudoers file: + Defaults closefrom_override +See full documentation for more details. + +Examples: + # execute a python script remotely, with stdin still hooked up correctly + mshell --cwd "~/work" -i key.pem --ssh ubuntu@somehost -- "python3 /dev/fd/4" 4< myscript.py + + # capture multiple outputs + mshell --ssh ubuntu@test -- "cat file1.txt > /dev/fd/3; cat file2.txt > /dev/fd/4" 3> file1.txt 4> file2.txt + + # execute a script, catpure stdout/stderr in fd-3 and fd-4 + # useful if you need to see stdout for interacting with ssh (password or host auth) + mshell --ssh user@host -- "test.sh > /dev/fd/3 2> /dev/fd/4" 3> test.stdout 4> test.stderr + + # run a script as root (via sudo), capture output + mshell --sudo-with-passfile pw.txt --ssh ubuntu@somehost -- "python3 /dev/fd/3 > /dev/fd/4" 3< myscript.py 4> script-output.txt < script-input.txt +` + fmt.Printf("%s\n\n", strings.TrimSpace(usage)) +} + +func main() { + base.SetBuildTime(BuildTime) + if len(os.Args) == 1 { + handleUsage() + return + } + firstArg := os.Args[1] + if firstArg == "--help" { + handleUsage() + return + } else if firstArg == "--version" { + fmt.Printf("mshell %s+%s\n", base.MShellVersion, base.BuildTime) + return + } else if firstArg == "--test-env" { + state, err := shexec.GetShellState() + if state != nil { + + } + if err != nil { + fmt.Fprintf(os.Stderr, "[error] %v\n", err) + os.Exit(1) + } + } else if firstArg == "--single" { + base.InitDebugLog("single") + handleSingle(false) + return + } else if firstArg == "--single-from-server" { + base.InitDebugLog("single") + handleSingle(true) + return + } else if firstArg == "--server" { + base.InitDebugLog("server") + rtnCode, err := server.RunServer() + if err != nil { + fmt.Fprintf(os.Stderr, "[error] %v\n", err) + } + if rtnCode != 0 { + os.Exit(rtnCode) + } + return + } else if firstArg == "--install" { + rtnCode, err := handleInstall() + if err != nil { + fmt.Fprintf(os.Stderr, "[error] %v\n", err) + } + os.Exit(rtnCode) + return + } else { + rtnCode, err := handleClient() + if err != nil { + fmt.Fprintf(os.Stderr, "[error] %v\n", err) + } + if rtnCode != 0 { + os.Exit(rtnCode) + } + return + } +} diff --git a/waveshell/pkg/base/base.go b/waveshell/pkg/base/base.go new file mode 100644 index 00000000..d9204f4f --- /dev/null +++ b/waveshell/pkg/base/base.go @@ -0,0 +1,381 @@ +package base + +import ( + "errors" + "fmt" + "io" + "io/fs" + "log" + "os" + "os/exec" + "path" + "path/filepath" + "strings" + "sync" + + "github.com/google/uuid" + "golang.org/x/mod/semver" +) + +const HomeVarName = "HOME" +const DefaultMShellHome = "~/.mshell" +const DefaultMShellName = "mshell" +const MShellPathVarName = "MSHELL_PATH" +const MShellHomeVarName = "MSHELL_HOME" +const MShellInstallBinVarName = "MSHELL_INSTALLBIN_PATH" +const SSHCommandVarName = "SSH_COMMAND" +const MShellDebugVarName = "MSHELL_DEBUG" +const SessionsDirBaseName = "sessions" +const MShellVersion = "v0.3.0" +const RemoteIdFile = "remoteid" +const DefaultMShellInstallBinDir = "/opt/mshell/bin" +const LogFileName = "mshell.log" +const ForceDebugLog = false + +const DebugFlag_LogRcFile = "logrc" +const LogRcFileName = "debug.rcfile" + +var sessionDirCache = make(map[string]string) +var baseLock = &sync.Mutex{} +var DebugLogEnabled = false +var DebugLogger *log.Logger +var BuildTime string = "0" + +type CommandFileNames struct { + PtyOutFile string + StdinFifo string + RunnerOutFile string +} + +type CommandKey string + +func SetBuildTime(build string) { + BuildTime = build +} + +func MakeCommandKey(sessionId string, cmdId string) CommandKey { + if sessionId == "" && cmdId == "" { + return CommandKey("") + } + return CommandKey(fmt.Sprintf("%s/%s", sessionId, cmdId)) +} + +func (ckey CommandKey) IsEmpty() bool { + return string(ckey) == "" +} + +func Logf(fmtStr string, args ...interface{}) { + if (!DebugLogEnabled && !ForceDebugLog) || DebugLogger == nil { + return + } + DebugLogger.Printf(fmtStr, args...) +} + +func InitDebugLog(prefix string) { + homeDir := GetMShellHomeDir() + err := os.MkdirAll(homeDir, 0777) + if err != nil { + return + } + logFile := path.Join(homeDir, LogFileName) + fd, err := os.OpenFile(logFile, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0600) + if err != nil { + return + } + DebugLogger = log.New(fd, prefix+" ", log.LstdFlags) + Logf("logger initialized\n") +} + +func SetEnableDebugLog(enable bool) { + DebugLogEnabled = enable +} + +// deprecated (use GetGroupId instead) +func (ckey CommandKey) GetSessionId() string { + return ckey.GetGroupId() +} + +func (ckey CommandKey) GetGroupId() string { + slashIdx := strings.Index(string(ckey), "/") + if slashIdx == -1 { + return "" + } + return string(ckey[0:slashIdx]) +} + +func (ckey CommandKey) GetCmdId() string { + slashIdx := strings.Index(string(ckey), "/") + if slashIdx == -1 { + return "" + } + return string(ckey[slashIdx+1:]) +} + +func (ckey CommandKey) Split() (string, string) { + fields := strings.SplitN(string(ckey), "/", 2) + if len(fields) < 2 { + return "", "" + } + return fields[0], fields[1] +} + +func (ckey CommandKey) Validate(typeStr string) error { + if typeStr == "" { + typeStr = "ck" + } + if ckey == "" { + return fmt.Errorf("%s has empty commandkey", typeStr) + } + sessionId, cmdId := ckey.Split() + if sessionId == "" { + return fmt.Errorf("%s does not have sessionid", typeStr) + } + _, err := uuid.Parse(sessionId) + if err != nil { + return fmt.Errorf("%s has invalid sessionid '%s'", typeStr, sessionId) + } + if cmdId == "" { + return fmt.Errorf("%s does not have cmdid", typeStr) + } + _, err = uuid.Parse(cmdId) + if err != nil { + return fmt.Errorf("%s has invalid cmdid '%s'", typeStr, cmdId) + } + return nil +} + +func HasDebugFlag(envMap map[string]string, flagName string) bool { + msDebug := envMap[MShellDebugVarName] + flags := strings.Split(msDebug, ",") + Logf("hasdebugflag[%s]: %s [%#v]\n", flagName, msDebug, flags) + for _, flag := range flags { + if strings.TrimSpace(flag) == flagName { + return true + } + } + return false +} + +func GetDebugRcFileName() string { + msHome := GetMShellHomeDir() + return path.Join(msHome, LogRcFileName) +} + +func GetHomeDir() string { + homeVar := os.Getenv(HomeVarName) + if homeVar == "" { + return "/" + } + return homeVar +} + +func GetMShellHomeDir() string { + homeVar := os.Getenv(MShellHomeVarName) + if homeVar != "" { + return homeVar + } + return ExpandHomeDir(DefaultMShellHome) +} + +func GetCommandFileNames(ck CommandKey) (*CommandFileNames, error) { + if err := ck.Validate("ck"); err != nil { + return nil, fmt.Errorf("cannot get command files: %w", err) + } + sessionId, cmdId := ck.Split() + sdir, err := EnsureSessionDir(sessionId) + if err != nil { + return nil, err + } + base := path.Join(sdir, cmdId) + return &CommandFileNames{ + PtyOutFile: base + ".ptyout", + StdinFifo: base + ".stdin", + RunnerOutFile: base + ".runout", + }, nil +} + +func CleanUpCmdFiles(sessionId string, cmdId string) error { + if cmdId == "" { + return fmt.Errorf("bad cmdid, cannot clean up") + } + sdir, err := EnsureSessionDir(sessionId) + if err != nil { + return err + } + cmdFileGlob := path.Join(sdir, cmdId+".*") + matches, err := filepath.Glob(cmdFileGlob) + if err != nil { + return err + } + for _, file := range matches { + rmErr := os.Remove(file) + if err == nil && rmErr != nil { + err = rmErr + } + } + return err +} + +func GetSessionsDir() string { + mhome := GetMShellHomeDir() + sdir := path.Join(mhome, SessionsDirBaseName) + return sdir +} + +func EnsureSessionDir(sessionId string) (string, error) { + if sessionId == "" { + return "", fmt.Errorf("Bad sessionid, cannot be empty") + } + baseLock.Lock() + sdir, ok := sessionDirCache[sessionId] + baseLock.Unlock() + if ok { + return sdir, nil + } + mhome := GetMShellHomeDir() + sdir = path.Join(mhome, SessionsDirBaseName, sessionId) + info, err := os.Stat(sdir) + if errors.Is(err, fs.ErrNotExist) { + err = os.MkdirAll(sdir, 0777) + if err != nil { + return "", fmt.Errorf("cannot make mshell session directory[%s]: %w", sdir, err) + } + info, err = os.Stat(sdir) + } + if err != nil { + return "", err + } + if !info.IsDir() { + return "", fmt.Errorf("session dir '%s' must be a directory", sdir) + } + baseLock.Lock() + sessionDirCache[sessionId] = sdir + baseLock.Unlock() + return sdir, nil +} + +func GetMShellPath() (string, error) { + msPath := os.Getenv(MShellPathVarName) // use MSHELL_PATH + if msPath != "" { + return exec.LookPath(msPath) + } + mhome := GetMShellHomeDir() + userMShellPath := path.Join(mhome, DefaultMShellName) // look in ~/.mshell + msPath, err := exec.LookPath(userMShellPath) + if err == nil { + return msPath, nil + } + return exec.LookPath(DefaultMShellName) // standard path lookup for 'mshell' +} + +func GetMShellSessionsDir() (string, error) { + mhome := GetMShellHomeDir() + return path.Join(mhome, SessionsDirBaseName), nil +} + +func ExpandHomeDir(pathStr string) string { + if pathStr != "~" && !strings.HasPrefix(pathStr, "~/") { + return pathStr + } + homeDir := GetHomeDir() + if pathStr == "~" { + return homeDir + } + return path.Join(homeDir, pathStr[2:]) +} + +func ValidGoArch(goos string, goarch string) bool { + return (goos == "darwin" || goos == "linux") && (goarch == "amd64" || goarch == "arm64") +} + +func GoArchOptFile(version string, goos string, goarch string) string { + installBinDir := os.Getenv(MShellInstallBinVarName) + if installBinDir == "" { + installBinDir = DefaultMShellInstallBinDir + } + versionStr := semver.MajorMinor(version) + if versionStr == "" { + versionStr = "unknown" + } + binBaseName := fmt.Sprintf("mshell-%s-%s.%s", versionStr, goos, goarch) + return fmt.Sprintf(path.Join(installBinDir, binBaseName)) +} + +func MShellBinaryFromOptDir(version string, goos string, goarch string) (io.ReadCloser, error) { + if !ValidGoArch(goos, goarch) { + return nil, fmt.Errorf("invalid goos/goarch combination: %s/%s", goos, goarch) + } + versionStr := semver.MajorMinor(version) + if versionStr == "" { + return nil, fmt.Errorf("invalid mshell version: %q", version) + } + fileName := GoArchOptFile(version, goos, goarch) + fd, err := os.Open(fileName) + if err != nil { + return nil, fmt.Errorf("cannot open mshell binary %q: %v", fileName, err) + } + return fd, nil +} + +func GetRemoteId() (string, error) { + mhome := GetMShellHomeDir() + homeInfo, err := os.Stat(mhome) + if errors.Is(err, fs.ErrNotExist) { + err = os.MkdirAll(mhome, 0777) + if err != nil { + return "", fmt.Errorf("cannot make mshell home directory[%s]: %w", mhome, err) + } + homeInfo, err = os.Stat(mhome) + } + if err != nil { + return "", fmt.Errorf("cannot stat mshell home directory[%s]: %w", mhome, err) + } + if !homeInfo.IsDir() { + return "", fmt.Errorf("mshell home directory[%s] is not a directory", mhome) + } + remoteIdFile := path.Join(mhome, RemoteIdFile) + fd, err := os.Open(remoteIdFile) + if errors.Is(err, fs.ErrNotExist) { + // write the file + remoteId := uuid.New().String() + err = os.WriteFile(remoteIdFile, []byte(remoteId), 0644) + if err != nil { + return "", fmt.Errorf("cannot write remoteid to '%s': %w", remoteIdFile, err) + } + return remoteId, nil + } else if err != nil { + return "", fmt.Errorf("cannot read remoteid file '%s': %w", remoteIdFile, err) + } else { + defer fd.Close() + contents, err := io.ReadAll(fd) + if err != nil { + return "", fmt.Errorf("cannot read remoteid file '%s': %w", remoteIdFile, err) + } + uuidStr := string(contents) + _, err = uuid.Parse(uuidStr) + if err != nil { + return "", fmt.Errorf("invalid uuid read from '%s': %w", remoteIdFile, err) + } + return uuidStr, nil + } +} + +func BoundInt(ival int, minVal int, maxVal int) int { + if ival < minVal { + return minVal + } + if ival > maxVal { + return maxVal + } + return ival +} + +func BoundInt64(ival int64, minVal int64, maxVal int64) int64 { + if ival < minVal { + return minVal + } + if ival > maxVal { + return maxVal + } + return ival +} diff --git a/waveshell/pkg/base/optsiter.go b/waveshell/pkg/base/optsiter.go new file mode 100644 index 00000000..f607c670 --- /dev/null +++ b/waveshell/pkg/base/optsiter.go @@ -0,0 +1,47 @@ +package base + +import "strings" + +type OptsIter struct { + Pos int + Opts []string +} + +func MakeOptsIter(opts []string) *OptsIter { + return &OptsIter{Opts: opts} +} + +func IsOption(argStr string) bool { + return strings.HasPrefix(argStr, "-") && argStr != "-" && !strings.HasPrefix(argStr, "-/") +} + +func (iter *OptsIter) HasNext() bool { + return iter.Pos <= len(iter.Opts)-1 +} + +func (iter *OptsIter) IsNextPlain() bool { + if !iter.HasNext() { + return false + } + return !IsOption(iter.Opts[iter.Pos]) +} + +func (iter *OptsIter) Next() string { + if iter.Pos >= len(iter.Opts) { + return "" + } + rtn := iter.Opts[iter.Pos] + iter.Pos++ + return rtn +} + +func (iter *OptsIter) Current() string { + if iter.Pos == 0 { + return "" + } + return iter.Opts[iter.Pos-1] +} + +func (iter *OptsIter) Rest() []string { + return iter.Opts[iter.Pos:] +} diff --git a/waveshell/pkg/binpack/binpack.go b/waveshell/pkg/binpack/binpack.go new file mode 100644 index 00000000..67e25337 --- /dev/null +++ b/waveshell/pkg/binpack/binpack.go @@ -0,0 +1,127 @@ +package binpack + +import ( + "encoding/binary" + "encoding/json" + "fmt" + "io" +) + +type Unpacker struct { + R FullByteReader + Err error +} + +type FullByteReader interface { + io.ByteReader + io.Reader +} + +func PackValue(w io.Writer, barr []byte) error { + viBuf := make([]byte, binary.MaxVarintLen64) + viLen := binary.PutUvarint(viBuf, uint64(len(barr))) + _, err := w.Write(viBuf[0:viLen]) + if err != nil { + return err + } + if len(barr) > 0 { + _, err = w.Write(barr) + if err != nil { + return err + } + } + return nil +} + +func PackStrArr(w io.Writer, strs []string) error { + barr, err := json.Marshal(strs) + if err != nil { + return err + } + return PackValue(w, barr) +} + +func PackInt(w io.Writer, ival int) error { + viBuf := make([]byte, binary.MaxVarintLen64) + l := binary.PutUvarint(viBuf, uint64(ival)) + _, err := w.Write(viBuf[0:l]) + return err +} + +func UnpackValue(r FullByteReader) ([]byte, error) { + lenVal, err := binary.ReadUvarint(r) + if err != nil { + return nil, err + } + if lenVal == 0 { + return nil, nil + } + rtnBuf := make([]byte, int(lenVal)) + _, err = io.ReadFull(r, rtnBuf) + if err != nil { + return nil, err + } + return rtnBuf, nil +} + +func UnpackStrArr(r FullByteReader) ([]string, error) { + barr, err := UnpackValue(r) + if err != nil { + return nil, err + } + var strs []string + err = json.Unmarshal(barr, &strs) + if err != nil { + return nil, err + } + return strs, nil +} + +func UnpackInt(r io.ByteReader) (int, error) { + ival64, err := binary.ReadVarint(r) + if err != nil { + return 0, err + } + return int(ival64), nil +} + +func (u *Unpacker) UnpackValue(name string) []byte { + if u.Err != nil { + return nil + } + rtn, err := UnpackValue(u.R) + if err != nil { + u.Err = fmt.Errorf("cannot unpack %s: %v", name, err) + } + return rtn +} + +func (u *Unpacker) UnpackInt(name string) int { + if u.Err != nil { + return 0 + } + rtn, err := UnpackInt(u.R) + if err != nil { + u.Err = fmt.Errorf("cannot unpack %s: %v", name, err) + } + return rtn +} + +func (u *Unpacker) UnpackStrArr(name string) []string { + if u.Err != nil { + return nil + } + rtn, err := UnpackStrArr(u.R) + if err != nil { + u.Err = fmt.Errorf("cannot unpack %s: %v", name, err) + } + return rtn +} + +func (u *Unpacker) Error() error { + return u.Err +} + +func MakeUnpacker(r FullByteReader) *Unpacker { + return &Unpacker{R: r} +} diff --git a/waveshell/pkg/cirfile/cirfile.go b/waveshell/pkg/cirfile/cirfile.go new file mode 100644 index 00000000..fef9a764 --- /dev/null +++ b/waveshell/pkg/cirfile/cirfile.go @@ -0,0 +1,570 @@ +package cirfile + +import ( + "context" + "fmt" + "io" + "os" + "syscall" + "time" +) + +// CBUF[version] [maxsize] [fileoffset] [startpos] [endpos] +const HeaderFmt1 = "CBUF%02d %19d %19d %19d %19d\n" // 87 bytes +const HeaderLen = 256 // set to 256 for future expandability +const FullHeaderFmt = "%-255s\n" // 256 bytes (255 + newline) +const CurrentVersion = 1 +const FilePosEmpty = -1 // sentinel, if startpos is set to -1, file is empty + +const InitialLockDelay = 10 * time.Millisecond +const InitialLockTries = 5 +const LockDelay = 100 * time.Millisecond + +// File objects are *not* multithread safe, operations must be externally synchronized +type File struct { + OSFile *os.File + Version byte + MaxSize int64 + FileOffset int64 + StartPos int64 + EndPos int64 + FileDataSize int64 // size of data (does not include header size) + FlockStatus int +} + +type Stat struct { + Location string + Version byte + MaxSize int64 + FileOffset int64 + DataSize int64 +} + +func (f *File) flock(ctx context.Context, lockType int) error { + err := syscall.Flock(int(f.OSFile.Fd()), lockType|syscall.LOCK_NB) + if err == nil { + f.FlockStatus = lockType + return nil + } + if err != syscall.EWOULDBLOCK { + return err + } + if ctx == nil { + return syscall.EWOULDBLOCK + } + // busy-wait with context + numWaits := 0 + for { + numWaits++ + var timeout time.Duration + if numWaits <= InitialLockTries { + timeout = InitialLockDelay + } else { + timeout = LockDelay + } + select { + case <-time.After(timeout): + break + case <-ctx.Done(): + return ctx.Err() + } + err = syscall.Flock(int(f.OSFile.Fd()), lockType|syscall.LOCK_NB) + if err == nil { + f.FlockStatus = lockType + return nil + } + if err != syscall.EWOULDBLOCK { + return err + } + } + return fmt.Errorf("could not acquire lock") +} + +func (f *File) unflock() { + if f.FlockStatus != 0 { + syscall.Flock(int(f.OSFile.Fd()), syscall.LOCK_UN) // ignore error (nothing to do about it anyway) + f.FlockStatus = 0 + } + return +} + +// does not read metadata because locking could block/fail. we want to be able +// to return a valid file struct without blocking. +func OpenCirFile(fileName string) (*File, error) { + fd, err := os.OpenFile(fileName, os.O_RDWR, 0777) + if err != nil { + return nil, err + } + finfo, err := fd.Stat() + if err != nil { + return nil, err + } + if finfo.Size() < HeaderLen { + return nil, fmt.Errorf("invalid cirfile, file length[%d] less than HeaderLen[%d]", finfo.Size(), HeaderLen) + } + rtn := &File{OSFile: fd} + return rtn, nil +} + +func StatCirFile(ctx context.Context, fileName string) (*Stat, error) { + file, err := OpenCirFile(fileName) + if err != nil { + return nil, err + } + defer file.Close() + fileOffset, dataSize, err := file.GetStartOffsetAndSize(ctx) + if err != nil { + return nil, err + } + return &Stat{ + Location: fileName, + Version: file.Version, + MaxSize: file.MaxSize, + FileOffset: fileOffset, + DataSize: dataSize, + }, nil +} + +// if the file already exists, it is an error. +// there is a race condition if two goroutines try to create the same file between Stat() and Create(), so +// they both might get no error, but only one file will be valid. if this is a concern, this call +// should be externally synchronized. +func CreateCirFile(fileName string, maxSize int64) (*File, error) { + if maxSize <= 0 { + return nil, fmt.Errorf("invalid maxsize[%d]", maxSize) + } + _, err := os.Stat(fileName) + if err == nil { + return nil, fmt.Errorf("file[%s] already exists", fileName) + } + if !os.IsNotExist(err) { + return nil, fmt.Errorf("cannot stat: %w", err) + } + fd, err := os.Create(fileName) + if err != nil { + return nil, err + } + rtn := &File{OSFile: fd, Version: CurrentVersion, MaxSize: maxSize, StartPos: FilePosEmpty} + err = rtn.flock(nil, syscall.LOCK_EX) + if err != nil { + return nil, err + } + defer rtn.unflock() + err = rtn.writeMeta() + if err != nil { + return nil, err + } + return rtn, nil +} + +func (f *File) Close() error { + return f.OSFile.Close() +} + +func (f *File) ReadMeta(ctx context.Context) error { + err := f.flock(ctx, syscall.LOCK_SH) + if err != nil { + return err + } + defer f.unflock() + return f.readMeta() +} + +func (f *File) hasShLock() bool { + return f.FlockStatus == syscall.LOCK_EX || f.FlockStatus == syscall.LOCK_SH +} + +func (f *File) hasExLock() bool { + return f.FlockStatus == syscall.LOCK_EX +} + +func (f *File) readMeta() error { + if f.OSFile == nil { + return fmt.Errorf("no *os.File") + } + if !f.hasShLock() { + return fmt.Errorf("writeMeta must hold LOCK_SH") + } + _, err := f.OSFile.Seek(0, 0) + if err != nil { + return fmt.Errorf("cannot seek file: %w", err) + } + finfo, err := f.OSFile.Stat() + if err != nil { + return fmt.Errorf("cannot stat file: %w", err) + } + if finfo.Size() < 256 { + return fmt.Errorf("invalid cbuf file size[%d] < 256", finfo.Size()) + } + f.FileDataSize = finfo.Size() - 256 + buf := make([]byte, 256) + _, err = io.ReadFull(f.OSFile, buf) + if err != nil { + return fmt.Errorf("error reading header: %w", err) + } + // currently only one version, so we don't need to have special logic here yet + _, err = fmt.Sscanf(string(buf), HeaderFmt1, &f.Version, &f.MaxSize, &f.FileOffset, &f.StartPos, &f.EndPos) + if err != nil { + return fmt.Errorf("sscanf error: %w", err) + } + if f.Version != CurrentVersion { + return fmt.Errorf("invalid cbuf version[%d]", f.Version) + } + // possible incomplete write, fix start/end pos to be within filesize + if f.FileDataSize == 0 { + f.StartPos = FilePosEmpty + f.EndPos = 0 + } else if f.StartPos >= f.FileDataSize && f.EndPos >= f.FileDataSize { + f.StartPos = FilePosEmpty + f.EndPos = 0 + } else if f.StartPos >= f.FileDataSize { + f.StartPos = 0 + } else if f.EndPos >= f.FileDataSize { + f.EndPos = f.FileDataSize - 1 + } + if f.MaxSize <= 0 || f.FileOffset < 0 || (f.StartPos < 0 && f.StartPos != FilePosEmpty) || f.StartPos >= f.MaxSize || f.EndPos < 0 || f.EndPos >= f.MaxSize { + return fmt.Errorf("invalid cbuf metadata version[%d] filedatasize[%d] maxsize[%d] fileoffset[%d] startpos[%d] endpos[%d]", f.Version, f.FileDataSize, f.MaxSize, f.FileOffset, f.StartPos, f.EndPos) + } + return nil +} + +// no error checking of meta values +func (f *File) writeMeta() error { + if f.OSFile == nil { + return fmt.Errorf("no *os.File") + } + if !f.hasExLock() { + return fmt.Errorf("writeMeta must hold LOCK_EX") + } + _, err := f.OSFile.Seek(0, 0) + if err != nil { + return fmt.Errorf("cannot seek file: %w", err) + } + metaStr := fmt.Sprintf(HeaderFmt1, f.Version, f.MaxSize, f.FileOffset, f.StartPos, f.EndPos) + fullMetaStr := fmt.Sprintf(FullHeaderFmt, metaStr) + _, err = f.OSFile.WriteString(fullMetaStr) + if err != nil { + return fmt.Errorf("write error: %w", err) + } + return nil +} + +// returns (fileOffset, datasize, error) +// datasize is the current amount of readable data held in the cirfile +func (f *File) GetStartOffsetAndSize(ctx context.Context) (int64, int64, error) { + err := f.flock(ctx, syscall.LOCK_SH) + if err != nil { + return 0, 0, err + } + defer f.unflock() + err = f.readMeta() + if err != nil { + return 0, 0, err + } + chunks := f.getFileChunks() + return f.FileOffset, totalChunksSize(chunks), nil +} + +type fileChunk struct { + StartPos int64 + Len int64 +} + +func totalChunksSize(chunks []fileChunk) int64 { + var rtn int64 + for _, chunk := range chunks { + rtn += chunk.Len + } + return rtn +} + +func advanceChunks(chunks []fileChunk, offset int64) []fileChunk { + if offset < 0 { + panic(fmt.Sprintf("invalid negative offset: %d", offset)) + } + if offset == 0 { + return chunks + } + var rtn []fileChunk + for _, chunk := range chunks { + if offset >= chunk.Len { + offset = offset - chunk.Len + continue + } + if offset == 0 { + rtn = append(rtn, chunk) + } else { + rtn = append(rtn, fileChunk{chunk.StartPos + offset, chunk.Len - offset}) + offset = 0 + } + } + return rtn +} + +func (f *File) getFileChunks() []fileChunk { + if f.StartPos == FilePosEmpty { + return nil + } + if f.EndPos >= f.StartPos { + return []fileChunk{fileChunk{f.StartPos, f.EndPos - f.StartPos + 1}} + } + return []fileChunk{ + fileChunk{f.StartPos, f.FileDataSize - f.StartPos}, + fileChunk{0, f.EndPos + 1}, + } +} + +func (f *File) getFreeChunks() []fileChunk { + if f.StartPos == FilePosEmpty { + return []fileChunk{fileChunk{0, f.MaxSize}} + } + if (f.EndPos == f.StartPos-1) || (f.StartPos == 0 && f.EndPos == f.MaxSize-1) { + return nil + } + if f.EndPos < f.StartPos { + return []fileChunk{fileChunk{f.EndPos + 1, f.StartPos - f.EndPos - 1}} + } + var rtn []fileChunk + if f.EndPos < f.MaxSize-1 { + rtn = append(rtn, fileChunk{f.EndPos + 1, f.MaxSize - f.EndPos - 1}) + } + if f.StartPos > 0 { + rtn = append(rtn, fileChunk{0, f.StartPos}) + } + return rtn +} + +// returns (offset, data, err) +func (f *File) ReadAll(ctx context.Context) (int64, []byte, error) { + err := f.flock(ctx, syscall.LOCK_SH) + if err != nil { + return 0, nil, err + } + defer f.unflock() + err = f.readMeta() + if err != nil { + return 0, nil, err + } + chunks := f.getFileChunks() + curSize := totalChunksSize(chunks) + buf := make([]byte, curSize) + realOffset, nr, err := f.internalReadNext(buf, 0) + return realOffset, buf[0:nr], err +} + +func (f *File) ReadAtWithMax(ctx context.Context, offset int64, maxSize int64) (int64, []byte, error) { + err := f.flock(ctx, syscall.LOCK_SH) + if err != nil { + return 0, nil, err + } + defer f.unflock() + err = f.readMeta() + if err != nil { + return 0, nil, err + } + chunks := f.getFileChunks() + curSize := totalChunksSize(chunks) + var buf []byte + if maxSize > curSize { + buf = make([]byte, curSize) + } else { + buf = make([]byte, maxSize) + } + realOffset, nr, err := f.internalReadNext(buf, offset) + return realOffset, buf[0:nr], err +} + +func (f *File) internalReadNext(buf []byte, offset int64) (int64, int, error) { + if offset < f.FileOffset { + offset = f.FileOffset + } + relativeOffset := offset - f.FileOffset + chunks := f.getFileChunks() + curSize := totalChunksSize(chunks) + if offset >= f.FileOffset+curSize { + return f.FileOffset + curSize, 0, nil + } + chunks = advanceChunks(chunks, relativeOffset) + numRead := 0 + for _, chunk := range chunks { + if numRead >= len(buf) { + break + } + toRead := len(buf) - numRead + if toRead > int(chunk.Len) { + toRead = int(chunk.Len) + } + nr, err := f.OSFile.ReadAt(buf[numRead:numRead+toRead], chunk.StartPos+HeaderLen) + if err != nil { + return offset, 0, err + } + numRead += nr + } + return offset, numRead, nil +} + +// returns (realOffset, numread, error) +// will only return io.EOF when len(data) == 0, otherwise will just do a short read +func (f *File) ReadNext(ctx context.Context, buf []byte, offset int64) (int64, int, error) { + err := f.flock(ctx, syscall.LOCK_SH) + if err != nil { + return 0, 0, err + } + defer f.unflock() + err = f.readMeta() + if err != nil { + return 0, 0, err + } + return f.internalReadNext(buf, offset) +} + +func (f *File) ensureFreeSpace(requiredSpace int64) error { + chunks := f.getFileChunks() + curSpace := f.MaxSize - totalChunksSize(chunks) + if curSpace >= requiredSpace { + return nil + } + neededSpace := requiredSpace - curSpace + if requiredSpace >= f.MaxSize || f.StartPos == FilePosEmpty { + f.StartPos = FilePosEmpty + f.EndPos = 0 + f.FileOffset += neededSpace + } else { + f.StartPos = (f.StartPos + neededSpace) % f.MaxSize + f.FileOffset += neededSpace + } + return f.writeMeta() +} + +// does not implement io.WriterAt (needs context) +func (f *File) WriteAt(ctx context.Context, buf []byte, writePos int64) error { + if writePos < 0 { + return fmt.Errorf("WriteAt got invalid writePos[%d]", writePos) + } + err := f.flock(ctx, syscall.LOCK_EX) + if err != nil { + return err + } + defer f.unflock() + err = f.readMeta() + if err != nil { + return err + } + chunks := f.getFileChunks() + currentSize := totalChunksSize(chunks) + if writePos < f.FileOffset { + negOffset := f.FileOffset - writePos + if negOffset >= int64(len(buf)) { + return nil + } + buf = buf[negOffset:] + writePos = f.FileOffset + } + if writePos > f.FileOffset+currentSize { + // fill gap with zero bytes + posOffset := writePos - (f.FileOffset + currentSize) + err = f.ensureFreeSpace(int64(posOffset)) + if err != nil { + return err + } + var zeroBuf []byte + if posOffset >= f.MaxSize { + zeroBuf = make([]byte, f.MaxSize) + } else { + zeroBuf = make([]byte, posOffset) + } + err = f.internalAppendData(zeroBuf) + if err != nil { + return err + } + // recalc chunks/currentSize + chunks = f.getFileChunks() + currentSize = totalChunksSize(chunks) + // after writing the zero bytes, writePos == f.FileOffset+currentSize (the rest is a straight append) + } + // now writePos >= f.FileOffset && writePos <= f.FileOffset+currentSize (check invariant) + if writePos < f.FileOffset || writePos > f.FileOffset+currentSize { + panic(fmt.Sprintf("invalid writePos, invariant violated writepos[%d] fileoffset[%d] currentsize[%d]", writePos, f.FileOffset, currentSize)) + } + // overwrite existing data (in chunks). advance by writePosOffset + writePosOffset := writePos - f.FileOffset + if writePosOffset < currentSize { + advChunks := advanceChunks(chunks, writePosOffset) + nw, err := f.writeToChunks(buf, advChunks, false) + if err != nil { + return err + } + buf = buf[nw:] + if len(buf) == 0 { + return nil + } + } + // buf contains what was unwritten. this unwritten data is now just a straight append + return f.internalAppendData(buf) +} + +// try writing to chunks, returns (nw, error) +func (f *File) writeToChunks(buf []byte, chunks []fileChunk, updatePos bool) (int64, error) { + var numWrite int64 + for _, chunk := range chunks { + if numWrite >= int64(len(buf)) { + break + } + if chunk.Len == 0 { + continue + } + toWrite := int64(len(buf)) - numWrite + if toWrite > chunk.Len { + toWrite = chunk.Len + } + nw, err := f.OSFile.WriteAt(buf[numWrite:numWrite+toWrite], chunk.StartPos+HeaderLen) + if err != nil { + return 0, err + } + if updatePos { + if chunk.StartPos+int64(nw) > f.FileDataSize { + f.FileDataSize = chunk.StartPos + int64(nw) + } + if f.StartPos == FilePosEmpty { + f.StartPos = chunk.StartPos + } + f.EndPos = chunk.StartPos + int64(nw) - 1 + } + numWrite += int64(nw) + } + return numWrite, nil +} + +func (f *File) internalAppendData(buf []byte) error { + err := f.ensureFreeSpace(int64(len(buf))) + if err != nil { + return err + } + if len(buf) >= int(f.MaxSize) { + buf = buf[len(buf)-int(f.MaxSize):] + } + chunks := f.getFreeChunks() + // don't track nw because we know we have enough free space to write entire buf + _, err = f.writeToChunks(buf, chunks, true) + if err != nil { + return err + } + err = f.writeMeta() + if err != nil { + return err + } + return nil +} + +func (f *File) AppendData(ctx context.Context, buf []byte) error { + err := f.flock(ctx, syscall.LOCK_EX) + if err != nil { + return err + } + defer f.unflock() + err = f.readMeta() + if err != nil { + return err + } + return f.internalAppendData(buf) +} diff --git a/waveshell/pkg/cirfile/cirfile_test.go b/waveshell/pkg/cirfile/cirfile_test.go new file mode 100644 index 00000000..d38cb09a --- /dev/null +++ b/waveshell/pkg/cirfile/cirfile_test.go @@ -0,0 +1,282 @@ +package cirfile + +import ( + "context" + "fmt" + "os" + "path" + "strings" + "syscall" + "testing" + "time" +) + +func validateFileSize(t *testing.T, name string, size int) { + finfo, err := os.Stat(name) + if err != nil { + t.Fatalf("error stating file[%s]: %v", name, err) + } + if int(finfo.Size()) != size { + t.Fatalf("invalid file[%s] expected[%d] got[%d]", name, size, finfo.Size()) + } +} + +func validateMeta(t *testing.T, desc string, f *File, startPos int64, endPos int64, dataSize int64, offset int64) { + if f.StartPos != startPos || f.EndPos != endPos || f.FileDataSize != dataSize || f.FileOffset != offset { + t.Fatalf("metadata error (%s): startpos[%d %d] endpos[%d %d] filedatasize[%d %d] fileoffset[%d %d]", desc, f.StartPos, startPos, f.EndPos, endPos, f.FileDataSize, dataSize, f.FileOffset, offset) + } +} + +func dumpFile(name string) { + barr, _ := os.ReadFile(name) + str := string(barr) + str = strings.ReplaceAll(str, "\x00", ".") + fmt.Printf("%s<<<\n%s\n>>>\n", name, str) +} + +func makeData(size int) string { + var rtn string + for { + if len(rtn) >= size { + break + } + needed := size - len(rtn) + if needed < 10 { + rtn += "123456789\n"[0:needed] + break + } + rtn += "123456789\n" + } + return rtn +} + +func TestCreate(t *testing.T) { + tempDir := t.TempDir() + f1Name := path.Join(tempDir, "f1.cf") + f, err := OpenCirFile(f1Name) + if err == nil || f != nil { + t.Fatalf("OpenCirFile f1.cf should fail (no file)") + } + f, err = CreateCirFile(f1Name, 100) + if err != nil { + t.Fatalf("CreateCirFile f1.cf failed: %v", err) + } + if f == nil { + t.Fatalf("CreateCirFile f1.cf returned nil") + } + err = f.ReadMeta(context.Background()) + if err != nil { + t.Fatalf("cannot readmeta from f1.cf: %v", err) + } + validateFileSize(t, f1Name, 256) + if f.Version != CurrentVersion || f.MaxSize != 100 || f.FileOffset != 0 || f.StartPos != FilePosEmpty || f.EndPos != 0 || f.FileDataSize != 0 || f.FlockStatus != 0 { + t.Fatalf("error with initial metadata #%v", f) + } + buf := make([]byte, 200) + realOffset, nr, err := f.ReadNext(context.Background(), buf, 0) + if realOffset != 0 || nr != 0 || err != nil { + t.Fatalf("error with empty read: real-offset[%d] nr[%d] err[%v]", realOffset, nr, err) + } + realOffset, nr, err = f.ReadNext(context.Background(), buf, 1000) + if realOffset != 0 || nr != 0 || err != nil { + t.Fatalf("error with empty read: real-offset[%d] nr[%d] err[%v]", realOffset, nr, err) + } + f2, err := CreateCirFile(f1Name, 100) + if err == nil || f2 != nil { + t.Fatalf("should be an error to create duplicate CirFile") + } +} + +func TestFile(t *testing.T) { + tempDir := t.TempDir() + f1Name := path.Join(tempDir, "f1.cf") + f, err := CreateCirFile(f1Name, 100) + if err != nil { + t.Fatalf("cannot create cirfile: %v", err) + } + err = f.AppendData(context.Background(), nil) + if err != nil { + t.Fatalf("cannot append data: %v", err) + } + validateFileSize(t, f1Name, HeaderLen) + validateMeta(t, "1", f, FilePosEmpty, 0, 0, 0) + err = f.AppendData(context.Background(), []byte("hello")) + if err != nil { + t.Fatalf("cannot append data: %v", err) + } + validateFileSize(t, f1Name, HeaderLen+5) + validateMeta(t, "2", f, 0, 4, 5, 0) + err = f.AppendData(context.Background(), []byte(" foo")) + if err != nil { + t.Fatalf("cannot append data: %v", err) + } + validateFileSize(t, f1Name, HeaderLen+9) + validateMeta(t, "3", f, 0, 8, 9, 0) + err = f.AppendData(context.Background(), []byte("\n"+makeData(20))) + if err != nil { + t.Fatalf("cannot append data: %v", err) + } + validateFileSize(t, f1Name, HeaderLen+30) + validateMeta(t, "4", f, 0, 29, 30, 0) + + data120 := makeData(120) + err = f.AppendData(context.Background(), []byte(data120)) + if err != nil { + t.Fatalf("cannot append data: %v", err) + } + validateFileSize(t, f1Name, HeaderLen+100) + validateMeta(t, "5", f, 0, 99, 100, 50) + err = f.AppendData(context.Background(), []byte("foo ")) + if err != nil { + t.Fatalf("cannot append data: %v", err) + } + validateFileSize(t, f1Name, HeaderLen+100) + validateMeta(t, "6", f, 4, 3, 100, 54) + + buf := make([]byte, 5) + realOffset, nr, err := f.ReadNext(context.Background(), buf, 0) + if err != nil { + t.Fatalf("cannot ReadNext: %v", err) + } + if realOffset != 54 { + t.Fatalf("wrong realoffset got[%d] expected[%d]", realOffset, 54) + } + if nr != 5 { + t.Fatalf("wrong nr got[%d] expected[%d]", nr, 5) + } + if string(buf[0:nr]) != "56789" { + t.Fatalf("wrong buf return got[%s] expected[%s]", string(buf[0:nr]), "56789") + } + realOffset, nr, err = f.ReadNext(context.Background(), buf, 60) + if err != nil { + t.Fatalf("cannot readnext: %v", err) + } + if realOffset != 60 && nr != 5 { + t.Fatalf("invalid rtn realoffset[%d] nr[%d]", realOffset, nr) + } + if string(buf[0:nr]) != "12345" { + t.Fatalf("invalid rtn buf[%s]", string(buf[0:nr])) + } + realOffset, nr, err = f.ReadNext(context.Background(), buf, 800) + if err != nil || realOffset != 154 || nr != 0 { + t.Fatalf("invalid past end read: err[%v] realoffset[%d] nr[%d]", err, realOffset, nr) + } + realOffset, nr, err = f.ReadNext(context.Background(), buf, 150) + if err != nil || realOffset != 150 || nr != 4 || string(buf[0:nr]) != "foo " { + t.Fatalf("invalid end read: err[%v] realoffset[%d] nr[%d] buf[%s]", err, realOffset, nr, string(buf[0:nr])) + } +} + +func TestFlock(t *testing.T) { + tempDir := t.TempDir() + f1Name := path.Join(tempDir, "f1.cf") + f, err := CreateCirFile(f1Name, 100) + if err != nil { + t.Fatalf("cannot create cirfile: %v", err) + } + fd2, err := os.OpenFile(f1Name, os.O_RDWR, 0777) + if err != nil { + t.Fatalf("cannot open file: %v", err) + } + err = syscall.Flock(int(fd2.Fd()), syscall.LOCK_EX) + if err != nil { + t.Fatalf("cannot lock fd: %v", err) + } + err = f.AppendData(nil, []byte("hello")) + if err != syscall.EWOULDBLOCK { + t.Fatalf("append should fail with EWOULDBLOCK") + } + timeoutCtx, _ := context.WithTimeout(context.Background(), 20*time.Millisecond) + startTs := time.Now() + err = f.ReadMeta(timeoutCtx) + if err != context.DeadlineExceeded { + t.Fatalf("readmeta should fail with context.DeadlineExceeded") + } + dur := time.Now().Sub(startTs) + if dur < 20*time.Millisecond { + t.Fatalf("readmeta should take at least 20ms") + } + syscall.Flock(int(fd2.Fd()), syscall.LOCK_UN) + err = f.ReadMeta(timeoutCtx) + if err != nil { + t.Fatalf("readmeta err: %v", err) + } + err = syscall.Flock(int(fd2.Fd()), syscall.LOCK_SH) + if err != nil { + t.Fatalf("cannot flock: %v", err) + } + err = f.AppendData(nil, []byte("hello")) + if err != syscall.EWOULDBLOCK { + t.Fatalf("append should fail with EWOULDBLOCK") + } + err = f.ReadMeta(timeoutCtx) + if err != nil { + t.Fatalf("readmeta err (should work because LOCK_SH): %v", err) + } + fd2.Close() + err = f.AppendData(nil, []byte("hello")) + if err != nil { + t.Fatalf("append error (should work fd2 was closed): %v", err) + } +} + +func TestWriteAt(t *testing.T) { + tempDir := t.TempDir() + f1Name := path.Join(tempDir, "f1.cf") + f, err := CreateCirFile(f1Name, 100) + if err != nil { + t.Fatalf("cannot create cirfile: %v", err) + } + err = f.WriteAt(nil, []byte("hello\nmike"), 4) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + err = f.WriteAt(nil, []byte("t"), 2) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + err = f.WriteAt(nil, []byte("more"), 30) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + err = f.WriteAt(nil, []byte("\n"), 19) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + dumpFile(f1Name) + err = f.WriteAt(nil, []byte("hello"), 200) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + buf := make([]byte, 10) + realOffset, nr, err := f.ReadNext(context.Background(), buf, 200) + if err != nil || realOffset != 200 || nr != 5 || string(buf[0:nr]) != "hello" { + t.Fatalf("invalid readnext: err[%v] realoffset[%d] nr[%d] buf[%s]", err, realOffset, nr, string(buf[0:nr])) + } + err = f.WriteAt(nil, []byte("0123456789\n"), 100) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + dumpFile(f1Name) + dataStr := makeData(200) + err = f.WriteAt(nil, []byte(dataStr), 50) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + dumpFile(f1Name) + + dataStr = makeData(1000) + err = f.WriteAt(nil, []byte(dataStr), 1002) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + err = f.WriteAt(nil, []byte("hello\n"), 2010) + if err != nil { + t.Fatalf("writeat error: %v", err) + } + err = f.AppendData(nil, []byte("foo\n")) + if err != nil { + t.Fatalf("appenddata error: %v", err) + } + dumpFile(f1Name) +} diff --git a/waveshell/pkg/cmdtail/cmdtail.go b/waveshell/pkg/cmdtail/cmdtail.go new file mode 100644 index 00000000..2578c3b5 --- /dev/null +++ b/waveshell/pkg/cmdtail/cmdtail.go @@ -0,0 +1,471 @@ +package cmdtail + +import ( + "encoding/base64" + "fmt" + "io" + "os" + "regexp" + "sync" + "time" + + "github.com/fsnotify/fsnotify" + "github.com/commandlinedev/apishell/pkg/base" + "github.com/commandlinedev/apishell/pkg/packet" +) + +const MaxDataBytes = 4096 +const FileTypePty = "ptyout" +const FileTypeRun = "runout" + +type Tailer struct { + Lock *sync.Mutex + WatchList map[base.CommandKey]CmdWatchEntry + Watcher *fsnotify.Watcher + Sender *packet.PacketSender + Gen FileNameGenerator + Sessions map[string]bool +} + +type TailPos struct { + ReqId string + Running bool // an active tailer sending data + TailPtyPos int64 + TailRunPos int64 + Follow bool +} + +type CmdWatchEntry struct { + CmdKey base.CommandKey + FilePtyLen int64 + FileRunLen int64 + Tails []TailPos + Done bool +} + +type FileNameGenerator interface { + PtyOutFile(ck base.CommandKey) string + RunOutFile(ck base.CommandKey) string + SessionDir(sessionId string) string +} + +func (w CmdWatchEntry) getTailPos(reqId string) (TailPos, bool) { + for _, pos := range w.Tails { + if pos.ReqId == reqId { + return pos, true + } + } + return TailPos{}, false +} + +func (w *CmdWatchEntry) updateTailPos(reqId string, newPos TailPos) { + for idx, pos := range w.Tails { + if pos.ReqId == reqId { + w.Tails[idx] = newPos + return + } + } + w.Tails = append(w.Tails, newPos) +} + +func (w *CmdWatchEntry) removeTailPos(reqId string) { + var newTails []TailPos + for _, pos := range w.Tails { + if pos.ReqId == reqId { + continue + } + newTails = append(newTails, pos) + } + w.Tails = newTails +} + +func (pos TailPos) IsCurrent(entry CmdWatchEntry) bool { + return pos.TailPtyPos >= entry.FilePtyLen && pos.TailRunPos >= entry.FileRunLen +} + +func (t *Tailer) updateTailPos_nolock(cmdKey base.CommandKey, reqId string, pos TailPos) { + entry, found := t.WatchList[cmdKey] + if !found { + return + } + entry.updateTailPos(reqId, pos) + t.WatchList[cmdKey] = entry +} + +func (t *Tailer) removeTailPos(cmdKey base.CommandKey, reqId string) { + t.Lock.Lock() + defer t.Lock.Unlock() + t.removeTailPos_nolock(cmdKey, reqId) +} + +func (t *Tailer) removeTailPos_nolock(cmdKey base.CommandKey, reqId string) { + entry, found := t.WatchList[cmdKey] + if !found { + return + } + entry.removeTailPos(reqId) + t.WatchList[cmdKey] = entry + if len(entry.Tails) == 0 { + t.removeWatch_nolock(cmdKey) + } +} + +func (t *Tailer) removeWatch_nolock(cmdKey base.CommandKey) { + // delete from watchlist, remove watches + delete(t.WatchList, cmdKey) + t.Watcher.Remove(t.Gen.PtyOutFile(cmdKey)) + t.Watcher.Remove(t.Gen.RunOutFile(cmdKey)) +} + +func (t *Tailer) getEntryAndPos_nolock(cmdKey base.CommandKey, reqId string) (CmdWatchEntry, TailPos, bool) { + entry, found := t.WatchList[cmdKey] + if !found { + return CmdWatchEntry{}, TailPos{}, false + } + pos, found := entry.getTailPos(reqId) + if !found { + return CmdWatchEntry{}, TailPos{}, false + } + return entry, pos, true +} + +func (t *Tailer) addSessionWatcher(sessionId string) error { + t.Lock.Lock() + defer t.Lock.Unlock() + + if t.Sessions[sessionId] { + return + } + sdir := t.Gen.SessionDir(sessionId) + err := t.Watcher.Add(sdir) + if err != nil { + return err + } + t.Sessions[sessionId] = true + return nil +} + +func (t *Tailer) removeSessionWatcher(sessionId string) { + t.Lock.Lock() + defer t.Lock.Unlock() + + if !t.Sessions[sessionId] { + return + } + sdir := t.Gen.SessionDir(sessionId) + t.Watcher.Remove(sdir) +} + +func MakeTailer(sender *packet.PacketSender, gen FileNameGenerator) (*Tailer, error) { + rtn := &Tailer{ + Lock: &sync.Mutex{}, + WatchList: make(map[base.CommandKey]CmdWatchEntry), + Sessions: make(map[string]bool), + Sender: sender, + Gen: gen, + } + var err error + rtn.Watcher, err = fsnotify.NewWatcher() + if err != nil { + return nil, err + } + return rtn, nil +} + +func (t *Tailer) readDataFromFile(fileName string, pos int64, maxBytes int) ([]byte, error) { + fd, err := os.Open(fileName) + defer fd.Close() + if err != nil { + return nil, err + } + buf := make([]byte, maxBytes) + nr, err := fd.ReadAt(buf, pos) + if err != nil && err != io.EOF { // ignore EOF error + return nil, err + } + return buf[0:nr], nil +} + +func (t *Tailer) makeCmdDataPacket(entry CmdWatchEntry, pos TailPos) (*packet.CmdDataPacketType, error) { + dataPacket := packet.MakeCmdDataPacket(pos.ReqId) + dataPacket.CK = entry.CmdKey + dataPacket.PtyPos = pos.TailPtyPos + dataPacket.RunPos = pos.TailRunPos + if entry.FilePtyLen > pos.TailPtyPos { + ptyData, err := t.readDataFromFile(t.Gen.PtyOutFile(entry.CmdKey), pos.TailPtyPos, MaxDataBytes) + if err != nil { + return nil, err + } + dataPacket.PtyData64 = base64.StdEncoding.EncodeToString(ptyData) + dataPacket.PtyDataLen = len(ptyData) + } + if entry.FileRunLen > pos.TailRunPos { + runData, err := t.readDataFromFile(t.Gen.RunOutFile(entry.CmdKey), pos.TailRunPos, MaxDataBytes) + if err != nil { + return nil, err + } + dataPacket.RunData64 = base64.StdEncoding.EncodeToString(runData) + dataPacket.RunDataLen = len(runData) + } + return dataPacket, nil +} + +// returns (data-packet, keepRunning) +func (t *Tailer) runSingleDataTransfer(key base.CommandKey, reqId string) (*packet.CmdDataPacketType, bool, error) { + t.Lock.Lock() + entry, pos, foundPos := t.getEntryAndPos_nolock(key, reqId) + t.Lock.Unlock() + if !foundPos { + return nil, false, nil + } + dataPacket, dataErr := t.makeCmdDataPacket(entry, pos) + + t.Lock.Lock() + defer t.Lock.Unlock() + entry, pos, foundPos = t.getEntryAndPos_nolock(key, reqId) + if !foundPos { + return nil, false, nil + } + // pos was updated between first and second get, throw out data-packet and re-run + if pos.TailPtyPos != dataPacket.PtyPos || pos.TailRunPos != dataPacket.RunPos { + return nil, true, nil + } + if dataErr != nil { + // error, so return error packet, and stop running + pos.Running = false + t.updateTailPos_nolock(key, reqId, pos) + return nil, false, dataErr + } + pos.TailPtyPos += int64(dataPacket.PtyDataLen) + pos.TailRunPos += int64(dataPacket.RunDataLen) + if pos.IsCurrent(entry) { + // we caught up, tail position equals file length + pos.Running = false + } + t.updateTailPos_nolock(key, reqId, pos) + return dataPacket, pos.Running, nil +} + +// returns (removed) +func (t *Tailer) checkRemove(cmdKey base.CommandKey, reqId string) bool { + t.Lock.Lock() + defer t.Lock.Unlock() + entry, pos, foundPos := t.getEntryAndPos_nolock(cmdKey, reqId) + if !foundPos { + return false + } + if !pos.IsCurrent(entry) { + return false + } + if !pos.Follow || entry.Done { + t.removeTailPos_nolock(cmdKey, reqId) + return true + } + return false +} + +func (t *Tailer) RunDataTransfer(key base.CommandKey, reqId string) { + for { + dataPacket, keepRunning, err := t.runSingleDataTransfer(key, reqId) + if dataPacket != nil { + t.Sender.SendPacket(dataPacket) + } + if err != nil { + t.removeTailPos(key, reqId) + t.Sender.SendErrorResponse(reqId, err) + break + } + if !keepRunning { + removed := t.checkRemove(key, reqId) + if removed { + t.Sender.SendResponse(reqId, true) + } + break + } + time.Sleep(10 * time.Millisecond) + } +} + +func (t *Tailer) tryStartRun_nolock(entry CmdWatchEntry, pos TailPos) { + if pos.Running { + return + } + if pos.IsCurrent(entry) { + return + } + pos.Running = true + t.updateTailPos_nolock(entry.CmdKey, pos.ReqId, pos) + go t.RunDataTransfer(entry.CmdKey, pos.ReqId) +} + +var updateFileRe = regexp.MustCompile("/([a-z0-9-]+)/([a-z0-9-]+)\\.(ptyout|runout)$") + +func (t *Tailer) updateFile(relFileName string) { + m := updateFileRe.FindStringSubmatch(relFileName) + if m == nil { + return + } + finfo, err := os.Stat(relFileName) + if err != nil { + t.Sender.SendPacket(packet.FmtMessagePacket("error trying to stat file '%s': %v", relFileName, err)) + return + } + cmdKey := base.MakeCommandKey(m[1], m[2]) + t.Lock.Lock() + defer t.Lock.Unlock() + entry, foundEntry := t.WatchList[cmdKey] + if !foundEntry { + return + } + fileType := m[3] + if fileType == FileTypePty { + entry.FilePtyLen = finfo.Size() + } else if fileType == FileTypeRun { + entry.FileRunLen = finfo.Size() + } + t.WatchList[cmdKey] = entry + for _, pos := range entry.Tails { + t.tryStartRun_nolock(entry, pos) + } +} + +func (t *Tailer) Run() { + for { + select { + case event, ok := <-t.Watcher.Events: + if !ok { + return + } + if event.Op&fsnotify.Write == fsnotify.Write { + t.updateFile(event.Name) + } + + case err, ok := <-t.Watcher.Errors: + if !ok { + return + } + // what to do with this error? just send a message + t.Sender.SendPacket(packet.FmtMessagePacket("error in tailer: %v", err)) + } + } + return +} + +func (t *Tailer) Close() error { + return t.Watcher.Close() +} + +func max(v1 int64, v2 int64) int64 { + if v1 > v2 { + return v1 + } + return v2 +} + +func (entry *CmdWatchEntry) fillFilePos(gen FileNameGenerator) { + ptyInfo, _ := os.Stat(gen.PtyOutFile(entry.CmdKey)) + if ptyInfo != nil { + entry.FilePtyLen = ptyInfo.Size() + } + runoutInfo, _ := os.Stat(gen.RunOutFile(entry.CmdKey)) + if runoutInfo != nil { + entry.FileRunLen = runoutInfo.Size() + } +} + +func (t *Tailer) KeyDone(key base.CommandKey) { + t.Lock.Lock() + defer t.Lock.Unlock() + entry, foundEntry := t.WatchList[key] + if !foundEntry { + return + } + entry.Done = true + var newTails []TailPos + for _, pos := range entry.Tails { + if pos.IsCurrent(entry) { + continue + } + newTails = append(newTails, pos) + } + entry.Tails = newTails + t.WatchList[key] = entry + if len(entry.Tails) == 0 { + t.removeWatch_nolock(key) + } + t.WatchList[key] = entry +} + +func (t *Tailer) RemoveWatch(pk *packet.UntailCmdPacketType) { + t.Lock.Lock() + defer t.Lock.Unlock() + t.removeTailPos_nolock(pk.CK, pk.ReqId) +} + +func (t *Tailer) AddFileWatches_nolock(key base.CommandKey, ptyOnly bool) error { + ptyName := t.Gen.PtyOutFile(key) + runName := t.Gen.RunOutFile(key) + fmt.Printf("WATCH> add %s\n", ptyName) + err := t.Watcher.Add(ptyName) + if err != nil { + return err + } + if ptyOnly { + return nil + } + err = t.Watcher.Add(runName) + if err != nil { + t.Watcher.Remove(ptyName) // best effort clean up + return err + } + return nil +} + +// returns (up-to-date/done, error) +func (t *Tailer) AddWatch(getPacket *packet.GetCmdPacketType) (bool, error) { + if err := getPacket.CK.Validate("getcmd"); err != nil { + return false, err + } + if getPacket.ReqId == "" { + return false, fmt.Errorf("getcmd, no reqid specified") + } + t.Lock.Lock() + defer t.Lock.Unlock() + key := getPacket.CK + entry, foundEntry := t.WatchList[key] + if !foundEntry { + // initialize entry, add watches + entry = CmdWatchEntry{CmdKey: key} + entry.fillFilePos(t.Gen) + } + pos, foundPos := entry.getTailPos(getPacket.ReqId) + if !foundPos { + // initialize a new tailpos + pos = TailPos{ReqId: getPacket.ReqId} + } + // update tailpos with new values from getpacket + pos.TailPtyPos = getPacket.PtyPos + pos.TailRunPos = getPacket.RunPos + pos.Follow = getPacket.Tail + // convert negative pos to positive + if pos.TailPtyPos < 0 { + pos.TailPtyPos = max(0, entry.FilePtyLen+pos.TailPtyPos) // + because negative + } + if pos.TailRunPos < 0 { + pos.TailRunPos = max(0, entry.FileRunLen+pos.TailRunPos) // + because negative + } + entry.updateTailPos(pos.ReqId, pos) + if !pos.Follow && pos.IsCurrent(entry) { + // don't add to t.WatchList, don't t.AddFileWatches_nolock, send rpc response + return true, nil + } + if !foundEntry { + err := t.AddFileWatches_nolock(key, getPacket.PtyOnly) + if err != nil { + return false, err + } + } + t.WatchList[key] = entry + t.tryStartRun_nolock(entry, pos) + return false, nil +} diff --git a/waveshell/pkg/mpio/bufreader.go b/waveshell/pkg/mpio/bufreader.go new file mode 100644 index 00000000..ab4f5709 --- /dev/null +++ b/waveshell/pkg/mpio/bufreader.go @@ -0,0 +1,144 @@ +package mpio + +import ( + "io" + "sync" + + "github.com/commandlinedev/apishell/pkg/packet" +) + +type FdReader struct { + CVar *sync.Cond + M *Multiplexer + FdNum int + Fd io.ReadCloser + BufSize int + Closed bool + ShouldCloseFd bool + IsPty bool +} + +func MakeFdReader(m *Multiplexer, fd io.ReadCloser, fdNum int, shouldCloseFd bool, isPty bool) *FdReader { + fr := &FdReader{ + CVar: sync.NewCond(&sync.Mutex{}), + M: m, + FdNum: fdNum, + Fd: fd, + BufSize: 0, + ShouldCloseFd: shouldCloseFd, + IsPty: isPty, + } + return fr +} + +func (r *FdReader) Close() { + r.CVar.L.Lock() + defer r.CVar.L.Unlock() + if r.Closed { + return + } + if r.Fd != nil && r.ShouldCloseFd { + r.Fd.Close() + } + r.CVar.Broadcast() +} + +func (r *FdReader) GetBufSize() int { + r.CVar.L.Lock() + defer r.CVar.L.Unlock() + return r.BufSize +} + +func (r *FdReader) NotifyAck(ackLen int) { + r.CVar.L.Lock() + defer r.CVar.L.Unlock() + if r.Closed { + return + } + r.BufSize -= ackLen + if r.BufSize < 0 { + r.BufSize = 0 + } + r.CVar.Broadcast() +} + +// !! inverse locking. must already hold the lock when you call this method. +// will *unlock*, send the packet, and then *relock* once it is done. +// this can prevent an unlikely deadlock where we are holding r.CVar.L and stuck on sender.SendCh +func (r *FdReader) sendPacket_unlock(pk packet.PacketType) { + r.CVar.L.Unlock() + defer r.CVar.L.Lock() + r.M.sendPacket(pk) +} + +// returns (success) +func (r *FdReader) WriteWait(data []byte, isEof bool) bool { + r.CVar.L.Lock() + defer r.CVar.L.Unlock() + for { + bufAvail := ReadBufSize - r.BufSize + if r.Closed { + return false + } + if bufAvail == 0 { + r.CVar.Wait() + continue + } + writeLen := min(bufAvail, len(data)) + pk := r.M.makeDataPacket(r.FdNum, data[0:writeLen], nil) + pk.Eof = isEof && (writeLen == len(data)) + r.BufSize += writeLen + data = data[writeLen:] + r.sendPacket_unlock(pk) + if len(data) == 0 { + return true + } + // do *not* do a CVar.Wait() here -- because we *unlocked* to send the packet, we should + // recheck the condition before waiting to avoid deadlock. + } +} + +func min(v1 int, v2 int) int { + if v1 <= v2 { + return v1 + } + return v2 +} + +func (r *FdReader) isClosed() bool { + r.CVar.L.Lock() + defer r.CVar.L.Unlock() + return r.Closed +} + +func (r *FdReader) ReadLoop(wg *sync.WaitGroup) { + defer r.Close() + if wg != nil { + defer wg.Done() + } + buf := make([]byte, 4096) + for { + nr, err := r.Fd.Read(buf) + if r.isClosed() { + return // should not send data or error if we already closed the fd + } + if nr > 0 || err == io.EOF { + isOpen := r.WriteWait(buf[0:nr], (err == io.EOF)) + if !isOpen { + return + } + if err == io.EOF { + return + } + } + if err != nil { + if r.IsPty { + r.WriteWait(nil, true) + return + } + errPk := r.M.makeDataPacket(r.FdNum, nil, err) + r.M.sendPacket(errPk) + return + } + } +} diff --git a/waveshell/pkg/mpio/bufwriter.go b/waveshell/pkg/mpio/bufwriter.go new file mode 100644 index 00000000..86d2efa3 --- /dev/null +++ b/waveshell/pkg/mpio/bufwriter.go @@ -0,0 +1,112 @@ +package mpio + +import ( + "fmt" + "io" + "sync" +) + +type FdWriter struct { + CVar *sync.Cond + M *Multiplexer + FdNum int + Buffer []byte + BufferLimit int + Fd io.WriteCloser + Eof bool + Closed bool + ShouldCloseFd bool + Desc string +} + +func MakeFdWriter(m *Multiplexer, fd io.WriteCloser, fdNum int, shouldCloseFd bool, desc string) *FdWriter { + fw := &FdWriter{ + CVar: sync.NewCond(&sync.Mutex{}), + Fd: fd, + M: m, + FdNum: fdNum, + ShouldCloseFd: shouldCloseFd, + Desc: desc, + BufferLimit: WriteBufSize, + } + return fw +} + +func (w *FdWriter) Close() { + w.CVar.L.Lock() + defer w.CVar.L.Unlock() + if w.Closed { + return + } + w.Closed = true + if w.Fd != nil && w.ShouldCloseFd { + w.Fd.Close() + } + w.Buffer = nil + w.CVar.Broadcast() +} + +func (w *FdWriter) WaitForData() ([]byte, bool) { + w.CVar.L.Lock() + defer w.CVar.L.Unlock() + for { + if len(w.Buffer) > 0 || w.Eof || w.Closed { + toWrite := w.Buffer + w.Buffer = nil + return toWrite, w.Eof + } + w.CVar.Wait() + } +} + +func (w *FdWriter) AddData(data []byte, eof bool) error { + w.CVar.L.Lock() + defer w.CVar.L.Unlock() + if w.Closed || w.Eof { + if len(data) == 0 { + return nil + } + return fmt.Errorf("write to closed file %q (fd:%d) eof[%v]", w.Desc, w.FdNum, w.Eof) + } + if len(data) > 0 { + if len(data)+len(w.Buffer) > w.BufferLimit { + return fmt.Errorf("write exceeds buffer size %q (fd:%d) bufsize=%d (max=%d)", w.Desc, w.FdNum, len(data)+len(w.Buffer), w.BufferLimit) + } + w.Buffer = append(w.Buffer, data...) + } + if eof { + w.Eof = true + } + w.CVar.Broadcast() + return nil +} + +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 + for len(data) > 0 { + if w.Closed { + return + } + chunkSize := min(len(data), MaxSingleWriteSize) + chunk := data[0:chunkSize] + nw, err := w.Fd.Write(chunk) + if nw > 0 || err != nil { + ack := w.M.makeDataAckPacket(w.FdNum, nw, err) + w.M.sendPacket(ack) + } + if err != nil { + return + } + data = data[chunkSize:] + } + if isEof { + return + } + } +} diff --git a/waveshell/pkg/mpio/mpio.go b/waveshell/pkg/mpio/mpio.go new file mode 100644 index 00000000..fe8fed09 --- /dev/null +++ b/waveshell/pkg/mpio/mpio.go @@ -0,0 +1,303 @@ +package mpio + +import ( + "encoding/base64" + "fmt" + "io" + "os" + "sync" + + "github.com/commandlinedev/apishell/pkg/base" + "github.com/commandlinedev/apishell/pkg/packet" +) + +const ReadBufSize = 128 * 1024 +const WriteBufSize = 128 * 1024 +const MaxSingleWriteSize = 4 * 1024 +const MaxTotalRunDataSize = 10 * ReadBufSize + +type Multiplexer struct { + Lock *sync.Mutex + CK base.CommandKey + FdReaders map[int]*FdReader // synchronized + FdWriters map[int]*FdWriter // synchronized + RunData map[int]*FdReader // synchronized + CloseAfterStart []*os.File // synchronized + + Sender *packet.PacketSender + Input *packet.PacketParser + Started bool + UPR packet.UnknownPacketReporter + + Debug bool +} + +func MakeMultiplexer(ck base.CommandKey, upr packet.UnknownPacketReporter) *Multiplexer { + if upr == nil { + upr = packet.DefaultUPR{} + } + return &Multiplexer{ + Lock: &sync.Mutex{}, + CK: ck, + FdReaders: make(map[int]*FdReader), + FdWriters: make(map[int]*FdWriter), + UPR: upr, + } +} + +func (m *Multiplexer) Close() { + m.Lock.Lock() + defer m.Lock.Unlock() + + for _, fr := range m.FdReaders { + fr.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() + if err != nil { + return nil, err + } + m.Lock.Lock() + defer m.Lock.Unlock() + m.FdReaders[fdNum] = MakeFdReader(m, pr, fdNum, true, false) + m.CloseAfterStart = append(m.CloseAfterStart, pw) + return pw, nil +} + +// returns the *reader* to connect to process, writer is put in FdWriters +func (m *Multiplexer) MakeWriterPipe(fdNum int, desc string) (*os.File, error) { + pr, pw, err := os.Pipe() + if err != nil { + return nil, err + } + m.Lock.Lock() + defer m.Lock.Unlock() + m.FdWriters[fdNum] = MakeFdWriter(m, pw, fdNum, true, desc) + m.CloseAfterStart = append(m.CloseAfterStart, pr) + return pr, nil +} + +// returns the *reader* to connect to process, writer is put in FdWriters +func (m *Multiplexer) MakeStaticWriterPipe(fdNum int, data []byte, bufferLimit int, desc string) (*os.File, error) { + pr, pw, err := os.Pipe() + if err != nil { + return nil, err + } + m.Lock.Lock() + defer m.Lock.Unlock() + fdWriter := MakeFdWriter(m, pw, fdNum, true, desc) + fdWriter.BufferLimit = bufferLimit + err = fdWriter.AddData(data, true) + if err != nil { + return nil, err + } + m.FdWriters[fdNum] = fdWriter + m.CloseAfterStart = append(m.CloseAfterStart, pr) + return pr, nil +} + +func (m *Multiplexer) MakeRawFdReader(fdNum int, fd io.ReadCloser, shouldClose bool, isPty bool) { + m.Lock.Lock() + defer m.Lock.Unlock() + m.FdReaders[fdNum] = MakeFdReader(m, fd, fdNum, shouldClose, isPty) +} + +func (m *Multiplexer) MakeRawFdWriter(fdNum int, fd io.WriteCloser, shouldClose bool, desc string) { + m.Lock.Lock() + defer m.Lock.Unlock() + m.FdWriters[fdNum] = MakeFdWriter(m, fd, fdNum, shouldClose, desc) +} + +func (m *Multiplexer) makeDataAckPacket(fdNum int, ackLen int, err error) *packet.DataAckPacketType { + ack := packet.MakeDataAckPacket() + ack.CK = m.CK + ack.FdNum = fdNum + ack.AckLen = ackLen + if err != nil { + ack.Error = err.Error() + } + return ack +} + +func (m *Multiplexer) makeDataPacket(fdNum int, data []byte, err error) *packet.DataPacketType { + pk := packet.MakeDataPacket() + pk.CK = m.CK + pk.FdNum = fdNum + pk.Data64 = base64.StdEncoding.EncodeToString(data) + if err != nil { + pk.Error = err.Error() + } + return pk +} + +func (m *Multiplexer) sendPacket(p packet.PacketType) { + m.Sender.SendPacket(p) +} + +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(wg) + } +} + +func (m *Multiplexer) launchReaders(wg *sync.WaitGroup) { + m.Lock.Lock() + defer m.Lock.Unlock() + if wg != nil { + wg.Add(len(m.FdReaders)) + } + for _, fr := range m.FdReaders { + go fr.ReadLoop(wg) + } +} + +func (m *Multiplexer) startIO(packetParser *packet.PacketParser, sender *packet.PacketSender) { + m.Lock.Lock() + defer m.Lock.Unlock() + if m.Started { + panic("Multiplexer is already running, cannot start again") + } + m.Input = packetParser + m.Sender = sender + m.Started = true +} + +func (m *Multiplexer) runPacketInputLoop() *packet.CmdDonePacketType { + defer m.HandleInputDone() + for pk := range m.Input.MainCh { + if m.Debug { + fmt.Printf("PK-M> %s\n", packet.AsString(pk)) + } + if pk.GetType() == packet.DataPacketStr { + dataPacket := pk.(*packet.DataPacketType) + err := m.processDataPacket(dataPacket) + if err != nil { + errPacket := m.makeDataAckPacket(dataPacket.FdNum, 0, err) + m.sendPacket(errPacket) + } + continue + } + if pk.GetType() == packet.DataAckPacketStr { + ackPacket := pk.(*packet.DataAckPacketType) + m.processAckPacket(ackPacket) + continue + } + if pk.GetType() == packet.CmdDonePacketStr { + donePacket := pk.(*packet.CmdDonePacketType) + return donePacket + } + m.UPR.UnknownPacket(pk) + } + return nil +} + +func (m *Multiplexer) WriteDataToFd(fdNum int, data []byte, isEof bool) error { + m.Lock.Lock() + defer m.Lock.Unlock() + fw := m.FdWriters[fdNum] + if fw == nil { + // add a closed FdWriter as a placeholder so we only send one error + fw := MakeFdWriter(m, nil, fdNum, false, "invalid-fd") + fw.Close() + m.FdWriters[fdNum] = fw + return fmt.Errorf("write to closed file (no fd)") + } + err := fw.AddData(data, isEof) + if err != nil { + fw.Close() + return err + } + return nil +} + +func (m *Multiplexer) processDataPacket(dataPacket *packet.DataPacketType) error { + realData, err := base64.StdEncoding.DecodeString(dataPacket.Data64) + if err != nil { + return fmt.Errorf("decoding base64 data: %w", err) + } + return m.WriteDataToFd(dataPacket.FdNum, realData, dataPacket.Eof) +} + +func (m *Multiplexer) processAckPacket(ackPacket *packet.DataAckPacketType) { + m.Lock.Lock() + defer m.Lock.Unlock() + fr := m.FdReaders[ackPacket.FdNum] + if fr == nil { + return + } + fr.NotifyAck(ackPacket.AckLen) +} + +func (m *Multiplexer) closeTempStartFds() { + m.Lock.Lock() + defer m.Lock.Unlock() + for _, fd := range m.CloseAfterStart { + fd.Close() + } + m.CloseAfterStart = nil +} + +func (m *Multiplexer) RunIOAndWait(packetParser *packet.PacketParser, sender *packet.PacketSender, waitOnReaders bool, waitOnWriters bool, waitForInputLoop bool) *packet.CmdDonePacketType { + m.startIO(packetParser, sender) + m.closeTempStartFds() + var wg sync.WaitGroup + 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/waveshell/pkg/packet/combined.go b/waveshell/pkg/packet/combined.go new file mode 100644 index 00000000..0f15fd93 --- /dev/null +++ b/waveshell/pkg/packet/combined.go @@ -0,0 +1,34 @@ +package packet + +type CombinedPacket struct { + Type string `json:"type"` + Success bool `json:"success"` + Ts int64 `json:"ts"` + Id string `json:"id,omitempty"` + + SessionId string `json:"sessionid"` + CmdId string `json:"cmdid"` + + PtyPos int64 `json:"ptypos"` + PtyLen int64 `json:"ptylen"` + RunPos int64 `json:"runpos"` + RunLen int64 `json:"runlen"` + + Error string `json:"error"` + NotFound bool `json:"notfound,omitempty"` + Tail bool `json:"tail,omitempty"` + Dir string `json:"dir"` + ChDir string `json:"chdir,omitempty"` + + Data string `json:"data"` + PtyData string `json:"ptydata"` + RunData string `json:"rundata"` + Message string `json:"message"` + Command string `json:"command"` + + ScHomeDir string `json:"schomedir"` + HomeDir string `json:"homedir"` + Env []string `json:"env"` + ExitCode int `json:"exitcode"` + RunnerPid int `json:"runnerpid"` +} diff --git a/waveshell/pkg/packet/packet.go b/waveshell/pkg/packet/packet.go new file mode 100644 index 00000000..30c36f77 --- /dev/null +++ b/waveshell/pkg/packet/packet.go @@ -0,0 +1,1178 @@ +package packet + +import ( + "bytes" + "context" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "reflect" + "sync" + + "github.com/commandlinedev/apishell/pkg/base" +) + +// single : run, >cmddata, >cmddone, data, <>dataack, run, >cmddata, >cmddone, run, >cmddata, >cmddone, data, <>dataack, cd, >getcmd, >untailcmd, >input, error, <>message, <>ping, streamfile, writefile, filedata*, = 127 || (b < 32 && b != 10 && b != 13) { + buf[idx] = '?' + } + } +} + +type SendError struct { + IsWriteError bool // fatal + IsMarshalError bool // not fatal + PacketType string + Err error +} + +func (e *SendError) Unwrap() error { + return e.Err +} + +func (e *SendError) Error() string { + if e.IsMarshalError { + return fmt.Sprintf("SendPacket marshal-error '%s' packet: %v", e.PacketType, e.Err) + } else if e.IsWriteError { + return fmt.Sprintf("SendPacket write-error packet[%s]: %v", e.PacketType, e.Err) + } else { + return e.Err.Error() + } +} + +func MarshalPacket(packet PacketType) ([]byte, error) { + if packet == nil { + return nil, fmt.Errorf("invalid nil packet") + } + jsonBytes, err := json.Marshal(packet) + if err != nil { + return nil, &SendError{IsMarshalError: true, PacketType: packet.GetType(), Err: err} + } + var outBuf bytes.Buffer + outBuf.WriteByte('\n') + outBuf.WriteString(fmt.Sprintf("##%d", len(jsonBytes))) + outBuf.Write(jsonBytes) + outBuf.WriteByte('\n') + outBytes := outBuf.Bytes() + sanitizeBytes(outBytes) + return outBytes, nil +} + +func SendPacket(w io.Writer, packet PacketType) error { + if packet == nil { + return nil + } + outBytes, err := MarshalPacket(packet) + if err != nil { + return err + } + if GlobalDebug { + base.Logf("SEND> %s\n", AsString(packet)) + } + _, err = w.Write(outBytes) + if err != nil { + return &SendError{IsWriteError: true, PacketType: packet.GetType(), Err: err} + } + return nil +} + +func SendCmdError(w io.Writer, ck base.CommandKey, err error) error { + return SendPacket(w, MakeCmdErrorPacket(ck, err)) +} + +type PacketSender struct { + Lock *sync.Mutex + SendCh chan PacketType + Done bool + DoneCh chan bool + ErrHandler func(*PacketSender, PacketType, error) + ExitErr error +} + +func MakePacketSender(output io.Writer, errHandler func(*PacketSender, PacketType, error)) *PacketSender { + sender := &PacketSender{ + Lock: &sync.Mutex{}, + SendCh: make(chan PacketType, PacketSenderQueueSize), + DoneCh: make(chan bool), + ErrHandler: errHandler, + } + go func() { + defer close(sender.DoneCh) + defer sender.Close() + for pk := range sender.SendCh { + err := SendPacket(output, pk) + if err != nil { + sender.goHandleError(pk, err) + if serr, ok := err.(*SendError); ok && serr.IsMarshalError { + // marshaler errors are recoverable + continue + } + // write errors are not recoverable + sender.Lock.Lock() + sender.ExitErr = err + sender.Lock.Unlock() + return + } + } + }() + return sender +} + +func (sender *PacketSender) goHandleError(pk PacketType, err error) { + sender.Lock.Lock() + defer sender.Lock.Unlock() + if sender.ErrHandler != nil { + go sender.ErrHandler(sender, pk, err) + } +} + +func MakeChannelPacketSender(packetCh chan PacketType) *PacketSender { + sender := &PacketSender{ + Lock: &sync.Mutex{}, + SendCh: make(chan PacketType, PacketSenderQueueSize), + DoneCh: make(chan bool), + } + go func() { + defer close(sender.DoneCh) + defer sender.Close() + for pk := range sender.SendCh { + packetCh <- pk + } + }() + return sender +} + +func (sender *PacketSender) Close() { + sender.Lock.Lock() + defer sender.Lock.Unlock() + if sender.Done { + return + } + sender.Done = true + close(sender.SendCh) +} + +// returns ExitErr if set +func (sender *PacketSender) WaitForDone() error { + <-sender.DoneCh + sender.Lock.Lock() + defer sender.Lock.Unlock() + return sender.ExitErr +} + +// this is "advisory", as there is a race condition between the loop closing and setting Done. +// that's okay because that's an impossible race condition anyway (you could enqueue the packet +// and then the connection dies, or it dies half way, etc.). this just stops blindly adding +// packets forever when the loop is done. +func (sender *PacketSender) checkStatus() error { + sender.Lock.Lock() + defer sender.Lock.Unlock() + if sender.Done { + return fmt.Errorf("cannot send packet, sender write loop is closed") + } + return nil +} + +func (sender *PacketSender) SendPacketCtx(ctx context.Context, pk PacketType) error { + err := sender.checkStatus() + if err != nil { + return err + } + select { + case sender.SendCh <- pk: + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + +func (sender *PacketSender) SendPacket(pk PacketType) error { + err := sender.checkStatus() + if err != nil { + return err + } + sender.SendCh <- pk + return nil +} + +func (sender *PacketSender) SendCmdError(ck base.CommandKey, err error) error { + return sender.SendPacket(MakeCmdErrorPacket(ck, err)) +} + +func (sender *PacketSender) SendErrorResponse(reqId string, err error) error { + pk := MakeErrorResponsePacket(reqId, err) + return sender.SendPacket(pk) +} + +func (sender *PacketSender) SendResponse(reqId string, data interface{}) error { + pk := MakeResponsePacket(reqId, data) + return sender.SendPacket(pk) +} + +func (sender *PacketSender) SendMessageFmt(fmtStr string, args ...interface{}) error { + return sender.SendPacket(MakeMessagePacket(fmt.Sprintf(fmtStr, args...))) +} + +type UnknownPacketReporter interface { + UnknownPacket(pk PacketType) +} + +type DefaultUPR struct{} + +func (DefaultUPR) UnknownPacket(pk PacketType) { + if pk.GetType() == CmdErrorPacketStr { + errPacket := pk.(*CmdErrorPacketType) + // at this point, just send the error packet to stderr rather than try to do something special + fmt.Fprintf(os.Stderr, "[error] %s\n", errPacket.Error) + } else if pk.GetType() == RawPacketStr { + rawPacket := pk.(*RawPacketType) + fmt.Fprintf(os.Stderr, "%s\n", rawPacket.Data) + } else if pk.GetType() == CmdStartPacketStr { + return // do nothing + } else { + fmt.Fprintf(os.Stderr, "[error] invalid packet received '%s'", AsExtType(pk)) + } +} + +type MessageUPR struct { + CK base.CommandKey + Sender *PacketSender +} + +func (upr MessageUPR) UnknownPacket(pk PacketType) { + msg := FmtMessagePacket("[error] invalid packet received %s", AsString(pk)) + msg.CK = upr.CK + upr.Sender.SendPacket(msg) +} + +// todo: clean hanging entries in RunMap when in server mode +type RunPacketBuilder struct { + RunMap map[base.CommandKey]*RunPacketType +} + +func MakeRunPacketBuilder() *RunPacketBuilder { + return &RunPacketBuilder{ + RunMap: make(map[base.CommandKey]*RunPacketType), + } +} + +// returns (consumed, fullRunPacket) +func (b *RunPacketBuilder) ProcessPacket(pk PacketType) (bool, *RunPacketType) { + if pk.GetType() == RunPacketStr { + runPacket := pk.(*RunPacketType) + if len(runPacket.RunData) == 0 { + return true, runPacket + } + b.RunMap[runPacket.CK] = runPacket + return true, nil + } + if pk.GetType() == DataEndPacketStr { + endPacket := pk.(*DataEndPacketType) + runPacket := b.RunMap[endPacket.CK] // might be nil + delete(b.RunMap, endPacket.CK) + return true, runPacket + } + if pk.GetType() == DataPacketStr { + dataPacket := pk.(*DataPacketType) + runPacket := b.RunMap[dataPacket.CK] + if runPacket == nil { + return false, nil + } + for idx, runData := range runPacket.RunData { + if runData.FdNum == dataPacket.FdNum { + // can ignore error, will get caught later with RunData.DataLen check + realData, _ := base64.StdEncoding.DecodeString(dataPacket.Data64) + runData.Data = append(runData.Data, realData...) + runPacket.RunData[idx] = runData + break + } + } + return true, nil + } + return false, nil +} diff --git a/waveshell/pkg/packet/parser.go b/waveshell/pkg/packet/parser.go new file mode 100644 index 00000000..f84be258 --- /dev/null +++ b/waveshell/pkg/packet/parser.go @@ -0,0 +1,238 @@ +package packet + +import ( + "bufio" + "context" + "io" + "strconv" + "strings" + "sync" +) + +type PacketParser struct { + Lock *sync.Mutex + MainCh chan PacketType + RpcMap map[string]*RpcEntry + RpcHandler bool + Err error +} + +type RpcEntry struct { + ReqId string + RespCh chan RpcResponsePacketType +} + +type RpcResponseIter struct { + ReqId string + Parser *PacketParser +} + +func (iter *RpcResponseIter) Next(ctx context.Context) (RpcResponsePacketType, error) { + // will unregister the rpc on ResponseDone + return iter.Parser.GetNextResponse(ctx, iter.ReqId) +} + +func (iter *RpcResponseIter) Close() { + iter.Parser.UnRegisterRpc(iter.ReqId) +} + +func CombinePacketParsers(p1 *PacketParser, p2 *PacketParser, rpcHandler bool) *PacketParser { + rtnParser := &PacketParser{ + Lock: &sync.Mutex{}, + MainCh: make(chan PacketType), + RpcMap: make(map[string]*RpcEntry), + RpcHandler: rpcHandler, + } + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + for pk := range p1.MainCh { + if rtnParser.RpcHandler { + sent := rtnParser.trySendRpcResponse(pk) + if sent { + continue + } + } + rtnParser.MainCh <- pk + } + }() + go func() { + defer wg.Done() + for pk := range p2.MainCh { + if rtnParser.RpcHandler { + sent := rtnParser.trySendRpcResponse(pk) + if sent { + continue + } + } + rtnParser.MainCh <- pk + } + }() + go func() { + wg.Wait() + close(rtnParser.MainCh) + }() + return rtnParser +} + +// should have already registered rpc +func (p *PacketParser) WaitForResponse(ctx context.Context, reqId string) RpcResponsePacketType { + entry := p.getRpcEntry(reqId) + if entry == nil { + return nil + } + defer p.UnRegisterRpc(reqId) + select { + case resp := <-entry.RespCh: + return resp + case <-ctx.Done(): + return nil + } +} + +func (p *PacketParser) GetResponseIter(reqId string) *RpcResponseIter { + return &RpcResponseIter{Parser: p, ReqId: reqId} +} + +func (p *PacketParser) GetNextResponse(ctx context.Context, reqId string) (RpcResponsePacketType, error) { + entry := p.getRpcEntry(reqId) + if entry == nil { + return nil, nil + } + select { + case resp := <-entry.RespCh: + if resp.GetResponseDone() { + p.UnRegisterRpc(reqId) + } + return resp, nil + case <-ctx.Done(): + return nil, ctx.Err() + } +} + +func (p *PacketParser) UnRegisterRpc(reqId string) { + p.Lock.Lock() + defer p.Lock.Unlock() + entry := p.RpcMap[reqId] + if entry != nil { + close(entry.RespCh) + delete(p.RpcMap, reqId) + } +} + +func (p *PacketParser) RegisterRpc(reqId string) chan RpcResponsePacketType { + return p.RegisterRpcSz(reqId, 2) +} + +func (p *PacketParser) RegisterRpcSz(reqId string, queueSize int) chan RpcResponsePacketType { + p.Lock.Lock() + defer p.Lock.Unlock() + ch := make(chan RpcResponsePacketType, queueSize) + entry := &RpcEntry{ReqId: reqId, RespCh: ch} + p.RpcMap[reqId] = entry + return ch +} + +func (p *PacketParser) getRpcEntry(reqId string) *RpcEntry { + p.Lock.Lock() + defer p.Lock.Unlock() + entry := p.RpcMap[reqId] + return entry +} + +func (p *PacketParser) trySendRpcResponse(pk PacketType) bool { + respPk, ok := pk.(RpcResponsePacketType) + if !ok { + return false + } + p.Lock.Lock() + defer p.Lock.Unlock() + entry := p.RpcMap[respPk.GetResponseId()] + if entry == nil { + return false + } + // nonblocking send + select { + case entry.RespCh <- respPk: + default: + } + return true +} + +func (p *PacketParser) GetErr() error { + p.Lock.Lock() + defer p.Lock.Unlock() + return p.Err +} + +func (p *PacketParser) SetErr(err error) { + p.Lock.Lock() + defer p.Lock.Unlock() + if p.Err == nil { + p.Err = err + } +} + +func MakePacketParser(input io.Reader, rpcHandler bool) *PacketParser { + parser := &PacketParser{ + Lock: &sync.Mutex{}, + MainCh: make(chan PacketType), + RpcMap: make(map[string]*RpcEntry), + RpcHandler: rpcHandler, + } + bufReader := bufio.NewReader(input) + go func() { + defer func() { + close(parser.MainCh) + }() + for { + line, err := bufReader.ReadString('\n') + if err == io.EOF { + return + } + if err != nil { + parser.SetErr(err) + return + } + if line == "\n" { + continue + } + // ##[len][json]\n + // ##14{"hello":true}\n + // ##N{...} + bracePos := strings.Index(line, "{") + if !strings.HasPrefix(line, "##") || bracePos == -1 { + parser.MainCh <- MakeRawPacket(line[:len(line)-1]) + continue + } + packetLen := -1 + if line[2:bracePos] != "N" { + packetLen, err = strconv.Atoi(line[2:bracePos]) + if err != nil || packetLen != len(line)-bracePos-1 { + parser.MainCh <- MakeRawPacket(line[:len(line)-1]) + continue + } + } + pk, err := ParseJsonPacket([]byte(line[bracePos:])) + if err != nil { + parser.MainCh <- MakeRawPacket(line[:len(line)-1]) + continue + } + if pk.GetType() == DonePacketStr { + return + } + if pk.GetType() == PingPacketStr { + continue + } + if parser.RpcHandler { + sent := parser.trySendRpcResponse(pk) + if sent { + continue + } + } + parser.MainCh <- pk + } + }() + return parser +} diff --git a/waveshell/pkg/packet/shellstate.go b/waveshell/pkg/packet/shellstate.go new file mode 100644 index 00000000..bcf59f27 --- /dev/null +++ b/waveshell/pkg/packet/shellstate.go @@ -0,0 +1,192 @@ +package packet + +import ( + "bytes" + "crypto/sha1" + "encoding/base64" + "encoding/json" + "fmt" + + "github.com/commandlinedev/apishell/pkg/binpack" + "github.com/commandlinedev/apishell/pkg/statediff" +) + +const ShellStatePackVersion = 0 +const ShellStateDiffPackVersion = 0 + +type ShellState struct { + Version string `json:"version"` // [type] [semver] + Cwd string `json:"cwd,omitempty"` + ShellVars []byte `json:"shellvars,omitempty"` + Aliases string `json:"aliases,omitempty"` + Funcs string `json:"funcs,omitempty"` + Error string `json:"error,omitempty"` + HashVal string `json:"-"` +} + +type ShellStateDiff struct { + Version string `json:"version"` // [type] [semver] + BaseHash string `json:"basehash"` + DiffHashArr []string `json:"diffhasharr,omitempty"` + Cwd string `json:"cwd,omitempty"` + VarsDiff []byte `json:"shellvarsdiff,omitempty"` // vardiff + AliasesDiff []byte `json:"aliasesdiff,omitempty"` // linediff + FuncsDiff []byte `json:"funcsdiff,omitempty"` // linediff + Error string `json:"error,omitempty"` + HashVal string `json:"-"` +} + +func (state ShellState) IsEmpty() bool { + return state.Version == "" && state.Cwd == "" && len(state.ShellVars) == 0 && state.Aliases == "" && state.Funcs == "" && state.Error == "" +} + +// returns base64 hash of data +func sha1Hash(data []byte) string { + hvalRaw := sha1.Sum(data) + hval := base64.StdEncoding.EncodeToString(hvalRaw[:]) + return hval +} + +// returns (SHA1, encoded-state) +func (state ShellState) EncodeAndHash() (string, []byte) { + var buf bytes.Buffer + binpack.PackInt(&buf, ShellStatePackVersion) + binpack.PackValue(&buf, []byte(state.Version)) + binpack.PackValue(&buf, []byte(state.Cwd)) + binpack.PackValue(&buf, state.ShellVars) + binpack.PackValue(&buf, []byte(state.Aliases)) + binpack.PackValue(&buf, []byte(state.Funcs)) + binpack.PackValue(&buf, []byte(state.Error)) + return sha1Hash(buf.Bytes()), buf.Bytes() +} + +func (state ShellState) MarshalJSON() ([]byte, error) { + _, encodedBytes := state.EncodeAndHash() + return json.Marshal(encodedBytes) +} + +// caches HashVal in struct +func (state *ShellState) GetHashVal(force bool) string { + if state.HashVal == "" || force { + state.HashVal, _ = state.EncodeAndHash() + } + return state.HashVal +} + +func (state *ShellState) DecodeShellState(barr []byte) error { + state.HashVal = sha1Hash(barr) + buf := bytes.NewBuffer(barr) + u := binpack.MakeUnpacker(buf) + version := u.UnpackInt("ShellState pack version") + if version != ShellStatePackVersion { + return fmt.Errorf("invalid ShellState pack version: %d", version) + } + state.Version = string(u.UnpackValue("ShellState.Version")) + state.Cwd = string(u.UnpackValue("ShellState.Cwd")) + state.ShellVars = u.UnpackValue("ShellState.ShellVars") + state.Aliases = string(u.UnpackValue("ShellState.Aliases")) + state.Funcs = string(u.UnpackValue("ShellState.Funcs")) + state.Error = string(u.UnpackValue("ShellState.Error")) + return u.Error() +} + +func (state *ShellState) UnmarshalJSON(jsonBytes []byte) error { + var barr []byte + err := json.Unmarshal(jsonBytes, &barr) + if err != nil { + return err + } + return state.DecodeShellState(barr) +} + +func (sdiff ShellStateDiff) EncodeAndHash() (string, []byte) { + var buf bytes.Buffer + binpack.PackInt(&buf, ShellStateDiffPackVersion) + binpack.PackValue(&buf, []byte(sdiff.Version)) + binpack.PackValue(&buf, []byte(sdiff.BaseHash)) + binpack.PackStrArr(&buf, sdiff.DiffHashArr) + binpack.PackValue(&buf, []byte(sdiff.Cwd)) + binpack.PackValue(&buf, sdiff.VarsDiff) + binpack.PackValue(&buf, sdiff.AliasesDiff) + binpack.PackValue(&buf, sdiff.FuncsDiff) + binpack.PackValue(&buf, []byte(sdiff.Error)) + return sha1Hash(buf.Bytes()), buf.Bytes() +} + +func (sdiff ShellStateDiff) MarshalJSON() ([]byte, error) { + _, encodedBytes := sdiff.EncodeAndHash() + return json.Marshal(encodedBytes) +} + +func (sdiff *ShellStateDiff) DecodeShellStateDiff(barr []byte) error { + sdiff.HashVal = sha1Hash(barr) + buf := bytes.NewBuffer(barr) + u := binpack.MakeUnpacker(buf) + version := u.UnpackInt("ShellState pack version") + if version != ShellStateDiffPackVersion { + return fmt.Errorf("invalid ShellStateDiff pack version: %d", version) + } + sdiff.Version = string(u.UnpackValue("ShellStateDiff.Version")) + sdiff.BaseHash = string(u.UnpackValue("ShellStateDiff.BaseHash")) + sdiff.DiffHashArr = u.UnpackStrArr("ShellStateDiff.DiffHashArr") + sdiff.Cwd = string(u.UnpackValue("ShellStateDiff.Cwd")) + sdiff.VarsDiff = u.UnpackValue("ShellStateDiff.VarsDiff") + sdiff.AliasesDiff = u.UnpackValue("ShellStateDiff.AliasesDiff") + sdiff.FuncsDiff = u.UnpackValue("ShellStateDiff.FuncsDiff") + sdiff.Error = string(u.UnpackValue("ShellStateDiff.Error")) + return u.Error() +} + +func (sdiff *ShellStateDiff) UnmarshalJSON(jsonBytes []byte) error { + var barr []byte + err := json.Unmarshal(jsonBytes, &barr) + if err != nil { + return err + } + return sdiff.DecodeShellStateDiff(barr) +} + +// caches HashVal in struct +func (sdiff *ShellStateDiff) GetHashVal(force bool) string { + if sdiff.HashVal == "" || force { + sdiff.HashVal, _ = sdiff.EncodeAndHash() + } + return sdiff.HashVal +} + +func (sdiff ShellStateDiff) Dump(vars bool, aliases bool, funcs bool) { + fmt.Printf("ShellStateDiff:\n") + fmt.Printf(" version: %s\n", sdiff.Version) + fmt.Printf(" base: %s\n", sdiff.BaseHash) + fmt.Printf(" vars: %d, aliases: %d, funcs: %d\n", len(sdiff.VarsDiff), len(sdiff.AliasesDiff), len(sdiff.FuncsDiff)) + if sdiff.Error != "" { + fmt.Printf(" error: %s\n", sdiff.Error) + } + if vars { + var mdiff statediff.MapDiffType + err := mdiff.Decode(sdiff.VarsDiff) + if err != nil { + fmt.Printf(" vars: error[%s]\n", err.Error()) + } else { + mdiff.Dump() + } + } + if aliases && len(sdiff.AliasesDiff) > 0 { + var ldiff statediff.LineDiffType + err := ldiff.Decode(sdiff.AliasesDiff) + if err != nil { + fmt.Printf(" aliases: error[%s]\n", err.Error()) + } else { + ldiff.Dump() + } + } + if funcs && len(sdiff.FuncsDiff) > 0 { + var ldiff statediff.LineDiffType + err := ldiff.Decode(sdiff.FuncsDiff) + if err != nil { + fmt.Printf(" funcs: error[%s]\n", err.Error()) + } else { + ldiff.Dump() + } + } +} diff --git a/waveshell/pkg/server/server.go b/waveshell/pkg/server/server.go new file mode 100644 index 00000000..3ccbb318 --- /dev/null +++ b/waveshell/pkg/server/server.go @@ -0,0 +1,747 @@ +package server + +import ( + "context" + "errors" + "fmt" + "io" + "io/fs" + "os" + "os/exec" + "path/filepath" + "sort" + "strings" + "sync" + "time" + + "github.com/alessio/shellescape" + "github.com/commandlinedev/apishell/pkg/base" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/apishell/pkg/shexec" +) + +const MaxFileDataPacketSize = 16 * 1024 +const WriteFileContextTimeout = 30 * time.Second +const cleanLoopTime = 5 * time.Second +const MaxWriteFileContextData = 100 + +// TODO create unblockable packet-sender (backed by an array) for clientproc +type MServer struct { + Lock *sync.Mutex + MainInput *packet.PacketParser + Sender *packet.PacketSender + ClientMap map[base.CommandKey]*shexec.ClientProc + Debug bool + StateMap map[string]*packet.ShellState // sha1->state + CurrentState string // sha1 + WriteErrorCh chan bool // closed if there is a I/O write error + WriteErrorChOnce *sync.Once + WriteFileContextMap map[string]*WriteFileContext + Done bool +} + +type WriteFileContext struct { + CVar *sync.Cond + Data []*packet.FileDataPacketType + LastActive time.Time + Err error + Done bool +} + +func (m *MServer) Close() { + m.Sender.Close() + m.Sender.WaitForDone() + m.Lock.Lock() + defer m.Lock.Unlock() + m.Done = true +} + +func (m *MServer) checkDone() bool { + m.Lock.Lock() + defer m.Lock.Unlock() + return m.Done +} + +func (m *MServer) getWriteFileContext(reqId string) *WriteFileContext { + m.Lock.Lock() + defer m.Lock.Unlock() + wfc := m.WriteFileContextMap[reqId] + if wfc == nil { + wfc = &WriteFileContext{ + CVar: sync.NewCond(&sync.Mutex{}), + LastActive: time.Now(), + } + m.WriteFileContextMap[reqId] = wfc + } + return wfc +} + +func (m *MServer) addFileDataPacket(pk *packet.FileDataPacketType) { + m.Lock.Lock() + wfc := m.WriteFileContextMap[pk.RespId] + m.Lock.Unlock() + if wfc == nil { + return + } + wfc.CVar.L.Lock() + defer wfc.CVar.L.Unlock() + if wfc.Done || wfc.Err != nil { + return + } + if len(wfc.Data) > MaxWriteFileContextData { + wfc.Err = errors.New("write-file buffer length exceeded") + wfc.Data = nil + wfc.CVar.Broadcast() + return + } + wfc.LastActive = time.Now() + wfc.Data = append(wfc.Data, pk) + wfc.CVar.Signal() +} + +func (wfc *WriteFileContext) setDone() { + wfc.CVar.L.Lock() + defer wfc.CVar.L.Unlock() + wfc.Done = true + wfc.Data = nil + wfc.CVar.Broadcast() +} + +func (m *MServer) cleanWriteFileContexts() { + now := time.Now() + var staleWfcs []*WriteFileContext + m.Lock.Lock() + for reqId, wfc := range m.WriteFileContextMap { + if now.Sub(wfc.LastActive) > WriteFileContextTimeout { + staleWfcs = append(staleWfcs, wfc) + delete(m.WriteFileContextMap, reqId) + } + } + m.Lock.Unlock() + + // we do this outside of m.Lock just in case there is some lock contention (end of WriteFile could theoretically be slow) + for _, wfc := range staleWfcs { + wfc.setDone() + } +} + +func (m *MServer) ProcessCommandPacket(pk packet.CommandPacketType) { + ck := pk.GetCK() + if ck == "" { + m.Sender.SendMessageFmt("received '%s' packet without ck", pk.GetType()) + return + } + m.Lock.Lock() + cproc := m.ClientMap[ck] + m.Lock.Unlock() + if cproc == nil { + m.Sender.SendCmdError(ck, fmt.Errorf("no client proc for ck '%s', pk=%s", ck, packet.AsString(pk))) + return + } + cproc.Input.SendPacket(pk) + return +} + +func runSingleCompGen(cwd string, compType string, prefix string) ([]string, bool, error) { + if !packet.IsValidCompGenType(compType) { + return nil, false, fmt.Errorf("invalid compgen type '%s'", compType) + } + compGenCmdStr := fmt.Sprintf("cd %s; compgen -A %s -- %s | sort | uniq | head -n %d", shellescape.Quote(cwd), shellescape.Quote(compType), shellescape.Quote(prefix), packet.MaxCompGenValues+1) + ecmd := exec.Command("bash", "-c", compGenCmdStr) + outputBytes, err := ecmd.Output() + if err != nil { + return nil, false, fmt.Errorf("compgen error: %w", err) + } + outputStr := string(outputBytes) + parts := strings.Split(outputStr, "\n") + if len(parts) > 0 && parts[len(parts)-1] == "" { + parts = parts[0 : len(parts)-1] + } + hasMore := false + if len(parts) > packet.MaxCompGenValues { + hasMore = true + parts = parts[0:packet.MaxCompGenValues] + } + return parts, hasMore, nil +} + +func appendSlashes(comps []string) { + for idx, comp := range comps { + comps[idx] = comp + "/" + } +} + +func strArrToMap(strs []string) map[string]bool { + rtn := make(map[string]bool) + for _, s := range strs { + rtn[s] = true + } + return rtn +} + +func (m *MServer) runMixedCompGen(compPk *packet.CompGenPacketType) { + // get directories and files, unique them and put slashes on directories for completion + reqId := compPk.GetReqId() + compDirs, hasMoreDirs, err := runSingleCompGen(compPk.Cwd, "directory", compPk.Prefix) + if err != nil { + m.Sender.SendErrorResponse(reqId, err) + return + } + compFiles, hasMoreFiles, err := runSingleCompGen(compPk.Cwd, compPk.CompType, compPk.Prefix) + if err != nil { + m.Sender.SendErrorResponse(reqId, err) + return + } + + dirMap := strArrToMap(compDirs) + // seed comps with dirs (but append slashes) + comps := compDirs + appendSlashes(comps) + // add files that are not directories (look up in dirMap) + for _, file := range compFiles { + if dirMap[file] { + continue + } + comps = append(comps, file) + } + sort.Strings(comps) // resort + m.Sender.SendResponse(reqId, map[string]interface{}{"comps": comps, "hasmore": (hasMoreFiles || hasMoreDirs)}) + return +} + +func (m *MServer) runCompGen(compPk *packet.CompGenPacketType) { + reqId := compPk.GetReqId() + if compPk.CompType == "file" || compPk.CompType == "command" { + m.runMixedCompGen(compPk) + return + } + comps, hasMore, err := runSingleCompGen(compPk.Cwd, compPk.CompType, compPk.Prefix) + if err != nil { + m.Sender.SendErrorResponse(reqId, err) + return + } + if compPk.CompType == "directory" { + appendSlashes(comps) + } + m.Sender.SendResponse(reqId, map[string]interface{}{"comps": comps, "hasmore": hasMore}) + return +} + +func (m *MServer) setCurrentState(state *packet.ShellState) { + if state == nil { + return + } + hval, _ := state.EncodeAndHash() + m.Lock.Lock() + defer m.Lock.Unlock() + m.StateMap[hval] = state + m.CurrentState = hval +} + +func (m *MServer) reinit(reqId string) { + initPk, err := shexec.MakeServerInitPacket() + if err != nil { + m.Sender.SendErrorResponse(reqId, fmt.Errorf("error creating init packet: %w", err)) + return + } + m.setCurrentState(initPk.State) + initPk.RespId = reqId + m.Sender.SendPacket(initPk) +} + +func makeTemp(path string, mode fs.FileMode) (*os.File, error) { + dirName := filepath.Dir(path) + baseName := filepath.Base(path) + baseTempName := baseName + ".tmp." + writeFd, err := os.CreateTemp(dirName, baseTempName) + if err != nil { + return nil, err + } + err = writeFd.Chmod(mode) + if err != nil { + writeFd.Close() + os.Remove(writeFd.Name()) + return nil, fmt.Errorf("error setting tempfile permissions: %w", err) + } + return writeFd, nil +} + +func checkFileWritable(path string) error { + finfo, err := os.Stat(path) // ok to follow symlinks + if errors.Is(err, fs.ErrNotExist) { + dirName := filepath.Dir(path) + dirInfo, err := os.Stat(dirName) + if err != nil { + return fmt.Errorf("file does not exist, error trying to stat parent directory: %w", err) + } + if !dirInfo.IsDir() { + return fmt.Errorf("file does not exist, parent path [%s] is not a directory", dirName) + } + return nil + } else { + if err != nil { + return fmt.Errorf("cannot stat: %w", err) + } + if finfo.IsDir() { + return fmt.Errorf("invalid path, cannot write a directory") + } + if (finfo.Mode() & fs.ModeSymlink) != 0 { + return fmt.Errorf("writefile does not support symlinks") // note this shouldn't happen because we're using Stat (not Lstat) + } + if (finfo.Mode() & (fs.ModeNamedPipe | fs.ModeSocket | fs.ModeDevice)) != 0 { + return fmt.Errorf("writefile does not support special files (named pipes, sockets, devices): mode=%v", finfo.Mode()) + } + writePerm := (finfo.Mode().Perm() & 0o222) + if writePerm == 0 { + return fmt.Errorf("file is not writable, perms: %v", finfo.Mode().Perm()) + } + return nil + } +} + +func copyFile(dstName string, srcName string) error { + srcFd, err := os.Open(srcName) + if err != nil { + return err + } + defer srcFd.Close() + dstFd, err := os.OpenFile(dstName, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o666) // use 666 because OpenFile respects umask + if err != nil { + return err + } + // we don't defer dstFd.Close() so we can return an error if dstFd.Close() returns an error + _, err = io.Copy(dstFd, srcFd) + if err != nil { + dstFd.Close() + return err + } + return dstFd.Close() +} + +func (m *MServer) writeFile(pk *packet.WriteFilePacketType, wfc *WriteFileContext) { + defer wfc.setDone() + if pk.Path == "" { + resp := packet.MakeWriteFileReadyPacket(pk.ReqId) + resp.Error = "invalid write-file request, no path specified" + m.Sender.SendPacket(resp) + return + } + err := checkFileWritable(pk.Path) + if err != nil { + resp := packet.MakeWriteFileReadyPacket(pk.ReqId) + resp.Error = err.Error() + m.Sender.SendPacket(resp) + return + } + var writeFd *os.File + if pk.UseTemp { + writeFd, err = os.CreateTemp("", "mshell.writefile.*") // "" means make this file in standard TempDir + if err != nil { + resp := packet.MakeWriteFileReadyPacket(pk.ReqId) + resp.Error = fmt.Sprintf("cannot create temp file: %v", err) + m.Sender.SendPacket(resp) + return + } + } else { + writeFd, err = os.OpenFile(pk.Path, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o666) // use 666 because OpenFile respects umask + if err != nil { + resp := packet.MakeWriteFileReadyPacket(pk.ReqId) + resp.Error = fmt.Sprintf("write-file could not open file: %v", err) + m.Sender.SendPacket(resp) + return + } + } + + // ok, so now writeFd is valid, send the "ready" response + resp := packet.MakeWriteFileReadyPacket(pk.ReqId) + m.Sender.SendPacket(resp) + + // now we wait for data (cond var) + // this Unlock() runs first (because it is a later defer) so we can still run wfc.setDone() safely + wfc.CVar.L.Lock() + defer wfc.CVar.L.Unlock() + var doneErr error + for { + if wfc.Done { + break + } + if wfc.Err != nil { + doneErr = wfc.Err + break + } + if len(wfc.Data) == 0 { + wfc.CVar.Wait() + continue + } + dataPk := wfc.Data[0] + wfc.Data = wfc.Data[1:] + if dataPk.Error != "" { + doneErr = fmt.Errorf("error received from client: %v", errors.New(dataPk.Error)) + break + } + if len(dataPk.Data) > 0 { + _, err := writeFd.Write(dataPk.Data) + if err != nil { + doneErr = fmt.Errorf("error writing data to file: %v", err) + break + } + } + if dataPk.Eof { + break + } + } + closeErr := writeFd.Close() + if doneErr == nil && closeErr != nil { + doneErr = fmt.Errorf("error closing file: %v", closeErr) + } + if pk.UseTemp { + if doneErr != nil { + os.Remove(writeFd.Name()) + } else { + // copy file between writeFd.Name() and pk.Path + copyErr := copyFile(pk.Path, writeFd.Name()) + if err != nil { + doneErr = fmt.Errorf("error writing file: %v", copyErr) + } + os.Remove(writeFd.Name()) + } + } + donePk := packet.MakeWriteFileDonePacket(pk.ReqId) + if doneErr != nil { + donePk.Error = doneErr.Error() + } + m.Sender.SendPacket(donePk) +} + +func (m *MServer) returnStreamFileNewFileResponse(pk *packet.StreamFilePacketType) { + // ok, file doesn't exist, so try to check the directory at least to see if we can write a file here + resp := packet.MakeStreamFileResponse(pk.ReqId) + defer func() { + if resp.Error == "" { + resp.Done = true + } + m.Sender.SendPacket(resp) + }() + dirName := filepath.Dir(pk.Path) + dirInfo, err := os.Stat(dirName) + if err != nil { + resp.Error = fmt.Sprintf("file does not exist, error trying to stat parent directory: %v", err) + return + } + if !dirInfo.IsDir() { + resp.Error = fmt.Sprintf("file does not exist, parent path [%s] is not a directory", dirName) + return + } + resp.Info = &packet.FileInfo{ + Name: pk.Path, + Size: 0, + ModTs: 0, + IsDir: false, + Perm: int(dirInfo.Mode().Perm()), + NotFound: true, + } + return +} + +func (m *MServer) streamFile(pk *packet.StreamFilePacketType) { + resp := packet.MakeStreamFileResponse(pk.ReqId) + finfo, err := os.Stat(pk.Path) + if errors.Is(err, fs.ErrNotExist) { + // special return + m.returnStreamFileNewFileResponse(pk) + return + } + if err != nil { + resp.Error = fmt.Sprintf("cannot stat file %q: %v", pk.Path, err) + m.Sender.SendPacket(resp) + return + } + resp.Info = &packet.FileInfo{ + Name: pk.Path, + Size: finfo.Size(), + ModTs: finfo.ModTime().UnixMilli(), + IsDir: finfo.IsDir(), + Perm: int(finfo.Mode().Perm()), + } + if pk.StatOnly { + resp.Done = true + m.Sender.SendPacket(resp) + return + } + // like the http Range header. range header is end inclusive. for us, endByte is non-inclusive (so we add 1) + var startByte, endByte int64 + if len(pk.ByteRange) == 0 { + endByte = finfo.Size() + } else if len(pk.ByteRange) == 1 && pk.ByteRange[0] >= 0 { + startByte = pk.ByteRange[0] + endByte = finfo.Size() + } else if len(pk.ByteRange) == 1 && pk.ByteRange[0] < 0 { + startByte = finfo.Size() + pk.ByteRange[0] // "+" since ByteRange[0] is less than 0 + endByte = finfo.Size() + } else if len(pk.ByteRange) == 2 { + startByte = pk.ByteRange[0] + endByte = pk.ByteRange[1] + 1 + } else { + resp.Error = fmt.Sprintf("invalid byte range (%d entries)", len(pk.ByteRange)) + m.Sender.SendPacket(resp) + return + } + if startByte < 0 { + startByte = 0 + } + if endByte > finfo.Size() { + endByte = finfo.Size() + } + if startByte >= endByte { + resp.Done = true + m.Sender.SendPacket(resp) + return + } + fd, err := os.Open(pk.Path) + if err != nil { + resp.Error = fmt.Sprintf("opening file: %v", err) + m.Sender.SendPacket(resp) + return + } + defer fd.Close() + m.Sender.SendPacket(resp) + var buffer [MaxFileDataPacketSize]byte + var sentDone bool + first := true + for ; startByte < endByte; startByte += MaxFileDataPacketSize { + if !first { + // throttle packet sending @ 1000 packets/s, or 16M/s + time.Sleep(1 * time.Millisecond) + } + first = false + readLen := int64Min(MaxFileDataPacketSize, endByte-startByte) + bufSlice := buffer[0:readLen] + nr, err := fd.ReadAt(bufSlice, startByte) + dataPk := packet.MakeFileDataPacket(pk.ReqId) + dataPk.Data = make([]byte, nr) + copy(dataPk.Data, bufSlice) + if err == io.EOF { + dataPk.Eof = true + } else if err != nil { + dataPk.Error = err.Error() + } + m.Sender.SendPacket(dataPk) + if dataPk.GetResponseDone() { + sentDone = true + break + } + } + if !sentDone { + dataPk := packet.MakeFileDataPacket(pk.ReqId) + dataPk.Eof = true + m.Sender.SendPacket(dataPk) + } + return +} + +func int64Min(v1 int64, v2 int64) int64 { + if v1 < v2 { + return v1 + } + return v2 +} + +func (m *MServer) ProcessRpcPacket(pk packet.RpcPacketType) { + reqId := pk.GetReqId() + if cdPk, ok := pk.(*packet.CdPacketType); ok { + err := os.Chdir(cdPk.Dir) + if err != nil { + m.Sender.SendErrorResponse(reqId, fmt.Errorf("cannot change directory: %w", err)) + return + } + m.Sender.SendResponse(reqId, true) + return + } + if compPk, ok := pk.(*packet.CompGenPacketType); ok { + go m.runCompGen(compPk) + return + } + if _, ok := pk.(*packet.ReInitPacketType); ok { + go m.reinit(reqId) + return + } + if streamPk, ok := pk.(*packet.StreamFilePacketType); ok { + go m.streamFile(streamPk) + return + } + if writePk, ok := pk.(*packet.WriteFilePacketType); ok { + wfc := m.getWriteFileContext(writePk.ReqId) + go m.writeFile(writePk, wfc) + return + } + m.Sender.SendErrorResponse(reqId, fmt.Errorf("invalid rpc type '%s'", pk.GetType())) + return +} + +func (m *MServer) getCurrentState() (string, *packet.ShellState) { + m.Lock.Lock() + defer m.Lock.Unlock() + return m.CurrentState, m.StateMap[m.CurrentState] +} + +func (m *MServer) clientPacketCallback(pk packet.PacketType) { + if pk.GetType() != packet.CmdDonePacketStr { + return + } + donePk := pk.(*packet.CmdDonePacketType) + if donePk.FinalState == nil { + return + } + stateHash, curState := m.getCurrentState() + if curState == nil { + return + } + diff, err := shexec.MakeShellStateDiff(*curState, stateHash, *donePk.FinalState) + if err != nil { + return + } + donePk.FinalState = nil + donePk.FinalStateDiff = &diff +} + +func (m *MServer) runCommand(runPacket *packet.RunPacketType) { + if err := runPacket.CK.Validate("packet"); err != nil { + m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("server run packets require valid ck: %s", err)) + return + } + ecmd, err := shexec.SSHOpts{}.MakeMShellSingleCmd(true) + if err != nil { + m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("server run packets require valid ck: %s", err)) + return + } + cproc, _, err := shexec.MakeClientProc(context.Background(), ecmd) + if err != nil { + m.Sender.SendErrorResponse(runPacket.ReqId, fmt.Errorf("starting mshell client: %s", err)) + return + } + m.Lock.Lock() + m.ClientMap[runPacket.CK] = cproc + m.Lock.Unlock() + go func() { + defer func() { + r := recover() + finalPk := packet.MakeCmdFinalPacket(runPacket.CK) + finalPk.Ts = time.Now().UnixMilli() + if r != nil { + finalPk.Error = fmt.Sprintf("%s", r) + } + m.Sender.SendPacket(finalPk) + m.Lock.Lock() + delete(m.ClientMap, runPacket.CK) + m.Lock.Unlock() + cproc.Close() + }() + shexec.SendRunPacketAndRunData(context.Background(), cproc.Input, runPacket) + cproc.ProxySingleOutput(runPacket.CK, m.Sender, m.clientPacketCallback) + }() +} + +func (m *MServer) packetSenderErrorHandler(sender *packet.PacketSender, pk packet.PacketType, err error) { + if serr, ok := err.(*packet.SendError); ok && serr.IsMarshalError { + msg := packet.MakeMessagePacket(err.Error()) + if cpk, ok := pk.(packet.CommandPacketType); ok { + msg.CK = cpk.GetCK() + } + sender.SendPacket(msg) + return + } else { + // I/O error: close the WriteErrorCh to signal that we are dead (cannot continue if we can't write output) + m.WriteErrorChOnce.Do(func() { + close(m.WriteErrorCh) + }) + } +} + +func (server *MServer) runReadLoop() { + builder := packet.MakeRunPacketBuilder() + for pk := range server.MainInput.MainCh { + if server.Debug { + fmt.Printf("PK> %s\n", packet.AsString(pk)) + } + ok, runPacket := builder.ProcessPacket(pk) + if ok { + if runPacket != nil { + server.runCommand(runPacket) + continue + } + continue + } + if cmdPk, ok := pk.(packet.CommandPacketType); ok { + server.ProcessCommandPacket(cmdPk) + continue + } + if rpcPk, ok := pk.(packet.RpcPacketType); ok { + server.ProcessRpcPacket(rpcPk) + continue + } + if fileDataPk, ok := pk.(*packet.FileDataPacketType); ok { + server.addFileDataPacket(fileDataPk) + continue + } + server.Sender.SendMessageFmt("invalid packet '%s' sent to mshell server", packet.AsString(pk)) + continue + } +} + +func RunServer() (int, error) { + debug := false + if len(os.Args) >= 3 && os.Args[2] == "--debug" { + debug = true + } + server := &MServer{ + Lock: &sync.Mutex{}, + ClientMap: make(map[base.CommandKey]*shexec.ClientProc), + StateMap: make(map[string]*packet.ShellState), + Debug: debug, + WriteErrorCh: make(chan bool), + WriteErrorChOnce: &sync.Once{}, + WriteFileContextMap: make(map[string]*WriteFileContext), + } + go func() { + for { + if server.checkDone() { + return + } + time.Sleep(cleanLoopTime) + server.cleanWriteFileContexts() + } + }() + if debug { + packet.GlobalDebug = true + } + server.MainInput = packet.MakePacketParser(os.Stdin, false) + server.Sender = packet.MakePacketSender(os.Stdout, server.packetSenderErrorHandler) + defer server.Close() + var err error + initPacket, err := shexec.MakeServerInitPacket() + if err != nil { + return 1, err + } + server.setCurrentState(initPacket.State) + server.Sender.SendPacket(initPacket) + ticker := time.NewTicker(1 * time.Minute) + go func() { + for range ticker.C { + server.Sender.SendPacket(packet.MakePingPacket()) + } + }() + defer ticker.Stop() + readLoopDoneCh := make(chan bool) + go func() { + defer close(readLoopDoneCh) + server.runReadLoop() + }() + select { + case <-readLoopDoneCh: + break + + case <-server.WriteErrorCh: + break + } + return 0, nil +} diff --git a/waveshell/pkg/shexec/client.go b/waveshell/pkg/shexec/client.go new file mode 100644 index 00000000..c6292b82 --- /dev/null +++ b/waveshell/pkg/shexec/client.go @@ -0,0 +1,132 @@ +package shexec + +import ( + "context" + "fmt" + "io" + "os/exec" + "time" + + "github.com/commandlinedev/apishell/pkg/base" + "github.com/commandlinedev/apishell/pkg/packet" + "golang.org/x/mod/semver" +) + +// TODO - track buffer sizes for sending input + +const NotFoundVersion = "v0.0" + +type ClientProc struct { + Cmd *exec.Cmd + InitPk *packet.InitPacketType + StartTs time.Time + StdinWriter io.WriteCloser + StdoutReader io.ReadCloser + StderrReader io.ReadCloser + Input *packet.PacketSender + Output *packet.PacketParser +} + +// returns (clientproc, initpk, error) +func MakeClientProc(ctx context.Context, ecmd *exec.Cmd) (*ClientProc, *packet.InitPacketType, error) { + inputWriter, err := ecmd.StdinPipe() + if err != nil { + return nil, nil, fmt.Errorf("creating stdin pipe: %v", err) + } + stdoutReader, err := ecmd.StdoutPipe() + if err != nil { + return nil, nil, fmt.Errorf("creating stdout pipe: %v", err) + } + stderrReader, err := ecmd.StderrPipe() + if err != nil { + return nil, nil, fmt.Errorf("creating stderr pipe: %v", err) + } + startTs := time.Now() + err = ecmd.Start() + if err != nil { + return nil, nil, fmt.Errorf("running local client: %w", err) + } + sender := packet.MakePacketSender(inputWriter, nil) + stdoutPacketParser := packet.MakePacketParser(stdoutReader, false) + stderrPacketParser := packet.MakePacketParser(stderrReader, false) + packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser, true) + cproc := &ClientProc{ + Cmd: ecmd, + StartTs: startTs, + StdinWriter: inputWriter, + StdoutReader: stdoutReader, + StderrReader: stderrReader, + Input: sender, + Output: packetParser, + } + + var pk packet.PacketType + select { + case pk = <-packetParser.MainCh: + case <-ctx.Done(): + cproc.Close() + return nil, nil, ctx.Err() + } + if pk != nil { + if pk.GetType() != packet.InitPacketStr { + cproc.Close() + return nil, nil, fmt.Errorf("invalid packet received from mshell client: %s", packet.AsString(pk)) + } + initPk := pk.(*packet.InitPacketType) + if initPk.NotFound { + cproc.Close() + return nil, initPk, fmt.Errorf("mshell client not found") + } + if semver.MajorMinor(initPk.Version) != semver.MajorMinor(base.MShellVersion) { + cproc.Close() + return nil, initPk, fmt.Errorf("invalid remote mshell version '%s', must be '=%s'", initPk.Version, semver.MajorMinor(base.MShellVersion)) + } + cproc.InitPk = initPk + } + if cproc.InitPk == nil { + cproc.Close() + return nil, nil, fmt.Errorf("no init packet received from mshell client") + } + return cproc, cproc.InitPk, nil +} + +func (cproc *ClientProc) Close() { + if cproc.Input != nil { + cproc.Input.Close() + } + if cproc.StdinWriter != nil { + cproc.StdinWriter.Close() + } + if cproc.StdoutReader != nil { + cproc.StdoutReader.Close() + } + if cproc.StderrReader != nil { + cproc.StderrReader.Close() + } + if cproc.Cmd != nil { + cproc.Cmd.Process.Kill() + } +} + +func (cproc *ClientProc) ProxySingleOutput(ck base.CommandKey, sender *packet.PacketSender, packetCallback func(packet.PacketType)) { + sentDonePk := false + for pk := range cproc.Output.MainCh { + if packetCallback != nil { + packetCallback(pk) + } + if pk.GetType() == packet.CmdDonePacketStr { + sentDonePk = true + } + sender.SendPacket(pk) + } + exitErr := cproc.Cmd.Wait() + if !sentDonePk { + endTs := time.Now() + cmdDuration := endTs.Sub(cproc.StartTs) + donePacket := packet.MakeCmdDonePacket(ck) + donePacket.Ts = endTs.UnixMilli() + donePacket.ExitCode = GetExitCode(exitErr) + donePacket.DurationMs = int64(cmdDuration / time.Millisecond) + sender.SendPacket(donePacket) + } +} diff --git a/waveshell/pkg/shexec/parser.go b/waveshell/pkg/shexec/parser.go new file mode 100644 index 00000000..5123d444 --- /dev/null +++ b/waveshell/pkg/shexec/parser.go @@ -0,0 +1,638 @@ +package shexec + +import ( + "bytes" + "fmt" + "io" + "regexp" + "sort" + "strings" + + "github.com/alessio/shellescape" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/apishell/pkg/simpleexpand" + "github.com/commandlinedev/apishell/pkg/statediff" + "mvdan.cc/sh/v3/expand" + "mvdan.cc/sh/v3/syntax" +) + +const ( + DeclTypeArray = "array" + DeclTypeAssocArray = "assoc" + DeclTypeInt = "int" + DeclTypeNormal = "normal" +) + +type ParseEnviron struct { + Env map[string]string +} + +func (e *ParseEnviron) Get(name string) expand.Variable { + val, ok := e.Env[name] + if !ok { + return expand.Variable{} + } + return expand.Variable{ + Exported: true, + Kind: expand.String, + Str: val, + } +} + +func (e *ParseEnviron) Each(fn func(name string, vr expand.Variable) bool) { + for key, _ := range e.Env { + rtn := fn(key, e.Get(key)) + if !rtn { + break + } + } +} + +func doCmdSubst(commandStr string, w io.Writer, word *syntax.CmdSubst) error { + return nil +} + +func doProcSubst(w *syntax.ProcSubst) (string, error) { + return "", nil +} + +func GetParserConfig(envMap map[string]string) *expand.Config { + cfg := &expand.Config{ + Env: &ParseEnviron{Env: envMap}, + GlobStar: false, + NullGlob: false, + NoUnset: false, + CmdSubst: func(w io.Writer, word *syntax.CmdSubst) error { return doCmdSubst("", w, word) }, + ProcSubst: doProcSubst, + ReadDir: nil, + } + return cfg +} + +func writeIndent(buf *bytes.Buffer, num int) { + for i := 0; i < num; i++ { + buf.WriteByte(' ') + } +} + +func makeSpaceStr(num int) string { + barr := make([]byte, num) + for i := 0; i < num; i++ { + barr[i] = ' ' + } + return string(barr) +} + +// https://wiki.bash-hackers.org/syntax/shellvars +var NoStoreVarNames = map[string]bool{ + "BASH": true, + "BASHOPTS": true, + "BASHPID": true, + "BASH_ALIASES": true, + "BASH_ARGC": true, + "BASH_ARGV": true, + "BASH_ARGV0": true, + "BASH_CMDS": true, + "BASH_COMMAND": true, + "BASH_EXECUTION_STRING": true, + "LINENO": true, + "BASH_LINENO": true, + "BASH_REMATCH": true, + "BASH_SOURCE": true, + "BASH_SUBSHELL": true, + "COPROC": true, + "DIRSTACK": true, + "EPOCHREALTIME": true, + "EPOCHSECONDS": true, + "FUNCNAME": true, + "HISTCMD": true, + "OLDPWD": true, + "PIPESTATUS": true, + "PPID": true, + "PWD": true, + "RANDOM": true, + "SECONDS": true, + "SHLVL": true, + "HISTFILE": true, + "HISTFILESIZE": true, + "HISTCONTROL": true, + "HISTIGNORE": true, + "HISTSIZE": true, + "HISTTIMEFORMAT": true, + "SRANDOM": true, + "COLUMNS": true, + + // we want these in our remote state object + // "EUID": true, + // "SHELLOPTS": true, + // "UID": true, + // "BASH_VERSINFO": true, + // "BASH_VERSION": true, +} + +type DeclareDeclType struct { + Args string + Name string + + // this holds the raw quoted value suitable for bash. this is *not* the real expanded variable value + Value string +} + +var declareDeclArgsRe = regexp.MustCompile("^[aAxrifx]*$") +var bashValidIdentifierRe = regexp.MustCompile("^[a-zA-Z_][a-zA-Z0-9_]*$") + +func (d *DeclareDeclType) Validate() error { + if len(d.Name) == 0 || !IsValidBashIdentifier(d.Name) { + return fmt.Errorf("invalid shell variable name (invalid bash identifier)") + } + if strings.Index(d.Value, "\x00") >= 0 { + return fmt.Errorf("invalid shell variable value (cannot contain 0 byte)") + } + if !declareDeclArgsRe.MatchString(d.Args) { + return fmt.Errorf("invalid shell variable type %s", shellescape.Quote(d.Args)) + } + return nil +} + +func (d *DeclareDeclType) Serialize() string { + return fmt.Sprintf("%s|%s=%s\x00", d.Args, d.Name, d.Value) +} + +func (d *DeclareDeclType) DeclareStmt() string { + var argsStr string + if d.Args == "" { + argsStr = "--" + } else { + argsStr = "-" + d.Args + } + return fmt.Sprintf("declare %s %s=%s", argsStr, d.Name, d.Value) +} + +// envline should be valid +func ParseDeclLine(envLine string) *DeclareDeclType { + eqIdx := strings.Index(envLine, "=") + if eqIdx == -1 { + return nil + } + namePart := envLine[0:eqIdx] + valPart := envLine[eqIdx+1:] + pipeIdx := strings.Index(namePart, "|") + if pipeIdx == -1 { + return nil + } + return &DeclareDeclType{ + Args: namePart[0:pipeIdx], + Name: namePart[pipeIdx+1:], + Value: valPart, + } +} + +// returns name => full-line +func parseDeclLineToKV(envLine string) (string, string) { + decl := ParseDeclLine(envLine) + if decl == nil { + return "", "" + } + return decl.Name, envLine +} + +func shellStateVarsToMap(shellVars []byte) map[string]string { + if len(shellVars) == 0 { + return nil + } + rtn := make(map[string]string) + vars := bytes.Split(shellVars, []byte{0}) + for _, varLine := range vars { + name, val := parseDeclLineToKV(string(varLine)) + if name == "" { + continue + } + rtn[name] = val + } + return rtn +} + +func strMapToShellStateVars(varMap map[string]string) []byte { + var buf bytes.Buffer + orderedKeys := getOrderedKeysStrMap(varMap) + for _, key := range orderedKeys { + val := varMap[key] + buf.WriteString(val) + buf.WriteByte(0) + } + return buf.Bytes() +} + +func getOrderedKeysStrMap(m map[string]string) []string { + keys := make([]string, 0, len(m)) + for key, _ := range m { + keys = append(keys, key) + } + sort.Strings(keys) + return keys +} + +func getOrderedKeysDeclMap(m map[string]*DeclareDeclType) []string { + keys := make([]string, 0, len(m)) + for key, _ := range m { + keys = append(keys, key) + } + sort.Strings(keys) + return keys +} + +func DeclMapFromState(state *packet.ShellState) map[string]*DeclareDeclType { + if state == nil { + return nil + } + rtn := make(map[string]*DeclareDeclType) + vars := bytes.Split(state.ShellVars, []byte{0}) + for _, varLine := range vars { + decl := ParseDeclLine(string(varLine)) + if decl != nil { + rtn[decl.Name] = decl + } + } + return rtn +} + +func SerializeDeclMap(declMap map[string]*DeclareDeclType) []byte { + var rtn bytes.Buffer + orderedKeys := getOrderedKeysDeclMap(declMap) + for _, key := range orderedKeys { + decl := declMap[key] + rtn.WriteString(decl.Serialize()) + } + return rtn.Bytes() +} + +func EnvMapFromState(state *packet.ShellState) map[string]string { + if state == nil { + return nil + } + rtn := make(map[string]string) + ectx := simpleexpand.SimpleExpandContext{} + vars := bytes.Split(state.ShellVars, []byte{0}) + for _, varLine := range vars { + decl := ParseDeclLine(string(varLine)) + if decl != nil && decl.IsExport() { + rtn[decl.Name], _ = simpleexpand.SimpleExpandPartialWord(ectx, decl.Value, false) + } + } + return rtn +} + +func ShellVarMapFromState(state *packet.ShellState) map[string]string { + if state == nil { + return nil + } + rtn := make(map[string]string) + ectx := simpleexpand.SimpleExpandContext{} + vars := bytes.Split(state.ShellVars, []byte{0}) + for _, varLine := range vars { + decl := ParseDeclLine(string(varLine)) + if decl != nil { + rtn[decl.Name], _ = simpleexpand.SimpleExpandPartialWord(ectx, decl.Value, false) + } + } + return rtn +} + +func DumpVarMapFromState(state *packet.ShellState) { + fmt.Printf("DUMP-STATE-VARS:\n") + if state == nil { + fmt.Printf(" nil\n") + return + } + vars := bytes.Split(state.ShellVars, []byte{0}) + for _, varLine := range vars { + fmt.Printf(" %s\n", varLine) + } +} + +func VarDeclsFromState(state *packet.ShellState) []*DeclareDeclType { + if state == nil { + return nil + } + var rtn []*DeclareDeclType + vars := bytes.Split(state.ShellVars, []byte{0}) + for _, varLine := range vars { + decl := ParseDeclLine(string(varLine)) + if decl != nil { + rtn = append(rtn, decl) + } + } + return rtn +} + +func IsValidBashIdentifier(s string) bool { + return bashValidIdentifierRe.MatchString(s) +} + +func (d *DeclareDeclType) IsExport() bool { + return strings.Index(d.Args, "x") >= 0 +} + +func (d *DeclareDeclType) IsReadOnly() bool { + return strings.Index(d.Args, "r") >= 0 +} + +func (d *DeclareDeclType) DataType() string { + if strings.Index(d.Args, "a") >= 0 { + return DeclTypeArray + } + if strings.Index(d.Args, "A") >= 0 { + return DeclTypeAssocArray + } + if strings.Index(d.Args, "i") >= 0 { + return DeclTypeInt + } + return DeclTypeNormal +} + +func parseDeclareStmt(stmt *syntax.Stmt, src string) (*DeclareDeclType, error) { + cmd := stmt.Cmd + decl, ok := cmd.(*syntax.DeclClause) + if !ok || decl.Variant.Value != "declare" || len(decl.Args) != 2 { + return nil, fmt.Errorf("invalid declare variant") + } + rtn := &DeclareDeclType{} + declArgs := decl.Args[0] + if !declArgs.Naked || len(declArgs.Value.Parts) != 1 { + return nil, fmt.Errorf("wrong number of declare args parts") + } + declArgsLit, ok := declArgs.Value.Parts[0].(*syntax.Lit) + if !ok { + return nil, fmt.Errorf("declare args is not a literal") + } + if !strings.HasPrefix(declArgsLit.Value, "-") { + return nil, fmt.Errorf("declare args not an argument (does not start with '-')") + } + if declArgsLit.Value == "--" { + rtn.Args = "" + } else { + rtn.Args = declArgsLit.Value[1:] + } + declAssign := decl.Args[1] + if declAssign.Name == nil { + return nil, fmt.Errorf("declare does not have a valid name") + } + rtn.Name = declAssign.Name.Value + if declAssign.Naked || declAssign.Index != nil || declAssign.Append { + return nil, fmt.Errorf("invalid decl format") + } + if declAssign.Value != nil { + rtn.Value = string(src[declAssign.Value.Pos().Offset():declAssign.Value.End().Offset()]) + } else if declAssign.Array != nil { + rtn.Value = string(src[declAssign.Array.Pos().Offset():declAssign.Array.End().Offset()]) + } else { + return nil, fmt.Errorf("invalid decl, not plain value or array") + } + err := rtn.normalize() + if err != nil { + return nil, err + } + if err = rtn.Validate(); err != nil { + return nil, err + } + return rtn, nil +} + +func parseDeclareOutput(state *packet.ShellState, declareBytes []byte, pvarBytes []byte) error { + declareStr := string(declareBytes) + r := bytes.NewReader(declareBytes) + parser := syntax.NewParser(syntax.Variant(syntax.LangBash)) + file, err := parser.Parse(r, "aliases") + if err != nil { + return err + } + var firstParseErr error + declMap := make(map[string]*DeclareDeclType) + for _, stmt := range file.Stmts { + decl, err := parseDeclareStmt(stmt, declareStr) + if err != nil { + if firstParseErr == nil { + firstParseErr = err + } + } + if decl != nil && !NoStoreVarNames[decl.Name] { + declMap[decl.Name] = decl + } + } + pvars := bytes.Split(pvarBytes, []byte{0}) + for _, pvarBA := range pvars { + pvarStr := string(pvarBA) + pvarFields := strings.SplitN(pvarStr, " ", 2) + if len(pvarFields) != 2 { + continue + } + if pvarFields[0] == "" { + continue + } + decl := &DeclareDeclType{Args: "x"} + decl.Name = "PROMPTVAR_" + pvarFields[0] + decl.Value = shellescape.Quote(pvarFields[1]) + declMap[decl.Name] = decl + } + state.ShellVars = SerializeDeclMap(declMap) // this writes out the decls in a canonical order + if firstParseErr != nil { + state.Error = firstParseErr.Error() + } + return nil +} + +func ParseShellStateOutput(outputBytes []byte) (*packet.ShellState, error) { + // 5 fields: version, cwd, env/vars, aliases, funcs + fields := bytes.Split(outputBytes, []byte{0, 0}) + if len(fields) != 6 { + return nil, fmt.Errorf("invalid shell state output, wrong number of fields, fields=%d", len(fields)) + } + rtn := &packet.ShellState{} + rtn.Version = strings.TrimSpace(string(fields[0])) + if strings.Index(rtn.Version, "bash") == -1 { + return nil, fmt.Errorf("invalid shell state output, only bash is supported") + } + rtn.Version = rtn.Version + cwdStr := string(fields[1]) + if strings.HasSuffix(cwdStr, "\r\n") { + cwdStr = cwdStr[0 : len(cwdStr)-2] + } else if strings.HasSuffix(cwdStr, "\n") { + cwdStr = cwdStr[0 : len(cwdStr)-1] + } + rtn.Cwd = string(cwdStr) + err := parseDeclareOutput(rtn, fields[2], fields[5]) + if err != nil { + return nil, err + } + rtn.Aliases = strings.ReplaceAll(string(fields[3]), "\r\n", "\n") + rtn.Funcs = strings.ReplaceAll(string(fields[4]), "\r\n", "\n") + rtn.Funcs = removeFunc(rtn.Funcs, "_mshell_exittrap") + return rtn, nil +} + +func removeFunc(funcs string, toRemove string) string { + lines := strings.Split(funcs, "\n") + var newLines []string + removeLine := fmt.Sprintf("%s ()", toRemove) + doingRemove := false + for _, line := range lines { + if line == removeLine { + doingRemove = true + continue + } + if doingRemove { + if line == "}" { + doingRemove = false + } + continue + } + newLines = append(newLines, line) + } + return strings.Join(newLines, "\n") +} + +func (d *DeclareDeclType) normalize() error { + if d.DataType() == DeclTypeAssocArray { + return d.normalizeAssocArrayDecl() + } + return nil +} + +// normalizes order of assoc array keys so value is stable +func (d *DeclareDeclType) normalizeAssocArrayDecl() error { + if d.DataType() != DeclTypeAssocArray { + return fmt.Errorf("invalid decltype passed to assocArrayDeclToStr: %s", d.DataType()) + } + varMap, err := assocArrayVarToMap(d) + if err != nil { + return err + } + keys := make([]string, 0, len(varMap)) + for key, _ := range varMap { + keys = append(keys, key) + } + sort.Strings(keys) + var buf bytes.Buffer + buf.WriteByte('(') + for _, key := range keys { + buf.WriteByte('[') + buf.WriteString(key) + buf.WriteByte(']') + buf.WriteByte('=') + buf.WriteString(varMap[key]) + buf.WriteByte(' ') + } + buf.WriteByte(')') + d.Value = buf.String() + return nil +} + +func assocArrayVarToMap(d *DeclareDeclType) (map[string]string, error) { + if d.DataType() != DeclTypeAssocArray { + return nil, fmt.Errorf("decl is not an assoc-array") + } + refStr := "X=" + d.Value + r := strings.NewReader(refStr) + parser := syntax.NewParser(syntax.Variant(syntax.LangBash)) + file, err := parser.Parse(r, "assocdecl") + if err != nil { + return nil, err + } + if len(file.Stmts) != 1 { + return nil, fmt.Errorf("invalid assoc-array parse (multiple stmts)") + } + stmt := file.Stmts[0] + callExpr, ok := stmt.Cmd.(*syntax.CallExpr) + if !ok || len(callExpr.Args) != 0 || len(callExpr.Assigns) != 1 { + return nil, fmt.Errorf("invalid assoc-array parse (bad expr)") + } + assign := callExpr.Assigns[0] + arrayExpr := assign.Array + if arrayExpr == nil { + return nil, fmt.Errorf("invalid assoc-array parse (no array expr)") + } + rtn := make(map[string]string) + for _, elem := range arrayExpr.Elems { + indexStr := refStr[elem.Index.Pos().Offset():elem.Index.End().Offset()] + valStr := refStr[elem.Value.Pos().Offset():elem.Value.End().Offset()] + rtn[indexStr] = valStr + } + return rtn, nil +} + +func strMapsEqual(m1 map[string]string, m2 map[string]string) bool { + if len(m1) != len(m2) { + return false + } + for key, val1 := range m1 { + val2, found := m2[key] + if !found || val1 != val2 { + return false + } + } + for key, _ := range m2 { + _, found := m1[key] + if !found { + return false + } + } + return true +} + +func DeclsEqual(compareName bool, d1 *DeclareDeclType, d2 *DeclareDeclType) bool { + if d1.IsExport() != d2.IsExport() { + return false + } + if d1.DataType() != d2.DataType() { + return false + } + if compareName && d1.Name != d2.Name { + return false + } + return d1.Value == d2.Value // this works even for assoc arrays because we normalize them when parsing +} + +func MakeShellStateDiff(oldState packet.ShellState, oldStateHash string, newState packet.ShellState) (packet.ShellStateDiff, error) { + var rtn packet.ShellStateDiff + rtn.BaseHash = oldStateHash + if oldState.Version != newState.Version { + return rtn, fmt.Errorf("cannot diff, states have different versions") + } + rtn.Version = newState.Version + if oldState.Cwd != newState.Cwd { + rtn.Cwd = newState.Cwd + } + rtn.Error = newState.Error + oldVars := shellStateVarsToMap(oldState.ShellVars) + newVars := shellStateVarsToMap(newState.ShellVars) + rtn.VarsDiff = statediff.MakeMapDiff(oldVars, newVars) + rtn.AliasesDiff = statediff.MakeLineDiff(oldState.Aliases, newState.Aliases) + rtn.FuncsDiff = statediff.MakeLineDiff(oldState.Funcs, newState.Funcs) + return rtn, nil +} + +func ApplyShellStateDiff(oldState packet.ShellState, diff packet.ShellStateDiff) (packet.ShellState, error) { + var rtnState packet.ShellState + var err error + rtnState.Version = oldState.Version + rtnState.Cwd = oldState.Cwd + if diff.Cwd != "" { + rtnState.Cwd = diff.Cwd + } + rtnState.Error = diff.Error + oldVars := shellStateVarsToMap(oldState.ShellVars) + newVars, err := statediff.ApplyMapDiff(oldVars, diff.VarsDiff) + if err != nil { + return rtnState, fmt.Errorf("applying mapdiff 'vars': %v", err) + } + rtnState.ShellVars = strMapToShellStateVars(newVars) + rtnState.Aliases, err = statediff.ApplyLineDiff(oldState.Aliases, diff.AliasesDiff) + if err != nil { + return rtnState, fmt.Errorf("applying diff 'aliases': %v", err) + } + rtnState.Funcs, err = statediff.ApplyLineDiff(oldState.Funcs, diff.FuncsDiff) + if err != nil { + return rtnState, fmt.Errorf("applying diff 'funcs': %v", err) + } + return rtnState, nil +} diff --git a/waveshell/pkg/shexec/shexec.go b/waveshell/pkg/shexec/shexec.go new file mode 100644 index 00000000..d5dc865f --- /dev/null +++ b/waveshell/pkg/shexec/shexec.go @@ -0,0 +1,1542 @@ +package shexec + +import ( + "bytes" + "context" + "encoding/base64" + "fmt" + "io" + "os" + "os/exec" + "os/signal" + "os/user" + "runtime" + "strconv" + "strings" + "sync" + "syscall" + "time" + + "github.com/alessio/shellescape" + "github.com/commandlinedev/apishell/pkg/base" + "github.com/commandlinedev/apishell/pkg/cirfile" + "github.com/commandlinedev/apishell/pkg/mpio" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/creack/pty" + "golang.org/x/mod/semver" + "golang.org/x/sys/unix" +) + +const DefaultTermRows = 24 +const DefaultTermCols = 80 +const MinTermRows = 2 +const MinTermCols = 10 +const MaxTermRows = 1024 +const MaxTermCols = 1024 +const MaxFdNum = 1023 +const FirstExtraFilesFdNum = 3 +const DefaultTermType = "xterm-256color" +const DefaultMaxPtySize = 1024 * 1024 +const MinMaxPtySize = 16 * 1024 +const MaxMaxPtySize = 100 * 1024 * 1024 +const MaxRunDataSize = 1024 * 1024 +const MaxTotalRunDataSize = 10 * MaxRunDataSize +const ShellVarName = "SHELL" + +const GetStateTimeout = 5 * time.Second + +const BaseBashOpts = `set +m; set +H; shopt -s extglob` + +var GetShellStateCmds = []string{ + `echo bash v${BASH_VERSINFO[0]}.${BASH_VERSINFO[1]}.${BASH_VERSINFO[2]};`, + `pwd;`, + `declare -p $(compgen -A variable);`, + `alias -p;`, + `declare -f;`, + `printf "GITBRANCH %s\x00" "$(git rev-parse --abbrev-ref HEAD 2>/dev/null)"`, +} + +const ClientCommandFmt = ` +PATH=$PATH:~/.mshell; +which mshell > /dev/null; +if [[ "$?" -ne 0 ]] +then + printf "\n##N{\"type\": \"init\", \"notfound\": true, \"uname\": \"%s|%s\"}\n" "$(uname -s)" "$(uname -m)" +else + mshell-[%VERSION%] --single +fi +` + +func MakeClientCommandStr() string { + return strings.ReplaceAll(ClientCommandFmt, "[%VERSION%]", semver.MajorMinor(base.MShellVersion)) +} + +const InstallCommandFmt = ` +printf "\n##N{\"type\": \"init\", \"notfound\": true, \"uname\": \"%s|%s\"}\n" "$(uname -s)" "$(uname -m)"; +mkdir -p ~/.mshell/; +cat > ~/.mshell/mshell.temp; +if [[ -s ~/.mshell/mshell.temp ]] +then + mv ~/.mshell/mshell.temp ~/.mshell/mshell-[%VERSION%]; + chmod a+x ~/.mshell/mshell-[%VERSION%]; + ~/.mshell/mshell-[%VERSION%] --single --version +fi +` + +func MakeInstallCommandStr() string { + return strings.ReplaceAll(InstallCommandFmt, "[%VERSION%]", semver.MajorMinor(base.MShellVersion)) +} + +const RunCommandFmt = `%s` +const RunSudoCommandFmt = `sudo -n -C %d bash /dev/fd/%d` +const RunSudoPasswordCommandFmt = `cat /dev/fd/%d | sudo -k -S -C %d bash -c "echo '[from-mshell]'; exec %d>&-; bash /dev/fd/%d < /dev/fd/%d"` + +type MShellBinaryReaderFn func(version string, goos string, goarch string) (io.ReadCloser, error) + +type ReturnStateBuf struct { + Lock *sync.Mutex + Buf []byte + Done bool + Err error + Reader *os.File + FdNum int + DoneCh chan bool +} + +func MakeReturnStateBuf() *ReturnStateBuf { + return &ReturnStateBuf{Lock: &sync.Mutex{}, DoneCh: make(chan bool)} +} + +type ShExecType struct { + Lock *sync.Mutex // only locks "Exited" field + StartTs time.Time + CK base.CommandKey + FileNames *base.CommandFileNames + Cmd *exec.Cmd + CmdPty *os.File + MaxPtySize int64 + Multiplexer *mpio.Multiplexer + Detached bool + DetachedOutput *packet.PacketSender + RunnerOutFd *os.File + MsgSender *packet.PacketSender // where to send out-of-band messages back to calling proceess + ReturnState *ReturnStateBuf + Exited bool // locked via Lock +} + +type StdContext struct{} + +func (StdContext) GetWriter(fdNum int) io.WriteCloser { + if fdNum == 0 { + return os.Stdin + } + if fdNum == 1 { + return os.Stdout + } + if fdNum == 2 { + return os.Stderr + } + fd := os.NewFile(uintptr(fdNum), fmt.Sprintf("/dev/fd/%d", fdNum)) + return fd +} + +func (StdContext) GetReader(fdNum int) io.ReadCloser { + if fdNum == 0 { + return os.Stdin + } + if fdNum == 1 { + return os.Stdout + } + if fdNum == 2 { + return os.Stdout + } + fd := os.NewFile(uintptr(fdNum), fmt.Sprintf("/dev/fd/%d", fdNum)) + return fd +} + +type FdContext interface { + GetWriter(fdNum int) io.WriteCloser + GetReader(fdNum int) io.ReadCloser +} + +type ShExecUPR struct { + ShExec *ShExecType + UPR packet.UnknownPacketReporter +} + +func GetShellStateCmd() string { + return strings.Join(GetShellStateCmds, ` printf "\x00\x00";`) +} + +func (s *ShExecType) processSpecialInputPacket(pk *packet.SpecialInputPacketType) error { + base.Logf("processSpecialInputPacket: %#v\n", pk) + if pk.WinSize != nil { + if s.CmdPty == nil { + return fmt.Errorf("cannot change winsize, cmd was not started with a pty") + } + winSize := &pty.Winsize{ + Rows: uint16(base.BoundInt(pk.WinSize.Rows, MinTermRows, MaxTermRows)), + Cols: uint16(base.BoundInt(pk.WinSize.Cols, MinTermCols, MaxTermCols)), + } + pty.Setsize(s.CmdPty, winSize) + s.Cmd.Process.Signal(syscall.SIGWINCH) + } + if pk.SigName != "" { + var signal syscall.Signal + sigNumInt, err := strconv.Atoi(pk.SigName) + if err == nil { + signal = syscall.Signal(sigNumInt) + } else { + signal = unix.SignalNum(pk.SigName) + } + if signal == 0 { + return fmt.Errorf("error signal %q not found, cannot send", pk.SigName) + } + s.SendSignal(syscall.Signal(signal)) + } + return nil +} + +func (s ShExecUPR) UnknownPacket(pk packet.PacketType) { + if pk.GetType() == packet.SpecialInputPacketStr { + inputPacket := pk.(*packet.SpecialInputPacketType) + err := s.ShExec.processSpecialInputPacket(inputPacket) + if err != nil && s.ShExec.MsgSender != nil { + msg := packet.MakeMessagePacket(err.Error()) + msg.CK = s.ShExec.CK + s.ShExec.MsgSender.SendPacket(msg) + } + return + } + if s.UPR != nil { + s.UPR.UnknownPacket(pk) + } +} + +func MakeShExec(ck base.CommandKey, upr packet.UnknownPacketReporter) *ShExecType { + return &ShExecType{ + Lock: &sync.Mutex{}, + StartTs: time.Now(), + CK: ck, + Multiplexer: mpio.MakeMultiplexer(ck, upr), + } +} + +func (c *ShExecType) Close() { + if c.CmdPty != nil { + c.CmdPty.Close() + } + c.Multiplexer.Close() + if c.DetachedOutput != nil { + c.DetachedOutput.Close() + c.DetachedOutput.WaitForDone() + } + if c.RunnerOutFd != nil { + c.RunnerOutFd.Close() + } + if c.ReturnState != nil { + c.ReturnState.Reader.Close() + } +} + +func (c *ShExecType) MakeCmdStartPacket(reqId string) *packet.CmdStartPacketType { + startPacket := packet.MakeCmdStartPacket(reqId) + startPacket.Ts = time.Now().UnixMilli() + startPacket.CK = c.CK + startPacket.Pid = c.Cmd.Process.Pid + startPacket.MShellPid = os.Getpid() + return startPacket +} + +func getEnvStrKey(envStr string) string { + eqIdx := strings.Index(envStr, "=") + if eqIdx == -1 { + return envStr + } + return envStr[0:eqIdx] +} + +func UpdateCmdEnv(cmd *exec.Cmd, envVars map[string]string) { + if len(envVars) == 0 { + return + } + found := make(map[string]bool) + var newEnv []string + for _, envStr := range cmd.Env { + envKey := getEnvStrKey(envStr) + newEnvVal, ok := envVars[envKey] + if ok { + if newEnvVal == "" { + continue + } + newEnv = append(newEnv, envKey+"="+newEnvVal) + found[envKey] = true + } else { + newEnv = append(newEnv, envStr) + } + } + for envKey, envVal := range envVars { + if found[envKey] { + continue + } + newEnv = append(newEnv, envKey+"="+envVal) + } + cmd.Env = newEnv +} + +// returns (pr, err) +func MakeSimpleStaticWriterPipe(data []byte) (*os.File, error) { + pr, pw, err := os.Pipe() + if err != nil { + return nil, err + } + go func() { + defer pw.Close() + pw.Write(data) + }() + return pr, err +} + +func MakeDetachedExecCmd(pk *packet.RunPacketType, cmdTty *os.File) (*exec.Cmd, error) { + state := pk.State + if state == nil { + state = &packet.ShellState{} + } + ecmd := exec.Command("bash", "-c", pk.Command) + if !pk.StateComplete { + ecmd.Env = os.Environ() + } + UpdateCmdEnv(ecmd, EnvMapFromState(state)) + UpdateCmdEnv(ecmd, MShellEnvVars(getTermType(pk))) + if state.Cwd != "" { + ecmd.Dir = base.ExpandHomeDir(state.Cwd) + } + if HasDupStdin(pk.Fds) { + return nil, fmt.Errorf("cannot detach command with dup stdin") + } + ecmd.Stdin = cmdTty + ecmd.Stdout = cmdTty + ecmd.Stderr = cmdTty + ecmd.SysProcAttr = &syscall.SysProcAttr{ + Setsid: true, + Setctty: true, + } + extraFiles := make([]*os.File, 0, MaxFdNum+1) + if len(pk.Fds) > 0 { + return nil, fmt.Errorf("invalid fd %d passed to detached command", pk.Fds[0].FdNum) + } + for _, runData := range pk.RunData { + if runData.FdNum >= len(extraFiles) { + extraFiles = extraFiles[:runData.FdNum+1] + } + var err error + extraFiles[runData.FdNum], err = MakeSimpleStaticWriterPipe(runData.Data) + if err != nil { + return nil, err + } + } + if len(extraFiles) > FirstExtraFilesFdNum { + ecmd.ExtraFiles = extraFiles[FirstExtraFilesFdNum:] + } + return ecmd, nil +} + +func MakeRunnerExec(ck base.CommandKey) (*exec.Cmd, error) { + msPath, err := base.GetMShellPath() + if err != nil { + return nil, err + } + ecmd := exec.Command(msPath, string(ck)) + return ecmd, nil +} + +// this will never return (unless there is an error creating/opening the file), as fifoFile will never EOF +func MakeAndCopyStdinFifo(dst *os.File, fifoName string) error { + os.Remove(fifoName) + err := syscall.Mkfifo(fifoName, 0600) // only read/write from user for security + if err != nil { + return fmt.Errorf("cannot make stdin-fifo '%s': %v", fifoName, err) + } + // rw is non-blocking, will keep the fifo "open" for the blocking reader + rwfd, err := os.OpenFile(fifoName, os.O_RDWR, 0600) + if err != nil { + return fmt.Errorf("cannot open stdin-fifo(1) '%s': %v", fifoName, err) + } + defer rwfd.Close() + fifoReader, err := os.Open(fifoName) // blocking open/reads (open won't block because of rwfd) + if err != nil { + return fmt.Errorf("cannot open stdin-fifo(2) '%s': %w", fifoName, err) + } + defer fifoReader.Close() + io.Copy(dst, fifoReader) + return nil +} + +func ValidateRunPacket(pk *packet.RunPacketType) error { + if pk.Type != packet.RunPacketStr { + return fmt.Errorf("run packet has wrong type: %s", pk.Type) + } + if pk.Detached { + err := pk.CK.Validate("run packet") + if err != nil { + return err + } + for _, rfd := range pk.Fds { + if rfd.Write { + return fmt.Errorf("cannot detach command with writable remote files fd=%d", rfd.FdNum) + } + if rfd.Read && rfd.DupStdin { + return fmt.Errorf("cannot detach command with dup stdin fd=%d", rfd.FdNum) + } + if rfd.Read { + return fmt.Errorf("cannot detach command with readable remote files fd=%d", rfd.FdNum) + } + } + totalRunData := 0 + for _, rd := range pk.RunData { + if rd.DataLen > MaxRunDataSize { + return fmt.Errorf("cannot detach command, constant rundata input too large fd=%d, len=%d, max=%d", rd.FdNum, rd.DataLen, mpio.ReadBufSize) + } + totalRunData += rd.DataLen + } + if totalRunData > MaxTotalRunDataSize { + return fmt.Errorf("cannot detach command, constant rundata input too large len=%d, max=%d", totalRunData, mpio.MaxTotalRunDataSize) + } + } + if pk.State != nil && pk.State.Cwd != "" { + realCwd := base.ExpandHomeDir(pk.State.Cwd) + dirInfo, err := os.Stat(realCwd) + if err != nil { + return fmt.Errorf("invalid cwd '%s' for command: %v", realCwd, err) + } + if !dirInfo.IsDir() { + return fmt.Errorf("invalid cwd '%s' for command, not a directory", realCwd) + } + } + for _, runData := range pk.RunData { + if runData.DataLen != len(runData.Data) { + return fmt.Errorf("rundata length mismatch, fd=%d, datalen=%d, expected=%d", runData.FdNum, len(runData.Data), runData.DataLen) + } + } + if pk.UsePty && HasDupStdin(pk.Fds) { + return fmt.Errorf("cannot use pty with command that has dup stdin") + } + return nil +} + +func GetWinsize(p *packet.RunPacketType) *pty.Winsize { + rows := DefaultTermRows + cols := DefaultTermCols + if p.TermOpts != nil { + rows = base.BoundInt(p.TermOpts.Rows, MinTermRows, MaxTermRows) + cols = base.BoundInt(p.TermOpts.Cols, MinTermCols, MaxTermCols) + } + return &pty.Winsize{Rows: uint16(rows), Cols: uint16(cols)} +} + +type SSHOpts struct { + SSHHost string + SSHOptsStr string + SSHIdentity string + SSHUser string + SSHPort int + SSHErrorsToTty bool + BatchMode bool +} + +type InstallOpts struct { + SSHOpts SSHOpts + ArchStr string + OptName string + Detect bool + CmdPty *os.File +} + +type ClientOpts struct { + SSHOpts SSHOpts + Command string + Fds []packet.RemoteFd + Cwd string + Debug bool + Sudo bool + SudoWithPass bool + SudoPw string + Detach bool + UsePty bool +} + +func (opts SSHOpts) MakeSSHInstallCmd() (*exec.Cmd, error) { + if opts.SSHHost == "" { + return nil, fmt.Errorf("no ssh host provided, can only install to a remote host") + } + cmdStr := MakeInstallCommandStr() + return opts.MakeSSHExecCmd(cmdStr), 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(fromServer bool) (*exec.Cmd, error) { + if opts.SSHHost == "" { + execFile, err := os.Executable() + if err != nil { + return nil, fmt.Errorf("cannot find local mshell executable: %w", err) + } + var ecmd *exec.Cmd + if fromServer { + ecmd = exec.Command(execFile, "--single-from-server") + } else { + ecmd = exec.Command(execFile, "--single") + } + return ecmd, nil + } + cmdStr := MakeClientCommandStr() + return opts.MakeSSHExecCmd(cmdStr), nil +} + +func (opts SSHOpts) MakeSSHExecCmd(remoteCommand string) *exec.Cmd { + remoteCommand = strings.TrimSpace(remoteCommand) + if opts.SSHHost == "" { + homeDir, _ := os.UserHomeDir() // ignore error + if homeDir == "" { + homeDir = "/" + } + ecmd := exec.Command("bash", "-c", remoteCommand) + ecmd.Dir = homeDir + return ecmd + } else { + var moreSSHOpts []string + if opts.SSHIdentity != "" { + identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHIdentity)) + moreSSHOpts = append(moreSSHOpts, identityOpt) + } + if opts.SSHUser != "" { + userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHUser)) + moreSSHOpts = append(moreSSHOpts, userOpt) + } + if opts.SSHPort != 0 { + portOpt := fmt.Sprintf("-p %d", opts.SSHPort) + moreSSHOpts = append(moreSSHOpts, portOpt) + } + if opts.SSHErrorsToTty { + errFdStr := "-E /dev/tty" + moreSSHOpts = append(moreSSHOpts, errFdStr) + } + if opts.BatchMode { + batchOpt := "-o 'BatchMode=yes'" + moreSSHOpts = append(moreSSHOpts, batchOpt) + } + // note that SSHOptsStr is *not* escaped + sshCmd := fmt.Sprintf("ssh %s %s %s %s", strings.Join(moreSSHOpts, " "), opts.SSHOptsStr, shellescape.Quote(opts.SSHHost), shellescape.Quote(remoteCommand)) + ecmd := exec.Command("bash", "-c", sshCmd) + return ecmd + } +} + +func (opts SSHOpts) MakeMShellSSHOpts() string { + var moreSSHOpts []string + if opts.SSHIdentity != "" { + identityOpt := fmt.Sprintf("-i %s", shellescape.Quote(opts.SSHIdentity)) + moreSSHOpts = append(moreSSHOpts, identityOpt) + } + if opts.SSHUser != "" { + userOpt := fmt.Sprintf("-l %s", shellescape.Quote(opts.SSHUser)) + moreSSHOpts = append(moreSSHOpts, userOpt) + } + if opts.SSHPort != 0 { + portOpt := fmt.Sprintf("-p %d", opts.SSHPort) + moreSSHOpts = append(moreSSHOpts, portOpt) + } + if opts.SSHOptsStr != "" { + optsOpt := fmt.Sprintf("--ssh-opts %s", shellescape.Quote(opts.SSHOptsStr)) + moreSSHOpts = append(moreSSHOpts, optsOpt) + } + if opts.SSHHost != "" { + sshArg := fmt.Sprintf("--ssh %s", shellescape.Quote(opts.SSHHost)) + moreSSHOpts = append(moreSSHOpts, sshArg) + } + return strings.Join(moreSSHOpts, " ") +} + +func GetTerminalSize() (int, int, error) { + fd, err := os.Open("/dev/tty") + if err != nil { + return 0, 0, err + } + defer fd.Close() + return pty.Getsize(fd) +} + +func (opts *ClientOpts) MakeRunPacket() (*packet.RunPacketType, error) { + runPacket := packet.MakeRunPacket() + runPacket.Detached = opts.Detach + runPacket.State = &packet.ShellState{} + runPacket.State.Cwd = opts.Cwd + runPacket.Fds = opts.Fds + if opts.UsePty { + runPacket.UsePty = true + runPacket.TermOpts = &packet.TermOpts{} + rows, cols, err := GetTerminalSize() + if err == nil { + runPacket.TermOpts.Rows = rows + runPacket.TermOpts.Cols = cols + } + term := os.Getenv("TERM") + if term != "" { + runPacket.TermOpts.Term = term + } + } + if !opts.Sudo { + // normal, non-sudo command + runPacket.Command = fmt.Sprintf(RunCommandFmt, opts.Command) + return runPacket, nil + } + if opts.SudoWithPass { + pwFdNum, err := AddRunData(runPacket, opts.SudoPw, "sudo pw") + if err != nil { + return nil, err + } + commandFdNum, err := AddRunData(runPacket, opts.Command, "command") + if err != nil { + return nil, err + } + commandStdinFdNum, err := NextFreeFdNum(runPacket) + if err != nil { + return nil, err + } + commandStdinRfd := packet.RemoteFd{FdNum: commandStdinFdNum, Read: true, DupStdin: true} + runPacket.Fds = append(runPacket.Fds, commandStdinRfd) + maxFdNum := MaxFdNumInPacket(runPacket) + runPacket.Command = fmt.Sprintf(RunSudoPasswordCommandFmt, pwFdNum, maxFdNum+1, pwFdNum, commandFdNum, commandStdinFdNum) + return runPacket, nil + } else { + commandFdNum, err := AddRunData(runPacket, opts.Command, "command") + if err != nil { + return nil, err + } + maxFdNum := MaxFdNumInPacket(runPacket) + runPacket.Command = fmt.Sprintf(RunSudoCommandFmt, maxFdNum+1, commandFdNum) + return runPacket, nil + } +} + +func AddRunData(pk *packet.RunPacketType, data string, dataType string) (int, error) { + if len(data) > MaxRunDataSize { + return 0, fmt.Errorf("%s too large, exceeds read buffer size size:%d", dataType, len(data)) + } + fdNum, err := NextFreeFdNum(pk) + if err != nil { + return 0, err + } + runData := packet.RunDataType{FdNum: fdNum, DataLen: len(data), Data: []byte(data)} + pk.RunData = append(pk.RunData, runData) + return fdNum, nil +} + +func NextFreeFdNum(pk *packet.RunPacketType) (int, error) { + fdMap := make(map[int]bool) + for _, fd := range pk.Fds { + fdMap[fd.FdNum] = true + } + for _, rd := range pk.RunData { + fdMap[rd.FdNum] = true + } + for i := 3; i <= MaxFdNum; i++ { + if !fdMap[i] { + return i, nil + } + } + return 0, fmt.Errorf("reached maximum number of fds, all fds between 3-%d are in use", MaxFdNum) +} + +func MaxFdNumInPacket(pk *packet.RunPacketType) int { + maxFdNum := 3 + for _, fd := range pk.Fds { + if fd.FdNum > maxFdNum { + maxFdNum = fd.FdNum + } + } + for _, rd := range pk.RunData { + if rd.FdNum > maxFdNum { + maxFdNum = rd.FdNum + } + } + return maxFdNum +} + +func ValidateRemoteFds(rfds []packet.RemoteFd) error { + dupMap := make(map[int]bool) + for _, rfd := range rfds { + if rfd.FdNum < 0 { + return fmt.Errorf("mshell negative fd numbers fd=%d", rfd.FdNum) + } + if rfd.FdNum < FirstExtraFilesFdNum { + return fmt.Errorf("mshell does not support re-opening fd=%d (0, 1, and 2, are always open)", rfd.FdNum) + } + if rfd.FdNum > MaxFdNum { + return fmt.Errorf("mshell does not support opening fd numbers above %d", MaxFdNum) + } + if dupMap[rfd.FdNum] { + return fmt.Errorf("mshell got duplicate entries for fd=%d", rfd.FdNum) + } + if rfd.Read && rfd.Write { + return fmt.Errorf("mshell does not support opening fd numbers for reading and writing, fd=%d", rfd.FdNum) + } + if !rfd.Read && !rfd.Write { + return fmt.Errorf("invalid fd=%d, neither reading or writing mode specified", rfd.FdNum) + } + dupMap[rfd.FdNum] = true + } + return nil +} + +func sendMShellBinary(input io.WriteCloser, mshellStream io.Reader) { + go func() { + defer input.Close() + io.Copy(input, mshellStream) + }() +} + +func RunInstallFromCmd(ctx context.Context, ecmd *exec.Cmd, tryDetect bool, mshellStream io.Reader, mshellReaderFn MShellBinaryReaderFn, msgFn func(string)) error { + inputWriter, err := ecmd.StdinPipe() + if err != nil { + return fmt.Errorf("creating stdin pipe: %v", err) + } + stdoutReader, err := ecmd.StdoutPipe() + if err != nil { + return fmt.Errorf("creating stdout pipe: %v", err) + } + stderrReader, err := ecmd.StderrPipe() + if err != nil { + return fmt.Errorf("creating stderr pipe: %v", err) + } + go func() { + io.Copy(os.Stderr, stderrReader) + }() + if mshellStream != nil { + sendMShellBinary(inputWriter, mshellStream) + } + packetParser := packet.MakePacketParser(stdoutReader, false) + err = ecmd.Start() + if err != nil { + return fmt.Errorf("running ssh command: %w", err) + } + firstInit := true + for { + var pk packet.PacketType + select { + case pk = <-packetParser.MainCh: + case <-ctx.Done(): + return ctx.Err() + } + if pk == nil { + return fmt.Errorf("no response packet received from client") + } + if pk.GetType() == packet.InitPacketStr && firstInit { + firstInit = false + initPacket := pk.(*packet.InitPacketType) + if !tryDetect { + continue // ignore + } + tryDetect = false + if initPacket.UName == "" { + return fmt.Errorf("cannot detect arch, no uname received from remote server") + } + goos, goarch, err := DetectGoArch(initPacket.UName) + if err != nil { + return fmt.Errorf("arch cannot be detected (might be incompatible with mshell): %w", err) + } + msgStr := fmt.Sprintf("mshell detected remote architecture as '%s.%s'\n", goos, goarch) + msgFn(msgStr) + detectedMSS, err := mshellReaderFn(base.MShellVersion, goos, goarch) + if err != nil { + return err + } + defer detectedMSS.Close() + sendMShellBinary(inputWriter, detectedMSS) + continue + } + if pk.GetType() == packet.InitPacketStr && !firstInit { + initPacket := pk.(*packet.InitPacketType) + if initPacket.Version == base.MShellVersion { + return nil + } + return fmt.Errorf("invalid version '%s' received from client, expecting '%s'", initPacket.Version, base.MShellVersion) + } + if pk.GetType() == packet.RawPacketStr { + rawPk := pk.(*packet.RawPacketType) + msgFn(fmt.Sprintf("%s\n", rawPk.Data)) + continue + } + return fmt.Errorf("invalid response packet '%s' received from client", pk.GetType()) + } + return fmt.Errorf("did not receive version string from client, install not successful") +} + +func RunInstallFromOpts(opts *InstallOpts) error { + ecmd, err := opts.SSHOpts.MakeSSHInstallCmd() + if err != nil { + return err + } + msgFn := func(str string) { + fmt.Printf("%s", str) + } + var mshellStream *os.File + if opts.OptName != "" { + mshellStream, err = os.Open(opts.OptName) + if err != nil { + return fmt.Errorf("cannot open mshell binary %q: %v", opts.OptName, err) + } + defer mshellStream.Close() + } + err = RunInstallFromCmd(context.Background(), ecmd, opts.Detect, mshellStream, base.MShellBinaryFromOptDir, msgFn) + if err != nil { + return err + } + mmVersion := semver.MajorMinor(base.MShellVersion) + fmt.Printf("mshell installed successfully at %s:~/.mshell/mshell%s\n", opts.SSHOpts.SSHHost, mmVersion) + return nil +} + +func HasDupStdin(fds []packet.RemoteFd) bool { + for _, rfd := range fds { + if rfd.Read && rfd.DupStdin { + return true + } + } + return false +} + +func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdContext, sshOpts SSHOpts, upr packet.UnknownPacketReporter, debug bool) (*packet.CmdDonePacketType, error) { + cmd := MakeShExec(runPacket.CK, upr) + ecmd, err := sshOpts.MakeMShellSingleCmd(false) + if err != nil { + return nil, err + } + 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) + } + if !HasDupStdin(runPacket.Fds) { + cmd.Multiplexer.MakeRawFdReader(0, fdContext.GetReader(0), false, false) + } + cmd.Multiplexer.MakeRawFdWriter(1, fdContext.GetWriter(1), false, "client") + cmd.Multiplexer.MakeRawFdWriter(2, fdContext.GetWriter(2), false, "client") + for _, rfd := range runPacket.Fds { + if rfd.Read && rfd.DupStdin { + cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, fdContext.GetReader(0), false, false) + continue + } + if rfd.Read { + fd := fdContext.GetReader(rfd.FdNum) + cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, fd, false, false) + } else if rfd.Write { + fd := fdContext.GetWriter(rfd.FdNum) + cmd.Multiplexer.MakeRawFdWriter(rfd.FdNum, fd, true, "client") + } + } + err = ecmd.Start() + if err != nil { + return nil, fmt.Errorf("running ssh command: %w", err) + } + defer cmd.Close() + stdoutPacketParser := packet.MakePacketParser(stdoutReader, false) + stderrPacketParser := packet.MakePacketParser(stderrReader, false) + packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser, false) + sender := packet.MakePacketSender(inputWriter, nil) + versionOk := false + for pk := range packetParser.MainCh { + 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) + mmVersion := semver.MajorMinor(base.MShellVersion) + if initPk.NotFound { + if sshOpts.SSHHost == "" { + return nil, fmt.Errorf("mshell-%s command not found on local server", mmVersion) + } + if initPk.UName == "" { + return nil, fmt.Errorf("mshell-%s command not found on remote server, no uname detected", mmVersion) + } + goos, goarch, err := DetectGoArch(initPk.UName) + if err != nil { + return nil, fmt.Errorf("mshell-%s command not found on remote server, architecture cannot be detected (might be incompatible with mshell): %w", mmVersion, err) + } + sshOptsStr := sshOpts.MakeMShellSSHOpts() + return nil, fmt.Errorf("mshell-%s command not found on remote server, can install with 'mshell --install %s %s.%s'", mmVersion, sshOptsStr, goos, goarch) + } + if semver.MajorMinor(initPk.Version) != semver.MajorMinor(base.MShellVersion) { + return nil, fmt.Errorf("invalid remote mshell version '%s', must be '=%s'", initPk.Version, semver.MajorMinor(base.MShellVersion)) + } + versionOk = true + if debug { + fmt.Printf("VERSION> %s\n", initPk.Version) + } + break + } + } + if !versionOk { + return nil, fmt.Errorf("did not receive version from remote mshell") + } + SendRunPacketAndRunData(context.Background(), sender, runPacket) + if debug { + cmd.Multiplexer.Debug = true + } + remoteDonePacket := cmd.Multiplexer.RunIOAndWait(packetParser, sender, false, true, true) + donePacket := cmd.WaitForCommand() + if remoteDonePacket != nil { + donePacket = remoteDonePacket + } + return donePacket, nil +} + +func min(v1 int, v2 int) int { + if v1 <= v2 { + return v1 + } + return v2 +} + +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 nil + } + for _, runData := range runPacket.RunData { + sendBuf := runData.Data + for len(sendBuf) > 0 { + chunkSize := min(len(sendBuf), mpio.MaxSingleWriteSize) + chunk := sendBuf[0:chunkSize] + dataPk := packet.MakeDataPacket() + dataPk.CK = runPacket.CK + dataPk.FdNum = runData.FdNum + dataPk.Data64 = base64.StdEncoding.EncodeToString(chunk) + dataPk.Eof = (len(chunk) == len(sendBuf)) + sendBuf = sendBuf[chunkSize:] + err = sender.SendPacketCtx(ctx, dataPk) + if err != nil { + return err + } + } + } + err = sender.SendPacketCtx(ctx, packet.MakeDataEndPacket(runPacket.CK)) + if err != nil { + return err + } + return nil +} + +func DetectGoArch(uname string) (string, string, error) { + fields := strings.SplitN(uname, "|", 2) + if len(fields) != 2 { + return "", "", fmt.Errorf("invalid uname string returned") + } + osVal := strings.TrimSpace(strings.ToLower(fields[0])) + archVal := strings.TrimSpace(strings.ToLower(fields[1])) + if osVal != "darwin" && osVal != "linux" { + return "", "", fmt.Errorf("invalid uname OS '%s', mshell only supports OS X (darwin) and linux", osVal) + } + goos := osVal + goarch := "" + if archVal == "x86_64" || archVal == "i686" || archVal == "amd64" { + goarch = "amd64" + } else if archVal == "aarch64" || archVal == "arm64" { + goarch = "arm64" + } + if goarch == "" { + return "", "", fmt.Errorf("invalid uname machine type '%s', mshell only supports aarch64 (amd64) and x86_64 (amd64)", archVal) + } + if !base.ValidGoArch(goos, goarch) { + return "", "", fmt.Errorf("invalid arch detected %s.%s", goos, goarch) + } + return goos, goarch, nil +} + +func (cmd *ShExecType) RunRemoteIOAndWait(packetParser *packet.PacketParser, sender *packet.PacketSender) { + defer cmd.Close() + if cmd.ReturnState != nil { + go cmd.ReturnState.Run() + } + cmd.Multiplexer.RunIOAndWait(packetParser, sender, true, false, false) + donePacket := cmd.WaitForCommand() + sender.SendPacket(donePacket) +} + +func getTermType(pk *packet.RunPacketType) string { + termType := DefaultTermType + if pk.TermOpts != nil && pk.TermOpts.Term != "" { + termType = pk.TermOpts.Term + } + return termType +} + +func makeRcFileStr(pk *packet.RunPacketType) string { + var rcBuf bytes.Buffer + rcBuf.WriteString(BaseBashOpts + "\n") + varDecls := VarDeclsFromState(pk.State) + for _, varDecl := range varDecls { + if varDecl.IsExport() || varDecl.IsReadOnly() { + continue + } + rcBuf.WriteString(varDecl.DeclareStmt()) + rcBuf.WriteString("\n") + } + if pk.State != nil && pk.State.Funcs != "" { + rcBuf.WriteString(pk.State.Funcs) + rcBuf.WriteString("\n") + } + if pk.State != nil && pk.State.Aliases != "" { + rcBuf.WriteString(pk.State.Aliases) + rcBuf.WriteString("\n") + } + return rcBuf.String() +} + +func makeExitTrap(fdNum int) string { + stateCmd := GetShellStateRedirectCommandStr(fdNum) + fmtStr := ` +_mshell_exittrap () { + %s +} +trap _mshell_exittrap EXIT +` + return fmt.Sprintf(fmtStr, stateCmd) +} + +func (s *ShExecType) SendSignal(sig syscall.Signal) { + base.Logf("signal start\n") + if s.Cmd == nil || s.Cmd.Process == nil || s.IsExited() { + return + } + pgroup := false + if s.Cmd.SysProcAttr != nil && (s.Cmd.SysProcAttr.Setsid || s.Cmd.SysProcAttr.Setpgid) { + pgroup = true + } + pid := s.Cmd.Process.Pid + if pgroup { + base.Logf("send signal %s to %d (pgroup)\n", sig, -pid) + syscall.Kill(-pid, sig) + } else { + base.Logf("send signal %s to %d (normal)\n", sig, pid) + syscall.Kill(pid, sig) + } +} + +func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fromServer bool) (rtnShExec *ShExecType, rtnErr error) { + state := pk.State + if state == nil { + state = &packet.ShellState{} + } + cmd := MakeShExec(pk.CK, nil) + defer func() { + // on error, call cmd.Close() + if rtnErr != nil { + cmd.Close() + } + }() + if fromServer { + msgUpr := packet.MessageUPR{CK: pk.CK, Sender: sender} + upr := ShExecUPR{ShExec: cmd, UPR: msgUpr} + cmd.Multiplexer.UPR = upr + cmd.MsgSender = sender + } + var rtnStateWriter *os.File + rcFileStr := makeRcFileStr(pk) + if pk.ReturnState { + pr, pw, err := os.Pipe() + if err != nil { + return nil, fmt.Errorf("cannot create returnstate pipe: %v", err) + } + cmd.ReturnState = MakeReturnStateBuf() + cmd.ReturnState.Reader = pr + cmd.ReturnState.FdNum = 20 + rtnStateWriter = pw + defer pw.Close() + trapCmdStr := makeExitTrap(cmd.ReturnState.FdNum) + rcFileStr += trapCmdStr + } + shellVarMap := ShellVarMapFromState(state) + if base.HasDebugFlag(shellVarMap, base.DebugFlag_LogRcFile) { + debugRcFileName := base.GetDebugRcFileName() + err := os.WriteFile(debugRcFileName, []byte(rcFileStr), 0600) + if err != nil { + base.Logf("error writing %s: %v\n", debugRcFileName, err) + } + } + rcFileFdNum, err := AddRunData(pk, rcFileStr, "rcfile") + if err != nil { + return nil, err + } + if pk.UsePty { + cmd.Cmd = exec.Command("bash", "--rcfile", fmt.Sprintf("/dev/fd/%d", rcFileFdNum), "-i", "-c", pk.Command) + } else { + cmd.Cmd = exec.Command("bash", "--rcfile", fmt.Sprintf("/dev/fd/%d", rcFileFdNum), "-c", pk.Command) + } + if !pk.StateComplete { + cmd.Cmd.Env = os.Environ() + } + UpdateCmdEnv(cmd.Cmd, EnvMapFromState(state)) + if state.Cwd != "" { + cmd.Cmd.Dir = base.ExpandHomeDir(state.Cwd) + } + err = ValidateRemoteFds(pk.Fds) + if err != nil { + return nil, err + } + var cmdPty *os.File + var cmdTty *os.File + if pk.UsePty { + cmdPty, cmdTty, err = pty.Open() + if err != nil { + return nil, fmt.Errorf("opening new pty: %w", err) + } + pty.Setsize(cmdPty, GetWinsize(pk)) + defer func() { + cmdTty.Close() + }() + cmd.CmdPty = cmdPty + UpdateCmdEnv(cmd.Cmd, MShellEnvVars(getTermType(pk))) + } + if cmdTty != nil { + cmd.Cmd.Stdin = cmdTty + cmd.Cmd.Stdout = cmdTty + cmd.Cmd.Stderr = cmdTty + cmd.Cmd.SysProcAttr = &syscall.SysProcAttr{ + Setsid: true, + Setctty: true, + } + cmd.Multiplexer.MakeRawFdWriter(0, cmdPty, false, "simple") + cmd.Multiplexer.MakeRawFdReader(1, cmdPty, false, true) + nullFd, err := os.Open("/dev/null") + if err != nil { + return nil, fmt.Errorf("cannot open /dev/null: %w", err) + } + cmd.Multiplexer.MakeRawFdReader(2, nullFd, true, false) + } else { + cmd.Cmd.Stdin, err = cmd.Multiplexer.MakeWriterPipe(0, "simple") + if err != nil { + return nil, err + } + cmd.Cmd.Stdout, err = cmd.Multiplexer.MakeReaderPipe(1) + if err != nil { + return nil, err + } + cmd.Cmd.Stderr, err = cmd.Multiplexer.MakeReaderPipe(2) + if err != nil { + return nil, err + } + } + extraFiles := make([]*os.File, 0, MaxFdNum+1) + for _, runData := range pk.RunData { + if runData.FdNum >= len(extraFiles) { + extraFiles = extraFiles[:runData.FdNum+1] + } + extraFiles[runData.FdNum], err = cmd.Multiplexer.MakeStaticWriterPipe(runData.FdNum, runData.Data, MaxRunDataSize, "simple-rundata") + if err != nil { + return nil, err + } + } + for _, rfd := range pk.Fds { + if rfd.FdNum >= len(extraFiles) { + extraFiles = extraFiles[:rfd.FdNum+1] + } + if rfd.Read { + // client file is open for reading, so we make a writer pipe + extraFiles[rfd.FdNum], err = cmd.Multiplexer.MakeWriterPipe(rfd.FdNum, "simple") + if err != nil { + return nil, err + } + } + if rfd.Write { + // client file is open for writing, so we make a reader pipe + extraFiles[rfd.FdNum], err = cmd.Multiplexer.MakeReaderPipe(rfd.FdNum) + if err != nil { + return nil, err + } + } + } + if cmd.ReturnState != nil { + if cmd.ReturnState.FdNum >= len(extraFiles) { + extraFiles = extraFiles[:cmd.ReturnState.FdNum+1] + } + extraFiles[cmd.ReturnState.FdNum] = rtnStateWriter + } + if len(extraFiles) > FirstExtraFilesFdNum { + cmd.Cmd.ExtraFiles = extraFiles[FirstExtraFilesFdNum:] + } + err = cmd.Cmd.Start() + if err != nil { + return nil, err + } + return cmd, nil +} + +// TODO limit size of read state buffer +func (rs *ReturnStateBuf) Run() { + buf := make([]byte, 1024) + defer func() { + rs.Lock.Lock() + defer rs.Lock.Unlock() + rs.Reader.Close() + rs.Done = true + close(rs.DoneCh) + }() + for { + n, readErr := rs.Reader.Read(buf) + if readErr == io.EOF { + break + } + if readErr != nil { + rs.Lock.Lock() + rs.Err = readErr + rs.Lock.Unlock() + break + } + rs.Lock.Lock() + rs.Buf = append(rs.Buf, buf[0:n]...) + rs.Lock.Unlock() + } +} + +// in detached run mode, we don't want mshell to die from signals +// since we want mshell to persist even if the mshell --server is terminated +func SetupSignalsForDetach() { + sigCh := make(chan os.Signal, 1) + signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM, syscall.SIGHUP, syscall.SIGPIPE) + go func() { + for range sigCh { + // do nothing + } + }() +} + +// in detached run mode, we don't want mshell to die from signals +// since we want mshell to persist even if the mshell --server is terminated +func IgnoreSigPipe() { + sigCh := make(chan os.Signal, 1) + signal.Notify(sigCh, syscall.SIGPIPE) + go func() { + for sig := range sigCh { + base.Logf("ignoring signal %v\n", sig) + } + }() +} + +func copyToCirFile(dest *cirfile.File, src io.Reader) error { + buf := make([]byte, 64*1024) + for { + var appendErr error + nr, readErr := src.Read(buf) + if nr > 0 { + appendErr = dest.AppendData(context.Background(), buf[0:nr]) + } + if readErr != nil && readErr != io.EOF { + return readErr + } + if appendErr != nil { + return appendErr + } + if readErr == io.EOF { + return nil + } + } +} + +func (cmd *ShExecType) DetachedWait(startPacket *packet.CmdStartPacketType) { + // after Start(), any output/errors must go to DetachedOutput + // close stdin, redirect stdout/stderr to /dev/null, but wait for cmdstart packet to get sent + cmd.DetachedOutput.SendPacket(startPacket) + err := os.Stdin.Close() + if err != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot close stdin: %w", err)) + } + err = unix.Dup2(int(cmd.RunnerOutFd.Fd()), int(os.Stdout.Fd())) + if err != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot dup2 stdin to runout: %w", err)) + } + err = unix.Dup2(int(cmd.RunnerOutFd.Fd()), int(os.Stderr.Fd())) + if err != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot dup2 stdin to runout: %w", err)) + } + ptyOutFile, err := cirfile.CreateCirFile(cmd.FileNames.PtyOutFile, cmd.MaxPtySize) + if err != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("cannot open ptyout file '%s': %w", cmd.FileNames.PtyOutFile, err)) + // don't return (command is already running) + } + ptyCopyDone := make(chan bool) + go func() { + // copy pty output to .ptyout file + defer close(ptyCopyDone) + defer ptyOutFile.Close() + copyErr := copyToCirFile(ptyOutFile, cmd.CmdPty) + if copyErr != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("copying pty output to ptyout file: %w", copyErr)) + } + }() + go func() { + // copy .stdin fifo contents to pty input + copyFifoErr := MakeAndCopyStdinFifo(cmd.CmdPty, cmd.FileNames.StdinFifo) + if copyFifoErr != nil { + cmd.DetachedOutput.SendCmdError(cmd.CK, fmt.Errorf("reading from stdin fifo: %w", copyFifoErr)) + } + }() + donePacket := cmd.WaitForCommand() + cmd.DetachedOutput.SendPacket(donePacket) + <-ptyCopyDone + cmd.Close() + return +} + +func RunCommandDetached(pk *packet.RunPacketType, sender *packet.PacketSender) (*ShExecType, *packet.CmdStartPacketType, error) { + fileNames, err := base.GetCommandFileNames(pk.CK) + if err != nil { + return nil, nil, err + } + runOutInfo, err := os.Stat(fileNames.RunnerOutFile) + if err == nil { // non-nil error will be caught by regular OpenFile below + // must have size 0 + if runOutInfo.Size() != 0 { + return nil, nil, fmt.Errorf("cmdkey '%s' was already used (runout len=%d)", pk.CK, runOutInfo.Size()) + } + } + cmdPty, cmdTty, err := pty.Open() + if err != nil { + return nil, nil, fmt.Errorf("opening new pty: %w", err) + } + pty.Setsize(cmdPty, GetWinsize(pk)) + defer func() { + cmdTty.Close() + }() + cmd := MakeShExec(pk.CK, nil) + cmd.FileNames = fileNames + cmd.CmdPty = cmdPty + cmd.Detached = true + cmd.MaxPtySize = DefaultMaxPtySize + if pk.TermOpts != nil && pk.TermOpts.MaxPtySize > 0 { + cmd.MaxPtySize = base.BoundInt64(pk.TermOpts.MaxPtySize, MinMaxPtySize, MaxMaxPtySize) + } + cmd.RunnerOutFd, err = os.OpenFile(fileNames.RunnerOutFile, os.O_TRUNC|os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0600) + if err != nil { + return nil, nil, fmt.Errorf("cannot open runout file '%s': %w", fileNames.RunnerOutFile, err) + } + cmd.DetachedOutput = packet.MakePacketSender(cmd.RunnerOutFd, nil) + ecmd, err := MakeDetachedExecCmd(pk, cmdTty) + if err != nil { + return nil, nil, err + } + cmd.Cmd = ecmd + SetupSignalsForDetach() + err = ecmd.Start() + if err != nil { + return nil, nil, fmt.Errorf("starting command: %w", err) + } + for _, fd := range ecmd.ExtraFiles { + if fd != cmdTty { + fd.Close() + } + } + startPacket := cmd.MakeCmdStartPacket(pk.ReqId) + return cmd, startPacket, nil +} + +func GetExitCode(err error) int { + if err == nil { + return 0 + } + if exitErr, ok := err.(*exec.ExitError); ok { + return exitErr.ExitCode() + } else { + return -1 + } +} + +func (c *ShExecType) ProcWait() error { + exitErr := c.Cmd.Wait() + base.Logf("procwait: %v\n", exitErr) + c.Lock.Lock() + c.Exited = true + c.Lock.Unlock() + return exitErr +} + +func (c *ShExecType) IsExited() bool { + c.Lock.Lock() + defer c.Lock.Unlock() + return c.Exited +} + +func (c *ShExecType) WaitForCommand() *packet.CmdDonePacketType { + donePacket := packet.MakeCmdDonePacket(c.CK) + exitErr := c.ProcWait() + if c.ReturnState != nil { + <-c.ReturnState.DoneCh + state, _ := ParseShellStateOutput(c.ReturnState.Buf) // TODO what to do with error? + donePacket.FinalState = state + } + endTs := time.Now() + cmdDuration := endTs.Sub(c.StartTs) + donePacket.Ts = endTs.UnixMilli() + donePacket.ExitCode = GetExitCode(exitErr) + donePacket.DurationMs = int64(cmdDuration / time.Millisecond) + if c.FileNames != nil { + os.Remove(c.FileNames.StdinFifo) // best effort (no need to check error) + } + return donePacket +} + +func MakeInitPacket() *packet.InitPacketType { + initPacket := packet.MakeInitPacket() + initPacket.Version = base.MShellVersion + initPacket.BuildTime = base.BuildTime + initPacket.HomeDir = base.GetHomeDir() + initPacket.MShellHomeDir = base.GetMShellHomeDir() + if user, _ := user.Current(); user != nil { + initPacket.User = user.Username + } + initPacket.HostName, _ = os.Hostname() + initPacket.UName = fmt.Sprintf("%s|%s", runtime.GOOS, runtime.GOARCH) + return initPacket +} + +func MakeServerInitPacket() (*packet.InitPacketType, error) { + var err error + initPacket := MakeInitPacket() + shellState, err := GetShellState() + if err != nil { + return nil, err + } + initPacket.State = shellState + initPacket.Shell = os.Getenv(ShellVarName) + initPacket.RemoteId, err = base.GetRemoteId() + if err != nil { + return nil, err + } + return initPacket, nil +} + +func ParseEnv0(env []byte) map[string]string { + envLines := bytes.Split(env, []byte{0}) + rtn := make(map[string]string) + for _, envLine := range envLines { + if len(envLine) == 0 { + continue + } + eqIdx := bytes.Index(envLine, []byte{'='}) + if eqIdx == -1 { + continue + } + varName := string(envLine[0:eqIdx]) + varVal := string(envLine[eqIdx+1:]) + rtn[varName] = varVal + } + return rtn +} + +func MakeEnv0(envMap map[string]string) []byte { + var buf bytes.Buffer + for envName, envVal := range envMap { + buf.WriteString(envName) + buf.WriteByte('=') + buf.WriteString(envVal) + buf.WriteByte(0) + } + return buf.Bytes() +} + +func getStderr(err error) string { + exitErr, ok := err.(*exec.ExitError) + if !ok { + return "" + } + if len(exitErr.Stderr) == 0 { + return "" + } + lines := strings.SplitN(string(exitErr.Stderr), "\n", 2) + if len(lines[0]) > 100 { + return lines[0][0:100] + } + return lines[0] +} + +func runSimpleCmdInPty(ecmd *exec.Cmd) ([]byte, error) { + ecmd.Env = os.Environ() + UpdateCmdEnv(ecmd, MShellEnvVars(DefaultTermType)) + cmdPty, cmdTty, err := pty.Open() + if err != nil { + return nil, fmt.Errorf("opening new pty: %w", err) + } + pty.Setsize(cmdPty, &pty.Winsize{Rows: DefaultTermRows, Cols: DefaultTermCols}) + ecmd.Stdin = cmdTty + ecmd.Stdout = cmdTty + ecmd.Stderr = cmdTty + ecmd.SysProcAttr = &syscall.SysProcAttr{} + ecmd.SysProcAttr.Setsid = true + ecmd.SysProcAttr.Setctty = true + err = ecmd.Start() + if err != nil { + cmdTty.Close() + cmdPty.Close() + return nil, err + } + cmdTty.Close() + defer cmdPty.Close() + ioDone := make(chan bool) + var outputBuf bytes.Buffer + go func() { + // ignore error (/dev/ptmx has read error when process is done) + io.Copy(&outputBuf, cmdPty) + close(ioDone) + }() + exitErr := ecmd.Wait() + if exitErr != nil { + return nil, exitErr + } + <-ioDone + return outputBuf.Bytes(), nil +} + +func GetShellStateRedirectCommandStr(outputFdNum int) string { + return fmt.Sprintf("cat <(%s) > /dev/fd/%d", GetShellStateCmd(), outputFdNum) +} + +func GetShellState() (*packet.ShellState, error) { + ctx, _ := context.WithTimeout(context.Background(), GetStateTimeout) + cmdStr := BaseBashOpts + "; " + GetShellStateCmd() + ecmd := exec.CommandContext(ctx, "bash", "-l", "-i", "-c", cmdStr) + outputBytes, err := runSimpleCmdInPty(ecmd) + if err != nil { + return nil, err + } + return ParseShellStateOutput(outputBytes) +} + +func MShellEnvVars(termType string) map[string]string { + rtn := make(map[string]string) + if termType != "" { + rtn["TERM"] = termType + } + rtn["MSHELL"], _ = os.Executable() + rtn["MSHELL_VERSION"] = base.MShellVersion + return rtn +} diff --git a/waveshell/pkg/simpleexpand/simpleexpand.go b/waveshell/pkg/simpleexpand/simpleexpand.go new file mode 100644 index 00000000..d92d35c4 --- /dev/null +++ b/waveshell/pkg/simpleexpand/simpleexpand.go @@ -0,0 +1,222 @@ +package simpleexpand + +import ( + "bytes" + "strings" + + "mvdan.cc/sh/v3/expand" + "mvdan.cc/sh/v3/syntax" +) + +type SimpleExpandContext struct { + HomeDir string +} + +type SimpleExpandInfo struct { + HasTilde bool // only ~ as the first character when SimpleExpandContext.HomeDir is set + HasVar bool // $x, $$, ${...} + HasGlob bool // *, ?, [, { + HasExtGlob bool // ?(...) ... ?*+@! + HasHistory bool // ! (anywhere) + HasSpecial bool // subshell, arith +} + +func expandHomeDir(info *SimpleExpandInfo, litVal string, multiPart bool, homeDir string) string { + if homeDir == "" { + return litVal + } + if litVal == "~" && !multiPart { + return homeDir + } + if strings.HasPrefix(litVal, "~/") { + info.HasTilde = true + return homeDir + litVal[1:] + } + return litVal +} + +func expandLiteral(buf *bytes.Buffer, info *SimpleExpandInfo, litVal string) { + var lastBackSlash bool + var lastExtGlob bool + var lastDollar bool + for _, ch := range litVal { + if ch == 0 { + break + } + if lastBackSlash { + lastBackSlash = false + if ch == '\n' { + // special case, backslash *and* newline are ignored + continue + } + buf.WriteRune(ch) + continue + } + if ch == '\\' { + lastBackSlash = true + lastExtGlob = false + lastDollar = false + continue + } + if ch == '*' || ch == '?' || ch == '[' || ch == '{' { + info.HasGlob = true + } + if ch == '`' { + info.HasSpecial = true + } + if ch == '!' { + info.HasHistory = true + } + if lastExtGlob && ch == '(' { + info.HasExtGlob = true + } + if lastDollar && (ch != ' ' && ch != '"' && ch != '\'' && ch != '(' || ch != '[') { + info.HasVar = true + } + if lastDollar && (ch == '(' || ch == '[') { + info.HasSpecial = true + } + lastExtGlob = (ch == '?' || ch == '*' || ch == '+' || ch == '@' || ch == '!') + lastDollar = (ch == '$') + buf.WriteRune(ch) + } + if lastBackSlash { + buf.WriteByte('\\') + } +} + +// also expands ~ +func expandLiteralPlus(buf *bytes.Buffer, info *SimpleExpandInfo, litVal string, multiPart bool, ectx SimpleExpandContext) { + litVal = expandHomeDir(info, litVal, multiPart, ectx.HomeDir) + expandLiteral(buf, info, litVal) +} + +func expandSQANSILiteral(buf *bytes.Buffer, litVal string) { + // no info specials + if strings.HasSuffix(litVal, "'") { + litVal = litVal[0 : len(litVal)-1] + } + str, _, _ := expand.Format(nil, litVal, nil) + buf.WriteString(str) +} + +func expandSQLiteral(buf *bytes.Buffer, litVal string) { + // no info specials + if strings.HasSuffix(litVal, "'") { + litVal = litVal[0 : len(litVal)-1] + } + buf.WriteString(litVal) +} + +// will also work for partial double quoted strings +func expandDQLiteral(buf *bytes.Buffer, info *SimpleExpandInfo, litVal string) { + var lastBackSlash bool + var lastDollar bool + for _, ch := range litVal { + if ch == 0 { + break + } + if lastBackSlash { + lastBackSlash = false + if ch == '"' || ch == '\\' || ch == '$' || ch == '`' { + buf.WriteRune(ch) + continue + } + buf.WriteRune('\\') + buf.WriteRune(ch) + continue + } + if ch == '\\' { + lastBackSlash = true + lastDollar = false + continue + } + if ch == '"' { + break + } + + // similar to expandLiteral, but no globbing + if ch == '`' { + info.HasSpecial = true + } + if ch == '!' { + info.HasHistory = true + } + if lastDollar && (ch != ' ' && ch != '"' && ch != '\'' && ch != '(' || ch != '[') { + info.HasVar = true + } + if lastDollar && (ch == '(' || ch == '[') { + info.HasSpecial = true + } + lastDollar = (ch == '$') + buf.WriteRune(ch) + } + // in a valid parsed DQ string, you cannot have a trailing backslash (because \" would not end the string) + // still putting the case here though in case we ever deal with incomplete strings (e.g. completion) + if lastBackSlash { + buf.WriteByte('\\') + } +} + +func simpleExpandWordInternal(buf *bytes.Buffer, info *SimpleExpandInfo, ectx SimpleExpandContext, parts []syntax.WordPart, sourceStr string, inDoubleQuote bool, level int) { + for partIdx, untypedPart := range parts { + switch part := untypedPart.(type) { + case *syntax.Lit: + if !inDoubleQuote && partIdx == 0 && level == 1 && ectx.HomeDir != "" { + expandLiteralPlus(buf, info, part.Value, len(parts) > 1, ectx) + } else if inDoubleQuote { + expandDQLiteral(buf, info, part.Value) + } else { + expandLiteral(buf, info, part.Value) + } + + case *syntax.SglQuoted: + if part.Dollar { + expandSQANSILiteral(buf, part.Value) + } else { + expandSQLiteral(buf, part.Value) + } + + case *syntax.DblQuoted: + simpleExpandWordInternal(buf, info, ectx, part.Parts, sourceStr, true, level+1) + + default: + rawStr := sourceStr[part.Pos().Offset():part.End().Offset()] + buf.WriteString(rawStr) + } + } +} + +// simple word expansion +// expands: literals, single-quoted strings, double-quoted strings (recursively) +// does *not* expand: params (variables), command substitution, arithmetic expressions, process substituions, globs +// for the not expands, they will show up as the literal string +// this is different than expand.Literal which will replace variables as empty string if they aren't defined. +// so "a"'foo'${bar}$x => "afoo${bar}$x", but expand.Literal would produce => "afoo" +// note will do ~ expansion (will not do ~user expansion) +func SimpleExpandWord(ectx SimpleExpandContext, word *syntax.Word, sourceStr string) (string, SimpleExpandInfo) { + var buf bytes.Buffer + var info SimpleExpandInfo + simpleExpandWordInternal(&buf, &info, ectx, word.Parts, sourceStr, false, 1) + return buf.String(), info +} + +func SimpleExpandPartialWord(ectx SimpleExpandContext, partialWord string, multiPart bool) (string, SimpleExpandInfo) { + var buf bytes.Buffer + var info SimpleExpandInfo + if partialWord == "" { + return "", info + } + if strings.HasPrefix(partialWord, "\"") { + expandDQLiteral(&buf, &info, partialWord[1:]) + } else if strings.HasPrefix(partialWord, "$\"") { + expandDQLiteral(&buf, &info, partialWord[2:]) + } else if strings.HasPrefix(partialWord, "'") { + expandSQLiteral(&buf, partialWord[1:]) + } else if strings.HasPrefix(partialWord, "$'") { + expandSQANSILiteral(&buf, partialWord[2:]) + } else { + expandLiteralPlus(&buf, &info, partialWord, multiPart, ectx) + } + return buf.String(), info +} diff --git a/waveshell/pkg/statediff/linediff.go b/waveshell/pkg/statediff/linediff.go new file mode 100644 index 00000000..fa7ce11b --- /dev/null +++ b/waveshell/pkg/statediff/linediff.go @@ -0,0 +1,188 @@ +package statediff + +import ( + "bytes" + "encoding/binary" + "fmt" + "strings" +) + +const LineDiffVersion = 0 + +type SingleLineEntry struct { + LineVal int + Run int +} + +type LineDiffType struct { + Lines []SingleLineEntry + NewData []string +} + +func (diff LineDiffType) Dump() { + fmt.Printf("DIFF:\n") + pos := 1 + for _, entry := range diff.Lines { + fmt.Printf(" %d-%d: %d\n", pos, pos+entry.Run, entry.LineVal) + pos += entry.Run + } + for idx, str := range diff.NewData { + fmt.Printf(" n%d: %s\n", idx+1, str) + } +} + +// simple encoding +// a 0 means read a line from NewData +// a non-zero number means read the 1-indexed line from OldData +func (diff LineDiffType) applyDiff(oldData []string) ([]string, error) { + rtn := make([]string, 0, len(diff.Lines)) + newDataPos := 0 + for _, entry := range diff.Lines { + if entry.LineVal == 0 { + for i := 0; i < entry.Run; i++ { + if newDataPos >= len(diff.NewData) { + return nil, fmt.Errorf("not enough newdata for diff") + } + rtn = append(rtn, diff.NewData[newDataPos]) + newDataPos++ + } + } else { + oldDataPos := entry.LineVal - 1 // 1-indexed + for i := 0; i < entry.Run; i++ { + realPos := oldDataPos + i + if realPos < 0 || realPos >= len(oldData) { + return nil, fmt.Errorf("diff index out of bounds %d old-data-len:%d", realPos, len(oldData)) + } + rtn = append(rtn, oldData[realPos]) + } + } + } + return rtn, nil +} + +func putUVarint(buf *bytes.Buffer, viBuf []byte, ival int) { + l := binary.PutUvarint(viBuf, uint64(ival)) + buf.Write(viBuf[0:l]) +} + +// simple encoding +// write varints. first version, then len, then len-number-of-varints, then fill the rest with newdata +// [version] [len-varint] [varint]xlen... newdata (bytes) +func (diff LineDiffType) Encode() []byte { + var buf bytes.Buffer + viBuf := make([]byte, binary.MaxVarintLen64) + putUVarint(&buf, viBuf, LineDiffVersion) + putUVarint(&buf, viBuf, len(diff.Lines)) + for _, entry := range diff.Lines { + putUVarint(&buf, viBuf, entry.LineVal) + putUVarint(&buf, viBuf, entry.Run) + } + for idx, str := range diff.NewData { + buf.WriteString(str) + if idx != len(diff.NewData)-1 { + buf.WriteByte('\n') + } + } + return buf.Bytes() +} + +func (rtn *LineDiffType) Decode(diffBytes []byte) error { + r := bytes.NewBuffer(diffBytes) + version, err := binary.ReadUvarint(r) + if err != nil { + return fmt.Errorf("invalid diff, cannot read version: %v", err) + } + if version != LineDiffVersion { + return fmt.Errorf("invalid diff, bad version: %d", version) + } + linesLen64, err := binary.ReadUvarint(r) + if err != nil { + return fmt.Errorf("invalid diff, cannot read lines length: %v", err) + } + linesLen := int(linesLen64) + rtn.Lines = make([]SingleLineEntry, linesLen) + for idx := 0; idx < linesLen; idx++ { + lineVal, err := binary.ReadUvarint(r) + if err != nil { + return fmt.Errorf("invalid diff, cannot read line %d: %v", idx, err) + } + lineRun, err := binary.ReadUvarint(r) + if err != nil { + return fmt.Errorf("invalid diff, cannot read line-run %d: %v", idx, err) + } + rtn.Lines[idx] = SingleLineEntry{LineVal: int(lineVal), Run: int(lineRun)} + } + restOfInput := string(r.Bytes()) + if len(restOfInput) > 0 { + rtn.NewData = strings.Split(restOfInput, "\n") + } + return nil +} + +func makeLineDiff(oldData []string, newData []string) LineDiffType { + var rtn LineDiffType + oldDataMap := make(map[string]int) // 1-indexed + for idx, str := range oldData { + if _, found := oldDataMap[str]; found { + continue + } + oldDataMap[str] = idx + 1 + } + var cur *SingleLineEntry + rtn.Lines = make([]SingleLineEntry, 0) + for _, str := range newData { + oldIdx, found := oldDataMap[str] + if cur != nil && cur.LineVal != 0 { + checkLine := cur.LineVal + cur.Run - 1 + if checkLine < len(oldData) && oldData[checkLine] == str { + cur.Run++ + continue + } + } else if cur != nil && cur.LineVal == 0 && !found { + cur.Run++ + rtn.NewData = append(rtn.NewData, str) + continue + } + if cur != nil { + rtn.Lines = append(rtn.Lines, *cur) + } + cur = &SingleLineEntry{Run: 1} + if found { + cur.LineVal = oldIdx + } else { + cur.LineVal = 0 + rtn.NewData = append(rtn.NewData, str) + } + } + if cur != nil { + rtn.Lines = append(rtn.Lines, *cur) + } + return rtn +} + +func MakeLineDiff(str1 string, str2 string) []byte { + if str1 == str2 { + return nil + } + str1Arr := strings.Split(str1, "\n") + str2Arr := strings.Split(str2, "\n") + diff := makeLineDiff(str1Arr, str2Arr) + return diff.Encode() +} + +func ApplyLineDiff(str1 string, diffBytes []byte) (string, error) { + if len(diffBytes) == 0 { + return str1, nil + } + var diff LineDiffType + err := diff.Decode(diffBytes) + if err != nil { + return "", err + } + str1Arr := strings.Split(str1, "\n") + str2Arr, err := diff.applyDiff(str1Arr) + if err != nil { + return "", err + } + return strings.Join(str2Arr, "\n"), nil +} diff --git a/waveshell/pkg/statediff/mapdiff.go b/waveshell/pkg/statediff/mapdiff.go new file mode 100644 index 00000000..47db9f4b --- /dev/null +++ b/waveshell/pkg/statediff/mapdiff.go @@ -0,0 +1,130 @@ +package statediff + +import ( + "bytes" + "encoding/binary" + "fmt" +) + +const MapDiffVersion = 0 + +// 0-bytes are not allowed in entries or keys (same as bash) + +type MapDiffType struct { + ToAdd map[string]string + ToRemove []string +} + +func (diff MapDiffType) Dump() { + fmt.Printf("VAR-DIFF\n") + for name, val := range diff.ToAdd { + fmt.Printf(" add[%s] %s\n", name, val) + } + for _, name := range diff.ToRemove { + fmt.Printf(" rem[%s]\n", name) + } +} + +func makeMapDiff(oldMap map[string]string, newMap map[string]string) MapDiffType { + var rtn MapDiffType + rtn.ToAdd = make(map[string]string) + for name, newVal := range newMap { + oldVal, found := oldMap[name] + if !found || oldVal != newVal { + rtn.ToAdd[name] = newVal + continue + } + } + for name, _ := range oldMap { + _, found := newMap[name] + if !found { + rtn.ToRemove = append(rtn.ToRemove, name) + } + } + return rtn +} + +func (diff MapDiffType) apply(oldMap map[string]string) map[string]string { + rtn := make(map[string]string) + for name, val := range oldMap { + rtn[name] = val + } + for name, val := range diff.ToAdd { + rtn[name] = val + } + for _, name := range diff.ToRemove { + delete(rtn, name) + } + return rtn +} + +func (diff MapDiffType) Encode() []byte { + var buf bytes.Buffer + viBuf := make([]byte, binary.MaxVarintLen64) + putUVarint(&buf, viBuf, MapDiffVersion) + putUVarint(&buf, viBuf, len(diff.ToAdd)) + for key, val := range diff.ToAdd { + buf.WriteString(key) + buf.WriteByte(0) + buf.WriteString(val) + buf.WriteByte(0) + } + for _, val := range diff.ToRemove { + buf.WriteString(val) + buf.WriteByte(0) + } + return buf.Bytes() +} + +func (diff *MapDiffType) Decode(diffBytes []byte) error { + r := bytes.NewBuffer(diffBytes) + version, err := binary.ReadUvarint(r) + if err != nil { + return fmt.Errorf("invalid diff, cannot read version: %v", err) + } + if version != MapDiffVersion { + return fmt.Errorf("invalid diff, bad version: %d", version) + } + mapLen64, err := binary.ReadUvarint(r) + if err != nil { + return fmt.Errorf("invalid diff, cannot map length: %v", err) + } + mapLen := int(mapLen64) + fields := bytes.Split(r.Bytes(), []byte{0}) + if len(fields) < 2*mapLen { + return fmt.Errorf("invalid diff, not enough fields, maplen:%d fields:%d", mapLen, len(fields)) + } + mapFields := fields[0 : 2*mapLen] + removeFields := fields[2*mapLen:] + diff.ToAdd = make(map[string]string) + for i := 0; i < len(mapFields); i += 2 { + diff.ToAdd[string(mapFields[i])] = string(mapFields[i+1]) + } + for _, removeVal := range removeFields { + if len(removeVal) == 0 { + continue + } + diff.ToRemove = append(diff.ToRemove, string(removeVal)) + } + return nil +} + +func MakeMapDiff(m1 map[string]string, m2 map[string]string) []byte { + diff := makeMapDiff(m1, m2) + if len(diff.ToAdd) == 0 && len(diff.ToRemove) == 0 { + return nil + } + return diff.Encode() +} + +func ApplyMapDiff(oldMap map[string]string, diffBytes []byte) (map[string]string, error) { + if len(diffBytes) == 0 { + return oldMap, nil + } + var diff MapDiffType + err := diff.Decode(diffBytes) + if err != nil { + return nil, err + } + return diff.apply(oldMap), nil +} diff --git a/waveshell/pkg/statediff/statediff_test.go b/waveshell/pkg/statediff/statediff_test.go new file mode 100644 index 00000000..34dc478c --- /dev/null +++ b/waveshell/pkg/statediff/statediff_test.go @@ -0,0 +1,99 @@ +package statediff + +import ( + "fmt" + "testing" +) + +const Str1 = ` +hello +line #2 +apple +grapes +banana +apple +` + +const Str2 = ` +line #2 +apple +grapes +banana +` + +const Str3 = ` +more +stuff +banana +coconut +` + +const Str4 = ` +more +stuff +banana2 +coconut +` + +func testLineDiff(t *testing.T, str1 string, str2 string) { + diffBytes := MakeLineDiff(str1, str2) + fmt.Printf("diff-len: %d\n", len(diffBytes)) + out, err := ApplyLineDiff(str1, diffBytes) + if err != nil { + t.Errorf("error in diff: %v", err) + return + } + if out != str2 { + t.Errorf("bad diff output") + } + var dt LineDiffType + err = dt.Decode(diffBytes) + if err != nil { + t.Errorf("error decoding diff: %v\n", err) + } +} + +func TestLineDiff(t *testing.T) { + testLineDiff(t, Str1, Str2) + testLineDiff(t, Str2, Str3) + testLineDiff(t, Str1, Str3) + testLineDiff(t, Str3, Str1) + testLineDiff(t, Str3, Str4) +} + +func strMapsEqual(m1 map[string]string, m2 map[string]string) bool { + if len(m1) != len(m2) { + return false + } + for key, val := range m1 { + val2, ok := m2[key] + if !ok || val != val2 { + return false + } + } + for key, val := range m2 { + val2, ok := m1[key] + if !ok || val != val2 { + return false + } + } + return true +} + +func TestMapDiff(t *testing.T) { + m1 := map[string]string{"a": "5", "b": "hello", "c": "mike"} + m2 := map[string]string{"a": "5", "b": "goodbye", "d": "more"} + diffBytes := MakeMapDiff(m1, m2) + fmt.Printf("mapdifflen: %d\n", len(diffBytes)) + var diff MapDiffType + diff.Decode(diffBytes) + diff.Dump() + mcheck, err := ApplyMapDiff(m1, diffBytes) + if err != nil { + t.Fatalf("error applying map diff: %v", err) + } + if !strMapsEqual(m2, mcheck) { + t.Errorf("maps not equal") + } + fmt.Printf("%v\n", mcheck) +} diff --git a/waveshell/scripthaus.md b/waveshell/scripthaus.md new file mode 100644 index 00000000..d6a26698 --- /dev/null +++ b/waveshell/scripthaus.md @@ -0,0 +1,18 @@ + +```bash +# @scripthaus command build +GO_LDFLAGS="-s -w -X main.BuildTime=$(date +'%Y%m%d%H%M')" +go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-darwin.amd64 main-waveshell.go +``` + +```bash +# @scripthaus command fullbuild +GO_LDFLAGS="-s -w -X main.BuildTime=$(date +'%Y%m%d%H%M')" +go build -ldflags="$GO_LDFLAGS" -o ~/.mshell/mshell-v0.2 main-waveshell.go +GOOS=linux GOARCH=amd64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-linux.amd64 main-waveshell.go +GOOS=linux GOARCH=arm64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-linux.arm64 main-waveshell.go +GOOS=darwin GOARCH=amd64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-darwin.amd64 main-waveshell.go +GOOS=darwin GOARCH=arm64 go build -ldflags="$GO_LDFLAGS" -o bin/mshell-v0.3-darwin.arm64 main-waveshell.go +``` + +