1use 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";
74const FLIGHT_TRAILING_METRICS_TIMEOUT: Duration = Duration::from_secs(5);
77
78#[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 pub fn update(&self, metrics: Option<RecordBatchMetrics>) {
124 *self.inner.metrics.write().unwrap() = metrics;
125 }
126
127 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 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 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 pub fn is_ready(&self) -> bool {
174 self.inner.ready.load(Ordering::Acquire)
175 }
176
177 pub fn get(&self) -> Option<RecordBatchMetrics> {
179 self.inner.metrics.read().unwrap().clone()
180 }
181
182 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 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#[derive(Debug)]
217pub struct OutputWithMetrics {
218 pub output: Output,
219 pub metrics: OutputMetrics,
220}
221
222impl OutputWithMetrics {
223 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 pub fn region_watermark_map(&self) -> Option<std::collections::HashMap<u64, u64>> {
238 self.metrics.region_watermark_map()
239 }
240
241 pub fn participating_regions(&self) -> Option<std::collections::BTreeSet<u64>> {
243 self.metrics.participating_regions()
244 }
245
246 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 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 catalog: String,
507 schema: String,
508 dbname: String,
511 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
546pub 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 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 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 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 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 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 pub fn set_timezone(&mut self, timezone: impl Into<String>) {
639 self.timezone = timezone.into();
640 }
641
642 pub fn set_auth(&mut self, auth: AuthScheme) {
644 self.ctx.auth_header = Some(AuthHeader {
645 auth_scheme: Some(auth),
646 });
647 }
648
649 pub fn flight_request(&self) -> DatabaseFlightRequest<'_> {
651 DatabaseFlightRequest {
652 database: self,
653 options: FlightRequestOptions::default(),
654 }
655 }
656
657 pub async fn insert(&self, requests: InsertRequests) -> Result<u32> {
659 self.handle(Request::Inserts(requests)).await
660 }
661
662 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 pub async fn row_inserts(&self, requests: RowInsertRequests) -> Result<u32> {
686 self.handle(Request::RowInserts(requests)).await
687 }
688
689 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 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 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 if is_grpc_retryable(&err) {
790 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 tracing_context: W3cTrace::new(),
824 }),
825 request: Some(request),
826 }
827 }
828
829 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 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 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 pub async fn logical_plan(&self, logical_plan: Vec<u8>) -> Result<Output> {
867 self.flight_request().logical_plan(logical_plan).await
868 }
869
870 #[cfg(feature = "testing")]
872 pub async fn create(&self, expr: CreateTableExpr) -> Result<Output> {
873 self.flight_request().create(expr).await
874 }
875
876 #[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 pub async fn do_put(&self, stream: FlightDataStream) -> Result<DoPutResponseStream> {
926 self.do_put_with_hints(stream, &[]).await
927 }
928
929 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 pub fn with_hints(mut self, hints: &[(&str, &str)]) -> Self {
972 self.options.hints = Database::encode_hints(hints);
973 self
974 }
975
976 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 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 pub fn with_timeout(mut self, timeout: Duration) -> Self {
990 self.options.timeout = Some(timeout);
991 self
992 }
993
994 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 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 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 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 #[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 #[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
1061pub 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 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 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}