diff --git a/btree/cache_internal_test.go b/btree/cache_internal_test.go index cfa4c8850..ef1817023 100644 --- a/btree/cache_internal_test.go +++ b/btree/cache_internal_test.go @@ -18,15 +18,31 @@ func Test_Get(t *testing.T) { Keys: []rune{'b'}, Pointers: []int{2}, } + n3 := Node[rune, int, int]{ + Keys: []rune{'c'}, + Pointers: []int{3}, + } + n4 := Node[rune, int, int]{ + Keys: []rune{'d'}, + Pointers: []int{4}, + } ptr1, err := st.Put(&n1) require.NoError(t, err) ptr2, err := st.Put(&n2) require.NoError(t, err) - - c := cachefor[rune, int, int](&st, 1) - n, err := st.Get(ptr1) + ptr3, err := st.Put(&n3) + require.NoError(t, err) + ptr4, err := st.Put(&n4) require.NoError(t, err) - require.NoError(t, c.Update(ptr1, n)) + + c := cachefor[rune, int, int](&st, 2) + + // warm up the cache + for _, ptrn := range []int{ptr1, ptr2, ptr3, ptr4} { + n, err := st.Get(ptrn) + require.NoError(t, err) + require.NoError(t, c.Update(ptrn, n)) + } flushErr := errors.New("Store.Update() failed") st.UpdateFn = func(ptr int, n *Node[rune, int, int]) error { return flushErr } diff --git a/btree/iter_internal_test.go b/btree/iter_internal_test.go index 2752de568..43e0a626f 100644 --- a/btree/iter_internal_test.go +++ b/btree/iter_internal_test.go @@ -504,6 +504,9 @@ func TestBackwardIter(t *testing.T) { return &st.store[ptr], nil } + // flush the cache + it.b.cache.lru.Close() + require.False(t, it.Next()) require.ErrorIs(t, it.Err(), getErr) }) diff --git a/caching/lru/lru.go b/caching/lru/lru.go index 241b196bc..5891fc77f 100644 --- a/caching/lru/lru.go +++ b/caching/lru/lru.go @@ -5,54 +5,68 @@ import ( "sync/atomic" ) -type node[K comparable] struct { +type node[K comparable, V any] struct { key K - next *node[K] + val V + next *node[K, V] + prev *node[K, V] } +// Cache implements a LRU caching strategy. All methods beside Close +// are safe for concurrent use. type Cache[K comparable, V any] struct { - mtx sync.RWMutex + mtx sync.Mutex size int target int onevict func(K, V) error - items map[K]V - head *node[K] - tail *node[K] + items map[K]*node[K, V] + head *node[K, V] + tail *node[K, V] hits atomic.Uint64 misses atomic.Uint64 } +// New constructs a cache of maximum size target. The minimum size is +// two, and any value minor than that will silently be bumped to at +// least 2. onevict is an optional callback that is called when +// evicting an item from the cache. func New[K comparable, V any](target int, onevict func(K, V) error) *Cache[K, V] { + // it's simpler to assume that we'll always have something in + // the linked list, rather than trying to cope with + // patological cases like zero or one. + target = max(target, 2) return &Cache[K, V]{ target: target, onevict: onevict, - items: make(map[K]V, target), + items: make(map[K]*node[K, V], target), } } func (c *Cache[K, V]) put(key K, val V) { c.size++ - c.items[key] = val - - n := &node[K]{key: key} + n := &node[K, V]{key: key, val: val} if c.head == nil { c.head = n c.tail = n } else { - c.tail.next = n - c.tail = n + n.next = c.head + c.head.prev = n + c.head = n } + + c.items[key] = n } -// assume that the item was just removed from the linked list -func (c *Cache[K, V]) flush(key K) error { - val := c.items[key] +// evictItem removes an item from the `items' map, without touching +// the list. +func (c *Cache[K, V]) evictItem(key K) error { + n := c.items[key] if c.onevict != nil { - if err := c.onevict(key, val); err != nil { + if err := c.onevict(key, n.val); err != nil { return err } } @@ -62,55 +76,95 @@ func (c *Cache[K, V]) flush(key K) error { return nil } +// Get retrieves a key from the cache, returning the value and true, +// or the zero value and false if not found. It also bumps the +// "recent-ness" of the key. func (c *Cache[K, V]) Get(key K) (V, bool) { - c.mtx.RLock() - val, ok := c.items[key] - c.mtx.RUnlock() + c.mtx.Lock() + n, ok := c.items[key] + if ok && n != c.head { + // move it as the first item in the list + prev := n.prev + next := n.next + prev.next = next + if next != nil { + next.prev = prev + } + + if c.tail == n { + c.tail = n.prev + c.tail.next = nil + } + + n.next = c.head + c.head.prev = n + c.head = n + n.prev = nil + } + c.mtx.Unlock() if ok { c.hits.Add(1) + return n.val, true } else { c.misses.Add(1) + var zero V + return zero, false } - - return val, ok } -// Put adds or overrides an element in the cache. +// Put adds or overrides an element in the cache, eventually evicting +// the less recent item in the cache. It can only fail if the +// `onevict` callback fails, in which case it aborts the operation. func (c *Cache[K, V]) Put(key K, val V) error { c.mtx.Lock() defer c.mtx.Unlock() if _, ok := c.items[key]; ok { - c.items[key] = val + c.items[key].val = val return nil } if c.size == c.target { - if err := c.flush(c.head.key); err != nil { + if err := c.evictItem(c.tail.key); err != nil { return err } - c.head = c.head.next + c.tail = c.tail.prev + if c.tail != nil { + c.tail.next.prev = nil + c.tail.next = nil + } } c.put(key, val) return nil } +// Close flushes all the items in the cache using the `onevict` +// callback, if provided. It is safe to reuse a cache after it has +// been closed, provided that the size and `onevict` callback are +// still fine: this method behaves like a reset. It can only fail if +// `onevict` fails, in which case it returns the last error. func (c *Cache[K, V]) Close() error { c.mtx.Lock() defer c.mtx.Unlock() var err error - for n := c.head; n != nil; n = n.next { - if e := c.flush(n.key); e != nil { - err = e + if c.onevict != nil { + for key, n := range c.items { + if e := c.onevict(key, n.val); e != nil { + err = e + } } } + clear(c.items) + c.size = 0 c.head = nil c.tail = nil return err } +// Stats return the number of cache hit, misses and the size of the +// cache. It does *not* reset them. func (c *Cache[K, V]) Stats() (hit, miss, size uint64) { return c.hits.Load(), c.misses.Load(), uint64(c.size) } diff --git a/caching/lru/lru_bench_test.go b/caching/lru/lru_bench_test.go new file mode 100644 index 000000000..f13686241 --- /dev/null +++ b/caching/lru/lru_bench_test.go @@ -0,0 +1,159 @@ +package lru + +import ( + "fmt" + "testing" +) + +func BenchmarkPut(b *testing.B) { + for _, target := range []int{16, 256, 4096} { + b.Run(fmt.Sprintf("target=%d", target), func(b *testing.B) { + cache := New[int, int](target, nil) + defer cache.Close() + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + if err := cache.Put(i, i); err != nil { + b.Fatal(err) + } + } + }) + } +} + +func BenchmarkPutUpdateExisting(b *testing.B) { + cache := New[int, int](256, nil) + defer cache.Close() + + for i := range 256 { + if err := cache.Put(i, i); err != nil { + b.Fatal(err) + } + } + + b.ReportAllocs() + for i := 0; b.Loop(); i++ { + if err := cache.Put(i%256, i); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkGetHit(b *testing.B) { + cache := New[int, int](256, nil) + defer cache.Close() + + for i := range 256 { + if err := cache.Put(i, i); err != nil { + b.Fatal(err) + } + } + + b.ReportAllocs() + for i := 0; b.Loop(); i++ { + _, _ = cache.Get(i % 256) + } +} + +func BenchmarkGetMiss(b *testing.B) { + cache := New[int, int](256, nil) + defer cache.Close() + + for i := range 256 { + if err := cache.Put(i, i); err != nil { + b.Fatal(err) + } + } + + b.ReportAllocs() + for i := 0; b.Loop(); i++ { + _, _ = cache.Get(i + 1_000_000) + } +} + +func BenchmarkGetParallel(b *testing.B) { + cache := New[int, int](256, nil) + defer cache.Close() + + for i := range 256 { + if err := cache.Put(i, i); err != nil { + b.Fatal(err) + } + } + + b.ReportAllocs() + b.ResetTimer() + b.RunParallel(func(pb *testing.PB) { + i := 0 + for pb.Next() { + _, _ = cache.Get(i % 256) + i++ + } + }) +} + +// BenchmarkPutEviction measures Put cost when every insert forces an +// eviction, i.e. the cache is kept permanently full. +func BenchmarkPutEviction(b *testing.B) { + const target = 256 + cache := New[int, int](target, nil) + defer cache.Close() + + for i := range target { + if err := cache.Put(i, i); err != nil { + b.Fatal(err) + } + } + + b.ReportAllocs() + for i := 0; b.Loop(); i++ { + if err := cache.Put(target+i, i); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkPutEvictionWithCallback(b *testing.B) { + const target = 256 + cache := New(target, func(K int, V int) error { return nil }) + defer cache.Close() + + for i := range target { + if err := cache.Put(i, i); err != nil { + b.Fatal(err) + } + } + + b.ReportAllocs() + for i := 0; b.Loop(); i++ { + if err := cache.Put(target+i, i); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkMixed simulates a workload with 90% reads and 10% writes, +// where writes cycle through a fixed key space and steadily trigger +// evictions once the cache is warm. +func BenchmarkMixed(b *testing.B) { + const target = 512 + cache := New[int, int](target, nil) + defer cache.Close() + + for i := range target { + if err := cache.Put(i, i); err != nil { + b.Fatal(err) + } + } + + b.ReportAllocs() + for i := 0; b.Loop(); i++ { + if i%10 == 0 { + if err := cache.Put(target+i, i); err != nil { + b.Fatal(err) + } + } else { + _, _ = cache.Get(i % target) + } + } +} diff --git a/caching/lru/lru_test.go b/caching/lru/lru_test.go index e7d0b6c47..e77399ea8 100644 --- a/caching/lru/lru_test.go +++ b/caching/lru/lru_test.go @@ -2,11 +2,34 @@ package lru import ( "errors" + "fmt" "testing" "github.com/stretchr/testify/require" ) +// requireConsistentList walks the cache's internal linked list from +// head to tail and checks that it forms a single well-formed chain of +// exactly c.size nodes with correctly paired prev/next pointers, no +// cycles, and a tail that matches where the walk actually ends. +func requireConsistentList[K comparable, V any](t *testing.T, c *Cache[K, V]) { + t.Helper() + + seen := make(map[K]bool, c.size) + var prev *node[K, V] + count := 0 + for n := c.head; n != nil; n = n.next { + count++ + require.LessOrEqualf(t, count, c.size, "list has more nodes than size=%d (cycle?)", c.size) + require.Falsef(t, seen[n.key], "key %v visited twice walking from head (cycle)", n.key) + seen[n.key] = true + require.Equal(t, prev, n.prev, "node %v has a mismatched prev pointer", n.key) + prev = n + } + require.Equal(t, c.size, count, "list length does not match c.size") + require.Equal(t, c.tail, prev, "c.tail does not match the last node reached from head") +} + func TestCache(t *testing.T) { t.Run("New Cache", func(t *testing.T) { cache := New[int, string](10, nil) @@ -63,6 +86,100 @@ func TestCache(t *testing.T) { require.Equal(t, "three", val) }) + t.Run("Eviction Beyond One Step", func(t *testing.T) { + cache := New[int, string](2, nil) + defer cache.Close() + + err := cache.Put(1, "one") + require.NoError(t, err) + err = cache.Put(2, "two") + require.NoError(t, err) + + // Repeatedly touch key 1 so it stays the most-recently-used + // entry, while a stream of new keys cycles through the other + // slot. + for i := 3; i <= 10; i++ { + _, ok := cache.Get(1) + require.True(t, ok) + + err := cache.Put(i, fmt.Sprintf("val-%d", i)) + require.NoError(t, err) + + // the key from the previous round, never touched again + // after it was inserted, must now be gone + _, ok = cache.Get(i - 1) + require.False(t, ok, "key %d should have been evicted", i-1) + + // key 1 and the newly inserted key must both still be + // present + val, ok := cache.Get(1) + require.True(t, ok) + require.Equal(t, "one", val) + + val, ok = cache.Get(i) + require.True(t, ok) + require.Equal(t, fmt.Sprintf("val-%d", i), val) + } + + _, _, size := cache.Stats() + require.Equal(t, uint64(2), size) + }) + + t.Run("Promote Middle Node", func(t *testing.T) { + cache := New[int, string](3, nil) + defer cache.Close() + + // With 3 resident keys, key 2 sits strictly between head and + // tail. + err := cache.Put(1, "one") + require.NoError(t, err) + err = cache.Put(2, "two") + require.NoError(t, err) + err = cache.Put(3, "three") + require.NoError(t, err) + requireConsistentList(t, cache) + + val, ok := cache.Get(2) + require.True(t, ok) + require.Equal(t, "two", val) + requireConsistentList(t, cache) + + // nothing should have been evicted by a Get; 1 and 3 must still + // be reachable even though 2 was spliced out from between them + val, ok = cache.Get(1) + require.True(t, ok) + require.Equal(t, "one", val) + requireConsistentList(t, cache) + + val, ok = cache.Get(3) + require.True(t, ok) + require.Equal(t, "three", val) + requireConsistentList(t, cache) + + // recency order is now (MRU -> LRU): 3, 1, 2. + // Filling the cache once more must evict 2, the one + // entry not touched since being spliced out. + err = cache.Put(4, "four") + require.NoError(t, err) + requireConsistentList(t, cache) + + _, ok = cache.Get(2) + require.False(t, ok, "key 2 should have been evicted") + + val, ok = cache.Get(1) + require.True(t, ok) + require.Equal(t, "one", val) + + val, ok = cache.Get(3) + require.True(t, ok) + require.Equal(t, "three", val) + + val, ok = cache.Get(4) + require.True(t, ok) + require.Equal(t, "four", val) + requireConsistentList(t, cache) + }) + t.Run("OnEvict Callback", func(t *testing.T) { evicted := make(map[int]string) onEvict := func(key int, val string) error { @@ -70,7 +187,7 @@ func TestCache(t *testing.T) { return nil } - cache := New[int, string](2, onEvict) + cache := New(2, onEvict) defer cache.Close() // Fill the cache @@ -93,7 +210,7 @@ func TestCache(t *testing.T) { return expectedErr } - cache := New[int, string](2, onEvict) + cache := New(2, onEvict) defer cache.Close() // Fill the cache @@ -136,6 +253,49 @@ func TestCache(t *testing.T) { require.Equal(t, uint64(2), size) }) + t.Run("Stats After Eviction", func(t *testing.T) { + cache := New[int, string](2, nil) + defer cache.Close() + + err := cache.Put(1, "one") + require.NoError(t, err) + err = cache.Put(2, "two") + require.NoError(t, err) + + // cache is now full; size must reflect that, with no hits/misses + // recorded yet + hits, misses, size := cache.Stats() + require.Equal(t, uint64(0), hits) + require.Equal(t, uint64(0), misses) + require.Equal(t, uint64(2), size) + + // this Put evicts key 1, so size must stay at 2, not grow to 3 + err = cache.Put(3, "three") + require.NoError(t, err) + + hits, misses, size = cache.Stats() + require.Equal(t, uint64(0), hits) + require.Equal(t, uint64(0), misses) + require.Equal(t, uint64(2), size) + + // a lookup for the evicted key must count as a miss and must not + // change size + _, ok := cache.Get(1) + require.False(t, ok) + + hits, misses, size = cache.Stats() + require.Equal(t, uint64(0), hits) + require.Equal(t, uint64(1), misses) + require.Equal(t, uint64(2), size) + + // evict once more (key 2) and confirm size still holds at 2 + err = cache.Put(4, "four") + require.NoError(t, err) + + _, _, size = cache.Stats() + require.Equal(t, uint64(2), size) + }) + t.Run("Update Existing", func(t *testing.T) { cache := New[int, string](3, nil) defer cache.Close()