mirror of
https://github.com/GreptimeTeam/greptimedb.git
synced 2026-10-03 02:25:35 +00:00
fix(client): yield Flight batches and affected rows without waiting for next message (#8918)
* fix(client): avoid Flight metrics lookahead stalls Signed-off-by: discord9 <55937128+discord9@users.noreply.github.com> * refactor(client): extract trailing Flight metrics task Signed-off-by: discord9 <55937128+discord9@users.noreply.github.com> * test(client): synchronize trailing metrics and bound cancellation Signed-off-by: discord9 <55937128+discord9@users.noreply.github.com> --------- Signed-off-by: discord9 <55937128+discord9@users.noreply.github.com>
This commit is contained in:
+303
-45
@@ -16,7 +16,7 @@ use std::collections::HashMap;
|
||||
use std::pin::Pin;
|
||||
use std::str::FromStr;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::{Arc, RwLock};
|
||||
use std::sync::{Arc, Mutex, RwLock};
|
||||
use std::task::{Context, Poll};
|
||||
use std::time::Duration;
|
||||
|
||||
@@ -55,6 +55,7 @@ use futures::future;
|
||||
use futures_util::{Stream, StreamExt, TryStreamExt};
|
||||
use prost::Message;
|
||||
use snafu::{IntoError, ResultExt};
|
||||
use tokio::sync::Notify;
|
||||
use tonic::metadata::{AsciiMetadataKey, AsciiMetadataValue, MetadataMap, MetadataValue};
|
||||
use tonic::transport::Channel;
|
||||
|
||||
@@ -70,20 +71,47 @@ type FlightDataStream = Pin<Box<dyn Stream<Item = FlightData> + Send>>;
|
||||
type DoPutResponseStream = Pin<Box<dyn Stream<Item = Result<DoPutResponse>>>>;
|
||||
|
||||
const HINTS_METADATA_KEY: &str = "x-greptime-hints";
|
||||
/// Maximum time to wait for the optional trailing metrics message after
|
||||
/// affected rows have already been delivered.
|
||||
const FLIGHT_TRAILING_METRICS_TIMEOUT: Duration = Duration::from_secs(5);
|
||||
|
||||
/// Terminal metrics associated with a query output.
|
||||
///
|
||||
/// For streaming outputs, metrics are only final after the stream is fully
|
||||
/// drained and [`Self::is_ready`] returns `true`.
|
||||
/// drained and [`Self::is_ready`] returns `true`. Affected-row outputs may
|
||||
/// briefly await a compatibility trailing metrics message.
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct OutputMetrics {
|
||||
inner: Arc<OutputMetricsInner>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
#[derive(Debug)]
|
||||
struct OutputMetricsInner {
|
||||
metrics: RwLock<Option<RecordBatchMetrics>>,
|
||||
completion_error: RwLock<Option<String>>,
|
||||
ready: AtomicBool,
|
||||
ready_notify: Notify,
|
||||
compatibility_task: Mutex<Option<tokio::task::AbortHandle>>,
|
||||
}
|
||||
|
||||
impl Default for OutputMetricsInner {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
metrics: RwLock::new(None),
|
||||
completion_error: RwLock::new(None),
|
||||
ready: AtomicBool::new(false),
|
||||
ready_notify: Notify::new(),
|
||||
compatibility_task: Mutex::new(None),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for OutputMetricsInner {
|
||||
fn drop(&mut self) {
|
||||
if let Some(handle) = self.compatibility_task.get_mut().unwrap().take() {
|
||||
handle.abort();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl OutputMetrics {
|
||||
@@ -98,10 +126,45 @@ impl OutputMetrics {
|
||||
|
||||
/// Marks the terminal metrics as final for this output.
|
||||
pub fn mark_ready(&self) {
|
||||
let _ = self
|
||||
if self
|
||||
.inner
|
||||
.ready
|
||||
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire);
|
||||
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
|
||||
.is_ok()
|
||||
{
|
||||
self.inner.ready_notify.notify_waiters();
|
||||
}
|
||||
}
|
||||
|
||||
/// Waits until terminal metrics are final.
|
||||
pub async fn wait_ready(&self) {
|
||||
loop {
|
||||
let notified = self.inner.ready_notify.notified();
|
||||
if self.is_ready() {
|
||||
return;
|
||||
}
|
||||
notified.await;
|
||||
}
|
||||
}
|
||||
|
||||
/// Returns an error encountered while completing the output, if any.
|
||||
pub fn completion_error(&self) -> Option<String> {
|
||||
self.inner.completion_error.read().unwrap().clone()
|
||||
}
|
||||
|
||||
fn set_completion_error(&self, error: impl Into<String>) {
|
||||
*self.inner.completion_error.write().unwrap() = Some(error.into());
|
||||
}
|
||||
|
||||
fn set_compatibility_task(&self, handle: tokio::task::AbortHandle) {
|
||||
let mut task = self.inner.compatibility_task.lock().unwrap();
|
||||
if !self.is_ready() {
|
||||
*task = Some(handle);
|
||||
}
|
||||
}
|
||||
|
||||
fn take_compatibility_task(&self) -> Option<tokio::task::AbortHandle> {
|
||||
self.inner.compatibility_task.lock().unwrap().take()
|
||||
}
|
||||
|
||||
/// Returns whether terminal metrics are final.
|
||||
@@ -148,7 +211,8 @@ impl OutputMetrics {
|
||||
///
|
||||
/// The contained [`OutputMetrics`] lets callers read stream terminal metrics
|
||||
/// after consuming `output`. For non-stream outputs, metrics are ready
|
||||
/// immediately.
|
||||
/// immediately. Flight affected-row outputs without inline metrics require
|
||||
/// [`OutputMetrics::wait_ready`] before compatibility trailing metrics can be read.
|
||||
#[derive(Debug)]
|
||||
pub struct OutputWithMetrics {
|
||||
pub output: Output,
|
||||
@@ -194,6 +258,55 @@ fn parse_terminal_metrics(metrics_json: &str) -> Result<RecordBatchMetrics> {
|
||||
})
|
||||
}
|
||||
|
||||
fn spawn_affected_rows_trailing_metrics_task<S>(
|
||||
terminal_metrics: &OutputMetrics,
|
||||
mut reader: FlightMessageReader<S>,
|
||||
) where
|
||||
S: Stream<Item = Result<FlightMessage>> + Send + Unpin + 'static,
|
||||
{
|
||||
let metrics_ref = Arc::downgrade(&terminal_metrics.inner);
|
||||
let remote_addr = reader.remote_addr().to_string();
|
||||
let task = common_runtime::spawn_global(async move {
|
||||
let result =
|
||||
tokio::time::timeout(FLIGHT_TRAILING_METRICS_TIMEOUT, reader.read_next()).await;
|
||||
let Some(inner) = metrics_ref.upgrade() else {
|
||||
return;
|
||||
};
|
||||
let metrics = OutputMetrics { inner };
|
||||
match result {
|
||||
Ok(Ok(Some(FlightMessage::Metrics(s)))) => match parse_terminal_metrics(&s) {
|
||||
Ok(metrics_json) => metrics.update(Some(metrics_json)),
|
||||
Err(error) => {
|
||||
metrics.set_completion_error(error.to_string());
|
||||
warn!(
|
||||
"Failed to decode trailing Flight metrics from {}: {}",
|
||||
remote_addr, error
|
||||
);
|
||||
}
|
||||
},
|
||||
Ok(Ok(None)) => {}
|
||||
Ok(Ok(Some(other))) => {
|
||||
let error = format!("Unexpected trailing Flight message: {other:?}");
|
||||
metrics.set_completion_error(error.clone());
|
||||
warn!("{} from {}", error, remote_addr);
|
||||
}
|
||||
Ok(Err(error)) => {
|
||||
let error = flight_stream_error(&remote_addr, error);
|
||||
metrics.set_completion_error(error.to_string());
|
||||
warn!("{}", error);
|
||||
}
|
||||
Err(_) => {
|
||||
let error = "Timed out waiting for trailing Flight metrics";
|
||||
metrics.set_completion_error(error);
|
||||
warn!("{} from {}", error, remote_addr);
|
||||
}
|
||||
}
|
||||
metrics.mark_ready();
|
||||
metrics.take_compatibility_task();
|
||||
});
|
||||
terminal_metrics.set_compatibility_task(task.abort_handle());
|
||||
}
|
||||
|
||||
struct StreamWithMetrics {
|
||||
stream: common_recordbatch::SendableRecordBatchStream,
|
||||
metrics: OutputMetrics,
|
||||
@@ -233,6 +346,9 @@ impl Stream for StreamWithMetrics {
|
||||
|
||||
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
||||
let polled = Pin::new(&mut self.stream).poll_next(cx);
|
||||
if let Poll::Ready(Some(Err(error))) = &polled {
|
||||
self.metrics.set_completion_error(error.to_string());
|
||||
}
|
||||
if let Poll::Ready(None) = &polled {
|
||||
self.sync_terminal_metrics();
|
||||
self.metrics.mark_ready();
|
||||
@@ -280,32 +396,12 @@ where
|
||||
FlightMessage::AffectedRows { rows, metrics } => {
|
||||
let terminal_metrics = OutputMetrics::new();
|
||||
if let Some(metrics) = metrics {
|
||||
// Inline metrics are authoritative. They complete the output
|
||||
// without touching the rest of the Flight stream.
|
||||
terminal_metrics.update(Some(parse_terminal_metrics(&metrics)?));
|
||||
}
|
||||
let next_message = reader
|
||||
.read_next()
|
||||
.await
|
||||
.map_err(|error| flight_stream_error(reader.remote_addr(), error))?;
|
||||
match next_message {
|
||||
None => terminal_metrics.mark_ready(),
|
||||
Some(FlightMessage::Metrics(s)) if terminal_metrics.get().is_none() => {
|
||||
terminal_metrics.update(Some(parse_terminal_metrics(&s)?));
|
||||
terminal_metrics.mark_ready();
|
||||
}
|
||||
Some(FlightMessage::Metrics(_)) => {
|
||||
return IllegalFlightMessagesSnafu {
|
||||
reason: "'AffectedRows' Flight metadata already carries Metrics and cannot be followed by another Metrics message",
|
||||
}
|
||||
.fail();
|
||||
}
|
||||
Some(other) => {
|
||||
return IllegalFlightMessagesSnafu {
|
||||
reason: format!(
|
||||
"'AffectedRows' Flight message can only be followed by a Metrics message, got {other:?}"
|
||||
),
|
||||
}
|
||||
.fail();
|
||||
}
|
||||
terminal_metrics.mark_ready();
|
||||
} else {
|
||||
spawn_affected_rows_trailing_metrics_task(&terminal_metrics, reader);
|
||||
}
|
||||
Ok(OutputWithMetrics {
|
||||
output: Output::new_with_affected_rows(rows),
|
||||
@@ -749,7 +845,9 @@ impl Database {
|
||||
/// Executes a SQL query and returns the output with terminal metrics.
|
||||
///
|
||||
/// For stream outputs, callers must consume the stream before reading final
|
||||
/// terminal metrics from [`OutputWithMetrics::metrics`].
|
||||
/// terminal metrics from [`OutputWithMetrics::metrics`]. For affected-row
|
||||
/// outputs without inline metrics, call [`OutputMetrics::wait_ready`] when
|
||||
/// compatibility trailing metrics are required.
|
||||
pub async fn sql_with_terminal_metrics<S>(
|
||||
&self,
|
||||
sql: S,
|
||||
@@ -986,6 +1084,7 @@ mod tests {
|
||||
use datatypes::schema::{ColumnSchema, Schema};
|
||||
use datatypes::vectors::Int32Vector;
|
||||
use futures_util::StreamExt;
|
||||
use tokio::sync::oneshot;
|
||||
use tonic::codegen::http::{HeaderMap, HeaderValue};
|
||||
use tonic::metadata::MetadataMap;
|
||||
use tonic::{Code, Status};
|
||||
@@ -1296,11 +1395,6 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_terminal_metrics_rejects_invalid_json() {
|
||||
assert!(parse_terminal_metrics("{not-json}").is_err());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_affected_rows_inline_metrics_are_parsed() {
|
||||
let output = output_from_flight_message_stream(
|
||||
@@ -1322,27 +1416,191 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_affected_rows_inline_metrics_rejects_trailing_metrics() {
|
||||
async fn test_affected_rows_inline_metrics_do_not_poll_trailer() {
|
||||
let metrics_json = terminal_metrics_json();
|
||||
let err = output_from_flight_message_stream(
|
||||
let output = output_from_flight_message_stream(
|
||||
"test-peer".to_string(),
|
||||
futures_util::stream::iter(vec![
|
||||
Ok(FlightMessage::AffectedRows {
|
||||
rows: 3,
|
||||
metrics: Some(metrics_json.clone()),
|
||||
metrics: Some(metrics_json),
|
||||
}),
|
||||
Ok(FlightMessage::Metrics(metrics_json)),
|
||||
Ok(FlightMessage::Metrics(terminal_metrics_json_with_seq(99))),
|
||||
] as Vec<Result<FlightMessage>>),
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
.unwrap();
|
||||
|
||||
assert!(
|
||||
err.to_string().contains("already carries Metrics"),
|
||||
"unexpected error: {err:?}"
|
||||
assert!(output.metrics.is_ready());
|
||||
assert_eq!(
|
||||
output.metrics.region_watermark_map(),
|
||||
Some(std::collections::HashMap::from([(7, 42)]))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_affected_rows_without_inline_metrics_becomes_ready_after_trailer() {
|
||||
// Hold the trailer until the not-ready state has been observed.
|
||||
let (trailer_tx, trailer_rx) = oneshot::channel();
|
||||
let trailer = futures_util::stream::once(trailer_rx).map(|ready| {
|
||||
ready.unwrap();
|
||||
Ok(FlightMessage::Metrics(terminal_metrics_json()))
|
||||
});
|
||||
let output = output_from_flight_message_stream(
|
||||
"test-peer".to_string(),
|
||||
futures_util::stream::iter(vec![Ok(FlightMessage::AffectedRows {
|
||||
rows: 3,
|
||||
metrics: None,
|
||||
})])
|
||||
.chain(trailer),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(!output.metrics.is_ready());
|
||||
trailer_tx.send(()).unwrap();
|
||||
tokio::time::timeout(Duration::from_secs(1), output.metrics.wait_ready())
|
||||
.await
|
||||
.expect("terminal metrics must become ready once the trailer arrives");
|
||||
assert!(output.metrics.completion_error().is_none());
|
||||
assert_eq!(
|
||||
output.metrics.region_watermark_map(),
|
||||
Some(std::collections::HashMap::from([(7, 42)]))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_affected_rows_malformed_trailer_sets_completion_error_and_ready() {
|
||||
let output = output_from_flight_message_stream(
|
||||
"test-peer".to_string(),
|
||||
futures_util::stream::iter(vec![
|
||||
Ok(FlightMessage::AffectedRows {
|
||||
rows: 3,
|
||||
metrics: None,
|
||||
}),
|
||||
Ok(FlightMessage::Metrics("{not-json}".to_string())),
|
||||
] as Vec<Result<FlightMessage>>),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
output.metrics.wait_ready().await;
|
||||
let error = output.metrics.completion_error().unwrap();
|
||||
assert!(error.contains("Invalid terminal metrics message"));
|
||||
assert!(output.metrics.is_ready());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_affected_rows_transport_trailer_error_sets_completion_error_and_ready() {
|
||||
let output = output_from_flight_message_stream(
|
||||
"test-peer".to_string(),
|
||||
futures_util::stream::iter(vec![
|
||||
Ok(FlightMessage::AffectedRows {
|
||||
rows: 3,
|
||||
metrics: None,
|
||||
}),
|
||||
Err(Status::unavailable("trailer read failed").into()),
|
||||
] as Vec<Result<FlightMessage>>),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
output.metrics.wait_ready().await;
|
||||
let error = output.metrics.completion_error().unwrap();
|
||||
assert!(error.contains("trailer read failed"));
|
||||
assert!(output.metrics.is_ready());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_affected_rows_compatibility_reader_is_cancelled_after_second_poll_begins() {
|
||||
struct DropProbe {
|
||||
first: Option<Result<FlightMessage>>,
|
||||
polled: Option<oneshot::Sender<()>>,
|
||||
dropped: Option<oneshot::Sender<()>>,
|
||||
}
|
||||
|
||||
impl Stream for DropProbe {
|
||||
type Item = Result<FlightMessage>;
|
||||
|
||||
fn poll_next(
|
||||
mut self: Pin<&mut Self>,
|
||||
_cx: &mut Context<'_>,
|
||||
) -> Poll<Option<Self::Item>> {
|
||||
if let Some(message) = self.first.take() {
|
||||
return Poll::Ready(Some(message));
|
||||
}
|
||||
if let Some(polled) = self.polled.take() {
|
||||
let _ = polled.send(());
|
||||
}
|
||||
Poll::Pending
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for DropProbe {
|
||||
fn drop(&mut self) {
|
||||
if let Some(dropped) = self.dropped.take() {
|
||||
let _ = dropped.send(());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let (polled_tx, polled_rx) = oneshot::channel();
|
||||
let (dropped_tx, dropped_rx) = oneshot::channel();
|
||||
let output = output_from_flight_message_stream(
|
||||
"test-peer".to_string(),
|
||||
DropProbe {
|
||||
first: Some(Ok(FlightMessage::AffectedRows {
|
||||
rows: 3,
|
||||
metrics: None,
|
||||
})),
|
||||
polled: Some(polled_tx),
|
||||
dropped: Some(dropped_tx),
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
polled_rx.await.unwrap();
|
||||
drop(output);
|
||||
// Require cancellation before the five-second compatibility timeout.
|
||||
tokio::time::timeout(Duration::from_secs(1), dropped_rx)
|
||||
.await
|
||||
.expect("the compatibility reader must be dropped once the output is dropped")
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_schema_record_batch_yields_before_pending_message_stream() {
|
||||
let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
|
||||
"v",
|
||||
ConcreteDataType::int32_datatype(),
|
||||
false,
|
||||
)]));
|
||||
let batch = RecordBatch::new(
|
||||
schema.clone(),
|
||||
vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
|
||||
)
|
||||
.unwrap();
|
||||
let messages = futures_util::stream::iter(vec![
|
||||
Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
|
||||
Ok(FlightMessage::RecordBatch(batch.into_df_record_batch())),
|
||||
] as Vec<Result<FlightMessage>>)
|
||||
.chain(futures_util::stream::pending());
|
||||
let output = output_from_flight_message_stream("test-peer".to_string(), messages)
|
||||
.await
|
||||
.unwrap();
|
||||
let OutputData::Stream(mut stream) = output.output.data else {
|
||||
panic!("expected stream output");
|
||||
};
|
||||
|
||||
let batch = tokio::time::timeout(Duration::from_secs(1), stream.next())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(batch.num_rows(), 1);
|
||||
assert!(!output.metrics.is_ready());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_invalid_terminal_metrics_after_record_batch_yields_batch_then_error() {
|
||||
let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
|
||||
|
||||
@@ -12,40 +12,18 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
use std::pin::Pin;
|
||||
|
||||
use arrow_flight::FlightData;
|
||||
use common_grpc::flight::{FlightDecoder, FlightMessage};
|
||||
use futures_util::stream::Peekable;
|
||||
use futures_util::{Stream, StreamExt};
|
||||
use snafu::{OptionExt, ResultExt};
|
||||
|
||||
use crate::Result;
|
||||
use crate::error::{ConvertFlightDataSnafu, Error, IllegalFlightMessagesSnafu};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(crate) enum FlightMessageKind {
|
||||
Schema,
|
||||
RecordBatch,
|
||||
AffectedRows,
|
||||
Metrics,
|
||||
}
|
||||
|
||||
impl From<&FlightMessage> for FlightMessageKind {
|
||||
fn from(message: &FlightMessage) -> Self {
|
||||
match message {
|
||||
FlightMessage::Schema(_) => Self::Schema,
|
||||
FlightMessage::RecordBatch(_) => Self::RecordBatch,
|
||||
FlightMessage::AffectedRows { .. } => Self::AffectedRows,
|
||||
FlightMessage::Metrics(_) => Self::Metrics,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) struct FlightMessageReader<S: Stream + Unpin> {
|
||||
/// Remote Flight peer associated with this response stream.
|
||||
remote_addr: String,
|
||||
messages: Peekable<S>,
|
||||
messages: S,
|
||||
}
|
||||
|
||||
impl<S> FlightMessageReader<S>
|
||||
@@ -55,7 +33,7 @@ where
|
||||
pub(crate) fn new(remote_addr: impl Into<String>, messages: S) -> Self {
|
||||
Self {
|
||||
remote_addr: remote_addr.into(),
|
||||
messages: messages.peekable(),
|
||||
messages,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -72,21 +50,6 @@ where
|
||||
pub(crate) async fn read_next(&mut self) -> Result<Option<FlightMessage>> {
|
||||
self.messages.next().await.transpose()
|
||||
}
|
||||
|
||||
pub(crate) async fn peek_next_message_kind(&mut self) -> Result<Option<FlightMessageKind>> {
|
||||
match Pin::new(&mut self.messages).peek().await {
|
||||
Some(Ok(message)) => Ok(Some(message.into())),
|
||||
None => Ok(None),
|
||||
Some(Err(_)) => match self.read_next().await {
|
||||
// `peek` only borrows the error; consume it to preserve the source error.
|
||||
Err(error) => Err(error),
|
||||
Ok(_) => IllegalFlightMessagesSnafu {
|
||||
reason: "Flight stream changed after peek".to_string(),
|
||||
}
|
||||
.fail(),
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn decode_flight_data(
|
||||
|
||||
+51
-62
@@ -48,7 +48,7 @@ use crate::error::{
|
||||
self, FlightGetSnafu, IllegalDatabaseResponseSnafu, IllegalFlightMessagesSnafu,
|
||||
MissingFieldSnafu, Result, ServerSnafu,
|
||||
};
|
||||
use crate::flight::{FlightMessageKind, FlightMessageReader, decode_flight_data};
|
||||
use crate::flight::{FlightMessageReader, decode_flight_data};
|
||||
use crate::{Client, metrics};
|
||||
|
||||
const FLIGHT_DO_GET_TIMEOUT: Duration = Duration::from_secs(10);
|
||||
@@ -247,9 +247,7 @@ where
|
||||
"poll_flight_data_stream"
|
||||
));
|
||||
|
||||
let mut stream_ended = false;
|
||||
|
||||
while !stream_ended {
|
||||
loop {
|
||||
let flight_message = match reader.read_next().await {
|
||||
Ok(Some(message)) => message,
|
||||
Ok(None) => break,
|
||||
@@ -262,59 +260,27 @@ where
|
||||
|
||||
match flight_message {
|
||||
FlightMessage::RecordBatch(record_batch) => {
|
||||
let result_to_yield =
|
||||
RecordBatch::from_df_record_batch(schema_cloned.clone(), record_batch);
|
||||
|
||||
// Metrics follow a batch so MergeScan can observe them before yielding it.
|
||||
match reader.peek_next_message_kind().await {
|
||||
Ok(Some(FlightMessageKind::Metrics)) => {
|
||||
let metrics_message = match reader.read_next().await {
|
||||
Ok(Some(FlightMessage::Metrics(metrics))) => metrics,
|
||||
Ok(Some(_) | None) => {
|
||||
yield IllegalFlightMessagesSnafu {
|
||||
reason: "Flight stream changed after peek",
|
||||
}
|
||||
.fail()
|
||||
.map_err(BoxedError::new)
|
||||
.context(ExternalSnafu);
|
||||
break;
|
||||
}
|
||||
Err(error) => {
|
||||
yield Err(BoxedError::new(flight_stream_error(
|
||||
&stream_addr,
|
||||
error,
|
||||
)))
|
||||
.context(ExternalSnafu);
|
||||
break;
|
||||
}
|
||||
};
|
||||
let metrics = serde_json::from_str(&metrics_message).ok().map(Arc::new);
|
||||
metrics_ref.swap(metrics);
|
||||
}
|
||||
Ok(Some(FlightMessageKind::RecordBatch)) => {}
|
||||
Ok(Some(FlightMessageKind::Schema | FlightMessageKind::AffectedRows)) => {
|
||||
yield IllegalFlightMessagesSnafu {
|
||||
reason: "A RecordBatch message can only be succeeded by a Metrics message or another RecordBatch message"
|
||||
}
|
||||
.fail()
|
||||
.map_err(BoxedError::new)
|
||||
.context(ExternalSnafu);
|
||||
break;
|
||||
}
|
||||
Ok(None) => stream_ended = true,
|
||||
Err(error) => {
|
||||
yield Err(BoxedError::new(flight_stream_error(&stream_addr, error)))
|
||||
.context(ExternalSnafu);
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
yield Ok(result_to_yield);
|
||||
// Deliver each batch immediately. In particular, do not
|
||||
// wait for a possible following Metrics message; it is
|
||||
// consumed on the next poll of this stream.
|
||||
yield Ok(RecordBatch::from_df_record_batch(
|
||||
schema_cloned.clone(),
|
||||
record_batch,
|
||||
));
|
||||
}
|
||||
FlightMessage::Metrics(s) => {
|
||||
// Metrics may arrive before the next RecordBatch.
|
||||
let m = serde_json::from_str(&s).ok().map(Arc::new);
|
||||
metrics_ref.swap(m);
|
||||
match serde_json::from_str(&s) {
|
||||
Ok(metrics) => {
|
||||
metrics_ref.swap(Some(Arc::new(metrics)));
|
||||
}
|
||||
Err(error) => {
|
||||
common_telemetry::warn!(
|
||||
"Failed to decode region Flight metrics: {}",
|
||||
error
|
||||
);
|
||||
}
|
||||
}
|
||||
continue;
|
||||
}
|
||||
_ => {
|
||||
@@ -719,7 +685,34 @@ mod test {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_record_batch_stream_updates_following_metrics_before_yielding_batch() {
|
||||
async fn test_record_batch_is_yielded_without_waiting_for_next_message() {
|
||||
let schema = test_schema();
|
||||
let batch = RecordBatch::new(
|
||||
schema.clone(),
|
||||
vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let messages = stream::iter(vec![
|
||||
Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
|
||||
Ok(FlightMessage::RecordBatch(batch.into_df_record_batch())),
|
||||
])
|
||||
.chain(stream::pending::<Result<FlightMessage>>());
|
||||
let mut recordbatches =
|
||||
recordbatches_from_flight_message_stream("test-peer".to_string(), messages)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let batch = tokio::time::timeout(Duration::from_secs(1), recordbatches.next())
|
||||
.await
|
||||
.expect("the first batch must not wait for lookahead")
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(batch.num_rows(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_malformed_region_metrics_are_non_fatal() {
|
||||
let schema = test_schema();
|
||||
let batch = RecordBatch::new(
|
||||
schema.clone(),
|
||||
@@ -730,23 +723,19 @@ mod test {
|
||||
"test-peer".to_string(),
|
||||
stream::iter(vec![
|
||||
Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
|
||||
Ok(FlightMessage::Metrics("{not-json}".to_string())),
|
||||
Ok(FlightMessage::RecordBatch(batch.into_df_record_batch())),
|
||||
Ok(FlightMessage::Metrics(test_metrics_json())),
|
||||
]),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let batch = recordbatches.next().await.unwrap().unwrap();
|
||||
assert_eq!(batch.num_rows(), 1);
|
||||
|
||||
let metrics = recordbatches.metrics().unwrap();
|
||||
assert_eq!(metrics.elapsed_compute, 7);
|
||||
assert_eq!(recordbatches.next().await.unwrap().unwrap().num_rows(), 1);
|
||||
assert!(recordbatches.next().await.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_record_batch_stream_preserves_peeked_record_batch() {
|
||||
async fn test_record_batch_stream_preserves_following_record_batch() {
|
||||
let schema = test_schema();
|
||||
let first_batch = RecordBatch::new(
|
||||
schema.clone(),
|
||||
|
||||
@@ -1400,6 +1400,12 @@ impl BatchingTask {
|
||||
match res {
|
||||
Ok(res) => {
|
||||
let (affected_rows, _) = res.output.extract_rows_and_cost();
|
||||
if matches!(&res.output.data, common_query::OutputData::AffectedRows(_)) {
|
||||
res.metrics.wait_ready().await;
|
||||
}
|
||||
if let Some(error) = res.metrics.completion_error() {
|
||||
warn!("Flow {flow_id} completed with terminal metrics error: {error}");
|
||||
}
|
||||
debug!(
|
||||
"Flow {flow_id} executed, affected_rows: {affected_rows:?}, elapsed: {:?}, watermark: {:?}",
|
||||
elapsed,
|
||||
|
||||
@@ -990,7 +990,7 @@ mod test {
|
||||
panic!("expected affected rows output");
|
||||
};
|
||||
assert_eq!(affected_rows, 9);
|
||||
assert!(result.metrics.is_ready());
|
||||
result.metrics.wait_ready().await;
|
||||
assert!(result.region_watermark_map().is_none());
|
||||
|
||||
let err = client
|
||||
@@ -1016,6 +1016,7 @@ mod test {
|
||||
panic!("expected affected rows output");
|
||||
};
|
||||
assert_eq!(affected_rows, 9);
|
||||
result.metrics.wait_ready().await;
|
||||
assert_eq!(
|
||||
result.region_watermark_map(),
|
||||
Some(std::collections::HashMap::from([previous_watermark]))
|
||||
|
||||
Reference in New Issue
Block a user