Coverage for src/lexigram/admin/auth/guards.py: 67%

218 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-21 15:04 +0800

1"""Authentication and authorization guards for lexigram-admin. 

2 

3Provides middleware and guard utilities for protecting routes. 

4Integrates with lexigram-auth session management. 

5""" 

6 

7from __future__ import annotations 

8 

9import base64 

10from dataclasses import dataclass 

11from functools import wraps 

12from typing import TYPE_CHECKING, Any 

13 

14from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint 

15from starlette.requests import Request 

16from starlette.responses import RedirectResponse, Response 

17 

18from lexigram.admin.auth.permissions import PermissionSet, get_user_permissions 

19from lexigram.admin.exceptions import ErrorCode, PermissionDeniedError 

20from lexigram.contracts import ( 

21 AuthorizerProtocol, 

22 AuthProviderProtocol, 

23) 

24from lexigram.contracts.web import RequestProtocol, ResponseProtocol 

25from lexigram.di.decorators import inject 

26from lexigram.logging import get_logger 

27from lexigram.result import Err, Ok, Result 

28from lexigram.serialization.backends import json as json_backend 

29 

30if TYPE_CHECKING: 

31 from collections.abc import Awaitable, Callable 

32 

33logger = get_logger(__name__) 

34 

35 

36@dataclass 

37class GuardConfig: 

38 """Configuration for authentication guards.""" 

39 

40 login_url: str = "/admin/login" 

41 logout_url: str = "/admin/logout" 

42 exempt_paths: tuple[str, ...] = ( 

43 "/admin/login", 

44 "/admin/static", 

45 "/admin/health", 

46 # Standalone pre-session flows (own CSRF + guest handling): 

47 "/admin/setup", 

48 "/admin/verify-email", 

49 "/admin/password-reset", 

50 ) 

51 # Whether to accept Authorization: Bearer <token> for admin APIs. 

52 # Default: False to enforce strict cookie-based admin sessions. 

53 allow_bearer_tokens: bool = False 

54 htmx_redirect_header: str = "HX-Redirect" 

55 

56 

57@inject 

58class AuthGuardMiddleware(BaseHTTPMiddleware): 

59 """Middleware that enforces authentication on admin routes. 

60 

61 Checks for valid session and loads user into request.state. 

62 Redirects unauthenticated requests to login page. 

63 

64 For HTMX requests, returns HX-Redirect header instead of 302. 

65 """ 

66 

67 def __init__( 

68 self, 

69 app: Any, 

70 auth_provider: AuthProviderProtocol | None = None, 

71 config: GuardConfig | None = None, 

72 authorizer: AuthorizerProtocol | None = None, 

73 ) -> None: 

74 super().__init__(app) 

75 self.auth_provider = auth_provider 

76 self.config = config or GuardConfig() 

77 self.authorizer = authorizer 

78 

79 async def dispatch( # type: ignore[override] 

80 self, 

81 request: RequestProtocol, # type: ignore[override] 

82 call_next: RequestResponseEndpoint, 

83 ) -> ResponseProtocol: 

84 # Skip auth for exempt paths 

85 if self._is_exempt(request.url.path): 

86 return await call_next(request) # type: ignore[return-value, arg-type] 

87 

88 # Check if user is already loaded by AdminAuthMiddleware 

89 user = getattr(request.state, "user", None) 

90 

91 # If not, try to load it (fallback/standalone usage) 

92 if user is None: 

93 user = await self._get_authenticated_user(request) 

94 

95 if not self._is_authenticated(user): 

96 # If request has Authorization: Bearer, return 401 instead of redirect 

97 auth_header = request.headers.get("Authorization", "") 

98 if auth_header.startswith("Bearer "): 

99 from starlette.responses import JSONResponse 

100 

101 return JSONResponse( # type: ignore[return-value] 

102 {"authenticated": False, "detail": "Invalid or missing token"}, 

103 status_code=401, 

104 ) 

105 return self._redirect_to_login(request) 

106 

107 # Ensure user and permissions are in state 

108 request.state.user = user 

109 if not hasattr(request.state, "permissions"): 

110 # Use injected authorizer if available 

111 if self.authorizer: 

112 try: 

113 request.state.permissions = get_user_permissions( 

114 user, 

115 self.authorizer, 

116 ) 

117 except ( 

118 ConnectionError, 

119 RuntimeError, 

120 ValueError, 

121 TypeError, 

122 AttributeError, 

123 ): 

124 # Authorization failed, skip permissions 

125 logger.debug( 

126 "Could not compute user permissions", 

127 exc_info=True, 

128 ) 

129 request.state.permissions = None 

