diff --git a/lib/system/guest_agent/exec.go b/lib/system/guest_agent/exec.go index 409754c65..1b8cbc8a6 100644 --- a/lib/system/guest_agent/exec.go +++ b/lib/system/guest_agent/exec.go @@ -80,9 +80,25 @@ func (s *guestServer) executeNoTTY(ctx context.Context, stream pb.GuestService_E // Mutex to protect concurrent stream.Send calls (gRPC streams are not thread-safe) var sendMu sync.Mutex - // Use WaitGroup to ensure all output is read before sending + // Stream output as it is produced so long-lived commands (relays, tails) + // deliver bytes before exit; drain both pipes before Wait closes them. var wg sync.WaitGroup - var stdoutData, stderrData []byte + pump := func(r io.Reader, wrap func([]byte) *pb.ExecResponse) { + defer wg.Done() + buf := make([]byte, 32*1024) + for { + n, err := r.Read(buf) + if n > 0 { + chunk := append([]byte(nil), buf[:n]...) + sendMu.Lock() + _ = stream.Send(wrap(chunk)) + sendMu.Unlock() + } + if err != nil { + return + } + } + } // Handle stdin in background go func() { @@ -98,52 +114,17 @@ func (s *guestServer) executeNoTTY(ctx context.Context, stream pb.GuestService_E } }() - // Read all stdout/stderr BEFORE calling Wait() - Wait() closes the pipes! - wg.Add(1) - go func() { - defer wg.Done() - data, _ := io.ReadAll(stdout) - stdoutData = data - }() - - wg.Add(1) - go func() { - defer wg.Done() - data, _ := io.ReadAll(stderr) - stderrData = data - }() - - // Wait for all reads to complete FIRST (before Wait closes pipes) + wg.Add(2) + go pump(stdout, func(b []byte) *pb.ExecResponse { + return &pb.ExecResponse{Response: &pb.ExecResponse_Stdout{Stdout: b}} + }) + go pump(stderr, func(b []byte) *pb.ExecResponse { + return &pb.ExecResponse{Response: &pb.ExecResponse_Stderr{Stderr: b}} + }) wg.Wait() - // Now safe to call Wait - pipes are fully drained waitErr := cmd.Wait() - // Now stream output in chunks (streaming compatible) - const chunkSize = 32 * 1024 - for i := 0; i < len(stdoutData); i += chunkSize { - end := i + chunkSize - if end > len(stdoutData) { - end = len(stdoutData) - } - sendMu.Lock() - stream.Send(&pb.ExecResponse{ - Response: &pb.ExecResponse_Stdout{Stdout: stdoutData[i:end]}, - }) - sendMu.Unlock() - } - for i := 0; i < len(stderrData); i += chunkSize { - end := i + chunkSize - if end > len(stderrData) { - end = len(stderrData) - } - sendMu.Lock() - stream.Send(&pb.ExecResponse{ - Response: &pb.ExecResponse_Stderr{Stderr: stderrData[i:end]}, - }) - sendMu.Unlock() - } - exitCode := int32(0) if cmd.ProcessState != nil { exitCode = int32(cmd.ProcessState.ExitCode()) diff --git a/lib/system/guest_agent/exec_stream_test.go b/lib/system/guest_agent/exec_stream_test.go new file mode 100644 index 000000000..c386131de --- /dev/null +++ b/lib/system/guest_agent/exec_stream_test.go @@ -0,0 +1,64 @@ +package main + +import ( + "context" + "io" + "strings" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" + "google.golang.org/grpc" + + pb "github.com/kernel/hypeman/lib/guest" +) + +// recordingExecStream is a fake Exec server stream that timestamps stdout chunks. +type recordingExecStream struct { + grpc.ServerStream + ctx context.Context + mu sync.Mutex + stdout []timedChunk +} + +type timedChunk struct { + at time.Time + data string +} + +func (s *recordingExecStream) Context() context.Context { return s.ctx } +func (s *recordingExecStream) Recv() (*pb.ExecRequest, error) { + return nil, io.EOF +} + +func (s *recordingExecStream) Send(resp *pb.ExecResponse) error { + if out := resp.GetStdout(); out != nil { + s.mu.Lock() + s.stdout = append(s.stdout, timedChunk{at: time.Now(), data: string(out)}) + s.mu.Unlock() + } + return nil +} + +func TestExecuteNoTTYStreamsOutputBeforeExit(t *testing.T) { + stream := &recordingExecStream{ctx: t.Context()} + start := time.Now() + err := (&guestServer{}).executeNoTTY(t.Context(), stream, &pb.ExecStart{ + Command: []string{"sh", "-c", "echo first; sleep 2; echo second"}, + }) + require.NoError(t, err) + exitedAt := time.Since(start) + + require.NotEmpty(t, stream.stdout) + require.True(t, strings.HasPrefix(stream.stdout[0].data, "first")) + firstAt := stream.stdout[0].at.Sub(start) + require.Less(t, firstAt, exitedAt-time.Second, + "first chunk arrived at %s but command exited at %s; output is being buffered", firstAt, exitedAt) + + var all strings.Builder + for _, c := range stream.stdout { + all.WriteString(c.data) + } + require.Equal(t, "first\nsecond\n", all.String()) +}