diff --git a/pkg/packet/packet.go b/pkg/packet/packet.go index a7674518..5a374a25 100644 --- a/pkg/packet/packet.go +++ b/pkg/packet/packet.go @@ -566,15 +566,36 @@ func (sender *PacketSender) SendMessage(fmtStr string, args ...interface{}) erro return sender.SendPacket(MakeMessagePacket(fmt.Sprintf(fmtStr, args...))) } -func PacketParser(input io.Reader) chan PacketType { +func CombinePacketParsers(p1 chan PacketType, p2 chan PacketType) chan PacketType { rtnCh := make(chan PacketType) - PacketParserAttach(input, rtnCh) + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + for v := range p1 { + rtnCh <- v + } + }() + go func() { + defer wg.Done() + for v := range p2 { + rtnCh <- v + } + }() + go func() { + wg.Wait() + close(rtnCh) + }() return rtnCh } -func PacketParserAttach(input io.Reader, rtnCh chan PacketType) { +func PacketParser(input io.Reader) chan PacketType { + rtnCh := make(chan PacketType) bufReader := bufio.NewReader(input) go func() { + defer func() { + close(rtnCh) + }() for { line, err := bufReader.ReadString('\n') if err == io.EOF { @@ -612,6 +633,7 @@ func PacketParserAttach(input io.Reader, rtnCh chan PacketType) { rtnCh <- pk } }() + return rtnCh } type ErrorReporter interface { diff --git a/pkg/shexec/shexec.go b/pkg/shexec/shexec.go index 82aa31b4..9bf98a56 100644 --- a/pkg/shexec/shexec.go +++ b/pkg/shexec/shexec.go @@ -298,8 +298,9 @@ func RunClientSSHCommandAndWait(opts *ClientOpts) (*packet.CmdDonePacketType, er return nil, fmt.Errorf("running ssh command: %w", err) } defer cmd.Close() - packetCh := packet.PacketParser(stdoutReader) - packet.PacketParserAttach(stderrReader, packetCh) + stdoutPacketCh := packet.PacketParser(stdoutReader) + stderrPacketCh := packet.PacketParser(stderrReader) + packetCh := packet.CombinePacketParsers(stdoutPacketCh, stderrPacketCh) sender := packet.MakePacketSender(inputWriter) for pk := range packetCh { if pk.GetType() == packet.RawPacketStr {