130 else: 

131 # No authorizer available, skip permissions 

132 request.state.permissions = None 

133 

134 return await call_next(request) # type: ignore[return-value, arg-type] 

135 

136 def _is_authenticated(self, user: Any) -> bool: 

137 """Check if user is traditionally authenticated (not guest).""" 

138 if user is None: 

139 return False 

140 

141 # Check common identity fields 

142 user_id = getattr(user, "user_id", None) or getattr(user, "id", None) 

143 if not user_id or user_id == "guest": 

144 return False 

145 

146 # Check activity 

147 return getattr(user, "is_active", True) 

148 

149 def _is_exempt(self, path: str) -> bool: 

150 """Check if path is exempt from auth.""" 

151 return any(path.startswith(exempt) for exempt in self.config.exempt_paths) 

152 

153 async def _get_authenticated_user(self, request: RequestProtocol) -> Any | None: 

154 """Get authenticated user from signed request session.""" 

155 if "session" in request.scope: # type: ignore[attr-defined] 

156 user_id = request.session.get("admin_user_id") # type: ignore[attr-defined] 

157 if user_id: 

158 try: 

159 if hasattr(self.auth_provider, "user_store"): 

160 return await self.auth_provider.user_store.get_user_by_id( # type: ignore[union-attr] 

161 user_id, 

162 ) 

163 return None 

164 except ( 

165 ConnectionError, 

166 RuntimeError, 

167 ValueError, 

168 TypeError, 

169 AttributeError, 

170 ): 

171 logger.debug("Session user resolution failed", exc_info=True) 

172 

173 # Try Authorization header (for API calls) 

174 # By default, admin routes do NOT accept bearer tokens to keep a strict 

175 # separation between admin sessions (cookie-based) and application 

176 # JWTs. This can be enabled explicitly via GuardConfig.allow_bearer_tokens. 

177 if self.config.allow_bearer_tokens: 

178 auth_header = request.headers.get("Authorization", "") 

179 if auth_header.startswith("Bearer "): 

180 token = auth_header[7:] 

181 try: 

182 # lexigram-auth uses authenticate_user or verify_token but AuthGuard usually checks tokens too 

183 if hasattr(self.auth_provider, "verify_token"): 

184 token_result = self.auth_provider.verify_token(token) # type: ignore[union-attr] 

185 if hasattr(token_result, "__await__"): 

186 token_result = await token_result 

187 # Handle Result[VerifiedToken, ...] (new API) 

188 if hasattr(token_result, "is_ok"): 

189 if token_result.is_ok(): 

190 verified = token_result.unwrap() # type: ignore[union-attr] 

191 return ( 

192 await self.auth_provider.user_store.get_user_by_id( # type: ignore[union-attr] 

193 verified.user_id, 

194 ) 

195 ) 

196 # Fallback: legacy dict payload (older providers) 

197 elif token_result and "sub" in token_result: # type: ignore[operator] 

198 return await self.auth_provider.user_store.get_user_by_id( # type: ignore[union-attr] 

199 token_result["sub"], # type: ignore[index] 

200 ) 

201 elif self.auth_provider is not None and hasattr( 

202 self.auth_provider, "validate_token" 

203 ): 

204 payload = self.auth_provider.validate_token(token) 

205 if hasattr(payload, "__await__"): 

206 payload = await payload 

207 if ( 

208 payload 

209 and "sub" in payload 

210 and self.auth_provider is not None 

211 ): 

212 user_store = getattr(self.auth_provider, "user_store", None) 

213 if user_store is not None: 

214 return await user_store.get_user_by_id( 

215 payload["sub"], 

216 ) 

217 return payload 

218 except ( 

219 ConnectionError, 

220 RuntimeError, 

221 ValueError, 

222 TypeError, 

223 AttributeError, 

224 ) as e: 

225 # Provide richer diagnostic logging so we can see token header issues (e.g., unexpected 'alg') 

226 try: 

227 header = None 

228 try: 

229 segment = token.split(".", 1)[0] 

230 padded = segment + "=" * (-len(segment) % 4) 

231 header = json_backend.loads( 

232 base64.urlsafe_b64decode(padded) 

233 ) 

234 except (ValueError, TypeError): 

235 header = None 

236 logger.warning( 

237 "Token validation failed: %s - header=%s", 

238 str(e), 

239 header, 

240 exc_info=True, 

241 ) 

242 except (OSError, ValueError, TypeError) as e: 

243 logger.warning("Token validation failed", exc_info=True) 

244 return None 

245 

246 def _redirect_to_login(self, request: RequestProtocol) -> ResponseProtocol: 

