1use std::time::Duration;
18
19use franken_snowflake_core::error::{SnowflakeError, SnowflakeErrorCode};
20use franken_snowflake_core::ids::StatementHandle;
21
22use crate::response::{QueryFailureStatus, QueryStatus, ResultSet};
23use crate::status::ResponseClass;
24
25pub const MIN_POLL_INTERVAL: Duration = Duration::from_millis(50);
29
30#[derive(Clone, Copy, Debug, PartialEq, Eq)]
33pub struct PollPlan {
34 pub max_polls: u32,
36 pub poll_interval: Duration,
42 pub partition_concurrency: usize,
45 pub row_cap: Option<usize>,
49}
50
51pub const DEFAULT_PARTITION_CONCURRENCY: usize = 4;
53pub const MAX_PARTITION_CONCURRENCY: usize = 16;
55
56impl Default for PollPlan {
57 fn default() -> Self {
58 Self {
59 max_polls: 120,
60 poll_interval: Duration::from_millis(1_000),
61 partition_concurrency: DEFAULT_PARTITION_CONCURRENCY,
62 row_cap: None,
63 }
64 }
65}
66
67impl PollPlan {
68 #[must_use]
71 pub fn with_max_polls(max_polls: u32) -> Self {
72 Self {
73 max_polls: max_polls.max(1),
74 ..Self::default()
75 }
76 }
77
78 #[must_use]
81 pub fn with_poll_interval(mut self, poll_interval: Duration) -> Self {
82 self.poll_interval = poll_interval.max(MIN_POLL_INTERVAL);
83 self
84 }
85
86 #[must_use]
89 pub fn effective_poll_interval(&self) -> Duration {
90 self.poll_interval.max(MIN_POLL_INTERVAL)
91 }
92
93 #[must_use]
95 pub fn with_partition_concurrency(mut self, concurrency: usize) -> Self {
96 self.partition_concurrency = concurrency.clamp(1, MAX_PARTITION_CONCURRENCY);
97 self
98 }
99
100 #[must_use]
103 pub fn with_row_cap(mut self, row_cap: Option<usize>) -> Self {
104 self.row_cap = row_cap;
105 self
106 }
107
108 #[must_use]
110 pub fn effective_partition_concurrency(&self) -> usize {
111 self.partition_concurrency
112 .clamp(1, MAX_PARTITION_CONCURRENCY)
113 }
114}
115
116#[derive(Clone, Debug, PartialEq)]
120pub struct CompletedStatement {
121 pub statement_handle: StatementHandle,
123 pub result_set: ResultSet,
125 pub rows: Vec<Vec<Option<String>>>,
127 pub fetched_partitions: u32,
129 pub total_partitions: u32,
131}
132
133impl CompletedStatement {
134 #[must_use]
138 pub fn is_partial(&self) -> bool {
139 self.fetched_partitions < self.total_partitions
140 }
141}
142
143#[allow(clippy::large_enum_variant)]
149#[derive(Clone, Debug, PartialEq)]
150pub enum Progress {
151 PollAgain(StatementHandle),
153 FetchPartition {
155 handle: StatementHandle,
157 partition: u32,
159 },
160 Complete(CompletedStatement),
162 TimedOut(QueryFailureStatus),
164 Failed(QueryFailureStatus),
166}
167
168#[derive(Clone, Debug, PartialEq, Eq)]
171pub struct LifecycleError {
172 pub code: LifecycleErrorCode,
174 pub message: String,
176}
177
178#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
180pub enum LifecycleErrorCode {
181 DecodeFailed,
183 UnexpectedStatus,
185 PollQuotaExhausted,
187 PartitionRowMismatch,
189}
190
191impl LifecycleError {
192 fn new(code: LifecycleErrorCode, message: impl Into<String>) -> Self {
193 Self {
194 code,
195 message: message.into(),
196 }
197 }
198
199 #[must_use]
201 pub fn into_snowflake_error(self) -> SnowflakeError {
202 let code = match self.code {
203 LifecycleErrorCode::DecodeFailed
204 | LifecycleErrorCode::UnexpectedStatus
205 | LifecycleErrorCode::PartitionRowMismatch => SnowflakeErrorCode::UpstreamError,
206 LifecycleErrorCode::PollQuotaExhausted => SnowflakeErrorCode::RetryBudgetExhausted,
207 };
208 SnowflakeError::new(code, self.message)
209 }
210}
211
212impl std::fmt::Display for LifecycleError {
213 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
214 write!(f, "{:?}: {}", self.code, self.message)
215 }
216}
217
218impl std::error::Error for LifecycleError {}
219
220#[allow(clippy::large_enum_variant)]
225#[derive(Clone, Debug)]
226enum Phase {
227 Pending,
229 Assembling {
231 result_set: ResultSet,
232 handle: StatementHandle,
233 total: u32,
234 next: u32,
235 rows: Vec<Vec<Option<String>>>,
236 },
237 Done,
239}
240
241#[derive(Clone, Debug)]
245pub struct StatementMachine {
246 poll_plan: PollPlan,
247 polls_done: u32,
248 phase: Phase,
249 drained_rows: usize,
252}
253
254impl StatementMachine {
255 #[must_use]
257 pub fn new(poll_plan: PollPlan) -> Self {
258 Self {
259 poll_plan,
260 polls_done: 0,
261 phase: Phase::Pending,
262 drained_rows: 0,
263 }
264 }
265
266 pub fn drain_rows(&mut self) -> Vec<Vec<Option<String>>> {
272 match &mut self.phase {
273 Phase::Assembling { rows, .. } => {
274 let drained = std::mem::take(rows);
275 self.drained_rows = self.drained_rows.saturating_add(drained.len());
276 drained
277 }
278 Phase::Pending | Phase::Done => Vec::new(),
279 }
280 }
281
282 #[must_use]
284 pub fn result_set(&self) -> Option<&ResultSet> {
285 match &self.phase {
286 Phase::Assembling { result_set, .. } => Some(result_set),
287 Phase::Pending | Phase::Done => None,
288 }
289 }
290
291 #[must_use]
293 pub const fn polls_done(&self) -> u32 {
294 self.polls_done
295 }
296
297 #[must_use]
299 pub fn assembling_window(&self) -> Option<(u32, u32)> {
300 match &self.phase {
301 Phase::Assembling { next, total, .. } => Some((*next, *total)),
302 Phase::Pending | Phase::Done => None,
303 }
304 }
305
306 #[must_use]
308 pub fn rows_assembled(&self) -> usize {
309 match &self.phase {
310 Phase::Assembling { rows, .. } => self.drained_rows.saturating_add(rows.len()),
311 Phase::Pending | Phase::Done => 0,
312 }
313 }
314
315 pub fn complete_early(&mut self) -> Result<CompletedStatement, LifecycleError> {
321 let Phase::Assembling {
322 result_set,
323 handle,
324 total,
325 next,
326 rows,
327 } = std::mem::replace(&mut self.phase, Phase::Done)
328 else {
329 return Err(LifecycleError::new(
330 LifecycleErrorCode::UnexpectedStatus,
331 "early completion requested outside the assembling phase",
332 ));
333 };
334 Ok(CompletedStatement {
335 statement_handle: handle,
336 result_set,
337 rows,
338 fetched_partitions: next,
339 total_partitions: total,
340 })
341 }
342
343 pub fn on_submit(
349 &mut self,
350 class: ResponseClass,
351 body: &[u8],
352 ) -> Result<Progress, LifecycleError> {
353 self.ensure_not_terminal()?;
354 match class {
355 ResponseClass::Completed => self.enter_terminal_result(parse_result_set(body)?),
356 ResponseClass::Running => {
357 let status = parse_query_status(body)?;
358 Ok(Progress::PollAgain(status.statement_handle))
359 }
360 ResponseClass::StatementTimeout => self.enter_terminal_timeout(parse_failure(body)?),
361 ResponseClass::StatementFailed => self.enter_terminal_failure(parse_failure(body)?),
362 ResponseClass::RateLimited | ResponseClass::Other(_) => Err(LifecycleError::new(
363 LifecycleErrorCode::UnexpectedStatus,
364 "submit returned a non-terminal, non-running status",
365 )),
366 }
367 }
368
369 pub fn on_poll(
375 &mut self,
376 class: ResponseClass,
377 body: &[u8],
378 ) -> Result<Progress, LifecycleError> {
379 self.ensure_not_terminal()?;
380 self.polls_done = self.polls_done.saturating_add(1);
381 match class {
382 ResponseClass::Completed => self.enter_terminal_result(parse_result_set(body)?),
383 ResponseClass::StatementTimeout => self.enter_terminal_timeout(parse_failure(body)?),
384 ResponseClass::StatementFailed => self.enter_terminal_failure(parse_failure(body)?),
385 ResponseClass::Running | ResponseClass::RateLimited => {
388 if self.polls_done > self.poll_plan.max_polls {
389 return Err(LifecycleError::new(
390 LifecycleErrorCode::PollQuotaExhausted,
391 format!(
392 "statement still running after {} polls",
393 self.poll_plan.max_polls
394 ),
395 ));
396 }
397 let status = parse_query_status(body)?;
398 Ok(Progress::PollAgain(status.statement_handle))
399 }
400 ResponseClass::Other(_) => Err(LifecycleError::new(
401 LifecycleErrorCode::UnexpectedStatus,
402 "poll returned an unexpected status",
403 )),
404 }
405 }
406
407 pub fn on_partition(
416 &mut self,
417 class: ResponseClass,
418 partition: u32,
419 body: &[u8],
420 ) -> Result<Progress, LifecycleError> {
421 if !matches!(class, ResponseClass::Completed) {
422 self.phase = Phase::Done;
423 return Err(LifecycleError::new(
424 LifecycleErrorCode::UnexpectedStatus,
425 format!("partition {partition} returned a non-200 status"),
426 ));
427 }
428 let Phase::Assembling {
429 result_set,
430 handle,
431 total,
432 next,
433 mut rows,
434 } = std::mem::replace(&mut self.phase, Phase::Done)
435 else {
436 return Err(LifecycleError::new(
437 LifecycleErrorCode::UnexpectedStatus,
438 "partition response arrived outside the assembling phase",
439 ));
440 };
441 if partition != next {
442 self.phase = Phase::Assembling {
443 result_set,
444 handle,
445 total,
446 next,
447 rows,
448 };
449 return Err(LifecycleError::new(
450 LifecycleErrorCode::UnexpectedStatus,
451 format!("expected partition {next}, received {partition}"),
452 ));
453 }
454
455 let mut partition_rows = parse_partition_rows(body)?;
456 validate_partition_row_count(&result_set, partition, partition_rows.len())?;
457 rows.append(&mut partition_rows);
458 let upcoming = next.saturating_add(1);
459 if upcoming >= total {
460 validate_total_row_count(
461 self.drained_rows.saturating_add(rows.len()),
462 result_set.result_set_meta_data.num_rows,
463 )?;
464 Ok(Progress::Complete(CompletedStatement {
465 statement_handle: handle,
466 result_set,
467 rows,
468 fetched_partitions: total,
469 total_partitions: total,
470 }))
471 } else {
472 let resume = handle.clone();
473 self.phase = Phase::Assembling {
474 result_set,
475 handle,
476 total,
477 next: upcoming,
478 rows,
479 };
480 Ok(Progress::FetchPartition {
481 handle: resume,
482 partition: upcoming,
483 })
484 }
485 }
486
487 fn ensure_not_terminal(&self) -> Result<(), LifecycleError> {
494 if matches!(self.phase, Phase::Done) {
495 return Err(LifecycleError::new(
496 LifecycleErrorCode::UnexpectedStatus,
497 "statement machine already reached a terminal state",
498 ));
499 }
500 Ok(())
501 }
502
503 fn enter_terminal_timeout(
504 &mut self,
505 failure: QueryFailureStatus,
506 ) -> Result<Progress, LifecycleError> {
507 self.phase = Phase::Done;
508 Ok(Progress::TimedOut(failure))
509 }
510
511 fn enter_terminal_failure(
512 &mut self,
513 failure: QueryFailureStatus,
514 ) -> Result<Progress, LifecycleError> {
515 self.phase = Phase::Done;
516 Ok(Progress::Failed(failure))
517 }
518
519 fn enter_terminal_result(&mut self, result_set: ResultSet) -> Result<Progress, LifecycleError> {
521 let handle = result_set.statement_handle.clone();
522 if result_set.is_multi_statement() {
526 self.phase = Phase::Done;
527 let rows = result_set.data.clone();
528 return Ok(Progress::Complete(CompletedStatement {
529 statement_handle: handle,
530 result_set,
531 rows,
532 fetched_partitions: 1,
533 total_partitions: 1,
534 }));
535 }
536 let total = partition_total(&result_set);
537 let rows = result_set.data.clone();
538 validate_partition_row_count(&result_set, 0, rows.len())?;
539 if total <= 1 {
540 self.phase = Phase::Done;
541 validate_total_row_count(rows.len(), result_set.result_set_meta_data.num_rows)?;
544 Ok(Progress::Complete(CompletedStatement {
545 statement_handle: handle,
546 result_set,
547 rows,
548 fetched_partitions: 1,
549 total_partitions: 1,
550 }))
551 } else {
552 let resume = handle.clone();
553 self.phase = Phase::Assembling {
554 result_set,
555 handle,
556 total,
557 next: 1,
558 rows,
559 };
560 Ok(Progress::FetchPartition {
561 handle: resume,
562 partition: 1,
563 })
564 }
565 }
566}
567
568#[must_use]
571fn partition_total(result_set: &ResultSet) -> u32 {
572 u32::try_from(result_set.result_set_meta_data.partition_info.len().max(1)).unwrap_or(u32::MAX)
573}
574
575fn parse_result_set(body: &[u8]) -> Result<ResultSet, LifecycleError> {
576 serde_json::from_slice(body)
577 .map_err(|error| LifecycleError::new(LifecycleErrorCode::DecodeFailed, error.to_string()))
578}
579
580fn parse_query_status(body: &[u8]) -> Result<QueryStatus, LifecycleError> {
581 serde_json::from_slice(body)
582 .map_err(|error| LifecycleError::new(LifecycleErrorCode::DecodeFailed, error.to_string()))
583}
584
585fn parse_failure(body: &[u8]) -> Result<QueryFailureStatus, LifecycleError> {
586 serde_json::from_slice(body)
587 .map_err(|error| LifecycleError::new(LifecycleErrorCode::DecodeFailed, error.to_string()))
588}
589
590pub fn parse_partition_rows(body: &[u8]) -> Result<Vec<Vec<Option<String>>>, LifecycleError> {
595 #[derive(serde::Deserialize)]
603 struct PartitionEnvelope {
604 data: Vec<Vec<Option<String>>>,
605 }
606 if let Ok(envelope) = serde_json::from_slice::<PartitionEnvelope>(body) {
607 return Ok(envelope.data);
608 }
609 serde_json::from_slice(body)
610 .map_err(|error| LifecycleError::new(LifecycleErrorCode::DecodeFailed, error.to_string()))
611}
612
613fn validate_partition_row_count(
614 result_set: &ResultSet,
615 partition: u32,
616 actual_rows: usize,
617) -> Result<(), LifecycleError> {
618 let Some(expected) = usize::try_from(partition)
619 .ok()
620 .and_then(|index| result_set.result_set_meta_data.partition_info.get(index))
621 .map(|info| info.row_count)
622 else {
623 return Ok(());
624 };
625 if expected < 0 {
626 return Err(LifecycleError::new(
627 LifecycleErrorCode::PartitionRowMismatch,
628 format!("partition {partition} rowCount is negative"),
629 ));
630 }
631 if i64::try_from(actual_rows).ok() != Some(expected) {
632 return Err(LifecycleError::new(
633 LifecycleErrorCode::PartitionRowMismatch,
634 format!("partition {partition} returned {actual_rows} rows but rowCount is {expected}"),
635 ));
636 }
637 Ok(())
638}
639
640fn validate_total_row_count(actual_rows: usize, expected: i64) -> Result<(), LifecycleError> {
641 if expected < 0 {
642 return Err(LifecycleError::new(
643 LifecycleErrorCode::PartitionRowMismatch,
644 "numRows is negative",
645 ));
646 }
647 if i64::try_from(actual_rows).ok() != Some(expected) {
648 return Err(LifecycleError::new(
649 LifecycleErrorCode::PartitionRowMismatch,
650 format!("assembled {actual_rows} rows but numRows is {expected}"),
651 ));
652 }
653 Ok(())
654}
655
656#[cfg(test)]
657mod tests {
658 use super::*;
659
660 #[test]
661 fn poll_plan_interval_is_always_sane() {
662 assert_eq!(
664 PollPlan::default().poll_interval,
665 Duration::from_millis(1_000)
666 );
667 assert!(PollPlan::default().effective_poll_interval() >= MIN_POLL_INTERVAL);
668 assert_eq!(
670 PollPlan::with_max_polls(5).poll_interval,
671 PollPlan::default().poll_interval
672 );
673 assert_eq!(
676 PollPlan::default()
677 .with_poll_interval(Duration::ZERO)
678 .poll_interval,
679 MIN_POLL_INTERVAL
680 );
681 let hand_set = PollPlan {
682 poll_interval: Duration::ZERO,
683 ..PollPlan::default()
684 };
685 assert_eq!(hand_set.effective_poll_interval(), MIN_POLL_INTERVAL);
686 }
687
688 #[test]
689 fn partition_total_treats_absent_or_single_info_as_inline() -> Result<(), String> {
690 let body = br#"{"resultSetMetaData":{"numRows":0,"format":"jsonv2","rowType":[]},
691 "data":[],"code":"090001","statementHandle":"h"}"#;
692 let result_set = parse_result_set(body).map_err(|error| error.to_string())?;
693 assert_eq!(partition_total(&result_set), 1);
694 Ok(())
695 }
696
697 #[test]
698 fn single_partition_completes_immediately() -> Result<(), String> {
699 let body = br#"{"resultSetMetaData":{"numRows":1,"format":"jsonv2",
700 "rowType":[{"name":"A","type":"TEXT","nullable":false}],
701 "partitionInfo":[{"rowCount":1,"compressedSize":1,"uncompressedSize":1}]},
702 "data":[["x"]],"code":"090001","statementHandle":"h"}"#;
703 let mut machine = StatementMachine::new(PollPlan::default());
704 match machine.on_submit(ResponseClass::Completed, body) {
705 Ok(Progress::Complete(done)) => {
706 assert_eq!(done.rows.len(), 1);
707 assert_eq!(done.statement_handle, StatementHandle::new("h"));
708 Ok(())
709 }
710 other => Err(format!("expected Complete, got {other:?}")),
711 }
712 }
713
714 #[test]
715 fn running_then_completed_polls_then_finishes() -> Result<(), String> {
716 let running = br#"{"code":"333334","statementHandle":"h2"}"#;
717 let completed = br#"{"resultSetMetaData":{"numRows":0,"format":"jsonv2","rowType":[]},
718 "data":[],"code":"090001","statementHandle":"h2"}"#;
719 let mut machine = StatementMachine::new(PollPlan::default());
720 match machine.on_submit(ResponseClass::Running, running) {
721 Ok(Progress::PollAgain(h)) => assert_eq!(h, StatementHandle::new("h2")),
722 other => return Err(format!("expected PollAgain, got {other:?}")),
723 }
724 match machine.on_poll(ResponseClass::Completed, completed) {
725 Ok(Progress::Complete(_)) => Ok(()),
726 other => Err(format!("expected Complete, got {other:?}")),
727 }
728 }
729
730 #[test]
731 fn poll_quota_is_enforced() -> Result<(), String> {
732 let running = br#"{"code":"333334","statementHandle":"h3"}"#;
733 let mut machine = StatementMachine::new(PollPlan::with_max_polls(2));
734 machine
735 .on_poll(ResponseClass::Running, running)
736 .map_err(|e| e.to_string())?;
737 machine
738 .on_poll(ResponseClass::Running, running)
739 .map_err(|e| e.to_string())?;
740 match machine.on_poll(ResponseClass::Running, running) {
741 Err(error) => {
742 assert_eq!(error.code, LifecycleErrorCode::PollQuotaExhausted);
743 Ok(())
744 }
745 Ok(progress) => Err(format!("expected quota error, got {progress:?}")),
746 }
747 }
748
749 #[test]
750 fn timeout_and_failure_are_distinct_terminal_states() -> Result<(), String> {
751 let timeout = br#"{"code":"000630","message":"timed out","statementHandle":"h"}"#;
752 let failure = br#"{"code":"001003","message":"bad sql","statementHandle":"h"}"#;
753 let mut machine = StatementMachine::new(PollPlan::default());
754 assert!(matches!(
755 machine.on_submit(ResponseClass::StatementTimeout, timeout),
756 Ok(Progress::TimedOut(_))
757 ));
758 let mut other = StatementMachine::new(PollPlan::default());
759 assert!(matches!(
760 other.on_submit(ResponseClass::StatementFailed, failure),
761 Ok(Progress::Failed(_))
762 ));
763 Ok(())
764 }
765
766 #[test]
767 fn timeout_and_failure_close_the_machine() -> Result<(), String> {
768 let timeout = br#"{"code":"000630","message":"timed out","statementHandle":"h"}"#;
769 let completed = br#"{"resultSetMetaData":{"numRows":0,"format":"jsonv2","rowType":[]},
770 "data":[],"code":"090001","statementHandle":"h"}"#;
771 let mut machine = StatementMachine::new(PollPlan::default());
772 assert!(matches!(
773 machine.on_submit(ResponseClass::StatementTimeout, timeout),
774 Ok(Progress::TimedOut(_))
775 ));
776 match machine.on_poll(ResponseClass::Completed, completed) {
777 Err(error) => {
778 assert_eq!(error.code, LifecycleErrorCode::UnexpectedStatus);
779 Ok(())
780 }
781 Ok(progress) => Err(format!(
782 "expected terminal machine refusal, got {progress:?}"
783 )),
784 }
785 }
786
787 #[test]
788 fn terminal_machine_refuses_every_reentry_path() -> Result<(), String> {
789 let failure = br#"{"code":"001003","message":"bad sql","statementHandle":"h"}"#;
790 let completed = br#"{"resultSetMetaData":{"numRows":0,"format":"jsonv2","rowType":[]},
791 "data":[],"code":"090001","statementHandle":"h"}"#;
792
793 let mut after_failure = StatementMachine::new(PollPlan::default());
796 assert!(matches!(
797 after_failure.on_submit(ResponseClass::StatementFailed, failure),
798 Ok(Progress::Failed(_))
799 ));
800 assert_eq!(
801 after_failure
802 .on_poll(ResponseClass::Completed, completed)
803 .map(|_| ())
804 .unwrap_err()
805 .code,
806 LifecycleErrorCode::UnexpectedStatus
807 );
808
809 let mut after_success = StatementMachine::new(PollPlan::default());
812 assert!(matches!(
813 after_success.on_submit(ResponseClass::Completed, completed),
814 Ok(Progress::Complete(_))
815 ));
816 assert_eq!(
817 after_success
818 .on_submit(ResponseClass::Completed, completed)
819 .map(|_| ())
820 .unwrap_err()
821 .code,
822 LifecycleErrorCode::UnexpectedStatus
823 );
824 assert_eq!(
825 after_success
826 .on_poll(ResponseClass::Completed, completed)
827 .map(|_| ())
828 .unwrap_err()
829 .code,
830 LifecycleErrorCode::UnexpectedStatus
831 );
832 assert_eq!(after_success.polls_done(), 0);
834 Ok(())
835 }
836
837 #[test]
838 fn multi_partition_assembles_rows_in_order() -> Result<(), String> {
839 let terminal = br#"{"resultSetMetaData":{"numRows":5,"format":"jsonv2",
841 "rowType":[{"name":"ID","type":"FIXED","nullable":false},
842 {"name":"NAME","type":"TEXT","nullable":false}],
843 "partitionInfo":[{"rowCount":2,"compressedSize":1,"uncompressedSize":1},
844 {"rowCount":2,"compressedSize":1,"uncompressedSize":1},
845 {"rowCount":1,"compressedSize":1,"uncompressedSize":1}]},
846 "data":[["1","a"],["2","b"]],"code":"090001","statementHandle":"hp"}"#;
847 let mut machine = StatementMachine::new(PollPlan::default());
848 let first = machine.on_submit(ResponseClass::Completed, terminal);
849 let handle = match first {
850 Ok(Progress::FetchPartition {
851 handle,
852 partition: 1,
853 }) => handle,
854 other => return Err(format!("expected FetchPartition 1, got {other:?}")),
855 };
856 assert_eq!(handle, StatementHandle::new("hp"));
857 match machine.on_partition(ResponseClass::Completed, 1, br#"[["3","c"],["4","d"]]"#) {
858 Ok(Progress::FetchPartition { partition: 2, .. }) => {}
859 other => return Err(format!("expected FetchPartition 2, got {other:?}")),
860 }
861 match machine.on_partition(ResponseClass::Completed, 2, br#"[["5","e"]]"#) {
862 Ok(Progress::Complete(done)) => {
863 assert_eq!(done.rows.len(), 5);
864 assert_eq!(
865 done.rows[4],
866 vec![Some("5".to_owned()), Some("e".to_owned())]
867 );
868 Ok(())
869 }
870 other => Err(format!("expected Complete, got {other:?}")),
871 }
872 }
873
874 #[test]
875 fn parse_partition_rows_accepts_live_object_data_form() -> Result<(), String> {
876 let rows = parse_partition_rows(br#"{"data":[["3","c"],["4","d"]]}"#)
881 .map_err(|error| error.to_string())?;
882 assert_eq!(
883 rows,
884 vec![
885 vec![Some("3".to_owned()), Some("c".to_owned())],
886 vec![Some("4".to_owned()), Some("d".to_owned())],
887 ]
888 );
889 Ok(())
890 }
891
892 #[test]
893 fn parse_partition_rows_still_accepts_bare_array_form() -> Result<(), String> {
894 let rows = parse_partition_rows(br#"[["5","e"]]"#).map_err(|error| error.to_string())?;
897 assert_eq!(rows, vec![vec![Some("5".to_owned()), Some("e".to_owned())]]);
898 Ok(())
899 }
900
901 #[test]
902 fn partition_info_decodes_when_compressed_size_is_omitted() -> Result<(), String> {
903 let body = br#"{"resultSetMetaData":{"numRows":2,"format":"jsonv2",
908 "rowType":[{"name":"A","type":"TEXT","nullable":false}],
909 "partitionInfo":[{"rowCount":2,"uncompressedSize":64}]},
910 "data":[["x"],["y"]],"code":"090001","statementHandle":"h"}"#;
911 let result_set = parse_result_set(body).map_err(|error| error.to_string())?;
912 let info = result_set
913 .result_set_meta_data
914 .partition_info
915 .first()
916 .ok_or("expected partition_info[0]")?;
917 assert_eq!(info.row_count, 2);
918 assert_eq!(info.compressed_size, None);
919 assert_eq!(info.uncompressed_size, Some(64));
920 Ok(())
921 }
922
923 #[test]
924 fn live_shaped_multi_partition_flow_decodes_end_to_end() -> Result<(), String> {
925 let terminal = br#"{"resultSetMetaData":{"numRows":3,"format":"jsonv2",
930 "rowType":[{"name":"ID","type":"FIXED","nullable":false}],
931 "partitionInfo":[{"rowCount":2,"uncompressedSize":16},
932 {"rowCount":1,"compressedSize":8,"uncompressedSize":16}]},
933 "data":[["1"],["2"]],"code":"090001","statementHandle":"hp"}"#;
934 let mut machine = StatementMachine::new(PollPlan::default());
935 let handle = match machine.on_submit(ResponseClass::Completed, terminal) {
936 Ok(Progress::FetchPartition {
937 handle,
938 partition: 1,
939 }) => handle,
940 other => return Err(format!("expected FetchPartition 1, got {other:?}")),
941 };
942 assert_eq!(handle, StatementHandle::new("hp"));
943 match machine.on_partition(ResponseClass::Completed, 1, br#"{"data":[["3"]]}"#) {
944 Ok(Progress::Complete(done)) => {
945 assert_eq!(done.rows.len(), 3);
946 assert_eq!(done.rows[2], vec![Some("3".to_owned())]);
947 Ok(())
948 }
949 other => Err(format!("expected Complete, got {other:?}")),
950 }
951 }
952
953 #[test]
954 fn row_count_mismatch_is_rejected() {
955 let terminal = br#"{"resultSetMetaData":{"numRows":99,"format":"jsonv2",
956 "rowType":[{"name":"A","type":"TEXT","nullable":false}],
957 "partitionInfo":[{"rowCount":1,"compressedSize":1,"uncompressedSize":1},
958 {"rowCount":1,"compressedSize":1,"uncompressedSize":1}]},
959 "data":[["x"]],"code":"090001","statementHandle":"hp"}"#;
960 let mut machine = StatementMachine::new(PollPlan::default());
961 let _ = machine.on_submit(ResponseClass::Completed, terminal);
962 let result = machine.on_partition(ResponseClass::Completed, 1, br#"[["y"]]"#);
963 assert!(matches!(
964 result,
965 Err(LifecycleError {
966 code: LifecycleErrorCode::PartitionRowMismatch,
967 ..
968 })
969 ));
970 }
971
972 #[test]
973 fn fetched_partition_row_count_mismatch_is_rejected_before_total_can_compensate() {
974 let terminal = br#"{"resultSetMetaData":{"numRows":3,"format":"jsonv2",
975 "rowType":[{"name":"A","type":"TEXT","nullable":false}],
976 "partitionInfo":[{"rowCount":1,"compressedSize":1,"uncompressedSize":1},
977 {"rowCount":1,"compressedSize":1,"uncompressedSize":1},
978 {"rowCount":1,"compressedSize":1,"uncompressedSize":1}]},
979 "data":[["inline"]],"code":"090001","statementHandle":"hp"}"#;
980 let mut machine = StatementMachine::new(PollPlan::default());
981 assert!(matches!(
982 machine.on_submit(ResponseClass::Completed, terminal),
983 Ok(Progress::FetchPartition { partition: 1, .. })
984 ));
985
986 let result = machine.on_partition(
987 ResponseClass::Completed,
988 1,
989 br#"[["too-many"],["would-hide-empty-next"]]"#,
990 );
991
992 assert!(matches!(
993 result,
994 Err(LifecycleError {
995 code: LifecycleErrorCode::PartitionRowMismatch,
996 ..
997 })
998 ));
999 }
1000
1001 #[test]
1002 fn empty_fetched_partition_with_positive_row_count_is_rejected() {
1003 let terminal = br#"{"resultSetMetaData":{"numRows":2,"format":"jsonv2",
1004 "rowType":[{"name":"A","type":"TEXT","nullable":false}],
1005 "partitionInfo":[{"rowCount":1,"compressedSize":1,"uncompressedSize":1},
1006 {"rowCount":1,"compressedSize":1,"uncompressedSize":1}]},
1007 "data":[["inline"]],"code":"090001","statementHandle":"hp"}"#;
1008 let mut machine = StatementMachine::new(PollPlan::default());
1009 assert!(matches!(
1010 machine.on_submit(ResponseClass::Completed, terminal),
1011 Ok(Progress::FetchPartition { partition: 1, .. })
1012 ));
1013
1014 let result = machine.on_partition(ResponseClass::Completed, 1, br#"[]"#);
1015
1016 assert!(matches!(
1017 result,
1018 Err(LifecycleError {
1019 code: LifecycleErrorCode::PartitionRowMismatch,
1020 ..
1021 })
1022 ));
1023 }
1024
1025 #[test]
1026 fn inline_partition_info_row_count_mismatch_is_rejected() {
1027 let terminal = br#"{"resultSetMetaData":{"numRows":1,"format":"jsonv2",
1028 "rowType":[{"name":"A","type":"TEXT","nullable":false}],
1029 "partitionInfo":[{"rowCount":2,"compressedSize":1,"uncompressedSize":1}]},
1030 "data":[["x"]],"code":"090001","statementHandle":"h"}"#;
1031 let mut machine = StatementMachine::new(PollPlan::default());
1032
1033 assert!(matches!(
1034 machine.on_submit(ResponseClass::Completed, terminal),
1035 Err(LifecycleError {
1036 code: LifecycleErrorCode::PartitionRowMismatch,
1037 ..
1038 })
1039 ));
1040 }
1041
1042 #[test]
1043 fn single_partition_row_count_mismatch_is_rejected() {
1044 let terminal = br#"{"resultSetMetaData":{"numRows":5,"format":"jsonv2",
1047 "rowType":[{"name":"A","type":"TEXT","nullable":false}],
1048 "partitionInfo":[{"rowCount":5,"compressedSize":1,"uncompressedSize":1}]},
1049 "data":[["x"]],"code":"090001","statementHandle":"h"}"#;
1050 let mut machine = StatementMachine::new(PollPlan::default());
1051 assert!(matches!(
1052 machine.on_submit(ResponseClass::Completed, terminal),
1053 Err(LifecycleError {
1054 code: LifecycleErrorCode::PartitionRowMismatch,
1055 ..
1056 })
1057 ));
1058 }
1059}