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