diff --git a/packages/orchestrator/internal/sandbox/sandbox.go b/packages/orchestrator/internal/sandbox/sandbox.go index 03302888f4..68cc59d0c7 100644 --- a/packages/orchestrator/internal/sandbox/sandbox.go +++ b/packages/orchestrator/internal/sandbox/sandbox.go @@ -15,7 +15,6 @@ import ( "go.opentelemetry.io/otel/metric" "go.opentelemetry.io/otel/trace" "go.uber.org/zap" - "golang.org/x/sync/errgroup" "github.com/e2b-dev/infra/packages/orchestrator/internal/cfg" "github.com/e2b-dev/infra/packages/orchestrator/internal/sandbox/block" @@ -387,148 +386,188 @@ func (f *Factory) ResumeSandbox( telemetry.ReportEvent(ctx, "created sandbox files") - var wg errgroup.Group + // Uffd initialization + fcUffdPath := sandboxFiles.SandboxUffdSocketPath() + uffdPromise := utils.NewPromise(func() (*uffd.Uffd, error) { + memfile, err := t.Memfile(ctx) + if err != nil { + return nil, fmt.Errorf("failed to get memfile: %w", err) + } + + telemetry.ReportEvent(ctx, "got template memfile") + + fcUffd, err := uffd.New(memfile, fcUffdPath) + if err != nil { + return nil, fmt.Errorf("failed to create uffd: %w", err) + } + + cleanup.AddNoContext(ctx, fcUffd.Close) + + return fcUffd, nil + }) + + // Prefetching + go func() { + memfile, err := t.Memfile(ctx) + if err != nil { + return + } + + meta, err := t.Metadata() + if err != nil { + return + } + + telemetry.ReportEvent(ctx, "got metadata") + + // Start background prefetcher as early as possible if prefetch mapping exists + // Fetching from source starts immediately; copying waits for uffd to be ready + if meta.Prefetch != nil && meta.Prefetch.Memory != nil { + fcUffd, err := uffdPromise.Wait(ctx) + if err != nil { + return + } + + telemetry.ReportEvent(ctx, "starting prefetcher") + l := logger.L().With(logger.WithSandboxID(runtime.SandboxID), logger.WithTemplateID(runtime.TemplateID), logger.WithTeamID(runtime.TeamID)) - var ipsCh chan networkSlotRes - wg.Go(func() error { - ipsCh = getNetworkSlotAsync(ctx, f.networkPool, cleanup, config.Network) - cleanup.Add(ctx, func(_ context.Context) error { - // Ensure the slot is received from chan before ResumeSandbox returns so the slot is cleaned up properly in cleanup - <-ipsCh + go func() { + p := prefetch.New( + l, + memfile, + fcUffd, + meta.Prefetch.Memory, + f.featureFlags, + ) + err := p.Start(execCtx) + if err != nil { + l.Error(ctx, "failed to start prefetcher", zap.Error(err)) + } + }() + } + }() + + // Slot initialization + ipsPromise := utils.NewPromise(func() (*network.Slot, error) { + slot, err := f.networkPool.Get(ctx, config.Network) + if err != nil { + return nil, fmt.Errorf("failed to get network slot: %w", err) + } + + cleanup.Add(ctx, func(ctx context.Context) error { + ctx, span := tracer.Start(ctx, "clean network-slot") + defer span.End() + + go func(ctx context.Context) { + returnErr := f.networkPool.Return(ctx, slot) + if returnErr != nil { + logger.L().Error(ctx, "failed to return network slot", zap.Error(returnErr)) + } + }(ctx) return nil }) - return nil + return slot, nil }) - var readonlyRootfs block.ReadonlyDevice - var rootfsOverlay rootfs.Provider - wg.Go(func() error { - var err error - - readonlyRootfs, err = t.Rootfs() + // Rootfs initialization + overlayPromise := utils.NewPromise(func() (rootfs.Provider, error) { + readonlyRootfs, err := t.Rootfs() if err != nil { - return fmt.Errorf("failed to get rootfs: %w", err) + return nil, fmt.Errorf("failed to get rootfs: %w", err) } telemetry.ReportEvent(ctx, "got template rootfs") - rootfsOverlay, err = rootfs.NewNBDProvider( + overlay, err := rootfs.NewNBDProvider( readonlyRootfs, sandboxFiles.SandboxCacheRootfsPath(f.config.StorageConfig), f.devicePool, f.featureFlags, ) if err != nil { - return fmt.Errorf("failed to create rootfs overlay: %w", err) + return nil, fmt.Errorf("failed to create rootfs overlay: %w", err) } - cleanup.Add(ctx, rootfsOverlay.Close) + cleanup.Add(ctx, overlay.Close) telemetry.ReportEvent(ctx, "created rootfs overlay") go func() { - runErr := rootfsOverlay.Start(execCtx) + runErr := overlay.Start(execCtx) if runErr != nil { logger.L().Error(ctx, "rootfs overlay error", zap.Error(runErr)) } }() - return nil + return overlay, nil }) - memfile, err := t.Memfile(ctx) - if err != nil { - return nil, fmt.Errorf("failed to get memfile: %w", err) - } - - telemetry.ReportEvent(ctx, "got template memfile") - - fcUffdPath := sandboxFiles.SandboxUffdSocketPath() - fcUffd, err := uffd.New(memfile, fcUffdPath) - if err != nil { - return nil, fmt.Errorf("failed to create uffd: %w", err) - } - cleanup.AddNoContext(ctx, fcUffd.Close) + // Memory initialization + memoryPromise := utils.NewPromise(func() (struct{}, error) { + fcUffd, err := uffdPromise.Wait(ctx) + if err != nil { + return struct{}{}, err + } - wg.Go(func() error { - err := serveMemory( + err = serveMemory( execCtx, cleanup, fcUffd, runtime.SandboxID, ) if err != nil { - return fmt.Errorf("failed to serve memory: %w", err) + return struct{}{}, fmt.Errorf("failed to serve memory: %w", err) } telemetry.ReportEvent(ctx, "started serving memory") - return nil + return struct{}{}, nil }) - var meta metadata.Template - wg.Go(func() error { - var err error - - meta, err = t.Metadata() - if err != nil { - return fmt.Errorf("failed to get metadata: %w", err) - } - - telemetry.ReportEvent(ctx, "got metadata") - - // Start background prefetcher as early as possible if prefetch mapping exists - // Fetching from source starts immediately; copying waits for uffd to be ready - if meta.Prefetch != nil && meta.Prefetch.Memory != nil { - telemetry.ReportEvent(ctx, "starting prefetcher") - l := logger.L().With(logger.WithSandboxID(runtime.SandboxID), logger.WithTemplateID(runtime.TemplateID), logger.WithTeamID(runtime.TeamID)) + // Wait for all resources to be initialized + ips, err := ipsPromise.Wait(ctx) + if err != nil { + return nil, fmt.Errorf("failed to get network slot: %w", err) + } - go func() { - p := prefetch.New( - l, - memfile, - fcUffd, - meta.Prefetch.Memory, - f.featureFlags, - ) - err := p.Start(execCtx) - if err != nil { - l.Error(ctx, "failed to start prefetcher", zap.Error(err)) - } - }() - } + telemetry.ReportEvent(ctx, "got network slot") - return nil - }) + overlay, err := overlayPromise.Wait(ctx) + if err != nil { + return nil, err + } - err = wg.Wait() + _, err = memoryPromise.Wait(ctx) if err != nil { - return nil, fmt.Errorf("failed to initialize resources: %w", err) + return nil, err } // ==== END of resources initialization ==== - ips := <-ipsCh - if ips.err != nil { - return nil, fmt.Errorf("failed to get network slot: %w", ips.err) + rootfs, err := t.Rootfs() + if err != nil { + return nil, fmt.Errorf("failed to get rootfs overlay: %w", err) } - telemetry.ReportEvent(ctx, "got network slot") + meta, err := t.Metadata() + if err != nil { + return nil, fmt.Errorf("failed to get metadata: %w", err) + } fcHandle, fcErr := fc.NewProcess( ctx, execCtx, f.config, - ips.slot, + ips, sandboxFiles, // The versions need to base exactly the same as the paused sandbox template because of the FC compatibility. config.FirecrackerConfig, - rootfsOverlay, + overlay, fc.RootfsPaths{ TemplateVersion: meta.Version, TemplateID: config.BaseTemplateID, - BuildID: readonlyRootfs.Header().Metadata.BaseBuildId.String(), + BuildID: rootfs.Header().Metadata.BaseBuildId.String(), }, ) if fcErr != nil { @@ -545,6 +584,11 @@ func (f *Factory) ResumeSandbox( telemetry.ReportEvent(ctx, "got snapfile") + fcUffd, err := uffdPromise.Wait(ctx) + if err != nil { + return nil, fmt.Errorf("failed to get uffd: %w", err) + } + uffdStartCtx, cancelUffdStartCtx := context.WithCancelCause(ctx) defer cancelUffdStartCtx(fmt.Errorf("uffd finished starting")) go func() { @@ -562,7 +606,7 @@ func (f *Factory) ResumeSandbox( fcUffdPath, snapfile, fcUffd.Ready(), - ips.slot, + ips, ) if fcStartErr != nil { return nil, fmt.Errorf("failed to start FC: %w", fcStartErr) @@ -571,8 +615,8 @@ func (f *Factory) ResumeSandbox( telemetry.ReportEvent(ctx, "initialized FC") resources := &Resources{ - Slot: ips.slot, - rootfs: rootfsOverlay, + Slot: ips, + rootfs: overlay, memory: fcUffd, } diff --git a/packages/shared/pkg/utils/promise.go b/packages/shared/pkg/utils/promise.go new file mode 100644 index 0000000000..9d3191e055 --- /dev/null +++ b/packages/shared/pkg/utils/promise.go @@ -0,0 +1,44 @@ +package utils + +import "context" + +// Promise represents an asynchronous computation that will eventually produce a value or error. +// The computation starts immediately when the promise is created. +// Multiple goroutines can safely wait on the same promise. +type Promise[T any] struct { + result *SetOnce[T] +} + +// NewPromise creates a new Promise that immediately starts executing the given function +// in a goroutine. The result (value or error) will be available via Wait. +func NewPromise[T any](fn func() (T, error)) *Promise[T] { + p := &Promise[T]{ + result: NewSetOnce[T](), + } + + go func() { + value, err := fn() + p.result.SetResult(value, err) + }() + + return p +} + +// Wait blocks until the promise is resolved and returns the result. +// Returns the value and nil error on success, or zero value and the error on failure. +// If the context is cancelled before the promise resolves, returns ctx.Err(). +func (p *Promise[T]) Wait(ctx context.Context) (T, error) { + return p.result.WaitWithContext(ctx) +} + +// Done returns a channel that's closed when the promise is resolved. +// This allows using Promise in select statements. +func (p *Promise[T]) Done() <-chan struct{} { + return p.result.Done +} + +// Result returns the current result without blocking. +// Returns NotSetError if the promise hasn't resolved yet. +func (p *Promise[T]) Result() (T, error) { + return p.result.Result() +} diff --git a/packages/shared/pkg/utils/promise_test.go b/packages/shared/pkg/utils/promise_test.go new file mode 100644 index 0000000000..5f309b94f6 --- /dev/null +++ b/packages/shared/pkg/utils/promise_test.go @@ -0,0 +1,151 @@ +package utils + +import ( + "context" + "errors" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestPromiseSuccess(t *testing.T) { + t.Parallel() + + p := NewPromise(func() (int, error) { + return 42, nil + }) + + value, err := p.Wait(context.Background()) + require.NoError(t, err) + assert.Equal(t, 42, value) +} + +func TestPromiseError(t *testing.T) { + t.Parallel() + + expectedErr := errors.New("test error") + p := NewPromise(func() (int, error) { + return 0, expectedErr + }) + + value, err := p.Wait(context.Background()) + require.ErrorIs(t, err, expectedErr) + assert.Equal(t, 0, value) +} + +func TestPromiseDelayedResult(t *testing.T) { + t.Parallel() + + p := NewPromise(func() (string, error) { + time.Sleep(50 * time.Millisecond) + + return "delayed", nil + }) + + value, err := p.Wait(context.Background()) + require.NoError(t, err) + assert.Equal(t, "delayed", value) +} + +func TestPromiseContextCancelled(t *testing.T) { + t.Parallel() + + p := NewPromise(func() (int, error) { + time.Sleep(1 * time.Second) + + return 42, nil + }) + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond) + defer cancel() + + _, err := p.Wait(ctx) + assert.ErrorIs(t, err, context.DeadlineExceeded) +} + +func TestPromiseMultipleWaiters(t *testing.T) { + t.Parallel() + + p := NewPromise(func() (int, error) { + time.Sleep(50 * time.Millisecond) + + return 42, nil + }) + + var wg sync.WaitGroup + results := make([]int, 5) + errs := make([]error, 5) + + for i := range 5 { + wg.Add(1) + go func(idx int) { + defer wg.Done() + results[idx], errs[idx] = p.Wait(context.Background()) + }(i) + } + + wg.Wait() + + for i := range 5 { + require.NoError(t, errs[i]) + assert.Equal(t, 42, results[i]) + } +} + +func TestPromiseDone(t *testing.T) { + t.Parallel() + + p := NewPromise(func() (int, error) { + time.Sleep(50 * time.Millisecond) + + return 42, nil + }) + + select { + case <-p.Done(): + t.Fatal("promise should not be done yet") + default: + // expected + } + + <-p.Done() + + value, err := p.Result() + require.NoError(t, err) + assert.Equal(t, 42, value) +} + +func TestPromiseResultBeforeResolve(t *testing.T) { + t.Parallel() + + p := NewPromise(func() (int, error) { + time.Sleep(100 * time.Millisecond) + + return 42, nil + }) + + _, err := p.Result() + assert.ErrorAs(t, err, &NotSetError{}) +} + +func TestPromiseResultAfterResolve(t *testing.T) { + t.Parallel() + + p := NewPromise(func() (int, error) { + return 42, nil + }) + + <-p.Done() + + value, err := p.Result() + require.NoError(t, err) + assert.Equal(t, 42, value) + + // Multiple calls should return the same result + value2, err2 := p.Result() + require.NoError(t, err2) + assert.Equal(t, 42, value2) +}