diff --git a/guest_linux.go b/guest_linux.go index 8e41a65..de1ae69 100644 --- a/guest_linux.go +++ b/guest_linux.go @@ -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) { @@ -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 } diff --git a/host_linux.go b/host_linux.go index a7433eb..d96e2ed 100644 --- a/host_linux.go +++ b/host_linux.go @@ -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 @@ -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 } diff --git a/internal/reflecttype/decode.go b/internal/reflecttype/decode.go index 3c6afa9..1a724b1 100644 --- a/internal/reflecttype/decode.go +++ b/internal/reflecttype/decode.go @@ -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() @@ -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 { @@ -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)) } diff --git a/internal/reflecttype/reflecttype.go b/internal/reflecttype/reflecttype.go index 0814173..2892888 100644 --- a/internal/reflecttype/reflecttype.go +++ b/internal/reflecttype/reflecttype.go @@ -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 @@ -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 { @@ -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) { diff --git a/internal/reflecttype/runtime.go b/internal/reflecttype/runtime.go index 9282e82..aa1e672 100644 --- a/internal/reflecttype/runtime.go +++ b/internal/reflecttype/runtime.go @@ -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) diff --git a/internal/reflectxtype/decode.go b/internal/reflectxtype/decode.go index f6f9d43..3e26448 100644 --- a/internal/reflectxtype/decode.go +++ b/internal/reflectxtype/decode.go @@ -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 { @@ -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 { @@ -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) { @@ -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)) } diff --git a/internal/reflectxtype/method.go b/internal/reflectxtype/method.go index 8ecbfb6..840fe92 100644 --- a/internal/reflectxtype/method.go +++ b/internal/reflectxtype/method.go @@ -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 { @@ -149,9 +149,9 @@ 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 @@ -159,7 +159,7 @@ func (t *ReflectType) SetMethods(callbacks []func([]reflect.Value) []reflect.Val } 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 { diff --git a/internal/reflectxtype/method_cache_test.go b/internal/reflectxtype/method_cache_test.go index 0416614..65cae1e 100644 --- a/internal/reflectxtype/method_cache_test.go +++ b/internal/reflectxtype/method_cache_test.go @@ -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) } @@ -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 { diff --git a/internal/reflectxtype/method_shared_test.go b/internal/reflectxtype/method_shared_test.go index b4923c4..217bbc0 100644 --- a/internal/reflectxtype/method_shared_test.go +++ b/internal/reflectxtype/method_shared_test.go @@ -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] { diff --git a/internal/reflectxtype/reflectxtype.go b/internal/reflectxtype/reflectxtype.go index e5b177b..b8f1b37 100644 --- a/internal/reflectxtype/reflectxtype.go +++ b/internal/reflectxtype/reflectxtype.go @@ -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 @@ -75,6 +75,7 @@ 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)) @@ -82,15 +83,15 @@ func export(previous *ReflectType, roots []reflect.Type) (*Snapshot, error) { 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) @@ -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) { @@ -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 diff --git a/internal/reflectxtype/roundtrip_test.go b/internal/reflectxtype/roundtrip_test.go index 0bbd564..78b5c51 100644 --- a/internal/reflectxtype/roundtrip_test.go +++ b/internal/reflectxtype/roundtrip_test.go @@ -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]) } @@ -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) diff --git a/internal/reflectxtype/runtime.go b/internal/reflectxtype/runtime.go index 7e5278e..974a657 100644 --- a/internal/reflectxtype/runtime.go +++ b/internal/reflectxtype/runtime.go @@ -5,7 +5,6 @@ package reflectxtype import ( "reflect" - "sync" "unsafe" ) @@ -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) diff --git a/internal/reflectxtype/struct_test.go b/internal/reflectxtype/struct_test.go index f03ea6b..6c9529a 100644 --- a/internal/reflectxtype/struct_test.go +++ b/internal/reflectxtype/struct_test.go @@ -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) diff --git a/internal/state/makefunc_linux_test.go b/internal/state/makefunc_linux_test.go index baafa83..bbc5605 100644 --- a/internal/state/makefunc_linux_test.go +++ b/internal/state/makefunc_linux_test.go @@ -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 { diff --git a/internal/state/methodvalue_linux_test.go b/internal/state/methodvalue_linux_test.go index 2defffe..6d61eb7 100644 --- a/internal/state/methodvalue_linux_test.go +++ b/internal/state/methodvalue_linux_test.go @@ -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 { diff --git a/internal/state/native_linux_test.go b/internal/state/native_linux_test.go index aede831..3a96eb5 100644 --- a/internal/state/native_linux_test.go +++ b/internal/state/native_linux_test.go @@ -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) } diff --git a/internal/state/reflectx_callback_linux_test.go b/internal/state/reflectx_callback_linux_test.go index da5e357..091a4c9 100644 --- a/internal/state/reflectx_callback_linux_test.go +++ b/internal/state/reflectx_callback_linux_test.go @@ -88,11 +88,14 @@ func TestReflectxSharedCallbackRoots(t *testing.T) { if snapshot.IDs[types[0]] != 0 || snapshot.IDs[types[1]] == 0 || snapshot.IDs[types[2]] == 0 { t.Fatal("shared method imported its historical receiver type") } - if len(snapshot.Methods) != 4 { - t.Fatalf("method implementations: got %d, want 4", len(snapshot.Methods)) + if methods := readMethodRecords(t, mem[:n]); len(methods) != 4 { + t.Fatalf("method implementations: got %d, want 4", len(methods)) } - for _, object := range host.saved.pending { - if object.obj.Type() == reflect.TypeFor[int]() && object.obj.Addr().Interface().(*int) == counters[0] { + if snapshot.Methods != nil { + t.Fatal("save retained exported method functions") + } + for _, object := range host.saved.objectsByID { + if object.obj.IsValid() && object.obj.Type() == reflect.TypeFor[int]() && object.obj.Addr().Interface().(*int) == counters[0] { t.Fatal("historical receiver's Own capture entered the graph") } } @@ -135,7 +138,7 @@ func TestReflectxSharedCallbackRoots(t *testing.T) { if err != nil { t.Fatal(err) } - if len(guest.saved.reflectxSnapshot.Methods) != 4 { + if len(readMethodRecords(t, mem[:n])) != 4 { t.Fatal("return added method implementations") } if _, err := host.Load(context.Background(), mem[:n], &src); err != nil { diff --git a/internal/state/reflectx_method_linux_test.go b/internal/state/reflectx_method_linux_test.go index db8e1f7..4a30e8b 100644 --- a/internal/state/reflectx_method_linux_test.go +++ b/internal/state/reflectx_method_linux_test.go @@ -257,7 +257,7 @@ func TestReflectxNativeMethodProcess(t *testing.T) { } } -func checkOriginalMethodRecords(t *testing.T, es *encodeState, data []byte) { +func readMethodRecords(t *testing.T, data []byte) multipleObjects { t.Helper() r := reader{mem: data} for range 2 { @@ -275,8 +275,14 @@ func checkOriginalMethodRecords(t *testing.T, es *encodeState, data []byte) { if !ok { t.Fatalf("method table is %T", encoded) } + return *methods +} + +func checkOriginalMethodRecords(t *testing.T, saved *savedState, data []byte) { + t.Helper() + methods := readMethodRecords(t, data) var native, dynamic int - for _, record := range *methods { + for _, record := range methods { switch value := record.(type) { case *functionValue: dynamic++ @@ -304,10 +310,13 @@ func checkOriginalMethodRecords(t *testing.T, es *encodeState, data []byte) { } callPC := reflect.ValueOf(reflect.Value{}.Call).Pointer() callSlicePC := reflect.ValueOf(reflect.Value{}.CallSlice).Pointer() - for _, obj := range es.pending { - pc := es.native.storage[obj.obj.Type()] + for i, obj := range saved.objectsByID { + if !obj.obj.IsValid() { + continue + } + pc := saved.native.storage[obj.obj.Type()] if pc == reflectxMethodCallPC || pc == callPC || pc == callSlicePC { - t.Fatalf("local method adapter entered the object graph: ID %d", obj.id) + t.Fatalf("local method adapter entered the object graph: ID %d", i+1) } } } diff --git a/internal/state/reflectx_roots_linux_test.go b/internal/state/reflectx_roots_linux_test.go index a18d2cf..5f605ef 100644 --- a/internal/state/reflectx_roots_linux_test.go +++ b/internal/state/reflectx_roots_linux_test.go @@ -70,8 +70,8 @@ func TestReflectxMethodRoots(t *testing.T) { if referenced { wantMethods, wantNumber, wantForeign = 2, 32, 21 } - if len(snapshot.Methods) != wantMethods { - t.Fatalf("method count: got %d, want %d", len(snapshot.Methods), wantMethods) + if methods := readMethodRecords(t, mem[:n]); len(methods) != wantMethods { + t.Fatalf("method count: got %d, want %d", len(methods), wantMethods) } var dst root if _, err := guest.Load(ctx, mem[:n], &dst); err != nil { diff --git a/internal/state/retention_test.go b/internal/state/retention_test.go new file mode 100644 index 0000000..dd1335a --- /dev/null +++ b/internal/state/retention_test.go @@ -0,0 +1,103 @@ +package state + +import ( + "bytes" + "context" + "runtime" + "testing" + "weak" +) + +func TestStateReleasesTransferResources(t *testing.T) { + for _, stream := range []bool{false, true} { + name := "Save" + if stream { + name = "SaveTo" + } + t.Run(name, func(t *testing.T) { + value := 10 + host := [2]*int{&value, &value} + var source, destination State + data, resources := saveWithResources(t, &source, &host, stream) + runtime.GC() + for _, pointer := range resources { + if pointer.Value() != nil { + t.Fatal("Save retained its context or output buffer") + } + } + var guest [2]*int + resources = loadWithResources(t, &destination, data, &guest) + runtime.GC() + for _, pointer := range resources { + if pointer.Value() != nil { + t.Fatal("Load retained its context or input buffer") + } + } + if guest[0] == &value || guest[0] != guest[1] || *guest[0] != 10 { + t.Fatal("Load lost isolation or aliases") + } + *guest[0] = 20 + guest[0], guest[1] = nil, nil + data, resources = saveWithResources(t, &destination, &guest, stream) + runtime.GC() + for _, pointer := range resources { + if pointer.Value() != nil { + t.Fatal("return Save retained its context or output buffer") + } + } + resources = loadWithResources(t, &source, data, &host) + runtime.GC() + for _, pointer := range resources { + if pointer.Value() != nil { + t.Fatal("writeback retained its context or input buffer") + } + } + if host != [2]*int{} || value != 20 { + t.Fatal("writeback lost a detached object's identity") + } + runtime.KeepAlive(&source) + runtime.KeepAlive(&destination) + }) + } +} + +// Keep the caller's resources outside the frame that runs GC, so compiler +// liveness cannot mask references retained by State itself. +// +//go:noinline +func saveWithResources(t *testing.T, graph *State, root any, stream bool) ([]byte, []weak.Pointer[byte]) { + t.Helper() + marker := make([]byte, 4096) + ctx := context.WithValue(context.Background(), struct{}{}, marker) + resources := []weak.Pointer[byte]{weak.Make(&marker[0])} + var data []byte + if stream { + var out bytes.Buffer + if _, _, err := graph.SaveTo(ctx, &out, root); err != nil { + t.Fatal(err) + } + data = out.Bytes() + } else { + mem := make([]byte, 1<<20) + n, _, err := graph.Save(ctx, mem, root) + if err != nil { + t.Fatal(err) + } + data = mem[:n] + } + resources = append(resources, weak.Make(&data[0])) + return bytes.Clone(data), resources +} + +//go:noinline +func loadWithResources(t *testing.T, graph *State, data []byte, root any) []weak.Pointer[byte] { + t.Helper() + marker := make([]byte, 4096) + ctx := context.WithValue(context.Background(), struct{}{}, marker) + input := bytes.Clone(data) + resources := []weak.Pointer[byte]{weak.Make(&marker[0]), weak.Make(&input[0])} + if _, err := graph.Load(ctx, input, root); err != nil { + t.Fatal(err) + } + return resources +} diff --git a/internal/state/roundtrip.go b/internal/state/roundtrip.go index 5ac67c0..a916075 100644 --- a/internal/state/roundtrip.go +++ b/internal/state/roundtrip.go @@ -5,6 +5,8 @@ import ( "io" "reflect" "sync" + + "github.com/xgo-dev/sandbox/internal/reflectxtype" ) // Separate States can retain the same host objects. Serialize their restores, @@ -16,8 +18,29 @@ var loadMu sync.Mutex // keeps it alive until the round trip finishes, including unlinked objects. // Calls on one State must not overlap, and the root's storage must not change. type State struct { - saved *encodeState - loaded *decodeState + saved *savedState + loaded *loadedState +} + +// Only object identity survives a transfer, not its encoded value or the +// traversal bookkeeping. The slice index is objectID-1; holes stay empty. +type objectState struct { + obj reflect.Value + how encodeStrategy +} + +type savedState struct { + objectsByID []objectState + native nativeState + reflectxSnapshot *reflectxtype.Snapshot +} + +type loadedState struct { + objectsByID []objectState + native nativeState + reflectx *reflectxtype.ReflectType + makeFuncs map[reflect.Value]reflect.Value + methodValues map[reflect.Value]reflect.Value } // Save writes the graph, preserving IDs from the preceding Load when present. @@ -41,7 +64,7 @@ func (s *State) save(ctx context.Context, mem []byte, out io.Writer, rootPtr any es.Save(reflect.ValueOf(rootPtr).Elem()) }) if err == nil { - s.saved, s.loaded = es, nil + s.saved, s.loaded = es.saved(), nil } return es.w.pos, es.stats, err } @@ -62,14 +85,17 @@ func (s *State) Load(ctx context.Context, mem []byte, rootPtr any) (Stats, error check := newDecodeState(ctx, mem) check.native = s.saved.native check.reflectxSnapshot = s.saved.reflectxSnapshot - for id, saved := range s.saved.pending { + for i, saved := range s.saved.objectsByID { + if !saved.obj.IsValid() { + continue + } value := reflect.New(saved.obj.Type()).Elem() if saved.how == encodeMapAsValue { value.Set(reflect.MakeMap(value.Type())) } else if saved.how == encodeChannelAsValue { value.Set(reflect.MakeChan(value.Type(), saved.obj.Cap())) } - check.addObject(id, value).how = saved.how + check.addObject(objectID(i+1), value).how = saved.how } check.Load(check.lookup(1).obj) ds = s.saved.decoder(ctx, mem) @@ -77,20 +103,55 @@ func (s *State) Load(ctx context.Context, mem []byte, rootPtr any) (Stats, error ds.Load(reflect.ValueOf(rootPtr).Elem()) }) if err == nil { - s.loaded, s.saved = ds, nil + s.loaded, s.saved = ds.loaded(), nil } return ds.stats, err } +func (es *encodeState) saved() *savedState { + s := &savedState{ + objectsByID: make([]objectState, es.lastID), + native: es.native, + } + for id, object := range es.pending { + s.objectsByID[id-1] = objectState{obj: object.obj, how: object.how} + } + if snapshot := es.reflectxSnapshot; snapshot != nil { + // Method functions have already entered the object graph. Open only + // needs the definitions and original types to validate a returned table. + s.reflectxSnapshot = &reflectxtype.Snapshot{Data: snapshot.Data, IDs: snapshot.IDs} + } + return s +} + +func (ds *decodeState) loaded() *loadedState { + s := &loadedState{ + objectsByID: make([]objectState, len(ds.objectsByID)), + native: ds.native, + reflectx: ds.reflectx, + makeFuncs: ds.makeFuncs, + methodValues: ds.methodValues, + } + for i, object := range ds.objectsByID { + if object != nil { + s.objectsByID[i] = objectState{obj: object.obj, how: object.how} + } + } + return s +} + // decoder reuses the objects from this save when loading its returned graph. -// Keeping es alive retains even objects the guest later unlinks from the root. +// Keeping s alive retains even objects the guest later unlinks from the root. // No source address is written to the stream; both sides use the existing IDs. -func (es *encodeState) decoder(ctx context.Context, mem []byte) *decodeState { +func (s *savedState) decoder(ctx context.Context, mem []byte) *decodeState { ds := newDecodeState(ctx, mem) - ds.native = es.native - ds.reflectxSnapshot = es.reflectxSnapshot - for id, saved := range es.pending { + ds.native = s.native + ds.reflectxSnapshot = s.reflectxSnapshot + for i, saved := range s.objectsByID { value := saved.obj + if !value.IsValid() { + continue + } if saved.how == encodeChannelAsValue { // Channels retain the independent-snapshot semantics of Load. // Never enqueue data into, or close, the original host channel. @@ -99,7 +160,7 @@ func (es *encodeState) decoder(ctx context.Context, mem []byte) *decodeState { channel.Set(reflect.MakeChan(typ, value.Cap())) value = channel } - ods := ds.addObject(id, value) + ods := ds.addObject(objectID(i+1), value) ods.how = saved.how } return ds @@ -107,13 +168,13 @@ func (es *encodeState) decoder(ctx context.Context, mem []byte) *decodeState { // encoder starts the return save with the decoded objects' existing IDs. // For example, swapping two pointers changes their refs, not their object IDs. -func (ds *decodeState) encoder(ctx context.Context, mem []byte) *encodeState { +func (s *loadedState) encoder(ctx context.Context, mem []byte) *encodeState { es := newEncodeState(ctx, mem) - es.native = ds.native - es.reflectx = ds.reflectx - es.lastID = objectID(len(ds.objectsByID)) - for _, decoded := range ds.objectsByID { - if decoded == nil { + es.native = s.native + es.reflectx = s.reflectx + es.lastID = objectID(len(s.objectsByID)) + for i, decoded := range s.objectsByID { + if !decoded.obj.IsValid() { continue } value := decoded.obj @@ -121,7 +182,7 @@ func (ds *decodeState) encoder(ctx context.Context, mem []byte) *encodeState { if decoded.how == encodeMapAsValue || decoded.how == encodeChannelAsValue { addr, size = value.Pointer(), 1 } - oes := &objectEncodeState{id: decoded.id, obj: value, how: decoded.how, retained: true} + oes := &objectEncodeState{id: objectID(i + 1), obj: value, how: decoded.how, retained: true} es.pending[oes.id] = oes es.deferred.PushBack(oes) if size == 0 && addr == dummyAddr { @@ -133,7 +194,7 @@ func (ds *decodeState) encoder(ctx context.Context, mem []byte) *encodeState { } // MakeFunc rebuilds a wrapper whose callback slot differs from the // decoded slot. Both locations must resolve to that same object ID. - for callback, fn := range ds.makeFuncs { + for callback, fn := range s.makeFuncs { seg, _ := es.values.Find(callback.Addr().Pointer()) if !seg.Ok() { Failf("MakeFunc callback object is missing") @@ -145,7 +206,7 @@ func (ds *decodeState) encoder(ctx context.Context, mem []byte) *encodeState { } // Rebuilt method wrappers keep the decoded environment's object ID, just // like MakeFunc callback slots above. - for env, fn := range ds.methodValues { + for env, fn := range s.methodValues { seg, _ := es.values.Find(env.Addr().Pointer()) if !seg.Ok() { Failf("method environment object is missing") diff --git a/internal/state/roundtrip_linux_test.go b/internal/state/roundtrip_linux_test.go index 775c085..018c4b9 100644 --- a/internal/state/roundtrip_linux_test.go +++ b/internal/state/roundtrip_linux_test.go @@ -44,7 +44,7 @@ func TestObjectIDClosureProcess(t *testing.T) { } guest.Native = nil runtime.GC() - 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 != objectID(len(loaded.objectsByID)) { t.Fatal("native or MakeFunc environments acquired new IDs") @@ -77,7 +77,7 @@ func TestObjectIDClosureProcess(t *testing.T) { if err != nil { t.Fatal(err) } - loadObjects(t, saved.decoder(ctx, output), &host) + loadObjects(t, saved.saved().decoder(ctx, output), &host) runtime.GC() if host.Native != nil || n != 15 || receiver.N != 23 || original() != 16 || host.Wrapped() != 18 || host.Alias() != 20 || fn() != 22 || host.Method(1) != 24 { t.Fatalf("host captures did not preserve identity: n=%d receiver=%d", n, receiver.N) diff --git a/internal/state/roundtrip_test.go b/internal/state/roundtrip_test.go index 4808d64..97242d8 100644 --- a/internal/state/roundtrip_test.go +++ b/internal/state/roundtrip_test.go @@ -57,12 +57,12 @@ func TestObjectIDRoundTrip(t *testing.T) { guest.Slice = append(guest.Slice, 4) guest.Value.SetInt(12) runtime.GC() - 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 <= objectID(len(loaded.objectsByID)) { t.Fatal("new objects did not receive new IDs") } - loadObjects(t, saved.decoder(ctx, output), &host) + loadObjects(t, saved.saved().decoder(ctx, output), &host) if host.A != b || host.B != a || a.Value != 12 || b.Value != 20 || a.Next != b || b.Next != a { t.Fatal("pointer swap, original identity or cycle was lost") } @@ -94,19 +94,19 @@ func TestObjectIDRepeatedRoundTrip(t *testing.T) { loadObjects(t, loaded, &guest) for i := int64(11); i < 15; i++ { guest[0].Value = i - 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("unchanged graph acquired new IDs") } - restored := saved.decoder(ctx, output) + restored := saved.saved().decoder(ctx, output) loadObjects(t, restored, &host) if host[0] != n || host[1] != n || n.Value != i { t.Fatal("original object was replaced") } - saved = restored.encoder(ctx, make([]byte, 1<<20)) + saved = restored.loaded().encoder(ctx, make([]byte, 1<<20)) input = saveObjects(t, saved, &host) - loaded = returned.decoder(ctx, input) + loaded = returned.saved().decoder(ctx, input) loadObjects(t, loaded, &guest) } } @@ -126,7 +126,7 @@ func TestObjectIDMergedStorage(t *testing.T) { *guest[0].(*int64) = 11 *guest[1].(**graphNode) = guest[3].(*graphNode) guest[4] = &graphNode{Value: 30, Next: guest[2].(*graphNode)} - returned := loaded.encoder(ctx, make([]byte, 1<<20)) + returned := loaded.loaded().encoder(ctx, make([]byte, 1<<20)) output := saveObjects(t, returned, &guest) for id := range saved.pending { if returned.pending[id] == nil { @@ -136,7 +136,7 @@ func TestObjectIDMergedStorage(t *testing.T) { if returned.lastID != saved.lastID+1 { t.Fatal("new object reused an old ID") } - loadObjects(t, saved.decoder(ctx, output), &host) + loadObjects(t, saved.saved().decoder(ctx, output), &host) if host[0] != &n.Value || host[1] != &n.Next || host[2] != n || host[3] != other || n.Value != 11 || n.Next != other || host[4].(*graphNode).Next != n { t.Fatal("merged fields or parent lost their identity") } @@ -151,12 +151,12 @@ func TestObjectIDRetainedStorage(t *testing.T) { ctx := context.Background() saved := newEncodeState(ctx, make([]byte, 1<<20)) input := saveObjects(t, saved, &host) - restored := saved.decoder(ctx, input) + restored := saved.saved().decoder(ctx, input) loadObjects(t, restored, &host) // A later save must not repurpose the field's existing ID for its // previously unseen parent. The other process allocated only an int. host.Whole = &array - next := restored.encoder(ctx, make([]byte, 1<<20)) + next := restored.loaded().encoder(ctx, make([]byte, 1<<20)) err := safely(func() { next.Save(reflect.ValueOf(&host).Elem()) }) if err == nil || !strings.Contains(err.Error(), "cannot change retained storage") { t.Fatalf("retained object ID was repurposed: %v", err) @@ -209,9 +209,9 @@ func TestObjectIDChannelSnapshot(t *testing.T) { guestChannel <- value guestChannel <- value close(guestChannel) - returned := loaded.encoder(ctx, make([]byte, 1<<20)) + returned := loaded.loaded().encoder(ctx, make([]byte, 1<<20)) output := saveObjects(t, returned, guest.Interface()) - loadObjects(t, saved.decoder(ctx, output), &host) + loadObjects(t, saved.saved().decoder(ctx, output), &host) if host.Channel == ch || host.Alias != host.Channel || len(ch) != 1 || len(host.Channel) != 2 || n.Value != 2 { t.Fatal("channel snapshot modified the original queue or lost aliases") } diff --git a/internal/state/subslice_test.go b/internal/state/subslice_test.go index fb7d462..d5b9935 100644 --- a/internal/state/subslice_test.go +++ b/internal/state/subslice_test.go @@ -67,7 +67,7 @@ func TestSubsliceRoundTrip(t *testing.T) { if err != nil { t.Fatal(err) } - if guestState.saved.lastID != hostState.saved.lastID { + if len(guestState.saved.objectsByID) != len(hostState.saved.objectsByID) { t.Fatal("subarray views allocated new object IDs on return") } if _, err := hostState.Load(ctx, mem[:n], &host); err != nil { diff --git a/internal/state/sync_linux_test.go b/internal/state/sync_linux_test.go index aca2f2d..59fb65d 100644 --- a/internal/state/sync_linux_test.go +++ b/internal/state/sync_linux_test.go @@ -35,9 +35,9 @@ func TestSyncPoolNewRoundTrip(t *testing.T) { } guest.Pool.Put("guest cache") host.Pool.Put("old destination cache") - returned := loaded.encoder(ctx, make([]byte, 1<<20)) + returned := loaded.loaded().encoder(ctx, make([]byte, 1<<20)) output := saveObjects(t, returned, &guest) - loadObjects(t, saved.decoder(ctx, output), host) + loadObjects(t, saved.saved().decoder(ctx, output), host) if host.Alias != &host.Pool || host.Pool.New == nil || host.Calls != 1 { t.Fatal("returned Pool lost its alias, New or captured mutation") }