From d29844d6a0df0c14df6e5f4e7db3a0940620d6e7 Mon Sep 17 00:00:00 2001 From: Rick Guo Date: Mon, 21 Sep 2026 13:44:35 +0800 Subject: [PATCH] Transfer reflectx method callbacks without historical receivers --- internal/state/decode.go | 29 ++-- internal/state/encode.go | 14 +- .../state/reflectx_callback_linux_test.go | 155 ++++++++++++++++++ internal/state/reflectx_method_linux_test.go | 29 ++-- 4 files changed, 199 insertions(+), 28 deletions(-) create mode 100644 internal/state/reflectx_callback_linux_test.go diff --git a/internal/state/decode.go b/internal/state/decode.go index 5fda6f5..135474e 100644 --- a/internal/state/decode.go +++ b/internal/state/decode.go @@ -825,22 +825,27 @@ func (ds *decodeState) Load(obj reflect.Value) { if err != nil { Failf("method functions: %w", err) } - methods, ok := encoded.(*arrayValue) - if !ok || len(methods.Contents) != ds.reflectx.MethodCount() { + methods, ok := encoded.(*multipleObjects) + if !ok || len(*methods) != ds.reflectx.MethodCount() { Failf("method function count does not match type table") } - callbacks := make([]func([]reflect.Value) []reflect.Value, len(methods.Contents)) - for i, record := range methods.Contents { - fn, ok := record.(*reflectedValue) - if !ok { + callbacks := make([]func([]reflect.Value) []reflect.Value, len(*methods)) + for i, record := range *methods { + switch fn := record.(type) { + case *functionValue: + // Publish the callback's closure storage now, before its captured + // objects are decoded. SetMethods rebuilds the receiver wrapper. + ds.decodeFunction(reflect.ValueOf(&callbacks[i]).Elem(), fn) + case *reflectedValue: + var function reflect.Value + ds.decodeObject(nil, reflect.ValueOf(&function).Elem(), fn) + if !function.IsValid() || function.Kind() != reflect.Func || function.IsNil() { + Failf("invalid restored method function") + } + callbacks[i] = reflectxMethod{function: function}.call + default: Failf("invalid method function %T", record) } - var function reflect.Value - ds.decodeObject(nil, reflect.ValueOf(&function).Elem(), fn) - if !function.IsValid() || function.Kind() != reflect.Func || function.IsNil() { - Failf("invalid restored method function") - } - callbacks[i] = reflectxMethod{function: function}.call } // Allocate closure storage first, install the method table, then fill // the environments. Interfaces decoded below see the final Ifn entries. diff --git a/internal/state/encode.go b/internal/state/encode.go index f42463b..5ca5f3f 100644 --- a/internal/state/encode.go +++ b/internal/state/encode.go @@ -906,7 +906,7 @@ func (es *encodeState) Save(obj reflect.Value) { var oes *objectEncodeState var snapshot *reflecttype.Snapshot var extended *reflectxtype.Snapshot - var methods arrayValue + var methods multipleObjects if err := safely(func() { for { for oes = es.deferred.Front(); oes != nil; oes = es.deferred.Front() { @@ -945,10 +945,16 @@ func (es *encodeState) Save(obj reflect.Value) { if err != nil { Failf("export reflectx types: %w", err) } - methods.Contents = make([]object, len(extended.Methods)) + methods = make(multipleObjects, len(extended.Methods)) for i, fn := range extended.Methods { fn = es.native.originalFunction(fn) - es.encodeObject(reflect.ValueOf(fn), encodeAsValue, &methods.Contents[i]) + if impl := makeFuncStorage(fn); impl != nil { + // A shared Tfn may still carry an earlier receiver type. The + // type table owns the signature; only the callback is state. + es.encodeFunction(reflect.ValueOf(impl.fn), &methods[i]) + } else { + es.encodeObject(reflect.ValueOf(fn), encodeAsValue, &methods[i]) + } } // Method signatures and environments can expose more types. Finish // both before assigning the final type IDs for the entire graph. @@ -998,7 +1004,7 @@ func (es *encodeState) Save(obj reflect.Value) { Failf("error writing reflectx type table header: %w", err) } es.w.writeBytes(reflectxData) - if len(methods.Contents) != 0 { + if len(methods) != 0 { if err := es.w.put(&methods); err != nil { Failf("writing method functions: %w", err) } diff --git a/internal/state/reflectx_callback_linux_test.go b/internal/state/reflectx_callback_linux_test.go new file mode 100644 index 0000000..da5e357 --- /dev/null +++ b/internal/state/reflectx_callback_linux_test.go @@ -0,0 +1,155 @@ +//go:build linux && (amd64 || arm64) + +package state + +import ( + "context" + "os" + "os/exec" + "reflect" + "testing" + + "github.com/goplus/reflectx" +) + +func TestReflectxSharedCallbackRoots(t *testing.T) { + for _, mode := range []string{"value", "pointer"} { + t.Run(mode, func(t *testing.T) { + // FuncId entries are process-global; keep the two fixtures independent. + if os.Getenv("SANDBOX_CALLBACK_ROOTS") != mode { + command := exec.Command(os.Args[0], "-test.run=^TestReflectxSharedCallbackRoots$/^"+mode+"$", "-test.v") + command.Env = append(os.Environ(), "SANDBOX_CALLBACK_ROOTS="+mode) + if output, err := command.CombinedOutput(); err != nil { + t.Fatalf("callback roots: %v\n%s", err, output) + } + return + } + ctx := reflectx.NewContext() + t.Cleanup(ctx.Reset) + pointer := mode == "pointer" + shared := 7 + read := func(args []reflect.Value) []reflect.Value { + value := args[0] + if pointer { + value = value.Elem() + } + shared++ + return []reflect.Value{reflect.ValueOf(int(value.Field(0).Int()) + shared)} + } + var types [3]reflect.Type + var values [3]any + var counters [3]*int + for i, name := range []string{"Historical", "Current", "Peer"} { + valueMethods := 3 + if pointer { + valueMethods = 0 + } + typ := ctx.NewMethodSet(reflectx.NamedTypeOf("example/callbackroots", name, reflect.TypeFor[struct{ N int }]()), valueMethods, 3) + counter := new(int) + *counter = 100 * (i + 1) + methods := []reflectx.Method{ + reflectx.MakeMethod("Read", "", pointer, reflect.TypeFor[func() int](), read), + reflectx.MakeMethod("Unused", "", pointer, reflect.TypeFor[func()](), nil), + reflectx.MakeMethod("Own", "", pointer, reflect.TypeFor[func() int](), func([]reflect.Value) []reflect.Value { + *counter++ + return []reflect.Value{reflect.ValueOf(*counter)} + }), + } + methods[0].FuncId, methods[1].FuncId = 1, 2 + if err := ctx.SetMethodSet(typ, methods, false); err != nil { + t.Fatal(err) + } + value := reflect.New(typ) + value.Elem().Field(0).SetInt(int64(10 * (i + 1))) + if !pointer { + value = value.Elem() + } + types[i], values[i], counters[i] = typ, value.Interface(), counter + } + first, _ := reflectx.MethodByName(reflect.TypeOf(values[0]), "Read") + second, _ := reflectx.MethodByName(reflect.TypeOf(values[1]), "Read") + if makeFuncStorage(first.Func) != makeFuncStorage(second.Func) || nativeReflectType(makeFuncStorage(second.Func).ftyp).In(0) != first.Type.In(0) { + t.Fatal("fixture must reuse a wrapper with the historical receiver signature") + } + type root struct { + Values [2]any + Counters [2]*int + Shared *int + } + src := root{[2]any{values[1], values[2]}, [2]*int{counters[1], counters[2]}, &shared} + mem := make([]byte, 1<<20) + for round := range 3 { + var host, guest State + n, _, err := host.Save(context.Background(), mem, &src) + if err != nil { + t.Fatal(err) + } + snapshot := host.saved.reflectxSnapshot + 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)) + } + for _, object := range host.saved.pending { + if object.obj.Type() == reflect.TypeFor[int]() && object.obj.Addr().Interface().(*int) == counters[0] { + t.Fatal("historical receiver's Own capture entered the graph") + } + } + _, before, _ := reflectx.IcallStat() + cached := reflectx.IcallCached() + var dst root + if _, err := guest.Load(context.Background(), mem[:n], &dst); err != nil { + t.Fatal(err) + } + _, after, _ := reflectx.IcallStat() + if after-before != 4 || reflectx.IcallCached() != cached { + t.Fatalf("method import allocated %d slots, want 4 without global cache additions", after-before) + } + wantShared := shared + for i, value := range dst.Values { + wantShared++ + if got := value.(interface{ Read() int }).Read(); got != 10*(i+2)+wantShared { + t.Fatalf("interface Read: got %d", got) + } + method, _ := reflectx.MethodByName(reflect.TypeOf(value), "Read") + wantShared++ + if got := method.Func.Call([]reflect.Value{reflect.ValueOf(value)})[0].Int(); got != int64(10*(i+2)+wantShared) { + t.Fatalf("reflection Read: got %d", got) + } + if got := value.(interface{ Own() int }).Own(); got != 100*(i+2)+round+1 || *dst.Counters[i] != got { + t.Fatalf("Own lost its capture alias: got %d", got) + } + unused, _ := reflectx.MethodByName(reflect.TypeOf(value), "Unused") + if !makeFuncCallback(unused.Func).IsNil() { + t.Fatal("nil callback acquired a wrapper") + } + if *counters[i+1] != 100*(i+2)+round { + t.Fatal("guest changed host captures before writeback") + } + } + if *dst.Shared != wantShared || shared != wantShared-4 { + t.Fatal("shared callbacks lost their captured alias or host isolation") + } + n, _, err = guest.Save(context.Background(), mem, &dst) + if err != nil { + t.Fatal(err) + } + if len(guest.saved.reflectxSnapshot.Methods) != 4 { + t.Fatal("return added method implementations") + } + if _, err := host.Load(context.Background(), mem[:n], &src); err != nil { + t.Fatal(err) + } + if src.Shared != &shared || shared != wantShared || *counters[0] != 100 { + t.Fatal("return changed shared capture identity or historical receiver state") + } + for i := range src.Values { + if src.Values[i] != values[i+1] || src.Counters[i] != counters[i+1] || *counters[i+1] != 100*(i+2)+round+1 { + t.Fatal("return lost receiver identity or capture writeback") + } + } + } + }) + } +} diff --git a/internal/state/reflectx_method_linux_test.go b/internal/state/reflectx_method_linux_test.go index 4423f3d..db8e1f7 100644 --- a/internal/state/reflectx_method_linux_test.go +++ b/internal/state/reflectx_method_linux_test.go @@ -271,27 +271,32 @@ func checkOriginalMethodRecords(t *testing.T, es *encodeState, data []byte) { if err != nil { t.Fatal(err) } - methods, ok := encoded.(*arrayValue) + methods, ok := encoded.(*multipleObjects) if !ok { t.Fatalf("method table is %T", encoded) } var native, dynamic int - for _, record := range methods.Contents { - value, ok := record.(*reflectedValue) - if !ok || value.Addressable { - t.Fatalf("method is not an original function value: %T", record) - } - fn, ok := value.Value.(*functionValue) - if !ok { - t.Fatalf("method payload is %T", value.Value) - } - if uintptr(fn.PC) == makeFuncPC { + for _, record := range *methods { + switch value := record.(type) { + case *functionValue: dynamic++ - } else { + if uintptr(value.PC) == makeFuncPC { + t.Fatal("dynamic method kept its outer MakeFunc wrapper") + } + case *reflectedValue: native++ + if value.Addressable { + t.Fatal("native method is not an original function value") + } + fn, ok := value.Value.(*functionValue) + if !ok { + t.Fatalf("method payload is %T", value.Value) + } if fn.Env.Root != 0 { t.Fatalf("native method acquired a wrapper environment: %v", fn.Env) } + default: + t.Fatalf("invalid method record: %T", record) } } if native == 0 || dynamic == 0 {