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#[derive(Debug, Clone, PartialEq, Eq)]
13pub struct AuthContext {
14 claims: TokenClaims,
15}
16
17impl AuthContext {
18 #[must_use]
20 pub const fn new(claims: TokenClaims) -> Self {
21 Self { claims }
22 }
23
24 #[must_use]
26 pub const fn claims(&self) -> &TokenClaims {
27 &self.claims
28 }
29}
30
31#[derive(Clone)]
33pub struct ServerAuth {
34 provider: Arc<dyn AuthProvider>,
35}
36
37impl ServerAuth {
38 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 #[must_use]
53 pub fn from_provider(provider: Box<dyn AuthProvider>) -> Self {
54 Self {
55 provider: Arc::from(provider),
56 }
57 }
58
59 #[must_use]
61 pub fn provider(&self) -> &dyn AuthProvider {
62 self.provider.as_ref()
63 }
64
65 #[must_use]
67 pub fn provider_arc(&self) -> Arc<dyn AuthProvider> {
68 self.provider.clone()
69 }
70
71 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 #[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 #[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 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 #[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 #[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 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 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 assert!(std::sync::Arc::ptr_eq(&auth.provider, &arc));
524 }
525
526 #[test]
529 fn authorize_picks_first_of_two_separate_authorization_headers() {
530 let auth = ServerAuth::new(b"test-signing-key-32-bytes-long!!").unwrap();
533
534 let mut headers = HeaderMap::new();
536 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 headers.append(
557 AUTHORIZATION,
558 HeaderValue::from_static("Bearer invalid-token-here"),
559 );
560
561 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 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 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 #[test]
604 fn server_auth_authorize_token_with_no_matching_token_provider() {
605 use shardline_server_core::auth::PassthroughProvider;
608 let provider = Box::new(PassthroughProvider);
609 let auth = super::ServerAuth::from_provider(provider);
610
611 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 #[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}