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
8 changes: 8 additions & 0 deletions internal/state/roundtrip.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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 {
Expand Down
65 changes: 65 additions & 0 deletions internal/state/session_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
}
}
Loading