| 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 |