From d006d95c6bc86fee820bd40202e0f3640e39d700 Mon Sep 17 00:00:00 2001 From: Rick Guo Date: Thu, 24 Sep 2026 10:29:53 +0800 Subject: [PATCH] state: preserve sync.Once completion across transfers --- internal/state/sync.go | 7 ++- internal/state/sync_test.go | 102 ++++++++++++++++++++++++++++++++---- 2 files changed, 97 insertions(+), 12 deletions(-) diff --git a/internal/state/sync.go b/internal/state/sync.go index 6f8cc98..386228d 100644 --- a/internal/state/sync.go +++ b/internal/state/sync.go @@ -10,8 +10,9 @@ import ( // syncFields projects synchronization wrappers onto their values. In // particular, atomic.Pointer[T].v must be viewed as *T so state relocates its // target instead of saving an untyped address. Pool.New and Cond.L retain their -// object graphs; other sync primitives are empty records. sync.Map uses its own -// entry codec. The source data graph is quiescent. +// object graphs; Once retains done, with a fresh mutex on load. Other sync +// primitives are empty records. sync.Map uses its own entry codec. The source +// data graph is quiescent, including any Once.Do call. func syncFields(obj reflect.Value) ([]reflect.Value, bool) { typ := obj.Type() pkg, name := typ.PkgPath(), typ.Name() @@ -20,6 +21,8 @@ func syncFields(obj reflect.Value) ([]reflect.Value, bool) { return []reflect.Value{obj.FieldByName("New")}, true case reflect.TypeFor[sync.Cond](): return []reflect.Value{obj.FieldByName("L")}, true + case reflect.TypeFor[sync.Once](): + return []reflect.Value{obj.FieldByName("done")}, true } if pkg == "sync" && typ != reflect.TypeFor[sync.Map]() || pkg == "internal/sync" && name == "Mutex" { return nil, true diff --git a/internal/state/sync_test.go b/internal/state/sync_test.go index 7c46ca6..5bb5a28 100644 --- a/internal/state/sync_test.go +++ b/internal/state/sync_test.go @@ -1,6 +1,8 @@ package state import ( + "bytes" + "context" "reflect" "runtime" "sync" @@ -212,17 +214,97 @@ func TestSyncPoolZero(t *testing.T) { } func TestSyncOnceZero(t *testing.T) { - src, dst := new(sync.Once), new(sync.Once) - src.Do(func() {}) - dst.Do(func() {}) - roundtrip(t, src, dst) - calls := 0 - dst.Do(func() { calls++ }) - dst.Do(func() { calls++ }) - src.Do(func() { t.Fatal("snapshot reset the source Once") }) - if calls != 1 { - t.Fatal("restored Once did not start fresh") + for _, status := range []string{"zero", "done", "panic"} { + t.Run(status, func(t *testing.T) { + for _, initialized := range []bool{false, true} { + src, dst := new(sync.Once), new(sync.Once) + switch status { + case "done": + src.Do(func() {}) + case "panic": + func() { + defer func() { + if got := recover(); got != "once panic" { + t.Fatalf("unexpected panic: %v", got) + } + }() + src.Do(func() { panic("once panic") }) + }() + } + if initialized { + dst.Do(func() {}) + } + // A reused destination must lose its old mutex state as well as done. + mutex := reflectValueRWAddr(reflect.ValueOf(dst).Elem().FieldByName("m")).Interface().(*sync.Mutex) + mutex.Lock() + roundtrip(t, src, dst) + if !mutex.TryLock() { + t.Fatal("restored Once retained the destination mutex state") + } + mutex.Unlock() + var calls atomic.Int32 + var workers sync.WaitGroup + for range 8 { + workers.Go(func() { dst.Do(func() { calls.Add(1) }) }) + } + workers.Wait() + want := int32(0) + if status == "zero" { + want = 1 + } + if calls.Load() != want { + t.Fatalf("destination initialized=%v: got %d calls, want %d", initialized, calls.Load(), want) + } + sourceCalls := int32(0) + src.Do(func() { sourceCalls++ }) + if sourceCalls != want { + t.Fatal("snapshot changed the source Once") + } + } + }) } + t.Run("writeback", func(t *testing.T) { + type root struct { + Once sync.Once + Alias *sync.Once + Value int + } + host := new(root) + host.Alias = &host.Once + var guest root + var source, destination State + ctx := context.Background() + mem := make([]byte, 1<<20) + for range 2 { + var input bytes.Buffer + if _, _, err := source.SaveTo(ctx, &input, host); err != nil { + t.Fatal(err) + } + if _, err := destination.Load(ctx, input.Bytes(), &guest); err != nil { + t.Fatal(err) + } + if guest.Alias != &guest.Once { + t.Fatal("Once lost its alias in the guest") + } + before := host.Value + guest.Alias.Do(func() { guest.Value++ }) + guest.Once.Do(func() { guest.Value++ }) + if guest.Value != 1 || host.Value != before { + t.Fatal("guest repeated initialization or modified the host") + } + n, _, err := destination.Save(ctx, mem, &guest) + if err != nil { + t.Fatal(err) + } + if _, err := source.Load(ctx, mem[:n], host); err != nil { + t.Fatal(err) + } + if host.Alias != &host.Once || host.Value != 1 { + t.Fatal("writeback lost Once identity or initialized data") + } + host.Alias.Do(func() { t.Fatal("writeback lost Once completion") }) + } + }) } func TestSyncWaitGroupZero(t *testing.T) {