1use crate::error::AuthError;
2use async_trait::async_trait;
3use http::request::Parts;
4use std::marker::PhantomData;
5
6#[async_trait]
11pub trait AuthenticationStrategy<I>: Send + Sync {
12 async fn authenticate(&self, parts: &Parts) -> Result<Option<I>, AuthError>;
19}
20
21#[async_trait]
23pub trait BasicAuthenticator: Send + Sync {
24 type Identity;
26 async fn authenticate(
28 &self,
29 username: &str,
30 password: &str,
31 ) -> Result<Option<Self::Identity>, AuthError>;
32}
33
34#[non_exhaustive]
36pub struct BasicStrategy<P, I> {
37 authenticator: P,
38 _marker: PhantomData<I>,
39}
40
41impl<P, I> BasicStrategy<P, I> {
42 pub fn new(authenticator: P) -> Self {
44 Self {
45 authenticator,
46 _marker: PhantomData,
47 }
48 }
49}
50
51#[async_trait]
52impl<P, I> AuthenticationStrategy<I> for BasicStrategy<P, I>
53where
54 P: BasicAuthenticator<Identity = I> + Send + Sync,
55 I: Send + Sync + 'static,
56{
57 async fn authenticate(&self, parts: &Parts) -> Result<Option<I>, AuthError> {
58 if let Some((username, password)) = utils::extract_basic_credentials(&parts.headers) {
59 self.authenticator.authenticate(&username, &password).await
60 } else {
61 Ok(None)
62 }
63 }
64}
65
66#[async_trait]
68pub trait TokenValidator: Send + Sync {
69 type Identity;
71 async fn validate(&self, token: &str) -> Result<Option<Self::Identity>, AuthError>;
73}
74
75#[non_exhaustive]
77pub struct TokenStrategy<V, I> {
78 validator: V,
79 _marker: PhantomData<I>,
80}
81
82impl<V, I> TokenStrategy<V, I> {
83 pub fn new(validator: V) -> Self {
85 Self {
86 validator,
87 _marker: PhantomData,
88 }
89 }
90}
91
92#[async_trait]
93impl<V, I> AuthenticationStrategy<I> for TokenStrategy<V, I>
94where
95 V: TokenValidator<Identity = I> + Send + Sync,
96 I: Send + Sync + 'static,
97{
98 async fn authenticate(&self, parts: &Parts) -> Result<Option<I>, AuthError> {
99 if let Some(token) = utils::extract_bearer_token(&parts.headers) {
100 self.validator.validate(token).await
101 } else {
102 Ok(None)
103 }
104 }
105}
106
107#[non_exhaustive]
109pub struct HeaderStrategy<F, I> {
110 header_name: http::header::HeaderName,
111 validator: F,
112 _marker: PhantomData<I>,
113}
114
115impl<F, I> HeaderStrategy<F, I> {
116 pub fn new(header_name: http::header::HeaderName, validator: F) -> Self {
118 Self {
119 header_name,
120 validator,
121 _marker: PhantomData,
122 }
123 }
124}
125
126#[async_trait]
127impl<F, I, Fut> AuthenticationStrategy<I> for HeaderStrategy<F, I>
128where
129 F: Fn(String) -> Fut + Send + Sync,
130 Fut: std::future::Future<Output = Result<Option<I>, AuthError>> + Send,
131 I: Send + Sync + 'static,
132{
133 async fn authenticate(&self, parts: &Parts) -> Result<Option<I>, AuthError> {
134 if let Some(value) = parts.headers.get(&self.header_name) {
135 if let Ok(value_str) = value.to_str() {
136 return (self.validator)(value_str.to_string()).await;
137 }
138 }
139 Ok(None)
140 }
141}
142
143#[async_trait]
145pub trait SessionProvider: Send + Sync {
146 type Identity;
148 async fn load_session(&self, session_id: &str) -> Result<Option<Self::Identity>, AuthError>;
150}
151
152#[non_exhaustive]
154pub struct SessionStrategy<P, I> {
155 provider: P,
156 cookie_name: String,
157 _marker: PhantomData<I>,
158}
159
160impl<P, I> SessionStrategy<P, I> {
161 pub fn new(provider: P, cookie_name: impl Into<String>) -> Self {
163 Self {
164 provider,
165 cookie_name: cookie_name.into(),
166 _marker: PhantomData,
167 }
168 }
169}
170
171#[async_trait]
172impl<P, I> AuthenticationStrategy<I> for SessionStrategy<P, I>
173where
174 P: SessionProvider<Identity = I> + Send + Sync,
175 I: Send + Sync + 'static,
176{
177 async fn authenticate(&self, parts: &Parts) -> Result<Option<I>, AuthError> {
178 if let Some(session_id) = utils::extract_cookie(&parts.headers, &self.cookie_name) {
179 self.provider.load_session(session_id).await
180 } else {
181 Ok(None)
182 }
183 }
184}
185
186pub mod utils {
188 use http::header::{HeaderMap, AUTHORIZATION};
189
190 pub fn extract_bearer_token(headers: &HeaderMap) -> Option<&str> {
192 headers
193 .get(AUTHORIZATION)?
194 .to_str()
195 .ok()?
196 .strip_prefix("Bearer ")
197 .map(|s| s.trim())
198 }
199
200 pub fn extract_basic_credentials(headers: &HeaderMap) -> Option<(String, String)> {
202 let auth_header = headers.get(AUTHORIZATION)?.to_str().ok()?;
203 if !auth_header.starts_with("Basic ") {
204 return None;
205 }
206 let encoded = auth_header.strip_prefix("Basic ")?.trim();
207 let decoded =
208 base64::Engine::decode(&base64::engine::general_purpose::STANDARD, encoded).ok()?;
209 let decoded_str = String::from_utf8(decoded).ok()?;
210 let mut parts = decoded_str.splitn(2, ':');
211 let username = parts.next()?.to_string();
212 let password = parts.next()?.to_string();
213 Some((username, password))
214 }
215
216 pub fn extract_cookie<'a>(headers: &'a http::HeaderMap, name: &str) -> Option<&'a str> {
218 let cookie_header = headers.get(http::header::COOKIE)?.to_str().ok()?;
219 for cookie in cookie_header.split(';') {
220 let mut parts = cookie.splitn(2, '=');
221 let k = parts.next()?.trim();
222 let v = parts.next()?.trim();
223 if k == name {
224 return Some(v);
225 }
226 }
227 None
228 }
229}
230
231#[cfg(test)]
232mod tests {
233 use super::*;
234 use http::{
235 header::{HeaderMap, HeaderName, HeaderValue, AUTHORIZATION, COOKIE},
236 Request,
237 };
238
239 #[derive(Debug, PartialEq)]
240 struct DummyIdentity(String);
241
242 struct DummyBasic;
243 #[async_trait]
244 impl BasicAuthenticator for DummyBasic {
245 type Identity = DummyIdentity;
246 async fn authenticate(
247 &self,
248 u: &str,
249 p: &str,
250 ) -> Result<Option<Self::Identity>, AuthError> {
251 if u == "user" && p == "pass" {
252 Ok(Some(DummyIdentity(u.to_string())))
253 } else {
254 Ok(None)
255 }
256 }
257 }
258
259 struct DummyToken;
260 #[async_trait]
261 impl TokenValidator for DummyToken {
262 type Identity = DummyIdentity;
263 async fn validate(&self, t: &str) -> Result<Option<Self::Identity>, AuthError> {
264 if t == "valid_token" {
265 Ok(Some(DummyIdentity("user".to_string())))
266 } else {
267 Ok(None)
268 }
269 }
270 }
271
272 struct DummySession;
273 #[async_trait]
274 impl SessionProvider for DummySession {
275 type Identity = DummyIdentity;
276 async fn load_session(&self, sid: &str) -> Result<Option<Self::Identity>, AuthError> {
277 if sid == "valid_sid" {
278 Ok(Some(DummyIdentity("user".to_string())))
279 } else {
280 Ok(None)
281 }
282 }
283 }
284
285 #[tokio::test]
286 async fn test_basic_strategy() {
287 let strategy = BasicStrategy::new(DummyBasic);
288 let mut req = Request::builder().uri("/").body(()).unwrap();
289
290 let res = strategy.authenticate(&req.into_parts().0).await.unwrap();
291 assert_eq!(res, None);
292
293 let mut req2 = Request::builder()
294 .uri("/")
295 .header(AUTHORIZATION, "Basic dXNlcjpwYXNz")
296 .body(())
297 .unwrap();
298 let res2 = strategy.authenticate(&req2.into_parts().0).await.unwrap();
299 assert_eq!(res2.unwrap(), DummyIdentity("user".to_string()));
300 }
301
302 #[tokio::test]
303 async fn test_token_strategy() {
304 let strategy = TokenStrategy::new(DummyToken);
305 let mut req = Request::builder().uri("/").body(()).unwrap();
306
307 let res = strategy.authenticate(&req.into_parts().0).await.unwrap();
308 assert_eq!(res, None);
309
310 let mut req2 = Request::builder()
311 .uri("/")
312 .header(AUTHORIZATION, "Bearer valid_token")
313 .body(())
314 .unwrap();
315 let res2 = strategy.authenticate(&req2.into_parts().0).await.unwrap();
316 assert_eq!(res2.unwrap(), DummyIdentity("user".to_string()));
317 }
318
319 #[tokio::test]
320 async fn test_header_strategy() {
321 let strategy = HeaderStrategy::new(
322 HeaderName::from_static("x-api-key"),
323 |key: String| async move {
324 if key == "secret" {
325 Ok(Some(DummyIdentity("user".to_string())))
326 } else {
327 Ok(None)
328 }
329 },
330 );
331 let mut req = Request::builder().uri("/").body(()).unwrap();
332
333 let res = strategy.authenticate(&req.into_parts().0).await.unwrap();
334 assert_eq!(res, None);
335
336 let mut req2 = Request::builder()
337 .uri("/")
338 .header("x-api-key", "secret")
339 .body(())
340 .unwrap();
341 let res2 = strategy.authenticate(&req2.into_parts().0).await.unwrap();
342 assert_eq!(res2.unwrap(), DummyIdentity("user".to_string()));
343 }
344
345 #[tokio::test]
346 async fn test_session_strategy() {
347 let strategy = SessionStrategy::new(DummySession, "sid");
348 let mut req = Request::builder().uri("/").body(()).unwrap();
349
350 let res = strategy.authenticate(&req.into_parts().0).await.unwrap();
351 assert_eq!(res, None);
352
353 let mut req2 = Request::builder()
354 .uri("/")
355 .header(COOKIE, "sid=valid_sid")
356 .body(())
357 .unwrap();
358 let res2 = strategy.authenticate(&req2.into_parts().0).await.unwrap();
359 assert_eq!(res2.unwrap(), DummyIdentity("user".to_string()));
360 }
361
362 #[test]
363 fn test_utils_extractors() {
364 let mut headers = HeaderMap::new();
365 assert_eq!(utils::extract_bearer_token(&headers), None);
366 headers.insert(AUTHORIZATION, HeaderValue::from_static("Bearer token"));
367 assert_eq!(utils::extract_bearer_token(&headers), Some("token"));
368
369 let mut headers2 = HeaderMap::new();
370 headers2.insert(
371 AUTHORIZATION,
372 HeaderValue::from_static("Basic dXNlcjpwYXNz"),
373 );
374 assert_eq!(
375 utils::extract_basic_credentials(&headers2),
376 Some(("user".to_string(), "pass".to_string()))
377 );
378
379 let mut headers3 = HeaderMap::new();
380 headers3.insert(COOKIE, HeaderValue::from_static("foo=bar; sid=123"));
381 assert_eq!(utils::extract_cookie(&headers3, "sid"), Some("123"));
382 }
383}