diff --git a/packages/orchestrator/internal/sandbox/uffd/handler.go b/packages/orchestrator/internal/sandbox/uffd/handler.go index 947ff168cf..0f9c98f995 100644 --- a/packages/orchestrator/internal/sandbox/uffd/handler.go +++ b/packages/orchestrator/internal/sandbox/uffd/handler.go @@ -7,6 +7,7 @@ import ( "fmt" "net" "os" + "sync" "syscall" "time" @@ -152,12 +153,15 @@ func (u *Uffd) handle(ctx context.Context, sandboxId string) error { u.readyCh <- struct{}{} + missingRequests := &sync.Map{} + err = Serve( ctx, uffd, m, u.memfile, u.fdExit, + missingRequests, zap.L().With(logger.WithSandboxID(sandboxId)), ) if err != nil { diff --git a/packages/orchestrator/internal/sandbox/uffd/helpers_test.go b/packages/orchestrator/internal/sandbox/uffd/helpers_test.go new file mode 100644 index 0000000000..ff1fc6d093 --- /dev/null +++ b/packages/orchestrator/internal/sandbox/uffd/helpers_test.go @@ -0,0 +1,204 @@ +package uffd + +import ( + "bytes" + "context" + "fmt" + "slices" + "sync" + "syscall" + "testing" + + "github.com/bits-and-blooms/bitset" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/e2b-dev/infra/packages/orchestrator/internal/sandbox/uffd/fdexit" + "github.com/e2b-dev/infra/packages/orchestrator/internal/sandbox/uffd/mapping" + "github.com/e2b-dev/infra/packages/orchestrator/internal/sandbox/uffd/testutils" + "github.com/e2b-dev/infra/packages/orchestrator/internal/sandbox/uffd/userfaultfd" +) + +type testConfig struct { + name string + // Page size of the memory area. + pagesize uint64 + // Number of pages in the memory area. + numberOfPages uint64 + // Operations to trigger on the memory area. + operations []operation +} + +type operationMode uint32 + +const ( + operationModeRead operationMode = 1 << iota + operationModeWrite +) + +type operation struct { + // Offset in bytes. Must be smaller than the (numberOfPages-1) * pagesize as it reads a page and it must be aligned to the pagesize from the testConfig. + offset int64 + mode operationMode +} + +type testHandler struct { + memoryArea *[]byte + pagesize uint64 + data *testutils.MemorySlicer + memoryMap mapping.Mappings + uffd uintptr + missingRequests *sync.Map + writeMu sync.Mutex +} + +func configureTest(t *testing.T, tt testConfig) (*testHandler, func()) { + t.Helper() + + cleanupList := []func(){} + + cleanup := func() { + slices.Reverse(cleanupList) + + for _, cleanup := range cleanupList { + cleanup() + } + } + + data := testutils.RandomPages(tt.pagesize, tt.numberOfPages) + + size, err := data.Size() + require.NoError(t, err) + + memoryArea, memoryStart, unmap, err := testutils.NewPageMmap(uint64(size), tt.pagesize) + require.NoError(t, err) + + cleanupList = append(cleanupList, func() { + unmap() + }) + + m := mapping.FcMappings([]mapping.GuestRegionUffdMapping{ + { + BaseHostVirtAddr: memoryStart, + Size: uintptr(size), + Offset: uintptr(0), + PageSize: uintptr(tt.pagesize), + }, + }) + + logger := testutils.NewTestLogger(t) + + fdExit, err := fdexit.New() + require.NoError(t, err) + + cleanupList = append(cleanupList, func() { + fdExit.Close() + }) + + uffd, err := userfaultfd.NewUserfaultfd(syscall.O_CLOEXEC | syscall.O_NONBLOCK) + require.NoError(t, err) + + cleanupList = append(cleanupList, func() { + userfaultfd.Close(uffd) + }) + + err = userfaultfd.ConfigureApi(uffd, tt.pagesize) + require.NoError(t, err) + + err = userfaultfd.Register(uffd, memoryStart, uint64(size), userfaultfd.UFFDIO_REGISTER_MODE_MISSING) + require.NoError(t, err) + + exitUffd := make(chan struct{}, 1) + + missingRequests := &sync.Map{} + + go func() { + err := Serve(t.Context(), int(uffd), m, data, fdExit, missingRequests, logger) + assert.NoError(t, err) + + exitUffd <- struct{}{} + }() + + cleanupList = append(cleanupList, func() { + signalExitErr := fdExit.SignalExit() + assert.NoError(t, signalExitErr) + + <-exitUffd + }) + + return &testHandler{ + memoryArea: &memoryArea, + memoryMap: m, + pagesize: tt.pagesize, + data: data, + uffd: uffd, + missingRequests: missingRequests, + }, cleanup +} + +func (h *testHandler) getAccessedOffsets() []uint { + offsets := []uint{} + + h.missingRequests.Range(func(key, _ any) bool { + offsets = append(offsets, uint(key.(int64))) + + return true + }) + + return offsets +} + +//go:noinline +func touchRead(b []byte) { + var dst [1]byte + _ = copy(dst[:], b[:1]) // forces a real read → MISSING fault +} + +func (h *testHandler) executeRead(ctx context.Context, op operation) error { + readBytes := (*h.memoryArea)[op.offset : op.offset+int64(h.pagesize)] + touchRead(readBytes) + + expectedBytes, err := h.data.Slice(ctx, op.offset, int64(h.pagesize)) + if err != nil { + return err + } + + if !bytes.Equal(readBytes, expectedBytes) { + idx, want, got := testutils.FirstDifferentByte(readBytes, expectedBytes) + + return fmt.Errorf("content mismatch: want '%x, got %x at index %d", want, got, idx) + } + + return nil +} + +func (h *testHandler) executeWrite(ctx context.Context, op operation) error { + bytesToWrite, err := h.data.Slice(ctx, op.offset, int64(h.pagesize)) + if err != nil { + return err + } + + // An unprotected parallel write to map results in undefined behavior—here usually manifesting as total freeze of the test. + h.writeMu.Lock() + defer h.writeMu.Unlock() + + n := copy((*h.memoryArea)[op.offset:op.offset+int64(h.pagesize)], bytesToWrite) + if n != int(h.pagesize) { + return fmt.Errorf("copy length mismatch: want %d, got %d", h.pagesize, n) + } + + return nil +} + +// Get a bitset of the offsets of the operations for the given mode. +func getOperationsOffsets(ops []operation, m operationMode) []uint { + b := bitset.New(0) + + for _, operation := range ops { + if operation.mode&m != 0 { + b.Set(uint(operation.offset)) + } + } + + return slices.Collect(b.EachSet()) +} diff --git a/packages/orchestrator/internal/sandbox/uffd/missing_test.go b/packages/orchestrator/internal/sandbox/uffd/missing_test.go new file mode 100644 index 0000000000..8c9a40d417 --- /dev/null +++ b/packages/orchestrator/internal/sandbox/uffd/missing_test.go @@ -0,0 +1,221 @@ +package uffd + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/sync/errgroup" + + "github.com/e2b-dev/infra/packages/shared/pkg/storage/header" +) + +func TestMissing(t *testing.T) { + tests := []testConfig{ + { + name: "standard 4k page, operation at start", + pagesize: header.PageSize, + numberOfPages: 32, + operations: []operation{ + { + offset: 0, + mode: operationModeRead, + }, + }, + }, + { + name: "standard 4k page, operation at middle", + pagesize: header.PageSize, + numberOfPages: 32, + operations: []operation{ + { + offset: 15 * header.PageSize, + mode: operationModeRead, + }, + }, + }, + { + name: "standard 4k page, operation at last page", + pagesize: header.PageSize, + numberOfPages: 32, + operations: []operation{ + { + offset: 31 * header.PageSize, + mode: operationModeRead, + }, + }, + }, + { + name: "standard 4k page, read after read", + pagesize: header.PageSize, + numberOfPages: 32, + operations: []operation{ + { + offset: 0, + mode: operationModeRead, + }, + { + offset: 0, + mode: operationModeRead, + }, + }, + }, + { + name: "hugepage, operation at start", + pagesize: header.HugepageSize, + numberOfPages: 8, + operations: []operation{ + { + offset: 0, + mode: operationModeRead, + }, + }, + }, + { + name: "hugepage, operation at middle", + pagesize: header.HugepageSize, + numberOfPages: 8, + operations: []operation{ + { + offset: 3 * header.HugepageSize, + mode: operationModeRead, + }, + }, + }, + { + name: "hugepage, operation at last page", + pagesize: header.HugepageSize, + numberOfPages: 8, + operations: []operation{ + { + offset: 7 * header.HugepageSize, + mode: operationModeRead, + }, + }, + }, + { + name: "hugepage, read after read", + pagesize: header.HugepageSize, + numberOfPages: 32, + operations: []operation{ + { + offset: 0, + mode: operationModeRead, + }, + { + offset: 0, + mode: operationModeRead, + }, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + h, cleanupFunc := configureTest(t, tt) + defer cleanupFunc() + + for _, operation := range tt.operations { + if operation.mode == operationModeRead { + err := h.executeRead(t.Context(), operation) + require.NoError(t, err) + } + } + + expectedAccessedOffsets := getOperationsOffsets(tt.operations, operationModeRead|operationModeWrite) + assert.Equal(t, expectedAccessedOffsets, h.getAccessedOffsets(), "checking which pages were faulted") + }) + } +} + +func TestParallelMissing(t *testing.T) { + t.Skipf("skipping for now because it freezes in debug mode") + + parallelOperations := 10_000_000 + + tt := testConfig{ + pagesize: header.PageSize, + numberOfPages: 2, + } + + h, cleanup := configureTest(t, tt) + t.Cleanup(cleanup) + + readOp := operation{ + offset: 0, + mode: operationModeRead, + } + + var verr errgroup.Group + + for range parallelOperations { + verr.Go(func() error { + return h.executeRead(t.Context(), readOp) + }) + } + + err := verr.Wait() + require.NoError(t, err) + + expectedAccessedOffsets := getOperationsOffsets([]operation{readOp}, operationModeRead) + assert.Equal(t, expectedAccessedOffsets, h.getAccessedOffsets(), "checking which pages were faulted") +} + +func TestParallelMissingWithPrefault(t *testing.T) { + parallelOperations := 10 + + tt := testConfig{ + pagesize: header.PageSize, + numberOfPages: 2, + } + + h, cleanup := configureTest(t, tt) + t.Cleanup(cleanup) + + readOp := operation{ + offset: 0, + mode: operationModeRead, + } + + err := h.executeRead(t.Context(), readOp) + require.NoError(t, err) + + var verr errgroup.Group + + for range parallelOperations { + verr.Go(func() error { + return h.executeRead(t.Context(), readOp) + }) + } + + err = verr.Wait() + require.NoError(t, err) + + expectedAccessedOffsets := getOperationsOffsets([]operation{readOp}, operationModeRead) + assert.Equal(t, expectedAccessedOffsets, h.getAccessedOffsets(), "checking which pages were faulted") +} + +func TestSerialMissing(t *testing.T) { + serialOperations := 1_000_000 + + tt := testConfig{ + pagesize: header.PageSize, + numberOfPages: 2, + } + + h, cleanup := configureTest(t, tt) + t.Cleanup(cleanup) + + readOp := operation{ + offset: 0, + mode: operationModeRead, + } + + for range serialOperations { + err := h.executeRead(t.Context(), readOp) + require.NoError(t, err) + } + + expectedAccessedOffsets := getOperationsOffsets([]operation{readOp}, operationModeRead) + assert.Equal(t, expectedAccessedOffsets, h.getAccessedOffsets(), "checking which pages were faulted") +} diff --git a/packages/orchestrator/internal/sandbox/uffd/missing_write_test.go b/packages/orchestrator/internal/sandbox/uffd/missing_write_test.go new file mode 100644 index 0000000000..f05337fb29 --- /dev/null +++ b/packages/orchestrator/internal/sandbox/uffd/missing_write_test.go @@ -0,0 +1,226 @@ +package uffd + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/sync/errgroup" + + "github.com/e2b-dev/infra/packages/shared/pkg/storage/header" +) + +func TestMissingWrite(t *testing.T) { + tests := []testConfig{ + { + name: "standard 4k page, operation at start", + pagesize: header.PageSize, + numberOfPages: 32, + operations: []operation{ + { + offset: 0, + mode: operationModeWrite, + }, + }, + }, + { + name: "standard 4k page, operation at middle", + pagesize: header.PageSize, + numberOfPages: 32, + operations: []operation{ + { + offset: 15 * header.PageSize, + mode: operationModeWrite, + }, + }, + }, + { + name: "standard 4k page, operation at last page", + pagesize: header.PageSize, + numberOfPages: 32, + operations: []operation{ + { + offset: 31 * header.PageSize, + mode: operationModeWrite, + }, + }, + }, + { + name: "standard 4k page, read after write", + pagesize: header.PageSize, + numberOfPages: 32, + operations: []operation{ + { + offset: 0, + mode: operationModeWrite, + }, + { + offset: 0, + mode: operationModeRead, + }, + }, + }, + { + name: "hugepage, operation at start", + pagesize: header.HugepageSize, + numberOfPages: 8, + operations: []operation{ + { + offset: 0, + mode: operationModeWrite, + }, + }, + }, + { + name: "hugepage, operation at middle", + pagesize: header.HugepageSize, + numberOfPages: 8, + operations: []operation{ + { + offset: 3 * header.HugepageSize, + mode: operationModeWrite, + }, + }, + }, + { + name: "hugepage, operation at last page", + pagesize: header.HugepageSize, + numberOfPages: 8, + operations: []operation{ + { + offset: 7 * header.HugepageSize, + mode: operationModeWrite, + }, + }, + }, + { + name: "hugepage, read after write", + pagesize: header.HugepageSize, + numberOfPages: 32, + operations: []operation{ + { + offset: 0, + mode: operationModeWrite, + }, + { + offset: 0, + mode: operationModeRead, + }, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + h, cleanupFunc := configureTest(t, tt) + defer cleanupFunc() + + for _, operation := range tt.operations { + if operation.mode == operationModeRead { + err := h.executeRead(t.Context(), operation) + require.NoError(t, err) + } + + if operation.mode == operationModeWrite { + err := h.executeWrite(t.Context(), operation) + require.NoError(t, err) + } + } + + expectedAccessedOffsets := getOperationsOffsets(tt.operations, operationModeRead|operationModeWrite) + assert.Equal(t, expectedAccessedOffsets, h.getAccessedOffsets(), "checking which pages were faulted") + }) + } +} + +func TestParallelMissingWrite(t *testing.T) { + t.Skipf("skipping for now because it freezes in debug mode") + + parallelOperations := 10_000_000 + + tt := testConfig{ + pagesize: header.PageSize, + numberOfPages: 2, + } + + h, cleanup := configureTest(t, tt) + t.Cleanup(cleanup) + + writeOp := operation{ + offset: 0, + mode: operationModeWrite, + } + + var verr errgroup.Group + + for range parallelOperations { + verr.Go(func() error { + return h.executeWrite(t.Context(), writeOp) + }) + } + + err := verr.Wait() + require.NoError(t, err) + + expectedAccessedOffsets := getOperationsOffsets([]operation{writeOp}, operationModeRead|operationModeWrite) + assert.Equal(t, expectedAccessedOffsets, h.getAccessedOffsets(), "checking which pages were faulted") +} + +func TestParallelMissingWriteWithPrefault(t *testing.T) { + parallelOperations := 1_000_000 + + tt := testConfig{ + pagesize: header.PageSize, + numberOfPages: 2, + } + + h, cleanup := configureTest(t, tt) + t.Cleanup(cleanup) + + writeOp := operation{ + offset: 0, + mode: operationModeWrite, + } + + err := h.executeWrite(t.Context(), writeOp) + require.NoError(t, err) + + var verr errgroup.Group + + for range parallelOperations { + verr.Go(func() error { + return h.executeWrite(t.Context(), writeOp) + }) + } + + err = verr.Wait() + require.NoError(t, err) + + expectedAccessedOffsets := getOperationsOffsets([]operation{writeOp}, operationModeRead|operationModeWrite) + assert.Equal(t, expectedAccessedOffsets, h.getAccessedOffsets(), "checking which pages were faulted") +} + +func TestSerialMissingWrite(t *testing.T) { + serialOperations := 1_000_000 + + tt := testConfig{ + pagesize: header.PageSize, + numberOfPages: 2, + } + + h, cleanup := configureTest(t, tt) + t.Cleanup(cleanup) + + writeOp := operation{ + offset: 0, + mode: operationModeWrite, + } + + for range serialOperations { + err := h.executeWrite(t.Context(), writeOp) + require.NoError(t, err) + } + + expectedAccessedOffsets := getOperationsOffsets([]operation{writeOp}, operationModeRead|operationModeWrite) + assert.Equal(t, expectedAccessedOffsets, h.getAccessedOffsets(), "checking which pages were faulted") +} diff --git a/packages/orchestrator/internal/sandbox/uffd/serve.go b/packages/orchestrator/internal/sandbox/uffd/serve.go index f493087220..e83734eca5 100644 --- a/packages/orchestrator/internal/sandbox/uffd/serve.go +++ b/packages/orchestrator/internal/sandbox/uffd/serve.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "sync" "syscall" "unsafe" @@ -32,6 +33,7 @@ func Serve( mappings mapping.Mappings, src block.Slicer, fdExit *fdexit.FdExit, + missingRequests *sync.Map, logger *zap.Logger, ) error { pollFds := []unix.PollFd{ @@ -41,8 +43,6 @@ func Serve( var eg errgroup.Group - missingPagesBeingHandled := map[int64]struct{}{} - eagainCounter := newEagainCounter(logger, "uffd: eagain during fd read (accumulated)") defer eagainCounter.Close() @@ -145,11 +145,11 @@ outerLoop: return fmt.Errorf("failed to map: %w", err) } - if _, ok := missingPagesBeingHandled[offset]; ok { + if _, ok := missingRequests.Load(offset); ok { continue } - missingPagesBeingHandled[offset] = struct{}{} + missingRequests.Store(offset, struct{}{}) eg.Go(func() error { defer func() { diff --git a/packages/orchestrator/internal/sandbox/uffd/testutils/diff_byte.go b/packages/orchestrator/internal/sandbox/uffd/testutils/diff_byte.go new file mode 100644 index 0000000000..68298ea6ea --- /dev/null +++ b/packages/orchestrator/internal/sandbox/uffd/testutils/diff_byte.go @@ -0,0 +1,20 @@ +package testutils + +// FirstDifferentByte returns the first byte index where a and b differ. +// It also returns the differing byte values (want, got). +// If slices are identical, it returns idx -1. +func FirstDifferentByte(a, b []byte) (idx int, want, got byte) { + smallerSize := min(len(a), len(b)) + + for i := range smallerSize { + if a[i] != b[i] { + return i, b[i], a[i] + } + } + + if len(a) != len(b) { + return smallerSize, 0, 0 + } + + return -1, 0, 0 +} diff --git a/packages/orchestrator/internal/sandbox/uffd/testutils/logger.go b/packages/orchestrator/internal/sandbox/uffd/testutils/logger.go new file mode 100644 index 0000000000..fa197edbd9 --- /dev/null +++ b/packages/orchestrator/internal/sandbox/uffd/testutils/logger.go @@ -0,0 +1,43 @@ +package testutils + +import ( + "testing" + + "go.uber.org/zap" + "go.uber.org/zap/zapcore" +) + +type testWriter struct { + t *testing.T +} + +func (w *testWriter) Write(p []byte) (n int, err error) { + w.t.Log(string(p)) + + return len(p), nil +} + +// NewTestLogger creates a new zap logger that logs all zap logs to the test output. +func NewTestLogger(t *testing.T) *zap.Logger { + t.Helper() + + encoderCfg := zap.NewDevelopmentEncoderConfig() + encoderCfg.EncodeLevel = zapcore.CapitalColorLevelEncoder + encoderCfg.CallerKey = zapcore.OmitKey + encoderCfg.ConsoleSeparator = " " + encoderCfg.TimeKey = "" + encoderCfg.MessageKey = "message" + encoderCfg.LevelKey = "level" + encoderCfg.NameKey = "logger" + encoderCfg.StacktraceKey = "stacktrace" + encoderCfg.EncodeTime = zapcore.RFC3339NanoTimeEncoder + encoderCfg.EncodeCaller = zapcore.ShortCallerEncoder + encoderCfg.EncodeDuration = zapcore.StringDurationEncoder + + encoder := zapcore.NewConsoleEncoder(encoderCfg) + + testSyncer := zapcore.AddSync(&testWriter{t}) + core := zapcore.NewCore(encoder, testSyncer, zap.DebugLevel) + + return zap.New(core, zap.AddCaller()) +} diff --git a/packages/orchestrator/internal/sandbox/uffd/testutils/memory_slicer.go b/packages/orchestrator/internal/sandbox/uffd/testutils/memory_slicer.go new file mode 100644 index 0000000000..3c67d4f05e --- /dev/null +++ b/packages/orchestrator/internal/sandbox/uffd/testutils/memory_slicer.go @@ -0,0 +1,35 @@ +package testutils + +import ( + "context" + + "github.com/e2b-dev/infra/packages/orchestrator/internal/sandbox/block" +) + +// MemorySlicer exposes byte slice via the Slicer interface. +// This is used for testing purposes. +type MemorySlicer struct { + content []byte + pagesize int64 +} + +var _ block.Slicer = (*MemorySlicer)(nil) + +func newMemorySlicer(content []byte, pagesize int64) *MemorySlicer { + return &MemorySlicer{ + content: content, + pagesize: pagesize, + } +} + +func (s *MemorySlicer) Slice(_ context.Context, offset, size int64) ([]byte, error) { + return s.content[offset : offset+size], nil +} + +func (s *MemorySlicer) Size() (int64, error) { + return int64(len(s.content)), nil +} + +func (s *MemorySlicer) Content() []byte { + return s.content +} diff --git a/packages/orchestrator/internal/sandbox/uffd/testutils/page_mmap.go b/packages/orchestrator/internal/sandbox/uffd/testutils/page_mmap.go new file mode 100644 index 0000000000..e85a234ec5 --- /dev/null +++ b/packages/orchestrator/internal/sandbox/uffd/testutils/page_mmap.go @@ -0,0 +1,47 @@ +package testutils + +import ( + "fmt" + "math" + "syscall" + "unsafe" + + "golang.org/x/sys/unix" + + "github.com/e2b-dev/infra/packages/shared/pkg/storage/header" +) + +func NewPageMmap(size, pagesize uint64) ([]byte, uintptr, func() error, error) { + if pagesize == header.PageSize { + return newMmap(size, header.PageSize, 0) + } + + if pagesize == header.HugepageSize { + return newMmap(size, header.HugepageSize, unix.MAP_HUGETLB|unix.MAP_HUGE_2MB) + } + + return nil, 0, nil, fmt.Errorf("unsupported page size: %d", pagesize) +} + +// Even though UFFD behaves differently with file backend memory (and hugetlbfs file backed), the FC uses MAP_PRIVATE|MAP_ANONYMOUS, so the following stub is correct to test for FC. +// - https://docs.kernel.org/admin-guide/mm/userfaultfd.html#write-protect-notifications +// - https://github.com/firecracker-microvm/firecracker/blob/a305f362d0e6f7ba926c73e65452cb51262a44d8/src/vmm/src/persist.rs#L499 +func newMmap(size, pagesize uint64, flags int) ([]byte, uintptr, func() error, error) { + l := int(math.Ceil(float64(size)/float64(pagesize)) * float64(pagesize)) + b, err := syscall.Mmap( + -1, + 0, + l, + syscall.PROT_READ|syscall.PROT_WRITE, + syscall.MAP_PRIVATE|syscall.MAP_ANONYMOUS|flags, + ) + if err != nil { + return nil, 0, nil, fmt.Errorf("failed to mmap: %w", err) + } + + closeMmap := func() error { + return syscall.Munmap(b) + } + + return b, uintptr(unsafe.Pointer(&b[0])), closeMmap, nil +} diff --git a/packages/orchestrator/internal/sandbox/uffd/testutils/random_data.go b/packages/orchestrator/internal/sandbox/uffd/testutils/random_data.go new file mode 100644 index 0000000000..c8bfe8ed1e --- /dev/null +++ b/packages/orchestrator/internal/sandbox/uffd/testutils/random_data.go @@ -0,0 +1,17 @@ +package testutils + +import ( + "crypto/rand" +) + +func RandomPages(pagesize, numberOfPages uint64) *MemorySlicer { + size := pagesize * numberOfPages + + n := int(size) + buf := make([]byte, n) + if _, err := rand.Read(buf); err != nil { + panic(err) + } + + return newMemorySlicer(buf, int64(pagesize)) +} diff --git a/packages/orchestrator/internal/sandbox/uffd/userfaultfd/syscalls.go b/packages/orchestrator/internal/sandbox/uffd/userfaultfd/syscalls.go new file mode 100644 index 0000000000..8359acefc3 --- /dev/null +++ b/packages/orchestrator/internal/sandbox/uffd/userfaultfd/syscalls.go @@ -0,0 +1,73 @@ +package userfaultfd + +import ( + "fmt" + "syscall" + "unsafe" + + "github.com/e2b-dev/infra/packages/shared/pkg/storage/header" +) + +// flags: syscall.O_CLOEXEC|syscall.O_NONBLOCK +func NewUserfaultfd(flags uintptr) (uintptr, error) { + uffd, _, errno := syscall.Syscall(NR_userfaultfd, flags, 0, 0) + if errno != 0 { + return 0, fmt.Errorf("userfaultfd syscall failed: %w", errno) + } + + return uffd, nil +} + +// features: UFFD_FEATURE_MISSING_HUGETLBFS +// This is already called by the FC +func ConfigureApi(fd uintptr, pagesize uint64) error { + var features CULong + + // Only set the hugepage feature if we're using hugepages + if pagesize == header.HugepageSize { + features |= UFFD_FEATURE_MISSING_HUGETLBFS + } + + api := NewUffdioAPI(UFFD_API, features) + ret, _, errno := syscall.Syscall(syscall.SYS_IOCTL, fd, UFFDIO_API, uintptr(unsafe.Pointer(&api))) + if errno != 0 { + return fmt.Errorf("UFFDIO_API ioctl failed: %w (ret=%d)", errno, ret) + } + + return nil +} + +// mode: UFFDIO_REGISTER_MODE_WP|UFFDIO_REGISTER_MODE_MISSING +// This is already called by the FC, but only with the UFFDIO_REGISTER_MODE_MISSING +// We need to call it with UFFDIO_REGISTER_MODE_WP when we use both missing and wp +func Register(fd uintptr, addr uintptr, size uint64, mode CULong) error { + register := NewUffdioRegister(CULong(addr), CULong(size), mode) + + ret, _, errno := syscall.Syscall(syscall.SYS_IOCTL, fd, UFFDIO_REGISTER, uintptr(unsafe.Pointer(®ister))) + if errno != 0 { + return fmt.Errorf("UFFDIO_REGISTER ioctl failed: %w (ret=%d)", errno, ret) + } + + return nil +} + +// mode: UFFDIO_COPY_MODE_WP +// When we use both missing and wp, we need to use UFFDIO_COPY_MODE_WP, otherwise copying would unprotect the page +func Copy(fd uintptr, addr uintptr, data []byte, pagesize uint64, mode CULong) error { + cpy := NewUffdioCopy(data, CULong(addr)&^CULong(pagesize-1), CULong(pagesize), mode, 0) + + if _, _, errno := syscall.Syscall(syscall.SYS_IOCTL, fd, UFFDIO_COPY, uintptr(unsafe.Pointer(&cpy))); errno != 0 { + return errno + } + + // Check if the copied size matches the requested pagesize + if uint64(cpy.copy) != pagesize { + return fmt.Errorf("UFFDIO_COPY copied %d bytes, expected %d", cpy.copy, pagesize) + } + + return nil +} + +func Close(fd uintptr) error { + return syscall.Close(int(fd)) +}