Coverage for src/lexigram/auth/authn/oauth2.py: 87%
148 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 00:58 +0800
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 00:58 +0800
1"""OAuth2 authentication flows.
3This module provides OAuth2 authentication support for Lexigram, enabling
4social login and OAuth2-based authentication with providers like Google,
5GitHub, Facebook, and other OAuth2-compatible identity providers.
7Example:
8 Setting up OAuth2 authentication::
10 from lexigram.auth.authn.oauth2 import OAuth2Authenticator
12 authenticator = OAuth2Authenticator(
13 provider="google",
14 client_id="your-client-id",
15 client_secret="your-client-secret",
16 redirect_uri="https://yourapp.com/auth/callback"
17 )
19 # Get authorization URL
20 auth_url = await authenticator.get_authorization_url()
22 # Exchange code for tokens
23 user_info = await authenticator.get_user_info(code)
24"""
26from __future__ import annotations
28from typing import TYPE_CHECKING, Any, cast
30from lexigram.auth.authn.oauth2_session import (
31 LexigramConnectResponse as LexigramConnectResponse,
32)
33from lexigram.auth.authn.oauth2_session import (
34 LexigramConnectSession as LexigramConnectSession,
35)
36from lexigram.auth.exceptions import OAuth2Error
37from lexigram.auth.types import OAuth2UserInfo
38from lexigram.logging import get_logger
40# Initialize logger early so runtime checks can log errors
41logger = get_logger(__name__)
43try:
44 from authlib.integrations.base_client import (
45 OAuthError,
46 )
47 from authlib.integrations.httpx_client import (
48 AsyncOAuth2Client,
49 )
51 HAS_AUTHLIB = True
52except ImportError:
53 # Provide placeholders so tests can patch these symbols even when authlib
54 # isn't installed in the test environment.
55 HAS_AUTHLIB = False
56 OAuthError = Exception
57 AsyncOAuth2Client = None
60if TYPE_CHECKING:
61 from lexigram.auth.models.user import User
62 from lexigram.contracts.web import HTTPClientProtocol
65class OAuth2IdentityProvider:
66 """OAuth2 identity provider configuration."""
68 def __init__(
69 self,
70 name: str,
71 client_id: str,
72 client_secret: str,
73 authorize_url: str,
74 access_token_url: str,
75 userinfo_url: str,
76 scope: str = "openid email profile",
77 redirect_uri: str | None = None,
78 require_pkce: bool = True,
79 ):
80 self.name = name
81 self.client_id = client_id
82 self.client_secret = client_secret
83 self.authorize_url = authorize_url
84 self.access_token_url = access_token_url
85 self.userinfo_url = userinfo_url
86 self.scope = scope
87 self.redirect_uri = redirect_uri
88 self.require_pkce = require_pkce
91class OAuth2Manager:
92 """OAuth2 authentication manager"""
94 def __init__(
95 self,
96 providers: dict[str, OAuth2IdentityProvider],
97 http_client: HTTPClientProtocol | None = None,
98 ) -> None:
99 # Allow construction even when authlib isn't installed so tests can
100 # assert that providers are configured. Actual network operations
101 # will raise if authlib is absent.
102 self.providers = providers
103 self._http_client = http_client
104 self._available = HAS_AUTHLIB
105 if not self._available:
106 logger.warning(
107 "authlib not installed; OAuth2Manager created in no-op mode for tests",
108 )
110 def __repr__(self) -> str:
111 """Return developer-friendly string representation."""
112 return f"OAuth2Manager(providers={list(self.providers)!r})"
114 async def get_authorization_url(
115 self,
116 provider_name: str,
117 state: str | None = None,
118 ) -> tuple[str, str, str]:
119 """Get OAuth2 authorization URL, PKCE code_verifier (S256), and state.
121 Returns a tuple of ``(authorization_url, code_verifier, state)``.
122 Callers **must** persist both ``code_verifier`` and ``state`` and
123 verify ``state`` matches on the callback to prevent CSRF attacks.
125 If *state* is not supplied a cryptographically random value is
126 generated automatically using :func:`secrets.token_urlsafe`."""
127 import base64
128 import hashlib
129 import secrets
131 if not self._available:
132 raise RuntimeError(
133 "authlib is required for OAuth2 operations; OAuth2Manager is in no-op mode",
134 )
136 provider = self.providers.get(provider_name)
137 if not provider:
138 raise ValueError(f"OAuth2 provider '{provider_name}' not configured")
140 # Auto-generate state when caller does not supply one
141 if state is None:
142 state = secrets.token_urlsafe(32)
143 logger.debug(
144 "OAuth2 state not provided; auto-generating secure random state",
145 provider=provider_name,
146 )
148 # Use custom session if http_client provided, otherwise use default httpx
149 session = None
150 if self._http_client:
151 session = LexigramConnectSession(self._http_client)
153 client = AsyncOAuth2Client(
154 client_id=provider.client_id,
155 client_secret=provider.client_secret,
156 redirect_uri=provider.redirect_uri,
157 session=session,
158 )
160 # Generate PKCE code verifier and challenge (S256)
161 code_verifier = secrets.token_urlsafe(64)
162 digest = hashlib.sha256(code_verifier.encode("ascii")).digest()
163 code_challenge = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
165 authorization_url, _ = client.create_authorization_url(
166 provider.authorize_url,
167 scope=provider.scope,
168 state=state,
169 code_challenge=code_challenge,
170 code_challenge_method="S256",
171 )
173 return str(authorization_url), code_verifier, state
175 async def exchange_code_for_token(
176 self,
177 provider_name: str,
178 code: str,
179 code_verifier: str | None = None,
180 ) -> dict[str, Any]:
181 """Exchange authorization code for access token, optionally with PKCE code_verifier"""
182 provider = self.providers.get(provider_name)
183 if not provider:
184 raise ValueError(f"OAuth2 provider '{provider_name}' not configured")
186 if provider.require_pkce and code_verifier is None:
187 raise OAuth2Error(
188 f"PKCE code_verifier is required for provider '{provider_name}'; "
189 "include the code_verifier returned by get_authorization_url()"
190 )
192 # Use custom session if http_client provided, otherwise use default httpx
193 session = None
194 if self._http_client:
195 session = LexigramConnectSession(self._http_client)
197 client = AsyncOAuth2Client(
198 client_id=provider.client_id,
199 client_secret=provider.client_secret,
200 redirect_uri=provider.redirect_uri,
201 session=session,
202 )
204 if not self._available:
205 raise RuntimeError(
206 "authlib is required for OAuth2 operations; OAuth2Manager is in no-op mode",
207 )
209 try:
210 fetch_kwargs = {"code": code}
211 if code_verifier:
212 fetch_kwargs["code_verifier"] = code_verifier
214 return cast(
215 "dict[str, Any]",
216 await client.fetch_token(provider.access_token_url, **fetch_kwargs),
217 )
218 except OAuthError as e:
219 logger.exception("OAuth2 token exchange failed")
220 raise ValueError(f"OAuth2 token exchange failed: {e!s}") from e
222 async def get_user_info(
223 self,
224 provider_name: str,
225 token: dict[str, Any],
226 ) -> dict[str, Any]:
227 """Get user information from OAuth2 provider"""
228 provider = self.providers.get(provider_name)
229 if not provider:
230 raise ValueError(f"OAuth2 provider '{provider_name}' not configured")
232 # Use custom session if http_client provided, otherwise use default httpx
233 session = None
234 if self._http_client:
235 session = LexigramConnectSession(self._http_client)
237 client = AsyncOAuth2Client(
238 client_id=provider.client_id,
239 client_secret=provider.client_secret,
240 token=token,
241 session=session,
242 )
244 if not self._available:
245 raise RuntimeError(
246 "authlib is required for OAuth2 operations; OAuth2Manager is in no-op mode",
247 )
249 try:
250 resp = await client.get(provider.userinfo_url)
251 resp.raise_for_status()
252 # Handle sync or async json() implementations (tests use mocks returning dict)
253 import asyncio
255 maybe_json = resp.json()
256 if asyncio.iscoroutine(maybe_json):
257 data = await maybe_json
258 else:
259 data = maybe_json
260 return cast("dict[str, Any]", data)
261 except (OAuthError, RuntimeError, OSError) as e:
262 logger.exception("Failed to get user info from OAuth2 provider")
263 raise ValueError(f"Failed to get user info: {e!s}") from e
264 except (KeyError, TypeError, ValueError) as e:
265 logger.exception("Unexpected error during OAuth2 user info retrieval")
266 raise ValueError(f"Unexpected error: {e!s}") from e
269class OAuth2AuthProvider:
270 """Complete OAuth2 provider with user provisioning."""
272 def __init__(
273 self,
274 oauth2_manager: OAuth2Manager,
275 user_store: Any,
276 oauth_identity_store: Any | None = None,
277 ) -> None:
278 self.oauth2_manager = oauth2_manager
279 self.user_store = user_store
280 self.oauth_identity_store = oauth_identity_store
282 async def authenticate_with_oauth2(
283 self,
284 provider_name: str,
285 code: str,
286 code_verifier: str | None = None,
287 ) -> User:
288 """Authenticate user via OAuth2 and provision if needed."""
290 # Exchange code for token (support PKCE via code_verifier)
291 if code_verifier is None:
292 token = await self.oauth2_manager.exchange_code_for_token(
293 provider_name,
294 code,
295 )
296 else:
297 token = await self.oauth2_manager.exchange_code_for_token(
298 provider_name,
299 code,
300 code_verifier=code_verifier,
301 )
303 # Get user info from provider
304 user_info = await self.oauth2_manager.get_user_info(provider_name, token)
306 # Map to OAuth2UserInfo
307 oauth_user = OAuth2UserInfo(
308 provider=provider_name,
309 provider_user_id=user_info["id"],
310 email=user_info.get("email"),
311 email_verified=bool(user_info.get("email_verified", False)),
312 username=(user_info.get("login") or user_info.get("username")),
313 # prefer login/username over display name so that the test above
314 # that expects ``testuser`` continues to pass
315 name=(
316 user_info.get("login")
317 or user_info.get("username")
318 or user_info.get("name")
319 or user_info.get("email")
320 ),
321 avatar_url=user_info.get("avatar_url"),
322 raw_data=user_info,
323 )
325 # Find or create user
326 return await self._find_or_create_oauth_user(oauth_user)
328 async def _find_or_create_oauth_user(self, oauth_user: OAuth2UserInfo) -> User:
329 """Find existing user or provision new one."""
331 # Try to find by email (only when the IdP verified the address)
332 if oauth_user.email:
333 if oauth_user.email_verified:
334 existing_user = await self.user_store.get_user_by_email(
335 oauth_user.email
336 )
337 if existing_user:
338 from typing import cast
340 return cast("User", existing_user)
341 identity_user = await self._find_by_oauth_identity(
342 oauth_user.provider,
343 oauth_user.provider_user_id,
344 )
345 if identity_user:
346 return identity_user
348 # Provision new user
349 # ``name`` used for login; for profile we keep the display name
350 # provided by the OAuth provider if available.
351 display_name = None
352 if oauth_user.raw_data:
353 display_name = oauth_user.raw_data.get("name")
354 if not display_name:
355 display_name = oauth_user.name
357 user = await self.user_store.create_user(
358 name=oauth_user.name or oauth_user.email,
359 email=oauth_user.email,
360 hashed_password=None, # No password for OAuth users
361 roles=["user"],
362 is_verified=oauth_user.email_verified,
363 profile={
364 "name": display_name,
365 "avatar_url": oauth_user.avatar_url,
366 "oauth_provider": oauth_user.provider,
367 },
368 )
370 # Link OAuth identity (if store available)
371 if self.oauth_identity_store:
372 await self._link_oauth_identity(
373 user.user_id,
374 oauth_user.provider,
375 oauth_user.provider_user_id,
376 )
378 logger.info(
379 "Provisioned new user %s via %s",
380 user.name,
381 oauth_user.provider,
382 )
384 from typing import cast
386 return cast("User", user)
388 async def _find_by_oauth_identity(
389 self,
390 provider: str,
391 provider_user_id: str,
392 ) -> User | None:
393 """Find user by OAuth identity."""
394 if not self.oauth_identity_store:
395 return None
397 identity = await self.oauth_identity_store.get_oauth_identity(
398 provider,
399 provider_user_id,
400 )
401 if identity:
402 from typing import cast
404 return cast("User", await self.user_store.get_user_by_id(identity.user_id))
405 return None
407 async def _link_oauth_identity(
408 self,
409 user_id: str,
410 provider: str,
411 provider_user_id: str,
412 ) -> None:
413 """Link OAuth identity to user."""
414 if not self.oauth_identity_store:
415 return
417 await self.oauth_identity_store.create_oauth_identity(
418 user_id=user_id,
419 provider=provider,
420 provider_user_id=provider_user_id,
421 )
424__all__ = [
425 "LexigramConnectResponse",
426 "LexigramConnectSession",
427 "OAuth2AuthProvider",
428 "OAuth2Manager",
429 "OAuth2Provider",
430]