Skip to main content

shardline_server/
auth.rs

1use std::fmt;
2use std::sync::Arc;
3
4use axum::http::{HeaderMap, header::AUTHORIZATION};
5use shardline_protocol::{MAX_TOKEN_STRING_BYTES, TokenClaims, TokenCodecError, TokenScope};
6use shardline_server_core::{AuthError, AuthProvider};
7use subtle::ConstantTimeEq;
8
9use crate::ServerError;
10
11/// Verified request authorization context.
12#[derive(Debug, Clone, PartialEq, Eq)]
13pub struct AuthContext {
14    claims: TokenClaims,
15}
16
17impl AuthContext {
18    /// Creates an authorization context from verified token claims.
19    #[must_use]
20    pub const fn new(claims: TokenClaims) -> Self {
21        Self { claims }
22    }
23
24    /// Returns the verified claims.
25    #[must_use]
26    pub const fn claims(&self) -> &TokenClaims {
27        &self.claims
28    }
29}
30
31/// Bearer-token verifier backed by a pluggable [`AuthProvider`].
32#[derive(Clone)]
33pub struct ServerAuth {
34    provider: Arc<dyn AuthProvider>,
35}
36
37impl ServerAuth {
38    /// Creates a bearer-token verifier from a signing key using the local
39    /// HMAC-SHA256 provider.
40    ///
41    /// # Errors
42    ///
43    /// Returns [`ServerError`] when the signing key is invalid.
44    pub fn new(signing_key: &[u8]) -> Result<Self, ServerError> {
45        let provider = shardline_server_core::auth::LocalHmacProvider::new(signing_key)?;
46        Ok(Self {
47            provider: Arc::new(provider),
48        })
49    }
50
51    /// Creates a bearer-token verifier from a boxed [`AuthProvider`].
52    #[must_use]
53    pub fn from_provider(provider: Box<dyn AuthProvider>) -> Self {
54        Self {
55            provider: Arc::from(provider),
56        }
57    }
58
59    /// Returns a reference to the underlying [`AuthProvider`].
60    #[must_use]
61    pub fn provider(&self) -> &dyn AuthProvider {
62        self.provider.as_ref()
63    }
64
65    /// Returns a clone of the underlying [`AuthProvider`] as an [`Arc`].
66    #[must_use]
67    pub fn provider_arc(&self) -> Arc<dyn AuthProvider> {
68        self.provider.clone()
69    }
70
71    /// Validates the request token and required scope.
72    ///
73    /// # Errors
74    ///
75    /// Returns [`ServerError`] when the authorization header is missing, malformed, or
76    /// insufficient for the requested scope.
77    pub fn authorize(
78        &self,
79        headers: &HeaderMap,
80        required_scope: TokenScope,
81    ) -> Result<AuthContext, ServerError> {
82        let header = headers
83            .get(AUTHORIZATION)
84            .ok_or(ServerError::MissingAuthorization)?;
85        let header = header
86            .to_str()
87            .map_err(|_error| ServerError::InvalidAuthorizationHeader)?;
88        let token = parse_bearer_token(header)?;
89        let claims = self.provider.verify_token(token)?;
90        if !scope_allows(claims.scope(), required_scope) {
91            return Err(ServerError::InsufficientScope);
92        }
93
94        Ok(AuthContext::new(claims))
95    }
96}
97
98impl fmt::Debug for ServerAuth {
99    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
100        f.debug_struct("ServerAuth")
101            .field("provider", &"<dyn AuthProvider>")
102            .finish()
103    }
104}
105
106fn parse_bearer_token(header: &str) -> Result<&str, ServerError> {
107    let Some(token) = header.strip_prefix("Bearer ") else {
108        return Err(ServerError::InvalidAuthorizationHeader);
109    };
110    if token.trim().is_empty() {
111        return Err(ServerError::InvalidAuthorizationHeader);
112    }
113    if token.len() > MAX_TOKEN_STRING_BYTES {
114        return Err(ServerError::InvalidAuthorizationHeader);
115    }
116    if token.bytes().any(|byte| byte.is_ascii_whitespace()) {
117        return Err(ServerError::InvalidAuthorizationHeader);
118    }
119
120    Ok(token)
121}
122
123pub(crate) fn authorize_static_bearer_token(
124    headers: &HeaderMap,
125    expected_token: &[u8],
126) -> Result<(), ServerError> {
127    let header = headers
128        .get(AUTHORIZATION)
129        .ok_or(ServerError::MissingAuthorization)?;
130    let header = header
131        .to_str()
132        .map_err(|_error| ServerError::InvalidAuthorizationHeader)?;
133    let token = parse_bearer_token(header)?;
134    let actual = token.as_bytes();
135
136    use sha2::{Digest, Sha256};
137    let actual_hash = Sha256::digest(actual);
138    let expected_hash = Sha256::digest(expected_token);
139    if bool::from(actual_hash.ct_eq(&expected_hash)) {
140        return Ok(());
141    }
142
143    Err(ServerError::InvalidAuthorizationHeader)
144}
145
146const fn scope_allows(actual_scope: TokenScope, required_scope: TokenScope) -> bool {
147    match required_scope {
148        TokenScope::Read => actual_scope.allows_read(),
149        TokenScope::Write => actual_scope.allows_write(),
150    }
151}
152
153impl From<TokenCodecError> for ServerError {
154    fn from(error: TokenCodecError) -> Self {
155        Self::InvalidToken(error)
156    }
157}
158
159impl From<AuthError> for ServerError {
160    fn from(error: AuthError) -> Self {
161        match error {
162            AuthError::InvalidToken => Self::InvalidToken(TokenCodecError::InvalidFormat),
163            AuthError::ExpiredToken => Self::InvalidToken(TokenCodecError::Expired),
164            AuthError::InsufficientScope => Self::InsufficientScope,
165            AuthError::ProviderError(msg) => Self::SigningKeyError(msg),
166        }
167    }
168}
169
170#[cfg(test)]
171mod tests {
172    use axum::http::{
173        HeaderMap,
174        header::{AUTHORIZATION, HeaderValue},
175    };
176    use shardline_protocol::{
177        RepositoryProvider, RepositoryScope, TokenClaims, TokenScope, TokenSigner,
178    };
179
180    use super::{MAX_TOKEN_STRING_BYTES, ServerAuth, authorize_static_bearer_token};
181    use crate::ServerError;
182
183    #[test]
184    fn server_auth_rejects_missing_header() {
185        let auth = ServerAuth::new(b"test-signing-key-32-bytes-long!!");
186        assert!(auth.is_ok());
187        let Ok(auth) = auth else {
188            return;
189        };
190
191        assert!(matches!(
192            auth.authorize(&HeaderMap::new(), TokenScope::Read),
193            Err(ServerError::MissingAuthorization)
194        ));
195    }
196
197    #[test]
198    fn server_auth_rejects_insufficient_scope() {
199        let auth = ServerAuth::new(b"test-signing-key-32-bytes-long!!");
200        assert!(auth.is_ok());
201        let Ok(auth) = auth else {
202            return;
203        };
204        let signer = TokenSigner::new(b"test-signing-key-32-bytes-long!!");
205        assert!(signer.is_ok());
206        let Ok(signer) = signer else {
207            return;
208        };
209        let repository =
210            RepositoryScope::new(RepositoryProvider::GitHub, "team", "assets", Some("main"));
211        assert!(repository.is_ok());
212        let Ok(repository) = repository else {
213            return;
214        };
215        let claims = TokenClaims::new(
216            "local",
217            "provider-user-1",
218            TokenScope::Read,
219            repository,
220            u64::MAX,
221        );
222        assert!(claims.is_ok());
223        let Ok(claims) = claims else {
224            return;
225        };
226        let token = signer.sign(&claims);
227        assert!(token.is_ok());
228        let Ok(token) = token else {
229            return;
230        };
231        let mut headers = HeaderMap::new();
232        let header_value = HeaderValue::from_str(&format!("Bearer {token}"));
233        assert!(header_value.is_ok());
234        let Ok(header_value) = header_value else {
235            return;
236        };
237        headers.insert(AUTHORIZATION, header_value);
238
239        assert!(matches!(
240            auth.authorize(&headers, TokenScope::Write),
241            Err(ServerError::InsufficientScope)
242        ));
243    }
244
245    #[test]
246    fn server_auth_rejects_oversized_bearer_token_before_decoding() {
247        let auth = ServerAuth::new(b"test-signing-key-32-bytes-long!!");
248        assert!(auth.is_ok());
249        let Ok(auth) = auth else {
250            return;
251        };
252        let token = "a".repeat(MAX_TOKEN_STRING_BYTES + 1);
253        let mut headers = HeaderMap::new();
254        let header_value = HeaderValue::from_str(&format!("Bearer {token}"));
255        assert!(header_value.is_ok());
256        let Ok(header_value) = header_value else {
257            return;
258        };
259        headers.insert(AUTHORIZATION, header_value);
260
261        assert!(matches!(
262            auth.authorize(&headers, TokenScope::Read),
263            Err(ServerError::InvalidAuthorizationHeader)
264        ));
265    }
266
267    #[test]
268    fn server_auth_rejects_bearer_token_with_whitespace() {
269        let auth = ServerAuth::new(b"test-signing-key-32-bytes-long!!");
270        assert!(auth.is_ok());
271        let Ok(auth) = auth else {
272            return;
273        };
274        let mut headers = HeaderMap::new();
275        headers.insert(
276            AUTHORIZATION,
277            HeaderValue::from_static("Bearer abc.def ghi"),
278        );
279
280        assert!(matches!(
281            auth.authorize(&headers, TokenScope::Read),
282            Err(ServerError::InvalidAuthorizationHeader)
283        ));
284    }
285
286    #[test]
287    fn server_auth_accepts_valid_write_token() {
288        let auth = ServerAuth::new(b"test-signing-key-32-bytes-long!!");
289        assert!(auth.is_ok());
290        let Ok(auth) = auth else {
291            return;
292        };
293        let signer = TokenSigner::new(b"test-signing-key-32-bytes-long!!");
294        assert!(signer.is_ok());
295        let Ok(signer) = signer else {
296            return;
297        };
298        let repository =
299            RepositoryScope::new(RepositoryProvider::GitHub, "team", "assets", Some("main"));
300        assert!(repository.is_ok());
301        let Ok(repository) = repository else {
302            return;
303        };
304        let claims = TokenClaims::new(
305            "local",
306            "provider-user-1",
307            TokenScope::Write,
308            repository,
309            u64::MAX,
310        );
311        assert!(claims.is_ok());
312        let Ok(claims) = claims else {
313            return;
314        };
315        let token = signer.sign(&claims);
316        assert!(token.is_ok());
317        let Ok(token) = token else {
318            return;
319        };
320        let mut headers = HeaderMap::new();
321        let header_value = HeaderValue::from_str(&format!("Bearer {token}"));
322        assert!(header_value.is_ok());
323        let Ok(header_value) = header_value else {
324            return;
325        };
326        headers.insert(AUTHORIZATION, header_value);
327
328        let context = auth.authorize(&headers, TokenScope::Read);
329
330        assert!(context.is_ok());
331        let Ok(context) = context else {
332            return;
333        };
334        assert_eq!(context.claims().subject(), "provider-user-1");
335        assert_eq!(context.claims().scope(), TokenScope::Write);
336    }
337
338    #[test]
339    fn static_bearer_token_rejects_missing_header() {
340        let result = authorize_static_bearer_token(&HeaderMap::new(), b"metrics-token");
341
342        assert!(matches!(result, Err(ServerError::MissingAuthorization)));
343    }
344
345    #[test]
346    fn static_bearer_token_rejects_wrong_value() {
347        let mut headers = HeaderMap::new();
348        headers.insert(
349            AUTHORIZATION,
350            HeaderValue::from_static("Bearer wrong-token"),
351        );
352
353        let result = authorize_static_bearer_token(&headers, b"metrics-token");
354
355        assert!(matches!(
356            result,
357            Err(ServerError::InvalidAuthorizationHeader)
358        ));
359    }
360
361    #[test]
362    fn static_bearer_token_accepts_matching_value() {
363        let mut headers = HeaderMap::new();
364        headers.insert(
365            AUTHORIZATION,
366            HeaderValue::from_static("Bearer metrics-token"),
367        );
368
369        let result = authorize_static_bearer_token(&headers, b"metrics-token");
370
371        assert!(result.is_ok());
372    }
373
374    // ── parse_bearer_token edge cases ──────────────────────────────────────
375
376    #[test]
377    fn parse_bearer_token_rejects_missing_bearer_prefix() {
378        use super::parse_bearer_token;
379        let result = parse_bearer_token("Basic token");
380        assert!(matches!(
381            result,
382            Err(ServerError::InvalidAuthorizationHeader)
383        ));
384    }
385
386    #[test]
387    fn parse_bearer_token_rejects_empty_token_after_prefix() {
388        use super::parse_bearer_token;
389        let result = parse_bearer_token("Bearer ");
390        assert!(matches!(
391            result,
392            Err(ServerError::InvalidAuthorizationHeader)
393        ));
394    }
395
396    #[test]
397    fn parse_bearer_token_rejects_whitespace_only_token() {
398        use super::parse_bearer_token;
399        let result = parse_bearer_token("Bearer   ");
400        assert!(matches!(
401            result,
402            Err(ServerError::InvalidAuthorizationHeader)
403        ));
404    }
405
406    #[test]
407    fn parse_bearer_token_rejects_token_with_whitespace() {
408        use super::parse_bearer_token;
409        let result = parse_bearer_token("Bearer abc def");
410        assert!(matches!(
411            result,
412            Err(ServerError::InvalidAuthorizationHeader)
413        ));
414    }
415
416    #[test]
417    fn parse_bearer_token_rejects_oversized_token() {
418        use super::{MAX_TOKEN_STRING_BYTES, parse_bearer_token};
419        let large = "a".repeat(MAX_TOKEN_STRING_BYTES + 1);
420        let header = format!("Bearer {large}");
421        let result = parse_bearer_token(&header);
422        assert!(matches!(
423            result,
424            Err(ServerError::InvalidAuthorizationHeader)
425        ));
426    }
427
428    #[test]
429    fn parse_bearer_token_accepts_valid_token() {
430        use super::parse_bearer_token;
431        let result = parse_bearer_token("Bearer valid-token-here");
432        assert!(result.is_ok());
433        assert_eq!(result.unwrap(), "valid-token-here");
434    }
435
436    // ── scope_allows ───────────────────────────────────────────────────────
437
438    #[test]
439    fn scope_allows_read_when_scope_is_read() {
440        assert!(super::scope_allows(TokenScope::Read, TokenScope::Read));
441    }
442
443    #[test]
444    fn scope_allows_write_when_scope_is_write() {
445        assert!(super::scope_allows(TokenScope::Write, TokenScope::Write));
446    }
447
448    #[test]
449    fn scope_allows_read_when_scope_is_write() {
450        // Write scope implicitly allows Read
451        assert!(super::scope_allows(TokenScope::Write, TokenScope::Read));
452    }
453
454    #[test]
455    fn scope_allows_rejects_write_when_scope_is_read() {
456        assert!(!super::scope_allows(TokenScope::Read, TokenScope::Write));
457    }
458
459    // ── AuthError conversion ───────────────────────────────────────────────
460
461    #[test]
462    fn from_auth_error_invalid_token() {
463        use shardline_server_core::AuthError;
464        let err: ServerError = AuthError::InvalidToken.into();
465        assert!(matches!(err, ServerError::InvalidToken(_)));
466    }
467
468    #[test]
469    fn from_auth_error_expired_token() {
470        use shardline_server_core::AuthError;
471        let err: ServerError = AuthError::ExpiredToken.into();
472        assert!(matches!(err, ServerError::InvalidToken(_)));
473    }
474
475    #[test]
476    fn from_auth_error_insufficient_scope() {
477        use shardline_server_core::AuthError;
478        let err: ServerError = AuthError::InsufficientScope.into();
479        assert!(matches!(err, ServerError::InsufficientScope));
480    }
481
482    #[test]
483    fn from_auth_error_provider_error() {
484        use shardline_server_core::AuthError;
485        let err: ServerError = AuthError::ProviderError("msg".to_owned()).into();
486        assert!(matches!(err, ServerError::SigningKeyError(_)));
487    }
488
489    // ── ServerAuth from_provider ───────────────────────────────────────────
490
491    #[test]
492    fn server_auth_from_provider_delegates() {
493        use shardline_server_core::auth::PassthroughProvider;
494        let provider = Box::new(PassthroughProvider);
495        let auth = ServerAuth::from_provider(provider);
496        // Verify it can authorize a request with a Bearer token
497        let mut headers = HeaderMap::new();
498        headers.insert(AUTHORIZATION, HeaderValue::from_static("Bearer any-token"));
499        let result = auth.authorize(&headers, TokenScope::Write);
500        assert!(result.is_ok());
501        let ctx = result.unwrap();
502        // PassthroughProvider uses "anonymous" as the default subject
503        assert!(ctx.claims().subject() == "anonymous" || ctx.claims().subject() == "passthrough");
504    }
505
506    #[test]
507    fn server_auth_debug_redacts_provider() {
508        use shardline_server_core::auth::PassthroughProvider;
509        let provider = Box::new(PassthroughProvider);
510        let auth = ServerAuth::from_provider(provider);
511        let debug = format!("{auth:?}");
512        assert!(!debug.contains("PassthroughProvider"));
513        assert!(debug.contains("<dyn AuthProvider>"));
514    }
515
516    #[test]
517    fn server_auth_provider_arc_returns_cloneable_arc() {
518        use shardline_server_core::auth::PassthroughProvider;
519        let provider = Box::new(PassthroughProvider);
520        let auth = ServerAuth::from_provider(provider);
521        let arc = auth.provider_arc();
522        // Verify the arc points to the same provider
523        assert!(std::sync::Arc::ptr_eq(&auth.provider, &arc));
524    }
525
526    // ── Repeated Authorization header tests ────────────────────────────────
527
528    #[test]
529    fn authorize_picks_first_of_two_separate_authorization_headers() {
530        // When a client sends two separate Authorization headers, `HeaderMap::get()`
531        // returns the first one.  This test validates that behavior.
532        let auth = ServerAuth::new(b"test-signing-key-32-bytes-long!!").unwrap();
533
534        // Create two Authorization headers via `append`.
535        let mut headers = HeaderMap::new();
536        // A valid token is appended first.
537        let signer = TokenSigner::new(b"test-signing-key-32-bytes-long!!").unwrap();
538        let repository =
539            RepositoryScope::new(RepositoryProvider::GitHub, "team", "assets", Some("main"))
540                .unwrap();
541        let claims = TokenClaims::new(
542            "local",
543            "provider-user-1",
544            TokenScope::Write,
545            repository,
546            u64::MAX,
547        )
548        .unwrap();
549        let valid_token = signer.sign(&claims).unwrap();
550
551        headers.append(
552            AUTHORIZATION,
553            HeaderValue::from_str(&format!("Bearer {valid_token}")).unwrap(),
554        );
555        // Append a second (invalid) header — should be ignored.
556        headers.append(
557            AUTHORIZATION,
558            HeaderValue::from_static("Bearer invalid-token-here"),
559        );
560
561        // Headers.get() returns the first entry — the valid token.
562        let result = auth.authorize(&headers, TokenScope::Read);
563        assert!(
564            result.is_ok(),
565            "first Authorization header should be used, got: {result:?}"
566        );
567    }
568
569    #[test]
570    fn authorize_rejects_comma_separated_bearer_in_one_header() {
571        // RFC 7230 §3.2.2 allows combining multiple header values into a single
572        // comma-separated value.  If a client sends
573        // `Authorization: Bearer token1, Bearer token2`, the token after
574        // comma + space contains whitespace and MUST be rejected.
575        let auth = ServerAuth::new(b"test-signing-key-32-bytes-long!!").unwrap();
576        let mut headers = HeaderMap::new();
577        headers.insert(
578            AUTHORIZATION,
579            HeaderValue::from_static("Bearer valid-token, Bearer invalid-token"),
580        );
581
582        let result = auth.authorize(&headers, TokenScope::Read);
583        assert!(
584            matches!(result, Err(ServerError::InvalidAuthorizationHeader)),
585            "comma-separated Bearer tokens should be rejected, got: {result:?}"
586        );
587    }
588
589    #[test]
590    fn parse_bearer_token_rejects_comma_space_in_token() {
591        // Even without the second "Bearer" prefix, a comma + space triggers the
592        // whitespace rejection in parse_bearer_token.
593        use super::parse_bearer_token;
594        let result = parse_bearer_token("Bearer token1, token2");
595        assert!(matches!(
596            result,
597            Err(ServerError::InvalidAuthorizationHeader)
598        ));
599    }
600
601    // ── ServerAuth authorize_token method ──────────────────────────────────
602
603    #[test]
604    fn server_auth_authorize_token_with_no_matching_token_provider() {
605        // When auth is built with Local provider and no token is in the header,
606        // it should go through the provider_path.
607        use shardline_server_core::auth::PassthroughProvider;
608        let provider = Box::new(PassthroughProvider);
609        let auth = super::ServerAuth::from_provider(provider);
610
611        // Empty headers → MissingAuthorization
612        let result = auth.authorize(&HeaderMap::new(), TokenScope::Read);
613        assert!(matches!(result, Err(ServerError::MissingAuthorization)));
614    }
615
616    #[test]
617    fn server_auth_authorize_token_with_invalid_scheme() {
618        use shardline_server_core::auth::PassthroughProvider;
619        let provider = Box::new(PassthroughProvider);
620        let auth = super::ServerAuth::from_provider(provider);
621
622        let mut headers = HeaderMap::new();
623        headers.insert(AUTHORIZATION, HeaderValue::from_static("Basic token"));
624        let result = auth.authorize(&headers, TokenScope::Read);
625        assert!(matches!(
626            result,
627            Err(ServerError::InvalidAuthorizationHeader)
628        ));
629    }
630
631    // ── scope_allows edge cases ────────────────────────────────────────────
632
633    #[test]
634    fn scope_allows_with_same_scope_read_read() {
635        assert!(super::scope_allows(TokenScope::Read, TokenScope::Read));
636    }
637
638    #[test]
639    fn scope_allows_with_same_scope_write_write() {
640        assert!(super::scope_allows(TokenScope::Write, TokenScope::Write));
641    }
642
643    #[test]
644    fn scope_allows_write_grants_read() {
645        assert!(super::scope_allows(TokenScope::Write, TokenScope::Read));
646    }
647
648    #[test]
649    fn scope_allows_read_denies_write() {
650        assert!(!super::scope_allows(TokenScope::Read, TokenScope::Write));
651    }
652}