diff --git a/src/sync/map.go b/src/sync/map.go index cd8a1967d..8b5c0cff7 100644 --- a/src/sync/map.go +++ b/src/sync/map.go @@ -70,3 +70,15 @@ func (m *Map) Range(f func(key, value interface{}) bool) { } } } + +// Swap replaces the value for the given key, and returns the old value if any. +func (m *Map) Swap(key, value any) (previous any, loaded bool) { + m.lock.Lock() + defer m.lock.Unlock() + if m.m == nil { + m.m = make(map[interface{}]interface{}) + } + previous, loaded = m.m[key] + m.m[key] = value + return +} diff --git a/src/sync/map_test.go b/src/sync/map_test.go index f493bdfb5..a41faa40f 100644 --- a/src/sync/map_test.go +++ b/src/sync/map_test.go @@ -17,3 +17,22 @@ func TestMapLoadAndDelete(t *testing.T) { t.Errorf("LoadAndDelete returned %v, %v, want nil, false", v, ok) } } + +func TestMapSwap(t *testing.T) { + var sm sync.Map + sm.Store("present", "value") + + if v, ok := sm.Swap("present", "value2"); !ok || v != "value" { + t.Errorf("Swap returned %v, %v, want value, true", v, ok) + } + if v, ok := sm.Load("present"); !ok || v != "value2" { + t.Errorf("Load after Swap returned %v, %v, want value2, true", v, ok) + } + + if v, ok := sm.Swap("new", "foo"); ok || v != nil { + t.Errorf("Swap returned %v, %v, want nil, false", v, ok) + } + if v, ok := sm.Load("present"); !ok || v != "value2" { + t.Errorf("Load after Swap returned %v, %v, want foo, true", v, ok) + } +}