use crate::traits::StorageError;
pub async fn refresh_oauth2_token(
client: &reqwest::Client,
auth_url: &str,
client_id: &str,
client_secret: &str,
refresh_token: &str,
provider_name: &str,
) -> Result<String, StorageError> {
let (access_token, _, _) = refresh_oauth2_token_details(client, auth_url, client_id, client_secret, refresh_token, provider_name).await?;
Ok(access_token)
}
pub async fn refresh_oauth2_token_details(
client: &reqwest::Client,
auth_url: &str,
client_id: &str,
client_secret: &str,
refresh_token: &str,
provider_name: &str,
) -> Result<(String, Option<String>, Option<u64>), StorageError> {
let params = [
("client_id", client_id),
("client_secret", client_secret),
("refresh_token", refresh_token),
("grant_type", "refresh_token"),
];
let res = client.post(auth_url)
.form(¶ms)
.send()
.await?;
if !res.status().is_success() {
return Err(translate_http_error(res, provider_name, "refresh_oauth2_token").await);
}
let json: serde_json::Value = res.json().await?;
let access_token = json["access_token"].as_str().ok_or_else(|| {
StorageError::Authentication(format!("Failed to retrieve {} access token: {:?}", provider_name, json))
})?.to_string();
let new_refresh_token = json["refresh_token"].as_str().map(|s| s.to_string());
let expires_in = json["expires_in"].as_u64();
Ok((access_token, new_refresh_token, expires_in))
}
pub type TokenRefreshCallback = std::sync::Arc<dyn Fn(&str) + Send + Sync>;
pub struct OAuthTokenManager {
client: reqwest::Client,
token_url: String,
client_id: String,
client_secret: String,
refresh_token: tokio::sync::RwLock<String>,
provider_name: String,
cache: tokio::sync::RwLock<Option<(String, std::time::Instant)>>,
on_refresh: Option<TokenRefreshCallback>,
}
impl OAuthTokenManager {
pub fn new(
client: reqwest::Client,
token_url: &str,
client_id: &str,
client_secret: &str,
refresh_token: &str,
provider_name: &str,
) -> Self {
Self::with_callback(client, token_url, client_id, client_secret, refresh_token, provider_name, None)
}
pub fn with_callback(
client: reqwest::Client,
token_url: &str,
client_id: &str,
client_secret: &str,
refresh_token: &str,
provider_name: &str,
on_refresh: Option<TokenRefreshCallback>,
) -> Self {
Self {
client,
token_url: token_url.to_string(),
client_id: client_id.to_string(),
client_secret: client_secret.to_string(),
refresh_token: tokio::sync::RwLock::new(refresh_token.to_string()),
provider_name: provider_name.to_string(),
cache: tokio::sync::RwLock::new(None),
on_refresh,
}
}
pub async fn get_access_token(&self) -> Result<String, StorageError> {
{
let cache = self.cache.read().await;
if let Some((ref token, expiry)) = *cache {
if std::time::Instant::now() + std::time::Duration::from_secs(30) < expiry {
return Ok(token.clone());
}
}
}
let mut cache = self.cache.write().await;
if let Some((ref token, expiry)) = *cache {
if std::time::Instant::now() + std::time::Duration::from_secs(30) < expiry {
return Ok(token.clone());
}
}
let current_refresh_token = self.refresh_token.read().await.clone();
let (token, new_refresh_token, expires_in) = refresh_oauth2_token_details(
&self.client,
&self.token_url,
&self.client_id,
&self.client_secret,
¤t_refresh_token,
&self.provider_name,
).await?;
if let Some(ref new_ref) = new_refresh_token {
let mut ref_guard = self.refresh_token.write().await;
*ref_guard = new_ref.clone();
if let Some(ref cb) = self.on_refresh {
cb(new_ref);
}
}
let ttl = expires_in.unwrap_or(3600);
let safety_margin = if ttl > 600 { 300 } else { ttl / 2 };
let expiry = std::time::Instant::now() + std::time::Duration::from_secs(ttl.saturating_sub(safety_margin));
*cache = Some((token.clone(), expiry));
Ok(token)
}
}
pub fn translate_status_code_error(status_code: u16, provider_name: &str, action: &str, detail: Option<&str>) -> StorageError {
let msg = match detail {
Some(d) if !d.trim().is_empty() => d.to_string(),
_ => format!("HTTP status {}", status_code),
};
match status_code {
429 => StorageError::RateLimit {
message: format!("Rate limit exceeded on {}: {}", provider_name, msg),
retry_after: None,
},
404 => StorageError::NotFound(format!("Resource not found on {}: {}", provider_name, msg)),
401 | 403 => StorageError::AuthenticationExpired(format!("Authentication expired or forbidden on {}: {}", provider_name, msg)),
409 => StorageError::Conflict(format!("Conflict on {}: {}", provider_name, msg)),
_ => StorageError::Provider {
message: format!("Failed to {} on {}: {}", action, provider_name, msg),
status: Some(status_code),
},
}
}
pub async fn translate_http_error(res: reqwest::Response, provider_name: &str, action: &str) -> StorageError {
let status = res.status();
let retry_after = res.headers().get(reqwest::header::RETRY_AFTER)
.and_then(|val| val.to_str().ok())
.and_then(|val_str| {
if let Ok(secs) = val_str.parse::<u64>() {
Some(std::time::Duration::from_secs(secs))
} else {
None
}
});
let body = res.text().await.unwrap_or_default();
let detail = if body.trim().is_empty() {
status.to_string()
} else {
body
};
let mut err = translate_status_code_error(status.as_u16(), provider_name, action, Some(&detail));
if let StorageError::RateLimit { retry_after: ref mut ra, .. } = err {
if ra.is_none() {
*ra = retry_after;
}
}
err
}
pub fn copy_buffered<R: std::io::Read, W: std::io::Write>(mut reader: R, mut writer: W) -> std::io::Result<u64> {
let mut buffer = [0u8; 16384];
let mut total_copied = 0;
loop {
let bytes_read = reader.read(&mut buffer)?;
if bytes_read == 0 {
break;
}
writer.write_all(&buffer[..bytes_read])?;
total_copied += bytes_read as u64;
}
Ok(total_copied)
}
pub fn apply_bearer_auth(req: reqwest::RequestBuilder, token: &str) -> reqwest::RequestBuilder {
req.bearer_auth(token)
}
pub fn get_secure_credential(service: &str, key: &str, fallback: &str) -> String {
if fallback.is_empty() || fallback.starts_with("PLACEHOLDER_") {
if let Ok(entry) = keyring::Entry::new(service, key) {
if let Ok(secret) = entry.get_password() {
return secret;
}
}
}
fallback.to_string()
}
pub use cloud_sync_core::path::{normalize_remote_path, format_relative_path, format_absolute_path, strip_destination_prefix, url_encode, url_encode_path, get_parent_and_filename};
#[macro_export]
macro_rules! impl_provider_builder {
($provider:ident, $builder:ident, $creds:ty, absolute) => {
$crate::impl_provider_builder!($provider, $builder, $creds);
impl $provider {
fn format_path<'a>(&self, remote_path: &'a str) -> std::borrow::Cow<'a, str> {
$crate::providers::utils::format_absolute_path(remote_path, self.credentials.common.destination_folder.as_deref())
}
}
};
($provider:ident, $builder:ident, $creds:ty, relative) => {
$crate::impl_provider_builder!($provider, $builder, $creds);
impl $provider {
fn format_path<'a>(&self, remote_path: &'a str) -> std::borrow::Cow<'a, str> {
$crate::providers::utils::format_relative_path(remote_path, self.credentials.common.destination_folder.as_deref())
}
}
};
($provider:ident, $builder:ident, $creds:ty) => {
impl $provider {
pub fn builder(credentials: $creds) -> $builder {
$builder::new(credentials)
}
pub fn new(credentials: $creds) -> Self {
Self::with_client_options(credentials, None, None)
}
}
impl $builder {
pub fn timeout(mut self, timeout: std::time::Duration) -> Self {
self.timeout = Some(timeout);
self
}
pub fn custom_headers(mut self, headers: reqwest::header::HeaderMap) -> Self {
self.custom_headers = Some(headers);
self
}
}
};
}
#[macro_export]
macro_rules! impl_oauth_token_helper {
($provider:ident) => {
impl $provider {
async fn get_access_token(&self) -> Result<String, StorageError> {
self.token_manager.get_access_token().await
}
}
};
}
pub fn build_http_client(
timeout: Option<std::time::Duration>,
custom_headers: Option<reqwest::header::HeaderMap>,
) -> reqwest::Client {
let mut builder = reqwest::Client::builder()
.timeout(timeout.unwrap_or(std::time::Duration::from_secs(600)))
.pool_max_idle_per_host(10);
if let Some(headers) = custom_headers {
builder = builder.default_headers(headers);
}
builder.build().unwrap_or_else(|_| reqwest::Client::new())
}
pub async fn get_upload_body(
local_path: &std::path::Path,
limiter: Option<crate::rate_limit::TokenBucket>,
) -> Result<(reqwest::Body, u64), StorageError> {
let file = tokio::fs::File::open(local_path).await?;
let metadata = file.metadata().await?;
let size = metadata.len();
let reader = RateLimitedReader::new(file, limiter);
let stream = ReaderStream::new(reader);
let body = reqwest::Body::wrap_stream(stream);
Ok((body, size))
}
pub async fn download_rate_limited(
res: reqwest::Response,
local_path: &std::path::Path,
limiter: Option<crate::rate_limit::TokenBucket>,
) -> Result<(), StorageError> {
if let Some(parent) = local_path.parent() {
tokio::fs::create_dir_all(parent).await?;
}
let mut file = tokio::fs::File::create(local_path).await?;
let byte_stream = res.bytes_stream();
let mut rate_limited_stream = RateLimitedStream::new(byte_stream, limiter);
use futures_util::stream::StreamExt;
use tokio::io::AsyncWriteExt;
while let Some(chunk_result) = rate_limited_stream.next().await {
let chunk = chunk_result.map_err(StorageError::Reqwest)?;
file.write_all(&chunk).await?;
}
file.flush().await?;
Ok(())
}
use crate::rate_limit::{RateLimitedReader, RateLimitedStream};
use tokio_util::io::ReaderStream;
pub fn is_transient_error(err: &StorageError) -> bool {
match err {
StorageError::RateLimit { .. } => true,
StorageError::Reqwest(e) => {
if e.is_timeout() || e.is_connect() {
return true;
}
if let Some(status) = e.status() {
status == reqwest::StatusCode::TOO_MANY_REQUESTS
|| status.is_server_error()
} else {
false
}
}
StorageError::ConnectionFailed(_) => true,
StorageError::Provider { status: Some(status_code), .. } => {
*status_code == 429 || *status_code == 502 || *status_code == 503 || *status_code == 504
}
StorageError::Provider { message, .. } => {
message.contains("429") || message.contains("503") || message.contains("504") || message.contains("502")
}
_ => false,
}
}
use std::sync::Mutex;
#[derive(Debug, Clone, Copy)]
pub struct RetryConfig {
pub max_attempts: usize,
pub initial_delay: std::time::Duration,
pub multiplier: f64,
}
impl Default for RetryConfig {
fn default() -> Self {
Self {
max_attempts: 5,
initial_delay: std::time::Duration::from_millis(500),
multiplier: 2.0,
}
}
}
static GLOBAL_RETRY_CONFIG: Mutex<Option<RetryConfig>> = Mutex::new(None);
pub fn set_global_retry_config(config: RetryConfig) {
if let Ok(mut lock) = GLOBAL_RETRY_CONFIG.lock() {
*lock = Some(config);
}
}
pub fn get_global_retry_config() -> RetryConfig {
if let Ok(lock) = GLOBAL_RETRY_CONFIG.lock() {
lock.clone().unwrap_or_default()
} else {
RetryConfig::default()
}
}
pub async fn execute_with_retry<T, F, Fut>(
provider_name: &str,
action: &str,
f: F,
) -> Result<T, StorageError>
where
F: Fn() -> Fut + Send + Sync,
Fut: std::future::Future<Output = Result<T, StorageError>> + Send,
{
let config = get_global_retry_config();
let max_attempts = config.max_attempts;
let mut attempt = 0;
let mut delay = if cfg!(test) {
std::time::Duration::from_millis(1)
} else {
config.initial_delay
};
let multiplier = config.multiplier;
loop {
match f().await {
Ok(val) => return Ok(val),
Err(e) => {
if is_transient_error(&e) && attempt < max_attempts - 1 {
attempt += 1;
let sleep_duration = match &e {
StorageError::RateLimit { retry_after: Some(d), .. } => *d,
_ => delay,
};
tracing::warn!(
"[{}] Transient error during {}: {}. Retrying in {:?} (attempt {}/{})",
provider_name,
action,
e,
sleep_duration,
attempt,
max_attempts
);
tokio::time::sleep(sleep_duration).await;
delay = std::time::Duration::from_secs_f64(delay.as_secs_f64() * multiplier);
continue;
}
return Err(e);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::StorageError;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use tempfile::tempdir;
use http;
#[tokio::test]
async fn test_execute_with_retry_success() {
let counter = Arc::new(AtomicUsize::new(0));
let counter_clone = counter.clone();
let result = execute_with_retry("test", "op", || {
let cnt = counter_clone.clone();
async move {
cnt.fetch_add(1, Ordering::SeqCst);
Ok::<_, StorageError>("success")
}
}).await;
assert_eq!(result.unwrap(), "success");
assert_eq!(counter.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn test_execute_with_retry_fail_non_transient() {
let counter = Arc::new(AtomicUsize::new(0));
let counter_clone = counter.clone();
let result = execute_with_retry("test", "op", || {
let cnt = counter_clone.clone();
async move {
cnt.fetch_add(1, Ordering::SeqCst);
Err::<(), StorageError>(StorageError::Authentication("Auth error".to_string()))
}
}).await;
assert!(matches!(result, Err(StorageError::Authentication(_))));
assert_eq!(counter.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn test_execute_with_retry_retry_then_success() {
let counter = Arc::new(AtomicUsize::new(0));
let counter_clone = counter.clone();
let result = execute_with_retry("test", "op", || {
let cnt = counter_clone.clone();
async move {
let current = cnt.fetch_add(1, Ordering::SeqCst);
if current < 2 {
Err(StorageError::RateLimit { message: "Rate limit".to_string(), retry_after: None })
} else {
Ok("success")
}
}
}).await;
assert_eq!(result.unwrap(), "success");
assert_eq!(counter.load(Ordering::SeqCst), 3); }
#[tokio::test]
async fn test_execute_with_retry_max_attempts() {
let counter = Arc::new(AtomicUsize::new(0));
let counter_clone = counter.clone();
let result = execute_with_retry("test", "op", || {
let cnt = counter_clone.clone();
async move {
cnt.fetch_add(1, Ordering::SeqCst);
Err::<(), StorageError>(StorageError::RateLimit { message: "Rate limit".to_string(), retry_after: None })
}
}).await;
assert!(matches!(result, Err(StorageError::RateLimit { .. })));
assert_eq!(counter.load(Ordering::SeqCst), 5); }
#[tokio::test]
async fn test_execute_with_retry_respects_retry_after() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
use std::time::Instant;
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/retry-test"))
.respond_with(
ResponseTemplate::new(429)
.insert_header("Retry-After", "1")
)
.up_to_n_times(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/retry-test"))
.respond_with(ResponseTemplate::new(200).set_body_string("success"))
.mount(&server)
.await;
let client = reqwest::Client::new();
let request_url = format!("{}/retry-test", server.uri());
let start = Instant::now();
let result = execute_with_retry("test_retry_after", "get", || {
let cl = client.clone();
let url = request_url.clone();
async move {
let res = cl.get(&url).send().await.map_err(StorageError::Reqwest)?;
if !res.status().is_success() {
return Err(translate_http_error(res, "test_retry_after", "get").await);
}
Ok("success")
}
}).await;
let elapsed = start.elapsed();
assert_eq!(result.unwrap(), "success");
assert!(elapsed.as_millis() >= 950, "Should respect Retry-After delay, took {:?}", elapsed);
}
#[tokio::test]
async fn test_translate_status_code_error() {
assert!(matches!(translate_status_code_error(429, "P", "A", None), StorageError::RateLimit { .. }));
assert!(matches!(translate_status_code_error(404, "P", "A", None), StorageError::NotFound(_)));
assert!(matches!(translate_status_code_error(401, "P", "A", None), StorageError::AuthenticationExpired(_)));
assert!(matches!(translate_status_code_error(409, "P", "A", None), StorageError::Conflict(_)));
assert!(matches!(translate_status_code_error(500, "P", "A", None), StorageError::Provider { status: Some(500), .. }));
}
#[test]
fn test_copy_buffered() {
let src = b"hello world buffer copy";
let mut dest = Vec::new();
let bytes = copy_buffered(&src[..], &mut dest).unwrap();
assert_eq!(bytes, src.len() as u64);
assert_eq!(dest, src);
}
#[test]
fn test_apply_bearer_auth() {
let client = reqwest::Client::new();
let builder = client.get("http://localhost");
let _builder = apply_bearer_auth(builder, "token123");
}
#[test]
fn test_build_http_client() {
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(reqwest::header::USER_AGENT, reqwest::header::HeaderValue::from_static("test-agent"));
let _client = build_http_client(Some(std::time::Duration::from_secs(5)), Some(headers));
}
#[tokio::test]
async fn test_retry_config_global() {
let original = get_global_retry_config();
let custom = RetryConfig {
max_attempts: 12,
initial_delay: std::time::Duration::from_millis(10),
multiplier: 1.5,
};
set_global_retry_config(custom);
let current = get_global_retry_config();
assert_eq!(current.max_attempts, 12);
assert_eq!(current.multiplier, 1.5);
set_global_retry_config(original);
}
#[tokio::test]
async fn test_oauth_token_manager() {
use wiremock::matchers::{method, path, body_string_contains};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
let response_body = r#"{"access_token":"new_access_token_123","refresh_token":"new_refresh_token_456","expires_in":3600}"#;
Mock::given(method("POST"))
.and(path("/token"))
.and(body_string_contains("refresh_token=old_refresh"))
.respond_with(ResponseTemplate::new(200).set_body_string(response_body))
.mount(&server)
.await;
let client = reqwest::Client::new();
let callback_called = Arc::new(AtomicUsize::new(0));
let callback_called_clone = callback_called.clone();
let on_refresh = Arc::new(move |new_token: &str| {
assert_eq!(new_token, "new_refresh_token_456");
callback_called_clone.fetch_add(1, Ordering::SeqCst);
});
let manager = OAuthTokenManager::with_callback(
client,
&format!("{}/token", server.uri()),
"client_id",
"client_secret",
"old_refresh",
"MockProvider",
Some(on_refresh),
);
let token = manager.get_access_token().await.unwrap();
assert_eq!(token, "new_access_token_123");
assert_eq!(callback_called.load(Ordering::SeqCst), 1);
let token_cached = manager.get_access_token().await.unwrap();
assert_eq!(token_cached, "new_access_token_123");
assert_eq!(callback_called.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn test_rate_limited_upload_download_bodies() {
use bytes::Bytes;
let temp_dir = tempdir().unwrap();
let file_path = temp_dir.path().join("upload.txt");
std::fs::write(&file_path, "hello rate limited upload body").unwrap();
let limiter = crate::rate_limit::TokenBucket::new(100);
let (_body, size) = get_upload_body(&file_path, Some(limiter.clone())).await.unwrap();
assert_eq!(size, 30);
let res = reqwest::Response::from(
http::Response::builder()
.status(200)
.body(Bytes::from("hello rate limited download body"))
.unwrap()
);
let download_path = temp_dir.path().join("download.txt");
download_rate_limited(res, &download_path, Some(limiter)).await.unwrap();
assert_eq!(std::fs::read_to_string(download_path).unwrap(), "hello rate limited download body");
}
}