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)