use std::future::Future;
use std::sync::{Mutex, PoisonError};
use std::time::Duration;
use std::pin::Pin;
use std::task::{Context, Poll};
use asupersync::Cx;
use asupersync::record::{ObligationAbortReason, ObligationKind};
use asupersync::runtime::obligation_mailbox::ObligationToken;
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_core::sql_lexer;
use franken_snowflake_http::{
AuthorizationDescriptor, CancelHttpResponse, PartitionBody, PartitionHttpRequest,
PollHttpRequest, PollHttpResponse, RawHttp, SnowflakeHttpClient, StatusClass,
SubmitHttpRequest, SubmitHttpResponse, TransportOutcome, TransportRoute,
run_with_cancellation_mask,
};
use crate::lifecycle::{
CompletedStatement, CostQuota, 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>>;
fn cancel_on_drop(&self, auth: AuthorizationDescriptor, statement_handle: StatementHandle) {
let _ = (auth, statement_handle);
}
}
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
}
fn cancel_on_drop(&self, auth: AuthorizationDescriptor, statement_handle: StatementHandle) {
let config = self.config().clone();
spawn_detached_cancel(move || async move {
let (Some(cx), Ok(client)) = (Cx::current(), SnowflakeHttpClient::for_runtime(config))
else {
return;
};
let _ = client
.cancel_orphaned_statement(&cx, auth, statement_handle)
.await;
});
}
}
static DROPPED_CANCELS: std::sync::Mutex<Vec<std::thread::JoinHandle<()>>> =
std::sync::Mutex::new(Vec::new());
fn spawn_detached_cancel<F, Fut>(cancel: F)
where
F: FnOnce() -> Fut + Send + 'static,
Fut: Future<Output = ()>,
{
let spawned = std::thread::Builder::new()
.name("fsnow-drop-cancel".to_owned())
.spawn(move || {
if let Ok(runtime) = asupersync::runtime::RuntimeBuilder::current_thread().build() {
runtime.block_on(cancel());
}
});
if let (Ok(thread), Ok(mut pending)) = (spawned, DROPPED_CANCELS.lock()) {
pending.retain(|thread| !thread.is_finished());
pending.push(thread);
}
}
pub fn wait_for_dropped_cancels(timeout: Duration) -> usize {
let deadline = std::time::Instant::now() + timeout;
loop {
let running = DROPPED_CANCELS.lock().map_or(0, |mut pending| {
pending.retain(|thread| !thread.is_finished());
pending.len()
});
if running == 0 || std::time::Instant::now() >= deadline {
return running;
}
std::thread::sleep(Duration::from_millis(20));
}
}
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 poll_quota: u32,
pub execution_timeout: Option<Duration>,
pub cost_quota: Option<CostQuota>,
pub execution: Option<Duration>,
}
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: Send {
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: Send {
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: Mutex<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()),
};
let mut cancels = self.cancels.lock().unwrap_or_else(PoisonError::into_inner);
cancels.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
}
fn cancel_on_drop(&self, auth: AuthorizationDescriptor, statement_handle: StatementHandle) {
self.inner.cancel_on_drop(auth, statement_handle);
}
}
#[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: Mutex::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;
let cancels = recorder
.cancels
.into_inner()
.unwrap_or_else(PoisonError::into_inner);
for event in cancels {
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 answer_early = !params.asynchronous
&& params.retry
&& params.request_id.is_some()
&& sql_lexer::is_side_effect_free_read(&request.statement);
let mut learn_handle = false;
let submit_response = loop {
let route = if learn_handle {
submit_route(&SubmitQueryParams {
asynchronous: true,
..params.clone()
})
} else {
submit_route(params)
};
let submit = SubmitHttpRequest {
route,
auth: auth.clone(),
body: body.clone(),
retry_resubmit: params.retry,
};
if !learn_handle && cx.checkpoint().is_err() {
return SnowflakeOutcome::cancelled(local_cancel_reason(cx));
}
let exchange = run_with_cancellation_mask(cx, client.submit_statement(cx, submit));
let answer = if answer_early && !learn_handle {
let Some(answer) = answered_before_cancel(cx, exchange).await else {
learn_handle = true;
continue;
};
answer
} else {
exchange.await
};
match answer {
SnowflakeOutcome::Ok(response) if response.status == StatusClass::Unauthorized => {
match refresh_after_unauthorized(provider, reauth_left, "submit") {
Ok(fresh) => *auth = fresh,
Err(error) => {
return failed_submit_outcome(
cx,
learn_handle,
SnowflakeOutcome::err(error),
);
}
}
}
SnowflakeOutcome::Ok(response) => break response,
SnowflakeOutcome::Err(error) => {
return failed_submit_outcome(cx, learn_handle, SnowflakeOutcome::err(error));
}
SnowflakeOutcome::Cancelled(reason) => {
return failed_submit_outcome(
cx,
learn_handle,
SnowflakeOutcome::cancelled(reason),
);
}
SnowflakeOutcome::Panicked(payload) => {
return failed_submit_outcome(
cx,
learn_handle,
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) => failed_submit_outcome(
cx,
learn_handle,
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 mut guard = DropGuard::new(client);
let outcome = drive_guarded(
cx, client, provider, start, poll_plan, stats, hooks, &mut guard,
)
.await;
guard.settle(&outcome);
outcome
}
#[allow(clippy::too_many_arguments)]
async fn drive_guarded<T: StatementTransport, A: AuthProvider>(
cx: &Cx,
client: &T,
provider: &mut A,
start: Start,
poll_plan: PollPlan,
stats: &mut DriverStats,
hooks: StatementHooks<'_>,
guard: &mut DropGuard<'_, T>,
) -> 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();
stats.poll_quota = poll_plan.max_polls;
stats.execution_timeout = poll_plan.execution_timeout;
stats.cost_quota = poll_plan.cost_quota;
let started = asupersync::time::wall_now();
let bounds = ExecutionBounds {
deadline: poll_plan.execution_timeout.map(|timeout| started + timeout),
cost: poll_plan
.cost_quota
.and_then(|quota| quota.time_bound())
.map(|bound| started + bound),
};
let cost_quota = poll_plan.cost_quota;
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 } => {
if let Some(quota) = cost_quota.filter(CostQuota::below_resume_minimum) {
return SnowflakeOutcome::err(SnowflakeError::new(
SnowflakeErrorCode::SafetyLimitExceeded,
format!(
"not submitted: resuming the suspended warehouse is billed a minute at \
least (~{} millionths of a credit), over the {}-millionth credit cap",
quota.estimate(Duration::ZERO),
quota.max_microcredits
),
));
}
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
}
};
guard.arm(cx, &auth, &progress);
let mut executing = true;
loop {
if executing {
stats.execution = Some(elapsed_since(started));
executing = matches!(progress, Progress::PollAgain(_));
}
match progress {
Progress::Complete(mut completed) => {
guard.finish();
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) => {
guard.finish();
return SnowflakeOutcome::err(terminal_failure_error(
SnowflakeErrorCode::StatementTimeout,
failure,
));
}
Progress::Failed(failure) => {
guard.finish();
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 let Some(reason) = bounds.passed() {
return cancel_locally(cx, client, &auth, &handle, reason).await;
}
if !std::mem::take(&mut poll_now)
&& let Err(reason) = wait_poll_interval(cx, bounds.cap(poll_interval)).await
{
return cancel_locally(cx, client, &auth, &handle, reason).await;
}
stats.execution = Some(elapsed_since(started));
if let Some(reason) = bounds.passed() {
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;
}
};
guard.track(&auth);
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;
guard.track(&auth);
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) => {
guard.finish();
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;
}
};
guard.track(&auth);
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;
}
};
guard.track(&auth);
}
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;
}
};
}
}
}
}
struct DropGuard<'t, T: StatementTransport> {
transport: &'t T,
armed: Option<(AuthorizationDescriptor, StatementHandle)>,
lease: Option<ObligationToken>,
}
impl<'t, T: StatementTransport> DropGuard<'t, T> {
const fn new(transport: &'t T) -> Self {
Self {
transport,
armed: None,
lease: None,
}
}
fn arm(&mut self, cx: &Cx, auth: &AuthorizationDescriptor, progress: &Progress) {
self.armed = match progress {
Progress::PollAgain(handle) | Progress::FetchPartition { handle, .. } => {
Some((auth.clone(), handle.clone()))
}
Progress::Complete(_) | Progress::TimedOut(_) | Progress::Failed(_) => None,
};
if self.armed.is_some() {
self.lease = cx
.try_register_obligation_checked(ObligationKind::Lease, cx.task_id())
.ok()
.flatten();
}
}
fn track(&mut self, auth: &AuthorizationDescriptor) {
if let Some((held, _)) = self.armed.as_mut() {
held.clone_from(auth);
}
}
fn finish(&mut self) {
self.armed = None;
if let Some(lease) = self.lease.take() {
let _ = lease.commit();
}
}
fn settle(&mut self, outcome: &StatementOutcome) {
self.armed = None;
if let Some(lease) = self.lease.take() {
let _ = match outcome {
SnowflakeOutcome::Ok(_) => lease.commit(),
SnowflakeOutcome::Cancelled(_) => lease.abort(ObligationAbortReason::Cancel),
SnowflakeOutcome::Err(_) | SnowflakeOutcome::Panicked(_) => {
lease.abort(ObligationAbortReason::Error)
}
};
}
}
}
impl<T: StatementTransport> Drop for DropGuard<'_, T> {
fn drop(&mut self) {
if let Some(lease) = self.lease.take() {
let _ = lease.abort(ObligationAbortReason::Cancel);
}
if let Some((auth, handle)) = self.armed.take() {
self.transport.cancel_on_drop(auth, handle);
}
}
}
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<_> = partitions
.map(|partition| {
let request = PartitionHttpRequest {
auth: auth.clone(),
statement_handle: handle.clone(),
partition,
};
Some(Box::pin(client.fetch_partition(cx, request)))
})
.collect();
let done = pending.iter().map(|_| None).collect();
JoinInOrder { pending, done }.await
}
struct JoinInOrder<F: Future> {
pending: Vec<Option<Pin<Box<F>>>>,
done: Vec<Option<F::Output>>,
}
impl<F: Future> Future for JoinInOrder<F>
where
F::Output: Unpin,
{
type Output = Vec<F::Output>;
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)
}
async fn answered_before_cancel<T>(cx: &Cx, exchange: impl Future<Output = T>) -> Option<T> {
let mut exchange = std::pin::pin!(exchange);
std::future::poll_fn(|task| {
if let Poll::Ready(answer) = exchange.as_mut().poll(task) {
return Poll::Ready(Some(answer));
}
if cx.is_cancel_requested() {
Poll::Ready(None)
} else {
Poll::Pending
}
})
.await
}
fn local_cancel_reason(cx: &Cx) -> CancelReason {
cx.cancel_reason()
.unwrap_or_else(CancelReason::parent_cancelled)
}
fn failed_submit_outcome<T>(
cx: &Cx,
learn_handle: bool,
failure: SnowflakeOutcome<T>,
) -> SnowflakeOutcome<T> {
if learn_handle && cx.is_cancel_requested() {
SnowflakeOutcome::cancelled(local_cancel_reason(cx))
} else {
failure
}
}
fn terminal_failure_error(
code: SnowflakeErrorCode,
failure: crate::response::QueryFailureStatus,
) -> SnowflakeError {
SnowflakeError::new(code, redact(&failure.message).into_owned())
}
fn elapsed_since(started: asupersync::Time) -> Duration {
Duration::from_nanos(asupersync::time::wall_now().duration_since(started))
}
#[derive(Clone, Copy)]
struct ExecutionBounds {
deadline: Option<asupersync::Time>,
cost: Option<asupersync::Time>,
}
impl ExecutionBounds {
fn passed(&self) -> Option<CancelReason> {
let now = asupersync::time::wall_now();
if self.cost.is_some_and(|cost| now >= cost) {
Some(CancelReason::cost_budget())
} else if self.deadline.is_some_and(|deadline| now >= deadline) {
Some(CancelReason::deadline())
} else {
None
}
}
fn cap(&self, delay: Duration) -> Duration {
let now = asupersync::time::wall_now();
[self.deadline, self.cost]
.into_iter()
.flatten()
.map(|bound| Duration::from_nanos(bound.duration_since(now)))
.fold(delay, Duration::min)
}
}
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::lab::runtime::InvariantViolation;
use asupersync::lab::{LabConfig, LabRuntime};
use asupersync::trace::{TraceData, TraceEventKind};
use asupersync::{Budget, CancelKind, PanicPayload, Time};
use franken_snowflake_core::outcome::{OutcomeKind, SnowflakeOutcomeExt};
use franken_snowflake_http::{
CompressionEvidence, ContentEncoding, SnowflakeAuthTokenType, SnowflakeEndpoint,
TransportConfig, 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>,
dropped_cancels: RefCell<Vec<StatementHandle>>,
submit_queries: RefCell<Vec<Vec<(&'static str, String)>>>,
cancel_mid_sync_submit: Cell<bool>,
sync_submit_answered: 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),
dropped_cancels: RefCell::new(Vec::new()),
submit_queries: RefCell::new(Vec::new()),
cancel_mid_sync_submit: Cell::new(false),
sync_submit_answered: 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 query = match &request.route {
TransportRoute::SubmitWithQuery { query } => query.clone(),
_ => Vec::new(),
};
let asynchronous = query.iter().any(|(key, _)| *key == "async");
self.submit_queries.borrow_mut().push(query);
if self.cancel_mid_sync_submit.get() && !asynchronous {
let mut interrupted = false;
std::future::poll_fn(|task| {
if interrupted {
return Poll::Ready(());
}
interrupted = true;
cx.cancel_with(CancelKind::User, Some("the caller gave up"));
task.waker().wake_by_ref();
Poll::Pending
})
.await;
self.sync_submit_answered.set(true);
}
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 cancel_on_drop(
&self,
_auth: AuthorizationDescriptor,
statement_handle: StatementHandle,
) {
self.dropped_cancels.borrow_mut().push(statement_handle);
}
}
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.polls, stats.partitions_fetched), (0, 0));
assert_eq!(stats.poll_quota, 5);
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>);
#[test]
fn the_driver_future_is_send_over_the_production_client() {
fn assert_send<F: Future + Send>(_: &F) {}
let endpoint = SnowflakeEndpoint::parse("https://xy12345.us-east-1.snowflakecomputing.com")
.expect("a valid account endpoint");
let client = SnowflakeHttpClient::for_runtime(TransportConfig::new(endpoint))
.expect("the native-roots client builds");
let cx = Cx::for_testing();
let mut auth = fake_auth();
let mut sink = CollectingSink {
batches: Vec::new(),
refuse: false,
};
let mut observer = CollectingObserver::default();
let single = run_statement_hooked(
&cx,
&client,
&mut auth,
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
PollPlan::default(),
StatementHooks {
sink: Some(&mut sink),
observer: Some(&mut observer),
},
);
assert_send(&single);
drop(single);
let multi = run_multi_statement_hooked(
&cx,
&client,
&mut auth,
SubmitStatementRequest::new("select 1; select 2"),
SubmitQueryParams::default(),
PollPlan::default(),
Some(&mut observer),
);
assert_send(&multi);
}
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 the_execution_timeout_cancels_a_statement_that_keeps_running() {
asupersync::test_utils::run_test(|| async {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
for _ in 0..200 {
transport
.polls
.borrow_mut()
.push(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
}
let cx = Cx::for_testing();
let (outcome, stats) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
PollPlan {
max_polls: 200,
poll_interval: Duration::from_millis(5),
..PollPlan::default()
}
.with_execution_timeout(Some(Duration::from_millis(40))),
)
.await;
assert!(
matches!(&outcome, SnowflakeOutcome::Cancelled(reason) if reason.is_kind(CancelKind::Deadline)),
"{outcome:?}"
);
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!(stats.polls < 200, "{stats:?}");
assert_eq!(stats.poll_quota, 200);
assert_eq!(stats.execution_timeout, Some(Duration::from_millis(40)));
});
}
#[test]
fn dropping_the_driver_mid_poll_cancels_the_statement() {
asupersync::test_utils::run_test(|| async {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
let cx = Cx::for_testing();
{
let mut running = std::pin::pin!(run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
PollPlan::default(),
));
let first = std::future::poll_fn(|task| {
Poll::Ready(running.as_mut().poll(task).is_pending())
})
.await;
assert!(first, "the statement is still running");
assert!(transport.dropped_cancels.borrow().is_empty());
} assert_eq!(*transport.dropped_cancels.borrow(), vec![fixture_handle()]);
assert!(transport.cancels_after_local.borrow().is_empty());
});
}
#[test]
fn a_statement_driven_to_its_end_is_not_cancelled_on_drop() {
asupersync::test_utils::run_test(|| async {
let completed =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
completed.polls.borrow_mut().push(Scripted::Ok(
StatusClass::Completed,
RESP_200_SINGLE.to_vec(),
));
let failed = FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
failed.polls.borrow_mut().push(Scripted::Err);
let refused = FakeTransport::new(Scripted::Err);
for transport in [&completed, &failed, &refused] {
let _ = run_statement_with_stats(
&Cx::for_testing(),
transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
fast_poll_plan(5),
)
.await;
assert!(transport.dropped_cancels.borrow().is_empty());
}
});
}
#[derive(Debug)]
struct LabRun {
violations: Vec<InvariantViolation>,
reserved: usize,
resolved: Vec<(TraceEventKind, Option<ObligationAbortReason>)>,
quiescent: bool,
leaks: u64,
oracle_failures: Vec<String>,
}
fn in_lab_task<R: Send + 'static>(body: impl FnOnce(&Cx) -> R + Send + 'static) -> (R, LabRun) {
let mut lab = LabRuntime::new(LabConfig::new(0x0045).panic_on_leak(false).max_steps(1_000));
let root = lab.state.create_root_region(Budget::INFINITE);
let slot = std::sync::Arc::new(std::sync::Mutex::new(None));
let filled = std::sync::Arc::clone(&slot);
let (task, _handle) = lab
.state
.create_task(root, Budget::INFINITE, async move {
let cx = Cx::current().expect("a lab task has a runtime context");
let result = body(&cx);
*filled.lock().expect("the slot is not poisoned") = Some(result);
})
.expect("the lab admits the task");
lab.scheduler.lock().schedule(task, 0);
let report = lab.run_until_quiescent_with_report();
let quiescent = lab.is_quiescent();
let result = slot
.lock()
.expect("the slot is not poisoned")
.take()
.expect("the task ran to its end");
let violations = lab.check_invariants();
let mut reserved = 0;
let mut resolved = Vec::new();
for event in lab.trace().snapshot() {
if let TraceData::Obligation {
kind: ObligationKind::Lease,
abort_reason,
..
} = event.data
{
if event.kind == TraceEventKind::ObligationReserve {
reserved += 1;
} else {
resolved.push((event.kind, abort_reason));
}
}
}
let run = LabRun {
violations,
reserved,
resolved,
quiescent,
leaks: lab.state.leak_count(),
oracle_failures: report
.oracle_report
.failures()
.iter()
.map(|failure| format!("{}: {:?}", failure.invariant, failure.violation))
.collect(),
};
(result, run)
}
fn assert_resolved(run: &LabRun, resolved: &[(TraceEventKind, Option<ObligationAbortReason>)]) {
assert!(run.violations.is_empty(), "{run:?}");
assert!(run.oracle_failures.is_empty(), "{run:?}");
assert_eq!(run.leaks, 0, "{run:?}");
assert!(run.quiescent, "{run:?}");
assert_eq!(run.reserved, resolved.len(), "{run:?}");
assert_eq!(run.resolved, resolved, "{run:?}");
}
fn poll_once<F: Future>(future: Pin<&mut F>) -> Poll<F::Output> {
future.poll(&mut Context::from_waker(std::task::Waker::noop()))
}
struct CancelOnSubmit(Cx);
impl DriverObserver for CancelOnSubmit {
fn event(&mut self, event: DriverEvent) {
if matches!(event, DriverEvent::Submitted { running: true, .. }) {
self.0
.cancel_with(CancelKind::User, Some("the caller gave up"));
}
}
}
#[test]
fn every_statement_lease_is_resolved_under_the_lab_obligation_oracle() {
use ObligationAbortReason::{Cancel, Error};
use TraceEventKind::{ObligationAbort, ObligationCommit};
const FIRST: &str = "01b2c3d4-0000-0000-0000-0000000000a1";
const SECOND: &str = "01b2c3d4-0000-0000-0000-0000000000a2";
let (dropped, run) = in_lab_task(|cx| {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
{
let running = std::pin::pin!(run_statement_with_stats(
cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
PollPlan::default(),
));
assert!(poll_once(running).is_pending(), "the statement runs");
} transport.dropped_cancels.take()
});
assert_eq!(dropped, [fixture_handle()]);
assert_resolved(&run, &[(ObligationAbort, Some(Cancel))]);
let (polled, run) = in_lab_task(|cx| {
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 mut auth = fake_auth();
let running = std::pin::pin!(run_multi_statement_hooked(
cx,
&transport,
&mut auth,
SubmitStatementRequest::new("select 'first'; select 'second'"),
SubmitQueryParams::default(),
fast_poll_plan(5),
None,
));
let Poll::Ready((outcome, _)) = poll_once(running) else {
panic!("both statements were already complete");
};
assert!(matches!(outcome, SnowflakeOutcome::Ok(_)), "{outcome:?}");
assert!(transport.dropped_cancels.borrow().is_empty());
transport.polled.take()
});
assert_eq!(polled, [FIRST, SECOND]);
assert_resolved(&run, &[(ObligationCommit, None), (ObligationCommit, None)]);
let (abandoned, run) = in_lab_task(|cx| {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Completed, multi_parent(&[FIRST])));
*transport.polls.borrow_mut() = vec![Scripted::Err];
let mut auth = fake_auth();
let running = std::pin::pin!(run_multi_statement_hooked(
cx,
&transport,
&mut auth,
SubmitStatementRequest::new("select 'first'; select 'second'"),
SubmitQueryParams::default(),
fast_poll_plan(5),
None,
));
let Poll::Ready((outcome, _)) = poll_once(running) else {
panic!("the failed poll ends the request");
};
assert!(matches!(outcome, SnowflakeOutcome::Err(_)), "{outcome:?}");
transport.orphan_cancels.take()
});
assert_eq!(abandoned, [StatementHandle::new(FIRST)]);
assert_resolved(&run, &[(ObligationAbort, Some(Error))]);
let (cancelled, run) = in_lab_task(|cx| {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
let mut auth = fake_auth();
let mut observer = CancelOnSubmit(cx.clone());
let running = std::pin::pin!(run_statement_hooked(
cx,
&transport,
&mut auth,
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
PollPlan::default(),
StatementHooks {
sink: None,
observer: Some(&mut observer),
},
));
let Poll::Ready((outcome, _)) = poll_once(running) else {
panic!("a cancelled statement ends at its next checkpoint");
};
assert!(
matches!(outcome, SnowflakeOutcome::Cancelled(_)),
"{outcome:?}"
);
transport.cancels_after_local.take()
});
assert_eq!(cancelled, [(fixture_handle(), CancelKind::User)]);
assert_resolved(&run, &[(ObligationAbort, Some(Cancel))]);
}
#[test]
fn a_statement_whose_driver_is_leaked_is_reported_as_a_leaked_lease() {
let (dropped, run) = in_lab_task(|cx| {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
let mut running = Box::pin(run_statement_with_stats(
cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
PollPlan::default(),
));
assert!(
poll_once(running.as_mut()).is_pending(),
"the statement runs"
);
std::mem::forget(running);
transport.dropped_cancels.take()
});
assert!(dropped.is_empty(), "the guard never ran");
assert_eq!(run.reserved, 1, "{run:?}");
assert_eq!(
run.resolved,
[(TraceEventKind::ObligationLeak, None)],
"{run:?}"
);
assert_eq!(run.leaks, 1, "{run:?}");
assert_eq!(run.oracle_failures.len(), 1, "{run:?}");
}
#[test]
fn the_cli_runtime_flavor_tracks_statement_leases() {
fn drive(leak: bool) -> Vec<StatementHandle> {
let runtime = asupersync::runtime::RuntimeBuilder::current_thread()
.build()
.expect("the runtime starts");
runtime.block_on(async move {
let cx = Cx::current().expect("block_on installs a context");
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
let mut running = Box::pin(run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
PollPlan::default(),
));
assert!(
poll_once(running.as_mut()).is_pending(),
"the statement runs"
);
if leak {
std::mem::forget(running);
} else {
drop(running);
}
transport.dropped_cancels.take()
})
}
assert_eq!(drive(false), [fixture_handle()]);
let payload = std::panic::catch_unwind(|| drive(true))
.expect_err("the runtime refuses to retire a task holding a lease");
let message = payload
.downcast_ref::<String>()
.map(String::as_str)
.or_else(|| payload.downcast_ref::<&str>().copied())
.unwrap_or_default();
assert!(
message.starts_with("obligation leak:") && message.contains(" Lease holder="),
"{message}"
);
}
#[test]
fn a_credit_cap_cancels_before_the_deadline_with_the_cost_kind() {
asupersync::test_utils::run_test(|| async {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
for _ in 0..200 {
transport
.polls
.borrow_mut()
.push(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
}
let quota = CostQuota {
microcredits_per_hour: 3_600_000,
max_microcredits: 40,
resumes_warehouse: false,
};
let (outcome, stats) = run_statement_with_stats(
&Cx::for_testing(),
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
PollPlan {
max_polls: 200,
poll_interval: Duration::from_millis(5),
..PollPlan::default()
}
.with_execution_timeout(Some(Duration::from_secs(10)))
.with_cost_quota(Some(quota)),
)
.await;
assert!(
matches!(&outcome, SnowflakeOutcome::Cancelled(reason) if reason.is_kind(CancelKind::CostBudget)),
"{outcome:?}"
);
let cancels = transport.cancels_after_local.borrow();
assert_eq!(cancels.len(), 1);
assert_eq!(cancels[0].1, CancelKind::CostBudget);
assert_eq!(stats.cost_quota, Some(quota));
let execution = stats.execution.unwrap_or_default();
assert!(
execution >= Duration::from_millis(40) && execution < Duration::from_secs(10),
"{stats:?}"
);
});
}
#[test]
fn a_credit_cap_below_the_resume_minimum_submits_nothing() {
asupersync::test_utils::run_test(|| async {
let transport =
FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
let quota = CostQuota {
microcredits_per_hour: 1_000_000,
max_microcredits: 1_000,
resumes_warehouse: true,
};
let (outcome, _) = run_statement_with_stats(
&Cx::for_testing(),
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
SubmitQueryParams::default(),
PollPlan::default().with_cost_quota(Some(quota)),
)
.await;
assert!(
matches!(&outcome, SnowflakeOutcome::Err(error) if error.code == SnowflakeErrorCode::SafetyLimitExceeded),
"{outcome:?}"
);
assert!(transport.auth_seen.borrow().is_empty(), "nothing was sent");
assert!(transport.cancels_after_local.borrow().is_empty());
});
}
#[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 deadline = asupersync::time::wall_now() + Duration::from_millis(25);
let cx = Cx::for_testing_with_budget(Budget::new().with_deadline(deadline));
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 an_already_expired_deadline_submits_nothing() {
asupersync::test_utils::run_test(|| async {
let transport =
FakeTransport::new(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::default(),
)
.await;
assert!(
matches!(&outcome, SnowflakeOutcome::Cancelled(reason) if reason.is_kind(CancelKind::Deadline)),
"{outcome:?}"
);
assert!(
transport.auth_seen.borrow().is_empty(),
"no request was sent"
);
assert!(transport.cancels_after_local.borrow().is_empty());
assert!(transport.orphan_cancels.borrow().is_empty());
});
}
fn idempotent_params() -> SubmitQueryParams {
SubmitQueryParams {
request_id: Some("0b5e7a9c-0000-4000-8000-000000000b5e".to_owned()),
retry: true,
..SubmitQueryParams::default()
}
}
fn cancel_mid_submit(
sql: &str,
params: SubmitQueryParams,
) -> (StatementOutcome, FakeTransport) {
let transport = FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
transport.cancel_mid_sync_submit.set(true);
let mut slot = None;
let (fake, ended) = (&transport, &mut slot);
asupersync::test_utils::run_test(move || async move {
let cx = Cx::for_testing();
let (outcome, _) = run_statement_with_stats(
&cx,
fake,
fake_auth(),
SubmitStatementRequest::new(sql),
params,
PollPlan::default(),
)
.await;
*ended = Some(outcome);
});
let outcome = slot.expect("the statement ran to an outcome");
(outcome, transport)
}
#[test]
fn a_read_cancelled_mid_submit_learns_its_handle_at_once() {
let (outcome, transport) = cancel_mid_submit("select system$wait(60)", idempotent_params());
assert!(
matches!(outcome, SnowflakeOutcome::Cancelled(_)),
"{outcome:?}"
);
assert!(
!transport.sync_submit_answered.get(),
"the synchronous answer was not waited for"
);
let request_id = (
"requestId",
idempotent_params().request_id.unwrap_or_default(),
);
assert_eq!(
*transport.submit_queries.borrow(),
vec![
vec![request_id.clone(), ("retry", "true".to_owned())],
vec![
request_id,
("retry", "true".to_owned()),
("async", "true".to_owned())
],
]
);
let cancels = transport.cancels_after_local.borrow();
assert_eq!(cancels.len(), 1);
assert_eq!(cancels[0].0, fixture_handle());
assert!(transport.orphan_cancels.borrow().is_empty());
}
#[test]
fn failed_handle_recovery_preserves_the_callers_cancel_outcome() {
for scripted in [
Scripted::Err,
Scripted::Panicked("recovery failed"),
Scripted::Ok(StatusClass::Unauthorized, Vec::new()),
Scripted::Ok(StatusClass::Completed, b"{}".to_vec()),
] {
asupersync::test_utils::run_test(move || async move {
let transport = FakeTransport::new(scripted);
transport.cancel_mid_sync_submit.set(true);
let cx = Cx::for_testing();
let (outcome, _) = run_statement_with_stats(
&cx,
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
idempotent_params(),
PollPlan::default(),
)
.await;
assert!(
matches!(&outcome, SnowflakeOutcome::Cancelled(reason) if reason.kind == CancelKind::User),
"{outcome:?}"
);
assert!(!transport.sync_submit_answered.get());
let queries = transport.submit_queries.borrow();
assert_eq!(queries.len(), 2);
assert_eq!(queries[0][0], queries[1][0]);
assert!(queries[1].contains(&("retry", "true".to_owned())));
assert!(queries[1].contains(&("async", "true".to_owned())));
assert!(transport.cancels_after_local.borrow().is_empty());
assert!(transport.orphan_cancels.borrow().is_empty());
assert!(transport.dropped_cancels.borrow().is_empty());
assert!(transport.polled.borrow().is_empty());
});
}
}
#[test]
fn an_uncancelled_submit_preserves_its_failure_outcome() {
for (scripted, panicked) in [
(Scripted::Err, false),
(Scripted::Panicked("submit failed"), true),
(Scripted::Ok(StatusClass::Unauthorized, Vec::new()), false),
(Scripted::Ok(StatusClass::Completed, b"{}".to_vec()), false),
] {
asupersync::test_utils::run_test(move || async move {
let transport = FakeTransport::new(scripted);
let (outcome, _) = run_statement_with_stats(
&Cx::for_testing(),
&transport,
fake_auth(),
SubmitStatementRequest::new("select 1"),
idempotent_params(),
PollPlan::default(),
)
.await;
assert!(
if panicked {
matches!(outcome, SnowflakeOutcome::Panicked(_))
} else {
matches!(outcome, SnowflakeOutcome::Err(_))
},
"{outcome:?}"
);
assert_eq!(transport.submit_queries.borrow().len(), 1);
assert!(transport.cancels_after_local.borrow().is_empty());
assert!(transport.orphan_cancels.borrow().is_empty());
});
}
}
#[test]
fn failed_handle_lookup_keeps_the_original_cancel_reason() {
let cx = Cx::for_testing();
cx.cancel_with(CancelKind::User, Some("the caller gave up"));
let outcome = failed_submit_outcome::<()>(
&cx,
true,
SnowflakeOutcome::cancelled(CancelReason::deadline()),
);
assert!(
matches!(outcome, SnowflakeOutcome::Cancelled(reason) if reason.kind == CancelKind::User)
);
let outcome = failed_submit_outcome::<()>(
&cx,
false,
SnowflakeOutcome::cancelled(CancelReason::deadline()),
);
assert!(
matches!(outcome, SnowflakeOutcome::Cancelled(reason) if reason.kind == CancelKind::Deadline)
);
}
#[test]
fn a_write_or_a_non_idempotent_read_keeps_waiting_for_the_answer() {
for (sql, params) in [
("insert into t values (1)", idempotent_params()),
("select 1; delete from t", idempotent_params()),
("select 1", SubmitQueryParams::default()),
] {
let (outcome, transport) = cancel_mid_submit(sql, params);
assert!(
matches!(outcome, SnowflakeOutcome::Cancelled(_)),
"{sql}: {outcome:?}"
);
assert!(transport.sync_submit_answered.get(), "{sql}");
assert_eq!(transport.submit_queries.borrow().len(), 1, "{sql}");
let cancels = transport.cancels_after_local.borrow();
assert_eq!(cancels.len(), 1, "{sql}");
assert_eq!(cancels[0].0, fixture_handle(), "{sql}");
}
}
#[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));
}
}