use std::ops::{Deref, DerefMut};
use base64::Engine;
use base64::engine::general_purpose::STANDARD as ENGINE;
use rama_core::extensions::Extensions;
use rama_core::telemetry::tracing;
use rama_core::username::{UsernameLabelParser, parse_username};
use rama_http_types::{HeaderName, HeaderValue};
use rama_net::user::{Basic, Bearer, RawToken, UserId};
use crate::{Error, HeaderDecode, HeaderEncode, TypedHeader};
#[derive(Clone, PartialEq, Debug)]
pub struct Authorization<C>(pub C);
impl<C> Authorization<C> {
pub fn new(credentials: C) -> Self {
Self(credentials)
}
pub fn credentials(&self) -> &C {
&self.0
}
pub fn into_inner(self) -> C {
self.0
}
}
impl<C> AsRef<C> for Authorization<C> {
fn as_ref(&self) -> &C {
&self.0
}
}
impl<C> Deref for Authorization<C> {
type Target = C;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl<C> DerefMut for Authorization<C> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.0
}
}
impl<C: Credentials> TypedHeader for Authorization<C> {
fn name() -> &'static HeaderName {
&::rama_http_types::header::AUTHORIZATION
}
}
impl<C: Credentials> HeaderDecode for Authorization<C> {
fn decode<'i, I: Iterator<Item = &'i HeaderValue>>(values: &mut I) -> Result<Self, Error> {
values
.next()
.and_then(|val| {
if C::SCHEME.is_empty() {
return C::decode(val).map(Authorization);
}
let slice = val.as_bytes();
if slice.len() > C::SCHEME.len()
&& slice[C::SCHEME.len()] == b' '
&& slice[..C::SCHEME.len()].eq_ignore_ascii_case(C::SCHEME.as_bytes())
{
C::decode(val).map(Authorization)
} else {
None
}
})
.ok_or_else(Error::invalid)
}
}
impl<C: Credentials> HeaderEncode for Authorization<C> {
fn encode<E: Extend<HeaderValue>>(&self, values: &mut E) {
values.extend(self.0.encode().map(|mut value| {
value.set_sensitive(true);
debug_assert!(
value.as_bytes().starts_with(C::SCHEME.as_bytes()),
"Credentials::encode should include its scheme: scheme = {:?}, encoded = {:?}",
C::SCHEME,
value,
);
value
}));
}
}
pub trait Credentials: Sized {
const SCHEME: &'static str;
fn decode(value: &HeaderValue) -> Option<Self>;
fn encode(&self) -> Option<HeaderValue>;
}
impl Credentials for Basic {
const SCHEME: &'static str = "Basic";
fn decode(value: &HeaderValue) -> Option<Self> {
let value = value.as_ref();
if value.len() <= Self::SCHEME.len() + 1 {
tracing::trace!(
"Basic credentials failed to decode: invalid scheme length in basic str"
);
return None;
}
if !value[..Self::SCHEME.len()].eq_ignore_ascii_case(Self::SCHEME.as_bytes()) {
tracing::trace!("Basic credentials failed to decode: invalid scheme in basic str");
return None;
}
let bytes = &value[Self::SCHEME.len() + 1..];
let Some(non_space_pos) = bytes.iter().position(|b| *b != b' ') else {
tracing::trace!(
"Basic credentials failed to decode: missing space separator in basic str"
);
return None;
};
let bytes = &bytes[non_space_pos..];
let bytes = ENGINE
.decode(bytes)
.inspect_err(|err| {
tracing::trace!("Basic credentials failed to decode: base64 decode: {err:?}");
})
.ok()?;
let decoded = String::from_utf8(bytes)
.inspect_err(|err| {
tracing::trace!("Basic credentials failed to decode: utf8 validation: {err:?}");
})
.ok()?;
decoded
.parse()
.inspect_err(|err| {
tracing::trace!("Basic credentials failed to decode: str parse: {err:?}");
})
.ok()
}
fn encode(&self) -> Option<HeaderValue> {
let mut encoded = format!("{} ", Self::SCHEME);
ENGINE.encode_string(self.to_string(), &mut encoded);
HeaderValue::try_from(encoded)
.inspect_err(|err| {
tracing::debug!("failed to encode basic value as header value: {err}");
})
.ok()
}
}
impl Credentials for Bearer {
const SCHEME: &'static str = "Bearer";
fn decode(value: &HeaderValue) -> Option<Self> {
let value = value.as_ref();
if value.len() <= Self::SCHEME.len() + 1 {
tracing::trace!("Bearer credentials failed to decode: invalid bearer scheme length");
return None;
}
if !value[..Self::SCHEME.len()].eq_ignore_ascii_case(Self::SCHEME.as_bytes()) {
tracing::trace!("Bearer credentials failed to decode: invalid bearer scheme");
return None;
}
let bytes = &value[Self::SCHEME.len() + 1..];
let Some(non_space_pos) = bytes.iter().position(|b| *b != b' ') else {
tracing::trace!("Bearer credentials failed to decode: no token found");
return None;
};
let bytes = &bytes[non_space_pos..];
let s = std::str::from_utf8(bytes)
.inspect_err(|err| {
tracing::trace!("Bearer credentials failed to decode: {err:?}");
})
.ok()?;
Self::try_from(s.to_owned())
.inspect_err(|err| {
tracing::trace!("Bearer credentials failed to decode: {err:?}");
})
.ok()
}
fn encode(&self) -> Option<HeaderValue> {
HeaderValue::try_from(format!("{} {}", Self::SCHEME, self.token()))
.inspect_err(|err| {
tracing::debug!("failed to encode bearer auth as header value: {err}");
})
.ok()
}
}
impl Credentials for RawToken {
const SCHEME: &'static str = "";
fn decode(value: &HeaderValue) -> Option<Self> {
let s = std::str::from_utf8(value.as_bytes())
.inspect_err(|err| {
tracing::trace!("RawToken credentials failed to decode: {err:?}");
})
.ok()?;
Self::try_from(s.to_owned())
.inspect_err(|err| {
tracing::trace!("RawToken credentials failed to decode: {err}");
})
.ok()
}
fn encode(&self) -> Option<HeaderValue> {
HeaderValue::try_from(self.token().to_owned())
.inspect_err(|err| {
tracing::debug!("failed to encode raw token as header value: {err}");
})
.ok()
}
}
pub trait Authority<C, L>: Send + Sync + 'static {
fn authorized(&self, credentials: C) -> impl Future<Output = Option<Extensions>> + Send + '_;
}
pub trait AuthoritySync<C, L>: Send + Sync + 'static {
fn authorized(&self, ext: &Extensions, credentials: &C) -> bool;
}
impl<A, C, L> Authority<C, L> for A
where
A: AuthoritySync<C, L>,
C: Credentials + Send + 'static,
L: 'static,
{
async fn authorized(&self, credentials: C) -> Option<Extensions> {
let ext = Extensions::new();
if self.authorized(&ext, &credentials) {
Some(ext)
} else {
None
}
}
}
impl<T: UsernameLabelParser> AuthoritySync<Self, T> for Basic {
fn authorized(&self, ext: &Extensions, credentials: &Self) -> bool {
let username = credentials.username();
let password = credentials.password();
if password != self.password() {
return false;
}
let parser_ext = Extensions::new();
let username = match parse_username(&parser_ext, T::default(), username) {
Ok(t) => t,
Err(err) => {
tracing::trace!("failed to parse username: {:?}", err);
return if self == credentials {
ext.insert(UserId::Username(username.to_owned()));
true
} else {
false
};
}
};
if username != self.username() {
return false;
}
ext.extend(&parser_ext);
ext.insert(UserId::Username(username));
true
}
}
impl<C, L, T, const N: usize> AuthoritySync<C, L> for [T; N]
where
C: Credentials + Send + 'static,
T: AuthoritySync<C, L>,
{
fn authorized(&self, ext: &Extensions, credentials: &C) -> bool {
self.iter().any(|t| t.authorized(ext, credentials))
}
}
impl<C, L, T> AuthoritySync<C, L> for Vec<T>
where
C: Credentials + Send + 'static,
T: AuthoritySync<C, L>,
{
fn authorized(&self, ext: &Extensions, credentials: &C) -> bool {
self.iter().any(|t| t.authorized(ext, credentials))
}
}
#[cfg(test)]
mod tests {
use rama_http_types::header::HeaderMap;
use rama_net::user::credentials::bearer;
use rama_utils::str::non_empty_str;
use super::super::{test_decode, test_encode};
use super::{Authorization, Basic, Bearer};
use crate::HeaderMapExt;
#[test]
fn basic_encode() {
let auth = Authorization::new(Basic::new(
non_empty_str!("Aladdin"),
non_empty_str!("open sesame"),
));
let headers = test_encode(auth);
assert_eq!(
headers["authorization"],
"Basic QWxhZGRpbjpvcGVuIHNlc2FtZQ==",
);
}
#[test]
fn basic_username_encode() {
let auth = Authorization::new(Basic::new_insecure(non_empty_str!("Aladdin")));
let headers = test_encode(auth);
assert_eq!(headers["authorization"], "Basic QWxhZGRpbjo=",);
}
#[test]
fn basic_roundtrip() {
let auth = Authorization::new(Basic::new(
non_empty_str!("Aladdin"),
non_empty_str!("open sesame"),
));
let mut h = HeaderMap::new();
h.typed_insert(&auth);
assert_eq!(h.typed_get(), Some(auth));
}
#[test]
fn basic_decode() {
let auth: Authorization<Basic> =
test_decode(&["Basic QWxhZGRpbjpvcGVuIHNlc2FtZQ=="]).unwrap();
assert_eq!(auth.0.username(), "Aladdin");
assert_eq!(auth.0.password(), Some("open sesame"));
}
#[test]
fn basic_decode_case_insensitive() {
let auth: Authorization<Basic> =
test_decode(&["basic QWxhZGRpbjpvcGVuIHNlc2FtZQ=="]).unwrap();
assert_eq!(auth.0.username(), "Aladdin");
assert_eq!(auth.0.password(), Some("open sesame"));
}
#[test]
fn basic_decode_extra_whitespaces() {
let auth: Authorization<Basic> =
test_decode(&["Basic QWxhZGRpbjpvcGVuIHNlc2FtZQ=="]).unwrap();
assert_eq!(auth.0.username(), "Aladdin");
assert_eq!(auth.0.password(), Some("open sesame"));
}
#[test]
fn basic_decode_no_password() {
let auth: Authorization<Basic> = test_decode(&["Basic QWxhZGRpbjo="]).unwrap();
assert_eq!(auth.0.username(), "Aladdin");
assert_eq!(auth.0.password(), None);
}
#[test]
fn bearer_encode() {
let auth = Authorization::new(bearer!("fpKL54jvWmEGVoRdCNjG"));
let headers = test_encode(auth);
assert_eq!(headers["authorization"], "Bearer fpKL54jvWmEGVoRdCNjG",);
}
#[test]
fn bearer_decode() {
let auth: Authorization<Bearer> = test_decode(&["Bearer fpKL54jvWmEGVoRdCNjG"]).unwrap();
assert_eq!(auth.0.token().as_bytes(), b"fpKL54jvWmEGVoRdCNjG");
}
#[test]
fn bearer_decode_case_insensitive() {
let auth: Authorization<Bearer> = test_decode(&["bearer fpKL54jvWmEGVoRdCNjG"]).unwrap();
assert_eq!(auth.0.token().as_bytes(), b"fpKL54jvWmEGVoRdCNjG");
}
#[test]
fn bearer_decode_extra_whitespaces() {
let auth: Authorization<Bearer> = test_decode(&["Bearer fpKL54jvWmEGVoRdCNjG"]).unwrap();
assert_eq!(auth.0.token().as_bytes(), b"fpKL54jvWmEGVoRdCNjG");
}
#[test]
fn regression_authorization_raw_token_roundtrip() {
use rama_net::user::RawToken;
let token = RawToken::try_from("fpKL54jvWmEGVoRdCNjG").unwrap();
let auth = Authorization::new(token.clone());
let headers = test_encode(auth);
assert_eq!(headers["authorization"], "fpKL54jvWmEGVoRdCNjG");
let decoded: Authorization<RawToken> = test_decode(&["fpKL54jvWmEGVoRdCNjG"]).unwrap();
assert_eq!(decoded.0, token);
}
#[test]
fn regression_authorization_raw_token_accepts_loose_alphabet() {
use rama_net::user::RawToken;
let decoded: Authorization<RawToken> =
test_decode(&["sk-live_abc=xyz,scope:read"]).unwrap();
assert_eq!(decoded.0.token(), "sk-live_abc=xyz,scope:read");
}
}
#[cfg(test)]
mod test_auth {
use super::*;
use rama_core::username::{UsernameLabels, UsernameOpaqueLabelParser};
use rama_net::user::credentials::basic;
#[tokio::test]
async fn basic_authorization() {
let auth = basic!("Aladdin", "open sesame");
let auths = vec![basic!("foo", "bar"), auth.clone()];
let ext = Authority::<_, ()>::authorized(&auths, auth).await.unwrap();
let user: &UserId = ext.get_ref().unwrap();
assert_eq!(user, "Aladdin");
}
#[tokio::test]
async fn basic_authorization_with_labels_found() {
let auths = vec![basic!("foo", "bar"), basic!("john", "secret")];
let ext = Authority::<_, UsernameOpaqueLabelParser>::authorized(
&auths,
basic!("john-green-red", "secret"),
)
.await
.unwrap();
let c: &UserId = ext.get_ref().unwrap();
assert_eq!(c, "john");
let labels: &UsernameLabels = ext.get_ref().unwrap();
assert_eq!(&labels.0, &vec!["green".to_owned(), "red".to_owned()]);
}
#[tokio::test]
async fn basic_authorization_with_labels_not_found() {
let auth = basic!("john", "secret");
let auths = vec![basic!("foo", "bar"), auth.clone()];
let ext = Authority::<_, UsernameOpaqueLabelParser>::authorized(&auths, auth)
.await
.unwrap();
let c: &UserId = ext.get_ref().unwrap();
assert_eq!(c, "john");
assert!(ext.get_ref::<UsernameLabels>().is_none());
}
}