diff --git a/backend/pkg/cache/disk/cache.go b/backend/pkg/cache/disk/cache.go index a4514957..83b32117 100644 --- a/backend/pkg/cache/disk/cache.go +++ b/backend/pkg/cache/disk/cache.go @@ -143,7 +143,10 @@ func (c *Cache) Set(key string, value []byte, ttl time.Duration) error { // Update memory tracker if elem, ok := c.items[key]; ok { - item := elem.Value.(*cacheItem) + item, ok := elem.Value.(*cacheItem) + if !ok { + return fmt.Errorf("cache: evict list entry for %q has invalid type %T", key, elem.Value) + } c.currentSize += size - item.size item.size = size item.expiredAt = expiredAt @@ -174,7 +177,11 @@ func (c *Cache) Get(key string) ([]byte, error) { return nil, ErrCacheMiss } - item := elem.Value.(*cacheItem) + item, ok := elem.Value.(*cacheItem) + if !ok { + c.mu.RUnlock() + return nil, ErrCacheMiss + } if !item.expiredAt.IsZero() && time.Now().After(item.expiredAt) { c.mu.RUnlock() return c.getAndDeleteIfExpired(key) @@ -224,7 +231,11 @@ func (c *Cache) getAndDeleteIfExpired(key string) ([]byte, error) { return nil, ErrCacheMiss } - item := elem.Value.(*cacheItem) + item, ok := elem.Value.(*cacheItem) + if !ok { + _ = c.deleteUnlocked(key) + return nil, ErrCacheMiss + } if !item.expiredAt.IsZero() && time.Now().After(item.expiredAt) { _ = c.deleteUnlocked(key) return nil, ErrCacheMiss @@ -254,8 +265,9 @@ func (c *Cache) Delete(key string) error { func (c *Cache) deleteUnlocked(key string) error { if elem, ok := c.items[key]; ok { - item := elem.Value.(*cacheItem) - c.currentSize -= item.size + if item, ok := elem.Value.(*cacheItem); ok { + c.currentSize -= item.size + } c.evictList.Remove(elem) delete(c.items, key) } @@ -308,7 +320,11 @@ func (c *Cache) evict() { for c.currentSize > c.maxSize && c.evictList.Len() > 0 { elem := c.evictList.Back() - item := elem.Value.(*cacheItem) + item, ok := elem.Value.(*cacheItem) + if !ok { + c.evictList.Remove(elem) + continue + } c.currentSize -= item.size c.evictList.Remove(elem) delete(c.items, item.key) @@ -400,7 +416,12 @@ func (c *Cache) cleanExpired() { now := time.Now() for key, elem := range c.items { - item := elem.Value.(*cacheItem) + item, ok := elem.Value.(*cacheItem) + if !ok { + c.evictList.Remove(elem) + delete(c.items, key) + continue + } if !item.expiredAt.IsZero() && now.After(item.expiredAt) { c.currentSize -= item.size c.evictList.Remove(elem) diff --git a/backend/pkg/cache/disk/cache_corruption_test.go b/backend/pkg/cache/disk/cache_corruption_test.go new file mode 100644 index 00000000..379c005a --- /dev/null +++ b/backend/pkg/cache/disk/cache_corruption_test.go @@ -0,0 +1,63 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package disk + +import ( + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +// 这些用例锁住「LRU 链表节点被污染时不得 panic」的行为:一旦 items 与 evictList +// 的不变量被破坏(例如后续改动误写节点),缓存必须降级为未命中/跳过, +// 而不是在读、写、删除与淘汰路径上崩掉整个进程。 + +// corruptEntry 写入一个键后把其链表节点值换成非法类型,返回缓存。 +func corruptEntry(t *testing.T, key string) *Cache { + t.Helper() + + c := New(t.TempDir()) + require.NoError(t, c.Set(key, []byte("payload"), time.Minute)) + + elem, ok := c.items[key] + require.True(t, ok, "entry must be tracked after Set") + elem.Value = "not-a-cacheItem" + return c +} + +func TestGetToleratesCorruptEvictEntry(t *testing.T) { + c := corruptEntry(t, "k") + + got, err := c.Get("k") + require.ErrorIs(t, err, ErrCacheMiss) + require.Nil(t, got) +} + +func TestSetOverCorruptEvictEntryReportsError(t *testing.T) { + c := corruptEntry(t, "k") + + err := c.Set("k", []byte("second"), time.Minute) + require.Error(t, err, "Set must report the corrupted tracker entry instead of panicking") + require.Contains(t, err.Error(), "invalid type") +} + +func TestDeleteToleratesCorruptEvictEntry(t *testing.T) { + c := corruptEntry(t, "k") + + require.NotPanics(t, func() { _ = c.Delete("k") }) + require.NotContains(t, c.items, "k") +} + +func TestEvictToleratesCorruptEvictEntry(t *testing.T) { + c := corruptEntry(t, "k") + // 让任意写入都触发淘汰扫描:扫到被污染的节点必须跳过而非 panic。 + c.UpdatePolicy(0, 0, true) + + require.NotPanics(t, func() { + for i := range 4 { + _ = c.Set(string(rune('a'+i)), []byte("x"), time.Minute) + } + }) +}