diff --git a/agent/reviewer_diff.py b/agent/reviewer_diff.py index c709d41f..a8e1f0f9 100644 --- a/agent/reviewer_diff.py +++ b/agent/reviewer_diff.py @@ -214,15 +214,18 @@ async def fetch_pr_diff( """ import httpx - headers = { - "Accept": "application/vnd.github.diff", - "Authorization": f"Bearer {token}", - "X-GitHub-Api-Version": "2022-11-28", - } + from .utils.github_http import github_client, github_request + url = f"https://api.github.com/repos/{owner}/{repo}/pulls/{pr_number}" try: - async with httpx.AsyncClient() as client: - response = await client.get(url, headers=headers, timeout=timeout) + async with github_client(token=token) as client: + response = await github_request( + client, + "GET", + url, + headers={"Accept": "application/vnd.github.diff"}, + timeout=timeout, + ) response.raise_for_status() except httpx.HTTPError: logger.exception("Failed to fetch PR diff for %s/%s#%s", owner, repo, pr_number) @@ -247,15 +250,12 @@ async def fetch_pr_metadata( """ import httpx - headers = { - "Accept": "application/vnd.github+json", - "Authorization": f"Bearer {token}", - "X-GitHub-Api-Version": "2022-11-28", - } + from .utils.github_http import github_client, github_request + url = f"https://api.github.com/repos/{owner}/{repo}/pulls/{pr_number}" try: - async with httpx.AsyncClient() as client: - response = await client.get(url, headers=headers, timeout=timeout) + async with github_client(token=token) as client: + response = await github_request(client, "GET", url, timeout=timeout) response.raise_for_status() payload = response.json() except httpx.HTTPError: diff --git a/agent/reviewer_publish.py b/agent/reviewer_publish.py index 6b5c2f67..ee491002 100644 --- a/agent/reviewer_publish.py +++ b/agent/reviewer_publish.py @@ -34,14 +34,19 @@ from .reviewer_findings import ( ) from .utils.dashboard_links import dashboard_thread_url from .utils.github_checks import CheckConclusion, complete_review_check_run +from .utils.github_http import ( + GITHUB_API_BASE, + GITHUB_GRAPHQL, + github_client, + github_request, +) from .utils.github_token import GitHubAuthError logger = logging.getLogger(__name__) -_GITHUB_API_BASE = "https://api.github.com" -_GITHUB_GRAPHQL = "https://api.github.com/graphql" -_GITHUB_HEADERS_VERSION = "2022-11-28" +_GITHUB_API_BASE = GITHUB_API_BASE +_GITHUB_GRAPHQL = GITHUB_GRAPHQL _OPEN_SWE_REVIEW_COMMENT_MARKER_RE = re.compile( r"", re.DOTALL, @@ -347,11 +352,9 @@ async def post_status_comment( ) -> int | None: """POST the live status comment to a PR. Returns its comment id or None.""" url = f"{_GITHUB_API_BASE}/repos/{owner}/{repo}/issues/{pr_number}/comments" - async with httpx.AsyncClient() as client: + async with github_client(token=token) as client: try: - response = await client.post( - url, headers=_github_headers(token), json={"body": body}, timeout=30 - ) + response = await github_request(client, "POST", url, json={"body": body}) response.raise_for_status() except httpx.HTTPError: logger.exception("Failed to post status comment for %s/%s#%s", owner, repo, pr_number) @@ -370,9 +373,9 @@ async def delete_status_comment( ) -> bool: """DELETE a status comment by id. Returns True on success (or if gone).""" url = f"{_GITHUB_API_BASE}/repos/{owner}/{repo}/issues/comments/{comment_id}" - async with httpx.AsyncClient() as client: + async with github_client(token=token) as client: try: - response = await client.delete(url, headers=_github_headers(token), timeout=30) + response = await github_request(client, "DELETE", url) if response.status_code == 404: # noqa: PLR2004 return True response.raise_for_status() @@ -505,12 +508,11 @@ async def open_swe_review_exists( """ marker = review_summary_marker(pr_number) url = f"{_GITHUB_API_BASE}/repos/{owner}/{repo}/pulls/{pr_number}/reviews" - headers = _github_headers(token) params: dict[str, Any] = {"per_page": 100, "page": 1} - async with httpx.AsyncClient() as client: + async with github_client(token=token) as client: while True: try: - response = await client.get(url, headers=headers, params=params, timeout=30) + response = await github_request(client, "GET", url, params=params) response.raise_for_status() except httpx.HTTPError: logger.exception( @@ -551,10 +553,9 @@ async def post_pull_request_review( "body": body, "comments": inline_comments, } - headers = _github_headers(token) - async with httpx.AsyncClient() as client: + async with github_client(token=token) as client: try: - response = await client.post(url, headers=headers, json=payload, timeout=30) + response = await github_request(client, "POST", url, json=payload) if response.status_code == 401: raise GitHubAuthError( f"GitHub returned 401 posting PR review for {owner}/{repo}#{pr_number}" @@ -631,13 +632,12 @@ async def fetch_review_comments( per-comment IDs in all paths; this paginates the canonical list endpoint. """ url = f"{_GITHUB_API_BASE}/repos/{owner}/{repo}/pulls/{pr_number}/reviews/{review_id}/comments" - headers = _github_headers(token) out: list[dict[str, Any]] = [] params: dict[str, Any] = {"per_page": 100, "page": 1} - async with httpx.AsyncClient() as client: + async with github_client(token=token) as client: while True: try: - response = await client.get(url, headers=headers, params=params, timeout=30) + response = await github_request(client, "GET", url, params=params) response.raise_for_status() except httpx.HTTPError: logger.exception( @@ -713,12 +713,13 @@ async def fetch_pr_review_threads( """ out: list[dict[str, Any]] = [] cursor: str | None = None - async with httpx.AsyncClient() as client: + async with github_client(token=token) as client: while len(out) < max_threads: try: - response = await client.post( + response = await github_request( + client, + "POST", _GITHUB_GRAPHQL, - headers={"Authorization": f"Bearer {token}"}, json={ "query": query, "variables": { @@ -729,7 +730,6 @@ async def fetch_pr_review_threads( "perThread": max_comments_per_thread, }, }, - timeout=30, ) response.raise_for_status() except httpx.HTTPError: @@ -836,12 +836,13 @@ async def fetch_review_thread_id_for_comment( } """ cursor: str | None = None - async with httpx.AsyncClient() as client: + async with github_client(token=token) as client: while True: try: - response = await client.post( + response = await github_request( + client, + "POST", _GITHUB_GRAPHQL, - headers={"Authorization": f"Bearer {token}"}, json={ "query": query, "variables": { @@ -851,7 +852,6 @@ async def fetch_review_thread_id_for_comment( "cursor": cursor, }, }, - timeout=30, ) response.raise_for_status() except httpx.HTTPError: @@ -899,13 +899,13 @@ async def resolve_review_thread(*, thread_node_id: str, token: str) -> bool: } } """ - async with httpx.AsyncClient() as client: + async with github_client(token=token) as client: try: - response = await client.post( + response = await github_request( + client, + "POST", _GITHUB_GRAPHQL, - headers={"Authorization": f"Bearer {token}"}, json={"query": mutation, "variables": {"threadId": thread_node_id}}, - timeout=30, ) response.raise_for_status() except httpx.HTTPError: @@ -935,14 +935,9 @@ async def reply_to_review_comment( f"{_GITHUB_API_BASE}/repos/{owner}/{repo}/pulls/" f"{pr_number}/comments/{review_comment_id}/replies" ) - async with httpx.AsyncClient() as client: + async with github_client(token=token) as client: try: - response = await client.post( - url, - headers=_github_headers(token), - json={"body": body}, - timeout=30, - ) + response = await github_request(client, "POST", url, json={"body": body}) if response.status_code == 401: raise GitHubAuthError( f"GitHub returned 401 replying to review comment {review_comment_id}" @@ -961,11 +956,3 @@ async def reply_to_review_comment( return None data = response.json() return data if isinstance(data, dict) else None - - -def _github_headers(token: str) -> dict[str, str]: - return { - "Authorization": f"Bearer {token}", - "Accept": "application/vnd.github+json", - "X-GitHub-Api-Version": _GITHUB_HEADERS_VERSION, - } diff --git a/agent/utils/github_checks.py b/agent/utils/github_checks.py index 35a3dd99..8ce86719 100644 --- a/agent/utils/github_checks.py +++ b/agent/utils/github_checks.py @@ -17,24 +17,25 @@ from typing import Literal import httpx +from .github_http import ( + GITHUB_API_BASE, + github_client, + github_headers, # re-exported for backward compatibility + github_request, +) + +__all__ = ["github_headers"] + logger = logging.getLogger(__name__) REVIEW_CHECK_RUN_NAME = "Open SWE Review" AUTOFIX_CHECK_RUN_NAME = "Open SWE Auto-fix" -_GITHUB_API_BASE = "https://api.github.com" +_GITHUB_API_BASE = GITHUB_API_BASE CheckConclusion = Literal["success", "neutral", "failure"] -def github_headers(token: str) -> dict[str, str]: - return { - "Authorization": f"Bearer {token}", - "Accept": "application/vnd.github+json", - "X-GitHub-Api-Version": "2022-11-28", - } - - def _utc_now_iso() -> str: return datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ") @@ -66,10 +67,8 @@ async def create_review_check_run( payload["details_url"] = details_url url = f"{_GITHUB_API_BASE}/repos/{owner}/{repo}/check-runs" try: - async with httpx.AsyncClient() as client: - response = await client.post( - url, headers=github_headers(token), json=payload, timeout=30 - ) + async with github_client(token=token) as client: + response = await github_request(client, "POST", url, json=payload) response.raise_for_status() except httpx.HTTPError: logger.exception( @@ -104,10 +103,8 @@ async def complete_review_check_run( "output": {"title": title, "summary": summary}, } try: - async with httpx.AsyncClient() as client: - response = await client.patch( - url, headers=github_headers(token), json=payload, timeout=30 - ) + async with github_client(token=token) as client: + response = await github_request(client, "PATCH", url, json=payload) response.raise_for_status() except httpx.HTTPError: logger.exception( @@ -147,10 +144,8 @@ async def post_autofix_status_check( payload["details_url"] = details_url url = f"{_GITHUB_API_BASE}/repos/{owner}/{repo}/check-runs" try: - async with httpx.AsyncClient() as client: - response = await client.post( - url, headers=github_headers(token), json=payload, timeout=30 - ) + async with github_client(token=token) as client: + response = await github_request(client, "POST", url, json=payload) response.raise_for_status() except httpx.HTTPError: logger.warning("Failed to post auto-fix status check for %s/%s@%s", owner, repo, head_sha) diff --git a/agent/utils/github_ci.py b/agent/utils/github_ci.py index 1f75f662..8934279e 100644 --- a/agent/utils/github_ci.py +++ b/agent/utils/github_ci.py @@ -16,11 +16,12 @@ from typing import Any import httpx -from .github_checks import REVIEW_CHECK_RUN_NAME, github_headers +from .github_checks import REVIEW_CHECK_RUN_NAME +from .github_http import GITHUB_API_BASE, github_client, github_request logger = logging.getLogger(__name__) -_GITHUB_API_BASE = "https://api.github.com" +_GITHUB_API_BASE = GITHUB_API_BASE # Check-run conclusions that mean "this CI step did not pass" and are worth an # auto-fix attempt. ``cancelled`` / ``stale`` / ``skipped`` are intentionally @@ -46,10 +47,8 @@ async def list_failing_check_runs( url = f"{_GITHUB_API_BASE}/repos/{owner}/{repo}/commits/{ref}/check-runs" params = {"per_page": "100", "filter": "latest"} try: - async with httpx.AsyncClient() as client: - response = await client.get( - url, headers=github_headers(token), params=params, timeout=30 - ) + async with github_client(token=token) as client: + response = await github_request(client, "GET", url, params=params) response.raise_for_status() except httpx.HTTPError: logger.warning( @@ -86,8 +85,8 @@ async def list_failing_statuses( """Return failing legacy commit statuses on ``ref`` (the ``status`` API).""" url = f"{_GITHUB_API_BASE}/repos/{owner}/{repo}/commits/{ref}/status" try: - async with httpx.AsyncClient() as client: - response = await client.get(url, headers=github_headers(token), timeout=30) + async with github_client(token=token) as client: + response = await github_request(client, "GET", url) response.raise_for_status() except httpx.HTTPError: logger.warning("Failed to read combined status for %s/%s@%s", owner, repo, ref) @@ -135,10 +134,8 @@ async def fetch_open_pr_for_branch( url = f"{_GITHUB_API_BASE}/repos/{owner}/{repo}/pulls" params = {"head": f"{owner}:{branch}", "state": "open", "per_page": "1"} try: - async with httpx.AsyncClient() as client: - response = await client.get( - url, headers=github_headers(token), params=params, timeout=30 - ) + async with github_client(token=token) as client: + response = await github_request(client, "GET", url, params=params) response.raise_for_status() except httpx.HTTPError: logger.warning("Failed to find open PR for %s/%s head=%s", owner, repo, branch) @@ -153,8 +150,8 @@ async def fetch_pr(*, owner: str, repo: str, pr_number: int, token: str) -> dict """Fetch full PR metadata (includes ``mergeable_state``).""" url = f"{_GITHUB_API_BASE}/repos/{owner}/{repo}/pulls/{pr_number}" try: - async with httpx.AsyncClient() as client: - response = await client.get(url, headers=github_headers(token), timeout=30) + async with github_client(token=token) as client: + response = await github_request(client, "GET", url) response.raise_for_status() except httpx.HTTPError: logger.warning("Failed to fetch PR %s/%s#%s", owner, repo, pr_number) @@ -167,8 +164,8 @@ async def head_commit_author_login(*, owner: str, repo: str, sha: str, token: st """Return the GitHub login that authored commit ``sha`` (or ``None``).""" url = f"{_GITHUB_API_BASE}/repos/{owner}/{repo}/commits/{sha}" try: - async with httpx.AsyncClient() as client: - response = await client.get(url, headers=github_headers(token), timeout=30) + async with github_client(token=token) as client: + response = await github_request(client, "GET", url) response.raise_for_status() except httpx.HTTPError: logger.debug("Failed to fetch commit %s/%s@%s for author check", owner, repo, sha) @@ -189,8 +186,8 @@ async def has_repo_write_permission(*, owner: str, repo: str, username: str, tok return False url = f"{_GITHUB_API_BASE}/repos/{owner}/{repo}/collaborators/{username}/permission" try: - async with httpx.AsyncClient() as client: - response = await client.get(url, headers=github_headers(token), timeout=30) + async with github_client(token=token) as client: + response = await github_request(client, "GET", url) response.raise_for_status() except httpx.HTTPError: logger.info("Could not verify %s's permission on %s/%s; denying", username, owner, repo) diff --git a/agent/utils/github_http.py b/agent/utils/github_http.py new file mode 100644 index 00000000..3d25f9df --- /dev/null +++ b/agent/utils/github_http.py @@ -0,0 +1,194 @@ +"""Shared GitHub HTTP helper with sane timeouts, retries, and rate-limit handling. + +All GitHub API calls in the reviewer publish path (and gradually everywhere else) +should go through ``github_request`` instead of raw ``httpx.AsyncClient`` calls. +This centralises: + +- **Timeouts**: httpx defaults to 5 s which is too aggressive for paginated + GitHub/GraphQL fetches. The default here is 30 s read / 10 s connect. +- **Retries**: exponential backoff with jitter for retryable HTTP status codes + and transport errors, gated by method idempotency to prevent duplicate writes. + See ``github_request`` for the full retry matrix. +- **Rate-limit awareness**: respects ``Retry-After`` headers and detects + GitHub secondary rate limits (403 with ``X-RateLimit-Remaining: 0`` or a + "secondary rate limit" body message), backing off before retrying. +""" + +from __future__ import annotations + +import asyncio +import logging +import random +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager +from typing import Any + +import httpx + +logger = logging.getLogger(__name__) + +GITHUB_API_BASE = "https://api.github.com" +GITHUB_GRAPHQL = "https://api.github.com/graphql" +GITHUB_HEADERS_VERSION = "2022-11-28" + +DEFAULT_TIMEOUT = httpx.Timeout(30.0, connect=10.0, pool=5.0) +DEFAULT_MAX_RETRIES = 3 + +_ALWAYS_RETRYABLE_STATUS = frozenset({429, 503}) +_IDEMPOTENT_RETRYABLE_STATUS = frozenset({502, 504}) +_SECONDARY_RATE_LIMIT_MARKERS = ("secondary rate limit", "rate limit") +_RETRYABLE_TRANSPORT_METHODS = frozenset({"GET", "HEAD", "OPTIONS", "PUT", "DELETE"}) + +_BASE_BACKOFF = 1.0 +_BACKOFF_MULTIPLIER = 2.0 +_MAX_BACKOFF = 60.0 +_JITTER_FACTOR = 0.25 + + +def github_headers(token: str) -> dict[str, str]: + """Standard GitHub API headers for a bearer token.""" + return { + "Authorization": f"Bearer {token}", + "Accept": "application/vnd.github+json", + "X-GitHub-Api-Version": GITHUB_HEADERS_VERSION, + } + + +def _is_secondary_rate_limit(response: httpx.Response) -> bool: + if response.status_code != 403: + return False + if response.headers.get("X-RateLimit-Remaining") == "0": + return True + body = (response.text or "").lower() + return any(marker in body for marker in _SECONDARY_RATE_LIMIT_MARKERS) + + +def _is_retryable_response(response: httpx.Response, method: str) -> bool: + if response.status_code in _ALWAYS_RETRYABLE_STATUS: + return True + if response.status_code in _IDEMPOTENT_RETRYABLE_STATUS: + return method.upper() in _RETRYABLE_TRANSPORT_METHODS + return _is_secondary_rate_limit(response) + + +def _retry_after_seconds(response: httpx.Response) -> float | None: + retry_after = response.headers.get("Retry-After") + if retry_after: + try: + return float(retry_after) + except ValueError: + return None + return None + + +def _compute_backoff(response: httpx.Response | None, attempt: int) -> float: + if response is not None: + retry_after = _retry_after_seconds(response) + if retry_after is not None: + return min(retry_after, _MAX_BACKOFF) + base = _BASE_BACKOFF * (_BACKOFF_MULTIPLIER**attempt) + jitter = base * random.uniform(-_JITTER_FACTOR, _JITTER_FACTOR) + return min(base + jitter, _MAX_BACKOFF) + + +@asynccontextmanager +async def github_client( + *, + token: str | None = None, + timeout: httpx.Timeout | float | None = None, + headers: dict[str, str] | None = None, +) -> AsyncIterator[httpx.AsyncClient]: + """Yield an ``httpx.AsyncClient`` with sane GitHub defaults. + + The token (when provided) is baked into the default headers so callers + don't need to pass headers on every request. A custom ``timeout`` can + override the default 30 s / 10 s-connect timeout. + """ + merged_headers: dict[str, str] = {} + if token: + merged_headers.update(github_headers(token)) + if headers: + merged_headers.update(headers) + async with httpx.AsyncClient( + headers=merged_headers or None, + timeout=timeout or DEFAULT_TIMEOUT, + ) as client: + yield client + + +async def github_request( + client: httpx.AsyncClient, + method: str, + url: str, + *, + max_retries: int = DEFAULT_MAX_RETRIES, + **kwargs: Any, +) -> httpx.Response: + """Execute a single GitHub API request with retries and rate-limit handling. + + Returns the ``httpx.Response`` for non-retryable status codes and for + retryable status codes that have exhausted retries (caller should call + ``raise_for_status()``). + + Retry matrix: + + | Condition | Idempotent (GET, PUT, DELETE…) | Non-idempotent (POST, PATCH) | + |-----------------------------------|--------------------------------|------------------------------| + | Transport error (timeout/reset) | Retry with backoff | Raise immediately | + | 429 / 503 / secondary rate limit | Retry with backoff | Retry with backoff | + | 502 / 504 (ambiguous gateway) | Retry with backoff | Raise immediately | + + 429 and 503 are safe to retry for any method: the server explicitly did + not process the request. 502/504 are ambiguous — the upstream may have + processed the write before the gateway returned an error — so they are + only retried for idempotent methods. Transport errors are only retried + for idempotent methods for the same reason. + """ + method_upper = method.upper() + retry_transport = method_upper in _RETRYABLE_TRANSPORT_METHODS + method_func = getattr(client, method.lower()) + last_exc: Exception | None = None + for attempt in range(max_retries + 1): + try: + response = await method_func(url, **kwargs) + except (httpx.TimeoutException, httpx.TransportError) as exc: + last_exc = exc + if retry_transport and attempt < max_retries: + delay = _compute_backoff(None, attempt) + logger.warning( + "GitHub API %s %s raised %s, retrying in %.1fs (attempt %d/%d)", + method, + url, + type(exc).__name__, + delay, + attempt + 1, + max_retries, + ) + await asyncio.sleep(delay) + continue + raise + + if _is_retryable_response(response, method): + if attempt < max_retries: + delay = _compute_backoff(response, attempt) + logger.warning( + "GitHub API %s %s returned %d, retrying in %.1fs (attempt %d/%d)", + method, + url, + response.status_code, + delay, + attempt + 1, + max_retries, + ) + await asyncio.sleep(delay) + continue + logger.warning( + "GitHub API %s %s returned %d after %d retries, giving up", + method, + url, + response.status_code, + max_retries, + ) + return response + + raise last_exc or httpx.HTTPError("Max retries exceeded") diff --git a/tests/test_github_checks.py b/tests/test_github_checks.py index cd6dcb0d..c79425b3 100644 --- a/tests/test_github_checks.py +++ b/tests/test_github_checks.py @@ -10,9 +10,16 @@ from agent.utils import github_checks class _FakeResponse: - def __init__(self, payload: dict[str, Any] | None = None, error: bool = False) -> None: + def __init__( + self, + payload: dict[str, Any] | None = None, + error: bool = False, + status_code: int = 200, + ) -> None: self._payload = payload or {} self._error = error + self.status_code = status_code + self.headers: dict[str, str] = {} def raise_for_status(self) -> None: if self._error: @@ -28,6 +35,9 @@ class _FakeAsyncClient: post_response: _FakeResponse = _FakeResponse({"id": 42}) patch_response: _FakeResponse = _FakeResponse({}) + def __init__(self, **kwargs: Any) -> None: + pass + async def __aenter__(self) -> _FakeAsyncClient: return self diff --git a/tests/test_github_ci.py b/tests/test_github_ci.py index 53defe94..443bcd49 100644 --- a/tests/test_github_ci.py +++ b/tests/test_github_ci.py @@ -14,6 +14,8 @@ class _FakeResponse: def __init__(self, payload: Any = None, error: bool = False) -> None: self._payload = payload if payload is not None else {} self._error = error + self.status_code = 200 + self.headers: dict[str, str] = {} def raise_for_status(self) -> None: if self._error: @@ -26,6 +28,9 @@ class _FakeResponse: class _FakeClient: response: _FakeResponse = _FakeResponse({}) + def __init__(self, **kwargs: Any) -> None: + pass + async def __aenter__(self) -> _FakeClient: return self diff --git a/tests/test_github_http.py b/tests/test_github_http.py new file mode 100644 index 00000000..e13c9adb --- /dev/null +++ b/tests/test_github_http.py @@ -0,0 +1,354 @@ +"""Unit tests for the shared GitHub HTTP helper.""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, patch + +import httpx +import pytest + +from agent.utils.github_http import ( + GITHUB_API_BASE, + GITHUB_GRAPHQL, + _compute_backoff, + _is_retryable_response, + _is_secondary_rate_limit, + _retry_after_seconds, + github_client, + github_headers, + github_request, +) + + +def test_github_headers_returns_standard_headers() -> None: + headers = github_headers("mytoken") + assert headers["Authorization"] == "Bearer mytoken" + assert headers["Accept"] == "application/vnd.github+json" + assert headers["X-GitHub-Api-Version"] == "2022-11-28" + + +def test_github_constants() -> None: + assert GITHUB_API_BASE == "https://api.github.com" + assert GITHUB_GRAPHQL == "https://api.github.com/graphql" + + +def _make_response(status_code: int, headers: dict[str, str] | None = None) -> httpx.Response: + return httpx.Response(status_code, headers=headers or {}) + + +class TestIsSecondaryRateLimit: + def test_403_with_rate_limit_remaining_zero(self) -> None: + resp = _make_response(403, {"X-RateLimit-Remaining": "0"}) + assert _is_secondary_rate_limit(resp) + + def test_403_with_secondary_rate_limit_body(self) -> None: + resp = httpx.Response(403, text="You have exceeded a secondary rate limit") + assert _is_secondary_rate_limit(resp) + + def test_403_with_rate_limit_body(self) -> None: + resp = httpx.Response(403, text="API rate limit exceeded") + assert _is_secondary_rate_limit(resp) + + def test_403_without_rate_limit_indicators(self) -> None: + resp = httpx.Response(403, text="Forbidden") + assert not _is_secondary_rate_limit(resp) + + def test_non_403_not_secondary_rate_limit(self) -> None: + resp = _make_response(429, {"X-RateLimit-Remaining": "0"}) + assert not _is_secondary_rate_limit(resp) + + +class TestIsRetryableResponse: + @pytest.mark.parametrize("status", [429, 503]) + def test_always_retryable_status_codes(self, status: int) -> None: + assert _is_retryable_response(_make_response(status), "POST") + assert _is_retryable_response(_make_response(status), "GET") + + @pytest.mark.parametrize("status", [502, 504]) + def test_idempotent_only_retryable_status_codes(self, status: int) -> None: + assert _is_retryable_response(_make_response(status), "GET") + assert not _is_retryable_response(_make_response(status), "POST") + + def test_secondary_rate_limit_is_retryable(self) -> None: + resp = httpx.Response(403, text="secondary rate limit") + assert _is_retryable_response(resp, "POST") + assert _is_retryable_response(resp, "GET") + + def test_200_not_retryable(self) -> None: + assert not _is_retryable_response(_make_response(200), "GET") + + def test_404_not_retryable(self) -> None: + assert not _is_retryable_response(_make_response(404), "GET") + + def test_422_not_retryable(self) -> None: + assert not _is_retryable_response(_make_response(422), "POST") + + +class TestRetryAfterSeconds: + def test_valid_retry_after(self) -> None: + resp = _make_response(429, {"Retry-After": "30"}) + assert _retry_after_seconds(resp) == 30.0 + + def test_no_retry_after_header(self) -> None: + resp = _make_response(429) + assert _retry_after_seconds(resp) is None + + def test_invalid_retry_after(self) -> None: + resp = _make_response(429, {"Retry-After": "not-a-number"}) + assert _retry_after_seconds(resp) is None + + +class TestComputeBackoff: + def test_uses_retry_after_when_present(self) -> None: + resp = _make_response(429, {"Retry-After": "15"}) + delay = _compute_backoff(resp, attempt=0) + assert delay == 15.0 + + def test_caps_retry_after_at_max(self) -> None: + resp = _make_response(429, {"Retry-After": "120"}) + delay = _compute_backoff(resp, attempt=0) + assert delay <= 60.0 + + def test_exponential_backoff_without_response(self) -> None: + delay = _compute_backoff(None, attempt=0) + assert 0.75 <= delay <= 1.25 + + def test_exponential_backoff_attempt_2(self) -> None: + delay = _compute_backoff(None, attempt=2) + base = 1.0 * (2**2) + assert base - base * 0.25 <= delay <= base + base * 0.25 + + def test_backoff_capped_at_max(self) -> None: + delay = _compute_backoff(None, attempt=10) + assert delay <= 60.0 + + +@pytest.mark.asyncio +async def test_github_client_yields_client_with_token_headers() -> None: + async with github_client(token="testtoken") as client: + assert isinstance(client, httpx.AsyncClient) + assert client.headers["Authorization"] == "Bearer testtoken" + assert client.headers["Accept"] == "application/vnd.github+json" + + +@pytest.mark.asyncio +async def test_github_client_without_token() -> None: + async with github_client() as client: + assert isinstance(client, httpx.AsyncClient) + assert "Authorization" not in client.headers + + +@pytest.mark.asyncio +async def test_github_request_retries_on_429() -> None: + responses = [ + _make_response(429, {"Retry-After": "0"}), + _make_response(200), + ] + client = AsyncMock() + client.get = AsyncMock(side_effect=responses) + + with patch("agent.utils.github_http.asyncio.sleep", new_callable=AsyncMock): + response = await github_request(client, "GET", "https://api.github.com/test") + + assert response.status_code == 200 + assert client.get.await_count == 2 + + +@pytest.mark.asyncio +async def test_github_request_retries_on_503() -> None: + responses = [ + _make_response(503), + _make_response(200), + ] + client = AsyncMock() + client.post = AsyncMock(side_effect=responses) + + with patch("agent.utils.github_http.asyncio.sleep", new_callable=AsyncMock): + response = await github_request(client, "POST", "https://api.github.com/test") + + assert response.status_code == 200 + assert client.post.await_count == 2 + + +@pytest.mark.asyncio +async def test_github_request_retries_on_secondary_rate_limit() -> None: + responses = [ + httpx.Response(403, text="secondary rate limit", headers={}), + _make_response(200), + ] + client = AsyncMock() + client.get = AsyncMock(side_effect=responses) + + with patch("agent.utils.github_http.asyncio.sleep", new_callable=AsyncMock): + response = await github_request(client, "GET", "https://api.github.com/test") + + assert response.status_code == 200 + assert client.get.await_count == 2 + + +@pytest.mark.asyncio +async def test_github_request_does_not_retry_on_404() -> None: + response_404 = _make_response(404) + client = AsyncMock() + client.get = AsyncMock(return_value=response_404) + + response = await github_request(client, "GET", "https://api.github.com/test") + + assert response.status_code == 404 + assert client.get.await_count == 1 + + +@pytest.mark.asyncio +async def test_github_request_does_not_retry_on_422() -> None: + response_422 = _make_response(422) + client = AsyncMock() + client.post = AsyncMock(return_value=response_422) + + response = await github_request(client, "POST", "https://api.github.com/test") + + assert response.status_code == 422 + assert client.post.await_count == 1 + + +@pytest.mark.asyncio +async def test_github_request_gives_up_after_max_retries() -> None: + response_429 = _make_response(429, {"Retry-After": "0"}) + client = AsyncMock() + client.get = AsyncMock(return_value=response_429) + + with patch("agent.utils.github_http.asyncio.sleep", new_callable=AsyncMock): + response = await github_request(client, "GET", "https://api.github.com/test", max_retries=2) + + assert response.status_code == 429 + assert client.get.await_count == 3 + + +@pytest.mark.asyncio +async def test_github_request_retries_on_timeout() -> None: + responses = [ + httpx.TimeoutException("timeout"), + _make_response(200), + ] + client = AsyncMock() + client.get = AsyncMock(side_effect=responses) + + with patch("agent.utils.github_http.asyncio.sleep", new_callable=AsyncMock): + response = await github_request(client, "GET", "https://api.github.com/test") + + assert response.status_code == 200 + assert client.get.await_count == 2 + + +@pytest.mark.asyncio +async def test_github_request_does_not_retry_transport_error_on_post() -> None: + """Transport errors on POST must not retry — the server may have already + processed the write, and retrying would duplicate the resource.""" + client = AsyncMock() + client.post = AsyncMock(side_effect=httpx.TimeoutException("timeout")) + + with pytest.raises(httpx.TimeoutException): + await github_request(client, "POST", "https://api.github.com/test") + + assert client.post.await_count == 1 + + +@pytest.mark.asyncio +async def test_github_request_retries_on_429_even_for_post() -> None: + """429 is safe to retry for any method — the server explicitly did not + process the request.""" + responses = [ + _make_response(429, {"Retry-After": "0"}), + _make_response(201), + ] + client = AsyncMock() + client.post = AsyncMock(side_effect=responses) + + with patch("agent.utils.github_http.asyncio.sleep", new_callable=AsyncMock): + response = await github_request(client, "POST", "https://api.github.com/test") + + assert response.status_code == 201 + assert client.post.await_count == 2 + + +@pytest.mark.asyncio +async def test_github_request_retries_on_503_even_for_post() -> None: + responses = [ + _make_response(503), + _make_response(201), + ] + client = AsyncMock() + client.post = AsyncMock(side_effect=responses) + + with patch("agent.utils.github_http.asyncio.sleep", new_callable=AsyncMock): + response = await github_request(client, "POST", "https://api.github.com/test") + + assert response.status_code == 201 + assert client.post.await_count == 2 + + +@pytest.mark.asyncio +async def test_github_request_does_not_retry_502_on_post() -> None: + """502 is ambiguous — the upstream may have processed the write before the + gateway returned an error. Must not retry for non-idempotent methods.""" + response_502 = _make_response(502) + client = AsyncMock() + client.post = AsyncMock(return_value=response_502) + + response = await github_request(client, "POST", "https://api.github.com/test") + + assert response.status_code == 502 + assert client.post.await_count == 1 + + +@pytest.mark.asyncio +async def test_github_request_does_not_retry_504_on_post() -> None: + """504 is ambiguous — the upstream may have processed the write before the + gateway timed out. Must not retry for non-idempotent methods.""" + response_504 = _make_response(504) + client = AsyncMock() + client.post = AsyncMock(return_value=response_504) + + response = await github_request(client, "POST", "https://api.github.com/test") + + assert response.status_code == 504 + assert client.post.await_count == 1 + + +@pytest.mark.asyncio +async def test_github_request_retries_502_on_get() -> None: + """502 is safe to retry for idempotent methods.""" + responses = [ + _make_response(502), + _make_response(200), + ] + client = AsyncMock() + client.get = AsyncMock(side_effect=responses) + + with patch("agent.utils.github_http.asyncio.sleep", new_callable=AsyncMock): + response = await github_request(client, "GET", "https://api.github.com/test") + + assert response.status_code == 200 + assert client.get.await_count == 2 + + +@pytest.mark.asyncio +async def test_github_request_raises_after_exhausting_transport_retries() -> None: + client = AsyncMock() + client.get = AsyncMock(side_effect=httpx.ConnectTimeout("timeout")) + + with patch("agent.utils.github_http.asyncio.sleep", new_callable=AsyncMock): + with pytest.raises(httpx.ConnectTimeout): + await github_request(client, "GET", "https://api.github.com/test", max_retries=1) + + assert client.get.await_count == 2 + + +@pytest.mark.asyncio +async def test_github_request_propagates_non_retryable_http_error() -> None: + client = AsyncMock() + client.get = AsyncMock(side_effect=httpx.HTTPError("boom")) + + with pytest.raises(httpx.HTTPError): + await github_request(client, "GET", "https://api.github.com/test") + + assert client.get.await_count == 1 diff --git a/tests/test_reviewer_publish.py b/tests/test_reviewer_publish.py index 5fe296ca..045f576e 100644 --- a/tests/test_reviewer_publish.py +++ b/tests/test_reviewer_publish.py @@ -480,7 +480,7 @@ async def test_resolve_review_thread_returns_true_on_success() -> None: client_cm.__aenter__.return_value = client_cm client_cm.post = AsyncMock(return_value=response) - with patch("agent.reviewer_publish.httpx.AsyncClient", return_value=client_cm): + with patch("agent.utils.github_http.httpx.AsyncClient", return_value=client_cm): ok = await resolve_review_thread(thread_node_id="T_1", token="t") assert ok is True @@ -498,7 +498,7 @@ async def test_fetch_pr_review_threads_handles_null_repository() -> None: client_cm.__aenter__.return_value = client_cm client_cm.post = AsyncMock(return_value=response) - with patch("agent.reviewer_publish.httpx.AsyncClient", return_value=client_cm): + with patch("agent.utils.github_http.httpx.AsyncClient", return_value=client_cm): threads = await fetch_pr_review_threads(owner="o", repo="r", pr_number=1, token="t") assert threads == [] @@ -518,7 +518,7 @@ async def test_post_pull_request_review_non_dict_body_surfaces_status_and_excerp client_cm.__aenter__.return_value = client_cm client_cm.post = AsyncMock(return_value=response) - with patch("agent.reviewer_publish.httpx.AsyncClient", return_value=client_cm): + with patch("agent.utils.github_http.httpx.AsyncClient", return_value=client_cm): result = await post_pull_request_review( owner="o", repo="r", @@ -549,7 +549,7 @@ async def test_resolve_review_thread_returns_false_on_graphql_errors() -> None: client_cm.__aenter__.return_value = client_cm client_cm.post = AsyncMock(return_value=response) - with patch("agent.reviewer_publish.httpx.AsyncClient", return_value=client_cm): + with patch("agent.utils.github_http.httpx.AsyncClient", return_value=client_cm): ok = await resolve_review_thread(thread_node_id="T_1", token="t") assert ok is False @@ -958,7 +958,7 @@ async def test_open_swe_review_exists_detects_summary_marker() -> None: client_cm.__aenter__.return_value = client_cm client_cm.get = AsyncMock(return_value=response) - with patch("agent.reviewer_publish.httpx.AsyncClient", return_value=client_cm): + with patch("agent.utils.github_http.httpx.AsyncClient", return_value=client_cm): exists = await open_swe_review_exists(owner="o", repo="r", pr_number=7, token="t") assert exists is True @@ -973,7 +973,7 @@ async def test_open_swe_review_exists_false_without_marker() -> None: client_cm.__aenter__.return_value = client_cm client_cm.get = AsyncMock(return_value=response) - with patch("agent.reviewer_publish.httpx.AsyncClient", return_value=client_cm): + with patch("agent.utils.github_http.httpx.AsyncClient", return_value=client_cm): exists = await open_swe_review_exists(owner="o", repo="r", pr_number=7, token="t") assert exists is False @@ -989,7 +989,7 @@ async def test_open_swe_review_exists_returns_none_on_http_error() -> None: client_cm.__aenter__.return_value = client_cm client_cm.get = AsyncMock(side_effect=httpx.HTTPError("boom")) - with patch("agent.reviewer_publish.httpx.AsyncClient", return_value=client_cm): + with patch("agent.utils.github_http.httpx.AsyncClient", return_value=client_cm): exists = await open_swe_review_exists(owner="o", repo="r", pr_number=7, token="t") assert exists is None @@ -1665,7 +1665,7 @@ async def test_fetch_pr_review_threads_parses_threads_and_comments() -> None: client_cm.__aenter__.return_value = client_cm client_cm.post = AsyncMock(return_value=response) - with patch("agent.reviewer_publish.httpx.AsyncClient", return_value=client_cm): + with patch("agent.utils.github_http.httpx.AsyncClient", return_value=client_cm): threads = await fetch_pr_review_threads(owner="o", repo="r", pr_number=1, token="t") assert len(threads) == 2 @@ -1688,7 +1688,7 @@ async def test_fetch_pr_review_threads_returns_empty_on_http_error() -> None: client_cm.__aenter__.return_value = client_cm client_cm.post = AsyncMock(side_effect=httpx.HTTPError("boom")) - with patch("agent.reviewer_publish.httpx.AsyncClient", return_value=client_cm): + with patch("agent.utils.github_http.httpx.AsyncClient", return_value=client_cm): threads = await fetch_pr_review_threads(owner="o", repo="r", pr_number=1, token="t") assert threads == [] @@ -1704,7 +1704,7 @@ async def test_reply_to_review_comment_posts_reply_payload() -> None: client_cm.__aenter__.return_value = client_cm client_cm.post = AsyncMock(return_value=response) - with patch("agent.reviewer_publish.httpx.AsyncClient", return_value=client_cm): + with patch("agent.utils.github_http.httpx.AsyncClient", return_value=client_cm): result = await reply_to_review_comment( owner="o", repo="r", @@ -1742,7 +1742,7 @@ async def test_post_pull_request_review_tags_unresolved_anchor_on_422() -> None: client_cm.__aenter__.return_value = client_cm client_cm.post = AsyncMock(return_value=response) - with patch("agent.reviewer_publish.httpx.AsyncClient", return_value=client_cm): + with patch("agent.utils.github_http.httpx.AsyncClient", return_value=client_cm): result = await post_pull_request_review( owner="o", repo="r", @@ -1781,7 +1781,7 @@ async def test_post_pull_request_review_tags_unresolved_anchor_on_line_error() - client_cm.__aenter__.return_value = client_cm client_cm.post = AsyncMock(return_value=response) - with patch("agent.reviewer_publish.httpx.AsyncClient", return_value=client_cm): + with patch("agent.utils.github_http.httpx.AsyncClient", return_value=client_cm): result = await post_pull_request_review( owner="o", repo="r", @@ -1817,7 +1817,7 @@ async def test_post_pull_request_review_does_not_tag_unrelated_422() -> None: client_cm.__aenter__.return_value = client_cm client_cm.post = AsyncMock(return_value=response) - with patch("agent.reviewer_publish.httpx.AsyncClient", return_value=client_cm): + with patch("agent.utils.github_http.httpx.AsyncClient", return_value=client_cm): result = await post_pull_request_review( owner="o", repo="r",