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
5 changes: 1 addition & 4 deletions guest_linux.go
Original file line number Diff line number Diff line change
Expand Up @@ -22,9 +22,6 @@ func guestEntry() {
fmt.Fprintln(os.Stderr, err)
os.Exit(1)
}
// runGuest has released its state graph and input mapping. Collect after
// those roots leave the stack, before the guest process exits.
runtime.GC()
}

func runGuest() (err error) {
Expand All @@ -45,11 +42,11 @@ func runGuest() (err error) {
if _, err := graph.Load(ctx, data, &fn); err != nil {
return err
}
runtime.GC()
fn()
_, err = writeStateImage(3, int64(len(data))+8, &graph, &fn)
if err != nil {
return err
}
runtime.KeepAlive(&graph)
return nil
}
4 changes: 3 additions & 1 deletion host_linux.go
Original file line number Diff line number Diff line change
Expand Up @@ -111,7 +111,7 @@ func (s *Sandbox) Run(fn func()) (err error) {
if err != nil {
return fmt.Errorf("sandbox export: %w", err)
}
defer runtime.KeepAlive(&graph)
runtime.GC()
executable, err := os.Executable()
if err != nil {
return err
Expand Down Expand Up @@ -167,6 +167,8 @@ func (s *Sandbox) Run(fn func()) (err error) {
if _, err := graph.Load(ctx, data, &fn); err != nil {
return fmt.Errorf("sandbox import: %w", err)
}
graph = state.State{}
runtime.GC()
return nil
}

Expand Down
4 changes: 3 additions & 1 deletion internal/reflecttype/decode.go
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@ func Open(data []byte) (result *ReflectType, err error) {
entries: make([][]byte, n),
types: make([]reflect.Type, n),
resolving: make([]bool, n),
static: indexStaticTypes().byLocation,
}
for i := range d.entries {
length := r.count()
Expand All @@ -55,6 +56,7 @@ type importer struct {
entries [][]byte
types []reflect.Type
resolving []bool
static map[staticLocation]reflect.Type
}

func (d *importer) resolve(id uint64) reflect.Type {
Expand Down Expand Up @@ -82,7 +84,7 @@ func (d *importer) resolve(id uint64) reflect.Type {
if module > math.MaxUint32 {
panic(fmt.Errorf("invalid static type module %d", module))
}
typ = staticTypes().byLocation[staticLocation{uint32(module), offset}]
typ = d.static[staticLocation{uint32(module), offset}]
if typ == nil {
panic(fmt.Errorf("static type module=%d offset=%#x is unavailable", module, offset))
}
Expand Down
7 changes: 4 additions & 3 deletions internal/reflecttype/reflecttype.go
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@ func Export() (*Snapshot, error) {
if runtime.Version() != "go1.26.6" {
return nil, fmt.Errorf("reflecttype requires go1.26.6, got %s", runtime.Version())
}
e := exporter{ids: make(map[reflect.Type]uint32), supported: make(map[reflect.Type]bool)}
e := exporter{ids: make(map[reflect.Type]uint32), supported: make(map[reflect.Type]bool), static: indexStaticTypes().byType}
for _, typ := range cachedTypes() {
if !e.supports(typ) {
continue
Expand Down Expand Up @@ -71,13 +71,14 @@ type exporter struct {
ids map[reflect.Type]uint32
entries [][]byte
supported map[reflect.Type]bool
static map[reflect.Type]staticLocation
}

func (e *exporter) supports(typ reflect.Type) (supported bool) {
if builtinTypes[typ.Kind()] == typ {
return true
}
if _, ok := staticTypes().byType[typ]; ok {
if _, ok := e.static[typ]; ok {
return true
}
if supported, ok := e.supported[typ]; ok {
Expand Down Expand Up @@ -157,7 +158,7 @@ func (e *exporter) encode(typ reflect.Type) ([]byte, error) {
if builtinTypes[kind] == typ {
return binary.AppendUvarint(nil, uint64(kind)), nil
}
if location, ok := staticTypes().byType[typ]; ok {
if location, ok := e.static[typ]; ok {
// A static entry is restored by location, but callers can also refer
// to its dependencies directly (e.g. reflect.Type inside []reflect.Type).
for _, dependency := range appendDependencies(nil, typ) {
Expand Down
2 changes: 0 additions & 2 deletions internal/reflecttype/runtime.go
Original file line number Diff line number Diff line change
Expand Up @@ -19,8 +19,6 @@ type staticTypeIndex struct {
byLocation map[staticLocation]reflect.Type
}

var staticTypes = sync.OnceValue(indexStaticTypes)

//go:linkname reflectTypeLinks reflect.typelinks
func reflectTypeLinks() ([]unsafe.Pointer, [][]int32)

Expand Down
14 changes: 12 additions & 2 deletions internal/reflectxtype/decode.go
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@ func open(data []byte, previous *Snapshot) (result *ReflectType, err error) {
mocks: make([]reflect.Type, n),
mocking: make([]bool, n),
ctx: reflectx.NewContext(),
static: indexStaticTypes().byLocation,
}
entries := make([][]byte, n)
for i := range d.definitions {
Expand Down Expand Up @@ -132,7 +133,15 @@ func open(data []byte, previous *Snapshot) (result *ReflectType, err error) {
for i := range d.definitions {
d.resolve(uint32(i + 1))
}
return &ReflectType{types: d.types, entries: entries, definitions: d.definitions, ctx: d.ctx, methodCount: d.methodCount, retained: retained}, nil
// Type construction is complete. Only concrete method definitions are
// needed for installation and for preserving IDs on the return export.
methods := make([][]method, n)
for i, def := range d.definitions {
if def.kind != reflect.Interface {
methods[i] = def.methods
}
}
return &ReflectType{types: d.types, entries: entries, methods: methods, ctx: d.ctx, methodCount: d.methodCount, retained: retained}, nil
}

type definition struct {
Expand Down Expand Up @@ -172,6 +181,7 @@ type importer struct {
mocking []bool
ctx *reflectx.Context
methodCount int
static map[staticLocation]reflect.Type
}

func (d *importer) parse(index int, r *typeReader) {
Expand All @@ -181,7 +191,7 @@ func (d *importer) parse(index int, r *typeReader) {
if module > math.MaxUint32 {
panic(fmt.Errorf("invalid static type module %d", module))
}
d.types[index] = staticTypes().byLocation[staticLocation{uint32(module), offset}]
d.types[index] = d.static[staticLocation{uint32(module), offset}]
if d.types[index] == nil {
panic(fmt.Errorf("unavailable static type module=%d offset=%#x", module, offset))
}
Expand Down
14 changes: 7 additions & 7 deletions internal/reflectxtype/method.go
Original file line number Diff line number Diff line change
Expand Up @@ -133,12 +133,12 @@ func (t *ReflectType) SetMethods(callbacks []func([]reflect.Value) []reflect.Val
}
}()
installed := make(map[int]methodEntries)
for i, def := range t.definitions {
if def.kind == reflect.Interface || len(def.methods) == 0 {
for i, definitions := range t.methods {
if len(definitions) == 0 {
continue
}
ids := make(map[methodIdentity]int, len(def.methods))
for _, method := range def.methods {
ids := make(map[methodIdentity]int, len(definitions))
for _, method := range definitions {
ids[methodIdentity{method.name, method.pkg, method.pointer}] = method.function
}
if i < t.retained {
Expand All @@ -149,17 +149,17 @@ func (t *ReflectType) SetMethods(callbacks []func([]reflect.Value) []reflect.Val
}
continue
}
methods := make([]reflectx.Method, len(def.methods))
methods := make([]reflectx.Method, len(definitions))
t.ctx.SetHasImethod(func(_ reflect.Type, m reflectx.Method) bool {
for _, method := range def.methods {
for _, method := range definitions {
if method.name == m.Name && method.pkg == m.PkgPath {
_, shared := installed[method.function]
return method.hasInterface && !shared
}
}
return false
})
for j, method := range def.methods {
for j, method := range definitions {
methods[j] = reflectx.Method{Name: method.name, PkgPath: method.pkg, Pointer: method.pointer, Type: t.types[method.typ-1], Func: callbacks[method.function-1]}
}
if err := t.ctx.SetMethodSet(t.types[i], methods, false); err != nil {
Expand Down
4 changes: 2 additions & 2 deletions internal/reflectxtype/method_cache_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -92,7 +92,7 @@ func TestMethodCacheIdentity(t *testing.T) {
var ids [2]int
callbacks := make([]func([]reflect.Value) []reflect.Value, guest.MethodCount())
for i, typ := range originals {
methods := guest.definitions[sent.IDs[typ]-1].methods
methods := guest.methods[sent.IDs[typ]-1]
if len(methods) != 1 || methods[0].name != "Read" {
t.Fatalf("%v lost its Read declaration: %v", typ, methods)
}
Expand Down Expand Up @@ -161,7 +161,7 @@ func TestMethodCacheIdentity(t *testing.T) {
}
for i, typ := range originals {
id := sent.IDs[typ]
if returned.IDs[restored[i]] != id || host.definitions[id-1].methods[0].function != ids[i] {
if returned.IDs[restored[i]] != id || host.methods[id-1][0].function != ids[i] {
t.Fatalf("%v changed its TypeID or MethodID on return", typ)
}
if got, err := host.Resolve(id); err != nil || got != typ {
Expand Down
2 changes: 1 addition & 1 deletion internal/reflectxtype/method_shared_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,7 @@ func TestSharedMethodEntries(t *testing.T) {
shared := make(map[string]int)
locals := make(map[int]bool)
for _, typ := range originals {
for _, method := range table.definitions[snapshot.IDs[typ]-1].methods {
for _, method := range table.methods[snapshot.IDs[typ]-1] {
switch method.name {
case "Local":
if locals[method.function] {
Expand Down
14 changes: 8 additions & 6 deletions internal/reflectxtype/reflectxtype.go
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@ type Snapshot struct {
type ReflectType struct {
types []reflect.Type
entries [][]byte
definitions []definition
methods [][]method
ctx *reflectx.Context
methodCount int
retained int
Expand Down Expand Up @@ -75,22 +75,23 @@ func export(previous *ReflectType, roots []reflect.Type) (*Snapshot, error) {
e := exporter{
ids: make(map[reflect.Type]uint32),
sharedMethods: make(map[methodEntries]int),
static: indexStaticTypes().byType,
}
if previous != nil {
e.entries = make([][]byte, len(previous.types))
e.methods = make([]reflect.Value, previous.methodCount)
e.retainedMethods = make(map[reflect.Type][]method)
for i, typ := range previous.types {
e.ids[typ] = uint32(i + 1)
def := previous.definitions[i]
if def.kind != reflect.Interface && len(def.methods) != 0 {
e.retainedMethods[typ] = def.methods
retained := previous.methods[i]
if len(retained) != 0 {
e.retainedMethods[typ] = retained
methods, _, _, entries := concreteMethodSet(typ)
indices := make(map[methodIdentity]int, len(methods))
for j, method := range methods {
indices[methodIdentity{method.Name, method.PkgPath, method.Pointer}] = j
}
for _, method := range def.methods {
for _, method := range retained {
index, ok := indices[methodIdentity{method.name, method.pkg, method.pointer}]
if !ok {
return nil, fmt.Errorf("retained method %s.%s changed for %v", method.pkg, method.name, typ)
Expand Down Expand Up @@ -137,6 +138,7 @@ type exporter struct {
methods []reflect.Value
retainedMethods map[reflect.Type][]method
sharedMethods map[methodEntries]int
static map[reflect.Type]staticLocation
}

func (e *exporter) intern(typ reflect.Type) (uint32, error) {
Expand All @@ -162,7 +164,7 @@ func (e *exporter) encode(typ reflect.Type) ([]byte, error) {
if builtinTypes[kind] == typ {
return binary.AppendUvarint(nil, uint64(kind)), nil
}
if location, ok := staticTypes().byType[typ]; ok {
if location, ok := e.static[typ]; ok {
data := binary.AppendUvarint(nil, 0)
data = binary.AppendUvarint(data, uint64(location.module))
return binary.AppendUvarint(data, location.offset), nil
Expand Down
10 changes: 5 additions & 5 deletions internal/reflectxtype/roundtrip_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -132,11 +132,11 @@ func TestRetainedMethodIDs(t *testing.T) {
t.Fatal(err)
}
callbacks := make([]func([]reflect.Value) []reflect.Value, guest.MethodCount())
def := guest.definitions[sent.IDs[typ]-1]
if len(def.methods) != len(order) {
t.Fatalf("got %d methods, want %d", len(def.methods), len(order))
retainedMethods := guest.methods[sent.IDs[typ]-1]
if len(retainedMethods) != len(order) {
t.Fatalf("got %d methods, want %d", len(retainedMethods), len(order))
}
for i, method := range def.methods {
for i, method := range retainedMethods {
if got := (identity{method.name, method.pkg}); got != order[i] {
t.Fatalf("source method %d: got %v, want %v", i, got, order[i])
}
Expand All @@ -160,7 +160,7 @@ func TestRetainedMethodIDs(t *testing.T) {
t.Fatal(err)
}
receiver := reflect.New(restored)
for i, method := range def.methods {
for i, method := range retainedMethods {
got := returned.Methods[method.function-1].Call([]reflect.Value{receiver})[0].Int()
if got != int64(i+22) {
t.Fatalf("method %s.%s ID %d returned %d, want %d", method.pkg, method.name, method.function, got, i+22)
Expand Down
3 changes: 0 additions & 3 deletions internal/reflectxtype/runtime.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,6 @@ package reflectxtype

import (
"reflect"
"sync"
"unsafe"
)

Expand All @@ -19,8 +18,6 @@ type staticTypeIndex struct {
byLocation map[staticLocation]reflect.Type
}

var staticTypes = sync.OnceValue(indexStaticTypes)

//go:linkname reflectTypeLinks reflect.typelinks
func reflectTypeLinks() ([]unsafe.Pointer, [][]int32)

Expand Down
2 changes: 1 addition & 1 deletion internal/reflectxtype/struct_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,7 @@ func TestStructTypeIsolation(t *testing.T) {
if embeddedFirst {
roots[0], roots[4] = roots[4], roots[0]
}
e := exporter{ids: make(map[reflect.Type]uint32)}
e := exporter{ids: make(map[reflect.Type]uint32), static: indexStaticTypes().byType}
for _, typ := range roots {
if _, err := e.intern(typ); err != nil {
t.Fatal(err)
Expand Down
2 changes: 1 addition & 1 deletion internal/state/makefunc_linux_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -104,7 +104,7 @@ func TestMakeFuncReinterpretedProcess(t *testing.T) {
if err != nil {
t.Fatal(err)
}
if graph.saved.lastID != objectID(len(loaded.objectsByID)) {
if len(graph.saved.objectsByID) != len(loaded.objectsByID) {
t.Fatal("MakeFunc signature views acquired new object IDs")
}
if err := os.WriteFile(path, mem[:n], 0600); err != nil {
Expand Down
4 changes: 2 additions & 2 deletions internal/state/methodvalue_linux_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -143,12 +143,12 @@ func TestReflectMethodValueRoundTrip(t *testing.T) {
if got := guest[0].(func(int) int)(2); got != 12 || r.N != 10 {
t.Fatal("guest method did not retain an independent receiver")
}
returned := loaded.encoder(ctx, make([]byte, 1<<20))
returned := loaded.loaded().encoder(ctx, make([]byte, 1<<20))
output := saveObjects(t, returned, &guest)
if returned.lastID != saved.lastID {
t.Fatal("rebuilt method environment acquired a new object ID")
}
loadObjects(t, saved.decoder(ctx, output), &host)
loadObjects(t, saved.saved().decoder(ctx, output), &host)
runtime.GC()
if host[3].(*nativeMethodReceiver) != r || r.N != 12 || fn(1) != 13 ||
host[1].(func(int) int)(2) != 15 || host[2].(reflect.Value).Call([]reflect.Value{reflect.ValueOf(3)})[0].Int() != 18 {
Expand Down
26 changes: 23 additions & 3 deletions internal/state/native_linux_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -147,10 +147,30 @@ func TestNativeUnreachableMethod(t *testing.T) {
if err != nil {
t.Fatal(err)
}
if len(graph.saved.pending) != 1 {
t.Fatalf("placeholder emitted %d objects, want only the function", len(graph.saved.pending))
if len(graph.saved.objectsByID) != 1 {
t.Fatalf("placeholder emitted %d objects, want only the function", len(graph.saved.objectsByID))
}
r := reader{mem: mem[:n]}
for range 2 {
length, objects, err := readHeader(&r)
if err != nil || objects {
t.Fatalf("type table header: objects=%t err=%v", objects, err)
}
r.readBytes(length)
}
count, objects, err := readHeader(&r)
if err != nil || !objects || count != 1 {
t.Fatalf("object header: count=%d objects=%t err=%v", count, objects, err)
}
id, err := r.get()
if err != nil || id != uintValue(1) {
t.Fatalf("root ID: %v, %v", id, err)
}
encoded, err := r.get()
if err != nil {
t.Fatal(err)
}
record, ok := graph.saved.pending[1].encoded.(*functionValue)
record, ok := encoded.(*functionValue)
if !ok || record.PC != uintValue(pc) || record.Env.Root != 0 {
t.Fatalf("placeholder record: %#v", record)
}
Expand Down
Loading
Loading