1use crate::model::DirectedReadOptions;
16use crate::model::execute_sql_request::QueryMode;
17use crate::model::execute_sql_request::QueryOptions;
18use crate::model::request_options::Priority;
19use crate::types::Type;
20use crate::value::Value;
21
22use google_cloud_gax::backoff_policy::BackoffPolicyArg;
23use google_cloud_gax::options::RequestOptions as GaxRequestOptions;
24use google_cloud_gax::retry_policy::RetryPolicyArg;
25use std::collections::BTreeMap;
26use std::time::Duration;
27
28#[derive(Clone, Debug)]
38pub struct StatementBuilder {
39 sql: String,
40 params: BTreeMap<String, Value>,
41 param_types: BTreeMap<String, Type>,
42 request_options: Option<crate::model::RequestOptions>,
43 directed_read_options: Option<DirectedReadOptions>,
44 query_options: Option<QueryOptions>,
45 query_mode: Option<QueryMode>,
46 last_statement: bool,
47 gax_options: GaxRequestOptions,
48}
49
50impl StatementBuilder {
51 pub(crate) fn new(sql: impl Into<String>) -> Self {
52 Self {
53 sql: sql.into(),
54 params: BTreeMap::new(),
55 param_types: BTreeMap::new(),
56 request_options: None,
57 directed_read_options: None,
58 query_options: None,
59 query_mode: None,
60 last_statement: false,
61 gax_options: GaxRequestOptions::default(),
62 }
63 }
64
65 pub fn add_param<T: Into<Value>>(mut self, name: impl Into<String>, value: T) -> Self {
73 self.params.insert(name.into(), value.into());
74 self
75 }
76
77 pub fn add_typed_param<T: Into<Value>>(
82 mut self,
83 name: impl Into<String>,
84 value: T,
85 param_type: Type,
86 ) -> Self {
87 let name = name.into();
88 self.params.insert(name.clone(), value.into());
89 self.param_types.insert(name, param_type);
90 self
91 }
92
93 pub fn set_request_tag(mut self, tag: impl Into<String>) -> Self {
105 self.request_options
106 .get_or_insert_with(crate::model::RequestOptions::default)
107 .request_tag = tag.into();
108 self
109 }
110
111 pub fn set_priority(mut self, priority: Priority) -> Self {
122 self.request_options
123 .get_or_insert_with(crate::model::RequestOptions::default)
124 .priority = priority;
125 self
126 }
127
128 pub fn set_directed_read_options(mut self, options: DirectedReadOptions) -> Self {
142 self.directed_read_options = Some(options);
143 self
144 }
145
146 pub fn set_query_options(mut self, options: QueryOptions) -> Self {
159 self.query_options = Some(options);
160 self
161 }
162
163 pub fn set_query_mode(mut self, mode: QueryMode) -> Self {
174 self.query_mode = Some(mode);
175 self
176 }
177
178 pub fn with_attempt_timeout(mut self, timeout: Duration) -> Self {
180 self.gax_options.set_attempt_timeout(timeout);
181 self
182 }
183
184 pub fn set_last_statement(mut self, last_statement: bool) -> Self {
208 self.last_statement = last_statement;
209 self
210 }
211
212 pub fn with_retry_policy(mut self, policy: impl Into<RetryPolicyArg>) -> Self {
214 self.gax_options.set_retry_policy(policy);
215 self
216 }
217
218 pub fn with_backoff_policy(mut self, policy: impl Into<BackoffPolicyArg>) -> Self {
220 self.gax_options.set_backoff_policy(policy);
221 self
222 }
223
224 pub fn build(self) -> Statement {
226 Statement {
227 sql: self.sql,
228 params: self.params,
229 param_types: self.param_types,
230 request_options: self.request_options,
231 directed_read_options: self.directed_read_options,
232 query_options: self.query_options,
233 query_mode: self.query_mode,
234 last_statement: self.last_statement,
235 gax_options: self.gax_options,
236 }
237 }
238}
239
240#[derive(Clone, Debug)]
264pub struct Statement {
265 pub(crate) sql: String,
266 pub(crate) params: BTreeMap<String, Value>,
267 pub(crate) param_types: BTreeMap<String, Type>,
268 pub(crate) request_options: Option<crate::model::RequestOptions>,
269 pub(crate) directed_read_options: Option<DirectedReadOptions>,
270 pub(crate) query_options: Option<QueryOptions>,
271 pub(crate) query_mode: Option<QueryMode>,
272 pub(crate) last_statement: bool,
273 pub(crate) gax_options: GaxRequestOptions,
274}
275
276impl Statement {
277 pub fn builder(sql: impl Into<String>) -> StatementBuilder {
279 StatementBuilder::new(sql)
280 }
281
282 pub fn sql(&self) -> &str {
284 &self.sql
285 }
286
287 pub(crate) fn gax_options(&self) -> &GaxRequestOptions {
288 &self.gax_options
289 }
290
291 pub(crate) fn with_gax_options(mut self, options: GaxRequestOptions) -> Self {
293 self.gax_options = options;
294 self
295 }
296
297 pub fn set_query_mode(mut self, mode: QueryMode) -> Self {
315 self.query_mode = Some(mode);
316 self
317 }
318
319 pub fn set_last_statement(mut self, last_statement: bool) -> Self {
343 self.last_statement = last_statement;
344 self
345 }
346
347 fn into_parts(
348 self,
349 ) -> (
350 String,
351 Option<wkt::Struct>,
352 std::collections::HashMap<String, crate::model::Type>,
353 ) {
354 let params: Option<wkt::Struct> = if self.params.is_empty() {
355 None
356 } else {
357 Some(
358 self.params
359 .into_iter()
360 .map(|(k, v)| (k, v.into_serde_value()))
361 .collect(),
362 )
363 };
364 let param_types: std::collections::HashMap<String, crate::model::Type> = self
365 .param_types
366 .into_iter()
367 .map(|(k, v)| (k, v.0))
368 .collect();
369 (self.sql, params, param_types)
370 }
371
372 pub(crate) fn into_request(self) -> crate::model::ExecuteSqlRequest {
373 let request_options = self.request_options.clone();
374 let directed_read_options = self.directed_read_options.clone();
375 let query_options = self.query_options.clone();
376 let query_mode = self.query_mode.clone();
377 let last_statement = self.last_statement;
378 let (sql, params, param_types) = self.into_parts();
379 crate::model::ExecuteSqlRequest::default()
380 .set_sql(sql)
381 .set_or_clear_params(params)
382 .set_param_types(param_types)
383 .set_or_clear_request_options(request_options)
384 .set_or_clear_directed_read_options(directed_read_options)
385 .set_or_clear_query_options(query_options)
386 .set_query_mode(query_mode.unwrap_or_default())
387 .set_last_statement(last_statement)
388 }
389
390 pub(crate) fn into_batch_statement(self) -> crate::model::execute_batch_dml_request::Statement {
391 let (sql, params, param_types) = self.into_parts();
392 crate::model::execute_batch_dml_request::Statement::default()
393 .set_sql(sql)
394 .set_or_clear_params(params)
395 .set_param_types(param_types)
396 }
397
398 pub(crate) fn into_partition_query_request(self) -> crate::model::PartitionQueryRequest {
399 let (sql, params, param_types) = self.into_parts();
400 crate::model::PartitionQueryRequest::default()
401 .set_sql(sql)
402 .set_or_clear_params(params)
403 .set_param_types(param_types)
404 }
405}
406
407impl From<StatementBuilder> for Statement {
408 fn from(builder: StatementBuilder) -> Self {
409 builder.build()
410 }
411}
412
413impl From<String> for Statement {
414 fn from(sql: String) -> Self {
415 Statement::builder(sql).build()
416 }
417}
418
419impl From<&str> for Statement {
420 fn from(sql: &str) -> Self {
421 Statement::builder(sql).build()
422 }
423}
424
425#[cfg(test)]
426mod tests {
427 use super::*;
428 use crate::to_value::ToValue;
429 use anyhow::Context;
430
431 #[test]
432 fn test_auto_traits() {
433 static_assertions::assert_impl_all!(Statement: Clone, std::fmt::Debug, Send, Sync);
434 static_assertions::assert_impl_all!(StatementBuilder: Clone, std::fmt::Debug, Send, Sync);
435 }
436
437 #[test]
438 fn test_untyped_param() {
439 let stmt = Statement::builder("SELECT * FROM users WHERE age > @age")
440 .add_param("age", 21)
441 .build();
442
443 assert_eq!(stmt.sql, "SELECT * FROM users WHERE age > @age");
444 assert_eq!(stmt.param_types.len(), 0);
445 assert_eq!(stmt.params.len(), 1);
446 assert_eq!(stmt.request_options, None);
447
448 let val = stmt.params.get("age").unwrap();
449 assert_eq!(val.as_string(), "21");
450 }
451
452 #[test]
453 fn test_param_direct_types() {
454 let id_str = "user-123";
455 let id_string = String::from("user-456");
456 let age_i64 = 42i64;
457 let active_bool = true;
458
459 let stmt = Statement::builder(
460 "SELECT * FROM users WHERE id = @id AND age = @age AND active = @active",
461 )
462 .add_param("id", id_str)
463 .add_param("id2", id_string)
464 .add_param("age", age_i64)
465 .add_param("active", active_bool)
466 .build();
467
468 assert_eq!(stmt.params.get("id").unwrap().as_string(), "user-123");
469 assert_eq!(stmt.params.get("id2").unwrap().as_string(), "user-456");
470 assert_eq!(stmt.params.get("age").unwrap().as_string(), "42");
471 assert!(stmt.params.get("active").unwrap().as_bool());
472 }
473
474 #[test]
475 #[allow(clippy::needless_borrows_for_generic_args)]
476 fn test_param_borrowed_types() {
477 use crate::types;
478 let id_str = "user-123";
479 let id_string = String::from("user-456");
480 let age_i64 = 42i64;
481 let active_bool = true;
482
483 let stmt = Statement::builder(
484 "SELECT * FROM users WHERE id = @id AND age = @age AND active = @active",
485 )
486 .add_param("id", &id_str)
487 .add_param("id2", &id_string)
488 .add_param("age", &age_i64)
489 .add_param("active", &active_bool)
490 .add_typed_param("role", &"admin", types::string())
491 .build();
492
493 assert_eq!(stmt.params.get("id").unwrap().as_string(), "user-123");
494 assert_eq!(stmt.params.get("id2").unwrap().as_string(), "user-456");
495 assert_eq!(stmt.params.get("age").unwrap().as_string(), "42");
496 assert!(stmt.params.get("active").unwrap().as_bool());
497 assert_eq!(stmt.params.get("role").unwrap().as_string(), "admin");
498 }
499
500 #[test]
501 fn test_param_owned_value() {
502 let value = 21i32.to_value();
503 let stmt = Statement::builder("SELECT * FROM users WHERE age > @age")
504 .add_param("age", value)
505 .build();
506
507 assert_eq!(stmt.param_types.len(), 0);
508 assert_eq!(stmt.params.len(), 1);
509 assert_eq!(
510 stmt.params
511 .get("age")
512 .expect("parameter 'age' should be present")
513 .as_string(),
514 "21"
515 );
516 }
517
518 #[test]
519 fn test_typed_param_owned_value() {
520 use crate::types;
521 let value = "user-123".to_value();
522 let stmt = Statement::builder("SELECT * FROM users WHERE id = @id")
523 .add_typed_param("id", value, types::string())
524 .build();
525
526 assert_eq!(stmt.param_types.len(), 1);
527 assert_eq!(
528 stmt.param_types
529 .get("id")
530 .expect("parameter type for 'id' should be present"),
531 &types::string()
532 );
533 assert_eq!(stmt.params.len(), 1);
534 assert_eq!(
535 stmt.params
536 .get("id")
537 .expect("parameter 'id' should be present")
538 .as_string(),
539 "user-123"
540 );
541 }
542
543 #[test]
544 fn test_param_turbofished_ref_none() {
545 let stmt_ref_turbofished = Statement::builder("SELECT * FROM users WHERE age > @age")
546 .add_param::<&Option<i64>>("age", &None)
547 .build();
548 let stmt_owned_none = Statement::builder("SELECT * FROM users WHERE age > @age")
549 .add_param::<Option<i64>>("age", None)
550 .build();
551 assert_eq!(stmt_ref_turbofished.params, stmt_owned_none.params);
552 }
553
554 #[test]
555 fn test_param_untyped_null() {
556 let stmt_null = Statement::builder("SELECT * FROM users WHERE age > @age")
557 .add_param("age", Value::null())
558 .build();
559 let stmt_unit_none = Statement::builder("SELECT * FROM users WHERE age > @age")
560 .add_param("age", None::<()>)
561 .build();
562 let stmt_value_none = Statement::builder("SELECT * FROM users WHERE age > @age")
563 .add_param("age", None::<Value>)
564 .build();
565 assert_eq!(stmt_null.params, stmt_unit_none.params);
566 assert_eq!(stmt_null.params, stmt_value_none.params);
567 }
568
569 #[test]
570 fn test_typed_param() {
571 use crate::types;
572 let stmt = Statement::builder("SELECT * FROM users WHERE id = @id")
573 .add_typed_param("id", "user-123", types::string())
574 .build();
575
576 assert_eq!(stmt.param_types.len(), 1);
577 assert_eq!(stmt.param_types.get("id").unwrap(), &types::string());
578
579 assert_eq!(stmt.params.len(), 1);
580 let val = stmt.params.get("id").unwrap();
581 assert_eq!(val.as_string(), "user-123");
582 }
583
584 #[test]
585 fn test_multiple_params() {
586 use crate::types;
587 let stmt = Statement::builder("SELECT * FROM users WHERE age > @age AND role = @role")
588 .add_param("age", 21)
589 .add_typed_param("role", "admin", types::string())
590 .build();
591
592 assert_eq!(stmt.params.len(), 2);
593 assert_eq!(stmt.param_types.len(), 1);
594 }
595
596 #[test]
597 fn test_from_string_conversions() {
598 let stmt_str: Statement = "SELECT 1".into();
599 let stmt_string: Statement = "SELECT 1".to_string().into();
600 assert_eq!(stmt_str.sql, "SELECT 1");
601 assert_eq!(stmt_string.sql, "SELECT 1");
602 assert!(stmt_str.params.is_empty());
603 assert!(stmt_string.params.is_empty());
604 assert!(stmt_str.param_types.is_empty());
605 assert!(stmt_string.param_types.is_empty());
606 assert!(stmt_str.request_options.is_none());
607 assert!(stmt_string.request_options.is_none());
608 }
609
610 #[test]
611 fn test_from_builder_conversion() {
612 use crate::types;
613 let builder = Statement::builder("SELECT * FROM users WHERE age > @age AND role = @role")
614 .add_param("age", 21)
615 .add_typed_param("role", "admin", types::string());
616
617 let stmt: Statement = builder.into();
618 assert_eq!(
619 stmt.sql,
620 "SELECT * FROM users WHERE age > @age AND role = @role"
621 );
622 assert_eq!(stmt.params.len(), 2);
623 assert_eq!(stmt.param_types.len(), 1);
624 }
625
626 #[test]
627 fn test_into_request() {
628 use crate::types;
629 let stmt = Statement::builder("SELECT * FROM users WHERE age > @age AND role = @role")
630 .add_param("age", 21)
631 .add_typed_param("role", "admin", types::string())
632 .build();
633
634 let req = stmt.into_request();
635
636 let params = req
637 .params
638 .expect("ExecuteSqlRequest parameters should be set after into_request conversion");
639 assert_eq!(params.len(), 2);
640 assert!(params.contains_key("age"));
641 assert!(params.contains_key("role"));
642
643 let param_types = req.param_types;
644 assert_eq!(param_types.len(), 1);
645 assert!(param_types.contains_key("role"));
646 }
647
648 #[test]
649 fn with_request_tag() {
650 let stmt = Statement::builder("SELECT * FROM users")
651 .set_request_tag("tag1")
652 .build();
653 assert_eq!(
654 stmt.request_options
655 .expect("request options missing")
656 .request_tag,
657 "tag1"
658 );
659 }
660
661 #[test]
662 fn with_priority() {
663 let stmt = Statement::builder("SELECT * FROM users")
664 .set_priority(Priority::High)
665 .build();
666 assert_eq!(
667 stmt.request_options
668 .expect("request options missing")
669 .priority,
670 Priority::High
671 );
672 }
673
674 #[test]
675 fn with_directed_read_options() {
676 let dro = DirectedReadOptions::default();
677 let stmt = Statement::builder("SELECT * FROM users")
678 .set_directed_read_options(dro.clone())
679 .build();
680 assert_eq!(stmt.directed_read_options, Some(dro));
681 }
682
683 #[test]
684 fn with_query_options() -> anyhow::Result<()> {
685 let query_options = QueryOptions::default().set_optimizer_version("1");
686 let stmt = Statement::builder("SELECT * FROM users")
687 .set_query_options(query_options.clone())
688 .build();
689 assert_eq!(
690 stmt.query_options
691 .as_ref()
692 .context("query options missing")?
693 .optimizer_version,
694 "1"
695 );
696
697 let req = stmt.into_request();
698 assert_eq!(
699 req.query_options
700 .context("query options missing in request")?
701 .optimizer_version,
702 "1"
703 );
704 Ok(())
705 }
706
707 #[test]
708 fn with_query_mode() -> anyhow::Result<()> {
709 let stmt = Statement::builder("SELECT * FROM users")
710 .set_query_mode(QueryMode::Plan)
711 .build();
712 assert_eq!(stmt.query_mode, Some(QueryMode::Plan));
713
714 let req = stmt.into_request();
715 assert_eq!(req.query_mode, QueryMode::Plan);
716 Ok(())
717 }
718
719 #[test]
720 fn statement_with_query_mode() -> anyhow::Result<()> {
721 let stmt = Statement::builder("SELECT * FROM users").build();
722 assert_eq!(stmt.query_mode, None);
723
724 let stmt = stmt.set_query_mode(QueryMode::Profile);
725 assert_eq!(stmt.query_mode, Some(QueryMode::Profile));
726
727 let req = stmt.into_request();
728 assert_eq!(req.query_mode, QueryMode::Profile);
729 Ok(())
730 }
731
732 #[test]
733 fn with_gax_options() -> anyhow::Result<()> {
734 use google_cloud_gax::exponential_backoff::ExponentialBackoff;
735 use google_cloud_gax::retry_policy::NeverRetry;
736 use std::time::Duration;
737
738 let stmt = Statement::builder("SELECT * FROM users")
739 .with_attempt_timeout(Duration::from_secs(10))
740 .with_retry_policy(NeverRetry)
741 .with_backoff_policy(ExponentialBackoff::default())
742 .build();
743
744 assert_eq!(
745 stmt.gax_options.attempt_timeout(),
746 &Some(Duration::from_secs(10))
747 );
748 assert!(stmt.gax_options.retry_policy().is_some());
749 assert!(stmt.gax_options.backoff_policy().is_some());
750
751 Ok(())
752 }
753
754 #[test]
755 fn set_last_statement() -> anyhow::Result<()> {
756 let stmt = Statement::builder("SELECT * FROM users").build();
757 assert!(!stmt.last_statement);
758
759 let req = stmt.into_request();
760 assert!(!req.last_statement);
761
762 let stmt = Statement::builder("SELECT * FROM users")
763 .set_last_statement(true)
764 .build();
765 assert!(stmt.last_statement);
766
767 let req = stmt.into_request();
768 assert!(req.last_statement);
769
770 let stmt = Statement::builder("SELECT * FROM users")
771 .build()
772 .set_last_statement(true);
773 assert!(stmt.last_statement);
774
775 let req = stmt.into_request();
776 assert!(req.last_statement);
777
778 Ok(())
779 }
780}