返回 last30days-skill
test_reddit_keyless_backoff.py
根目录 / tests / test_reddit_keyless_backoff.py
1 """Tests for the shared keyless-Reddit throttle (U4) and 429 in-lane retry."""
2
3 import os
4 from unittest import mock
5
6 from lib import http, reddit_listing, render, schema
7
8
9 class TestRateLimiter:
10 def test_burst_does_not_sleep(self):
11 # A full bucket lets `burst` calls through immediately.
12 limiter = http.RateLimiter(rate_per_sec=5.0, burst=3)
13 with mock.patch.object(http.time, "monotonic", return_value=100.0), \
14 mock.patch.object(http.time, "sleep") as slept:
15 limiter.acquire()
16 limiter.acquire()
17 limiter.acquire()
18 slept.assert_not_called()
19
20 def test_sleeps_when_bucket_empty(self):
21 # burst=1: first call passes, second (same instant) must wait ~1/rate.
22 limiter = http.RateLimiter(rate_per_sec=2.0, burst=1)
23 times = iter([100.0, 100.0, 100.0, 100.5])
24 with mock.patch.object(http.time, "monotonic", side_effect=lambda: next(times)), \
25 mock.patch.object(http.time, "sleep") as slept:
26 limiter.acquire() # consumes the one token
27 limiter.acquire() # bucket empty -> sleep, then refilled token consumed
28 slept.assert_called()
29 waited = slept.call_args.args[0]
30 assert abs(waited - 0.5) < 1e-6 # (1 token deficit) / 2 per sec
31
32 def test_refill_over_time_avoids_sleep(self):
33 limiter = http.RateLimiter(rate_per_sec=2.0, burst=1)
34 # Second call 1s later: bucket refilled (2/s * 1s capped at burst=1) -> no sleep.
35 times = iter([100.0, 101.0])
36 with mock.patch.object(http.time, "monotonic", side_effect=lambda: next(times)), \
37 mock.patch.object(http.time, "sleep") as slept:
38 limiter.acquire()
39 limiter.acquire()
40 slept.assert_not_called()
41
42
43 class TestRedditKeylessGetText:
44 def test_acquires_limiter_then_delegates(self):
45 with mock.patch.object(http.REDDIT_KEYLESS_LIMITER, "acquire") as acq, \
46 mock.patch.object(http, "get_text", return_value="body") as gt:
47 out = http.reddit_keyless_get_text("https://www.reddit.com/svc/shreddit/search/?q=x", accept="text/html")
48 assert out == "body"
49 acq.assert_called_once()
50 gt.assert_called_once()
51
52
53 class TestRedditKeylessRateKnob:
54 def test_default_rate_is_one_per_sec_small_burst(self):
55 limiter = http.make_reddit_keyless_limiter(environ={})
56 assert limiter.rate == 1.0
57 assert limiter.capacity == 2
58
59 def test_env_override(self):
60 limiter = http.make_reddit_keyless_limiter(
61 environ={http.REDDIT_KEYLESS_RATE_ENV: "0.25"}
62 )
63 assert limiter.rate == 0.25
64 assert limiter.capacity == 2
65
66 def test_invalid_and_nonpositive_fall_back_to_default(self):
67 for raw in ("fast", "", "-1", "0", "nan", "inf"):
68 limiter = http.make_reddit_keyless_limiter(
69 environ={http.REDDIT_KEYLESS_RATE_ENV: raw}
70 )
71 assert limiter.rate == http.DEFAULT_REDDIT_KEYLESS_RATE, raw
72
73 def test_process_env_syncs_onto_shared_limiter(self, monkeypatch):
74 monkeypatch.setattr(http.REDDIT_KEYLESS_LIMITER, "rate", 1.0)
75 monkeypatch.setenv(http.REDDIT_KEYLESS_RATE_ENV, "0.5")
76 with mock.patch.object(http.REDDIT_KEYLESS_LIMITER, "acquire"), \
77 mock.patch.object(http, "get_text", return_value="ok"):
78 http.reddit_keyless_get_text("https://www.reddit.com/svc/shreddit/search/?q=x")
79 assert http.REDDIT_KEYLESS_LIMITER.rate == 0.5
80
81 def test_env_file_value_is_exported_for_limiter(self, tmp_path, monkeypatch):
82 from lib import env
83
84 config_file = tmp_path / ".env"
85 config_file.write_text(f"{http.REDDIT_KEYLESS_RATE_ENV}=0.25\n", encoding="utf-8")
86 config_file.chmod(0o600)
87 monkeypatch.setattr(env, "CONFIG_DIR", tmp_path)
88 monkeypatch.setattr(env, "CONFIG_FILE", config_file)
89 monkeypatch.setenv("LAST30DAYS_CONFIG_DIR", str(tmp_path))
90 monkeypatch.delenv(http.REDDIT_KEYLESS_RATE_ENV, raising=False)
91 monkeypatch.chdir(tmp_path)
92 with mock.patch.object(env, "_load_keychain", return_value={}), \
93 mock.patch.object(env, "_load_pass", return_value={}):
94 config = env.get_config()
95 assert config[http.REDDIT_KEYLESS_RATE_ENV] == "0.25"
96 assert http.parse_reddit_keyless_rate(
97 config[http.REDDIT_KEYLESS_RATE_ENV]
98 ) == 0.25
99 assert os.environ.get(http.REDDIT_KEYLESS_RATE_ENV) == "0.25"
100
101
102 def _record_status(code: int, reason: str) -> None:
103 http._record_failure(http.HTTPError(f"HTTP {code}: {reason}", code))
104
105
106 class TestRedditKeyless429Retry:
107 def test_429_then_200_recovers_via_single_retry(self):
108 bodies = [None, "<html>ok</html>"]
109
110 def fake_get(*_args, **_kwargs):
111 val = bodies.pop(0)
112 if val is None:
113 _record_status(429, "Too Many Requests")
114 return val
115
116 with mock.patch.object(http, "get_text", side_effect=fake_get), \
117 mock.patch.object(http.REDDIT_KEYLESS_LIMITER, "acquire") as acq, \
118 mock.patch.object(http.time, "sleep") as slept, \
119 mock.patch.object(http.random, "uniform", return_value=0.0):
120 text, err = http.reddit_keyless_get_text_retry_429(
121 "https://www.reddit.com/svc/shreddit/search/?q=x"
122 )
123 assert text.startswith("<html")
124 assert err is None
125 assert acq.call_count == 2
126 slept.assert_called_once()
127 assert bodies == []
128
129 def test_429_then_429_records_failure_and_stops(self):
130 def fake_get(*_args, **_kwargs):
131 _record_status(429, "Too Many Requests")
132 return None
133
134 with mock.patch.object(http, "get_text", side_effect=fake_get) as gt, \
135 mock.patch.object(http.REDDIT_KEYLESS_LIMITER, "acquire"), \
136 mock.patch.object(http.time, "sleep"), \
137 mock.patch.object(http.random, "uniform", return_value=0.0), \
138 http.capture_failures() as failures:
139 text, err = http.reddit_keyless_get_text_retry_429(
140 "https://www.reddit.com/svc/shreddit/search/?q=x"
141 )
142 assert text is None
143 assert err is not None and "429" in err
144 assert gt.call_count == 2
145 assert any(f.status_code == 429 for f in failures)
146
147 def test_non_429_miss_is_not_retried(self):
148 def fake_get(*_args, **_kwargs):
149 _record_status(403, "Forbidden")
150 return None
151
152 with mock.patch.object(http, "get_text", side_effect=fake_get) as gt, \
153 mock.patch.object(http.REDDIT_KEYLESS_LIMITER, "acquire"), \
154 mock.patch.object(http.time, "sleep") as slept, \
155 http.capture_failures() as failures:
156 text, err = http.reddit_keyless_get_text_retry_429(
157 "https://www.reddit.com/svc/shreddit/search/?q=x"
158 )
159 assert text is None
160 assert "403" in (err or "")
161 assert gt.call_count == 1
162 slept.assert_not_called()
163 assert any(f.status_code == 403 for f in failures)
164
165 def test_listing_fetch_recovers_after_one_429(self):
166 from pathlib import Path
167
168 html = (
169 Path(__file__).resolve().parent.parent
170 / "fixtures"
171 / "reddit_listing_cards_sample.html"
172 ).read_text(encoding="utf-8")
173 bodies = [None, html]
174
175 def fake_get(*_args, **_kwargs):
176 val = bodies.pop(0)
177 if val is None:
178 _record_status(429, "Too Many Requests")
179 return val
180
181 with mock.patch.object(http, "get_text", side_effect=fake_get), \
182 mock.patch.object(http.REDDIT_KEYLESS_LIMITER, "acquire"), \
183 mock.patch.object(http.time, "sleep"):
184 items, error = reddit_listing._fetch_one_with_status(
185 "technology", "hot", "netherlands"
186 )
187 assert error is None
188 assert items
189 assert bodies == []
190
191 def test_listing_fetch_records_double_429(self):
192 def fake_get(*_args, **_kwargs):
193 _record_status(429, "Too Many Requests")
194 return None
195
196 with mock.patch.object(http, "get_text", side_effect=fake_get) as gt, \
197 mock.patch.object(http.REDDIT_KEYLESS_LIMITER, "acquire"), \
198 mock.patch.object(http.time, "sleep"), \
199 http.capture_failures() as failures:
200 items, error = reddit_listing._fetch_one_with_status(
201 "technology", "hot", "x"
202 )
203 assert items == []
204 assert error is not None and "429" in error
205 assert gt.call_count == 2
206 assert any(f.status_code == 429 for f in failures)
207
208
209 class TestPartialOutcomeWording:
210 def test_rate_limited_partial_does_not_read_as_cutoff(self):
211 outcome = schema.SourceOutcome(
212 source="reddit",
213 state=schema.PARTIAL,
214 items_returned=8,
215 detail="HTTP 429: Too Many Requests",
216 )
217 text = render._format_outcome(outcome)
218 assert "partial after" not in text
219 assert "8 items returned" in text
220 assert "some requests rate-limited" in text
221 assert "HTTP 429" in text
222
223 def test_non_rate_limit_partial_keeps_count_without_429_claim(self):
224 outcome = schema.SourceOutcome(
225 source="instagram",
226 state=schema.PARTIAL,
227 items_returned=1,
228 detail="HTTP 400: Bad Request",
229 )
230 text = render._format_outcome(outcome)
231 assert "partial after" not in text
232 assert "1 item returned" in text
233 assert "rate-limited" not in text
234 assert "HTTP 400" in text
235
235 lines PYTHON