use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, OnceLock};
use std::time::Duration;
use hotdata::auth::{async_trait, BearerTokenError, BearerTokenProvider};
use hotdata::models::QueryRequest;
use hotdata::{Client, Configuration, RetryPolicy, UploadOptions};
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, Request, ResponseTemplate};
fn fast_retry() -> RetryPolicy {
RetryPolicy {
max_retries: 3,
base_backoff: Duration::from_millis(1),
max_backoff: Duration::from_millis(5),
deadline: Duration::from_secs(10),
jitter: 0.0,
}
}
#[derive(Debug)]
struct SequenceProvider {
values: Vec<String>,
calls: AtomicUsize,
}
impl SequenceProvider {
fn new(values: &[&str]) -> Arc<Self> {
Arc::new(Self {
values: values.iter().map(|s| (*s).to_owned()).collect(),
calls: AtomicUsize::new(0),
})
}
fn call_count(&self) -> usize {
self.calls.load(Ordering::SeqCst)
}
}
#[async_trait]
impl BearerTokenProvider for SequenceProvider {
async fn bearer_value(&self) -> Result<String, BearerTokenError> {
let i = self.calls.fetch_add(1, Ordering::SeqCst);
Ok(self.values[i.min(self.values.len() - 1)].clone())
}
}
#[derive(Debug)]
struct FailingProvider;
#[async_trait]
impl BearerTokenProvider for FailingProvider {
async fn bearer_value(&self) -> Result<String, BearerTokenError> {
Err(BearerTokenError::Malformed(
"refresh token expired".to_owned(),
))
}
}
fn config_for(base_url: &str) -> Configuration {
Configuration {
base_path: base_url.to_owned(),
user_agent: Some("hotdata-rust-test".to_owned()),
..Configuration::default()
}
}
fn recorded_bearers(requests: &[Request]) -> Vec<Option<String>> {
requests
.iter()
.map(|r| {
r.headers
.get("authorization")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_owned())
})
.collect()
}
async fn mount_workspaces(server: &MockServer) {
Mock::given(method("GET"))
.and(path("/v1/workspaces"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"ok": true,
"workspaces": [],
})))
.mount(server)
.await;
}
#[tokio::test]
async fn provider_value_reaches_the_wire() {
let server = MockServer::start().await;
mount_workspaces(&server).await;
let provider = SequenceProvider::new(&["provided-token"]);
let mut config = config_for(&server.uri());
config.token_provider = Some(provider.clone());
let client = Client::from_configuration(config);
client
.workspaces()
.list(None)
.await
.expect("list_workspaces should succeed");
let requests = server.received_requests().await.expect("requests recorded");
assert_eq!(
recorded_bearers(&requests),
vec![Some("Bearer provided-token".to_owned())],
"the provider's value must be sent as the Authorization bearer"
);
assert_eq!(provider.call_count(), 1);
}
#[tokio::test]
async fn provider_is_consulted_per_request() {
let server = MockServer::start().await;
mount_workspaces(&server).await;
let provider = SequenceProvider::new(&["token-first", "token-second"]);
let mut config = config_for(&server.uri());
config.token_provider = Some(provider.clone());
let client = Client::from_configuration(config);
client.workspaces().list(None).await.expect("first call");
client.workspaces().list(None).await.expect("second call");
let requests = server.received_requests().await.expect("requests recorded");
assert_eq!(
recorded_bearers(&requests),
vec![
Some("Bearer token-first".to_owned()),
Some("Bearer token-second".to_owned()),
],
"each request must carry the value the provider returned for it"
);
assert_eq!(
provider.call_count(),
2,
"the provider must be asked once per request"
);
}
#[tokio::test]
async fn provider_takes_precedence_over_static_token() {
let server = MockServer::start().await;
mount_workspaces(&server).await;
let mut config = config_for(&server.uri());
config.bearer_access_token = Some("static-token".to_owned());
config.token_provider = Some(SequenceProvider::new(&["provided-token"]));
let client = Client::from_configuration(config);
client.workspaces().list(None).await.expect("call succeeds");
let requests = server.received_requests().await.expect("requests recorded");
assert_eq!(
recorded_bearers(&requests),
vec![Some("Bearer provided-token".to_owned())]
);
}
#[tokio::test]
async fn static_bearer_still_works_without_a_provider() {
let server = MockServer::start().await;
mount_workspaces(&server).await;
let mut config = config_for(&server.uri());
config.bearer_access_token = Some("static-token".to_owned());
assert!(config.token_provider.is_none());
let client = Client::from_configuration(config);
client.workspaces().list(None).await.expect("call succeeds");
let requests = server.received_requests().await.expect("requests recorded");
assert_eq!(
recorded_bearers(&requests),
vec![Some("Bearer static-token".to_owned())]
);
}
#[tokio::test]
async fn builder_installs_no_provider() {
let server = MockServer::start().await;
mount_workspaces(&server).await;
let client = Client::builder()
.api_token("hd_opaque")
.workspace_id("ws_x")
.base_url(server.uri())
.build()
.expect("build should succeed");
assert!(
client.configuration().token_provider.is_none(),
"the builder must not install a token provider"
);
assert_eq!(
client.configuration().bearer_access_token.as_deref(),
Some("hd_opaque")
);
client.workspaces().list(None).await.expect("call succeeds");
let requests = server.received_requests().await.expect("requests recorded");
assert_eq!(
recorded_bearers(&requests),
vec![Some("Bearer hd_opaque".to_owned())],
"exactly one request, carrying the API token — no exchange round trip"
);
}
static LOG_BUF: OnceLock<Mutex<Vec<String>>> = OnceLock::new();
fn log_buf() -> &'static Mutex<Vec<String>> {
LOG_BUF.get_or_init(|| Mutex::new(Vec::new()))
}
struct CaptureLogger;
impl log::Log for CaptureLogger {
fn enabled(&self, _meta: &log::Metadata) -> bool {
true
}
fn log(&self, record: &log::Record) {
log_buf()
.lock()
.unwrap()
.push(format!("{} {}", record.level(), record.args()));
}
fn flush(&self) {}
}
static LOGGER: CaptureLogger = CaptureLogger;
#[tokio::test]
async fn failing_provider_logs_and_sends_no_bearer() {
log::set_logger(&LOGGER).expect("logger installs once");
log::set_max_level(log::LevelFilter::Warn);
let server = MockServer::start().await;
mount_workspaces(&server).await;
let mut config = config_for(&server.uri());
config.bearer_access_token = Some("static-token".to_owned());
config.token_provider = Some(Arc::new(FailingProvider));
let client = Client::from_configuration(config);
client
.workspaces()
.list(None)
.await
.expect("the request is still sent, just unauthenticated");
let requests = server.received_requests().await.expect("requests recorded");
assert_eq!(
recorded_bearers(&requests),
vec![None],
"a failed resolution must send no Authorization header at all"
);
let logged = log_buf().lock().unwrap().join("\n");
assert!(
logged.contains("bearer token resolution failed"),
"the failure must be logged; captured records were:\n{logged}"
);
assert!(
logged.contains("refresh token expired"),
"the underlying cause must reach the log; captured records were:\n{logged}"
);
}
#[tokio::test]
async fn upload_file_resolves_create_and_finalize_through_the_provider() {
let server = MockServer::start().await;
let storage_url = format!("{}/storage/single", server.uri());
let contents = b"hello per-request bearer";
Mock::given(method("POST"))
.and(path("/v1/uploads"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"finalize_token": "ftok_single",
"headers": {},
"mode": "single",
"upload_id": "upl_single",
"url": storage_url,
})))
.mount(&server)
.await;
Mock::given(method("PUT"))
.and(path("/storage/single"))
.respond_with(ResponseTemplate::new(200).insert_header("ETag", "\"single-etag\""))
.mount(&server)
.await;
Mock::given(method("POST"))
.and(path("/v1/uploads/upl_single/finalize"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"created_at": "2026-06-25T00:00:00Z",
"size_bytes": contents.len(),
"status": "ready",
"upload_id": "upl_single",
})))
.mount(&server)
.await;
let provider = SequenceProvider::new(&["token-at-create", "token-at-finalize"]);
let mut config = config_for(&server.uri());
config.token_provider = Some(provider.clone());
let client = Client::from_configuration(config);
let file = std::env::temp_dir().join(format!(
"hotdata-bearer-provider-{}",
uuid::Uuid::new_v4().simple()
));
std::fs::write(&file, contents).expect("writing the temp upload file should succeed");
let result = client.upload_file(&file, UploadOptions::default()).await;
let _ = std::fs::remove_file(&file);
result.expect("single upload should succeed");
let requests = server.received_requests().await.expect("requests recorded");
let bearer_of = |p: &str| -> Option<String> {
requests
.iter()
.find(|r| r.url.path() == p)
.unwrap_or_else(|| panic!("a request to {p} should have been made"))
.headers
.get("authorization")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_owned())
};
assert_eq!(
bearer_of("/v1/uploads"),
Some("Bearer token-at-create".to_owned()),
"create-session must resolve through the provider"
);
assert_eq!(
bearer_of("/v1/uploads/upl_single/finalize"),
Some("Bearer token-at-finalize".to_owned()),
"finalize must resolve through the provider independently of create-session"
);
assert_eq!(
provider.call_count(),
2,
"exactly the two API legs resolve a bearer"
);
assert_eq!(
bearer_of("/storage/single"),
None,
"the storage PUT must carry no Authorization header"
);
}
fn sql(text: &str) -> QueryRequest {
QueryRequest {
sql: text.to_owned(),
..Default::default()
}
}
#[tokio::test]
async fn send_query_resolves_per_request() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/query"))
.respond_with(ResponseTemplate::new(200).set_body_string(r#"{"unparseable":true}"#))
.mount(&server)
.await;
let provider = SequenceProvider::new(&["query-first", "query-second"]);
let mut config = config_for(&server.uri());
config.token_provider = Some(provider.clone());
let client = Client::from_configuration(config);
let _ = client.query(sql("SELECT 1")).await;
let _ = client.query(sql("SELECT 2")).await;
let requests = server.received_requests().await.expect("requests recorded");
assert_eq!(
recorded_bearers(&requests),
vec![
Some("Bearer query-first".to_owned()),
Some("Bearer query-second".to_owned()),
],
"send_query must resolve a fresh bearer for each POST /v1/query"
);
assert_eq!(provider.call_count(), 2);
}
#[tokio::test]
async fn submit_query_resolves_per_request() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/query"))
.respond_with(ResponseTemplate::new(200).set_body_string(r#"{"unparseable":true}"#))
.mount(&server)
.await;
let provider = SequenceProvider::new(&["submit-first", "submit-second"]);
let mut config = config_for(&server.uri());
config.token_provider = Some(provider.clone());
let client = Client::from_configuration(config);
let _ = client.submit_query(sql("SELECT 1"), None).await;
let _ = client.submit_query(sql("SELECT 2"), None).await;
let requests = server.received_requests().await.expect("requests recorded");
assert_eq!(
recorded_bearers(&requests),
vec![
Some("Bearer submit-first".to_owned()),
Some("Bearer submit-second".to_owned()),
],
"submit_query must resolve a fresh bearer for each call"
);
assert_eq!(provider.call_count(), 2);
}
#[cfg(feature = "arrow")]
#[tokio::test]
async fn arrow_fetch_resolves_per_request() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/results/res_1"))
.respond_with(ResponseTemplate::new(200).set_body_string("not-arrow-ipc"))
.mount(&server)
.await;
let provider = SequenceProvider::new(&["arrow-first", "arrow-second"]);
let mut config = config_for(&server.uri());
config.token_provider = Some(provider.clone());
let client = Client::from_configuration(config);
let _ = client.get_result_arrow("res_1", "dbid_1", None, None).await;
let _ = client.get_result_arrow("res_1", "dbid_1", None, None).await;
let requests = server.received_requests().await.expect("requests recorded");
assert_eq!(
recorded_bearers(&requests),
vec![
Some("Bearer arrow-first".to_owned()),
Some("Bearer arrow-second".to_owned()),
],
"the Arrow fetch must resolve a fresh bearer for each call"
);
assert_eq!(provider.call_count(), 2);
}
#[tokio::test]
async fn retry_after_429_re_resolves_the_bearer() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/workspaces"))
.respond_with(ResponseTemplate::new(429))
.up_to_n_times(1)
.with_priority(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/v1/workspaces"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"ok": true,
"workspaces": [],
})))
.with_priority(2)
.mount(&server)
.await;
let provider = SequenceProvider::new(&["stale-token", "refreshed-token"]);
let mut config = config_for(&server.uri());
config.token_provider = Some(provider.clone());
config.retry = fast_retry();
let client = Client::from_configuration(config);
client
.workspaces()
.list(None)
.await
.expect("the retry should succeed");
let requests = server.received_requests().await.expect("requests recorded");
assert_eq!(
recorded_bearers(&requests),
vec![
Some("Bearer stale-token".to_owned()),
Some("Bearer refreshed-token".to_owned()),
],
"the 429 retry must re-resolve instead of replaying attempt 0's bearer"
);
assert_eq!(
provider.call_count(),
2,
"one resolve for the initial attempt, one for the retry"
);
}
#[tokio::test]
async fn retry_after_429_replays_static_bearer_unchanged() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/workspaces"))
.respond_with(ResponseTemplate::new(429))
.up_to_n_times(1)
.with_priority(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/v1/workspaces"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"ok": true,
"workspaces": [],
})))
.with_priority(2)
.mount(&server)
.await;
let mut config = config_for(&server.uri());
config.bearer_access_token = Some("static-token".to_owned());
config.retry = fast_retry();
let client = Client::from_configuration(config);
client
.workspaces()
.list(None)
.await
.expect("the retry should succeed");
let requests = server.received_requests().await.expect("requests recorded");
assert_eq!(
recorded_bearers(&requests),
vec![
Some("Bearer static-token".to_owned()),
Some("Bearer static-token".to_owned()),
],
"with no provider the static bearer is replayed, as in 0.12.0"
);
}
#[tokio::test]
async fn storage_part_put_never_gains_a_bearer_on_retry() {
let server = MockServer::start().await;
let part_size = 5usize;
let contents: Vec<u8> = (0u8..8).collect();
let part_urls: Vec<String> = (1..=2)
.map(|i| format!("{}/storage/part/{i}", server.uri()))
.collect();
Mock::given(method("POST"))
.and(path("/v1/uploads"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"finalize_token": "ftok_retry",
"headers": {},
"mode": "multipart",
"part_size": part_size,
"part_urls": part_urls,
"upload_id": "upl_retry",
})))
.mount(&server)
.await;
Mock::given(method("PUT"))
.and(path("/storage/part/1"))
.respond_with(ResponseTemplate::new(429))
.up_to_n_times(1)
.with_priority(1)
.mount(&server)
.await;
Mock::given(method("PUT"))
.and(path("/storage/part/1"))
.respond_with(ResponseTemplate::new(200).insert_header("ETag", "\"etag-part-1\""))
.with_priority(2)
.mount(&server)
.await;
Mock::given(method("PUT"))
.and(path("/storage/part/2"))
.respond_with(ResponseTemplate::new(200).insert_header("ETag", "\"etag-part-2\""))
.mount(&server)
.await;
Mock::given(method("POST"))
.and(path("/v1/uploads/upl_retry/finalize"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"created_at": "2026-06-25T00:00:00Z",
"size_bytes": contents.len(),
"status": "ready",
"upload_id": "upl_retry",
})))
.mount(&server)
.await;
let provider = SequenceProvider::new(&["tok-a", "tok-b", "tok-c", "tok-d"]);
let mut config = config_for(&server.uri());
config.token_provider = Some(provider.clone());
config.retry = fast_retry();
let client = Client::from_configuration(config);
let file = std::env::temp_dir().join(format!(
"hotdata-bearer-retry-{}",
uuid::Uuid::new_v4().simple()
));
std::fs::write(&file, &contents).expect("writing the temp upload file should succeed");
let result = client.upload_file(&file, UploadOptions::default()).await;
let _ = std::fs::remove_file(&file);
result.expect("the upload should succeed through the storage retry");
let requests = server.received_requests().await.expect("requests recorded");
let part_puts: Vec<_> = requests
.iter()
.filter(|r| r.url.path().starts_with("/storage/part/"))
.collect();
assert_eq!(
part_puts.len(),
3,
"expected part 1 to retry after its 429, got {} part PUTs",
part_puts.len()
);
for put in &part_puts {
assert!(
put.headers.get("authorization").is_none(),
"part PUT to {} must carry no Authorization header",
put.url.path()
);
}
}