Coverage for src/lexigram/auth/authz/guards.py: 97%

77 statements  

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

1"""Authorization guards and route protection decorators. 

2 

3Provides :class:`AuthorizationGuard` and :class:`RouteGuard` for RBAC/ABAC 

4checks, and decorator helpers (:func:`require_auth`, :func:`require_roles`, 

5:func:`require_permissions`, :func:`optional_auth`) that protect route 

6handlers. 

7 

8All guards delegate actual permission evaluation to 

9:class:`~lexigram.auth.authz.service.AuthorizationService`, which is either 

10injected at construction time or created as a default instance. 

11 

12Example:: 

13 

14 from lexigram.auth.authz.guards import require_auth, require_roles 

15 

16 @require_auth(roles=["admin"]) 

17 async def admin_endpoint(request): 

18 ... 

19 

20 @require_roles("editor", "admin") 

21 async def editor_endpoint(request): 

22 ... 

23""" 

24 

25from __future__ import annotations 

26 

27from functools import wraps 

28from typing import TYPE_CHECKING, Any 

29 

30from starlette.responses import JSONResponse 

31 

32from lexigram.auth.authz.service import AuthorizationService 

33from lexigram.logging import get_logger 

34 

35if TYPE_CHECKING: 

36 from collections.abc import Callable 

37 

38 from lexigram.auth.models.user import User 

39 

40logger = get_logger(__name__) 

41 

42 

43def _find_request(args: tuple, kwargs: dict) -> Any: 

44 """Extract the first Starlette-like request from positional or keyword args.""" 

45 for arg in args: 

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

47 return arg 

48 return kwargs.get("request") 

49 

50 

51class AuthorizationGuard: 

52 """GuardProtocol that checks user roles and/or permissions. 

53 

54 Args: 

55 roles: Required role names (any match is sufficient). 

56 permissions: Required permission strings (all must match). 

57 auth_service: Optional pre-built :class:`AuthorizationService` instance. 

58 When ``None``, a default instance is created on first use. 

59 """ 

60 

61 def __init__( 

62 self, 

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

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

65 auth_service: AuthorizationService | None = None, 

66 ) -> None: 

67 self.roles: list[str] = list(roles or []) 

68 self.permissions: list[str] = list(permissions or []) 

69 self._auth_service: AuthorizationService | None = auth_service 

70 

71 async def check_authorization(self, user: User | None) -> bool: 

72 """Return ``True`` if *user* satisfies all role and permission requirements. 

73 

74 Args: 

75 user: The authenticated user, or ``None`` for anonymous requests. 

76 

77 Returns: 

78 ``False`` if user is ``None`` or inactive; delegates to 

79 :class:`~lexigram.auth.authz.service.AuthorizationService` for 

80 actual role/permission checks. 

81 """ 

82 if user is None: 

83 return False 

84 if hasattr(user, "is_active") and not user.is_active: 

85 return False 

86 

87 service = ( 

88 self._auth_service 

89 if self._auth_service is not None 

90 else AuthorizationService() 

91 ) 

92 

93 if self.roles: 

94 if not service.has_any_role(user, list(self.roles)): 

95 return False 

96 

97 if self.permissions: 

98 for perm in self.permissions: 

99 if not await service.can(user, perm, perm): 

100 return False 

101 

102 return True 

103 

104 def get_error_message(self) -> str: 

105 """Build a human-readable description of the requirements. 

106 

107 Returns: 

108 A string listing the required roles and/or permissions. 

109 """ 

110 parts: list[str] = [] 

111 if self.roles: 

112 parts.append(f"Required roles: {', '.join(sorted(self.roles))}") 

113 if self.permissions: 

114 parts.append(f"Required permissions: {', '.join(sorted(self.permissions))}") 

115 return "; ".join(parts) if parts else "Authorization required" 

116 

117 

118class RouteGuard: 

119 """Wraps an :class:`AuthorizationGuard` and provides HTTP response helpers. 

120 

121 Args: 

122 guard: The underlying :class:`AuthorizationGuard` to delegate to. 

123 """ 

