Coverage for src/lexigram/auth/web/middleware/auth.py: 34%

198 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-26 00:58 +0800

1"""Authentication middleware for web applications""" 

2 

3from __future__ import annotations 

4 

5from functools import wraps 

6import hashlib 

7from typing import TYPE_CHECKING, Any, cast 

8 

9from lexigram.auth.config import AuthMiddlewareConfig 

10from lexigram.logging import get_logger 

11from lexigram.primitives.context import USER_ID, Context 

12 

13if TYPE_CHECKING: 

14 from collections.abc import Callable 

15 

16 from lexigram.auth.authn.services import LoginAttemptTracker 

17 from lexigram.auth.models.user import User 

18 from lexigram.contracts import AuthProviderProtocol 

19 from lexigram.contracts.web import RequestProtocol as Request 

20 

21from datetime import UTC 

22 

23# Re-export guards for convenience 

24from lexigram.auth.authz.guards import optional_auth, require_permissions, require_roles 

25from lexigram.auth.web.middleware.api_key_authenticator import ApiKeyAuthenticator 

26from lexigram.auth.web.middleware.jwt_authenticator import JwtAuthenticator 

27from lexigram.auth.web.middleware.response_handler import AuthResponseHandler 

28from lexigram.auth.web.middleware.session_authenticator import SessionAuthenticator 

29from lexigram.auth.web.middleware.session_validator import SessionValidator 

30from lexigram.auth.web.middleware.throttle import RateLimitMiddleware 

31from lexigram.auth.web.middleware.token_cache import TokenCache 

32from lexigram.auth.web.middleware.token_extractor import TokenExtractor 

33 

34logger = get_logger(__name__) 

35 

36 

37class AuthMiddleware: 

38 """Middleware for handling authentication and authorization - Pure ASGI implementation.""" 

39 

40 def __init__( 

41 self, 

42 auth_provider: AuthProviderProtocol, 

43 config: AuthMiddlewareConfig | None = None, 

44 ctx: Context | None = None, 

45 attempt_tracker: LoginAttemptTracker | None = None, 

46 ): 

47 self.auth_provider = auth_provider 

48 self.config = config or AuthMiddlewareConfig() 

49 self._ctx = ctx 

50 

51 # Initialize extracted components 

52 self.token_extractor = TokenExtractor(self.config) 

53 self.session_validator = SessionValidator(self.config, self.auth_provider) 

54 self.token_cache = TokenCache() 

55 self.api_key_authenticator = ApiKeyAuthenticator(self.auth_provider) 

56 self.session_authenticator = SessionAuthenticator(self.auth_provider) 

57 self.jwt_authenticator = JwtAuthenticator(self.auth_provider) 

58 self.response_handler = AuthResponseHandler() 

59 

60 self.attempt_tracker = attempt_tracker 

61 

62 # Initialize Rate Limiter 

63 cache_service = getattr(self.auth_provider, "cache_service", None) 

64 # Safely extract rate_limit value (avoid accessing Field descriptor) 

65 rate_limit_val = getattr(self.config, "login_rate_limit", None) 

66 if not isinstance(rate_limit_val, str): 

67 rate_limit_val = "5/minute" 

68 self.rate_limiter = RateLimitMiddleware( 

69 app=None, # Will be managed manually 

70 cache_service=cache_service, 

71 rate_limit=rate_limit_val, 

72 ) 

73 

74 def should_skip_auth(self, path: str) -> bool: 

75 """Check if authentication should be skipped for this path""" 

76 return self.session_validator.should_skip_auth(path) 

77 

78 def extract_token(self, request: Request) -> str | None: 

79 """Extract authentication token from request""" 

80 return self.token_extractor.extract_token(request) 

81 

82 async def authenticate_request(self, request: Any) -> User | None: 

83 """Authenticate the request and return user if valid""" 

84 token = self.extract_token(request) 

85 if not token: 

86 logger.info("AuthMiddleware.authenticate_request: no token extracted") 

87 return None 

88 

89 # Check lockout before doing any authentication work 

90 if self.attempt_tracker is not None: 

