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),
)