Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
148 changes: 148 additions & 0 deletions internal/cache/redis.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,148 @@
// Package cache wraps the Redis client with a typed GetOrSet helper that
// collapses concurrent identical requests via singleflight and fails open
// when Redis is unavailable.
//
// Designed for the §13 eventual-consistency surfaces (billing/usage,
// team/summary) where:
//
// - The per-team aggregation is expensive enough that N concurrent
// dashboard tabs should NOT trigger N DB scans — singleflight collapses
// them to one in-process compute + one cache write.
// - A Redis outage MUST NOT break the read endpoint (the underlying DB is
// still authoritative). GetOrSet falls through to fn on every Redis
// error so the user sees data, just without the cache amortisation.
// - Hot-path callers prefer a typed result (struct, not []byte). The
// generic `T any` parameter keeps callers off encoding/json directly.
//
// Real-time paths (POST /db/new quota checks, webhook handlers) MUST NOT
// use this helper — they read fresh per the §13 freshness matrix.
package cache

import (
"context"
"encoding/json"
"errors"
"fmt"
"log/slog"
"time"

"github.com/redis/go-redis/v9"
"golang.org/x/sync/singleflight"
)

// group is the per-process singleflight that collapses concurrent calls to
// GetOrSet sharing the same key. Keys live in one global namespace so callers
// must scope them (e.g. "billing:usage:" + teamID).
var group singleflight.Group

// GetOrSet returns the cached value for key when present and fresh.
//
// Miss path (cache empty or returns a NOT FOUND): runs fn under singleflight,
// stores the encoded result with TTL ttl, returns the result.
//
// Failure modes (intentional fail-open semantics):
//
// - Redis GET errored — log + skip cache, run fn, return its result without
// attempting another SET (the cache layer is currently broken; don't
// hammer it). This matches the "Redis down → fall through" cell in the
// §13 freshness matrix.
// - JSON unmarshal of the cached value failed — treat as miss. Most likely
// cause is a serialised value shape change across deploys; the next SET
// after fn runs heals the cache entry.
// - fn returned an error — propagate it without touching the cache.
// - Redis SET errored on the way back — log + return the freshly-computed
// value anyway. The next call will re-attempt the SET.
//
// Negative caching (fn returned a zero-value T) is allowed and uses the same
// ttl — callers that want a shorter negative TTL should branch outside.
func GetOrSet[T any](
ctx context.Context,
rdb *redis.Client,
key string,
ttl time.Duration,
fn func(context.Context) (T, error),
) (T, error) {
var zero T

// Fast path: try the cache. A nil client means cache is disabled — go
// straight to fn without using singleflight (no point — there's nothing
// to collapse on).
if rdb != nil {
raw, err := rdb.Get(ctx, key).Bytes()
switch {
case err == nil:
var out T
if jerr := json.Unmarshal(raw, &out); jerr == nil {
return out, nil
}
// Corrupt cache entry — treat as miss, log so the shape skew is
// visible. Don't return the unmarshal error to the caller.
slog.Warn("cache.get_unmarshal_failed", "key", key, "error", "json decode")
case errors.Is(err, redis.Nil):
// True miss — fall through to fn under singleflight.
default:
// Redis is unreachable / down. Fail open: run fn without the
// cache wrapper and skip the SET path entirely so we don't
// hammer a flapping Redis. Bypassing singleflight here means
// N concurrent callers will all hit the DB during an outage,
// which is acceptable — the cache being down IS the
// degradation, the DB is the source of truth.
slog.Warn("cache.get_failed_fail_open", "key", key, "error", err.Error())
return fn(ctx)
}
}

// Miss path: collapse concurrent callers to one fn invocation.
//
// singleflight returns (value, error, shared). We ignore `shared`; both
// the leader and the followers see the same value+error pair. The leader
// is the only one that touches Redis SET — followers piggyback on the
// returned value.
v, err, _ := group.Do(key, func() (interface{}, error) {
out, fnErr := fn(ctx)
if fnErr != nil {
return out, fnErr
}
if rdb != nil {
encoded, jerr := json.Marshal(out)
if jerr != nil {
// Encoding failure is a programmer error (T can't be
// marshalled). Don't poison the cache; log + return the
// value so the request still succeeds.
slog.Warn("cache.set_marshal_failed", "key", key, "error", jerr.Error())
return out, nil
}
if setErr := rdb.Set(ctx, key, encoded, ttl).Err(); setErr != nil {
// Same fail-open as GET: log but return the value.
slog.Warn("cache.set_failed", "key", key, "error", setErr.Error())
}
}
return out, nil
})

if err != nil {
return zero, err
}
// singleflight returns the leader's value via interface{}. The type
// parameter T is the same for every caller of this key, so the assertion
// is safe under normal use; a panic here would indicate two callers
// using the same cache key with different T (a bug in caller code).
out, ok := v.(T)
if !ok {
return zero, fmt.Errorf("cache.GetOrSet: type mismatch for key %q", key)
}
return out, nil
}

