use std::cell::RefCell;
use std::future::Future;
use std::time::Duration;
use std::pin::Pin;
use std::task::{Context, Poll};
use asupersync::Cx;
use franken_snowflake_core::cancel::CancelReason;
use franken_snowflake_core::error::{SnowflakeError, SnowflakeErrorCode};
use franken_snowflake_core::ids::StatementHandle;
use franken_snowflake_core::outcome::SnowflakeOutcome;
use franken_snowflake_core::redact::redact;
use franken_snowflake_http::{
AuthorizationDescriptor, CancelHttpResponse, PartitionBody, PartitionHttpRequest,
PollHttpRequest, PollHttpResponse, RawHttp, SnowflakeHttpClient, StatusClass,
SubmitHttpRequest, SubmitHttpResponse, TransportOutcome, TransportRoute,
};
use crate::lifecycle::{
CompletedStatement, MIN_POLL_INTERVAL, PollPlan, Progress, StatementMachine,
};
use crate::request::{SubmitQueryParams, SubmitStatementRequest};
use crate::response::ResultSet;
use crate::status::ResponseClass;
pub type StatementOutcome = SnowflakeOutcome<CompletedStatement>;
pub trait StatementTransport {
fn submit_statement(
&self,
cx: &Cx,
request: SubmitHttpRequest,
) -> impl Future<Output = TransportOutcome<SubmitHttpResponse>>;
fn poll_statement(
&self,
cx: &Cx,
request: PollHttpRequest,
) -> impl Future<Output = TransportOutcome<PollHttpResponse>>;
fn fetch_partition(
&self,
cx: &Cx,
request: PartitionHttpRequest,
) -> impl Future<Output = TransportOutcome<PartitionBody>>;
fn cancel_after_local_cancel(
&self,
cx: &Cx,
auth: AuthorizationDescriptor,
statement_handle: StatementHandle,
reason: CancelReason,
) -> impl Future<Output = TransportOutcome<CancelHttpResponse>>;
fn cancel_orphaned_statement(
&self,
cx: &Cx,
auth: AuthorizationDescriptor,
statement_handle: StatementHandle,
) -> impl Future<Output = TransportOutcome<CancelHttpResponse>>;
}
impl<H: RawHttp> StatementTransport for SnowflakeHttpClient<H> {
async fn submit_statement(
&self,
cx: &Cx,
request: SubmitHttpRequest,
) -> TransportOutcome<SubmitHttpResponse> {
Self::submit_statement(self, cx, request).await
}
async fn poll_statement(
&self,
cx: &Cx,
request: PollHttpRequest,
) -> TransportOutcome<PollHttpResponse> {
Self::poll_statement(self, cx, request).await
}
async fn fetch_partition(
&self,
cx: &Cx,
request: PartitionHttpRequest,
) -> TransportOutcome<PartitionBody> {
Self::fetch_partition(self, cx, request).await
}
async fn cancel_after_local_cancel(
&self,
cx: &Cx,
auth: AuthorizationDescriptor,
statement_handle: StatementHandle,
reason: CancelReason,
) -> TransportOutcome<CancelHttpResponse> {
Self::cancel_after_local_cancel(self, cx, auth, statement_handle, reason).await
}
async fn cancel_orphaned_statement(
&self,
cx: &Cx,
auth: AuthorizationDescriptor,
statement_handle: StatementHandle,
) -> TransportOutcome<CancelHttpResponse> {
Self::cancel_orphaned_statement(self, cx, auth, statement_handle).await
}
}
pub trait AuthProvider {
fn descriptor(&mut self) -> Result<AuthorizationDescriptor, SnowflakeError>;
fn on_unauthorized(&mut self) -> Result<bool, SnowflakeError>;
}
impl AuthProvider for AuthorizationDescriptor {
fn descriptor(&mut self) -> Result<AuthorizationDescriptor, SnowflakeError> {
Ok(self.clone())
}
fn on_unauthorized(&mut self) -> Result<bool, SnowflakeError> {
Ok(false)
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct DriverStats {
pub polls: u32,
pub partitions_fetched: u32,
}
pub async fn run_statement<T: StatementTransport>(
cx: &Cx,
client: &T,
auth: AuthorizationDescriptor,
request: SubmitStatementRequest,
params: SubmitQueryParams,
poll_plan: PollPlan,
) -> StatementOutcome {
run_statement_with_stats(cx, client, auth, request, params, poll_plan)
.await
.0
}
pub async fn run_statement_with_stats<T: StatementTransport>(
cx: &Cx,
client: &T,
auth: AuthorizationDescriptor,
request: SubmitStatementRequest,
params: SubmitQueryParams,
poll_plan: PollPlan,
) -> (StatementOutcome, DriverStats) {
let mut frozen = auth;
run_statement_with_auth(cx, client, &mut frozen, request, params, poll_plan).await
}
pub async fn run_statement_with_auth<T: StatementTransport, A: AuthProvider>(
cx: &Cx,
client: &T,
auth: &mut A,
request: SubmitStatementRequest,
params: SubmitQueryParams,
poll_plan: PollPlan,
) -> (StatementOutcome, DriverStats) {
let mut stats = DriverStats::default();
let outcome = drive(
cx,
client,
auth,
Start::Submit { request, params },
poll_plan,
&mut stats,
StatementHooks::default(),
)
.await;
(outcome, stats)
}
pub trait RowSink {
fn accept(
&mut self,
result_set: &ResultSet,
rows: Vec<Vec<Option<String>>>,
) -> Result<(), SnowflakeError>;
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum DriverEvent {
Submitted {
statement_handle: Option<String>,
running: bool,
},
Polled {
polls: u32,
},
PartitionFetched {
index: u32,
rows: u64,
bytes: u64,
},
Completed {
rows: i64,
partitions: u32,
},
RemoteCancel {
statement_handle: String,
acknowledged: bool,
detail: String,
},
}
pub trait DriverObserver {
fn event(&mut self, event: DriverEvent);
}
#[derive(Default)]
pub struct StatementHooks<'a> {
pub sink: Option<&'a mut dyn RowSink>,
pub observer: Option<&'a mut dyn DriverObserver>,
}
pub async fn run_statement_streaming<T: StatementTransport, A: AuthProvider>(
cx: &Cx,
client: &T,
auth: &mut A,
request: SubmitStatementRequest,
params: SubmitQueryParams,
poll_plan: PollPlan,
sink: &mut dyn RowSink,
) -> (StatementOutcome, DriverStats) {
let hooks = StatementHooks {
sink: Some(sink),
observer: None,
};
run_statement_hooked(cx, client, auth, request, params, poll_plan, hooks).await
}
pub async fn run_statement_hooked<T: StatementTransport, A: AuthProvider>(
cx: &Cx,
client: &T,
auth: &mut A,
request: SubmitStatementRequest,
params: SubmitQueryParams,
poll_plan: PollPlan,
hooks: StatementHooks<'_>,
) -> (StatementOutcome, DriverStats) {
let mut stats = DriverStats::default();
let outcome = drive(
cx,
client,
auth,
Start::Submit { request, params },
poll_plan,
&mut stats,
hooks,
)
.await;
(outcome, stats)
}
#[derive(Clone, Debug, PartialEq)]
pub struct MultiStatementResult {
pub parent: CompletedStatement,
pub statements: Vec<CompletedStatement>,
}
pub async fn run_multi_statement_hooked<T: StatementTransport, A: AuthProvider>(
cx: &Cx,
client: &T,
auth: &mut A,
request: SubmitStatementRequest,
params: SubmitQueryParams,
poll_plan: PollPlan,
mut observer: Option<&mut dyn DriverObserver>,
) -> (SnowflakeOutcome<MultiStatementResult>, DriverStats) {
let mut stats = DriverStats::default();
let parent = match drive(
cx,
client,
auth,
Start::Submit { request, params },
poll_plan,
&mut stats,
StatementHooks {
sink: None,
observer: observer
.as_mut()
.map(|observer| &mut **observer as &mut dyn DriverObserver),
},
)
.await
{
SnowflakeOutcome::Ok(parent) => parent,
SnowflakeOutcome::Err(error) => return (SnowflakeOutcome::err(error), stats),
SnowflakeOutcome::Cancelled(reason) => return (SnowflakeOutcome::cancelled(reason), stats),
SnowflakeOutcome::Panicked(payload) => return (SnowflakeOutcome::panicked(payload), stats),
};
let handles = parent
.result_set
.statement_handles
.clone()
.unwrap_or_default();
if handles.is_empty() {
return (
SnowflakeOutcome::err(SnowflakeError::new(
SnowflakeErrorCode::UpstreamError,
"the SQL API answered a multi-statement request without statementHandles",
)),
stats,
);
}
let mut statements = Vec::with_capacity(handles.len());
for handle in handles {
match drive(
cx,
client,
auth,
Start::Handle(handle),
poll_plan,
&mut stats,
StatementHooks {
sink: None,
observer: observer
.as_mut()
.map(|observer| &mut **observer as &mut dyn DriverObserver),
},
)
.await
{
SnowflakeOutcome::Ok(done) => statements.push(done),
SnowflakeOutcome::Err(error) => return (SnowflakeOutcome::err(error), stats),
SnowflakeOutcome::Cancelled(reason) => {
return (SnowflakeOutcome::cancelled(reason), stats);
}
SnowflakeOutcome::Panicked(payload) => {
return (SnowflakeOutcome::panicked(payload), stats);
}
}
}
(
SnowflakeOutcome::ok(MultiStatementResult { parent, statements }),
stats,
)
}
enum Start {
Submit {
request: SubmitStatementRequest,
params: SubmitQueryParams,
},
Handle(StatementHandle),
}
fn notify(observer: &mut Option<&mut dyn DriverObserver>, event: DriverEvent) {
if let Some(observer) = observer.as_mut() {
observer.event(event);
}
}
fn progress_handle(progress: &Progress) -> Option<String> {
match progress {
Progress::PollAgain(handle) | Progress::FetchPartition { handle, .. } => {
Some(handle.as_str().to_owned())
}
Progress::Complete(done) => Some(done.statement_handle.as_str().to_owned()),
Progress::TimedOut(_) | Progress::Failed(_) => None,
}
}
fn unauthorized_error(step: &str, detail: &str) -> SnowflakeError {
SnowflakeError::new(
SnowflakeErrorCode::CredentialExpired,
format!("SQL API returned 401 Unauthorized on {step}: {detail}"),
)
}
fn refresh_after_unauthorized<A: AuthProvider>(
provider: &mut A,
reauth_left: &mut u8,
step: &str,
) -> Result<AuthorizationDescriptor, SnowflakeError> {
if *reauth_left == 0 {
return Err(unauthorized_error(
step,
"the re-signed credential was rejected again; not retrying further",
));
}
*reauth_left = reauth_left.saturating_sub(1);
match provider.on_unauthorized()? {
true => provider.descriptor(),
false => Err(unauthorized_error(
step,
"this credential lane cannot re-sign mid-flight; issue a fresh token and retry",
)),
}
}
struct CancelRecorder<'t, T> {
inner: &'t T,
cancels: RefCell<Vec<DriverEvent>>,
}
impl<T> CancelRecorder<'_, T> {
fn record(&self, handle: &StatementHandle, outcome: &TransportOutcome<CancelHttpResponse>) {
let (acknowledged, detail) = match outcome {
SnowflakeOutcome::Ok(response) => (
response.status == StatusClass::Completed,
status_class_label(response.status).to_owned(),
),
SnowflakeOutcome::Err(error) => (false, redact(&error.message).into_owned()),
SnowflakeOutcome::Cancelled(reason) => (
false,
format!(
"the cancel request was itself cancelled ({:?})",
reason.kind
),
),
SnowflakeOutcome::Panicked(_) => (false, "the cancel request panicked".to_owned()),
};
self.cancels.borrow_mut().push(DriverEvent::RemoteCancel {
statement_handle: handle.as_str().to_owned(),
acknowledged,
detail,
});
}
}
const fn status_class_label(status: StatusClass) -> &'static str {
match status {
StatusClass::Completed => "completed",
StatusClass::Running => "running",
StatusClass::StatementTimeout => "statement_timeout",
StatusClass::QueryFailure => "query_failure",
StatusClass::RateLimited => "rate_limited",
StatusClass::ServerErrorRetryable => "server_error",
StatusClass::Unauthorized => "unauthorized",
StatusClass::Unexpected => "unexpected",
}
}
impl<T: StatementTransport> StatementTransport for CancelRecorder<'_, T> {
async fn submit_statement(
&self,
cx: &Cx,
request: SubmitHttpRequest,
) -> TransportOutcome<SubmitHttpResponse> {
self.inner.submit_statement(cx, request).await
}
async fn poll_statement(
&self,
cx: &Cx,
request: PollHttpRequest,
) -> TransportOutcome<PollHttpResponse> {
self.inner.poll_statement(cx, request).await
}
async fn fetch_partition(
&self,
cx: &Cx,
request: PartitionHttpRequest,
) -> TransportOutcome<PartitionBody> {
self.inner.fetch_partition(cx, request).await
}
async fn cancel_after_local_cancel(
&self,
cx: &Cx,
auth: AuthorizationDescriptor,
statement_handle: StatementHandle,
reason: CancelReason,
) -> TransportOutcome<CancelHttpResponse> {
let outcome = self
.inner
.cancel_after_local_cancel(cx, auth, statement_handle.clone(), reason)
.await;
self.record(&statement_handle, &outcome);
outcome
}
async fn cancel_orphaned_statement(
&self,
cx: &Cx,
auth: AuthorizationDescriptor,
statement_handle: StatementHandle,
) -> TransportOutcome<CancelHttpResponse> {
let outcome = self
.inner
.cancel_orphaned_statement(cx, auth, statement_handle.clone())
.await;
self.record(&statement_handle, &outcome);
outcome
}
}
#[allow(clippy::too_many_arguments)]
async fn drive<T: StatementTransport, A: AuthProvider>(
cx: &Cx,
client: &T,
provider: &mut A,
start: Start,
poll_plan: PollPlan,
stats: &mut DriverStats,
hooks: StatementHooks<'_>,
) -> StatementOutcome {
let recorder = CancelRecorder {
inner: client,
cancels: RefCell::new(Vec::new()),
};
let StatementHooks { sink, mut observer } = hooks;
let outcome = drive_statement(
cx,
&recorder,
provider,
start,
poll_plan,
stats,
StatementHooks {
sink: sink.map(|sink| sink as &mut dyn RowSink),
observer: observer
.as_mut()
.map(|observer| &mut **observer as &mut dyn DriverObserver),
},
)
.await;
for event in recorder.cancels.take() {
notify(&mut observer, event);
}
outcome
}
#[allow(clippy::too_many_arguments)]
async fn submit<T: StatementTransport, A: AuthProvider>(
cx: &Cx,
client: &T,
provider: &mut A,
auth: &mut AuthorizationDescriptor,
reauth_left: &mut u8,
request: &SubmitStatementRequest,
params: &SubmitQueryParams,
machine: &mut StatementMachine,
) -> SnowflakeOutcome<Progress> {
let body = match serde_json::to_vec(request) {
Ok(body) => body,
Err(error) => {
return SnowflakeOutcome::err(SnowflakeError::new(
SnowflakeErrorCode::UsageError,
format!("failed to serialize submit body: {error}"),
));
}
};
let submit_response = loop {
let submit = SubmitHttpRequest {
route: submit_route(params),
auth: auth.clone(),
body: body.clone(),
retry_resubmit: params.retry,
};
match client.submit_statement(cx, submit).await {
SnowflakeOutcome::Ok(response) if response.status == StatusClass::Unauthorized => {
match refresh_after_unauthorized(provider, reauth_left, "submit") {
Ok(fresh) => *auth = fresh,
Err(error) => return SnowflakeOutcome::err(error),
}
}
SnowflakeOutcome::Ok(response) => break response,
SnowflakeOutcome::Err(error) => return SnowflakeOutcome::err(error),
SnowflakeOutcome::Cancelled(reason) => return SnowflakeOutcome::cancelled(reason),
SnowflakeOutcome::Panicked(payload) => return SnowflakeOutcome::panicked(payload),
}
};
*reauth_left = 1;
match machine.on_submit(
response_class(submit_response.status),
&submit_response.body,
) {
Ok(progress) => SnowflakeOutcome::ok(progress),
Err(error) => SnowflakeOutcome::err(error.into_snowflake_error()),
}
}
async fn drive_statement<T: StatementTransport, A: AuthProvider>(
cx: &Cx,
client: &T,
provider: &mut A,
start: Start,
poll_plan: PollPlan,
stats: &mut DriverStats,
hooks: StatementHooks<'_>,
) -> StatementOutcome {
let StatementHooks {
mut sink,
mut observer,
} = hooks;
let mut auth = match provider.descriptor() {
Ok(auth) => auth,
Err(error) => return SnowflakeOutcome::err(error),
};
let mut reauth_left: u8 = 1;
let poll_interval = poll_plan.effective_poll_interval();
let mut machine = StatementMachine::new(poll_plan);
let mut poll_now = false;
let mut progress = match start {
Start::Handle(handle) => {
poll_now = true;
Progress::PollAgain(handle)
}
Start::Submit { request, params } => {
let progress = match submit(
cx,
client,
provider,
&mut auth,
&mut reauth_left,
&request,
¶ms,
&mut machine,
)
.await
{
SnowflakeOutcome::Ok(progress) => progress,
SnowflakeOutcome::Err(error) => return SnowflakeOutcome::err(error),
SnowflakeOutcome::Cancelled(reason) => return SnowflakeOutcome::cancelled(reason),
SnowflakeOutcome::Panicked(payload) => return SnowflakeOutcome::panicked(payload),
};
notify(
&mut observer,
DriverEvent::Submitted {
statement_handle: progress_handle(&progress),
running: matches!(progress, Progress::PollAgain(_)),
},
);
progress
}
};
loop {
match progress {
Progress::Complete(mut completed) => {
notify(
&mut observer,
DriverEvent::Completed {
rows: completed.result_set.total_rows(),
partitions: completed.fetched_partitions,
},
);
if let Some(sink) = sink.as_mut() {
let rows = std::mem::take(&mut completed.rows);
if let Err(error) = sink.accept(&completed.result_set, rows) {
return SnowflakeOutcome::err(error);
}
}
return SnowflakeOutcome::ok(completed);
}
Progress::TimedOut(failure) => {
return SnowflakeOutcome::err(terminal_failure_error(
SnowflakeErrorCode::StatementTimeout,
failure,
));
}
Progress::Failed(failure) => {
return SnowflakeOutcome::err(terminal_failure_error(
SnowflakeErrorCode::StatementFailed,
failure,
));
}
Progress::PollAgain(handle) => {
if cx.checkpoint().is_err() {
return cancel_locally(cx, client, &auth, &handle, local_cancel_reason(cx))
.await;
}
if !std::mem::take(&mut poll_now)
&& let Err(reason) = wait_poll_interval(cx, poll_interval).await
{
return cancel_locally(cx, client, &auth, &handle, reason).await;
}
stats.polls = stats.polls.saturating_add(1);
auth = match provider.descriptor() {
Ok(fresh) => fresh,
Err(error) => {
return abandon_with_error(cx, client, &auth, &handle, error).await;
}
};
let poll = client
.poll_statement(
cx,
PollHttpRequest {
auth: auth.clone(),
statement_handle: handle.clone(),
},
)
.await;
let response = match poll {
SnowflakeOutcome::Ok(response) => response,
SnowflakeOutcome::Err(error) => {
return abandon_with_error(cx, client, &auth, &handle, error).await;
}
SnowflakeOutcome::Cancelled(reason) => {
return cancel_locally(cx, client, &auth, &handle, reason).await;
}
SnowflakeOutcome::Panicked(payload) => {
return abandon_with_outcome(
cx,
client,
&auth,
&handle,
SnowflakeOutcome::panicked(payload),
)
.await;
}
};
if response.status == StatusClass::Unauthorized {
match refresh_after_unauthorized(provider, &mut reauth_left, "poll") {
Ok(fresh) => {
auth = fresh;
progress = Progress::PollAgain(handle);
continue;
}
Err(error) => {
return abandon_with_error(cx, client, &auth, &handle, error).await;
}
}
}
reauth_left = 1;
progress = match machine.on_poll(response_class(response.status), &response.body) {
Ok(progress) => progress,
Err(error) => {
return abandon_with_error(
cx,
client,
&auth,
&handle,
error.into_snowflake_error(),
)
.await;
}
};
notify(&mut observer, DriverEvent::Polled { polls: stats.polls });
}
Progress::FetchPartition { handle, partition } => {
if cx.checkpoint().is_err() {
return cancel_locally(cx, client, &auth, &handle, local_cancel_reason(cx))
.await;
}
if let Some(sink) = sink.as_mut() {
let rows = machine.drain_rows();
if !rows.is_empty()
&& let Some(result_set) = machine.result_set()
&& let Err(error) = sink.accept(result_set, rows)
{
return abandon_with_error(cx, client, &auth, &handle, error).await;
}
}
let (next, total) = machine
.assembling_window()
.unwrap_or((partition, partition.saturating_add(1)));
if let Some(cap) = poll_plan.row_cap
&& machine.rows_assembled() >= cap
{
return match machine.complete_early() {
Ok(mut done) => {
notify(
&mut observer,
DriverEvent::Completed {
rows: done.result_set.total_rows(),
partitions: done.fetched_partitions,
},
);
if let Some(sink) = sink.as_mut() {
let rows = std::mem::take(&mut done.rows);
if let Err(error) = sink.accept(&done.result_set, rows) {
return SnowflakeOutcome::err(error);
}
}
SnowflakeOutcome::ok(done)
}
Err(error) => {
abandon_with_error(
cx,
client,
&auth,
&handle,
error.into_snowflake_error(),
)
.await
}
};
}
auth = match provider.descriptor() {
Ok(fresh) => fresh,
Err(error) => {
return abandon_with_error(cx, client, &auth, &handle, error).await;
}
};
let window =
u32::try_from(poll_plan.effective_partition_concurrency()).unwrap_or(u32::MAX);
let window_end = next.saturating_add(window).min(total);
let window_auth = auth.clone();
let fetched =
fetch_window(cx, client, &window_auth, &handle, next..window_end).await;
stats.partitions_fetched = stats
.partitions_fetched
.saturating_add(window_end.saturating_sub(next));
let mut after_window = None;
for (offset, fetch) in fetched.into_iter().enumerate() {
let index = next.saturating_add(u32::try_from(offset).unwrap_or(u32::MAX));
let mut response = match fetch {
SnowflakeOutcome::Ok(response) => response,
SnowflakeOutcome::Err(error) => {
return abandon_with_error(cx, client, &auth, &handle, error).await;
}
SnowflakeOutcome::Cancelled(reason) => {
return cancel_locally(cx, client, &auth, &handle, reason).await;
}
SnowflakeOutcome::Panicked(payload) => {
return abandon_with_outcome(
cx,
client,
&auth,
&handle,
SnowflakeOutcome::panicked(payload),
)
.await;
}
};
if response.status == StatusClass::Unauthorized {
if auth == window_auth {
auth = match refresh_after_unauthorized(
provider,
&mut reauth_left,
"partition fetch",
) {
Ok(fresh) => fresh,
Err(error) => {
return abandon_with_error(cx, client, &auth, &handle, error)
.await;
}
};
}
stats.partitions_fetched = stats.partitions_fetched.saturating_add(1);
let refetch = client
.fetch_partition(
cx,
PartitionHttpRequest {
auth: auth.clone(),
statement_handle: handle.clone(),
partition: index,
},
)
.await;
response = match refetch {
SnowflakeOutcome::Ok(response)
if response.status == StatusClass::Unauthorized =>
{
return abandon_with_error(
cx,
client,
&auth,
&handle,
unauthorized_error(
"partition fetch",
"the re-signed credential was rejected again; not retrying further",
),
)
.await;
}
SnowflakeOutcome::Ok(response) => response,
SnowflakeOutcome::Err(error) => {
return abandon_with_error(cx, client, &auth, &handle, error).await;
}
SnowflakeOutcome::Cancelled(reason) => {
return cancel_locally(cx, client, &auth, &handle, reason).await;
}
SnowflakeOutcome::Panicked(payload) => {
return abandon_with_outcome(
cx,
client,
&auth,
&handle,
SnowflakeOutcome::panicked(payload),
)
.await;
}
};
}
reauth_left = 1;
let partition_rows = machine
.result_set()
.and_then(|result_set| {
result_set
.result_set_meta_data
.partition_info
.get(usize::try_from(index).unwrap_or(usize::MAX))
})
.map_or(0, |info| u64::try_from(info.row_count).unwrap_or(0));
let partition_bytes = u64::try_from(response.body.len()).unwrap_or(u64::MAX);
let partition_progress = match machine.on_partition(
response_class(response.status),
index,
&response.body,
) {
Ok(progress) => progress,
Err(error) => {
return abandon_with_error(
cx,
client,
&auth,
&handle,
error.into_snowflake_error(),
)
.await;
}
};
notify(
&mut observer,
DriverEvent::PartitionFetched {
index,
rows: partition_rows,
bytes: partition_bytes,
},
);
after_window = Some(partition_progress);
}
progress = match after_window {
Some(progress) => progress,
None => {
return abandon_with_error(
cx,
client,
&auth,
&handle,
SnowflakeError::new(
SnowflakeErrorCode::Internal,
format!("empty partition window {next}..{window_end} of {total}"),
),
)
.await;
}
};
}
}
}
}
type BoxedFetch<'a> = Pin<Box<dyn Future<Output = TransportOutcome<PartitionBody>> + 'a>>;
async fn fetch_window<T: StatementTransport>(
cx: &Cx,
client: &T,
auth: &AuthorizationDescriptor,
handle: &StatementHandle,
partitions: std::ops::Range<u32>,
) -> Vec<TransportOutcome<PartitionBody>> {
let pending: Vec<Option<BoxedFetch<'_>>> = partitions
.map(|partition| {
let request = PartitionHttpRequest {
auth: auth.clone(),
statement_handle: handle.clone(),
partition,
};
let fetch: BoxedFetch<'_> = Box::pin(client.fetch_partition(cx, request));
Some(fetch)
})
.collect();
let done = pending.iter().map(|_| None).collect();
JoinInOrder { pending, done }.await
}
struct JoinInOrder<'a, T> {
pending: Vec<Option<Pin<Box<dyn Future<Output = T> + 'a>>>>,
done: Vec<Option<T>>,
}
impl<T: Unpin> Future for JoinInOrder<'_, T> {
type Output = Vec<T>;
fn poll(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
let mut all_done = true;
for (slot, done) in this.pending.iter_mut().zip(this.done.iter_mut()) {
if let Some(future) = slot.as_mut() {
match future.as_mut().poll(context) {
Poll::Ready(value) => {
*done = Some(value);
*slot = None;
}
Poll::Pending => all_done = false,
}
}
}
if all_done {
Poll::Ready(this.done.iter_mut().filter_map(Option::take).collect())
} else {
Poll::Pending
}
}
}
async fn abandon_with_error<T: StatementTransport>(
cx: &Cx,
client: &T,
auth: &AuthorizationDescriptor,
handle: &StatementHandle,
error: SnowflakeError,
) -> StatementOutcome {
abandon_with_outcome(cx, client, auth, handle, SnowflakeOutcome::err(error)).await
}
async fn abandon_with_outcome<T: StatementTransport>(
cx: &Cx,
client: &T,
auth: &AuthorizationDescriptor,
handle: &StatementHandle,
outcome: StatementOutcome,
) -> StatementOutcome {
let _ = client
.cancel_orphaned_statement(cx, auth.clone(), handle.clone())
.await;
outcome
}
async fn cancel_locally<T: StatementTransport>(
cx: &Cx,
client: &T,
auth: &AuthorizationDescriptor,
handle: &StatementHandle,
reason: CancelReason,
) -> StatementOutcome {
let _ = client
.cancel_after_local_cancel(cx, auth.clone(), handle.clone(), reason.clone())
.await;
SnowflakeOutcome::cancelled(reason)
}
fn local_cancel_reason(cx: &Cx) -> CancelReason {
cx.cancel_reason()
.unwrap_or_else(CancelReason::parent_cancelled)
}
fn terminal_failure_error(
code: SnowflakeErrorCode,
failure: crate::response::QueryFailureStatus,
) -> SnowflakeError {
SnowflakeError::new(code, redact(&failure.message).into_owned())
}
async fn wait_poll_interval(cx: &Cx, delay: Duration) -> Result<(), CancelReason> {
let mut remaining = delay;
while !remaining.is_zero() {
if cx.checkpoint().is_err() {
return Err(local_cancel_reason(cx));
}
let slice = remaining.min(MIN_POLL_INTERVAL);
if asupersync::time::budget_sleep(cx, slice, cx.now_for_observability())
.await
.is_err()
{
let _ = cx.checkpoint();
return Err(local_cancel_reason(cx));
}
if cx.checkpoint().is_err() {
return Err(local_cancel_reason(cx));
}
remaining = remaining.saturating_sub(slice);
}
Ok(())
}
fn submit_route(params: &SubmitQueryParams) -> TransportRoute {
let query = params.to_query_pairs();
if query.is_empty() {
TransportRoute::Submit
} else {
TransportRoute::SubmitWithQuery { query }
}
}
const fn response_class(status: StatusClass) -> ResponseClass {
match status {
StatusClass::Completed => ResponseClass::Completed,
StatusClass::Running => ResponseClass::Running,
StatusClass::StatementTimeout => ResponseClass::StatementTimeout,
StatusClass::QueryFailure => ResponseClass::StatementFailed,
StatusClass::RateLimited => ResponseClass::RateLimited,
StatusClass::ServerErrorRetryable => ResponseClass::Other(503),
StatusClass::Unauthorized => ResponseClass::Other(401),
StatusClass::Unexpected => ResponseClass::Other(0),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::response::QueryFailureStatus;
use asupersync::{Budget, CancelKind, PanicPayload, Time};
use franken_snowflake_core::outcome::{OutcomeKind, SnowflakeOutcomeExt};
use franken_snowflake_http::{
CompressionEvidence, ContentEncoding, SnowflakeAuthTokenType, TransportError,
TransportErrorCode,
};
use std::cell::{Cell, RefCell};
use std::collections::{BTreeMap, VecDeque};
const RESP_202: &[u8] = include_bytes!("../tests/fixtures/resp_202_running.json");
const RESP_200_MULTI: &[u8] =
include_bytes!("../tests/fixtures/resp_200_resultset_multi_partition.json");
const RESP_200_SINGLE: &[u8] =
include_bytes!("../tests/fixtures/resp_200_resultset_single_partition.json");
#[derive(Clone)]
enum Scripted {
Ok(StatusClass, Vec<u8>),
Err,
Panicked(&'static str),
}
struct FakeTransport {
submit: Scripted,
submit_first: RefCell<Option<Scripted>>,
polls: RefCell<Vec<Scripted>>,
polled: RefCell<Vec<String>>,
partitions: RefCell<BTreeMap<u32, VecDeque<Scripted>>>,
cancels_after_local: RefCell<Vec<(StatementHandle, CancelKind)>>,
orphan_cancels: RefCell<Vec<StatementHandle>>,
orphan_cancel_auth: RefCell<Vec<String>>,
orphan_cancel_result: Scripted,
orphan_cleanup_finished: Cell<bool>,
auth_seen: RefCell<Vec<String>>,
partition_events: RefCell<Vec<(&'static str, u32)>>,
yield_once: Cell<bool>,
}
impl FakeTransport {
fn new(submit: Scripted) -> Self {
Self {
submit,
submit_first: RefCell::new(None),
polls: RefCell::new(Vec::new()),
polled: RefCell::new(Vec::new()),
partitions: RefCell::new(BTreeMap::new()),
cancels_after_local: RefCell::new(Vec::new()),
orphan_cancels: RefCell::new(Vec::new()),
orphan_cancel_auth: RefCell::new(Vec::new()),
orphan_cancel_result: Scripted::Ok(StatusClass::Completed, Vec::new()),
orphan_cleanup_finished: Cell::new(false),
auth_seen: RefCell::new(Vec::new()),
partition_events: RefCell::new(Vec::new()),
yield_once: Cell::new(false),
}
}
fn script_partition(&self, partition: u32, scripted: Scripted) {
self.partitions
.borrow_mut()
.entry(partition)
.or_default()
.push_back(scripted);
}
fn events(&self) -> Vec<String> {
self.partition_events
.borrow()
.iter()
.map(|(kind, partition)| format!("{kind}{partition}"))
.collect()
}
fn transport_error() -> SnowflakeError {
TransportError::new(TransportErrorCode::NetworkError, "connection reset")
.into_snowflake_error()
}
}
impl StatementTransport for FakeTransport {
async fn submit_statement(
&self,
_cx: &Cx,
request: SubmitHttpRequest,
) -> TransportOutcome<SubmitHttpResponse> {
self.auth_seen
.borrow_mut()
.push(request.auth.redacted_fingerprint().to_owned());
let scripted = self
.submit_first
.borrow_mut()
.take()
.unwrap_or_else(|| self.submit.clone());
match scripted {
Scripted::Ok(status, body) => {
TransportOutcome::ok(SubmitHttpResponse { status, body })
}
Scripted::Err => TransportOutcome::err(Self::transport_error()),
Scripted::Panicked(message) => {
TransportOutcome::panicked(PanicPayload::new(message))
}
}
}
async fn poll_statement(
&self,
_cx: &Cx,
request: PollHttpRequest,
) -> TransportOutcome<PollHttpResponse> {
self.auth_seen
.borrow_mut()
.push(request.auth.redacted_fingerprint().to_owned());
self.polled
.borrow_mut()
.push(request.statement_handle.as_str().to_owned());
let next = self.polls.borrow_mut().remove(0);
match next {
Scripted::Ok(status, body) => {
TransportOutcome::ok(PollHttpResponse { status, body })
}
Scripted::Err => TransportOutcome::err(Self::transport_error()),
Scripted::Panicked(message) => {
TransportOutcome::panicked(PanicPayload::new(message))
}
}
}
async fn fetch_partition(
&self,
_cx: &Cx,
request: PartitionHttpRequest,
) -> TransportOutcome<PartitionBody> {
self.auth_seen
.borrow_mut()
.push(request.auth.redacted_fingerprint().to_owned());
self.partition_events
.borrow_mut()
.push(("start", request.partition));
if self.yield_once.get() {
YieldOnce { yielded: false }.await;
}
self.partition_events
.borrow_mut()
.push(("done", request.partition));
let next = self
.partitions
.borrow_mut()
.get_mut(&request.partition)
.and_then(VecDeque::pop_front);
match next {
Some(Scripted::Ok(status, body)) => TransportOutcome::ok(PartitionBody {
status,
compression: CompressionEvidence {
content_encoding: ContentEncoding::Identity,
compressed_bytes: body.len() as u64,
uncompressed_bytes: body.len() as u64,
},
body,
}),
Some(Scripted::Err) | None => TransportOutcome::err(Self::transport_error()),
Some(Scripted::Panicked(message)) => {
TransportOutcome::panicked(PanicPayload::new(message))
}
}
}
async fn cancel_after_local_cancel(
&self,
_cx: &Cx,
_auth: AuthorizationDescriptor,
statement_handle: StatementHandle,
reason: CancelReason,
) -> TransportOutcome<CancelHttpResponse> {
self.cancels_after_local
.borrow_mut()
.push((statement_handle, reason.kind));
TransportOutcome::cancelled(reason)
}
async fn cancel_orphaned_statement(
&self,
_cx: &Cx,
auth: AuthorizationDescriptor,
statement_handle: StatementHandle,
) -> TransportOutcome<CancelHttpResponse> {
self.orphan_cancels.borrow_mut().push(statement_handle);
self.orphan_cancel_auth
.borrow_mut()
.push(auth.redacted_fingerprint().to_owned());
if self.yield_once.get() {
YieldOnce { yielded: false }.await;
}
self.orphan_cleanup_finished.set(true);
match &self.orphan_cancel_result {
Scripted::Ok(status, body) => TransportOutcome::ok(CancelHttpResponse {
status: *status,
body: body.clone(),
}),
Scripted::Err => TransportOutcome::err(Self::transport_error()),
Scripted::Panicked(message) => {
TransportOutcome::panicked(PanicPayload::new(*message))
}
}
}
}
fn fake_auth() -> AuthorizationDescriptor {
AuthorizationDescriptor::bearer(
SnowflakeAuthTokenType::ProgrammaticAccessToken,
"fake-token",
"cred_test",
)
}
struct YieldOnce {
yielded: bool,
}
impl Future for YieldOnce {
type Output = ();
fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<()> {
if self.yielded {
Poll::Ready(())
} else {
self.yielded = true;
context.waker().wake_by_ref();
Poll::Pending
}
}
}
fn multi_partition_body(inline_rows: usize, partition_rows: &[usize]) -> Vec<u8> {
let mut partition_info = vec![serde_json::json!({ "rowCount": inline_rows })];
partition_info.extend(
partition_rows
.iter()
.map(|rows| serde_json::json!({ "rowCount": rows, "uncompressedSize": 1 })),
);
let total: usize = inline_rows + partition_rows.iter().sum::<usize>();
let body = serde_json::json!({
"resultSetMetaData": {
"numRows": total,
"format": "jsonv2",
"rowType": [{ "name": "P", "type": "TEXT", "nullable": false }],
"partitionInfo": partition_info
},
"data": (0..inline_rows).map(|_| vec!["p0"]).collect::<Vec<_>>(),
"code": "090001",
"statementHandle": "01b2c3d4-0000-0000-0000-000000000002",
"sqlState": "00000",
"message": "Statement executed successfully.",
"createdOn": 1_700_000_000_000_u64
});
serde_json::to_vec(&body).unwrap_or_default()
}
fn partition_body(partition: u32, rows: usize) -> Scripted {
let body = serde_json::json!({
"data": (0..rows).map(|_| vec![format!("p{partition}")]).collect::<Vec<_>>()
});
Scripted::Ok(
StatusClass::Completed,
serde_json::to_vec(&body).unwrap_or_default(),
)
}
fn windowed_transport(inline_rows: usize, partition_rows: &[usize]) -> FakeTransport {
let transport = FakeTransport::new(Scripted::Ok(
StatusClass::Completed,
multi_partition_body(inline_rows, partition_rows),
));
for (offset, rows) in partition_rows.iter().enumerate() {
let partition = u32::try_from(offset + 1).unwrap_or(u32::MAX);
transport.script_partition(partition, partition_body(partition, *rows));
}
transport
}
fn column_values(done: &CompletedStatement) -> Vec<String> {
done.rows
.iter()
.map(|row| row[0].clone().unwrap_or_default())
.collect()
}
fn multi_parent(handles: &[&str]) -> Vec<u8> {
let body = serde_json::json!({
"resultSetMetaData": {
"numRows": 1,
"format": "jsonv2",
"rowType": [{ "name": "multiple statement execution", "type": "text", "nullable": false }],
"partitionInfo": [{ "rowCount": 1 }, { "rowCount": 5 }]
},
"data": [["Multiple statements executed successfully."]],
"code": "090001",
"statementHandle": "01b2c3d4-0000-0000-0000-0000000000a0",
"statementHandles": handles,
});
serde_json::to_vec(&body).unwrap_or_default()
}
fn one_value_result(handle: &str, value: &str) -> Scripted {
let body = serde_json::json!({
"resultSetMetaData": {
"numRows": 1,
"format": "jsonv2",
"rowType": [{ "name": "V", "type": "TEXT", "nullable": false }]
},
"data": [[value]],
"code": "090001",
"statementHandle": handle,
});
Scripted::Ok(
StatusClass::Completed,
serde_json::to_vec(&body).unwrap_or_default(),
)
}
#[test]
fn a_multi_statement_request_fetches_each_statement_by_handle_in_order() {
asupersync::test_utils::run_test(|| async {
const FIRST: &str = "01b2c3d4-0000-0000-0000-0000000000a1";
const SECOND: &str = "01b2c3d4-0000-0000-0000-0000000000a2";
let transport = FakeTransport::new(Scripted::Ok(
StatusClass::Completed,
multi_parent(&[FIRST, SECOND]),
));
*transport.polls.borrow_mut() = vec![
one_value_result(FIRST, "first"),
one_value_result(SECOND, "second"),
];
let cx = Cx::for_testing();
let (outcome, stats) = run_multi_statement_hooked(
&cx,
&transport,
&mut fake_auth(),
SubmitStatementRequest::new("select 'first'; select 'second'"),
SubmitQueryParams::default(),
fast_poll_plan(5),
None,
)
.await;
let result = match outcome {
SnowflakeOutcome::Ok(result) => result,
other => panic!("expected both statements, got {other:?}"),
};
let values: Vec<Vec<String>> = result.statements.iter().map(column_values).collect();
assert_eq!(values, [["first"], ["second"]]);
assert_eq!(
result.parent.rows,
[[Some(
"Multiple statements executed successfully.".to_owned()
)]]
);
assert_eq!(transport.polled.borrow().as_slice(), [FIRST, SECOND]);
assert_eq!(stats.polls, 2);
assert!(transport.events().is_empty(), "{:?}", transport.events());
assert!(transport.orphan_cancels.borrow().is_empty());
});
}
#[test]
fn a_multi_statement_parent_without_handles_is_an_upstream_error() {
asupersync::test_utils::run_test(|| async {
let transport = FakeTransport::new(Scripted::Ok(
StatusClass::Completed,
RESP_200_SINGLE.to_vec(),
));
let cx = Cx::for_testing();
let (outcome, _) = run_multi_statement_hooked(
&cx,
&transport,
&mut fake_auth(),
SubmitStatementRequest::new("select 1; select 2"),
SubmitQueryParams::default(),
fast_poll_plan(5),
None,
)
.await;
let error = match outcome {
SnowflakeOutcome::Err(error) => error,
other => panic!("expected an error, got {other:?}"),
};
assert_eq!(error.code, SnowflakeErrorCode::UpstreamError);
assert!(transport.polled.borrow().is_empty());
});
}
#[test]
fn a_failing_statement_fails_the_multi_statement_request_before_any_fetch() {
asupersync::test_utils::run_test(|| async {
let failure = serde_json::json!({
"code": "100132",
"sqlState": "P0000",
"message": "JavaScript execution error: Uncaught Execution of multiple statements failed on statement \"select * from missing_table\"",
"statementHandle": "01b2c3d4-0000-0000-0000-0000000000a0",
});
let transport = FakeTransport::new(Scripted::Ok(
StatusClass::QueryFailure,
serde_json::to_vec(&failure).unwrap_or_default(),
));
let cx = Cx::for_testing();
let (outcome, _) = run_multi_statement_hooked(
&cx,
&transport,
&mut fake_auth(),
SubmitStatementRequest::new("select 1; select * from missing_table"),
SubmitQueryParams::default(),
fast_poll_plan(5),
None,
)
.await;
let error = match outcome {
SnowflakeOutcome::Err(error) => error,
other => panic!("expected the statement failure, got {other:?}"),
};
assert_eq!(error.code, SnowflakeErrorCode::StatementFailed);
assert!(error.message.contains("missing_table"), "{}", error.message);
assert!(transport.polled.borrow().is_empty());
});
}
#[test]
fn window_fetches_partitions_concurrently_and_assembles_in_order() {
asupersync::test_utils::run_test(|| async {
let transport = windowed_transport(1, &[1, 1, 1, 1, 1]);
transport.yield_once.set(true);
let cx = Cx::for_testing();
let (outcome, stats) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5).with_partition_concurrency(3),
)
.await;
let done = match outcome {
SnowflakeOutcome::Ok(done) => done,
other => {
assert!(
matches!(other, SnowflakeOutcome::Ok(_)),
"expected completion, got {other:?}"
);
return;
}
};
assert_eq!(
column_values(&done),
vec!["p0", "p1", "p2", "p3", "p4", "p5"]
);
assert_eq!(done.fetched_partitions, 6);
assert_eq!(done.total_partitions, 6);
assert!(!done.is_partial());
assert_eq!(stats.partitions_fetched, 5);
assert_eq!(
transport.events(),
vec![
"start1", "start2", "start3", "done1", "done2", "done3", "start4", "start5",
"done4", "done5"
]
);
});
}
#[test]
fn partition_concurrency_one_is_strictly_sequential() {
asupersync::test_utils::run_test(|| async {
let transport = windowed_transport(1, &[1, 1, 1]);
transport.yield_once.set(true);
let cx = Cx::for_testing();
let (outcome, _) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5).with_partition_concurrency(1),
)
.await;
assert!(matches!(outcome, SnowflakeOutcome::Ok(_)), "{outcome:?}");
assert_eq!(
transport.events(),
vec!["start1", "done1", "start2", "done2", "start3", "done3"]
);
});
}
#[test]
fn row_cap_stops_fetching_early_and_reports_a_partial_prefix() {
asupersync::test_utils::run_test(|| async {
let transport = windowed_transport(1, &[1, 1, 1, 1, 1]);
let cx = Cx::for_testing();
let (outcome, stats) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5)
.with_partition_concurrency(1)
.with_row_cap(Some(3)),
)
.await;
let done = match outcome {
SnowflakeOutcome::Ok(done) => done,
other => {
assert!(
matches!(other, SnowflakeOutcome::Ok(_)),
"expected completion, got {other:?}"
);
return;
}
};
assert_eq!(column_values(&done), vec!["p0", "p1", "p2"]);
assert!(done.is_partial());
assert_eq!(done.fetched_partitions, 3);
assert_eq!(done.total_partitions, 6);
assert_eq!(done.result_set.result_set_meta_data.num_rows, 6);
assert_eq!(
stats.partitions_fetched, 2,
"partitions 3..5 were never fetched"
);
assert_eq!(
transport.events(),
vec!["start1", "done1", "start2", "done2"]
);
assert!(transport.orphan_cancels.borrow().is_empty());
});
}
#[test]
fn row_cap_with_a_window_stops_after_the_window_that_crossed_it() {
asupersync::test_utils::run_test(|| async {
let transport = windowed_transport(1, &[1, 1, 1, 1, 1]);
let cx = Cx::for_testing();
let (outcome, stats) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5)
.with_partition_concurrency(3)
.with_row_cap(Some(3)),
)
.await;
let done = match outcome {
SnowflakeOutcome::Ok(done) => done,
other => {
assert!(
matches!(other, SnowflakeOutcome::Ok(_)),
"expected completion, got {other:?}"
);
return;
}
};
assert_eq!(column_values(&done), vec!["p0", "p1", "p2", "p3"]);
assert!(done.is_partial());
assert_eq!(done.fetched_partitions, 4);
assert_eq!(stats.partitions_fetched, 3);
});
}
#[test]
fn row_cap_never_cuts_a_result_that_fits() {
asupersync::test_utils::run_test(|| async {
let transport = windowed_transport(1, &[1, 1]);
let cx = Cx::for_testing();
let (outcome, _) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5).with_row_cap(Some(1_000)),
)
.await;
let done = match outcome {
SnowflakeOutcome::Ok(done) => done,
other => {
assert!(
matches!(other, SnowflakeOutcome::Ok(_)),
"expected completion, got {other:?}"
);
return;
}
};
assert!(!done.is_partial());
assert_eq!(done.rows.len(), 3);
});
}
#[test]
fn window_401_resigns_once_and_refetches_every_rejected_partition() {
asupersync::test_utils::run_test(|| async {
let transport = FakeTransport::new(Scripted::Ok(
StatusClass::Completed,
multi_partition_body(1, &[1, 1, 1]),
));
transport.script_partition(1, unauthorized());
transport.script_partition(1, partition_body(1, 1));
transport.script_partition(2, partition_body(2, 1));
transport.script_partition(3, unauthorized());
transport.script_partition(3, partition_body(3, 1));
let mut auth = FakeAuth::resigning();
let cx = Cx::for_testing();
let (outcome, stats) = run_statement_with_auth(
&cx,
&transport,
&mut auth,
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5).with_partition_concurrency(3),
)
.await;
let done = match outcome {
SnowflakeOutcome::Ok(done) => done,
other => {
assert!(
matches!(other, SnowflakeOutcome::Ok(_)),
"expected completion, got {other:?}"
);
return;
}
};
assert_eq!(column_values(&done), vec!["p0", "p1", "p2", "p3"]);
assert_eq!(
auth.resigns, 1,
"one re-sign covers every rejection in the window"
);
assert_eq!(
stats.partitions_fetched, 5,
"3 window fetches + 2 refetches"
);
assert_eq!(
*transport.auth_seen.borrow(),
vec![
"cred_gen0",
"cred_gen0",
"cred_gen0",
"cred_gen0",
"cred_gen1",
"cred_gen1"
],
"submit + window used gen0; both refetches used the re-signed gen1"
);
assert!(transport.orphan_cancels.borrow().is_empty());
});
}
#[test]
fn partition_rejected_again_after_the_resign_is_terminal_with_an_orphan_cancel() {
asupersync::test_utils::run_test(|| async {
let transport = FakeTransport::new(Scripted::Ok(
StatusClass::Completed,
multi_partition_body(1, &[1]),
));
transport.script_partition(1, unauthorized());
transport.script_partition(1, unauthorized());
let mut auth = FakeAuth::resigning();
let cx = Cx::for_testing();
let (outcome, stats) = run_statement_with_auth(
&cx,
&transport,
&mut auth,
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5),
)
.await;
let error = match outcome {
SnowflakeOutcome::Err(error) => error,
other => {
assert!(
matches!(other, SnowflakeOutcome::Err(_)),
"expected a typed error, got {other:?}"
);
return;
}
};
assert_eq!(error.code, SnowflakeErrorCode::CredentialExpired);
assert_eq!(auth.resigns, 1);
assert_eq!(stats.partitions_fetched, 2);
assert_eq!(transport.orphan_cancels.borrow().len(), 1);
});
}
#[test]
fn one_failed_fetch_in_a_window_abandons_the_statement_after_the_window_settles() {
asupersync::test_utils::run_test(|| async {
let transport = windowed_transport(1, &[1, 1, 1]);
transport.yield_once.set(true);
transport
.partitions
.borrow_mut()
.insert(2, VecDeque::from([Scripted::Err]));
let cx = Cx::for_testing();
let (outcome, stats) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5).with_partition_concurrency(3),
)
.await;
assert!(matches!(outcome, SnowflakeOutcome::Err(_)), "{outcome:?}");
assert_eq!(
transport.events(),
vec!["start1", "start2", "start3", "done1", "done2", "done3"]
);
assert_eq!(stats.partitions_fetched, 3);
assert_eq!(transport.orphan_cancels.borrow().len(), 1);
});
}
struct FakeAuth {
can_resign: bool,
generation: u32,
resigns: u32,
}
impl FakeAuth {
fn resigning() -> Self {
Self {
can_resign: true,
generation: 0,
resigns: 0,
}
}
fn frozen_lane() -> Self {
Self {
can_resign: false,
generation: 0,
resigns: 0,
}
}
}
impl AuthProvider for FakeAuth {
fn descriptor(&mut self) -> Result<AuthorizationDescriptor, SnowflakeError> {
Ok(AuthorizationDescriptor::bearer(
SnowflakeAuthTokenType::KeypairJwt,
format!("jwt-gen-{}", self.generation),
format!("cred_gen{}", self.generation),
))
}
fn on_unauthorized(&mut self) -> Result<bool, SnowflakeError> {
if !self.can_resign {
return Ok(false);
}
self.generation += 1;
self.resigns += 1;
Ok(true)
}
}
fn unauthorized() -> Scripted {
Scripted::Ok(
StatusClass::Unauthorized,
b"{\"message\":\"JWT token is invalid.\"}".to_vec(),
)
}
#[test]
fn poll_401_resigns_once_and_retries_with_the_new_token() {
asupersync::test_utils::run_test(|| async {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
transport.polls.borrow_mut().push(unauthorized());
transport
.polls
.borrow_mut()
.push(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
transport.polls.borrow_mut().push(Scripted::Ok(
StatusClass::Completed,
RESP_200_SINGLE.to_vec(),
));
let mut auth = FakeAuth::resigning();
let cx = Cx::for_testing();
let (outcome, stats) = run_statement_with_auth(
&cx,
&transport,
&mut auth,
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5),
)
.await;
assert!(matches!(outcome, SnowflakeOutcome::Ok(_)), "{outcome:?}");
assert_eq!(auth.resigns, 1);
assert_eq!(stats.polls, 3, "the retried poll is a real GET");
assert_eq!(
*transport.auth_seen.borrow(),
vec!["cred_gen0", "cred_gen0", "cred_gen1", "cred_gen1"],
"submit + first poll used gen0; the retry and the next poll used the re-signed gen1"
);
assert!(transport.orphan_cancels.borrow().is_empty());
assert!(transport.cancels_after_local.borrow().is_empty());
});
}
#[test]
fn submit_401_resigns_and_resubmits_without_a_cancel() {
asupersync::test_utils::run_test(|| async {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
*transport.submit_first.borrow_mut() = Some(unauthorized());
transport.polls.borrow_mut().push(Scripted::Ok(
StatusClass::Completed,
RESP_200_SINGLE.to_vec(),
));
let mut auth = FakeAuth::resigning();
let cx = Cx::for_testing();
let (outcome, stats) = run_statement_with_auth(
&cx,
&transport,
&mut auth,
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5),
)
.await;
assert!(matches!(outcome, SnowflakeOutcome::Ok(_)), "{outcome:?}");
assert_eq!(auth.resigns, 1);
assert_eq!(stats.polls, 1);
assert_eq!(
*transport.auth_seen.borrow(),
vec!["cred_gen0", "cred_gen1", "cred_gen1"]
);
assert!(transport.orphan_cancels.borrow().is_empty());
});
}
#[test]
fn poll_401_on_a_lane_that_cannot_resign_is_typed_and_cancels_the_orphan() {
asupersync::test_utils::run_test(|| async {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
transport.polls.borrow_mut().push(unauthorized());
let mut auth = FakeAuth::frozen_lane();
let cx = Cx::for_testing();
let (outcome, _) = run_statement_with_auth(
&cx,
&transport,
&mut auth,
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5),
)
.await;
let error = match outcome {
SnowflakeOutcome::Err(error) => error,
other => {
assert!(
matches!(other, SnowflakeOutcome::Err(_)),
"expected a typed error, got {other:?}"
);
return;
}
};
assert_eq!(error.code, SnowflakeErrorCode::CredentialExpired);
assert!(error.message.contains("401"), "{}", error.message);
assert_eq!(auth.resigns, 0);
assert_eq!(transport.orphan_cancels.borrow().len(), 1);
});
}
#[test]
fn frozen_descriptor_entry_point_treats_401_as_terminal() {
asupersync::test_utils::run_test(|| async {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
transport.polls.borrow_mut().push(unauthorized());
let cx = Cx::for_testing();
let (outcome, _) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5),
)
.await;
let error = match outcome {
SnowflakeOutcome::Err(error) => error,
other => {
assert!(
matches!(other, SnowflakeOutcome::Err(_)),
"expected a typed error, got {other:?}"
);
return;
}
};
assert_eq!(error.code, SnowflakeErrorCode::CredentialExpired);
assert_eq!(transport.orphan_cancels.borrow().len(), 1);
});
}
#[test]
fn two_consecutive_401s_stop_after_exactly_one_resign() {
asupersync::test_utils::run_test(|| async {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
transport.polls.borrow_mut().push(unauthorized());
transport.polls.borrow_mut().push(unauthorized());
let mut auth = FakeAuth::resigning();
let cx = Cx::for_testing();
let (outcome, stats) = run_statement_with_auth(
&cx,
&transport,
&mut auth,
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5),
)
.await;
let error = match outcome {
SnowflakeOutcome::Err(error) => error,
other => {
assert!(
matches!(other, SnowflakeOutcome::Err(_)),
"expected a typed error, got {other:?}"
);
return;
}
};
assert_eq!(error.code, SnowflakeErrorCode::CredentialExpired);
assert!(
error.message.contains("rejected again"),
"{}",
error.message
);
assert_eq!(auth.resigns, 1, "exactly one re-sign, no loop");
assert_eq!(stats.polls, 2);
assert_eq!(transport.orphan_cancels.borrow().len(), 1);
assert!(transport.polls.borrow().is_empty());
});
}
#[test]
fn partition_401_resigns_once_and_refetches() {
asupersync::test_utils::run_test(|| async {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
transport.polls.borrow_mut().push(Scripted::Ok(
StatusClass::Completed,
RESP_200_MULTI.to_vec(),
));
let multi: serde_json::Value =
serde_json::from_slice(RESP_200_MULTI).unwrap_or_default();
let partitions = multi["resultSetMetaData"]["partitionInfo"]
.as_array()
.cloned()
.unwrap_or_default();
let partition_count = partitions.len();
transport.script_partition(1, unauthorized());
for (index, info) in partitions.iter().enumerate().skip(1) {
let rows = info["rowCount"].as_u64().unwrap_or(0);
let body = format!(
r#"{{"data":[{}]}}"#,
(0..rows)
.map(|_| r#"["2024-01-02","ENTITY","2.50"]"#)
.collect::<Vec<_>>()
.join(",")
);
transport.script_partition(
u32::try_from(index).unwrap_or(u32::MAX),
Scripted::Ok(StatusClass::Completed, body.into_bytes()),
);
}
let mut auth = FakeAuth::resigning();
let cx = Cx::for_testing();
let (outcome, stats) = run_statement_with_auth(
&cx,
&transport,
&mut auth,
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5),
)
.await;
assert!(matches!(outcome, SnowflakeOutcome::Ok(_)), "{outcome:?}");
assert_eq!(auth.resigns, 1);
assert_eq!(
stats.partitions_fetched as usize, partition_count,
"one extra fetch for the retry"
);
assert!(transport.orphan_cancels.borrow().is_empty());
});
}
fn fast_poll_plan(max_polls: u32) -> PollPlan {
PollPlan {
max_polls,
poll_interval: Duration::ZERO,
..PollPlan::default()
}
}
fn fixture_handle() -> StatementHandle {
StatementHandle::new("01b2c3d4-0000-0000-0000-000000000002")
}
#[test]
fn submit_panic_without_a_handle_does_not_attempt_cleanup() {
asupersync::test_utils::run_test(|| async {
let transport = FakeTransport::new(Scripted::Panicked("submit panic"));
let cx = Cx::current().unwrap_or_else(Cx::for_testing);
let (outcome, stats) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5),
)
.await;
let payload = match outcome {
SnowflakeOutcome::Panicked(payload) => payload,
other => {
assert!(
matches!(other, SnowflakeOutcome::Panicked(_)),
"expected the submit panic, got {other:?}"
);
return;
}
};
assert_eq!(payload.message(), "submit panic");
assert_eq!(stats, DriverStats::default());
assert!(transport.orphan_cancels.borrow().is_empty());
assert!(transport.cancels_after_local.borrow().is_empty());
assert!(!transport.orphan_cleanup_finished.get());
});
}
#[test]
fn poll_panic_awaits_cleanup_and_preserves_the_original_payload() {
asupersync::test_utils::run_test(|| async {
for cleanup in [
Scripted::Ok(StatusClass::Completed, Vec::new()),
Scripted::Err,
Scripted::Panicked("secondary cleanup panic"),
] {
let mut transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
transport.orphan_cancel_result = cleanup;
transport.yield_once.set(true);
transport
.polls
.borrow_mut()
.push(Scripted::Panicked("poll panic"));
let cx = Cx::current().unwrap_or_else(Cx::for_testing);
let (outcome, stats) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5),
)
.await;
let payload = match outcome {
SnowflakeOutcome::Panicked(payload) => payload,
other => {
assert!(
matches!(other, SnowflakeOutcome::Panicked(_)),
"expected the original poll panic, got {other:?}"
);
continue;
}
};
assert_eq!(payload.message(), "poll panic");
assert_eq!(stats.polls, 1);
assert_eq!(
transport.orphan_cancels.borrow().as_slice(),
&[fixture_handle()]
);
assert_eq!(
transport.orphan_cancel_auth.borrow().as_slice(),
&["cred_test"]
);
assert!(transport.orphan_cleanup_finished.get());
assert!(transport.cancels_after_local.borrow().is_empty());
}
});
}
#[test]
fn partition_panic_drains_the_window_and_yielding_cleanup_before_returning() {
let transport = windowed_transport(1, &[1, 1, 1]);
transport.yield_once.set(true);
transport
.partitions
.borrow_mut()
.insert(2, VecDeque::from([Scripted::Panicked("partition panic")]));
let cx = Cx::for_testing();
let mut driver = std::pin::pin!(run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5).with_partition_concurrency(3),
));
let mut context = Context::from_waker(std::task::Waker::noop());
assert!(driver.as_mut().poll(&mut context).is_pending());
assert_eq!(transport.events(), vec!["start1", "start2", "start3"]);
assert!(transport.orphan_cancels.borrow().is_empty());
assert!(driver.as_mut().poll(&mut context).is_pending());
assert_eq!(
transport.events(),
vec!["start1", "start2", "start3", "done1", "done2", "done3"]
);
assert_eq!(
transport.orphan_cancels.borrow().as_slice(),
&[fixture_handle()]
);
assert!(!transport.orphan_cleanup_finished.get());
let poll_result = driver.as_mut().poll(&mut context);
assert!(
matches!(poll_result, Poll::Ready(_)),
"driver did not return after cleanup completed"
);
let (outcome, stats) = match poll_result {
Poll::Ready(ready) => ready,
Poll::Pending => return,
};
let payload = match outcome {
SnowflakeOutcome::Panicked(payload) => payload,
other => {
assert!(
matches!(other, SnowflakeOutcome::Panicked(_)),
"expected the partition panic, got {other:?}"
);
return;
}
};
assert_eq!(payload.message(), "partition panic");
assert_eq!(stats.partitions_fetched, 3);
assert_eq!(transport.orphan_cancels.borrow().len(), 1);
assert!(transport.orphan_cleanup_finished.get());
assert!(transport.cancels_after_local.borrow().is_empty());
}
#[test]
fn partition_retry_panic_cleans_up_with_the_refreshed_credential() {
asupersync::test_utils::run_test(|| async {
let transport = FakeTransport::new(Scripted::Ok(
StatusClass::Completed,
multi_partition_body(1, &[1]),
));
transport.yield_once.set(true);
transport.script_partition(1, unauthorized());
transport.script_partition(1, Scripted::Panicked("retry panic"));
let mut auth = FakeAuth::resigning();
let cx = Cx::current().unwrap_or_else(Cx::for_testing);
let (outcome, stats) = run_statement_with_auth(
&cx,
&transport,
&mut auth,
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5),
)
.await;
let payload = match outcome {
SnowflakeOutcome::Panicked(payload) => payload,
other => {
assert!(
matches!(other, SnowflakeOutcome::Panicked(_)),
"expected the retry panic, got {other:?}"
);
return;
}
};
assert_eq!(payload.message(), "retry panic");
assert_eq!(stats.partitions_fetched, 2);
assert_eq!(auth.resigns, 1);
assert_eq!(
transport.orphan_cancels.borrow().as_slice(),
&[fixture_handle()]
);
assert_eq!(
transport.orphan_cancel_auth.borrow().as_slice(),
&["cred_gen1"]
);
assert!(transport.orphan_cleanup_finished.get());
assert!(transport.cancels_after_local.borrow().is_empty());
});
}
#[test]
fn poll_transport_error_after_submit_fires_an_orphan_cancel() {
asupersync::test_utils::run_test(|| async {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
transport.polls.borrow_mut().push(Scripted::Err);
let cx = Cx::for_testing();
let (outcome, stats) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5),
)
.await;
assert!(matches!(outcome, SnowflakeOutcome::Err(_)));
assert_eq!(stats.polls, 1);
assert_eq!(
transport.orphan_cancels.borrow().as_slice(),
&[fixture_handle()],
"a transport error after the handle exists must cancel the orphaned statement"
);
assert!(transport.cancels_after_local.borrow().is_empty());
});
}
#[test]
fn undecodable_poll_body_after_submit_fires_an_orphan_cancel() {
asupersync::test_utils::run_test(|| async {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
transport
.polls
.borrow_mut()
.push(Scripted::Ok(StatusClass::Completed, b"not json".to_vec()));
let cx = Cx::for_testing();
let (outcome, _) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5),
)
.await;
assert!(matches!(outcome, SnowflakeOutcome::Err(_)));
assert_eq!(transport.orphan_cancels.borrow().len(), 1);
});
}
struct CollectingSink {
batches: Vec<Vec<Vec<Option<String>>>>,
refuse: bool,
}
impl RowSink for CollectingSink {
fn accept(
&mut self,
_result_set: &ResultSet,
rows: Vec<Vec<Option<String>>>,
) -> Result<(), SnowflakeError> {
if self.refuse {
return Err(SnowflakeError::new(
SnowflakeErrorCode::UsageError,
"the sink refused the rows",
));
}
self.batches.push(rows);
Ok(())
}
}
fn streaming_transport() -> FakeTransport {
let transport = FakeTransport::new(Scripted::Ok(
StatusClass::Completed,
RESP_200_MULTI.to_vec(),
));
transport.script_partition(
1,
Scripted::Ok(
StatusClass::Completed,
br#"{"data":[["p1a","x"],["p1b","x"]]}"#.to_vec(),
),
);
transport.script_partition(
2,
Scripted::Ok(
StatusClass::Completed,
br#"{"data":[["p2a","x"]]}"#.to_vec(),
),
);
transport
}
#[test]
fn streaming_hands_rows_to_the_sink_one_window_at_a_time() {
asupersync::test_utils::run_test(|| async {
let transport = streaming_transport();
let mut sink = CollectingSink {
batches: Vec::new(),
refuse: false,
};
let mut auth = fake_auth();
let cx = Cx::for_testing();
let (outcome, _) = run_statement_streaming(
&cx,
&transport,
&mut auth,
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5).with_partition_concurrency(1),
&mut sink,
)
.await;
let SnowflakeOutcome::Ok(done) = outcome else {
panic!("streaming run failed: {outcome:?}");
};
assert!(done.rows.is_empty(), "every row went to the sink");
assert_eq!(done.fetched_partitions, 3);
let firsts: Vec<String> = sink
.batches
.iter()
.flatten()
.map(|row| row.first().cloned().flatten().unwrap_or_default())
.collect();
assert_eq!(firsts.len(), 5, "inline 2 + 2 + 1");
assert_eq!(&firsts[2..], ["p1a", "p1b", "p2a"]);
assert_eq!(
sink.batches.len(),
3,
"one batch per partition with window 1"
);
assert!(
sink.batches.iter().all(|batch| batch.len() <= 2),
"no batch holds more than one partition"
);
assert!(transport.orphan_cancels.borrow().is_empty());
});
}
#[derive(Default)]
struct CollectingObserver(Vec<DriverEvent>);
impl DriverObserver for CollectingObserver {
fn event(&mut self, event: DriverEvent) {
self.0.push(event);
}
}
#[test]
fn the_observer_sees_each_remote_cancel_and_its_answer() {
asupersync::test_utils::run_test(|| async {
let run = |transport: FakeTransport| async move {
let mut observer = CollectingObserver::default();
let mut auth = fake_auth();
let cx = Cx::for_testing();
let hooks = StatementHooks {
sink: None,
observer: Some(&mut observer),
};
let (outcome, _) = run_statement_hooked(
&cx,
&transport,
&mut auth,
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5),
hooks,
)
.await;
(outcome, observer.0)
};
let cancels = |events: &[DriverEvent]| {
events
.iter()
.filter_map(|event| match event {
DriverEvent::RemoteCancel {
statement_handle,
acknowledged,
detail,
} => Some((statement_handle.clone(), *acknowledged, detail.clone())),
_ => None,
})
.collect::<Vec<_>>()
};
let acknowledged =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
acknowledged.polls.borrow_mut().push(Scripted::Err);
let (outcome, events) = run(acknowledged).await;
assert!(matches!(outcome, SnowflakeOutcome::Err(_)), "{outcome:?}");
assert_eq!(
cancels(&events),
vec![(
fixture_handle().as_str().to_owned(),
true,
"completed".to_owned()
)],
"{events:?}"
);
let mut refused =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
refused.polls.borrow_mut().push(Scripted::Err);
refused.orphan_cancel_result = Scripted::Err;
let (_, events) = run(refused).await;
let recorded = cancels(&events);
assert_eq!(recorded.len(), 1, "{events:?}");
assert!(!recorded[0].1, "{recorded:?}");
let (outcome, events) = run(streaming_transport()).await;
assert!(matches!(outcome, SnowflakeOutcome::Ok(_)), "{outcome:?}");
assert!(cancels(&events).is_empty(), "{events:?}");
});
}
#[test]
fn the_observer_sees_submit_partitions_and_completion() {
asupersync::test_utils::run_test(|| async {
let transport = streaming_transport();
let mut observer = CollectingObserver::default();
let mut auth = fake_auth();
let cx = Cx::for_testing();
let hooks = StatementHooks {
sink: None,
observer: Some(&mut observer),
};
let (outcome, _) = run_statement_hooked(
&cx,
&transport,
&mut auth,
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5).with_partition_concurrency(1),
hooks,
)
.await;
assert!(matches!(outcome, SnowflakeOutcome::Ok(_)), "{outcome:?}");
let events = observer.0;
assert!(
matches!(
&events[0],
DriverEvent::Submitted {
running: false,
statement_handle: Some(_)
}
),
"{events:?}"
);
assert_eq!(
events[1],
DriverEvent::PartitionFetched {
index: 1,
rows: 2,
bytes: u64::try_from(br#"{"data":[["p1a","x"],["p1b","x"]]}"#.len())
.unwrap_or(0),
}
);
assert!(
matches!(
events[2],
DriverEvent::PartitionFetched {
index: 2,
rows: 1,
..
}
),
"{events:?}"
);
assert_eq!(
events[3],
DriverEvent::Completed {
rows: 5,
partitions: 3
}
);
assert_eq!(events.len(), 4, "{events:?}");
});
}
#[test]
fn a_failing_sink_cancels_the_statement() {
asupersync::test_utils::run_test(|| async {
let transport = streaming_transport();
let mut sink = CollectingSink {
batches: Vec::new(),
refuse: true,
};
let mut auth = fake_auth();
let cx = Cx::for_testing();
let (outcome, stats) = run_statement_streaming(
&cx,
&transport,
&mut auth,
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5).with_partition_concurrency(1),
&mut sink,
)
.await;
assert!(matches!(outcome, SnowflakeOutcome::Err(_)), "{outcome:?}");
assert_eq!(stats.partitions_fetched, 0, "stopped before any fetch");
assert_eq!(transport.orphan_cancels.borrow().len(), 1);
});
}
#[test]
fn partition_fetch_error_fires_an_orphan_cancel() {
asupersync::test_utils::run_test(|| async {
let transport = FakeTransport::new(Scripted::Ok(
StatusClass::Completed,
RESP_200_MULTI.to_vec(),
));
transport.script_partition(1, Scripted::Err);
let cx = Cx::for_testing();
let (outcome, stats) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5).with_partition_concurrency(1),
)
.await;
assert!(matches!(outcome, SnowflakeOutcome::Err(_)));
assert_eq!(stats.partitions_fetched, 1);
assert_eq!(transport.orphan_cancels.borrow().len(), 1);
});
}
#[test]
fn deadline_during_poll_routes_through_the_policy_cancel_with_deadline_kind() {
asupersync::test_utils::run_test(|| async {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
for _ in 0..10 {
transport
.polls
.borrow_mut()
.push(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
}
let cx = Cx::for_testing_with_budget(Budget::new().with_deadline(Time::from_millis(1)));
let (outcome, _) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
PollPlan {
max_polls: 50,
poll_interval: Duration::from_millis(5),
..PollPlan::default()
},
)
.await;
assert!(matches!(outcome, SnowflakeOutcome::Cancelled(_)));
let cancels = transport.cancels_after_local.borrow();
assert_eq!(cancels.len(), 1);
assert_eq!(cancels[0].0, fixture_handle());
assert_eq!(cancels[0].1, CancelKind::Deadline);
assert!(transport.orphan_cancels.borrow().is_empty());
});
}
#[test]
fn happy_path_reports_polls_and_partitions_without_any_cancel() {
asupersync::test_utils::run_test(|| async {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
transport
.polls
.borrow_mut()
.push(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
transport.polls.borrow_mut().push(Scripted::Ok(
StatusClass::Completed,
RESP_200_MULTI.to_vec(),
));
let multi: serde_json::Value =
serde_json::from_slice(RESP_200_MULTI).unwrap_or_default();
let partitions = multi["resultSetMetaData"]["partitionInfo"]
.as_array()
.cloned()
.unwrap_or_default();
let partition_count = partitions.len();
for (index, info) in partitions.iter().enumerate().skip(1) {
let rows = info["rowCount"].as_u64().unwrap_or(0);
let body = format!(
r#"{{"data":[{}]}}"#,
(0..rows)
.map(|_| r#"["2024-01-02","ENTITY","2.50"]"#)
.collect::<Vec<_>>()
.join(",")
);
transport.script_partition(
u32::try_from(index).unwrap_or(u32::MAX),
Scripted::Ok(StatusClass::Completed, body.into_bytes()),
);
}
let cx = Cx::for_testing();
let (outcome, stats) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5),
)
.await;
assert!(matches!(outcome, SnowflakeOutcome::Ok(_)), "{outcome:?}");
assert_eq!(stats.polls, 2);
assert_eq!(stats.partitions_fetched as usize, partition_count - 1);
assert!(transport.orphan_cancels.borrow().is_empty());
assert!(transport.cancels_after_local.borrow().is_empty());
});
}
#[test]
fn response_class_maps_each_transport_status() {
assert_eq!(
response_class(StatusClass::Completed),
ResponseClass::Completed
);
assert_eq!(response_class(StatusClass::Running), ResponseClass::Running);
assert_eq!(
response_class(StatusClass::StatementTimeout),
ResponseClass::StatementTimeout
);
assert_eq!(
response_class(StatusClass::QueryFailure),
ResponseClass::StatementFailed
);
assert_eq!(
response_class(StatusClass::RateLimited),
ResponseClass::RateLimited
);
}
#[test]
fn submit_route_requires_request_id_and_retry_for_resubmit() {
let plain = SubmitQueryParams::default();
assert!(matches!(submit_route(&plain), TransportRoute::Submit));
let resubmit = SubmitQueryParams {
request_id: Some("req-1".to_owned()),
retry: true,
..SubmitQueryParams::default()
};
assert!(submit_route(&resubmit).has_retry_contract());
let no_id = SubmitQueryParams {
retry: true,
..SubmitQueryParams::default()
};
assert!(!submit_route(&no_id).has_retry_contract());
}
#[test]
fn submit_route_golden_preserves_async_and_nullable_query_params() {
let params = SubmitQueryParams {
request_id: Some("req-async-nullable".to_owned()),
retry: true,
asynchronous: true,
nullable: Some(false),
};
let expected_pairs = params.to_query_pairs();
let route = submit_route(¶ms);
assert!(matches!(
&route,
TransportRoute::SubmitWithQuery { query } if query == &expected_pairs
));
assert!(route.has_retry_contract());
assert_eq!(
route.path_and_query(),
"/api/v2/statements?requestId=req-async-nullable&retry=true&async=true&nullable=false"
);
}
#[test]
fn wait_poll_interval_preserves_deadline_attribution() {
asupersync::test_utils::run_test(|| async {
let cx = Cx::for_testing_with_budget(Budget::new().with_deadline(Time::from_millis(1)));
let reason = wait_poll_interval(&cx, Duration::from_millis(10))
.await
.expect_err("deadline should expire during poll wait");
assert_eq!(reason.kind, CancelKind::Deadline);
});
}
#[test]
fn terminal_statement_failures_keep_precise_error_projection() {
let timeout = QueryFailureStatus {
code: "000630".to_owned(),
sql_state: Some("57014".to_owned()),
message: "Statement reached its statement timeout and was canceled.".to_owned(),
statement_handle: Some(StatementHandle::new("timeout-handle")),
};
let timeout_error = terminal_failure_error(SnowflakeErrorCode::StatementTimeout, timeout);
let timeout_outcome: StatementOutcome = SnowflakeOutcome::err(timeout_error.clone());
assert_eq!(timeout_error.code, SnowflakeErrorCode::StatementTimeout);
assert_eq!(timeout_outcome.outcome_kind(), OutcomeKind::Timeout);
let failure = QueryFailureStatus {
code: "001003".to_owned(),
sql_state: Some("42000".to_owned()),
message: "SQL compilation error.".to_owned(),
statement_handle: Some(StatementHandle::new("failed-handle")),
};
let failure_error = terminal_failure_error(SnowflakeErrorCode::StatementFailed, failure);
let failure_outcome: StatementOutcome = SnowflakeOutcome::err(failure_error.clone());
assert_eq!(failure_error.code, SnowflakeErrorCode::StatementFailed);
assert_eq!(failure_outcome.outcome_kind(), OutcomeKind::Error);
}
#[test]
fn terminal_statement_failures_redact_secret_shaped_upstream_messages() {
let raw_token = "sfpat_driverFailureEcho001";
let failure = QueryFailureStatus {
code: "001003".to_owned(),
sql_state: Some("42000".to_owned()),
message: format!("SQL compilation error near literal '{raw_token}'"),
statement_handle: Some(StatementHandle::new("failed-handle")),
};
let error = terminal_failure_error(SnowflakeErrorCode::StatementFailed, failure);
assert_eq!(error.code, SnowflakeErrorCode::StatementFailed);
assert!(error.message.contains("[REDACTED]"));
assert!(!error.message.contains(raw_token));
}
}