| @@ -0,0 +1,137 @@ |
| |
1 | package app |
| |
2 | |
| |
3 | import ( |
| |
4 | "net/http" |
| |
5 | "net/http/httptest" |
| |
6 | "sync" |
| |
7 | "testing" |
| |
8 | "time" |
| |
9 | ) |
| |
10 | |
| |
11 | // stubTransport records how many requests reached it and returns an empty 200. |
| |
12 | type stubTransport struct { |
| |
13 | mu sync.Mutex |
| |
14 | n int |
| |
15 | } |
| |
16 | |
| |
17 | func (s *stubTransport) RoundTrip(req *http.Request) (*http.Response, error) { |
| |
18 | s.mu.Lock() |
| |
19 | s.n++ |
| |
20 | s.mu.Unlock() |
| |
21 | return httptest.NewRecorder().Result(), nil |
| |
22 | } |
| |
23 | |
| |
24 | func newTestThrottle(base http.RoundTripper, gap time.Duration, max int) *daThrottle { |
| |
25 | return &daThrottle{base: base, sem: make(chan struct{}, max)} |
| |
26 | } |
| |
27 | |
| |
28 | // DeviantArt requests must be spaced by at least daMinInterval. |
| |
29 | func TestThrottleRateLimitsDeviantArt(t *testing.T) { |
| |
30 | stub := &stubTransport{} |
| |
31 | tr := newTestThrottle(stub, daMinInterval, daMaxConcurrent) |
| |
32 | |
| |
33 | start := time.Now() |
| |
34 | const n = 3 |
| |
35 | for i := 0; i < n; i++ { |
| |
36 | req, _ := http.NewRequest("GET", "https://www.deviantart.com/_puppy/x", nil) |
| |
37 | if _, err := tr.RoundTrip(req); err != nil { |
| |
38 | t.Fatalf("unexpected error: %v", err) |
| |
39 | } |
| |
40 | } |
| |
41 | elapsed := time.Since(start) |
| |
42 | |
| |
43 | // n requests => at least (n-1) gaps between them. |
| |
44 | if want := time.Duration(n-1) * daMinInterval; elapsed < want { |
| |
45 | t.Errorf("DA requests were not throttled: %d requests took %v, want >= %v", n, elapsed, want) |
| |
46 | } |
| |
47 | if stub.n != n { |
| |
48 | t.Errorf("expected all %d requests to reach the base transport, got %d", n, stub.n) |
| |
49 | } |
| |
50 | } |
| |
51 | |
| |
52 | // Non-DA hosts (e.g. the wixmp image CDN) must not be slowed down. |
| |
53 | func TestThrottleSkipsOtherHosts(t *testing.T) { |
| |
54 | stub := &stubTransport{} |
| |
55 | tr := newTestThrottle(stub, daMinInterval, daMaxConcurrent) |
| |
56 | |
| |
57 | start := time.Now() |
| |
58 | for i := 0; i < 5; i++ { |
| |
59 | req, _ := http.NewRequest("GET", "https://images-wixmp-ed30a86b8c4ca887773594c2.wixmp.com/f/x.jpg", nil) |
| |
60 | if _, err := tr.RoundTrip(req); err != nil { |
| |
61 | t.Fatalf("unexpected error: %v", err) |
| |
62 | } |
| |
63 | } |
| |
64 | |
| |
65 | if elapsed := time.Since(start); elapsed >= daMinInterval { |
| |
66 | t.Errorf("non-DA host was throttled: 5 requests took %v, want < %v", elapsed, daMinInterval) |
| |
67 | } |
| |
68 | if stub.n != 5 { |
| |
69 | t.Errorf("expected 5 requests through, got %d", stub.n) |
| |
70 | } |
| |
71 | } |
| |
72 | |
| |
73 | // Concurrent callers must never exceed daMaxConcurrent in-flight DA requests. |
| |
74 | func TestThrottleCapsConcurrency(t *testing.T) { |
| |
75 | var ( |
| |
76 | mu sync.Mutex |
| |
77 | inFlight int |
| |
78 | peak int |
| |
79 | ) |
| |
80 | counting := roundTripFunc(func(req *http.Request) (*http.Response, error) { |
| |
81 | mu.Lock() |
| |
82 | inFlight++ |
| |
83 | if inFlight > peak { |
| |
84 | peak = inFlight |
| |
85 | } |
| |
86 | mu.Unlock() |
| |
87 | |
| |
88 | time.Sleep(20 * time.Millisecond) // hold the slot |
| |
89 | |
| |
90 | mu.Lock() |
| |
91 | inFlight-- |
| |
92 | mu.Unlock() |
| |
93 | return httptest.NewRecorder().Result(), nil |
| |
94 | }) |
| |
95 | |
| |
96 | tr := newTestThrottle(counting, daMinInterval, daMaxConcurrent) |
| |
97 | |
| |
98 | var wg sync.WaitGroup |
| |
99 | for i := 0; i < 6; i++ { |
| |
100 | wg.Add(1) |
| |
101 | go func() { |
| |
102 | defer wg.Done() |
| |
103 | req, _ := http.NewRequest("GET", "https://www.deviantart.com/_puppy/x", nil) |
| |
104 | tr.RoundTrip(req) |
| |
105 | }() |
| |
106 | } |
| |
107 | wg.Wait() |
| |
108 | |
| |
109 | if peak > daMaxConcurrent { |
| |
110 | t.Errorf("concurrency cap breached: peak %d in-flight DA requests, max %d", peak, daMaxConcurrent) |
| |
111 | } |
| |
112 | } |
| |
113 | |
| |
114 | type roundTripFunc func(*http.Request) (*http.Response, error) |
| |
115 | |
| |
116 | func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) } |
| |
117 | |
| |
118 | // InstallDAThrottle must preserve proxy-from-environment so HTTPS_PROXY (VPN |
| |
119 | // egress) keeps working, and must not panic on a repeat call. |
| |
120 | func TestInstallDAThrottlePreservesProxy(t *testing.T) { |
| |
121 | orig := http.DefaultTransport |
| |
122 | defer func() { http.DefaultTransport = orig }() |
| |
123 | |
| |
124 | InstallDAThrottle() |
| |
125 | |
| |
126 | th, ok := http.DefaultTransport.(*daThrottle) |
| |
127 | if !ok { |
| |
128 | t.Fatalf("DefaultTransport was not wrapped, got %T", http.DefaultTransport) |
| |
129 | } |
| |
130 | base, ok := th.base.(*http.Transport) |
| |
131 | if !ok { |
| |
132 | t.Fatalf("base transport is not *http.Transport, got %T", th.base) |
| |
133 | } |
| |
134 | if base.Proxy == nil { |
| |
135 | t.Error("base transport lost its Proxy func: HTTPS_PROXY / VPN egress would break") |
| |
136 | } |
| |
137 | } |