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
4 changes: 3 additions & 1 deletion guest_linux.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ package sandbox

import (
"context"
"errors"
"fmt"
"os"
"runtime"
Expand All @@ -30,10 +31,11 @@ func runGuest() (err error) {
}
}()
defer unix.Close(3)
data, err := readStateImage(3, 0)
data, unmap, err := readStateImage(3, 0)
if err != nil {
return err
}
defer func() { err = errors.Join(err, unmap()) }()
var graph state.State
var fn func()
ctx := context.Background()
Expand Down
11 changes: 9 additions & 2 deletions host_linux.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ import "C"
import (
"context"
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
Expand Down Expand Up @@ -85,7 +86,7 @@ func sandboxInspect(owner C.uintptr_t, event *C.struct_syscall_event) {
// Captures must be exclusively owned for the duration of Run. Calls may overlap
// when their captured graphs are independent. Configuration must not change
// until all active calls return.
func (s *Sandbox) Run(fn func()) error {
func (s *Sandbox) Run(fn func()) (err error) {
if fn == nil {
return fmt.Errorf("sandbox: nil function")
}
Expand Down Expand Up @@ -153,10 +154,16 @@ func (s *Sandbox) Run(fn func()) error {
if inspectionErr != nil {
return inspectionErr
}
data, err := readStateImage(fd, resultOffset)
// Seal before either decode pass, including against any fd the guest
// transferred elsewhere. Existing writable mappings make this fail.
if _, err := unix.FcntlInt(uintptr(fd), unix.F_ADD_SEALS, unix.F_SEAL_WRITE|unix.F_SEAL_SEAL); err != nil {
return fmt.Errorf("sandbox result seal: %w", err)
}
data, unmap, err := readStateImage(fd, resultOffset)
if err != nil {
return fmt.Errorf("sandbox result: %w", err)
}
defer func() { err = errors.Join(err, unmap()) }()
if _, err := graph.Load(ctx, data, &fn); err != nil {
return fmt.Errorf("sandbox import: %w", err)
}
Expand Down
111 changes: 55 additions & 56 deletions image_linux.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,9 @@
package sandbox

import (
"bufio"
"context"
"encoding/binary"
"errors"
"fmt"
"io"
"math"
Expand All @@ -19,9 +19,9 @@ func newStateImage() (int, error) {
if err != nil {
return -1, err
}
// The result may grow, but neither participant may truncate a live mapping
// or change the seals after the descriptor is handed to the guest.
if _, err := unix.FcntlInt(uintptr(fd), unix.F_ADD_SEALS, unix.F_SEAL_SHRINK|unix.F_SEAL_SEAL); err != nil {
// The guest may append its result, but cannot truncate a live mapping.
// The host adds the final write and seal locks before reading that result.
if _, err := unix.FcntlInt(uintptr(fd), unix.F_ADD_SEALS, unix.F_SEAL_SHRINK); err != nil {
unix.Close(fd)
return -1, err
}
Expand All @@ -32,90 +32,89 @@ func newStateImage() (int, error) {
// The input and result occupy consecutive images in the same memfd. The host
// retains the result offset independently of the guest-writable input header.
func writeStateImage(fd int, offset int64, graph *state.State, root any) (int64, error) {
data := make([]byte, 1<<20)
var n int
for {
var err error
n, _, err = graph.Save(context.Background(), data, root)
if err == nil {
break
}
if !errors.Is(err, io.ErrShortBuffer) {
return 0, err
}
if len(data) > (math.MaxInt-16)/2 {
return 0, fmt.Errorf("sandbox image exceeds addressable memory")
}
data = make([]byte, 2*len(data))
}
// Reserve the following header as well. A guest that exits without writing
// its result leaves this zero, so the host cannot accept the input as output.
if offset < 0 || n > math.MaxInt-16 || offset > math.MaxInt64-int64(n)-16 {
if offset < 0 || offset > math.MaxInt64-16 {
return 0, fmt.Errorf("sandbox image size overflow")
}
size := int64(n) + 8
var stat unix.Stat_t
if err := unix.Fstat(fd, &stat); err != nil {
output := imageWriter{fd: fd, offset: offset}
var header [8]byte
if _, err := output.Write(header[:]); err != nil {
return 0, err
}
if end := offset + size + 8; end > stat.Size {
if err := unix.Ftruncate(fd, end); err != nil {
return 0, err
}
buffer := bufio.NewWriter(&output)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

bufio.NewWriter uses the default 4 KB buffer, so the payload flushes to pwrite every 4 KB. Since the encoder emits many tiny writes (per-object tag + varints), buffering is load-bearing here — but for large state images a bigger buffer (e.g. bufio.NewWriterSize(&output, 64<<10)) would materially cut the number of pwrite syscalls at negligible memory cost. Tuning only, not a correctness issue.

n, _, err := graph.SaveTo(context.Background(), buffer, root)
if err != nil {
return 0, err
}
mapOffset := offset - offset%int64(unix.Getpagesize())
delta := int(offset - mapOffset)
if n > math.MaxInt-delta-16 {
return 0, fmt.Errorf("sandbox mapping size overflow")
if err := buffer.Flush(); err != nil {
return 0, err
}
mem, err := unix.Mmap(fd, mapOffset, delta+n+16, unix.PROT_READ|unix.PROT_WRITE, unix.MAP_SHARED)
if err != nil {
// An unwritten result must keep its following header unpublished.

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

This comment reads as a prohibition ("must keep ... unpublished"), but the next line actively writes an all-zero header at the following image's offset. The mechanism is proactive zeroing so a reader walking to the next offset sees a zero-length (rejected) image. Consider rewording to describe the action, e.g. "Zero the following image's header so the next result reads as unpublished until it is written."

if _, err := output.Write(header[:]); err != nil {
return 0, err
}
image := mem[delta:]
binary.LittleEndian.PutUint64(image, 0)
copy(image[8:], data[:n])
binary.LittleEndian.PutUint64(image[8+n:], 0)
size := int64(n) + 8
// Publish the size only after the payload is complete. The reader must also
// wait for ownership to transfer; this header is not a synchronization lock.
binary.LittleEndian.PutUint64(image, uint64(size))
if err := unix.Munmap(mem); err != nil {
binary.LittleEndian.PutUint64(header[:], uint64(size))
output.offset = offset
if _, err := output.Write(header[:]); err != nil {
return 0, err
}
return size, nil
}

// The sender must have finished before reading. Copy before decoding, and keep
// no references to the shared mapping while restoring the host object graph.
func readStateImage(fd int, offset int64) ([]byte, error) {
// imageWriter keeps the input and result independent of the shared fd offset.
type imageWriter struct {
fd int
offset int64
}

func (w *imageWriter) Write(p []byte) (int, error) {
if w.offset < 0 || w.offset > math.MaxInt64-int64(len(p)) {
return 0, fmt.Errorf("sandbox image size overflow")
}
for {
n, err := unix.Pwrite(w.fd, p, w.offset)
if err == unix.EINTR {
continue
}
n = max(n, 0)
w.offset += int64(n)
if err == nil && n != len(p) {
err = io.ErrShortWrite

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

A short Pwrite (err == nil && n != len(p)) is turned into io.ErrShortWrite and returned immediately, after advancing w.offset by the partial n. For a memfd on tmpfs this effectively never happens, so it's safe in practice — but since imageWriter is a general io.Writer fed through bufio, the more robust choice is to loop on the remainder (p = p[n:]) and only give up when Pwrite returns n == 0, nil. As-is, the partial-write branch is essentially unreachable/untested and would advance the offset past what was persisted. A one-line comment noting the reliance on tmpfs semantics would help if you keep it.

}
return n, err
}
}

// The sender must have finished before reading. The host also seals returned
// images against writes. The caller must keep the mapping until its state
// round trip is complete, then release it using the returned function.
func readStateImage(fd int, offset int64) ([]byte, func() error, error) {
var header [8]byte
n, err := unix.Pread(fd, header[:], offset)
if err != nil {
return nil, err
return nil, nil, err
}
if n != len(header) {
return nil, io.ErrUnexpectedEOF
return nil, nil, io.ErrUnexpectedEOF
}
size := binary.LittleEndian.Uint64(header[:])
var stat unix.Stat_t
if err := unix.Fstat(fd, &stat); err != nil {
return nil, err
return nil, nil, err
}
if offset < 0 || offset > stat.Size || size < 8 || size > uint64(stat.Size-offset) {
return nil, fmt.Errorf("invalid or incomplete sandbox image length %d", size)
return nil, nil, fmt.Errorf("invalid or incomplete sandbox image length %d", size)
}
mapOffset := offset - offset%int64(unix.Getpagesize())
delta := int(offset - mapOffset)
if size > uint64(math.MaxInt-delta) {
return nil, fmt.Errorf("sandbox mapping size overflow")
return nil, nil, fmt.Errorf("sandbox mapping size overflow")
}
mem, err := unix.Mmap(fd, mapOffset, delta+int(size), unix.PROT_READ, unix.MAP_SHARED)
if err != nil {
return nil, err
}
data := append([]byte(nil), mem[delta+8:delta+int(size)]...)
if err := unix.Munmap(mem); err != nil {
return nil, err
return nil, nil, err
}
return data, nil
return mem[delta+8 : delta+int(size)], func() error { return unix.Munmap(mem) }, nil
}
109 changes: 99 additions & 10 deletions image_linux_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
package sandbox

import (
"bytes"
"context"
"encoding/binary"
"errors"
Expand Down Expand Up @@ -37,13 +38,20 @@ func TestStateImageRoundTrip(t *testing.T) {
if err != nil {
t.Fatal(err)
}
if _, err := readStateImage(fd, offset); err == nil {
if _, _, err := readStateImage(fd, offset); err == nil {
t.Fatal("unpublished output accepted")
}
data, err := readStateImage(fd, 0)
data, unmapInput, err := readStateImage(fd, 0)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
if unmapInput != nil {
if err := unmapInput(); err != nil {
t.Error(err)
}
}
})
if int64(len(data))+8 != offset {
t.Fatalf("input header: payload=%d offset=%d", len(data), offset)
}
Expand All @@ -59,10 +67,23 @@ func TestStateImageRoundTrip(t *testing.T) {
if err != nil {
t.Fatal(err)
}
result, err := readStateImage(fd, offset)
// Guest mappings are released before the host seals the result.
if err := unmapInput(); err != nil {
t.Fatal(err)
}
unmapInput = nil
// A changed input header must not redirect the host's result read.
if _, err := unix.Pwrite(fd, make([]byte, 8), 0); err != nil {
t.Fatal(err)
}
if _, err := unix.FcntlInt(uintptr(fd), unix.F_ADD_SEALS, unix.F_SEAL_WRITE|unix.F_SEAL_SEAL); err != nil {
t.Fatal(err)
}
result, unmapResult, err := readStateImage(fd, offset)
if err != nil {
t.Fatal(err)
}
defer unmapResult()
if int64(len(result))+8 != resultSize {
t.Fatal("result header does not describe its actual size")
}
Expand All @@ -73,17 +94,17 @@ func TestStateImageRoundTrip(t *testing.T) {
if stat.Size != offset+resultSize+8 {
t.Fatalf("memfd size=%d, want %d", stat.Size, offset+resultSize+8)
}
// A changed input header must not redirect the host's result read.
if _, err := unix.Pwrite(fd, make([]byte, 8), 0); err != nil {
t.Fatal(err)
}
if _, err := host.Load(context.Background(), result, &input); err != nil {
t.Fatal(err)
}
if input != output {
t.Fatal("result was not written back")
}
if _, err := readStateImage(fd, offset); err != nil {
_, unmapAgain, err := readStateImage(fd, offset)
if err != nil {
t.Fatal(err)
}
if err := unmapAgain(); err != nil {
t.Fatal(err)
}
})
Expand All @@ -96,7 +117,7 @@ func TestStateImageLengths(t *testing.T) {
t.Fatal(err)
}
defer unix.Close(fd)
if _, err := readStateImage(fd, 0); !errors.Is(err, io.ErrUnexpectedEOF) {
if _, _, err := readStateImage(fd, 0); !errors.Is(err, io.ErrUnexpectedEOF) {
t.Fatalf("empty file: %v", err)
}
if err := unix.Ftruncate(fd, 32); err != nil {
Expand All @@ -108,7 +129,7 @@ func TestStateImageLengths(t *testing.T) {
if _, err := unix.Pwrite(fd, header[:], 0); err != nil {
t.Fatal(err)
}
if _, err := readStateImage(fd, 0); err == nil {
if _, _, err := readStateImage(fd, 0); err == nil {
t.Fatalf("accepted invalid length %d", size)
}
}
Expand All @@ -119,3 +140,71 @@ func TestStateImageLengths(t *testing.T) {
t.Fatalf("image could not grow: %v", err)
}
}

func TestStateImageReadMapping(t *testing.T) {
fd, err := newStateImage()
if err != nil {
t.Fatal(err)
}
defer unix.Close(fd)
var graph state.State
value := "before"
if _, err := writeStateImage(fd, 0, &graph, &value); err != nil {
t.Fatal(err)
}
data, unmap, err := readStateImage(fd, 0)
if err != nil {
t.Fatal(err)
}
defer unmap()
index := bytes.Index(data, []byte(value))
if index < 0 {
t.Fatal("string missing from image")
}
// An unsealed input view must refer to the file, not a heap copy.
if _, err := unix.Pwrite(fd, []byte("after!"), int64(8+index)); err != nil {
t.Fatal(err)
}
if string(data[index:index+len(value)]) != "after!" {
t.Fatal("read image does not share the mapped file")
}
}

func TestStateImageWriteSeal(t *testing.T) {
fd, err := newStateImage()
if err != nil {
t.Fatal(err)
}
defer unix.Close(fd)
if err := unix.Ftruncate(fd, int64(unix.Getpagesize())); err != nil {
t.Fatal(err)
}
alias, err := unix.Dup(fd)
if err != nil {
t.Fatal(err)
}
defer unix.Close(alias)
writable, err := unix.Mmap(alias, 0, unix.Getpagesize(), unix.PROT_READ|unix.PROT_WRITE, unix.MAP_SHARED)
if err != nil {
t.Fatal(err)
}
_, sealErr := unix.FcntlInt(uintptr(fd), unix.F_ADD_SEALS, unix.F_SEAL_WRITE|unix.F_SEAL_SEAL)
if err := unix.Munmap(writable); err != nil {
t.Fatal(err)
}
if !errors.Is(sealErr, unix.EBUSY) {
t.Fatalf("writable mapping was not rejected: %v", sealErr)
}
if _, err := unix.FcntlInt(uintptr(fd), unix.F_ADD_SEALS, unix.F_SEAL_WRITE|unix.F_SEAL_SEAL); err != nil {
t.Fatal(err)
}
if _, err := unix.Pwrite(alias, []byte{1}, 0); !errors.Is(err, unix.EPERM) {
t.Fatalf("write through duplicate fd: %v", err)
}
if mem, err := unix.Mmap(alias, 0, unix.Getpagesize(), unix.PROT_READ|unix.PROT_WRITE, unix.MAP_SHARED); !errors.Is(err, unix.EPERM) {
if err == nil {
unix.Munmap(mem)
}
t.Fatalf("new writable mapping: %v", err)
}
}
Loading
Loading