evalstate/diffusers-pr-api
0
1from __future__ import annotations2 3import json4import urllib.error5import urllib.request6from collections.abc import Callable, Iterable7from typing import Any8 9from slop_farmer.data.http import urlopen_with_retry10 11 12class GhReplicaApiRequestError(RuntimeError):13 """Raised when ghreplica returns a non-recoverable HTTP response."""14 15 def __init__(self, status_code: int, path: str, detail: str):16 self.status_code = status_code17 self.path = path18 self.detail = detail19 super().__init__(f"ghreplica API request failed: {status_code} {path} {detail}")20 21 22class GhReplicaProbeUnavailableError(RuntimeError):23 """Raised when ghreplica cannot yet serve a live probe payload."""24 25 def __init__(self, detail: str, *, status_code: int = 503):26 self.status_code = status_code27 super().__init__(detail)28 29 30class GhrProbeClient:31 provider = "ghreplica"32 33 def __init__(34 self,35 *,36 base_url: str,37 timeout: int = 180,38 max_retries: int = 5,39 log: Callable[[str], None] | None = None,40 ):41 self.base_url = base_url.rstrip("/")42 self.timeout = timeout43 self.max_retries = max_retries44 self.log = log45 46 def _request_json(self, path: str) -> Any:47 request = urllib.request.Request(f"{self.base_url}{path}")48 request.add_header("Accept", "application/json")49 try:50 with urlopen_with_retry(51 request,52 timeout=self.timeout,53 max_retries=self.max_retries,54 log=self.log,55 label=path,56 ) as response:57 payload = response.read().decode("utf-8")58 except urllib.error.HTTPError as exc:59 detail = exc.read().decode("utf-8", errors="replace")60 raise GhReplicaApiRequestError(exc.code, path, detail) from exc61 return json.loads(payload)62 63 def _request_json_or_none(self, path: str) -> Any | None:64 try:65 return self._request_json(path)66 except GhReplicaApiRequestError as exc:67 if exc.status_code == 404:68 return None69 raise70 71 def get_pull_request(self, owner: str, repo: str, number: int) -> dict[str, Any]:72 try:73 payload = self._request_json(f"/v1/github/repos/{owner}/{repo}/pulls/{number}")74 except GhReplicaApiRequestError as exc:75 if exc.status_code == 404:76 raise GhReplicaProbeUnavailableError(77 f"PR #{number} was not found in ghreplica.",78 status_code=404,79 ) from exc80 raise81 if not isinstance(payload, dict):82 raise RuntimeError(f"Expected dict payload for pull request, got {type(payload)!r}")83 return payload84 85 def iter_pull_files(self, owner: str, repo: str, number: int) -> Iterable[dict[str, Any]]:86 try:87 payload = self._request_json(f"/v1/changes/repos/{owner}/{repo}/pulls/{number}/files")88 except GhReplicaApiRequestError as exc:89 if exc.status_code != 404:90 raise91 status = self.get_pull_request_status(owner, repo, number)92 if isinstance(status, dict):93 detail_bits = []94 for key in (95 "indexed",96 "backfill_in_progress",97 "changed_files",98 "indexed_file_count",99 ):100 if key in status:101 detail_bits.append(f"{key}={status[key]}")102 suffix = f" ({', '.join(detail_bits)})" if detail_bits else ""103 raise GhReplicaProbeUnavailableError(104 f"PR #{number} is not available in ghreplica yet{suffix}.",105 status_code=503,106 ) from exc107 raise GhReplicaProbeUnavailableError(108 f"PR #{number} was not found in ghreplica changed-file replica.",109 status_code=404,110 ) from exc111 rows = payload if isinstance(payload, list) else payload.get("files")112 if not isinstance(rows, list):113 raise RuntimeError(114 f"Expected list payload for pull request files, got {type(payload)!r}"115 )116 for row in rows:117 if not isinstance(row, dict):118 continue119 additions = int(row.get("additions") or 0)120 deletions = int(row.get("deletions") or 0)121 yield {122 "sha": row.get("sha"),123 "filename": row.get("filename") or row.get("path"),124 "status": row.get("status"),125 "additions": additions,126 "deletions": deletions,127 "changes": row.get("changes") or additions + deletions,128 "blob_url": row.get("blob_url"),129 "raw_url": row.get("raw_url"),130 "contents_url": row.get("contents_url"),131 "previous_filename": row.get("previous_filename"),132 "patch": row.get("patch"),133 }134 135 def get_pull_request_status(self, owner: str, repo: str, number: int) -> dict[str, Any] | None:136 payload = self._request_json_or_none(137 f"/v1/changes/repos/{owner}/{repo}/pulls/{number}/status"138 )139 if payload is None:140 return None141 if not isinstance(payload, dict):142 raise RuntimeError(143 f"Expected dict payload for pull request status, got {type(payload)!r}"144 )145 return payload146 