91 client_ip = getattr(getattr(request, "client", None), "host", token) 

92 if await self.attempt_tracker.is_locked(client_ip): 

93 logger.warning( 

94 "AuthMiddleware.authenticate_request: client locked out", 

95 extra={"client_ip": client_ip}, 

96 ) 

97 return None 

98 

99 # Check token cache first (avoid JWT decode + DB lookup) 

100 cached_user = await self.token_cache.get(token) 

101 if cached_user: 

102 token_hash = hashlib.sha256(token.encode()).hexdigest() 

103 logger.info( 

104 "AuthMiddleware.authenticate_request: Token cache HIT for token_hash=%s", 

105 token_hash[:10], 

106 ) 

107 return cast("User | None", cached_user) 

108 

109 # Try API key authentication 

110 user = await self.api_key_authenticator.authenticate(token, request) 

111 if user: 

112 await self.token_cache.set(token, user) 

113 return cast("User | None", user) 

114 

115 # Try session authentication 

116 user = await self.session_authenticator.authenticate(request) 

117 if user: 

118 return cast("User | None", user) 

119 

120 # Try JWT authentication — the authenticator now extracts the token 

121 # from the request internally (signature changed in an earlier refactor). 

122 user = await self.jwt_authenticator.authenticate(request) 

123 if user: 

124 await self.token_cache.set(token, user) 

125 

126 # Record failed attempt when all authenticators returned None 

127 if user is None and self.attempt_tracker is not None: 

128 client_ip = getattr(getattr(request, "client", None), "host", token) 

129 await self.attempt_tracker.record_failure(client_ip) 

130 logger.debug( 

131 "AuthMiddleware.authenticate_request: recorded failed attempt", 

132 extra={"client_ip": client_ip}, 

133 ) 

134 

135 return cast("User | None", user) 

136 

137 def check_authorization(self, user: User) -> bool: 

138 """Check if user is authorized based on roles/permissions""" 

139 return self.session_validator.check_authorization(user) 

140 

141 async def __call__( 

142 self, 

143 scope: dict[str, Any], 

144 receive: Callable, 

145 send: Callable, 

146 ) -> None: 

147 """Pure ASGI middleware entry point - OPT-AUTH-2.""" 

148 # Only handle HTTP requests 

149 if scope.get("type") != "http": 

150 await self.app(scope, receive, send) 

151 return 

152 

153 # ASGI framework-binding layer: StarletteRequest constructs a request 

154 # from the raw ASGI scope/receive callables. This is intentionally 

155 # Starlette-specific — it cannot be replaced with RequestProtocol, 

156 # which is a structural protocol for type annotations only. 

157 from starlette.requests import Request as StarletteRequest 

158 

159 request = StarletteRequest(scope, receive) 

160 

161 # Skip authentication for excluded paths 

162 logger.debug("AuthMiddleware: checking path=%s", request.url.path) 

163 if self.should_skip_auth(request.url.path): 

164 await self.app(scope, receive, send) 

165 return 

166 

167 # Skip authentication for OPTIONS method (CORS preflight) 

168 logger.debug( 

169 "AuthMiddleware: method=%s, path=%s", request.method, request.url.path 

170 ) 

171 if request.method == "OPTIONS": 

172 logger.debug("AuthMiddleware: skipping auth for OPTIONS request") 

173 await self.app(scope, receive, send) 

174 return 

175 

176 # Authenticate request 

177 user = await self.authenticate_request(request) 

178 

179 # Store user in scope (Starlette style) 

180 scope["user"] = user 

181 

182 # Initialize state dict if not present 

183 if "state" not in scope: 

184 scope["state"] = {} 

185 scope["state"]["user"] = user 

186 scope["state"]["user_id"] = str(user.user_id) if user is not None else None 

187 

188 # Store user in runtime context for unified access across HTTP/WebSocket/Tasks 

189 from lexigram.di.resolution.context import get_resolver 

190 

191 resolver = get_resolver(scope) 

192 if resolver and self._ctx is not None and user is not None: 

