From c29c4a9a2dda697f1fc8e99bef6bf261169e0dd1 Mon Sep 17 00:00:00 2001 From: sawka Date: Thu, 31 Aug 2023 22:03:38 -0700 Subject: [PATCH] PE-41 remote file api (#1) * new remote file streaming API packets. implemented 'stat' for remote files * introduce filedata packets. allow streaming RPCs. fix RPC bug with combined packet parsers. implement file streaming for filestream RPC. * checkpoint on adding write-file * completely untested write-file impl -- writefilecontext, condition var for signaling new data packets, cleanup goroutine, ready/done states. * better error messages, also unlock MServer before calling done on wfcs * fix bug with perm json tag. change constant name --- main-mshell.go | 2 +- pkg/packet/packet.go | 214 ++++++++++++++++++++++--- pkg/packet/parser.go | 89 ++++++++--- pkg/server/server.go | 363 +++++++++++++++++++++++++++++++++++++++++-- pkg/shexec/client.go | 6 +- pkg/shexec/shexec.go | 8 +- 6 files changed, 612 insertions(+), 70 deletions(-) diff --git a/main-mshell.go b/main-mshell.go index 98ed90f4..31a80f21 100644 --- a/main-mshell.go +++ b/main-mshell.go @@ -159,7 +159,7 @@ func readFullRunPacket(packetParser *packet.PacketParser) (*packet.RunPacketType } func handleSingle(fromServer bool) { - packetParser := packet.MakePacketParser(os.Stdin) + packetParser := packet.MakePacketParser(os.Stdin, false) sender := packet.MakePacketSender(os.Stdout, nil) defer func() { sender.Close() diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index 8241e69e..c35ec856 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -26,33 +26,42 @@ import ( // server : run, >cmddata, >cmddone, data, <>dataack, cd, >getcmd, >untailcmd, >input, error, <>message, <>ping, streamfile, writefile, filedata*, state - CurrentState string // sha1 - WriteErrorCh chan bool // closed if there is a I/O write error - WriteErrorChOnce *sync.Once + 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) { @@ -164,6 +254,224 @@ func (m *MServer) reinit(reqId string) { m.Sender.SendPacket(initPk) } +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 + } + finfo, err := os.Stat(pk.Path) + if err == nil && finfo.IsDir() { + err = fmt.Errorf("invalid path, cannot write a directory") + } + if err == nil { + writePerm := (finfo.Mode().Perm() & 0o222) + if writePerm == 0 { + err = fmt.Errorf("file is not writable, perms: %v", finfo.Mode().Perm()) + } + } + if err != nil { + resp := packet.MakeWriteFileReadyPacket(pk.ReqId) + resp.Error = err.Error() + m.Sender.SendPacket(resp) + return + } + + var writeFd *os.File + if pk.UseTemp { + dirName := filepath.Dir(pk.Path) + dirFInfo, err := os.Stat(dirName) + if err == nil { + writePerm := (dirFInfo.Mode().Perm() & 0o222) + if writePerm == 0 { + err = fmt.Errorf("file-write tempmode is set, but parent directory is not writeable, perms: %v", dirFInfo.Mode().Perm()) + } + } + if err != nil { + resp := packet.MakeWriteFileReadyPacket(pk.ReqId) + resp.Error = err.Error() + m.Sender.SendPacket(resp) + return + } + baseName := filepath.Base(pk.Path) + writeFd, err = os.CreateTemp(dirName, baseName+".tmp.") + if err != nil { + resp := packet.MakeWriteFileReadyPacket(pk.ReqId) + resp.Error = fmt.Sprintf("write-file could not open tempfile: %v", err) + m.Sender.SendPacket(resp) + return + } + } else { + writeFd, err = os.OpenFile(pk.Path, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o777) // use 777 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 { + renameErr := os.Rename(writeFd.Name(), pk.Path) + if renameErr != nil { + doneErr = fmt.Errorf("error renaming temp file: %v", renameErr) + // rename failed, try to remove temp file still + os.Remove(writeFd.Name()) + } + } + } + donePk := packet.MakeWriteFileDonePacket(pk.ReqId) + if doneErr != nil { + donePk.Error = doneErr.Error() + } + m.Sender.SendPacket(donePk) +} + +func (m *MServer) streamFile(pk *packet.StreamFilePacketType) { + resp := packet.MakeStreamFileResponse(pk.ReqId) + finfo, err := os.Stat(pk.Path) + 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 { @@ -183,6 +491,15 @@ func (m *MServer) ProcessRpcPacket(pk packet.RpcPacketType) { 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 } @@ -288,6 +605,10 @@ func (server *MServer) runReadLoop() { 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 } @@ -299,17 +620,27 @@ func RunServer() (int, error) { 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{}, + 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) + server.MainInput = packet.MakePacketParser(os.Stdin, false) server.Sender = packet.MakePacketSender(os.Stdout, server.packetSenderErrorHandler) defer server.Close() var err error diff --git a/pkg/shexec/client.go b/pkg/shexec/client.go index efb3414a..c6292b82 100644 --- a/pkg/shexec/client.go +++ b/pkg/shexec/client.go @@ -47,9 +47,9 @@ func MakeClientProc(ctx context.Context, ecmd *exec.Cmd) (*ClientProc, *packet.I return nil, nil, fmt.Errorf("running local client: %w", err) } sender := packet.MakePacketSender(inputWriter, nil) - stdoutPacketParser := packet.MakePacketParser(stdoutReader) - stderrPacketParser := packet.MakePacketParser(stderrReader) - packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser) + stdoutPacketParser := packet.MakePacketParser(stdoutReader, false) + stderrPacketParser := packet.MakePacketParser(stderrReader, false) + packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser, true) cproc := &ClientProc{ Cmd: ecmd, StartTs: startTs, diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 7b38e0a7..63337fa8 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -727,7 +727,7 @@ func RunInstallFromCmd(ctx context.Context, ecmd *exec.Cmd, tryDetect bool, mshe if mshellStream != nil { sendMShellBinary(inputWriter, mshellStream) } - packetParser := packet.MakePacketParser(stdoutReader) + packetParser := packet.MakePacketParser(stdoutReader, false) err = ecmd.Start() if err != nil { return fmt.Errorf("running ssh command: %w", err) @@ -860,9 +860,9 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdCon return nil, fmt.Errorf("running ssh command: %w", err) } defer cmd.Close() - stdoutPacketParser := packet.MakePacketParser(stdoutReader) - stderrPacketParser := packet.MakePacketParser(stderrReader) - packetParser := packet.CombinePacketParsers(stdoutPacketParser, stderrPacketParser) + 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 {