oidc.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261
  1. """OIDC Authorization Code + PKCE protocol helpers with strict token validation."""
  2. from __future__ import annotations
  3. import base64
  4. import hashlib
  5. import ipaddress
  6. import json
  7. import secrets
  8. import socket
  9. from collections.abc import Callable, Mapping
  10. from dataclasses import dataclass
  11. from datetime import UTC, datetime, timedelta
  12. from typing import Any
  13. from urllib.parse import urlencode, urlsplit
  14. import jwt
  15. import requests
  16. import urllib3
  17. from app.core.system.enterprise_identity import (
  18. IdentityPolicyError,
  19. IdentityUpstreamError,
  20. IdpConfig,
  21. )
  22. def _digest(value: str) -> str:
  23. return hashlib.sha256(value.encode()).hexdigest()
  24. @dataclass(frozen=True)
  25. class AuthorizationStart:
  26. url: str
  27. state: str
  28. nonce: str
  29. code_verifier: str
  30. code_challenge: str
  31. code_challenge_method: str = "S256"
  32. class OidcClient:
  33. def __init__(self, repository: Any, *, clock: Callable[[], datetime] | None = None,
  34. flow_lifetime: timedelta = timedelta(minutes=5)) -> None:
  35. self.repository = repository
  36. self.clock = clock or (lambda: datetime.now(UTC))
  37. self.flow_lifetime = flow_lifetime
  38. def begin(self, config: IdpConfig, redirect_uri: str) -> AuthorizationStart:
  39. config.validate()
  40. if redirect_uri not in config.redirect_uris:
  41. raise IdentityPolicyError("redirect URI is not allowlisted")
  42. state, nonce, verifier = (secrets.token_urlsafe(32) for _ in range(3))
  43. challenge = base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode()
  44. self.repository.put_flow({"state_hash": _digest(state), "provider_uid": config.provider_uid,
  45. "provider_version": config.version, "nonce_hash": _digest(nonce),
  46. "verifier_hash": _digest(verifier), "redirect_uri": redirect_uri,
  47. "expires_at": self.clock() + self.flow_lifetime, "consumed_at": None})
  48. params = {"response_type": "code", "client_id": config.client_id, "redirect_uri": redirect_uri,
  49. "scope": "openid profile", "state": state, "nonce": nonce,
  50. "code_challenge": challenge, "code_challenge_method": "S256"}
  51. return AuthorizationStart(f"{config.authorization_endpoint}?{urlencode(params)}", state, nonce, verifier, challenge)
  52. def verify_callback(self, config: IdpConfig, *, state: str, redirect_uri: str, code_verifier: str,
  53. id_token: str, jwks: Mapping[str, Any], commit: bool = True) -> dict[str, Any]:
  54. flow = self.repository.consume_flow(_digest(state), commit=commit)
  55. if not flow:
  56. raise IdentityPolicyError("state is invalid or already consumed")
  57. if self.clock() >= flow["expires_at"] or flow["provider_uid"] != config.provider_uid or flow["provider_version"] != config.version:
  58. raise IdentityPolicyError("state is expired or bound to another provider")
  59. if redirect_uri != flow["redirect_uri"] or redirect_uri not in config.redirect_uris or _digest(code_verifier) != flow["verifier_hash"]:
  60. raise IdentityPolicyError("callback binding or PKCE verifier is invalid")
  61. try:
  62. header = jwt.get_unverified_header(id_token)
  63. algorithm = header.get("alg")
  64. if algorithm not in config.algorithms or algorithm == "none":
  65. raise IdentityPolicyError("ID token algorithm is not allowed")
  66. candidates = [key for key in jwks.get("keys", ()) if key.get("kid") == header.get("kid") and key.get("alg", algorithm) == algorithm]
  67. if len(candidates) != 1:
  68. raise IdentityPolicyError("ID token signing key is ambiguous or missing")
  69. public_key = jwt.algorithms.get_default_algorithms()[algorithm].from_jwk(candidates[0])
  70. claims = jwt.decode(id_token, public_key, algorithms=[algorithm], audience=config.client_id,
  71. issuer=config.issuer, options={"verify_exp": False,
  72. "require": ["iss", "aud", "sub", "nonce", "iat", "exp"]})
  73. except IdentityPolicyError:
  74. raise
  75. except jwt.PyJWTError as exc:
  76. raise IdentityPolicyError("ID token validation failed") from exc
  77. if _digest(str(claims["nonce"])) != flow["nonce_hash"]:
  78. raise IdentityPolicyError("ID token nonce mismatch")
  79. now = int(self.clock().timestamp())
  80. if int(claims["exp"]) <= now:
  81. raise IdentityPolicyError("ID token expired")
  82. if int(claims["iat"]) > now + 30:
  83. raise IdentityPolicyError("ID token issued in the future")
  84. return dict(claims)
  85. class OidcTransport:
  86. """Bounded HTTPS transport whose production connection is pinned to validated IPs."""
  87. def __init__(self, http: Any | None = None, *, connect_timeout: float = 3.0,
  88. read_timeout: float = 7.0, max_response_bytes: int = 1024 * 1024,
  89. resolver: Callable[..., Any] = socket.getaddrinfo,
  90. pool_factory: Callable[..., Any] = urllib3.HTTPSConnectionPool,
  91. allowed_hosts: set[str] | None = None) -> None:
  92. self.http = http
  93. self.connect_timeout = connect_timeout
  94. self.read_timeout = read_timeout
  95. self.timeout = (connect_timeout, read_timeout)
  96. self.max_response_bytes = max_response_bytes
  97. self.resolver = resolver
  98. self.pool_factory = pool_factory
  99. self.allowed_hosts = {host.lower() for host in allowed_hosts} if allowed_hosts else None
  100. def _validate_destination(self, url: str) -> tuple[Any, tuple[str, ...]]:
  101. parsed = urlsplit(url)
  102. if parsed.scheme != "https" or not parsed.hostname:
  103. raise IdentityPolicyError("OIDC destination must be an absolute HTTPS URL")
  104. hostname = parsed.hostname.lower()
  105. if self.allowed_hosts is not None and hostname not in self.allowed_hosts:
  106. raise IdentityPolicyError("enterprise identity provider host is not server-allowlisted")
  107. try:
  108. addresses = {
  109. result[4][0]
  110. for result in self.resolver(hostname, parsed.port or 443, type=socket.SOCK_STREAM)
  111. }
  112. except (OSError, socket.gaierror) as exc:
  113. raise IdentityUpstreamError("enterprise identity provider DNS resolution failed") from exc
  114. if not addresses:
  115. raise IdentityUpstreamError("enterprise identity provider DNS returned no addresses")
  116. try:
  117. parsed_addresses = [ipaddress.ip_address(address) for address in addresses]
  118. except ValueError as exc:
  119. raise IdentityUpstreamError("enterprise identity provider DNS response is invalid") from exc
  120. if any(not address.is_global for address in parsed_addresses):
  121. raise IdentityPolicyError("enterprise identity provider resolved to a non-public address")
  122. return parsed, tuple(sorted(str(address) for address in parsed_addresses))
  123. def _decode_json_response(self, *, status: int, headers: Mapping[str, Any], chunks: Any) -> Mapping[str, Any]:
  124. if 300 <= status < 400:
  125. raise IdentityUpstreamError("enterprise identity provider redirect is forbidden")
  126. if status < 200 or status >= 300:
  127. raise IdentityUpstreamError("enterprise identity provider returned an HTTP error")
  128. content_type = str(headers.get("Content-Type", "")).lower()
  129. if "application/json" not in content_type:
  130. raise IdentityUpstreamError("enterprise identity provider response must be JSON")
  131. content_length = headers.get("Content-Length")
  132. if content_length:
  133. try:
  134. parsed_length = int(content_length)
  135. except (TypeError, ValueError) as exc:
  136. raise IdentityUpstreamError("enterprise identity provider response length is invalid") from exc
  137. if parsed_length < 0:
  138. raise IdentityUpstreamError("enterprise identity provider response length is invalid")
  139. if parsed_length > self.max_response_bytes:
  140. raise IdentityUpstreamError("enterprise identity provider response is too large")
  141. body = bytearray()
  142. try:
  143. for chunk in chunks:
  144. if chunk:
  145. body.extend(chunk)
  146. if len(body) > self.max_response_bytes:
  147. raise IdentityUpstreamError("enterprise identity provider response is too large")
  148. except (requests.RequestException, urllib3.exceptions.HTTPError, OSError) as exc:
  149. raise IdentityUpstreamError("enterprise identity provider response read failed") from exc
  150. try:
  151. payload = json.loads(body.decode("utf-8"))
  152. except (UnicodeDecodeError, json.JSONDecodeError) as exc:
  153. raise IdentityUpstreamError("enterprise identity provider returned invalid JSON") from exc
  154. if not isinstance(payload, Mapping):
  155. raise IdentityUpstreamError("enterprise identity provider JSON object is required")
  156. return payload
  157. def _pinned_request_json(self, method: str, parsed: Any, addresses: tuple[str, ...],
  158. data: Mapping[str, Any] | None) -> Mapping[str, Any]:
  159. hostname = parsed.hostname.lower()
  160. port = parsed.port or 443
  161. request_target = parsed.path or "/"
  162. if parsed.query:
  163. request_target += f"?{parsed.query}"
  164. host_header = hostname if port == 443 else f"{hostname}:{port}"
  165. body = urlencode(data).encode("utf-8") if data is not None else None
  166. headers = {"Host": host_header, "Accept": "application/json"}
  167. if body is not None:
  168. headers["Content-Type"] = "application/x-www-form-urlencoded"
  169. last_error: Exception | None = None
  170. for address in addresses:
  171. pool = self.pool_factory(
  172. host=address,
  173. port=port,
  174. server_hostname=hostname,
  175. assert_hostname=hostname,
  176. cert_reqs="CERT_REQUIRED",
  177. ca_certs=requests.certs.where(),
  178. timeout=urllib3.Timeout(connect=self.connect_timeout, read=self.read_timeout),
  179. retries=False,
  180. maxsize=1,
  181. block=True,
  182. )
  183. try:
  184. response = pool.urlopen(
  185. method.upper(), request_target, body=body, headers=headers,
  186. redirect=False, retries=False, assert_same_host=False,
  187. preload_content=False, decode_content=True,
  188. timeout=urllib3.Timeout(connect=self.connect_timeout, read=self.read_timeout),
  189. )
  190. except (urllib3.exceptions.HTTPError, OSError) as exc:
  191. last_error = exc
  192. pool.close()
  193. continue
  194. try:
  195. return self._decode_json_response(
  196. status=int(response.status), headers=response.headers,
  197. chunks=response.stream(amt=64 * 1024, decode_content=True),
  198. )
  199. finally:
  200. response.release_conn()
  201. pool.close()
  202. raise IdentityUpstreamError("enterprise identity provider connection failed") from last_error
  203. def _request_json(self, method: str, url: str, **kwargs: Any) -> Mapping[str, Any]:
  204. parsed, addresses = self._validate_destination(url)
  205. if self.http is None:
  206. return self._pinned_request_json(method, parsed, addresses, kwargs.get("data"))
  207. try:
  208. response = getattr(self.http, method)(url, timeout=self.timeout, allow_redirects=False,
  209. stream=True, **kwargs)
  210. except requests.RequestException as exc:
  211. raise IdentityUpstreamError("enterprise identity provider request failed") from exc
  212. try:
  213. try:
  214. status = int(getattr(response, "status_code", 0))
  215. except (TypeError, ValueError) as exc:
  216. raise IdentityUpstreamError("enterprise identity provider HTTP status is invalid") from exc
  217. return self._decode_json_response(
  218. status=status, headers=getattr(response, "headers", {}),
  219. chunks=response.iter_content(chunk_size=64 * 1024),
  220. )
  221. finally:
  222. close = getattr(response, "close", None)
  223. if callable(close):
  224. close()
  225. def exchange_code(self, config: IdpConfig, *, code: str, redirect_uri: str, verifier: str) -> str:
  226. payload = self._request_json(
  227. "post", config.token_endpoint,
  228. data={"grant_type": "authorization_code", "code": code, "redirect_uri": redirect_uri,
  229. "client_id": config.client_id, "client_secret": config.resolve_secret(),
  230. "code_verifier": verifier},
  231. )
  232. id_token = payload.get("id_token")
  233. if not isinstance(id_token, str) or not id_token:
  234. raise IdentityPolicyError("token endpoint omitted ID token")
  235. return id_token
  236. def fetch_jwks(self, config: IdpConfig) -> Mapping[str, Any]:
  237. payload = self._request_json("get", config.jwks_uri)
  238. if not isinstance(payload, Mapping) or not isinstance(payload.get("keys"), list):
  239. raise IdentityPolicyError("JWKS response is invalid")
  240. return payload