1use std::collections::{HashMap, HashSet};
2use std::pin::Pin;
3use std::task::{Context, Poll};
4use std::time::Duration;
5
6use backon::ExponentialBuilder;
7use backon::Retryable;
8use futures::Stream;
9use http::header::{ACCEPT_ENCODING, USER_AGENT};
10use http::StatusCode;
11use iterable::*;
12use reqwest::header::HeaderValue;
13use reqwest::{RequestBuilder, Response, Url};
14use tokio::sync::RwLock;
15use tracing::*;
16
17use crate::auth::Auth;
18use crate::build_dataset;
19use crate::error::TrinoRetryResult;
20use crate::error::{Error, Result};
21use crate::header::*;
22use crate::models::Column;
23use crate::models::QueryResultData;
24#[cfg(feature = "spooling")]
25use crate::models::SpooledData;
26use crate::retry::RetryPolicy;
27use crate::selected_role::SelectedRole;
28use crate::session::{Session, SessionBuilder};
29#[cfg(feature = "spooling")]
30use crate::spooling::decompress_segment_bytes;
31#[cfg(feature = "spooling")]
32use crate::spooling::{SegmentFetcher, SpoolingEncoding};
33use crate::ssl::Ssl;
34use crate::transaction::TransactionId;
35use crate::{DataSet, QueryResult, Row, Trino};
36
37pub struct Client {
49 client: reqwest::Client,
50 session: RwLock<Session>,
51 auth: Option<Auth>,
52 retry: RetryPolicy,
53 url: Url,
54 #[cfg(feature = "spooling")]
55 segment_fetcher: SegmentFetcher,
56}
57
58pub struct ClientBuilder {
75 session: SessionBuilder,
76 auth: Option<Auth>,
77 auth_http_insecure: bool,
78 retry: RetryPolicy,
79 ssl: Option<Ssl>,
80 no_verify: bool,
81 #[cfg(feature = "spooling")]
82 segment_fetcher: Option<SegmentFetcher>,
83 #[cfg(feature = "spooling")]
84 max_concurrent_segments: Option<usize>,
85}
86
87#[derive(Debug)]
89pub struct ExecuteResult {
90 pub output_uri: Option<String>,
92 pub update_type: Option<String>,
94 pub update_count: Option<u64>,
96}
97
98impl ClientBuilder {
99 pub fn new(user: impl ToString, host: impl ToString) -> Self {
104 let builder = SessionBuilder::new(user, host);
105 Self {
106 session: builder,
107 auth: None,
108 auth_http_insecure: false,
109 retry: RetryPolicy::default(),
110 ssl: None,
111 no_verify: false,
112 #[cfg(feature = "spooling")]
113 segment_fetcher: None,
114 #[cfg(feature = "spooling")]
115 max_concurrent_segments: None,
116 }
117 }
118
119 pub fn port(mut self, s: u16) -> Self {
120 self.session.port = s;
121 self
122 }
123
124 pub fn secure(mut self, s: bool) -> Self {
125 self.session.secure = s;
126 self
127 }
128
129 pub fn no_verify(mut self, nv: bool) -> Self {
130 self.no_verify = nv;
131 self
132 }
133
134 pub fn source(mut self, s: impl ToString) -> Self {
135 self.session.source = s.to_string();
136 self
137 }
138
139 pub fn trace_token(mut self, s: impl ToString) -> Self {
140 self.session.trace_token = Some(s.to_string());
141 self
142 }
143
144 pub fn client_tags(mut self, s: HashSet<String>) -> Self {
145 self.session.client_tags = s;
146 self
147 }
148
149 pub fn client_tag(mut self, s: impl ToString) -> Self {
150 self.session.client_tags.insert(s.to_string());
151 self
152 }
153
154 pub fn client_info(mut self, s: impl ToString) -> Self {
155 self.session.client_info = Some(s.to_string());
156 self
157 }
158
159 pub fn catalog(mut self, s: impl ToString) -> Self {
160 self.session.catalog = Some(s.to_string());
161 self
162 }
163
164 pub fn schema(mut self, s: impl ToString) -> Self {
165 self.session.schema = Some(s.to_string());
166 self
167 }
168
169 pub fn path(mut self, s: impl ToString) -> Self {
170 self.session.path = Some(s.to_string());
171 self
172 }
173
174 pub fn resource_estimates(mut self, s: HashMap<String, String>) -> Self {
175 self.session.resource_estimates = s;
176 self
177 }
178
179 pub fn resource_estimate(mut self, k: impl ToString, v: impl ToString) -> Self {
180 self.session
181 .resource_estimates
182 .insert(k.to_string(), v.to_string());
183 self
184 }
185
186 pub fn properties(mut self, s: HashMap<String, String>) -> Self {
187 self.session.properties = s;
188 self
189 }
190
191 pub fn property(mut self, k: impl ToString, v: impl ToString) -> Self {
192 self.session.properties.insert(k.to_string(), v.to_string());
193 self
194 }
195
196 pub fn prepared_statements(mut self, s: HashMap<String, String>) -> Self {
197 self.session.prepared_statements = s;
198 self
199 }
200
201 pub fn prepared_statement(mut self, k: impl ToString, v: impl ToString) -> Self {
202 self.session
203 .prepared_statements
204 .insert(k.to_string(), v.to_string());
205 self
206 }
207
208 pub fn extra_credentials(mut self, s: HashMap<String, String>) -> Self {
209 self.session.extra_credentials = s;
210 self
211 }
212
213 pub fn extra_credential(mut self, k: impl ToString, v: impl ToString) -> Self {
214 self.session
215 .extra_credentials
216 .insert(k.to_string(), v.to_string());
217 self
218 }
219
220 pub fn transaction_id(mut self, s: TransactionId) -> Self {
221 self.session.transaction_id = s;
222 self
223 }
224
225 pub fn client_request_timeout(mut self, s: Duration) -> Self {
226 self.session.client_request_timeout = s;
227 self
228 }
229
230 pub fn compression_disabled(mut self, s: bool) -> Self {
231 self.session.compression_disabled = s;
232 self
233 }
234
235 #[cfg(feature = "spooling")]
236 pub fn segment_fetcher(mut self, segment_fetcher: SegmentFetcher) -> Self {
237 self.segment_fetcher = Some(segment_fetcher);
238 self
239 }
240
241 #[cfg(feature = "spooling")]
242 pub fn max_concurrent_segments(mut self, count: usize) -> Self {
245 self.max_concurrent_segments = Some(count);
246 self
247 }
248
249 #[cfg(feature = "spooling")]
250 pub fn spooling_encoding(mut self, encoding: impl ToString) -> Self {
253 let encoding_str = encoding.to_string();
254
255 match SpoolingEncoding::try_from(encoding_str.as_str()) {
256 Ok(_) => {
257 self.session.spooling_encoding = Some(encoding_str);
258 }
259 Err(_) => {
260 tracing::warn!(
261 "Invalid spooling encoding '{}', using default 'json+zstd'. Valid values: json, json+zstd, json+lz4",
262 encoding_str
263 );
264 self.session.spooling_encoding = Some("json+zstd".to_string());
265 }
266 }
267
268 self
269 }
270
271 pub fn auth(mut self, s: Auth) -> Self {
274 self.auth = Some(s);
275 self
276 }
277
278 pub fn auth_http_insecure(mut self, ahi: bool) -> Self {
279 self.auth_http_insecure = ahi;
280 self
281 }
282
283 pub fn max_attempt(mut self, s: usize) -> Self {
284 self.retry.max_retries = s;
285 self
286 }
287
288 pub fn retry_policy(mut self, policy: RetryPolicy) -> Self {
290 self.retry = policy;
291 self
292 }
293
294 pub fn ssl(mut self, ssl: Ssl) -> Self {
295 self.ssl = Some(ssl);
296 self
297 }
298
299 pub fn build(self) -> Result<Client> {
300 let session = self.session.build()?;
301 let retry = self.retry.clone();
302
303 if (self.auth.is_some() && session.url.scheme() == "http") && !self.auth_http_insecure {
304 return Err(Error::BasicAuthWithHttp);
305 }
306
307 let mut client_builder =
308 reqwest::ClientBuilder::new().timeout(session.client_request_timeout);
309
310 if self.no_verify {
311 client_builder = client_builder.danger_accept_invalid_certs(true);
312 }
313
314 if let Some(ssl) = &self.ssl {
315 if let Some(root) = &ssl.root_cert {
316 client_builder = client_builder.add_root_certificate(root.0.clone());
317 }
318 }
319
320 let client = client_builder.build()?;
321
322 #[cfg(feature = "spooling")]
323 let segment_fetcher = self.segment_fetcher.unwrap_or_else(|| {
324 let mut fetcher = SegmentFetcher::new(client.clone());
325 if let Some(max_concurrent) = self.max_concurrent_segments {
326 fetcher = fetcher.with_max_concurrent(max_concurrent);
327 }
328 fetcher
329 });
330
331 let cli = Client {
332 auth: self.auth,
333 url: session.url.clone(),
334 session: RwLock::new(session),
335 client,
336 retry,
337 #[cfg(feature = "spooling")]
338 segment_fetcher,
339 };
340
341 Ok(cli)
342 }
343}
344
345fn add_prepare_header(mut builder: RequestBuilder, session: &Session) -> RequestBuilder {
346 builder = builder.header(HEADER_USER, &session.user);
348 builder = builder.header(USER_AGENT, "trino-rust-client");
350 if session.compression_disabled {
351 builder = builder.header(ACCEPT_ENCODING, "identity")
352 }
353 builder
354}
355
356fn add_session_header(mut builder: RequestBuilder, session: &Session) -> RequestBuilder {
357 builder = add_prepare_header(builder, session);
358 builder = builder.header(HEADER_SOURCE, &session.source);
359
360 if let Some(v) = &session.trace_token {
361 builder = builder.header(HEADER_TRACE_TOKEN, v);
362 }
363
364 if !session.client_tags.is_empty() {
365 builder = builder.header(HEADER_CLIENT_TAGS, session.client_tags.by_ref().join(","));
366 }
367
368 if let Some(v) = &session.client_info {
369 builder = builder.header(HEADER_CLIENT_INFO, v);
370 }
371
372 if let Some(v) = &session.catalog {
373 builder = builder.header(HEADER_CATALOG, v);
374 }
375
376 if let Some(v) = &session.schema {
377 builder = builder.header(HEADER_SCHEMA, v);
378 }
379
380 if let Some(v) = &session.path {
381 builder = builder.header(HEADER_PATH, v);
382 }
383 if let Some(v) = &session.timezone {
384 builder = builder.header(HEADER_TIME_ZONE, v.to_string())
385 }
386 builder = add_header_map(builder, HEADER_SESSION, &session.properties);
388 builder = add_header_map(
389 builder,
390 HEADER_RESOURCE_ESTIMATE,
391 &session.resource_estimates,
392 );
393 builder = add_header_map(
394 builder,
395 HEADER_ROLE,
396 &session
397 .roles
398 .by_ref()
399 .map_kv(|(k, v)| (k.to_string(), v.to_string())),
400 );
401 builder = add_header_map(builder, HEADER_EXTRA_CREDENTIAL, &session.extra_credentials);
402 builder = add_header_map(
403 builder,
404 HEADER_PREPARED_STATEMENT,
405 &session.prepared_statements,
406 );
407 builder = builder.header(HEADER_TRANSACTION, session.transaction_id.as_header_value());
408 builder = builder.header(HEADER_CLIENT_CAPABILITIES, "PATH,PARAMETRIC_DATETIME");
409
410 #[cfg(feature = "spooling")]
412 {
413 if let Some(encoding) = &session.spooling_encoding {
414 builder = builder.header(HEADER_SPOOLING, encoding);
415 }
416 }
417
418 builder
419}
420
421fn add_header_map<'a>(
422 mut builder: RequestBuilder,
423 header: &str,
424 map: impl IntoIterator<Item = (&'a String, &'a String)>,
425) -> RequestBuilder {
426 for (k, v) in map {
427 let kv = encode_kv(k, v);
428 builder = builder.header(header, kv);
429 }
430 builder
431}
432
433macro_rules! set_header {
434 ($session:expr, $header:expr, $resp:expr) => {
435 set_header!($session, $header, $resp, |x: &str| Some(Some(
436 x.to_string()
437 )));
438 };
439
440 ($session:expr, $header:expr, $resp:expr, $from_str:expr) => {
441 if let Some(v) = $resp.headers().get($header) {
442 match v.to_str() {
443 Ok(s) => {
444 if let Some(s) = $from_str(s) {
445 $session = s;
446 }
447 }
448 Err(e) => warn!("parse header {} failed, reason: {}", $header, e),
449 }
450 }
451 };
452}
453
454macro_rules! clear_header {
455 ($session:expr, $header:expr, $resp:expr) => {
456 if let Some(_) = $resp.headers().get($header) {
457 $session = Default::default();
458 }
459 };
460}
461
462macro_rules! set_header_map {
463 ($session:expr, $header:expr, $resp:expr) => {
464 set_header_map!($session, $header, $resp, |x: &str| Some(x.to_string()));
465 };
466 ($session:expr, $header:expr, $resp:expr, $from_str:expr) => {
467 for v in $resp.headers().get_all($header) {
468 if let Some((k, v)) = decode_kv_from_header(v) {
469 if let Some(parsed) = $from_str(&v) {
470 $session.insert(k, parsed);
471 } else {
472 warn!("parse header {} value '{}' failed, ignoring", $header, v)
473 }
474 } else {
475 warn!("decode '{:?}' failed", v)
476 }
477 }
478 };
479}
480
481macro_rules! clear_header_map {
482 ($session:expr, $header:expr, $resp:expr) => {
483 for v in $resp.headers().get_all($header) {
484 match v.to_str() {
485 Ok(s) => {
486 $session.remove(s);
487 }
488 Err(e) => warn!("parse header {} failed, reason: {}", $header, e),
489 }
490 }
491 };
492}
493
494fn transient_status(code: &StatusCode) -> bool {
495 matches!(
496 *code,
497 StatusCode::BAD_GATEWAY | StatusCode::SERVICE_UNAVAILABLE | StatusCode::GATEWAY_TIMEOUT
498 )
499}
500
501fn need_retry_fetch(e: &Error) -> bool {
506 match e {
507 Error::HttpError(e) => {
508 e.is_timeout() || e.is_connect() || e.status().as_ref().is_some_and(transient_status)
509 }
510 Error::HttpNotOk(code, _) => transient_status(code),
511 _ => false,
512 }
513}
514
515fn need_retry_submit(e: &Error) -> bool {
521 match e {
522 Error::HttpError(e) => e.is_connect(),
524 Error::HttpNotOk(code, _) => *code == StatusCode::SERVICE_UNAVAILABLE,
526 _ => false,
527 }
528}
529
530struct CancelOnDrop {
533 client: reqwest::Client,
534 url: String,
535 auth: Option<Auth>,
536}
537
538pub struct RowStream<'a, T> {
557 columns: Vec<Column>,
558 cancel: Option<CancelOnDrop>,
559 span: tracing::Span,
563 inner: Pin<Box<dyn Stream<Item = Result<T>> + Send + 'a>>,
564}
565
566impl<T> RowStream<'_, T> {
567 pub fn columns(&self) -> &[Column] {
570 &self.columns
571 }
572}
573
574impl<T> Stream for RowStream<'_, T> {
575 type Item = Result<T>;
576
577 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
578 let me = self.get_mut();
579 let _enter = me.span.enter();
580 let polled = me.inner.as_mut().poll_next(cx);
581 if let Poll::Ready(None) = polled {
582 me.cancel = None;
584 }
585 polled
586 }
587}
588
589impl<T> Drop for RowStream<'_, T> {
590 fn drop(&mut self) {
591 let Some(cancel) = self.cancel.take() else {
592 return;
593 };
594 if let Ok(handle) = tokio::runtime::Handle::try_current() {
596 handle.spawn(async move {
597 let mut req = cancel.client.delete(&cancel.url);
598 if let Some(auth) = &cancel.auth {
599 req = match auth {
600 Auth::Basic(u, p) => req.basic_auth(u, p.as_ref()),
601 Auth::Jwt(t) => req.bearer_auth(t),
602 Auth::OAuth2(state) => match state.cached_token() {
608 Some(t) => req.bearer_auth(t),
609 None => req,
610 },
611 };
612 }
613 let _ = req.send().await;
614 });
615 }
616 }
617}
618
619impl Client {
620 pub async fn stream<'a, T>(&'a self, sql: impl Into<String>) -> Result<RowStream<'a, T>>
657 where
658 T: Trino + Send + 'static,
659 for<'de> T: serde::Deserialize<'de>,
660 {
661 let sql = sql.into();
662
663 let mut res = self.get_retry::<T>(sql).await?;
667 let span = tracing::info_span!("query_stream", query_id = %res.id);
672 loop {
673 if let Some(error) = res.error.take() {
674 return Err(error.into());
675 }
676 if res.columns.is_some() || res.data.is_some() {
677 break;
678 }
679 match res.next_uri.clone() {
680 Some(url) => res = self.get_next_retry::<T>(&url).await?,
681 None => break,
682 }
683 }
684
685 let columns = res.columns.clone().unwrap_or_default();
686
687 let cancel = Some(CancelOnDrop {
690 client: self.client.clone(),
691 url: format!("{}v1/query/{}", self.url, res.id),
692 auth: self.auth.clone(),
693 });
694
695 let inner = async_stream::try_stream! {
696 let mut res = res;
699 #[cfg(feature = "spooling")]
701 let mut raw_columns: Option<Vec<Column>> = res.columns.clone();
702
703 loop {
704 if let Some(error) = res.error.take() {
705 Err(Error::from(error))?;
706 }
707
708 #[cfg(feature = "spooling")]
709 if raw_columns.is_none() {
710 raw_columns = res.columns.clone();
711 }
712
713 if let Some(data) = res.data.take() {
714 match data {
715 QueryResultData::Direct(rows) => {
716 for row in rows {
717 yield row;
718 }
719 }
720 #[cfg(feature = "spooling")]
721 QueryResultData::Spooled(spooled) => {
722 let cols = raw_columns.clone().or_else(|| res.columns.clone());
723 let ds = self.fetch_spooled_data::<T>(spooled, cols).await?;
724 for row in ds.into_vec() {
725 yield row;
726 }
727 }
728 #[cfg(not(feature = "spooling"))]
729 QueryResultData::Spooled(_) => {
730 Err(Error::Protocol(
731 "Server sent spooled data but 'spooling' feature is not enabled. \
732 Add features = [\"spooling\"] to your trino-rust-client dependency in Cargo.toml.".to_string(),
733 ))?;
734 }
735 }
736 }
737
738 match res.next_uri.take() {
739 Some(url) => {
740 res = self.get_next_retry::<T>(&url).await?;
741 }
742 None => break,
743 }
744 }
745 };
746
747 Ok(RowStream {
748 columns,
749 cancel,
750 span,
751 inner: Box::pin(inner),
752 })
753 }
754
755 #[tracing::instrument(skip_all, fields(query_id = tracing::field::Empty))]
761 pub async fn get_all<T>(&self, sql: impl Into<String>) -> Result<DataSet<T>>
762 where
763 T: Trino + 'static,
764 for<'de> T: serde::Deserialize<'de> + serde::Serialize,
765 {
766 let res = self.get_retry(sql.into()).await?;
767 tracing::Span::current().record("query_id", res.id.as_str());
768
769 let mut columns = res.columns;
771
772 match res.data {
773 Some(QueryResultData::Direct(rows)) => {
774 let mut all_rows = rows;
776
777 let mut next = res.next_uri;
778 while let Some(url) = &next {
779 let mut res = self.get_next_retry(url).await?;
780 next = res.next_uri;
781
782 if columns.is_none() {
784 columns = res.columns.take();
785 }
786
787 if let Some(error) = res.error {
788 return Err(error.into());
789 }
790
791 if let Some(data) = res.data {
792 match data {
793 QueryResultData::Direct(rows) => {
794 all_rows.extend(rows);
795 }
796 #[cfg(feature = "spooling")]
797 QueryResultData::Spooled(_) => {
798 return Err(Error::Protocol(
799 "Cannot mix Direct and Spooled protocols in same query".to_string(),
800 ));
801 }
802 #[cfg(not(feature = "spooling"))]
803 QueryResultData::Spooled(_) => {
804 return Err(Error::Protocol(
805 "Server sent spooled data but 'spooling' feature is not enabled. \
806 Add features = [\"spooling\"] to your trino-rust-client dependency in Cargo.toml.".to_string(),
807 ));
808 }
809 }
810 }
811 }
812
813 build_dataset(all_rows, columns)
814 }
815 #[cfg(feature = "spooling")]
816 Some(QueryResultData::Spooled(spooled)) => {
817 let mut dataset = self
818 .fetch_spooled_data::<T>(spooled, columns.clone())
819 .await?;
820
821 let mut next = res.next_uri;
822 while let Some(url) = &next {
823 let mut res = self.get_next_retry::<T>(url).await?;
824 next = res.next_uri;
825
826 if columns.is_none() {
827 columns = res.columns.take();
828 }
829
830 if let Some(error) = res.error {
831 return Err(error.into());
832 }
833
834 if let Some(data) = res.data {
835 match data {
836 QueryResultData::Direct(_) => {
837 return Err(Error::Protocol(
838 "Cannot mix Direct and Spooled protocols in same query".to_string(),
839 ));
840 }
841 QueryResultData::Spooled(spooled) => {
842 tracing::info!("🗄️ Received SPOOLED protocol data - fetching from S3/MinIO");
843 let cols_for_spooled = columns.clone().or_else(|| res.columns.take());
844 let next_dataset = self
845 .fetch_spooled_data::<T>(spooled, cols_for_spooled)
846 .await?;
847 dataset.merge(next_dataset);
848 }
849 }
850 }
851 }
852
853 Ok(dataset)
854 }
855 #[cfg(not(feature = "spooling"))]
856 Some(QueryResultData::Spooled(_)) => {
857 Err(Error::Protocol(
858 "Server sent spooled data but 'spooling' feature is not enabled. \
859 Add features = [\"spooling\"] to your trino-rust-client dependency in Cargo.toml.".to_string(),
860 ))
861 }
862 None => {
863 let mut next = res.next_uri;
865 let mut protocol_detected = false;
866 let mut all_rows: Vec<T> = Vec::new();
867 #[cfg(feature = "spooling")]
868 let mut dataset: Option<DataSet<T>> = None;
869
870 while let Some(url) = &next {
871 let mut res = self.get_next_retry::<T>(url).await?;
872 next = res.next_uri;
873
874 if columns.is_none() {
875 columns = res.columns.take();
876 }
877
878 if let Some(error) = res.error {
879 return Err(error.into());
880 }
881
882 if let Some(data) = res.data {
883 match data {
884 QueryResultData::Direct(rows) => {
885 if !protocol_detected {
886 protocol_detected = true;
887 }
888 all_rows.extend(rows);
889 }
890 #[cfg(feature = "spooling")]
891 QueryResultData::Spooled(spooled) => {
892 if !protocol_detected {
893 protocol_detected = true;
894 let cols_for_spooled = columns.clone().or_else(|| res.columns.take());
895 dataset = Some(self.fetch_spooled_data::<T>(spooled, cols_for_spooled).await?);
896 } else {
897 let cols_for_spooled = columns.clone().or_else(|| res.columns.take());
898 let next_dataset = self.fetch_spooled_data::<T>(spooled, cols_for_spooled).await?;
899 if let Some(ref mut ds) = dataset {
900 ds.merge(next_dataset);
901 }
902 }
903 }
904 #[cfg(not(feature = "spooling"))]
905 QueryResultData::Spooled(_) => {
906 return Err(Error::Protocol(
907 "Server sent spooled data but 'spooling' feature is not enabled. \
908 Add features = [\"spooling\"] to your trino-rust-client dependency in Cargo.toml.".to_string(),
909 ));
910 }
911 }
912 }
913 }
914
915 #[cfg(feature = "spooling")]
916 if let Some(ds) = dataset {
917 Ok(ds)
918 } else {
919 build_dataset(all_rows, columns)
920 }
921 #[cfg(not(feature = "spooling"))]
922 build_dataset(all_rows, columns)
923 }
924 }
925 }
926
927 #[cfg(feature = "spooling")]
928 async fn fetch_spooled_data<T: Trino + 'static>(
929 &self,
930 spooled: SpooledData,
931 columns: Option<Vec<crate::models::Column>>,
932 ) -> Result<DataSet<T>> {
933 let segment_bytes = self
934 .segment_fetcher
935 .fetch_segments(spooled.segments)
936 .await?;
937
938 let dataset = self.decode_segments::<T>(&spooled.encoding, segment_bytes, columns)?;
939
940 Ok(dataset)
941 }
942
943 #[cfg(feature = "spooling")]
944 fn decode_segments<T: Trino + 'static>(
945 &self,
946 encoding: &str,
947 segment_bytes: Vec<Vec<u8>>,
948 columns: Option<Vec<crate::models::Column>>,
949 ) -> Result<DataSet<T>> {
950 let cols = columns.ok_or_else(|| {
951 Error::Protocol("Column metadata required for spooling protocol".to_string())
952 })?;
953
954 let mut all_rows: Vec<Vec<serde_json::Value>> = Vec::new();
955
956 let encoding = SpoolingEncoding::try_from(encoding).map_err(|e| {
957 Error::Decode(format!(
958 "Failed to parse encoding: {}. Only 'json' based formats are supported.",
959 e
960 ))
961 })?;
962
963 for bytes in segment_bytes {
964 let json_str = decompress_segment_bytes(&bytes, &encoding)?;
965
966 let mut rows: Vec<Vec<serde_json::Value>> = serde_json::from_str(&json_str)
967 .map_err(|e| Error::Decode(format!("Failed to parse segment JSON: {}", e)))?;
968
969 all_rows.append(&mut rows);
970 }
971
972 let json_obj = serde_json::json!({
973 "columns": cols,
974 "data": all_rows
975 });
976
977 let dataset: DataSet<T> = serde_json::from_value(json_obj)
978 .map_err(|e| Error::Decode(format!("Failed to deserialize DataSet: {}", e)))?;
979
980 Ok(dataset)
981 }
982
983 #[tracing::instrument(skip_all, fields(query_id = tracing::field::Empty))]
990 pub async fn execute(&self, sql: impl Into<String>) -> Result<ExecuteResult> {
991 let res = self.get_retry::<Row>(sql.into()).await?;
993 tracing::Span::current().record("query_id", res.id.as_str());
994
995 let mut next = res.next_uri;
996 let mut final_uri = next.clone();
997
998 while let Some(url) = &next {
1001 let res = self.get_next_retry::<Row>(url).await?;
1002
1003 let next_uri = res.next_uri;
1004
1005 if next_uri.is_some() {
1007 final_uri = next_uri.clone();
1008 }
1009 next = next_uri;
1010 }
1011
1012 let url = final_uri.ok_or_else(|| {
1013 Error::InternalError("No next URI available for execution result".to_string())
1014 })?;
1015
1016 let result = self.try_get_retry_result(&url).await?;
1018
1019 if let Some(error) = result.error {
1020 return Err(error.into());
1021 }
1022
1023 Ok(ExecuteResult {
1024 output_uri: None,
1025 update_type: result.update_type,
1026 update_count: result.update_count,
1027 })
1028 }
1029
1030 pub async fn transaction_id(&self) -> TransactionId {
1032 self.session.read().await.transaction_id.clone()
1033 }
1034
1035 pub async fn set_transaction_id(&self, id: TransactionId) {
1041 self.session.write().await.transaction_id = id;
1042 }
1043
1044 pub async fn begin_transaction(&self) -> Result<()> {
1085 let active = self.session.read().await.transaction_id.is_active();
1088 if active {
1089 return Err(Error::Transaction(
1090 "a transaction is already active; Trino does not support nested transactions"
1091 .to_string(),
1092 ));
1093 }
1094 self.execute("START TRANSACTION").await?;
1095
1096 if !self.session.read().await.transaction_id.is_active() {
1102 return Err(Error::Transaction(
1103 "START TRANSACTION succeeded but no usable transaction id was returned in \
1104 X-Trino-Started-Transaction-Id; the session is not in a transaction"
1105 .to_string(),
1106 ));
1107 }
1108 Ok(())
1109 }
1110
1111 pub async fn commit(&self) -> Result<()> {
1117 self.end_transaction("COMMIT").await
1118 }
1119
1120 pub async fn rollback(&self) -> Result<()> {
1126 self.end_transaction("ROLLBACK").await
1127 }
1128
1129 async fn end_transaction(&self, statement: &str) -> Result<()> {
1145 let active = self.session.read().await.transaction_id.is_active();
1146 if !active {
1147 return Err(Error::Transaction(format!(
1148 "no active transaction to {}",
1149 statement.to_lowercase()
1150 )));
1151 }
1152 self.execute(statement).await?;
1153 Ok(())
1154 }
1155
1156 async fn try_get_retry_result(&self, url: &str) -> Result<TrinoRetryResult> {
1157 let response = self.client.get(url).send().await?;
1158
1159 let result = response.json::<TrinoRetryResult>().await?;
1160
1161 Ok(result)
1162 }
1163
1164 fn retry_policy(&self) -> ExponentialBuilder {
1165 self.retry.backoff()
1166 }
1167
1168 async fn get_retry<T>(&self, sql: String) -> Result<QueryResult<T>>
1169 where
1170 T: Trino + 'static,
1171 for<'de> T: serde::Deserialize<'de>,
1172 {
1173 let result = || async { self.get::<T>(sql.clone()).await };
1174
1175 result
1177 .retry(self.retry_policy())
1178 .when(need_retry_submit)
1179 .await
1180 }
1181
1182 async fn get_next_retry<T>(&self, url: &str) -> Result<QueryResult<T>>
1183 where
1184 T: Trino + 'static,
1185 for<'de> T: serde::Deserialize<'de>,
1186 {
1187 let result = || async { self.get_next(url).await };
1188
1189 result
1191 .retry(self.retry_policy())
1192 .when(need_retry_fetch)
1193 .await
1194 }
1195
1196 pub async fn get<T>(&self, sql: impl Into<String>) -> Result<QueryResult<T>>
1203 where
1204 T: Trino + 'static,
1205 for<'de> T: serde::Deserialize<'de>,
1206 {
1207 let req = self
1208 .client
1209 .post(format!("{}v1/statement", self.url))
1210 .body(sql.into());
1211 let req = {
1212 let session = self.session.read().await;
1213 add_session_header(req, &session)
1214 };
1215
1216 self.send(req, StatusCode::OK, |resp| async {
1217 let text = resp.text().await?;
1218
1219 let data: QueryResult<T> = serde_json::from_str(&text)
1220 .map_err(|e| Error::Decode(format!("Failed to parse response: {}", e)))?;
1221 Ok(data)
1222 })
1223 .await
1224 }
1225
1226 pub async fn get_next<T>(&self, url: &str) -> Result<QueryResult<T>>
1229 where
1230 T: Trino + 'static,
1231 for<'de> T: serde::Deserialize<'de>,
1232 {
1233 let req = self.client.get(url);
1234 let req = {
1235 let session = self.session.read().await;
1236 add_prepare_header(req, &session)
1237 };
1238
1239 self.send(req, StatusCode::OK, |resp| async {
1240 let text = resp.text().await?;
1241 let data: QueryResult<T> = serde_json::from_str(&text)
1242 .map_err(|e| Error::Decode(format!("Failed to parse response: {}", e)))?;
1243 Ok(data)
1244 })
1245 .await
1246 }
1247
1248 pub async fn cancel(&self, query_id: &str) -> Result<()> {
1251 let url = format!("{}v1/query/{}", self.url, query_id);
1252 let req = self.client.delete(url);
1253 let req = {
1254 let session = self.session.read().await;
1255 add_prepare_header(req, &session)
1256 };
1257
1258 self.send(req, StatusCode::NO_CONTENT, |_| async { Ok(()) })
1259 .await
1260 }
1261
1262 fn auth_req(&self, req: RequestBuilder) -> RequestBuilder {
1263 if let Some(auth) = self.auth.as_ref() {
1264 match auth {
1265 Auth::Basic(u, p) => req.basic_auth(u, p.as_ref()),
1266 Auth::Jwt(t) => req.bearer_auth(t),
1267 Auth::OAuth2(state) => match state.cached_token() {
1271 Some(t) => req.bearer_auth(t),
1272 None => req,
1273 },
1274 }
1275 } else {
1276 req
1277 }
1278 }
1279
1280 async fn send<R, F, Fut>(
1281 &self,
1282 req: RequestBuilder,
1283 expected_status: StatusCode,
1284 handle_response: F,
1285 ) -> Result<R>
1286 where
1287 F: FnOnce(Response) -> Fut,
1288 Fut: std::future::Future<Output = Result<R>>,
1289 {
1290 let sent_token = match self.auth.as_ref() {
1293 Some(Auth::OAuth2(state)) => state.cached_token(),
1294 _ => None,
1295 };
1296 let retry_req = req.try_clone();
1301 let resp = self.auth_req(req).send().await?;
1302
1303 if resp.status() == StatusCode::UNAUTHORIZED {
1304 if let (Some(Auth::OAuth2(state)), Some(retry_req)) = (self.auth.as_ref(), retry_req) {
1305 if let Some(challenge) = resp
1311 .headers()
1312 .get_all(reqwest::header::WWW_AUTHENTICATE)
1313 .iter()
1314 .filter_map(|v| v.to_str().ok())
1315 .find_map(crate::auth::parse_www_authenticate)
1316 {
1317 self.acquire_oauth2_token(state, &challenge, sent_token)
1318 .await?;
1319 let resp = self.auth_req(retry_req).send().await?;
1320 return self
1321 .finish_send(resp, expected_status, handle_response)
1322 .await;
1323 }
1324 }
1325 }
1326
1327 self.finish_send(resp, expected_status, handle_response)
1328 .await
1329 }
1330
1331 async fn finish_send<R, F, Fut>(
1334 &self,
1335 resp: Response,
1336 expected_status: StatusCode,
1337 handle_response: F,
1338 ) -> Result<R>
1339 where
1340 F: FnOnce(Response) -> Fut,
1341 Fut: std::future::Future<Output = Result<R>>,
1342 {
1343 let status = resp.status();
1344 if status != expected_status {
1345 let data = resp.text().await.unwrap_or("".to_string());
1346 Err(Error::HttpNotOk(status, data))
1347 } else {
1348 self.update_session(&resp).await;
1349 handle_response(resp).await
1350 }
1351 }
1352
1353 async fn acquire_oauth2_token(
1357 &self,
1358 state: &std::sync::Arc<crate::auth::OAuth2State>,
1359 challenge: &crate::auth::Challenge,
1360 sent_token: Option<String>,
1361 ) -> Result<()> {
1362 let _guard = state.acquire.lock().await;
1363 if state.cached_token() != sent_token {
1366 return Ok(());
1367 }
1368 let token = crate::auth::run_flow(&self.client, state, challenge).await?;
1369 *state.token.write().unwrap() = Some(token);
1370 Ok(())
1371 }
1372
1373 async fn update_session(&self, resp: &Response) {
1374 let mut session = self.session.write().await;
1375
1376 set_header!(session.catalog, HEADER_SET_CATALOG, resp);
1377 set_header!(session.schema, HEADER_SET_SCHEMA, resp);
1378 set_header!(session.path, HEADER_SET_PATH, resp);
1379
1380 set_header_map!(session.properties, HEADER_SET_SESSION, resp);
1381 clear_header_map!(session.properties, HEADER_CLEAR_SESSION, resp);
1382
1383 set_header_map!(session.roles, HEADER_SET_ROLE, resp, SelectedRole::from_str);
1384
1385 set_header_map!(session.prepared_statements, HEADER_ADDED_PREPARE, resp);
1386 clear_header_map!(
1387 session.prepared_statements,
1388 HEADER_DEALLOCATED_PREPARE,
1389 resp
1390 );
1391
1392 if let Some(v) = resp.headers().get(HEADER_STARTED_TRANSACTION_ID) {
1393 match v.to_str() {
1394 Ok(s) => session.transaction_id = TransactionId::from_header_value(s),
1395 Err(e) => warn!(
1396 "parse header {} failed, reason: {}",
1397 HEADER_STARTED_TRANSACTION_ID, e
1398 ),
1399 }
1400 }
1401 clear_header!(session.transaction_id, HEADER_CLEAR_TRANSACTION_ID, resp);
1402 }
1403}
1404
1405fn encode_kv(k: &str, v: &str) -> String {
1409 url::form_urlencoded::Serializer::new(String::new())
1410 .append_pair(k, v)
1411 .finish()
1412}
1413
1414fn decode_kv_from_header(input: &HeaderValue) -> Option<(String, String)> {
1415 let kvs = url::form_urlencoded::parse(input.as_bytes()).collect::<Vec<_>>();
1416 if kvs.is_empty() {
1417 None
1418 } else {
1419 Some((kvs[0].0.to_string(), kvs[0].1.to_string()))
1420 }
1421}
1422
1423#[cfg(test)]
1424mod tests {
1425 use http::StatusCode;
1426 use reqwest::header::HeaderValue;
1427
1428 use super::*;
1429 use crate::client::{decode_kv_from_header, need_retry_fetch, need_retry_submit};
1430 use crate::error::Error;
1431 use crate::transaction::TransactionId;
1432
1433 #[test]
1434 fn test_decode_kv_from_header_plus_sign_to_space() {
1435 let header_value = HeaderValue::from_static("statement=show+tables");
1436 let result = decode_kv_from_header(&header_value);
1437 assert!(result.is_some());
1438 let (key, value) = result.unwrap();
1439 assert_eq!(key, "statement");
1440 assert_eq!(value, "show tables");
1441 }
1442
1443 #[test]
1444 fn test_decode_kv_from_header_percent_encoding() {
1445 let header_value = HeaderValue::from_static("statement=show%20tables");
1446 let result = decode_kv_from_header(&header_value);
1447 assert!(result.is_some());
1448 let (key, value) = result.unwrap();
1449 assert_eq!(key, "statement");
1450 assert_eq!(value, "show tables");
1451 }
1452
1453 fn http_not_ok(code: StatusCode) -> Error {
1454 Error::HttpNotOk(code, String::new())
1455 }
1456
1457 #[test]
1458 fn fetch_retries_all_transient_statuses() {
1459 for code in [
1461 StatusCode::BAD_GATEWAY,
1462 StatusCode::SERVICE_UNAVAILABLE,
1463 StatusCode::GATEWAY_TIMEOUT,
1464 ] {
1465 assert!(need_retry_fetch(&http_not_ok(code)), "{code}");
1466 }
1467 for code in [
1469 StatusCode::BAD_REQUEST,
1470 StatusCode::UNAUTHORIZED,
1471 StatusCode::INTERNAL_SERVER_ERROR,
1472 ] {
1473 assert!(!need_retry_fetch(&http_not_ok(code)), "{code}");
1474 }
1475 assert!(!need_retry_fetch(&Error::Protocol(
1476 "mixed protocols".into()
1477 )));
1478 assert!(!need_retry_fetch(&Error::InconsistentData));
1479 }
1480
1481 #[test]
1482 fn submit_only_retries_definitely_unprocessed() {
1483 assert!(need_retry_submit(&http_not_ok(
1485 StatusCode::SERVICE_UNAVAILABLE
1486 )));
1487 assert!(!need_retry_submit(&http_not_ok(StatusCode::BAD_GATEWAY)));
1489 assert!(!need_retry_submit(&http_not_ok(
1490 StatusCode::GATEWAY_TIMEOUT
1491 )));
1492 assert!(!need_retry_submit(&http_not_ok(
1493 StatusCode::INTERNAL_SERVER_ERROR
1494 )));
1495 }
1496
1497 #[tokio::test]
1498 async fn transaction_id_defaults_to_no_transaction() {
1499 let client = ClientBuilder::new("user", "localhost").build().unwrap();
1500 assert_eq!(client.transaction_id().await, TransactionId::NoTransaction);
1501 }
1502
1503 #[tokio::test]
1504 async fn set_transaction_id_is_observable() {
1505 let client = ClientBuilder::new("user", "localhost").build().unwrap();
1506 let id = TransactionId::Id("17cbc429-462a-4da3-9a06-02b6507d0d01".to_string());
1507 client.set_transaction_id(id.clone()).await;
1508 assert_eq!(client.transaction_id().await, id);
1509 }
1510
1511 #[tokio::test]
1512 async fn begin_transaction_rejects_nesting() {
1513 let client = ClientBuilder::new("user", "localhost").build().unwrap();
1514 client
1515 .set_transaction_id(TransactionId::Id("abc".to_string()))
1516 .await;
1517
1518 let err = client.begin_transaction().await.unwrap_err();
1519 assert!(
1520 matches!(err, Error::Transaction(_)),
1521 "expected Error::Transaction, got {err:?}"
1522 );
1523 }
1524
1525 #[tokio::test]
1526 async fn commit_without_transaction_is_rejected() {
1527 let client = ClientBuilder::new("user", "localhost").build().unwrap();
1528 let err = client.commit().await.unwrap_err();
1529 assert!(
1530 matches!(err, Error::Transaction(_)),
1531 "expected Error::Transaction, got {err:?}"
1532 );
1533 }
1534
1535 #[tokio::test]
1536 async fn rollback_without_transaction_is_rejected() {
1537 let client = ClientBuilder::new("user", "localhost").build().unwrap();
1538 let err = client.rollback().await.unwrap_err();
1539 assert!(
1540 matches!(err, Error::Transaction(_)),
1541 "expected Error::Transaction, got {err:?}"
1542 );
1543 }
1544}