use std::time::{Duration, SystemTime, UNIX_EPOCH};
use axum_core::extract::Request;
use base64::Engine;
use crate::{config::DEFAULT_TOKEN_TTL, error::OidcError};
pub fn extract_bearer_token(req: &Request) -> Result<String, OidcError> {
req.headers()
.get("authorization")
.and_then(|val| val.to_str().ok())
.and_then(|val| {
if val.starts_with("Bearer ") {
Some(val.trim_start_matches("Bearer ").to_string())
} else {
None
}
})
.ok_or(OidcError::MissingToken)
}
pub fn extract_token_ttl(token: &str) -> Duration {
let parts: Vec<&str> = token.split('.').collect();
if parts.len() != 3 {
return DEFAULT_TOKEN_TTL;
}
let Ok(payload) = base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(parts[1]) else {
return DEFAULT_TOKEN_TTL;
};
let Ok(claims) = serde_json::from_slice::<serde_json::Value>(&payload) else {
return DEFAULT_TOKEN_TTL;
};
claims.get("exp").and_then(serde_json::Value::as_u64).map_or_else(
|| DEFAULT_TOKEN_TTL,
|exp| {
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
if exp > now {
Duration::from_secs(exp - now)
} else {
Duration::from_secs(0) }
},
)
}
#[cfg(test)]
mod tests {
use super::*;
use axum_core::body::Body;
use http::{HeaderMap, HeaderValue};
#[test]
fn test_extract_bearer_token_success() {
let mut headers = HeaderMap::new();
headers.insert("authorization", HeaderValue::from_static("Bearer test-token"));
let req = Request::builder()
.body(Body::empty())
.unwrap()
.into_parts().0;
let mut req = Request::from_parts(req, Body::empty());
*req.headers_mut() = headers;
let result = extract_bearer_token(&req);
assert!(result.is_ok());
assert_eq!(result.unwrap(), "test-token");
}
#[test]
fn test_extract_bearer_token_missing() {
let req = Request::builder()
.body(Body::empty())
.unwrap();
let result = extract_bearer_token(&req);
assert!(matches!(result, Err(OidcError::MissingToken)));
}
#[test]
fn test_extract_token_ttl_invalid_format() {
let result = extract_token_ttl("invalid-token");
assert_eq!(result, DEFAULT_TOKEN_TTL);
}
}