use bytes::Bytes;
use http::Extensions;
use reqwest_middleware::reqwest::header::HeaderValue;
use reqwest_middleware::reqwest::{Request, Response, header};
use reqwest_middleware::{Middleware, Next};
pub struct BasicAuthMiddleware {
username: String,
password: Option<String>,
}
impl BasicAuthMiddleware {
pub fn new(username: String, password: Option<String>) -> Self {
Self { username, password }
}
fn basic_auth(&self) -> HeaderValue {
use base64::prelude::BASE64_STANDARD;
use base64::write::EncoderWriter;
use std::io::Write;
let mut buf = b"Basic ".to_vec();
{
let mut encoder = EncoderWriter::new(&mut buf, &BASE64_STANDARD);
let _ = write!(encoder, "{}:", self.username);
if let Some(password) = &self.password {
let _ = write!(encoder, "{password}");
}
}
let mut header = HeaderValue::from_maybe_shared(Bytes::from(buf))
.expect("base64 is always valid HeaderValue");
header.set_sensitive(true);
header
}
}
#[async_trait::async_trait]
impl Middleware for BasicAuthMiddleware {
async fn handle(
&self,
mut req: Request,
extensions: &mut Extensions,
next: Next<'_>,
) -> reqwest_middleware::Result<Response> {
req.headers_mut()
.insert(header::AUTHORIZATION, self.basic_auth());
next.run(req, extensions).await
}
}