app/httpclient_test.go
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}