Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
69 changes: 25 additions & 44 deletions lib/system/guest_agent/exec.go
Original file line number Diff line number Diff line change
Expand Up @@ -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() {
Expand All @@ -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())
Expand Down
64 changes: 64 additions & 0 deletions lib/system/guest_agent/exec_stream_test.go
Original file line number Diff line number Diff line change
@@ -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())
}