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

1"""OAuth2 authentication flows. 

2 

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. 

6 

7Example: 

8 Setting up OAuth2 authentication:: 

9 

10 from lexigram.auth.authn.oauth2 import OAuth2Authenticator 

11 

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 ) 

18 

19 # Get authorization URL 

20 auth_url = await authenticator.get_authorization_url() 

21 

22 # Exchange code for tokens 

23 user_info = await authenticator.get_user_info(code) 

24""" 

25 

26from __future__ import annotations 

27 

28from typing import TYPE_CHECKING, Any, cast 

29 

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 

39 

40# Initialize logger early so runtime checks can log errors 

41logger = get_logger(__name__) 

42 

43try: 

44 from authlib.integrations.base_client import ( 

45 OAuthError, 

46 ) 

47 from authlib.integrations.httpx_client import ( 

48 AsyncOAuth2Client, 

49 ) 

50 

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 

58 

59 

60if TYPE_CHECKING: 

61 from lexigram.auth.models.user import User 

62 from lexigram.contracts.web import HTTPClientProtocol 

63 

64 

65class OAuth2IdentityProvider: 

66 """OAuth2 identity provider configuration.""" 

67 

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 

89 

90 

91class OAuth2Manager: 

92 """OAuth2 authentication manager""" 

93 

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 ) 

109 

110 def __repr__(self) -> str: 

111 """Return developer-friendly string representation.""" 

112 return f"OAuth2Manager(providers={list(self.providers)!r})" 

113 

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. 

120 

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. 

124 

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 

130 

131 if not self._available: 

132 raise RuntimeError( 

133 "authlib is required for OAuth2 operations; OAuth2Manager is in no-op mode", 

134 ) 

135 

136 provider = self.providers.get(provider_name) 

137 if not provider: 

138 raise ValueError(f"OAuth2 provider '{provider_name}' not configured") 

139 

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 ) 

147 

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) 

152 

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 ) 

159 

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") 

164 

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 ) 

172 

173 return str(authorization_url), code_verifier, state 

174 

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") 

185 

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 ) 

191 

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) 

196 

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 ) 

203 

204 if not self._available: 

205 raise RuntimeError( 

206 "authlib is required for OAuth2 operations; OAuth2Manager is in no-op mode", 

207 ) 

208 

209 try: 

210 fetch_kwargs = {"code": code} 

211 if code_verifier: 

212 fetch_kwargs["code_verifier"] = code_verifier 

213 

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 

221 

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") 

231 

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) 

236 

237 client = AsyncOAuth2Client( 

238 client_id=provider.client_id, 

239 client_secret=provider.client_secret, 

240 token=token, 

241 session=session, 

242 ) 

243 

244 if not self._available: 

245 raise RuntimeError( 

246 "authlib is required for OAuth2 operations; OAuth2Manager is in no-op mode", 

247 ) 

248 

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 

254 

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 

267 

268 

269class OAuth2AuthProvider: 

270 """Complete OAuth2 provider with user provisioning.""" 

271 

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 

281 

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.""" 

289 

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 ) 

302 

303 # Get user info from provider 

304 user_info = await self.oauth2_manager.get_user_info(provider_name, token) 

305 

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 ) 

324 

325 # Find or create user 

326 return await self._find_or_create_oauth_user(oauth_user) 

327 

328 async def _find_or_create_oauth_user(self, oauth_user: OAuth2UserInfo) -> User: 

329 """Find existing user or provision new one.""" 

330 

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 

339 

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 

347 

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 

356 

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 ) 

369 

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 ) 

377 

378 logger.info( 

379 "Provisioned new user %s via %s", 

380 user.name, 

381 oauth_user.provider, 

382 ) 

383 

384 from typing import cast 

385 

386 return cast("User", user) 

387 

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 

396 

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 

403 

404 return cast("User", await self.user_store.get_user_by_id(identity.user_id)) 

405 return None 

406 

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 

416 

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 ) 

422 

423 

424__all__ = [ 

425 "LexigramConnectResponse", 

426 "LexigramConnectSession", 

427 "OAuth2AuthProvider", 

428 "OAuth2Manager", 

429 "OAuth2Provider", 

430]