app/apicache.go

v1.5.2
skunky-art/app/apicache.go history · blame · raw

274 lines · 7341 bytes

  1package app
  2
  3import (
  4	"bytes"
  5	"container/list"
  6	"io"
  7	"net/http"
  8	"strings"
  9	"sync"
 10	"time"
 11
 12	"golang.org/x/sync/singleflight"
 13)
 14
 15// apiCache holds DeviantArt API responses so that repeat and concurrent
 16// requests for one URL cost one upstream call. It sits in front of the
 17// throttle: a hit never touches DeviantArt or the throttle's budget.
 18//
 19// The store is bounded by body bytes with least-recently-used eviction. It
 20// is safe for concurrent use.
 21type apiCache struct {
 22	maxBytes int64
 23	ttl      time.Duration
 24	// stale is how long past ttl an entry is kept to be served when upstream
 25	// is blocked or unreachable. Zero keeps nothing past ttl.
 26	stale time.Duration
 27	now   func() time.Time
 28
 29	mu        sync.Mutex
 30	entries   map[string]*cacheEntry
 31	lru       *list.List // front is most recently used
 32	held      int64
 33	hits      int64
 34	misses    int64
 35	staleHits int64
 36
 37	// blockedUntil is set when DeviantArt answers with a block. Until then a
 38	// miss with nothing stale to serve gets blockResp back without an
 39	// upstream call, so a banned instance stops hammering the WAF.
 40	blockedUntil time.Time
 41	blockResp    *cacheEntry
 42
 43	flight singleflight.Group
 44}
 45
 46// blockBackoff is how long upstream is left alone after a block response.
 47const blockBackoff = time.Minute
 48
 49// cacheEntry is one buffered response. header is a clone of the upstream
 50// header; body is the whole body, read once.
 51type cacheEntry struct {
 52	key     string
 53	status  int
 54	header  http.Header
 55	body    []byte
 56	expires time.Time
 57	elem    *list.Element
 58}
 59
 60func newAPICache(maxBytes int64, ttl, stale time.Duration) *apiCache {
 61	return &apiCache{
 62		maxBytes: maxBytes,
 63		ttl:      ttl,
 64		stale:    stale,
 65		now:      time.Now,
 66		entries:  map[string]*cacheEntry{},
 67		lru:      list.New(),
 68	}
 69}
 70
 71// cacheable reports whether a request is one the cache handles: a GET to
 72// DeviantArt's API or its group search page. The session bootstrap (/_puppy
 73// with no path), the homepage, avatars and media all pass through.
 74func cacheable(req *http.Request) bool {
 75	if req.Method != http.MethodGet || req.URL.Host != "www.deviantart.com" {
 76		return false
 77	}
 78	p := req.URL.Path
 79	return (strings.HasPrefix(p, "/_puppy/") && len(p) > len("/_puppy/")) ||
 80		strings.HasPrefix(p, "/groups/")
 81}
 82
 83// cacheKey is the URL without csrf_token, which changes every twelve hours
 84// and would otherwise empty the cache on each refresh.
 85func cacheKey(req *http.Request) string {
 86	u := *req.URL
 87	q := u.Query()
 88	q.Del("csrf_token")
 89	u.RawQuery = q.Encode()
 90	return u.String()
 91}
 92
 93// transport returns a RoundTripper that answers from this cache and sends
 94// misses to base. Several transports may share one cache.
 95func (c *apiCache) transport(base http.RoundTripper) http.RoundTripper {
 96	return &cachedTransport{cache: c, base: base}
 97}
 98
 99type cachedTransport struct {
100	cache *apiCache
101	base  http.RoundTripper
102}
103
104// RoundTrip serves a fresh hit from memory. A miss is fetched once per key
105// however many callers are waiting, buffered, stored if it is a 200, and
106// handed to every waiter as its own response.
107//
108// When upstream fails or answers with a block, a stale entry is served
109// instead if one is still held, so a short ban does not take the popular
110// pages down. A block also starts a backoff during which misses with nothing
111// stale get the block response back without an upstream call.
112func (t *cachedTransport) RoundTrip(req *http.Request) (*http.Response, error) {
113	if !cacheable(req) {
114		return t.base.RoundTrip(req)
115	}
116	key := cacheKey(req)
117	old, fresh := t.cache.get(key)
118	if fresh {
119		return old.response(req), nil
120	}
121	if blocked := t.cache.blockedResponse(); blocked != nil {
122		if old != nil {
123			t.cache.countStale()
124			return old.response(req), nil
125		}
126		return blocked.response(req), nil
127	}
128
129	v, err, _ := t.cache.flight.Do(key, func() (any, error) {
130		resp, err := t.base.RoundTrip(req)
131		if err != nil {
132			return nil, err
133		}
134		defer func() { _ = resp.Body.Close() }()
135		body, err := io.ReadAll(resp.Body)
136		if err != nil {
137			return nil, err
138		}
139		e := &cacheEntry{key: key, status: resp.StatusCode, header: resp.Header.Clone(), body: body}
140		switch e.status {
141		case http.StatusOK:
142			t.cache.put(e)
143		case http.StatusForbidden, http.StatusTooManyRequests:
144			t.cache.block(e)
145		}
146		return e, nil
147	})
148	if err != nil {
149		if old != nil {
150			t.cache.countStale()
151			return old.response(req), nil
152		}
153		return nil, err
154	}
155	e, ok := v.(*cacheEntry)
156	if !ok {
157		return nil, io.ErrUnexpectedEOF
158	}
159	if e.status != http.StatusOK && old != nil {
160		t.cache.countStale()
161		return old.response(req), nil
162	}
163	return e.response(req), nil
164}
165
166// block records a block response and starts the backoff.
167func (c *apiCache) block(e *cacheEntry) {
168	c.mu.Lock()
169	defer c.mu.Unlock()
170	c.blockedUntil = c.now().Add(blockBackoff)
171	c.blockResp = e
172}
173
174// blockedResponse returns the last block response while the backoff runs,
175// or nil once it is over.
176func (c *apiCache) blockedResponse() *cacheEntry {
177	c.mu.Lock()
178	defer c.mu.Unlock()
179	if c.blockResp != nil && c.now().Before(c.blockedUntil) {
180		return c.blockResp
181	}
182	return nil
183}
184
185func (c *apiCache) countStale() {
186	c.mu.Lock()
187	c.staleHits++
188	c.mu.Unlock()
189}
190
191// response builds a fresh http.Response over the buffered body, so each
192// caller can read and close its own.
193func (e *cacheEntry) response(req *http.Request) *http.Response {
194	return &http.Response{
195		Status:        http.StatusText(e.status),
196		StatusCode:    e.status,
197		Proto:         "HTTP/1.1",
198		ProtoMajor:    1,
199		ProtoMinor:    1,
200		Header:        e.header.Clone(),
201		Body:          io.NopCloser(bytes.NewReader(e.body)),
202		ContentLength: int64(len(e.body)),
203		Request:       req,
204	}
205}
206
207// get returns the entry for key and whether it is still fresh. A fresh hit is
208// marked most recently used. An entry past ttl but within the stale window is
209// returned as not fresh, for the caller to fall back on; one past the stale
210// window is dropped.
211func (c *apiCache) get(key string) (*cacheEntry, bool) {
212	c.mu.Lock()
213	defer c.mu.Unlock()
214
215	e := c.entries[key]
216	if e == nil {
217		c.misses++
218		return nil, false
219	}
220	now := c.now()
221	if now.Before(e.expires) {
222		c.lru.MoveToFront(e.elem)
223		c.hits++
224		return e, true
225	}
226	c.misses++
227	if now.Before(e.expires.Add(c.stale)) {
228		return e, false
229	}
230	c.remove(e)
231	return nil, false
232}
233
234// put stores e, evicting from the least recently used end until it fits. A
235// body larger than the whole bound is not stored.
236func (c *apiCache) put(e *cacheEntry) {
237	size := int64(len(e.body))
238	if size > c.maxBytes {
239		return
240	}
241
242	c.mu.Lock()
243	defer c.mu.Unlock()
244
245	if old := c.entries[e.key]; old != nil {
246		c.remove(old)
247	}
248	for c.held+size > c.maxBytes {
249		back, ok := c.lru.Back().Value.(*cacheEntry)
250		if !ok {
251			break
252		}
253		c.remove(back)
254	}
255	e.expires = c.now().Add(c.ttl)
256	e.elem = c.lru.PushFront(e)
257	c.entries[e.key] = e
258	c.held += size
259}
260
261// remove drops e. The caller holds mu.
262func (c *apiCache) remove(e *cacheEntry) {
263	c.lru.Remove(e.elem)
264	delete(c.entries, e.key)
265	c.held -= int64(len(e.body))
266}
267
268// stats reports the counters for the hourly log line. stale counts the
269// misses that were answered from an expired entry because upstream failed.
270func (c *apiCache) stats() (hits, misses, stale int64, entries int, held int64) {
271	c.mu.Lock()
272	defer c.mu.Unlock()
273	return c.hits, c.misses, c.staleHits, len(c.entries), c.held
274}