llm-ready-data / app /services /google_oauth_service.py
Soumik Bose
optimization 404
bd469c1
Raw
History Blame Contribute Delete
13.4 kB
from __future__ import annotations
import asyncio
import time
from typing import Any, Dict, Optional
from urllib.parse import urlencode
import httpx
import jwt
from app.config import get_settings
from app.core.logger import get_logger
from app.core.thread_pool import run_in_executor
from app.models.schemas import GoogleOAuthUserInfo
from app.utils.http_utils import SharedAsyncClient
_logger = get_logger(__name__)
_settings = get_settings()
# Retryable transient HTTP status codes when talking to Google.
_RETRYABLE_STATUS = frozenset({429, 500, 502, 503, 504})
class GoogleOAuthError(Exception):
"""Raised for upstream Google OAuth failures that map to a client-facing error."""
def __init__(self, message: str, status_code: int = 400) -> None:
super().__init__(message)
self.message = message
self.status_code = status_code
class GoogleJWKSCache:
"""Asynchronously fetches and caches Google's public JWKS keys.
Keys are cached for a short TTL to avoid hammering Google's cert endpoint,
and re-fetched on demand if the requested ``kid`` is missing (key rotation).
"""
def __init__(self) -> None:
self._keys: Dict[str, Dict[str, Any]] = {}
self._last_fetched: float = 0.0
self._lock = asyncio.Lock()
def _is_expired(self) -> bool:
return (
not self._keys
or time.time() - self._last_fetched >= _settings.google_oauth_jwks_ttl_seconds
)
async def get_public_key(self, client: httpx.AsyncClient, kid: str) -> Any:
"""Return the RSA public key for the given key id (``kid``)."""
if self._is_expired():
await self._refresh(client)
jwk = self._keys.get(kid)
if jwk is None:
# Key may have rotated; force a refresh before giving up.
await self._refresh(client)
jwk = self._keys.get(kid)
if jwk is None:
raise GoogleOAuthError(
f"Google signing key not found for kid '{kid}'", status_code=401
)
return jwt.algorithms.RSAAlgorithm.from_jwk(jwk)
async def _refresh(self, client: httpx.AsyncClient) -> None:
async with self._lock:
if not self._is_expired():
return
try:
response = await client.get(
_settings.google_oauth_jwks_url, timeout=10.0
)
response.raise_for_status()
data = response.json()
self._keys = {
key["kid"]: key
for key in data.get("keys", [])
if isinstance(key, dict) and "kid" in key
}
self._last_fetched = time.time()
_logger.info("Google JWKS refreshed (%d keys)", len(self._keys))
except httpx.HTTPError as exc:
_logger.error("Failed to refresh Google JWKS: %s", exc)
raise GoogleOAuthError(
"Unable to fetch Google signing keys", status_code=502
) from exc
class GoogleOAuthService:
"""Abstraction over the Google OAuth 2.0 / OpenID Connect flow."""
def __init__(self) -> None:
self._http = SharedAsyncClient(
timeout=httpx.Timeout(_settings.google_oauth_timeout),
follow_redirects=False,
)
self._jwks = GoogleJWKSCache()
# ------------------------------------------------------------------
# HTTP client management
# ------------------------------------------------------------------
async def _get_client(self) -> httpx.AsyncClient:
return await self._http.get()
async def close(self) -> None:
await self._http.close()
# ------------------------------------------------------------------
# Authorization URL (Step 1)
# ------------------------------------------------------------------
def build_auth_url(
self,
*,
client_id: str,
redirect_uri: str,
state: str,
scope: str,
prompt: Optional[str] = None,
access_type: Optional[str] = None,
login_hint: Optional[str] = None,
include_granted_scopes: bool = False,
) -> str:
params: Dict[str, Any] = {
"client_id": client_id,
"redirect_uri": redirect_uri,
"response_type": "code",
"scope": scope,
"state": state,
}
if prompt:
params["prompt"] = prompt
if access_type:
params["access_type"] = access_type
if login_hint:
params["login_hint"] = login_hint
if include_granted_scopes:
params["include_granted_scopes"] = "true"
return f"{_settings.google_oauth_auth_url}?{urlencode(params)}"
# ------------------------------------------------------------------
# Token endpoint helpers
# ------------------------------------------------------------------
async def _post_token(
self, data: Dict[str, str], *, context: str
) -> Dict[str, Any]:
client = await self._get_client()
last_error: Optional[str] = None
for attempt in range(1 + _settings.google_oauth_max_retries):
try:
response = await client.post(_settings.google_oauth_token_url, data=data)
if response.status_code == 200:
return response.json()
if response.status_code in _RETRYABLE_STATUS:
last_error = (
f"Google OAuth upstream error HTTP {response.status_code}"
)
_logger.warning(
"%s: transient HTTP %s (attempt %d/%d)",
context,
response.status_code,
attempt + 1,
1 + _settings.google_oauth_max_retries,
)
else:
return self._parse_token_error(response, context=context)
except httpx.TimeoutException:
last_error = "Google OAuth request timed out"
_logger.warning(
"%s: timeout (attempt %d/%d)",
context,
attempt + 1,
1 + _settings.google_oauth_max_retries,
)
except httpx.RequestError as exc:
last_error = f"Google OAuth request failed: {exc}"
_logger.warning(
"%s: %s (attempt %d/%d)",
context,
last_error,
attempt + 1,
1 + _settings.google_oauth_max_retries,
)
if attempt < _settings.google_oauth_max_retries:
await asyncio.sleep(2 ** attempt)
raise GoogleOAuthError(last_error or "Unknown Google OAuth error", status_code=502)
@staticmethod
def _parse_token_error(
response: httpx.Response, *, context: str
) -> Dict[str, Any]:
try:
body = response.json()
error: str = body.get("error", "") or ""
error_description: str = body.get("error_description", "") or ""
except Exception:
error = ""
error_description = ""
body = {}
_logger.warning("%s: Google token error '%s': %s", context, error, error_description)
if error == "invalid_grant":
message = error_description or (
"The provided code or refresh token is invalid, expired, or has been revoked."
)
raise GoogleOAuthError(message, status_code=401)
if error == "invalid_client":
raise GoogleOAuthError(
error_description or "Invalid client_id or client_secret.", status_code=401
)
if error == "invalid_request":
raise GoogleOAuthError(
error_description or "Malformed OAuth request.", status_code=400
)
if error == "access_denied":
raise GoogleOAuthError(
error_description or "Access denied by the user.", status_code=403
)
# Google does not return a structured error body for some failures.
if response.status_code >= 500:
raise GoogleOAuthError(
"Google OAuth service is temporarily unavailable.", status_code=502
)
if response.status_code == 429:
raise GoogleOAuthError(
"Google OAuth rate limit exceeded. Please retry later.", status_code=429
)
message = error_description or f"Google OAuth error (HTTP {response.status_code})."
return {"error": True, "message": message, "raw": body}
# ------------------------------------------------------------------
# Code exchange (Step 2)
# ------------------------------------------------------------------
async def exchange_code(
self,
*,
client_id: str,
client_secret: str,
code: str,
redirect_uri: str,
) -> Dict[str, Any]:
data = {
"client_id": client_id,
"client_secret": client_secret,
"code": code,
"grant_type": "authorization_code",
"redirect_uri": redirect_uri,
}
tokens = await self._post_token(data, context="authorization code exchange")
self._raise_for_token_error(tokens)
return tokens
async def refresh_access_token(
self,
*,
client_id: str,
client_secret: str,
refresh_token: str,
) -> Dict[str, Any]:
data = {
"client_id": client_id,
"client_secret": client_secret,
"refresh_token": refresh_token,
"grant_type": "refresh_token",
}
tokens = await self._post_token(data, context="refresh token exchange")
self._raise_for_token_error(tokens)
return tokens
@staticmethod
def _raise_for_token_error(tokens: Dict[str, Any]) -> None:
if tokens.get("error"):
raise GoogleOAuthError(tokens.get("message", "Unknown OAuth error"), status_code=400)
if not tokens.get("access_token"):
raise GoogleOAuthError(
"Google did not return an access token.", status_code=502
)
# ------------------------------------------------------------------
# ID token verification
# ------------------------------------------------------------------
async def verify_id_token(
self, id_token: str, client_id: str
) -> GoogleOAuthUserInfo:
client = await self._get_client()
try:
unverified = jwt.get_unverified_header(id_token)
kid = unverified.get("kid")
alg = unverified.get("alg")
except jwt.DecodeError as exc:
raise GoogleOAuthError("Malformed ID token.", status_code=400) from exc
if not kid or alg != "RS256":
raise GoogleOAuthError(
"ID token does not use an RS256 signature.", status_code=401
)
try:
public_key = await self._jwks.get_public_key(client, kid)
except GoogleOAuthError:
raise
try:
payload = await run_in_executor(
jwt.decode,
id_token,
key=public_key,
algorithms=["RS256"],
audience=client_id,
issuer=["accounts.google.com", "https://accounts.google.com"],
options={
"verify_exp": True,
"verify_iat": True,
"verify_aud": True,
"verify_iss": True,
"require": ["exp", "iat", "sub", "email"],
},
)
except jwt.ExpiredSignatureError as exc:
raise GoogleOAuthError("ID token has expired.", status_code=401) from exc
except jwt.InvalidAudienceError as exc:
raise GoogleOAuthError(
"ID token audience does not match the provided client_id.", status_code=401
) from exc
except jwt.InvalidIssuerError as exc:
raise GoogleOAuthError(
"ID token issuer is not Google.", status_code=401
) from exc
except jwt.PyJWTError as exc:
_logger.warning("ID token verification failed: %s", exc)
raise GoogleOAuthError("Invalid ID token.", status_code=401) from exc
return self._user_info_from_claims(payload)
# ------------------------------------------------------------------
# Userinfo (fallback profile source)
# ------------------------------------------------------------------
@staticmethod
def _user_info_from_claims(payload: Dict[str, Any]) -> GoogleOAuthUserInfo:
return GoogleOAuthUserInfo(
sub=payload.get("sub", ""),
email=payload.get("email", ""),
email_verified=bool(payload.get("email_verified", False)),
name=payload.get("name"),
given_name=payload.get("given_name"),
family_name=payload.get("family_name"),
picture=payload.get("picture"),
locale=payload.get("locale"),
)