1use std::cell::RefCell;
18use std::future::Future;
19use std::time::Duration;
20
21use std::pin::Pin;
22use std::task::{Context, Poll};
23
24use asupersync::Cx;
25use franken_snowflake_core::cancel::CancelReason;
26use franken_snowflake_core::error::{SnowflakeError, SnowflakeErrorCode};
27use franken_snowflake_core::ids::StatementHandle;
28use franken_snowflake_core::outcome::SnowflakeOutcome;
29use franken_snowflake_core::redact::redact;
30use franken_snowflake_http::{
31 AuthorizationDescriptor, CancelHttpResponse, PartitionBody, PartitionHttpRequest,
32 PollHttpRequest, PollHttpResponse, RawHttp, SnowflakeHttpClient, StatusClass,
33 SubmitHttpRequest, SubmitHttpResponse, TransportOutcome, TransportRoute,
34};
35
36use crate::lifecycle::{
37 CompletedStatement, MIN_POLL_INTERVAL, PollPlan, Progress, StatementMachine,
38};
39use crate::request::{SubmitQueryParams, SubmitStatementRequest};
40use crate::response::ResultSet;
41use crate::status::ResponseClass;
42
43pub type StatementOutcome = SnowflakeOutcome<CompletedStatement>;
46
47pub trait StatementTransport {
51 fn submit_statement(
53 &self,
54 cx: &Cx,
55 request: SubmitHttpRequest,
56 ) -> impl Future<Output = TransportOutcome<SubmitHttpResponse>>;
57 fn poll_statement(
59 &self,
60 cx: &Cx,
61 request: PollHttpRequest,
62 ) -> impl Future<Output = TransportOutcome<PollHttpResponse>>;
63 fn fetch_partition(
65 &self,
66 cx: &Cx,
67 request: PartitionHttpRequest,
68 ) -> impl Future<Output = TransportOutcome<PartitionBody>>;
69 fn cancel_after_local_cancel(
71 &self,
72 cx: &Cx,
73 auth: AuthorizationDescriptor,
74 statement_handle: StatementHandle,
75 reason: CancelReason,
76 ) -> impl Future<Output = TransportOutcome<CancelHttpResponse>>;
77 fn cancel_orphaned_statement(
79 &self,
80 cx: &Cx,
81 auth: AuthorizationDescriptor,
82 statement_handle: StatementHandle,
83 ) -> impl Future<Output = TransportOutcome<CancelHttpResponse>>;
84}
85
86impl<H: RawHttp> StatementTransport for SnowflakeHttpClient<H> {
87 async fn submit_statement(
88 &self,
89 cx: &Cx,
90 request: SubmitHttpRequest,
91 ) -> TransportOutcome<SubmitHttpResponse> {
92 Self::submit_statement(self, cx, request).await
93 }
94
95 async fn poll_statement(
96 &self,
97 cx: &Cx,
98 request: PollHttpRequest,
99 ) -> TransportOutcome<PollHttpResponse> {
100 Self::poll_statement(self, cx, request).await
101 }
102
103 async fn fetch_partition(
104 &self,
105 cx: &Cx,
106 request: PartitionHttpRequest,
107 ) -> TransportOutcome<PartitionBody> {
108 Self::fetch_partition(self, cx, request).await
109 }
110
111 async fn cancel_after_local_cancel(
112 &self,
113 cx: &Cx,
114 auth: AuthorizationDescriptor,
115 statement_handle: StatementHandle,
116 reason: CancelReason,
117 ) -> TransportOutcome<CancelHttpResponse> {
118 Self::cancel_after_local_cancel(self, cx, auth, statement_handle, reason).await
119 }
120
121 async fn cancel_orphaned_statement(
122 &self,
123 cx: &Cx,
124 auth: AuthorizationDescriptor,
125 statement_handle: StatementHandle,
126 ) -> TransportOutcome<CancelHttpResponse> {
127 Self::cancel_orphaned_statement(self, cx, auth, statement_handle).await
128 }
129}
130
131pub trait AuthProvider {
143 fn descriptor(&mut self) -> Result<AuthorizationDescriptor, SnowflakeError>;
145 fn on_unauthorized(&mut self) -> Result<bool, SnowflakeError>;
148}
149
150impl AuthProvider for AuthorizationDescriptor {
152 fn descriptor(&mut self) -> Result<AuthorizationDescriptor, SnowflakeError> {
153 Ok(self.clone())
154 }
155
156 fn on_unauthorized(&mut self) -> Result<bool, SnowflakeError> {
157 Ok(false)
158 }
159}
160
161#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
163pub struct DriverStats {
164 pub polls: u32,
166 pub partitions_fetched: u32,
168}
169
170pub async fn run_statement<T: StatementTransport>(
177 cx: &Cx,
178 client: &T,
179 auth: AuthorizationDescriptor,
180 request: SubmitStatementRequest,
181 params: SubmitQueryParams,
182 poll_plan: PollPlan,
183) -> StatementOutcome {
184 run_statement_with_stats(cx, client, auth, request, params, poll_plan)
185 .await
186 .0
187}
188
189pub async fn run_statement_with_stats<T: StatementTransport>(
191 cx: &Cx,
192 client: &T,
193 auth: AuthorizationDescriptor,
194 request: SubmitStatementRequest,
195 params: SubmitQueryParams,
196 poll_plan: PollPlan,
197) -> (StatementOutcome, DriverStats) {
198 let mut frozen = auth;
199 run_statement_with_auth(cx, client, &mut frozen, request, params, poll_plan).await
200}
201
202pub async fn run_statement_with_auth<T: StatementTransport, A: AuthProvider>(
206 cx: &Cx,
207 client: &T,
208 auth: &mut A,
209 request: SubmitStatementRequest,
210 params: SubmitQueryParams,
211 poll_plan: PollPlan,
212) -> (StatementOutcome, DriverStats) {
213 let mut stats = DriverStats::default();
214 let outcome = drive(
215 cx,
216 client,
217 auth,
218 Start::Submit { request, params },
219 poll_plan,
220 &mut stats,
221 StatementHooks::default(),
222 )
223 .await;
224 (outcome, stats)
225}
226
227pub trait RowSink {
230 fn accept(
237 &mut self,
238 result_set: &ResultSet,
239 rows: Vec<Vec<Option<String>>>,
240 ) -> Result<(), SnowflakeError>;
241}
242
243#[derive(Clone, Debug, PartialEq, Eq)]
246pub enum DriverEvent {
247 Submitted {
250 statement_handle: Option<String>,
252 running: bool,
254 },
255 Polled {
257 polls: u32,
259 },
260 PartitionFetched {
262 index: u32,
264 rows: u64,
266 bytes: u64,
268 },
269 Completed {
271 rows: i64,
273 partitions: u32,
275 },
276 RemoteCancel {
279 statement_handle: String,
281 acknowledged: bool,
283 detail: String,
285 },
286}
287
288pub trait DriverObserver {
290 fn event(&mut self, event: DriverEvent);
292}
293
294#[derive(Default)]
296pub struct StatementHooks<'a> {
297 pub sink: Option<&'a mut dyn RowSink>,
299 pub observer: Option<&'a mut dyn DriverObserver>,
301}
302
303pub async fn run_statement_streaming<T: StatementTransport, A: AuthProvider>(
308 cx: &Cx,
309 client: &T,
310 auth: &mut A,
311 request: SubmitStatementRequest,
312 params: SubmitQueryParams,
313 poll_plan: PollPlan,
314 sink: &mut dyn RowSink,
315) -> (StatementOutcome, DriverStats) {
316 let hooks = StatementHooks {
317 sink: Some(sink),
318 observer: None,
319 };
320 run_statement_hooked(cx, client, auth, request, params, poll_plan, hooks).await
321}
322
323pub async fn run_statement_hooked<T: StatementTransport, A: AuthProvider>(
325 cx: &Cx,
326 client: &T,
327 auth: &mut A,
328 request: SubmitStatementRequest,
329 params: SubmitQueryParams,
330 poll_plan: PollPlan,
331 hooks: StatementHooks<'_>,
332) -> (StatementOutcome, DriverStats) {
333 let mut stats = DriverStats::default();
334 let outcome = drive(
335 cx,
336 client,
337 auth,
338 Start::Submit { request, params },
339 poll_plan,
340 &mut stats,
341 hooks,
342 )
343 .await;
344 (outcome, stats)
345}
346
347#[derive(Clone, Debug, PartialEq)]
349pub struct MultiStatementResult {
350 pub parent: CompletedStatement,
353 pub statements: Vec<CompletedStatement>,
355}
356
357pub async fn run_multi_statement_hooked<T: StatementTransport, A: AuthProvider>(
364 cx: &Cx,
365 client: &T,
366 auth: &mut A,
367 request: SubmitStatementRequest,
368 params: SubmitQueryParams,
369 poll_plan: PollPlan,
370 mut observer: Option<&mut dyn DriverObserver>,
371) -> (SnowflakeOutcome<MultiStatementResult>, DriverStats) {
372 let mut stats = DriverStats::default();
373 let parent = match drive(
374 cx,
375 client,
376 auth,
377 Start::Submit { request, params },
378 poll_plan,
379 &mut stats,
380 StatementHooks {
381 sink: None,
382 observer: observer
383 .as_mut()
384 .map(|observer| &mut **observer as &mut dyn DriverObserver),
385 },
386 )
387 .await
388 {
389 SnowflakeOutcome::Ok(parent) => parent,
390 SnowflakeOutcome::Err(error) => return (SnowflakeOutcome::err(error), stats),
391 SnowflakeOutcome::Cancelled(reason) => return (SnowflakeOutcome::cancelled(reason), stats),
392 SnowflakeOutcome::Panicked(payload) => return (SnowflakeOutcome::panicked(payload), stats),
393 };
394 let handles = parent
395 .result_set
396 .statement_handles
397 .clone()
398 .unwrap_or_default();
399 if handles.is_empty() {
400 return (
401 SnowflakeOutcome::err(SnowflakeError::new(
402 SnowflakeErrorCode::UpstreamError,
403 "the SQL API answered a multi-statement request without statementHandles",
404 )),
405 stats,
406 );
407 }
408 let mut statements = Vec::with_capacity(handles.len());
409 for handle in handles {
410 match drive(
411 cx,
412 client,
413 auth,
414 Start::Handle(handle),
415 poll_plan,
416 &mut stats,
417 StatementHooks {
418 sink: None,
419 observer: observer
420 .as_mut()
421 .map(|observer| &mut **observer as &mut dyn DriverObserver),
422 },
423 )
424 .await
425 {
426 SnowflakeOutcome::Ok(done) => statements.push(done),
427 SnowflakeOutcome::Err(error) => return (SnowflakeOutcome::err(error), stats),
428 SnowflakeOutcome::Cancelled(reason) => {
429 return (SnowflakeOutcome::cancelled(reason), stats);
430 }
431 SnowflakeOutcome::Panicked(payload) => {
432 return (SnowflakeOutcome::panicked(payload), stats);
433 }
434 }
435 }
436 (
437 SnowflakeOutcome::ok(MultiStatementResult { parent, statements }),
438 stats,
439 )
440}
441
442enum Start {
444 Submit {
446 request: SubmitStatementRequest,
447 params: SubmitQueryParams,
448 },
449 Handle(StatementHandle),
452}
453
454fn notify(observer: &mut Option<&mut dyn DriverObserver>, event: DriverEvent) {
455 if let Some(observer) = observer.as_mut() {
456 observer.event(event);
457 }
458}
459
460fn progress_handle(progress: &Progress) -> Option<String> {
461 match progress {
462 Progress::PollAgain(handle) | Progress::FetchPartition { handle, .. } => {
463 Some(handle.as_str().to_owned())
464 }
465 Progress::Complete(done) => Some(done.statement_handle.as_str().to_owned()),
466 Progress::TimedOut(_) | Progress::Failed(_) => None,
467 }
468}
469
470fn unauthorized_error(step: &str, detail: &str) -> SnowflakeError {
472 SnowflakeError::new(
473 SnowflakeErrorCode::CredentialExpired,
474 format!("SQL API returned 401 Unauthorized on {step}: {detail}"),
475 )
476}
477
478fn refresh_after_unauthorized<A: AuthProvider>(
482 provider: &mut A,
483 reauth_left: &mut u8,
484 step: &str,
485) -> Result<AuthorizationDescriptor, SnowflakeError> {
486 if *reauth_left == 0 {
487 return Err(unauthorized_error(
488 step,
489 "the re-signed credential was rejected again; not retrying further",
490 ));
491 }
492 *reauth_left = reauth_left.saturating_sub(1);
493 match provider.on_unauthorized()? {
494 true => provider.descriptor(),
495 false => Err(unauthorized_error(
496 step,
497 "this credential lane cannot re-sign mid-flight; issue a fresh token and retry",
498 )),
499 }
500}
501
502struct CancelRecorder<'t, T> {
505 inner: &'t T,
506 cancels: RefCell<Vec<DriverEvent>>,
507}
508
509impl<T> CancelRecorder<'_, T> {
510 fn record(&self, handle: &StatementHandle, outcome: &TransportOutcome<CancelHttpResponse>) {
511 let (acknowledged, detail) = match outcome {
512 SnowflakeOutcome::Ok(response) => (
513 response.status == StatusClass::Completed,
514 status_class_label(response.status).to_owned(),
515 ),
516 SnowflakeOutcome::Err(error) => (false, redact(&error.message).into_owned()),
517 SnowflakeOutcome::Cancelled(reason) => (
518 false,
519 format!(
520 "the cancel request was itself cancelled ({:?})",
521 reason.kind
522 ),
523 ),
524 SnowflakeOutcome::Panicked(_) => (false, "the cancel request panicked".to_owned()),
525 };
526 self.cancels.borrow_mut().push(DriverEvent::RemoteCancel {
527 statement_handle: handle.as_str().to_owned(),
528 acknowledged,
529 detail,
530 });
531 }
532}
533
534const fn status_class_label(status: StatusClass) -> &'static str {
535 match status {
536 StatusClass::Completed => "completed",
537 StatusClass::Running => "running",
538 StatusClass::StatementTimeout => "statement_timeout",
539 StatusClass::QueryFailure => "query_failure",
540 StatusClass::RateLimited => "rate_limited",
541 StatusClass::ServerErrorRetryable => "server_error",
542 StatusClass::Unauthorized => "unauthorized",
543 StatusClass::Unexpected => "unexpected",
544 }
545}
546
547impl<T: StatementTransport> StatementTransport for CancelRecorder<'_, T> {
548 async fn submit_statement(
549 &self,
550 cx: &Cx,
551 request: SubmitHttpRequest,
552 ) -> TransportOutcome<SubmitHttpResponse> {
553 self.inner.submit_statement(cx, request).await
554 }
555
556 async fn poll_statement(
557 &self,
558 cx: &Cx,
559 request: PollHttpRequest,
560 ) -> TransportOutcome<PollHttpResponse> {
561 self.inner.poll_statement(cx, request).await
562 }
563
564 async fn fetch_partition(
565 &self,
566 cx: &Cx,
567 request: PartitionHttpRequest,
568 ) -> TransportOutcome<PartitionBody> {
569 self.inner.fetch_partition(cx, request).await
570 }
571
572 async fn cancel_after_local_cancel(
573 &self,
574 cx: &Cx,
575 auth: AuthorizationDescriptor,
576 statement_handle: StatementHandle,
577 reason: CancelReason,
578 ) -> TransportOutcome<CancelHttpResponse> {
579 let outcome = self
580 .inner
581 .cancel_after_local_cancel(cx, auth, statement_handle.clone(), reason)
582 .await;
583 self.record(&statement_handle, &outcome);
584 outcome
585 }
586
587 async fn cancel_orphaned_statement(
588 &self,
589 cx: &Cx,
590 auth: AuthorizationDescriptor,
591 statement_handle: StatementHandle,
592 ) -> TransportOutcome<CancelHttpResponse> {
593 let outcome = self
594 .inner
595 .cancel_orphaned_statement(cx, auth, statement_handle.clone())
596 .await;
597 self.record(&statement_handle, &outcome);
598 outcome
599 }
600}
601
602#[allow(clippy::too_many_arguments)]
606async fn drive<T: StatementTransport, A: AuthProvider>(
607 cx: &Cx,
608 client: &T,
609 provider: &mut A,
610 start: Start,
611 poll_plan: PollPlan,
612 stats: &mut DriverStats,
613 hooks: StatementHooks<'_>,
614) -> StatementOutcome {
615 let recorder = CancelRecorder {
616 inner: client,
617 cancels: RefCell::new(Vec::new()),
618 };
619 let StatementHooks { sink, mut observer } = hooks;
620 let outcome = drive_statement(
621 cx,
622 &recorder,
623 provider,
624 start,
625 poll_plan,
626 stats,
627 StatementHooks {
630 sink: sink.map(|sink| sink as &mut dyn RowSink),
631 observer: observer
632 .as_mut()
633 .map(|observer| &mut **observer as &mut dyn DriverObserver),
634 },
635 )
636 .await;
637 for event in recorder.cancels.take() {
638 notify(&mut observer, event);
639 }
640 outcome
641}
642
643#[allow(clippy::too_many_arguments)]
646async fn submit<T: StatementTransport, A: AuthProvider>(
647 cx: &Cx,
648 client: &T,
649 provider: &mut A,
650 auth: &mut AuthorizationDescriptor,
651 reauth_left: &mut u8,
652 request: &SubmitStatementRequest,
653 params: &SubmitQueryParams,
654 machine: &mut StatementMachine,
655) -> SnowflakeOutcome<Progress> {
656 let body = match serde_json::to_vec(request) {
657 Ok(body) => body,
658 Err(error) => {
659 return SnowflakeOutcome::err(SnowflakeError::new(
660 SnowflakeErrorCode::UsageError,
661 format!("failed to serialize submit body: {error}"),
662 ));
663 }
664 };
665 let submit_response = loop {
666 let submit = SubmitHttpRequest {
667 route: submit_route(params),
668 auth: auth.clone(),
669 body: body.clone(),
670 retry_resubmit: params.retry,
671 };
672 match client.submit_statement(cx, submit).await {
673 SnowflakeOutcome::Ok(response) if response.status == StatusClass::Unauthorized => {
674 match refresh_after_unauthorized(provider, reauth_left, "submit") {
677 Ok(fresh) => *auth = fresh,
678 Err(error) => return SnowflakeOutcome::err(error),
679 }
680 }
681 SnowflakeOutcome::Ok(response) => break response,
682 SnowflakeOutcome::Err(error) => return SnowflakeOutcome::err(error),
683 SnowflakeOutcome::Cancelled(reason) => return SnowflakeOutcome::cancelled(reason),
684 SnowflakeOutcome::Panicked(payload) => return SnowflakeOutcome::panicked(payload),
685 }
686 };
687 *reauth_left = 1;
688 match machine.on_submit(
689 response_class(submit_response.status),
690 &submit_response.body,
691 ) {
692 Ok(progress) => SnowflakeOutcome::ok(progress),
693 Err(error) => SnowflakeOutcome::err(error.into_snowflake_error()),
694 }
695}
696
697async fn drive_statement<T: StatementTransport, A: AuthProvider>(
698 cx: &Cx,
699 client: &T,
700 provider: &mut A,
701 start: Start,
702 poll_plan: PollPlan,
703 stats: &mut DriverStats,
704 hooks: StatementHooks<'_>,
705) -> StatementOutcome {
706 let StatementHooks {
707 mut sink,
708 mut observer,
709 } = hooks;
710 let mut auth = match provider.descriptor() {
711 Ok(auth) => auth,
712 Err(error) => return SnowflakeOutcome::err(error),
713 };
714 let mut reauth_left: u8 = 1;
716 let poll_interval = poll_plan.effective_poll_interval();
719 let mut machine = StatementMachine::new(poll_plan);
720 let mut poll_now = false;
722 let mut progress = match start {
723 Start::Handle(handle) => {
724 poll_now = true;
725 Progress::PollAgain(handle)
726 }
727 Start::Submit { request, params } => {
728 let progress = match submit(
729 cx,
730 client,
731 provider,
732 &mut auth,
733 &mut reauth_left,
734 &request,
735 ¶ms,
736 &mut machine,
737 )
738 .await
739 {
740 SnowflakeOutcome::Ok(progress) => progress,
741 SnowflakeOutcome::Err(error) => return SnowflakeOutcome::err(error),
742 SnowflakeOutcome::Cancelled(reason) => return SnowflakeOutcome::cancelled(reason),
743 SnowflakeOutcome::Panicked(payload) => return SnowflakeOutcome::panicked(payload),
744 };
745 notify(
746 &mut observer,
747 DriverEvent::Submitted {
748 statement_handle: progress_handle(&progress),
749 running: matches!(progress, Progress::PollAgain(_)),
750 },
751 );
752 progress
753 }
754 };
755
756 loop {
757 match progress {
758 Progress::Complete(mut completed) => {
759 notify(
760 &mut observer,
761 DriverEvent::Completed {
762 rows: completed.result_set.total_rows(),
763 partitions: completed.fetched_partitions,
764 },
765 );
766 if let Some(sink) = sink.as_mut() {
769 let rows = std::mem::take(&mut completed.rows);
770 if let Err(error) = sink.accept(&completed.result_set, rows) {
771 return SnowflakeOutcome::err(error);
772 }
773 }
774 return SnowflakeOutcome::ok(completed);
775 }
776 Progress::TimedOut(failure) => {
777 return SnowflakeOutcome::err(terminal_failure_error(
778 SnowflakeErrorCode::StatementTimeout,
779 failure,
780 ));
781 }
782 Progress::Failed(failure) => {
783 return SnowflakeOutcome::err(terminal_failure_error(
784 SnowflakeErrorCode::StatementFailed,
785 failure,
786 ));
787 }
788 Progress::PollAgain(handle) => {
789 if cx.checkpoint().is_err() {
790 return cancel_locally(cx, client, &auth, &handle, local_cancel_reason(cx))
791 .await;
792 }
793 if !std::mem::take(&mut poll_now)
799 && let Err(reason) = wait_poll_interval(cx, poll_interval).await
800 {
801 return cancel_locally(cx, client, &auth, &handle, reason).await;
802 }
803 stats.polls = stats.polls.saturating_add(1);
804 auth = match provider.descriptor() {
807 Ok(fresh) => fresh,
808 Err(error) => {
809 return abandon_with_error(cx, client, &auth, &handle, error).await;
810 }
811 };
812 let poll = client
813 .poll_statement(
814 cx,
815 PollHttpRequest {
816 auth: auth.clone(),
817 statement_handle: handle.clone(),
818 },
819 )
820 .await;
821 let response = match poll {
822 SnowflakeOutcome::Ok(response) => response,
823 SnowflakeOutcome::Err(error) => {
824 return abandon_with_error(cx, client, &auth, &handle, error).await;
825 }
826 SnowflakeOutcome::Cancelled(reason) => {
827 return cancel_locally(cx, client, &auth, &handle, reason).await;
828 }
829 SnowflakeOutcome::Panicked(payload) => {
830 return abandon_with_outcome(
831 cx,
832 client,
833 &auth,
834 &handle,
835 SnowflakeOutcome::panicked(payload),
836 )
837 .await;
838 }
839 };
840 if response.status == StatusClass::Unauthorized {
841 match refresh_after_unauthorized(provider, &mut reauth_left, "poll") {
842 Ok(fresh) => {
843 auth = fresh;
844 progress = Progress::PollAgain(handle);
845 continue;
846 }
847 Err(error) => {
848 return abandon_with_error(cx, client, &auth, &handle, error).await;
849 }
850 }
851 }
852 reauth_left = 1;
853 progress = match machine.on_poll(response_class(response.status), &response.body) {
854 Ok(progress) => progress,
855 Err(error) => {
856 return abandon_with_error(
857 cx,
858 client,
859 &auth,
860 &handle,
861 error.into_snowflake_error(),
862 )
863 .await;
864 }
865 };
866 notify(&mut observer, DriverEvent::Polled { polls: stats.polls });
867 }
868 Progress::FetchPartition { handle, partition } => {
869 if cx.checkpoint().is_err() {
870 return cancel_locally(cx, client, &auth, &handle, local_cancel_reason(cx))
871 .await;
872 }
873 if let Some(sink) = sink.as_mut() {
876 let rows = machine.drain_rows();
877 if !rows.is_empty()
878 && let Some(result_set) = machine.result_set()
879 && let Err(error) = sink.accept(result_set, rows)
880 {
881 return abandon_with_error(cx, client, &auth, &handle, error).await;
882 }
883 }
884 let (next, total) = machine
885 .assembling_window()
886 .unwrap_or((partition, partition.saturating_add(1)));
887 if let Some(cap) = poll_plan.row_cap
890 && machine.rows_assembled() >= cap
891 {
892 return match machine.complete_early() {
893 Ok(mut done) => {
894 notify(
895 &mut observer,
896 DriverEvent::Completed {
897 rows: done.result_set.total_rows(),
898 partitions: done.fetched_partitions,
899 },
900 );
901 if let Some(sink) = sink.as_mut() {
902 let rows = std::mem::take(&mut done.rows);
903 if let Err(error) = sink.accept(&done.result_set, rows) {
904 return SnowflakeOutcome::err(error);
905 }
906 }
907 SnowflakeOutcome::ok(done)
908 }
909 Err(error) => {
910 abandon_with_error(
911 cx,
912 client,
913 &auth,
914 &handle,
915 error.into_snowflake_error(),
916 )
917 .await
918 }
919 };
920 }
921 auth = match provider.descriptor() {
922 Ok(fresh) => fresh,
923 Err(error) => {
924 return abandon_with_error(cx, client, &auth, &handle, error).await;
925 }
926 };
927 let window =
928 u32::try_from(poll_plan.effective_partition_concurrency()).unwrap_or(u32::MAX);
929 let window_end = next.saturating_add(window).min(total);
930 let window_auth = auth.clone();
931 let fetched =
932 fetch_window(cx, client, &window_auth, &handle, next..window_end).await;
933 stats.partitions_fetched = stats
934 .partitions_fetched
935 .saturating_add(window_end.saturating_sub(next));
936 let mut after_window = None;
937 for (offset, fetch) in fetched.into_iter().enumerate() {
938 let index = next.saturating_add(u32::try_from(offset).unwrap_or(u32::MAX));
939 let mut response = match fetch {
940 SnowflakeOutcome::Ok(response) => response,
941 SnowflakeOutcome::Err(error) => {
942 return abandon_with_error(cx, client, &auth, &handle, error).await;
943 }
944 SnowflakeOutcome::Cancelled(reason) => {
945 return cancel_locally(cx, client, &auth, &handle, reason).await;
946 }
947 SnowflakeOutcome::Panicked(payload) => {
948 return abandon_with_outcome(
949 cx,
950 client,
951 &auth,
952 &handle,
953 SnowflakeOutcome::panicked(payload),
954 )
955 .await;
956 }
957 };
958 if response.status == StatusClass::Unauthorized {
959 if auth == window_auth {
963 auth = match refresh_after_unauthorized(
964 provider,
965 &mut reauth_left,
966 "partition fetch",
967 ) {
968 Ok(fresh) => fresh,
969 Err(error) => {
970 return abandon_with_error(cx, client, &auth, &handle, error)
971 .await;
972 }
973 };
974 }
975 stats.partitions_fetched = stats.partitions_fetched.saturating_add(1);
976 let refetch = client
977 .fetch_partition(
978 cx,
979 PartitionHttpRequest {
980 auth: auth.clone(),
981 statement_handle: handle.clone(),
982 partition: index,
983 },
984 )
985 .await;
986 response = match refetch {
987 SnowflakeOutcome::Ok(response)
988 if response.status == StatusClass::Unauthorized =>
989 {
990 return abandon_with_error(
991 cx,
992 client,
993 &auth,
994 &handle,
995 unauthorized_error(
996 "partition fetch",
997 "the re-signed credential was rejected again; not retrying further",
998 ),
999 )
1000 .await;
1001 }
1002 SnowflakeOutcome::Ok(response) => response,
1003 SnowflakeOutcome::Err(error) => {
1004 return abandon_with_error(cx, client, &auth, &handle, error).await;
1005 }
1006 SnowflakeOutcome::Cancelled(reason) => {
1007 return cancel_locally(cx, client, &auth, &handle, reason).await;
1008 }
1009 SnowflakeOutcome::Panicked(payload) => {
1010 return abandon_with_outcome(
1011 cx,
1012 client,
1013 &auth,
1014 &handle,
1015 SnowflakeOutcome::panicked(payload),
1016 )
1017 .await;
1018 }
1019 };
1020 }
1021 reauth_left = 1;
1022 let partition_rows = machine
1024 .result_set()
1025 .and_then(|result_set| {
1026 result_set
1027 .result_set_meta_data
1028 .partition_info
1029 .get(usize::try_from(index).unwrap_or(usize::MAX))
1030 })
1031 .map_or(0, |info| u64::try_from(info.row_count).unwrap_or(0));
1032 let partition_bytes = u64::try_from(response.body.len()).unwrap_or(u64::MAX);
1033 let partition_progress = match machine.on_partition(
1035 response_class(response.status),
1036 index,
1037 &response.body,
1038 ) {
1039 Ok(progress) => progress,
1040 Err(error) => {
1041 return abandon_with_error(
1042 cx,
1043 client,
1044 &auth,
1045 &handle,
1046 error.into_snowflake_error(),
1047 )
1048 .await;
1049 }
1050 };
1051 notify(
1052 &mut observer,
1053 DriverEvent::PartitionFetched {
1054 index,
1055 rows: partition_rows,
1056 bytes: partition_bytes,
1057 },
1058 );
1059 after_window = Some(partition_progress);
1060 }
1061 progress = match after_window {
1062 Some(progress) => progress,
1063 None => {
1064 return abandon_with_error(
1067 cx,
1068 client,
1069 &auth,
1070 &handle,
1071 SnowflakeError::new(
1072 SnowflakeErrorCode::Internal,
1073 format!("empty partition window {next}..{window_end} of {total}"),
1074 ),
1075 )
1076 .await;
1077 }
1078 };
1079 }
1080 }
1081 }
1082}
1083
1084type BoxedFetch<'a> = Pin<Box<dyn Future<Output = TransportOutcome<PartitionBody>> + 'a>>;
1090
1091async fn fetch_window<T: StatementTransport>(
1092 cx: &Cx,
1093 client: &T,
1094 auth: &AuthorizationDescriptor,
1095 handle: &StatementHandle,
1096 partitions: std::ops::Range<u32>,
1097) -> Vec<TransportOutcome<PartitionBody>> {
1098 let pending: Vec<Option<BoxedFetch<'_>>> = partitions
1099 .map(|partition| {
1100 let request = PartitionHttpRequest {
1101 auth: auth.clone(),
1102 statement_handle: handle.clone(),
1103 partition,
1104 };
1105 let fetch: BoxedFetch<'_> = Box::pin(client.fetch_partition(cx, request));
1106 Some(fetch)
1107 })
1108 .collect();
1109 let done = pending.iter().map(|_| None).collect();
1110 JoinInOrder { pending, done }.await
1111}
1112
1113struct JoinInOrder<'a, T> {
1117 pending: Vec<Option<Pin<Box<dyn Future<Output = T> + 'a>>>>,
1118 done: Vec<Option<T>>,
1119}
1120
1121impl<T: Unpin> Future for JoinInOrder<'_, T> {
1122 type Output = Vec<T>;
1123
1124 fn poll(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
1125 let this = self.get_mut();
1126 let mut all_done = true;
1127 for (slot, done) in this.pending.iter_mut().zip(this.done.iter_mut()) {
1128 if let Some(future) = slot.as_mut() {
1129 match future.as_mut().poll(context) {
1130 Poll::Ready(value) => {
1131 *done = Some(value);
1132 *slot = None;
1133 }
1134 Poll::Pending => all_done = false,
1135 }
1136 }
1137 }
1138 if all_done {
1139 Poll::Ready(this.done.iter_mut().filter_map(Option::take).collect())
1140 } else {
1141 Poll::Pending
1142 }
1143 }
1144}
1145
1146async fn abandon_with_error<T: StatementTransport>(
1151 cx: &Cx,
1152 client: &T,
1153 auth: &AuthorizationDescriptor,
1154 handle: &StatementHandle,
1155 error: SnowflakeError,
1156) -> StatementOutcome {
1157 abandon_with_outcome(cx, client, auth, handle, SnowflakeOutcome::err(error)).await
1158}
1159
1160async fn abandon_with_outcome<T: StatementTransport>(
1164 cx: &Cx,
1165 client: &T,
1166 auth: &AuthorizationDescriptor,
1167 handle: &StatementHandle,
1168 outcome: StatementOutcome,
1169) -> StatementOutcome {
1170 let _ = client
1171 .cancel_orphaned_statement(cx, auth.clone(), handle.clone())
1172 .await;
1173 outcome
1174}
1175
1176async fn cancel_locally<T: StatementTransport>(
1179 cx: &Cx,
1180 client: &T,
1181 auth: &AuthorizationDescriptor,
1182 handle: &StatementHandle,
1183 reason: CancelReason,
1184) -> StatementOutcome {
1185 let _ = client
1188 .cancel_after_local_cancel(cx, auth.clone(), handle.clone(), reason.clone())
1189 .await;
1190 SnowflakeOutcome::cancelled(reason)
1191}
1192
1193fn local_cancel_reason(cx: &Cx) -> CancelReason {
1194 cx.cancel_reason()
1195 .unwrap_or_else(CancelReason::parent_cancelled)
1196}
1197
1198fn terminal_failure_error(
1199 code: SnowflakeErrorCode,
1200 failure: crate::response::QueryFailureStatus,
1201) -> SnowflakeError {
1202 SnowflakeError::new(code, redact(&failure.message).into_owned())
1203}
1204
1205async fn wait_poll_interval(cx: &Cx, delay: Duration) -> Result<(), CancelReason> {
1209 let mut remaining = delay;
1210 while !remaining.is_zero() {
1211 if cx.checkpoint().is_err() {
1212 return Err(local_cancel_reason(cx));
1213 }
1214
1215 let slice = remaining.min(MIN_POLL_INTERVAL);
1216 if asupersync::time::budget_sleep(cx, slice, cx.now_for_observability())
1217 .await
1218 .is_err()
1219 {
1220 let _ = cx.checkpoint();
1225 return Err(local_cancel_reason(cx));
1226 }
1227
1228 if cx.checkpoint().is_err() {
1229 return Err(local_cancel_reason(cx));
1230 }
1231 remaining = remaining.saturating_sub(slice);
1232 }
1233 Ok(())
1234}
1235
1236fn submit_route(params: &SubmitQueryParams) -> TransportRoute {
1238 let query = params.to_query_pairs();
1239 if query.is_empty() {
1240 TransportRoute::Submit
1241 } else {
1242 TransportRoute::SubmitWithQuery { query }
1243 }
1244}
1245
1246const fn response_class(status: StatusClass) -> ResponseClass {
1251 match status {
1252 StatusClass::Completed => ResponseClass::Completed,
1253 StatusClass::Running => ResponseClass::Running,
1254 StatusClass::StatementTimeout => ResponseClass::StatementTimeout,
1255 StatusClass::QueryFailure => ResponseClass::StatementFailed,
1256 StatusClass::RateLimited => ResponseClass::RateLimited,
1257 StatusClass::ServerErrorRetryable => ResponseClass::Other(503),
1258 StatusClass::Unauthorized => ResponseClass::Other(401),
1261 StatusClass::Unexpected => ResponseClass::Other(0),
1262 }
1263}
1264
1265#[cfg(test)]
1266mod tests {
1267 use super::*;
1268 use crate::response::QueryFailureStatus;
1269 use asupersync::{Budget, CancelKind, PanicPayload, Time};
1270 use franken_snowflake_core::outcome::{OutcomeKind, SnowflakeOutcomeExt};
1271 use franken_snowflake_http::{
1272 CompressionEvidence, ContentEncoding, SnowflakeAuthTokenType, TransportError,
1273 TransportErrorCode,
1274 };
1275 use std::cell::{Cell, RefCell};
1276 use std::collections::{BTreeMap, VecDeque};
1277
1278 const RESP_202: &[u8] = include_bytes!("../tests/fixtures/resp_202_running.json");
1279 const RESP_200_MULTI: &[u8] =
1280 include_bytes!("../tests/fixtures/resp_200_resultset_multi_partition.json");
1281 const RESP_200_SINGLE: &[u8] =
1282 include_bytes!("../tests/fixtures/resp_200_resultset_single_partition.json");
1283
1284 #[derive(Clone)]
1287 enum Scripted {
1288 Ok(StatusClass, Vec<u8>),
1289 Err,
1290 Panicked(&'static str),
1291 }
1292
1293 struct FakeTransport {
1295 submit: Scripted,
1296 submit_first: RefCell<Option<Scripted>>,
1298 polls: RefCell<Vec<Scripted>>,
1299 polled: RefCell<Vec<String>>,
1301 partitions: RefCell<BTreeMap<u32, VecDeque<Scripted>>>,
1304 cancels_after_local: RefCell<Vec<(StatementHandle, CancelKind)>>,
1305 orphan_cancels: RefCell<Vec<StatementHandle>>,
1306 orphan_cancel_auth: RefCell<Vec<String>>,
1307 orphan_cancel_result: Scripted,
1308 orphan_cleanup_finished: Cell<bool>,
1309 auth_seen: RefCell<Vec<String>>,
1311 partition_events: RefCell<Vec<(&'static str, u32)>>,
1314 yield_once: Cell<bool>,
1317 }
1318
1319 impl FakeTransport {
1320 fn new(submit: Scripted) -> Self {
1321 Self {
1322 submit,
1323 submit_first: RefCell::new(None),
1324 polls: RefCell::new(Vec::new()),
1325 polled: RefCell::new(Vec::new()),
1326 partitions: RefCell::new(BTreeMap::new()),
1327 cancels_after_local: RefCell::new(Vec::new()),
1328 orphan_cancels: RefCell::new(Vec::new()),
1329 orphan_cancel_auth: RefCell::new(Vec::new()),
1330 orphan_cancel_result: Scripted::Ok(StatusClass::Completed, Vec::new()),
1331 orphan_cleanup_finished: Cell::new(false),
1332 auth_seen: RefCell::new(Vec::new()),
1333 partition_events: RefCell::new(Vec::new()),
1334 yield_once: Cell::new(false),
1335 }
1336 }
1337
1338 fn script_partition(&self, partition: u32, scripted: Scripted) {
1339 self.partitions
1340 .borrow_mut()
1341 .entry(partition)
1342 .or_default()
1343 .push_back(scripted);
1344 }
1345
1346 fn events(&self) -> Vec<String> {
1347 self.partition_events
1348 .borrow()
1349 .iter()
1350 .map(|(kind, partition)| format!("{kind}{partition}"))
1351 .collect()
1352 }
1353
1354 fn transport_error() -> SnowflakeError {
1355 TransportError::new(TransportErrorCode::NetworkError, "connection reset")
1356 .into_snowflake_error()
1357 }
1358 }
1359
1360 impl StatementTransport for FakeTransport {
1361 async fn submit_statement(
1362 &self,
1363 _cx: &Cx,
1364 request: SubmitHttpRequest,
1365 ) -> TransportOutcome<SubmitHttpResponse> {
1366 self.auth_seen
1367 .borrow_mut()
1368 .push(request.auth.redacted_fingerprint().to_owned());
1369 let scripted = self
1370 .submit_first
1371 .borrow_mut()
1372 .take()
1373 .unwrap_or_else(|| self.submit.clone());
1374 match scripted {
1375 Scripted::Ok(status, body) => {
1376 TransportOutcome::ok(SubmitHttpResponse { status, body })
1377 }
1378 Scripted::Err => TransportOutcome::err(Self::transport_error()),
1379 Scripted::Panicked(message) => {
1380 TransportOutcome::panicked(PanicPayload::new(message))
1381 }
1382 }
1383 }
1384
1385 async fn poll_statement(
1386 &self,
1387 _cx: &Cx,
1388 request: PollHttpRequest,
1389 ) -> TransportOutcome<PollHttpResponse> {
1390 self.auth_seen
1391 .borrow_mut()
1392 .push(request.auth.redacted_fingerprint().to_owned());
1393 self.polled
1394 .borrow_mut()
1395 .push(request.statement_handle.as_str().to_owned());
1396 let next = self.polls.borrow_mut().remove(0);
1397 match next {
1398 Scripted::Ok(status, body) => {
1399 TransportOutcome::ok(PollHttpResponse { status, body })
1400 }
1401 Scripted::Err => TransportOutcome::err(Self::transport_error()),
1402 Scripted::Panicked(message) => {
1403 TransportOutcome::panicked(PanicPayload::new(message))
1404 }
1405 }
1406 }
1407
1408 async fn fetch_partition(
1409 &self,
1410 _cx: &Cx,
1411 request: PartitionHttpRequest,
1412 ) -> TransportOutcome<PartitionBody> {
1413 self.auth_seen
1414 .borrow_mut()
1415 .push(request.auth.redacted_fingerprint().to_owned());
1416 self.partition_events
1417 .borrow_mut()
1418 .push(("start", request.partition));
1419 if self.yield_once.get() {
1420 YieldOnce { yielded: false }.await;
1421 }
1422 self.partition_events
1423 .borrow_mut()
1424 .push(("done", request.partition));
1425 let next = self
1426 .partitions
1427 .borrow_mut()
1428 .get_mut(&request.partition)
1429 .and_then(VecDeque::pop_front);
1430 match next {
1431 Some(Scripted::Ok(status, body)) => TransportOutcome::ok(PartitionBody {
1432 status,
1433 compression: CompressionEvidence {
1434 content_encoding: ContentEncoding::Identity,
1435 compressed_bytes: body.len() as u64,
1436 uncompressed_bytes: body.len() as u64,
1437 },
1438 body,
1439 }),
1440 Some(Scripted::Err) | None => TransportOutcome::err(Self::transport_error()),
1441 Some(Scripted::Panicked(message)) => {
1442 TransportOutcome::panicked(PanicPayload::new(message))
1443 }
1444 }
1445 }
1446
1447 async fn cancel_after_local_cancel(
1448 &self,
1449 _cx: &Cx,
1450 _auth: AuthorizationDescriptor,
1451 statement_handle: StatementHandle,
1452 reason: CancelReason,
1453 ) -> TransportOutcome<CancelHttpResponse> {
1454 self.cancels_after_local
1455 .borrow_mut()
1456 .push((statement_handle, reason.kind));
1457 TransportOutcome::cancelled(reason)
1458 }
1459
1460 async fn cancel_orphaned_statement(
1461 &self,
1462 _cx: &Cx,
1463 auth: AuthorizationDescriptor,
1464 statement_handle: StatementHandle,
1465 ) -> TransportOutcome<CancelHttpResponse> {
1466 self.orphan_cancels.borrow_mut().push(statement_handle);
1467 self.orphan_cancel_auth
1468 .borrow_mut()
1469 .push(auth.redacted_fingerprint().to_owned());
1470 if self.yield_once.get() {
1471 YieldOnce { yielded: false }.await;
1472 }
1473 self.orphan_cleanup_finished.set(true);
1474 match &self.orphan_cancel_result {
1475 Scripted::Ok(status, body) => TransportOutcome::ok(CancelHttpResponse {
1476 status: *status,
1477 body: body.clone(),
1478 }),
1479 Scripted::Err => TransportOutcome::err(Self::transport_error()),
1480 Scripted::Panicked(message) => {
1481 TransportOutcome::panicked(PanicPayload::new(*message))
1482 }
1483 }
1484 }
1485 }
1486
1487 fn fake_auth() -> AuthorizationDescriptor {
1488 AuthorizationDescriptor::bearer(
1489 SnowflakeAuthTokenType::ProgrammaticAccessToken,
1490 "fake-token",
1491 "cred_test",
1492 )
1493 }
1494
1495 struct YieldOnce {
1498 yielded: bool,
1499 }
1500
1501 impl Future for YieldOnce {
1502 type Output = ();
1503
1504 fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<()> {
1505 if self.yielded {
1506 Poll::Ready(())
1507 } else {
1508 self.yielded = true;
1509 context.waker().wake_by_ref();
1510 Poll::Pending
1511 }
1512 }
1513 }
1514
1515 fn multi_partition_body(inline_rows: usize, partition_rows: &[usize]) -> Vec<u8> {
1518 let mut partition_info = vec![serde_json::json!({ "rowCount": inline_rows })];
1519 partition_info.extend(
1520 partition_rows
1521 .iter()
1522 .map(|rows| serde_json::json!({ "rowCount": rows, "uncompressedSize": 1 })),
1523 );
1524 let total: usize = inline_rows + partition_rows.iter().sum::<usize>();
1525 let body = serde_json::json!({
1526 "resultSetMetaData": {
1527 "numRows": total,
1528 "format": "jsonv2",
1529 "rowType": [{ "name": "P", "type": "TEXT", "nullable": false }],
1530 "partitionInfo": partition_info
1531 },
1532 "data": (0..inline_rows).map(|_| vec!["p0"]).collect::<Vec<_>>(),
1533 "code": "090001",
1534 "statementHandle": "01b2c3d4-0000-0000-0000-000000000002",
1535 "sqlState": "00000",
1536 "message": "Statement executed successfully.",
1537 "createdOn": 1_700_000_000_000_u64
1538 });
1539 serde_json::to_vec(&body).unwrap_or_default()
1540 }
1541
1542 fn partition_body(partition: u32, rows: usize) -> Scripted {
1543 let body = serde_json::json!({
1544 "data": (0..rows).map(|_| vec![format!("p{partition}")]).collect::<Vec<_>>()
1545 });
1546 Scripted::Ok(
1547 StatusClass::Completed,
1548 serde_json::to_vec(&body).unwrap_or_default(),
1549 )
1550 }
1551
1552 fn windowed_transport(inline_rows: usize, partition_rows: &[usize]) -> FakeTransport {
1555 let transport = FakeTransport::new(Scripted::Ok(
1556 StatusClass::Completed,
1557 multi_partition_body(inline_rows, partition_rows),
1558 ));
1559 for (offset, rows) in partition_rows.iter().enumerate() {
1560 let partition = u32::try_from(offset + 1).unwrap_or(u32::MAX);
1561 transport.script_partition(partition, partition_body(partition, *rows));
1562 }
1563 transport
1564 }
1565
1566 fn column_values(done: &CompletedStatement) -> Vec<String> {
1567 done.rows
1568 .iter()
1569 .map(|row| row[0].clone().unwrap_or_default())
1570 .collect()
1571 }
1572
1573 fn multi_parent(handles: &[&str]) -> Vec<u8> {
1574 let body = serde_json::json!({
1575 "resultSetMetaData": {
1576 "numRows": 1,
1577 "format": "jsonv2",
1578 "rowType": [{ "name": "multiple statement execution", "type": "text", "nullable": false }],
1579 "partitionInfo": [{ "rowCount": 1 }, { "rowCount": 5 }]
1582 },
1583 "data": [["Multiple statements executed successfully."]],
1584 "code": "090001",
1585 "statementHandle": "01b2c3d4-0000-0000-0000-0000000000a0",
1586 "statementHandles": handles,
1587 });
1588 serde_json::to_vec(&body).unwrap_or_default()
1589 }
1590
1591 fn one_value_result(handle: &str, value: &str) -> Scripted {
1592 let body = serde_json::json!({
1593 "resultSetMetaData": {
1594 "numRows": 1,
1595 "format": "jsonv2",
1596 "rowType": [{ "name": "V", "type": "TEXT", "nullable": false }]
1597 },
1598 "data": [[value]],
1599 "code": "090001",
1600 "statementHandle": handle,
1601 });
1602 Scripted::Ok(
1603 StatusClass::Completed,
1604 serde_json::to_vec(&body).unwrap_or_default(),
1605 )
1606 }
1607
1608 #[test]
1611 fn a_multi_statement_request_fetches_each_statement_by_handle_in_order() {
1612 asupersync::test_utils::run_test(|| async {
1613 const FIRST: &str = "01b2c3d4-0000-0000-0000-0000000000a1";
1614 const SECOND: &str = "01b2c3d4-0000-0000-0000-0000000000a2";
1615 let transport = FakeTransport::new(Scripted::Ok(
1616 StatusClass::Completed,
1617 multi_parent(&[FIRST, SECOND]),
1618 ));
1619 *transport.polls.borrow_mut() = vec![
1620 one_value_result(FIRST, "first"),
1621 one_value_result(SECOND, "second"),
1622 ];
1623 let cx = Cx::for_testing();
1624 let (outcome, stats) = run_multi_statement_hooked(
1625 &cx,
1626 &transport,
1627 &mut fake_auth(),
1628 SubmitStatementRequest::new("select 'first'; select 'second'"),
1629 SubmitQueryParams::default(),
1630 fast_poll_plan(5),
1631 None,
1632 )
1633 .await;
1634 let result = match outcome {
1635 SnowflakeOutcome::Ok(result) => result,
1636 other => panic!("expected both statements, got {other:?}"),
1637 };
1638 let values: Vec<Vec<String>> = result.statements.iter().map(column_values).collect();
1639 assert_eq!(values, [["first"], ["second"]]);
1640 assert_eq!(
1641 result.parent.rows,
1642 [[Some(
1643 "Multiple statements executed successfully.".to_owned()
1644 )]]
1645 );
1646 assert_eq!(transport.polled.borrow().as_slice(), [FIRST, SECOND]);
1647 assert_eq!(stats.polls, 2);
1648 assert!(transport.events().is_empty(), "{:?}", transport.events());
1650 assert!(transport.orphan_cancels.borrow().is_empty());
1651 });
1652 }
1653
1654 #[test]
1655 fn a_multi_statement_parent_without_handles_is_an_upstream_error() {
1656 asupersync::test_utils::run_test(|| async {
1657 let transport = FakeTransport::new(Scripted::Ok(
1658 StatusClass::Completed,
1659 RESP_200_SINGLE.to_vec(),
1660 ));
1661 let cx = Cx::for_testing();
1662 let (outcome, _) = run_multi_statement_hooked(
1663 &cx,
1664 &transport,
1665 &mut fake_auth(),
1666 SubmitStatementRequest::new("select 1; select 2"),
1667 SubmitQueryParams::default(),
1668 fast_poll_plan(5),
1669 None,
1670 )
1671 .await;
1672 let error = match outcome {
1673 SnowflakeOutcome::Err(error) => error,
1674 other => panic!("expected an error, got {other:?}"),
1675 };
1676 assert_eq!(error.code, SnowflakeErrorCode::UpstreamError);
1677 assert!(transport.polled.borrow().is_empty());
1678 });
1679 }
1680
1681 #[test]
1682 fn a_failing_statement_fails_the_multi_statement_request_before_any_fetch() {
1683 asupersync::test_utils::run_test(|| async {
1684 let failure = serde_json::json!({
1685 "code": "100132",
1686 "sqlState": "P0000",
1687 "message": "JavaScript execution error: Uncaught Execution of multiple statements failed on statement \"select * from missing_table\"",
1688 "statementHandle": "01b2c3d4-0000-0000-0000-0000000000a0",
1689 });
1690 let transport = FakeTransport::new(Scripted::Ok(
1691 StatusClass::QueryFailure,
1692 serde_json::to_vec(&failure).unwrap_or_default(),
1693 ));
1694 let cx = Cx::for_testing();
1695 let (outcome, _) = run_multi_statement_hooked(
1696 &cx,
1697 &transport,
1698 &mut fake_auth(),
1699 SubmitStatementRequest::new("select 1; select * from missing_table"),
1700 SubmitQueryParams::default(),
1701 fast_poll_plan(5),
1702 None,
1703 )
1704 .await;
1705 let error = match outcome {
1706 SnowflakeOutcome::Err(error) => error,
1707 other => panic!("expected the statement failure, got {other:?}"),
1708 };
1709 assert_eq!(error.code, SnowflakeErrorCode::StatementFailed);
1710 assert!(error.message.contains("missing_table"), "{}", error.message);
1711 assert!(transport.polled.borrow().is_empty());
1712 });
1713 }
1714
1715 #[test]
1716 fn window_fetches_partitions_concurrently_and_assembles_in_order() {
1717 asupersync::test_utils::run_test(|| async {
1718 let transport = windowed_transport(1, &[1, 1, 1, 1, 1]);
1719 transport.yield_once.set(true);
1720 let cx = Cx::for_testing();
1721 let (outcome, stats) = run_statement_with_stats(
1722 &cx,
1723 &transport,
1724 fake_auth(),
1725 SubmitStatementRequest::new("select 1"),
1726 SubmitQueryParams::default(),
1727 fast_poll_plan(5).with_partition_concurrency(3),
1728 )
1729 .await;
1730 let done = match outcome {
1731 SnowflakeOutcome::Ok(done) => done,
1732 other => {
1733 assert!(
1734 matches!(other, SnowflakeOutcome::Ok(_)),
1735 "expected completion, got {other:?}"
1736 );
1737 return;
1738 }
1739 };
1740 assert_eq!(
1741 column_values(&done),
1742 vec!["p0", "p1", "p2", "p3", "p4", "p5"]
1743 );
1744 assert_eq!(done.fetched_partitions, 6);
1745 assert_eq!(done.total_partitions, 6);
1746 assert!(!done.is_partial());
1747 assert_eq!(stats.partitions_fetched, 5);
1748 assert_eq!(
1751 transport.events(),
1752 vec![
1753 "start1", "start2", "start3", "done1", "done2", "done3", "start4", "start5",
1754 "done4", "done5"
1755 ]
1756 );
1757 });
1758 }
1759
1760 #[test]
1761 fn partition_concurrency_one_is_strictly_sequential() {
1762 asupersync::test_utils::run_test(|| async {
1763 let transport = windowed_transport(1, &[1, 1, 1]);
1764 transport.yield_once.set(true);
1765 let cx = Cx::for_testing();
1766 let (outcome, _) = run_statement_with_stats(
1767 &cx,
1768 &transport,
1769 fake_auth(),
1770 SubmitStatementRequest::new("select 1"),
1771 SubmitQueryParams::default(),
1772 fast_poll_plan(5).with_partition_concurrency(1),
1773 )
1774 .await;
1775 assert!(matches!(outcome, SnowflakeOutcome::Ok(_)), "{outcome:?}");
1776 assert_eq!(
1777 transport.events(),
1778 vec!["start1", "done1", "start2", "done2", "start3", "done3"]
1779 );
1780 });
1781 }
1782
1783 #[test]
1784 fn row_cap_stops_fetching_early_and_reports_a_partial_prefix() {
1785 asupersync::test_utils::run_test(|| async {
1786 let transport = windowed_transport(1, &[1, 1, 1, 1, 1]);
1788 let cx = Cx::for_testing();
1789 let (outcome, stats) = run_statement_with_stats(
1790 &cx,
1791 &transport,
1792 fake_auth(),
1793 SubmitStatementRequest::new("select 1"),
1794 SubmitQueryParams::default(),
1795 fast_poll_plan(5)
1796 .with_partition_concurrency(1)
1797 .with_row_cap(Some(3)),
1798 )
1799 .await;
1800 let done = match outcome {
1801 SnowflakeOutcome::Ok(done) => done,
1802 other => {
1803 assert!(
1804 matches!(other, SnowflakeOutcome::Ok(_)),
1805 "expected completion, got {other:?}"
1806 );
1807 return;
1808 }
1809 };
1810 assert_eq!(column_values(&done), vec!["p0", "p1", "p2"]);
1811 assert!(done.is_partial());
1812 assert_eq!(done.fetched_partitions, 3);
1813 assert_eq!(done.total_partitions, 6);
1814 assert_eq!(done.result_set.result_set_meta_data.num_rows, 6);
1815 assert_eq!(
1816 stats.partitions_fetched, 2,
1817 "partitions 3..5 were never fetched"
1818 );
1819 assert_eq!(
1820 transport.events(),
1821 vec!["start1", "done1", "start2", "done2"]
1822 );
1823 assert!(transport.orphan_cancels.borrow().is_empty());
1824 });
1825 }
1826
1827 #[test]
1828 fn row_cap_with_a_window_stops_after_the_window_that_crossed_it() {
1829 asupersync::test_utils::run_test(|| async {
1830 let transport = windowed_transport(1, &[1, 1, 1, 1, 1]);
1831 let cx = Cx::for_testing();
1832 let (outcome, stats) = run_statement_with_stats(
1833 &cx,
1834 &transport,
1835 fake_auth(),
1836 SubmitStatementRequest::new("select 1"),
1837 SubmitQueryParams::default(),
1838 fast_poll_plan(5)
1839 .with_partition_concurrency(3)
1840 .with_row_cap(Some(3)),
1841 )
1842 .await;
1843 let done = match outcome {
1844 SnowflakeOutcome::Ok(done) => done,
1845 other => {
1846 assert!(
1847 matches!(other, SnowflakeOutcome::Ok(_)),
1848 "expected completion, got {other:?}"
1849 );
1850 return;
1851 }
1852 };
1853 assert_eq!(column_values(&done), vec!["p0", "p1", "p2", "p3"]);
1854 assert!(done.is_partial());
1855 assert_eq!(done.fetched_partitions, 4);
1856 assert_eq!(stats.partitions_fetched, 3);
1857 });
1858 }
1859
1860 #[test]
1861 fn row_cap_never_cuts_a_result_that_fits() {
1862 asupersync::test_utils::run_test(|| async {
1863 let transport = windowed_transport(1, &[1, 1]);
1864 let cx = Cx::for_testing();
1865 let (outcome, _) = run_statement_with_stats(
1866 &cx,
1867 &transport,
1868 fake_auth(),
1869 SubmitStatementRequest::new("select 1"),
1870 SubmitQueryParams::default(),
1871 fast_poll_plan(5).with_row_cap(Some(1_000)),
1872 )
1873 .await;
1874 let done = match outcome {
1875 SnowflakeOutcome::Ok(done) => done,
1876 other => {
1877 assert!(
1878 matches!(other, SnowflakeOutcome::Ok(_)),
1879 "expected completion, got {other:?}"
1880 );
1881 return;
1882 }
1883 };
1884 assert!(!done.is_partial());
1885 assert_eq!(done.rows.len(), 3);
1886 });
1887 }
1888
1889 #[test]
1890 fn window_401_resigns_once_and_refetches_every_rejected_partition() {
1891 asupersync::test_utils::run_test(|| async {
1892 let transport = FakeTransport::new(Scripted::Ok(
1893 StatusClass::Completed,
1894 multi_partition_body(1, &[1, 1, 1]),
1895 ));
1896 transport.script_partition(1, unauthorized());
1898 transport.script_partition(1, partition_body(1, 1));
1899 transport.script_partition(2, partition_body(2, 1));
1900 transport.script_partition(3, unauthorized());
1901 transport.script_partition(3, partition_body(3, 1));
1902 let mut auth = FakeAuth::resigning();
1903 let cx = Cx::for_testing();
1904 let (outcome, stats) = run_statement_with_auth(
1905 &cx,
1906 &transport,
1907 &mut auth,
1908 SubmitStatementRequest::new("select 1"),
1909 SubmitQueryParams::default(),
1910 fast_poll_plan(5).with_partition_concurrency(3),
1911 )
1912 .await;
1913 let done = match outcome {
1914 SnowflakeOutcome::Ok(done) => done,
1915 other => {
1916 assert!(
1917 matches!(other, SnowflakeOutcome::Ok(_)),
1918 "expected completion, got {other:?}"
1919 );
1920 return;
1921 }
1922 };
1923 assert_eq!(column_values(&done), vec!["p0", "p1", "p2", "p3"]);
1924 assert_eq!(
1925 auth.resigns, 1,
1926 "one re-sign covers every rejection in the window"
1927 );
1928 assert_eq!(
1929 stats.partitions_fetched, 5,
1930 "3 window fetches + 2 refetches"
1931 );
1932 assert_eq!(
1933 *transport.auth_seen.borrow(),
1934 vec![
1935 "cred_gen0",
1936 "cred_gen0",
1937 "cred_gen0",
1938 "cred_gen0",
1939 "cred_gen1",
1940 "cred_gen1"
1941 ],
1942 "submit + window used gen0; both refetches used the re-signed gen1"
1943 );
1944 assert!(transport.orphan_cancels.borrow().is_empty());
1945 });
1946 }
1947
1948 #[test]
1949 fn partition_rejected_again_after_the_resign_is_terminal_with_an_orphan_cancel() {
1950 asupersync::test_utils::run_test(|| async {
1951 let transport = FakeTransport::new(Scripted::Ok(
1952 StatusClass::Completed,
1953 multi_partition_body(1, &[1]),
1954 ));
1955 transport.script_partition(1, unauthorized());
1956 transport.script_partition(1, unauthorized());
1957 let mut auth = FakeAuth::resigning();
1958 let cx = Cx::for_testing();
1959 let (outcome, stats) = run_statement_with_auth(
1960 &cx,
1961 &transport,
1962 &mut auth,
1963 SubmitStatementRequest::new("select 1"),
1964 SubmitQueryParams::default(),
1965 fast_poll_plan(5),
1966 )
1967 .await;
1968 let error = match outcome {
1969 SnowflakeOutcome::Err(error) => error,
1970 other => {
1971 assert!(
1972 matches!(other, SnowflakeOutcome::Err(_)),
1973 "expected a typed error, got {other:?}"
1974 );
1975 return;
1976 }
1977 };
1978 assert_eq!(error.code, SnowflakeErrorCode::CredentialExpired);
1979 assert_eq!(auth.resigns, 1);
1980 assert_eq!(stats.partitions_fetched, 2);
1981 assert_eq!(transport.orphan_cancels.borrow().len(), 1);
1982 });
1983 }
1984
1985 #[test]
1986 fn one_failed_fetch_in_a_window_abandons_the_statement_after_the_window_settles() {
1987 asupersync::test_utils::run_test(|| async {
1988 let transport = windowed_transport(1, &[1, 1, 1]);
1989 transport.yield_once.set(true);
1990 transport
1992 .partitions
1993 .borrow_mut()
1994 .insert(2, VecDeque::from([Scripted::Err]));
1995 let cx = Cx::for_testing();
1996 let (outcome, stats) = run_statement_with_stats(
1997 &cx,
1998 &transport,
1999 fake_auth(),
2000 SubmitStatementRequest::new("select 1"),
2001 SubmitQueryParams::default(),
2002 fast_poll_plan(5).with_partition_concurrency(3),
2003 )
2004 .await;
2005 assert!(matches!(outcome, SnowflakeOutcome::Err(_)), "{outcome:?}");
2006 assert_eq!(
2009 transport.events(),
2010 vec!["start1", "start2", "start3", "done1", "done2", "done3"]
2011 );
2012 assert_eq!(stats.partitions_fetched, 3);
2013 assert_eq!(transport.orphan_cancels.borrow().len(), 1);
2014 });
2015 }
2016
2017 struct FakeAuth {
2020 can_resign: bool,
2021 generation: u32,
2022 resigns: u32,
2023 }
2024
2025 impl FakeAuth {
2026 fn resigning() -> Self {
2027 Self {
2028 can_resign: true,
2029 generation: 0,
2030 resigns: 0,
2031 }
2032 }
2033
2034 fn frozen_lane() -> Self {
2035 Self {
2036 can_resign: false,
2037 generation: 0,
2038 resigns: 0,
2039 }
2040 }
2041 }
2042
2043 impl AuthProvider for FakeAuth {
2044 fn descriptor(&mut self) -> Result<AuthorizationDescriptor, SnowflakeError> {
2045 Ok(AuthorizationDescriptor::bearer(
2046 SnowflakeAuthTokenType::KeypairJwt,
2047 format!("jwt-gen-{}", self.generation),
2048 format!("cred_gen{}", self.generation),
2049 ))
2050 }
2051
2052 fn on_unauthorized(&mut self) -> Result<bool, SnowflakeError> {
2053 if !self.can_resign {
2054 return Ok(false);
2055 }
2056 self.generation += 1;
2057 self.resigns += 1;
2058 Ok(true)
2059 }
2060 }
2061
2062 fn unauthorized() -> Scripted {
2063 Scripted::Ok(
2064 StatusClass::Unauthorized,
2065 b"{\"message\":\"JWT token is invalid.\"}".to_vec(),
2066 )
2067 }
2068
2069 #[test]
2070 fn poll_401_resigns_once_and_retries_with_the_new_token() {
2071 asupersync::test_utils::run_test(|| async {
2072 let transport =
2073 FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2074 transport.polls.borrow_mut().push(unauthorized());
2075 transport
2076 .polls
2077 .borrow_mut()
2078 .push(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2079 transport.polls.borrow_mut().push(Scripted::Ok(
2080 StatusClass::Completed,
2081 RESP_200_SINGLE.to_vec(),
2082 ));
2083 let mut auth = FakeAuth::resigning();
2084 let cx = Cx::for_testing();
2085 let (outcome, stats) = run_statement_with_auth(
2086 &cx,
2087 &transport,
2088 &mut auth,
2089 SubmitStatementRequest::new("select 1"),
2090 SubmitQueryParams::default(),
2091 fast_poll_plan(5),
2092 )
2093 .await;
2094 assert!(matches!(outcome, SnowflakeOutcome::Ok(_)), "{outcome:?}");
2095 assert_eq!(auth.resigns, 1);
2096 assert_eq!(stats.polls, 3, "the retried poll is a real GET");
2097 assert_eq!(
2098 *transport.auth_seen.borrow(),
2099 vec!["cred_gen0", "cred_gen0", "cred_gen1", "cred_gen1"],
2100 "submit + first poll used gen0; the retry and the next poll used the re-signed gen1"
2101 );
2102 assert!(transport.orphan_cancels.borrow().is_empty());
2103 assert!(transport.cancels_after_local.borrow().is_empty());
2104 });
2105 }
2106
2107 #[test]
2108 fn submit_401_resigns_and_resubmits_without_a_cancel() {
2109 asupersync::test_utils::run_test(|| async {
2110 let transport =
2111 FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2112 *transport.submit_first.borrow_mut() = Some(unauthorized());
2113 transport.polls.borrow_mut().push(Scripted::Ok(
2114 StatusClass::Completed,
2115 RESP_200_SINGLE.to_vec(),
2116 ));
2117 let mut auth = FakeAuth::resigning();
2118 let cx = Cx::for_testing();
2119 let (outcome, stats) = run_statement_with_auth(
2120 &cx,
2121 &transport,
2122 &mut auth,
2123 SubmitStatementRequest::new("select 1"),
2124 SubmitQueryParams::default(),
2125 fast_poll_plan(5),
2126 )
2127 .await;
2128 assert!(matches!(outcome, SnowflakeOutcome::Ok(_)), "{outcome:?}");
2129 assert_eq!(auth.resigns, 1);
2130 assert_eq!(stats.polls, 1);
2131 assert_eq!(
2132 *transport.auth_seen.borrow(),
2133 vec!["cred_gen0", "cred_gen1", "cred_gen1"]
2134 );
2135 assert!(transport.orphan_cancels.borrow().is_empty());
2136 });
2137 }
2138
2139 #[test]
2140 fn poll_401_on_a_lane_that_cannot_resign_is_typed_and_cancels_the_orphan() {
2141 asupersync::test_utils::run_test(|| async {
2142 let transport =
2143 FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2144 transport.polls.borrow_mut().push(unauthorized());
2145 let mut auth = FakeAuth::frozen_lane();
2146 let cx = Cx::for_testing();
2147 let (outcome, _) = run_statement_with_auth(
2148 &cx,
2149 &transport,
2150 &mut auth,
2151 SubmitStatementRequest::new("select 1"),
2152 SubmitQueryParams::default(),
2153 fast_poll_plan(5),
2154 )
2155 .await;
2156 let error = match outcome {
2157 SnowflakeOutcome::Err(error) => error,
2158 other => {
2159 assert!(
2160 matches!(other, SnowflakeOutcome::Err(_)),
2161 "expected a typed error, got {other:?}"
2162 );
2163 return;
2164 }
2165 };
2166 assert_eq!(error.code, SnowflakeErrorCode::CredentialExpired);
2167 assert!(error.message.contains("401"), "{}", error.message);
2168 assert_eq!(auth.resigns, 0);
2169 assert_eq!(transport.orphan_cancels.borrow().len(), 1);
2170 });
2171 }
2172
2173 #[test]
2174 fn frozen_descriptor_entry_point_treats_401_as_terminal() {
2175 asupersync::test_utils::run_test(|| async {
2176 let transport =
2177 FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2178 transport.polls.borrow_mut().push(unauthorized());
2179 let cx = Cx::for_testing();
2180 let (outcome, _) = run_statement_with_stats(
2181 &cx,
2182 &transport,
2183 fake_auth(),
2184 SubmitStatementRequest::new("select 1"),
2185 SubmitQueryParams::default(),
2186 fast_poll_plan(5),
2187 )
2188 .await;
2189 let error = match outcome {
2190 SnowflakeOutcome::Err(error) => error,
2191 other => {
2192 assert!(
2193 matches!(other, SnowflakeOutcome::Err(_)),
2194 "expected a typed error, got {other:?}"
2195 );
2196 return;
2197 }
2198 };
2199 assert_eq!(error.code, SnowflakeErrorCode::CredentialExpired);
2200 assert_eq!(transport.orphan_cancels.borrow().len(), 1);
2201 });
2202 }
2203
2204 #[test]
2205 fn two_consecutive_401s_stop_after_exactly_one_resign() {
2206 asupersync::test_utils::run_test(|| async {
2207 let transport =
2208 FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2209 transport.polls.borrow_mut().push(unauthorized());
2210 transport.polls.borrow_mut().push(unauthorized());
2211 let mut auth = FakeAuth::resigning();
2212 let cx = Cx::for_testing();
2213 let (outcome, stats) = run_statement_with_auth(
2214 &cx,
2215 &transport,
2216 &mut auth,
2217 SubmitStatementRequest::new("select 1"),
2218 SubmitQueryParams::default(),
2219 fast_poll_plan(5),
2220 )
2221 .await;
2222 let error = match outcome {
2223 SnowflakeOutcome::Err(error) => error,
2224 other => {
2225 assert!(
2226 matches!(other, SnowflakeOutcome::Err(_)),
2227 "expected a typed error, got {other:?}"
2228 );
2229 return;
2230 }
2231 };
2232 assert_eq!(error.code, SnowflakeErrorCode::CredentialExpired);
2233 assert!(
2234 error.message.contains("rejected again"),
2235 "{}",
2236 error.message
2237 );
2238 assert_eq!(auth.resigns, 1, "exactly one re-sign, no loop");
2239 assert_eq!(stats.polls, 2);
2240 assert_eq!(transport.orphan_cancels.borrow().len(), 1);
2241 assert!(transport.polls.borrow().is_empty());
2242 });
2243 }
2244
2245 #[test]
2246 fn partition_401_resigns_once_and_refetches() {
2247 asupersync::test_utils::run_test(|| async {
2248 let transport =
2249 FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2250 transport.polls.borrow_mut().push(Scripted::Ok(
2251 StatusClass::Completed,
2252 RESP_200_MULTI.to_vec(),
2253 ));
2254 let multi: serde_json::Value =
2255 serde_json::from_slice(RESP_200_MULTI).unwrap_or_default();
2256 let partitions = multi["resultSetMetaData"]["partitionInfo"]
2257 .as_array()
2258 .cloned()
2259 .unwrap_or_default();
2260 let partition_count = partitions.len();
2261 transport.script_partition(1, unauthorized());
2264 for (index, info) in partitions.iter().enumerate().skip(1) {
2265 let rows = info["rowCount"].as_u64().unwrap_or(0);
2266 let body = format!(
2267 r#"{{"data":[{}]}}"#,
2268 (0..rows)
2269 .map(|_| r#"["2024-01-02","ENTITY","2.50"]"#)
2270 .collect::<Vec<_>>()
2271 .join(",")
2272 );
2273 transport.script_partition(
2274 u32::try_from(index).unwrap_or(u32::MAX),
2275 Scripted::Ok(StatusClass::Completed, body.into_bytes()),
2276 );
2277 }
2278 let mut auth = FakeAuth::resigning();
2279 let cx = Cx::for_testing();
2280 let (outcome, stats) = run_statement_with_auth(
2281 &cx,
2282 &transport,
2283 &mut auth,
2284 SubmitStatementRequest::new("select 1"),
2285 SubmitQueryParams::default(),
2286 fast_poll_plan(5),
2287 )
2288 .await;
2289 assert!(matches!(outcome, SnowflakeOutcome::Ok(_)), "{outcome:?}");
2290 assert_eq!(auth.resigns, 1);
2291 assert_eq!(
2292 stats.partitions_fetched as usize, partition_count,
2293 "one extra fetch for the retry"
2294 );
2295 assert!(transport.orphan_cancels.borrow().is_empty());
2296 });
2297 }
2298
2299 fn fast_poll_plan(max_polls: u32) -> PollPlan {
2300 PollPlan {
2301 max_polls,
2302 poll_interval: Duration::ZERO,
2303 ..PollPlan::default()
2304 }
2305 }
2306
2307 fn fixture_handle() -> StatementHandle {
2308 StatementHandle::new("01b2c3d4-0000-0000-0000-000000000002")
2309 }
2310
2311 #[test]
2312 fn submit_panic_without_a_handle_does_not_attempt_cleanup() {
2313 asupersync::test_utils::run_test(|| async {
2314 let transport = FakeTransport::new(Scripted::Panicked("submit panic"));
2315 let cx = Cx::current().unwrap_or_else(Cx::for_testing);
2316 let (outcome, stats) = run_statement_with_stats(
2317 &cx,
2318 &transport,
2319 fake_auth(),
2320 SubmitStatementRequest::new("select 1"),
2321 SubmitQueryParams::default(),
2322 fast_poll_plan(5),
2323 )
2324 .await;
2325 let payload = match outcome {
2326 SnowflakeOutcome::Panicked(payload) => payload,
2327 other => {
2328 assert!(
2329 matches!(other, SnowflakeOutcome::Panicked(_)),
2330 "expected the submit panic, got {other:?}"
2331 );
2332 return;
2333 }
2334 };
2335 assert_eq!(payload.message(), "submit panic");
2336 assert_eq!(stats, DriverStats::default());
2337 assert!(transport.orphan_cancels.borrow().is_empty());
2338 assert!(transport.cancels_after_local.borrow().is_empty());
2339 assert!(!transport.orphan_cleanup_finished.get());
2340 });
2341 }
2342
2343 #[test]
2344 fn poll_panic_awaits_cleanup_and_preserves_the_original_payload() {
2345 asupersync::test_utils::run_test(|| async {
2346 for cleanup in [
2347 Scripted::Ok(StatusClass::Completed, Vec::new()),
2348 Scripted::Err,
2349 Scripted::Panicked("secondary cleanup panic"),
2350 ] {
2351 let mut transport =
2352 FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2353 transport.orphan_cancel_result = cleanup;
2354 transport.yield_once.set(true);
2355 transport
2356 .polls
2357 .borrow_mut()
2358 .push(Scripted::Panicked("poll panic"));
2359 let cx = Cx::current().unwrap_or_else(Cx::for_testing);
2360 let (outcome, stats) = run_statement_with_stats(
2361 &cx,
2362 &transport,
2363 fake_auth(),
2364 SubmitStatementRequest::new("select 1"),
2365 SubmitQueryParams::default(),
2366 fast_poll_plan(5),
2367 )
2368 .await;
2369 let payload = match outcome {
2370 SnowflakeOutcome::Panicked(payload) => payload,
2371 other => {
2372 assert!(
2373 matches!(other, SnowflakeOutcome::Panicked(_)),
2374 "expected the original poll panic, got {other:?}"
2375 );
2376 continue;
2377 }
2378 };
2379 assert_eq!(payload.message(), "poll panic");
2380 assert_eq!(stats.polls, 1);
2381 assert_eq!(
2382 transport.orphan_cancels.borrow().as_slice(),
2383 &[fixture_handle()]
2384 );
2385 assert_eq!(
2386 transport.orphan_cancel_auth.borrow().as_slice(),
2387 &["cred_test"]
2388 );
2389 assert!(transport.orphan_cleanup_finished.get());
2390 assert!(transport.cancels_after_local.borrow().is_empty());
2391 }
2392 });
2393 }
2394
2395 #[test]
2396 fn partition_panic_drains_the_window_and_yielding_cleanup_before_returning() {
2397 let transport = windowed_transport(1, &[1, 1, 1]);
2398 transport.yield_once.set(true);
2399 transport
2400 .partitions
2401 .borrow_mut()
2402 .insert(2, VecDeque::from([Scripted::Panicked("partition panic")]));
2403 let cx = Cx::for_testing();
2404 let mut driver = std::pin::pin!(run_statement_with_stats(
2405 &cx,
2406 &transport,
2407 fake_auth(),
2408 SubmitStatementRequest::new("select 1"),
2409 SubmitQueryParams::default(),
2410 fast_poll_plan(5).with_partition_concurrency(3),
2411 ));
2412 let mut context = Context::from_waker(std::task::Waker::noop());
2413
2414 assert!(driver.as_mut().poll(&mut context).is_pending());
2415 assert_eq!(transport.events(), vec!["start1", "start2", "start3"]);
2416 assert!(transport.orphan_cancels.borrow().is_empty());
2417
2418 assert!(driver.as_mut().poll(&mut context).is_pending());
2421 assert_eq!(
2422 transport.events(),
2423 vec!["start1", "start2", "start3", "done1", "done2", "done3"]
2424 );
2425 assert_eq!(
2426 transport.orphan_cancels.borrow().as_slice(),
2427 &[fixture_handle()]
2428 );
2429 assert!(!transport.orphan_cleanup_finished.get());
2430
2431 let poll_result = driver.as_mut().poll(&mut context);
2432 assert!(
2433 matches!(poll_result, Poll::Ready(_)),
2434 "driver did not return after cleanup completed"
2435 );
2436 let (outcome, stats) = match poll_result {
2437 Poll::Ready(ready) => ready,
2438 Poll::Pending => return,
2439 };
2440 let payload = match outcome {
2441 SnowflakeOutcome::Panicked(payload) => payload,
2442 other => {
2443 assert!(
2444 matches!(other, SnowflakeOutcome::Panicked(_)),
2445 "expected the partition panic, got {other:?}"
2446 );
2447 return;
2448 }
2449 };
2450 assert_eq!(payload.message(), "partition panic");
2451 assert_eq!(stats.partitions_fetched, 3);
2452 assert_eq!(transport.orphan_cancels.borrow().len(), 1);
2453 assert!(transport.orphan_cleanup_finished.get());
2454 assert!(transport.cancels_after_local.borrow().is_empty());
2455 }
2456
2457 #[test]
2458 fn partition_retry_panic_cleans_up_with_the_refreshed_credential() {
2459 asupersync::test_utils::run_test(|| async {
2460 let transport = FakeTransport::new(Scripted::Ok(
2461 StatusClass::Completed,
2462 multi_partition_body(1, &[1]),
2463 ));
2464 transport.yield_once.set(true);
2465 transport.script_partition(1, unauthorized());
2466 transport.script_partition(1, Scripted::Panicked("retry panic"));
2467 let mut auth = FakeAuth::resigning();
2468 let cx = Cx::current().unwrap_or_else(Cx::for_testing);
2469 let (outcome, stats) = run_statement_with_auth(
2470 &cx,
2471 &transport,
2472 &mut auth,
2473 SubmitStatementRequest::new("select 1"),
2474 SubmitQueryParams::default(),
2475 fast_poll_plan(5),
2476 )
2477 .await;
2478 let payload = match outcome {
2479 SnowflakeOutcome::Panicked(payload) => payload,
2480 other => {
2481 assert!(
2482 matches!(other, SnowflakeOutcome::Panicked(_)),
2483 "expected the retry panic, got {other:?}"
2484 );
2485 return;
2486 }
2487 };
2488 assert_eq!(payload.message(), "retry panic");
2489 assert_eq!(stats.partitions_fetched, 2);
2490 assert_eq!(auth.resigns, 1);
2491 assert_eq!(
2492 transport.orphan_cancels.borrow().as_slice(),
2493 &[fixture_handle()]
2494 );
2495 assert_eq!(
2496 transport.orphan_cancel_auth.borrow().as_slice(),
2497 &["cred_gen1"]
2498 );
2499 assert!(transport.orphan_cleanup_finished.get());
2500 assert!(transport.cancels_after_local.borrow().is_empty());
2501 });
2502 }
2503
2504 #[test]
2505 fn poll_transport_error_after_submit_fires_an_orphan_cancel() {
2506 asupersync::test_utils::run_test(|| async {
2507 let transport =
2508 FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2509 transport.polls.borrow_mut().push(Scripted::Err);
2510 let cx = Cx::for_testing();
2511 let (outcome, stats) = run_statement_with_stats(
2512 &cx,
2513 &transport,
2514 fake_auth(),
2515 SubmitStatementRequest::new("select 1"),
2516 SubmitQueryParams::default(),
2517 fast_poll_plan(5),
2518 )
2519 .await;
2520 assert!(matches!(outcome, SnowflakeOutcome::Err(_)));
2521 assert_eq!(stats.polls, 1);
2522 assert_eq!(
2523 transport.orphan_cancels.borrow().as_slice(),
2524 &[fixture_handle()],
2525 "a transport error after the handle exists must cancel the orphaned statement"
2526 );
2527 assert!(transport.cancels_after_local.borrow().is_empty());
2528 });
2529 }
2530
2531 #[test]
2532 fn undecodable_poll_body_after_submit_fires_an_orphan_cancel() {
2533 asupersync::test_utils::run_test(|| async {
2534 let transport =
2535 FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2536 transport
2537 .polls
2538 .borrow_mut()
2539 .push(Scripted::Ok(StatusClass::Completed, b"not json".to_vec()));
2540 let cx = Cx::for_testing();
2541 let (outcome, _) = run_statement_with_stats(
2542 &cx,
2543 &transport,
2544 fake_auth(),
2545 SubmitStatementRequest::new("select 1"),
2546 SubmitQueryParams::default(),
2547 fast_poll_plan(5),
2548 )
2549 .await;
2550 assert!(matches!(outcome, SnowflakeOutcome::Err(_)));
2551 assert_eq!(transport.orphan_cancels.borrow().len(), 1);
2552 });
2553 }
2554
2555 struct CollectingSink {
2557 batches: Vec<Vec<Vec<Option<String>>>>,
2558 refuse: bool,
2559 }
2560
2561 impl RowSink for CollectingSink {
2562 fn accept(
2563 &mut self,
2564 _result_set: &ResultSet,
2565 rows: Vec<Vec<Option<String>>>,
2566 ) -> Result<(), SnowflakeError> {
2567 if self.refuse {
2568 return Err(SnowflakeError::new(
2569 SnowflakeErrorCode::UsageError,
2570 "the sink refused the rows",
2571 ));
2572 }
2573 self.batches.push(rows);
2574 Ok(())
2575 }
2576 }
2577
2578 fn streaming_transport() -> FakeTransport {
2579 let transport = FakeTransport::new(Scripted::Ok(
2580 StatusClass::Completed,
2581 RESP_200_MULTI.to_vec(),
2582 ));
2583 transport.script_partition(
2584 1,
2585 Scripted::Ok(
2586 StatusClass::Completed,
2587 br#"{"data":[["p1a","x"],["p1b","x"]]}"#.to_vec(),
2588 ),
2589 );
2590 transport.script_partition(
2591 2,
2592 Scripted::Ok(
2593 StatusClass::Completed,
2594 br#"{"data":[["p2a","x"]]}"#.to_vec(),
2595 ),
2596 );
2597 transport
2598 }
2599
2600 #[test]
2604 fn streaming_hands_rows_to_the_sink_one_window_at_a_time() {
2605 asupersync::test_utils::run_test(|| async {
2606 let transport = streaming_transport();
2607 let mut sink = CollectingSink {
2608 batches: Vec::new(),
2609 refuse: false,
2610 };
2611 let mut auth = fake_auth();
2612 let cx = Cx::for_testing();
2613 let (outcome, _) = run_statement_streaming(
2614 &cx,
2615 &transport,
2616 &mut auth,
2617 SubmitStatementRequest::new("select 1"),
2618 SubmitQueryParams::default(),
2619 fast_poll_plan(5).with_partition_concurrency(1),
2620 &mut sink,
2621 )
2622 .await;
2623 let SnowflakeOutcome::Ok(done) = outcome else {
2624 panic!("streaming run failed: {outcome:?}");
2625 };
2626 assert!(done.rows.is_empty(), "every row went to the sink");
2627 assert_eq!(done.fetched_partitions, 3);
2628 let firsts: Vec<String> = sink
2629 .batches
2630 .iter()
2631 .flatten()
2632 .map(|row| row.first().cloned().flatten().unwrap_or_default())
2633 .collect();
2634 assert_eq!(firsts.len(), 5, "inline 2 + 2 + 1");
2635 assert_eq!(&firsts[2..], ["p1a", "p1b", "p2a"]);
2636 assert_eq!(
2637 sink.batches.len(),
2638 3,
2639 "one batch per partition with window 1"
2640 );
2641 assert!(
2642 sink.batches.iter().all(|batch| batch.len() <= 2),
2643 "no batch holds more than one partition"
2644 );
2645 assert!(transport.orphan_cancels.borrow().is_empty());
2646 });
2647 }
2648
2649 #[derive(Default)]
2651 struct CollectingObserver(Vec<DriverEvent>);
2652
2653 impl DriverObserver for CollectingObserver {
2654 fn event(&mut self, event: DriverEvent) {
2655 self.0.push(event);
2656 }
2657 }
2658
2659 #[test]
2662 fn the_observer_sees_each_remote_cancel_and_its_answer() {
2663 asupersync::test_utils::run_test(|| async {
2664 let run = |transport: FakeTransport| async move {
2665 let mut observer = CollectingObserver::default();
2666 let mut auth = fake_auth();
2667 let cx = Cx::for_testing();
2668 let hooks = StatementHooks {
2669 sink: None,
2670 observer: Some(&mut observer),
2671 };
2672 let (outcome, _) = run_statement_hooked(
2673 &cx,
2674 &transport,
2675 &mut auth,
2676 SubmitStatementRequest::new("select 1"),
2677 SubmitQueryParams::default(),
2678 fast_poll_plan(5),
2679 hooks,
2680 )
2681 .await;
2682 (outcome, observer.0)
2683 };
2684 let cancels = |events: &[DriverEvent]| {
2685 events
2686 .iter()
2687 .filter_map(|event| match event {
2688 DriverEvent::RemoteCancel {
2689 statement_handle,
2690 acknowledged,
2691 detail,
2692 } => Some((statement_handle.clone(), *acknowledged, detail.clone())),
2693 _ => None,
2694 })
2695 .collect::<Vec<_>>()
2696 };
2697
2698 let acknowledged =
2701 FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2702 acknowledged.polls.borrow_mut().push(Scripted::Err);
2703 let (outcome, events) = run(acknowledged).await;
2704 assert!(matches!(outcome, SnowflakeOutcome::Err(_)), "{outcome:?}");
2705 assert_eq!(
2706 cancels(&events),
2707 vec![(
2708 fixture_handle().as_str().to_owned(),
2709 true,
2710 "completed".to_owned()
2711 )],
2712 "{events:?}"
2713 );
2714
2715 let mut refused =
2717 FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2718 refused.polls.borrow_mut().push(Scripted::Err);
2719 refused.orphan_cancel_result = Scripted::Err;
2720 let (_, events) = run(refused).await;
2721 let recorded = cancels(&events);
2722 assert_eq!(recorded.len(), 1, "{events:?}");
2723 assert!(!recorded[0].1, "{recorded:?}");
2724
2725 let (outcome, events) = run(streaming_transport()).await;
2727 assert!(matches!(outcome, SnowflakeOutcome::Ok(_)), "{outcome:?}");
2728 assert!(cancels(&events).is_empty(), "{events:?}");
2729 });
2730 }
2731
2732 #[test]
2736 fn the_observer_sees_submit_partitions_and_completion() {
2737 asupersync::test_utils::run_test(|| async {
2738 let transport = streaming_transport();
2739 let mut observer = CollectingObserver::default();
2740 let mut auth = fake_auth();
2741 let cx = Cx::for_testing();
2742 let hooks = StatementHooks {
2743 sink: None,
2744 observer: Some(&mut observer),
2745 };
2746 let (outcome, _) = run_statement_hooked(
2747 &cx,
2748 &transport,
2749 &mut auth,
2750 SubmitStatementRequest::new("select 1"),
2751 SubmitQueryParams::default(),
2752 fast_poll_plan(5).with_partition_concurrency(1),
2753 hooks,
2754 )
2755 .await;
2756 assert!(matches!(outcome, SnowflakeOutcome::Ok(_)), "{outcome:?}");
2757 let events = observer.0;
2758 assert!(
2759 matches!(
2760 &events[0],
2761 DriverEvent::Submitted {
2762 running: false,
2763 statement_handle: Some(_)
2764 }
2765 ),
2766 "{events:?}"
2767 );
2768 assert_eq!(
2769 events[1],
2770 DriverEvent::PartitionFetched {
2771 index: 1,
2772 rows: 2,
2773 bytes: u64::try_from(br#"{"data":[["p1a","x"],["p1b","x"]]}"#.len())
2774 .unwrap_or(0),
2775 }
2776 );
2777 assert!(
2778 matches!(
2779 events[2],
2780 DriverEvent::PartitionFetched {
2781 index: 2,
2782 rows: 1,
2783 ..
2784 }
2785 ),
2786 "{events:?}"
2787 );
2788 assert_eq!(
2789 events[3],
2790 DriverEvent::Completed {
2791 rows: 5,
2792 partitions: 3
2793 }
2794 );
2795 assert_eq!(events.len(), 4, "{events:?}");
2796 });
2797 }
2798
2799 #[test]
2801 fn a_failing_sink_cancels_the_statement() {
2802 asupersync::test_utils::run_test(|| async {
2803 let transport = streaming_transport();
2804 let mut sink = CollectingSink {
2805 batches: Vec::new(),
2806 refuse: true,
2807 };
2808 let mut auth = fake_auth();
2809 let cx = Cx::for_testing();
2810 let (outcome, stats) = run_statement_streaming(
2811 &cx,
2812 &transport,
2813 &mut auth,
2814 SubmitStatementRequest::new("select 1"),
2815 SubmitQueryParams::default(),
2816 fast_poll_plan(5).with_partition_concurrency(1),
2817 &mut sink,
2818 )
2819 .await;
2820 assert!(matches!(outcome, SnowflakeOutcome::Err(_)), "{outcome:?}");
2821 assert_eq!(stats.partitions_fetched, 0, "stopped before any fetch");
2822 assert_eq!(transport.orphan_cancels.borrow().len(), 1);
2823 });
2824 }
2825
2826 #[test]
2827 fn partition_fetch_error_fires_an_orphan_cancel() {
2828 asupersync::test_utils::run_test(|| async {
2829 let transport = FakeTransport::new(Scripted::Ok(
2830 StatusClass::Completed,
2831 RESP_200_MULTI.to_vec(),
2832 ));
2833 transport.script_partition(1, Scripted::Err);
2834 let cx = Cx::for_testing();
2835 let (outcome, stats) = run_statement_with_stats(
2836 &cx,
2837 &transport,
2838 fake_auth(),
2839 SubmitStatementRequest::new("select 1"),
2840 SubmitQueryParams::default(),
2841 fast_poll_plan(5).with_partition_concurrency(1),
2842 )
2843 .await;
2844 assert!(matches!(outcome, SnowflakeOutcome::Err(_)));
2845 assert_eq!(stats.partitions_fetched, 1);
2846 assert_eq!(transport.orphan_cancels.borrow().len(), 1);
2847 });
2848 }
2849
2850 #[test]
2851 fn deadline_during_poll_routes_through_the_policy_cancel_with_deadline_kind() {
2852 asupersync::test_utils::run_test(|| async {
2853 let transport =
2854 FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2855 for _ in 0..10 {
2857 transport
2858 .polls
2859 .borrow_mut()
2860 .push(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2861 }
2862 let cx = Cx::for_testing_with_budget(Budget::new().with_deadline(Time::from_millis(1)));
2863 let (outcome, _) = run_statement_with_stats(
2864 &cx,
2865 &transport,
2866 fake_auth(),
2867 SubmitStatementRequest::new("select 1"),
2868 SubmitQueryParams::default(),
2869 PollPlan {
2870 max_polls: 50,
2871 poll_interval: Duration::from_millis(5),
2872 ..PollPlan::default()
2873 },
2874 )
2875 .await;
2876 assert!(matches!(outcome, SnowflakeOutcome::Cancelled(_)));
2877 let cancels = transport.cancels_after_local.borrow();
2878 assert_eq!(cancels.len(), 1);
2879 assert_eq!(cancels[0].0, fixture_handle());
2880 assert_eq!(cancels[0].1, CancelKind::Deadline);
2881 assert!(transport.orphan_cancels.borrow().is_empty());
2882 });
2883 }
2884
2885 #[test]
2886 fn happy_path_reports_polls_and_partitions_without_any_cancel() {
2887 asupersync::test_utils::run_test(|| async {
2888 let transport =
2889 FakeTransport::new(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2890 transport
2891 .polls
2892 .borrow_mut()
2893 .push(Scripted::Ok(StatusClass::Running, RESP_202.to_vec()));
2894 transport.polls.borrow_mut().push(Scripted::Ok(
2895 StatusClass::Completed,
2896 RESP_200_MULTI.to_vec(),
2897 ));
2898 let multi: serde_json::Value =
2899 serde_json::from_slice(RESP_200_MULTI).unwrap_or_default();
2900 let partitions = multi["resultSetMetaData"]["partitionInfo"]
2901 .as_array()
2902 .cloned()
2903 .unwrap_or_default();
2904 let partition_count = partitions.len();
2905 for (index, info) in partitions.iter().enumerate().skip(1) {
2909 let rows = info["rowCount"].as_u64().unwrap_or(0);
2910 let body = format!(
2911 r#"{{"data":[{}]}}"#,
2912 (0..rows)
2913 .map(|_| r#"["2024-01-02","ENTITY","2.50"]"#)
2914 .collect::<Vec<_>>()
2915 .join(",")
2916 );
2917 transport.script_partition(
2918 u32::try_from(index).unwrap_or(u32::MAX),
2919 Scripted::Ok(StatusClass::Completed, body.into_bytes()),
2920 );
2921 }
2922 let cx = Cx::for_testing();
2923 let (outcome, stats) = run_statement_with_stats(
2924 &cx,
2925 &transport,
2926 fake_auth(),
2927 SubmitStatementRequest::new("select 1"),
2928 SubmitQueryParams::default(),
2929 fast_poll_plan(5),
2930 )
2931 .await;
2932 assert!(matches!(outcome, SnowflakeOutcome::Ok(_)), "{outcome:?}");
2933 assert_eq!(stats.polls, 2);
2934 assert_eq!(stats.partitions_fetched as usize, partition_count - 1);
2935 assert!(transport.orphan_cancels.borrow().is_empty());
2936 assert!(transport.cancels_after_local.borrow().is_empty());
2937 });
2938 }
2939
2940 #[test]
2941 fn response_class_maps_each_transport_status() {
2942 assert_eq!(
2943 response_class(StatusClass::Completed),
2944 ResponseClass::Completed
2945 );
2946 assert_eq!(response_class(StatusClass::Running), ResponseClass::Running);
2947 assert_eq!(
2948 response_class(StatusClass::StatementTimeout),
2949 ResponseClass::StatementTimeout
2950 );
2951 assert_eq!(
2954 response_class(StatusClass::QueryFailure),
2955 ResponseClass::StatementFailed
2956 );
2957 assert_eq!(
2958 response_class(StatusClass::RateLimited),
2959 ResponseClass::RateLimited
2960 );
2961 }
2962
2963 #[test]
2964 fn submit_route_requires_request_id_and_retry_for_resubmit() {
2965 let plain = SubmitQueryParams::default();
2966 assert!(matches!(submit_route(&plain), TransportRoute::Submit));
2967
2968 let resubmit = SubmitQueryParams {
2969 request_id: Some("req-1".to_owned()),
2970 retry: true,
2971 ..SubmitQueryParams::default()
2972 };
2973 assert!(submit_route(&resubmit).has_retry_contract());
2974
2975 let no_id = SubmitQueryParams {
2977 retry: true,
2978 ..SubmitQueryParams::default()
2979 };
2980 assert!(!submit_route(&no_id).has_retry_contract());
2981 }
2982
2983 #[test]
2984 fn submit_route_golden_preserves_async_and_nullable_query_params() {
2985 let params = SubmitQueryParams {
2986 request_id: Some("req-async-nullable".to_owned()),
2987 retry: true,
2988 asynchronous: true,
2989 nullable: Some(false),
2990 };
2991 let expected_pairs = params.to_query_pairs();
2992
2993 let route = submit_route(¶ms);
2994 assert!(matches!(
2995 &route,
2996 TransportRoute::SubmitWithQuery { query } if query == &expected_pairs
2997 ));
2998 assert!(route.has_retry_contract());
2999 assert_eq!(
3000 route.path_and_query(),
3001 "/api/v2/statements?requestId=req-async-nullable&retry=true&async=true&nullable=false"
3002 );
3003 }
3004
3005 #[test]
3006 fn wait_poll_interval_preserves_deadline_attribution() {
3007 asupersync::test_utils::run_test(|| async {
3008 let cx = Cx::for_testing_with_budget(Budget::new().with_deadline(Time::from_millis(1)));
3009
3010 let reason = wait_poll_interval(&cx, Duration::from_millis(10))
3011 .await
3012 .expect_err("deadline should expire during poll wait");
3013
3014 assert_eq!(reason.kind, CancelKind::Deadline);
3015 });
3016 }
3017
3018 #[test]
3019 fn terminal_statement_failures_keep_precise_error_projection() {
3020 let timeout = QueryFailureStatus {
3021 code: "000630".to_owned(),
3022 sql_state: Some("57014".to_owned()),
3023 message: "Statement reached its statement timeout and was canceled.".to_owned(),
3024 statement_handle: Some(StatementHandle::new("timeout-handle")),
3025 };
3026 let timeout_error = terminal_failure_error(SnowflakeErrorCode::StatementTimeout, timeout);
3027 let timeout_outcome: StatementOutcome = SnowflakeOutcome::err(timeout_error.clone());
3028 assert_eq!(timeout_error.code, SnowflakeErrorCode::StatementTimeout);
3029 assert_eq!(timeout_outcome.outcome_kind(), OutcomeKind::Timeout);
3030
3031 let failure = QueryFailureStatus {
3032 code: "001003".to_owned(),
3033 sql_state: Some("42000".to_owned()),
3034 message: "SQL compilation error.".to_owned(),
3035 statement_handle: Some(StatementHandle::new("failed-handle")),
3036 };
3037 let failure_error = terminal_failure_error(SnowflakeErrorCode::StatementFailed, failure);
3038 let failure_outcome: StatementOutcome = SnowflakeOutcome::err(failure_error.clone());
3039 assert_eq!(failure_error.code, SnowflakeErrorCode::StatementFailed);
3040 assert_eq!(failure_outcome.outcome_kind(), OutcomeKind::Error);
3041 }
3042
3043 #[test]
3044 fn terminal_statement_failures_redact_secret_shaped_upstream_messages() {
3045 let raw_token = "sfpat_driverFailureEcho001";
3046 let failure = QueryFailureStatus {
3047 code: "001003".to_owned(),
3048 sql_state: Some("42000".to_owned()),
3049 message: format!("SQL compilation error near literal '{raw_token}'"),
3050 statement_handle: Some(StatementHandle::new("failed-handle")),
3051 };
3052
3053 let error = terminal_failure_error(SnowflakeErrorCode::StatementFailed, failure);
3054
3055 assert_eq!(error.code, SnowflakeErrorCode::StatementFailed);
3056 assert!(error.message.contains("[REDACTED]"));
3057 assert!(!error.message.contains(raw_token));
3058 }
3059}