diff --git a/.golangci.yml b/.golangci.yml index 385b1c1..f8b8344 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -61,6 +61,15 @@ linters: # Not all errors need to be checked - errcheck + # It's OK to not initialize some struct fields + - exhaustruct + + # This is not a library, so it's OK to use dynamic errors + - err113 + + # Inline error handling keeps assignment and checking together + - noinlineerr + issues: # Don't hide multiple issues that belong to one class since GitHub annotations can handle them all nicely. max-issues-per-linter: 0 diff --git a/internal/rpc/agent.pb.go b/internal/rpc/agent.pb.go index c1b27fa..4842296 100644 --- a/internal/rpc/agent.pb.go +++ b/internal/rpc/agent.pb.go @@ -22,6 +22,55 @@ const ( _ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20) ) +type ExecRequest_Signal int32 + +const ( + ExecRequest_SIGNAL_UNSPECIFIED ExecRequest_Signal = 0 + ExecRequest_SIGNAL_SIGTERM ExecRequest_Signal = 1 + ExecRequest_SIGNAL_SIGKILL ExecRequest_Signal = 2 +) + +// Enum value maps for ExecRequest_Signal. +var ( + ExecRequest_Signal_name = map[int32]string{ + 0: "SIGNAL_UNSPECIFIED", + 1: "SIGNAL_SIGTERM", + 2: "SIGNAL_SIGKILL", + } + ExecRequest_Signal_value = map[string]int32{ + "SIGNAL_UNSPECIFIED": 0, + "SIGNAL_SIGTERM": 1, + "SIGNAL_SIGKILL": 2, + } +) + +func (x ExecRequest_Signal) Enum() *ExecRequest_Signal { + p := new(ExecRequest_Signal) + *p = x + return p +} + +func (x ExecRequest_Signal) String() string { + return protoimpl.X.EnumStringOf(x.Descriptor(), protoreflect.EnumNumber(x)) +} + +func (ExecRequest_Signal) Descriptor() protoreflect.EnumDescriptor { + return file_rpc_agent_proto_enumTypes[0].Descriptor() +} + +func (ExecRequest_Signal) Type() protoreflect.EnumType { + return &file_rpc_agent_proto_enumTypes[0] +} + +func (x ExecRequest_Signal) Number() protoreflect.EnumNumber { + return protoreflect.EnumNumber(x) +} + +// Deprecated: Use ExecRequest_Signal.Descriptor instead. +func (ExecRequest_Signal) EnumDescriptor() ([]byte, []int) { + return file_rpc_agent_proto_rawDescGZIP(), []int{0, 0} +} + type ExecRequest struct { state protoimpl.MessageState `protogen:"open.v1"` // Types that are valid to be assigned to Type: @@ -29,6 +78,7 @@ type ExecRequest struct { // *ExecRequest_Command_ // *ExecRequest_StandardInput // *ExecRequest_TerminalResize + // *ExecRequest_Signal_ Type isExecRequest_Type `protobuf_oneof:"type"` unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache @@ -98,6 +148,15 @@ func (x *ExecRequest) GetTerminalResize() *TerminalSize { return nil } +func (x *ExecRequest) GetSignal() ExecRequest_Signal { + if x != nil { + if x, ok := x.Type.(*ExecRequest_Signal_); ok { + return x.Signal + } + } + return ExecRequest_SIGNAL_UNSPECIFIED +} + type isExecRequest_Type interface { isExecRequest_Type() } @@ -114,12 +173,18 @@ type ExecRequest_TerminalResize struct { TerminalResize *TerminalSize `protobuf:"bytes,3,opt,name=terminal_resize,json=terminalResize,proto3,oneof"` } +type ExecRequest_Signal_ struct { + Signal ExecRequest_Signal `protobuf:"varint,4,opt,name=signal,proto3,enum=ExecRequest_Signal,oneof"` +} + func (*ExecRequest_Command_) isExecRequest_Type() {} func (*ExecRequest_StandardInput) isExecRequest_Type() {} func (*ExecRequest_TerminalResize) isExecRequest_Type() {} +func (*ExecRequest_Signal_) isExecRequest_Type() {} + type ExecResponse struct { state protoimpl.MessageState `protogen:"open.v1"` // Types that are valid to be assigned to Type: @@ -127,6 +192,7 @@ type ExecResponse struct { // *ExecResponse_Exit_ // *ExecResponse_StandardOutput // *ExecResponse_StandardError + // *ExecResponse_Started_ Type isExecResponse_Type `protobuf_oneof:"type"` unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache @@ -196,6 +262,15 @@ func (x *ExecResponse) GetStandardError() *IOChunk { return nil } +func (x *ExecResponse) GetStarted() *ExecResponse_Started { + if x != nil { + if x, ok := x.Type.(*ExecResponse_Started_); ok { + return x.Started + } + } + return nil +} + type isExecResponse_Type interface { isExecResponse_Type() } @@ -212,12 +287,18 @@ type ExecResponse_StandardError struct { StandardError *IOChunk `protobuf:"bytes,3,opt,name=standard_error,json=standardError,proto3,oneof"` } +type ExecResponse_Started_ struct { + Started *ExecResponse_Started `protobuf:"bytes,4,opt,name=started,proto3,oneof"` +} + func (*ExecResponse_Exit_) isExecResponse_Type() {} func (*ExecResponse_StandardOutput) isExecResponse_Type() {} func (*ExecResponse_StandardError) isExecResponse_Type() {} +func (*ExecResponse_Started_) isExecResponse_Type() {} + type TerminalSize struct { state protoimpl.MessageState `protogen:"open.v1"` Rows uint32 `protobuf:"varint,1,opt,name=rows,proto3" json:"rows,omitempty"` @@ -538,15 +619,52 @@ func (x *ExecResponse_Exit) GetCode() int32 { return 0 } +type ExecResponse_Started struct { + state protoimpl.MessageState `protogen:"open.v1"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *ExecResponse_Started) Reset() { + *x = ExecResponse_Started{} + mi := &file_rpc_agent_proto_msgTypes[9] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *ExecResponse_Started) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*ExecResponse_Started) ProtoMessage() {} + +func (x *ExecResponse_Started) ProtoReflect() protoreflect.Message { + mi := &file_rpc_agent_proto_msgTypes[9] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use ExecResponse_Started.ProtoReflect.Descriptor instead. +func (*ExecResponse_Started) Descriptor() ([]byte, []int) { + return file_rpc_agent_proto_rawDescGZIP(), []int{1, 1} +} + var File_rpc_agent_proto protoreflect.FileDescriptor const file_rpc_agent_proto_rawDesc = "" + "\n" + - "\x0frpc/agent.proto\x1a\x1bgoogle/protobuf/empty.proto\"\xeb\x03\n" + + "\x0frpc/agent.proto\x1a\x1bgoogle/protobuf/empty.proto\"\xe4\x04\n" + "\vExecRequest\x120\n" + "\acommand\x18\x01 \x01(\v2\x14.ExecRequest.CommandH\x00R\acommand\x121\n" + "\x0estandard_input\x18\x02 \x01(\v2\b.IOChunkH\x00R\rstandardInput\x128\n" + - "\x0fterminal_resize\x18\x03 \x01(\v2\r.TerminalSizeH\x00R\x0eterminalResize\x1a\xb4\x02\n" + + "\x0fterminal_resize\x18\x03 \x01(\v2\r.TerminalSizeH\x00R\x0eterminalResize\x12-\n" + + "\x06signal\x18\x04 \x01(\x0e2\x13.ExecRequest.SignalH\x00R\x06signal\x1a\xb4\x02\n" + "\aCommand\x12\x12\n" + "\x04name\x18\x01 \x01(\tR\x04name\x12\x12\n" + "\x04args\x18\x02 \x03(\tR\x04args\x12 \n" + @@ -558,14 +676,20 @@ const file_rpc_agent_proto_rawDesc = "" + "\aworkdir\x18\b \x01(\tR\aworkdir\x1a6\n" + "\bEnvEntry\x12\x10\n" + "\x03key\x18\x01 \x01(\tR\x03key\x12\x14\n" + - "\x05value\x18\x02 \x01(\tR\x05value:\x028\x01B\x06\n" + - "\x04type\"\xc4\x01\n" + + "\x05value\x18\x02 \x01(\tR\x05value:\x028\x01\"H\n" + + "\x06Signal\x12\x16\n" + + "\x12SIGNAL_UNSPECIFIED\x10\x00\x12\x12\n" + + "\x0eSIGNAL_SIGTERM\x10\x01\x12\x12\n" + + "\x0eSIGNAL_SIGKILL\x10\x02B\x06\n" + + "\x04type\"\x82\x02\n" + "\fExecResponse\x12(\n" + "\x04exit\x18\x01 \x01(\v2\x12.ExecResponse.ExitH\x00R\x04exit\x123\n" + "\x0fstandard_output\x18\x02 \x01(\v2\b.IOChunkH\x00R\x0estandardOutput\x121\n" + - "\x0estandard_error\x18\x03 \x01(\v2\b.IOChunkH\x00R\rstandardError\x1a\x1a\n" + + "\x0estandard_error\x18\x03 \x01(\v2\b.IOChunkH\x00R\rstandardError\x121\n" + + "\astarted\x18\x04 \x01(\v2\x15.ExecResponse.StartedH\x00R\astarted\x1a\x1a\n" + "\x04Exit\x12\x12\n" + - "\x04code\x18\x01 \x01(\x05R\x04codeB\x06\n" + + "\x04code\x18\x01 \x01(\x05R\x04code\x1a\t\n" + + "\aStartedB\x06\n" + "\x04type\"6\n" + "\fTerminalSize\x12\x12\n" + "\x04rows\x18\x01 \x01(\rR\x04rows\x12\x12\n" + @@ -591,36 +715,41 @@ func file_rpc_agent_proto_rawDescGZIP() []byte { return file_rpc_agent_proto_rawDescData } -var file_rpc_agent_proto_msgTypes = make([]protoimpl.MessageInfo, 9) +var file_rpc_agent_proto_enumTypes = make([]protoimpl.EnumInfo, 1) +var file_rpc_agent_proto_msgTypes = make([]protoimpl.MessageInfo, 10) var file_rpc_agent_proto_goTypes = []any{ - (*ExecRequest)(nil), // 0: ExecRequest - (*ExecResponse)(nil), // 1: ExecResponse - (*TerminalSize)(nil), // 2: TerminalSize - (*IOChunk)(nil), // 3: IOChunk - (*ResolveIPRequest)(nil), // 4: ResolveIPRequest - (*ResolveIPResponse)(nil), // 5: ResolveIPResponse - (*ExecRequest_Command)(nil), // 6: ExecRequest.Command - nil, // 7: ExecRequest.Command.EnvEntry - (*ExecResponse_Exit)(nil), // 8: ExecResponse.Exit + (ExecRequest_Signal)(0), // 0: ExecRequest.Signal + (*ExecRequest)(nil), // 1: ExecRequest + (*ExecResponse)(nil), // 2: ExecResponse + (*TerminalSize)(nil), // 3: TerminalSize + (*IOChunk)(nil), // 4: IOChunk + (*ResolveIPRequest)(nil), // 5: ResolveIPRequest + (*ResolveIPResponse)(nil), // 6: ResolveIPResponse + (*ExecRequest_Command)(nil), // 7: ExecRequest.Command + nil, // 8: ExecRequest.Command.EnvEntry + (*ExecResponse_Exit)(nil), // 9: ExecResponse.Exit + (*ExecResponse_Started)(nil), // 10: ExecResponse.Started } var file_rpc_agent_proto_depIdxs = []int32{ - 6, // 0: ExecRequest.command:type_name -> ExecRequest.Command - 3, // 1: ExecRequest.standard_input:type_name -> IOChunk - 2, // 2: ExecRequest.terminal_resize:type_name -> TerminalSize - 8, // 3: ExecResponse.exit:type_name -> ExecResponse.Exit - 3, // 4: ExecResponse.standard_output:type_name -> IOChunk - 3, // 5: ExecResponse.standard_error:type_name -> IOChunk - 2, // 6: ExecRequest.Command.terminal_size:type_name -> TerminalSize - 7, // 7: ExecRequest.Command.env:type_name -> ExecRequest.Command.EnvEntry - 0, // 8: Agent.Exec:input_type -> ExecRequest - 4, // 9: Agent.ResolveIP:input_type -> ResolveIPRequest - 1, // 10: Agent.Exec:output_type -> ExecResponse - 5, // 11: Agent.ResolveIP:output_type -> ResolveIPResponse - 10, // [10:12] is the sub-list for method output_type - 8, // [8:10] is the sub-list for method input_type - 8, // [8:8] is the sub-list for extension type_name - 8, // [8:8] is the sub-list for extension extendee - 0, // [0:8] is the sub-list for field type_name + 7, // 0: ExecRequest.command:type_name -> ExecRequest.Command + 4, // 1: ExecRequest.standard_input:type_name -> IOChunk + 3, // 2: ExecRequest.terminal_resize:type_name -> TerminalSize + 0, // 3: ExecRequest.signal:type_name -> ExecRequest.Signal + 9, // 4: ExecResponse.exit:type_name -> ExecResponse.Exit + 4, // 5: ExecResponse.standard_output:type_name -> IOChunk + 4, // 6: ExecResponse.standard_error:type_name -> IOChunk + 10, // 7: ExecResponse.started:type_name -> ExecResponse.Started + 3, // 8: ExecRequest.Command.terminal_size:type_name -> TerminalSize + 8, // 9: ExecRequest.Command.env:type_name -> ExecRequest.Command.EnvEntry + 1, // 10: Agent.Exec:input_type -> ExecRequest + 5, // 11: Agent.ResolveIP:input_type -> ResolveIPRequest + 2, // 12: Agent.Exec:output_type -> ExecResponse + 6, // 13: Agent.ResolveIP:output_type -> ResolveIPResponse + 12, // [12:14] is the sub-list for method output_type + 10, // [10:12] is the sub-list for method input_type + 10, // [10:10] is the sub-list for extension type_name + 10, // [10:10] is the sub-list for extension extendee + 0, // [0:10] is the sub-list for field type_name } func init() { file_rpc_agent_proto_init() } @@ -632,24 +761,27 @@ func file_rpc_agent_proto_init() { (*ExecRequest_Command_)(nil), (*ExecRequest_StandardInput)(nil), (*ExecRequest_TerminalResize)(nil), + (*ExecRequest_Signal_)(nil), } file_rpc_agent_proto_msgTypes[1].OneofWrappers = []any{ (*ExecResponse_Exit_)(nil), (*ExecResponse_StandardOutput)(nil), (*ExecResponse_StandardError)(nil), + (*ExecResponse_Started_)(nil), } type x struct{} out := protoimpl.TypeBuilder{ File: protoimpl.DescBuilder{ GoPackagePath: reflect.TypeOf(x{}).PkgPath(), RawDescriptor: unsafe.Slice(unsafe.StringData(file_rpc_agent_proto_rawDesc), len(file_rpc_agent_proto_rawDesc)), - NumEnums: 0, - NumMessages: 9, + NumEnums: 1, + NumMessages: 10, NumExtensions: 0, NumServices: 1, }, GoTypes: file_rpc_agent_proto_goTypes, DependencyIndexes: file_rpc_agent_proto_depIdxs, + EnumInfos: file_rpc_agent_proto_enumTypes, MessageInfos: file_rpc_agent_proto_msgTypes, }.Build() File_rpc_agent_proto = out.File diff --git a/internal/rpc/exec.go b/internal/rpc/exec.go index 0e934ab..36f59c7 100644 --- a/internal/rpc/exec.go +++ b/internal/rpc/exec.go @@ -4,23 +4,30 @@ import ( "context" "errors" "fmt" - "github.com/creack/pty" - "github.com/samber/lo" - "go.uber.org/zap" - "golang.org/x/sync/errgroup" - "google.golang.org/grpc" "io" "os" "os/exec" "slices" "strings" + "sync" "syscall" + + "github.com/creack/pty" + "github.com/samber/lo" + "go.uber.org/zap" + "golang.org/x/sync/errgroup" + "google.golang.org/grpc" ) const ( standardStreamsBufferSize = 4096 eofChar = 0x04 + + // execRuntimeFailureExitCode matches Docker's exit code for runtime failures before a process starts. + execRuntimeFailureExitCode = 125 + // signalExitCodeOffset is the base for shell-style exit codes of processes terminated by signals. + signalExitCodeOffset = 128 ) func (rpc *RPC) Exec(stream grpc.BidiStreamingServer[ExecRequest, ExecResponse]) error { @@ -58,8 +65,18 @@ func (rpc *RPC) Exec(stream grpc.BidiStreamingServer[ExecRequest, ExecResponse]) cmd.SysProcAttr = &syscall.SysProcAttr{Setsid: true} if err := cmd.Start(); err != nil { + zap.S().Warnf("failed to start %s: %v", formatCommandAndArgs(firstExecRequestCommand.Command.GetName(), + firstExecRequestCommand.Command.GetArgs()), err) + + return sendStartFailure(stream) + } + + // Explicitly notify the client that the process was started + err = sendStartSuccess(stream) + if err != nil { return err } + if cmd.Process != nil { if err := cmd.Process.Release(); err != nil { return err @@ -114,9 +131,20 @@ func (rpc *RPC) Exec(stream grpc.BidiStreamingServer[ExecRequest, ExecResponse]) err = cmd.Start() } + + if err != nil { + zap.S().Warnf("failed to start %s: %v", formatCommandAndArgs(firstExecRequestCommand.Command.GetName(), + firstExecRequestCommand.Command.GetArgs()), err) + + return sendStartFailure(stream) + } + + // Explicitly notify the client that the process was started + err = sendStartSuccess(stream) if err != nil { return err } + if ptmx != nil { defer ptmx.Close() } @@ -183,12 +211,41 @@ func (rpc *RPC) Exec(stream grpc.BidiStreamingServer[ExecRequest, ExecResponse]) }); err != nil { fromClientErrCh <- err + return + } + case *ExecRequest_Signal_: + var signal syscall.Signal + + switch typedAction.Signal { + case ExecRequest_SIGNAL_SIGTERM: + signal = syscall.SIGTERM + case ExecRequest_SIGNAL_SIGKILL: + signal = syscall.SIGKILL + default: + fromClientErrCh <- fmt.Errorf("unsupported exec signal %q", typedAction.Signal.String()) + + return + } + + if err := cmd.Process.Signal(signal); err != nil && !errors.Is(err, os.ErrProcessDone) { + fromClientErrCh <- err + return } } } }() + // Serialize responses from the stdout and stderr goroutines + var sendMutex sync.Mutex + + sendResponse := func(response *ExecResponse) error { + sendMutex.Lock() + defer sendMutex.Unlock() + + return stream.Send(response) + } + group, _ := errgroup.WithContext(stream.Context()) // Handle standard output from the command @@ -210,7 +267,7 @@ func (rpc *RPC) Exec(stream grpc.BidiStreamingServer[ExecRequest, ExecResponse]) return err } - if err := stream.Send(&ExecResponse{ + if err := sendResponse(&ExecResponse{ Type: &ExecResponse_StandardOutput{ StandardOutput: &IOChunk{ Data: slices.Clone(buf[:n]), @@ -240,7 +297,7 @@ func (rpc *RPC) Exec(stream grpc.BidiStreamingServer[ExecRequest, ExecResponse]) return err } - if err := stream.Send(&ExecResponse{ + if err := sendResponse(&ExecResponse{ Type: &ExecResponse_StandardError{ StandardError: &IOChunk{ Data: slices.Clone(buf[:n]), @@ -262,11 +319,16 @@ func (rpc *RPC) Exec(stream grpc.BidiStreamingServer[ExecRequest, ExecResponse]) if err := cmd.Wait(); err != nil { var exitError *exec.ExitError - if errors.As(err, &exitError) { - exitCode = exitError.ExitCode() - } else { + if !errors.As(err, &exitError) { return err } + + // ExitCode returns -1 for signals; report the containerd-compatible 128 + signal instead + exitCode = exitError.ExitCode() + + if waitStatus, ok := exitError.Sys().(syscall.WaitStatus); ok && waitStatus.Signaled() { + exitCode = signalExitCodeOffset + int(waitStatus.Signal()) + } } return stream.Send(&ExecResponse{ @@ -278,6 +340,24 @@ func (rpc *RPC) Exec(stream grpc.BidiStreamingServer[ExecRequest, ExecResponse]) }) } +func sendStartSuccess(stream grpc.BidiStreamingServer[ExecRequest, ExecResponse]) error { + return stream.Send(&ExecResponse{ + Type: &ExecResponse_Started_{ + Started: &ExecResponse_Started{}, + }, + }) +} + +func sendStartFailure(stream grpc.BidiStreamingServer[ExecRequest, ExecResponse]) error { + return stream.Send(&ExecResponse{ + Type: &ExecResponse_Exit_{ + Exit: &ExecResponse_Exit{ + Code: execRuntimeFailureExitCode, + }, + }, + }) +} + func applyExecOverrides(cmd *exec.Cmd, command *ExecRequest_Command) { if command.Workdir != "" { cmd.Dir = command.Workdir diff --git a/internal/rpc/exec_test.go b/internal/rpc/exec_test.go new file mode 100644 index 0000000..c4cd7a1 --- /dev/null +++ b/internal/rpc/exec_test.go @@ -0,0 +1,196 @@ +// In-process stream scaffolding intentionally favors direct test construction. +// +//nolint:containedctx,testpackage,wsl_v5 +package rpc + +import ( + "context" + "io" + "syscall" + "testing" + "time" + + "github.com/stretchr/testify/require" + "google.golang.org/grpc" +) + +const execTestTimeout = 5 * time.Second + +type execTestStream struct { + grpc.ServerStream + + ctx context.Context + requests chan *ExecRequest + responses chan *ExecResponse +} + +var _ grpc.BidiStreamingServer[ExecRequest, ExecResponse] = (*execTestStream)(nil) + +func newExecTestStream(ctx context.Context) *execTestStream { + return &execTestStream{ + ctx: ctx, + requests: make(chan *ExecRequest, 8), + responses: make(chan *ExecResponse, 8), + } +} + +func (stream *execTestStream) Send(response *ExecResponse) error { + select { + case stream.responses <- response: + return nil + case <-stream.ctx.Done(): + return stream.ctx.Err() + } +} + +func (stream *execTestStream) Recv() (*ExecRequest, error) { + select { + case request, ok := <-stream.requests: + if !ok { + return nil, io.EOF + } + return request, nil + case <-stream.ctx.Done(): + return nil, stream.ctx.Err() + } +} + +func (stream *execTestStream) Context() context.Context { return stream.ctx } + +func TestExecSendsStartedBeforeOutputAndExit(t *testing.T) { + stream, result := startExecTest(t, &ExecRequest_Command{ + Name: "/bin/sh", + Args: []string{"-c", "printf hello"}, + }) + + first := receiveExecResponse(t, stream) + require.NotNil(t, first.GetStarted()) + + var output []byte + for { + response := receiveExecResponse(t, stream) + switch response := response.GetType().(type) { + case *ExecResponse_StandardOutput: + output = append(output, response.StandardOutput.GetData()...) + case *ExecResponse_Exit_: + require.EqualValues(t, 0, response.Exit.GetCode()) + require.Equal(t, []byte("hello"), output) + require.NoError(t, receiveExecResult(t, result)) + return + default: + t.Fatalf("unexpected exec response %T", response) + } + } +} + +func TestExecReportsStartFailureBeforeStarted(t *testing.T) { + tests := []struct { + name string + command *ExecRequest_Command + }{ + { + name: "missing executable", + command: &ExecRequest_Command{ + Name: "/definitely/missing/tart-guest-agent-test-command", + }, + }, + { + name: "missing workdir", + command: &ExecRequest_Command{ + Name: "/bin/sh", + Workdir: "/definitely/missing/tart-guest-agent-test-workdir", + }, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + stream, result := startExecTest(t, test.command) + response := receiveExecResponse(t, stream) + require.Nil(t, response.GetStarted()) + require.EqualValues(t, execRuntimeFailureExitCode, response.GetExit().GetCode()) + require.NoError(t, receiveExecResult(t, result)) + }) + } +} + +func TestExecSignalsProcess(t *testing.T) { + tests := []struct { + name string + signal ExecRequest_Signal + code int32 + }{ + { + name: "SIGTERM", + signal: ExecRequest_SIGNAL_SIGTERM, + code: int32(signalExitCodeOffset + syscall.SIGTERM), + }, + { + name: "SIGKILL", + signal: ExecRequest_SIGNAL_SIGKILL, + code: int32(signalExitCodeOffset + syscall.SIGKILL), + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + stream, result := startExecTest(t, &ExecRequest_Command{ + Name: "/bin/sleep", + Args: []string{"30"}, + }) + require.NotNil(t, receiveExecResponse(t, stream).GetStarted()) + + stream.requests <- &ExecRequest{ + Type: &ExecRequest_Signal_{Signal: test.signal}, + } + + response := receiveExecResponse(t, stream) + require.NotNil(t, response.GetExit()) + require.Equal(t, test.code, response.GetExit().GetCode()) + require.NoError(t, receiveExecResult(t, result)) + }) + } +} + +func startExecTest( + t *testing.T, + command *ExecRequest_Command, +) (*execTestStream, <-chan error) { + t.Helper() + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + stream := newExecTestStream(ctx) + result := make(chan error, 1) + go func() { + result <- (&RPC{}).Exec(stream) + }() + stream.requests <- &ExecRequest{ + Type: &ExecRequest_Command_{Command: command}, + } + return stream, result +} + +func receiveExecResponse(t *testing.T, stream *execTestStream) *ExecResponse { + t.Helper() + + select { + case response := <-stream.responses: + return response + case <-time.After(execTestTimeout): + t.Fatal("timed out waiting for exec response") + return nil + } +} + +func receiveExecResult(t *testing.T, result <-chan error) error { + t.Helper() + + select { + case err := <-result: + return err + case <-time.After(execTestTimeout): + t.Fatal("timed out waiting for Exec to return") + return nil + } +} diff --git a/proto/rpc/agent.proto b/proto/rpc/agent.proto index 638fdd7..7328dbb 100644 --- a/proto/rpc/agent.proto +++ b/proto/rpc/agent.proto @@ -10,6 +10,12 @@ service Agent { } message ExecRequest { + enum Signal { + SIGNAL_UNSPECIFIED = 0; + SIGNAL_SIGTERM = 1; + SIGNAL_SIGKILL = 2; + } + message Command { string name = 1; repeated string args = 2; @@ -25,6 +31,7 @@ message ExecRequest { Command command = 1; IOChunk standard_input = 2; TerminalSize terminal_resize = 3; + Signal signal = 4; } } @@ -33,10 +40,15 @@ message ExecResponse { int32 code = 1; } + message Started { + // nothing for now + } + oneof type { Exit exit = 1; IOChunk standard_output = 2; IOChunk standard_error = 3; + Started started = 4; } }