Skip to content
Merged
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
142 changes: 140 additions & 2 deletions packages/orchestrator/pkg/sandbox/cgroup/manager.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ import (
"time"

"go.uber.org/zap"
"golang.org/x/sys/unix"

"github.com/e2b-dev/infra/packages/shared/pkg/logger"
)
Expand All @@ -26,6 +27,9 @@ const (
// NoCgroupFD is a sentinel value indicating that no cgroup file descriptor
// is available (e.g. cgroup accounting is disabled or the FD has been released).
NoCgroupFD = -1

cgroupKillTimeout = 2 * time.Second
cgroupKillPollInterval = 100 * time.Millisecond
)

// Stats contains resource usage statistics from a cgroup
Expand Down Expand Up @@ -103,6 +107,88 @@ func (h *CgroupHandle) GetStats(ctx context.Context) (*Stats, error) {
return h.manager.getStatsForPath(ctx, h.path, h.memoryPeakFile)
}

// Kill terminates all processes currently in this cgroup.
// Safe to call multiple times. Returns nil if the cgroup is already empty or gone.
func (h *CgroupHandle) Kill(ctx context.Context) error {
if h == nil || h.noop || h.removed {
return nil
}

return h.kill(ctx)
}

func (h *CgroupHandle) kill(ctx context.Context) error {
if h == nil || h.noop {
return nil
}

populated, err := h.populated()
if err != nil {
return err
}
if !populated {
return nil
}

if err := os.WriteFile(filepath.Join(h.path, "cgroup.kill"), []byte("1"), 0); err != nil {
if os.IsNotExist(err) {
return nil
}

return fmt.Errorf("failed to write cgroup.kill: %w", err)
}

events, err := os.Open(filepath.Join(h.path, "cgroup.events"))
if os.IsNotExist(err) {
return nil
}
if err != nil {
return fmt.Errorf("failed to open cgroup.events: %w", err)
}
defer events.Close()

deadline := time.Now().Add(cgroupKillTimeout)

for {
populated, err := cgroupEventsPopulated(events)
if err != nil {
return err
}
if !populated {
return nil
}

if err := ctx.Err(); err != nil {
return err
}

remaining := time.Until(deadline)
if remaining <= 0 {
return fmt.Errorf("cgroup %s still has processes after cgroup.kill", h.cgroupName)
}

pollTimeout := min(remaining, cgroupKillPollInterval)
pollTimeoutMillis := int(pollTimeout / time.Millisecond)
if pollTimeoutMillis == 0 {
pollTimeoutMillis = 1
}

fds := []unix.PollFd{{
Fd: int32(events.Fd()),
Events: unix.POLLPRI | unix.POLLERR,
}}

_, err = unix.Poll(fds, pollTimeoutMillis)
if err != nil {
if errors.Is(err, unix.EINTR) {
continue
}

return fmt.Errorf("failed to poll cgroup.events: %w", err)
}
}
}

// Remove closes all open FDs and deletes the cgroup directory.
// The handle should not be used after calling Remove.
// Safe to call multiple times. Returns error if removal fails
Expand Down Expand Up @@ -144,8 +230,8 @@ func (h *CgroupHandle) Remove(ctx context.Context) error {
zap.String("path", h.path),
zap.Error(rmErr))

