From 2014e95fb2fd067dd93d8c7aca3794129cffc3a0 Mon Sep 17 00:00:00 2001 From: Jakub Dobry Date: Mon, 12 Jan 2026 15:32:51 +0100 Subject: [PATCH 1/2] promise poc --- .../orchestrator/internal/sandbox/sandbox.go | 208 +++++++++++------- packages/shared/pkg/utils/promise.go | 44 ++++ packages/shared/pkg/utils/promise_test.go | 145 ++++++++++++ 3 files changed, 315 insertions(+), 82 deletions(-) create mode 100644 packages/shared/pkg/utils/promise.go create mode 100644 packages/shared/pkg/utils/promise_test.go 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..4bf20d90d3 --- /dev/null +++ b/packages/shared/pkg/utils/promise_test.go @@ -0,0 +1,145 @@ +package utils + +import ( + "context" + "errors" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" +) + +func TestPromiseSuccess(t *testing.T) { + t.Parallel() + + p := NewPromise(func() (int, error) { + return 42, nil + }) + + value, err := p.Wait(context.Background()) + assert.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()) + assert.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()) + assert.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 := 0; i < 5; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + results[idx], errs[idx] = p.Wait(context.Background()) + }(i) + } + + wg.Wait() + + for i := 0; i < 5; i++ { + assert.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() + assert.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() + assert.NoError(t, err) + assert.Equal(t, 42, value) + + // Multiple calls should return the same result + value2, err2 := p.Result() + assert.NoError(t, err2) + assert.Equal(t, 42, value2) +} From 7608780cce993993ce8e1f0c0dbc16a4fc7a1ecf Mon Sep 17 00:00:00 2001 From: ValentaTomas Date: Mon, 12 Jan 2026 06:53:09 -0800 Subject: [PATCH 2/2] Fix lint --- packages/shared/pkg/utils/promise_test.go | 24 ++++++++++++++--------- 1 file changed, 15 insertions(+), 9 deletions(-) diff --git a/packages/shared/pkg/utils/promise_test.go b/packages/shared/pkg/utils/promise_test.go index 4bf20d90d3..5f309b94f6 100644 --- a/packages/shared/pkg/utils/promise_test.go +++ b/packages/shared/pkg/utils/promise_test.go @@ -8,6 +8,7 @@ import ( "time" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestPromiseSuccess(t *testing.T) { @@ -18,7 +19,7 @@ func TestPromiseSuccess(t *testing.T) { }) value, err := p.Wait(context.Background()) - assert.NoError(t, err) + require.NoError(t, err) assert.Equal(t, 42, value) } @@ -31,7 +32,7 @@ func TestPromiseError(t *testing.T) { }) value, err := p.Wait(context.Background()) - assert.ErrorIs(t, err, expectedErr) + require.ErrorIs(t, err, expectedErr) assert.Equal(t, 0, value) } @@ -40,11 +41,12 @@ func TestPromiseDelayedResult(t *testing.T) { p := NewPromise(func() (string, error) { time.Sleep(50 * time.Millisecond) + return "delayed", nil }) value, err := p.Wait(context.Background()) - assert.NoError(t, err) + require.NoError(t, err) assert.Equal(t, "delayed", value) } @@ -53,6 +55,7 @@ func TestPromiseContextCancelled(t *testing.T) { p := NewPromise(func() (int, error) { time.Sleep(1 * time.Second) + return 42, nil }) @@ -68,6 +71,7 @@ func TestPromiseMultipleWaiters(t *testing.T) { p := NewPromise(func() (int, error) { time.Sleep(50 * time.Millisecond) + return 42, nil }) @@ -75,7 +79,7 @@ func TestPromiseMultipleWaiters(t *testing.T) { results := make([]int, 5) errs := make([]error, 5) - for i := 0; i < 5; i++ { + for i := range 5 { wg.Add(1) go func(idx int) { defer wg.Done() @@ -85,8 +89,8 @@ func TestPromiseMultipleWaiters(t *testing.T) { wg.Wait() - for i := 0; i < 5; i++ { - assert.NoError(t, errs[i]) + for i := range 5 { + require.NoError(t, errs[i]) assert.Equal(t, 42, results[i]) } } @@ -96,6 +100,7 @@ func TestPromiseDone(t *testing.T) { p := NewPromise(func() (int, error) { time.Sleep(50 * time.Millisecond) + return 42, nil }) @@ -109,7 +114,7 @@ func TestPromiseDone(t *testing.T) { <-p.Done() value, err := p.Result() - assert.NoError(t, err) + require.NoError(t, err) assert.Equal(t, 42, value) } @@ -118,6 +123,7 @@ func TestPromiseResultBeforeResolve(t *testing.T) { p := NewPromise(func() (int, error) { time.Sleep(100 * time.Millisecond) + return 42, nil }) @@ -135,11 +141,11 @@ func TestPromiseResultAfterResolve(t *testing.T) { <-p.Done() value, err := p.Result() - assert.NoError(t, err) + require.NoError(t, err) assert.Equal(t, 42, value) // Multiple calls should return the same result value2, err2 := p.Result() - assert.NoError(t, err2) + require.NoError(t, err2) assert.Equal(t, 42, value2) }