From d576bcc906fa8710cd31ef7250e3d142ba215777 Mon Sep 17 00:00:00 2001 From: cmyui Date: Sat, 27 Jun 2026 16:30:28 -0400 Subject: [PATCH] Use osu API rate limit headers --- app/adapters/osu_api_backoff.py | 83 +++++++++++++++++++++++++++++---- app/adapters/osu_api_v2/api.py | 15 ++++++ app/oauth.py | 10 ++++ tests/test_osu_api_backoff.py | 81 ++++++++++++++++++++++++++++++++ 4 files changed, 181 insertions(+), 8 deletions(-) diff --git a/app/adapters/osu_api_backoff.py b/app/adapters/osu_api_backoff.py index 11cf85a..e594902 100644 --- a/app/adapters/osu_api_backoff.py +++ b/app/adapters/osu_api_backoff.py @@ -13,6 +13,9 @@ DEFAULT_RATE_LIMIT_COOLDOWN_SECONDS = 60 DEFAULT_OSU_API_REQUESTS_PER_MINUTE = 600 DEFAULT_OSU_API_BURST_SIZE = 10 +RATE_LIMIT_LIMIT_HEADER = "X-Ratelimit-Limit" +RATE_LIMIT_REMAINING_HEADER = "X-Ratelimit-Remaining" +RATE_LIMIT_RESET_HEADER = "X-Ratelimit-Reset" logger = logging.getLogger(__name__) @@ -154,6 +157,7 @@ def record_failure( endpoint: str, cooldown_seconds: float | None = None, force_open: bool = False, + log_extra: dict[str, object] | None = None, ) -> None: cooldown_seconds = cooldown_seconds or self._failure_cooldown_seconds @@ -176,16 +180,20 @@ def record_failure( self._cooldown_seconds = cooldown_seconds self._probe_in_flight = False + extra: dict[str, object] = { + "upstream": upstream, + "endpoint": endpoint, + "previous_state": previous_state, + "circuit_state": CircuitState.OPEN, + "consecutive_failures": self._consecutive_failures, + "cooldown_seconds": cooldown_seconds, + } + if log_extra is not None: + extra.update(log_extra) + logger.warning( "Opened osu! API circuit", - extra={ - "upstream": upstream, - "endpoint": endpoint, - "previous_state": previous_state, - "circuit_state": CircuitState.OPEN, - "consecutive_failures": self._consecutive_failures, - "cooldown_seconds": cooldown_seconds, - }, + extra=extra, ) def apply_if_rate_limited( @@ -209,6 +217,37 @@ def apply_if_rate_limited( f"{upstream} returned {response.status_code}; backing off", ) + def record_rate_limit_headers( + self, + response: httpx.Response, + *, + upstream: str, + endpoint: str, + ) -> None: + remaining = _get_int_header(response, RATE_LIMIT_REMAINING_HEADER) + if remaining is None or remaining > 0: + return + + limit = _get_int_header(response, RATE_LIMIT_LIMIT_HEADER) + cooldown_seconds = _get_rate_limit_reset_cooldown_seconds(response) + if cooldown_seconds is None: + cooldown_seconds = DEFAULT_RATE_LIMIT_COOLDOWN_SECONDS + + self.record_failure( + upstream=upstream, + endpoint=endpoint, + cooldown_seconds=cooldown_seconds, + force_open=True, + log_extra={ + "reason": "rate_limit_headers_exhausted", + "rate_limit": { + "limit": limit, + "remaining": remaining, + "reset": response.headers.get(RATE_LIMIT_RESET_HEADER), + }, + }, + ) + def _evaluate_state(self) -> CircuitState: if self._state != CircuitState.OPEN: return self._state @@ -254,6 +293,10 @@ def _raise_if_rate_limited(self, *, upstream: str) -> None: def _get_rate_limit_cooldown_seconds(response: httpx.Response) -> int: retry_after = response.headers.get("Retry-After") if retry_after is None: + reset_cooldown_seconds = _get_rate_limit_reset_cooldown_seconds(response) + if reset_cooldown_seconds is not None: + return reset_cooldown_seconds + return DEFAULT_RATE_LIMIT_COOLDOWN_SECONDS try: @@ -266,6 +309,30 @@ def _get_rate_limit_cooldown_seconds(response: httpx.Response) -> int: return max(1, int((retry_at - datetime.now(timezone.utc)).total_seconds())) +def _get_rate_limit_reset_cooldown_seconds(response: httpx.Response) -> int | None: + reset = response.headers.get(RATE_LIMIT_RESET_HEADER) + if reset is None: + return None + + try: + reset_at = int(reset) + except ValueError: + return None + + return max(1, reset_at - int(time.time())) + + +def _get_int_header(response: httpx.Response, header: str) -> int | None: + value = response.headers.get(header) + if value is None: + return None + + try: + return int(value) + except ValueError: + return None + + def _parse_retry_after_datetime(retry_after: str) -> datetime | None: try: retry_at = parsedate_to_datetime(retry_after) diff --git a/app/adapters/osu_api_v2/api.py b/app/adapters/osu_api_v2/api.py index 8de08a9..5604d4e 100644 --- a/app/adapters/osu_api_v2/api.py +++ b/app/adapters/osu_api_v2/api.py @@ -52,6 +52,11 @@ async def get_beatmap(beatmap_id: int) -> BeatmapExtended | None: upstream="osu! API v2", endpoint=endpoint, ) + osu_api_v2_backoff.record_rate_limit_headers( + response, + upstream="osu! API v2", + endpoint=endpoint, + ) if response.status_code in (404, 451): osu_api_v2_backoff.record_success( upstream="osu! API v2", @@ -89,6 +94,11 @@ async def get_beatmapset(beatmapset_id: int) -> BeatmapsetExtended | None: upstream="osu! API v2", endpoint=endpoint, ) + osu_api_v2_backoff.record_rate_limit_headers( + response, + upstream="osu! API v2", + endpoint=endpoint, + ) if response.status_code in (404, 451): osu_api_v2_backoff.record_success( upstream="osu! API v2", @@ -158,6 +168,11 @@ async def search_beatmapsets( upstream="osu! API v2", endpoint=endpoint, ) + osu_api_v2_backoff.record_rate_limit_headers( + response, + upstream="osu! API v2", + endpoint=endpoint, + ) response.raise_for_status() osu_api_response_data = response.json() assert osu_api_response_data is not None diff --git a/app/oauth.py b/app/oauth.py index 228a2ce..a9b9e82 100644 --- a/app/oauth.py +++ b/app/oauth.py @@ -73,6 +73,11 @@ async def async_auth_flow( upstream="osu! API v2", endpoint="oauth/token", ) + self.backoff.record_rate_limit_headers( + refresh_response, + upstream="osu! API v2", + endpoint="oauth/token", + ) refresh_response_data = refresh_response.json() if "access_token" not in refresh_response_data: logging.warning( @@ -98,6 +103,11 @@ async def async_auth_flow( upstream="osu! API v2", endpoint="oauth/token", ) + self.backoff.record_rate_limit_headers( + refresh_response, + upstream="osu! API v2", + endpoint="oauth/token", + ) refresh_response_data = refresh_response.json() if "access_token" not in refresh_response_data: logging.warning( diff --git a/tests/test_osu_api_backoff.py b/tests/test_osu_api_backoff.py index a3a27c5..50aeb77 100644 --- a/tests/test_osu_api_backoff.py +++ b/tests/test_osu_api_backoff.py @@ -100,6 +100,87 @@ def test_skipped_calls_do_not_extend_cooldown(self) -> None: mock_time.monotonic.return_value = started_at + 61 self.assertEqual(backoff.state, CircuitState.HALF_OPEN) + def test_zero_remaining_rate_limit_header_opens_circuit(self) -> None: + backoff = OsuApiBackoff() + + backoff.record_rate_limit_headers( + httpx.Response( + 200, + headers={ + "X-Ratelimit-Limit": "1200", + "X-Ratelimit-Remaining": "0", + }, + ), + upstream="osu! API v2", + endpoint="beatmaps", + ) + + self.assertEqual(backoff.state, CircuitState.OPEN) + with self.assertRaises(OsuApiBackoffError): + backoff.raise_if_unavailable(upstream="osu! API v2") + + def test_positive_remaining_rate_limit_header_does_not_open_circuit(self) -> None: + backoff = OsuApiBackoff() + + backoff.record_rate_limit_headers( + httpx.Response( + 200, + headers={ + "X-Ratelimit-Limit": "1200", + "X-Ratelimit-Remaining": "1", + }, + ), + upstream="osu! API v2", + endpoint="beatmaps", + ) + + self.assertEqual(backoff.state, CircuitState.CLOSED) + + def test_rate_limit_reset_header_controls_cooldown(self) -> None: + backoff = OsuApiBackoff() + started_at = time.monotonic() + + with patch("app.adapters.osu_api_backoff.time") as mock_time: + mock_time.monotonic.return_value = started_at + mock_time.time.return_value = 1000 + backoff.record_rate_limit_headers( + httpx.Response( + 200, + headers={ + "X-Ratelimit-Remaining": "0", + "X-Ratelimit-Reset": "1030", + }, + ), + upstream="osu! API v2", + endpoint="beatmaps", + ) + + mock_time.monotonic.return_value = started_at + 29 + self.assertEqual(backoff.state, CircuitState.OPEN) + + mock_time.monotonic.return_value = started_at + 30 + self.assertEqual(backoff.state, CircuitState.HALF_OPEN) + + def test_rate_limit_response_uses_reset_header_without_retry_after(self) -> None: + backoff = OsuApiBackoff() + started_at = time.monotonic() + + with patch("app.adapters.osu_api_backoff.time") as mock_time: + mock_time.monotonic.return_value = started_at + mock_time.time.return_value = 1000 + with self.assertRaises(OsuApiBackoffError): + backoff.apply_if_rate_limited( + httpx.Response(429, headers={"X-Ratelimit-Reset": "1030"}), + upstream="osu! API v2", + endpoint="beatmaps", + ) + + mock_time.monotonic.return_value = started_at + 29 + self.assertEqual(backoff.state, CircuitState.OPEN) + + mock_time.monotonic.return_value = started_at + 30 + self.assertEqual(backoff.state, CircuitState.HALF_OPEN) + def test_in_flight_success_does_not_close_open_circuit(self) -> None: backoff = OsuApiBackoff()