247 """Create redirect response to login page.""" 

248 # Build redirect URL with return path 

249 return_to = request.url.path 

250 if request.url.query: 

251 return_to = f"{return_to}?{request.url.query}" 

252 

253 login_url = f"{self.config.login_url}?next={return_to}" 

254 

255 # For HTMX requests, use HX-Redirect header 

256 if request.headers.get("HX-Request"): 

257 response = Response(status_code=200) 

258 response.headers[self.config.htmx_redirect_header] = login_url 

259 return response # type: ignore[return-value] 

260 

261 return RedirectResponse(url=login_url, status_code=302) # type: ignore[return-value] 

262 

263 

264class PermissionGuard: 

265 """GuardProtocol that checks permissions on specific routes. 

266 

267 Usage with @use_guards decorator: 

268 class UserController(Controller): 

269 @get("/admin/users") 

270 @use_guards(PermissionGuard("users.list")) 

271 async def list_users(self, request: Request) -> ...: ... 

272 

273 Usage as standalone callable: 

274 guard = PermissionGuard("users.delete") 

275 result = await guard(request) 

276 if result.is_err(): 

277 raise result.unwrap_err() 

278 """ 

279 

280 def __init__( 

281 self, 

282 *permissions: str, 

283 require_all: bool = False, 

284 message: str | None = None, 

285 authorizer: AuthorizerProtocol | None = None, 

286 ): 

287 self.permissions = permissions 

288 self.require_all = require_all 

289 self.message = message 

290 self._authorizer = authorizer 

291 

292 async def __call__( 

293 self, request: RequestProtocol 

294 ) -> Result[None, PermissionDeniedError]: 

295 """Check permissions. Returns Ok(None) on success, Err(PermissionDeniedError) on denial.""" 

296 user = getattr(request.state, "user", None) 

297 

298 if user is None: 

299 return Err(PermissionDeniedError(message="Authentication required")) 

300 

301 user_perms: PermissionSet = getattr(request.state, "permissions", None) # type: ignore[assignment] 

302 if user_perms is None: 

303 # Use injected authorizer or fallback to request.state.permissions if already set 

304 authorizer = self._authorizer 

305 if authorizer is None: 

306 # Permissions should have been set by middleware 

307 return Err( 

308 PermissionDeniedError( 

309 message="Authorization service unavailable", 

310 ) 

311 ) 

312 user_perms = get_user_permissions(user, authorizer) 

313 

314 if self.require_all: 

315 if not user_perms.has_all(*self.permissions): 

316 missing = list( 

317 filter(lambda p: not user_perms.has(p), self.permissions), 

318 ) 

319 return Err( 

320 PermissionDeniedError( 

321 message=self.message 

322 or f"Missing permissions: {', '.join(missing)}", 

323 required_permission=str(self.permissions), 

324 ) 

325 ) 

326 elif not user_perms.has_any(*self.permissions): 

327 return Err( 

328 PermissionDeniedError( 

329 message=self.message 

330 or f"Requires permission: {' or '.join(self.permissions)}", 

331 required_permission=str(self.permissions), 

332 ) 

333 ) 

334 return Ok(None) 

335 

336 def __matmul__(self, func: Callable) -> Callable: 

337 """Allow usage as @guard decorator via @ operator.""" 

338 return self.wrap(func) 

339 

340 def wrap( 

341 self, 

342 func: Callable[..., Awaitable[Any]], 

343 ) -> Callable[..., Awaitable[Any]]: 

344 """Wrap a function with permission check. 

345 

346 Raises PermissionDeniedError if the guard check fails. 

347 """ 

348 

349 @wraps(func) 

350 async def wrapper(request: Request, *args, **kwargs) -> Any: 

351 result = await self(request) # type: ignore[arg-type] 

352 if result.is_err(): 

353 raise result.unwrap_err() 

354 return await func(request, *args, **kwargs) 

355 

356 return wrapper 

357 

358 

359class RoleGuard: 

360 """GuardProtocol that checks roles on specific routes.""" 

361 

362 def __init__( 

363 self, 

364 *roles: str, 

365 require_all: bool = False, 

366 message: str | None = None, 

367 ): 

368 self.roles = roles 

369 self.require_all = require_all 

370 self.message = message 

371 

372 async def __call__( 

373 self, request: RequestProtocol 

374 ) -> Result[None, PermissionDeniedError]: 

375 """Check roles. Returns Ok(None) on success, Err(PermissionDeniedError) on denial.""" 

376 user = getattr(request.state, "user", None) 

377 

378 if user is None: 

379 return Err(PermissionDeniedError(message="Authentication required")) 

