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
29 changes: 17 additions & 12 deletions internal/state/decode.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
14 changes: 10 additions & 4 deletions internal/state/encode.go
Original file line number Diff line number Diff line change
Expand Up @@ -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() {
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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)
}
Expand Down
155 changes: 155 additions & 0 deletions internal/state/reflectx_callback_linux_test.go
Original file line number Diff line number Diff line change
@@ -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")
}
}
}
})
}
}
29 changes: 17 additions & 12 deletions internal/state/reflectx_method_linux_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
Loading