Skip to main content

trino_rust_client/
client.rs

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
37// TODO:
38// allow_redirects
39// proxies
40
41/// A configured Trino client.
42///
43/// Created with [`ClientBuilder`]. Cheap to share: it wraps a connection-pooled
44/// HTTP client, so build one and reuse it for all queries. The main entry
45/// points are [`get_all`](Client::get_all) (buffer the result),
46/// [`stream`](Client::stream) (stream it lazily) and [`execute`](Client::execute)
47/// (run a statement).
48pub 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
58/// Builder for a [`Client`].
59///
60/// Start with [`ClientBuilder::new`], chain the setters you need, then call
61/// [`build`](ClientBuilder::build).
62///
63/// ```no_run
64/// # use trino_rust_client::client::ClientBuilder;
65/// # fn run() -> Result<(), Box<dyn std::error::Error>> {
66/// let client = ClientBuilder::new("user", "trino.example.com")
67///     .port(8443)
68///     .secure(true)
69///     .catalog("hive")
70///     .schema("default")
71///     .build()?;
72/// # Ok(()) }
73/// ```
74pub 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/// Outcome of a statement run with [`Client::execute`].
88#[derive(Debug)]
89pub struct ExecuteResult {
90    /// URI of the output, when the statement produces one.
91    pub output_uri: Option<String>,
92    /// The kind of update (e.g. `INSERT`, `CREATE TABLE`), if reported.
93    pub update_type: Option<String>,
94    /// Number of rows affected, if reported.
95    pub update_count: Option<u64>,
96}
97
98impl ClientBuilder {
99    /// Start building a client for the given Trino `user` and `host`.
100    ///
101    /// Defaults: port 8080, plain HTTP, no authentication. Use the setters to
102    /// change them, then call [`build`](ClientBuilder::build).
103    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    /// Set the maximum number of concurrent segment fetches
243    /// Default is based on available CPU parallelism (minimum 1)
244    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    /// Set the spooling encoding format. Supported values: "json", "json+zstd", "json+lz4".
251    /// Defaults to "json+zstd" if not specified.
252    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    ////////////////////////////////////////////////////////////////////////////////////////////////
272
273    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    /// Set the full retry/backoff policy for transient failures.
289    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    //FIXME : set trino user from jwt ?
347    builder = builder.header(HEADER_USER, &session.user);
348    // TODO: difference with session.source?
349    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    // TODO: add locale
387    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    // Add spooling header when feature is enabled
411    #[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
501/// Retry predicate for **idempotent** requests (fetching result pages via
502/// `GET nextUri`). Any transient failure is safe to retry: gateway/availability
503/// responses (HTTP 502/503/504) and low-level connect/timeout errors. Query,
504/// decode, protocol and other errors are terminal.
505fn 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
515/// Retry predicate for **query submission** (`POST /v1/statement`). Only retry
516/// when the request was definitely NOT processed by the server, so a
517/// non-idempotent statement (e.g. `INSERT`/`UPDATE`/DDL via [`Client::execute`])
518/// is never submitted twice. A timeout — or a 502/504 from an intermediary — is
519/// ambiguous (the query may already be running) and is treated as terminal.
520fn need_retry_submit(e: &Error) -> bool {
521    match e {
522        // Connection was never established, so the request was not sent.
523        Error::HttpError(e) => e.is_connect(),
524        // 503 means the coordinator rejected the request without processing it.
525        Error::HttpNotOk(code, _) => *code == StatusCode::SERVICE_UNAVAILABLE,
526        _ => false,
527    }
528}
529
530/// Everything needed to fire a best-effort query cancellation from
531/// [`RowStream`]'s `Drop`, without borrowing the [`Client`].
532struct CancelOnDrop {
533    client: reqwest::Client,
534    url: String,
535    auth: Option<Auth>,
536}
537
538/// A lazy stream of query rows, with the result columns resolved up front.
539///
540/// Created by [`Client::stream`]. The result columns are available immediately
541/// via [`RowStream::columns`]; rows are then produced lazily, page by page, by
542/// the [`Stream`] implementation — the whole result set is never buffered in
543/// memory.
544///
545/// `RowStream` is [`Unpin`], so it can be polled directly (e.g. with
546/// [`StreamExt::next`](futures::StreamExt::next)) without `pin!`, and [`Send`],
547/// so it can be held across `.await` inside a spawned task.
548///
549/// # Cancellation
550/// Dropping a `RowStream` before it is exhausted best-effort cancels the query
551/// on the Trino coordinator (a fire-and-forget `DELETE`), so early termination
552/// (`take`, `break`, an error, a dropped task) does not leave the query running
553/// server-side and holding coordinator resources. Cancellation is skipped once
554/// the query has finished normally, and requires a Tokio runtime to be active
555/// at drop time.
556pub struct RowStream<'a, T> {
557    columns: Vec<Column>,
558    cancel: Option<CancelOnDrop>,
559    // Entered on every poll so events emitted while streaming (page fetches,
560    // segment downloads) carry the query_id — the span from `stream()` itself
561    // would otherwise close as soon as the RowStream is handed back.
562    span: tracing::Span,
563    inner: Pin<Box<dyn Stream<Item = Result<T>> + Send + 'a>>,
564}
565
566impl<T> RowStream<'_, T> {
567    /// The result columns (name, Trino type name and full type signature),
568    /// resolved before the first row is produced.
569    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            // The query finished normally — there is nothing to cancel.
583            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        // Fire-and-forget; only possible from within a running Tokio runtime.
595        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                        // Only ever the cached token: a `Drop` must not block on
603                        // an interactive login, so unlike `Client::auth_req` this
604                        // never runs the OAuth2 flow. With no cached token — or an
605                        // expired one — the cancellation is simply lost and the
606                        // coordinator times the query out on its own.
607                        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    /// Execute `sql` and stream the resulting rows lazily, page by page, without
621    /// buffering the whole result set in memory.
622    ///
623    /// Trino returns results as a chain of pages linked by `nextUri`. This method
624    /// first drives the query far enough to resolve the result schema (so
625    /// [`RowStream::columns`] is available up front), then hands back a
626    /// [`RowStream`] that follows the remaining pages on demand, yielding each
627    /// row as it is decoded. Prefer it over [`Client::get_all`] for large result
628    /// sets.
629    ///
630    /// Both the Direct and (with the `spooling` feature) Spooled protocols are
631    /// supported. With spooling, rows are still materialized one segment at a
632    /// time rather than for the entire query, keeping peak memory bounded.
633    ///
634    /// Unlike [`Client::get_all`], this does not reject a query that mixes the
635    /// Direct and Spooled protocols across pages; each page is decoded according
636    /// to its own protocol.
637    ///
638    /// The returned stream borrows `self`, so it must not outlive the [`Client`].
639    ///
640    /// # Example
641    /// ```no_run
642    /// # use trino_rust_client::{client::ClientBuilder, Row};
643    /// # async fn run() -> Result<(), Box<dyn std::error::Error>> {
644    /// use futures::StreamExt;
645    ///
646    /// let client = ClientBuilder::new("user", "localhost").port(8080).build()?;
647    /// let mut rows = client.stream::<Row>("SELECT 1").await?;
648    /// println!("columns: {:?}", rows.columns());
649    /// while let Some(row) = rows.next().await {
650    ///     let row = row?;
651    ///     // use row
652    /// }
653    /// # Ok(())
654    /// # }
655    /// ```
656    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        // Prime the query until the schema is known: follow pages until one
664        // carries `columns` (or the query finishes without any). Errors on these
665        // early pages are surfaced eagerly.
666        let mut res = self.get_retry::<T>(sql).await?;
667        // Span stored on the RowStream and entered on each `poll_next`, so
668        // events emitted while streaming carry the query_id. (Entering it here
669        // across the priming `.await`s would be the guard-across-await
670        // anti-pattern; priming emits little, so it is left unspanned.)
671        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        // Capture what is needed to cancel the query on early drop, without
688        // borrowing `self` (so the cancel can be spawned as a 'static task).
689        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            // `res` already holds the first schema-bearing page (with its data,
697            // if any); keep decoding from there.
698            let mut res = res;
699            // Track raw columns across pages so later spooled pages can be decoded.
700            #[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    /// Run `sql` and return the whole result set as a [`DataSet`].
756    ///
757    /// The entire result is buffered in memory — for large results prefer
758    /// [`stream`](Client::stream). `T` is a `#[derive(Trino)]` row struct, or
759    /// [`Row`] for a dynamically-typed result.
760    #[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        // Store columns from responses (used for Direct protocol DataSet construction)
770        let mut columns = res.columns;
771
772        match res.data {
773            Some(QueryResultData::Direct(rows)) => {
774                // Direct protocol: accumulate Vec<T>, convert to DataSet at the end
775                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                    // Collect columns from any response that has them
783                    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                // No initial data, wait for next response to detect protocol
864                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    /**
984     * Execute a SQL statement and return the result.
985     * If the TRINO query returns an error, the method returns an error of type `Error::Query`
986     * @param sql The SQL statement to execute
987     * @return [`Result<ExecuteResult>`]` The result of the execution
988     * */
989    #[tracing::instrument(skip_all, fields(query_id = tracing::field::Empty))]
990    pub async fn execute(&self, sql: impl Into<String>) -> Result<ExecuteResult> {
991        // try the sql first
992        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        // Trino attempts several times to execute a query before marking it as failed.
999        // At the end, retrieve the URL of the last request to get the result
1000        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 not None, update final_uri
1006            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        // Parse the final URI to get TrinoRetryResult
1017        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    /// The transaction this client's session is currently bound to.
1031    pub async fn transaction_id(&self) -> TransactionId {
1032        self.session.read().await.transaction_id.clone()
1033    }
1034
1035    /// Bind the session to a transaction.
1036    ///
1037    /// Normally unnecessary — [`begin_transaction`](Self::begin_transaction)
1038    /// captures the identifier Trino issues. Use this to adopt a transaction
1039    /// started elsewhere.
1040    pub async fn set_transaction_id(&self, id: TransactionId) {
1041        self.session.write().await.transaction_id = id;
1042    }
1043
1044    /// Start a transaction.
1045    ///
1046    /// Issues `START TRANSACTION` and captures the identifier Trino returns, so
1047    /// statements issued afterwards on this client run inside the transaction
1048    /// until [`commit`](Self::commit) or [`rollback`](Self::rollback).
1049    ///
1050    /// # Concurrency
1051    ///
1052    /// A transaction is a property of the whole client, so treat a client as
1053    /// single-threaded for as long as one is open. Statements already in flight
1054    /// when the transaction starts do not join it, and statements issued
1055    /// concurrently from another task will run inside it whether or not that
1056    /// was intended.
1057    ///
1058    /// The nesting check below is best-effort, not atomic: the session lock is
1059    /// released before `START TRANSACTION` is sent (holding it would deadlock
1060    /// against the write lock taken when the response is processed). Two tasks
1061    /// calling this concurrently can therefore both pass the check and open two
1062    /// transactions, of which only the last is retained — the other is orphaned
1063    /// on the coordinator until it times out. Use a separate client per
1064    /// transaction if you need concurrency.
1065    ///
1066    /// # Errors
1067    ///
1068    /// Returns [`Error::Transaction`] if a transaction is already active —
1069    /// Trino does not support nested transactions.
1070    ///
1071    /// Also returns [`Error::Transaction`] if the statement succeeded but no
1072    /// usable identifier came back in `X-Trino-Started-Transaction-Id`. A
1073    /// healthy coordinator always sends it, so this signals something between
1074    /// client and coordinator dropping or rewriting the header. The
1075    /// transaction may be open on the coordinator, and because its identifier
1076    /// never reached the client it cannot be committed or rolled back — it
1077    /// stays open until the coordinator times it out. Surfacing that as an
1078    /// error is what lets `Ok(())` mean a transaction is genuinely active.
1079    ///
1080    /// When the statement itself fails the transaction may nevertheless have
1081    /// been started, since the identifier is captured before the statement
1082    /// finishes. Call [`rollback`](Self::rollback) to discard it; that also
1083    /// clears an identifier the coordinator has already expired.
1084    pub async fn begin_transaction(&self) -> Result<()> {
1085        // Bind the guard to a local: holding it across `execute` would deadlock
1086        // against the write lock `update_session` takes.
1087        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        // `execute` drains every page and each response passes through
1097        // `update_session`, so the identifier has had every chance to arrive by
1098        // now. If it still has not, the session is not in a transaction and
1099        // reporting success would put every later statement back on `NONE` —
1100        // the silent failure this API exists to prevent.
1101        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    /// Commit the active transaction.
1112    ///
1113    /// # Errors
1114    ///
1115    /// Returns [`Error::Transaction`] if no transaction is active.
1116    pub async fn commit(&self) -> Result<()> {
1117        self.end_transaction("COMMIT").await
1118    }
1119
1120    /// Roll back the active transaction.
1121    ///
1122    /// # Errors
1123    ///
1124    /// Returns [`Error::Transaction`] if no transaction is active.
1125    pub async fn rollback(&self) -> Result<()> {
1126        self.end_transaction("ROLLBACK").await
1127    }
1128
1129    /// Shared implementation of [`commit`](Self::commit) and
1130    /// [`rollback`](Self::rollback).
1131    ///
1132    /// Trino answers either with `X-Trino-Clear-Transaction-Id`, which
1133    /// `update_session` turns back into `TransactionId::NoTransaction`.
1134    ///
1135    /// Deliberately without the post-condition [`begin_transaction`]
1136    /// (Self::begin_transaction) carries. If the clear header never arrives the
1137    /// session keeps an identifier the coordinator has already retired, but
1138    /// that fails loudly: the next statement sends the dead id and Trino
1139    /// rejects it. The `begin_transaction` case is the dangerous one because
1140    /// there the session falls back to `NONE`, which Trino happily accepts as
1141    /// "no transaction" and executes outside any transaction. Callers that
1142    /// need to recover here can reset the session with
1143    /// [`set_transaction_id`](Self::set_transaction_id).
1144    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        // Submission is not idempotent — retry only when definitely not processed.
1176        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        // Page fetches are idempotent GETs — any transient failure is retryable.
1190        result
1191            .retry(self.retry_policy())
1192            .when(need_retry_fetch)
1193            .await
1194    }
1195
1196    /// Submit `sql` and return the first result page.
1197    ///
1198    /// Low-level building block: the returned [`QueryResult`] may carry a
1199    /// `next_uri` that you must follow with [`get_next`](Client::get_next) to
1200    /// retrieve the rest. Most callers should use [`get_all`](Client::get_all)
1201    /// or [`stream`](Client::stream), which handle pagination.
1202    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    /// Fetch the next result page from a `next_uri` returned by a previous
1227    /// [`get`](Client::get) / `get_next` call.
1228    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    /// Cancel a running query by its id, releasing its resources on the
1249    /// coordinator.
1250    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                // Tokens are acquired lazily: with nothing cached the request
1268                // goes out unauthenticated, and `send` runs the login flow on
1269                // the resulting 401 challenge before retrying once.
1270                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        // Capture the token we are about to authenticate with (if any) so the
1291        // single-flight refresh can tell whether another task already rotated it.
1292        let sent_token = match self.auth.as_ref() {
1293            Some(Auth::OAuth2(state)) => state.cached_token(),
1294            _ => None,
1295        };
1296        // Clone the UN-authed builder up front so an OAuth2 401 can be retried
1297        // with exactly one fresh token header (`bearer_auth` appends, so auth is
1298        // applied only after the clone is taken). Bodies here are `String`s, so
1299        // `try_clone` always succeeds.
1300        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                // A coordinator with several authentication types configured
1306                // (e.g. `http-server.authentication.type=PASSWORD,OAUTH2`) sends
1307                // one `WWW-Authenticate` header per type, in configuration
1308                // order — so `Basic realm="Trino"` may well precede the Bearer
1309                // challenge. Scan all of them for the OAuth2 one.
1310                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    /// Shared response-status handling (extracted so both the first and the
1332    /// retried OAuth2 request go through the same path).
1333    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    /// Acquire an OAuth2 token under a single-flight lock: if another task
1354    /// already refreshed while we waited, reuse that token instead of opening a
1355    /// second browser.
1356    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        // Someone else finished the flow (rotating the token this request was
1364        // sent with) while we waited for the lock — reuse it.
1365        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
1405////////////////////////////////////////////////////////////////////////////////////////////////
1406// helper functions
1407
1408fn 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        // Idempotent page fetches retry every transient gateway/availability status.
1460        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        // Client errors and non-transient 5xx fail fast.
1468        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        // Submission is non-idempotent: only 503 (rejected, not processed) is retried.
1484        assert!(need_retry_submit(&http_not_ok(
1485            StatusCode::SERVICE_UNAVAILABLE
1486        )));
1487        // 502/504 are ambiguous (a proxy may have forwarded the query) -> terminal.
1488        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}