Skip to main content

client/
database.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::collections::HashMap;
16use std::pin::Pin;
17use std::str::FromStr;
18use std::sync::atomic::{AtomicBool, Ordering};
19use std::sync::{Arc, Mutex, RwLock};
20use std::task::{Context, Poll};
21use std::time::Duration;
22
23use api::v1::auth_header::AuthScheme;
24#[cfg(feature = "testing")]
25use api::v1::ddl_request::Expr as DdlExpr;
26use api::v1::greptime_database_client::GreptimeDatabaseClient;
27use api::v1::greptime_request::Request;
28use api::v1::query_request::Query;
29#[cfg(feature = "testing")]
30use api::v1::{AlterTableExpr, CreateTableExpr, DdlRequest};
31use api::v1::{
32    AuthHeader, Basic, GreptimeRequest, InsertRequests, QueryRequest, RequestHeader,
33    RowInsertRequests,
34};
35use arc_swap::ArcSwapOption;
36use arrow_flight::{FlightData, Ticket};
37use async_stream::stream;
38use base64::Engine;
39use base64::prelude::BASE64_STANDARD;
40use common_catalog::build_db_string;
41use common_catalog::consts::{DEFAULT_CATALOG_NAME, DEFAULT_SCHEMA_NAME};
42use common_error::ext::{BoxedError, ErrorExt};
43use common_grpc::flight::do_put::DoPutResponse;
44use common_grpc::flight::{
45    FLOW_EXTENSIONS_METADATA_KEY, FlightDecoder, FlightMessage, SNAPSHOT_SEQS_METADATA_KEY,
46};
47use common_query::Output;
48use common_recordbatch::adapter::RecordBatchMetrics;
49use common_recordbatch::error::ExternalSnafu;
50use common_recordbatch::{OrderOption, RecordBatch, RecordBatchStream, RecordBatchStreamWrapper};
51use common_telemetry::tracing::Span;
52use common_telemetry::tracing_context::W3cTrace;
53use common_telemetry::{error, warn};
54use futures::future;
55use futures_util::{Stream, StreamExt, TryStreamExt};
56use prost::Message;
57use snafu::{IntoError, ResultExt};
58use tokio::sync::Notify;
59use tonic::metadata::{AsciiMetadataKey, AsciiMetadataValue, MetadataMap, MetadataValue};
60use tonic::transport::Channel;
61
62use crate::error::{
63    ConvertFlightDataSnafu, Error, FlightGetSnafu, FlightStreamSnafu, IllegalFlightMessagesSnafu,
64    InvalidTonicMetadataValueSnafu,
65};
66use crate::flight::{FlightMessageReader, decode_flight_data};
67use crate::{Client, Result, error, from_grpc_response};
68
69type FlightDataStream = Pin<Box<dyn Stream<Item = FlightData> + Send>>;
70
71type DoPutResponseStream = Pin<Box<dyn Stream<Item = Result<DoPutResponse>>>>;
72
73const HINTS_METADATA_KEY: &str = "x-greptime-hints";
74/// Maximum time to wait for the optional trailing metrics message after
75/// affected rows have already been delivered.
76const FLIGHT_TRAILING_METRICS_TIMEOUT: Duration = Duration::from_secs(5);
77
78/// Terminal metrics associated with a query output.
79///
80/// For streaming outputs, metrics are only final after the stream is fully
81/// drained and [`Self::is_ready`] returns `true`. Affected-row outputs may
82/// briefly await a compatibility trailing metrics message.
83#[derive(Debug, Clone, Default)]
84pub struct OutputMetrics {
85    inner: Arc<OutputMetricsInner>,
86}
87
88#[derive(Debug)]
89struct OutputMetricsInner {
90    metrics: RwLock<Option<RecordBatchMetrics>>,
91    completion_error: RwLock<Option<String>>,
92    ready: AtomicBool,
93    ready_notify: Notify,
94    compatibility_task: Mutex<Option<tokio::task::AbortHandle>>,
95}
96
97impl Default for OutputMetricsInner {
98    fn default() -> Self {
99        Self {
100            metrics: RwLock::new(None),
101            completion_error: RwLock::new(None),
102            ready: AtomicBool::new(false),
103            ready_notify: Notify::new(),
104            compatibility_task: Mutex::new(None),
105        }
106    }
107}
108
109impl Drop for OutputMetricsInner {
110    fn drop(&mut self) {
111        if let Some(handle) = self.compatibility_task.get_mut().unwrap().take() {
112            handle.abort();
113        }
114    }
115}
116
117impl OutputMetrics {
118    fn new() -> Self {
119        Self::default()
120    }
121
122    /// Replaces the current terminal metrics snapshot.
123    pub fn update(&self, metrics: Option<RecordBatchMetrics>) {
124        *self.inner.metrics.write().unwrap() = metrics;
125    }
126
127    /// Marks the terminal metrics as final for this output.
128    pub fn mark_ready(&self) {
129        if self
130            .inner
131            .ready
132            .compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
133            .is_ok()
134        {
135            self.inner.ready_notify.notify_waiters();
136        }
137    }
138
139    /// Waits until terminal metrics are final.
140    pub async fn wait_ready(&self) {
141        loop {
142            let notified = self.inner.ready_notify.notified();
143            if self.is_ready() {
144                return;
145            }
146            notified.await;
147        }
148    }
149
150    /// Returns an error encountered while completing the output, if any.
151    pub fn completion_error(&self) -> Option<String> {
152        self.inner.completion_error.read().unwrap().clone()
153    }
154
155    fn set_completion_error(&self, error: impl Into<String>) {
156        *self.inner.completion_error.write().unwrap() = Some(error.into());
157    }
158
159    fn set_compatibility_task(&self, handle: tokio::task::AbortHandle) {
160        let mut task = self.inner.compatibility_task.lock().unwrap();
161        if !self.is_ready() {
162            *task = Some(handle);
163        }
164    }
165
166    fn take_compatibility_task(&self) -> Option<tokio::task::AbortHandle> {
167        self.inner.compatibility_task.lock().unwrap().take()
168    }
169
170    /// Returns whether terminal metrics are final.
171    ///
172    /// Streaming outputs become ready only after the stream reaches EOF.
173    pub fn is_ready(&self) -> bool {
174        self.inner.ready.load(Ordering::Acquire)
175    }
176
177    /// Returns the latest terminal metrics snapshot, if any.
178    pub fn get(&self) -> Option<RecordBatchMetrics> {
179        self.inner.metrics.read().unwrap().clone()
180    }
181
182    /// Returns proved per-region watermarks.
183    ///
184    /// Entries whose watermark is `None` are intentionally omitted because they
185    /// represent participating regions whose terminal sequence bound was not
186    /// provable.
187    pub fn region_watermark_map(&self) -> Option<std::collections::HashMap<u64, u64>> {
188        Some(
189            self.get()?
190                .region_watermarks
191                .into_iter()
192                .filter_map(|entry| entry.watermark.map(|seq| (entry.region_id, seq)))
193                .collect::<std::collections::HashMap<_, _>>(),
194        )
195    }
196
197    /// Returns all regions that participated in terminal metric collection,
198    /// including entries whose watermark is `None`.
199    pub fn participating_regions(&self) -> Option<std::collections::BTreeSet<u64>> {
200        Some(
201            self.get()?
202                .region_watermarks
203                .into_iter()
204                .map(|entry| entry.region_id)
205                .collect::<std::collections::BTreeSet<_>>(),
206        )
207    }
208}
209
210/// Query output together with a handle for its terminal metrics.
211///
212/// The contained [`OutputMetrics`] lets callers read stream terminal metrics
213/// after consuming `output`. For non-stream outputs, metrics are ready
214/// immediately. Flight affected-row outputs without inline metrics require
215/// [`OutputMetrics::wait_ready`] before compatibility trailing metrics can be read.
216#[derive(Debug)]
217pub struct OutputWithMetrics {
218    pub output: Output,
219    pub metrics: OutputMetrics,
220}
221
222impl OutputWithMetrics {
223    /// Wraps an output with a terminal metrics handle.
224    ///
225    /// Stream outputs update the handle as the stream is consumed. Non-stream
226    /// outputs are marked ready immediately.
227    pub fn from_output(output: Output) -> Self {
228        let terminal_metrics = OutputMetrics::new();
229        let output = attach_terminal_metrics(output, &terminal_metrics);
230        Self {
231            output,
232            metrics: terminal_metrics,
233        }
234    }
235
236    /// Returns proved per-region watermarks from the terminal metrics.
237    pub fn region_watermark_map(&self) -> Option<std::collections::HashMap<u64, u64>> {
238        self.metrics.region_watermark_map()
239    }
240
241    /// Returns all regions participating in terminal metric collection.
242    pub fn participating_regions(&self) -> Option<std::collections::BTreeSet<u64>> {
243        self.metrics.participating_regions()
244    }
245
246    /// Drops the terminal metrics handle and returns the original output.
247    pub fn into_output(self) -> Output {
248        self.output
249    }
250}
251
252fn parse_terminal_metrics(metrics_json: &str) -> Result<RecordBatchMetrics> {
253    serde_json::from_str(metrics_json).map_err(|e| {
254        IllegalFlightMessagesSnafu {
255            reason: format!("Invalid terminal metrics message: {e}"),
256        }
257        .build()
258    })
259}
260
261fn spawn_affected_rows_trailing_metrics_task<S>(
262    terminal_metrics: &OutputMetrics,
263    mut reader: FlightMessageReader<S>,
264) where
265    S: Stream<Item = Result<FlightMessage>> + Send + Unpin + 'static,
266{
267    let metrics_ref = Arc::downgrade(&terminal_metrics.inner);
268    let remote_addr = reader.remote_addr().to_string();
269    let task = common_runtime::spawn_global(async move {
270        let result =
271            tokio::time::timeout(FLIGHT_TRAILING_METRICS_TIMEOUT, reader.read_next()).await;
272        let Some(inner) = metrics_ref.upgrade() else {
273            return;
274        };
275        let metrics = OutputMetrics { inner };
276        match result {
277            Ok(Ok(Some(FlightMessage::Metrics(s)))) => match parse_terminal_metrics(&s) {
278                Ok(metrics_json) => metrics.update(Some(metrics_json)),
279                Err(error) => {
280                    metrics.set_completion_error(error.to_string());
281                    warn!(
282                        "Failed to decode trailing Flight metrics from {}: {}",
283                        remote_addr, error
284                    );
285                }
286            },
287            Ok(Ok(None)) => {}
288            Ok(Ok(Some(other))) => {
289                let error = format!("Unexpected trailing Flight message: {other:?}");
290                metrics.set_completion_error(error.clone());
291                warn!("{} from {}", error, remote_addr);
292            }
293            Ok(Err(error)) => {
294                let error = flight_stream_error(&remote_addr, error);
295                metrics.set_completion_error(error.to_string());
296                warn!("{}", error);
297            }
298            Err(_) => {
299                let error = "Timed out waiting for trailing Flight metrics";
300                metrics.set_completion_error(error);
301                warn!("{} from {}", error, remote_addr);
302            }
303        }
304        metrics.mark_ready();
305        metrics.take_compatibility_task();
306    });
307    terminal_metrics.set_compatibility_task(task.abort_handle());
308}
309
310struct StreamWithMetrics {
311    stream: common_recordbatch::SendableRecordBatchStream,
312    metrics: OutputMetrics,
313}
314
315impl StreamWithMetrics {
316    fn new(stream: common_recordbatch::SendableRecordBatchStream, metrics: OutputMetrics) -> Self {
317        Self { stream, metrics }
318    }
319
320    fn sync_terminal_metrics(&self) {
321        self.metrics.update(self.stream.metrics());
322    }
323}
324
325impl RecordBatchStream for StreamWithMetrics {
326    fn name(&self) -> &str {
327        self.stream.name()
328    }
329
330    fn schema(&self) -> datatypes::schema::SchemaRef {
331        self.stream.schema()
332    }
333
334    fn output_ordering(&self) -> Option<&[OrderOption]> {
335        self.stream.output_ordering()
336    }
337
338    fn metrics(&self) -> Option<RecordBatchMetrics> {
339        self.sync_terminal_metrics();
340        self.metrics.get()
341    }
342}
343
344impl Stream for StreamWithMetrics {
345    type Item = common_recordbatch::error::Result<RecordBatch>;
346
347    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
348        let polled = Pin::new(&mut self.stream).poll_next(cx);
349        if let Poll::Ready(Some(Err(error))) = &polled {
350            self.metrics.set_completion_error(error.to_string());
351        }
352        if let Poll::Ready(None) = &polled {
353            self.sync_terminal_metrics();
354            self.metrics.mark_ready();
355        }
356        polled
357    }
358
359    fn size_hint(&self) -> (usize, Option<usize>) {
360        self.stream.size_hint()
361    }
362}
363
364fn attach_terminal_metrics(output: Output, terminal_metrics: &OutputMetrics) -> Output {
365    let Output { data, meta } = output;
366    let data = match data {
367        common_query::OutputData::Stream(stream) => {
368            terminal_metrics.update(stream.metrics());
369            common_query::OutputData::Stream(Box::pin(StreamWithMetrics::new(
370                stream,
371                terminal_metrics.clone(),
372            )))
373        }
374        other => {
375            terminal_metrics.mark_ready();
376            other
377        }
378    };
379    Output::new(data, meta)
380}
381
382async fn output_from_flight_message_stream<S>(
383    remote_addr: String,
384    flight_message_stream: S,
385) -> Result<OutputWithMetrics>
386where
387    S: Stream<Item = Result<FlightMessage>> + Send + Unpin + 'static,
388{
389    let mut reader = FlightMessageReader::new(remote_addr, flight_message_stream);
390    let first_flight_message = reader
391        .read_first()
392        .await
393        .map_err(|error| flight_stream_error(reader.remote_addr(), error))?;
394
395    match first_flight_message {
396        FlightMessage::AffectedRows { rows, metrics } => {
397            let terminal_metrics = OutputMetrics::new();
398            if let Some(metrics) = metrics {
399                // Inline metrics are authoritative. They complete the output
400                // without touching the rest of the Flight stream.
401                terminal_metrics.update(Some(parse_terminal_metrics(&metrics)?));
402                terminal_metrics.mark_ready();
403            } else {
404                spawn_affected_rows_trailing_metrics_task(&terminal_metrics, reader);
405            }
406            Ok(OutputWithMetrics {
407                output: Output::new_with_affected_rows(rows),
408                metrics: terminal_metrics,
409            })
410        }
411        FlightMessage::RecordBatch(_) | FlightMessage::Metrics(_) => IllegalFlightMessagesSnafu {
412            reason: "The first flight message cannot be a RecordBatch or Metrics message",
413        }
414        .fail(),
415        FlightMessage::Schema(schema) => {
416            let metrics = Arc::new(ArcSwapOption::from(None));
417            let metrics_ref = metrics.clone();
418            let schema = Arc::new(
419                datatypes::schema::Schema::try_from(schema).context(error::ConvertSchemaSnafu)?,
420            );
421            let schema_cloned = schema.clone();
422            let stream = Box::pin(stream!({
423                loop {
424                    let flight_message = match reader.read_next().await {
425                        Ok(Some(message)) => message,
426                        Ok(None) => break,
427                        Err(error) => {
428                            yield Err(BoxedError::new(flight_stream_error(
429                                reader.remote_addr(),
430                                error,
431                            )))
432                            .context(ExternalSnafu);
433                            break;
434                        }
435                    };
436                    match flight_message {
437                        FlightMessage::RecordBatch(arrow_batch) => {
438                            yield Ok(RecordBatch::from_df_record_batch(
439                                schema_cloned.clone(),
440                                arrow_batch,
441                            ))
442                        }
443                        FlightMessage::Metrics(s) => {
444                            match parse_terminal_metrics(&s) {
445                                Ok(m) => {
446                                    metrics_ref.swap(Some(Arc::new(m)));
447                                }
448                                Err(e) => {
449                                    yield Err(BoxedError::new(e)).context(ExternalSnafu);
450                                }
451                            };
452                        }
453                        FlightMessage::AffectedRows { .. } | FlightMessage::Schema(_) => {
454                            yield IllegalFlightMessagesSnafu {
455                                reason: format!(
456                                    "A Schema message must be succeeded exclusively by a set of RecordBatch messages, flight_message: {:?}",
457                                    flight_message
458                                )
459                            }
460                            .fail()
461                            .map_err(BoxedError::new)
462                            .context(ExternalSnafu);
463                            break;
464                        }
465                    }
466                }
467            }));
468            let record_batch_stream = RecordBatchStreamWrapper {
469                schema,
470                stream,
471                output_ordering: None,
472                metrics,
473                span: Span::current(),
474            };
475            Ok(OutputWithMetrics::from_output(Output::new_with_stream(
476                Box::pin(record_batch_stream),
477            )))
478        }
479    }
480}
481
482fn flight_stream_error(addr: &str, error: Error) -> Error {
483    let tonic_code = error.tonic_code().unwrap_or(tonic::Code::Unknown);
484    let message = error.to_string();
485    if error.status_code().should_log_error() {
486        error!(
487            error; "Failed to receive Flight data, addr: {}, code: {}",
488            addr,
489            tonic_code
490        );
491    }
492
493    FlightStreamSnafu {
494        addr: addr.to_string(),
495        tonic_code,
496        message,
497    }
498    .into_error(BoxedError::new(error))
499}
500
501#[derive(Clone, Debug, Default)]
502pub struct Database {
503    // The "catalog" and "schema" to be used in processing the requests at the server side.
504    // They are the "hint" or "context", just like how the "database" in "USE" statement is treated in MySQL.
505    // They will be carried in the request header.
506    catalog: String,
507    schema: String,
508    // The dbname follows naming rule as out mysql, postgres and http
509    // protocol. The server treat dbname in priority of catalog/schema.
510    dbname: String,
511    // The time zone indicates the time zone where the user is located.
512    // Some queries need to be aware of the user's time zone to perform some specific actions.
513    timezone: String,
514
515    client: Client,
516    ctx: FlightContext,
517}
518
519#[derive(Default)]
520struct FlightRequestOptions {
521    hints: Option<String>,
522    flow_extensions: Option<String>,
523    snapshot_seqs: Option<String>,
524    timeout: Option<Duration>,
525}
526
527impl FlightRequestOptions {
528    fn apply_to<T>(self, request: &mut tonic::Request<T>) -> Result<()> {
529        let metadata = request.metadata_mut();
530        if let Some(hints) = self.hints {
531            Database::put_metadata_value(metadata, HINTS_METADATA_KEY, hints)?;
532        }
533        if let Some(flow_extensions) = self.flow_extensions {
534            Database::put_metadata_value(metadata, FLOW_EXTENSIONS_METADATA_KEY, flow_extensions)?;
535        }
536        if let Some(snapshot_seqs) = self.snapshot_seqs {
537            Database::put_metadata_value(metadata, SNAPSHOT_SEQS_METADATA_KEY, snapshot_seqs)?;
538        }
539        if let Some(timeout) = self.timeout {
540            request.set_timeout(timeout);
541        }
542        Ok(())
543    }
544}
545
546/// A single Flight DoGet request to a [`Database`].
547///
548/// The builder carries request-scoped metadata and options. It does not modify
549/// the underlying [`Database`], so its configuration cannot affect later RPCs.
550pub struct DatabaseFlightRequest<'a> {
551    database: &'a Database,
552    options: FlightRequestOptions,
553}
554
555pub struct DatabaseClient {
556    pub addr: String,
557    pub inner: GreptimeDatabaseClient<Channel>,
558}
559
560impl DatabaseClient {
561    /// Returns a closure that logs the error when the request fails.
562    pub fn inspect_err<'a>(&'a self, context: &'a str) -> impl Fn(&tonic::Status) + 'a {
563        let addr = &self.addr;
564        move |status| {
565            error!("Failed to {context} request, peer: {addr}, status: {status:?}");
566        }
567    }
568}
569
570fn make_database_client(client: &Client) -> Result<DatabaseClient> {
571    let (addr, channel) = client.find_channel()?;
572    Ok(DatabaseClient {
573        addr,
574        inner: GreptimeDatabaseClient::new(channel)
575            .max_decoding_message_size(client.max_grpc_recv_message_size())
576            .max_encoding_message_size(client.max_grpc_send_message_size()),
577    })
578}
579
580impl Database {
581    /// Create database service client using catalog and schema
582    pub fn new(catalog: impl Into<String>, schema: impl Into<String>, client: Client) -> Self {
583        Self {
584            catalog: catalog.into(),
585            schema: schema.into(),
586            dbname: String::default(),
587            timezone: String::default(),
588            client,
589            ctx: FlightContext::default(),
590        }
591    }
592
593    /// Create database service client using dbname.
594    ///
595    /// This API is designed for external usage. `dbname` is:
596    ///
597    /// - the name of database when using GreptimeDB standalone or cluster
598    /// - the name provided by GreptimeCloud or other multi-tenant GreptimeDB
599    ///   environment
600    pub fn new_with_dbname(dbname: impl Into<String>, client: Client) -> Self {
601        Self {
602            catalog: String::default(),
603            schema: String::default(),
604            timezone: String::default(),
605            dbname: dbname.into(),
606            client,
607            ctx: FlightContext::default(),
608        }
609    }
610
611    /// Set the catalog for the database client.
612    pub fn set_catalog(&mut self, catalog: impl Into<String>) {
613        self.catalog = catalog.into();
614    }
615
616    fn catalog_or_default(&self) -> &str {
617        if self.catalog.is_empty() {
618            DEFAULT_CATALOG_NAME
619        } else {
620            &self.catalog
621        }
622    }
623
624    /// Set the schema for the database client.
625    pub fn set_schema(&mut self, schema: impl Into<String>) {
626        self.schema = schema.into();
627    }
628
629    fn schema_or_default(&self) -> &str {
630        if self.schema.is_empty() {
631            DEFAULT_SCHEMA_NAME
632        } else {
633            &self.schema
634        }
635    }
636
637    /// Set the timezone for the database client.
638    pub fn set_timezone(&mut self, timezone: impl Into<String>) {
639        self.timezone = timezone.into();
640    }
641
642    /// Set the auth scheme for the database client.
643    pub fn set_auth(&mut self, auth: AuthScheme) {
644        self.ctx.auth_header = Some(AuthHeader {
645            auth_scheme: Some(auth),
646        });
647    }
648
649    /// Creates a builder for a single Flight DoGet request.
650    pub fn flight_request(&self) -> DatabaseFlightRequest<'_> {
651        DatabaseFlightRequest {
652            database: self,
653            options: FlightRequestOptions::default(),
654        }
655    }
656
657    /// Make an InsertRequests request to the database.
658    pub async fn insert(&self, requests: InsertRequests) -> Result<u32> {
659        self.handle(Request::Inserts(requests)).await
660    }
661
662    /// Make an InsertRequests request to the database with hints.
663    pub async fn insert_with_hints(
664        &self,
665        requests: InsertRequests,
666        hints: &[(&str, &str)],
667    ) -> Result<u32> {
668        let mut client = make_database_client(&self.client)?;
669        let request = self.to_rpc_request(Request::Inserts(requests));
670
671        let mut request = tonic::Request::new(request);
672        let metadata = request.metadata_mut();
673        Self::put_hints(metadata, hints)?;
674
675        let response = client
676            .inner
677            .handle(request)
678            .await
679            .inspect_err(client.inspect_err("insert_with_hints"))?
680            .into_inner();
681        from_grpc_response(response)
682    }
683
684    /// Make a RowInsertRequests request to the database.
685    pub async fn row_inserts(&self, requests: RowInsertRequests) -> Result<u32> {
686        self.handle(Request::RowInserts(requests)).await
687    }
688
689    /// Make a RowInsertRequests request to the database with hints.
690    pub async fn row_inserts_with_hints(
691        &self,
692        requests: RowInsertRequests,
693        hints: &[(&str, &str)],
694    ) -> Result<u32> {
695        let mut client = make_database_client(&self.client)?;
696        let request = self.to_rpc_request(Request::RowInserts(requests));
697
698        let mut request = tonic::Request::new(request);
699        let metadata = request.metadata_mut();
700        Self::put_hints(metadata, hints)?;
701
702        let response = client
703            .inner
704            .handle(request)
705            .await
706            .inspect_err(client.inspect_err("row_inserts_with_hints"))?
707            .into_inner();
708        from_grpc_response(response)
709    }
710
711    fn put_hints(metadata: &mut MetadataMap, hints: &[(&str, &str)]) -> Result<()> {
712        let Some(value) = Self::encode_hints(hints) else {
713            return Ok(());
714        };
715
716        Self::put_metadata_value(metadata, HINTS_METADATA_KEY, value)
717    }
718
719    fn encode_hints(hints: &[(&str, &str)]) -> Option<String> {
720        hints
721            .iter()
722            .map(|(k, v)| format!("{}={}", k, v))
723            .reduce(|a, b| format!("{},{}", a, b))
724    }
725
726    fn encode_flow_extensions(flow_extensions: &[(&str, &str)]) -> Option<String> {
727        (!flow_extensions.is_empty()).then(|| {
728            serde_json::to_string(&flow_extensions.to_vec())
729                .expect("flow extension pairs should serialize")
730        })
731    }
732
733    fn encode_snapshot_seqs(snapshot_seqs: &HashMap<u64, u64>) -> Option<String> {
734        (!snapshot_seqs.is_empty()).then(|| {
735            serde_json::to_string(snapshot_seqs).expect("snapshot sequence map should serialize")
736        })
737    }
738
739    fn put_metadata_value(
740        metadata: &mut MetadataMap,
741        key: &'static str,
742        value: String,
743    ) -> Result<()> {
744        let key = AsciiMetadataKey::from_static(key);
745        let value = AsciiMetadataValue::from_str(&value).context(InvalidTonicMetadataValueSnafu)?;
746        metadata.insert(key, value);
747        Ok(())
748    }
749
750    /// Make a request to the database.
751    pub async fn handle(&self, request: Request) -> Result<u32> {
752        let mut client = make_database_client(&self.client)?;
753        let request = self.to_rpc_request(request);
754        let response = client
755            .inner
756            .handle(request)
757            .await
758            .inspect_err(client.inspect_err("handle"))?
759            .into_inner();
760        from_grpc_response(response)
761    }
762
763    /// Retry if connection fails, max_retries is the max number of retries, so the total wait time
764    /// is `max_retries * GRPC_CONN_TIMEOUT`
765    pub async fn handle_with_retry(
766        &self,
767        request: Request,
768        max_retries: u32,
769        hints: &[(&str, &str)],
770    ) -> Result<u32> {
771        let mut client = make_database_client(&self.client)?;
772        let mut retries = 0;
773
774        let request = self.to_rpc_request(request);
775
776        loop {
777            let mut tonic_request = tonic::Request::new(request.clone());
778            let metadata = tonic_request.metadata_mut();
779            Self::put_hints(metadata, hints)?;
780            let raw_response = client
781                .inner
782                .handle(tonic_request)
783                .await
784                .inspect_err(client.inspect_err("handle"));
785            match (raw_response, retries < max_retries) {
786                (Ok(resp), _) => return from_grpc_response(resp.into_inner()),
787                (Err(err), true) => {
788                    // determine if the error is retryable
789                    if is_grpc_retryable(&err) {
790                        // retry
791                        retries += 1;
792                        warn!("Retrying {} times with error = {:?}", retries, err);
793                        continue;
794                    } else {
795                        error!(
796                            err; "Failed to send request to grpc handle, retries = {}, not retryable error, aborting",
797                            retries
798                        );
799                        return Err(err.into());
800                    }
801                }
802                (Err(err), false) => {
803                    error!(
804                        err; "Failed to send request to grpc handle after {} retries",
805                        retries,
806                    );
807                    return Err(err.into());
808                }
809            }
810        }
811    }
812
813    #[inline]
814    fn to_rpc_request(&self, request: Request) -> GreptimeRequest {
815        GreptimeRequest {
816            header: Some(RequestHeader {
817                catalog: self.catalog.clone(),
818                schema: self.schema.clone(),
819                authorization: self.ctx.auth_header.clone(),
820                dbname: self.dbname.clone(),
821                timezone: self.timezone.clone(),
822                // TODO(Taylor-lagrange): add client grpc tracing
823                tracing_context: W3cTrace::new(),
824            }),
825            request: Some(request),
826        }
827    }
828
829    /// Executes a SQL query without any hints.
830    pub async fn sql<S>(&self, sql: S) -> Result<Output>
831    where
832        S: AsRef<str>,
833    {
834        self.flight_request().sql(sql).await
835    }
836
837    /// Executes a SQL query with optional hints for query optimization.
838    pub async fn sql_with_hint<S>(&self, sql: S, hints: &[(&str, &str)]) -> Result<Output>
839    where
840        S: AsRef<str>,
841    {
842        self.flight_request().with_hints(hints).sql(sql).await
843    }
844
845    /// Executes a SQL query and returns the output with terminal metrics.
846    ///
847    /// For stream outputs, callers must consume the stream before reading final
848    /// terminal metrics from [`OutputWithMetrics::metrics`]. For affected-row
849    /// outputs without inline metrics, call [`OutputMetrics::wait_ready`] when
850    /// compatibility trailing metrics are required.
851    pub async fn sql_with_terminal_metrics<S>(
852        &self,
853        sql: S,
854        hints: &[(&str, &str)],
855    ) -> Result<OutputWithMetrics>
856    where
857        S: AsRef<str>,
858    {
859        self.flight_request()
860            .with_hints(hints)
861            .sql_with_terminal_metrics(sql)
862            .await
863    }
864
865    /// Executes a logical plan directly without SQL parsing.
866    pub async fn logical_plan(&self, logical_plan: Vec<u8>) -> Result<Output> {
867        self.flight_request().logical_plan(logical_plan).await
868    }
869
870    /// Creates a new table using the provided table expression.
871    #[cfg(feature = "testing")]
872    pub async fn create(&self, expr: CreateTableExpr) -> Result<Output> {
873        self.flight_request().create(expr).await
874    }
875
876    /// Alters an existing table using the provided alter expression.
877    #[cfg(feature = "testing")]
878    pub async fn alter(&self, expr: AlterTableExpr) -> Result<Output> {
879        self.flight_request().alter(expr).await
880    }
881
882    async fn do_get(
883        &self,
884        request: Request,
885        options: FlightRequestOptions,
886    ) -> Result<OutputWithMetrics> {
887        let request = self.to_rpc_request(request);
888        let request = Ticket {
889            ticket: request.encode_to_vec().into(),
890        };
891
892        let mut request = tonic::Request::new(request);
893        options.apply_to(&mut request)?;
894
895        let mut client = self.client.make_flight_client(false, false)?;
896        let remote_addr = client.addr().to_string();
897
898        let response = client.mut_inner().do_get(request).await.or_else(|e| {
899            let tonic_code = e.code();
900            let e: Error = e.into();
901            error!(
902                "Failed to do Flight get, addr: {}, code: {}, source: {:?}",
903                client.addr(),
904                tonic_code,
905                e
906            );
907            Err(BoxedError::new(e)).with_context(|_| FlightGetSnafu {
908                addr: remote_addr.clone(),
909                tonic_code,
910            })
911        })?;
912
913        let flight_data_stream = response.into_inner();
914        let mut decoder = FlightDecoder::default();
915
916        let flight_message_stream = flight_data_stream.filter_map(move |flight_data| {
917            future::ready(decode_flight_data(&mut decoder, flight_data))
918        });
919
920        output_from_flight_message_stream(remote_addr, flight_message_stream).await
921    }
922
923    /// Ingest a stream of [RecordBatch]es that belong to a table, using Arrow Flight's "`DoPut`"
924    /// method. The return value is also a stream, produces [DoPutResponse]s.
925    pub async fn do_put(&self, stream: FlightDataStream) -> Result<DoPutResponseStream> {
926        self.do_put_with_hints(stream, &[]).await
927    }
928
929    /// Ingest a stream of [RecordBatch]es using Arrow Flight's `DoPut` with request hints.
930    pub async fn do_put_with_hints(
931        &self,
932        stream: FlightDataStream,
933        hints: &[(&str, &str)],
934    ) -> Result<DoPutResponseStream> {
935        let mut request = tonic::Request::new(stream);
936        Self::put_hints(request.metadata_mut(), hints)?;
937
938        if let Some(AuthHeader {
939            auth_scheme: Some(AuthScheme::Basic(Basic { username, password })),
940        }) = &self.ctx.auth_header
941        {
942            let encoded = BASE64_STANDARD.encode(format!("{username}:{password}"));
943            let value = MetadataValue::from_str(&format!("Basic {encoded}"))
944                .context(InvalidTonicMetadataValueSnafu)?;
945            request.metadata_mut().insert("x-greptime-auth", value);
946        }
947
948        let db_to_put = if !self.dbname.is_empty() {
949            &self.dbname
950        } else {
951            &build_db_string(self.catalog_or_default(), self.schema_or_default())
952        };
953        request.metadata_mut().insert(
954            "x-greptime-db-name",
955            MetadataValue::from_str(db_to_put).context(InvalidTonicMetadataValueSnafu)?,
956        );
957
958        let mut client = self.client.make_control_flight_client(false, false)?;
959        let response = client.mut_inner().do_put(request).await?;
960        let response = response
961            .into_inner()
962            .map_err(Into::into)
963            .and_then(|x| future::ready(DoPutResponse::try_from(x).context(ConvertFlightDataSnafu)))
964            .boxed();
965        Ok(response)
966    }
967}
968
969impl<'a> DatabaseFlightRequest<'a> {
970    /// Adds query optimization hints to this Flight request.
971    pub fn with_hints(mut self, hints: &[(&str, &str)]) -> Self {
972        self.options.hints = Database::encode_hints(hints);
973        self
974    }
975
976    /// Adds Flow extensions to this Flight request.
977    pub fn with_flow_extensions(mut self, flow_extensions: &[(&str, &str)]) -> Self {
978        self.options.flow_extensions = Database::encode_flow_extensions(flow_extensions);
979        self
980    }
981
982    /// Adds snapshot sequence fences to this Flight request.
983    pub fn with_snapshot_seqs(mut self, snapshot_seqs: &HashMap<u64, u64>) -> Self {
984        self.options.snapshot_seqs = Database::encode_snapshot_seqs(snapshot_seqs);
985        self
986    }
987
988    /// Sets a timeout for this Flight request only.
989    pub fn with_timeout(mut self, timeout: Duration) -> Self {
990        self.options.timeout = Some(timeout);
991        self
992    }
993
994    /// Executes a SQL query.
995    pub async fn sql<S>(self, sql: S) -> Result<Output>
996    where
997        S: AsRef<str>,
998    {
999        let request = Request::Query(QueryRequest {
1000            query: Some(Query::Sql(sql.as_ref().to_string())),
1001        });
1002        self.do_get(request)
1003            .await
1004            .map(OutputWithMetrics::into_output)
1005    }
1006
1007    /// Executes a SQL query and returns terminal metrics.
1008    pub async fn sql_with_terminal_metrics<S>(self, sql: S) -> Result<OutputWithMetrics>
1009    where
1010        S: AsRef<str>,
1011    {
1012        self.query_with_terminal_metrics(QueryRequest {
1013            query: Some(Query::Sql(sql.as_ref().to_string())),
1014        })
1015        .await
1016    }
1017
1018    /// Executes a logical plan directly without SQL parsing.
1019    pub async fn logical_plan(self, logical_plan: Vec<u8>) -> Result<Output> {
1020        self.query_with_terminal_metrics(QueryRequest {
1021            query: Some(Query::LogicalPlan(logical_plan)),
1022        })
1023        .await
1024        .map(OutputWithMetrics::into_output)
1025    }
1026
1027    /// Executes a query and returns terminal metrics.
1028    pub async fn query_with_terminal_metrics(
1029        self,
1030        request: QueryRequest,
1031    ) -> Result<OutputWithMetrics> {
1032        self.do_get(Request::Query(request)).await
1033    }
1034
1035    /// Creates a new table using the provided table expression.
1036    #[cfg(feature = "testing")]
1037    pub async fn create(self, expr: CreateTableExpr) -> Result<Output> {
1038        self.do_get(Request::Ddl(DdlRequest {
1039            expr: Some(DdlExpr::CreateTable(expr)),
1040        }))
1041        .await
1042        .map(OutputWithMetrics::into_output)
1043    }
1044
1045    /// Alters an existing table using the provided alter expression.
1046    #[cfg(feature = "testing")]
1047    pub async fn alter(self, expr: AlterTableExpr) -> Result<Output> {
1048        self.do_get(Request::Ddl(DdlRequest {
1049            expr: Some(DdlExpr::AlterTable(expr)),
1050        }))
1051        .await
1052        .map(OutputWithMetrics::into_output)
1053    }
1054
1055    async fn do_get(self, request: Request) -> Result<OutputWithMetrics> {
1056        let Self { database, options } = self;
1057        database.do_get(request, options).await
1058    }
1059}
1060
1061/// by grpc standard, only `Unavailable` is retryable, see: https://github.com/grpc/grpc/blob/master/doc/statuscodes.md#status-codes-and-their-use-in-grpc
1062pub fn is_grpc_retryable(err: &tonic::Status) -> bool {
1063    matches!(err.code(), tonic::Code::Unavailable)
1064}
1065
1066#[derive(Default, Debug, Clone)]
1067struct FlightContext {
1068    auth_header: Option<AuthHeader>,
1069}
1070
1071#[cfg(test)]
1072mod tests {
1073    use std::sync::Arc;
1074    use std::task::{Context, Poll};
1075
1076    use api::v1::auth_header::AuthScheme;
1077    use api::v1::{AuthHeader, Basic};
1078    use common_error::ext::{ErrorExt, RetryHint};
1079    use common_error::status_code::StatusCode;
1080    use common_error::{GREPTIME_DB_HEADER_ERROR_CODE, GREPTIME_DB_HEADER_ERROR_RETRY_HINT};
1081    use common_query::OutputData;
1082    use common_recordbatch::{OrderOption, RecordBatch, RecordBatchStream};
1083    use datatypes::prelude::{ConcreteDataType, VectorRef};
1084    use datatypes::schema::{ColumnSchema, Schema};
1085    use datatypes::vectors::Int32Vector;
1086    use futures_util::StreamExt;
1087    use tokio::sync::oneshot;
1088    use tonic::codegen::http::{HeaderMap, HeaderValue};
1089    use tonic::metadata::MetadataMap;
1090    use tonic::{Code, Status};
1091
1092    use super::*;
1093    use crate::error::TonicSnafu;
1094
1095    struct MockMetricsStream {
1096        schema: datatypes::schema::SchemaRef,
1097        batch: Option<RecordBatch>,
1098        metrics: RecordBatchMetrics,
1099        terminal_metrics_only: bool,
1100    }
1101
1102    impl Stream for MockMetricsStream {
1103        type Item = common_recordbatch::error::Result<RecordBatch>;
1104
1105        fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
1106            Poll::Ready(self.batch.take().map(Ok))
1107        }
1108    }
1109
1110    impl RecordBatchStream for MockMetricsStream {
1111        fn name(&self) -> &str {
1112            "MockMetricsStream"
1113        }
1114
1115        fn schema(&self) -> datatypes::schema::SchemaRef {
1116            self.schema.clone()
1117        }
1118
1119        fn output_ordering(&self) -> Option<&[OrderOption]> {
1120            None
1121        }
1122
1123        fn metrics(&self) -> Option<RecordBatchMetrics> {
1124            if self.terminal_metrics_only && self.batch.is_some() {
1125                return None;
1126            }
1127            Some(self.metrics.clone())
1128        }
1129    }
1130
1131    fn terminal_metrics_json() -> String {
1132        terminal_metrics_json_with_seq(42)
1133    }
1134
1135    fn terminal_metrics_json_with_seq(seq: u64) -> String {
1136        serde_json::to_string(&RecordBatchMetrics {
1137            region_watermarks: vec![common_recordbatch::adapter::RegionWatermarkEntry {
1138                region_id: 7,
1139                watermark: Some(seq),
1140            }],
1141            ..Default::default()
1142        })
1143        .unwrap()
1144    }
1145
1146    #[test]
1147    fn test_put_flow_extensions_preserves_comma_bearing_values() {
1148        let mut metadata = MetadataMap::new();
1149        Database::put_metadata_value(
1150            &mut metadata,
1151            FLOW_EXTENSIONS_METADATA_KEY,
1152            Database::encode_flow_extensions(&[
1153                ("flow.return_region_seq", "true"),
1154                ("flow.incremental_after_seqs", r#"{"1":10,"2":20}"#),
1155            ])
1156            .unwrap(),
1157        )
1158        .unwrap();
1159
1160        let value = metadata
1161            .get(FLOW_EXTENSIONS_METADATA_KEY)
1162            .unwrap()
1163            .to_str()
1164            .unwrap();
1165        let decoded: Vec<(String, String)> = serde_json::from_str(value).unwrap();
1166        assert_eq!(
1167            decoded,
1168            vec![
1169                ("flow.return_region_seq".to_string(), "true".to_string()),
1170                (
1171                    "flow.incremental_after_seqs".to_string(),
1172                    r#"{"1":10,"2":20}"#.to_string()
1173                ),
1174            ]
1175        );
1176    }
1177
1178    #[test]
1179    fn test_put_snapshot_seqs_preserves_u64_precision() {
1180        let mut metadata = MetadataMap::new();
1181        let snapshot_seqs = std::collections::HashMap::from([
1182            (u64::MAX, u64::MAX - 1),
1183            (9_007_199_254_740_993_u64, 9_007_199_254_740_995_u64),
1184        ]);
1185
1186        Database::put_metadata_value(
1187            &mut metadata,
1188            SNAPSHOT_SEQS_METADATA_KEY,
1189            Database::encode_snapshot_seqs(&snapshot_seqs).unwrap(),
1190        )
1191        .unwrap();
1192
1193        let value = metadata
1194            .get(SNAPSHOT_SEQS_METADATA_KEY)
1195            .unwrap()
1196            .to_str()
1197            .unwrap();
1198        let decoded: std::collections::HashMap<u64, u64> = serde_json::from_str(value).unwrap();
1199        assert_eq!(decoded, snapshot_seqs);
1200    }
1201
1202    #[test]
1203    fn test_flight_request_builder_applies_request_options() {
1204        let database = Database::new("greptime", "public", Client::default());
1205        let snapshot_seqs = HashMap::from([(42, 99)]);
1206        let request = database
1207            .flight_request()
1208            .with_hints(&[("query_parallelism", "1")])
1209            .with_flow_extensions(&[("flow.return_region_seq", "true")])
1210            .with_snapshot_seqs(&snapshot_seqs)
1211            .with_timeout(Duration::from_millis(50));
1212        let mut tonic_request = tonic::Request::new(());
1213
1214        request.options.apply_to(&mut tonic_request).unwrap();
1215
1216        let metadata = tonic_request.metadata();
1217        assert_eq!(
1218            metadata.get(HINTS_METADATA_KEY).unwrap(),
1219            "query_parallelism=1"
1220        );
1221        assert_eq!(
1222            serde_json::from_str::<Vec<(String, String)>>(
1223                metadata
1224                    .get(FLOW_EXTENSIONS_METADATA_KEY)
1225                    .unwrap()
1226                    .to_str()
1227                    .unwrap(),
1228            )
1229            .unwrap(),
1230            vec![("flow.return_region_seq".to_string(), "true".to_string())]
1231        );
1232        assert_eq!(
1233            serde_json::from_str::<HashMap<u64, u64>>(
1234                metadata
1235                    .get(SNAPSHOT_SEQS_METADATA_KEY)
1236                    .unwrap()
1237                    .to_str()
1238                    .unwrap(),
1239            )
1240            .unwrap(),
1241            snapshot_seqs
1242        );
1243        assert!(metadata.get("grpc-timeout").is_some());
1244    }
1245
1246    #[test]
1247    fn test_flight_ctx() {
1248        let mut ctx = FlightContext::default();
1249        assert!(ctx.auth_header.is_none());
1250
1251        let basic = AuthScheme::Basic(Basic {
1252            username: "u".to_string(),
1253            password: "p".to_string(),
1254        });
1255
1256        ctx.auth_header = Some(AuthHeader {
1257            auth_scheme: Some(basic),
1258        });
1259
1260        assert!(matches!(
1261            ctx.auth_header,
1262            Some(AuthHeader {
1263                auth_scheme: Some(AuthScheme::Basic(_)),
1264            })
1265        ));
1266    }
1267
1268    #[test]
1269    fn test_from_tonic_status() {
1270        let expected = TonicSnafu {
1271            code: StatusCode::Internal,
1272            msg: "blabla".to_string(),
1273            tonic_code: Code::Internal,
1274            retry_hint: RetryHint::NonRetryable,
1275        }
1276        .build();
1277
1278        let status = Status::new(Code::Internal, "blabla");
1279        let actual: Error = status.into();
1280
1281        assert_eq!(expected.to_string(), actual.to_string());
1282        assert_eq!(expected.retry_hint(), actual.retry_hint());
1283        assert_eq!(expected.should_retry(), actual.should_retry());
1284    }
1285
1286    #[test]
1287    fn test_flight_stream_error_preserves_addr_and_message() {
1288        let error = flight_stream_error(
1289            "127.0.0.1:4001",
1290            Status::out_of_range("message length too large").into(),
1291        );
1292
1293        assert!(matches!(
1294            &error,
1295            Error::FlightStream {
1296                addr,
1297                tonic_code: Code::OutOfRange,
1298                message,
1299                ..
1300            } if addr == "127.0.0.1:4001" && message == "message length too large"
1301        ));
1302        assert_eq!(
1303            "Failed to receive Flight data from 127.0.0.1:4001, code: Operation was attempted past the valid range: message length too large",
1304            error.to_string(),
1305        );
1306    }
1307
1308    #[test]
1309    fn test_from_tonic_status_with_retry_hint() {
1310        let mut headers = HeaderMap::new();
1311        headers.insert(
1312            GREPTIME_DB_HEADER_ERROR_CODE,
1313            HeaderValue::from(StatusCode::Internal as u32),
1314        );
1315        headers.insert(
1316            GREPTIME_DB_HEADER_ERROR_RETRY_HINT,
1317            HeaderValue::from_static(RetryHint::Retryable.as_str()),
1318        );
1319        let status =
1320            Status::with_metadata(Code::Internal, "blabla", MetadataMap::from_headers(headers));
1321
1322        let actual: Error = status.into();
1323
1324        assert_eq!(actual.retry_hint(), RetryHint::Retryable);
1325        assert!(actual.should_retry());
1326    }
1327
1328    #[test]
1329    fn test_from_tonic_status_fallback() {
1330        let mut headers = HeaderMap::new();
1331        headers.insert(
1332            GREPTIME_DB_HEADER_ERROR_CODE,
1333            HeaderValue::from(StatusCode::InvalidArguments as u32),
1334        );
1335        let status =
1336            Status::with_metadata(Code::Internal, "blabla", MetadataMap::from_headers(headers));
1337
1338        let actual: Error = status.into();
1339
1340        assert_eq!(actual.retry_hint(), RetryHint::NonRetryable);
1341        assert!(!actual.should_retry());
1342    }
1343
1344    #[test]
1345    fn test_should_retry_preserves_transport_retry() {
1346        let status = Status::new(Code::Unavailable, "blabla");
1347        let actual: Error = status.into();
1348
1349        assert!(actual.should_retry());
1350    }
1351
1352    #[tokio::test]
1353    async fn test_query_with_terminal_metrics_tracks_terminal_only_metrics() {
1354        let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1355            "v",
1356            ConcreteDataType::int32_datatype(),
1357            false,
1358        )]));
1359        let batch = RecordBatch::new(
1360            schema.clone(),
1361            vec![Arc::new(Int32Vector::from_slice([1, 2])) as VectorRef],
1362        )
1363        .unwrap();
1364        let output = Output::new_with_stream(Box::pin(MockMetricsStream {
1365            schema,
1366            batch: Some(batch),
1367            metrics: RecordBatchMetrics {
1368                region_watermarks: vec![common_recordbatch::adapter::RegionWatermarkEntry {
1369                    region_id: 7,
1370                    watermark: Some(42),
1371                }],
1372                ..Default::default()
1373            },
1374            terminal_metrics_only: true,
1375        }));
1376
1377        let result = OutputWithMetrics::from_output(output);
1378        let terminal_metrics = result.metrics.clone();
1379        assert!(!terminal_metrics.is_ready());
1380        assert!(terminal_metrics.get().is_none());
1381
1382        let OutputData::Stream(mut stream) = result.output.data else {
1383            panic!("expected stream output");
1384        };
1385        while stream.next().await.is_some() {}
1386
1387        assert!(terminal_metrics.is_ready());
1388        assert_eq!(
1389            terminal_metrics.participating_regions(),
1390            Some(std::collections::BTreeSet::from([7_u64]))
1391        );
1392        assert_eq!(
1393            terminal_metrics.region_watermark_map(),
1394            Some(std::collections::HashMap::from([(7_u64, 42_u64)]))
1395        );
1396    }
1397
1398    #[tokio::test]
1399    async fn test_affected_rows_inline_metrics_are_parsed() {
1400        let output = output_from_flight_message_stream(
1401            "test-peer".to_string(),
1402            futures_util::stream::iter(vec![Ok(FlightMessage::AffectedRows {
1403                rows: 3,
1404                metrics: Some(terminal_metrics_json()),
1405            })] as Vec<Result<FlightMessage>>),
1406        )
1407        .await
1408        .unwrap();
1409
1410        assert!(matches!(output.output.data, OutputData::AffectedRows(3)));
1411        assert!(output.metrics.is_ready());
1412        assert_eq!(
1413            output.metrics.region_watermark_map(),
1414            Some(std::collections::HashMap::from([(7, 42)]))
1415        );
1416    }
1417
1418    #[tokio::test]
1419    async fn test_affected_rows_inline_metrics_do_not_poll_trailer() {
1420        let metrics_json = terminal_metrics_json();
1421        let output = output_from_flight_message_stream(
1422            "test-peer".to_string(),
1423            futures_util::stream::iter(vec![
1424                Ok(FlightMessage::AffectedRows {
1425                    rows: 3,
1426                    metrics: Some(metrics_json),
1427                }),
1428                Ok(FlightMessage::Metrics(terminal_metrics_json_with_seq(99))),
1429            ] as Vec<Result<FlightMessage>>),
1430        )
1431        .await
1432        .unwrap();
1433
1434        assert!(output.metrics.is_ready());
1435        assert_eq!(
1436            output.metrics.region_watermark_map(),
1437            Some(std::collections::HashMap::from([(7, 42)]))
1438        );
1439    }
1440
1441    #[tokio::test]
1442    async fn test_affected_rows_without_inline_metrics_becomes_ready_after_trailer() {
1443        // Hold the trailer until the not-ready state has been observed.
1444        let (trailer_tx, trailer_rx) = oneshot::channel();
1445        let trailer = futures_util::stream::once(trailer_rx).map(|ready| {
1446            ready.unwrap();
1447            Ok(FlightMessage::Metrics(terminal_metrics_json()))
1448        });
1449        let output = output_from_flight_message_stream(
1450            "test-peer".to_string(),
1451            futures_util::stream::iter(vec![Ok(FlightMessage::AffectedRows {
1452                rows: 3,
1453                metrics: None,
1454            })])
1455            .chain(trailer),
1456        )
1457        .await
1458        .unwrap();
1459
1460        assert!(!output.metrics.is_ready());
1461        trailer_tx.send(()).unwrap();
1462        tokio::time::timeout(Duration::from_secs(1), output.metrics.wait_ready())
1463            .await
1464            .expect("terminal metrics must become ready once the trailer arrives");
1465        assert!(output.metrics.completion_error().is_none());
1466        assert_eq!(
1467            output.metrics.region_watermark_map(),
1468            Some(std::collections::HashMap::from([(7, 42)]))
1469        );
1470    }
1471
1472    #[tokio::test]
1473    async fn test_affected_rows_malformed_trailer_sets_completion_error_and_ready() {
1474        let output = output_from_flight_message_stream(
1475            "test-peer".to_string(),
1476            futures_util::stream::iter(vec![
1477                Ok(FlightMessage::AffectedRows {
1478                    rows: 3,
1479                    metrics: None,
1480                }),
1481                Ok(FlightMessage::Metrics("{not-json}".to_string())),
1482            ] as Vec<Result<FlightMessage>>),
1483        )
1484        .await
1485        .unwrap();
1486
1487        output.metrics.wait_ready().await;
1488        let error = output.metrics.completion_error().unwrap();
1489        assert!(error.contains("Invalid terminal metrics message"));
1490        assert!(output.metrics.is_ready());
1491    }
1492
1493    #[tokio::test]
1494    async fn test_affected_rows_transport_trailer_error_sets_completion_error_and_ready() {
1495        let output = output_from_flight_message_stream(
1496            "test-peer".to_string(),
1497            futures_util::stream::iter(vec![
1498                Ok(FlightMessage::AffectedRows {
1499                    rows: 3,
1500                    metrics: None,
1501                }),
1502                Err(Status::unavailable("trailer read failed").into()),
1503            ] as Vec<Result<FlightMessage>>),
1504        )
1505        .await
1506        .unwrap();
1507
1508        output.metrics.wait_ready().await;
1509        let error = output.metrics.completion_error().unwrap();
1510        assert!(error.contains("trailer read failed"));
1511        assert!(output.metrics.is_ready());
1512    }
1513
1514    #[tokio::test]
1515    async fn test_affected_rows_compatibility_reader_is_cancelled_after_second_poll_begins() {
1516        struct DropProbe {
1517            first: Option<Result<FlightMessage>>,
1518            polled: Option<oneshot::Sender<()>>,
1519            dropped: Option<oneshot::Sender<()>>,
1520        }
1521
1522        impl Stream for DropProbe {
1523            type Item = Result<FlightMessage>;
1524
1525            fn poll_next(
1526                mut self: Pin<&mut Self>,
1527                _cx: &mut Context<'_>,
1528            ) -> Poll<Option<Self::Item>> {
1529                if let Some(message) = self.first.take() {
1530                    return Poll::Ready(Some(message));
1531                }
1532                if let Some(polled) = self.polled.take() {
1533                    let _ = polled.send(());
1534                }
1535                Poll::Pending
1536            }
1537        }
1538
1539        impl Drop for DropProbe {
1540            fn drop(&mut self) {
1541                if let Some(dropped) = self.dropped.take() {
1542                    let _ = dropped.send(());
1543                }
1544            }
1545        }
1546
1547        let (polled_tx, polled_rx) = oneshot::channel();
1548        let (dropped_tx, dropped_rx) = oneshot::channel();
1549        let output = output_from_flight_message_stream(
1550            "test-peer".to_string(),
1551            DropProbe {
1552                first: Some(Ok(FlightMessage::AffectedRows {
1553                    rows: 3,
1554                    metrics: None,
1555                })),
1556                polled: Some(polled_tx),
1557                dropped: Some(dropped_tx),
1558            },
1559        )
1560        .await
1561        .unwrap();
1562        polled_rx.await.unwrap();
1563        drop(output);
1564        // Require cancellation before the five-second compatibility timeout.
1565        tokio::time::timeout(Duration::from_secs(1), dropped_rx)
1566            .await
1567            .expect("the compatibility reader must be dropped once the output is dropped")
1568            .unwrap();
1569    }
1570
1571    #[tokio::test]
1572    async fn test_schema_record_batch_yields_before_pending_message_stream() {
1573        let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1574            "v",
1575            ConcreteDataType::int32_datatype(),
1576            false,
1577        )]));
1578        let batch = RecordBatch::new(
1579            schema.clone(),
1580            vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
1581        )
1582        .unwrap();
1583        let messages = futures_util::stream::iter(vec![
1584            Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
1585            Ok(FlightMessage::RecordBatch(batch.into_df_record_batch())),
1586        ] as Vec<Result<FlightMessage>>)
1587        .chain(futures_util::stream::pending());
1588        let output = output_from_flight_message_stream("test-peer".to_string(), messages)
1589            .await
1590            .unwrap();
1591        let OutputData::Stream(mut stream) = output.output.data else {
1592            panic!("expected stream output");
1593        };
1594
1595        let batch = tokio::time::timeout(Duration::from_secs(1), stream.next())
1596            .await
1597            .unwrap()
1598            .unwrap()
1599            .unwrap();
1600        assert_eq!(batch.num_rows(), 1);
1601        assert!(!output.metrics.is_ready());
1602    }
1603
1604    #[tokio::test]
1605    async fn test_invalid_terminal_metrics_after_record_batch_yields_batch_then_error() {
1606        let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1607            "v",
1608            ConcreteDataType::int32_datatype(),
1609            false,
1610        )]));
1611        let batch = RecordBatch::new(
1612            schema.clone(),
1613            vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
1614        )
1615        .unwrap();
1616        let output = output_from_flight_message_stream(
1617            "test-peer".to_string(),
1618            futures_util::stream::iter(vec![
1619                Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
1620                Ok(FlightMessage::RecordBatch(batch.into_df_record_batch())),
1621                Ok(FlightMessage::Metrics("{not-json}".to_string())),
1622            ] as Vec<Result<FlightMessage>>),
1623        )
1624        .await
1625        .unwrap();
1626        let terminal_metrics = output.metrics.clone();
1627        let OutputData::Stream(mut record_batch_stream) = output.output.data else {
1628            panic!("expected stream output");
1629        };
1630
1631        let batch = record_batch_stream.next().await.unwrap().unwrap();
1632        assert_eq!(batch.num_rows(), 1);
1633
1634        let err = record_batch_stream.next().await.unwrap().unwrap_err();
1635        assert_eq!("External error", err.to_string());
1636        assert!(
1637            format!("{err:?}").contains("Invalid terminal metrics message"),
1638            "unexpected error: {err:?}"
1639        );
1640        assert!(record_batch_stream.next().await.is_none());
1641        assert!(terminal_metrics.is_ready());
1642        assert!(terminal_metrics.get().is_none());
1643    }
1644
1645    #[tokio::test]
1646    async fn test_record_batch_stream_continues_after_partial_metrics() {
1647        let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1648            "v",
1649            ConcreteDataType::int32_datatype(),
1650            false,
1651        )]));
1652        let first_batch = RecordBatch::new(
1653            schema.clone(),
1654            vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
1655        )
1656        .unwrap();
1657        let second_batch = RecordBatch::new(
1658            schema.clone(),
1659            vec![Arc::new(Int32Vector::from_slice([2])) as VectorRef],
1660        )
1661        .unwrap();
1662        let output = output_from_flight_message_stream(
1663            "test-peer".to_string(),
1664            futures_util::stream::iter(vec![
1665                Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
1666                Ok(FlightMessage::RecordBatch(
1667                    first_batch.into_df_record_batch(),
1668                )),
1669                Ok(FlightMessage::Metrics(terminal_metrics_json_with_seq(1))),
1670                Ok(FlightMessage::RecordBatch(
1671                    second_batch.into_df_record_batch(),
1672                )),
1673                Ok(FlightMessage::Metrics(terminal_metrics_json_with_seq(2))),
1674            ] as Vec<Result<FlightMessage>>),
1675        )
1676        .await
1677        .unwrap();
1678        let terminal_metrics = output.metrics.clone();
1679        let OutputData::Stream(mut record_batch_stream) = output.output.data else {
1680            panic!("expected stream output");
1681        };
1682
1683        let first_batch = record_batch_stream.next().await.unwrap().unwrap();
1684        assert_eq!(first_batch.num_rows(), 1);
1685        let second_batch = record_batch_stream.next().await.unwrap().unwrap();
1686        assert_eq!(second_batch.num_rows(), 1);
1687        assert!(record_batch_stream.next().await.is_none());
1688
1689        assert!(terminal_metrics.is_ready());
1690        assert_eq!(
1691            terminal_metrics.region_watermark_map(),
1692            Some(std::collections::HashMap::from([(7, 2)]))
1693        );
1694    }
1695
1696    #[test]
1697    fn test_output_metrics_distinguishes_empty_region_watermarks_from_absence() {
1698        let metrics = OutputMetrics::default();
1699        metrics.update(Some(RecordBatchMetrics::default()));
1700
1701        assert_eq!(
1702            metrics.participating_regions(),
1703            Some(std::collections::BTreeSet::new())
1704        );
1705        assert_eq!(
1706            metrics.region_watermark_map(),
1707            Some(std::collections::HashMap::new())
1708        );
1709    }
1710}