1use std::{convert::Infallible, fmt::Debug, future::Future, marker::PhantomData, pin::Pin};
31
32use axum::{
33 extract::{FromRequestParts, Request},
34 http::{StatusCode, request::Parts},
35 response::{IntoResponse, Response},
36 routing::MethodRouter,
37};
38use tower::{Layer, Service};
39
40#[cfg(any(feature = "jwt", feature = "cookie", feature = "basic-auth"))]
41use crate::session::Session;
42
43pub struct RbacLayer<R: RBAC> {
47 required: AuthType<R>,
48}
49
50impl<R: RBAC> Clone for RbacLayer<R> {
51 fn clone(&self) -> Self {
52 RbacLayer {
53 required: self.required.clone(),
54 }
55 }
56}
57
58impl<R: RBAC, S> Layer<S> for RbacLayer<R> {
59 type Service = RbacService<R, S>;
60
61 fn layer(&self, inner: S) -> RbacService<R, S> {
62 RbacService {
63 required: self.required.clone(),
64 inner,
65 }
66 }
67}
68
69pub struct RbacService<R: RBAC, S> {
71 required: AuthType<R>,
72 inner: S,
73}
74
75impl<R: RBAC, S: Clone> Clone for RbacService<R, S> {
76 fn clone(&self) -> Self {
77 RbacService {
78 required: self.required.clone(),
79 inner: self.inner.clone(),
80 }
81 }
82}
83
84type BoxFuture<T> = Pin<Box<dyn Future<Output = T> + Send>>;
85
86impl<R, S> Service<Request> for RbacService<R, S>
87where
88 R: RBAC + 'static,
89 S: Service<Request, Response = Response> + Clone + Send + 'static,
90 S::Error: Send,
91 S::Future: Send,
92{
93 type Response = Response;
94 type Error = S::Error;
95 type Future = BoxFuture<Result<Response, S::Error>>;
96
97 fn poll_ready(
98 &mut self,
99 cx: &mut std::task::Context<'_>,
100 ) -> std::task::Poll<Result<(), Self::Error>> {
101 self.inner.poll_ready(cx)
102 }
103
104 fn call(&mut self, mut req: Request) -> Self::Future {
105 let required = self.required.clone();
106 let mut inner = self.inner.clone();
107 Box::pin(async move {
108 let Some(session) = Session::<R::Resource>::from_extensions(req.extensions_mut())
109 else {
110 return Ok(StatusCode::UNAUTHORIZED.into_response());
111 };
112
113 let user_roles: Vec<R> = R::extract_roles(&session).into_iter().copied().collect();
114
115 let ok = match &required {
116 AuthType::RequiresAll(roles) => roles.iter().all(|r| user_roles.contains(r)),
117 AuthType::Allows(roles) => user_roles.iter().any(|r| roles.contains(r)),
118 };
119
120 if !ok {
121 return Ok(StatusCode::FORBIDDEN.into_response());
122 }
123
124 session.insert_into(req.extensions_mut());
125 inner.call(req).await
126 })
127 }
128}
129
130pub trait RBAC: Send + Sync + 'static + Clone + Eq + Copy + Debug {
133 type Resource: Clone + Send + Sync + 'static;
135
136 fn extract_roles(resource: &Self::Resource) -> impl IntoIterator<Item = &Self>;
138}
139
140#[derive(Clone)]
141enum AuthType<T: RBAC> {
142 RequiresAll(Vec<T>),
143 Allows(Vec<T>),
144}
145
146pub trait RBACExt {
148 fn requires<T: RBAC>(self, role: T) -> Self;
150 fn requires_all<T: RBAC>(self, roles: impl Into<Vec<T>>) -> Self;
152 fn allows<T: RBAC>(self, roles: impl Into<Vec<T>>) -> Self;
154}
155
156impl<S: Clone + 'static> RBACExt for MethodRouter<S, Infallible> {
157 fn requires<T: RBAC>(self, role: T) -> Self {
158 self.layer(RbacLayer {
159 required: AuthType::RequiresAll(vec![role]),
160 })
161 }
162
163 fn requires_all<T: RBAC>(self, roles: impl Into<Vec<T>>) -> Self {
164 self.layer(RbacLayer {
165 required: AuthType::RequiresAll(roles.into()),
166 })
167 }
168
169 fn allows<T: RBAC>(self, roles: impl Into<Vec<T>>) -> Self {
170 self.layer(RbacLayer {
171 required: AuthType::Allows(roles.into()),
172 })
173 }
174}
175
176#[doc(hidden)]
178pub struct RolesExtractor<T: RBAC> {
179 pub roles: Vec<T>,
180 _p: PhantomData<T>,
181}
182
183#[cfg(any(feature = "jwt", feature = "cookie", feature = "basic-auth"))]
184fn extract_roles__<R: RBAC + Copy>(parts: &mut Parts) -> Option<Vec<R>> {
185 let session = Session::<R::Resource>::from_extensions(&mut parts.extensions)?;
186 let roles: Vec<R> = R::extract_roles(&session).into_iter().copied().collect();
187 session.insert_into(&mut parts.extensions);
188 Some(roles)
189}
190
191#[cfg(not(any(feature = "jwt", feature = "cookie", feature = "basic-auth")))]
192fn extract_roles__<R: RBAC + Copy>(parts: &mut Parts) -> Option<Vec<R>> {
193 None
194}
195
196impl<S: Send + Sync, R: RBAC + Copy> FromRequestParts<S> for RolesExtractor<R> {
197 type Rejection = StatusCode;
198
199 async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
200 let Some(roles) = extract_roles__(parts) else {
201 return Err(StatusCode::UNAUTHORIZED);
202 };
203
204 Ok(RolesExtractor {
205 roles,
206 _p: PhantomData,
207 })
208 }
209}
210
211pub fn __requires<T: RBAC>(resource: RolesExtractor<T>, roles: &[T]) -> Option<Response> {
212 if roles.iter().all(|r| resource.roles.contains(r)) {
213 None
214 } else {
215 Some(StatusCode::FORBIDDEN.into_response())
216 }
217}
218
219pub fn __allows<T: RBAC>(resource: RolesExtractor<T>, roles: &[T]) -> Option<Response> {
220 if resource.roles.iter().any(|r| roles.contains(r)) {
221 None
222 } else {
223 Some(StatusCode::FORBIDDEN.into_response())
224 }
225}
226
227#[cfg(feature = "macros")]
228pub use axum_security_macros::{allows, requires};
229
230#[doc(hidden)]
231pub mod __private {
232 pub use super::__allows;
233 pub use super::__requires;
234 pub use super::RolesExtractor;
235}
236
237#[cfg(test)]
238mod tests {
239 use super::*;
240
241 #[derive(Clone, Copy, Eq, PartialEq, Debug)]
242 enum Role {
243 Admin,
244 Mod,
245 #[allow(dead_code)]
246 User,
247 }
248
249 #[derive(Clone)]
250 struct FakeUser {
251 roles: Vec<Role>,
252 }
253
254 impl RBAC for Role {
255 type Resource = FakeUser;
256 fn extract_roles(r: &FakeUser) -> impl IntoIterator<Item = &Role> {
257 &r.roles
258 }
259 }
260
261 fn make_extractor(roles: Vec<Role>) -> RolesExtractor<Role> {
262 RolesExtractor {
263 roles,
264 _p: PhantomData,
265 }
266 }
267
268 #[test]
269 fn requires_exact_match() {
270 let ext = make_extractor(vec![Role::Admin]);
271 assert!(
272 __requires(ext, &[Role::Admin]).is_none(),
273 "Admin with required=[Admin] should pass"
274 );
275 }
276
277 #[test]
278 fn requires_missing_role() {
279 let ext = make_extractor(vec![Role::Admin]);
280 assert!(
281 __requires(ext, &[Role::Admin, Role::Mod]).is_some(),
282 "Admin with required=[Admin, Mod] should fail"
283 );
284 }
285
286 #[test]
287 fn requires_superset_passes() {
288 let ext = make_extractor(vec![Role::Admin, Role::Mod]);
289 assert!(
290 __requires(ext, &[Role::Admin]).is_none(),
291 "[Admin, Mod] with required=[Admin] should pass"
292 );
293 }
294
295 #[test]
296 fn allows_match() {
297 let ext = make_extractor(vec![Role::Mod]);
298 assert!(
299 __allows(ext, &[Role::Admin, Role::Mod]).is_none(),
300 "Mod with any=[Admin, Mod] should pass"
301 );
302 }
303
304 #[test]
305 fn allows_no_match() {
306 let ext = make_extractor(vec![Role::User]);
307 assert!(
308 __allows(ext, &[Role::Admin, Role::Mod]).is_some(),
309 "User with any=[Admin, Mod] should fail"
310 );
311 }
312}