use async_trait::async_trait;
use bytes::Bytes;
use percent_encoding::{AsciiSet, NON_ALPHANUMERIC, utf8_percent_encode};
use crate::error::Error;
use crate::http::header::{AUTHORIZATION, COOKIE};
use crate::http::{HeaderValue, Request};
use crate::types::SensitiveString;
pub const SESSION_COOKIE: &str = "session_token";
const COOKIE_VALUE: &AsciiSet = &NON_ALPHANUMERIC
.remove(b'-')
.remove(b'_')
.remove(b'.')
.remove(b'~');
pub(crate) fn cookie_header(name: &str, value: &str) -> Result<HeaderValue, Error> {
let encoded = utf8_percent_encode(value, COOKIE_VALUE);
let mut header = HeaderValue::from_str(&format!("{name}={encoded}"))
.map_err(|_| Error::auth(format!("{name} is not a valid cookie value")))?;
header.set_sensitive(true);
Ok(header)
}
#[async_trait]
pub trait TokenProvider: Send + Sync {
async fn access_token(&self) -> Result<String, Error>;
}
#[derive(Debug, Clone)]
pub struct StaticTokenProvider {
pub token: SensitiveString,
}
impl StaticTokenProvider {
pub fn new(token: impl Into<SensitiveString>) -> StaticTokenProvider {
StaticTokenProvider {
token: token.into(),
}
}
}
#[async_trait]
impl TokenProvider for StaticTokenProvider {
async fn access_token(&self) -> Result<String, Error> {
if self.token.is_empty() {
Err(Error::auth("no token configured"))
} else {
Ok(self.token.expose().to_string())
}
}
}
#[async_trait]
pub trait AuthStrategy: Send + Sync {
async fn authenticate(&self, request: &mut Request<Bytes>) -> Result<(), Error>;
}
pub struct BearerAuth<P: TokenProvider> {
provider: P,
}
impl<P: TokenProvider> BearerAuth<P> {
pub fn new(provider: P) -> BearerAuth<P> {
BearerAuth { provider }
}
}
#[async_trait]
impl<P: TokenProvider> AuthStrategy for BearerAuth<P> {
async fn authenticate(&self, request: &mut Request<Bytes>) -> Result<(), Error> {
let token = self.provider.access_token().await?;
let mut value = HeaderValue::from_str(&format!("Bearer {token}"))
.map_err(|_| Error::auth("access token is not a valid header value"))?;
value.set_sensitive(true);
request.headers_mut().insert(AUTHORIZATION, value);
Ok(())
}
}
pub struct CookieAuth<P: TokenProvider> {
provider: P,
}
impl<P: TokenProvider> CookieAuth<P> {
pub fn new(provider: P) -> CookieAuth<P> {
CookieAuth { provider }
}
}
#[async_trait]
impl<P: TokenProvider> AuthStrategy for CookieAuth<P> {
async fn authenticate(&self, request: &mut Request<Bytes>) -> Result<(), Error> {
let token = self.provider.access_token().await?;
request
.headers_mut()
.insert(COOKIE, cookie_header(SESSION_COOKIE, &token)?);
Ok(())
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct NoAuth;
#[async_trait]
impl AuthStrategy for NoAuth {
async fn authenticate(&self, _request: &mut Request<Bytes>) -> Result<(), Error> {
Ok(())
}
}