use std::sync::PoisonError;
use std::time::Duration;
use delta_kernel::object_store::Result as ObjectStoreResult;
use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
use super::generic_error;
pub trait AuthHeaderProvider: std::fmt::Debug + Send + Sync {
fn headers(&self) -> ObjectStoreResult<HeaderMap>;
}
#[derive(Debug, Clone)]
pub struct StaticHeaderProvider {
headers: HeaderMap,
}
impl StaticHeaderProvider {
pub fn new(headers: HeaderMap) -> Self {
Self { headers }
}
pub fn from_pairs(
pairs: impl IntoIterator<Item = (String, String)>,
) -> ObjectStoreResult<Self> {
Ok(Self {
headers: headers_from_pairs(pairs)?,
})
}
}
impl AuthHeaderProvider for StaticHeaderProvider {
fn headers(&self) -> ObjectStoreResult<HeaderMap> {
Ok(self.headers.clone())
}
}
pub fn headers_from_pairs(
pairs: impl IntoIterator<Item = (String, String)>,
) -> ObjectStoreResult<HeaderMap> {
let mut headers = HeaderMap::new();
for (name, value) in pairs {
let name = HeaderName::from_bytes(name.as_bytes()).map_err(generic_error)?;
let value = HeaderValue::from_str(&value).map_err(generic_error)?;
headers.insert(name, value);
}
Ok(headers)
}
const HEADER_REFRESH_BUFFER: Duration = Duration::from_secs(30);
pub struct RefreshingHeaderProvider {
produce: Box<dyn Fn() -> ObjectStoreResult<(HeaderMap, Option<Duration>)> + Send + Sync>,
cached: std::sync::Mutex<Option<(HeaderMap, std::time::Instant)>>,
}
impl std::fmt::Debug for RefreshingHeaderProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RefreshingHeaderProvider")
.finish_non_exhaustive()
}
}
impl RefreshingHeaderProvider {
pub fn new(
produce: impl Fn() -> ObjectStoreResult<(HeaderMap, Option<Duration>)> + Send + Sync + 'static,
) -> Self {
Self {
produce: Box::new(produce),
cached: std::sync::Mutex::new(None),
}
}
}
impl AuthHeaderProvider for RefreshingHeaderProvider {
fn headers(&self) -> ObjectStoreResult<HeaderMap> {
let mut cached = self.cached.lock().unwrap_or_else(PoisonError::into_inner);
if let Some((headers, deadline)) = cached.as_ref() {
if deadline.saturating_duration_since(std::time::Instant::now()) > HEADER_REFRESH_BUFFER
{
return Ok(headers.clone());
}
}
let (headers, ttl) = (self.produce)()?;
*cached = ttl
.and_then(|ttl| Some((headers.clone(), std::time::Instant::now().checked_add(ttl)?)));
Ok(headers)
}
}