use anyhow::{Context, Result};
use reqwest::{Client, Response};
use url::Url;
use crate::drive::auth::{DriveCredentials, DriveSession};
use crate::drive::error::DriveError;
use crate::request_log;
use crate::utils::env::{EnvSource, SystemEnv};
use crate::utils::http::{connect_timeout, read_timeout, retry_if};
pub struct DriveClient {
client: Client,
base_url: String,
session: DriveSession,
}
impl std::fmt::Debug for DriveClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("DriveClient")
.field("base_url", &self.base_url)
.finish_non_exhaustive()
}
}
impl DriveClient {
const DEFAULT_BASE_URL: &'static str = "https://www.googleapis.com";
pub fn new(base_url: &str, credentials: &DriveCredentials) -> Result<Self> {
let client = Client::builder()
.connect_timeout(connect_timeout())
.read_timeout(read_timeout())
.build()
.context("Failed to build HTTP client")?;
let session = DriveSession::new(client.clone(), credentials);
Ok(Self {
client,
base_url: base_url.trim_end_matches('/').to_string(),
session,
})
}
pub fn from_credentials(credentials: &DriveCredentials) -> Result<Self> {
Self::from_credentials_with(&SystemEnv, credentials)
}
pub(crate) fn from_credentials_with(
env: &impl EnvSource,
credentials: &DriveCredentials,
) -> Result<Self> {
let base_url = env
.var(crate::drive::auth::DRIVE_API_URL)
.filter(|s| !s.is_empty())
.unwrap_or_else(|| Self::DEFAULT_BASE_URL.to_string());
Self::new(&base_url, credentials)
}
#[must_use]
pub fn base_url(&self) -> &str {
&self.base_url
}
pub(crate) fn api_url(base_url: &str, path: &str) -> Result<Url> {
Url::parse(&format!("{base_url}{path}")).context("Invalid Drive base URL")
}
pub(crate) async fn parse_response<T: serde::de::DeserializeOwned>(
&self,
response: Response,
context: &'static str,
) -> Result<T> {
if !response.status().is_success() {
return Err(Self::response_to_error(response).await.into());
}
response.json().await.context(context)
}
pub(crate) async fn get_parsed<T: serde::de::DeserializeOwned>(
&self,
url: &str,
context: &'static str,
) -> Result<T> {
let response = self.get_json(url).await?;
self.parse_response(response, context).await
}
pub async fn get_json(&self, url: &str) -> Result<Response> {
self.send_authorized(url, "GET", |client, token| {
client
.get(url)
.bearer_auth(token)
.header("Accept", "application/json")
})
.await
}
pub async fn get_bytes(&self, url: &str) -> Result<Response> {
self.send_authorized(url, "GET", |client, token| {
client.get(url).bearer_auth(token)
})
.await
}
pub async fn post_json<T: serde::Serialize + Sync + ?Sized>(
&self,
url: &str,
body: &T,
) -> Result<Response> {
self.send_authorized(url, "POST", |client, token| {
client
.post(url)
.bearer_auth(token)
.header("Content-Type", "application/json")
.json(body)
})
.await
}
pub async fn patch_json<T: serde::Serialize + Sync + ?Sized>(
&self,
url: &str,
body: &T,
) -> Result<Response> {
self.send_authorized(url, "PATCH", |client, token| {
client
.patch(url)
.bearer_auth(token)
.header("Content-Type", "application/json")
.json(body)
})
.await
}
async fn send_authorized<F>(
&self,
url: &str,
method: &'static str,
build: F,
) -> Result<Response>
where
F: Fn(&Client, &str) -> reqwest::RequestBuilder + Send + Sync,
{
let token = self
.session
.access_token()
.await
.context("Failed to obtain a Drive access token")?;
let response = self
.send_once(url, method, &build, token.expose_secret())
.await?;
if response.status().as_u16() != 401 {
return Ok(response);
}
let refreshed = self
.session
.force_refresh(&token)
.await
.context("Failed to refresh the Drive access token after a 401")?;
self.send_once(url, method, &build, refreshed.expose_secret())
.await
}
async fn send_once<F>(
&self,
url: &str,
method: &'static str,
build: &F,
token: &str,
) -> Result<Response>
where
F: Fn(&Client, &str) -> reqwest::RequestBuilder + Send + Sync,
{
retry_if(
|| build(&self.client, token),
|started, result| {
request_log::record_http_result("drive", method, url, started, result);
},
|status, body| status == 429 || is_drive_quota_exceeded(status, body),
)
.await
.with_context(|| format!("Failed to send {method} request to Drive API"))
}
pub async fn response_to_error(response: Response) -> DriveError {
let status = response.status().as_u16();
let raw = response.text().await.unwrap_or_default();
let value = serde_json::from_str::<serde_json::Value>(&raw).ok();
let reason = value.as_ref().and_then(drive_error_reason);
let body = value
.as_ref()
.and_then(drive_error_message)
.map(|message| match &reason {
Some(r) => format!("{message} (reason: {r})"),
None => message,
})
.unwrap_or(raw);
DriveError::ApiRequestFailed {
status,
body,
reason,
}
}
}
fn drive_error_reason(value: &serde_json::Value) -> Option<String> {
value
.get("error")
.and_then(|e| e.get("errors"))
.and_then(|e| e.as_array())
.and_then(|a| a.first())
.and_then(|e| e.get("reason"))
.and_then(|r| r.as_str())
.map(str::to_string)
}
fn drive_error_message(value: &serde_json::Value) -> Option<String> {
value
.get("error")?
.get("message")?
.as_str()
.map(str::to_string)
}
fn is_drive_quota_exceeded(status: u16, body: &[u8]) -> bool {
if status != 403 {
return false;
}
let Ok(text) = std::str::from_utf8(body) else {
return false;
};
let Ok(value) = serde_json::from_str::<serde_json::Value>(text) else {
return false;
};
matches!(
drive_error_reason(&value).as_deref(),
Some("userRateLimitExceeded")
)
}
#[cfg(test)]
pub(crate) mod test_support {
use super::DriveClient;
use crate::drive::auth::{DriveCredentials, DriveSession};
pub(crate) fn replace_session(
client: &mut DriveClient,
credentials: &DriveCredentials,
token_endpoint: &str,
) {
client.session = DriveSession::new_with_token_endpoint(
client.client.clone(),
credentials,
token_endpoint,
);
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
use crate::drive::auth::DriveScope;
use crate::utils::secret::Secret;
fn test_credentials() -> DriveCredentials {
DriveCredentials {
client_id: "client-1".to_string(),
client_secret: Secret::new("secret-1"),
refresh_token: Secret::new("refresh-1"),
scope: DriveScope::ReadOnly,
}
}
#[test]
fn is_drive_quota_exceeded_false_on_non_utf8_body() {
assert!(!is_drive_quota_exceeded(403, &[0xff, 0xfe]));
}
#[test]
fn is_drive_quota_exceeded_false_on_non_403_status() {
assert!(!is_drive_quota_exceeded(
429,
br#"{"error":{"errors":[{"reason":"userRateLimitExceeded"}]}}"#
));
}
#[test]
fn is_drive_quota_exceeded_true_on_matching_403_reason() {
assert!(is_drive_quota_exceeded(
403,
br#"{"error":{"errors":[{"reason":"userRateLimitExceeded"}]}}"#
));
}
#[test]
fn new_client_strips_trailing_slash() {
let client = DriveClient::new("https://www.googleapis.com/", &test_credentials()).unwrap();
assert_eq!(client.base_url(), "https://www.googleapis.com");
}
#[test]
fn new_client_preserves_clean_url() {
let client = DriveClient::new("https://www.googleapis.com", &test_credentials()).unwrap();
assert_eq!(client.base_url(), "https://www.googleapis.com");
}
#[test]
fn from_credentials_uses_drive_api_host() {
let env = crate::test_support::env::MapEnv::new();
let client = DriveClient::from_credentials_with(&env, &test_credentials()).unwrap();
assert_eq!(client.base_url(), "https://www.googleapis.com");
}
#[test]
fn from_credentials_honours_api_url_override() {
let env = crate::test_support::env::MapEnv::new().with(
crate::drive::auth::DRIVE_API_URL,
"http://proxy.example:8080",
);
let client = DriveClient::from_credentials_with(&env, &test_credentials()).unwrap();
assert_eq!(client.base_url(), "http://proxy.example:8080");
}
#[test]
fn from_credentials_ignores_empty_api_url_override() {
let env =
crate::test_support::env::MapEnv::new().with(crate::drive::auth::DRIVE_API_URL, "");
let client = DriveClient::from_credentials_with(&env, &test_credentials()).unwrap();
assert_eq!(client.base_url(), "https://www.googleapis.com");
}
#[test]
fn client_debug_never_mentions_session_field() {
let client = DriveClient::new("https://www.googleapis.com", &test_credentials()).unwrap();
let debug = format!("{client:?}");
assert!(!debug.contains("secret-1"));
assert!(!debug.contains("refresh-1"));
assert!(!debug.contains("session"));
assert!(debug.contains("DriveClient"));
}
async fn client_with_bootstrapped_token(server: &wiremock::MockServer) -> DriveClient {
wiremock::Mock::given(wiremock::matchers::method("POST"))
.and(wiremock::matchers::path("/token"))
.respond_with(
wiremock::ResponseTemplate::new(200).set_body_json(serde_json::json!({
"access_token": "bootstrap-token",
"expires_in": 3600,
})),
)
.up_to_n_times(1)
.with_priority(1)
.mount(server)
.await;
let mut client = DriveClient::new(&server.uri(), &test_credentials()).unwrap();
client.session = DriveSession::new_with_token_endpoint(
client.client.clone(),
&test_credentials(),
&format!("{}/token", server.uri()),
);
client
}
#[tokio::test]
async fn get_json_sends_bearer_auth_header() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.and(wiremock::matchers::header(
"Authorization",
"Bearer bootstrap-token",
))
.respond_with(
wiremock::ResponseTemplate::new(200).set_body_json(serde_json::json!({"ok": true})),
)
.expect(1)
.mount(&server)
.await;
let resp = client
.get_json(&format!("{}/test", server.uri()))
.await
.unwrap();
assert!(resp.status().is_success());
}
#[tokio::test]
async fn get_bytes_sends_bearer_auth_header_without_json_accept_header() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.and(wiremock::matchers::header(
"Authorization",
"Bearer bootstrap-token",
))
.respond_with(
wiremock::ResponseTemplate::new(200).set_body_bytes(b"raw bytes".to_vec()),
)
.expect(1)
.mount(&server)
.await;
let resp = client
.get_bytes(&format!("{}/test", server.uri()))
.await
.unwrap();
assert!(resp.status().is_success());
let bytes = resp.bytes().await.unwrap();
assert_eq!(bytes.as_ref(), b"raw bytes");
}
#[tokio::test]
async fn get_bytes_retries_on_429() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.respond_with(wiremock::ResponseTemplate::new(429).append_header("Retry-After", "0"))
.up_to_n_times(1)
.with_priority(1)
.mount(&server)
.await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.respond_with(wiremock::ResponseTemplate::new(200).set_body_bytes(b"ok".to_vec()))
.with_priority(2)
.mount(&server)
.await;
let resp = client
.get_bytes(&format!("{}/test", server.uri()))
.await
.unwrap();
assert_eq!(resp.status().as_u16(), 200);
}
#[tokio::test]
async fn get_bytes_refreshes_and_retries_once_on_401() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("POST"))
.and(wiremock::matchers::path("/token"))
.respond_with(
wiremock::ResponseTemplate::new(200).set_body_json(serde_json::json!({
"access_token": "refreshed-token",
"expires_in": 3600,
})),
)
.up_to_n_times(1)
.with_priority(2)
.mount(&server)
.await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.and(wiremock::matchers::header(
"Authorization",
"Bearer bootstrap-token",
))
.respond_with(wiremock::ResponseTemplate::new(401))
.expect(1)
.mount(&server)
.await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.and(wiremock::matchers::header(
"Authorization",
"Bearer refreshed-token",
))
.respond_with(wiremock::ResponseTemplate::new(200).set_body_bytes(b"ok".to_vec()))
.expect(1)
.mount(&server)
.await;
let resp = client
.get_bytes(&format!("{}/test", server.uri()))
.await
.unwrap();
assert_eq!(resp.status().as_u16(), 200);
}
#[tokio::test]
async fn get_bytes_propagates_network_errors() {
let client = DriveClient::new("http://127.0.0.1:1", &test_credentials()).unwrap();
let result = client.get_bytes("http://127.0.0.1:1/test").await;
assert!(result.is_err());
}
#[tokio::test]
async fn post_json_sends_body_and_bearer_auth() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("POST"))
.and(wiremock::matchers::path("/test"))
.and(wiremock::matchers::header(
"Authorization",
"Bearer bootstrap-token",
))
.and(wiremock::matchers::body_json(serde_json::json!({"k": "v"})))
.respond_with(wiremock::ResponseTemplate::new(200))
.expect(1)
.mount(&server)
.await;
let resp = client
.post_json(
&format!("{}/test", server.uri()),
&serde_json::json!({"k": "v"}),
)
.await
.unwrap();
assert!(resp.status().is_success());
}
#[tokio::test]
async fn patch_json_sends_body_and_bearer_auth() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("PATCH"))
.and(wiremock::matchers::path("/test"))
.and(wiremock::matchers::header(
"Authorization",
"Bearer bootstrap-token",
))
.and(wiremock::matchers::body_json(
serde_json::json!({"name": "new-name"}),
))
.respond_with(wiremock::ResponseTemplate::new(200))
.expect(1)
.mount(&server)
.await;
let resp = client
.patch_json(
&format!("{}/test", server.uri()),
&serde_json::json!({"name": "new-name"}),
)
.await
.unwrap();
assert!(resp.status().is_success());
}
#[tokio::test]
async fn get_json_retries_on_429() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.respond_with(wiremock::ResponseTemplate::new(429).append_header("Retry-After", "0"))
.up_to_n_times(1)
.with_priority(1)
.mount(&server)
.await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.respond_with(wiremock::ResponseTemplate::new(200))
.with_priority(2)
.mount(&server)
.await;
let resp = client
.get_json(&format!("{}/test", server.uri()))
.await
.unwrap();
assert_eq!(resp.status().as_u16(), 200);
}
#[tokio::test]
async fn get_json_retries_403_user_rate_limit_exceeded_then_succeeds() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.respond_with(
wiremock::ResponseTemplate::new(403)
.append_header("Retry-After", "0")
.set_body_json(serde_json::json!({
"error": {"message": "User Rate Limit Exceeded", "errors": [{"reason": "userRateLimitExceeded"}]}
})),
)
.up_to_n_times(1)
.with_priority(1)
.mount(&server)
.await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.respond_with(wiremock::ResponseTemplate::new(200))
.with_priority(2)
.mount(&server)
.await;
let resp = client
.get_json(&format!("{}/test", server.uri()))
.await
.unwrap();
assert_eq!(resp.status().as_u16(), 200);
}
#[tokio::test]
async fn get_json_does_not_retry_insufficient_permissions_403() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.respond_with(
wiremock::ResponseTemplate::new(403).set_body_json(serde_json::json!({
"error": {"message": "Insufficient Permission", "errors": [{"reason": "insufficientPermissions"}]}
})),
)
.expect(1)
.mount(&server)
.await;
let resp = client
.get_json(&format!("{}/test", server.uri()))
.await
.unwrap();
assert_eq!(resp.status().as_u16(), 403);
}
#[tokio::test]
async fn get_json_refreshes_and_retries_once_on_401() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("POST"))
.and(wiremock::matchers::path("/token"))
.respond_with(
wiremock::ResponseTemplate::new(200).set_body_json(serde_json::json!({
"access_token": "refreshed-token",
"expires_in": 3600,
})),
)
.up_to_n_times(1)
.with_priority(2)
.mount(&server)
.await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.and(wiremock::matchers::header(
"Authorization",
"Bearer bootstrap-token",
))
.respond_with(wiremock::ResponseTemplate::new(401))
.expect(1)
.mount(&server)
.await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.and(wiremock::matchers::header(
"Authorization",
"Bearer refreshed-token",
))
.respond_with(wiremock::ResponseTemplate::new(200))
.expect(1)
.mount(&server)
.await;
let resp = client
.get_json(&format!("{}/test", server.uri()))
.await
.unwrap();
assert_eq!(resp.status().as_u16(), 200);
}
#[tokio::test]
async fn get_json_does_not_retry_a_second_time_on_persistent_401() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("POST"))
.and(wiremock::matchers::path("/token"))
.respond_with(
wiremock::ResponseTemplate::new(200).set_body_json(serde_json::json!({
"access_token": "still-rejected-token",
"expires_in": 3600,
})),
)
.up_to_n_times(1)
.with_priority(2)
.mount(&server)
.await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.respond_with(
wiremock::ResponseTemplate::new(401).set_body_string("still unauthorized"),
)
.expect(2)
.mount(&server)
.await;
let resp = client
.get_json(&format!("{}/test", server.uri()))
.await
.unwrap();
assert_eq!(resp.status().as_u16(), 401);
}
#[tokio::test]
async fn response_to_error_extracts_drive_message_and_reason() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.respond_with(
wiremock::ResponseTemplate::new(403)
.append_header("Retry-After", "0")
.set_body_json(serde_json::json!({
"error": {
"message": "User Rate Limit Exceeded",
"errors": [{"reason": "userRateLimitExceeded"}],
}
})),
)
.mount(&server)
.await;
let resp = client
.get_json(&format!("{}/test", server.uri()))
.await
.unwrap();
let err = DriveClient::response_to_error(resp).await;
let msg = err.to_string();
assert!(msg.contains("User Rate Limit Exceeded"));
assert!(msg.contains("userRateLimitExceeded"));
assert_eq!(err.reason(), Some("userRateLimitExceeded"));
}
#[tokio::test]
async fn response_to_error_omits_reason_suffix_when_absent() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.respond_with(
wiremock::ResponseTemplate::new(400).set_body_json(serde_json::json!({
"error": {
"message": "Invalid request",
}
})),
)
.mount(&server)
.await;
let resp = client
.get_json(&format!("{}/test", server.uri()))
.await
.unwrap();
let err = DriveClient::response_to_error(resp).await;
let msg = err.to_string();
assert!(msg.contains("Invalid request"));
assert!(!msg.contains("reason:"));
assert_eq!(err.reason(), None);
}
#[tokio::test]
async fn response_to_error_falls_back_to_raw_body_when_not_drive_shaped() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.respond_with(wiremock::ResponseTemplate::new(500).set_body_string("internal error"))
.mount(&server)
.await;
let resp = client
.get_json(&format!("{}/test", server.uri()))
.await
.unwrap();
let err = DriveClient::response_to_error(resp).await;
assert!(err.to_string().contains("internal error"));
}
#[tokio::test]
async fn get_json_propagates_network_errors() {
let client = DriveClient::new("http://127.0.0.1:1", &test_credentials()).unwrap();
let result = client.get_json("http://127.0.0.1:1/test").await;
assert!(result.is_err());
}
#[tokio::test]
async fn get_parsed_errors_on_malformed_json_response() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.respond_with(wiremock::ResponseTemplate::new(200).set_body_string("not json"))
.mount(&server)
.await;
let result: Result<serde_json::Value> = client
.get_parsed(&format!("{}/test", server.uri()), "test context")
.await;
assert!(result.is_err());
}
#[tokio::test]
async fn get_parsed_errors_on_non_success_status_without_parsing_the_body() {
let server = wiremock::MockServer::start().await;
let client = client_with_bootstrapped_token(&server).await;
wiremock::Mock::given(wiremock::matchers::method("GET"))
.and(wiremock::matchers::path("/test"))
.respond_with(
wiremock::ResponseTemplate::new(404).set_body_json(serde_json::json!({
"error": {"message": "File not found"}
})),
)
.mount(&server)
.await;
let result: Result<serde_json::Value> = client
.get_parsed(&format!("{}/test", server.uri()), "test context")
.await;
let err = result.unwrap_err();
assert!(err.to_string().contains("File not found"));
}
}