use std::fmt;
use std::future::Future;
use std::path::{Path, PathBuf};
use std::time::Duration;
use asupersync::http::compress::{
DecompressionLimit, Decompressor, GzipDecompressor, IdentityDecompressor,
};
use asupersync::http::h1::Http1Client;
use asupersync::http::{
Client as AsupersyncHttpClient, ClientError as AsupersyncClientError, Method, ParsedUrl,
Request, Response, Scheme, StatusCode,
};
use asupersync::net::TcpStream;
use asupersync::tls::{Certificate, TlsConnector, TlsConnectorBuilder};
use asupersync::{CancelKind, Cx, Time};
use franken_snowflake_core::budget::Budget;
use franken_snowflake_core::cancel::{CancelPolicy, CancelReason, cancel_policy};
use franken_snowflake_core::ids::{RequestId, StatementHandle};
use franken_snowflake_core::redact::{REDACTION_PLACEHOLDER, redact};
use serde::{Deserialize, Serialize, Serializer, ser::SerializeStruct};
pub mod capture;
pub const VERSION: &str = env!("CARGO_PKG_VERSION");
const SQL_API_STATEMENTS_PATH: &str = "/api/v2/statements";
const HEADER_AUTHORIZATION: &str = "Authorization";
const HEADER_TOKEN_TYPE: &str = "X-Snowflake-Authorization-Token-Type";
const HEADER_CONTENT_TYPE: &str = "Content-Type";
const HEADER_ACCEPT: &str = "Accept";
const HEADER_ACCEPT_ENCODING: &str = "Accept-Encoding";
const HEADER_USER_AGENT: &str = "User-Agent";
pub const USER_AGENT: &str = concat!("franken-snowflake/", env!("CARGO_PKG_VERSION"));
const HEADER_CONTENT_ENCODING: &str = "Content-Encoding";
const JSON_MEDIA_TYPE: &str = "application/json";
const PARTITION_ACCEPT_ENCODING: &str = "gzip, identity";
pub const SNOWFLAKE_SQL_API_RESUBMIT_DOC_URL: &str = "https://docs.snowflake.com/en/developer-guide/sql-api/submitting-requests#resubmitting-a-request-to-execute-sql-statements";
pub const SNOWFLAKE_SQL_API_RESUBMIT_DOC_CONSULTED: &str = "2026-06-25";
pub type TransportOutcome<T> = franken_snowflake_core::outcome::SnowflakeOutcome<T>;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct TransportCaps;
#[derive(Clone)]
pub struct SnowflakeHttpClient<H = LiveHttp> {
config: TransportConfig,
client: H,
}
pub trait RawHttp {
fn send(
&self,
cx: &Cx,
method: Method,
url: String,
headers: Vec<(String, String)>,
body: Vec<u8>,
timeout: Option<Duration>,
) -> impl Future<Output = Result<Response, AsupersyncClientError>>;
}
impl RawHttp for AsupersyncHttpClient {
async fn send(
&self,
cx: &Cx,
method: Method,
url: String,
headers: Vec<(String, String)>,
body: Vec<u8>,
timeout: Option<Duration>,
) -> Result<Response, AsupersyncClientError> {
let mut builder = self
.request_builder(method, url)
.headers(
headers
.iter()
.map(|(name, value)| (name.as_str(), value.as_str())),
)
.body(body);
if let Some(timeout) = timeout {
builder = builder.timeout(timeout);
}
builder.send(cx).await
}
}
#[derive(Clone)]
pub struct PemBundleHttp {
connector: TlsConnector,
max_body_bytes: usize,
}
impl PemBundleHttp {
pub fn from_pem_file(path: &Path, limits: &BodyLimits) -> Result<Self, TransportError> {
let refused =
|reason: String| TransportError::new(TransportErrorCode::TlsRootPolicyRefused, reason);
let roots = Certificate::from_pem_file(path).map_err(|error| {
refused(format!(
"cannot read the CA bundle {}: {error}",
path.display()
))
})?;
if roots.is_empty() {
return Err(refused(format!(
"the CA bundle {} holds no PEM certificate",
path.display()
)));
}
let connector = TlsConnectorBuilder::new()
.add_root_certificates(roots)
.alpn_protocols(vec![b"http/1.1".to_vec()])
.build()
.map_err(|error| {
refused(format!(
"the CA bundle {} is unusable: {error}",
path.display()
))
})?;
Ok(Self {
connector,
max_body_bytes: max_response_body_bytes(limits),
})
}
}
fn max_response_body_bytes(limits: &BodyLimits) -> usize {
[
limits.max_submit_response_bytes,
limits.max_poll_response_bytes,
limits.max_partition_compressed_bytes,
]
.into_iter()
.max()
.and_then(|bytes| usize::try_from(bytes).ok())
.unwrap_or(usize::MAX)
}
impl RawHttp for PemBundleHttp {
async fn send(
&self,
cx: &Cx,
method: Method,
url: String,
headers: Vec<(String, String)>,
body: Vec<u8>,
timeout: Option<Duration>,
) -> Result<Response, AsupersyncClientError> {
let parsed = ParsedUrl::parse(&url)?;
if parsed.scheme != Scheme::Https {
return Err(AsupersyncClientError::InvalidUrl(format!(
"{url}: a CA bundle applies to https endpoints only"
)));
}
let exchange = async {
if cx.checkpoint().is_err() {
return Err(AsupersyncClientError::Cancelled);
}
let tcp = TcpStream::connect(format!("{}:{}", parsed.host, parsed.port))
.await
.map_err(AsupersyncClientError::ConnectError)?;
let domain = parsed.host.trim_start_matches('[').trim_end_matches(']');
let tls = self
.connector
.connect(domain, tcp)
.await
.map_err(|error| AsupersyncClientError::TlsError(error.to_string()))?;
if cx.checkpoint().is_err() {
return Err(AsupersyncClientError::Cancelled);
}
let request = Request::builder(method, parsed.path.clone())
.header("Host", parsed.authority())
.headers(
headers
.into_iter()
.filter(|(name, _)| !name.eq_ignore_ascii_case("host")),
)
.body(body)
.build();
let (response, _connection, _body_withheld) =
Http1Client::request_with_io_and_max_body_size(tls, request, self.max_body_bytes)
.await?;
Ok(response)
};
let Some(limit) = timeout else {
return exchange.await;
};
match asupersync::time::timeout(
asupersync::time::wall_now(),
limit,
std::pin::pin!(exchange),
)
.await
{
Ok(result) => result,
Err(_elapsed) => Err(AsupersyncClientError::DeadlineExceeded),
}
}
}
#[derive(Clone)]
pub enum LiveHttp {
NativeRoots(AsupersyncHttpClient),
PemBundle(PemBundleHttp),
Capturing(Box<LiveHttp>, std::sync::Arc<capture::TranscriptRecorder>),
}
impl RawHttp for LiveHttp {
async fn send(
&self,
cx: &Cx,
method: Method,
url: String,
headers: Vec<(String, String)>,
body: Vec<u8>,
timeout: Option<Duration>,
) -> Result<Response, AsupersyncClientError> {
match self {
Self::NativeRoots(client) => client.send(cx, method, url, headers, body, timeout).await,
Self::PemBundle(client) => client.send(cx, method, url, headers, body, timeout).await,
Self::Capturing(inner, recorder) => {
let request = (
method.as_str().to_owned(),
url.clone(),
headers.clone(),
body.clone(),
);
let result = Box::pin(inner.send(cx, method, url, headers, body, timeout)).await;
if let Ok(response) = &result {
recorder.record(&request.0, &request.1, &request.2, &request.3, response);
}
result
}
}
}
}
impl<H> fmt::Debug for SnowflakeHttpClient<H> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SnowflakeHttpClient")
.field("config", &self.config)
.field("client", &"<asupersync-http-client>")
.finish()
}
}
impl SnowflakeHttpClient {
pub fn for_runtime(config: TransportConfig) -> Result<Self, TransportError> {
let client = match &config.tls_roots {
TlsRootPolicy::NativeRoots => LiveHttp::NativeRoots(
AsupersyncHttpClient::builder()
.max_body_size(max_response_body_bytes(&config.limits))
.build(),
),
TlsRootPolicy::ExplicitPemBundle(path) => {
LiveHttp::PemBundle(PemBundleHttp::from_pem_file(path, &config.limits)?)
}
TlsRootPolicy::TestOnlyInsecureDisabledByDefault => {
return Err(TransportError::new(
TransportErrorCode::TlsRootPolicyRefused,
"certificate verification cannot be disabled for a live transport",
));
}
};
Ok(Self { config, client })
}
#[must_use]
pub fn capturing(self, recorder: std::sync::Arc<capture::TranscriptRecorder>) -> Self {
Self {
config: self.config,
client: LiveHttp::Capturing(Box::new(self.client), recorder),
}
}
}
impl<H: RawHttp> SnowflakeHttpClient<H> {
#[must_use]
pub fn new(config: TransportConfig, client: H) -> Self {
Self { config, client }
}
#[must_use]
pub const fn config(&self) -> &TransportConfig {
&self.config
}
pub fn submit_plan(&self, request: &SubmitHttpRequest) -> Result<WireRequest, TransportError> {
self.wire_request(
Method::Post,
request.route.clone(),
request.body.clone(),
&request.auth,
request.retry_resubmit,
)
}
pub fn poll_plan(&self, request: &PollHttpRequest) -> Result<WireRequest, TransportError> {
let route = TransportRoute::Poll {
handle: request.statement_handle.clone(),
};
self.wire_request(Method::Get, route, Vec::new(), &request.auth, false)
}
pub fn partition_plan(
&self,
request: &PartitionHttpRequest,
) -> Result<WireRequest, TransportError> {
let route = TransportRoute::Partition {
handle: request.statement_handle.clone(),
partition: request.partition,
};
self.wire_request(Method::Get, route, Vec::new(), &request.auth, false)
}
pub fn cancel_plan(&self, request: &CancelHttpRequest) -> Result<WireRequest, TransportError> {
let route = TransportRoute::Cancel {
handle: request.statement_handle.clone(),
};
self.wire_request(Method::Post, route, Vec::new(), &request.auth, false)
}
pub async fn submit_statement(
&self,
cx: &Cx,
request: SubmitHttpRequest,
) -> TransportOutcome<SubmitHttpResponse> {
self.execute(cx, request, TransportRouteKind::Submit).await
}
pub async fn poll_statement(
&self,
cx: &Cx,
request: PollHttpRequest,
) -> TransportOutcome<PollHttpResponse> {
self.execute(cx, request, TransportRouteKind::Poll).await
}
pub async fn fetch_partition(
&self,
cx: &Cx,
request: PartitionHttpRequest,
) -> TransportOutcome<PartitionBody> {
self.execute(cx, request, TransportRouteKind::Partition)
.await
}
pub async fn cancel_statement(
&self,
cx: &Cx,
request: CancelHttpRequest,
) -> TransportOutcome<CancelHttpResponse> {
self.execute(cx, request, TransportRouteKind::Cancel).await
}
pub async fn cancel_after_local_cancel(
&self,
cleanup_cx: &Cx,
auth: AuthorizationDescriptor,
statement_handle: StatementHandle,
reason: CancelReason,
) -> TransportOutcome<CancelHttpResponse> {
let policy = cancel_policy(reason.kind);
if !franken_snowflake_core::cancel::attempts_remote_cancel(policy) {
return TransportOutcome::cancelled(reason);
}
debug_assert!(matches!(
policy,
CancelPolicy::RemoteCancelAndReceipt
| CancelPolicy::BoundedDrain
| CancelPolicy::RetryOrDegrade
));
run_with_cancellation_mask(
cleanup_cx,
self.cancel_statement(
cleanup_cx,
CancelHttpRequest {
auth,
statement_handle,
reason_kind: reason.kind,
},
),
)
.await
}
pub async fn cancel_orphaned_statement(
&self,
cleanup_cx: &Cx,
auth: AuthorizationDescriptor,
statement_handle: StatementHandle,
) -> TransportOutcome<CancelHttpResponse> {
run_with_cancellation_mask(
cleanup_cx,
self.cancel_statement(
cleanup_cx,
CancelHttpRequest {
auth,
statement_handle,
reason_kind: CancelKind::FailFast,
},
),
)
.await
}
pub async fn stream_partitions<S>(
&self,
cx: &Cx,
request: PartitionStreamRequest,
sink: &mut S,
) -> TransportOutcome<PartitionStreamSummary>
where
S: PartitionSink,
{
if cx.checkpoint().is_err() {
return TransportOutcome::cancelled(cancel_reason_or(
cx,
CancelReason::parent_cancelled,
));
}
let summary = match request.plan() {
Ok(summary) => summary,
Err(error) => return TransportOutcome::err(error.into_snowflake_error()),
};
for partition in request.seed_partitions.iter().cloned() {
if cx.checkpoint().is_err() {
let reason = cancel_reason_or(cx, CancelReason::parent_cancelled);
return self
.cancel_stream_after_local_cancel(cx, &request, &summary, reason)
.await;
}
if let Err(error) = sink.accept(cx, partition).await {
return TransportOutcome::err(error.into_snowflake_error());
}
if cx.checkpoint().is_err() {
let reason = cancel_reason_or(cx, CancelReason::parent_cancelled);
return self
.cancel_stream_after_local_cancel(cx, &request, &summary, reason)
.await;
}
}
for partition in request.first_partition..request.end_partition_exclusive {
if cx.checkpoint().is_err() {
let reason = cancel_reason_or(cx, CancelReason::parent_cancelled);
return self
.cancel_stream_after_local_cancel(cx, &request, &summary, reason)
.await;
}
let body: PartitionBody = match self
.execute(
cx,
PlannedTransportRequest::partition(
request.auth.clone(),
request.statement_handle.clone(),
partition,
request.child_budget,
),
TransportRouteKind::Partition,
)
.await
{
TransportOutcome::Ok(body) => body,
TransportOutcome::Err(error) => return TransportOutcome::err(error),
TransportOutcome::Cancelled(reason) => return TransportOutcome::cancelled(reason),
TransportOutcome::Panicked(payload) => return TransportOutcome::panicked(payload),
};
let decoded = DecodedPartition {
partition,
body: body.body,
compression: body.compression,
};
if let Err(error) = sink.accept(cx, decoded).await {
return TransportOutcome::err(error.into_snowflake_error());
}
if cx.checkpoint().is_err() {
let reason = cancel_reason_or(cx, CancelReason::parent_cancelled);
return self
.cancel_stream_after_local_cancel(cx, &request, &summary, reason)
.await;
}
}
TransportOutcome::ok(summary)
}
async fn cancel_stream_after_local_cancel(
&self,
cx: &Cx,
request: &PartitionStreamRequest,
_summary: &PartitionStreamSummary,
reason: CancelReason,
) -> TransportOutcome<PartitionStreamSummary> {
if request.remote_cancel_on_local_cancel {
let cleanup = self
.cancel_after_local_cancel(
cx,
request.auth.clone(),
request.statement_handle.clone(),
reason.clone(),
)
.await;
stream_cancel_cleanup_outcome(cleanup, reason)
} else {
TransportOutcome::cancelled(reason)
}
}
}
fn stream_cancel_cleanup_outcome<T>(
cleanup: TransportOutcome<T>,
reason: CancelReason,
) -> TransportOutcome<PartitionStreamSummary> {
match cleanup {
TransportOutcome::Panicked(payload) => TransportOutcome::panicked(payload),
TransportOutcome::Ok(_) | TransportOutcome::Err(_) | TransportOutcome::Cancelled(_) => {
TransportOutcome::cancelled(reason)
}
}
}
impl<H: RawHttp> SnowflakeHttpClient<H> {
async fn execute<R, T>(
&self,
cx: &Cx,
request: R,
route_kind: TransportRouteKind,
) -> TransportOutcome<T>
where
R: Into<PlannedTransportRequest>,
T: FromResponseBody,
{
let mut retry_spent_ms = 0_u64;
let planned = request.into();
let wire = match self.wire_request(
route_kind.method(),
planned.route.clone(),
planned.body.clone(),
&planned.auth,
planned.retry_resubmit,
) {
Ok(wire) => wire,
Err(error) => return TransportOutcome::err(error.into_snowflake_error()),
};
let mut attempt = 1_u32;
loop {
if cx.checkpoint().is_err() {
return TransportOutcome::cancelled(cancel_reason_or(
cx,
CancelReason::parent_cancelled,
));
}
let attempt_budget =
self.config
.retry
.attempt_budget(cx.budget(), planned.budget, attempt);
let budget_now = asupersync::time::wall_now();
if let Some(reason) = budget_exhaustion_reason_at(attempt_budget, budget_now) {
return TransportOutcome::cancelled(reason);
}
let result = self
.client
.send(
cx,
route_kind.method(),
wire.url.clone(),
wire.headers
.iter()
.map(|h| (h.name.clone(), h.value.clone()))
.collect(),
wire.body.clone(),
sooner(
budget_timeout_at(attempt_budget, budget_now),
self.config.attempt_timeout_for(route_kind),
),
)
.await;
match result {
Ok(response) => {
let retry_after_ms = retry_after_ms(response.headers.as_slice());
let retryable = is_retryable_status(response.status);
if retryable {
if !route_allows_automatic_retry(&wire.route) {
return TransportOutcome::err(
TransportError::new(
TransportErrorCode::NonIdempotentSubmitRetryRefused,
format!(
"{} returned retryable HTTP status {}, but automatic retry requires requestId plus retry=true",
route_kind.as_str(),
response.status
),
)
.into_snowflake_error(),
);
}
if let Some(retry) = self.config.retry.next_retry(
planned.request_id.as_ref(),
route_kind,
attempt,
retry_after_ms,
retry_spent_ms,
) {
if let Err(reason) = wait_retry_delay(cx, retry.delay).await {
return TransportOutcome::cancelled(reason);
}
attempt = attempt.saturating_add(1);
retry_spent_ms = retry.spent_after_ms;
continue;
}
return TransportOutcome::err(
TransportError::new(
TransportErrorCode::RetryBudgetExhausted,
format!(
"{} exhausted retry budget after {attempt} attempts",
route_kind.as_str()
),
)
.into_snowflake_error(),
);
}
return match T::from_response(response, self.config.limits, planned.partition) {
Ok(value) => TransportOutcome::ok(value),
Err(error) => TransportOutcome::err(error.into_snowflake_error()),
};
}
Err(error) if error.is_cancelled() => {
return TransportOutcome::cancelled(cancel_reason_or(
cx,
CancelReason::parent_cancelled,
));
}
Err(AsupersyncClientError::DeadlineExceeded) => {
return TransportOutcome::cancelled(cancel_reason_or(
cx,
CancelReason::deadline,
));
}
Err(error) => {
if route_allows_automatic_retry(&wire.route)
&& let Some(retry) = self.config.retry.next_retry(
planned.request_id.as_ref(),
route_kind,
attempt,
None,
retry_spent_ms,
)
{
if let Err(reason) = wait_retry_delay(cx, retry.delay).await {
return TransportOutcome::cancelled(reason);
}
attempt = attempt.saturating_add(1);
retry_spent_ms = retry.spent_after_ms;
continue;
}
return TransportOutcome::err(
TransportError::new(
TransportErrorCode::NetworkError,
format!("Asupersync HTTP client error: {error}"),
)
.into_snowflake_error(),
);
}
}
}
}
fn wire_request(
&self,
method: Method,
route: TransportRoute,
body: Vec<u8>,
auth: &AuthorizationDescriptor,
retry_resubmit: bool,
) -> Result<WireRequest, TransportError> {
let route_kind = route.kind();
self.config.limits.enforce(route_kind, body.len())?;
if matches!(route_kind, TransportRouteKind::Submit)
&& retry_resubmit
&& !route.has_retry_contract()
{
return Err(TransportError::new(
TransportErrorCode::NonIdempotentSubmitRetryRefused,
"submit retry requires requestId plus retry=true",
));
}
let mut headers = auth.wire_headers()?;
headers.push(Header::new(HEADER_ACCEPT, JSON_MEDIA_TYPE)?);
headers.push(Header::new(HEADER_USER_AGENT, USER_AGENT)?);
if matches!(route_kind, TransportRouteKind::Partition) {
headers.push(Header::new(
HEADER_ACCEPT_ENCODING,
PARTITION_ACCEPT_ENCODING,
)?);
}
if matches!(method, Method::Post) {
headers.push(Header::new(HEADER_CONTENT_TYPE, JSON_MEDIA_TYPE)?);
}
Ok(WireRequest {
method,
url: self.config.endpoint.route_url(&route),
route,
headers,
body,
})
}
}
pub const DEFAULT_ATTEMPT_TIMEOUT: Duration = Duration::from_secs(300);
pub const DEFAULT_CANCEL_ATTEMPT_TIMEOUT: Duration = Duration::from_secs(10);
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct TransportConfig {
pub endpoint: SnowflakeEndpoint,
pub tls_roots: TlsRootPolicy,
pub pool: PoolConfig,
pub limits: BodyLimits,
pub retry: RetryPolicy,
pub log: AttemptLogPolicy,
pub attempt_timeout: Option<Duration>,
pub cancel_attempt_timeout: Option<Duration>,
}
impl TransportConfig {
#[must_use]
pub fn new(endpoint: SnowflakeEndpoint) -> Self {
Self {
endpoint,
tls_roots: TlsRootPolicy::NativeRoots,
pool: PoolConfig::default(),
limits: BodyLimits::default(),
retry: RetryPolicy::default(),
log: AttemptLogPolicy::default(),
attempt_timeout: Some(DEFAULT_ATTEMPT_TIMEOUT),
cancel_attempt_timeout: Some(DEFAULT_CANCEL_ATTEMPT_TIMEOUT),
}
}
#[must_use]
pub fn attempt_timeout_for(&self, route_kind: TransportRouteKind) -> Option<Duration> {
match route_kind {
TransportRouteKind::Cancel => sooner(self.attempt_timeout, self.cancel_attempt_timeout),
TransportRouteKind::Submit
| TransportRouteKind::Poll
| TransportRouteKind::Partition => self.attempt_timeout,
}
}
}
fn sooner(left: Option<Duration>, right: Option<Duration>) -> Option<Duration> {
match (left, right) {
(Some(left), Some(right)) => Some(left.min(right)),
(left, None) => left,
(None, right) => right,
}
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub struct SnowflakeEndpoint {
base_url: String,
host: String,
}
impl SnowflakeEndpoint {
pub fn parse(raw: impl AsRef<str>) -> Result<Self, TransportError> {
let (base_url, host) = franken_snowflake_core::endpoint::validate_endpoint(raw.as_ref())
.map_err(|reason| {
TransportError::new(TransportErrorCode::InvalidSnowflakeHost, reason)
})?;
Ok(Self { base_url, host })
}
#[must_use]
pub fn base_url(&self) -> &str {
&self.base_url
}
#[must_use]
pub fn host(&self) -> &str {
&self.host
}
#[cfg(feature = "testkit-endpoint")]
pub fn parse_testkit_loopback(raw: &str) -> Result<Self, TransportError> {
let refused = || {
TransportError::new(
TransportErrorCode::InvalidSnowflakeHost,
"a testkit endpoint must be https://127.0.0.1:<port> or https://localhost:<port>",
)
};
let authority = raw
.strip_prefix("https://")
.map(|rest| rest.strip_suffix('/').unwrap_or(rest))
.ok_or_else(refused)?;
let (host, port) = authority.rsplit_once(':').ok_or_else(refused)?;
let port_ok = port.parse::<u16>().is_ok_and(|port| port != 0);
if !matches!(host, "127.0.0.1" | "localhost") || !port_ok {
return Err(refused());
}
Ok(Self {
base_url: format!("https://{host}:{port}"),
host: host.to_owned(),
})
}
#[must_use]
pub fn route_url(&self, route: &TransportRoute) -> String {
format!("{}{}", self.base_url, route.path_and_query())
}
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub enum TlsRootPolicy {
NativeRoots,
ExplicitPemBundle(PathBuf),
TestOnlyInsecureDisabledByDefault,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct PoolConfig {
pub max_idle_per_endpoint: usize,
pub max_in_flight_per_endpoint: usize,
}
impl Default for PoolConfig {
fn default() -> Self {
Self {
max_idle_per_endpoint: 8,
max_in_flight_per_endpoint: 16,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct BodyLimits {
pub max_submit_response_bytes: u64,
pub max_poll_response_bytes: u64,
pub max_partition_compressed_bytes: u64,
pub max_partition_uncompressed_bytes: u64,
pub max_submit_request_bytes: u64,
}
impl BodyLimits {
pub fn enforce(self, route: TransportRouteKind, body_len: usize) -> Result<(), TransportError> {
let max = match route {
TransportRouteKind::Submit => self.max_submit_request_bytes,
TransportRouteKind::Poll
| TransportRouteKind::Partition
| TransportRouteKind::Cancel => u64::MAX,
};
if body_len as u64 > max {
Err(TransportError::new(
TransportErrorCode::BodyLimitExceeded,
format!("request body has {body_len} bytes, limit is {max}"),
))
} else {
Ok(())
}
}
pub fn enforce_partition_sizes(
self,
compressed: u64,
uncompressed: u64,
) -> Result<(), TransportError> {
if compressed > self.max_partition_compressed_bytes {
return Err(TransportError::new(
TransportErrorCode::BodyLimitExceeded,
format!(
"compressed partition has {compressed} bytes, limit is {}",
self.max_partition_compressed_bytes
),
));
}
if uncompressed > self.max_partition_uncompressed_bytes {
return Err(TransportError::new(
TransportErrorCode::BodyLimitExceeded,
format!(
"uncompressed partition has {uncompressed} bytes, limit is {}",
self.max_partition_uncompressed_bytes
),
));
}
Ok(())
}
pub fn enforce_response_size(
self,
route: TransportRouteKind,
body_len: usize,
) -> Result<(), TransportError> {
let max = match route {
TransportRouteKind::Submit => self.max_submit_response_bytes,
TransportRouteKind::Poll => self.max_poll_response_bytes,
TransportRouteKind::Cancel => self.max_poll_response_bytes,
TransportRouteKind::Partition => self.max_partition_compressed_bytes,
};
if body_len as u64 > max {
Err(TransportError::new(
TransportErrorCode::BodyLimitExceeded,
format!("response body has {body_len} bytes, limit is {max}"),
))
} else {
Ok(())
}
}
}
impl Default for BodyLimits {
fn default() -> Self {
Self {
max_submit_response_bytes: 8 * 1024 * 1024,
max_poll_response_bytes: 8 * 1024 * 1024,
max_partition_compressed_bytes: 64 * 1024 * 1024,
max_partition_uncompressed_bytes: 512 * 1024 * 1024,
max_submit_request_bytes: 2 * 1024 * 1024,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct RetryPolicy {
pub max_attempts: u32,
pub base_delay_ms: u64,
pub max_delay_ms: u64,
pub total_budget_ms: u64,
pub respect_retry_after: bool,
pub deterministic_jitter: bool,
}
impl RetryPolicy {
#[must_use]
pub fn delay_for(
self,
request_id: Option<&RequestId>,
route: TransportRouteKind,
attempt: u32,
retry_after_ms: Option<u64>,
spent_ms: u64,
) -> Option<Duration> {
if attempt == 0 || attempt >= self.max_attempts || spent_ms >= self.total_budget_ms {
return None;
}
let from_header = retry_after_ms.filter(|_| self.respect_retry_after);
let exponential = self
.base_delay_ms
.saturating_mul(2_u64.saturating_pow(attempt.saturating_sub(1)))
.min(self.max_delay_ms);
let mut delay = from_header.unwrap_or(exponential);
if from_header.is_none() && self.deterministic_jitter {
delay = delay.saturating_add(deterministic_jitter_ms(request_id, route, attempt));
}
let remaining = self.total_budget_ms.saturating_sub(spent_ms);
if delay > remaining {
None
} else {
Some(Duration::from_millis(delay))
}
}
fn next_retry(
self,
request_id: Option<&RequestId>,
route: TransportRouteKind,
attempt: u32,
retry_after_ms: Option<u64>,
spent_ms: u64,
) -> Option<RetryDecision> {
let delay = self.delay_for(request_id, route, attempt, retry_after_ms, spent_ms)?;
Some(RetryDecision {
delay,
spent_after_ms: spent_ms.saturating_add(duration_millis_saturating(delay)),
})
}
#[must_use]
pub fn attempt_budget(self, ambient: Budget, route_budget: Budget, attempt: u32) -> Budget {
let remaining_attempts = self
.max_attempts
.saturating_sub(attempt.saturating_sub(1))
.max(1);
ambient.meet(route_budget).meet(
Budget::new()
.with_poll_quota(remaining_attempts)
.with_priority(0),
)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct RetryDecision {
delay: Duration,
spent_after_ms: u64,
}
fn duration_millis_saturating(duration: Duration) -> u64 {
u64::try_from(duration.as_millis()).unwrap_or(u64::MAX)
}
fn budget_timeout_at(budget: Budget, now: Time) -> Option<Duration> {
budget
.deadline
.map(|deadline| Duration::from_nanos(deadline.duration_since(now)))
}
fn budget_exhaustion_reason_at(budget: Budget, now: Time) -> Option<CancelReason> {
if budget_timeout_at(budget, now).is_some_and(|timeout| timeout.is_zero()) {
return Some(CancelReason::deadline());
}
if budget.poll_quota == 0 {
return Some(CancelReason::poll_quota());
}
if matches!(budget.cost_quota, Some(0)) {
return Some(CancelReason::cost_budget());
}
None
}
impl Default for RetryPolicy {
fn default() -> Self {
Self {
max_attempts: 4,
base_delay_ms: 100,
max_delay_ms: 2_000,
total_budget_ms: 10_000,
respect_retry_after: true,
deterministic_jitter: true,
}
}
}
fn deterministic_jitter_ms(
request_id: Option<&RequestId>,
route: TransportRouteKind,
attempt: u32,
) -> u64 {
let mut hash = 0xcbf2_9ce4_8422_2325_u64;
for byte in request_id
.map_or("no-request-id", RequestId::as_str)
.bytes()
.chain(route.as_str().bytes())
.chain(attempt.to_le_bytes())
{
hash ^= u64::from(byte);
hash = hash.wrapping_mul(0x0000_0100_0000_01b3);
}
hash % 31
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct AttemptLogPolicy {
pub enabled: bool,
pub hash_statement_handles: bool,
pub redact_account_identifiers: bool,
}
impl Default for AttemptLogPolicy {
fn default() -> Self {
Self {
enabled: true,
hash_statement_handles: true,
redact_account_identifiers: true,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
pub enum SnowflakeAuthTokenType {
ProgrammaticAccessToken,
KeypairJwt,
OAuth,
}
impl SnowflakeAuthTokenType {
#[must_use]
pub const fn as_header_value(self) -> &'static str {
match self {
Self::ProgrammaticAccessToken => "PROGRAMMATIC_ACCESS_TOKEN",
Self::KeypairJwt => "KEYPAIR_JWT",
Self::OAuth => "OAUTH",
}
}
}
#[derive(Clone, PartialEq, Eq)]
pub struct AuthorizationDescriptor {
token_type: SnowflakeAuthTokenType,
bearer_token: String,
redacted_fingerprint: String,
}
impl AuthorizationDescriptor {
#[must_use]
pub fn bearer(
token_type: SnowflakeAuthTokenType,
bearer_token: impl Into<String>,
redacted_fingerprint: impl Into<String>,
) -> Self {
Self {
token_type,
bearer_token: bearer_token.into(),
redacted_fingerprint: redacted_fingerprint.into(),
}
}
#[must_use]
pub fn redacted_fingerprint(&self) -> &str {
&self.redacted_fingerprint
}
fn wire_headers(&self) -> Result<Vec<Header>, TransportError> {
Ok(vec![
Header::new(
HEADER_AUTHORIZATION,
format!("Bearer {}", self.bearer_token),
)?,
Header::new(HEADER_TOKEN_TYPE, self.token_type.as_header_value())?,
])
}
}
impl fmt::Debug for AuthorizationDescriptor {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("AuthorizationDescriptor")
.field("token_type", &self.token_type)
.field("redacted_fingerprint", &self.redacted_fingerprint)
.finish_non_exhaustive()
}
}
#[derive(Clone, PartialEq, Eq, Hash, Deserialize)]
pub struct Header {
pub name: String,
pub value: String,
}
impl fmt::Debug for Header {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Header")
.field("name", &self.name)
.field("value", &redacted_header_value(&self.name, &self.value))
.finish()
}
}
impl Serialize for Header {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
let mut state = serializer.serialize_struct("Header", 2)?;
state.serialize_field("name", &self.name)?;
state.serialize_field("value", &redacted_header_value(&self.name, &self.value))?;
state.end()
}
}
impl Header {
pub fn new(name: impl Into<String>, value: impl Into<String>) -> Result<Self, TransportError> {
let name = name.into();
let value = value.into();
if !is_header_name(&name) || !is_header_value(&value) {
return Err(TransportError::new(
TransportErrorCode::HeaderRejected,
"header contains invalid characters",
));
}
Ok(Self { name, value })
}
}
fn redacted_header_value(name: &str, value: &str) -> String {
let value = redact(value).into_owned();
if !name.eq_ignore_ascii_case(HEADER_AUTHORIZATION) {
return value;
}
if value.contains(REDACTION_PLACEHOLDER) {
return value;
}
let mut words = value.split_whitespace();
match (words.next(), words.next()) {
(Some(scheme), Some(_)) => format!("{scheme} {REDACTION_PLACEHOLDER}"),
_ => REDACTION_PLACEHOLDER.to_owned(),
}
}
fn is_header_name(name: &str) -> bool {
!name.is_empty()
&& name
.bytes()
.all(|b| b.is_ascii_alphanumeric() || matches!(b, b'-'))
}
fn is_header_value(value: &str) -> bool {
value.bytes().all(|b| matches!(b, b'\t' | b' '..=b'~'))
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub enum TransportRoute {
Submit,
SubmitWithQuery {
query: Vec<(&'static str, String)>,
},
SubmitRetry {
request_id: RequestId,
},
Poll {
handle: StatementHandle,
},
Partition {
handle: StatementHandle,
partition: u32,
},
Cancel {
handle: StatementHandle,
},
}
impl TransportRoute {
#[must_use]
pub const fn kind(&self) -> TransportRouteKind {
match self {
Self::Submit | Self::SubmitWithQuery { .. } | Self::SubmitRetry { .. } => {
TransportRouteKind::Submit
}
Self::Poll { .. } => TransportRouteKind::Poll,
Self::Partition { .. } => TransportRouteKind::Partition,
Self::Cancel { .. } => TransportRouteKind::Cancel,
}
}
#[must_use]
pub fn has_retry_contract(&self) -> bool {
match self {
Self::SubmitRetry { .. } => true,
Self::SubmitWithQuery { query } => submit_query_has_retry_contract(query),
Self::Submit | Self::Poll { .. } | Self::Partition { .. } | Self::Cancel { .. } => {
false
}
}
}
#[must_use]
pub fn path_and_query(&self) -> String {
match self {
Self::Submit => SQL_API_STATEMENTS_PATH.to_owned(),
Self::SubmitWithQuery { query } => render_submit_query(query),
Self::SubmitRetry { request_id } => {
format!(
"{SQL_API_STATEMENTS_PATH}?requestId={}&retry=true",
percent_encode_query_component(request_id.as_str())
)
}
Self::Poll { handle } => {
format!(
"{SQL_API_STATEMENTS_PATH}/{}",
percent_encode_path_segment(handle.as_str())
)
}
Self::Partition { handle, partition } => {
format!(
"{SQL_API_STATEMENTS_PATH}/{}?partition={partition}",
percent_encode_path_segment(handle.as_str())
)
}
Self::Cancel { handle } => {
format!(
"{SQL_API_STATEMENTS_PATH}/{}/cancel",
percent_encode_path_segment(handle.as_str())
)
}
}
}
}
fn render_submit_query(query: &[(&'static str, String)]) -> String {
if query.is_empty() {
return SQL_API_STATEMENTS_PATH.to_owned();
}
let mut rendered = String::from(SQL_API_STATEMENTS_PATH);
for (index, (key, value)) in query.iter().enumerate() {
rendered.push(if index == 0 { '?' } else { '&' });
rendered.push_str(&percent_encode_query_component(key));
rendered.push('=');
rendered.push_str(&percent_encode_query_component(value));
}
rendered
}
fn percent_encode_query_component(value: &str) -> String {
let mut encoded = String::new();
for byte in value.bytes() {
if is_query_unreserved(byte) {
encoded.push(char::from(byte));
} else {
encoded.push('%');
encoded.push(hex_digit(byte >> 4));
encoded.push(hex_digit(byte & 0x0f));
}
}
encoded
}
fn percent_encode_path_segment(value: &str) -> String {
percent_encode_query_component(value)
}
const fn is_query_unreserved(byte: u8) -> bool {
byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'.' | b'_' | b'~')
}
const fn hex_digit(nibble: u8) -> char {
match nibble {
0..=9 => (b'0' + nibble) as char,
_ => (b'A' + (nibble - 10)) as char,
}
}
fn submit_query_has_retry_contract(query: &[(&'static str, String)]) -> bool {
query
.iter()
.any(|(key, value)| *key == "requestId" && !value.is_empty())
&& query
.iter()
.any(|(key, value)| *key == "retry" && value == "true")
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum TransportRouteKind {
Submit,
Poll,
Partition,
Cancel,
}
impl TransportRouteKind {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::Submit => "submit",
Self::Poll => "poll",
Self::Partition => "partition",
Self::Cancel => "cancel",
}
}
#[must_use]
pub const fn method(self) -> Method {
match self {
Self::Submit | Self::Cancel => Method::Post,
Self::Poll | Self::Partition => Method::Get,
}
}
}
#[derive(Clone, PartialEq, Eq)]
pub struct WireRequest {
pub method: Method,
pub url: String,
pub route: TransportRoute,
pub headers: Vec<Header>,
pub body: Vec<u8>,
}
impl fmt::Debug for WireRequest {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WireRequest")
.field("method", &self.method)
.field("url", &self.url)
.field("route", &self.route)
.field("headers", &self.headers)
.field("body_len", &self.body.len())
.finish()
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum StatusClass {
Completed,
Running,
StatementTimeout,
QueryFailure,
RateLimited,
ServerErrorRetryable,
Unauthorized,
Unexpected,
}
#[must_use]
pub fn classify_status(status: StatusCode) -> StatusClass {
match status.as_u16() {
200 => StatusClass::Completed,
202 => StatusClass::Running,
401 => StatusClass::Unauthorized,
408 => StatusClass::StatementTimeout,
422 => StatusClass::QueryFailure,
429 => StatusClass::RateLimited,
500..=599 => StatusClass::ServerErrorRetryable,
_ => StatusClass::Unexpected,
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct AttemptLogEvent {
pub schema: &'static str,
pub trace_id: String,
pub route: TransportRouteKind,
pub method: String,
pub attempt: u32,
pub request_fingerprint: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub statement_handle_hash: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub partition: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub status: Option<u16>,
pub retryable: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub retry_after_ms: Option<u64>,
pub elapsed_ms: u64,
pub compressed_bytes: u64,
pub uncompressed_bytes: u64,
pub outcome: AttemptOutcome,
#[serde(skip_serializing_if = "Option::is_none")]
pub error_code: Option<TransportErrorCode>,
}
impl AttemptLogEvent {
pub const SCHEMA: &'static str = "franken_snowflake.transport_attempt.v1";
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum AttemptOutcome {
Ok,
Error,
Cancelled,
}
#[derive(Clone, PartialEq, Eq)]
pub struct DecodedPartition {
pub partition: u32,
pub body: Vec<u8>,
pub compression: CompressionEvidence,
}
impl fmt::Debug for DecodedPartition {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("DecodedPartition")
.field("partition", &self.partition)
.field("body_len", &self.body.len())
.field("compression", &self.compression)
.finish()
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct CompressionEvidence {
pub content_encoding: ContentEncoding,
pub compressed_bytes: u64,
pub uncompressed_bytes: u64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ContentEncoding {
Identity,
Gzip,
}
impl ContentEncoding {
pub fn parse(raw: Option<&str>) -> Result<Self, TransportError> {
match raw.map(str::trim).filter(|value| !value.is_empty()) {
None => Ok(Self::Identity),
Some(value) if value.eq_ignore_ascii_case("identity") => Ok(Self::Identity),
Some(value) if value.eq_ignore_ascii_case("gzip") => Ok(Self::Gzip),
Some(_) => Err(TransportError::new(
TransportErrorCode::UnsupportedContentEncoding,
"unsupported partition content-encoding",
)),
}
}
}
struct DecodedPartitionResponse {
status: StatusClass,
body: Vec<u8>,
compression: CompressionEvidence,
}
fn decode_partition_response(
response: Response,
limits: BodyLimits,
partition: u32,
) -> Result<DecodedPartitionResponse, TransportError> {
let encoding = ContentEncoding::parse(response_header(
response.headers.as_slice(),
HEADER_CONTENT_ENCODING,
))?;
let compressed_bytes = response.body.len() as u64;
limits.enforce_partition_sizes(compressed_bytes, 0)?;
let max_uncompressed =
usize::try_from(limits.max_partition_uncompressed_bytes).map_err(|_| {
TransportError::new(
TransportErrorCode::BodyLimitExceeded,
"partition uncompressed limit exceeds local addressable memory",
)
})?;
let mut decoded = Vec::new();
let decode_result = match encoding {
ContentEncoding::Identity => {
let mut decompressor = IdentityDecompressor::new(Some(max_uncompressed));
decompressor
.decompress(response.body.as_slice(), &mut decoded)
.and_then(|()| decompressor.finish(&mut decoded))
}
ContentEncoding::Gzip => {
let mut decompressor = GzipDecompressor::new(DecompressionLimit::new(max_uncompressed));
decompressor
.decompress(response.body.as_slice(), &mut decoded)
.and_then(|()| decompressor.finish(&mut decoded))
}
};
if let Err(error) = decode_result {
let code = match encoding {
ContentEncoding::Identity => TransportErrorCode::BodyLimitExceeded,
ContentEncoding::Gzip => TransportErrorCode::GzipDecodeFailed,
};
return Err(TransportError::new(
code,
format!("partition {partition} decode failed: {error}"),
));
}
let uncompressed_bytes = decoded.len() as u64;
limits.enforce_partition_sizes(compressed_bytes, uncompressed_bytes)?;
Ok(DecodedPartitionResponse {
status: classify_status(StatusCode(response.status)),
body: decoded,
compression: CompressionEvidence {
content_encoding: encoding,
compressed_bytes,
uncompressed_bytes,
},
})
}
fn response_header<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> {
headers
.iter()
.find(|(header_name, _)| header_name.eq_ignore_ascii_case(name))
.map(|(_, value)| value.as_str())
}
fn retry_after_ms(headers: &[(String, String)]) -> Option<u64> {
let raw = response_header(headers, "Retry-After")?.trim();
raw.parse::<u64>()
.ok()
.and_then(|seconds| seconds.checked_mul(1_000))
.or_else(|| retry_after_http_date_ms(raw))
}
fn retry_after_http_date_ms(raw: &str) -> Option<u64> {
let target_seconds = parse_http_date_unix_seconds(raw)?;
let now_unix_seconds = i64::try_from(
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.ok()?
.as_secs(),
)
.ok()?;
let delta_seconds = target_seconds.saturating_sub(now_unix_seconds);
if delta_seconds <= 0 {
return Some(0);
}
Some((delta_seconds as u64).saturating_mul(1_000))
}
fn parse_http_date_unix_seconds(raw: &str) -> Option<i64> {
parse_imf_fixdate(raw)
.or_else(|| parse_rfc850_date(raw))
.or_else(|| parse_asctime_date(raw))
}
fn parse_imf_fixdate(raw: &str) -> Option<i64> {
let (_, rest) = raw.split_once(", ")?;
let mut parts = rest.split_ascii_whitespace();
let day = parts.next()?.parse::<u32>().ok()?;
let month = parse_http_month(parts.next()?)?;
let year = parts.next()?.parse::<i32>().ok()?;
let (hour, minute, second) = parse_http_time(parts.next()?)?;
if parts.next()? != "GMT" || parts.next().is_some() {
return None;
}
unix_seconds_utc(year, month, day, hour, minute, second)
}
fn parse_rfc850_date(raw: &str) -> Option<i64> {
let (_, rest) = raw.split_once(", ")?;
let mut parts = rest.split_ascii_whitespace();
let date = parts.next()?;
let mut date_parts = date.split('-');
let day = date_parts.next()?.parse::<u32>().ok()?;
let month = parse_http_month(date_parts.next()?)?;
let year_two_digits = date_parts.next()?.parse::<u32>().ok()?;
if date_parts.next().is_some() {
return None;
}
let year = if year_two_digits >= 70 {
1900 + i32::try_from(year_two_digits).ok()?
} else {
2000 + i32::try_from(year_two_digits).ok()?
};
let (hour, minute, second) = parse_http_time(parts.next()?)?;
if parts.next()? != "GMT" || parts.next().is_some() {
return None;
}
unix_seconds_utc(year, month, day, hour, minute, second)
}
fn parse_asctime_date(raw: &str) -> Option<i64> {
let mut parts = raw.split_ascii_whitespace();
let _weekday = parts.next()?;
let month = parse_http_month(parts.next()?)?;
let day = parts.next()?.parse::<u32>().ok()?;
let (hour, minute, second) = parse_http_time(parts.next()?)?;
let year = parts.next()?.parse::<i32>().ok()?;
if parts.next().is_some() {
return None;
}
unix_seconds_utc(year, month, day, hour, minute, second)
}
fn parse_http_month(month: &str) -> Option<u32> {
match month {
"Jan" => Some(1),
"Feb" => Some(2),
"Mar" => Some(3),
"Apr" => Some(4),
"May" => Some(5),
"Jun" => Some(6),
"Jul" => Some(7),
"Aug" => Some(8),
"Sep" => Some(9),
"Oct" => Some(10),
"Nov" => Some(11),
"Dec" => Some(12),
_ => None,
}
}
fn parse_http_time(time: &str) -> Option<(u32, u32, u32)> {
let mut parts = time.split(':');
let hour = parts.next()?.parse::<u32>().ok()?;
let minute = parts.next()?.parse::<u32>().ok()?;
let second = parts.next()?.parse::<u32>().ok()?;
if parts.next().is_some() || hour > 23 || minute > 59 || second > 60 {
return None;
}
Some((hour, minute, second))
}
fn unix_seconds_utc(
year: i32,
month: u32,
day: u32,
hour: u32,
minute: u32,
second: u32,
) -> Option<i64> {
if !(1..=12).contains(&month)
|| day == 0
|| day > days_in_month(year, month)
|| hour > 23
|| minute > 59
|| second > 60
{
return None;
}
let second = second.min(59);
days_from_civil(year, month, day)
.checked_mul(86_400)?
.checked_add(i64::from(hour) * 3_600 + i64::from(minute) * 60 + i64::from(second))
}
fn days_in_month(year: i32, month: u32) -> u32 {
match month {
1 | 3 | 5 | 7 | 8 | 10 | 12 => 31,
4 | 6 | 9 | 11 => 30,
2 if is_leap_year(year) => 29,
2 => 28,
_ => 0,
}
}
fn is_leap_year(year: i32) -> bool {
(year % 4 == 0 && year % 100 != 0) || year % 400 == 0
}
fn days_from_civil(year: i32, month: u32, day: u32) -> i64 {
let mut year = i64::from(year);
let month = i64::from(month);
let day = i64::from(day);
year -= i64::from(month <= 2);
let era = if year >= 0 { year } else { year - 399 } / 400;
let year_of_era = year - era * 400;
let month_prime = month + if month > 2 { -3 } else { 9 };
let day_of_year = (153 * month_prime + 2) / 5 + day - 1;
let day_of_era = year_of_era * 365 + year_of_era / 4 - year_of_era / 100 + day_of_year;
era * 146_097 + day_of_era - 719_468
}
fn is_retryable_status(status: u16) -> bool {
matches!(
classify_status(StatusCode(status)),
StatusClass::RateLimited | StatusClass::ServerErrorRetryable
)
}
fn route_allows_automatic_retry(route: &TransportRoute) -> bool {
route.has_retry_contract()
|| matches!(
route,
TransportRoute::Poll { .. }
| TransportRoute::Partition { .. }
| TransportRoute::Cancel { .. }
)
}
fn cancel_reason_or(cx: &Cx, fallback: fn() -> CancelReason) -> CancelReason {
cx.cancel_reason().unwrap_or_else(fallback)
}
async fn wait_retry_delay(cx: &Cx, delay: Duration) -> Result<(), CancelReason> {
if cx.checkpoint().is_err() {
return Err(cancel_reason_or(cx, CancelReason::parent_cancelled));
}
if asupersync::time::budget_sleep(cx, delay, cx.now_for_observability())
.await
.is_err()
{
return Err(cancel_reason_or(cx, CancelReason::deadline));
}
if cx.checkpoint().is_err() {
return Err(cancel_reason_or(cx, CancelReason::parent_cancelled));
}
Ok(())
}
pub async fn run_with_cancellation_mask<T>(cx: &Cx, future: impl Future<Output = T>) -> T {
let mut future = Box::pin(future);
std::future::poll_fn(|task| cx.masked(|| future.as_mut().poll(task))).await
}
pub trait PartitionSink {
fn accept(
&mut self,
cx: &Cx,
partition: DecodedPartition,
) -> impl std::future::Future<Output = Result<(), TransportError>>;
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct PartitionStreamRequest {
pub auth: AuthorizationDescriptor,
pub statement_handle: StatementHandle,
pub first_partition: u32,
pub end_partition_exclusive: u32,
pub max_concurrent_fetches: usize,
pub child_budget: Budget,
pub remote_cancel_on_local_cancel: bool,
pub seed_partitions: Vec<DecodedPartition>,
}
impl PartitionStreamRequest {
pub fn plan(&self) -> Result<PartitionStreamSummary, TransportError> {
if self.end_partition_exclusive < self.first_partition {
return Err(TransportError::new(
TransportErrorCode::InvalidPartitionPlan,
"partition end must be greater than or equal to start",
));
}
if self.max_concurrent_fetches == 0 {
return Err(TransportError::new(
TransportErrorCode::InvalidPartitionPlan,
"partition fetch concurrency must be at least one",
));
}
let mut previous_seed = None;
for seed in &self.seed_partitions {
if seed.partition >= self.first_partition {
return Err(TransportError::new(
TransportErrorCode::InvalidPartitionPlan,
"seed partitions must precede the fetch range",
));
}
if previous_seed.is_some_and(|previous| seed.partition <= previous) {
return Err(TransportError::new(
TransportErrorCode::InvalidPartitionPlan,
"seed partitions must be strictly increasing",
));
}
previous_seed = Some(seed.partition);
}
let accepted_seed_partitions = u32::try_from(self.seed_partitions.len()).map_err(|_| {
TransportError::new(
TransportErrorCode::InvalidPartitionPlan,
"too many seed partitions",
)
})?;
Ok(PartitionStreamSummary {
statement_handle: self.statement_handle.clone(),
planned_partitions: self.end_partition_exclusive - self.first_partition,
accepted_seed_partitions,
max_concurrent_fetches: self.max_concurrent_fetches,
})
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct PartitionStreamSummary {
pub statement_handle: StatementHandle,
pub planned_partitions: u32,
pub accepted_seed_partitions: u32,
pub max_concurrent_fetches: usize,
}
#[derive(Clone, PartialEq, Eq)]
pub struct SubmitHttpRequest {
pub route: TransportRoute,
pub auth: AuthorizationDescriptor,
pub body: Vec<u8>,
pub retry_resubmit: bool,
}
impl fmt::Debug for SubmitHttpRequest {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SubmitHttpRequest")
.field("route", &self.route)
.field("auth", &self.auth)
.field("body_len", &self.body.len())
.field("retry_resubmit", &self.retry_resubmit)
.finish()
}
}
#[derive(Clone, PartialEq, Eq)]
pub struct SubmitHttpResponse {
pub status: StatusClass,
pub body: Vec<u8>,
}
impl fmt::Debug for SubmitHttpResponse {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SubmitHttpResponse")
.field("status", &self.status)
.field("body_len", &self.body.len())
.finish()
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct PollHttpRequest {
pub auth: AuthorizationDescriptor,
pub statement_handle: StatementHandle,
}
#[derive(Clone, PartialEq, Eq)]
pub struct PollHttpResponse {
pub status: StatusClass,
pub body: Vec<u8>,
}
impl fmt::Debug for PollHttpResponse {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("PollHttpResponse")
.field("status", &self.status)
.field("body_len", &self.body.len())
.finish()
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct PartitionHttpRequest {
pub auth: AuthorizationDescriptor,
pub statement_handle: StatementHandle,
pub partition: u32,
}
#[derive(Clone, PartialEq, Eq)]
pub struct PartitionBody {
pub status: StatusClass,
pub body: Vec<u8>,
pub compression: CompressionEvidence,
}
impl fmt::Debug for PartitionBody {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("PartitionBody")
.field("status", &self.status)
.field("body_len", &self.body.len())
.field("compression", &self.compression)
.finish()
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct CancelHttpRequest {
pub auth: AuthorizationDescriptor,
pub statement_handle: StatementHandle,
pub reason_kind: CancelKind,
}
#[derive(Clone, PartialEq, Eq)]
pub struct CancelHttpResponse {
pub status: StatusClass,
pub body: Vec<u8>,
}
impl fmt::Debug for CancelHttpResponse {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("CancelHttpResponse")
.field("status", &self.status)
.field("body_len", &self.body.len())
.finish()
}
}
struct PlannedTransportRequest {
route: TransportRoute,
auth: AuthorizationDescriptor,
body: Vec<u8>,
retry_resubmit: bool,
request_id: Option<RequestId>,
partition: Option<u32>,
budget: Budget,
}
impl PlannedTransportRequest {
fn partition(
auth: AuthorizationDescriptor,
statement_handle: StatementHandle,
partition: u32,
budget: Budget,
) -> Self {
Self {
route: TransportRoute::Partition {
handle: statement_handle,
partition,
},
auth,
body: Vec::new(),
retry_resubmit: false,
request_id: None,
partition: Some(partition),
budget,
}
}
}
impl From<SubmitHttpRequest> for PlannedTransportRequest {
fn from(value: SubmitHttpRequest) -> Self {
let request_id = match &value.route {
TransportRoute::SubmitRetry { request_id } => Some(request_id.clone()),
TransportRoute::SubmitWithQuery { query } => query
.iter()
.find_map(|(key, value)| (*key == "requestId").then(|| RequestId::new(value))),
TransportRoute::Submit
| TransportRoute::Poll { .. }
| TransportRoute::Partition { .. }
| TransportRoute::Cancel { .. } => None,
};
Self {
route: value.route,
auth: value.auth,
body: value.body,
retry_resubmit: value.retry_resubmit,
request_id,
partition: None,
budget: Budget::unlimited(),
}
}
}
impl From<PollHttpRequest> for PlannedTransportRequest {
fn from(value: PollHttpRequest) -> Self {
Self {
route: TransportRoute::Poll {
handle: value.statement_handle,
},
auth: value.auth,
body: Vec::new(),
retry_resubmit: false,
request_id: None,
partition: None,
budget: Budget::unlimited(),
}
}
}
impl From<PartitionHttpRequest> for PlannedTransportRequest {
fn from(value: PartitionHttpRequest) -> Self {
Self {
route: TransportRoute::Partition {
handle: value.statement_handle,
partition: value.partition,
},
auth: value.auth,
body: Vec::new(),
retry_resubmit: false,
request_id: None,
partition: Some(value.partition),
budget: Budget::unlimited(),
}
}
}
impl From<CancelHttpRequest> for PlannedTransportRequest {
fn from(value: CancelHttpRequest) -> Self {
Self {
route: TransportRoute::Cancel {
handle: value.statement_handle,
},
auth: value.auth,
body: Vec::new(),
retry_resubmit: false,
request_id: None,
partition: None,
budget: Budget::unlimited(),
}
}
}
trait FromResponseBody: Sized {
fn from_response(
response: Response,
limits: BodyLimits,
partition: Option<u32>,
) -> Result<Self, TransportError>;
}
impl FromResponseBody for SubmitHttpResponse {
fn from_response(
response: Response,
limits: BodyLimits,
_partition: Option<u32>,
) -> Result<Self, TransportError> {
limits.enforce_response_size(TransportRouteKind::Submit, response.body.len())?;
Ok(Self {
status: classify_status(StatusCode(response.status)),
body: response.body,
})
}
}
impl FromResponseBody for PollHttpResponse {
fn from_response(
response: Response,
limits: BodyLimits,
_partition: Option<u32>,
) -> Result<Self, TransportError> {
limits.enforce_response_size(TransportRouteKind::Poll, response.body.len())?;
Ok(Self {
status: classify_status(StatusCode(response.status)),
body: response.body,
})
}
}
impl FromResponseBody for PartitionBody {
fn from_response(
response: Response,
limits: BodyLimits,
partition: Option<u32>,
) -> Result<Self, TransportError> {
let decoded = decode_partition_response(response, limits, partition.unwrap_or_default())?;
Ok(Self {
status: decoded.status,
body: decoded.body,
compression: decoded.compression,
})
}
}
impl FromResponseBody for CancelHttpResponse {
fn from_response(
response: Response,
limits: BodyLimits,
_partition: Option<u32>,
) -> Result<Self, TransportError> {
limits.enforce_response_size(TransportRouteKind::Cancel, response.body.len())?;
Ok(Self {
status: classify_status(StatusCode(response.status)),
body: response.body,
})
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum TransportErrorCode {
TlsRootPolicyRefused,
InvalidSnowflakeHost,
BodyLimitExceeded,
HeaderRejected,
RetryBudgetExhausted,
CancelledDuringConnect,
CancelAfterSubmitFailed,
HttpStatusUnexpected,
ResponseDecodeFailed,
GzipDecodeFailed,
UnsupportedContentEncoding,
NonIdempotentSubmitRetryRefused,
InvalidPartitionPlan,
NetworkError,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct TransportError {
pub code: TransportErrorCode,
pub message: String,
}
impl TransportError {
#[must_use]
pub fn new(code: TransportErrorCode, message: impl Into<String>) -> Self {
let message = message.into();
Self {
code,
message: redact(&message).into_owned(),
}
}
#[must_use]
pub fn into_snowflake_error(self) -> franken_snowflake_core::error::SnowflakeError {
use franken_snowflake_core::error::{SnowflakeError, SnowflakeErrorCode};
let code = match self.code {
TransportErrorCode::RetryBudgetExhausted => SnowflakeErrorCode::RetryBudgetExhausted,
TransportErrorCode::HttpStatusUnexpected | TransportErrorCode::ResponseDecodeFailed => {
SnowflakeErrorCode::UpstreamError
}
TransportErrorCode::NetworkError
| TransportErrorCode::TlsRootPolicyRefused
| TransportErrorCode::InvalidSnowflakeHost
| TransportErrorCode::CancelledDuringConnect
| TransportErrorCode::CancelAfterSubmitFailed => SnowflakeErrorCode::NetworkError,
TransportErrorCode::BodyLimitExceeded
| TransportErrorCode::HeaderRejected
| TransportErrorCode::GzipDecodeFailed
| TransportErrorCode::UnsupportedContentEncoding
| TransportErrorCode::NonIdempotentSubmitRetryRefused
| TransportErrorCode::InvalidPartitionPlan => SnowflakeErrorCode::UsageError,
};
SnowflakeError::new(code, self.message)
}
}
impl fmt::Display for TransportError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{:?}: {}", self.code, self.message)
}
}
impl std::error::Error for TransportError {}
#[cfg(test)]
mod tests {
use super::*;
use std::cell::RefCell;
use std::collections::VecDeque;
fn endpoint() -> SnowflakeEndpoint {
SnowflakeEndpoint {
base_url: "https://xy12345.us-east-1.snowflakecomputing.com".to_string(),
host: "xy12345.us-east-1.snowflakecomputing.com".to_string(),
}
}
fn auth() -> AuthorizationDescriptor {
auth_with_token("secret-token")
}
fn auth_with_token(token: &str) -> AuthorizationDescriptor {
AuthorizationDescriptor::bearer(
SnowflakeAuthTokenType::ProgrammaticAccessToken,
token,
"sha256:abc123",
)
}
fn header(name: &str, value: &str) -> Vec<(String, String)> {
vec![(name.to_string(), value.to_string())]
}
#[test]
fn pem_bundle_roots_are_loaded_or_refused() {
let dir = std::env::temp_dir().join(format!("fsnow-pem-{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let limits = BodyLimits::default();
let missing = PemBundleHttp::from_pem_file(&dir.join("absent.pem"), &limits)
.err()
.unwrap();
assert_eq!(missing.code, TransportErrorCode::TlsRootPolicyRefused);
let empty = dir.join("empty.pem");
std::fs::write(&empty, "not a certificate\n").unwrap();
let refused = PemBundleHttp::from_pem_file(&empty, &limits).err().unwrap();
assert_eq!(refused.code, TransportErrorCode::TlsRootPolicyRefused);
assert!(refused.message.contains("empty.pem"), "{refused}");
let mut config = TransportConfig::new(endpoint());
config.tls_roots = TlsRootPolicy::TestOnlyInsecureDisabledByDefault;
let error = SnowflakeHttpClient::for_runtime(config.clone())
.err()
.unwrap();
assert_eq!(error.code, TransportErrorCode::TlsRootPolicyRefused);
config.tls_roots = TlsRootPolicy::ExplicitPemBundle(empty);
assert!(SnowflakeHttpClient::for_runtime(config.clone()).is_err());
config.tls_roots = TlsRootPolicy::NativeRoots;
assert!(SnowflakeHttpClient::for_runtime(config).is_ok());
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn live_transports_read_up_to_the_configured_body_limit() {
let limits = BodyLimits::default();
assert_eq!(max_response_body_bytes(&limits), 64 * 1024 * 1024);
assert!(max_response_body_bytes(&limits) > 16 * 1024 * 1024);
let tight = BodyLimits {
max_partition_compressed_bytes: 1024,
max_poll_response_bytes: 2048,
max_submit_response_bytes: 512,
..limits
};
assert_eq!(max_response_body_bytes(&tight), 2048);
}
#[cfg(feature = "testkit-endpoint")]
#[test]
fn testkit_loopback_accepts_only_loopback_https_with_a_port() {
let ok = SnowflakeEndpoint::parse_testkit_loopback("https://127.0.0.1:8443/").unwrap();
assert_eq!(ok.base_url(), "https://127.0.0.1:8443");
assert!(SnowflakeEndpoint::parse_testkit_loopback("https://localhost:9").is_ok());
for refused in [
"http://127.0.0.1:8443",
"https://127.0.0.1",
"https://127.0.0.1:0",
"https://10.0.0.1:8443",
"https://xy123.snowflakecomputing.com:443",
"https://127.0.0.1:8443/path",
"https://user@127.0.0.1:8443",
] {
assert!(
SnowflakeEndpoint::parse_testkit_loopback(refused).is_err(),
"{refused} must be refused"
);
}
assert!(SnowflakeEndpoint::parse("https://127.0.0.1:8443").is_err());
}
#[test]
fn endpoint_requires_https_and_rejects_credentials() -> Result<(), String> {
assert!(SnowflakeEndpoint::parse("http://xy123.snowflakecomputing.com").is_err());
assert!(SnowflakeEndpoint::parse("https://user@xy123.snowflakecomputing.com").is_err());
assert!(SnowflakeEndpoint::parse("https://xy123.snowflakecomputing.com?role=x").is_err());
assert!(SnowflakeEndpoint::parse("https://xy123.snowflakecomputing.com/proxy").is_err());
assert!(SnowflakeEndpoint::parse("https://attacker.example.com").is_err());
let parsed = SnowflakeEndpoint::parse("https://xy123.snowflakecomputing.com/")
.map_err(|e| format!("valid endpoint failed: {e}"))?;
assert_eq!(parsed.base_url(), "https://xy123.snowflakecomputing.com");
let parsed_upper = SnowflakeEndpoint::parse("HTTPS://XY123.SNOWFLAKECOMPUTING.COM/")
.map_err(|e| format!("valid uppercase endpoint failed: {e}"))?;
assert_eq!(
parsed_upper.base_url(),
"https://xy123.snowflakecomputing.com"
);
assert_eq!(parsed_upper.host(), "xy123.snowflakecomputing.com");
Ok(())
}
#[test]
fn route_urls_are_canonical() {
let request_id = RequestId::new("req-123");
let handle = StatementHandle::new("stmt-456");
assert_eq!(
endpoint().route_url(&TransportRoute::SubmitRetry { request_id }),
"https://xy12345.us-east-1.snowflakecomputing.com/api/v2/statements?requestId=req-123&retry=true"
);
assert_eq!(
endpoint().route_url(&TransportRoute::Partition {
handle,
partition: 7,
}),
"https://xy12345.us-east-1.snowflakecomputing.com/api/v2/statements/stmt-456?partition=7"
);
}
#[test]
fn route_components_are_percent_encoded() {
let request_id = RequestId::new("req-123&retry=false");
let handle = StatementHandle::new("stmt/456?partition=9#frag");
assert_eq!(
endpoint().route_url(&TransportRoute::SubmitRetry { request_id }),
"https://xy12345.us-east-1.snowflakecomputing.com/api/v2/statements?requestId=req-123%26retry%3Dfalse&retry=true"
);
assert_eq!(
endpoint().route_url(&TransportRoute::Poll {
handle: handle.clone(),
}),
"https://xy12345.us-east-1.snowflakecomputing.com/api/v2/statements/stmt%2F456%3Fpartition%3D9%23frag"
);
assert_eq!(
endpoint().route_url(&TransportRoute::Cancel { handle }),
"https://xy12345.us-east-1.snowflakecomputing.com/api/v2/statements/stmt%2F456%3Fpartition%3D9%23frag/cancel"
);
}
#[test]
fn submit_query_values_are_percent_encoded() {
let route = TransportRoute::SubmitWithQuery {
query: vec![
("requestId", "req-123&retry=false".to_owned()),
("retry", "true".to_owned()),
("nullable", "false".to_owned()),
],
};
assert_eq!(
endpoint().route_url(&route),
"https://xy12345.us-east-1.snowflakecomputing.com/api/v2/statements?requestId=req-123%26retry%3Dfalse&retry=true&nullable=false"
);
assert!(route.has_retry_contract());
}
#[test]
fn auth_debug_redacts_secret_bearer() {
let rendered = format!("{:?}", auth());
assert!(rendered.contains("sha256:abc123"));
assert!(!rendered.contains("secret-token"));
}
#[test]
fn auth_headers_wire_token_type_without_logging_secret() -> Result<(), String> {
let headers = auth()
.wire_headers()
.map_err(|e| format!("wire_headers failed: {e}"))?;
assert!(
headers
.iter()
.any(|h| { h.name == HEADER_AUTHORIZATION && h.value == "Bearer secret-token" })
);
assert!(
headers
.iter()
.any(|h| { h.name == HEADER_TOKEN_TYPE && h.value == "PROGRAMMATIC_ACCESS_TOKEN" })
);
Ok(())
}
#[test]
fn wire_request_and_header_debug_redact_authorization_bearer() -> Result<(), String> {
let token = "sfpat_http_debug_secret_123";
let client = SnowflakeHttpClient::new(
TransportConfig::new(endpoint()),
AsupersyncHttpClient::new(),
);
let request = SubmitHttpRequest {
route: TransportRoute::Submit,
auth: auth_with_token(token),
body: b"{}".to_vec(),
retry_resubmit: false,
};
let plan = client
.submit_plan(&request)
.map_err(|e| format!("submit plan failed: {e}"))?;
let authorization = plan
.headers
.iter()
.find(|header| header.name == HEADER_AUTHORIZATION)
.ok_or_else(|| "authorization header not found".to_string())?;
assert_eq!(authorization.value, format!("Bearer {token}"));
for rendered in [
format!("{authorization:?}"),
format!("{:?}", plan.headers),
format!("{plan:?}"),
serde_json::to_string(authorization)
.map_err(|e| format!("header json serialization failed: {e}"))?,
] {
assert!(
!rendered.contains(token),
"diagnostic surface leaked bearer token: {rendered}"
);
assert!(rendered.contains(REDACTION_PLACEHOLDER));
}
Ok(())
}
#[test]
fn body_bearing_debug_surfaces_only_report_lengths() {
let secret_body = b"sfpat_http_body_secret_123".to_vec();
let decimal_secret_prefix = "115, 102, 112, 97, 116";
let compression = CompressionEvidence {
content_encoding: ContentEncoding::Identity,
compressed_bytes: secret_body.len() as u64,
uncompressed_bytes: secret_body.len() as u64,
};
let decoded = DecodedPartition {
partition: 0,
body: secret_body.clone(),
compression,
};
let stream = PartitionStreamRequest {
auth: auth(),
statement_handle: StatementHandle::new("stmt-1"),
first_partition: 1,
end_partition_exclusive: 1,
max_concurrent_fetches: 1,
child_budget: Budget::unlimited(),
remote_cancel_on_local_cancel: false,
seed_partitions: vec![decoded.clone()],
};
for rendered in [
format!(
"{:?}",
SubmitHttpRequest {
route: TransportRoute::Submit,
auth: auth(),
body: secret_body.clone(),
retry_resubmit: false,
}
),
format!(
"{:?}",
SubmitHttpResponse {
status: StatusClass::Completed,
body: secret_body.clone(),
}
),
format!(
"{:?}",
PollHttpResponse {
status: StatusClass::Running,
body: secret_body.clone(),
}
),
format!(
"{:?}",
PartitionBody {
status: StatusClass::Completed,
body: secret_body.clone(),
compression,
}
),
format!(
"{:?}",
CancelHttpResponse {
status: StatusClass::Completed,
body: secret_body.clone(),
}
),
format!("{decoded:?}"),
format!("{stream:?}"),
] {
assert!(
!rendered.contains("sfpat_http_body_secret_123"),
"debug surface leaked body as text: {rendered}"
);
assert!(
!rendered.contains(decimal_secret_prefix),
"debug surface leaked body as reconstructable bytes: {rendered}"
);
assert!(rendered.contains("body_len"));
}
}
#[test]
fn transport_error_constructor_redacts_secret_shaped_messages() -> Result<(), String> {
let token = "ghp_httpTransportSecret0123";
let error = TransportError::new(
TransportErrorCode::NetworkError,
format!("connect failed with token={token}"),
);
for rendered in [
error.message.clone(),
error.to_string(),
serde_json::to_string(&error).map_err(|e| format!("error json failed: {e}"))?,
format!("{error:?}"),
] {
assert!(
!rendered.contains(token),
"transport error leaked secret-shaped token: {rendered}"
);
assert!(rendered.contains(REDACTION_PLACEHOLDER));
}
Ok(())
}
#[test]
fn body_limits_refuse_oversize_submit_request() {
let limits = BodyLimits {
max_submit_request_bytes: 3,
..BodyLimits::default()
};
assert!(limits.enforce(TransportRouteKind::Submit, 4).is_err());
assert!(limits.enforce(TransportRouteKind::Poll, 4).is_ok());
}
#[test]
fn partition_limits_count_both_compressed_and_uncompressed() {
let limits = BodyLimits {
max_partition_compressed_bytes: 10,
max_partition_uncompressed_bytes: 20,
..BodyLimits::default()
};
assert!(limits.enforce_partition_sizes(10, 20).is_ok());
assert!(limits.enforce_partition_sizes(11, 20).is_err());
assert!(limits.enforce_partition_sizes(10, 21).is_err());
}
#[test]
fn retry_after_wins_before_exponential_jitter() -> Result<(), String> {
let policy = RetryPolicy::default();
let request_id = RequestId::new("req-123");
assert_eq!(
policy.delay_for(Some(&request_id), TransportRouteKind::Poll, 1, Some(500), 0),
Some(Duration::from_millis(500))
);
let exponential = policy
.delay_for(Some(&request_id), TransportRouteKind::Poll, 2, None, 0)
.ok_or_else(|| "expected exponential delay".to_string())?;
assert!(exponential >= Duration::from_millis(200));
assert!(exponential <= Duration::from_millis(230));
Ok(())
}
#[test]
fn retry_after_delta_seconds_form_is_clock_independent() {
assert_eq!(retry_after_ms(&header("Retry-After", "5")), Some(5_000));
assert_eq!(retry_after_ms(&header("Retry-After", "0")), Some(0));
assert_eq!(retry_after_ms(&header("Retry-After", "not-a-value")), None);
}
#[test]
fn retry_after_http_date_formats_parse_to_unix_seconds() {
for raw in [
"Wed, 21 Oct 2015 07:28:00 GMT",
"Wednesday, 21-Oct-15 07:28:00 GMT",
"Wed Oct 21 07:28:00 2015",
] {
assert_eq!(parse_http_date_unix_seconds(raw), Some(1_445_412_480));
}
}
#[test]
fn retry_after_http_date_is_measured_against_absolute_wall_clock() -> Result<(), String> {
for raw in [
"Wed, 21 Oct 2015 07:28:00 GMT",
"Wednesday, 21-Oct-15 07:28:00 GMT",
"Wed Oct 21 07:28:00 2015",
] {
assert_eq!(retry_after_ms(&header("Retry-After", raw)), Some(0));
}
let delay = retry_after_ms(&header("Retry-After", "Fri, 31 Dec 2100 23:59:59 GMT"))
.ok_or_else(|| "expected future date delay".to_string())?;
assert!(
delay > 0,
"future date must yield a positive delay, got {delay}"
);
Ok(())
}
#[test]
fn retry_budget_stops_before_sleeping_past_total_budget() {
let policy = RetryPolicy {
total_budget_ms: 50,
..RetryPolicy::default()
};
assert_eq!(
policy.delay_for(None, TransportRouteKind::Poll, 1, Some(100), 0),
None
);
}
#[test]
fn retry_decision_charges_each_delay_once() -> Result<(), String> {
let policy = RetryPolicy {
max_attempts: 4,
base_delay_ms: 25,
max_delay_ms: 100,
total_budget_ms: 75,
respect_retry_after: false,
deterministic_jitter: false,
};
let first = policy
.next_retry(None, TransportRouteKind::Poll, 1, None, 0)
.ok_or_else(|| "expected first retry".to_string())?;
assert_eq!(first.delay, Duration::from_millis(25));
assert_eq!(first.spent_after_ms, 25);
let second = policy
.next_retry(
None,
TransportRouteKind::Poll,
2,
None,
first.spent_after_ms,
)
.ok_or_else(|| "expected second retry".to_string())?;
assert_eq!(second.delay, Duration::from_millis(50));
assert_eq!(second.spent_after_ms, 75);
assert_eq!(
policy.next_retry(
None,
TransportRouteKind::Poll,
3,
None,
second.spent_after_ms
),
None
);
Ok(())
}
#[test]
fn attempt_budget_meets_parent_child_and_attempt_quota_once() {
let parent = Budget::new()
.with_deadline(Time::from_secs(30))
.with_poll_quota(10)
.with_cost_quota(100)
.with_priority(1);
let child = Budget::new()
.with_deadline(Time::from_secs(20))
.with_poll_quota(8)
.with_cost_quota(50)
.with_priority(9);
let policy = RetryPolicy {
max_attempts: 4,
..RetryPolicy::default()
};
let effective = policy.attempt_budget(parent, child, 2);
let expected = parent
.meet(child)
.meet(Budget::new().with_poll_quota(3).with_priority(0));
assert_eq!(effective, expected);
assert_eq!(effective.deadline, Some(Time::from_secs(20)));
assert_eq!(effective.poll_quota, 3);
assert_eq!(effective.cost_quota, Some(50));
assert_eq!(effective.priority, 9);
}
#[test]
fn partition_stream_child_budget_flows_to_internal_partition_plan() {
let parent = Budget::new().with_poll_quota(10).with_cost_quota(100);
let child = Budget::new().with_poll_quota(2).with_cost_quota(40);
let planned =
PlannedTransportRequest::partition(auth(), StatementHandle::new("stmt-1"), 7, child);
assert_eq!(planned.partition, Some(7));
assert_eq!(planned.budget, child);
assert_eq!(
RetryPolicy::default().attempt_budget(parent, planned.budget, 1),
parent
.meet(child)
.meet(Budget::new().with_poll_quota(RetryPolicy::default().max_attempts))
);
}
#[test]
fn effective_budget_drives_timeout_and_exhaustion_reason() -> Result<(), String> {
let now = Time::from_secs(10);
let bounded = Budget::new().with_deadline(Time::from_secs(12));
assert_eq!(
budget_timeout_at(bounded, now),
Some(Duration::from_secs(2))
);
assert!(
budget_exhaustion_reason_at(Budget::new().with_deadline(Time::from_secs(10)), now)
.ok_or_else(|| "expected deadline exhausted reason".to_string())?
.is_kind(CancelKind::Deadline)
);
assert!(
budget_exhaustion_reason_at(Budget::new().with_poll_quota(0), now)
.ok_or_else(|| "expected poll quota exhausted reason".to_string())?
.is_kind(CancelKind::PollQuota)
);
assert!(
budget_exhaustion_reason_at(Budget::new().with_cost_quota(0), now)
.ok_or_else(|| "expected cost budget exhausted reason".to_string())?
.is_kind(CancelKind::CostBudget)
);
Ok(())
}
#[test]
fn status_classification_keeps_202_and_429_distinct() {
assert_eq!(classify_status(StatusCode(200)), StatusClass::Completed);
assert_eq!(classify_status(StatusCode(202)), StatusClass::Running);
assert_eq!(classify_status(StatusCode(401)), StatusClass::Unauthorized);
assert_eq!(classify_status(StatusCode(403)), StatusClass::Unexpected);
assert_eq!(
classify_status(StatusCode(408)),
StatusClass::StatementTimeout
);
assert_eq!(classify_status(StatusCode(422)), StatusClass::QueryFailure);
assert_eq!(classify_status(StatusCode(429)), StatusClass::RateLimited);
assert_eq!(
classify_status(StatusCode(503)),
StatusClass::ServerErrorRetryable
);
}
#[test]
fn statement_timeout_408_is_not_retried() {
assert!(!is_retryable_status(408));
assert!(is_retryable_status(429));
assert!(is_retryable_status(500));
assert!(is_retryable_status(503));
assert!(!is_retryable_status(200));
assert!(!is_retryable_status(422));
assert!(!is_retryable_status(404));
}
#[test]
fn content_encoding_fails_closed() -> Result<(), String> {
assert_eq!(
ContentEncoding::parse(None).map_err(|e| format!("identity parse failed: {e}"))?,
ContentEncoding::Identity
);
assert_eq!(
ContentEncoding::parse(Some("gzip")).map_err(|e| format!("gzip parse failed: {e}"))?,
ContentEncoding::Gzip
);
assert!(ContentEncoding::parse(Some("br")).is_err());
Ok(())
}
#[test]
fn submit_retry_requires_idempotency_contract() {
let client = SnowflakeHttpClient::new(
TransportConfig::new(endpoint()),
AsupersyncHttpClient::new(),
);
let request = SubmitHttpRequest {
route: TransportRoute::Submit,
auth: auth(),
body: b"{}".to_vec(),
retry_resubmit: true,
};
assert!(client.submit_plan(&request).is_err());
let request = SubmitHttpRequest {
route: TransportRoute::SubmitRetry {
request_id: RequestId::new("req-123"),
},
auth: auth(),
body: b"{}".to_vec(),
retry_resubmit: true,
};
assert!(client.submit_plan(&request).is_ok());
}
#[test]
fn plain_submit_post_is_not_blindly_retried_but_cancel_is_retry_safe() {
assert_eq!(SNOWFLAKE_SQL_API_RESUBMIT_DOC_CONSULTED, "2026-06-25");
assert!(
SNOWFLAKE_SQL_API_RESUBMIT_DOC_URL
.contains("/developer-guide/sql-api/submitting-requests")
);
assert!(is_retryable_status(503));
assert!(is_retryable_status(429));
assert!(!route_allows_automatic_retry(&TransportRoute::Submit));
assert!(route_allows_automatic_retry(&TransportRoute::SubmitRetry {
request_id: RequestId::new("req-123")
}));
assert!(route_allows_automatic_retry(&TransportRoute::Poll {
handle: StatementHandle::new("stmt-1")
}));
assert!(route_allows_automatic_retry(&TransportRoute::Partition {
handle: StatementHandle::new("stmt-1"),
partition: 1,
}));
assert!(route_allows_automatic_retry(&TransportRoute::Cancel {
handle: StatementHandle::new("stmt-1")
}));
}
#[test]
fn wire_plan_adds_json_headers() -> Result<(), String> {
let client = SnowflakeHttpClient::new(
TransportConfig::new(endpoint()),
AsupersyncHttpClient::new(),
);
let request = SubmitHttpRequest {
route: TransportRoute::Submit,
auth: auth(),
body: b"{}".to_vec(),
retry_resubmit: false,
};
let plan = client
.submit_plan(&request)
.map_err(|e| format!("submit plan failed: {e}"))?;
assert_eq!(plan.method, Method::Post);
assert!(plan.headers.iter().any(|h| h.name == HEADER_ACCEPT));
assert!(plan.headers.iter().any(|h| h.name == HEADER_CONTENT_TYPE));
Ok(())
}
#[test]
fn every_route_names_the_application_in_user_agent() -> Result<(), String> {
let client = SnowflakeHttpClient::new(
TransportConfig::new(endpoint()),
AsupersyncHttpClient::new(),
);
let handle = StatementHandle::new("stmt-1");
let routes = [
TransportRoute::Submit,
TransportRoute::Poll {
handle: handle.clone(),
},
TransportRoute::Partition {
handle: handle.clone(),
partition: 1,
},
TransportRoute::Cancel { handle },
];
for route in routes {
let kind = route.kind();
let body = if kind.method() == Method::Post {
b"{}".to_vec()
} else {
Vec::new()
};
let wire = client
.wire_request(kind.method(), route, body, &auth(), false)
.map_err(|error| format!("{kind:?}: {error}"))?;
let agents: Vec<&str> = wire
.headers
.iter()
.filter(|header| header.name.eq_ignore_ascii_case("user-agent"))
.map(|header| header.value.as_str())
.collect();
assert_eq!(agents, [USER_AGENT], "{kind:?}");
}
assert_eq!(USER_AGENT, format!("franken-snowflake/{VERSION}"));
Ok(())
}
#[test]
fn partition_wire_plan_advertises_gzip() -> Result<(), String> {
let client = SnowflakeHttpClient::new(
TransportConfig::new(endpoint()),
AsupersyncHttpClient::new(),
);
let request = PartitionHttpRequest {
auth: auth(),
statement_handle: StatementHandle::new("stmt-1"),
partition: 2,
};
let plan = client
.partition_plan(&request)
.map_err(|e| format!("partition plan failed: {e}"))?;
assert!(
plan.headers
.iter()
.any(|h| h.name == HEADER_ACCEPT_ENCODING && h.value == PARTITION_ACCEPT_ENCODING)
);
Ok(())
}
#[test]
fn gzip_partition_response_is_decoded_with_evidence() -> Result<(), String> {
use asupersync::http::compress::{Compressor, GzipCompressor};
let mut compressed = Vec::new();
let mut compressor = GzipCompressor::new();
compressor
.compress(br#"{"data":[["one"]]}"#, &mut compressed)
.map_err(|e| format!("compress failed: {e}"))?;
compressor
.finish(&mut compressed)
.map_err(|e| format!("finish failed: {e}"))?;
let response =
Response::new(200, "OK", compressed.clone()).with_header("content-encoding", "gzip");
let body = PartitionBody::from_response(response, BodyLimits::default(), Some(1))
.map_err(|e| format!("decode failed: {e}"))?;
assert_eq!(body.body, br#"{"data":[["one"]]}"#);
assert_eq!(body.compression.content_encoding, ContentEncoding::Gzip);
assert_eq!(body.compression.compressed_bytes, compressed.len() as u64);
assert_eq!(
body.compression.uncompressed_bytes,
br#"{"data":[["one"]]}"#.len() as u64
);
Ok(())
}
#[test]
fn compressed_partition_limit_is_checked_before_gzip_decode() {
let response = Response::new(200, "OK", b"not a valid gzip body".to_vec())
.with_header("content-encoding", "gzip");
let limits = BodyLimits {
max_partition_compressed_bytes: 3,
..BodyLimits::default()
};
let error = PartitionBody::from_response(response, limits, Some(1))
.expect_err("compressed limit should fail before gzip decode");
assert_eq!(error.code, TransportErrorCode::BodyLimitExceeded);
assert!(error.message.contains("compressed partition"));
}
struct ScriptedRaw {
responses: RefCell<VecDeque<Result<Response, AsupersyncClientError>>>,
requests: RefCell<Vec<RecordedRequest>>,
cancel_in_flight: Option<Cx>,
}
#[derive(Clone, Debug)]
struct RecordedRequest {
method: Method,
url: String,
headers: Vec<(String, String)>,
body: Vec<u8>,
timeout: Option<Duration>,
}
impl ScriptedRaw {
fn new(responses: Vec<Result<Response, AsupersyncClientError>>) -> Self {
Self {
responses: RefCell::new(responses.into()),
requests: RefCell::new(Vec::new()),
cancel_in_flight: None,
}
}
fn requests(&self) -> Vec<RecordedRequest> {
self.requests.borrow().clone()
}
}
impl RawHttp for ScriptedRaw {
async fn send(
&self,
_cx: &Cx,
method: Method,
url: String,
headers: Vec<(String, String)>,
body: Vec<u8>,
timeout: Option<Duration>,
) -> Result<Response, AsupersyncClientError> {
self.requests.borrow_mut().push(RecordedRequest {
method,
url,
headers,
body,
timeout,
});
if let Some(cx) = &self.cancel_in_flight {
cx.cancel_with(
CancelKind::User,
Some("cancelled while the request was in flight"),
);
}
self.responses.borrow_mut().pop_front().unwrap_or_else(|| {
Err(AsupersyncClientError::Io(std::io::Error::other(
"scripted transport exhausted",
)))
})
}
}
fn fast_retry_config(max_attempts: u32) -> TransportConfig {
TransportConfig {
retry: RetryPolicy {
max_attempts,
base_delay_ms: 1,
max_delay_ms: 1,
total_budget_ms: 1_000,
respect_retry_after: true,
deterministic_jitter: true,
},
..TransportConfig::new(endpoint())
}
}
fn scripted_client(
max_attempts: u32,
responses: Vec<Result<Response, AsupersyncClientError>>,
) -> SnowflakeHttpClient<ScriptedRaw> {
SnowflakeHttpClient::new(fast_retry_config(max_attempts), ScriptedRaw::new(responses))
}
fn ok_json(status: u16, body: &str) -> Result<Response, AsupersyncClientError> {
Ok(Response::new(status, "scripted", body.as_bytes().to_vec())
.with_header("content-type", "application/json"))
}
fn network_error() -> Result<Response, AsupersyncClientError> {
Err(AsupersyncClientError::Io(std::io::Error::other(
"connection reset by peer",
)))
}
fn retry_submit(request_id: &str) -> SubmitHttpRequest {
SubmitHttpRequest {
route: TransportRoute::SubmitRetry {
request_id: RequestId::new(request_id),
},
auth: auth(),
body: b"{\"statement\":\"select 1\"}".to_vec(),
retry_resubmit: true,
}
}
fn poll_request() -> PollHttpRequest {
PollHttpRequest {
auth: auth(),
statement_handle: StatementHandle::new("stmt-poll-1"),
}
}
#[test]
fn execute_retries_a_429_by_resubmitting_the_same_request_id_with_retry_true() {
asupersync::test_utils::run_test(|| async {
let client = scripted_client(
4,
vec![
Ok(Response::new(429, "Too Many Requests", Vec::new())
.with_header("Retry-After", "0")),
ok_json(202, "{\"statementHandle\":\"h1\"}"),
],
);
let cx = Cx::for_testing();
let outcome = client
.submit_statement(&cx, retry_submit("req-idem-1"))
.await;
let response = match outcome {
TransportOutcome::Ok(response) => response,
other => {
assert!(
matches!(other, TransportOutcome::Ok(_)),
"expected a submit response, got {other:?}"
);
return;
}
};
assert_eq!(response.status, StatusClass::Running);
let requests = client.client.requests();
assert_eq!(requests.len(), 2, "one 429 then one resubmit");
for request in &requests {
assert_eq!(request.method, Method::Post);
assert!(
request.url.contains("requestId=req-idem-1"),
"{}",
request.url
);
assert!(request.url.contains("retry=true"), "{}", request.url);
assert!(
request.headers.iter().any(
|(name, value)| name == "Authorization" && value.starts_with("Bearer ")
),
"bearer header present on every attempt"
);
}
assert_eq!(requests[0].url, requests[1].url);
assert_eq!(requests[0].body, requests[1].body);
assert!(
requests
.iter()
.all(|request| request.timeout == Some(DEFAULT_ATTEMPT_TIMEOUT))
);
});
}
#[test]
fn execute_refuses_to_auto_retry_a_bare_submit_on_a_retryable_status() {
asupersync::test_utils::run_test(|| async {
let client = scripted_client(
4,
vec![Ok(Response::new(503, "Service Unavailable", Vec::new()))],
);
let cx = Cx::for_testing();
let outcome = client
.submit_statement(
&cx,
SubmitHttpRequest {
route: TransportRoute::Submit,
auth: auth(),
body: b"{}".to_vec(),
retry_resubmit: false,
},
)
.await;
let error = match outcome {
TransportOutcome::Err(error) => error,
other => {
assert!(
matches!(other, TransportOutcome::Err(_)),
"expected a typed refusal, got {other:?}"
);
return;
}
};
assert!(
error.message.contains("requestId plus retry=true"),
"{}",
error.message
);
assert_eq!(
client.client.requests().len(),
1,
"no second attempt was sent"
);
});
}
#[test]
fn execute_exhausts_the_retry_budget_after_max_attempts() {
asupersync::test_utils::run_test(|| async {
let client = scripted_client(
3,
vec![
Ok(Response::new(503, "Service Unavailable", Vec::new())),
Ok(Response::new(503, "Service Unavailable", Vec::new())),
Ok(Response::new(503, "Service Unavailable", Vec::new())),
ok_json(200, "{}"),
],
);
let cx = Cx::for_testing();
let outcome = client.poll_statement(&cx, poll_request()).await;
let error = match outcome {
TransportOutcome::Err(error) => error,
other => {
assert!(
matches!(other, TransportOutcome::Err(_)),
"expected retry exhaustion, got {other:?}"
);
return;
}
};
assert!(
error
.message
.contains("exhausted retry budget after 3 attempts"),
"{}",
error.message
);
assert_eq!(client.client.requests().len(), 3);
});
}
#[test]
fn execute_retries_a_network_error_on_a_poll_route_then_succeeds() {
asupersync::test_utils::run_test(|| async {
let client = scripted_client(
4,
vec![
network_error(),
ok_json(202, "{\"statementHandle\":\"h1\"}"),
],
);
let cx = Cx::for_testing();
let outcome = client.poll_statement(&cx, poll_request()).await;
let response = match outcome {
TransportOutcome::Ok(response) => response,
other => {
assert!(
matches!(other, TransportOutcome::Ok(_)),
"expected a poll response, got {other:?}"
);
return;
}
};
assert_eq!(response.status, StatusClass::Running);
let requests = client.client.requests();
assert_eq!(requests.len(), 2);
assert!(
requests[1].url.ends_with("/api/v2/statements/stmt-poll-1"),
"{}",
requests[1].url
);
assert_eq!(requests[1].method, Method::Get);
});
}
#[test]
fn every_exchange_carries_a_bound() {
assert_eq!(
sooner(Some(Duration::from_secs(9)), Some(Duration::from_secs(3))),
Some(Duration::from_secs(3))
);
assert_eq!(
sooner(None, Some(Duration::from_secs(3))),
Some(Duration::from_secs(3))
);
assert_eq!(
sooner(Some(Duration::from_secs(9)), None),
Some(Duration::from_secs(9))
);
assert_eq!(sooner(None, None), None);
asupersync::test_utils::run_test(|| async {
let cx = Cx::for_testing();
let bounded = scripted_client(1, vec![ok_json(200, "{}")]);
let _ = bounded.poll_statement(&cx, poll_request()).await;
assert_eq!(
bounded.client.requests()[0].timeout,
Some(DEFAULT_ATTEMPT_TIMEOUT)
);
let mut config = fast_retry_config(1);
config.attempt_timeout = None;
let unbounded =
SnowflakeHttpClient::new(config, ScriptedRaw::new(vec![ok_json(200, "{}")]));
let _ = unbounded.poll_statement(&cx, poll_request()).await;
assert_eq!(unbounded.client.requests()[0].timeout, None);
let cancel = CancelHttpRequest {
auth: auth(),
statement_handle: StatementHandle::new("stmt-cancel-1"),
reason_kind: CancelKind::User,
};
let canceller = scripted_client(1, vec![ok_json(200, "{}")]);
let _ = canceller.cancel_statement(&cx, cancel.clone()).await;
assert_eq!(
canceller.client.requests()[0].timeout,
Some(DEFAULT_CANCEL_ATTEMPT_TIMEOUT)
);
let mut config = fast_retry_config(1);
config.attempt_timeout = Some(Duration::from_secs(2));
let tight =
SnowflakeHttpClient::new(config, ScriptedRaw::new(vec![ok_json(200, "{}")]));
let _ = tight.cancel_statement(&cx, cancel).await;
assert_eq!(
tight.client.requests()[0].timeout,
Some(Duration::from_secs(2))
);
});
}
#[test]
fn execute_maps_a_client_deadline_to_a_deadline_cancel() {
asupersync::test_utils::run_test(|| async {
let client = scripted_client(4, vec![Err(AsupersyncClientError::DeadlineExceeded)]);
let cx = Cx::for_testing();
let outcome = client.poll_statement(&cx, poll_request()).await;
let reason = match outcome {
TransportOutcome::Cancelled(reason) => reason,
other => {
assert!(
matches!(other, TransportOutcome::Cancelled(_)),
"expected a deadline cancel, got {other:?}"
);
return;
}
};
assert_eq!(reason.kind, CancelKind::Deadline);
assert_eq!(client.client.requests().len(), 1);
});
}
#[test]
fn execute_reports_an_in_flight_user_cancel_and_the_masked_cleanup_still_sends_the_remote_cancel()
{
asupersync::test_utils::run_test(|| async {
let cx = Cx::for_testing();
let mut raw = ScriptedRaw::new(vec![
Err(AsupersyncClientError::Cancelled),
ok_json(200, "{\"status\":\"cancelled\"}"),
]);
raw.cancel_in_flight = Some(cx.clone());
let client = SnowflakeHttpClient::new(fast_retry_config(4), raw);
let outcome = client.poll_statement(&cx, poll_request()).await;
let reason = match outcome {
TransportOutcome::Cancelled(reason) => reason,
other => {
assert!(
matches!(other, TransportOutcome::Cancelled(_)),
"expected a cancellation, got {other:?}"
);
return;
}
};
assert_eq!(reason.kind, CancelKind::User);
let cleanup = client
.cancel_after_local_cancel(&cx, auth(), StatementHandle::new("stmt-poll-1"), reason)
.await;
assert!(
matches!(cleanup, TransportOutcome::Ok(_)),
"the masked cleanup must reach the endpoint and be acknowledged: {cleanup:?}"
);
let requests = client.client.requests();
assert_eq!(requests.len(), 2, "poll + remote cancel");
assert_eq!(requests[1].method, Method::Post);
assert!(
requests[1]
.url
.ends_with("/api/v2/statements/stmt-poll-1/cancel"),
"{}",
requests[1].url
);
});
}
#[test]
fn execute_short_circuits_on_an_already_cancelled_context() {
asupersync::test_utils::run_test(|| async {
let client = scripted_client(4, vec![ok_json(202, "{}")]);
let cx = Cx::for_testing();
cx.cancel_with(CancelKind::User, Some("operator hit ctrl-c"));
let outcome = client.poll_statement(&cx, poll_request()).await;
assert!(
matches!(outcome, TransportOutcome::Cancelled(_)),
"{outcome:?}"
);
assert!(
client.client.requests().is_empty(),
"nothing is sent after a cancel"
);
});
}
#[test]
fn execute_decodes_a_gzip_partition_end_to_end_and_asks_for_gzip() {
use asupersync::http::compress::{Compressor, GzipCompressor};
asupersync::test_utils::run_test(|| async {
let plain = br#"{"data":[["1","alpha"]]}"#;
let mut compressed = Vec::new();
let mut compressor = GzipCompressor::new();
let compress_res = compressor.compress(plain, &mut compressed);
assert!(compress_res.is_ok(), "compress failed: {compress_res:?}");
let finish_res = compressor.finish(&mut compressed);
assert!(finish_res.is_ok(), "finish failed: {finish_res:?}");
let client = scripted_client(
4,
vec![Ok(Response::new(200, "OK", compressed.clone())
.with_header("content-encoding", "gzip"))],
);
let cx = Cx::for_testing();
let outcome = client
.fetch_partition(
&cx,
PartitionHttpRequest {
auth: auth(),
statement_handle: StatementHandle::new("stmt-1"),
partition: 3,
},
)
.await;
let body = match outcome {
TransportOutcome::Ok(body) => body,
other => {
assert!(
matches!(other, TransportOutcome::Ok(_)),
"expected a decoded partition, got {other:?}"
);
return;
}
};
assert_eq!(body.body, plain);
assert_eq!(body.compression.content_encoding, ContentEncoding::Gzip);
assert_eq!(body.compression.compressed_bytes, compressed.len() as u64);
let requests = client.client.requests();
assert_eq!(requests.len(), 1);
assert!(
requests[0]
.url
.ends_with("/api/v2/statements/stmt-1?partition=3"),
"{}",
requests[0].url
);
assert!(
requests[0]
.headers
.iter()
.any(|(name, value)| name == HEADER_ACCEPT_ENCODING
&& value == PARTITION_ACCEPT_ENCODING),
"partition fetches advertise gzip"
);
});
}
fn poll_ready<F: Future>(future: F) -> F::Output {
let waker = std::task::Waker::noop();
let mut task = std::task::Context::from_waker(waker);
let mut future = std::pin::pin!(future);
match future.as_mut().poll(&mut task) {
std::task::Poll::Ready(output) => output,
std::task::Poll::Pending => unreachable!("test future unexpectedly pending"),
}
}
#[test]
fn cancellation_mask_defers_cleanup_checkpoint() {
let cx = Cx::for_testing();
cx.cancel_with(CancelKind::User, Some("cleanup"));
let checkpoint_ok = poll_ready(run_with_cancellation_mask(&cx, async {
cx.checkpoint().is_ok()
}));
assert!(checkpoint_ok);
assert!(cx.checkpoint().is_err());
}
#[test]
fn stream_cancel_cleanup_success_still_reports_cancelled() {
let reason = CancelReason::user("partition stream cancelled");
let cleanup = TransportOutcome::ok(CancelHttpResponse {
status: StatusClass::Completed,
body: br#"{"code":"090001"}"#.to_vec(),
});
let outcome = stream_cancel_cleanup_outcome(cleanup, reason.clone());
match outcome {
TransportOutcome::Cancelled(actual) => assert_eq!(actual.kind, reason.kind),
other => {
assert!(
matches!(other, TransportOutcome::Cancelled(_)),
"cancelled stream cleanup must not report success, got {other:?}"
);
}
}
}
struct CancellingSink {
cx: Cx,
accepted: u32,
}
impl PartitionSink for CancellingSink {
fn accept(
&mut self,
_cx: &Cx,
_partition: DecodedPartition,
) -> impl Future<Output = Result<(), TransportError>> {
self.accepted = self.accepted.saturating_add(1);
self.cx
.cancel_with(CancelKind::User, Some("partition sink cancelled"));
std::future::ready(Ok(()))
}
}
#[test]
fn partition_stream_observes_cancel_after_seed_accept() {
let cx = Cx::for_testing();
let client = SnowflakeHttpClient::new(
TransportConfig::new(endpoint()),
AsupersyncHttpClient::new(),
);
let request = PartitionStreamRequest {
auth: auth(),
statement_handle: StatementHandle::new("stmt-1"),
first_partition: 1,
end_partition_exclusive: 1,
max_concurrent_fetches: 1,
child_budget: Budget::unlimited(),
remote_cancel_on_local_cancel: false,
seed_partitions: vec![DecodedPartition {
partition: 0,
body: b"[]".to_vec(),
compression: CompressionEvidence {
content_encoding: ContentEncoding::Identity,
compressed_bytes: 2,
uncompressed_bytes: 2,
},
}],
};
let mut sink = CancellingSink {
cx: cx.clone(),
accepted: 0,
};
let outcome = poll_ready(client.stream_partitions(&cx, request, &mut sink));
let cancel_kind = match outcome {
TransportOutcome::Cancelled(reason) => Some(reason.kind),
TransportOutcome::Ok(_) | TransportOutcome::Err(_) | TransportOutcome::Panicked(_) => {
None
}
};
assert_eq!(sink.accepted, 1);
assert_eq!(cancel_kind, Some(CancelKind::User));
}
#[test]
fn partition_stream_plan_validates_concurrency() -> Result<(), String> {
let request = PartitionStreamRequest {
auth: auth(),
statement_handle: StatementHandle::new("stmt-1"),
first_partition: 1,
end_partition_exclusive: 3,
max_concurrent_fetches: 2,
child_budget: Budget::unlimited(),
remote_cancel_on_local_cancel: true,
seed_partitions: Vec::new(),
};
assert_eq!(
request
.plan()
.map_err(|e| format!("plan failed: {e}"))?
.planned_partitions,
2
);
let invalid = PartitionStreamRequest {
max_concurrent_fetches: 0,
..request
};
assert!(invalid.plan().is_err());
Ok(())
}
#[test]
fn partition_stream_plan_rejects_seed_overlap_and_reordering() -> Result<(), String> {
let request = PartitionStreamRequest {
auth: auth(),
statement_handle: StatementHandle::new("stmt-1"),
first_partition: 2,
end_partition_exclusive: 4,
max_concurrent_fetches: 2,
child_budget: Budget::unlimited(),
remote_cancel_on_local_cancel: true,
seed_partitions: vec![decoded_partition(0), decoded_partition(1)],
};
let summary = request
.plan()
.map_err(|e| format!("valid seeds before fetch range failed: {e}"))?;
assert_eq!(summary.accepted_seed_partitions, 2);
let overlaps = PartitionStreamRequest {
first_partition: 1,
seed_partitions: vec![decoded_partition(1)],
..request.clone()
};
assert!(overlaps.plan().is_err());
let duplicate = PartitionStreamRequest {
seed_partitions: vec![decoded_partition(0), decoded_partition(0)],
..request.clone()
};
assert!(duplicate.plan().is_err());
let out_of_order = PartitionStreamRequest {
seed_partitions: vec![decoded_partition(1), decoded_partition(0)],
..request
};
assert!(out_of_order.plan().is_err());
Ok(())
}
fn decoded_partition(partition: u32) -> DecodedPartition {
DecodedPartition {
partition,
body: b"[]".to_vec(),
compression: CompressionEvidence {
content_encoding: ContentEncoding::Identity,
compressed_bytes: 2,
uncompressed_bytes: 2,
},
}
}
}