diff --git a/bigcache_test.go b/bigcache_test.go index 15f1dae..d3fe265 100644 --- a/bigcache_test.go +++ b/bigcache_test.go @@ -1041,7 +1041,63 @@ func TestEntryBiggerThanMaxShardSizeError(t *testing.T) { err := cache.Set("key1", blob('a', 1024*1025)) // then - assertEqual(t, "entry is bigger than max shard size", err.Error()) + assertEqual(t, ErrEntryTooBig, err) +} + +func TestEntryBiggerThanMaxShardSizeDoesNotAllocateOrEvict(t *testing.T) { + t.Parallel() + + // given + cache, err := New(context.Background(), Config{ + Shards: 1, + LifeWindow: 5 * time.Second, + MaxEntriesInWindow: 1, + MaxEntrySize: 256, + HardMaxCacheSize: 1, + }) + noError(t, err) + noError(t, cache.Set("existing", []byte("value"))) + initialEntryBufferSize := len(cache.shards[0].entryBuffer) + initialQueueCapacity := cache.shards[0].entries.Capacity() + + // when + err = cache.Set("too-large", blob('a', 1024*1025)) + + // then + assertEqual(t, ErrEntryTooBig, err) + assertEqual(t, initialEntryBufferSize, len(cache.shards[0].entryBuffer)) + assertEqual(t, initialQueueCapacity, cache.shards[0].entries.Capacity()) + value, getErr := cache.Get("existing") + noError(t, getErr) + assertEqual(t, []byte("value"), value) +} + +func TestAppendEntryBiggerThanMaxShardSizeDoesNotAllocateOrEvict(t *testing.T) { + t.Parallel() + + // given + cache, err := New(context.Background(), Config{ + Shards: 1, + LifeWindow: 5 * time.Second, + MaxEntriesInWindow: 1, + MaxEntrySize: 256, + HardMaxCacheSize: 1, + }) + noError(t, err) + noError(t, cache.Set("existing", []byte("value"))) + initialEntryBufferSize := len(cache.shards[0].entryBuffer) + initialQueueCapacity := cache.shards[0].entries.Capacity() + + // when + err = cache.Append("existing", blob('a', 1024*1025)) + + // then + assertEqual(t, ErrEntryTooBig, err) + assertEqual(t, initialEntryBufferSize, len(cache.shards[0].entryBuffer)) + assertEqual(t, initialQueueCapacity, cache.shards[0].entries.Capacity()) + value, getErr := cache.Get("existing") + noError(t, getErr) + assertEqual(t, []byte("value"), value) } func TestHashCollision(t *testing.T) { diff --git a/entry_not_found_error.go b/errors.go similarity index 63% rename from entry_not_found_error.go rename to errors.go index c219dd8..c50d970 100644 --- a/entry_not_found_error.go +++ b/errors.go @@ -5,4 +5,7 @@ import "errors" var ( // ErrEntryNotFound is an error type struct which is returned when entry was not found for provided key ErrEntryNotFound = errors.New("Entry not found") //nolint:staticcheck // keep for backward compatibility + + // ErrEntryTooBig is returned when an entry cannot fit in a shard queue. + ErrEntryTooBig = errors.New("entry is bigger than max shard size") ) diff --git a/queue/bytes_queue.go b/queue/bytes_queue.go index 4cc43fe..e4c0041 100644 --- a/queue/bytes_queue.go +++ b/queue/bytes_queue.go @@ -106,6 +106,19 @@ func (q *BytesQueue) Push(data []byte) (int, error) { return index, nil } +// CanFit reports whether data of the given length can fit in the queue when +// the queue is empty without exceeding its maximum capacity. +func (q *BytesQueue) CanFit(dataLen int) bool { + if dataLen < 0 { + return false + } + if q.maxCapacity == 0 { + return true + } + + return getNeededSize(dataLen) <= q.maxCapacity-leftMarginIndex +} + func (q *BytesQueue) allocateAdditionalMemory(minimum int) { start := time.Now() if q.capacity < minimum { diff --git a/queue/bytes_queue_test.go b/queue/bytes_queue_test.go index dc7bdd7..43b78b6 100644 --- a/queue/bytes_queue_test.go +++ b/queue/bytes_queue_test.go @@ -403,6 +403,17 @@ func TestMaxSizeLimit(t *testing.T) { assertEqual(t, blob('b', 5), pop(queue)) } +func TestCanFit(t *testing.T) { + t.Parallel() + + // given + queue := NewBytesQueue(30, 50, false) + + // then + assertEqual(t, true, queue.CanFit(48)) + assertEqual(t, false, queue.CanFit(49)) +} + func TestPushEntryAfterAllocateAdditionMemory(t *testing.T) { t.Parallel() diff --git a/shard.go b/shard.go index 73ef86f..6cc9c61 100644 --- a/shard.go +++ b/shard.go @@ -1,7 +1,6 @@ package bigcache import ( - "errors" "sync" "sync/atomic" @@ -118,6 +117,10 @@ func (s *cacheShard) getValidWrapEntry(key string, hashedKey uint64) ([]byte, er } func (s *cacheShard) set(key string, hashedKey uint64, entry []byte) error { + if !s.entries.CanFit(len(entry) + len(key) + headersSizeInBytes) { + return ErrEntryTooBig + } + currentTimestamp := uint64(s.clock.Epoch()) s.lock.Lock() @@ -146,12 +149,16 @@ func (s *cacheShard) set(key string, hashedKey uint64, entry []byte) error { } if s.removeOldestEntry(NoSpace) != nil { s.lock.Unlock() - return errors.New("entry is bigger than max shard size") + return ErrEntryTooBig } } } func (s *cacheShard) addNewWithoutLock(key string, hashedKey uint64, entry []byte) error { + if !s.entries.CanFit(len(entry) + len(key) + headersSizeInBytes) { + return ErrEntryTooBig + } + currentTimestamp := uint64(s.clock.Epoch()) if !s.cleanEnabled { @@ -168,7 +175,7 @@ func (s *cacheShard) addNewWithoutLock(key string, hashedKey uint64, entry []byt return nil } if s.removeOldestEntry(NoSpace) != nil { - return errors.New("entry is bigger than max shard size") + return ErrEntryTooBig } } } @@ -192,7 +199,7 @@ func (s *cacheShard) setWrappedEntryWithoutLock(currentTimestamp uint64, w []byt return nil } if s.removeOldestEntry(NoSpace) != nil { - return errors.New("entry is bigger than max shard size") + return ErrEntryTooBig } } } @@ -210,6 +217,10 @@ func (s *cacheShard) append(key string, hashedKey uint64, entry []byte) error { s.lock.Unlock() return err } + if !s.entries.CanFit(len(wrappedEntry) + len(entry)) { + s.lock.Unlock() + return ErrEntryTooBig + } currentTimestamp := uint64(s.clock.Epoch())