Source code for fsh_lib.oauth.oidc

from __future__ import annotations

from typing import TYPE_CHECKING, Any, Self

import anyio
import jwt
from jwt import PyJWKClient
from pydantic import BaseModel, ConfigDict, Field

from fsh_lib.oauth.client import OAuthError

if TYPE_CHECKING:
    from collections.abc import Sequence

    from fsh_lib.oauth.providers import ProviderConfig

DEFAULT_ALGORITHMS: tuple[str, ...] = ("RS256",)


[docs] class OidcIdentity(BaseModel): model_config = ConfigDict(frozen=True, extra="ignore") sub: str email: str | None = None email_verified: bool = False name: str | None = None raw: dict[str, Any] = Field(default_factory=dict)
class IdTokenVerifier: def __init__( self, *, jwks_uri: str, audience: str, issuer: str | None = None, algorithms: Sequence[str] = DEFAULT_ALGORITHMS, leeway: float = 0.0, verify_email: bool = False, ) -> None: self._audience = audience self._issuer = issuer self._algorithms = list(algorithms) self._leeway = leeway self._verify_email = verify_email self._jwk_client = PyJWKClient(jwks_uri) @classmethod def from_provider( cls, provider: ProviderConfig, *, client_id: str, algorithms: Sequence[str] = DEFAULT_ALGORITHMS, leeway: float = 0.0, ) -> Self: if provider.jwks_uri is None: msg = "provider has no jwks_uri; cannot verify ID tokens" raise OAuthError(msg) return cls( jwks_uri=provider.jwks_uri, audience=client_id, issuer=provider.issuer, algorithms=algorithms, leeway=leeway, verify_email=provider.verify_email, ) async def verify( self, id_token: str | None, *, nonce: str | None = None, ) -> OidcIdentity: if not id_token: msg = "no id_token to verify (was 'openid' among the scopes?)" raise OAuthError(msg) try: signing_key = await anyio.to_thread.run_sync( self._jwk_client.get_signing_key_from_jwt, id_token, ) claims: dict[str, Any] = jwt.decode( id_token, signing_key.key, algorithms=self._algorithms, audience=self._audience, issuer=self._issuer, leeway=self._leeway, ) except jwt.PyJWTError as exc: msg = f"ID token verification failed: {exc}" raise OAuthError(msg) from exc if "exp" not in claims: msg = "ID token has no exp claim" raise OAuthError(msg) if nonce is not None and claims.get("nonce") != nonce: msg = "ID token nonce mismatch" raise OAuthError(msg) claimed = claims.get("email_verified") verified = ( claimed if self._verify_email and isinstance(claimed, bool) else not self._verify_email ) return OidcIdentity.model_validate( {**claims, "email_verified": verified, "raw": claims}, )