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
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 00:58 +0800
1"""Authentication middleware for web applications"""
3from __future__ import annotations
5from functools import wraps
6import hashlib
7from typing import TYPE_CHECKING, Any, cast
9from lexigram.auth.config import AuthMiddlewareConfig
10from lexigram.logging import get_logger
11from lexigram.primitives.context import USER_ID, Context
13if TYPE_CHECKING:
14 from collections.abc import Callable
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
21from datetime import UTC
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
34logger = get_logger(__name__)
37class AuthMiddleware:
38 """Middleware for handling authentication and authorization - Pure ASGI implementation."""
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
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()
60 self.attempt_tracker = attempt_tracker
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 )
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)
78 def extract_token(self, request: Request) -> str | None:
79 """Extract authentication token from request"""
80 return self.token_extractor.extract_token(request)
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
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
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)
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)
115 # Try session authentication
116 user = await self.session_authenticator.authenticate(request)
117 if user:
118 return cast("User | None", user)
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)
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 )
135 return cast("User | None", user)
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)
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
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
159 request = StarletteRequest(scope, receive)
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
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
176 # Authenticate request
177 user = await self.authenticate_request(request)
179 # Store user in scope (Starlette style)
180 scope["user"] = user
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
188 # Store user in runtime context for unified access across HTTP/WebSocket/Tasks
189 from lexigram.di.resolution.context import get_resolver
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))
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.
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)
207 # Temporarily set app and call
208 self.rate_limiter.app = next_app
209 await self.rate_limiter(scope, receive, send)
210 return
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
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
235 # Continue with request - capture response to add headers
236 response_started = False
237 response_headers = []
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)
246 await self.app(scope, receive, send_wrapper)
248 def set_app(self, app: Callable) -> None:
249 """Set the ASGI app to wrap."""
250 self.app = app
253# RateLimitMiddleware has been moved to lexigram.auth.web.middleware.throttle
256class AuthRouter:
257 """Router extension with authentication helpers"""
259 def __init__(self, auth_provider: AuthProviderProtocol):
260 self.auth_provider = auth_provider
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"""
270 def decorator(func: Callable) -> Callable:
271 @wraps(func)
272 async def wrapper(*args: Any, **kwargs: Any) -> Any:
273 request = _extract_request(*args, **kwargs)
275 if not request:
276 raise ValueError(
277 "Could not find request object in function arguments",
278 )
280 # Check authentication
281 user = getattr(request.state, "user", None)
282 if not optional and not user:
283 rf = await _get_response_factory(request)
285 return rf.json(
286 status_code=401,
287 content={
288 "error": "unauthorized",
289 "message": "Authentication required",
290 },
291 )
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)
298 return rf.json(
299 status_code=403,
300 content={
301 "error": "forbidden",
302 "message": "Insufficient roles",
303 },
304 )
306 if permissions and not self.auth_provider.has_any_permission(
307 user,
308 permissions,
309 ):
310 rf = await _get_response_factory(request)
312 return rf.json(
313 status_code=403,
314 content={
315 "error": "forbidden",
316 "message": "Insufficient permissions",
317 },
318 )
320 return await func(*args, **kwargs)
322 return wrapper
324 return decorator
326 def get_current_user(self, request: Request) -> User | None:
327 """Get current authenticated user from request"""
328 return getattr(request.state, "user", None)
331# Convenience helpers and functions for common auth patterns
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")
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
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 )
353 return cast("AuthProviderProtocol", await resolver.resolve(AuthProviderProtocol))
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
361 resolver = get_resolver(context)
362 if not resolver:
363 return None
365 return await resolver.resolve(ResponseFactoryProtocol)
368def require_mfa(
369 max_age_seconds: int = 300,
370) -> Callable[[Callable[..., Any]], Callable[..., Any]]:
371 """Decorator to require MFA verification (step-up).
373 Ensures:
374 1. User has MFA enabled.
375 2. Session was verified with MFA within the last `max_age_seconds`.
376 """
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)
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"})
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)
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 )
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
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
417 if not is_verified:
418 rf = await _get_response_factory(request)
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 )
429 return await func(*args, **kwargs)
431 return wrapper
433 return decorator
436__all__ = [
437 "AuthMiddleware",
438 "AuthMiddlewareConfig",
439 "AuthRouter",
440 "logger",
441 "optional_auth",
442 "require_mfa",
443 "require_permissions",
444 "require_roles",
445]