From d89c178821645ce6585950fb3ca9f5157ff8e02b Mon Sep 17 00:00:00 2001 From: jdl Date: Sun, 14 Jun 2026 17:31:52 +0200 Subject: [PATCH] Initial commit. --- go.mod | 3 + keyedmutex.go | 82 ++++++++++++------------ keyedmutex_test.go | 154 +++++++++++++++++++++------------------------ 3 files changed, 116 insertions(+), 123 deletions(-) create mode 100644 go.mod diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..5cb6856 --- /dev/null +++ b/go.mod @@ -0,0 +1,3 @@ +module git.crumpington.com/lib/keyedmutex + +go 1.25.1 diff --git a/keyedmutex.go b/keyedmutex.go index 699e4b1..7bc8494 100644 --- a/keyedmutex.go +++ b/keyedmutex.go @@ -1,67 +1,69 @@ package keyedmutex import ( - "container/list" "sync" ) +type keyedLock struct { + lock sync.Mutex + count int +} + type KeyedMutex[K comparable] struct { - mu *sync.Mutex - waitList map[K]*list.List + mu sync.Mutex + byKey map[K]*keyedLock } -func New[K comparable]() KeyedMutex[K] { - return KeyedMutex[K]{ - mu: new(sync.Mutex), - waitList: map[K]*list.List{}, +func New[K comparable]() *KeyedMutex[K] { + return &KeyedMutex[K]{ + byKey: map[K]*keyedLock{}, } } -func (m KeyedMutex[K]) Lock(key K) { - if ch := m.lock(key); ch != nil { - <-ch - } -} - -func (m KeyedMutex[K]) lock(key K) chan struct{} { +func (m *KeyedMutex[K]) getLock(key K) *keyedLock { m.mu.Lock() defer m.mu.Unlock() - if waitList, ok := m.waitList[key]; ok { - ch := make(chan struct{}) - waitList.PushBack(ch) - return ch + item, ok := m.byKey[key] + if !ok { + item = &keyedLock{} + m.byKey[key] = item } - - m.waitList[key] = list.New() - return nil + item.count++ + return item } -func (m KeyedMutex[K]) TryLock(key K) bool { +func (m *KeyedMutex[K]) release(key K, unlock bool) { m.mu.Lock() defer m.mu.Unlock() - if _, ok := m.waitList[key]; ok { - return false - } - - m.waitList[key] = list.New() - return true -} - -func (m KeyedMutex[K]) Unlock(key K) { - m.mu.Lock() - defer m.mu.Unlock() - - waitList, ok := m.waitList[key] + item, ok := m.byKey[key] if !ok { panic("unlock of unlocked mutex") } - if waitList.Len() == 0 { - delete(m.waitList, key) - } else { - ch := waitList.Remove(waitList.Front()).(chan struct{}) - ch <- struct{}{} + item.count-- + if unlock { + item.lock.Unlock() + } + + if item.count == 0 { + delete(m.byKey, key) } } + +func (m *KeyedMutex[K]) Lock(key K) { + m.getLock(key).lock.Lock() +} + +func (m *KeyedMutex[K]) TryLock(key K) bool { + if ok := m.getLock(key).lock.TryLock(); !ok { + m.release(key, false) + return false + } + return true +} + +func (m *KeyedMutex[K]) Unlock(key K) { + m.release(key, true) +} diff --git a/keyedmutex_test.go b/keyedmutex_test.go index 14fdaf0..da0ee2f 100644 --- a/keyedmutex_test.go +++ b/keyedmutex_test.go @@ -1,81 +1,92 @@ package keyedmutex - import ( "sync" "testing" - "time" ) -func TestKeyedMutex(t *testing.T) { - checkState := func(t *testing.T, m KeyedMutex[string], keys ...string) { - if len(m.waitList) != len(keys) { - t.Fatal(m.waitList, keys) - } - - for _, key := range keys { - if _, ok := m.waitList[key]; !ok { - t.Fatal(key) - } - } - } - +func TestLock_CleansUpMap(t *testing.T) { m := New[string]() - checkState(t, m) - m.Lock("a") - checkState(t, m, "a") - m.Lock("b") - checkState(t, m, "a", "b") - m.Lock("c") - checkState(t, m, "a", "b", "c") - - if m.TryLock("a") { - t.Fatal("a") - } - if m.TryLock("b") { - t.Fatal("b") - } - if m.TryLock("c") { - t.Fatal("c") - } - - if !m.TryLock("d") { - t.Fatal("d") - } - - checkState(t, m, "a", "b", "c", "d") - - if !m.TryLock("e") { - t.Fatal("e") - } - checkState(t, m, "a", "b", "c", "d", "e") - - m.Unlock("c") - checkState(t, m, "a", "b", "d", "e") m.Unlock("a") - checkState(t, m, "b", "d", "e") - m.Unlock("e") - checkState(t, m, "b", "d") + if len(m.byKey) != 0 { + t.Fatalf("expected empty byKey, got %v", m.byKey) + } +} - wg := sync.WaitGroup{} - for i := 0; i < 8; i++ { +func TestTryLock_SucceedsOnFreeKey(t *testing.T) { + m := New[string]() + if !m.TryLock("a") { + t.Fatal("expected TryLock to succeed on free key") + } + m.Unlock("a") +} + +func TestTryLock_FailsOnLockedKey(t *testing.T) { + m := New[string]() + m.Lock("a") + if m.TryLock("a") { + t.Fatal("expected TryLock to fail on locked key") + } + m.Unlock("a") +} + +func TestTryLock_FailureDoesNotCorruptLock(t *testing.T) { + // A failed TryLock must not call Unlock on the key's mutex. + // The old bug did this unconditionally, so the Unlock below would + // double-unlock and panic. + m := New[string]() + m.Lock("a") + m.TryLock("a") // must return false and leave the lock intact + m.Unlock("a") // must not panic +} + +func TestTryLock_FailureDecrementsCount(t *testing.T) { + // A failed TryLock must undo its getLock increment so that the + // original holder's Unlock cleans up the map entry. + m := New[string]() + m.Lock("a") + m.TryLock("a") // fails; a leaked count would leave a stale map entry + m.Unlock("a") + if len(m.byKey) != 0 { + t.Fatalf("expected empty byKey after unlock, got %v", m.byKey) + } +} + +func TestMultipleKeys_AreIndependent(t *testing.T) { + m := New[string]() + m.Lock("a") + if !m.TryLock("b") { + t.Fatal("expected TryLock on different key to succeed while 'a' is locked") + } + m.Unlock("b") + m.Unlock("a") +} + +func TestConcurrentLock_MutualExclusion(t *testing.T) { + m := New[string]() + const N = 100 + var wg sync.WaitGroup + var shared int // intentionally non-atomic: race detector catches improper access + + for range N { wg.Add(1) go func() { defer wg.Done() - m.Lock("b") - m.Unlock("b") + m.Lock("a") + shared++ + m.Unlock("a") }() } - time.Sleep(100 * time.Millisecond) - m.Unlock("b") wg.Wait() - checkState(t, m, "d") - - m.Unlock("d") - checkState(t, m) + if shared != N { + t.Fatalf("expected %d, got %d", N, shared) + } + if len(m.byKey) != 0 { + t.Fatalf("expected empty byKey after all goroutines done, got %v", m.byKey) + } } func TestKeyedMutex_unlockUnlocked(t *testing.T) { @@ -93,31 +104,8 @@ func BenchmarkUncontendedMutex(b *testing.B) { m := New[string]() key := "xyz" - for i := 0; i < b.N; i++ { + for b.Loop() { m.Lock(key) m.Unlock(key) } } - -func BenchmarkContendedMutex(b *testing.B) { - m := New[string]() - key := "xyz" - - m.Lock(key) - - wg := sync.WaitGroup{} - for i := 0; i < b.N; i++ { - wg.Add(1) - go func() { - defer wg.Done() - m.Lock(key) - m.Unlock(key) - }() - } - - time.Sleep(time.Second) - - b.ResetTimer() - m.Unlock(key) - wg.Wait() -}