// Invalidate deletes a cache key. Use it from write paths that change the
// underlying aggregate (e.g. a deploy completing should invalidate
// billing:usage:<team>). A nil client is a no-op so callers can wire this
// in without conditional checks.
func Invalidate(ctx context.Context, rdb *redis.Client, key string) {
if rdb == nil {
return
}
if err := rdb.Del(ctx, key).Err(); err != nil {
slog.Warn("cache.invalidate_failed", "key", key, "error", err.Error())
}
}
244 changes: 244 additions & 0 deletions internal/cache/redis_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,244 @@
package cache_test

import (
"context"
"errors"
"sync"
"sync/atomic"
"testing"
"time"

"github.com/alicebob/miniredis/v2"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"

"instant.dev/internal/cache"
)

// newMiniRedis returns a *redis.Client backed by an in-memory miniredis
// instance plus a cleanup func. Used everywhere we need a real-shaped
// Redis without a Docker container.
func newMiniRedis(t *testing.T) (*redis.Client, func()) {
t.Helper()
mr, err := miniredis.Run()
require.NoError(t, err)
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
return rdb, func() {
rdb.Close()
mr.Close()
}
}

type usagePayload struct {
Postgres int64 `json:"postgres"`
Redis int64 `json:"redis"`
}

// TestGetOrSet_MissRunsFnOnceAndCaches verifies the basic Redis-miss path:
// the first call runs fn, the second call short-circuits to the cache.
func TestGetOrSet_MissRunsFnOnceAndCaches(t *testing.T) {
rdb, cleanup := newMiniRedis(t)
defer cleanup()

var calls atomic.Int32
fn := func(_ context.Context) (usagePayload, error) {
calls.Add(1)
return usagePayload{Postgres: 100, Redis: 50}, nil
}

ctx := context.Background()
v1, err := cache.GetOrSet(ctx, rdb, "test:k1", 60*time.Second, fn)
require.NoError(t, err)
assert.Equal(t, usagePayload{Postgres: 100, Redis: 50}, v1)

v2, err := cache.GetOrSet(ctx, rdb, "test:k1", 60*time.Second, fn)
require.NoError(t, err)
assert.Equal(t, usagePayload{Postgres: 100, Redis: 50}, v2)

assert.Equal(t, int32(1), calls.Load(), "fn should have run exactly once across both calls")
}

// TestGetOrSet_SingleflightCollapsesConcurrentCallers — the headline §10.20
// guarantee: N concurrent identical requests collapse to 1 fn invocation.
// Without singleflight, N callers would race past the empty-cache check and
// all run fn before any of them got to SET. With singleflight, the leader
// runs fn and the followers receive its result.
func TestGetOrSet_SingleflightCollapsesConcurrentCallers(t *testing.T) {
rdb, cleanup := newMiniRedis(t)
defer cleanup()

const concurrency = 20
var calls atomic.Int32
// gate holds fn open until every goroutine is in flight, so they all
// observe the same "cache empty" snapshot. Without it the test races —
// goroutine #N might run after goroutine #1 already set the cache.
gate := make(chan struct{})
fn := func(_ context.Context) (usagePayload, error) {
<-gate
calls.Add(1)
// A small sleep makes the singleflight window visible — the leader
// is still inside fn when followers arrive. Without it the timing
// can occasionally let a follower miss the inflight entry.
time.Sleep(20 * time.Millisecond)
return usagePayload{Postgres: 42}, nil
}

ctx := context.Background()
results := make(chan usagePayload, concurrency)
errs := make(chan error, concurrency)

var wg sync.WaitGroup
for i := 0; i < concurrency; i++ {
wg.Add(1)
go func() {
defer wg.Done()
v, err := cache.GetOrSet(ctx, rdb, "test:sf", 60*time.Second, fn)
results <- v
errs <- err
}()
}
// Let every goroutine reach the gate before any of them runs fn.
time.Sleep(50 * time.Millisecond)
close(gate)
wg.Wait()
close(results)
close(errs)

for err := range errs {
require.NoError(t, err)
}
for v := range results {
assert.Equal(t, usagePayload{Postgres: 42}, v)
}
assert.Equal(t, int32(1), calls.Load(), "singleflight should collapse %d concurrent callers to 1 fn invocation", concurrency)
}