380 

381 user_roles = set(getattr(user, "roles", []) or []) 

382 

383 if self.require_all: 

384 if not all(r in user_roles for r in self.roles): 

385 missing = list(filter(lambda r: r not in user_roles, self.roles)) 

386 return Err( 

387 PermissionDeniedError( 

388 message=self.message 

389 or f"Requires all roles: {', '.join(missing)}", 

390 ) 

391 ) 

392 elif not user_roles.intersection(self.roles): 

393 return Err( 

394 PermissionDeniedError( 

395 message=self.message or f"Requires role: {' or '.join(self.roles)}", 

396 ) 

397 ) 

398 return Ok(None) 

399 

400 

401def require_auth(func: Callable[..., Awaitable[Any]]) -> Callable[..., Awaitable[Any]]: 

402 """Simple decorator to require authentication. 

403 

404 Just checks that user exists in request.state. 

405 """ 

406 

407 @wraps(func) 

408 async def wrapper(request: Request, *args, **kwargs) -> Any: 

409 user = getattr(request.state, "user", None) 

410 if user is None: 

411 raise PermissionDeniedError(message="Authentication required") 

412 return await func(request, *args, **kwargs) 

413 

414 return wrapper 

415 

416 

417def csrf_protect(func: Callable[..., Awaitable[Any]]) -> Callable[..., Awaitable[Any]]: 

418 """Decorator to require valid CSRF token for state-changing operations. 

419 

420 Checks for CSRF token in: 

421 1. X-CSRF-Token header 

422 2. csrf_token form field 

423 

424 HTMX requests automatically include the token via hx-headers. 

425 """ 

426 

427 @wraps(func) 

428 async def wrapper(request: Request, *args, **kwargs) -> Any: 

429 # Skip for safe methods 

430 if request.method in ("GET", "HEAD", "OPTIONS"): 

431 return await func(request, *args, **kwargs) 

432 

433 # Get expected token from session 

434 session = getattr(request.state, "session", None) 

435 expected_token = getattr(session, "csrf_token", None) if session else None 

436 

437 if expected_token is None: 

438 # No CSRF protection configured 

439 logger.warning("CSRF protection skipped - no token in session") 

440 return await func(request, *args, **kwargs) 

441 

442 # Get submitted token 

443 submitted_token = request.headers.get("X-CSRF-Token") 

444 

445 if not submitted_token: 

446 # Try form data 

447 try: 

448 form = request.scope.get("admin_form_data") 

449 if form is None: 

450 form = await request.form() 

451 submitted_token = form.get("csrf_token") # type: ignore[assignment] 

452 except ( 

453 ConnectionError, 

454 RuntimeError, 

455 ValueError, 

456 TypeError, 

457 AttributeError, 

458 ): 

459 pass 

460 

461 if not submitted_token or submitted_token != expected_token: 

462 raise PermissionDeniedError( 

463 message="Invalid or missing CSRF token", 

464 code=ErrorCode.AUTH_INVALID_TOKEN, 

465 ) 

466 

467 return await func(request, *args, **kwargs) 

468 

469 return wrapper 

470 

471 

472class CompositeGuard: 

473 """Combine multiple guards with AND/OR logic. 

474 

475 Usage: 

476 guard = CompositeGuard( 

477 PermissionGuard("users.list"), 

478 RoleGuard("admin"), 

479 logic="or" # User needs permission OR role 

480 ) 

481 """ 

482 

483 def __init__( 

484 self, 

485 *guards: PermissionGuard | RoleGuard, 

486 logic: str = "and", # "and" or "or" 

487 ): 

488 self.guards = guards 

489 self.logic = logic 

490 

491 async def __call__( 

492 self, request: RequestProtocol 

493 ) -> Result[None, PermissionDeniedError]: 

494 """Execute guards based on logic. 

495 

496 Returns Ok(None) when guard(s) pass. Returns Err(PermissionDeniedError) 

497 on denial. For "and" logic, the first failure short-circuits. For "or" 

498 logic, the last failure is returned if all guards deny. 

499 """ 

500 if self.logic == "and": 

501 # All guards must pass — short-circuit on first failure 

502 for guard in self.guards: 

503 result = await guard(request) 

504 if result.is_err(): 

505 return result 

506 return Ok(None) 

507 # At least one guard must pass 

508 last_failure: Result[None, PermissionDeniedError] = Err( 

509 PermissionDeniedError(message="All guards denied access") 

510 ) 

511 for guard in self.guards: 

512 result = await guard(request) 

513 if result.is_ok(): 

514 return Ok(None) 

515 last_failure = result 

516 return last_failure