Skip to content
Closed
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
3 changes: 3 additions & 0 deletions iac/modules/job-orchestrator/jobs/orchestrator.hcl
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,9 @@ job "orchestrator-${latest_orchestrator_job_id}" {
attempts = 0
}

# Matches the Nomad client max_kill_timeout.
kill_timeout = "24h"

resources {
memory = 1024
memory_max = -1
Expand Down
54 changes: 44 additions & 10 deletions packages/orchestrator/pkg/factories/run.go
Original file line number Diff line number Diff line change
Expand Up @@ -119,6 +119,16 @@ func (e serviceDoneError) Error() string {
return fmt.Sprintf("service %s finished", e.name)
}

func isServiceDoneError(err error) bool {
var sde serviceDoneError

return errors.As(err, &sde)
}

func isIgnorableSyncError(err error) bool {
return errors.Is(err, syscall.EINVAL)
}

// Run starts the orchestrator, blocking until shutdown.
// Returns true on clean shutdown.
func Run(opts Options) bool {
Expand Down Expand Up @@ -234,7 +244,7 @@ func run(config cfg.Config, opts Options) (success bool) {
// there's a panic.
defer func(g *errgroup.Group) {
err := g.Wait()
if err != nil {
if err != nil && !isServiceDoneError(err) {
log.Printf("error while shutting down: %v", err)
success = false
}
Expand Down Expand Up @@ -275,7 +285,7 @@ func run(config cfg.Config, opts Options) (success bool) {
}))
defer func(l logger.Logger) {
err := l.Sync()
if err != nil {
if err != nil && !isIgnorableSyncError(err) {
log.Printf("error while shutting down logger: %v", err)
success = false
}
Expand All @@ -293,7 +303,7 @@ func run(config cfg.Config, opts Options) (success bool) {
)
defer func(l logger.Logger) {
err := l.Sync()
if err != nil {
if err != nil && !isIgnorableSyncError(err) {
log.Printf("error while shutting down sandbox logger: %v", err)
success = false
}
Expand All @@ -311,7 +321,7 @@ func run(config cfg.Config, opts Options) (success bool) {
)
defer func(l logger.Logger) {
err := l.Sync()
if err != nil {
if err != nil && !isIgnorableSyncError(err) {
log.Printf("error while shutting down sandbox logger: %v", err)
success = false
}
Expand Down Expand Up @@ -575,8 +585,8 @@ func run(config cfg.Config, opts Options) (success bool) {
if err != nil {
logger.L().Fatal(ctx, "failed to create orchestrator server", zap.Error(err))
}
closers = append(closers, closer{"orchestrator server", func(context.Context) error {
return orchestratorService.Close()
closers = append(closers, closer{"orchestrator server", func(closeCtx context.Context) error {
return orchestratorService.Close(closeCtx)
}})

// template manager sandbox logger
Expand All @@ -591,8 +601,7 @@ func run(config cfg.Config, opts Options) (success bool) {
)
closers = append(closers, closer{
"template manager sandbox logger", func(context.Context) error {
// Sync returns EINVAL when path is /dev/stdout (for example)
if err := tmplSbxLoggerExternal.Sync(); err != nil && !errors.Is(err, syscall.EINVAL) {
if err := tmplSbxLoggerExternal.Sync(); err != nil && !isIgnorableSyncError(err) {
return err
}

Expand Down Expand Up @@ -773,6 +782,8 @@ func run(config cfg.Config, opts Options) (success bool) {
}
}

orchestratorService.StartDraining(ctx)

// Wait for services to be drained before closing them
if tmpl != nil {
err := tmpl.Wait(closeCtx)
Expand All @@ -782,6 +793,30 @@ func run(config cfg.Config, opts Options) (success bool) {
}
}

forceStopSandboxes := func() {
forceErr := orchestratorService.ForceStopSandboxes(context.WithoutCancel(ctx))
if forceErr != nil {
logger.L().Error(ctx, "forced sandbox shutdown failed", zap.Error(forceErr))
success = false
}
}
Comment thread
wj-e2b marked this conversation as resolved.

logger.L().Info(ctx, "Starting sandbox drain phase",
zap.Bool("forced", config.ForceStop),
zap.Int("sandbox_count", sandboxes.Count()),
)

if config.ForceStop {
forceStopSandboxes()
Comment thread
wj-e2b marked this conversation as resolved.
} else {
err := orchestratorService.DrainSandboxes(closeCtx)
if err != nil {
logger.L().Warn(ctx, "sandbox drain phase did not complete gracefully; forcing sandbox shutdown", zap.Error(err))

forceStopSandboxes()
}
}

slices.Reverse(closers)
for _, closer := range closers {
clog := globalLogger.With(zap.String("service", closer.name), zap.Bool("forced", config.ForceStop))
Expand All @@ -793,8 +828,7 @@ func run(config cfg.Config, opts Options) (success bool) {
}

logger.L().Info(ctx, "Waiting for services to finish")
var sde serviceDoneError
if err := g.Wait(); err != nil && !errors.As(err, &sde) {
if err := g.Wait(); err != nil && !isServiceDoneError(err) {
logger.L().Error(ctx, "service group error", zap.Error(err))
success = false
}
Expand Down
67 changes: 67 additions & 0 deletions packages/orchestrator/pkg/sandbox/factory_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,67 @@
//go:build linux

package sandbox

import (
"context"
"testing"
"time"

"github.com/stretchr/testify/require"
)

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

factory := testFactory()
factory.StartDraining(t.Context())

_, err := factory.enterSandboxStart()
require.ErrorIs(t, err, ErrFactoryDraining)
}

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

factory := testFactory()
release, err := factory.enterSandboxStart()
require.NoError(t, err)

waitCtx, cancel := context.WithTimeout(t.Context(), time.Second)
defer cancel()

done := make(chan error, 1)
go func() {
done <- factory.WaitSandboxStarts(waitCtx)
}()

select {
case err := <-done:
require.Failf(t, "WaitSandboxStarts returned before start left", "err: %v", err)
case <-time.After(25 * time.Millisecond):
}

release()
require.NoError(t, <-done)
}

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

factory := testFactory()
release, err := factory.enterSandboxStart()
require.NoError(t, err)
defer release()

waitCtx, cancel := context.WithCancel(t.Context())
cancel()

require.ErrorIs(t, factory.WaitSandboxStarts(waitCtx), context.Canceled)
}

func testFactory() *Factory {
return &Factory{
Sandboxes: NewSandboxesMap(),
drainDone: make(chan struct{}),
}
}
87 changes: 65 additions & 22 deletions packages/orchestrator/pkg/sandbox/fc/process.go
Original file line number Diff line number Diff line change
Expand Up @@ -661,21 +661,25 @@ func (p *Process) Stop(ctx context.Context) error {
logger.L().Warn(ctx, "failed to remove fc metrics FIFO", zap.Error(removeErr), logger.WithSandboxID(p.files.SandboxID))
}

// Check if process has already exited.
pid := p.cmd.Process.Pid

// Check if process has already exited. The parent exiting is not enough for
// cleanup: descendants in the same process group can still hold resources.
select {
case <-p.Exit.Done():
logger.L().Info(ctx, "fc process already exited", logger.WithSandboxID(p.files.SandboxID))

return nil
if !processGroupExists(pid) {
return nil
}
default:
}

// this function should never fail b/c a previous context was canceled.
ctx = context.WithoutCancel(ctx)

err := p.cmd.Process.Signal(syscall.SIGTERM)
err := signalProcessGroup(pid, syscall.SIGTERM)
if err != nil {
if errors.Is(err, os.ErrProcessDone) {
if errors.Is(err, os.ErrProcessDone) && !processGroupExists(pid) {
logger.L().Info(ctx, "fc process already exited", logger.WithSandboxID(p.files.SandboxID))

return nil
Expand All @@ -684,40 +688,79 @@ func (p *Process) Stop(ctx context.Context) error {
logger.L().Warn(ctx, "failed to send SIGTERM to fc process", zap.Error(err), logger.WithSandboxID(p.files.SandboxID))
}

go func() {
termDeadline := time.NewTimer(10 * time.Second)
defer termDeadline.Stop()
poll := time.NewTicker(50 * time.Millisecond)
defer poll.Stop()

for processGroupExists(pid) {
select {
// Wait 10 sec for the FC process to exit, if it doesn't, send SIGKILL.
case <-time.After(10 * time.Second):
// Check process status right before Kill — the pre-SIGTERM status
// captured above is 10s stale and no longer useful here.
status, stateErr := getProcessStatus(p.cmd.Process.Pid)
case <-termDeadline.C:
status, stateErr := getProcessStatus(pid)
if errors.Is(stateErr, process.ErrorProcessNotRunning) {
// Process already exited, no need to send SIGKILL.
return
logger.L().Info(ctx, "fc parent process exited before SIGKILL; checking process group", logger.WithSandboxID(p.files.SandboxID))
} else if stateErr != nil {
logger.L().Warn(ctx, "failed to get fc process status before SIGKILL", zap.Error(stateErr), logger.WithSandboxID(p.files.SandboxID))
}

err := p.cmd.Process.Kill()
if err == nil {
logger.L().Info(ctx, "sent SIGKILL to fc process because it was not responding to SIGTERM for 10 seconds",
killErr := signalProcessGroup(pid, syscall.SIGKILL)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔒 Agentic Security Review
Severity: MEDIUM
The shutdown path still sends SIGKILL to -pid after a delay based only on the stored numeric PID. If the original Firecracker process group exits and that PID is reused before the delayed kill path runs, this can terminate an unrelated process group.

Impact: tenant-triggerable sandbox churn can cause cross-sandbox availability impact by killing the wrong process group.

Fix in Cursor Fix in Web

Reviewed by Cursor Security Reviewer for commit 3c8efb1. Configure here.

if killErr == nil {
logger.L().Info(ctx, "sent SIGKILL to fc process group because it was not responding to SIGTERM for 10 seconds",
zap.Strings("status", status),
logger.WithSandboxID(p.files.SandboxID),
)
}
if err != nil && !errors.Is(err, os.ErrProcessDone) {
logger.L().Warn(ctx, "failed to send SIGKILL to fc process", zap.Error(err), logger.WithSandboxID(p.files.SandboxID))
if killErr != nil && !errors.Is(killErr, os.ErrProcessDone) {
logger.L().Warn(ctx, "failed to send SIGKILL to fc process", zap.Error(killErr), logger.WithSandboxID(p.files.SandboxID))
}

// If the FC process exited, we can return.
case <-p.Exit.Done():
return
killDeadline := time.NewTimer(time.Second)
for processGroupExists(pid) {
select {
case <-killDeadline.C:
return fmt.Errorf("fc process group %d still exists after SIGKILL", pid)
case <-poll.C:
}
}
killDeadline.Stop()

return nil
case <-poll.C:
}
}()
}

return nil
}

func signalProcessGroup(pid int, signal syscall.Signal) error {
if pid <= 0 {
return os.ErrProcessDone
}

// Firecracker is launched with Setsid, so the process PID is also the process
// group ID. Signal the group so unshare/bash/ip descendants cannot keep the VM
// mount namespace or Firecracker process alive after shutdown.
if err := syscall.Kill(-pid, signal); err != nil {
if errors.Is(err, syscall.ESRCH) {
return os.ErrProcessDone
}

return err
}

return nil
}

func processGroupExists(pid int) bool {
if pid <= 0 {
return false
}

err := syscall.Kill(-pid, 0)

return err == nil || errors.Is(err, syscall.EPERM)
}

func (p *Process) Pause(ctx context.Context) error {
ctx, childSpan := tracer.Start(ctx, "pause-fc")
defer childSpan.End()
Expand Down
37 changes: 37 additions & 0 deletions packages/orchestrator/pkg/sandbox/fc/process_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,37 @@
//go:build linux

package fc

import (
"context"
"os/exec"
"syscall"
"testing"
"time"

"github.com/stretchr/testify/require"
)

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

ctx, cancel := context.WithTimeout(t.Context(), time.Minute)
defer cancel()

cmd := exec.CommandContext(ctx, "sleep", "60")
cmd.SysProcAttr = &syscall.SysProcAttr{Setsid: true}
require.NoError(t, cmd.Start())

pid := cmd.Process.Pid
t.Cleanup(func() {
_ = syscall.Kill(-pid, syscall.SIGKILL)
_ = cmd.Wait()
})

require.True(t, processGroupExists(pid))
require.NoError(t, signalProcessGroup(pid, syscall.SIGKILL))
require.Error(t, cmd.Wait())
require.Eventually(t, func() bool {
return !processGroupExists(pid)
}, time.Second, 10*time.Millisecond)
}
Loading
Loading