Skip to main content

axum_security/rbac/
mod.rs

1//! Role-based access control (RBAC).
2//!
3//! Define a role enum, implement [`RBAC`] to extract roles from your user type,
4//! then protect routes using the [`RBACExt`] trait on [`MethodRouter`].
5//!
6//! Requires an active session (from the `jwt`, `cookie`, or `basic-auth` feature).
7//! Returns `401` if no session is present, `403` if the user lacks the required roles.
8//!
9//! # Example
10//!
11//! ```rust,ignore
12//! use axum::{Router, routing::get};
13//! use axum_security::rbac::{RBAC, RBACExt};
14//!
15//! #[derive(Clone, Copy, Debug, PartialEq, Eq)]
16//! enum Role { Admin, User }
17//!
18//! impl RBAC for Role {
19//!     type Resource = MyUser;
20//!     fn extract_roles(user: &MyUser) -> impl IntoIterator<Item = &Role> {
21//!         &user.roles
22//!     }
23//! }
24//!
25//! let app = Router::new()
26//!     .route("/admin", get(admin_handler).requires(Role::Admin))
27//!     .route("/any", get(handler).allows([Role::Admin, Role::User]));
28//! ```
29
30use 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
43/// Tower [`Layer`] that enforces role requirements on a route.
44///
45/// Prefer using the [`RBACExt`] trait methods instead of constructing this directly.
46pub 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
69/// The [`Service`] created by [`RbacLayer`]. You don't need to construct this directly.
70pub 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
130/// Trait for role types. Implement this on your role enum to define how roles
131/// are extracted from a user (the `Resource` type).
132pub trait RBAC: Send + Sync + 'static + Clone + Eq + Copy + Debug {
133    /// The user type that holds roles.
134    type Resource: Clone + Send + Sync + 'static;
135
136    /// Return the roles that `resource` has.
137    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
146/// Extension trait on [`MethodRouter`] for adding role requirements.
147pub trait RBACExt {
148    /// Require the user to have this single role.
149    fn requires<T: RBAC>(self, role: T) -> Self;
150    /// Require the user to have **all** of the given roles.
151    fn requires_all<T: RBAC>(self, roles: impl Into<Vec<T>>) -> Self;
152    /// Require the user to have **any** of the given roles.
153    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/// Extractor that provides the user's roles. Used internally by the `requires` macros.
177#[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}