diff --git a/hashmap.go b/hashmap.go index 3bf83c1..c0b8694 100644 --- a/hashmap.go +++ b/hashmap.go @@ -48,6 +48,10 @@ func (m *Map[Key, Value]) Get(key Key) (Value, bool) { hash := m.hasher(key) for element := m.store.Load().item(hash); element != nil; element = element.Next() { + if element.deleted.Load() != 0 { + continue + } + if element.keyHash == hash && element.key == key { return element.Value(), true } diff --git a/hashmap_test.go b/hashmap_test.go index f0eb327..5d356aa 100644 --- a/hashmap_test.go +++ b/hashmap_test.go @@ -482,3 +482,27 @@ func TestGetOrInsertHangIssue67(_ *testing.T) { wg.Wait() } + +func TestGetOrInsertConcurrentIssue81(t *testing.T) { + m := New[int, bool]() + var wg sync.WaitGroup + n := 1000 + + for i := range n { + wg.Add(1) + go func(val int) { + defer wg.Done() + m.GetOrInsert(val, true) + }(i) + } + wg.Wait() + + count := 0 + m.Range(func(_ int, value bool) bool { + count++ + return true + }) + + assert.Equal(t, n, count) + assert.Equal(t, uintptr(n), m.Len()) +} diff --git a/list.go b/list.go index 596b2cf..e6e0626 100644 --- a/list.go +++ b/list.go @@ -76,7 +76,7 @@ func (l *List[Key, Value]) Delete(element *ListElement[Key, Value]) { } func (l *List[Key, Value]) search(searchStart *ListElement[Key, Value], hash uintptr, key Key) (left, found, right *ListElement[Key, Value]) { - if searchStart != nil && hash < searchStart.keyHash { // key would remain left from item? { + if searchStart != nil && (hash < searchStart.keyHash || searchStart.deleted.Load() != 0) { // key would remain left from item or item deleted? searchStart = nil // start search at head } @@ -92,7 +92,9 @@ func (l *List[Key, Value]) search(searchStart *ListElement[Key, Value], hash uin for { if hash == found.keyHash && key == found.key { // key hash already exists, compare keys - return nil, found, nil + if found.deleted.Load() == 0 { + return nil, found, nil + } } if hash < found.keyHash { // new item needs to be inserted before the found value @@ -116,12 +118,21 @@ func (l *List[Key, Value]) insertAt(element, left, right *ListElement[Key, Value left = l.head } + if left != l.head && left.deleted.Load() != 0 { + return false // left was deleted concurrently + } + element.next.Store(right) if !left.next.CompareAndSwap(right, element) { return false // item was modified concurrently } + if left != l.head && left.deleted.Load() != 0 { + element.deleted.Store(1) + return false // left was deleted concurrently + } + l.count.Add(1) return true }