if err := os.WriteFile(filepath.Join(h.path, "cgroup.kill"), []byte("1"), 0); err != nil && !os.IsNotExist(err) {
logger.L().Warn(ctx, "failed to write cgroup.kill",
if err := h.kill(ctx); err != nil {
logger.L().Warn(ctx, "failed to kill cgroup processes",
zap.String("cgroup_name", h.cgroupName),
zap.String("path", h.path),
zap.Error(err))
Expand Down Expand Up @@ -192,6 +278,58 @@ func (h *CgroupHandle) CgroupName() string {
return h.cgroupName
}

func (h *CgroupHandle) populated() (bool, error) {
data, err := os.ReadFile(filepath.Join(h.path, "cgroup.events"))
if os.IsNotExist(err) {
return false, nil
}
if err != nil {
return false, fmt.Errorf("failed to read cgroup.events: %w", err)
}

return parseCgroupEventsPopulated(data)
}

func cgroupEventsPopulated(file *os.File) (bool, error) {
if _, err := file.Seek(0, io.SeekStart); err != nil {
if os.IsNotExist(err) {
return false, nil
}

return false, fmt.Errorf("failed to seek cgroup.events: %w", err)
}

data, err := io.ReadAll(file)
if os.IsNotExist(err) {
return false, nil
}
if err != nil {
return false, fmt.Errorf("failed to read cgroup.events: %w", err)
}

return parseCgroupEventsPopulated(data)
}

func parseCgroupEventsPopulated(data []byte) (bool, error) {
for line := range strings.SplitSeq(string(data), "\n") {
fields := strings.Fields(line)
if len(fields) != 2 || fields[0] != "populated" {
continue
}

switch fields[1] {
case "0":
return false, nil
case "1":
return true, nil
default:
return false, fmt.Errorf("invalid populated value in cgroup.events: %q", fields[1])
}
}

return false, errors.New("missing populated value in cgroup.events")
}

// Manager handles initialization and creation of cgroups
// Individual cgroup operations are performed through CgroupHandle
type Manager interface {
Expand Down
147 changes: 123 additions & 24 deletions packages/orchestrator/pkg/sandbox/cgroup/manager_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,32 @@ import (
"github.com/stretchr/testify/require"
)

func requireWritableCgroup(t *testing.T) {
t.Helper()

if os.Geteuid() != 0 {
t.Skip("test requires root privileges")
}

probePath := filepath.Join(cgroupV2MountPoint, fmt.Sprintf("e2b-probe-%d-%d", os.Getpid(), time.Now().UnixNano()))
if err := os.Mkdir(probePath, 0o755); err != nil {
t.Skipf("test requires writable cgroup v2 filesystem: %v", err)
}

subtreeControlPath := filepath.Join(probePath, "cgroup.subtree_control")
if _, err := os.Stat(subtreeControlPath); err != nil {
_ = os.Remove(probePath)
t.Skipf("test requires usable cgroup v2 control files: %v", err)
}
if err := os.WriteFile(subtreeControlPath, []byte("+cpu +memory"), 0); err != nil {
_ = os.Remove(probePath)
t.Skipf("test requires writable cgroup v2 subtree control: %v", err)
}
if err := os.Remove(probePath); err != nil {
t.Skipf("test requires removable cgroup v2 filesystem: %v", err)
}
}

func TestNewManager(t *testing.T) {
t.Parallel()

Expand All @@ -31,9 +57,7 @@ func TestNewManager(t *testing.T) {
func TestManagerInitialize(t *testing.T) {
t.Parallel()

if os.Geteuid() != 0 {
t.Skip("test requires root privileges")
}
requireWritableCgroup(t)

ctx := t.Context()

Expand All @@ -58,9 +82,7 @@ func TestManagerInitialize(t *testing.T) {
func TestCgroupHandleLifecycle(t *testing.T) {
t.Parallel()

if os.Geteuid() != 0 {
t.Skip("test requires root privileges")
}
requireWritableCgroup(t)

ctx := t.Context()
mgr, err := NewManager()
Expand Down Expand Up @@ -96,9 +118,7 @@ func TestCgroupHandleLifecycle(t *testing.T) {
func TestCgroupHandleWithProcessCreation(t *testing.T) {
t.Parallel()

if os.Geteuid() != 0 {
t.Skip("test requires root privileges")
}
requireWritableCgroup(t)

ctx := t.Context()
mgr, err := NewManager()
Expand Down Expand Up @@ -151,9 +171,7 @@ func TestCgroupHandleWithProcessCreation(t *testing.T) {
func TestCgroupHandleNoRaceOnQuickExit(t *testing.T) {
t.Parallel()

if os.Geteuid() != 0 {
t.Skip("test requires root privileges")
}
requireWritableCgroup(t)

ctx := t.Context()
mgr, err := NewManager()
Expand Down Expand Up @@ -188,9 +206,7 @@ func TestCgroupHandleNoRaceOnQuickExit(t *testing.T) {
func TestCgroupHandleGetStats(t *testing.T) {
t.Parallel()

if os.Geteuid() != 0 {
t.Skip("test requires root privileges")
}
requireWritableCgroup(t)

ctx := t.Context()
mgr, err := NewManager()
Expand Down Expand Up @@ -237,9 +253,7 @@ func TestCgroupHandleGetStats(t *testing.T) {
func TestCgroupHandleGetStatsNonExistent(t *testing.T) {
t.Parallel()

if os.Geteuid() != 0 {
t.Skip("test requires root privileges")
}
requireWritableCgroup(t)

ctx := t.Context()
mgr, err := NewManager()
Expand All @@ -262,12 +276,99 @@ func TestCgroupHandleGetStatsNonExistent(t *testing.T) {
assert.Contains(t, err.Error(), "failed to read cpu.stat")
}

func TestCgroupHandleRemoveNonExistent(t *testing.T) {
func TestCgroupHandleKillNoProcesses(t *testing.T) {
t.Parallel()

if os.Geteuid() != 0 {
t.Skip("test requires root privileges")
cgroupPath := t.TempDir()
require.NoError(t, os.WriteFile(filepath.Join(cgroupPath, "cgroup.events"), []byte("populated 0\n"), 0o644))

handle := &CgroupHandle{
cgroupName: "test-empty-kill",
path: cgroupPath,
}

require.NoError(t, handle.Kill(t.Context()))
}

func TestCgroupHandlePopulated(t *testing.T) {
t.Parallel()

for _, tc := range []struct {
name string
content string
write bool
want bool
}{
{
name: "populated",
content: "populated 1\n",
write: true,
want: true,
},
{
name: "empty",
content: "populated 0\n",
write: true,
},
{
name: "missing file",
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()

cgroupPath := t.TempDir()
if tc.write {
require.NoError(t, os.WriteFile(filepath.Join(cgroupPath, "cgroup.events"), []byte(tc.content), 0o644))
}

handle := &CgroupHandle{path: cgroupPath}

got, err := handle.populated()
require.NoError(t, err)
assert.Equal(t, tc.want, got)
})
}
}

func TestCgroupHandleKillTerminatesProcesses(t *testing.T) {
t.Parallel()

requireWritableCgroup(t)

ctx := t.Context()
mgr, err := NewManager()
require.NoError(t, err)

err = mgr.Initialize(ctx)
require.NoError(t, err)

handle, err := mgr.Create(ctx, "test-cgroup-kill")
require.NoError(t, err)
defer handle.Remove(ctx)
defer handle.ReleaseCgroupFD()

cmd := exec.CommandContext(ctx, "sleep", "60")
cmd.SysProcAttr = &syscall.SysProcAttr{
UseCgroupFD: true,
CgroupFD: handle.GetFD(),
}

require.NoError(t, cmd.Start())
t.Cleanup(func() {
_ = cmd.Process.Kill()
_ = cmd.Wait()
})

require.NoError(t, handle.ReleaseCgroupFD())
require.NoError(t, handle.Kill(ctx))
require.Error(t, cmd.Wait())
}

func TestCgroupHandleRemoveNonExistent(t *testing.T) {
t.Parallel()

requireWritableCgroup(t)

ctx := t.Context()
mgr, err := NewManager()
Expand Down Expand Up @@ -329,9 +430,7 @@ burst_usec 0`
func TestCgroupHandlePeakReset(t *testing.T) {
t.Parallel()

if os.Geteuid() != 0 {
t.Skip("test requires root privileges")
}
requireWritableCgroup(t)

ctx := t.Context()
mgr, err := NewManager()
Expand Down
Loading
Loading