354 lines
14 KiB
Python
354 lines
14 KiB
Python
"""Reusable HTTP API clients for Gitea and Vikunja."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
import time
|
|
from typing import Any
|
|
|
|
import requests
|
|
|
|
from .config import DEFAULT_TIMEOUT
|
|
from .exceptions import APIError
|
|
|
|
logger = logging.getLogger("grm")
|
|
|
|
# Retry configuration for transient errors (429, 5xx, connection errors)
|
|
MAX_RETRIES = 3
|
|
RETRY_BACKOFF_BASE = 2 # seconds: 2, 4, 8
|
|
RETRY_STATUS_CODES = {429, 500, 502, 503, 504}
|
|
|
|
|
|
def _parse_error(e: requests.HTTPError) -> tuple[int, str]:
|
|
"""Extract status code and message from an HTTPError response."""
|
|
response = getattr(e, "response", None)
|
|
status = response.status_code if response is not None else 0
|
|
try:
|
|
body: dict[str, Any] = response.json() if response is not None else {}
|
|
message: str = body.get("message", str(e))
|
|
except Exception:
|
|
message = str(e)
|
|
return status, message
|
|
|
|
|
|
def _is_retryable(e: Exception) -> bool:
|
|
"""Check if an exception is a transient error worth retrying."""
|
|
if isinstance(e, requests.ConnectionError):
|
|
return True
|
|
if isinstance(e, requests.HTTPError):
|
|
status, _ = _parse_error(e)
|
|
return status in RETRY_STATUS_CODES
|
|
return isinstance(e, requests.Timeout)
|
|
|
|
|
|
class GiteaClient:
|
|
"""Low-level Gitea REST API client with connection pooling."""
|
|
|
|
def __init__(self, base_url: str, token: str, owner: str, repo: str) -> None:
|
|
self._base_url = base_url.rstrip("/")
|
|
self._owner = owner
|
|
self._repo = repo
|
|
self._session = requests.Session()
|
|
self._session.headers.update(
|
|
{
|
|
"Authorization": f"token {token}",
|
|
"Content-Type": "application/json",
|
|
}
|
|
)
|
|
|
|
def _url(self, path: str) -> str:
|
|
return f"{self._base_url}/repos/{self._owner}/{self._repo}{path}"
|
|
|
|
def _request(self, method: str, path: str, **kwargs: Any) -> requests.Response:
|
|
url = self._url(path)
|
|
last_exc: Exception | None = None
|
|
for attempt in range(MAX_RETRIES):
|
|
try:
|
|
response = self._session.request(method, url, timeout=DEFAULT_TIMEOUT, **kwargs)
|
|
response.raise_for_status()
|
|
return response
|
|
except requests.HTTPError as e:
|
|
status, message = _parse_error(e)
|
|
if _is_retryable(e) and attempt < MAX_RETRIES - 1:
|
|
wait = RETRY_BACKOFF_BASE ** (attempt + 1)
|
|
logger.warning(
|
|
"Transient HTTP %d on %s %s, retrying in %ds (attempt %d/%d)",
|
|
status,
|
|
method,
|
|
path,
|
|
wait,
|
|
attempt + 1,
|
|
MAX_RETRIES,
|
|
)
|
|
time.sleep(wait)
|
|
last_exc = e
|
|
continue
|
|
raise APIError(status, message) from e
|
|
except (requests.ConnectionError, requests.Timeout) as e:
|
|
if attempt < MAX_RETRIES - 1:
|
|
wait = RETRY_BACKOFF_BASE ** (attempt + 1)
|
|
logger.warning(
|
|
"Connection error on %s %s, retrying in %ds (attempt %d/%d)",
|
|
method,
|
|
path,
|
|
wait,
|
|
attempt + 1,
|
|
MAX_RETRIES,
|
|
)
|
|
time.sleep(wait)
|
|
last_exc = e
|
|
continue
|
|
raise APIError(0, str(e)) from e
|
|
# Should not reach here, but just in case
|
|
if last_exc: # pragma: no cover
|
|
raise APIError(0, str(last_exc)) from last_exc
|
|
raise APIError(0, "Max retries exceeded") # pragma: no cover
|
|
|
|
# -- repo settings --
|
|
|
|
def update_repo_settings(self, settings: dict[str, Any]) -> dict[str, Any]:
|
|
"""Update repository settings (e.g. auto-delete branch after merge)."""
|
|
r = self._request("PATCH", "", json=settings)
|
|
return r.json()
|
|
|
|
# -- branch protection --
|
|
|
|
def list_branch_protections(self) -> list[dict[str, Any]]:
|
|
r = self._request("GET", "/branch_protections")
|
|
return r.json()
|
|
|
|
def create_branch_protection(self, config: dict[str, Any]) -> dict[str, Any]:
|
|
r = self._request("POST", "/branch_protections", json=config)
|
|
return r.json()
|
|
|
|
def update_branch_protection(self, branch: str, config: dict[str, Any]) -> dict[str, Any]:
|
|
r = self._request("PATCH", f"/branch_protections/{branch}", json=config)
|
|
return r.json()
|
|
|
|
def ensure_branch_protection(self, branch: str, config: dict[str, Any]) -> dict[str, Any]:
|
|
"""Idempotent: create or update branch protection for the given branch."""
|
|
existing = self.list_branch_protections()
|
|
for p in existing:
|
|
if p.get("branch_name") == branch:
|
|
update_config = {k: v for k, v in config.items() if k != "branch_name"}
|
|
return self.update_branch_protection(branch, update_config)
|
|
return self.create_branch_protection(config)
|
|
|
|
# -- labels --
|
|
|
|
def list_labels(self) -> list[dict[str, Any]]:
|
|
r = self._request("GET", "/labels")
|
|
return r.json()
|
|
|
|
def create_label(self, name: str, color: str, description: str = "") -> dict[str, Any]:
|
|
r = self._request(
|
|
"POST",
|
|
"/labels",
|
|
json={"name": name, "color": color, "description": description},
|
|
)
|
|
return r.json()
|
|
|
|
def ensure_label(self, name: str, color: str, description: str = "") -> dict[str, Any] | None:
|
|
"""Idempotent: create label if it doesn't already exist."""
|
|
labels = self.list_labels()
|
|
for label in labels:
|
|
if label["name"] == name:
|
|
return None
|
|
return self.create_label(name, color, description)
|
|
|
|
def create_issue(self, title: str, body: str = "", labels: list[int] | None = None) -> dict[str, Any]:
|
|
"""Create a new issue in the repository.
|
|
|
|
Args:
|
|
labels: List of label IDs (integers, not names).
|
|
"""
|
|
payload: dict[str, Any] = {"title": title, "body": body}
|
|
if labels:
|
|
payload["labels"] = labels
|
|
r = self._request("POST", "/issues", json=payload)
|
|
return r.json()
|
|
|
|
# -- pulls / releases --
|
|
|
|
def get_pr_labels(self, pr_number: str | int) -> list[dict[str, Any]]:
|
|
"""Fetch labels currently attached to a pull request."""
|
|
r = self._request("GET", f"/issues/{pr_number}/labels")
|
|
return r.json()
|
|
|
|
def merge_pr(self, pr_number: str | int, merge_title: str) -> None:
|
|
payload = {"Do": "squash", "MergeTitleField": merge_title}
|
|
self._request("POST", f"/pulls/{pr_number}/merge", json=payload)
|
|
|
|
def get_commit_status(self, sha: str) -> list[dict[str, Any]]:
|
|
"""Fetch all status check contexts reported for a commit.
|
|
|
|
Uses the combined status endpoint (/commits/{sha}/status) which
|
|
returns one entry per context (the latest), deduplicated server-side.
|
|
The plural endpoint (/commits/{sha}/statuses) returns every historical
|
|
entry including stale "pending" ones that never got updated.
|
|
"""
|
|
r = self._request("GET", f"/commits/{sha}/status")
|
|
data = r.json()
|
|
return data.get("statuses", [])
|
|
|
|
def get_pr(self, pr_number: str | int) -> dict[str, Any]:
|
|
"""Fetch pull request details including mergeable state."""
|
|
r = self._request("GET", f"/pulls/{pr_number}")
|
|
return r.json()
|
|
|
|
def get_pr_files(self, pr_number: str | int) -> list[dict[str, Any]]:
|
|
"""Fetch the list of files changed in a pull request."""
|
|
r = self._request("GET", f"/pulls/{pr_number}/files")
|
|
return r.json()
|
|
|
|
def get_pr_commits(self, pr_number: str | int) -> list[dict[str, Any]]:
|
|
"""Fetch the commits included in a pull request."""
|
|
r = self._request("GET", f"/pulls/{pr_number}/commits")
|
|
return r.json()
|
|
|
|
def get_pr_reviews(self, pr_number: str | int) -> list[dict[str, Any]]:
|
|
"""Fetch reviews posted on a pull request."""
|
|
r = self._request("GET", f"/pulls/{pr_number}/reviews")
|
|
return r.json()
|
|
|
|
def create_review(
|
|
self,
|
|
pr_number: str | int,
|
|
event: str = "COMMENT",
|
|
body: str = "",
|
|
comments: list[dict[str, Any]] | None = None,
|
|
) -> dict[str, Any]:
|
|
"""Post a review on a pull request.
|
|
|
|
Args:
|
|
event: ``APPROVED``, ``REQUEST_CHANGES``, or ``COMMENT``.
|
|
body: Top-level review body text.
|
|
comments: Line-level comments with ``path``, ``body``,
|
|
``new_position`` (and optionally ``old_position``).
|
|
"""
|
|
# Map common event names to Gitea API values
|
|
event_map = {"APPROVE": "APPROVED", "REQUEST_CHANGES": "REQUEST_CHANGES", "COMMENT": "COMMENT"}
|
|
gitea_event = event_map.get(event, event)
|
|
payload: dict[str, Any] = {"event": gitea_event, "body": body}
|
|
if comments:
|
|
payload["comments"] = comments
|
|
r = self._request("POST", f"/pulls/{pr_number}/reviews", json=payload)
|
|
return r.json()
|
|
|
|
def create_release(
|
|
self,
|
|
tag: str,
|
|
name: str = "",
|
|
body: str = "",
|
|
draft: bool = False,
|
|
prerelease: bool = False,
|
|
) -> dict[str, Any]:
|
|
payload = {
|
|
"tag_name": tag,
|
|
"name": name or tag,
|
|
"body": body,
|
|
"draft": draft,
|
|
"prerelease": prerelease,
|
|
}
|
|
r = self._request("POST", "/releases", json=payload)
|
|
return r.json()
|
|
|
|
def get_release_by_tag(self, tag: str) -> dict[str, Any] | None:
|
|
"""Fetch a release by its tag name. Returns None if not found."""
|
|
try:
|
|
r = self._request("GET", f"/releases/tags/{tag}")
|
|
return r.json()
|
|
except APIError:
|
|
return None
|
|
|
|
def create_release_idempotent(
|
|
self,
|
|
tag: str,
|
|
name: str = "",
|
|
body: str = "",
|
|
draft: bool = False,
|
|
prerelease: bool = False,
|
|
) -> dict[str, Any]:
|
|
"""Create a release, or return the existing one if it already exists.
|
|
|
|
This is idempotent — safe to call multiple times for the same tag.
|
|
"""
|
|
existing = self.get_release_by_tag(tag)
|
|
if existing:
|
|
logger.info("Release for tag %s already exists (ID %s), skipping creation.", tag, existing.get("id"))
|
|
return existing
|
|
return self.create_release(tag=tag, name=name, body=body, draft=draft, prerelease=prerelease)
|
|
|
|
|
|
class VikunjaClient:
|
|
"""Low-level Vikunja REST API client with connection pooling."""
|
|
|
|
def __init__(self, base_url: str, token: str) -> None:
|
|
self._base_url = base_url.rstrip("/")
|
|
self._session = requests.Session()
|
|
self._session.headers.update({"Authorization": f"Bearer {token}"})
|
|
|
|
def _request(self, method: str, path: str, **kwargs: Any) -> requests.Response:
|
|
url = f"{self._base_url}{path}"
|
|
last_exc: Exception | None = None
|
|
for attempt in range(MAX_RETRIES):
|
|
try:
|
|
response = self._session.request(method, url, timeout=DEFAULT_TIMEOUT, **kwargs)
|
|
response.raise_for_status()
|
|
return response
|
|
except requests.HTTPError as e:
|
|
status, message = _parse_error(e)
|
|
if _is_retryable(e) and attempt < MAX_RETRIES - 1:
|
|
wait = RETRY_BACKOFF_BASE ** (attempt + 1)
|
|
logger.warning(
|
|
"Transient HTTP %d on %s %s, retrying in %ds (attempt %d/%d)",
|
|
status,
|
|
method,
|
|
path,
|
|
wait,
|
|
attempt + 1,
|
|
MAX_RETRIES,
|
|
)
|
|
time.sleep(wait)
|
|
last_exc = e
|
|
continue
|
|
raise APIError(status, message) from e
|
|
except (requests.ConnectionError, requests.Timeout) as e:
|
|
if attempt < MAX_RETRIES - 1:
|
|
wait = RETRY_BACKOFF_BASE ** (attempt + 1)
|
|
logger.warning(
|
|
"Connection error on %s %s, retrying in %ds (attempt %d/%d)",
|
|
method,
|
|
path,
|
|
wait,
|
|
attempt + 1,
|
|
MAX_RETRIES,
|
|
)
|
|
time.sleep(wait)
|
|
last_exc = e
|
|
continue
|
|
raise APIError(0, str(e)) from e
|
|
if last_exc: # pragma: no cover
|
|
raise APIError(0, str(last_exc)) from last_exc
|
|
raise APIError(0, "Max retries exceeded") # pragma: no cover
|
|
|
|
def list_tasks(self, **params: Any) -> list[dict[str, Any]]:
|
|
r = self._request("GET", "/tasks", params=params)
|
|
return r.json()
|
|
|
|
def get_task(self, task_id: int) -> dict[str, Any]:
|
|
"""Fetch a single task by its numeric ID."""
|
|
r = self._request("GET", f"/tasks/{task_id}")
|
|
return r.json()
|
|
|
|
def list_project_tasks(self, project_id: int, **params: Any) -> list[dict[str, Any]]:
|
|
"""List tasks in a specific project (more efficient than listing all tasks)."""
|
|
r = self._request("GET", f"/projects/{project_id}/tasks", params=params)
|
|
return r.json()
|
|
|
|
def post_comment(self, task_id: int, comment: str) -> None:
|
|
self._request("PUT", f"/tasks/{task_id}/comments", json={"comment": comment})
|
|
|
|
def update_task(self, task_id: int, **fields: Any) -> None:
|
|
self._request("POST", f"/tasks/{task_id}", json=fields)
|