app/httpclient_test.go

v1.5.6
skunky-art/app/httpclient_test.go history · blame · raw

248 lines · 7805 bytes

  1package app
  2
  3import (
  4	"context"
  5	"errors"
  6	"net/http"
  7	"net/http/httptest"
  8	"sync"
  9	"testing"
 10	"time"
 11
 12	"github.com/krazywarez/devianter"
 13)
 14
 15// stubTransport records how many requests reached it and returns an empty 200.
 16type stubTransport struct {
 17	mu sync.Mutex
 18	n  int
 19}
 20
 21func (s *stubTransport) RoundTrip(req *http.Request) (*http.Response, error) {
 22	s.mu.Lock()
 23	s.n++
 24	s.mu.Unlock()
 25	return httptest.NewRecorder().Result(), nil
 26}
 27
 28func newTestThrottle(base http.RoundTripper, gap time.Duration, maxConcurrent int) *daThrottle {
 29	return &daThrottle{base: base, sem: make(chan struct{}, maxConcurrent)}
 30}
 31
 32// DeviantArt requests must be spaced by at least daMinInterval.
 33func TestThrottleRateLimitsDeviantArt(t *testing.T) {
 34	stub := &stubTransport{}
 35	tr := newTestThrottle(stub, daMinInterval, daMaxConcurrent)
 36
 37	start := time.Now()
 38	const n = 3
 39	for range n {
 40		req, _ := http.NewRequest("GET", "https://www.deviantart.com/_puppy/x", nil)
 41		if _, err := tr.RoundTrip(req); err != nil {
 42			t.Fatalf("unexpected error: %v", err)
 43		}
 44	}
 45	elapsed := time.Since(start)
 46
 47	// n requests => at least (n-1) gaps between them.
 48	if want := time.Duration(n-1) * daMinInterval; elapsed < want {
 49		t.Errorf("DA requests were not throttled: %d requests took %v, want >= %v", n, elapsed, want)
 50	}
 51	if stub.n != n {
 52		t.Errorf("expected all %d requests to reach the base transport, got %d", n, stub.n)
 53	}
 54}
 55
 56// Non-DA hosts (e.g. the wixmp image CDN) must not be slowed down.
 57func TestThrottleSkipsOtherHosts(t *testing.T) {
 58	stub := &stubTransport{}
 59	tr := newTestThrottle(stub, daMinInterval, daMaxConcurrent)
 60
 61	start := time.Now()
 62	for range 5 {
 63		req, _ := http.NewRequest("GET", "https://images-wixmp-ed30a86b8c4ca887773594c2.wixmp.com/f/x.jpg", nil)
 64		if _, err := tr.RoundTrip(req); err != nil {
 65			t.Fatalf("unexpected error: %v", err)
 66		}
 67	}
 68
 69	if elapsed := time.Since(start); elapsed >= daMinInterval {
 70		t.Errorf("non-DA host was throttled: 5 requests took %v, want < %v", elapsed, daMinInterval)
 71	}
 72	if stub.n != 5 {
 73		t.Errorf("expected 5 requests through, got %d", stub.n)
 74	}
 75}
 76
 77// Concurrent callers must never exceed daMaxConcurrent in-flight DA requests.
 78func TestThrottleCapsConcurrency(t *testing.T) {
 79	var (
 80		mu       sync.Mutex
 81		inFlight int
 82		peak     int
 83	)
 84	counting := roundTripFunc(func(req *http.Request) (*http.Response, error) {
 85		mu.Lock()
 86		inFlight++
 87		if inFlight > peak {
 88			peak = inFlight
 89		}
 90		mu.Unlock()
 91
 92		time.Sleep(20 * time.Millisecond) // hold the slot
 93
 94		mu.Lock()
 95		inFlight--
 96		mu.Unlock()
 97		return httptest.NewRecorder().Result(), nil
 98	})
 99
100	tr := newTestThrottle(counting, daMinInterval, daMaxConcurrent)
101
102	var wg sync.WaitGroup
103	for range 6 {
104		wg.Go(func() {
105			req, _ := http.NewRequest("GET", "https://www.deviantart.com/_puppy/x", nil)
106			_, _ = tr.RoundTrip(req)
107		})
108	}
109	wg.Wait()
110
111	if peak > daMaxConcurrent {
112		t.Errorf("concurrency cap breached: peak %d in-flight DA requests, max %d", peak, daMaxConcurrent)
113	}
114}
115
116type roundTripFunc func(*http.Request) (*http.Response, error)
117
118func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) }
119
120// InstallDAThrottle must preserve proxy-from-environment so HTTPS_PROXY (VPN
121// egress) keeps working, and must not panic on a repeat call.
122func TestInstallDAThrottlePreservesProxy(t *testing.T) {
123	orig, origCache := http.DefaultTransport, daCache
124	defer func() { http.DefaultTransport, daCache = orig, origCache }()
125
126	InstallDAThrottle()
127
128	// With api-cache on (the default) the cache is outermost and the throttle
129	// sits inside it; with it off the throttle is outermost.
130	rt := http.DefaultTransport
131	if CFG.APICache.Enabled {
132		ct, ok := rt.(*cachedTransport)
133		if !ok {
134			t.Fatalf("DefaultTransport is not the cache, got %T", rt)
135		}
136		rt = ct.base
137	}
138	th, ok := rt.(*daThrottle)
139	if !ok {
140		t.Fatalf("DefaultTransport was not wrapped by the throttle, got %T", rt)
141	}
142	base, ok := th.base.(*http.Transport)
143	if !ok {
144		t.Fatalf("base transport is not *http.Transport, got %T", th.base)
145	}
146	if base.Proxy == nil {
147		t.Error("base transport lost its Proxy func: HTTPS_PROXY / VPN egress would break")
148	}
149}
150
151// countingRT counts base round trips for the throttle tests.
152type countingRT struct {
153	mu    sync.Mutex
154	calls int
155}
156
157func (c *countingRT) RoundTrip(r *http.Request) (*http.Response, error) {
158	c.mu.Lock()
159	c.calls++
160	c.mu.Unlock()
161	return &http.Response{StatusCode: 200, Body: http.NoBody, Request: r}, nil
162}
163
164func daRequest(ctx context.Context) *http.Request {
165	req, _ := http.NewRequestWithContext(ctx, http.MethodGet, "https://www.deviantart.com/_puppy/x", nil)
166	return req
167}
168
169// TestThrottleShedsADeepQueue pins load shedding: with more requests queued
170// than the client timeout can absorb, a new one is refused at once with
171// errUpstreamBusy and never reaches upstream.
172func TestThrottleShedsADeepQueue(t *testing.T) {
173	interval := daMinInterval
174	daMinInterval = time.Second
175	defer func() { daMinInterval = interval }()
176	base := &countingRT{}
177	th := &daThrottle{base: base, sem: make(chan struct{}, 1)}
178	th.waiting.Store(int64(maxQueueWait/time.Second) + 1)
179
180	start := time.Now()
181	_, err := th.RoundTrip(daRequest(context.Background()))
182
183	if !errors.Is(err, errUpstreamBusy) {
184		t.Fatalf("err = %v, want errUpstreamBusy", err)
185	}
186	if time.Since(start) > 100*time.Millisecond || base.calls != 0 {
187		t.Errorf("shed request took %v and made %d upstream calls, want immediate and none", time.Since(start), base.calls)
188	}
189}
190
191// TestThrottleDropsACancelledRequestWaitingForASlot pins that a request whose
192// client has gone does not sit in the queue: with the only slot held, a
193// cancelled context returns at once.
194func TestThrottleDropsACancelledRequestWaitingForASlot(t *testing.T) {
195	base := &countingRT{}
196	th := &daThrottle{base: base, sem: make(chan struct{}, 1)}
197	th.sem <- struct{}{} // hold the only slot
198	ctx, cancel := context.WithCancel(context.Background())
199	cancel()
200
201	start := time.Now()
202	_, err := th.RoundTrip(daRequest(ctx))
203
204	if !errors.Is(err, context.Canceled) {
205		t.Fatalf("err = %v, want context.Canceled", err)
206	}
207	if time.Since(start) > 100*time.Millisecond || base.calls != 0 {
208		t.Errorf("took %v with %d upstream calls, want immediate and none", time.Since(start), base.calls)
209	}
210}
211
212// TestThrottleCancelledDuringIntervalKeepsTheSlot pins that a request
213// cancelled while waiting out the interval does not consume it: the next
214// live request starts as soon as the original interval allows.
215func TestThrottleCancelledDuringIntervalKeepsTheSlot(t *testing.T) {
216	interval := daMinInterval
217	daMinInterval = 300 * time.Millisecond
218	defer func() { daMinInterval = interval }()
219	base := &countingRT{}
220	th := &daThrottle{base: base, sem: make(chan struct{}, 1)}
221
222	if _, err := th.RoundTrip(daRequest(context.Background())); err != nil {
223		t.Fatal(err)
224	}
225	first := th.last
226
227	ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond)
228	defer cancel()
229	if _, err := th.RoundTrip(daRequest(ctx)); !errors.Is(err, context.DeadlineExceeded) {
230		t.Fatalf("err = %v, want context.DeadlineExceeded", err)
231	}
232	if th.last != first {
233		t.Error("a cancelled request advanced the interval clock")
234	}
235	if base.calls != 1 {
236		t.Errorf("%d upstream calls, want 1: the cancelled request must not go upstream", base.calls)
237	}
238}
239
240// TestErrorPageMapsAShedRequestTo503 pins the user-facing side: a shed
241// request is a 503 with Retry-After, not a DeviantArt error.
242func TestErrorPageMapsAShedRequestTo503(t *testing.T) {
243	rec := httptest.NewRecorder()
244	skunkyart{Writer: rec, Host: "http://localhost"}.Error(devianter.Error{Error: "devianter: Get ...: " + errUpstreamBusy.Error()})
245	if rec.Code != 503 || rec.Header().Get("Retry-After") != "5" {
246		t.Errorf("status %d Retry-After %q, want 503 and 5", rec.Code, rec.Header().Get("Retry-After"))
247	}
248}