diff --git a/README.md b/README.md index 301eb49..b39dc07 100644 --- a/README.md +++ b/README.md @@ -150,12 +150,50 @@ 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 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. 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. 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 { - 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..86d8891 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,560 @@ 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 = memoryWriterMinimumCharge + 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 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"` } 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) } +// 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 + 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 + metrics memoryMetricRecorder +} + +// 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 + } + 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: configuredInflightBytes, + metrics: newMemoryMetrics(), }, 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 +} + +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 +} - entry, exists := nsEntries[key] - if !exists { - return nil, os.ErrNotExist +func memoryEntryCharge(namespace Namespace, data []byte, headers http.Header) int64 { + charge := int64(cap(data)) + memoryMetadataCharge(namespace, headers) + return max(charge, int64(memoryEntryMinimumCharge)) +} + +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 +} + +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 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 reserveBounded(&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 { + if m.state.limitBytes == 0 { + m.state.retainedCharge.Add(amount) + return true + } + return reserveBounded(&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 { + // 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 + } + 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) { + // 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 + 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) + } + 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 +582,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() - - nsEntries, nsExists := m.entries[m.namespace] - if !nsExists { - return nil, nil, os.ErrNotExist - } + shard := m.shard(m.namespace, key) + shard.mu.RLock() + defer shard.mu.RUnlock() - 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 +608,77 @@ 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 +} + +var _ io.WriterTo = (*memoryReader)(nil) + +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) 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 + } + 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 +687,57 @@ 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 { + 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) - 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) { + m.state.metrics.recordDecline(ctx, memoryDeclineWriterReservation) + 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 +746,224 @@ 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 +} - totalObjects := int64(0) - for _, nsEntries := range m.entries { - totalObjects += int64(len(nsEntries)) +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 +} + +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 { + 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...) + 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) } + if w.expectedLength >= 0 { + maximum = min(maximum, w.expectedLength) + } + 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 +} + +type memoryReservationResult uint8 + +const ( + memoryReservationSucceeded memoryReservationResult = iota + memoryReservationInflightLimited + memoryReservationHardLimited +) + +func (w *memoryWriter) reserve(amount int64) bool { + if amount <= 0 { + return true + } + result := w.tryReserve(amount) + if result == memoryReservationSucceeded { + return true + } + if result != memoryReservationHardLimited { + return false + } + w.cache.trimToTarget(w.ctx, w.namespace, w.key) + return w.tryReserve(amount) == memoryReservationSucceeded +} + +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 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 memoryReservationHardLimited } } + w.reservedBytes += amount + return memoryReservationSucceeded +} - freedSpace := int64(0) - for _, e := range entries { - if freedSpace >= neededSpace { - break +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) + } + 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) } - m.currentSize.Add(-e.size) - delete(m.entries[e.namespace], e.key) - freedSpace += e.size + w.reservedBytes = 0 } } -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) 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) Write(p []byte) (int, error) { - if w.closed { - return 0, errors.New("writer closed") +func (w *memoryWriter) discard() { + w.releaseReservation() + w.data = nil + w.discarded = true +} + +func (w *memoryWriter) decline(reason memoryDeclineReason) { + if w.discarded { + return } - n, err := w.buf.Write(p) - return n, errors.WithStack(err) + w.cache.state.metrics.recordDecline(w.ctx, reason) + w.discard() } func (w *memoryWriter) Abort(err error) error { @@ -273,70 +976,73 @@ 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.decline(memoryDeclineContentLengthMismatch) + 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 { + 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) } -// 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..2ee2ecb --- /dev/null +++ b/internal/cache/memory_internal_test.go @@ -0,0 +1,1346 @@ +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 +} + +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) + 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 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") + 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 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 { + name string + config MemoryConfig + } + 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) + 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 TestMemoryUnlimitedRetentionSurvivesInflightPressure(t *testing.T) { + memory := newMemoryTestCacheWithConfig(t, MemoryConfig{InflightLimitMB: 1, MaxTTL: time.Hour}) + 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()) + } + before, err := memory.Stats(t.Context()) + assert.NoError(t, err) + 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) + written, err := pressure.Write(make([]byte, 1024*1024-memoryWriterMinimumCharge)) + assert.NoError(t, err) + 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) { + 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: 8, InflightLimitMB: 2, 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) + assert.True(t, actualObjects > 1) +} + +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: 8, InflightLimitMB: 2, 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 + } + } + }) +} + +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) + 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.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..b96a9c6 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) { @@ -714,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 7a3eff8..c9edb89 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" @@ -79,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 @@ -354,6 +377,138 @@ 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: 8, InflightLimitMB: 2, MaxTTL: time.Hour}) + assert.NoError(t, err) + 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, 2*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 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 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})