diff --git a/api/runtime/boot/v1/helpers.go b/api/runtime/boot/v1/helpers.go index 92d0d6246d..e8ce0350be 100644 --- a/api/runtime/boot/v1/helpers.go +++ b/api/runtime/boot/v1/helpers.go @@ -44,18 +44,17 @@ func (p *BootstrapParams) AddExtension(msg proto.Message) error { } // FindExtension finds an extension matching the type of dst and unmarshals it. -// Returns true if found, false if not found. -func (p *BootstrapParams) FindExtension(dst proto.Message) error { +func (p *BootstrapParams) FindExtension(dst proto.Message) (bool, error) { name := dst.ProtoReflect().Descriptor().FullName() for _, ext := range p.Extensions { if ext.GetValue().MessageIs(dst) { if err := ext.GetValue().UnmarshalTo(dst); err != nil { - return fmt.Errorf("failed to unmarshal extension %q: %w", name, err) + return false, fmt.Errorf("failed to unmarshal extension %q: %w", name, err) } - return nil + return true, nil } } - return nil + return false, nil } diff --git a/api/runtime/boot/v1/helpers_test.go b/api/runtime/boot/v1/helpers_test.go index ae7cec5c38..00bb0d9623 100644 --- a/api/runtime/boot/v1/helpers_test.go +++ b/api/runtime/boot/v1/helpers_test.go @@ -21,43 +21,47 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + options "github.com/containerd/containerd/api/types/runc/options" "google.golang.org/protobuf/types/known/anypb" ) func TestExtensions(t *testing.T) { params := &BootstrapParams{} - err := params.AddExtension(&RuncV2Extensions{SpecAnnotations: map[string]string{"test": "value"}}) + err := params.AddExtension(&options.Options{ShimCgroup: "test-cgroup"}) require.NoError(t, err) - got := &RuncV2Extensions{} - err = params.FindExtension(got) + got := &options.Options{} + found, err := params.FindExtension(got) require.NoError(t, err) - assert.Equal(t, "value", got.SpecAnnotations["test"]) + assert.True(t, found) + assert.Equal(t, "test-cgroup", got.ShimCgroup) } func TestExtensionNotFound(t *testing.T) { params := &BootstrapParams{} - got := &RuncV2Extensions{} - err := params.FindExtension(got) + got := &options.Options{} + found, err := params.FindExtension(got) require.NoError(t, err) + assert.False(t, found) } func TestAddExtensionWithAny(t *testing.T) { params := &BootstrapParams{} - ext := &RuncV2Extensions{SpecAnnotations: map[string]string{"test": "annotation"}} + ext := &options.Options{ShimCgroup: "test-cgroup"} anyVal, err := anypb.New(ext) require.NoError(t, err) err = params.AddExtension(anyVal) require.NoError(t, err) - got := &RuncV2Extensions{} - err = params.FindExtension(got) + got := &options.Options{} + found, err := params.FindExtension(got) require.NoError(t, err) - assert.Equal(t, "annotation", got.SpecAnnotations["test"]) + assert.True(t, found) + assert.Equal(t, "test-cgroup", got.ShimCgroup) - assert.Contains(t, params.Extensions[0].Value.TypeUrl, "RuncV2Extensions") + assert.Contains(t, params.Extensions[0].Value.TypeUrl, "Options") } diff --git a/cmd/containerd-shim-runc-v2/manager/manager_linux.go b/cmd/containerd-shim-runc-v2/manager/manager_linux.go index b33a2a191d..d156198d2e 100644 --- a/cmd/containerd-shim-runc-v2/manager/manager_linux.go +++ b/cmd/containerd-shim-runc-v2/manager/manager_linux.go @@ -263,26 +263,26 @@ func (manager) Start(ctx context.Context, opts *shim.BootstrapParams) (_ *shim.B go cmd.Wait() var runcOpts options.Options - if err := opts.FindExtension(&runcOpts); err != nil { + if found, err := opts.FindExtension(&runcOpts); err != nil { return nil, fmt.Errorf("failed to fetch runc options: %w", err) - } - - if shimCgroup := runcOpts.GetShimCgroup(); shimCgroup != "" { - if cgroups.Mode() == cgroups.Unified { - cg, err := cgroupsv2.Load(shimCgroup) - if err != nil { - return nil, fmt.Errorf("failed to load cgroup %s: %w", shimCgroup, err) - } - if err := cg.AddProc(uint64(cmd.Process.Pid)); err != nil { - return nil, fmt.Errorf("failed to join cgroup %s: %w", shimCgroup, err) - } - } else { - cg, err := cgroup1.Load(cgroup1.StaticPath(shimCgroup)) - if err != nil { - return nil, fmt.Errorf("failed to load cgroup %s: %w", shimCgroup, err) - } - if err := cg.AddProc(uint64(cmd.Process.Pid)); err != nil { - return nil, fmt.Errorf("failed to join cgroup %s: %w", shimCgroup, err) + } else if found { + if shimCgroup := runcOpts.GetShimCgroup(); shimCgroup != "" { + if cgroups.Mode() == cgroups.Unified { + cg, err := cgroupsv2.Load(shimCgroup) + if err != nil { + return nil, fmt.Errorf("failed to load cgroup %s: %w", shimCgroup, err) + } + if err := cg.AddProc(uint64(cmd.Process.Pid)); err != nil { + return nil, fmt.Errorf("failed to join cgroup %s: %w", shimCgroup, err) + } + } else { + cg, err := cgroup1.Load(cgroup1.StaticPath(shimCgroup)) + if err != nil { + return nil, fmt.Errorf("failed to load cgroup %s: %w", shimCgroup, err) + } + if err := cg.AddProc(uint64(cmd.Process.Pid)); err != nil { + return nil, fmt.Errorf("failed to join cgroup %s: %w", shimCgroup, err) + } } } } diff --git a/vendor/github.com/containerd/containerd/api/runtime/boot/v1/helpers.go b/vendor/github.com/containerd/containerd/api/runtime/boot/v1/helpers.go index 92d0d6246d..e8ce0350be 100644 --- a/vendor/github.com/containerd/containerd/api/runtime/boot/v1/helpers.go +++ b/vendor/github.com/containerd/containerd/api/runtime/boot/v1/helpers.go @@ -44,18 +44,17 @@ func (p *BootstrapParams) AddExtension(msg proto.Message) error { } // FindExtension finds an extension matching the type of dst and unmarshals it. -// Returns true if found, false if not found. -func (p *BootstrapParams) FindExtension(dst proto.Message) error { +func (p *BootstrapParams) FindExtension(dst proto.Message) (bool, error) { name := dst.ProtoReflect().Descriptor().FullName() for _, ext := range p.Extensions { if ext.GetValue().MessageIs(dst) { if err := ext.GetValue().UnmarshalTo(dst); err != nil { - return fmt.Errorf("failed to unmarshal extension %q: %w", name, err) + return false, fmt.Errorf("failed to unmarshal extension %q: %w", name, err) } - return nil + return true, nil } } - return nil + return false, nil }