From cb065190ab1ea4854f2138b1a9c2e73b897d1f04 Mon Sep 17 00:00:00 2001 From: Ravi Agarwal Date: Mon, 31 Aug 2026 21:37:48 -0700 Subject: [PATCH 1/3] memory: bound cache admission and eviction Replace the process-wide quadratic eviction path with sharded, fixed-work CLOCK admission. Revalidate planned victims with hit generations so competing planners cannot erase a recent reference. Account retained entries, incomplete writers, and reader-pinned generations against one configured ceiling. Grow writer buffers lazily and reserve only net capacity so declared lengths cannot allocate memory before body data and streaming chunk boundaries do not change admission. Reject invalid limit conversions and stop cancelled eviction plans before committing victims. Decline only the optional memory copy when bounded admission cannot obtain capacity, preserving authoritative-tier and client delivery. Cover replacement, cancellation, concurrent admission, tier fallback, accounting transitions, shutdown, and writer growth. Add cardinality and configured-path benchmarks that distinguish successful admissions from bounded declines. Verification: bin/just fmt; bin/just lint; bin/just test. --- README.md | 29 +- internal/cache/memory.go | 992 +++++++++++++++++---- internal/cache/memory_internal_test.go | 1113 ++++++++++++++++++++++++ internal/cache/tiered_test.go | 28 + 4 files changed, 1987 insertions(+), 175 deletions(-) create mode 100644 internal/cache/memory_internal_test.go diff --git a/README.md b/README.md index 301eb49..51e1312 100644 --- a/README.md +++ b/README.md @@ -150,12 +150,35 @@ copies before falling through to the authoritative tier. ### Memory -In-memory LRU cache. +In-memory sharded CLOCK cache with bounded admission and eviction work. Cache +hits lock only the shard containing the requested key. Each admission and trim +examines at most 64 victims, regardless of cache cardinality. + +`limit-mb` is a hard retained-memory accounting ceiling, not a process RSS +limit. Accounting includes object buffers, estimated metadata, and buffers held +by active readers. `Stats.Capacity` reports this ceiling; `Stats.Size` reports +payload bytes and can differ because it excludes charged metadata and spare +buffer capacity. Go runtime and allocator overhead can make RSS differ from +both values. + +Incomplete writes remain unbounded when `inflight-limit-mb` is zero, preserving +the behavior of configurations written before this option existed. A positive +value limits aggregate incomplete writes and reserves that amount inside +`limit-mb`: retained entries are trimmed toward `limit-mb - +inflight-limit-mb`, and retained plus incomplete accounting cannot exceed +`limit-mb`. Writes that cannot obtain capacity within the bounded admission +work bypass the memory tier without interrupting other cache tiers. Declared +content lengths are validated against these limits, but buffers grow only as +body bytes arrive. Buffer growth transfers the existing accounting reservation +to the larger capacity; the allocator may briefly retain both allocations, so +process RSS can transiently exceed the accounting ceiling by the old buffer's +capacity. ```hcl memory { - limit-mb = 1024 # default - max-ttl = "1h" # default + limit-mb = 1024 # default + inflight-limit-mb = 0 # disabled for compatibility + max-ttl = "1h" # default } ``` diff --git a/internal/cache/memory.go b/internal/cache/memory.go index 3b86f07..e62fce4 100644 --- a/internal/cache/memory.go +++ b/internal/cache/memory.go @@ -1,11 +1,11 @@ package cache import ( - "bytes" "context" "fmt" "io" "maps" + "math" "net/http" "os" "strconv" @@ -19,65 +19,566 @@ import ( "github.com/block/cachew/internal/logging" ) +const ( + memoryShardCount = 16 + maxMemoryEvictionsPerWrite = 64 + maxMemoryEvictionScansPerShard = 64 + maxMemoryReservationRetries = 8 + memoryEntryMinimumCharge = 4 * 1024 + memoryEntryBaseCharge = 512 + memoryHeaderEntryCharge = 64 + memoryHeaderValueCharge = 16 + memoryWriterMinimumCharge = 4 * 1024 + memoryWriterInitialCapacity = 64 * 1024 + memoryBytesPerMegabyte = 1024 * 1024 + fnv64Offset = 14695981039346656037 + fnv64Prime = 1099511628211 +) + +// RegisterMemory preserves the existing "memory" HCL backend while replacing its admission policy. func RegisterMemory(r *Registry) { Register( r, "memory", - "Caches objects in memory, with a maximum size limit and LRU eviction", + "Caches objects in memory, with retained-size accounting and bounded CLOCK eviction", NewMemory, ) } +// MemoryConfig keeps incomplete-write protection opt-in so existing zero-valued configurations remain compatible. type MemoryConfig struct { - LimitMB int `hcl:"limit-mb,optional" help:"Maximum size of the disk cache in megabytes (defaults to 1GB)." default:"1024"` - MaxTTL time.Duration `hcl:"max-ttl,optional" help:"Maximum time-to-live for entries in the disk cache (defaults to 1 hour)." default:"1h"` + LimitMB int `hcl:"limit-mb,optional" help:"Maximum retained size of the memory cache in megabytes (defaults to 1GB); positive inflight-limit-mb shares this budget." default:"1024"` + InflightLimitMB int `hcl:"inflight-limit-mb,optional" help:"Maximum aggregate incomplete writes in megabytes (0 disables the sub-limit for compatibility)." default:"0"` + MaxTTL time.Duration `hcl:"max-ttl,optional" help:"Maximum time-to-live for entries in the memory cache (defaults to 1 hour)." default:"1h"` } type memoryEntry struct { - data []byte - expiresAt time.Time - headers http.Header + namespace Namespace + key Key + data []byte + expiresAt time.Time + headers http.Header + charge int64 + referenceEpoch atomic.Uint64 + clockEpoch uint64 + readers atomic.Int64 + previous *memoryEntry + next *memoryEntry + retired atomic.Bool + released bool +} + +type memoryShard struct { + mu sync.RWMutex + entries map[Namespace]map[Key]*memoryEntry + evictionHead *memoryEntry + evictionTail *memoryEntry + evictionHand *memoryEntry +} + +func (s *memoryShard) entry(namespace Namespace, key Key) (*memoryEntry, bool) { + namespaceEntries, ok := s.entries[namespace] + if !ok { + return nil, false + } + entry, ok := namespaceEntries[key] + return entry, ok +} + +func (s *memoryShard) append(entry *memoryEntry) { + entry.previous = s.evictionTail + entry.next = nil + if s.evictionTail == nil { + s.evictionHead = entry + } else { + s.evictionTail.next = entry + } + s.evictionTail = entry + if s.evictionHand == nil { + s.evictionHand = entry + } +} + +func (s *memoryShard) remove(entry *memoryEntry) { + nextHand := entry.next + if entry.previous == nil { + s.evictionHead = entry.next + } else { + entry.previous.next = entry.next + } + if entry.next == nil { + s.evictionTail = entry.previous + } else { + entry.next.previous = entry.previous + } + entry.previous = nil + entry.next = nil + if s.evictionHand == entry { + s.evictionHand = nextHand + if s.evictionHand == nil { + s.evictionHand = s.evictionHead + } + } +} + +func (s *memoryShard) insert(entry *memoryEntry) { + namespaceEntries := s.entries[entry.namespace] + if namespaceEntries == nil { + namespaceEntries = make(map[Key]*memoryEntry) + s.entries[entry.namespace] = namespaceEntries + } + namespaceEntries[entry.key] = entry + s.append(entry) } +type memoryState struct { + shards []memoryShard + limitBytes int64 + retainedTarget int64 + inflightLimit int64 + retainedCharge atomic.Int64 + inflightCharge atomic.Int64 + hardLimitCharge atomic.Int64 + payloadSize atomic.Int64 + objectCount atomic.Int64 + evictionCursor atomic.Uint32 + closed atomic.Bool +} + +// Memory shares capacity across namespace views so each view cannot consume limit-mb independently. type Memory struct { - config MemoryConfig - namespace Namespace - mu *sync.RWMutex - entries map[Namespace]map[Key]*memoryEntry // namespace -> key -> entry - currentSize *atomic.Int64 + config MemoryConfig + namespace Namespace + state *memoryState } -func NewMemory(ctx context.Context, config MemoryConfig) (*Memory, error) { - logging.FromContext(ctx).InfoContext(ctx, "Constructing in-memory Cache", "limit-mb", config.LimitMB, "max-ttl", config.MaxTTL) - return &Memory{ - config: config, - mu: &sync.RWMutex{}, - entries: make(map[Namespace]map[Key]*memoryEntry), - currentSize: &atomic.Int64{}, +func memoryMegabytes(field string, value int) (int64, error) { + if value < 0 { + return 0, errors.Errorf("%s must be non-negative", field) + } + if int64(value) > math.MaxInt64/memoryBytesPerMegabyte { + return 0, errors.Errorf("%s is too large", field) + } + return int64(value) * memoryBytesPerMegabyte, nil +} + +func newMemoryState(config MemoryConfig) (*memoryState, error) { + limitBytes, err := memoryMegabytes("limit-mb", config.LimitMB) + if err != nil { + return nil, err + } + configuredInflightBytes, err := memoryMegabytes("inflight-limit-mb", config.InflightLimitMB) + if err != nil { + return nil, err + } + inflightLimit := memoryInflightLimit(configuredInflightBytes, limitBytes) + retainedTarget := memoryRetainedTarget(limitBytes, inflightLimit) + shards := make([]memoryShard, memoryShardCount) + for index := range shards { + shards[index].entries = make(map[Namespace]map[Key]*memoryEntry) + } + return &memoryState{ + shards: shards, limitBytes: limitBytes, + retainedTarget: retainedTarget, inflightLimit: inflightLimit, }, nil } -func (m *Memory) String() string { return fmt.Sprintf("memory:%dMB", m.config.LimitMB) } +// NewMemory rejects invalid byte conversions before a bad limit can silently become unlimited. +func NewMemory(ctx context.Context, config MemoryConfig) (*Memory, error) { + state, err := newMemoryState(config) + if err != nil { + return nil, errors.WithStack(err) + } + logging.FromContext(ctx).InfoContext(ctx, "Constructing in-memory Cache", "limit-mb", config.LimitMB, + "inflight-limit-mb", config.InflightLimitMB, "max-ttl", config.MaxTTL) + return &Memory{config: config, state: state}, nil +} -func (m *Memory) backendType() BackendType { return backendMemory } +func memoryShardIndex(namespace Namespace, key Key) int { + hash := uint64(fnv64Offset) + for index := range len(namespace) { + hash ^= uint64(namespace[index]) + hash *= fnv64Prime + } + hash ^= 0xff + hash *= fnv64Prime + for _, value := range key { + hash ^= uint64(value) + hash *= fnv64Prime + } + return int(hash % memoryShardCount) +} -func (m *Memory) Stat(_ context.Context, key Key, opts ...Option) (http.Header, error) { - m.mu.RLock() - defer m.mu.RUnlock() +func (m *Memory) shard(namespace Namespace, key Key) *memoryShard { + return &m.state.shards[memoryShardIndex(namespace, key)] +} - nsEntries, nsExists := m.entries[m.namespace] - if !nsExists { - return nil, os.ErrNotExist +func expectedContentLength(headers http.Header) int64 { + contentLength, err := strconv.ParseInt(headers.Get("Content-Length"), 10, 64) + if err != nil || contentLength < 0 { + return -1 } + return contentLength +} - entry, exists := nsEntries[key] - if !exists { - return nil, os.ErrNotExist +func memoryMetadataCharge(namespace Namespace, headers http.Header) int64 { + charge := int64(memoryEntryBaseCharge + len(namespace)) + for name, values := range headers { + charge += int64(memoryHeaderEntryCharge + len(name)) + for _, value := range values { + charge += int64(memoryHeaderValueCharge + len(value)) + } + } + return charge +} + +func memoryEntryCharge(namespace Namespace, data []byte, headers http.Header) int64 { + charge := int64(cap(data)) + memoryMetadataCharge(namespace, headers) + return max(charge, int64(memoryEntryMinimumCharge)) +} + +func memoryInflightLimit(configuredBytes, limitBytes int64) int64 { + if configuredBytes <= 0 { + return 0 + } + if limitBytes > 0 { + return min(configuredBytes, limitBytes) + } + return configuredBytes +} + +func memoryRetainedTarget(limitBytes, inflightLimit int64) int64 { + if limitBytes <= 0 { + return 0 + } + return limitBytes - min(limitBytes, max(inflightLimit, 0)) +} + +func reserveBounded(counter *atomic.Int64, limit, amount int64) bool { + if amount <= 0 { + return true + } + for range maxMemoryReservationRetries { + current := counter.Load() + if amount > limit-current { + return false + } + if counter.CompareAndSwap(current, current+amount) { + return true + } + } + return false +} + +func reserveCounter(counter *atomic.Int64, limit, amount int64) bool { + if amount <= 0 { + return true + } + if limit <= 0 { + counter.Add(amount) + return true + } + return reserveBounded(counter, limit, amount) +} + +type memoryPlannedEviction struct { + shard *memoryShard + entry *memoryEntry + referenceEpoch uint64 +} + +func (m *Memory) releaseRetiredLocked(entry *memoryEntry) { + if entry.released || !entry.retired.Load() || entry.readers.Load() != 0 { + return + } + entry.released = true + m.state.retainedCharge.Add(-entry.charge) + m.state.hardLimitCharge.Add(-entry.charge) + entry.data = nil + entry.headers = nil +} + +func (m *Memory) removeActiveLocked(shard *memoryShard, entry *memoryEntry) { + shard.remove(entry) + namespaceEntries := shard.entries[entry.namespace] + delete(namespaceEntries, entry.key) + if len(namespaceEntries) == 0 { + delete(shard.entries, entry.namespace) + } + entry.retired.Store(true) + m.state.objectCount.Add(-1) + m.state.payloadSize.Add(-int64(len(entry.data))) + m.releaseRetiredLocked(entry) +} + +func (m *Memory) replaceActiveLocked(shard *memoryShard, oldEntry, newEntry *memoryEntry) { + oldPayloadSize := len(oldEntry.data) + shard.remove(oldEntry) + oldEntry.retired.Store(true) + oldEntry.released = true + oldEntry.data = nil + oldEntry.headers = nil + shard.entries[newEntry.namespace][newEntry.key] = newEntry + shard.append(newEntry) + m.state.payloadSize.Add(int64(len(newEntry.data) - oldPayloadSize)) +} + +func (m *Memory) insertActiveLocked(shard *memoryShard, entry *memoryEntry) { + shard.insert(entry) + m.state.objectCount.Add(1) + m.state.payloadSize.Add(int64(len(entry.data))) +} + +func (m *Memory) reserveRetained(retainedLimit, amount int64) bool { + if !reserveCounter(&m.state.retainedCharge, retainedLimit, amount) { + return false + } + if reserveCounter(&m.state.hardLimitCharge, m.state.limitBytes, amount) { + return true + } + m.state.retainedCharge.Add(-amount) + return false +} + +type memoryAdmissionMode uint8 + +const ( + memoryAdmissionNeedsAllocation memoryAdmissionMode = iota + memoryAdmissionHasAllocation +) + +func (m *Memory) reserveAdmission(mode memoryAdmissionMode, retainedLimit, amount int64) bool { + if mode == memoryAdmissionHasAllocation { + return reserveCounter(&m.state.retainedCharge, retainedLimit, amount) + } + return m.reserveRetained(retainedLimit, amount) +} + +func (m *Memory) tryAdmission( + ctx context.Context, + entry *memoryEntry, + retainedLimit int64, + mode memoryAdmissionMode, +) (bool, error) { + shard := m.shard(entry.namespace, entry.key) + shard.mu.Lock() + defer shard.mu.Unlock() + if m.state.closed.Load() { + return false, errors.WithStack(os.ErrClosed) + } + if err := ctx.Err(); err != nil { + return false, errors.WithStack(err) + } + oldEntry, replacing := shard.entry(entry.namespace, entry.key) + if replacing && oldEntry.readers.Load() == 0 { + delta := entry.charge - oldEntry.charge + if delta > 0 && !m.reserveAdmission(mode, retainedLimit, delta) { + return false, nil + } + m.replaceActiveLocked(shard, oldEntry, entry) + if delta < 0 { + m.state.retainedCharge.Add(delta) + } + if mode == memoryAdmissionHasAllocation { + m.state.hardLimitCharge.Add(-oldEntry.charge) + } else if delta < 0 { + m.state.hardLimitCharge.Add(delta) + } + return true, nil + } + if !m.reserveAdmission(mode, retainedLimit, entry.charge) { + return false, nil + } + if oldEntry != nil { + m.removeActiveLocked(shard, oldEntry) + } + m.insertActiveLocked(shard, entry) + return true, nil +} + +func (m *Memory) admitReserved(ctx context.Context, entry *memoryEntry) (bool, error) { + if m.state.closed.Load() { + return false, errors.WithStack(os.ErrClosed) + } + if m.state.limitBytes > 0 && entry.charge > m.state.limitBytes { + return false, nil + } + admitted, err := m.tryAdmission(ctx, entry, m.state.limitBytes, memoryAdmissionHasAllocation) + if admitted || err != nil || m.state.limitBytes <= 0 { + if admitted && m.state.limitBytes > 0 { + m.trimToTarget(ctx, entry.namespace, entry.key) + } + return admitted, err + } + m.trimToTarget(ctx, entry.namespace, entry.key) + admitted, err = m.tryAdmission(ctx, entry, m.state.limitBytes, memoryAdmissionHasAllocation) + if admitted { + m.trimToTarget(ctx, entry.namespace, entry.key) + } + return admitted, err +} + +func (m *Memory) planEvictions( + needed int64, + protectedNamespace Namespace, + protectedKey Key, + planned []memoryPlannedEviction, +) int { + if needed <= 0 { + return 0 + } + now := time.Now() + start := int((m.state.evictionCursor.Add(1) - 1) % memoryShardCount) + plannedEntries := 0 + plannedSize := int64(0) + for offset := range memoryShardCount { + shard := &m.state.shards[(start+offset)%memoryShardCount] + shard.mu.Lock() + scanned := 0 + candidate := shard.evictionHand + startCandidate := candidate + for candidate != nil && scanned < maxMemoryEvictionScansPerShard { + scanned++ + nextCandidate := candidate.next + if nextCandidate == nil { + nextCandidate = shard.evictionHead + } + if candidate.namespace != protectedNamespace || candidate.key != protectedKey { + referenceEpoch := candidate.referenceEpoch.Load() + expired := !now.Before(candidate.expiresAt) + if candidate.readers.Load() == 0 { + if expired || candidate.clockEpoch == referenceEpoch { + planned[plannedEntries] = memoryPlannedEviction{ + shard: shard, entry: candidate, referenceEpoch: referenceEpoch, + } + plannedEntries++ + plannedSize += candidate.charge + } else { + candidate.clockEpoch = referenceEpoch + } + } + } + candidate = nextCandidate + if plannedSize >= needed || plannedEntries == maxMemoryEvictionsPerWrite || candidate == startCandidate { + break + } + } + shard.evictionHand = candidate + shard.mu.Unlock() + if plannedSize >= needed || plannedEntries == maxMemoryEvictionsPerWrite { + break + } + } + return plannedEntries +} + +func (m *Memory) commitMemoryEvictionPlan(ctx context.Context, planned []memoryPlannedEviction, target int64) { + now := time.Now() + for start := 0; start < len(planned); { + shard := planned[start].shard + shard.mu.Lock() + if ctx.Err() != nil { + shard.mu.Unlock() + return + } + end := start + for end < len(planned) && planned[end].shard == shard { + if m.state.retainedCharge.Load() <= target { + shard.mu.Unlock() + return + } + plannedEntry := planned[end] + entry := plannedEntry.entry + activeEntry, active := shard.entry(entry.namespace, entry.key) + if active && activeEntry == entry && entry.readers.Load() == 0 { + referenceEpoch := entry.referenceEpoch.Load() + if !now.Before(entry.expiresAt) || referenceEpoch == plannedEntry.referenceEpoch { + m.removeActiveLocked(shard, entry) + } else { + entry.clockEpoch = referenceEpoch + } + } + end++ + } + shard.mu.Unlock() + start = end + } +} + +func (m *Memory) trimToTarget(ctx context.Context, protectedNamespace Namespace, protectedKey Key) { + needed := m.state.retainedCharge.Load() - m.state.retainedTarget + if needed <= 0 { + return + } + var planBuffer [maxMemoryEvictionsPerWrite]memoryPlannedEviction + plannedEntries := m.planEvictions(needed, protectedNamespace, protectedKey, planBuffer[:]) + m.commitMemoryEvictionPlan(ctx, planBuffer[:plannedEntries], m.state.retainedTarget) +} + +func (m *Memory) trimForAdmission(ctx context.Context, entry *memoryEntry) { + amount := entry.charge + shard := m.shard(entry.namespace, entry.key) + shard.mu.RLock() + if oldEntry, replacing := shard.entry(entry.namespace, entry.key); replacing && oldEntry.readers.Load() == 0 { + amount = max(entry.charge-oldEntry.charge, 0) + } + shard.mu.RUnlock() + target := max(m.state.limitBytes-amount, 0) + needed := m.state.retainedCharge.Load() - target + if needed <= 0 { + return + } + var planBuffer [maxMemoryEvictionsPerWrite]memoryPlannedEviction + plannedEntries := m.planEvictions(needed, entry.namespace, entry.key, planBuffer[:]) + m.commitMemoryEvictionPlan(ctx, planBuffer[:plannedEntries], target) +} + +func (m *Memory) admit(ctx context.Context, entry *memoryEntry) (bool, error) { + if m.state.closed.Load() { + return false, errors.WithStack(os.ErrClosed) + } + if m.state.limitBytes > 0 && entry.charge > m.state.limitBytes { + return false, nil + } + admitted, err := m.tryAdmission(ctx, entry, m.state.retainedTarget, memoryAdmissionNeedsAllocation) + if err != nil || m.state.limitBytes <= 0 { + return admitted, err + } + if admitted { + m.trimToTarget(ctx, entry.namespace, entry.key) + return true, nil + } + if err := ctx.Err(); err != nil { + return false, errors.WithStack(err) + } + if admitted, err = m.tryAdmission(ctx, entry, m.state.limitBytes, memoryAdmissionNeedsAllocation); admitted || err != nil { + return admitted, err } + m.trimForAdmission(ctx, entry) + admitted, err = m.tryAdmission(ctx, entry, m.state.limitBytes, memoryAdmissionNeedsAllocation) + if !admitted || err != nil { + return admitted, err + } + m.trimToTarget(ctx, entry.namespace, entry.key) + return true, nil +} + +// String includes the hard limit so differently sized tiers remain distinguishable in diagnostics. +func (m *Memory) String() string { return fmt.Sprintf("memory:%dMB", m.config.LimitMB) } + +func (m *Memory) backendType() BackendType { return backendMemory } + +// Stat counts metadata-only hits as CLOCK references so they receive the same recency protection as Open. +func (m *Memory) Stat(_ context.Context, key Key, opts ...Option) (http.Header, error) { + shard := m.shard(m.namespace, key) + shard.mu.RLock() + defer shard.mu.RUnlock() - if time.Now().After(entry.expiresAt) { + entry, exists := shard.entry(m.namespace, key) + if !exists || time.Now().After(entry.expiresAt) { return nil, os.ErrNotExist } + entry.referenceEpoch.Add(1) headers := maps.Clone(entry.headers) headers.Set("Content-Length", strconv.Itoa(len(entry.data))) @@ -87,23 +588,17 @@ func (m *Memory) Stat(_ context.Context, key Key, opts ...Option) (http.Header, return headers, nil } +// Open pins its entry generation so replacement and eviction cannot invalidate an active response body. func (m *Memory) Open(_ context.Context, key Key, opts ...Option) (io.ReadCloser, http.Header, error) { - m.mu.RLock() - defer m.mu.RUnlock() + shard := m.shard(m.namespace, key) + shard.mu.RLock() + defer shard.mu.RUnlock() - nsEntries, nsExists := m.entries[m.namespace] - if !nsExists { - return nil, nil, os.ErrNotExist - } - - entry, exists := nsEntries[key] - if !exists { - return nil, nil, os.ErrNotExist - } - - if time.Now().After(entry.expiresAt) { + entry, exists := shard.entry(m.namespace, key) + if !exists || time.Now().After(entry.expiresAt) { return nil, nil, os.ErrNotExist } + entry.referenceEpoch.Add(1) headers := maps.Clone(entry.headers) headers.Set("Content-Length", strconv.Itoa(len(entry.data))) @@ -119,16 +614,56 @@ func (m *Memory) Open(_ context.Context, key Key, opts ...Option) (io.ReadCloser if partial { data = data[start : start+length] } - return io.NopCloser(bytes.NewReader(data)), headers, nil + entry.readers.Add(1) + return &memoryReader{data: data, cache: m, shard: shard, entry: entry}, headers, nil } +type memoryReader struct { + data []byte + offset int + cache *Memory + shard *memoryShard + entry *memoryEntry + closed atomic.Bool +} + +func (r *memoryReader) Read(p []byte) (int, error) { + if r.closed.Load() { + return 0, os.ErrClosed + } + if r.offset >= len(r.data) { + return 0, io.EOF + } + n := copy(p, r.data[r.offset:]) + r.offset += n + return n, nil +} + +func (r *memoryReader) Close() error { + if r.closed.Swap(true) { + return nil + } + r.data = nil + if r.entry.readers.Add(-1) != 0 || !r.entry.retired.Load() { + return nil + } + r.shard.mu.Lock() + r.cache.releaseRetiredLocked(r.entry) + r.shard.mu.Unlock() + return nil +} + +// Create buffers privately so incomplete bodies stay invisible and declined memory admission remains non-fatal. func (m *Memory) Create(ctx context.Context, key Key, headers http.Header, ttl time.Duration, opts ...Option) (Writer, error) { + if m.state.closed.Load() { + return nil, errors.WithStack(os.ErrClosed) + } if ttl == 0 { ttl = m.config.MaxTTL } now := time.Now() - // Clone (to avoid concurrent map writes) and drop transport headers. + contentLength := expectedContentLength(headers) clonedHeaders := httputil.FilterHeaders(headers, httputil.TransportHeaders...) if clonedHeaders.Get("Last-Modified") == "" { clonedHeaders.Set("Last-Modified", now.UTC().Format(http.TimeFormat)) @@ -137,40 +672,54 @@ func (m *Memory) Create(ctx context.Context, key Key, headers http.Header, ttl t return nil, err } + metadataCharge := memoryMetadataCharge(m.namespace, clonedHeaders) + baseCharge := max(metadataCharge, int64(memoryWriterMinimumCharge)) + if contentLength >= 0 && m.state.limitBytes > 0 && contentLength > m.state.limitBytes-metadataCharge { + return &noOpWriter{}, nil + } + if contentLength >= 0 && m.state.inflightLimit > 0 && contentLength > m.state.inflightLimit-baseCharge { + return &noOpWriter{}, nil + } ctx, cancel := context.WithCancelCause(ctx) - writer := &memoryWriter{ - cache: m, - namespace: m.namespace, - key: key, - buf: &bytes.Buffer{}, - expiresAt: now.Add(ttl), - headers: clonedHeaders, - ctx: ctx, - cancel: cancel, + cache: m, + namespace: m.namespace, + key: key, + expiresAt: now.Add(ttl), + headers: clonedHeaders, + limitBytes: m.state.limitBytes, + inflightLimit: m.state.inflightLimit, + budgeted: m.state.inflightLimit > 0, + baseCharge: baseCharge, + expectedLength: contentLength, + ctx: ctx, + cancel: cancel, + } + if !writer.reserve(baseCharge) { + cancel(nil) + return &noOpWriter{}, nil } - return writer, nil } +// Delete defers reclaiming reader-pinned storage until the final reader closes. func (m *Memory) Delete(_ context.Context, key Key) error { - m.mu.Lock() - defer m.mu.Unlock() - - nsEntries, nsExists := m.entries[m.namespace] - if !nsExists { - return os.ErrNotExist + shard := m.shard(m.namespace, key) + shard.mu.Lock() + defer shard.mu.Unlock() + if m.state.closed.Load() { + return errors.WithStack(os.ErrClosed) } - entry, exists := nsEntries[key] + entry, exists := shard.entry(m.namespace, key) if !exists { return os.ErrNotExist } - m.currentSize.Add(-int64(len(entry.data))) - delete(nsEntries, key) + m.removeActiveLocked(shard, entry) return nil } +// Invalidate treats a missing entry as success because callers use it to discard optional stale copies. func (m *Memory) Invalidate(ctx context.Context, key Key) error { err := m.Delete(ctx, key) if errors.Is(err, os.ErrNotExist) { @@ -179,88 +728,190 @@ func (m *Memory) Invalidate(ctx context.Context, key Key) error { return errors.WithStack(err) } +// Close stops admission immediately while reader-pinned generations remain valid until their readers close. func (m *Memory) Close() error { - m.mu.Lock() - defer m.mu.Unlock() - - m.entries = nil + if !m.state.closed.CompareAndSwap(false, true) { + return nil + } + for index := range m.state.shards { + shard := &m.state.shards[index] + shard.mu.Lock() + for shard.evictionHead != nil { + m.removeActiveLocked(shard, shard.evictionHead) + } + shard.entries = make(map[Namespace]map[Key]*memoryEntry) + shard.mu.Unlock() + } return nil } +// Stats uses atomics so metrics collection never waits behind cache traffic on a shard lock. func (m *Memory) Stats(_ context.Context) (Stats, error) { - m.mu.RLock() - defer m.mu.RUnlock() + return Stats{ + Objects: m.state.objectCount.Load(), + Size: m.state.payloadSize.Load(), + Capacity: m.state.limitBytes, + }, nil +} + +type memoryWriter struct { + cache *Memory + namespace Namespace + key Key + data []byte + expiresAt time.Time + headers http.Header + limitBytes int64 + inflightLimit int64 + budgeted bool + baseCharge int64 + reservedBytes int64 + expectedLength int64 + discarded bool + closed bool + ctx context.Context + cancel context.CancelCauseFunc +} - totalObjects := int64(0) - for _, nsEntries := range m.entries { - totalObjects += int64(len(nsEntries)) +func (w *memoryWriter) Write(p []byte) (int, error) { + if w.closed { + return 0, errors.New("writer closed") + } + if w.discarded { + return len(p), nil } + buffered := int64(len(w.data)) + tooLarge := w.limitBytes > 0 && int64(len(p)) > w.limitBytes-w.baseCharge-buffered + longerThanDeclared := w.expectedLength >= 0 && int64(len(p)) > w.expectedLength-buffered + needed := buffered + int64(len(p)) + if tooLarge || longerThanDeclared || !w.ensureCapacity(needed) { + w.discard() + return len(p), nil + } + w.data = append(w.data, p...) + return len(p), nil +} - return Stats{ - Objects: totalObjects, - Size: m.currentSize.Load(), - Capacity: int64(m.config.LimitMB) * 1024 * 1024, - }, nil +func memoryBufferCapacity(size int64) (int, bool) { + capacity := int(size) + return capacity, capacity >= 0 && int64(capacity) == size } -func (m *Memory) evictOldest(neededSpace int64) { - type entryInfo struct { - namespace Namespace - key Key - size int64 - expiresAt time.Time +func (w *memoryWriter) maximumBodyCapacity() int64 { + maximum := int64(math.MaxInt64) + if w.limitBytes > 0 { + maximum = min(maximum, w.limitBytes-w.baseCharge) } + if w.inflightLimit > 0 { + maximum = min(maximum, w.inflightLimit-w.baseCharge) + } + return maximum +} - var entries []entryInfo - for ns, nsEntries := range m.entries { - for k, e := range nsEntries { - entries = append(entries, entryInfo{ - namespace: ns, - key: k, - size: int64(len(e.data)), - expiresAt: e.expiresAt, - }) +func (w *memoryWriter) nextCapacity(needed int64) (int64, bool) { + maximum := w.maximumBodyCapacity() + if needed < 0 || needed > maximum { + return 0, false + } + current := int64(cap(w.data)) + initial := min(int64(memoryWriterInitialCapacity), maximum) + if w.expectedLength >= 0 { + initial = min(initial, w.expectedLength) + } + next := max(needed, initial) + if current > 0 { + doubled := maximum + if current <= maximum-current { + doubled = current * 2 } + next = max(next, min(doubled, maximum)) } + return next, true +} - // Sort by expiry time (earliest first) - for i := 0; i < len(entries); i++ { - for j := i + 1; j < len(entries); j++ { - if entries[i].expiresAt.After(entries[j].expiresAt) { - entries[i], entries[j] = entries[j], entries[i] - } - } +func (w *memoryWriter) ensureCapacity(needed int64) bool { + if needed <= int64(cap(w.data)) { + return true + } + next, ok := w.nextCapacity(needed) + if !ok { + return false } + oldCapacity := int64(cap(w.data)) + additionalCapacity := next - oldCapacity + if !w.reserve(additionalCapacity) { + return false + } + capacity, ok := memoryBufferCapacity(next) + if !ok { + w.release(additionalCapacity) + return false + } + grown := make([]byte, len(w.data), capacity) + copy(grown, w.data) + w.data = grown + return true +} - freedSpace := int64(0) - for _, e := range entries { - if freedSpace >= neededSpace { - break - } - m.currentSize.Add(-e.size) - delete(m.entries[e.namespace], e.key) - freedSpace += e.size +func (w *memoryWriter) reserve(amount int64) bool { + if amount <= 0 { + return true + } + if w.tryReserve(amount) { + return true } + if !w.budgeted { + return false + } + w.cache.trimToTarget(w.ctx, w.namespace, w.key) + return w.tryReserve(amount) } -type memoryWriter struct { - cache *Memory - namespace Namespace - key Key - buf *bytes.Buffer - expiresAt time.Time - headers http.Header - closed bool - ctx context.Context - cancel context.CancelCauseFunc +func (w *memoryWriter) tryReserve(amount int64) bool { + if !reserveCounter(&w.cache.state.inflightCharge, w.inflightLimit, amount) { + return false + } + if w.budgeted && !reserveCounter(&w.cache.state.hardLimitCharge, w.limitBytes, amount) { + w.cache.state.inflightCharge.Add(-amount) + return false + } + w.reservedBytes += amount + return true } -func (w *memoryWriter) Write(p []byte) (int, error) { - if w.closed { - return 0, errors.New("writer closed") +func (w *memoryWriter) release(amount int64) { + if amount <= 0 { + return + } + w.cache.state.inflightCharge.Add(-amount) + if w.budgeted { + w.cache.state.hardLimitCharge.Add(-amount) } - n, err := w.buf.Write(p) - return n, errors.WithStack(err) + w.reservedBytes -= amount +} + +func (w *memoryWriter) releaseReservation() { + if w.reservedBytes > 0 { + w.cache.state.inflightCharge.Add(-w.reservedBytes) + if w.budgeted { + w.cache.state.hardLimitCharge.Add(-w.reservedBytes) + } + w.reservedBytes = 0 + } +} + +func (w *memoryWriter) transferReservation(charge int64) { + w.cache.state.inflightCharge.Add(-w.reservedBytes) + if excess := w.reservedBytes - charge; excess > 0 { + w.cache.state.hardLimitCharge.Add(-excess) + } + w.reservedBytes = 0 +} + +func (w *memoryWriter) discard() { + w.releaseReservation() + w.data = nil + w.discarded = true } func (w *memoryWriter) Abort(err error) error { @@ -273,70 +924,67 @@ func (w *memoryWriter) Close() error { return nil } w.closed = true - - // Check if context was cancelled + defer w.releaseReservation() if err := w.ctx.Err(); err != nil { + w.discard() return errors.Wrap(err, "create operation cancelled") } - - w.cache.mu.Lock() - defer w.cache.mu.Unlock() - - newSize := int64(w.buf.Len()) - limitBytes := int64(w.cache.config.LimitMB) * 1024 * 1024 - - // Ensure namespace map exists - if w.cache.entries[w.namespace] == nil { - w.cache.entries[w.namespace] = make(map[Key]*memoryEntry) - } - nsEntries := w.cache.entries[w.namespace] - - // Remove old entry size if it exists - oldSize := int64(0) - if oldEntry, exists := nsEntries[w.key]; exists { - oldSize = int64(len(oldEntry.data)) + if w.discarded { + return nil } - - // Evict entries if needed to make room - if limitBytes > 0 { - neededSpace := w.cache.currentSize.Load() - oldSize + newSize - limitBytes - if neededSpace > 0 { - w.cache.evictOldest(neededSpace) - } + if w.expectedLength >= 0 && int64(len(w.data)) != w.expectedLength { + w.discard() + return nil } - w.cache.currentSize.Add(-oldSize) - // Copy the buffer data to avoid holding a reference to the buffer's internal slice - data := make([]byte, w.buf.Len()) - copy(data, w.buf.Bytes()) - w.buf.Reset() - nsEntries[w.key] = &memoryEntry{ + data := w.data + w.data = nil + entry := &memoryEntry{ + namespace: w.namespace, + key: w.key, data: data, expiresAt: w.expiresAt, headers: w.headers, } - w.cache.currentSize.Add(newSize) - - return nil + entry.charge = memoryEntryCharge(entry.namespace, entry.data, entry.headers) + if !w.budgeted { + _, err := w.cache.admit(w.ctx, entry) + return errors.WithStack(err) + } + if entry.charge > w.reservedBytes && !w.reserve(entry.charge-w.reservedBytes) { + return nil + } + admitted, err := w.cache.admitReserved(w.ctx, entry) + if admitted { + w.transferReservation(entry.charge) + } + return errors.WithStack(err) } -// Namespace creates a namespaced view of the memory cache. +// Namespace reuses one state object so every protocol namespace shares the configured capacity. func (m *Memory) Namespace(namespace Namespace) Cache { - c := *m - c.namespace = namespace - return &c + view := *m + view.namespace = namespace + return &view } -// ListNamespaces returns all unique namespaces in the memory cache. +// ListNamespaces excludes the default namespace because only explicit cache partitions are discoverable. func (m *Memory) ListNamespaces(_ context.Context) ([]string, error) { - m.mu.RLock() - defer m.mu.RUnlock() - - namespaces := make([]string, 0, len(m.entries)) - for ns := range m.entries { - if ns != "" { - namespaces = append(namespaces, string(ns)) + namespaces := make(map[Namespace]struct{}) + for index := range m.state.shards { + shard := &m.state.shards[index] + shard.mu.RLock() + for namespace, namespaceEntries := range shard.entries { + if namespace != "" && len(namespaceEntries) > 0 { + namespaces[namespace] = struct{}{} + } } + shard.mu.RUnlock() + } + + result := make([]string, 0, len(namespaces)) + for namespace := range namespaces { + result = append(result, string(namespace)) } - return namespaces, nil + return result, nil } diff --git a/internal/cache/memory_internal_test.go b/internal/cache/memory_internal_test.go new file mode 100644 index 0000000..2f832aa --- /dev/null +++ b/internal/cache/memory_internal_test.go @@ -0,0 +1,1113 @@ +package cache + +import ( + "bytes" + "context" + "encoding/binary" + "fmt" + "io" + "log/slog" + "math" + "net/http" + "os" + "strconv" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/alecthomas/assert/v2" + "github.com/alecthomas/errors" + + "github.com/block/cachew/internal/logging" +) + +func newMemoryTestCache(t *testing.T) *Memory { + t.Helper() + return newMemoryTestCacheWithConfig(t, MemoryConfig{LimitMB: 1, MaxTTL: time.Hour}) +} + +func newMemoryTestCacheWithConfig(t *testing.T, config MemoryConfig) *Memory { + t.Helper() + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + memory, err := NewMemory(ctx, config) + assert.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, memory.Close()) }) + return memory +} + +func memoryKeysForShard(t *testing.T, namespace Namespace, shardIndex, count int) []Key { + t.Helper() + keys := make([]Key, 0, count) + for candidate := 0; len(keys) < count; candidate++ { + key := NewKey(fmt.Sprintf("memory-shard-%d", candidate)) + if memoryShardIndex(namespace, key) == shardIndex { + keys = append(keys, key) + } + } + return keys +} + +func writeMemoryTestEntry(t *testing.T, cache Cache, key Key, data []byte, ttl time.Duration) { + t.Helper() + writer, err := cache.Create(t.Context(), key, http.Header{}, ttl) + assert.NoError(t, err) + written, err := writer.Write(data) + assert.NoError(t, err) + assert.Equal(t, len(data), written) + assert.NoError(t, writer.Close()) +} + +func newMemoryTestEntry(namespace Namespace, key Key, data []byte) *memoryEntry { + entry := &memoryEntry{ + namespace: namespace, + key: key, + data: data, + expiresAt: time.Now().Add(time.Hour), + headers: http.Header{}, + } + entry.charge = memoryEntryCharge(namespace, data, entry.headers) + return entry +} + +func memoryTestWriter(t testing.TB, writer Writer) *memoryWriter { + t.Helper() + memoryWriter, ok := writer.(*memoryWriter) + assert.True(t, ok) + return memoryWriter +} + +func memoryTestReader(t testing.TB, reader io.ReadCloser) *memoryReader { + t.Helper() + memoryReader, ok := reader.(*memoryReader) + assert.True(t, ok) + return memoryReader +} + +func admitMemoryTestEntry(ctx context.Context, t testing.TB, memory *Memory, entry *memoryEntry) { + t.Helper() + admitted, err := memory.admit(ctx, entry) + assert.NoError(t, err) + assert.True(t, admitted) +} + +func assertMemoryAccounting( + t *testing.T, + memory *Memory, + readers []*memoryReader, + writers []*memoryWriter, +) { + t.Helper() + entries := make(map[*memoryEntry]struct{}) + payloadSize := int64(0) + objectCount := int64(0) + for index := range memory.state.shards { + shard := &memory.state.shards[index] + shard.mu.RLock() + for _, namespaceEntries := range shard.entries { + for _, entry := range namespaceEntries { + entries[entry] = struct{}{} + payloadSize += int64(len(entry.data)) + objectCount++ + } + } + shard.mu.RUnlock() + } + for _, reader := range readers { + if !reader.entry.released { + entries[reader.entry] = struct{}{} + } + } + + retainedSize := int64(0) + for entry := range entries { + retainedSize += entry.charge + } + inflightSize := int64(0) + hardLimitCharge := retainedSize + for _, writer := range writers { + inflightSize += writer.reservedBytes + if writer.budgeted { + hardLimitCharge += writer.reservedBytes + } + } + + assert.Equal(t, retainedSize, memory.state.retainedCharge.Load()) + assert.Equal(t, inflightSize, memory.state.inflightCharge.Load()) + assert.Equal(t, hardLimitCharge, memory.state.hardLimitCharge.Load()) + assert.Equal(t, payloadSize, memory.state.payloadSize.Load()) + assert.Equal(t, objectCount, memory.state.objectCount.Load()) + if memory.state.limitBytes > 0 { + assert.True(t, hardLimitCharge <= memory.state.limitBytes) + } + if memory.state.inflightLimit > 0 { + assert.True(t, inflightSize <= memory.state.inflightLimit) + } +} + +type admissionCancellationContext struct { + context.Context + firstCheck chan struct{} + cancelled atomic.Bool + once sync.Once +} + +func (c *admissionCancellationContext) Err() error { + if c.cancelled.Load() { + return context.Canceled + } + c.once.Do(func() { close(c.firstCheck) }) + return nil +} + +func TestMemoryShardDoesNotBlockUnrelatedReads(t *testing.T) { + memory := newMemoryTestCache(t) + namespace := Namespace("sharded") + keys := memoryKeysForShard(t, namespace, 0, 1) + otherKeys := memoryKeysForShard(t, namespace, 1, 1) + cache := memory.Namespace(namespace) + writeMemoryTestEntry(t, cache, otherKeys[0], []byte("hit"), time.Hour) + lockedShard := memory.shard(namespace, keys[0]) + lockedShard.mu.Lock() + defer lockedShard.mu.Unlock() + + result := make(chan error, 1) + go func() { + _, err := cache.Stat(t.Context(), otherKeys[0]) + result <- err + }() + + select { + case err := <-result: + assert.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("cache hit on an unrelated shard blocked") + } +} + +func TestMemoryAdmissionReclaimsCapacityAcrossShards(t *testing.T) { + memory := newMemoryTestCache(t) + namespace := Namespace("cross-shard") + fullShardKeys := memoryKeysForShard(t, namespace, 0, 2) + otherShardKey := memoryKeysForShard(t, namespace, 1, 1)[0] + cache := memory.Namespace(namespace) + for _, key := range fullShardKeys { + writeMemoryTestEntry(t, cache, key, make([]byte, 480*1024), time.Hour) + } + + writeMemoryTestEntry(t, cache, otherShardKey, make([]byte, 128*1024), time.Hour) + + reader, _, err := cache.Open(t.Context(), otherShardKey) + assert.NoError(t, err) + assert.NoError(t, reader.Close()) +} + +func TestMemoryEvictionRetainsRecentlyReadEntry(t *testing.T) { + memory := newMemoryTestCache(t) + namespace := Namespace("clock") + keys := memoryKeysForShard(t, namespace, 0, 3) + cache := memory.Namespace(namespace) + entryData := make([]byte, 400*1024) + writeMemoryTestEntry(t, cache, keys[0], entryData, time.Hour) + writeMemoryTestEntry(t, cache, keys[1], entryData, time.Hour) + reader, _, err := cache.Open(t.Context(), keys[0]) + assert.NoError(t, err) + assert.NoError(t, reader.Close()) + + writeMemoryTestEntry(t, cache, keys[2], entryData, time.Hour) + + reader, _, err = cache.Open(t.Context(), keys[0]) + assert.NoError(t, err) + assert.NoError(t, reader.Close()) + _, _, err = cache.Open(t.Context(), keys[1]) + assert.IsError(t, err, os.ErrNotExist) +} + +func TestMemoryEvictionPrefersColdEntryFromAnotherShard(t *testing.T) { + memory := newMemoryTestCache(t) + namespace := Namespace("cross-shard-clock") + hotKey := memoryKeysForShard(t, namespace, 0, 1)[0] + coldKey := memoryKeysForShard(t, namespace, 1, 1)[0] + newKey := memoryKeysForShard(t, namespace, 2, 1)[0] + cache := memory.Namespace(namespace) + entryData := make([]byte, 400*1024) + writeMemoryTestEntry(t, cache, hotKey, entryData, time.Hour) + writeMemoryTestEntry(t, cache, coldKey, entryData, time.Hour) + reader, _, err := cache.Open(t.Context(), hotKey) + assert.NoError(t, err) + assert.NoError(t, reader.Close()) + memory.state.evictionCursor.Store(0) + + writeMemoryTestEntry(t, cache, newKey, entryData, time.Hour) + + for _, key := range []Key{hotKey, newKey} { + reader, _, err = cache.Open(t.Context(), key) + assert.NoError(t, err) + assert.NoError(t, reader.Close()) + } + _, _, err = cache.Open(t.Context(), coldKey) + assert.IsError(t, err, os.ErrNotExist) +} + +func TestMemoryZeroLengthEntriesConsumeCapacity(t *testing.T) { + memory := newMemoryTestCache(t) + cache := memory.Namespace("metadata") + for index := range 400 { + writeMemoryTestEntry(t, cache, NewKey(fmt.Sprintf("zero-%d", index)), nil, time.Hour) + } + + stats, err := cache.Stats(t.Context()) + assert.NoError(t, err) + assert.True(t, stats.Objects < 400) + assert.Equal(t, int64(0), stats.Size) + assert.True(t, memory.state.retainedCharge.Load() > 0) + assert.True(t, memory.state.retainedCharge.Load() <= stats.Capacity) +} + +func TestMemoryInflightBuffersAreBounded(t *testing.T) { + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + memory, err := NewMemory(ctx, MemoryConfig{LimitMB: 4, InflightLimitMB: 1, MaxTTL: time.Hour}) + assert.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, memory.Close()) }) + writers := make([]Writer, 0, 16) + for index := range 16 { + writer, err := memory.Create(t.Context(), NewKey(fmt.Sprintf("inflight-%d", index)), nil, time.Hour) + assert.NoError(t, err) + _, err = writer.Write(make([]byte, 256*1024)) + assert.NoError(t, err) + writers = append(writers, writer) + } + + activeWriters := 0 + reservedBytes := int64(0) + for _, writer := range writers { + memoryWriter := memoryTestWriter(t, writer) + if memoryWriter.discarded { + continue + } + activeWriters++ + reservedBytes += memoryWriter.reservedBytes + assert.Equal(t, memoryWriter.baseCharge+int64(cap(memoryWriter.data)), memoryWriter.reservedBytes) + } + assert.True(t, activeWriters > 0) + assert.True(t, activeWriters < len(writers)) + assert.Equal(t, reservedBytes, memory.state.inflightCharge.Load()) + assert.True(t, memory.state.inflightCharge.Load() <= memory.state.inflightLimit) + for _, writer := range writers { + assert.NoError(t, writer.Close()) + } + assert.Equal(t, int64(0), memory.state.inflightCharge.Load()) + stats, err := memory.Stats(t.Context()) + assert.NoError(t, err) + assert.Equal(t, int64(activeWriters), stats.Objects) +} + +func TestMemoryDefaultDoesNotLimitInflightBuffers(t *testing.T) { + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + memory, err := NewMemory(ctx, MemoryConfig{LimitMB: 1, MaxTTL: time.Hour}) + assert.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, memory.Close()) }) + writers := make([]Writer, 0, 2) + for index := range 2 { + writer, err := memory.Create(t.Context(), NewKey(fmt.Sprintf("unbounded-inflight-%d", index)), nil, time.Hour) + assert.NoError(t, err) + _, err = writer.Write(make([]byte, 600*1024)) + assert.NoError(t, err) + memoryWriter := memoryTestWriter(t, writer) + assert.False(t, memoryWriter.discarded) + writers = append(writers, writer) + } + assert.Equal(t, int64(0), memory.state.inflightLimit) + assert.True(t, memory.state.inflightCharge.Load() > memory.state.limitBytes) + for _, writer := range writers { + assert.NoError(t, writer.Close()) + } + assert.Equal(t, int64(0), memory.state.inflightCharge.Load()) + assert.True(t, memory.state.retainedCharge.Load() <= memory.state.limitBytes) +} + +func TestMemoryConfiguredInflightSharesHardBudget(t *testing.T) { + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + memory, err := NewMemory(ctx, MemoryConfig{LimitMB: 4, InflightLimitMB: 1, MaxTTL: time.Hour}) + assert.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, memory.Close()) }) + oldKey := NewKey("retained-before-inflight") + admitted, err := memory.tryAdmission( + ctx, + newMemoryTestEntry("", oldKey, make([]byte, 3500*1024)), + memory.state.limitBytes, + memoryAdmissionNeedsAllocation, + ) + assert.NoError(t, err) + assert.True(t, admitted) + + newKey := NewKey("configured-inflight") + writer, err := memory.Create(ctx, newKey, nil, time.Hour) + assert.NoError(t, err) + _, err = writer.Write(make([]byte, 768*1024)) + assert.NoError(t, err) + memoryWriter := memoryTestWriter(t, writer) + assert.False(t, memoryWriter.discarded) + assert.True(t, memory.state.hardLimitCharge.Load() <= memory.state.limitBytes) + assert.True(t, memory.state.inflightCharge.Load() <= memory.state.inflightLimit) + assert.NoError(t, writer.Close()) + + reader, _, err := memory.Open(ctx, newKey) + assert.NoError(t, err) + assert.NoError(t, reader.Close()) + assert.Equal(t, int64(0), memory.state.inflightCharge.Load()) + assert.True(t, memory.state.hardLimitCharge.Load() <= memory.state.limitBytes) + assert.True(t, memory.state.retainedCharge.Load() <= memory.state.limitBytes) +} + +func TestMemorySlowAdmissionDoesNotBlockUnrelatedHits(t *testing.T) { + memory := newMemoryTestCache(t) + namespace := Namespace("concurrent-hits") + blockedKey := memoryKeysForShard(t, namespace, 0, 1)[0] + hitKey := memoryKeysForShard(t, namespace, 1, 1)[0] + newKey := memoryKeysForShard(t, namespace, 2, 1)[0] + deleteKey := memoryKeysForShard(t, namespace, 3, 1)[0] + cache := memory.Namespace(namespace) + entryData := make([]byte, 400*1024) + writeMemoryTestEntry(t, cache, blockedKey, entryData, time.Hour) + writeMemoryTestEntry(t, cache, hitKey, entryData, time.Hour) + writeMemoryTestEntry(t, cache, deleteKey, make([]byte, 64*1024), time.Hour) + writer, err := cache.Create(t.Context(), newKey, nil, time.Hour) + assert.NoError(t, err) + _, err = writer.Write(entryData) + assert.NoError(t, err) + memory.state.evictionCursor.Store(0) + blockedShard := memory.shard(namespace, blockedKey) + blockedShard.mu.Lock() + + closed := make(chan error, 1) + go func() { + closed <- writer.Close() + }() + deadline := time.Now().Add(time.Second) + for memory.state.evictionCursor.Load() == 0 { + if time.Now().After(deadline) { + blockedShard.mu.Unlock() + t.Fatal("admission did not reach the capacity path") + } + time.Sleep(time.Millisecond) + } + + hitReader, _, err := cache.Open(t.Context(), hitKey) + assert.NoError(t, err) + assert.NoError(t, hitReader.Close()) + deleted := make(chan error, 1) + go func() { deleted <- cache.Delete(t.Context(), deleteKey) }() + select { + case err := <-deleted: + assert.NoError(t, err) + case <-time.After(time.Second): + blockedShard.mu.Unlock() + t.Fatal("delete on an unrelated shard blocked") + } + blockedShard.mu.Unlock() + select { + case err := <-closed: + assert.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("admission did not finish after the shard was released") + } +} + +func TestMemoryAdmissionAtCapacityEvictsExistingEntry(t *testing.T) { + memory := newMemoryTestCache(t) + namespace := Namespace("eviction") + keys := memoryKeysForShard(t, namespace, 0, 3) + cache := memory.Namespace(namespace) + writeMemoryTestEntry(t, cache, keys[0], make([]byte, 512*1024), 3*time.Hour) + writeMemoryTestEntry(t, cache, keys[1], make([]byte, 512*1024), time.Minute) + writeMemoryTestEntry(t, cache, keys[2], make([]byte, 128*1024), 2*time.Hour) + + _, err := cache.Stat(t.Context(), keys[2]) + assert.NoError(t, err) + stats, err := cache.Stats(t.Context()) + assert.NoError(t, err) + assert.Equal(t, int64(2), stats.Objects) + assert.True(t, stats.Size >= int64(640*1024)) + assert.True(t, stats.Size <= stats.Capacity) +} + +func TestMemoryReplacementKeepsUnrelatedEntries(t *testing.T) { + memory := newMemoryTestCache(t) + namespace := Namespace("replacement") + keys := memoryKeysForShard(t, namespace, 0, 2) + cache := memory.Namespace(namespace) + writeMemoryTestEntry(t, cache, keys[0], make([]byte, 40*1024), time.Minute) + writeMemoryTestEntry(t, cache, keys[1], make([]byte, 20*1024), 2*time.Hour) + replacement := make([]byte, 40*1024) + replacement[0] = 1 + writeMemoryTestEntry(t, cache, keys[0], replacement, 3*time.Hour) + + reader, _, err := cache.Open(t.Context(), keys[0]) + assert.NoError(t, err) + data, err := io.ReadAll(reader) + assert.NoError(t, err) + assert.NoError(t, reader.Close()) + assert.Equal(t, byte(1), data[0]) + _, err = cache.Stat(t.Context(), keys[1]) + assert.NoError(t, err) + stats, err := cache.Stats(t.Context()) + assert.NoError(t, err) + assert.Equal(t, int64(2), stats.Objects) + assert.True(t, stats.Size >= int64(60*1024)) + assert.True(t, stats.Size <= stats.Capacity) +} + +func TestMemoryReplacementKeepsReaderChargedUntilClose(t *testing.T) { + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + memory, err := NewMemory(ctx, MemoryConfig{LimitMB: 2, MaxTTL: time.Hour}) + assert.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, memory.Close()) }) + key := NewKey("reader-replacement") + oldData := make([]byte, 600*1024) + writeMemoryTestEntry(t, memory, key, oldData, time.Hour) + oldReader, _, err := memory.Open(t.Context(), key) + assert.NoError(t, err) + oneGeneration := memory.state.retainedCharge.Load() + newData := make([]byte, len(oldData)) + newData[0] = 1 + + writeMemoryTestEntry(t, memory, key, newData, time.Hour) + assert.True(t, memory.state.retainedCharge.Load() > oneGeneration) + storedOld, err := io.ReadAll(oldReader) + assert.NoError(t, err) + assert.Equal(t, byte(0), storedOld[0]) + assert.NoError(t, oldReader.Close()) + assert.Equal(t, oneGeneration, memory.state.retainedCharge.Load()) + assert.Equal(t, oneGeneration, memory.state.hardLimitCharge.Load()) + + newReader, _, err := memory.Open(t.Context(), key) + assert.NoError(t, err) + storedNew, err := io.ReadAll(newReader) + assert.NoError(t, err) + assert.NoError(t, newReader.Close()) + assert.Equal(t, byte(1), storedNew[0]) +} + +func TestMemoryConfiguredBudgetAccountingAcrossTransitions(t *testing.T) { + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + memory, err := NewMemory(ctx, MemoryConfig{LimitMB: 4, InflightLimitMB: 1, MaxTTL: time.Hour}) + assert.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, memory.Close()) }) + key := NewKey("accounting-transitions") + + writer, err := memory.Create(ctx, key, nil, time.Hour) + assert.NoError(t, err) + bufferedWriter := memoryTestWriter(t, writer) + _, err = writer.Write(make([]byte, 256*1024)) + assert.NoError(t, err) + assertMemoryAccounting(t, memory, nil, []*memoryWriter{bufferedWriter}) + assert.NoError(t, writer.Close()) + assertMemoryAccounting(t, memory, nil, []*memoryWriter{bufferedWriter}) + + reader, _, err := memory.Open(ctx, key) + assert.NoError(t, err) + pinnedReader := memoryTestReader(t, reader) + replacement, err := memory.Create(ctx, key, nil, time.Hour) + assert.NoError(t, err) + replacementWriter := memoryTestWriter(t, replacement) + _, err = replacement.Write(make([]byte, 128*1024)) + assert.NoError(t, err) + assertMemoryAccounting(t, memory, []*memoryReader{pinnedReader}, []*memoryWriter{replacementWriter}) + assert.NoError(t, replacement.Close()) + assertMemoryAccounting(t, memory, []*memoryReader{pinnedReader}, []*memoryWriter{replacementWriter}) + + oldData, err := io.ReadAll(reader) + assert.NoError(t, err) + assert.Equal(t, 256*1024, len(oldData)) + assert.NoError(t, reader.Close()) + assertMemoryAccounting(t, memory, []*memoryReader{pinnedReader}, nil) + + aborted, err := memory.Create(ctx, NewKey("aborted-accounting"), nil, time.Hour) + assert.NoError(t, err) + abortedWriter := memoryTestWriter(t, aborted) + _, err = aborted.Write(make([]byte, 64*1024)) + assert.NoError(t, err) + assertMemoryAccounting(t, memory, nil, []*memoryWriter{abortedWriter}) + assert.IsError(t, aborted.Abort(errors.New("abort accounting write")), context.Canceled) + assertMemoryAccounting(t, memory, nil, []*memoryWriter{abortedWriter}) + + assert.NoError(t, memory.Delete(ctx, key)) + assertMemoryAccounting(t, memory, nil, nil) + + shutdownKey := NewKey("shutdown-accounting") + writeMemoryTestEntry(t, memory, shutdownKey, make([]byte, 32*1024), time.Hour) + shutdownReader, _, err := memory.Open(ctx, shutdownKey) + assert.NoError(t, err) + pinnedShutdownReader := memoryTestReader(t, shutdownReader) + assert.NoError(t, memory.Close()) + assertMemoryAccounting(t, memory, []*memoryReader{pinnedShutdownReader}, nil) + assert.NoError(t, shutdownReader.Close()) + assertMemoryAccounting(t, memory, []*memoryReader{pinnedShutdownReader}, nil) +} + +func TestMemoryStatsUsePayloadBytesWithoutTakingShardLocks(t *testing.T) { + memory := newMemoryTestCache(t) + namespace := Namespace("stats") + key := memoryKeysForShard(t, namespace, 0, 1)[0] + cache := memory.Namespace(namespace) + payload := []byte("payload") + writeMemoryTestEntry(t, cache, key, payload, time.Hour) + shard := memory.shard(namespace, key) + shard.mu.Lock() + type statsResult struct { + stats Stats + err error + } + result := make(chan statsResult, 1) + go func() { + stats, err := cache.Stats(t.Context()) + result <- statsResult{stats: stats, err: err} + }() + var response statsResult + select { + case response = <-result: + case <-time.After(time.Second): + shard.mu.Unlock() + t.Fatal("stats blocked on a shard lock") + } + shard.mu.Unlock() + stats, err := response.stats, response.err + assert.NoError(t, err) + assert.Equal(t, int64(1), stats.Objects) + assert.Equal(t, int64(len(payload)), stats.Size) + assert.True(t, memory.state.retainedCharge.Load() > stats.Size) +} + +func TestMemoryUnlimitedAccountingReturnsToZero(t *testing.T) { + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + memory, err := NewMemory(ctx, MemoryConfig{LimitMB: 0, MaxTTL: time.Hour}) + assert.NoError(t, err) + key := NewKey("unlimited") + writeMemoryTestEntry(t, memory, key, []byte("payload"), time.Hour) + assert.True(t, memory.state.retainedCharge.Load() > 0) + + assert.NoError(t, memory.Delete(t.Context(), key)) + assert.Equal(t, int64(0), memory.state.retainedCharge.Load()) + assert.Equal(t, int64(0), memory.state.hardLimitCharge.Load()) + assert.NoError(t, memory.Close()) +} + +func TestMemoryOversizedWriteDoesNotReplaceExistingEntry(t *testing.T) { + memory := newMemoryTestCache(t) + namespace := Namespace("oversized") + key := memoryKeysForShard(t, namespace, 0, 1)[0] + cache := memory.Namespace(namespace) + writeMemoryTestEntry(t, cache, key, []byte("old"), time.Hour) + + limitBytes := int64(memory.config.LimitMB) * 1024 * 1024 + writeMemoryTestEntry(t, cache, key, make([]byte, limitBytes+1), time.Hour) + + reader, _, err := cache.Open(t.Context(), key) + assert.NoError(t, err) + data, err := io.ReadAll(reader) + assert.NoError(t, err) + assert.NoError(t, reader.Close()) + assert.Equal(t, "old", string(data)) +} + +func TestMemoryKnownOversizedWriteIsNotAdmitted(t *testing.T) { + memory := newMemoryTestCache(t) + key := NewKey("known-oversized") + limitBytes := int64(memory.config.LimitMB) * 1024 * 1024 + w, err := memory.Create(t.Context(), key, http.Header{"Content-Length": {strconv.FormatInt(limitBytes+1, 10)}}, time.Hour) + assert.NoError(t, err) + _, err = w.Write([]byte("partial prefix")) + assert.NoError(t, err) + assert.NoError(t, w.Close()) + _, _, err = memory.Open(t.Context(), key) + assert.IsError(t, err, os.ErrNotExist) +} + +func TestMemoryPartialWriteIsNotAdmitted(t *testing.T) { + memory := newMemoryTestCache(t) + key := NewKey("partial") + w, err := memory.Create(t.Context(), key, http.Header{"Content-Length": {"10"}}, time.Hour) + assert.NoError(t, err) + _, err = w.Write([]byte("abc")) + assert.NoError(t, err) + assert.NoError(t, w.Close()) + _, _, err = memory.Open(t.Context(), key) + assert.IsError(t, err, os.ErrNotExist) +} + +func TestMemoryDeclaredLengthDoesNotAllocateBeforeWrite(t *testing.T) { + memory := newMemoryTestCacheWithConfig(t, MemoryConfig{MaxTTL: time.Hour}) + declaredLength := 16 * 1024 * 1024 + writer, err := memory.Create(t.Context(), NewKey("lazy-declared-length"), http.Header{ + "Content-Length": {strconv.Itoa(declaredLength)}, + }, time.Hour) + assert.NoError(t, err) + memoryWriter := memoryTestWriter(t, writer) + assert.Equal(t, 0, cap(memoryWriter.data)) + assert.Equal(t, memoryWriter.baseCharge, memoryWriter.reservedBytes) + assert.NoError(t, writer.Close()) + assert.Equal(t, int64(0), memory.state.inflightCharge.Load()) +} + +func TestMemoryUnknownLengthAdmissionDoesNotDependOnWriteChunks(t *testing.T) { + memory := newMemoryTestCacheWithConfig(t, MemoryConfig{LimitMB: 2, InflightLimitMB: 1, MaxTTL: time.Hour}) + key := NewKey("chunked-unknown-length") + writer, err := memory.Create(t.Context(), key, http.Header{}, time.Hour) + assert.NoError(t, err) + payload := bytes.Repeat([]byte{0x7a}, 640*1024) + for offset := 0; offset < len(payload); offset += 64 * 1024 { + written, err := writer.Write(payload[offset : offset+64*1024]) + assert.NoError(t, err) + assert.Equal(t, 64*1024, written) + } + assert.NoError(t, writer.Close()) + + reader, _, err := memory.Open(t.Context(), key) + assert.NoError(t, err) + stored, err := io.ReadAll(reader) + assert.NoError(t, err) + assert.NoError(t, reader.Close()) + assert.Equal(t, payload, stored) +} + +func TestMemoryRejectsInvalidLimits(t *testing.T) { + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + type invalidLimitTest struct { + name string + config MemoryConfig + } + tests := []invalidLimitTest{ + {name: "negative retained limit", config: MemoryConfig{LimitMB: -1}}, + {name: "negative inflight limit", config: MemoryConfig{InflightLimitMB: -1}}, + } + if strconv.IntSize == 64 { + overflowMB := int(math.MaxInt64/(1024*1024) + 1) + tests = append(tests, + invalidLimitTest{name: "overflowing retained limit", config: MemoryConfig{LimitMB: overflowMB}}, + invalidLimitTest{name: "overflowing inflight limit", config: MemoryConfig{InflightLimitMB: overflowMB}}, + ) + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + memory, err := NewMemory(ctx, test.config) + if memory != nil { + assert.NoError(t, memory.Close()) + } + assert.Error(t, err) + }) + } +} + +func TestMemoryCancelledAdmissionDoesNotEvictEntries(t *testing.T) { + memory := newMemoryTestCache(t) + const namespace Namespace = "cancelled-admission" + oldKeys := []Key{ + memoryKeysForShard(t, namespace, 0, 1)[0], + memoryKeysForShard(t, namespace, 1, 1)[0], + } + newKey := memoryKeysForShard(t, namespace, 2, 1)[0] + cache := memory.Namespace(namespace) + for _, key := range oldKeys { + writeMemoryTestEntry(t, cache, key, make([]byte, 480*1024), time.Hour) + } + + ctx, cancel := context.WithCancel(t.Context()) + blockedShard := memory.shard(namespace, oldKeys[0]) + memory.state.evictionCursor.Store(0) + blockedShard.mu.Lock() + shardLocked := true + t.Cleanup(func() { + if shardLocked { + blockedShard.mu.Unlock() + } + }) + type admissionResult struct { + admitted bool + err error + } + result := make(chan admissionResult, 1) + go func() { + admitted, err := memory.admit(ctx, newMemoryTestEntry(namespace, newKey, make([]byte, 128*1024))) + result <- admissionResult{admitted: admitted, err: err} + }() + deadline := time.Now().Add(time.Second) + for memory.state.evictionCursor.Load() == 0 { + if time.Now().After(deadline) { + blockedShard.mu.Unlock() + shardLocked = false + t.Fatal("admission did not reach eviction planning") + } + time.Sleep(time.Millisecond) + } + cancel() + blockedShard.mu.Unlock() + shardLocked = false + + var response admissionResult + select { + case response = <-result: + case <-time.After(time.Second): + t.Fatal("cancelled admission did not return") + } + assert.False(t, response.admitted) + assert.IsError(t, response.err, context.Canceled) + for _, key := range oldKeys { + reader, _, err := cache.Open(t.Context(), key) + assert.NoError(t, err) + assert.NoError(t, reader.Close()) + } + _, _, err := cache.Open(t.Context(), newKey) + assert.IsError(t, err, os.ErrNotExist) +} + +func TestMemoryCancellationWhileWaitingForAdmissionIsNotAdmitted(t *testing.T) { + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + memory, err := NewMemory(ctx, MemoryConfig{LimitMB: 2, InflightLimitMB: 1, MaxTTL: time.Hour}) + assert.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, memory.Close()) }) + key := NewKey("cancel-race") + w, err := memory.Create(t.Context(), key, nil, time.Hour) + assert.NoError(t, err) + _, err = w.Write([]byte("cancelled")) + assert.NoError(t, err) + controlledContext := &admissionCancellationContext{Context: t.Context(), firstCheck: make(chan struct{})} + memoryWriter := memoryTestWriter(t, w) + memoryWriter.ctx = controlledContext + + shard := memory.shard("", key) + shard.mu.Lock() + closed := make(chan error, 1) + go func() { closed <- w.Close() }() + <-controlledContext.firstCheck + controlledContext.cancelled.Store(true) + shard.mu.Unlock() + + assert.IsError(t, <-closed, context.Canceled) + _, _, err = memory.Open(t.Context(), key) + assert.IsError(t, err, os.ErrNotExist) + assert.Equal(t, int64(0), memory.state.inflightCharge.Load()) + assert.Equal(t, int64(0), memory.state.hardLimitCharge.Load()) +} + +func TestMemoryConcurrentAdmissionKeepsAccountingBounded(t *testing.T) { + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + memory, err := NewMemory(ctx, MemoryConfig{LimitMB: 1, InflightLimitMB: 1, MaxTTL: time.Hour}) + assert.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, memory.Close()) }) + const writes = 64 + keys := make([]Key, writes) + writeErrors := make(chan error, writes) + var wg sync.WaitGroup + for index := range writes { + keys[index] = NewKey(fmt.Sprintf("concurrent-%d", index)) + wg.Go(func() { + writer, err := memory.Create(t.Context(), keys[index], http.Header{}, time.Hour) + if err == nil { + _, err = writer.Write(make([]byte, 32*1024)) + } + if err == nil { + err = writer.Close() + } + writeErrors <- err + }) + } + wg.Wait() + close(writeErrors) + for err := range writeErrors { + assert.NoError(t, err) + } + + stats, err := memory.Stats(t.Context()) + assert.NoError(t, err) + assert.Equal(t, int64(0), memory.state.inflightCharge.Load()) + assert.True(t, memory.state.hardLimitCharge.Load() <= memory.state.limitBytes) + actualObjects := int64(0) + actualSize := int64(0) + for _, key := range keys { + reader, _, err := memory.Open(t.Context(), key) + if errors.Is(err, os.ErrNotExist) { + continue + } + assert.NoError(t, err) + data, err := io.ReadAll(reader) + assert.NoError(t, err) + assert.NoError(t, reader.Close()) + actualObjects++ + actualSize += int64(len(data)) + } + assert.Equal(t, actualObjects, stats.Objects) + assert.Equal(t, actualSize, stats.Size) +} + +func TestMemoryConcurrentReplacementDoesNotExposeCapacity(t *testing.T) { + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + memory, err := NewMemory(ctx, MemoryConfig{LimitMB: 0, MaxTTL: time.Hour}) + assert.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, memory.Close()) }) + const namespace Namespace = "replacement-accounting" + payload := make([]byte, 64*1024) + keys := make([]Key, memoryShardCount) + for shardIndex := range memoryShardCount { + keys[shardIndex] = memoryKeysForShard(t, namespace, shardIndex, 1)[0] + admitMemoryTestEntry(t.Context(), t, memory, newMemoryTestEntry(namespace, keys[shardIndex], payload)) + } + memory.state.limitBytes = memory.state.retainedCharge.Load() + retainedCeiling := memory.state.limitBytes + memoryEntryMinimumCharge + memory.state.limitBytes = retainedCeiling + memory.state.retainedTarget = retainedCeiling - memoryEntryMinimumCharge + + workContext, cancelWork := context.WithCancel(t.Context()) + start := make(chan struct{}) + var stopped atomic.Bool + var observedMaximum atomic.Int64 + var replacements atomic.Int64 + var admissions atomic.Int64 + var wait sync.WaitGroup + for _, key := range keys { + wait.Go(func() { + <-start + for !stopped.Load() { + admitted, err := memory.admit(workContext, newMemoryTestEntry(namespace, key, payload)) + if err != nil { + return + } + if admitted { + replacements.Add(1) + } + } + }) + } + for worker := range 8 { + wait.Go(func() { + <-start + for index := 0; !stopped.Load(); index++ { + key := NewKey(fmt.Sprintf("replacement-admission-%d-%d", worker, index)) + admitted, err := memory.admit(workContext, newMemoryTestEntry(namespace, key, nil)) + if err != nil { + return + } + if admitted { + admissions.Add(1) + } + } + }) + } + + close(start) + deadline := time.Now().Add(250 * time.Millisecond) + for time.Now().Before(deadline) { + size := memory.state.hardLimitCharge.Load() + for maximum := observedMaximum.Load(); size > maximum; maximum = observedMaximum.Load() { + if observedMaximum.CompareAndSwap(maximum, size) { + break + } + } + if size > retainedCeiling { + break + } + } + stopped.Store(true) + cancelWork() + wait.Wait() + assert.True(t, replacements.Load() > 0) + assert.True(t, admissions.Load() > 0) + assert.True(t, observedMaximum.Load() <= retainedCeiling, + "retained charge exceeded ceiling: got %d, ceiling %d", observedMaximum.Load(), retainedCeiling) +} + +func TestMemoryUncommittedPlanDoesNotHideEntries(t *testing.T) { + memory := newMemoryTestCache(t) + const namespace Namespace = "transactional-visibility" + keys := make([]Key, memoryShardCount) + for shardIndex := range memoryShardCount { + keys[shardIndex] = memoryKeysForShard(t, namespace, shardIndex, 1)[0] + admitMemoryTestEntry(t.Context(), t, memory, newMemoryTestEntry(namespace, keys[shardIndex], nil)) + } + memory.state.limitBytes = memory.state.retainedCharge.Load() + var planBuffer [maxMemoryEvictionsPerWrite]memoryPlannedEviction + plannedEntries := memory.planEvictions(32*1024, "", Key{}, planBuffer[:]) + assert.True(t, plannedEntries > 0) + candidate := planBuffer[0].entry + cache := memory.Namespace(candidate.namespace) + reader, _, err := cache.Open(t.Context(), candidate.key) + assert.NoError(t, err) + assert.NoError(t, reader.Close()) +} + +func TestMemoryConcurrentPlanPreservesPostPlanReference(t *testing.T) { + memory := newMemoryTestCache(t) + key := NewKey("concurrent-plan-reference") + entry := newMemoryTestEntry("", key, []byte("value")) + admitMemoryTestEntry(t.Context(), t, memory, entry) + + var firstPlan [maxMemoryEvictionsPerWrite]memoryPlannedEviction + firstPlanSize := memory.planEvictions(entry.charge, "protected", Key{}, firstPlan[:]) + assert.Equal(t, 1, firstPlanSize) + + reader, _, err := memory.Open(t.Context(), key) + assert.NoError(t, err) + assert.NoError(t, reader.Close()) + + var concurrentPlan [maxMemoryEvictionsPerWrite]memoryPlannedEviction + concurrentPlanSize := memory.planEvictions(entry.charge, "protected", Key{}, concurrentPlan[:]) + assert.Equal(t, 0, concurrentPlanSize) + + memory.commitMemoryEvictionPlan(t.Context(), firstPlan[:firstPlanSize], 0) + _, err = memory.Stat(t.Context(), key) + assert.NoError(t, err) +} + +func TestMemoryLargeEntryCanDisplaceSmallEntries(t *testing.T) { + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + memory, err := NewMemory(ctx, MemoryConfig{LimitMB: 2, InflightLimitMB: 1, MaxTTL: time.Hour}) + assert.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, memory.Close()) }) + cache := memory.Namespace("mixed-entry-sizes") + for index := range 256 { + writeMemoryTestEntry(t, cache, NewKey(fmt.Sprintf("small-%d", index)), nil, time.Hour) + } + largeKey := NewKey("large") + writeMemoryTestEntry(t, cache, largeKey, make([]byte, 512*1024), time.Hour) + + reader, _, err := cache.Open(t.Context(), largeKey) + assert.NoError(t, err) + data, err := io.ReadAll(reader) + assert.NoError(t, err) + assert.NoError(t, reader.Close()) + assert.Equal(t, 512*1024, len(data)) + assert.True(t, memory.state.hardLimitCharge.Load() <= memory.state.limitBytes) +} + +func BenchmarkMemoryAdmissionAtCapacity(b *testing.B) { + for _, entryCount := range []int{1_000, 10_000, 100_000} { + b.Run(fmt.Sprintf("entries=%d", entryCount), func(b *testing.B) { + expiresAt := time.Now().Add(time.Hour) + memory := newMemoryBenchmarkCache(b, entryCount, expiresAt) + b.ReportAllocs() + b.ResetTimer() + for index := range b.N { + admitMemoryTestEntry(b.Context(), b, memory, newMemoryBenchmarkEntry(entryCount+index, expiresAt)) + } + }) + } +} + +func BenchmarkMemoryParallelAdmissionAtCapacity(b *testing.B) { + const entryCount = 10_000 + expiresAt := time.Now().Add(time.Hour) + memory := newMemoryBenchmarkCache(b, entryCount, expiresAt) + var nextKey atomic.Uint64 + var attempts atomic.Uint64 + var admittedCount atomic.Uint64 + nextKey.Store(entryCount) + b.ReportAllocs() + b.ResetTimer() + b.RunParallel(func(pb *testing.PB) { + for pb.Next() { + attempts.Add(1) + index := int(nextKey.Add(1)) + admitted, err := memory.admit(b.Context(), newMemoryBenchmarkEntry(index, expiresAt)) + if err != nil { + b.Error(err) + return + } + if admitted { + admittedCount.Add(1) + } + } + }) + b.ReportMetric(100*float64(admittedCount.Load())/float64(attempts.Load()), "accepted-%") +} + +func BenchmarkMemoryConfiguredCreateAtCapacity(b *testing.B) { + for _, entryCount := range []int{1_000, 10_000, 100_000} { + b.Run(fmt.Sprintf("entries=%d", entryCount), func(b *testing.B) { + expiresAt := time.Now().Add(time.Hour) + memory := newConfiguredMemoryBenchmarkCache(b, entryCount, expiresAt) + headers := http.Header{"Content-Length": {"1"}} + payload := []byte{1} + b.ReportAllocs() + b.ResetTimer() + for index := range b.N { + entry := newMemoryBenchmarkEntry(entryCount+index, expiresAt) + writer, err := memory.Create(b.Context(), entry.key, headers, time.Hour, WithETag("benchmark")) + assert.NoError(b, err) + memoryTestWriter(b, writer) + written, err := writer.Write(payload) + assert.NoError(b, err) + assert.Equal(b, len(payload), written) + assert.NoError(b, writer.Close()) + shard := memory.shard("", entry.key) + shard.mu.RLock() + _, admitted := shard.entry("", entry.key) + shard.mu.RUnlock() + assert.True(b, admitted) + } + }) + } +} + +func newMemoryBenchmarkCache(b *testing.B, entryCount int, expiresAt time.Time) *Memory { + b.Helper() + state, err := newMemoryState(MemoryConfig{}) + assert.NoError(b, err) + state.retainedTarget = int64(entryCount * memoryEntryMinimumCharge) + state.limitBytes = state.retainedTarget + memoryEntryMinimumCharge + memory := &Memory{state: state} + for index := range entryCount { + admitMemoryTestEntry(b.Context(), b, memory, newMemoryBenchmarkEntry(index, expiresAt)) + } + return memory +} + +func newConfiguredMemoryBenchmarkCache(b *testing.B, entryCount int, expiresAt time.Time) *Memory { + b.Helper() + config := MemoryConfig{LimitMB: 1, InflightLimitMB: 1, MaxTTL: time.Hour} + state, err := newMemoryState(config) + assert.NoError(b, err) + state.retainedTarget = int64(entryCount * memoryEntryMinimumCharge) + state.limitBytes = state.retainedTarget + 1024*1024 + memory := &Memory{config: config, state: state} + for index := range entryCount { + admitMemoryTestEntry(b.Context(), b, memory, newMemoryBenchmarkEntry(index, expiresAt)) + } + return memory +} + +func newMemoryBenchmarkEntry(index int, expiresAt time.Time) *memoryEntry { + var key Key + binary.LittleEndian.PutUint64(key[:8], uint64(index)) + entry := &memoryEntry{namespace: "benchmark", key: key, data: []byte{1}, expiresAt: expiresAt} + entry.charge = memoryEntryCharge(entry.namespace, entry.data, nil) + return entry +} + +func BenchmarkMemoryParallelHotHits(b *testing.B) { + _, ctx := logging.Configure(b.Context(), logging.Config{Level: slog.LevelError}) + memory, err := NewMemory(ctx, MemoryConfig{LimitMB: 1, MaxTTL: time.Hour}) + assert.NoError(b, err) + key := NewKey("parallel-hot-hit") + writer, err := memory.Create(ctx, key, http.Header{"Content-Length": {"1"}}, time.Hour) + assert.NoError(b, err) + _, err = writer.Write([]byte{1}) + assert.NoError(b, err) + assert.NoError(b, writer.Close()) + b.Cleanup(func() { assert.NoError(b, memory.Close()) }) + b.ReportAllocs() + b.ResetTimer() + + b.RunParallel(func(pb *testing.PB) { + for pb.Next() { + reader, _, err := memory.Open(ctx, key) + if err != nil { + b.Error(err) + return + } + if err := reader.Close(); err != nil { + b.Error(err) + return + } + } + }) +} diff --git a/internal/cache/tiered_test.go b/internal/cache/tiered_test.go index 7a3eff8..aa9e6dd 100644 --- a/internal/cache/tiered_test.go +++ b/internal/cache/tiered_test.go @@ -8,6 +8,7 @@ import ( "log/slog" "net/http" "os" + "strconv" "strings" "sync" "sync/atomic" @@ -354,6 +355,33 @@ func TestTieredCreateUsesSameETagInEveryTier(t *testing.T) { assert.Equal(t, lowerHeaders.Get(cache.ETagKey), upperHeaders.Get(cache.ETagKey)) } +func TestTieredCreateContinuesWhenMemoryTierDeclinesAdmission(t *testing.T) { + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + constrained, err := cache.NewMemory(ctx, cache.MemoryConfig{LimitMB: 1, InflightLimitMB: 1, MaxTTL: time.Hour}) + assert.NoError(t, err) + authoritative, err := cache.NewMemory(ctx, cache.MemoryConfig{LimitMB: 2, MaxTTL: time.Hour}) + assert.NoError(t, err) + tiered := newTiered(ctx, constrained, authoritative) + t.Cleanup(func() { assert.NoError(t, tiered.Close()) }) + key := cache.NewKey("declined-memory-admission") + content := make([]byte, 1024*1024) + content[0] = 1 + headers := http.Header{"Content-Length": {strconv.Itoa(len(content))}} + + writer, err := tiered.Create(ctx, key, headers, time.Hour) + assert.NoError(t, err) + written, err := writer.Write(content) + assert.NoError(t, err) + assert.Equal(t, len(content), written) + assert.NoError(t, writer.Close()) + + _, _, err = constrained.Open(ctx, key) + assert.IsError(t, err, os.ErrNotExist) + reader, _, err := authoritative.Open(ctx, key) + assert.NoError(t, err) + assert.Equal(t, content, readAllAndClose(t, reader)) +} + func TestTieredRequiresMetadataStore(t *testing.T) { _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelDebug}) lower, err := cache.NewMemory(ctx, cache.MemoryConfig{LimitMB: 1024, MaxTTL: time.Hour}) From 647d9824d725a77255c9cf853cfcec2c6c79d68b Mon Sep 17 00:00:00 2001 From: Ravi Agarwal Date: Mon, 31 Aug 2026 23:38:32 -0700 Subject: [PATCH 2/3] memory: harden admission and cleanup paths Cap declared-length buffers at their promised size, reject finite inflight limits that leave no retained capacity, and keep unknown-length starter allocations aligned with the per-entry accounting floor. Restore the io.WriterTo hit path and make Tiered.Create abort both completed and late writers on cancellation or backend failure, preventing reservation leaks and nil-writer returns. Record low-cardinality admission-decline reasons, clarify bounded CLOCK and capacity semantics, and cover the lifecycle and configuration edges under race. Verification: bin/just fmt; bin/just lint; bin/just test. --- README.md | 53 ++++-- internal/cache/memory.go | 145 ++++++++++----- internal/cache/memory_internal_test.go | 233 ++++++++++++++++++++++++- internal/cache/memory_metrics.go | 46 +++++ internal/cache/tiered.go | 63 +++++-- internal/cache/tiered_test.go | 107 +++++++++++- 6 files changed, 561 insertions(+), 86 deletions(-) create mode 100644 internal/cache/memory_metrics.go diff --git a/README.md b/README.md index 51e1312..b39dc07 100644 --- a/README.md +++ b/README.md @@ -151,28 +151,43 @@ copies before falling through to the authoritative tier. ### Memory In-memory sharded CLOCK cache with bounded admission and eviction work. Cache -hits lock only the shard containing the requested key. Each admission and trim -examines at most 64 victims, regardless of cache cardinality. - -`limit-mb` is a hard retained-memory accounting ceiling, not a process RSS +hits lock only the shard containing the requested key. Each trim pass scans at +most 64 entries per shard and commits at most 64 victims, regardless of cache +cardinality. Recently referenced entries receive a CLOCK second chance. If one +bounded pass cannot find enough cold entries, Cachew declines the optional +memory copy instead of extending the scan or blocking unrelated hits; later +admissions continue from the advanced CLOCK hands. + +`limit-mb` is a hard accounted-memory ceiling, not a process RSS limit. Accounting includes object buffers, estimated metadata, and buffers held -by active readers. `Stats.Capacity` reports this ceiling; `Stats.Size` reports -payload bytes and can differ because it excludes charged metadata and spare -buffer capacity. Go runtime and allocator overhead can make RSS differ from -both values. +by active readers. Every retained entry and incomplete writer has a minimum +4 KiB charge so collections of tiny objects cannot leave Go object and map +overhead unbounded. This means a 1 GiB cache retains at most roughly 262,000 +objects even when their payloads are smaller. `Stats.Capacity` reports the hard +accounting ceiling; `Stats.Size` reports payload bytes and can differ because it +excludes charged metadata and spare buffer capacity. Go runtime and allocator +overhead can make RSS differ from both values. `limit-mb = 0` disables the hard +ceiling and permits unlimited retained accounting. Incomplete writes remain unbounded when `inflight-limit-mb` is zero, preserving -the behavior of configurations written before this option existed. A positive -value limits aggregate incomplete writes and reserves that amount inside -`limit-mb`: retained entries are trimmed toward `limit-mb - -inflight-limit-mb`, and retained plus incomplete accounting cannot exceed -`limit-mb`. Writes that cannot obtain capacity within the bounded admission -work bypass the memory tier without interrupting other cache tiers. Declared -content lengths are validated against these limits, but buffers grow only as -body bytes arrive. Buffer growth transfers the existing accounting reservation -to the larger capacity; the allocator may briefly retain both allocations, so -process RSS can transiently exceed the accounting ceiling by the old buffer's -capacity. +the behavior of configurations written before this option existed. For a +finite `limit-mb`, a positive `inflight-limit-mb` must be smaller and reserves +that amount inside the hard ceiling: retained entries are trimmed toward +`limit-mb - inflight-limit-mb`, and retained plus incomplete accounting cannot +exceed `limit-mb`. With unlimited retention, a positive inflight limit still +bounds incomplete writes independently. Writes that cannot obtain capacity +within the bounded admission work bypass the memory tier without interrupting +other cache tiers. The `cachew.memory.admission_declines_total` counter reports +these events by low-cardinality `reason`. + +Declared content lengths are validated against the limits, but buffers grow +only as body bytes arrive and never beyond the declared length. Unknown-length +bodies use a 4 KiB minimum growth allocation for smaller writes and then grow +geometrically; this can retain spare capacity but avoids another full-body copy +at publication, and all spare capacity remains charged. Buffer growth transfers +the existing accounting reservation to the larger capacity; the allocator may +briefly retain both allocations, so process RSS can transiently exceed the +accounting ceiling by the old buffer's capacity. ```hcl memory { diff --git a/internal/cache/memory.go b/internal/cache/memory.go index e62fce4..bdd8272 100644 --- a/internal/cache/memory.go +++ b/internal/cache/memory.go @@ -29,7 +29,7 @@ const ( memoryHeaderEntryCharge = 64 memoryHeaderValueCharge = 16 memoryWriterMinimumCharge = 4 * 1024 - memoryWriterInitialCapacity = 64 * 1024 + memoryWriterInitialCapacity = memoryWriterMinimumCharge memoryBytesPerMegabyte = 1024 * 1024 fnv64Offset = 14695981039346656037 fnv64Prime = 1099511628211 @@ -47,8 +47,8 @@ func RegisterMemory(r *Registry) { // MemoryConfig keeps incomplete-write protection opt-in so existing zero-valued configurations remain compatible. type MemoryConfig struct { - LimitMB int `hcl:"limit-mb,optional" help:"Maximum retained size of the memory cache in megabytes (defaults to 1GB); positive inflight-limit-mb shares this budget." default:"1024"` - InflightLimitMB int `hcl:"inflight-limit-mb,optional" help:"Maximum aggregate incomplete writes in megabytes (0 disables the sub-limit for compatibility)." default:"0"` + LimitMB int `hcl:"limit-mb,optional" help:"Maximum accounted memory in megabytes (defaults to 1GB, 0 is unlimited); positive inflight-limit-mb shares this budget." default:"1024"` + InflightLimitMB int `hcl:"inflight-limit-mb,optional" help:"Maximum aggregate incomplete writes in megabytes (0 disables the sub-limit); must be smaller than a finite limit-mb." default:"0"` MaxTTL time.Duration `hcl:"max-ttl,optional" help:"Maximum time-to-live for entries in the memory cache (defaults to 1 hour)." default:"1h"` } @@ -131,6 +131,9 @@ func (s *memoryShard) insert(entry *memoryEntry) { s.append(entry) } +// With finite limits, the hard-budget counter overlaps retained and inflight +// charges only when an inflight sub-limit is configured, preventing independent +// reservations from crossing the process-wide accounting ceiling. type memoryState struct { shards []memoryShard limitBytes int64 @@ -143,6 +146,7 @@ type memoryState struct { objectCount atomic.Int64 evictionCursor atomic.Uint32 closed atomic.Bool + metrics memoryMetricRecorder } // Memory shares capacity across namespace views so each view cannot consume limit-mb independently. @@ -171,15 +175,21 @@ func newMemoryState(config MemoryConfig) (*memoryState, error) { if err != nil { return nil, err } - inflightLimit := memoryInflightLimit(configuredInflightBytes, limitBytes) - retainedTarget := memoryRetainedTarget(limitBytes, inflightLimit) + if limitBytes > 0 && configuredInflightBytes >= limitBytes { + return nil, errors.New("inflight-limit-mb must be less than limit-mb when limit-mb is finite") + } + retainedTarget := int64(0) + if limitBytes > 0 { + retainedTarget = limitBytes - configuredInflightBytes + } shards := make([]memoryShard, memoryShardCount) for index := range shards { shards[index].entries = make(map[Namespace]map[Key]*memoryEntry) } return &memoryState{ shards: shards, limitBytes: limitBytes, - retainedTarget: retainedTarget, inflightLimit: inflightLimit, + retainedTarget: retainedTarget, inflightLimit: configuredInflightBytes, + metrics: newMemoryMetrics(), }, nil } @@ -237,23 +247,6 @@ func memoryEntryCharge(namespace Namespace, data []byte, headers http.Header) in return max(charge, int64(memoryEntryMinimumCharge)) } -func memoryInflightLimit(configuredBytes, limitBytes int64) int64 { - if configuredBytes <= 0 { - return 0 - } - if limitBytes > 0 { - return min(configuredBytes, limitBytes) - } - return configuredBytes -} - -func memoryRetainedTarget(limitBytes, inflightLimit int64) int64 { - if limitBytes <= 0 { - return 0 - } - return limitBytes - min(limitBytes, max(inflightLimit, 0)) -} - func reserveBounded(counter *atomic.Int64, limit, amount int64) bool { if amount <= 0 { return true @@ -270,17 +263,6 @@ func reserveBounded(counter *atomic.Int64, limit, amount int64) bool { return false } -func reserveCounter(counter *atomic.Int64, limit, amount int64) bool { - if amount <= 0 { - return true - } - if limit <= 0 { - counter.Add(amount) - return true - } - return reserveBounded(counter, limit, amount) -} - type memoryPlannedEviction struct { shard *memoryShard entry *memoryEntry @@ -330,10 +312,15 @@ func (m *Memory) insertActiveLocked(shard *memoryShard, entry *memoryEntry) { } func (m *Memory) reserveRetained(retainedLimit, amount int64) bool { - if !reserveCounter(&m.state.retainedCharge, retainedLimit, amount) { + if m.state.limitBytes == 0 { + m.state.retainedCharge.Add(amount) + m.state.hardLimitCharge.Add(amount) + return true + } + if !reserveBounded(&m.state.retainedCharge, retainedLimit, amount) { return false } - if reserveCounter(&m.state.hardLimitCharge, m.state.limitBytes, amount) { + if reserveBounded(&m.state.hardLimitCharge, m.state.limitBytes, amount) { return true } m.state.retainedCharge.Add(-amount) @@ -349,7 +336,11 @@ const ( func (m *Memory) reserveAdmission(mode memoryAdmissionMode, retainedLimit, amount int64) bool { if mode == memoryAdmissionHasAllocation { - return reserveCounter(&m.state.retainedCharge, retainedLimit, amount) + if m.state.limitBytes == 0 { + m.state.retainedCharge.Add(amount) + return true + } + return reserveBounded(&m.state.retainedCharge, retainedLimit, amount) } return m.reserveRetained(retainedLimit, amount) } @@ -424,6 +415,9 @@ func (m *Memory) planEvictions( protectedKey Key, planned []memoryPlannedEviction, ) int { + // Planning leaves candidates visible to readers. The later commit phase + // revalidates each reference epoch so a hit between these phases wins its + // CLOCK second chance without holding multiple shard locks at once. if needed <= 0 { return 0 } @@ -473,6 +467,9 @@ func (m *Memory) planEvictions( } func (m *Memory) commitMemoryEvictionPlan(ctx context.Context, planned []memoryPlannedEviction, target int64) { + // Pointer identity rejects replacements and the epoch check rejects + // post-plan hits, so only the exact cold generation originally planned can + // be removed. now := time.Now() for start := 0; start < len(planned); { shard := planned[start].shard @@ -551,9 +548,6 @@ func (m *Memory) admit(ctx context.Context, entry *memoryEntry) (bool, error) { if err := ctx.Err(); err != nil { return false, errors.WithStack(err) } - if admitted, err = m.tryAdmission(ctx, entry, m.state.limitBytes, memoryAdmissionNeedsAllocation); admitted || err != nil { - return admitted, err - } m.trimForAdmission(ctx, entry) admitted, err = m.tryAdmission(ctx, entry, m.state.limitBytes, memoryAdmissionNeedsAllocation) if !admitted || err != nil { @@ -627,6 +621,8 @@ type memoryReader struct { closed atomic.Bool } +var _ io.WriterTo = (*memoryReader)(nil) + func (r *memoryReader) Read(p []byte) (int, error) { if r.closed.Load() { return 0, os.ErrClosed @@ -639,6 +635,25 @@ func (r *memoryReader) Read(p []byte) (int, error) { return n, nil } +func (r *memoryReader) WriteTo(destination io.Writer) (int64, error) { + if r.closed.Load() { + return 0, os.ErrClosed + } + if r.offset >= len(r.data) { + return 0, nil + } + remaining := len(r.data) - r.offset + written, err := destination.Write(r.data[r.offset:]) + if written < 0 || written > remaining { + return 0, errors.Errorf("invalid Write count %d", written) + } + r.offset += written + if written != remaining && err == nil { + err = io.ErrShortWrite + } + return int64(written), errors.WithStack(err) +} + func (r *memoryReader) Close() error { if r.closed.Swap(true) { return nil @@ -675,9 +690,11 @@ func (m *Memory) Create(ctx context.Context, key Key, headers http.Header, ttl t metadataCharge := memoryMetadataCharge(m.namespace, clonedHeaders) baseCharge := max(metadataCharge, int64(memoryWriterMinimumCharge)) if contentLength >= 0 && m.state.limitBytes > 0 && contentLength > m.state.limitBytes-metadataCharge { + m.state.metrics.recordDecline(ctx, memoryDeclineDeclaredHardLimit) return &noOpWriter{}, nil } if contentLength >= 0 && m.state.inflightLimit > 0 && contentLength > m.state.inflightLimit-baseCharge { + m.state.metrics.recordDecline(ctx, memoryDeclineDeclaredInflightLimit) return &noOpWriter{}, nil } ctx, cancel := context.WithCancelCause(ctx) @@ -696,6 +713,7 @@ func (m *Memory) Create(ctx context.Context, key Key, headers http.Header, ttl t cancel: cancel, } if !writer.reserve(baseCharge) { + m.state.metrics.recordDecline(ctx, memoryDeclineWriterReservation) cancel(nil) return &noOpWriter{}, nil } @@ -784,8 +802,16 @@ func (w *memoryWriter) Write(p []byte) (int, error) { tooLarge := w.limitBytes > 0 && int64(len(p)) > w.limitBytes-w.baseCharge-buffered longerThanDeclared := w.expectedLength >= 0 && int64(len(p)) > w.expectedLength-buffered needed := buffered + int64(len(p)) - if tooLarge || longerThanDeclared || !w.ensureCapacity(needed) { - w.discard() + if tooLarge { + w.decline(memoryDeclineBodyHardLimit) + return len(p), nil + } + if longerThanDeclared { + w.decline(memoryDeclineContentLengthMismatch) + return len(p), nil + } + if !w.ensureCapacity(needed) { + w.decline(memoryDeclineWriterReservation) return len(p), nil } w.data = append(w.data, p...) @@ -805,6 +831,9 @@ func (w *memoryWriter) maximumBodyCapacity() int64 { if w.inflightLimit > 0 { maximum = min(maximum, w.inflightLimit-w.baseCharge) } + if w.expectedLength >= 0 { + maximum = min(maximum, w.expectedLength) + } return maximum } @@ -868,12 +897,18 @@ func (w *memoryWriter) reserve(amount int64) bool { } func (w *memoryWriter) tryReserve(amount int64) bool { - if !reserveCounter(&w.cache.state.inflightCharge, w.inflightLimit, amount) { + if w.inflightLimit == 0 { + w.cache.state.inflightCharge.Add(amount) + } else if !reserveBounded(&w.cache.state.inflightCharge, w.inflightLimit, amount) { return false } - if w.budgeted && !reserveCounter(&w.cache.state.hardLimitCharge, w.limitBytes, amount) { - w.cache.state.inflightCharge.Add(-amount) - return false + if w.budgeted { + if w.limitBytes == 0 { + w.cache.state.hardLimitCharge.Add(amount) + } else if !reserveBounded(&w.cache.state.hardLimitCharge, w.limitBytes, amount) { + w.cache.state.inflightCharge.Add(-amount) + return false + } } w.reservedBytes += amount return true @@ -914,6 +949,14 @@ func (w *memoryWriter) discard() { w.discarded = true } +func (w *memoryWriter) decline(reason memoryDeclineReason) { + if w.discarded { + return + } + w.cache.state.metrics.recordDecline(w.ctx, reason) + w.discard() +} + func (w *memoryWriter) Abort(err error) error { w.cancel(err) return w.Close() @@ -933,7 +976,7 @@ func (w *memoryWriter) Close() error { return nil } if w.expectedLength >= 0 && int64(len(w.data)) != w.expectedLength { - w.discard() + w.decline(memoryDeclineContentLengthMismatch) return nil } @@ -948,15 +991,21 @@ func (w *memoryWriter) Close() error { } entry.charge = memoryEntryCharge(entry.namespace, entry.data, entry.headers) if !w.budgeted { - _, err := w.cache.admit(w.ctx, entry) + admitted, err := w.cache.admit(w.ctx, entry) + if !admitted && err == nil { + w.cache.state.metrics.recordDecline(w.ctx, memoryDeclineAdmissionLimit) + } return errors.WithStack(err) } if entry.charge > w.reservedBytes && !w.reserve(entry.charge-w.reservedBytes) { + w.cache.state.metrics.recordDecline(w.ctx, memoryDeclineWriterReservation) return nil } admitted, err := w.cache.admitReserved(w.ctx, entry) if admitted { w.transferReservation(entry.charge) + } else if err == nil { + w.cache.state.metrics.recordDecline(w.ctx, memoryDeclineAdmissionLimit) } return errors.WithStack(err) } diff --git a/internal/cache/memory_internal_test.go b/internal/cache/memory_internal_test.go index 2f832aa..6b13f54 100644 --- a/internal/cache/memory_internal_test.go +++ b/internal/cache/memory_internal_test.go @@ -84,6 +84,48 @@ func memoryTestReader(t testing.TB, reader io.ReadCloser) *memoryReader { return memoryReader } +type writeCountingBuffer struct { + bytes.Buffer + writes int +} + +func (w *writeCountingBuffer) Write(p []byte) (int, error) { + w.writes++ + return w.Buffer.Write(p) +} + +type writeOnlyDiscard struct{} + +func (writeOnlyDiscard) Write(p []byte) (int, error) { return len(p), nil } + +type recordingMemoryMetrics struct { + mu sync.Mutex + declines map[memoryDeclineReason]int +} + +func (r *recordingMemoryMetrics) recordDecline(_ context.Context, reason memoryDeclineReason) { + r.mu.Lock() + defer r.mu.Unlock() + if r.declines == nil { + r.declines = make(map[memoryDeclineReason]int) + } + r.declines[reason]++ +} + +func (r *recordingMemoryMetrics) declineCount(reason memoryDeclineReason) int { + r.mu.Lock() + defer r.mu.Unlock() + return r.declines[reason] +} + +func memoryWithRecordingMetrics(t *testing.T, config MemoryConfig) (*Memory, *recordingMemoryMetrics) { + t.Helper() + memory := newMemoryTestCacheWithConfig(t, config) + metrics := &recordingMemoryMetrics{} + memory.state.metrics = metrics + return memory, metrics +} + func admitMemoryTestEntry(ctx context.Context, t testing.TB, memory *Memory, entry *memoryEntry) { t.Helper() admitted, err := memory.admit(ctx, entry) @@ -650,6 +692,47 @@ func TestMemoryDeclaredLengthDoesNotAllocateBeforeWrite(t *testing.T) { assert.Equal(t, int64(0), memory.state.inflightCharge.Load()) } +func TestMemoryDeclaredLengthCapsFinalBufferCapacity(t *testing.T) { + memory := newMemoryTestCacheWithConfig(t, MemoryConfig{LimitMB: 4, InflightLimitMB: 2, MaxTTL: time.Hour}) + declaredLength := 1024*1024 + 1 + payload := bytes.Repeat([]byte{0x5a}, declaredLength) + writer, err := memory.Create(t.Context(), NewKey("exact-declared-capacity"), http.Header{ + "Content-Length": {strconv.Itoa(declaredLength)}, + }, time.Hour) + assert.NoError(t, err) + written, err := writer.Write(payload[:1024*1024]) + assert.NoError(t, err) + assert.Equal(t, 1024*1024, written) + written, err = writer.Write(payload[1024*1024:]) + assert.NoError(t, err) + assert.Equal(t, 1, written) + memoryWriter := memoryTestWriter(t, writer) + assert.Equal(t, declaredLength, cap(memoryWriter.data)) + assert.NoError(t, writer.Close()) +} + +func TestMemoryReaderPreservesWriterToFastPath(t *testing.T) { + memory := newMemoryTestCache(t) + key := NewKey("writer-to-fast-path") + payload := bytes.Repeat([]byte{0x6b}, 96*1024) + writeMemoryTestEntry(t, memory, key, payload, time.Hour) + reader, _, err := memory.Open(t.Context(), key) + assert.NoError(t, err) + prefix := make([]byte, 1024) + read, err := reader.Read(prefix) + assert.NoError(t, err) + assert.Equal(t, len(prefix), read) + writerTo, ok := reader.(io.WriterTo) + assert.True(t, ok) + destination := &writeCountingBuffer{} + written, err := writerTo.WriteTo(destination) + assert.NoError(t, err) + assert.Equal(t, int64(len(payload)-len(prefix)), written) + assert.Equal(t, 1, destination.writes) + assert.Equal(t, payload[len(prefix):], destination.Bytes()) + assert.NoError(t, reader.Close()) +} + func TestMemoryUnknownLengthAdmissionDoesNotDependOnWriteChunks(t *testing.T) { memory := newMemoryTestCacheWithConfig(t, MemoryConfig{LimitMB: 2, InflightLimitMB: 1, MaxTTL: time.Hour}) key := NewKey("chunked-unknown-length") @@ -671,6 +754,88 @@ func TestMemoryUnknownLengthAdmissionDoesNotDependOnWriteChunks(t *testing.T) { assert.Equal(t, payload, stored) } +func TestMemorySmallUnknownLengthUsesMinimumCapacity(t *testing.T) { + memory := newMemoryTestCache(t) + writer, err := memory.Create(t.Context(), NewKey("small-unknown-length"), nil, time.Hour) + assert.NoError(t, err) + written, err := writer.Write([]byte{1}) + assert.NoError(t, err) + assert.Equal(t, 1, written) + assert.Equal(t, memoryWriterInitialCapacity, cap(memoryTestWriter(t, writer).data)) + assert.NoError(t, writer.Close()) +} + +func TestMemoryRecordsAdmissionDeclineReasons(t *testing.T) { + t.Run("declared hard limit", func(t *testing.T) { + memory, metrics := memoryWithRecordingMetrics(t, MemoryConfig{LimitMB: 1, MaxTTL: time.Hour}) + writer, err := memory.Create(t.Context(), NewKey("declared-hard-limit"), http.Header{ + "Content-Length": {strconv.Itoa(2 * 1024 * 1024)}, + }, time.Hour) + assert.NoError(t, err) + assert.NoError(t, writer.Close()) + assert.Equal(t, 1, metrics.declineCount(memoryDeclineDeclaredHardLimit)) + }) + + t.Run("declared inflight limit", func(t *testing.T) { + memory, metrics := memoryWithRecordingMetrics(t, MemoryConfig{LimitMB: 4, InflightLimitMB: 1, MaxTTL: time.Hour}) + writer, err := memory.Create(t.Context(), NewKey("declared-inflight-limit"), http.Header{ + "Content-Length": {strconv.Itoa(1024 * 1024)}, + }, time.Hour) + assert.NoError(t, err) + assert.NoError(t, writer.Close()) + assert.Equal(t, 1, metrics.declineCount(memoryDeclineDeclaredInflightLimit)) + }) + + t.Run("writer reservation", func(t *testing.T) { + memory, metrics := memoryWithRecordingMetrics(t, MemoryConfig{LimitMB: 4, InflightLimitMB: 1, MaxTTL: time.Hour}) + memory.state.inflightCharge.Store(memory.state.inflightLimit) + writer, err := memory.Create(t.Context(), NewKey("writer-reservation"), nil, time.Hour) + assert.NoError(t, err) + assert.NoError(t, writer.Close()) + memory.state.inflightCharge.Store(0) + assert.Equal(t, 1, metrics.declineCount(memoryDeclineWriterReservation)) + }) + + t.Run("body hard limit", func(t *testing.T) { + memory, metrics := memoryWithRecordingMetrics(t, MemoryConfig{LimitMB: 1, MaxTTL: time.Hour}) + writer, err := memory.Create(t.Context(), NewKey("body-hard-limit"), nil, time.Hour) + assert.NoError(t, err) + written, err := writer.Write(make([]byte, 2*1024*1024)) + assert.NoError(t, err) + assert.Equal(t, 2*1024*1024, written) + assert.NoError(t, writer.Close()) + assert.Equal(t, 1, metrics.declineCount(memoryDeclineBodyHardLimit)) + }) + + t.Run("content length mismatch", func(t *testing.T) { + memory, metrics := memoryWithRecordingMetrics(t, MemoryConfig{LimitMB: 1, MaxTTL: time.Hour}) + writer, err := memory.Create(t.Context(), NewKey("content-length-mismatch"), http.Header{ + "Content-Length": {"1"}, + }, time.Hour) + assert.NoError(t, err) + assert.NoError(t, writer.Close()) + assert.Equal(t, 1, metrics.declineCount(memoryDeclineContentLengthMismatch)) + }) + + t.Run("admission limit", func(t *testing.T) { + memory, metrics := memoryWithRecordingMetrics(t, MemoryConfig{LimitMB: 1, MaxTTL: time.Hour}) + keys := make([]Key, 0, memory.state.limitBytes/memoryEntryMinimumCharge) + for index := range cap(keys) { + key := NewKey(fmt.Sprintf("admission-limit-%d", index)) + keys = append(keys, key) + writeMemoryTestEntry(t, memory, key, nil, time.Hour) + } + for _, key := range keys { + _, err := memory.Stat(t.Context(), key) + assert.NoError(t, err) + } + writer, err := memory.Create(t.Context(), NewKey("declined-admission"), nil, time.Hour) + assert.NoError(t, err) + assert.NoError(t, writer.Close()) + assert.Equal(t, 1, metrics.declineCount(memoryDeclineAdmissionLimit)) + }) +} + func TestMemoryRejectsInvalidLimits(t *testing.T) { _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) type invalidLimitTest struct { @@ -680,6 +845,8 @@ func TestMemoryRejectsInvalidLimits(t *testing.T) { tests := []invalidLimitTest{ {name: "negative retained limit", config: MemoryConfig{LimitMB: -1}}, {name: "negative inflight limit", config: MemoryConfig{InflightLimitMB: -1}}, + {name: "inflight limit equals retained limit", config: MemoryConfig{LimitMB: 8, InflightLimitMB: 8}}, + {name: "inflight limit exceeds retained limit", config: MemoryConfig{LimitMB: 8, InflightLimitMB: 9}}, } if strconv.IntSize == 64 { overflowMB := int(math.MaxInt64/(1024*1024) + 1) @@ -699,6 +866,32 @@ func TestMemoryRejectsInvalidLimits(t *testing.T) { } } +func TestMemoryUnlimitedRetentionSupportsInflightLimit(t *testing.T) { + memory := newMemoryTestCacheWithConfig(t, MemoryConfig{InflightLimitMB: 1, MaxTTL: time.Hour}) + first, err := memory.Create(t.Context(), NewKey("unlimited-retention-first"), nil, time.Hour) + assert.NoError(t, err) + second, err := memory.Create(t.Context(), NewKey("unlimited-retention-second"), nil, time.Hour) + assert.NoError(t, err) + for _, writer := range []Writer{first, second} { + written, err := writer.Write(make([]byte, 768*1024)) + assert.NoError(t, err) + assert.Equal(t, 768*1024, written) + } + assert.True(t, memory.state.inflightCharge.Load() <= memory.state.inflightLimit) + assert.NoError(t, first.Close()) + assert.NoError(t, second.Close()) + third, err := memory.Create(t.Context(), NewKey("unlimited-retention-third"), nil, time.Hour) + assert.NoError(t, err) + written, err := third.Write(make([]byte, 768*1024)) + assert.NoError(t, err) + assert.Equal(t, 768*1024, written) + assert.NoError(t, third.Close()) + stats, err := memory.Stats(t.Context()) + assert.NoError(t, err) + assert.Equal(t, int64(2), stats.Objects) + assert.True(t, memory.state.retainedCharge.Load() > memory.state.inflightLimit) +} + func TestMemoryCancelledAdmissionDoesNotEvictEntries(t *testing.T) { memory := newMemoryTestCache(t) const namespace Namespace = "cancelled-admission" @@ -792,7 +985,7 @@ func TestMemoryCancellationWhileWaitingForAdmissionIsNotAdmitted(t *testing.T) { func TestMemoryConcurrentAdmissionKeepsAccountingBounded(t *testing.T) { _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) - memory, err := NewMemory(ctx, MemoryConfig{LimitMB: 1, InflightLimitMB: 1, MaxTTL: time.Hour}) + memory, err := NewMemory(ctx, MemoryConfig{LimitMB: 8, InflightLimitMB: 2, MaxTTL: time.Hour}) assert.NoError(t, err) t.Cleanup(func() { assert.NoError(t, memory.Close()) }) const writes = 64 @@ -838,6 +1031,7 @@ func TestMemoryConcurrentAdmissionKeepsAccountingBounded(t *testing.T) { } assert.Equal(t, actualObjects, stats.Objects) assert.Equal(t, actualSize, stats.Size) + assert.True(t, actualObjects > 1) } func TestMemoryConcurrentReplacementDoesNotExposeCapacity(t *testing.T) { @@ -1063,7 +1257,7 @@ func newMemoryBenchmarkCache(b *testing.B, entryCount int, expiresAt time.Time) func newConfiguredMemoryBenchmarkCache(b *testing.B, entryCount int, expiresAt time.Time) *Memory { b.Helper() - config := MemoryConfig{LimitMB: 1, InflightLimitMB: 1, MaxTTL: time.Hour} + config := MemoryConfig{LimitMB: 8, InflightLimitMB: 2, MaxTTL: time.Hour} state, err := newMemoryState(config) assert.NoError(b, err) state.retainedTarget = int64(entryCount * memoryEntryMinimumCharge) @@ -1111,3 +1305,38 @@ func BenchmarkMemoryParallelHotHits(b *testing.B) { } }) } + +func BenchmarkMemoryParallelHotHitCopy(b *testing.B) { + _, ctx := logging.Configure(b.Context(), logging.Config{Level: slog.LevelError}) + memory, err := NewMemory(ctx, MemoryConfig{LimitMB: 8, MaxTTL: time.Hour}) + assert.NoError(b, err) + key := NewKey("parallel-hot-hit-copy") + payload := bytes.Repeat([]byte{0x4d}, 1024*1024) + writer, err := memory.Create(ctx, key, http.Header{"Content-Length": {strconv.Itoa(len(payload))}}, time.Hour) + assert.NoError(b, err) + _, err = writer.Write(payload) + assert.NoError(b, err) + assert.NoError(b, writer.Close()) + b.Cleanup(func() { assert.NoError(b, memory.Close()) }) + b.SetBytes(int64(len(payload))) + b.ReportAllocs() + b.ResetTimer() + destination := writeOnlyDiscard{} + b.RunParallel(func(pb *testing.PB) { + for pb.Next() { + reader, _, err := memory.Open(ctx, key) + if err != nil { + b.Error(err) + return + } + if _, err := io.Copy(destination, reader); err != nil { + b.Error(err) + return + } + if err := reader.Close(); err != nil { + b.Error(err) + return + } + } + }) +} diff --git a/internal/cache/memory_metrics.go b/internal/cache/memory_metrics.go new file mode 100644 index 0000000..ed98dfb --- /dev/null +++ b/internal/cache/memory_metrics.go @@ -0,0 +1,46 @@ +package cache + +import ( + "context" + + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/attribute" + "go.opentelemetry.io/otel/metric" + + cachewmetrics "github.com/block/cachew/internal/metrics" +) + +type memoryDeclineReason string + +const ( + memoryDeclineDeclaredHardLimit memoryDeclineReason = "declared_hard_limit" + memoryDeclineDeclaredInflightLimit memoryDeclineReason = "declared_inflight_limit" + memoryDeclineWriterReservation memoryDeclineReason = "writer_reservation" + memoryDeclineBodyHardLimit memoryDeclineReason = "body_hard_limit" + memoryDeclineContentLengthMismatch memoryDeclineReason = "content_length_mismatch" + memoryDeclineAdmissionLimit memoryDeclineReason = "admission_limit" +) + +type memoryMetricRecorder interface { + recordDecline(context.Context, memoryDeclineReason) +} + +type memoryMetrics struct { + declines metric.Int64Counter +} + +func newMemoryMetrics() memoryMetricRecorder { + meter := otel.Meter("cachew.memory") + return &memoryMetrics{ + declines: cachewmetrics.NewMetric[metric.Int64Counter]( + meter, + "cachew.memory.admission_declines_total", + "{declines}", + "Memory-tier writes declined without interrupting authoritative cache tiers, by reason", + ), + } +} + +func (m *memoryMetrics) recordDecline(ctx context.Context, reason memoryDeclineReason) { + m.declines.Add(ctx, 1, metric.WithAttributes(attribute.String("reason", string(reason)))) +} diff --git a/internal/cache/tiered.go b/internal/cache/tiered.go index a14f778..4e32dea 100644 --- a/internal/cache/tiered.go +++ b/internal/cache/tiered.go @@ -97,7 +97,6 @@ func (t Tiered) Create(ctx context.Context, key Key, headers http.Header, ttl ti return nil, err } - // The first error will cancel all outstanding writes. ctx, cancel := context.WithCancelCause(ctx) tw := &tieredWriter{ @@ -108,26 +107,62 @@ func (t Tiered) Create(ctx context.Context, key Key, headers http.Header, ttl ti etag: quotedETag, replaceETag: replaceETag, } - // Note: we can't use errgroup here because we do not want to cancel the context on Wait(). - wg := sync.WaitGroup{} + type createResult struct { + index int + writer Writer + err error + } + // An unbuffered result forces a writer that finishes after cancellation to + // take the context branch and abort itself instead of becoming orphaned. + results := make(chan createResult) for i, cache := range t.caches { - wg.Go(func() { + go func() { w, err := cache.Create(ctx, key, headers, ttl, createOpts...) - if err != nil { + result := createResult{index: i, writer: w, err: err} + select { + case results <- result: + case <-ctx.Done(): + if w != nil { + _ = w.Abort(context.Cause(ctx)) //nolint:errcheck // The caller has already returned; abort still releases local resources. + } + } + }() + } + for range t.caches { + select { + case result := <-results: + tw.writers[result.index] = result.writer + if result.err != nil { + cancel(result.err) + return nil, errors.Join(errors.WithStack(result.err), abortTieredWriters(tw.writers, result.err)) + } + if result.writer == nil { + err := errors.New("cache returned a nil writer") cancel(err) + return nil, errors.Join(err, abortTieredWriters(tw.writers, err)) } - tw.writers[i] = w - }) + case <-ctx.Done(): + cause := context.Cause(ctx) + return nil, errors.Join(errors.WithStack(cause), abortTieredWriters(tw.writers, cause)) + } } - done := make(chan struct{}) - go func() { wg.Wait(); close(done) }() - select { - case <-done: - return tw, nil + if cause := context.Cause(ctx); cause != nil { + return nil, errors.Join(errors.WithStack(cause), abortTieredWriters(tw.writers, cause)) + } + return tw, nil +} - case <-ctx.Done(): - return nil, errors.WithStack(context.Cause(ctx)) +func abortTieredWriters(writers []Writer, cause error) error { + wg := sync.WaitGroup{} + errs := make([]error, len(writers)) + for i, writer := range writers { + if writer == nil { + continue + } + wg.Go(func() { errs[i] = errors.WithStack(writer.Abort(cause)) }) } + wg.Wait() + return errors.Join(errs...) } func (t Tiered) replacementETag(ctx context.Context, key Key, newETag string) (bool, error) { diff --git a/internal/cache/tiered_test.go b/internal/cache/tiered_test.go index aa9e6dd..74e19cb 100644 --- a/internal/cache/tiered_test.go +++ b/internal/cache/tiered_test.go @@ -80,6 +80,28 @@ func (c failingCache) Open(_ context.Context, _ cache.Key, _ ...cache.Option) (i return nil, nil, c.err } +type createFuncCache struct { + cache.Cache + create func(context.Context, cache.Key, http.Header, time.Duration, ...cache.Option) (cache.Writer, error) +} + +func (c createFuncCache) Create( + ctx context.Context, + key cache.Key, + headers http.Header, + ttl time.Duration, + opts ...cache.Option, +) (cache.Writer, error) { + return c.create(ctx, key, headers, ttl, opts...) +} + +func newRecordingNoOpWriter(t *testing.T, committed, aborted chan struct{}) cache.Writer { + t.Helper() + writer, err := cache.NoOpCache().Create(t.Context(), cache.NewKey("recording-noop"), nil, time.Hour) + assert.NoError(t, err) + return &recordingWriter{Writer: writer, committed: committed, aborted: aborted} +} + type statFailingCache struct { cache.Cache err error @@ -357,14 +379,14 @@ func TestTieredCreateUsesSameETagInEveryTier(t *testing.T) { func TestTieredCreateContinuesWhenMemoryTierDeclinesAdmission(t *testing.T) { _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) - constrained, err := cache.NewMemory(ctx, cache.MemoryConfig{LimitMB: 1, InflightLimitMB: 1, MaxTTL: time.Hour}) + constrained, err := cache.NewMemory(ctx, cache.MemoryConfig{LimitMB: 8, InflightLimitMB: 2, MaxTTL: time.Hour}) assert.NoError(t, err) - authoritative, err := cache.NewMemory(ctx, cache.MemoryConfig{LimitMB: 2, MaxTTL: time.Hour}) + authoritative, err := cache.NewMemory(ctx, cache.MemoryConfig{LimitMB: 4, MaxTTL: time.Hour}) assert.NoError(t, err) tiered := newTiered(ctx, constrained, authoritative) t.Cleanup(func() { assert.NoError(t, tiered.Close()) }) key := cache.NewKey("declined-memory-admission") - content := make([]byte, 1024*1024) + content := make([]byte, 2*1024*1024) content[0] = 1 headers := http.Header{"Content-Length": {strconv.Itoa(len(content))}} @@ -382,6 +404,85 @@ func TestTieredCreateContinuesWhenMemoryTierDeclinesAdmission(t *testing.T) { assert.Equal(t, content, readAllAndClose(t, reader)) } +func TestTieredCreateCancellationAbortsCreatedAndLateWriters(t *testing.T) { + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + ctx, cancel := context.WithCancel(ctx) + created := make(chan struct{}) + lateCreateStarted := make(chan struct{}) + releaseLateCreate := make(chan struct{}) + firstAborted := make(chan struct{}) + lateAborted := make(chan struct{}) + firstWriter := newRecordingNoOpWriter(t, make(chan struct{}), firstAborted) + lateWriter := newRecordingNoOpWriter(t, make(chan struct{}), lateAborted) + first := createFuncCache{ + Cache: cache.NoOpCache(), + create: func(context.Context, cache.Key, http.Header, time.Duration, ...cache.Option) (cache.Writer, error) { + close(created) + return firstWriter, nil + }, + } + late := createFuncCache{ + Cache: cache.NoOpCache(), + create: func(context.Context, cache.Key, http.Header, time.Duration, ...cache.Option) (cache.Writer, error) { + close(lateCreateStarted) + <-releaseLateCreate + return lateWriter, nil + }, + } + tiered := newTiered(ctx, first, late) + result := make(chan error, 1) + go func() { + _, err := tiered.Create(ctx, cache.NewKey("cancelled-tiered-create"), nil, time.Hour) + result <- err + }() + <-created + <-lateCreateStarted + cancel() + assert.IsError(t, <-result, context.Canceled) + select { + case <-firstAborted: + case <-time.After(time.Second): + t.Fatal("created writer was not aborted") + } + close(releaseLateCreate) + select { + case <-lateAborted: + case <-time.After(time.Second): + t.Fatal("late writer was not aborted") + } +} + +func TestTieredCreateErrorAbortsOtherWriters(t *testing.T) { + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + createFailed := errors.New("create failed") + created := make(chan struct{}) + aborted := make(chan struct{}) + successfulWriter := newRecordingNoOpWriter(t, make(chan struct{}), aborted) + successful := createFuncCache{ + Cache: cache.NoOpCache(), + create: func(context.Context, cache.Key, http.Header, time.Duration, ...cache.Option) (cache.Writer, error) { + close(created) + return successfulWriter, nil + }, + } + failing := createFuncCache{ + Cache: cache.NoOpCache(), + create: func(context.Context, cache.Key, http.Header, time.Duration, ...cache.Option) (cache.Writer, error) { + <-created + return nil, createFailed + }, + } + tiered := newTiered(ctx, successful, failing) + writer, err := tiered.Create(ctx, cache.NewKey("failed-tiered-create"), nil, time.Hour) + assert.Zero(t, writer) + assert.IsError(t, err, createFailed) + select { + case <-aborted: + case <-time.After(time.Second): + t.Fatal("successful writer was not aborted") + } +} + func TestTieredRequiresMetadataStore(t *testing.T) { _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelDebug}) lower, err := cache.NewMemory(ctx, cache.MemoryConfig{LimitMB: 1024, MaxTTL: time.Hour}) From ec489cba0e7c9dc98a83d0f504f034e65df19509 Mon Sep 17 00:00:00 2001 From: Ravi Agarwal Date: Tue, 1 Sep 2026 00:05:28 -0700 Subject: [PATCH 3/3] memory: preserve cache on inflight saturation Distinguish inflight sub-limit exhaustion from hard-budget exhaustion so a declined writer only runs CLOCK trimming when eviction can actually make room. Cover unlimited retention while another writer fully occupies the inflight budget. Release Tiered create contexts after child writers close, and report the write-to-discard benchmark as latency and allocation work rather than physical copy throughput. Verification: bin/just fmt; bin/just lint; bin/just test; focused race tests repeated 20 times; go vet ./internal/cache. --- internal/cache/memory.go | 23 ++++++++++----- internal/cache/memory_internal_test.go | 40 ++++++++++++++------------ internal/cache/tiered.go | 1 + internal/cache/tiered_test.go | 26 +++++++++++++++++ 4 files changed, 65 insertions(+), 25 deletions(-) diff --git a/internal/cache/memory.go b/internal/cache/memory.go index bdd8272..86d8891 100644 --- a/internal/cache/memory.go +++ b/internal/cache/memory.go @@ -882,36 +882,45 @@ func (w *memoryWriter) ensureCapacity(needed int64) bool { return true } +type memoryReservationResult uint8 + +const ( + memoryReservationSucceeded memoryReservationResult = iota + memoryReservationInflightLimited + memoryReservationHardLimited +) + func (w *memoryWriter) reserve(amount int64) bool { if amount <= 0 { return true } - if w.tryReserve(amount) { + result := w.tryReserve(amount) + if result == memoryReservationSucceeded { return true } - if !w.budgeted { + if result != memoryReservationHardLimited { return false } w.cache.trimToTarget(w.ctx, w.namespace, w.key) - return w.tryReserve(amount) + return w.tryReserve(amount) == memoryReservationSucceeded } -func (w *memoryWriter) tryReserve(amount int64) bool { +func (w *memoryWriter) tryReserve(amount int64) memoryReservationResult { if w.inflightLimit == 0 { w.cache.state.inflightCharge.Add(amount) } else if !reserveBounded(&w.cache.state.inflightCharge, w.inflightLimit, amount) { - return false + return memoryReservationInflightLimited } if w.budgeted { if w.limitBytes == 0 { w.cache.state.hardLimitCharge.Add(amount) } else if !reserveBounded(&w.cache.state.hardLimitCharge, w.limitBytes, amount) { w.cache.state.inflightCharge.Add(-amount) - return false + return memoryReservationHardLimited } } w.reservedBytes += amount - return true + return memoryReservationSucceeded } func (w *memoryWriter) release(amount int64) { diff --git a/internal/cache/memory_internal_test.go b/internal/cache/memory_internal_test.go index 6b13f54..2ee2ecb 100644 --- a/internal/cache/memory_internal_test.go +++ b/internal/cache/memory_internal_test.go @@ -866,30 +866,35 @@ func TestMemoryRejectsInvalidLimits(t *testing.T) { } } -func TestMemoryUnlimitedRetentionSupportsInflightLimit(t *testing.T) { +func TestMemoryUnlimitedRetentionSurvivesInflightPressure(t *testing.T) { memory := newMemoryTestCacheWithConfig(t, MemoryConfig{InflightLimitMB: 1, MaxTTL: time.Hour}) - first, err := memory.Create(t.Context(), NewKey("unlimited-retention-first"), nil, time.Hour) - assert.NoError(t, err) - second, err := memory.Create(t.Context(), NewKey("unlimited-retention-second"), nil, time.Hour) - assert.NoError(t, err) - for _, writer := range []Writer{first, second} { + for _, key := range []Key{NewKey("unlimited-retention-first"), NewKey("unlimited-retention-second")} { + writer, err := memory.Create(t.Context(), key, nil, time.Hour) + assert.NoError(t, err) written, err := writer.Write(make([]byte, 768*1024)) assert.NoError(t, err) assert.Equal(t, 768*1024, written) + assert.NoError(t, writer.Close()) } - assert.True(t, memory.state.inflightCharge.Load() <= memory.state.inflightLimit) - assert.NoError(t, first.Close()) - assert.NoError(t, second.Close()) - third, err := memory.Create(t.Context(), NewKey("unlimited-retention-third"), nil, time.Hour) + before, err := memory.Stats(t.Context()) assert.NoError(t, err) - written, err := third.Write(make([]byte, 768*1024)) + assert.Equal(t, int64(2), before.Objects) + assert.True(t, memory.state.retainedCharge.Load() > memory.state.inflightLimit) + + pressure, err := memory.Create(t.Context(), NewKey("unlimited-inflight-pressure"), nil, time.Hour) assert.NoError(t, err) - assert.Equal(t, 768*1024, written) - assert.NoError(t, third.Close()) - stats, err := memory.Stats(t.Context()) + written, err := pressure.Write(make([]byte, 1024*1024-memoryWriterMinimumCharge)) assert.NoError(t, err) - assert.Equal(t, int64(2), stats.Objects) - assert.True(t, memory.state.retainedCharge.Load() > memory.state.inflightLimit) + assert.Equal(t, 1024*1024-memoryWriterMinimumCharge, written) + assert.Equal(t, memory.state.inflightLimit, memory.state.inflightCharge.Load()) + + declined, err := memory.Create(t.Context(), NewKey("unlimited-inflight-declined"), nil, time.Hour) + assert.NoError(t, err) + assert.NoError(t, declined.Close()) + after, err := memory.Stats(t.Context()) + assert.NoError(t, err) + assert.Equal(t, before, after) + assert.IsError(t, pressure.Abort(errors.New("release inflight pressure")), context.Canceled) } func TestMemoryCancelledAdmissionDoesNotEvictEntries(t *testing.T) { @@ -1306,7 +1311,7 @@ func BenchmarkMemoryParallelHotHits(b *testing.B) { }) } -func BenchmarkMemoryParallelHotHitCopy(b *testing.B) { +func BenchmarkMemoryParallelHotHitWriteToDiscard(b *testing.B) { _, ctx := logging.Configure(b.Context(), logging.Config{Level: slog.LevelError}) memory, err := NewMemory(ctx, MemoryConfig{LimitMB: 8, MaxTTL: time.Hour}) assert.NoError(b, err) @@ -1318,7 +1323,6 @@ func BenchmarkMemoryParallelHotHitCopy(b *testing.B) { assert.NoError(b, err) assert.NoError(b, writer.Close()) b.Cleanup(func() { assert.NoError(b, memory.Close()) }) - b.SetBytes(int64(len(payload))) b.ReportAllocs() b.ResetTimer() destination := writeOnlyDiscard{} diff --git a/internal/cache/tiered.go b/internal/cache/tiered.go index 4e32dea..b96a9c6 100644 --- a/internal/cache/tiered.go +++ b/internal/cache/tiered.go @@ -749,6 +749,7 @@ func (t *tieredWriter) Close() error { return nil } t.closed = true + defer t.cancel(nil) wg := sync.WaitGroup{} errs := make([]error, len(t.writers)) diff --git a/internal/cache/tiered_test.go b/internal/cache/tiered_test.go index 74e19cb..c9edb89 100644 --- a/internal/cache/tiered_test.go +++ b/internal/cache/tiered_test.go @@ -483,6 +483,32 @@ func TestTieredCreateErrorAbortsOtherWriters(t *testing.T) { } } +func TestTieredWriterCloseReleasesCreateContext(t *testing.T) { + _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelError}) + createdContext := make(chan context.Context, 1) + recording := createFuncCache{ + Cache: cache.NoOpCache(), + create: func( + ctx context.Context, + key cache.Key, + headers http.Header, + ttl time.Duration, + opts ...cache.Option, + ) (cache.Writer, error) { + createdContext <- ctx + return cache.NoOpCache().Create(ctx, key, headers, ttl, opts...) + }, + } + tiered := newTiered(ctx, recording, cache.NoOpCache()) + t.Cleanup(func() { assert.NoError(t, tiered.Close()) }) + writer, err := tiered.Create(ctx, cache.NewKey("released-create-context"), nil, time.Hour) + assert.NoError(t, err) + writerContext := <-createdContext + assert.Zero(t, context.Cause(writerContext)) + assert.NoError(t, writer.Close()) + assert.IsError(t, context.Cause(writerContext), context.Canceled) +} + func TestTieredRequiresMetadataStore(t *testing.T) { _, ctx := logging.Configure(t.Context(), logging.Config{Level: slog.LevelDebug}) lower, err := cache.NewMemory(ctx, cache.MemoryConfig{LimitMB: 1024, MaxTTL: time.Hour})