diff --git a/guest_linux.go b/guest_linux.go index e69c740..2a539a4 100644 --- a/guest_linux.go +++ b/guest_linux.go @@ -4,6 +4,7 @@ package sandbox import ( "context" + "errors" "fmt" "os" "runtime" @@ -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() diff --git a/host_linux.go b/host_linux.go index bf07989..a7433eb 100644 --- a/host_linux.go +++ b/host_linux.go @@ -12,6 +12,7 @@ import "C" import ( "context" "encoding/json" + "errors" "fmt" "os" "path/filepath" @@ -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") } @@ -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) } diff --git a/image_linux.go b/image_linux.go index 5e57377..c256e27 100644 --- a/image_linux.go +++ b/image_linux.go @@ -3,9 +3,9 @@ package sandbox import ( + "bufio" "context" "encoding/binary" - "errors" "fmt" "io" "math" @@ -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 } @@ -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) + 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. + 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 + } + 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 } diff --git a/image_linux_test.go b/image_linux_test.go index fb09cf6..dc81b81 100644 --- a/image_linux_test.go +++ b/image_linux_test.go @@ -3,6 +3,7 @@ package sandbox import ( + "bytes" "context" "encoding/binary" "errors" @@ -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) } @@ -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") } @@ -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) } }) @@ -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 { @@ -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) } } @@ -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) + } +} diff --git a/internal/state/memory.go b/internal/state/memory.go index 4716ff5..15ae39f 100644 --- a/internal/state/memory.go +++ b/internal/state/memory.go @@ -5,11 +5,11 @@ package state import "io" -// writer fills caller-owned memory. Exhaustion returns an error without -// growing the buffer, so a shared mapping remains the destination throughout. +// writer encodes into caller-owned memory or an output stream. type writer struct { mem []byte pos int + out io.Writer } func (w *writer) put(obj object) error { @@ -17,6 +17,17 @@ func (w *writer) put(obj object) error { } func (w *writer) writeBytes(p []byte) { + if w.out != nil { + n, err := w.out.Write(p) + w.pos += n + if err != nil { + panic(err) + } + if n != len(p) { + panic(io.ErrShortWrite) + } + return + } if len(p) > len(w.mem)-w.pos { panic(io.ErrShortBuffer) } @@ -24,6 +35,17 @@ func (w *writer) writeBytes(p []byte) { } func (w *writer) writeString(s string) { + if w.out != nil { + n, err := io.WriteString(w.out, s) + w.pos += n + if err != nil { + panic(err) + } + if n != len(s) { + panic(io.ErrShortWrite) + } + return + } if len(s) > len(w.mem)-w.pos { panic(io.ErrShortBuffer) } diff --git a/internal/state/memory_test.go b/internal/state/memory_test.go index a69e98b..32579c0 100644 --- a/internal/state/memory_test.go +++ b/internal/state/memory_test.go @@ -65,6 +65,14 @@ func TestMemoryObjects(t *testing.T) { if backing[0] != 0xa5 || backing[len(backing)-1] != 0xa5 || &w.mem[0] != &backing[1] { t.Fatal("writer changed its backing memory or wrote beyond it") } + var stream bytes.Buffer + streamed := writer{out: &stream} + if err := streamed.put(obj); err != nil { + t.Fatal(err) + } + if streamed.pos != w.pos || !bytes.Equal(stream.Bytes(), w.mem[:w.pos]) { + t.Fatal("stream output differs from memory output") + } r := reader{mem: w.mem[:w.pos]} got, err := r.get() if err != nil || r.pos != w.pos { diff --git a/internal/state/roundtrip.go b/internal/state/roundtrip.go index 82e10d6..e169671 100644 --- a/internal/state/roundtrip.go +++ b/internal/state/roundtrip.go @@ -2,6 +2,7 @@ package state import ( "context" + "io" "reflect" ) @@ -16,11 +17,22 @@ type State struct { // Save writes the graph, preserving IDs from the preceding Load when present. func (s *State) Save(ctx context.Context, mem []byte, rootPtr any) (int, Stats, error) { + return s.save(ctx, mem, nil, rootPtr) +} + +// SaveTo writes the same graph records as Save to an output stream. +// The caller owns flushing and closing the stream. +func (s *State) SaveTo(ctx context.Context, out io.Writer, rootPtr any) (int, Stats, error) { + return s.save(ctx, nil, out, rootPtr) +} + +func (s *State) save(ctx context.Context, mem []byte, out io.Writer, rootPtr any) (int, Stats, error) { es := newEncodeState(ctx, mem) err := safely(func() { if s.loaded != nil { es = s.loaded.encoder(ctx, mem) } + es.w.out = out es.Save(reflect.ValueOf(rootPtr).Elem()) }) if err == nil {