| @@ -1,11 +1,15 @@ |
| 1 | 1 | package app |
| 2 | 2 | |
| 3 | 3 | import ( |
| 4 | "context" |
| 5 | "errors" |
| 4 | 6 | "net/http" |
| 5 | 7 | "net/http/httptest" |
| 6 | 8 | "sync" |
| 7 | 9 | "testing" |
| 8 | 10 | "time" |
| 11 | |
| 12 | "github.com/krazywarez/devianter" |
| 9 | 13 | ) |
| 10 | 14 | |
| 11 | 15 | // stubTransport records how many requests reached it and returns an empty 200. |
| @@ -143,3 +147,102 @@ func TestInstallDAThrottlePreservesProxy(t *testing.T) { |
| 143 | 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 | } |