From 3c19b8257ef1cf096d7dbd81fdd982fb4ccc4d0e Mon Sep 17 00:00:00 2001 From: Matthias Bertschy Date: Tue, 15 Sep 2026 12:51:53 +0200 Subject: [PATCH 1/2] fix: synchronize backing-map access in concurrent maps Signed-off-by: Matthias Bertschy --- .github/workflows/go.yml | 2 +- safe_concurrency_test.go | 100 +++++++++++++++++++++++++++++++++++++++ safe_map.go | 27 ++--------- safe_slice_map.go | 6 +-- 4 files changed, 107 insertions(+), 28 deletions(-) create mode 100644 safe_concurrency_test.go diff --git a/.github/workflows/go.yml b/.github/workflows/go.yml index 359f314..8e27690 100644 --- a/.github/workflows/go.yml +++ b/.github/workflows/go.yml @@ -24,7 +24,7 @@ jobs: run: go build -v ./... - name: Test - run: go test -v -cover ./... -coverprofile coverage.out -coverpkg ./... + run: go test -race -v -cover ./... -coverprofile coverage.out -coverpkg ./... - name: Upload coverage to Codecov run: bash <(curl -s https://codecov.io/bash) diff --git a/safe_concurrency_test.go b/safe_concurrency_test.go new file mode 100644 index 0000000..ea20d8e --- /dev/null +++ b/safe_concurrency_test.go @@ -0,0 +1,100 @@ +package maps + +import ( + "sync" + "testing" +) + +func TestSafeMapsConcurrentAccess(t *testing.T) { + constructors := []struct { + name string + new func() MapI[string, int] + }{ + {"SafeMap", func() MapI[string, int] { return new(SafeMap[string, int]) }}, + {"SafeSliceMap", func() MapI[string, int] { return new(SafeSliceMap[string, int]) }}, + } + operations := []struct { + name string + run func(MapI[string, int]) + }{ + {"Load", func(m MapI[string, int]) { m.Load("key") }}, + {"Get", func(m MapI[string, int]) { m.Get("key") }}, + {"Has", func(m MapI[string, int]) { m.Has("key") }}, + {"Keys", func(m MapI[string, int]) { m.Keys() }}, + {"Values", func(m MapI[string, int]) { m.Values() }}, + {"Len", func(m MapI[string, int]) { m.Len() }}, + {"Range", func(m MapI[string, int]) { + m.Range(func(string, int) bool { return true }) + }}, + {"All", func(m MapI[string, int]) { + for range m.All() { + } + }}, + {"KeysIter", func(m MapI[string, int]) { + for range m.KeysIter() { + } + }}, + {"ValuesIter", func(m MapI[string, int]) { + for range m.ValuesIter() { + } + }}, + {"Clear", func(m MapI[string, int]) { m.Clear() }}, + {"Copy", func(m MapI[string, int]) { m.Copy(StdMap[string, int]{"key": 1}) }}, + } + + for _, constructor := range constructors { + t.Run(constructor.name, func(t *testing.T) { + for _, operation := range operations { + t.Run(operation.name, func(t *testing.T) { + m := constructor.new() + start := make(chan struct{}) + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + <-start + for range 256 { + operation.run(m) + } + }() + go func() { + defer wg.Done() + <-start + for range 256 { + m.Set("key", 1) + m.Clear() + } + }() + close(start) + wg.Wait() + + // The map must remain usable after concurrent clearing and initialization. + m.Clear() + m.Copy(StdMap[string, int]{"key": 1}) + if value, ok := m.Load("key"); !ok || value != 1 { + t.Fatalf("Load(key) = (%d, %t), want (1, true)", value, ok) + } + }) + } + }) + } +} + +func TestSafeMapsNilIteration(t *testing.T) { + var m *SafeMap[string, int] + var sm *SafeSliceMap[string, int] + unexpected := func(string, int) bool { + t.Fatal("nil map must not invoke the callback") + return false + } + m.Range(unexpected) + m.All()(unexpected) + sm.Range(unexpected) + sm.All()(unexpected) + for range sm.KeysIter() { + t.Fatal("nil map must not yield keys") + } + for range sm.ValuesIter() { + t.Fatal("nil map must not yield values") + } +} diff --git a/safe_map.go b/safe_map.go index b391ee2..e0c3b01 100644 --- a/safe_map.go +++ b/safe_map.go @@ -35,9 +35,6 @@ func NewSafeMap[K comparable, V any](sources ...map[K]V) *SafeMap[K, V] { // Clear resets the map to an empty map. func (m *SafeMap[K, V]) Clear() { - if m.items == nil { - return - } m.mu.Lock() m.items = nil m.mu.Unlock() @@ -69,9 +66,6 @@ func (m *SafeMap[K, V]) Has(k K) (exists bool) { // Load returns the value based on its key, and a boolean indicating whether it exists in the map. // This is the same interface as sync.Map.Load(). func (m *SafeMap[K, V]) Load(k K) (v V, ok bool) { - if m.items == nil { - return - } m.mu.RLock() if m.items != nil { v, ok = m.items[k] @@ -91,9 +85,6 @@ func (m *SafeMap[K, V]) Delete(k K) (v V) { // Values returns a slice of the values. It will return a nil slice if the map is empty. // Multiple calls to Values will result in the same list of values, but may be in a different order. func (m *SafeMap[K, V]) Values() (v []V) { - if m.items == nil { - return - } m.mu.RLock() v = m.items.Values() m.mu.RUnlock() @@ -103,9 +94,6 @@ func (m *SafeMap[K, V]) Values() (v []V) { // Keys returns a slice of the keys. It will return a nil slice if the map is empty. // Multiple calls to Keys will result in the same list of keys, but may be in a different order. func (m *SafeMap[K, V]) Keys() (keys []K) { - if m.items == nil { - return nil - } m.mu.RLock() keys = m.items.Keys() m.mu.RUnlock() @@ -114,9 +102,6 @@ func (m *SafeMap[K, V]) Keys() (keys []K) { // Len returns the number of items in the map func (m *SafeMap[K, V]) Len() (l int) { - if m.items == nil { - return - } m.mu.RLock() l = m.items.Len() m.mu.RUnlock() @@ -128,7 +113,7 @@ func (m *SafeMap[K, V]) Len() (l int) { // During this process, the map will be locked, so do not pass a function that will take // significant amounts of time, nor will call into other methods of the SafeMap which might also need a lock. func (m *SafeMap[K, V]) Range(f func(k K, v V) bool) { - if m == nil || m.items == nil { + if m == nil { return } m.mu.RLock() @@ -144,11 +129,11 @@ func (m *SafeMap[K, V]) Merge(in MapI[K, V]) { // Copy copies the keys and values of in into this map, overwriting any duplicates. func (m *SafeMap[K, V]) Copy(in MapI[K, V]) { + m.mu.Lock() + defer m.mu.Unlock() if m.items == nil { m.items = make(map[K]V, in.Len()) } - m.mu.Lock() - defer m.mu.Unlock() m.items.Copy(in) } @@ -210,9 +195,6 @@ func (m *SafeMap[K, V]) All() iter.Seq2[K, V] { // does not call back functions in SafeMap which will also require a lock. func (m *SafeMap[K, V]) KeysIter() iter.Seq[K] { return func(yield func(K) bool) { - if m.items == nil { - return - } m.mu.RLock() defer m.mu.RUnlock() for k := range m.items { @@ -228,9 +210,6 @@ func (m *SafeMap[K, V]) KeysIter() iter.Seq[K] { // not attempt to call other functions in SafeMap which also need a lock. func (m *SafeMap[K, V]) ValuesIter() iter.Seq[V] { return func(yield func(V) bool) { - if m.items == nil { - return - } m.mu.RLock() defer m.mu.RUnlock() for _, v := range m.items { diff --git a/safe_slice_map.go b/safe_slice_map.go index a9990cd..80f3679 100644 --- a/safe_slice_map.go +++ b/safe_slice_map.go @@ -191,7 +191,7 @@ func (m *SafeSliceMap[K, V]) Copy(in MapI[K, V]) { // The workaround is to call Keys() and iterate over the returned copy of the keys, but making sure // your function can handle the situation where the key no longer exists in the slice. func (m *SafeSliceMap[K, V]) Range(f func(key K, value V) bool) { - if m == nil || m.sm.items == nil { // prevent unnecessary lock + if m == nil { return } m.mu.RLock() @@ -244,7 +244,7 @@ func (m *SafeSliceMap[K, V]) All() iter.Seq2[K, V] { // significant amounts of time, nor will call into other methods of the SafeSliceMap which might also need a lock. func (m *SafeSliceMap[K, V]) KeysIter() iter.Seq[K] { return func(yield func(K) bool) { - if m == nil || m.sm.items == nil { + if m == nil { return } m.mu.RLock() @@ -258,7 +258,7 @@ func (m *SafeSliceMap[K, V]) KeysIter() iter.Seq[K] { // significant amounts of time, nor will call into other methods of the SafeSliceMap which might also need a lock. func (m *SafeSliceMap[K, V]) ValuesIter() iter.Seq[V] { return func(yield func(V) bool) { - if m == nil || m.sm.items == nil { + if m == nil { return } m.mu.RLock() From e9181ad404216d4e1523d9ac2790c3c91e8bdd7c Mon Sep 17 00:00:00 2001 From: Matthias Bertschy Date: Tue, 15 Sep 2026 13:13:45 +0200 Subject: [PATCH 2/2] perf: preserve empty-map fast paths with atomic initialization state Signed-off-by: Matthias Bertschy --- safe_benchmark_test.go | 108 ++++++++++++++++++++++++ safe_concurrency_test.go | 7 ++ safe_initialization_test.go | 162 ++++++++++++++++++++++++++++++++++++ safe_map.go | 40 ++++++++- safe_slice_map.go | 29 +++++-- 5 files changed, 338 insertions(+), 8 deletions(-) create mode 100644 safe_benchmark_test.go create mode 100644 safe_initialization_test.go diff --git a/safe_benchmark_test.go b/safe_benchmark_test.go new file mode 100644 index 0000000..eab4c24 --- /dev/null +++ b/safe_benchmark_test.go @@ -0,0 +1,108 @@ +package maps + +import ( + "runtime" + "testing" +) + +var emptyBenchValue int +var emptyBenchFound bool + +// BenchmarkEmptyLoad distinguishes nil backing storage from an allocated empty map. +func BenchmarkEmptyLoad(b *testing.B) { + for _, state := range []string{"Zero", "Cleared", "Allocated", "Populated"} { + b.Run(state, func(b *testing.B) { + m := new(SafeMap[string, int]) + switch state { + case "Cleared": + m.Set("key", 1) + m.Clear() + case "Allocated": + m.Set("key", 1) + m.Delete("key") + case "Populated": + m.Set("key", 1) + } + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + emptyBenchValue, emptyBenchFound = m.Load("key") + } + }) + } +} + +func BenchmarkZeroLen(b *testing.B) { + m := new(SafeMap[string, int]) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + emptyBenchValue = m.Len() + } +} + +func BenchmarkZeroClear(b *testing.B) { + m := new(SafeMap[string, int]) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + m.Clear() + } +} + +func BenchmarkZeroLoadParallel(b *testing.B) { + m := new(SafeMap[string, int]) + b.ReportAllocs() + b.ResetTimer() + b.RunParallel(func(pb *testing.PB) { + var v int + var ok bool + for pb.Next() { + v, ok = m.Load("missing") + } + runtime.KeepAlive(v) + runtime.KeepAlive(ok) + }) +} + +func BenchmarkPopulatedStore(b *testing.B) { + m := new(SafeMap[string, int]) + m.Set("key", 1) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + m.Set("key", i) + } +} + +func BenchmarkSetClearCycle(b *testing.B) { + m := new(SafeMap[string, int]) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + m.Set("key", i) + m.Clear() + } +} + +func BenchmarkSliceZeroRange(b *testing.B) { + m := new(SafeSliceMap[string, int]) + yield := func(string, int) bool { return true } + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + m.Range(yield) + } +} + +func BenchmarkSliceZeroRangeParallel(b *testing.B) { + m := new(SafeSliceMap[string, int]) + yield := func(string, int) bool { return true } + b.ReportAllocs() + b.ResetTimer() + b.RunParallel(func(pb *testing.PB) { + for pb.Next() { + m.Range(yield) + } + }) +} diff --git a/safe_concurrency_test.go b/safe_concurrency_test.go index ea20d8e..8369b05 100644 --- a/safe_concurrency_test.go +++ b/safe_concurrency_test.go @@ -1,6 +1,7 @@ package maps import ( + "encoding/json" "sync" "testing" ) @@ -40,6 +41,12 @@ func TestSafeMapsConcurrentAccess(t *testing.T) { }}, {"Clear", func(m MapI[string, int]) { m.Clear() }}, {"Copy", func(m MapI[string, int]) { m.Copy(StdMap[string, int]{"key": 1}) }}, + {"UnmarshalJSON", func(m MapI[string, int]) { + _ = m.(json.Unmarshaler).UnmarshalJSON([]byte(`{"key":1}`)) + }}, + {"UnmarshalNull", func(m MapI[string, int]) { + _ = m.(json.Unmarshaler).UnmarshalJSON([]byte(`null`)) + }}, } for _, constructor := range constructors { diff --git a/safe_initialization_test.go b/safe_initialization_test.go new file mode 100644 index 0000000..19dd339 --- /dev/null +++ b/safe_initialization_test.go @@ -0,0 +1,162 @@ +package maps + +import ( + "encoding" + "encoding/json" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func assertSafeMapContents(t *testing.T, m MapI[string, int], want map[string]int) { + t.Helper() + assert.Equal(t, len(want), m.Len()) + for key, value := range want { + got, ok := m.Load(key) + assert.True(t, ok) + assert.Equal(t, value, got) + assert.Equal(t, value, m.Get(key)) + assert.True(t, m.Has(key)) + } + _, ok := m.Load("missing") + assert.False(t, ok) + got := make(map[string]int) + m.Range(func(k string, v int) bool { got[k] = v; return true }) + assert.Equal(t, want, got) + got = make(map[string]int) + for k, v := range m.All() { + got[k] = v + } + assert.Equal(t, want, got) + var keys []string + var values []int + for k := range m.KeysIter() { + keys = append(keys, k) + } + for v := range m.ValuesIter() { + values = append(values, v) + } + expected := StdMap[string, int](want) + assert.ElementsMatch(t, expected.Keys(), keys) + assert.ElementsMatch(t, expected.Values(), values) + assert.ElementsMatch(t, expected.Keys(), m.Keys()) + assert.ElementsMatch(t, expected.Values(), m.Values()) +} + +func TestSafeMapsInitialization(t *testing.T) { + for _, constructor := range []struct { + name string + new func() MapI[string, int] + }{ + {"SafeMap", func() MapI[string, int] { return new(SafeMap[string, int]) }}, + {"SafeSliceMap", func() MapI[string, int] { return new(SafeSliceMap[string, int]) }}, + } { + t.Run(constructor.name, func(t *testing.T) { + source := constructor.new() + source.Set("key", 1) + binary, err := source.(encoding.BinaryMarshaler).MarshalBinary() + require.NoError(t, err) + for _, populate := range []struct { + name string + run func(*testing.T, MapI[string, int]) + }{ + {"Set", func(t *testing.T, m MapI[string, int]) { m.Set("key", 1) }}, + {"Copy", func(t *testing.T, m MapI[string, int]) { m.Copy(StdMap[string, int]{"key": 1}) }}, + {"Merge", func(t *testing.T, m MapI[string, int]) { m.Merge(StdMap[string, int]{"key": 1}) }}, + {"Insert", func(t *testing.T, m MapI[string, int]) { + m.Copy(StdMap[string, int]{}) + m.Insert(StdMap[string, int]{"key": 1}.All()) + }}, + {"JSON", func(t *testing.T, m MapI[string, int]) { + require.NoError(t, m.(json.Unmarshaler).UnmarshalJSON([]byte(`{"key":1}`))) + }}, + {"Binary", func(t *testing.T, m MapI[string, int]) { + require.NoError(t, m.(encoding.BinaryUnmarshaler).UnmarshalBinary(binary)) + }}, + } { + t.Run(populate.name, func(t *testing.T) { + m := constructor.new() + assertSafeMapContents(t, m, map[string]int{}) + for range 2 { + populate.run(t, m) + assertSafeMapContents(t, m, map[string]int{"key": 1}) + m.Clear() + assertSafeMapContents(t, m, map[string]int{}) + } + }) + } + }) + } +} + +func TestSafeMapsConstructedContents(t *testing.T) { + source := StdMap[string, int]{"key": 1} + for _, m := range []MapI[string, int]{ + NewSafeMap[string, int](source), + NewSafeSliceMap[string, int](source), + NewSafeMap[string, int](source).Clone(), + NewSafeSliceMap[string, int](source).Clone(), + CollectSafeMap(source.All()), + CollectSafeSliceMap(source.All()), + } { + assertSafeMapContents(t, m, map[string]int{"key": 1}) + m.Clear() + assertSafeMapContents(t, m, map[string]int{}) + m.Set("key", 2) + assertSafeMapContents(t, m, map[string]int{"key": 2}) + } +} + +func TestSafeMapPartialDecode(t *testing.T) { + var m SafeMap[string, int] + err := m.UnmarshalJSON([]byte(`{"key":1,"bad":"not an int"}`)) + require.Error(t, err) + assertSafeMapContents(t, &m, map[string]int{"key": 1, "bad": 0}) + require.NoError(t, m.UnmarshalJSON([]byte(`null`))) + assertSafeMapContents(t, &m, map[string]int{}) + m.Set("key", 2) + assertSafeMapContents(t, &m, map[string]int{"key": 2}) + require.Error(t, m.UnmarshalBinary([]byte("invalid gob"))) + assertSafeMapContents(t, &m, map[string]int{}) + m.Set("key", 3) + assertSafeMapContents(t, &m, map[string]int{"key": 3}) +} + +func TestSafeSliceMapInitializationTransitions(t *testing.T) { + var m SafeSliceMap[string, int] + for range 2 { + m.SetAt(0, "key", 1) + assertSafeMapContents(t, &m, map[string]int{"key": 1}) + m.Clear() + assertSafeMapContents(t, &m, map[string]int{}) + } + m.Set("key", 2) + require.Error(t, m.UnmarshalJSON([]byte(`{"key":"not an int"}`))) + assertSafeMapContents(t, &m, map[string]int{"key": 2}) + require.Error(t, m.UnmarshalBinary([]byte("invalid gob"))) + assertSafeMapContents(t, &m, map[string]int{"key": 2}) + require.NoError(t, m.UnmarshalJSON([]byte(`null`))) + assertSafeMapContents(t, &m, map[string]int{}) + m.SetAt(0, "key", 3) + assertSafeMapContents(t, &m, map[string]int{"key": 3}) +} + +func TestSafeSliceMapInitializationAfterPanic(t *testing.T) { + for _, operation := range []struct { + name string + set func(*SafeSliceMap[any, int], any, int) + }{ + {"Set", func(m *SafeSliceMap[any, int], key any, value int) { m.Set(key, value) }}, + {"SetAt", func(m *SafeSliceMap[any, int], key any, value int) { m.SetAt(0, key, value) }}, + } { + t.Run(operation.name, func(t *testing.T) { + var m SafeSliceMap[any, int] + assert.Panics(t, func() { operation.set(&m, []int{1}, 1) }) + operation.set(&m, "key", 2) + got := make(map[any]int) + m.Range(func(key any, value int) bool { got[key] = value; return true }) + assert.Equal(t, map[any]int{"key": 2}, got) + }) + } +} diff --git a/safe_map.go b/safe_map.go index e0c3b01..aaa619b 100644 --- a/safe_map.go +++ b/safe_map.go @@ -3,6 +3,7 @@ package maps import ( "iter" "sync" + "sync/atomic" ) // SafeMap is a go map that is safe for concurrent use and that uses a standard set of functions @@ -21,6 +22,9 @@ import ( type SafeMap[K comparable, V any] struct { mu sync.RWMutex items StdMap[K, V] + // initialized permits an empty fast path without reading items outside mu. + // Writers publish initialization/reset before unlocking; actual map access still requires mu. + initialized atomic.Bool } // NewSafeMap creates a new SafeMap. @@ -35,8 +39,12 @@ func NewSafeMap[K comparable, V any](sources ...map[K]V) *SafeMap[K, V] { // Clear resets the map to an empty map. func (m *SafeMap[K, V]) Clear() { + if !m.initialized.Load() { + return + } m.mu.Lock() m.items = nil + m.initialized.Store(false) m.mu.Unlock() } @@ -45,6 +53,7 @@ func (m *SafeMap[K, V]) Set(k K, v V) { m.mu.Lock() if m.items == nil { m.items = map[K]V{k: v} + m.initialized.Store(true) } else { m.items[k] = v } @@ -66,6 +75,9 @@ func (m *SafeMap[K, V]) Has(k K) (exists bool) { // Load returns the value based on its key, and a boolean indicating whether it exists in the map. // This is the same interface as sync.Map.Load(). func (m *SafeMap[K, V]) Load(k K) (v V, ok bool) { + if !m.initialized.Load() { + return + } m.mu.RLock() if m.items != nil { v, ok = m.items[k] @@ -85,6 +97,9 @@ func (m *SafeMap[K, V]) Delete(k K) (v V) { // Values returns a slice of the values. It will return a nil slice if the map is empty. // Multiple calls to Values will result in the same list of values, but may be in a different order. func (m *SafeMap[K, V]) Values() (v []V) { + if !m.initialized.Load() { + return + } m.mu.RLock() v = m.items.Values() m.mu.RUnlock() @@ -94,6 +109,9 @@ func (m *SafeMap[K, V]) Values() (v []V) { // Keys returns a slice of the keys. It will return a nil slice if the map is empty. // Multiple calls to Keys will result in the same list of keys, but may be in a different order. func (m *SafeMap[K, V]) Keys() (keys []K) { + if !m.initialized.Load() { + return nil + } m.mu.RLock() keys = m.items.Keys() m.mu.RUnlock() @@ -102,6 +120,9 @@ func (m *SafeMap[K, V]) Keys() (keys []K) { // Len returns the number of items in the map func (m *SafeMap[K, V]) Len() (l int) { + if !m.initialized.Load() { + return + } m.mu.RLock() l = m.items.Len() m.mu.RUnlock() @@ -113,7 +134,7 @@ func (m *SafeMap[K, V]) Len() (l int) { // During this process, the map will be locked, so do not pass a function that will take // significant amounts of time, nor will call into other methods of the SafeMap which might also need a lock. func (m *SafeMap[K, V]) Range(f func(k K, v V) bool) { - if m == nil { + if m == nil || !m.initialized.Load() { return } m.mu.RLock() @@ -133,6 +154,7 @@ func (m *SafeMap[K, V]) Copy(in MapI[K, V]) { defer m.mu.Unlock() if m.items == nil { m.items = make(map[K]V, in.Len()) + m.initialized.Store(true) } m.items.Copy(in) } @@ -156,7 +178,9 @@ func (m *SafeMap[K, V]) MarshalBinary() ([]byte, error) { func (m *SafeMap[K, V]) UnmarshalBinary(data []byte) (err error) { m.mu.Lock() defer m.mu.Unlock() - return m.items.UnmarshalBinary(data) + err = m.items.UnmarshalBinary(data) + m.initialized.Store(m.items != nil) + return } // MarshalJSON implements the json.Marshaler interface to convert the map into a JSON object. @@ -171,7 +195,9 @@ func (m *SafeMap[K, V]) MarshalJSON() (out []byte, err error) { func (m *SafeMap[K, V]) UnmarshalJSON(in []byte) (err error) { m.mu.Lock() defer m.mu.Unlock() - return m.items.UnmarshalJSON(in) + err = m.items.UnmarshalJSON(in) + m.initialized.Store(m.items != nil) + return } // String outputs the map as a string. @@ -195,6 +221,9 @@ func (m *SafeMap[K, V]) All() iter.Seq2[K, V] { // does not call back functions in SafeMap which will also require a lock. func (m *SafeMap[K, V]) KeysIter() iter.Seq[K] { return func(yield func(K) bool) { + if !m.initialized.Load() { + return + } m.mu.RLock() defer m.mu.RUnlock() for k := range m.items { @@ -210,6 +239,9 @@ func (m *SafeMap[K, V]) KeysIter() iter.Seq[K] { // not attempt to call other functions in SafeMap which also need a lock. func (m *SafeMap[K, V]) ValuesIter() iter.Seq[V] { return func(yield func(V) bool) { + if !m.initialized.Load() { + return + } m.mu.RLock() defer m.mu.RUnlock() for _, v := range m.items { @@ -238,6 +270,7 @@ func CollectSafeMap[K comparable, V any](seq iter.Seq2[K, V]) *SafeMap[K, V] { for k, v := range seq { m.items[k] = v } + m.initialized.Store(true) return m } @@ -248,6 +281,7 @@ func (m *SafeMap[K, V]) Clone() *SafeMap[K, V] { m.mu.RLock() defer m.mu.RUnlock() m1.items = m.items.Clone() + m1.initialized.Store(m1.items != nil) return m1 } diff --git a/safe_slice_map.go b/safe_slice_map.go index 80f3679..c5e21f9 100644 --- a/safe_slice_map.go +++ b/safe_slice_map.go @@ -6,6 +6,7 @@ import ( "slices" "strings" "sync" + "sync/atomic" ) // SafeSliceMap is a go map that uses a slice to save the order of its keys so that the map can @@ -30,6 +31,9 @@ import ( type SafeSliceMap[K comparable, V any] struct { mu sync.RWMutex sm SliceMap[K, V] + // initialized permits empty iteration without reading sm outside mu. + // Writers publish initialization/reset before unlocking; actual map access still requires mu. + initialized atomic.Bool } // NewSafeSliceMap creates a new SafeSliceMap. @@ -59,6 +63,10 @@ func (m *SafeSliceMap[K, V]) SetSortFunc(f func(key1, key2 K, val1, val2 V) bool func (m *SafeSliceMap[K, V]) Set(key K, val V) { m.mu.Lock() defer m.mu.Unlock() + if m.sm.items == nil { + // Set can allocate storage before panicking on an invalid key. + m.initialized.Store(true) + } m.sm.Set(key, val) } @@ -68,6 +76,10 @@ func (m *SafeSliceMap[K, V]) Set(key K, val V) { func (m *SafeSliceMap[K, V]) SetAt(index int, key K, val V) { m.mu.Lock() defer m.mu.Unlock() + if m.sm.items == nil { + // Publish before SetAt can allocate storage and then panic. + m.initialized.Store(true) + } m.sm.SetAt(index, key, val) } @@ -149,7 +161,9 @@ func (m *SafeSliceMap[K, V]) MarshalBinary() (data []byte, err error) { func (m *SafeSliceMap[K, V]) UnmarshalBinary(data []byte) (err error) { m.mu.Lock() defer m.mu.Unlock() - return m.sm.UnmarshalBinary(data) + err = m.sm.UnmarshalBinary(data) + m.initialized.Store(m.sm.items != nil) + return } // MarshalJSON implements the json.Marshaler interface to convert the map into a JSON object. @@ -166,7 +180,9 @@ func (m *SafeSliceMap[K, V]) MarshalJSON() (data []byte, err error) { func (m *SafeSliceMap[K, V]) UnmarshalJSON(data []byte) (err error) { m.mu.Lock() defer m.mu.Unlock() - return m.sm.UnmarshalJSON(data) + err = m.sm.UnmarshalJSON(data) + m.initialized.Store(m.sm.items != nil) + return } // Merge the given map into the current one. @@ -191,7 +207,7 @@ func (m *SafeSliceMap[K, V]) Copy(in MapI[K, V]) { // The workaround is to call Keys() and iterate over the returned copy of the keys, but making sure // your function can handle the situation where the key no longer exists in the slice. func (m *SafeSliceMap[K, V]) Range(f func(key K, value V) bool) { - if m == nil { + if m == nil || !m.initialized.Load() { return } m.mu.RLock() @@ -213,6 +229,7 @@ func (m *SafeSliceMap[K, V]) Equal(m2 MapI[K, V]) bool { func (m *SafeSliceMap[K, V]) Clear() { m.mu.Lock() m.sm.Clear() + m.initialized.Store(false) m.mu.Unlock() } @@ -244,7 +261,7 @@ func (m *SafeSliceMap[K, V]) All() iter.Seq2[K, V] { // significant amounts of time, nor will call into other methods of the SafeSliceMap which might also need a lock. func (m *SafeSliceMap[K, V]) KeysIter() iter.Seq[K] { return func(yield func(K) bool) { - if m == nil { + if m == nil || !m.initialized.Load() { return } m.mu.RLock() @@ -258,7 +275,7 @@ func (m *SafeSliceMap[K, V]) KeysIter() iter.Seq[K] { // significant amounts of time, nor will call into other methods of the SafeSliceMap which might also need a lock. func (m *SafeSliceMap[K, V]) ValuesIter() iter.Seq[V] { return func(yield func(V) bool) { - if m == nil { + if m == nil || !m.initialized.Load() { return } m.mu.RLock() @@ -285,6 +302,7 @@ func CollectSafeSliceMap[K comparable, V any](seq iter.Seq2[K, V]) *SafeSliceMap for k, v := range seq { m.sm.Set(k, v) } + m.initialized.Store(m.sm.items != nil) return m } @@ -297,6 +315,7 @@ func (m *SafeSliceMap[K, V]) Clone() *SafeSliceMap[K, V] { m1.sm.items = m.sm.items.Clone() m1.sm.order = slices.Clone(m.sm.order) m1.sm.lessF = m.sm.lessF + m1.initialized.Store(m1.sm.items != nil) return m1 }