From e397ac51e2dee5e96a2a1efa370f32e1650b1a10 Mon Sep 17 00:00:00 2001 From: Rick Guo Date: Mon, 21 Sep 2026 11:53:55 +0800 Subject: [PATCH] Serialize state loads across shared host objects --- internal/state/roundtrip.go | 8 +++++ internal/state/session_test.go | 65 ++++++++++++++++++++++++++++++++++ 2 files changed, 73 insertions(+) diff --git a/internal/state/roundtrip.go b/internal/state/roundtrip.go index e169671..5ac67c0 100644 --- a/internal/state/roundtrip.go +++ b/internal/state/roundtrip.go @@ -4,8 +4,13 @@ import ( "context" "io" "reflect" + "sync" ) +// Separate States can retain the same host objects. Serialize their restores, +// including validation and hooks, so two Loads cannot write the same map. +var loadMu sync.Mutex + // State retains object identities across alternating Save and Load calls. // The zero value is ready for use. Each participant owns its own State and // keeps it alive until the round trip finishes, including unlinked objects. @@ -48,6 +53,9 @@ func (s *State) save(ctx context.Context, mem []byte, out io.Writer, rootPtr any // and confined to their decoded graph. Their external side effects cannot be // rolled back, and a hook that fails only during writeback can partially apply. func (s *State) Load(ctx context.Context, mem []byte, rootPtr any) (Stats, error) { + loadMu.Lock() + defer loadMu.Unlock() + ds := newDecodeState(ctx, mem) err := safely(func() { if s.saved != nil { diff --git a/internal/state/session_test.go b/internal/state/session_test.go index b4e43af..b923037 100644 --- a/internal/state/session_test.go +++ b/internal/state/session_test.go @@ -76,3 +76,68 @@ func TestStateRejectedReturnDoesNotWriteBack(t *testing.T) { t.Fatal("valid return after rejected data lost object identities") } } + +func TestStateConcurrentSharedMapLoad(t *testing.T) { + const workers, entries = 8, 1024 + shared := make(map[int]int, entries) + for i := 0; i < entries; i++ { + shared[i] = 0 + } + var hosts [workers]State + var roots [workers]map[int]int + var images [workers][]byte + ctx := context.Background() + for i := range hosts { + roots[i] = shared + mem := make([]byte, 1<<20) + n, _, err := hosts[i].Save(ctx, mem, &roots[i]) + if err != nil { + t.Fatal(err) + } + var guest State + var returned map[int]int + if _, err := guest.Load(ctx, mem[:n], &returned); err != nil { + t.Fatal(err) + } + for key := range returned { + returned[key] = i + 1 + } + n, _, err = guest.Save(ctx, mem, &returned) + if err != nil { + t.Fatal(err) + } + images[i] = mem[:n] + } + start := make(chan struct{}) + errors := make(chan error, workers) + for i := range hosts { + go func() { + <-start + _, err := hosts[i].Load(ctx, images[i], &roots[i]) + errors <- err + }() + } + close(start) + for range hosts { + if err := <-errors; err != nil { + t.Error(err) + } + } + // The last complete snapshot wins; entry locking does not merge changes + // made by different guests or prescribe which guest writes back last. + want := shared[0] + if want < 1 || want > workers || len(shared) != entries { + t.Fatalf("invalid restored map: first=%d len=%d", want, len(shared)) + } + for key := 0; key < entries; key++ { + if got, ok := shared[key]; !ok || got != want { + t.Fatalf("interleaved snapshots at key %d: got %d (present=%v), want %d", key, got, ok, want) + } + } + shared[entries] = want + for i, root := range roots { + if root[entries] != want { + t.Fatalf("state %d replaced the shared map", i) + } + } +}