193 self._ctx.set(USER_ID, str(user.user_id)) 

194 

195 # Apply Rate Limiting for auth endpoints 

196 if request.url.path in ["/auth/login", "/auth/register"]: 

197 # We use the raw ASGI call pattern to integrate it 

198 # But we need a 'send' that we can wrap if we want to detect success 

199 # For now, let's keep it simple: just call it. 

200 # However, RateLimitMiddleware expects app to be a callable. 

201 # We can just delegate to our app if allowed. 

202 

203 # Re-usable wrapper to continue if not throttled 

204 async def next_app(s: dict[str, Any], r: Any, sn: Any) -> Any: 

205 await self.app(s, r, sn) 

206 

207 # Temporarily set app and call 

208 self.rate_limiter.app = next_app 

209 await self.rate_limiter(scope, receive, send) 

210 return 

211 

212 # Check authorization if user is required 

213 if ( 

214 not self.config.optional_auth 

215 or self.config.roles_required 

216 or self.config.permissions_required 

217 ): 

218 if not user: 

219 # No user but auth is required 

220 response = await self.response_handler.unauthorized_response( 

221 "Authentication required", 

222 request=request, 

223 ) 

224 await response(scope, receive, send) 

225 return 

226 

227 if not self.check_authorization(user): 

228 # User doesn't have required roles/permissions 

229 response = await self.response_handler.forbidden_response( 

230 "Insufficient permissions", 

231 ) 

232 await response(scope, receive, send) 

233 return 

234 

235 # Continue with request - capture response to add headers 

236 response_started = False 

237 response_headers = [] 

238 

239 async def send_wrapper(message: dict[str, Any]) -> None: 

240 nonlocal response_started, response_headers 

241 if message.get("type") == "http.response.start": 

242 response_started = True 

243 response_headers = list(message.get("headers", [])) 

244 await send(message) 

245 

246 await self.app(scope, receive, send_wrapper) 

247 

248 def set_app(self, app: Callable) -> None: 

249 """Set the ASGI app to wrap.""" 

250 self.app = app 

251 

252 

253# RateLimitMiddleware has been moved to lexigram.auth.web.middleware.throttle 

254 

255 

256class AuthRouter: 

257 """Router extension with authentication helpers""" 

258 

259 def __init__(self, auth_provider: AuthProviderProtocol): 

260 self.auth_provider = auth_provider 

261 

262 def require_auth( 

263 self, 

264 roles: list[str] | None = None, 

265 permissions: list[str] | None = None, 

266 optional: bool = False, 

267 ) -> Callable[[Callable[..., Any]], Callable[..., Any]]: 

268 """Decorator to require authentication and authorization for routes""" 

269 

270 def decorator(func: Callable) -> Callable: 

271 @wraps(func) 

272 async def wrapper(*args: Any, **kwargs: Any) -> Any: 

273 request = _extract_request(*args, **kwargs) 

274 

275 if not request: 

276 raise ValueError( 

277 "Could not find request object in function arguments", 

278 ) 

279 

280 # Check authentication 

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

282 if not optional and not user: 

283 rf = await _get_response_factory(request) 

284 

285 return rf.json( 

286 status_code=401, 

287 content={ 

288 "error": "unauthorized", 

289 "message": "Authentication required", 

290 }, 

291 ) 

292 

293 # Check authorization 

294 if user: 

295 if roles and not self.auth_provider.has_any_role(user, roles): 

296 rf = await _get_response_factory(request) 

297 

298 return rf.json( 

299 status_code=403, 

300 content={ 

301 "error": "forbidden", 

302 "message": "Insufficient roles", 

303 }, 

304 ) 

305 

306 if permissions and not self.auth_provider.has_any_permission( 

307 user, 

308 permissions, 

309 ): 

310 rf = await _get_response_factory(request) 

311 

312 return rf.json( 

313 status_code=403, 

314 content={ 

315 "error": "forbidden", 

316 "message": "Insufficient permissions", 

317 }, 

318 ) 

319 

320 return await func(*args, **kwargs) 

321 

322 return wrapper 

323 

