Skip to content
Merged
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
102 changes: 95 additions & 7 deletions src/openhound_github/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@
from datetime import datetime, timedelta, timezone
from threading import Lock
from typing import Iterator
from urllib.parse import urlparse
from weakref import WeakKeyDictionary

import requests
from dlt.common.configuration import configspec
Expand All @@ -17,6 +19,19 @@
logger = logging.getLogger(__name__)


def _normalized_http_origin(url: str) -> tuple[str, str, int | None]:
parsed = urlparse(url)
scheme = parsed.scheme.lower()
if scheme != "https" or not parsed.hostname:
raise ValueError("GitHub API URI must be an absolute HTTPS URL")

port = parsed.port
if scheme == "https" and port == 443:
port = None

return scheme, parsed.hostname.lower(), port


class AccountConfig(BaseModel):
id: int
login: str | None = None
Expand Down Expand Up @@ -63,7 +78,8 @@ def __init__(
private_key_path: str,
api_uri: str = "https://api.github.com/",
):
self.api_uri = api_uri
_normalized_http_origin(api_uri)
self.api_uri = f"{api_uri.rstrip('/')}/"
self.jwt_issuer = jwt_issuer
self.private_key_path = private_key_path
self.client = RESTClient(
Expand Down Expand Up @@ -158,11 +174,28 @@ def __init__(
self,
installation: GithubInstallation,
refresh_margin_seconds: int = 300,
api_uri: str | None = None,
):
self.installation = installation
self.refresh_margin_seconds = refresh_margin_seconds
installation_api_uri = getattr(
installation, "api_uri", "https://api.github.com/"
)
selected_api_uri = api_uri or installation_api_uri
installation_api_origin = _normalized_http_origin(installation_api_uri)
selected_api_origin = _normalized_http_origin(selected_api_uri)
if selected_api_origin != installation_api_origin:
raise ValueError(
"GitHub App auth API URI origin must match installation API URI origin"
)

self.api_uri = selected_api_uri
self._api_origin = selected_api_origin
self.access_token: str | None = None
self.expires_at: datetime | None = None
self._response_refreshed_requests: WeakKeyDictionary[
requests.PreparedRequest, str
] = WeakKeyDictionary()
self._token_lock = Lock()

def _should_refresh(self) -> bool:
Expand All @@ -172,6 +205,14 @@ def _should_refresh(self) -> bool:
refresh_at = self.expires_at - timedelta(seconds=self.refresh_margin_seconds)
return datetime.now(timezone.utc) >= refresh_at

def _refresh_token(self) -> None:
logger.info(
f"Refreshing access token for {self.installation.installation_id}"
)
get_token = self.installation.token
self.access_token = get_token.token
self.expires_at = get_token.expires_at

def token(self, force_refresh: bool = False) -> str | None:
if (
not force_refresh
Expand All @@ -182,15 +223,62 @@ def token(self, force_refresh: bool = False) -> str | None:

with self._token_lock:
if (force_refresh or self._should_refresh()) or self.access_token is None:
logger.info(
f"Refreshing access token for {self.installation.installation_id}"
)
get_token = self.installation.token
self.access_token = get_token.token
self.expires_at = get_token.expires_at
self._refresh_token()

return self.access_token

def refresh_request(self, request: requests.PreparedRequest) -> bool:
"""Repair a rejected same-origin request without stampeding token issuance."""
try:
request_origin = _normalized_http_origin(request.url or "")
except ValueError:
return False

if request_origin != self._api_origin:
return False

request_authorization = request.headers.get("Authorization")
if (
not request_authorization
or not request_authorization.startswith("Bearer ")
or not request_authorization.removeprefix("Bearer ").strip()
):
return False

with self._token_lock:
current_authorization = (
f"Bearer {self.access_token}" if self.access_token is not None else None
)
should_refresh = self._should_refresh()
repaired_authorization = self._response_refreshed_requests.get(request)
if (
not should_refresh
and request_authorization == current_authorization
and request_authorization == repaired_authorization
):
return False

if (
self.access_token is None
or should_refresh
or request_authorization == current_authorization
):
try:
self._refresh_token()
except Exception:
logger.warning(
"Failed to refresh GitHub App installation token for "
"installation %s during request retry",
self.installation.installation_id,
)
return False

replacement_authorization = f"Bearer {self.access_token}"
request.headers["Authorization"] = replacement_authorization
self._response_refreshed_requests[request] = replacement_authorization

return True

def __call__(self, request: requests.PreparedRequest) -> requests.PreparedRequest:
request.headers["Authorization"] = f"Bearer {self.token()}"
return request
Expand Down
19 changes: 18 additions & 1 deletion src/openhound_github/helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@
)
from requests import Request

from openhound_github.auth import GitHubAppInstallationAuth

logger = logging.getLogger(__name__)


Expand Down Expand Up @@ -160,6 +162,22 @@ def retry_policy(

headers = response.headers
now = int(time.time())
message = _response_message(response).lower()

# DLT retries the same prepared request after long Retry-After sleeps.
if (
response.status_code == 401
and "bad credentials" in message
and isinstance(auth, GitHubAppInstallationAuth)
and response.request is not None
):
if not auth.refresh_request(response.request):
return False
logger.warning(
"GitHub App installation token rejected, retrying request with refreshed token"
)
return True

if (
response.status_code == 200
and headers.get("x-ratelimit-resource") == "graphql"
Expand All @@ -178,7 +196,6 @@ def retry_policy(
return True
return False

message = _response_message(response).lower()
if response.status_code not in (403, 429):
return False

Expand Down
21 changes: 18 additions & 3 deletions src/openhound_github/source.py
Original file line number Diff line number Diff line change
Expand Up @@ -172,19 +172,24 @@ def token_client(token: str) -> RESTClient:
github_app_session = GithubApp(
jwt_issuer=jwt_issuer,
private_key_path=credentials.key_path,
api_uri=host,
)
for installation in github_app_session.installations:
if installation.target_type == "Organization":
org_installation = GithubInstallation(
installation_id=installation.id,
jwt_issuer=jwt_issuer,
private_key_path=credentials.key_path,
api_uri=host,
)
ctx.organizations.append(
OrgContext(
org_name=installation.account.login,
client=client(
GitHubAppInstallationAuth(installation=org_installation)
GitHubAppInstallationAuth(
installation=org_installation,
api_uri=host,
)
),
enterprise_name=credentials.enterprise_name,
github_deployment_id=github_deployment_id,
Expand All @@ -196,9 +201,13 @@ def token_client(token: str) -> RESTClient:
installation_id=installation.id,
jwt_issuer=jwt_issuer,
private_key_path=credentials.key_path,
api_uri=host,
)
ctx.client = client(
GitHubAppInstallationAuth(installation=es_installation)
GitHubAppInstallationAuth(
installation=es_installation,
api_uri=host,
)
)

return (*enterprise_resources(ctx), *organization_resources(ctx))
Expand All @@ -213,11 +222,17 @@ def token_client(token: str) -> RESTClient:
installation_id=credentials.install_id,
jwt_issuer=credentials.client_id,
private_key_path=credentials.key_path,
api_uri=host,
)
ctx.organizations.append(
OrgContext(
org_name=credentials.org_name,
client=client(GitHubAppInstallationAuth(installation=org_installation)),
client=client(
GitHubAppInstallationAuth(
installation=org_installation,
api_uri=host,
)
),
github_deployment_id=github_deployment_id,
github_web_origin=github_web_origin,
)
Expand Down
28 changes: 27 additions & 1 deletion tests/test_app_auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from openhound_github import auth
from openhound_github.auth import (
AccountConfig,
GitHubAppInstallationAuth,
GithubSession,
InstallationResponse,
resolve_github_app_jwt_issuer,
Expand Down Expand Up @@ -97,14 +98,38 @@ def test_legacy_installation_response_does_not_require_client_id() -> None:
assert installation.app_id == 123456


def test_github_app_installation_auth_rejects_mismatched_api_origins() -> None:
installation = SimpleNamespace(
installation_id="12345",
api_uri="https://ghe.example/api/v3/",
)

with pytest.raises(ValueError, match="must match installation API URI origin"):
GitHubAppInstallationAuth(
installation=installation,
api_uri="https://api.github.com/",
)


def test_github_session_rejects_plaintext_api_uri() -> None:
with pytest.raises(ValueError, match="absolute HTTPS URL"):
GithubSession(
jwt_issuer="123456",
private_key_path="/tmp/github-app.pem",
api_uri="http://ghe.example/api/v3/",
)


def test_enterprise_source_reuses_selected_issuer_for_installation_tokens(
monkeypatch: pytest.MonkeyPatch,
) -> None:
source_module = importlib.import_module("openhound_github.source")
captured_issuers: list[str] = []

class FakeGithubApp:
def __init__(self, jwt_issuer: str, private_key_path: str) -> None:
def __init__(
self, jwt_issuer: str, private_key_path: str, api_uri: str
) -> None:
captured_issuers.append(jwt_issuer)
self.installations = (
SimpleNamespace(
Expand All @@ -125,6 +150,7 @@ def __init__(
installation_id: int,
jwt_issuer: str,
private_key_path: str,
api_uri: str,
) -> None:
captured_issuers.append(jwt_issuer)

Expand Down
Loading
Loading