diff --git a/pkg/mpio/bufwriter.go b/pkg/mpio/bufwriter.go index 977a2493..f2f3cbbb 100644 --- a/pkg/mpio/bufwriter.go +++ b/pkg/mpio/bufwriter.go @@ -17,19 +17,23 @@ type FdWriter struct { 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) *FdWriter { +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 } @@ -68,11 +72,11 @@ func (w *FdWriter) AddData(data []byte, eof bool) error { if len(data) == 0 { return nil } - return fmt.Errorf("write to closed file eof[%v]", w.Eof) + 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) > WriteBufSize { - return fmt.Errorf("write exceeds buffer size bufsize=%d (max=%d)", len(data)+len(w.Buffer), WriteBufSize) + 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...) } diff --git a/pkg/mpio/mpio.go b/pkg/mpio/mpio.go index f16a411c..e27d9dcd 100644 --- a/pkg/mpio/mpio.go +++ b/pkg/mpio/mpio.go @@ -95,27 +95,28 @@ func (m *Multiplexer) MakeReaderPipe(fdNum int) (*os.File, error) { } // returns the *reader* to connect to process, writer is put in FdWriters -func (m *Multiplexer) MakeWriterPipe(fdNum int) (*os.File, error) { +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) + 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) (*os.File, error) { +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) + fdWriter := MakeFdWriter(m, pw, fdNum, true, desc) + fdWriter.BufferLimit = bufferLimit err = fdWriter.AddData(data, true) if err != nil { return nil, err @@ -131,10 +132,10 @@ func (m *Multiplexer) MakeRawFdReader(fdNum int, fd io.ReadCloser, shouldClose b m.FdReaders[fdNum] = MakeFdReader(m, fd, fdNum, shouldClose, isPty) } -func (m *Multiplexer) MakeRawFdWriter(fdNum int, fd io.WriteCloser, shouldClose bool) { +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) + m.FdWriters[fdNum] = MakeFdWriter(m, fd, fdNum, shouldClose, desc) } func (m *Multiplexer) makeDataAckPacket(fdNum int, ackLen int, err error) *packet.DataAckPacketType { @@ -225,22 +226,18 @@ func (m *Multiplexer) runPacketInputLoop() *packet.CmdDonePacketType { 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) - } +func (m *Multiplexer) WriteDataToFd(fdNum int, data []byte, isEof bool) error { m.Lock.Lock() defer m.Lock.Unlock() - fw := m.FdWriters[dataPacket.FdNum] + fw := m.FdWriters[fdNum] if fw == nil { // add a closed FdWriter as a placeholder so we only send one error - fw := MakeFdWriter(m, nil, dataPacket.FdNum, false) + fw := MakeFdWriter(m, nil, fdNum, false, "invalid-fd") fw.Close() - m.FdWriters[dataPacket.FdNum] = fw + m.FdWriters[fdNum] = fw return fmt.Errorf("write to closed file (no fd)") } - err = fw.AddData(realData, dataPacket.Eof) + err := fw.AddData(data, isEof) if err != nil { fw.Close() return err @@ -248,6 +245,14 @@ func (m *Multiplexer) processDataPacket(dataPacket *packet.DataPacketType) error 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() diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 8019e4e7..3ac000f3 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -840,8 +840,8 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdCon if !HasDupStdin(runPacket.Fds) { cmd.Multiplexer.MakeRawFdReader(0, fdContext.GetReader(0), false, false) } - cmd.Multiplexer.MakeRawFdWriter(1, fdContext.GetWriter(1), false) - cmd.Multiplexer.MakeRawFdWriter(2, fdContext.GetWriter(2), 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) @@ -852,7 +852,7 @@ func RunClientSSHCommandAndWait(runPacket *packet.RunPacketType, fdContext FdCon cmd.Multiplexer.MakeRawFdReader(rfd.FdNum, fd, false, false) } else if rfd.Write { fd := fdContext.GetWriter(rfd.FdNum) - cmd.Multiplexer.MakeRawFdWriter(rfd.FdNum, fd, true) + cmd.Multiplexer.MakeRawFdWriter(rfd.FdNum, fd, true, "client") } } err = ecmd.Start() @@ -1123,7 +1123,7 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro Setsid: true, Setctty: true, } - cmd.Multiplexer.MakeRawFdWriter(0, cmdPty, false) + cmd.Multiplexer.MakeRawFdWriter(0, cmdPty, false, "simple") cmd.Multiplexer.MakeRawFdReader(1, cmdPty, false, true) nullFd, err := os.Open("/dev/null") if err != nil { @@ -1131,7 +1131,7 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro } cmd.Multiplexer.MakeRawFdReader(2, nullFd, true, false) } else { - cmd.Cmd.Stdin, err = cmd.Multiplexer.MakeWriterPipe(0) + cmd.Cmd.Stdin, err = cmd.Multiplexer.MakeWriterPipe(0, "simple") if err != nil { return nil, err } @@ -1149,7 +1149,7 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro if runData.FdNum >= len(extraFiles) { extraFiles = extraFiles[:runData.FdNum+1] } - extraFiles[runData.FdNum], err = cmd.Multiplexer.MakeStaticWriterPipe(runData.FdNum, runData.Data) + extraFiles[runData.FdNum], err = cmd.Multiplexer.MakeStaticWriterPipe(runData.FdNum, runData.Data, MaxRunDataSize, "simple-rundata") if err != nil { return nil, err } @@ -1160,7 +1160,7 @@ func RunCommandSimple(pk *packet.RunPacketType, sender *packet.PacketSender, fro } if rfd.Read { // client file is open for reading, so we make a writer pipe - extraFiles[rfd.FdNum], err = cmd.Multiplexer.MakeWriterPipe(rfd.FdNum) + extraFiles[rfd.FdNum], err = cmd.Multiplexer.MakeWriterPipe(rfd.FdNum, "simple") if err != nil { return nil, err }