// TestGetOrSet_RedisDownFailsOpen verifies that when Redis errors on GET,
// GetOrSet falls through to fn and returns its result. The cache being
// unreachable must never break the read path.
func TestGetOrSet_RedisDownFailsOpen(t *testing.T) {
// Point at a closed port — the dial will fail fast.
rdb := redis.NewClient(&redis.Options{
Addr: "127.0.0.1:1", // reserved low port, refuses connections
DialTimeout: 50 * time.Millisecond,
})
defer rdb.Close()

var calls atomic.Int32
fn := func(_ context.Context) (usagePayload, error) {
calls.Add(1)
return usagePayload{Postgres: 7}, nil
}

ctx := context.Background()
v, err := cache.GetOrSet(ctx, rdb, "test:down", 60*time.Second, fn)
require.NoError(t, err)
assert.Equal(t, usagePayload{Postgres: 7}, v)
assert.Equal(t, int32(1), calls.Load(), "fn must run when redis is down")

// A second call must also reach fn — we bypass singleflight on the
// Redis-down path to avoid hammering a flapping cache, and the cache
// itself can't serve the entry. (See §10.20 fail-open contract.)
v2, err := cache.GetOrSet(ctx, rdb, "test:down", 60*time.Second, fn)
require.NoError(t, err)
assert.Equal(t, usagePayload{Postgres: 7}, v2)
assert.Equal(t, int32(2), calls.Load())
}

// TestGetOrSet_NilClientPassesThrough — a nil *redis.Client means "no cache
// configured"; GetOrSet should still call fn and return its result. Useful
// in tests and in dev configs where Redis isn't wired.
func TestGetOrSet_NilClientPassesThrough(t *testing.T) {
var calls atomic.Int32
fn := func(_ context.Context) (usagePayload, error) {
calls.Add(1)
return usagePayload{Postgres: 1}, nil
}
v, err := cache.GetOrSet(context.Background(), nil, "test:nil", 60*time.Second, fn)
require.NoError(t, err)
assert.Equal(t, usagePayload{Postgres: 1}, v)
assert.Equal(t, int32(1), calls.Load())
}

// TestGetOrSet_FnErrorPropagates — a fn error must not be cached and must
// surface to the caller verbatim.
func TestGetOrSet_FnErrorPropagates(t *testing.T) {
rdb, cleanup := newMiniRedis(t)
defer cleanup()

sentinel := errors.New("aggregate failed")
fn := func(_ context.Context) (usagePayload, error) {
return usagePayload{}, sentinel
}

_, err := cache.GetOrSet(context.Background(), rdb, "test:err", 60*time.Second, fn)
require.Error(t, err)
assert.ErrorIs(t, err, sentinel)

// Confirm the cache was NOT populated.
_, ferr := rdb.Get(context.Background(), "test:err").Bytes()
assert.ErrorIs(t, ferr, redis.Nil)
}

// TestGetOrSet_ZeroValueCachesNegative — fn returning a zero-value T is a
// valid result (e.g. a team with no resources). It must still be cached so
// the next caller doesn't re-run the aggregate.
func TestGetOrSet_ZeroValueCachesNegative(t *testing.T) {
rdb, cleanup := newMiniRedis(t)
defer cleanup()

var calls atomic.Int32
fn := func(_ context.Context) (usagePayload, error) {
calls.Add(1)
return usagePayload{}, nil
}
ctx := context.Background()
_, err := cache.GetOrSet(ctx, rdb, "test:empty", 60*time.Second, fn)
require.NoError(t, err)
_, err = cache.GetOrSet(ctx, rdb, "test:empty", 60*time.Second, fn)
require.NoError(t, err)
assert.Equal(t, int32(1), calls.Load(), "zero-value results must still be cached")
}

// TestGetOrSet_CorruptCacheEntryFallsThrough — if a cache entry was
// serialised under an older shape, json.Unmarshal returns an error and
// GetOrSet treats it as a miss. The next SET heals the entry.
func TestGetOrSet_CorruptCacheEntryFallsThrough(t *testing.T) {
rdb, cleanup := newMiniRedis(t)
defer cleanup()

// Plant a value that doesn't decode as usagePayload.
require.NoError(t, rdb.Set(context.Background(), "test:corrupt", "not-json", time.Minute).Err())

var calls atomic.Int32
fn := func(_ context.Context) (usagePayload, error) {
calls.Add(1)
return usagePayload{Postgres: 999}, nil
}
v, err := cache.GetOrSet(context.Background(), rdb, "test:corrupt", time.Minute, fn)
require.NoError(t, err)
assert.Equal(t, usagePayload{Postgres: 999}, v)
assert.Equal(t, int32(1), calls.Load())
}

// TestInvalidate_DeletesKey ensures Invalidate clears the cache and a nil
// client is a no-op.
func TestInvalidate_DeletesKey(t *testing.T) {
rdb, cleanup := newMiniRedis(t)
defer cleanup()

fn := func(_ context.Context) (usagePayload, error) {
return usagePayload{Postgres: 5}, nil
}
ctx := context.Background()
_, err := cache.GetOrSet(ctx, rdb, "test:inv", time.Minute, fn)
require.NoError(t, err)

cache.Invalidate(ctx, rdb, "test:inv")
_, err = rdb.Get(ctx, "test:inv").Bytes()
assert.ErrorIs(t, err, redis.Nil)

// nil client → no panic.
cache.Invalidate(ctx, nil, "test:inv")
}
Loading