| @@ -1,11 +1,15 @@ |
| 1 | package app |
1 | package app |
| 2 | |
2 | |
| 3 | import ( |
3 | import ( |
| |
4 | "context" |
| |
5 | "errors" |
| 4 | "net/http" |
6 | "net/http" |
| 5 | "net/http/httptest" |
7 | "net/http/httptest" |
| 6 | "sync" |
8 | "sync" |
| 7 | "testing" |
9 | "testing" |
| 8 | "time" |
10 | "time" |
| |
11 | |
| |
12 | "github.com/krazywarez/devianter" |
| 9 | ) |
13 | ) |
| 10 | |
14 | |
| 11 | // stubTransport records how many requests reached it and returns an empty 200. |
15 | // stubTransport records how many requests reached it and returns an empty 200. |
| @@ -143,3 +147,102 @@ func TestInstallDAThrottlePreservesProxy(t *testing.T) { |
| 143 | t.Error("base transport lost its Proxy func: HTTPS_PROXY / VPN egress would break") |
147 | t.Error("base transport lost its Proxy func: HTTPS_PROXY / VPN egress would break") |
| 144 | } |
148 | } |
| 145 | } |
149 | } |
| |
150 | |
| |
151 | // countingRT counts base round trips for the throttle tests. |
| |
152 | type countingRT struct { |
| |
153 | mu sync.Mutex |
| |
154 | calls int |
| |
155 | } |
| |
156 | |
| |
157 | func (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 | |
| |
164 | func 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. |
| |
172 | func 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. |
| |
194 | func 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. |
| |
215 | func 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. |
| |
242 | func 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 | } |