Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
83 changes: 75 additions & 8 deletions app/adapters/osu_api_backoff.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)

Expand Down Expand Up @@ -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

Expand All @@ -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(
Expand All @@ -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
Expand Down Expand Up @@ -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:
Expand All @@ -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)
Expand Down
15 changes: 15 additions & 0 deletions app/adapters/osu_api_v2/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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
Expand Down
10 changes: 10 additions & 0 deletions app/oauth.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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(
Expand Down
81 changes: 81 additions & 0 deletions tests/test_osu_api_backoff.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand Down
Loading