324 return decorator 

325 

326 def get_current_user(self, request: Request) -> User | None: 

327 """Get current authenticated user from request""" 

328 return getattr(request.state, "user", None) 

329 

330 

331# Convenience helpers and functions for common auth patterns 

332 

333 

334def _extract_request(*args: Any, **kwargs: Any) -> Any: 

335 """Extract the Starlette-like request object from positional or keyword args.""" 

336 for arg in args: 

337 if hasattr(arg, "state") and hasattr(arg, "headers"): 

338 return arg 

339 return kwargs.get("request") 

340 

341 

342async def _get_auth_provider(context: Any | None = None) -> AuthProviderProtocol: 

343 """Resolve `AuthProvider` from dynamic context or global container.""" 

344 from lexigram.contracts.auth import AuthProviderProtocol 

345 from lexigram.di.resolution.context import get_resolver 

346 

347 resolver = get_resolver(context) 

348 if not resolver: 

349 raise RuntimeError( 

350 "No DI resolver found in current context. Ensure application is initialized.", 

351 ) 

352 

353 return cast("AuthProviderProtocol", await resolver.resolve(AuthProviderProtocol)) 

354 

355 

356async def _get_response_factory(context: Any | None = None) -> Any: 

357 """Resolve `ResponseFactoryProtocol` from global container.""" 

358 from lexigram.contracts.web import ResponseFactoryProtocol 

359 from lexigram.di.resolution.context import get_resolver 

360 

361 resolver = get_resolver(context) 

362 if not resolver: 

363 return None 

364 

365 return await resolver.resolve(ResponseFactoryProtocol) 

366 

367 

368def require_mfa( 

369 max_age_seconds: int = 300, 

370) -> Callable[[Callable[..., Any]], Callable[..., Any]]: 

371 """Decorator to require MFA verification (step-up). 

372 

373 Ensures: 

374 1. User has MFA enabled. 

375 2. Session was verified with MFA within the last `max_age_seconds`. 

376 """ 

377 

378 def decorator(func: Callable) -> Callable: 

379 @wraps(func) 

380 async def wrapper(*args: Any, **kwargs: Any) -> Any: 

381 request = _extract_request(*args, **kwargs) 

382 if not request: 

383 return await func(*args, **kwargs) 

384 

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

386 if not user: 

387 rf = await _get_response_factory(request) 

388 return rf.json(status_code=401, content={"error": "unauthorized"}) 

389 

390 # Resolve AuthProvider/MFAManager 

391 auth_provider = await _get_auth_provider(request) 

392 if auth_provider.mfa_manager: 

393 mfa = await auth_provider.mfa_manager.get_mfa(user.user_id) 

394 if not mfa or not mfa.is_enabled: 

395 rf = await _get_response_factory(request) 

396 

397 return rf.json( 

398 status_code=403, 

399 content={ 

400 "error": "mfa_required", 

401 "message": "MFA must be enabled for this operation", 

402 }, 

403 ) 

404 

405 # Check session step-up status 

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

407 is_verified = False 

408 if session and session.mfa_verified_at: 

409 from datetime import datetime 

410 

411 age = ( 

412 datetime.now(UTC) - session.mfa_verified_at.replace(tzinfo=UTC) 

413 ).total_seconds() 

414 if age < max_age_seconds: 

415 is_verified = True 

416 

417 if not is_verified: 

418 rf = await _get_response_factory(request) 

419 

420 return rf.json( 

421 status_code=401, 

422 content={ 

423 "error": "mfa_verification_required", 

424 "message": "Step-up authentication required", 

425 "stepup_url": "/api/v1/auth/mfa/verify", 

426 }, 

427 ) 

428 

429 return await func(*args, **kwargs) 

430 

431 return wrapper 

432 

433 return decorator 

434 

435 

436__all__ = [ 

437 "AuthMiddleware", 

438 "AuthMiddlewareConfig", 

439 "AuthRouter", 

440 "logger", 

441 "optional_auth", 

442 "require_mfa", 

443 "require_permissions", 

444 "require_roles", 

445]