Source code for fsh_lib.oauth.client

from __future__ import annotations

import base64
import hashlib
import os
import secrets
from datetime import datetime, timedelta
from typing import TYPE_CHECKING, Any, Self
from urllib.parse import urlencode, urlsplit, urlunsplit

import httpx
from pydantic import BaseModel, ConfigDict, Field, ValidationError

if TYPE_CHECKING:
    from collections.abc import Mapping, Sequence

    from fsh_lib.oauth.providers import ClientCredentials, ProviderConfig

_ENTROPY_BYTES = 32
_DEFAULT_TIMEOUT = 10.0


[docs] class OAuthError(RuntimeError): pass
[docs] class TokenResponse(BaseModel): model_config = ConfigDict(frozen=True, extra="ignore") access_token: str token_type: str = "Bearer" # noqa: S105 -- a token type, not a secret expires_in: int | None = None refresh_token: str | None = None scope: str | None = None id_token: str | None = None raw: dict[str, Any] = Field(default_factory=dict) def expires_at( self, now: datetime, ) -> datetime | None: if self.expires_in is None: return None return now + timedelta(seconds=self.expires_in)
[docs] class AuthorizationState(BaseModel): model_config = ConfigDict(frozen=True) state: str code_verifier: str | None = None nonce: str | None = None
def _b64url(raw: bytes) -> str: return base64.urlsafe_b64encode(raw).rstrip(b"=").decode("ascii") def generate_state() -> str: return _b64url(secrets.token_bytes(_ENTROPY_BYTES)) def generate_pkce() -> tuple[str, str]: verifier = _b64url(secrets.token_bytes(_ENTROPY_BYTES)) challenge = _b64url(hashlib.sha256(verifier.encode("ascii")).digest()) return verifier, challenge class OAuthClient: def __init__( self, provider: ProviderConfig, credentials: ClientCredentials, *, client: httpx.AsyncClient | None = None, timeout: float = _DEFAULT_TIMEOUT, ) -> None: self._provider = provider self._credentials = credentials self._owns_client = client is None self._client = client or httpx.AsyncClient(timeout=timeout) async def __aenter__(self) -> Self: return self async def __aexit__(self, *_exc: object) -> None: await self.aclose() def _construct_redirect_uri(self) -> str: base = os.environ.get("API_URL", "") return base + self._credentials.redirect_path async def aclose(self) -> None: if self._owns_client: await self._client.aclose() def authorization_url( self, *, state: str, scopes: Sequence[str] | None = None, code_challenge: str | None = None, nonce: str | None = None, extra_params: Mapping[str, str] | None = None, ) -> str: chosen = ( tuple(scopes) if scopes is not None else (self._provider.default_scopes) ) params: dict[str, str] = { "response_type": "code", "client_id": self._credentials.client_id, "redirect_uri": self._construct_redirect_uri(), "scope": " ".join(chosen), "state": state, } if code_challenge is not None: params["code_challenge"] = code_challenge params["code_challenge_method"] = "S256" if nonce is not None: params["nonce"] = nonce if extra_params: params.update(extra_params) return _with_query(self._provider.authorization_endpoint, params) def begin( self, *, scopes: Sequence[str] | None = None, use_pkce: bool = True, use_nonce: bool = False, extra_params: Mapping[str, str] | None = None, ) -> tuple[str, AuthorizationState]: verifier, challenge = generate_pkce() if use_pkce else (None, None) nonce = generate_state() if use_nonce else None state = generate_state() url = self.authorization_url( state=state, scopes=scopes, code_challenge=challenge, nonce=nonce, extra_params=extra_params, ) return url, AuthorizationState( state=state, code_verifier=verifier, nonce=nonce, ) async def exchange_code( self, code: str, *, code_verifier: str | None = None, redirect_uri: str | None = None, ) -> TokenResponse: data = { "grant_type": "authorization_code", "code": code, "redirect_uri": redirect_uri or self._construct_redirect_uri(), } if code_verifier is not None: data["code_verifier"] = code_verifier return await self._token_request(data) async def refresh(self, refresh_token: str) -> TokenResponse: return await self._token_request( { "grant_type": "refresh_token", "refresh_token": refresh_token, }, ) async def client_credentials( self, *, scopes: Sequence[str] | None = None, ) -> TokenResponse: chosen = tuple(scopes) if scopes else (".default",) return await self._token_request( { "grant_type": "client_credentials", "scope": " ".join(chosen), }, ) def for_tenant(self, tenant: str) -> OAuthClient: return OAuthClient( self._provider.for_tenant(tenant), self._credentials, client=self._client, ) def admin_consent_url(self, *, state: str) -> str: params = { "client_id": self._credentials.client_id, "redirect_uri": self._construct_redirect_uri(), "state": state, } return _with_query(self._provider.admin_consent_endpoint, params) async def fetch_userinfo(self, access_token: str) -> dict[str, Any]: if self._provider.userinfo_endpoint is None: msg = "provider has no userinfo_endpoint configured" raise OAuthError(msg) try: response = await self._client.get( self._provider.userinfo_endpoint, headers={"Authorization": f"Bearer {access_token}"}, ) response.raise_for_status() claims: dict[str, Any] = response.json() except httpx.HTTPError as exc: msg = f"userinfo request failed: {exc}" raise OAuthError(msg) from exc except ValueError as exc: msg = "userinfo response is not valid JSON" raise OAuthError(msg) from exc return claims async def _token_request( self, data: dict[str, str], ) -> TokenResponse: body = dict(data) secret = self._client_secret() extra: dict[str, Any] = {} auth_method = self._provider.token_endpoint_auth_method if auth_method == "client_secret_basic": if secret is None: msg = "client_secret is required for client_secret_basic" raise OAuthError(msg) extra["auth"] = (self._credentials.client_id, secret) else: body["client_id"] = self._credentials.client_id if secret is not None: body["client_secret"] = secret try: response = await self._client.post( self._provider.token_endpoint, data=body, headers={"Accept": "application/json"}, **extra, ) response.raise_for_status() payload = response.json() except httpx.HTTPError as exc: msg = f"token request failed: {exc}" raise OAuthError(msg) from exc except ValueError as exc: msg = "token response is not valid JSON" raise OAuthError(msg) from exc if not isinstance(payload, dict): msg = "token response is not a JSON object" raise OAuthError(msg) try: return TokenResponse.model_validate({**payload, "raw": payload}) except ValidationError as exc: msg = f"malformed token response: {exc}" raise OAuthError(msg) from exc def _client_secret(self) -> str | None: env = self._credentials.client_secret_env if env is None: return None return os.environ.get(env) def _with_query(url: str, params: Mapping[str, str]) -> str: parts = urlsplit(url) existing = parts.query encoded = urlencode(params) query = f"{existing}&{encoded}" if existing else encoded return urlunsplit( (parts.scheme, parts.netloc, parts.path, query, parts.fragment), )