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

215 statements  

« prev     ^ index     » next       coverage.py v7.13.5, created at 2026-08-13 22:14 +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 

9from dataclasses import dataclass 

10from functools import wraps 

11from typing import TYPE_CHECKING, Any 

12 

13from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint 

14from starlette.requests import Request 

15from starlette.responses import RedirectResponse, Response 

16 

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

18from lexigram.admin.exceptions import ErrorCode, PermissionDeniedError 

19from lexigram.contracts import ( 

20 AuthorizerProtocol, 

21 AuthProviderProtocol, 

22) 

23from lexigram.contracts.web import RequestProtocol, ResponseProtocol 

24from lexigram.di.decorators import inject 

25from lexigram.logging import get_logger 

26from lexigram.result import Err, Ok, Result 

27 

28if TYPE_CHECKING: 

29 from collections.abc import Awaitable, Callable 

30 

31logger = get_logger(__name__) 

32 

33 

34@dataclass 

35class GuardConfig: 

36 """Configuration for authentication guards.""" 

37 

38 login_url: str = "/admin/login" 

39 logout_url: str = "/admin/logout" 

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

41 "/admin/login", 

42 "/admin/static", 

43 "/admin/health", 

44 ) 

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

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

47 allow_bearer_tokens: bool = False 

48 htmx_redirect_header: str = "HX-Redirect" 

49 

50 

51@inject 

52class AuthGuardMiddleware(BaseHTTPMiddleware): 

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

54 

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

56 Redirects unauthenticated requests to login page. 

57 

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

59 """ 

60 

61 def __init__( 

62 self, 

63 app: Any, 

64 auth_provider: AuthProviderProtocol | None = None, 

65 config: GuardConfig | None = None, 

66 authorizer: AuthorizerProtocol | None = None, 

67 ) -> None: 

68 super().__init__(app) 

69 self.auth_provider = auth_provider 

70 self.config = config or GuardConfig() 

71 self.authorizer = authorizer 

72 

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

74 self, 

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

76 call_next: RequestResponseEndpoint, 

77 ) -> ResponseProtocol: 

78 # Skip auth for exempt paths 

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

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

81 

82 # Check if user is already loaded by AdminAuthMiddleware 

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

84 

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

86 if user is None: 

87 user = await self._get_authenticated_user(request) 

88 

89 if not self._is_authenticated(user): 

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

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

92 if auth_header.startswith("Bearer "): 

93 from starlette.responses import JSONResponse 

94 

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

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

97 status_code=401, 

98 ) 

99 return self._redirect_to_login(request) 

100 

101 # Ensure user and permissions are in state 

102 request.state.user = user 

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

104 # Use injected authorizer if available 

105 if self.authorizer: 

106 try: 

107 request.state.permissions = get_user_permissions( 

108 user, 

109 self.authorizer, 

110 ) 

111 except ( 

112 ConnectionError, 

113 RuntimeError, 

114 ValueError, 

115 TypeError, 

116 AttributeError, 

117 ): 

118 # Authorization failed, skip permissions 

119 logger.debug( 

120 "Could not compute user permissions", 

121 exc_info=True, 

122 ) 

123 request.state.permissions = None 

124 else: 

125 # No authorizer available, skip permissions 

126 request.state.permissions = None 

127 

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

129 

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

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

132 if user is None: 

133 return False 

134 

135 # Check common identity fields 

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

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

138 return False 

139 

140 # Check activity 

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

142 

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

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

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

146 

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

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

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

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

151 if user_id: 

152 try: 

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

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

155 user_id, 

156 ) 

157 return None 

158 except ( 

159 ConnectionError, 

160 RuntimeError, 

161 ValueError, 

162 TypeError, 

163 AttributeError, 

164 ): 

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

166 

167 # Try Authorization header (for API calls) 

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

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

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

171 if self.config.allow_bearer_tokens: 

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

173 if auth_header.startswith("Bearer "): 

174 token = auth_header[7:] 

175 try: 

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

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

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

179 if hasattr(token_result, "__await__"): 

180 token_result = await token_result 

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

182 if hasattr(token_result, "is_ok"): 

183 if token_result.is_ok(): 

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

185 return ( 

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

187 verified.user_id, 

188 ) 

189 ) 

190 # Fallback: legacy dict payload (older providers) 

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

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

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

194 ) 

195 elif self.auth_provider is not None and hasattr( 

196 self.auth_provider, "validate_token" 

197 ): 

198 payload = self.auth_provider.validate_token(token) 

199 if hasattr(payload, "__await__"): 

200 payload = await payload 

201 if ( 

202 payload 

203 and "sub" in payload 

204 and self.auth_provider is not None 

205 ): 

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

207 if user_store is not None: 

208 return await user_store.get_user_by_id( 

209 payload["sub"], 

210 ) 

211 return payload 

212 except ( 

213 ConnectionError, 

214 RuntimeError, 

215 ValueError, 

216 TypeError, 

217 AttributeError, 

218 ) as e: 

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

220 try: 

221 from jose import jwt as jose_jwt # type: ignore[import-untyped] 

222 

223 header = None 

224 try: 

225 header = jose_jwt.get_unverified_header(token) 

226 except ValueError: 

227 header = None 

228 logger.warning( 

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

230 str(e), 

231 header, 

232 exc_info=True, 

233 ) 

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

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

236 return None 

237 

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

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

240 # Build redirect URL with return path 

241 return_to = request.url.path 

242 if request.url.query: 

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

244 

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

246 

247 # For HTMX requests, use HX-Redirect header 

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

249 response = Response(status_code=200) 

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

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

252 

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

254 

255 

256class PermissionGuard: 

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

258 

259 Usage with @use_guards decorator: 

260 class UserController(Controller): 

261 @get("/admin/users") 

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

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

264 

265 Usage as standalone callable: 

266 guard = PermissionGuard("users.delete") 

267 result = await guard(request) 

268 if result.is_err(): 

269 raise result.unwrap_err() 

270 """ 

