Source code for fsh_lib.oauth.registry

from __future__ import annotations

from dataclasses import dataclass
from typing import TYPE_CHECKING

from fastapi import HTTPException, status

from fsh_lib.oauth.client import OAuthClient, OAuthError
from fsh_lib.oauth.oidc import DEFAULT_ALGORITHMS, IdTokenVerifier
from fsh_lib.oauth.providers import OAuthConfigError

if TYPE_CHECKING:
    from collections.abc import Iterator, Mapping, Sequence
    from typing import Self

    import httpx

    from fsh_lib.oauth.providers import ClientCredentials, ProviderConfig


[docs] @dataclass(frozen=True) class Provider: client: OAuthClient verifier: IdTokenVerifier | None = None app_only: bool = False app_scopes: tuple[str, ...] = () def require_verifier(self) -> IdTokenVerifier: if self.verifier is None: msg = "provider has no ID-token verifier (not an OIDC provider?)" raise OAuthError(msg) return self.verifier def app_only_client(self, tenant: str) -> OAuthClient: if not self.app_only: msg = "provider is not configured for app-only access" raise OAuthError(msg) return self.client.for_tenant(tenant)
class ProviderRegistry: def __init__(self, providers: Mapping[str, Provider] | None = None) -> None: self._providers: dict[str, Provider] = dict(providers or {}) def register( self, slug: str, config: ProviderConfig, credentials: ClientCredentials, *, client: httpx.AsyncClient | None = None, verify_id_tokens: bool = True, algorithms: Sequence[str] = DEFAULT_ALGORITHMS, leeway: float = 0.0, app_only: bool = False, app_scopes: Sequence[str] = (), ) -> Provider: self._check_free(slug) oauth = OAuthClient(config, credentials, client=client) verifier = ( IdTokenVerifier.from_provider( config, client_id=credentials.client_id, algorithms=algorithms, leeway=leeway, ) if verify_id_tokens and config.jwks_uri is not None else None ) provider = Provider( client=oauth, verifier=verifier, app_only=app_only, app_scopes=tuple(app_scopes), ) self._providers[slug] = provider return provider def add(self, slug: str, provider: Provider) -> None: self._check_free(slug) self._providers[slug] = provider def _check_free(self, slug: str) -> None: if slug in self._providers: msg = f"OAuth provider slug already registered: {slug!r}" raise OAuthConfigError(msg) def get(self, slug: str) -> Provider: provider = self._providers.get(slug) if provider is None: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail=f"Unknown OAuth provider: {slug}", ) return provider async def aclose_all(self) -> None: for provider in self._providers.values(): await provider.client.aclose() async def __aenter__(self) -> Self: return self async def __aexit__(self, *_exc: object) -> None: await self.aclose_all() def slugs(self) -> list[str]: return list(self._providers) def __getitem__(self, slug: str) -> Provider: return self._providers[slug] def __contains__(self, slug: object) -> bool: return slug in self._providers def __iter__(self) -> Iterator[str]: return iter(self._providers) def __len__(self) -> int: return len(self._providers)