mod error;
pub mod handlers;
pub mod pagination;
mod response;
pub mod retry;
pub use cirrus_auth as auth;
pub use reqwest;
pub use auth::{AuthError, AuthSession, SharedAuth};
pub use bytes::Bytes;
pub use error::{CirrusError, CirrusResult, SalesforceError};
pub use handlers::bulk::{BulkIngestSpec, BulkQuerySpec};
pub use handlers::composite::{
BatchRequest, BatchSubrequest, CompositeRequest, CompositeSubrequest,
};
pub use handlers::metadata::{
DeployMessage, DeployOptions, DeployRequest, DeployResultDetails, DeployResultInnerDetails,
DeployStatus, MetadataHandler, RunTestResults, TestLevel,
};
pub use handlers::sobjects::BlobUploadSpec;
pub use pagination::Records;
pub use response::LimitInfo;
pub use response::{
ApiVersion, BatchResponse, BatchSubresult, BulkColumnDelimiter, BulkIngestJob, BulkJobState,
BulkJobStateChange, BulkLineEnding, BulkOperation, BulkQueryJob, BulkQueryResults,
CompositeError, CompositeResponse, CompositeSubresponse, CompositeTreeResponse,
CompositeTreeResult, DescribeGlobal, EventLogFileRecord, ExecuteAnonymousResult, Limit,
OrgLimits, QueryResult, SObjectCollectionResult, SObjectCreateResult, SObjectMetadata,
SearchResult,
};
pub use retry::RetryPolicy;
use reqwest::header::{HeaderMap, HeaderValue, USER_AGENT};
use serde::Serialize;
use serde::de::DeserializeOwned;
use std::sync::{Arc, RwLock};
pub const DEFAULT_API_VERSION: &str = "v66.0";
pub(crate) const DEFAULT_USER_AGENT: &str = concat!(
"cirrus/",
env!("CARGO_PKG_VERSION"),
" (Rust SDK for Salesforce)"
);
#[derive(Clone)]
pub struct Cirrus {
client: reqwest::Client,
auth: SharedAuth,
api_version: String,
retry_policy: RetryPolicy,
last_limit_info: Arc<RwLock<Option<LimitInfo>>>,
}
impl std::fmt::Debug for Cirrus {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Cirrus")
.field("api_version", &self.api_version)
.field("instance_url", &self.auth.instance_url())
.field("retry_policy", &self.retry_policy)
.finish_non_exhaustive()
}
}
enum AuthRetry {
Retry,
Done,
}
impl Cirrus {
pub fn builder() -> CirrusBuilder {
CirrusBuilder::default()
}
pub fn api_version(&self) -> &str {
&self.api_version
}
pub fn http_client(&self) -> &reqwest::Client {
&self.client
}
pub fn auth(&self) -> &SharedAuth {
&self.auth
}
pub fn retry_policy(&self) -> &RetryPolicy {
&self.retry_policy
}
pub fn last_limit_info(&self) -> Option<LimitInfo> {
self.last_limit_info.read().ok().and_then(|guard| *guard)
}
fn update_limit_info(&self, headers: &reqwest::header::HeaderMap) {
let Some(value) = headers.get("Sforce-Limit-Info") else {
return;
};
let Ok(s) = value.to_str() else { return };
let Some(info) = LimitInfo::parse(s) else {
return;
};
tracing::debug!(
target: "cirrus::limit_info",
used = info.used,
allowed = info.allowed,
"captured Sforce-Limit-Info",
);
if let Ok(mut guard) = self.last_limit_info.write() {
*guard = Some(info);
}
}
pub(crate) fn resolve_url(&self, path: &str) -> String {
if path.starts_with("http://") || path.starts_with("https://") {
path.to_string()
} else if path.starts_with('/') {
let rest = path.trim_start_matches('/');
format!("{}/{}", self.auth.instance_url(), rest)
} else {
format!(
"{}/services/data/{}/{}",
self.auth.instance_url(),
self.api_version,
path
)
}
}
pub(crate) fn versioned_segments(&self, segments: &[&str]) -> CirrusResult<String> {
let base = format!(
"{}/services/data/{}/",
self.auth.instance_url(),
self.api_version
);
let mut url = url::Url::parse(&base)?;
url.path_segments_mut()
.map_err(|()| CirrusError::InvalidResponse("instance URL is not hierarchical".into()))?
.pop_if_empty()
.extend(segments);
Ok(url.to_string())
}
pub async fn get<R: DeserializeOwned>(&self, path: &str) -> CirrusResult<R> {
let url = self.resolve_url(path);
self.send::<R, (), ()>(reqwest::Method::GET, &url, None, None)
.await
}
pub async fn get_with_query<R, Q>(&self, path: &str, query: &Q) -> CirrusResult<R>
where
R: DeserializeOwned,
Q: Serialize + ?Sized,
{
let url = self.resolve_url(path);
self.send::<R, Q, ()>(reqwest::Method::GET, &url, Some(query), None)
.await
}
pub async fn post<R, B>(&self, path: &str, body: &B) -> CirrusResult<R>
where
R: DeserializeOwned,
B: Serialize + ?Sized,
{
let url = self.resolve_url(path);
self.send::<R, (), B>(reqwest::Method::POST, &url, None, Some(body))
.await
}
pub async fn put<R, B>(&self, path: &str, body: &B) -> CirrusResult<R>
where
R: DeserializeOwned,
B: Serialize + ?Sized,
{
let url = self.resolve_url(path);
self.send::<R, (), B>(reqwest::Method::PUT, &url, None, Some(body))
.await
}
pub async fn patch<R, B>(&self, path: &str, body: &B) -> CirrusResult<R>
where
R: DeserializeOwned,
B: Serialize + ?Sized,
{
let url = self.resolve_url(path);
self.send::<R, (), B>(reqwest::Method::PATCH, &url, None, Some(body))
.await
}
pub async fn delete<R: DeserializeOwned>(&self, path: &str) -> CirrusResult<R> {
let url = self.resolve_url(path);
self.send::<R, (), ()>(reqwest::Method::DELETE, &url, None, None)
.await
}
pub(crate) async fn send_at<R, Q, B>(
&self,
method: reqwest::Method,
url: &str,
query: Option<&Q>,
body: Option<&B>,
) -> CirrusResult<R>
where
R: DeserializeOwned,
Q: Serialize + ?Sized,
B: Serialize + ?Sized,
{
self.send(method, url, query, body).await
}
pub async fn request_builder(
&self,
method: reqwest::Method,
path: &str,
) -> CirrusResult<reqwest::RequestBuilder> {
let url = self.resolve_url(path);
let token = self.auth.access_token().await?;
Ok(self.client.request(method, url).bearer_auth(&*token))
}
pub async fn execute(&self, request: reqwest::Request) -> CirrusResult<reqwest::Response> {
self.client.execute(request).await.map_err(Into::into)
}
async fn auth_retry_decision(
&self,
is_retryable_401: bool,
token: &str,
auth_retried: bool,
) -> CirrusResult<AuthRetry> {
if auth_retried || !is_retryable_401 {
return Ok(AuthRetry::Done);
}
tracing::warn!(
target: "cirrus::auth",
"received 401; invalidating cached token and retrying once",
);
self.auth.invalidate(token).await;
let fresh = self.auth.access_token().await?;
if *fresh == *token {
tracing::warn!(
target: "cirrus::auth",
"auth session returned same token after invalidate; surfacing 401 (likely static auth or scope/permission issue)",
);
return Ok(AuthRetry::Done);
}
Ok(AuthRetry::Retry)
}
async fn dispatch<T, MakeReq, Parse>(
&self,
method: &reqwest::Method,
make_request: MakeReq,
parse: Parse,
) -> CirrusResult<T>
where
MakeReq: Fn(&str) -> CirrusResult<reqwest::RequestBuilder>,
Parse: Fn(u16, reqwest::header::HeaderMap, bytes::Bytes) -> CirrusResult<T>,
{
let mut auth_retried = false;
let mut attempt: u32 = 0;
loop {
let token = self.auth.access_token().await?;
let result: CirrusResult<T> = loop {
let request = make_request(&token)?;
match request.send().await {
Ok(response) => {
let status = response.status().as_u16();
let headers = response.headers().clone();
self.update_limit_info(&headers);
if retry::should_retry_status(&self.retry_policy, method, status, attempt) {
let _ = response.bytes().await;
let retry_after = retry::parse_retry_after(&headers);
let delay =
retry::compute_delay(&self.retry_policy, attempt, retry_after);
tokio::time::sleep(delay).await;
attempt += 1;
continue;
}
match response.bytes().await {
Ok(bytes) => break parse(status, headers, bytes),
Err(e) => break Err(e.into()),
}
}
Err(e) => {
let err: CirrusError = e.into();
if retry::should_retry_network(&self.retry_policy, method, &err, attempt) {
let delay = retry::compute_delay(&self.retry_policy, attempt, None);
tokio::time::sleep(delay).await;
attempt += 1;
continue;
}
break Err(err);
}
}
};
let is_retryable_401 = matches!(&result, Err(CirrusError::Api { status: 401, .. }));
match self
.auth_retry_decision(is_retryable_401, &token, auth_retried)
.await?
{
AuthRetry::Retry => {
auth_retried = true;
attempt = 0;
continue;
}
AuthRetry::Done => return result,
}
}
}
async fn send<R, Q, B>(
&self,
method: reqwest::Method,
url: &str,
query: Option<&Q>,
body: Option<&B>,
) -> CirrusResult<R>
where
R: DeserializeOwned,
Q: Serialize + ?Sized,
B: Serialize + ?Sized,
{
self.dispatch(
&method,
|token: &str| {
let mut request = self.client.request(method.clone(), url).bearer_auth(token);
if let Some(q) = query {
request = request.query(q);
}
if let Some(b) = body {
request = request.json(b);
}
Ok(request)
},
|status, _headers, bytes| response::parse_response_bytes(status, &bytes),
)
.await
}
pub(crate) async fn send_with_body<R>(
&self,
method: reqwest::Method,
path: &str,
body: bytes::Bytes,
content_type: &str,
) -> CirrusResult<R>
where
R: DeserializeOwned,
{
let url = self.resolve_url(path);
self.dispatch(
&method,
|token: &str| {
Ok(self
.client
.request(method.clone(), &url)
.bearer_auth(token)
.header(reqwest::header::CONTENT_TYPE, content_type)
.body(body.clone()))
},
|status, _headers, bytes| response::parse_response_bytes(status, &bytes),
)
.await
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn send_multipart<R>(
&self,
method: reqwest::Method,
path: &str,
json_part_name: &str,
json_bytes: Vec<u8>,
blob_part_name: &str,
blob_filename: &str,
blob_content_type: &str,
blob: bytes::Bytes,
) -> CirrusResult<R>
where
R: DeserializeOwned,
{
let url = self.resolve_url(path);
self.dispatch(
&method,
|token: &str| {
let json_part = reqwest::multipart::Part::bytes(json_bytes.clone())
.mime_str("application/json")
.map_err(|e| {
CirrusError::InvalidHeader(format!("invalid JSON part content-type: {e}"))
})?;
let blob_part = reqwest::multipart::Part::stream(reqwest::Body::from(blob.clone()))
.file_name(blob_filename.to_string())
.mime_str(blob_content_type)
.map_err(|e| {
CirrusError::InvalidHeader(format!("invalid blob part content-type: {e}"))
})?;
let form = reqwest::multipart::Form::new()
.part(json_part_name.to_string(), json_part)
.part(blob_part_name.to_string(), blob_part);
Ok(self
.client
.request(method.clone(), &url)
.bearer_auth(token)
.multipart(form))
},
|status, _headers, bytes| response::parse_response_bytes(status, &bytes),
)
.await
}
pub(crate) async fn fetch_raw(
&self,
method: reqwest::Method,
path: &str,
accept: &str,
query: Option<&[(&str, &str)]>,
) -> CirrusResult<(reqwest::header::HeaderMap, bytes::Bytes)> {
let url = self.resolve_url(path);
self.dispatch(
&method,
|token: &str| {
let mut request = self
.client
.request(method.clone(), &url)
.bearer_auth(token)
.header(reqwest::header::ACCEPT, accept);
if let Some(q) = query {
request = request.query(q);
}
Ok(request)
},
|status, headers, bytes| {
if (200..300).contains(&status) {
Ok((headers, bytes))
} else {
Err(response::parse_error_response(status, &bytes))
}
},
)
.await
}
pub(crate) async fn send_with_headers(
&self,
method: reqwest::Method,
path: &str,
query: Option<&[(&str, &str)]>,
extra_headers: &[(&str, &str)],
) -> CirrusResult<(u16, bytes::Bytes)> {
let url = self.resolve_url(path);
self.dispatch(
&method,
|token: &str| {
let mut request = self.client.request(method.clone(), &url).bearer_auth(token);
for (name, value) in extra_headers {
request = request.header(*name, *value);
}
if let Some(q) = query {
request = request.query(q);
}
Ok(request)
},
|status, _headers, bytes| {
if (200..300).contains(&status) || status == 304 {
Ok((status, bytes))
} else {
Err(response::parse_error_response(status, &bytes))
}
},
)
.await
}
}
#[derive(Default)]
pub struct CirrusBuilder {
auth: Option<SharedAuth>,
api_version: Option<String>,
user_agent: Option<String>,
http_client: Option<reqwest::Client>,
retry_policy: Option<RetryPolicy>,
}
impl CirrusBuilder {
pub fn auth(mut self, auth: SharedAuth) -> Self {
self.auth = Some(auth);
self
}
pub fn api_version(mut self, version: impl Into<String>) -> Self {
self.api_version = Some(version.into());
self
}
pub fn user_agent(mut self, ua: impl Into<String>) -> Self {
self.user_agent = Some(ua.into());
self
}
pub fn http_client(mut self, client: reqwest::Client) -> Self {
self.http_client = Some(client);
self
}
pub fn retry_policy(mut self, policy: RetryPolicy) -> Self {
self.retry_policy = Some(policy);
self
}
pub fn build(self) -> CirrusResult<Cirrus> {
let auth = self.auth.ok_or(CirrusError::MissingField("auth"))?;
let client = if let Some(c) = self.http_client {
c
} else {
let ua = self.user_agent.as_deref().unwrap_or(DEFAULT_USER_AGENT);
let mut headers = HeaderMap::new();
headers.insert(
USER_AGENT,
HeaderValue::from_str(ua).map_err(|e| CirrusError::InvalidHeader(e.to_string()))?,
);
reqwest::Client::builder()
.default_headers(headers)
.build()
.map_err(CirrusError::HttpClient)?
};
Ok(Cirrus {
client,
auth,
api_version: self
.api_version
.unwrap_or_else(|| DEFAULT_API_VERSION.to_string()),
retry_policy: self.retry_policy.unwrap_or_default(),
last_limit_info: Arc::new(RwLock::new(None)),
})
}
pub async fn build_with_latest_version(self) -> CirrusResult<Cirrus> {
let bootstrap = self.build()?;
let latest = bootstrap.latest_api_version().await?;
Ok(Cirrus {
api_version: latest,
..bootstrap
})
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
use crate::auth::StaticTokenAuth;
use std::sync::Arc;
fn fixture(instance: &str) -> Cirrus {
let auth = Arc::new(StaticTokenAuth::new("tok", instance));
Cirrus::builder().auth(auth).build().unwrap()
}
#[test]
fn build_requires_auth() {
let err = Cirrus::builder().build().unwrap_err();
assert!(matches!(err, CirrusError::MissingField("auth")));
}
#[test]
fn resolve_url_versioned_for_relative_path() {
let sf = fixture("https://my.salesforce.com");
let url = sf.resolve_url("limits");
assert_eq!(url, "https://my.salesforce.com/services/data/v66.0/limits");
}
#[test]
fn resolve_url_versioned_for_nested_relative_path() {
let sf = fixture("https://my.salesforce.com");
let url = sf.resolve_url("sobjects/Account/001");
assert_eq!(
url,
"https://my.salesforce.com/services/data/v66.0/sobjects/Account/001"
);
}
#[test]
fn resolve_url_instance_rooted_for_leading_slash() {
let sf = fixture("https://my.salesforce.com");
let url = sf.resolve_url("/services/data");
assert_eq!(url, "https://my.salesforce.com/services/data");
}
#[test]
fn resolve_url_passthrough_for_https_url() {
let sf = fixture("https://my.salesforce.com");
let absolute = "https://other.example.com/some/path";
assert_eq!(sf.resolve_url(absolute), absolute);
}
#[test]
fn resolve_url_passthrough_for_http_url() {
let sf = fixture("https://my.salesforce.com");
let absolute = "http://localhost:1234/path";
assert_eq!(sf.resolve_url(absolute), absolute);
}
#[test]
fn api_version_can_be_overridden() {
let auth = Arc::new(StaticTokenAuth::new("tok", "https://my.salesforce.com"));
let sf = Cirrus::builder()
.auth(auth)
.api_version("v61.0")
.build()
.unwrap();
assert_eq!(sf.api_version(), "v61.0");
assert!(sf.resolve_url("x").contains("/v61.0/"));
}
mod escape_hatch {
use super::*;
use serde_json::{Value, json};
use wiremock::matchers::{body_json, header, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
fn server_fixture(uri: String) -> Cirrus {
let auth = Arc::new(StaticTokenAuth::new("tok", uri));
Cirrus::builder().auth(auth).build().unwrap()
}
#[tokio::test]
async fn get_resolves_relative_path_as_versioned() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.and(header("authorization", "Bearer tok"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"ok": true})))
.mount(&server)
.await;
let sf = server_fixture(server.uri());
let v: Value = sf.get("limits").await.unwrap();
assert_eq!(v["ok"], true);
}
#[tokio::test]
async fn get_resolves_leading_slash_as_instance_rooted() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/apexrest/foo"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"called": "apex"})))
.mount(&server)
.await;
let sf = server_fixture(server.uri());
let v: Value = sf.get("/services/apexrest/foo").await.unwrap();
assert_eq!(v["called"], "apex");
}
#[tokio::test]
async fn get_passes_through_absolute_url() {
let other = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/some/other/path"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"hit": "other"})))
.mount(&other)
.await;
let sf = server_fixture("https://unused.invalid".to_string());
let target = format!("{}/some/other/path", other.uri());
let v: Value = sf.get(&target).await.unwrap();
assert_eq!(v["hit"], "other");
}
#[tokio::test]
async fn post_sends_json_body() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/services/data/v66.0/composite/batch"))
.and(body_json(json!({"batchRequests": []})))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"results": []})))
.mount(&server)
.await;
let sf = server_fixture(server.uri());
let v: Value = sf
.post("composite/batch", &json!({"batchRequests": []}))
.await
.unwrap();
assert!(v["results"].is_array());
}
#[tokio::test]
async fn put_sends_json_body() {
let server = MockServer::start().await;
Mock::given(method("PUT"))
.and(path("/services/data/v66.0/custom/resource"))
.and(body_json(json!({"k": "v"})))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"updated": true})))
.mount(&server)
.await;
let sf = server_fixture(server.uri());
let v: Value = sf.put("custom/resource", &json!({"k": "v"})).await.unwrap();
assert_eq!(v["updated"], true);
}
#[tokio::test]
async fn patch_sends_json_body() {
let server = MockServer::start().await;
Mock::given(method("PATCH"))
.and(path("/services/data/v66.0/sobjects/Account/001"))
.and(body_json(json!({"Name": "X"})))
.respond_with(ResponseTemplate::new(204))
.mount(&server)
.await;
let sf = server_fixture(server.uri());
sf.patch::<(), _>("sobjects/Account/001", &json!({"Name": "X"}))
.await
.unwrap();
}
#[tokio::test]
async fn delete_handles_204() {
let server = MockServer::start().await;
Mock::given(method("DELETE"))
.and(path("/services/data/v66.0/sobjects/Account/001"))
.respond_with(ResponseTemplate::new(204))
.mount(&server)
.await;
let sf = server_fixture(server.uri());
sf.delete::<()>("sobjects/Account/001").await.unwrap();
}
#[tokio::test]
async fn request_builder_pre_injects_bearer_auth() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.and(header("authorization", "Bearer tok"))
.and(header("x-custom", "added-by-caller"))
.respond_with(ResponseTemplate::new(200))
.mount(&server)
.await;
let sf = server_fixture(server.uri());
let resp = sf
.request_builder(reqwest::Method::GET, "limits")
.await
.unwrap()
.header("X-Custom", "added-by-caller")
.send()
.await
.unwrap();
assert_eq!(resp.status().as_u16(), 200);
}
#[tokio::test]
async fn execute_runs_caller_built_request() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/raw/path"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"raw": true})))
.mount(&server)
.await;
let sf = server_fixture(server.uri());
let url = format!("{}/raw/path", server.uri());
let req = sf.http_client().get(&url).build().unwrap();
let resp = sf.execute(req).await.unwrap();
assert_eq!(resp.status().as_u16(), 200);
let body: Value = resp.json().await.unwrap();
assert_eq!(body["raw"], true);
}
}
mod retry_and_limits {
use super::*;
use crate::RetryPolicy;
use serde_json::{Value, json};
use std::sync::Arc;
use std::time::Duration;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
fn fast_retry_policy() -> RetryPolicy {
RetryPolicy {
base_delay: Duration::ZERO,
max_delay: Duration::ZERO,
jitter: false,
..RetryPolicy::default()
}
}
fn fixture_with_policy(uri: String, policy: RetryPolicy) -> Cirrus {
let auth = Arc::new(StaticTokenAuth::new("tok", uri));
Cirrus::builder()
.auth(auth)
.retry_policy(policy)
.build()
.unwrap()
}
#[tokio::test]
async fn retries_429_until_success() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(ResponseTemplate::new(429))
.up_to_n_times(2)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"ok": true})))
.mount(&server)
.await;
let sf = fixture_with_policy(server.uri(), fast_retry_policy());
let v: Value = sf.get("limits").await.unwrap();
assert_eq!(v["ok"], true);
}
#[tokio::test]
async fn retries_503_until_success() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(ResponseTemplate::new(503))
.up_to_n_times(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"ok": true})))
.mount(&server)
.await;
let sf = fixture_with_policy(server.uri(), fast_retry_policy());
let v: Value = sf.get("limits").await.unwrap();
assert_eq!(v["ok"], true);
}
#[tokio::test]
async fn surfaces_error_after_max_retries_exhausted() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(ResponseTemplate::new(503).set_body_json(json!([{
"errorCode": "SERVER_UNAVAILABLE",
"message": "Service Unavailable"
}])))
.expect(4)
.mount(&server)
.await;
let sf = fixture_with_policy(server.uri(), fast_retry_policy());
let err = sf.get::<Value>("limits").await.unwrap_err();
match err {
CirrusError::Api { status, .. } => assert_eq!(status, 503),
other => panic!("expected Api error, got {other:?}"),
}
}
#[tokio::test]
async fn does_not_retry_4xx_caller_errors() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(ResponseTemplate::new(404).set_body_json(json!([{
"errorCode": "NOT_FOUND",
"message": "not found"
}])))
.expect(1)
.mount(&server)
.await;
let sf = fixture_with_policy(server.uri(), fast_retry_policy());
let err = sf.get::<Value>("limits").await.unwrap_err();
assert!(matches!(err, CirrusError::Api { status: 404, .. }));
}
#[tokio::test]
async fn does_not_retry_500_on_post() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/services/data/v66.0/sobjects/Account"))
.respond_with(ResponseTemplate::new(500).set_body_json(json!([{
"errorCode": "INTERNAL_ERROR",
"message": "boom"
}])))
.expect(1)
.mount(&server)
.await;
let sf = fixture_with_policy(server.uri(), fast_retry_policy());
let err = sf
.post::<Value, _>("sobjects/Account", &json!({"Name": "Acme"}))
.await
.unwrap_err();
assert!(matches!(err, CirrusError::Api { status: 500, .. }));
}
#[tokio::test]
async fn does_not_retry_503_on_post() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/services/data/v66.0/sobjects/Account"))
.respond_with(ResponseTemplate::new(503))
.expect(1)
.mount(&server)
.await;
let sf = fixture_with_policy(server.uri(), fast_retry_policy());
let err = sf
.post::<Value, _>("sobjects/Account", &json!({"Name": "Acme"}))
.await
.unwrap_err();
assert!(matches!(err, CirrusError::Api { status: 503, .. }));
}
#[tokio::test]
async fn retries_500_on_get_when_idempotent_5xx_enabled() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(ResponseTemplate::new(500))
.up_to_n_times(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"ok": true})))
.mount(&server)
.await;
let sf = fixture_with_policy(server.uri(), fast_retry_policy());
let v: Value = sf.get("limits").await.unwrap();
assert_eq!(v["ok"], true);
}
#[tokio::test]
async fn none_policy_disables_retries_entirely() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(ResponseTemplate::new(429))
.expect(1)
.mount(&server)
.await;
let sf = fixture_with_policy(server.uri(), RetryPolicy::none());
let err = sf.get::<Value>("limits").await.unwrap_err();
assert!(matches!(err, CirrusError::Api { status: 429, .. }));
}
#[tokio::test]
async fn captures_sforce_limit_info_on_response() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(
ResponseTemplate::new(200)
.set_body_json(json!({"ok": true}))
.insert_header("Sforce-Limit-Info", "api-usage=42/15000"),
)
.mount(&server)
.await;
let sf = fixture_with_policy(server.uri(), RetryPolicy::none());
assert!(sf.last_limit_info().is_none());
let _: Value = sf.get("limits").await.unwrap();
let info = sf.last_limit_info().expect("limit info should be set");
assert_eq!(info.used, 42);
assert_eq!(info.allowed, 15000);
assert_eq!(info.remaining(), 14958);
}
#[tokio::test]
async fn limit_info_updates_on_subsequent_requests() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(
ResponseTemplate::new(200)
.set_body_json(json!({"ok": true}))
.insert_header("Sforce-Limit-Info", "api-usage=10/100"),
)
.up_to_n_times(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(
ResponseTemplate::new(200)
.set_body_json(json!({"ok": true}))
.insert_header("Sforce-Limit-Info", "api-usage=11/100"),
)
.mount(&server)
.await;
let sf = fixture_with_policy(server.uri(), RetryPolicy::none());
let _: Value = sf.get("limits").await.unwrap();
assert_eq!(sf.last_limit_info().unwrap().used, 10);
let _: Value = sf.get("limits").await.unwrap();
assert_eq!(sf.last_limit_info().unwrap().used, 11);
}
#[tokio::test]
async fn malformed_limit_info_header_is_ignored() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(
ResponseTemplate::new(200)
.set_body_json(json!({"ok": true}))
.insert_header("Sforce-Limit-Info", "junk-data=oops"),
)
.mount(&server)
.await;
let sf = fixture_with_policy(server.uri(), RetryPolicy::none());
let _: Value = sf.get("limits").await.unwrap();
assert!(sf.last_limit_info().is_none());
}
#[tokio::test]
async fn retry_after_header_overrides_backoff() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(ResponseTemplate::new(429).insert_header("Retry-After", "0"))
.up_to_n_times(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"ok": true})))
.mount(&server)
.await;
let sf = fixture_with_policy(server.uri(), fast_retry_policy());
let v: Value = sf.get("limits").await.unwrap();
assert_eq!(v["ok"], true);
}
}
mod auth_refresh {
use super::*;
use crate::auth::{AuthResult, AuthSession, SharedAuth};
use async_trait::async_trait;
use serde_json::{Value, json};
use std::borrow::Cow;
use std::sync::Arc;
use std::sync::Mutex;
use std::sync::atomic::{AtomicUsize, Ordering};
use wiremock::matchers::{header, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
struct RotatingAuth {
instance_url: String,
tokens: Vec<String>,
access_count: AtomicUsize,
invalidations: Mutex<Vec<String>>,
}
impl RotatingAuth {
fn new(instance_url: impl Into<String>, tokens: Vec<&str>) -> Self {
Self {
instance_url: instance_url.into(),
tokens: tokens.into_iter().map(String::from).collect(),
access_count: AtomicUsize::new(0),
invalidations: Mutex::new(Vec::new()),
}
}
}
#[async_trait]
impl AuthSession for RotatingAuth {
async fn access_token(&self) -> AuthResult<Cow<'_, str>> {
let n = self.access_count.fetch_add(1, Ordering::SeqCst);
let idx = n.min(self.tokens.len() - 1);
Ok(Cow::Borrowed(&self.tokens[idx]))
}
fn instance_url(&self) -> &str {
&self.instance_url
}
async fn invalidate(&self, stale_token: &str) {
if let Ok(mut g) = self.invalidations.lock() {
g.push(stale_token.to_string());
}
}
}
fn fixture(_uri: String, auth: SharedAuth) -> Cirrus {
Cirrus::builder()
.auth(auth)
.retry_policy(crate::RetryPolicy::none()) .build()
.unwrap()
}
#[tokio::test]
async fn refreshes_token_on_401_and_retries_once() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.and(header("authorization", "Bearer old"))
.respond_with(ResponseTemplate::new(401).set_body_json(json!([{
"errorCode": "INVALID_SESSION_ID",
"message": "Session expired or invalid"
}])))
.expect(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.and(header("authorization", "Bearer new"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"ok": true})))
.expect(1)
.mount(&server)
.await;
let auth = Arc::new(RotatingAuth::new(server.uri(), vec!["old", "new"]));
let sf = fixture(server.uri(), auth.clone());
let v: Value = sf.get("limits").await.unwrap();
assert_eq!(v["ok"], true);
let inv = auth.invalidations.lock().unwrap();
assert_eq!(inv.len(), 1);
assert_eq!(inv[0], "old");
}
#[tokio::test]
async fn transient_retry_budget_resets_after_auth_refresh() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.and(header("authorization", "Bearer old"))
.respond_with(ResponseTemplate::new(503))
.up_to_n_times(1)
.expect(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.and(header("authorization", "Bearer old"))
.respond_with(ResponseTemplate::new(401).set_body_json(json!([{
"errorCode": "INVALID_SESSION_ID",
"message": "Session expired or invalid"
}])))
.expect(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.and(header("authorization", "Bearer new"))
.respond_with(ResponseTemplate::new(503))
.up_to_n_times(1)
.expect(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.and(header("authorization", "Bearer new"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"ok": true})))
.expect(1)
.mount(&server)
.await;
let auth = Arc::new(RotatingAuth::new(server.uri(), vec!["old", "new"]));
let sf = Cirrus::builder()
.auth(auth.clone())
.retry_policy(crate::RetryPolicy {
max_retries: 1,
base_delay: std::time::Duration::ZERO,
max_delay: std::time::Duration::ZERO,
jitter: false,
..crate::RetryPolicy::default()
})
.build()
.unwrap();
let v: Value = sf.get("limits").await.unwrap();
assert_eq!(v["ok"], true);
let inv = auth.invalidations.lock().unwrap();
assert_eq!(*inv, vec!["old"]);
}
#[tokio::test]
async fn surfaces_401_when_refresh_returns_same_token() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(ResponseTemplate::new(401).set_body_json(json!([{
"errorCode": "INVALID_SESSION_ID",
"message": "..."
}])))
.expect(1)
.mount(&server)
.await;
let auth = Arc::new(RotatingAuth::new(server.uri(), vec!["only"]));
let sf = fixture(server.uri(), auth);
let err = sf.get::<Value>("limits").await.unwrap_err();
assert!(matches!(err, CirrusError::Api { status: 401, .. }));
}
#[tokio::test]
async fn second_401_after_refresh_surfaces_without_third_attempt() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(ResponseTemplate::new(401).set_body_json(json!([{
"errorCode": "INSUFFICIENT_ACCESS",
"message": "..."
}])))
.expect(2)
.mount(&server)
.await;
let auth = Arc::new(RotatingAuth::new(server.uri(), vec!["t1", "t2"]));
let sf = fixture(server.uri(), auth);
let err = sf.get::<Value>("limits").await.unwrap_err();
assert!(matches!(err, CirrusError::Api { status: 401, .. }));
}
#[tokio::test]
async fn does_not_invalidate_on_non_401_errors() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/services/data/v66.0/limits"))
.respond_with(ResponseTemplate::new(403).set_body_json(json!([{
"errorCode": "INSUFFICIENT_ACCESS",
"message": "..."
}])))
.expect(1)
.mount(&server)
.await;
let auth = Arc::new(RotatingAuth::new(server.uri(), vec!["t1", "t2"]));
let sf = fixture(server.uri(), auth.clone());
let _ = sf.get::<Value>("limits").await;
let inv = auth.invalidations.lock().unwrap();
assert!(inv.is_empty());
}
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod property_tests {
use super::*;
use crate::auth::StaticTokenAuth;
use proptest::prelude::*;
use std::sync::Arc;
fn fixture(instance: &str) -> Cirrus {
let auth = Arc::new(StaticTokenAuth::new("tok", instance));
Cirrus::builder().auth(auth).build().unwrap()
}
fn path_segment() -> impl Strategy<Value = String> {
"[A-Za-z0-9_./%=&+-]{1,32}".prop_filter("no double-slash runs in well-formed paths", |s| {
!s.contains("//")
})
}
fn nonempty_segment() -> impl Strategy<Value = String> {
"[A-Za-z0-9_-]{1,32}"
}
proptest! {
#[test]
fn resolve_url_never_emits_double_slash(path in path_segment()) {
let sf = fixture("https://my.salesforce.com");
let url = sf.resolve_url(&path);
let after_scheme = url.split_once("://").map(|(_, rest)| rest).unwrap_or(&url);
prop_assert!(
!after_scheme.contains("//"),
"resolve_url({path:?}) produced double slash: {url}",
);
prop_assert!(url::Url::parse(&url).is_ok(), "url should parse: {url}");
}
#[test]
fn resolve_url_passes_through_absolute_urls(host in "[a-z0-9-]{1,20}", path in path_segment()) {
let sf = fixture("https://my.salesforce.com");
let absolute = format!("https://{host}.example.com/{path}");
prop_assert_eq!(sf.resolve_url(&absolute), absolute);
}
#[test]
fn resolve_url_instance_rooted_skips_version(
rest in path_segment().prop_filter("rest follows the leading slash", |s| !s.starts_with('/')),
) {
let sf = fixture("https://my.salesforce.com");
let url = sf.resolve_url(&format!("/{rest}"));
prop_assert_eq!(url, format!("https://my.salesforce.com/{rest}"));
}
#[test]
fn versioned_segments_round_trip(
seg1 in nonempty_segment(),
seg2 in nonempty_segment(),
) {
let sf = fixture("https://my.salesforce.com");
let url_str = sf.versioned_segments(&[&seg1, &seg2]).unwrap();
let parsed = url::Url::parse(&url_str).unwrap();
let segments: Vec<&str> = parsed
.path_segments()
.map(|s| s.collect())
.unwrap_or_default();
prop_assert_eq!(segments.len(), 5, "got segments {:?} from {}", segments, url_str);
prop_assert_eq!(segments[0], "services");
prop_assert_eq!(segments[1], "data");
prop_assert_eq!(segments[3], seg1);
prop_assert_eq!(segments[4], seg2);
}
#[test]
fn versioned_segments_never_emits_double_slash(
segs in proptest::collection::vec(nonempty_segment(), 1..6),
) {
let sf = fixture("https://my.salesforce.com");
let refs: Vec<&str> = segs.iter().map(String::as_str).collect();
let url = sf.versioned_segments(&refs).unwrap();
let after_scheme = url.split_once("://").map(|(_, rest)| rest).unwrap_or(&url);
prop_assert!(
!after_scheme.contains("//"),
"got double slash in {url}",
);
}
}
#[test]
fn versioned_segments_percent_encodes_reserved_slash() {
let sf = fixture("https://my.salesforce.com");
let url = sf
.versioned_segments(&["sobjects", "Account", "Ext_Id__c", "abc/def"])
.unwrap();
let parsed = url::Url::parse(&url).unwrap();
let segs: Vec<&str> = parsed.path_segments().unwrap().collect();
assert_eq!(
segs.len(),
7,
"expected 7 segments (services, data, version, sobjects, Account, Ext_Id__c, abc%2Fdef), got {segs:?}",
);
assert!(
segs[6].contains("%2F") || segs[6].contains("%2f"),
"slash in external-ID value must be percent-encoded; got {:?}",
segs[6],
);
}
}