124 

125 def __init__(self, guard: AuthorizationGuard) -> None: 

126 self._guard = guard 

127 

128 async def check_access(self, user: User | None) -> bool: 

129 """Delegate to the underlying guard's ``check_authorization``. 

130 

131 Args: 

132 user: The authenticated user, or ``None``. 

133 """ 

134 return await self._guard.check_authorization(user) 

135 

136 async def get_deny_response(self) -> JSONResponse: 

137 """Build a 403 JSON response describing what was required. 

138 

139 Returns: 

140 A :class:`starlette.responses.JSONResponse` with status 403. 

141 """ 

142 return JSONResponse( 

143 {"error": "forbidden", "message": self._guard.get_error_message()}, 

144 status_code=403, 

145 ) 

146 

147 

148def require_auth( 

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

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

151 optional: bool = False, 

152) -> Callable[[Callable], Callable]: 

153 """Decorator that protects a route handler with authentication and RBAC/ABAC checks. 

154 

155 Args: 

156 roles: Required role names. Any single matching role is sufficient. 

157 permissions: Required permission strings. All must be satisfied. 

158 optional: When ``True``, allow unauthenticated requests to proceed 

159 (user will simply be ``None`` in ``request.state``). 

160 

161 Returns: 

162 A decorator that wraps the route handler. 

163 

164 Raises: 

165 ValueError: If no request object can be found in the handler's arguments. 

166 """ 

167 guard = AuthorizationGuard(roles=roles, permissions=permissions) 

168 

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

170 @wraps(func) 

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

172 request = _find_request(args, kwargs) 

173 

174 if request is None: 

175 if optional: 

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

177 raise ValueError("Could not find request object in handler arguments") 

178 

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

180 

181 if user is None: 

182 if optional: 

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

184 return JSONResponse( 

185 {"error": "unauthorized", "message": "Authentication required"}, 

186 status_code=401, 

187 ) 

188 

189 # No role/permission requirements → auth alone is sufficient 

190 if not roles and not permissions: 

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

192 

193 if not await guard.check_authorization(user): 

194 return JSONResponse( 

195 {"error": "forbidden", "message": guard.get_error_message()}, 

196 status_code=403, 

197 ) 

198 

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

200 

201 return wrapper 

202 

203 return decorator 

204 

205 

206def require_roles(*roles: str) -> Callable[[Callable], Callable]: 

207 """Shorthand decorator requiring the user to have at least one of *roles*. 

208 

209 Args: 

210 *roles: Role names that are acceptable. 

211 

212 Returns: 

213 A decorator equivalent to ``require_auth(roles=list(roles))``. 

214 """ 

215 return require_auth(roles=list(roles)) 

216 

217 

218def require_permissions(*permissions: str) -> Callable[[Callable], Callable]: 

219 """Shorthand decorator requiring the user to hold all of *permissions*. 

220 

221 Args: 

222 *permissions: Permission strings that must all be satisfied. 

223 

224 Returns: 

225 A decorator equivalent to ``require_auth(permissions=list(permissions))``. 

226 """ 

227 return require_auth(permissions=list(permissions)) 

228 

229 

230def optional_auth(func: Callable) -> Callable: 

231 """Decorator that attaches optional auth: user is populated if present, never blocked. 

232 

233 Can be used directly without invocation (unlike ``require_auth`` which always 

234 requires parentheses):: 

235 

236 @optional_auth 

237 async def public_endpoint(request): 

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

239 ... 

240 

241 Args: 

242 func: The async route handler to wrap. 

243 

244 Returns: 

245 The wrapped handler which always passes through. 

246 """ 

247 

248 @wraps(func) 

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

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

251 

252 return wrapper 

253 

254 

255__all__ = [ 

256 "AuthorizationGuard", 

257 "RouteGuard", 

258 "optional_auth", 

259 "require_auth", 

260 "require_permissions", 

261 "require_roles", 

262]