use core::fmt;
use crate::util::log::{debug, warn};
const BEARER_PREFIX: &str = "Bearer ";
const MAX_HEADER_LEN: usize = 4096;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum BearerErrorKind {
MissingPrefix,
EmptyToken,
InvalidCharacters,
TooLong,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct BearerError {
kind: BearerErrorKind,
}
impl BearerError {
#[must_use]
pub fn is_missing_prefix(&self) -> bool {
self.kind == BearerErrorKind::MissingPrefix
}
#[must_use]
pub fn is_empty_token(&self) -> bool {
self.kind == BearerErrorKind::EmptyToken
}
#[must_use]
pub fn is_invalid_characters(&self) -> bool {
self.kind == BearerErrorKind::InvalidCharacters
}
#[must_use]
pub fn is_too_long(&self) -> bool {
self.kind == BearerErrorKind::TooLong
}
}
impl fmt::Display for BearerError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self.kind {
BearerErrorKind::MissingPrefix => {
write!(f, "bearer: missing \"Bearer \" prefix")
}
BearerErrorKind::EmptyToken => {
write!(f, "bearer: token is empty")
}
BearerErrorKind::InvalidCharacters => {
write!(
f,
"bearer: token contains non-printable or non-ASCII characters"
)
}
BearerErrorKind::TooLong => {
write!(
f,
"bearer: header exceeds maximum length of {MAX_HEADER_LEN} bytes"
)
}
}
}
}
impl std::error::Error for BearerError {}
pub fn extract_bearer_token(header: &str) -> Result<&str, BearerError> {
if header.len() > MAX_HEADER_LEN {
warn!("bearer: extraction failed (too long)");
return Err(BearerError {
kind: BearerErrorKind::TooLong,
});
}
let token = header.strip_prefix(BEARER_PREFIX).ok_or_else(|| {
warn!("bearer: extraction failed (missing prefix)");
BearerError {
kind: BearerErrorKind::MissingPrefix,
}
})?;
if token.is_empty() {
warn!("bearer: extraction failed (empty token)");
return Err(BearerError {
kind: BearerErrorKind::EmptyToken,
});
}
if !token.bytes().all(|b| (0x20..=0x7E).contains(&b)) {
warn!("bearer: extraction failed (invalid characters)");
return Err(BearerError {
kind: BearerErrorKind::InvalidCharacters,
});
}
debug!(len = token.len(), "bearer: token extracted");
Ok(token)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn valid_bearer_token() {
let token = extract_bearer_token("Bearer abc123").unwrap();
assert_eq!(token, "abc123");
}
#[test]
fn valid_bearer_token_with_special_chars() {
let token = extract_bearer_token("Bearer eyJhbGciOiJIUzI1NiJ9.payload.sig").unwrap();
assert_eq!(token, "eyJhbGciOiJIUzI1NiJ9.payload.sig");
}
#[test]
fn valid_bearer_token_with_spaces_in_token() {
let token = extract_bearer_token("Bearer token with spaces").unwrap();
assert_eq!(token, "token with spaces");
}
#[test]
fn valid_bearer_token_single_char() {
let token = extract_bearer_token("Bearer x").unwrap();
assert_eq!(token, "x");
}
#[test]
fn missing_prefix_entirely() {
let err = extract_bearer_token("abc123").unwrap_err();
assert!(err.is_missing_prefix());
}
#[test]
fn missing_prefix_no_space() {
let err = extract_bearer_token("Bearerabc123").unwrap_err();
assert!(err.is_missing_prefix());
}
#[test]
fn lowercase_bearer_rejected() {
let err = extract_bearer_token("bearer abc123").unwrap_err();
assert!(err.is_missing_prefix());
}
#[test]
fn uppercase_bearer_rejected() {
let err = extract_bearer_token("BEARER abc123").unwrap_err();
assert!(err.is_missing_prefix());
}
#[test]
fn empty_header() {
let err = extract_bearer_token("").unwrap_err();
assert!(err.is_missing_prefix());
}
#[test]
fn just_bearer_prefix_no_token() {
let err = extract_bearer_token("Bearer ").unwrap_err();
assert!(err.is_empty_token());
}
#[test]
fn empty_token_after_prefix() {
let err = extract_bearer_token("Bearer ").unwrap_err();
assert!(err.is_empty_token());
}
#[test]
fn token_with_null_byte() {
let err = extract_bearer_token("Bearer abc\x00def").unwrap_err();
assert!(err.is_invalid_characters());
}
#[test]
fn token_with_newline() {
let err = extract_bearer_token("Bearer abc\ndef").unwrap_err();
assert!(err.is_invalid_characters());
}
#[test]
fn token_with_tab() {
let err = extract_bearer_token("Bearer abc\tdef").unwrap_err();
assert!(err.is_invalid_characters());
}
#[test]
fn token_with_del() {
let err = extract_bearer_token("Bearer abc\x7Fdef").unwrap_err();
assert!(err.is_invalid_characters());
}
#[test]
fn token_with_non_ascii() {
let err = extract_bearer_token("Bearer caf\u{00e9}").unwrap_err();
assert!(err.is_invalid_characters());
}
#[test]
fn token_at_max_length_accepted() {
let token_len = MAX_HEADER_LEN - BEARER_PREFIX.len();
let header = format!("Bearer {}", "a".repeat(token_len));
assert_eq!(header.len(), MAX_HEADER_LEN);
let result = extract_bearer_token(&header);
assert!(result.is_ok());
assert_eq!(result.unwrap().len(), token_len);
}
#[test]
fn token_exceeding_max_length_rejected() {
let header = format!("Bearer {}", "a".repeat(MAX_HEADER_LEN));
let err = extract_bearer_token(&header).unwrap_err();
assert!(err.is_too_long());
}
#[test]
fn error_display_messages() {
let cases = [
(
BearerError {
kind: BearerErrorKind::MissingPrefix,
},
"bearer: missing \"Bearer \" prefix",
),
(
BearerError {
kind: BearerErrorKind::EmptyToken,
},
"bearer: token is empty",
),
(
BearerError {
kind: BearerErrorKind::InvalidCharacters,
},
"bearer: token contains non-printable or non-ASCII characters",
),
(
BearerError {
kind: BearerErrorKind::TooLong,
},
&format!("bearer: header exceeds maximum length of {MAX_HEADER_LEN} bytes"),
),
];
for (err, expected) in &cases {
assert_eq!(err.to_string(), *expected);
}
}
#[test]
fn error_implements_std_error() {
let err: Box<dyn std::error::Error> = Box::new(BearerError {
kind: BearerErrorKind::MissingPrefix,
});
let _ = err.to_string();
}
}