271 

272 def __init__( 

273 self, 

274 *permissions: str, 

275 require_all: bool = False, 

276 message: str | None = None, 

277 authorizer: AuthorizerProtocol | None = None, 

278 ): 

279 self.permissions = permissions 

280 self.require_all = require_all 

281 self.message = message 

282 self._authorizer = authorizer 

283 

284 async def __call__( 

285 self, request: RequestProtocol 

286 ) -> Result[None, PermissionDeniedError]: 

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

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

289 

290 if user is None: 

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

292 

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

294 if user_perms is None: 

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

296 authorizer = self._authorizer 

297 if authorizer is None: 

298 # Permissions should have been set by middleware 

299 return Err( 

300 PermissionDeniedError( 

301 message="Authorization service unavailable", 

302 ) 

303 ) 

304 user_perms = get_user_permissions(user, authorizer) 

305 

306 if self.require_all: 

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

308 missing = list( 

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

310 ) 

311 return Err( 

312 PermissionDeniedError( 

313 message=self.message 

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

315 required_permission=str(self.permissions), 

316 ) 

317 ) 

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

319 return Err( 

320 PermissionDeniedError( 

321 message=self.message 

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

323 required_permission=str(self.permissions), 

324 ) 

325 ) 

326 return Ok(None) 

327 

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

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

330 return self.wrap(func) 

331 

332 def wrap( 

333 self, 

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

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

336 """Wrap a function with permission check. 

337 

338 Raises PermissionDeniedError if the guard check fails. 

339 """ 

340 

341 @wraps(func) 

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

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

344 if result.is_err(): 

345 raise result.unwrap_err() 

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

347 

348 return wrapper 

349 

350 

351class RoleGuard: 

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

353 

354 def __init__( 

355 self, 

356 *roles: str, 

357 require_all: bool = False, 

358 message: str | None = None, 

359 ): 

360 self.roles = roles 

361 self.require_all = require_all 

362 self.message = message 

363 

364 async def __call__( 

365 self, request: RequestProtocol 

366 ) -> Result[None, PermissionDeniedError]: 

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

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

369 

370 if user is None: 

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

372 

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

374 

375 if self.require_all: 

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

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

378 return Err( 

379 PermissionDeniedError( 

380 message=self.message 

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

382 ) 

383 ) 

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

385 return Err( 

386 PermissionDeniedError( 

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

388 ) 

389 ) 

390 return Ok(None) 

391 

392 

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

394 """Simple decorator to require authentication. 

395 

396 Just checks that user exists in request.state. 

397 """ 

398 

399 @wraps(func) 

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

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

402 if user is None: 

403 raise PermissionDeniedError(message="Authentication required") 

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

405 

406 return wrapper 

407 

408 

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

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

411 

412 Checks for CSRF token in: 

413 1. X-CSRF-Token header 

414 2. csrf_token form field 

415 

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

417 """ 

418 

419 @wraps(func) 

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

421 # Skip for safe methods 

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

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

424 

425 # Get expected token from session 

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

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

428 

429 if expected_token is None: 

430 # No CSRF protection configured 

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

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

433 

434 # Get submitted token 

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

436 

437 if not submitted_token: 

438 # Try form data 

439 try: 

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

441 if form is None: 

442 form = await request.form() 

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

444 except ( 

445 ConnectionError, 

446 RuntimeError, 

447 ValueError, 

448 TypeError, 

449 AttributeError, 

450 ): 

451 pass 

452 

453 if not submitted_token or submitted_token != expected_token: 

454 raise PermissionDeniedError( 

455 message="Invalid or missing CSRF token", 

456 code=ErrorCode.AUTH_INVALID_TOKEN, 

457 ) 

458 

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

460 

461 return wrapper 

462 

463 

464class CompositeGuard: 

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

466 

467 Usage: 

468 guard = CompositeGuard( 

469 PermissionGuard("users.list"), 

470 RoleGuard("admin"), 

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

472 ) 

473 """ 

474 

475 def __init__( 

476 self, 

477 *guards: PermissionGuard | RoleGuard, 

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

479 ): 

480 self.guards = guards 

481 self.logic = logic 

482 

483 async def __call__( 

484 self, request: RequestProtocol 

485 ) -> Result[None, PermissionDeniedError]: 

486 """Execute guards based on logic. 

487 

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

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

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

491 """ 

492 if self.logic == "and": 

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

494 for guard in self.guards: 

495 result = await guard(request) 

496 if result.is_err(): 

497 return result 

498 return Ok(None) 

499 # At least one guard must pass 

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

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

502 ) 

503 for guard in self.guards: 

504 result = await guard(request) 

505 if result.is_ok(): 

506 return Ok(None) 

507 last_failure = result 

508 return last_failure