mirror of
https://github.com/GreptimeTeam/greptimedb.git
synced 2026-09-12 16:32:16 +00:00
fix(query): preserve typed errors through distributed Flight streams
Signed-off-by: discord9 <discord9@163.com>
This commit is contained in:
@@ -277,6 +277,8 @@ pub fn datafusion_status_code<T: ErrorExt + 'static>(
|
||||
DataFusionError::External(e) => {
|
||||
if let Some(ext) = (*e).downcast_ref::<T>() {
|
||||
ext.status_code()
|
||||
} else if let Some(ext) = (*e).downcast_ref::<BoxedError>() {
|
||||
ext.status_code()
|
||||
} else {
|
||||
default_status.unwrap_or(StatusCode::EngineExecuteQuery)
|
||||
}
|
||||
@@ -285,3 +287,49 @@ pub fn datafusion_status_code<T: ErrorExt + 'static>(
|
||||
_ => default_status.unwrap_or(StatusCode::EngineExecuteQuery),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use common_error::ext::PlainError;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_datafusion_status_code_external_errors() {
|
||||
let boxed_error = || {
|
||||
DataFusionError::External(Box::new(BoxedError::new(PlainError::new(
|
||||
"neutral error".to_string(),
|
||||
StatusCode::RequestOutdated,
|
||||
))))
|
||||
};
|
||||
assert_eq!(
|
||||
datafusion_status_code::<Error>(&boxed_error(), None),
|
||||
StatusCode::RequestOutdated
|
||||
);
|
||||
assert_eq!(
|
||||
datafusion_status_code::<Error>(&boxed_error(), Some(StatusCode::PlanQuery)),
|
||||
StatusCode::RequestOutdated
|
||||
);
|
||||
|
||||
let direct_error = DataFusionError::External(Box::new(Error::DynFilterPayloadTooLarge {
|
||||
payload_size_bytes: 2,
|
||||
max_payload_bytes: 1,
|
||||
location: Location::default(),
|
||||
}));
|
||||
assert_eq!(
|
||||
datafusion_status_code::<Error>(&direct_error, None),
|
||||
StatusCode::PlanQuery
|
||||
);
|
||||
|
||||
let unknown_error =
|
||||
|| DataFusionError::External(Box::new(std::io::Error::other("neutral error")));
|
||||
assert_eq!(
|
||||
datafusion_status_code::<Error>(&unknown_error(), None),
|
||||
StatusCode::EngineExecuteQuery
|
||||
);
|
||||
assert_eq!(
|
||||
datafusion_status_code::<Error>(&unknown_error(), Some(StatusCode::PlanQuery)),
|
||||
StatusCode::PlanQuery
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -205,6 +205,15 @@ impl ErrorExt for Error {
|
||||
| Error::PhysicalExpr { .. }
|
||||
| Error::RecordBatchSliceIndexOverflow { .. } => StatusCode::Internal,
|
||||
|
||||
Error::PollStream {
|
||||
error: datafusion::error::DataFusionError::External(source),
|
||||
..
|
||||
} => source
|
||||
.downcast_ref::<BoxedError>()
|
||||
.map_or(StatusCode::EngineExecuteQuery, |source| {
|
||||
source.status_code()
|
||||
}),
|
||||
|
||||
Error::PollStream { .. } => StatusCode::EngineExecuteQuery,
|
||||
|
||||
Error::ArrowCompute { .. } => StatusCode::IllegalState,
|
||||
@@ -242,3 +251,46 @@ impl ErrorExt for Error {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use common_error::ext::PlainError;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn poll_stream_status_code_preserves_direct_external_boxed_error() {
|
||||
let cases = [
|
||||
(StatusCode::RequestOutdated, StatusCode::RequestOutdated),
|
||||
(StatusCode::Unknown, StatusCode::Unknown),
|
||||
];
|
||||
|
||||
for (source_status, expected_status) in cases {
|
||||
let error = Error::PollStream {
|
||||
error: datafusion::error::DataFusionError::External(Box::new(BoxedError::new(
|
||||
PlainError::new("neutral error".to_string(), source_status),
|
||||
))),
|
||||
location: Location::default(),
|
||||
};
|
||||
assert_eq!(error.status_code(), expected_status);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn poll_stream_status_code_defaults_for_unrecognized_datafusion_errors() {
|
||||
let errors = [
|
||||
datafusion::error::DataFusionError::External(Box::new(std::io::Error::other(
|
||||
"neutral io error",
|
||||
))),
|
||||
datafusion::error::DataFusionError::Internal("neutral internal error".to_string()),
|
||||
];
|
||||
|
||||
for error in errors {
|
||||
let error = Error::PollStream {
|
||||
error,
|
||||
location: Location::default(),
|
||||
};
|
||||
assert_eq!(error.status_code(), StatusCode::EngineExecuteQuery);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -25,6 +25,7 @@ use arrow_schema::{
|
||||
};
|
||||
use async_stream::stream;
|
||||
use common_catalog::parse_catalog_and_schema_from_db_string;
|
||||
use common_error::ext::BoxedError;
|
||||
use common_plugins::GREPTIME_EXEC_READ_COST;
|
||||
use common_query::request::QueryRequest;
|
||||
use common_recordbatch::adapter::{RecordBatchMetrics, region_scan_output_bytes};
|
||||
@@ -746,7 +747,7 @@ impl MergeScanExec {
|
||||
}
|
||||
let mut stream = do_get_result.map_err(|e| {
|
||||
MERGE_SCAN_ERRORS_TOTAL.inc();
|
||||
DataFusionError::External(Box::new(e))
|
||||
DataFusionError::External(Box::new(BoxedError::new(e)))
|
||||
})?;
|
||||
|
||||
if let Some(subscriber_rollback) = subscriber_rollback.as_mut() {
|
||||
@@ -815,7 +816,8 @@ impl MergeScanExec {
|
||||
let poll_elapsed = poll_timer.elapsed();
|
||||
poll_duration += poll_elapsed;
|
||||
|
||||
let batch = batch.map_err(|e| DataFusionError::External(Box::new(e)))?;
|
||||
let batch = batch
|
||||
.map_err(|e| DataFusionError::External(Box::new(BoxedError::new(e))))?;
|
||||
let df_batch = batch.into_df_record_batch();
|
||||
if !Arc::ptr_eq(&advertised_schema, df_batch.schema_ref()) {
|
||||
validate_remote_schema(
|
||||
@@ -1411,6 +1413,8 @@ mod tests {
|
||||
use arrow_schema::{DataType as TestArrowDataType, Field, TimeUnit};
|
||||
use async_trait::async_trait;
|
||||
use common_base::Plugins;
|
||||
use common_error::ext::{ErrorExt, PlainError};
|
||||
use common_error::status_code::StatusCode;
|
||||
use common_meta::peer::Peer;
|
||||
use common_query::request::{
|
||||
INITIAL_REMOTE_DYN_FILTER_REGISTRATIONS_EXTENSION_KEY, InitialDynFilterRegs,
|
||||
@@ -1435,6 +1439,7 @@ mod tests {
|
||||
use session::ReadPreference;
|
||||
use session::context::QueryContext;
|
||||
use session::query_id::QueryId;
|
||||
use snafu::IntoError;
|
||||
use table::table::scan::REGION_SCAN_EXEC_NAME;
|
||||
use table::table_name::TableName;
|
||||
use tokio::sync::{Notify, oneshot};
|
||||
@@ -1813,7 +1818,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn failed_do_get_rolls_back_new_subscriber_without_starting_fanout() {
|
||||
async fn failed_do_get_preserves_status_code_and_rolls_back_subscriber() {
|
||||
let handler = Arc::new(FailingRegionQueryHandler::default());
|
||||
let query_ctx = QueryContext::arc();
|
||||
let state = Arc::new(QueryEngineState::new(
|
||||
@@ -1869,8 +1874,12 @@ mod tests {
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let mut stream = exec.to_stream(task_ctx, 0).unwrap();
|
||||
assert!(stream.next().await.unwrap().is_err());
|
||||
let mut stream = common_recordbatch::adapter::RecordBatchStreamAdapter::try_new(
|
||||
exec.to_stream(task_ctx, 0).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
let error = stream.next().await.unwrap().unwrap_err();
|
||||
assert_eq!(error.status_code(), StatusCode::RequestOutdated);
|
||||
assert_eq!(handler.do_get_calls.load(Ordering::SeqCst), 1);
|
||||
assert!(handler.saw_subscriber.load(Ordering::SeqCst));
|
||||
|
||||
@@ -1880,6 +1889,30 @@ mod tests {
|
||||
assert!(!entries[0].fanout_started_for_test());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn merge_scan_later_stream_error_preserves_status_code() {
|
||||
let region_id = RegionId::new(1024, 1);
|
||||
let handler = Arc::new(TestRegionQueryHandler::with_responses(vec![(
|
||||
region_id,
|
||||
int64_schema(&["a", "b"]),
|
||||
vec![Err(common_recordbatch::error::ExternalSnafu.into_error(
|
||||
BoxedError::new(PlainError::new(
|
||||
"neutral stream error".to_string(),
|
||||
StatusCode::RequestOutdated,
|
||||
)),
|
||||
))],
|
||||
)]));
|
||||
let exec =
|
||||
merge_scan_exec_with_handler(vec![region_id], expected_int64_schema(), handler, 1);
|
||||
let mut stream = common_recordbatch::adapter::RecordBatchStreamAdapter::try_new(
|
||||
exec.to_stream(Arc::new(TaskContext::default()), 0).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let error = stream.next().await.unwrap().unwrap_err();
|
||||
assert_eq!(error.status_code(), StatusCode::RequestOutdated);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn aborting_pending_do_get_poll_rolls_back_subscriber_without_starting_fanout() {
|
||||
let handler = Arc::new(PendingDoGetHandler::default());
|
||||
@@ -2080,10 +2113,9 @@ mod tests {
|
||||
assert_eq!(registry_manager.registry_count(), 0);
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct TestRegionResponse {
|
||||
advertised_schema: Arc<Schema>,
|
||||
batches: Vec<RecordBatch>,
|
||||
batches: Vec<common_recordbatch::error::Result<RecordBatch>>,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
@@ -2100,7 +2132,7 @@ mod tests {
|
||||
region_id,
|
||||
TestRegionResponse {
|
||||
advertised_schema: batch.schema.clone(),
|
||||
batches: vec![batch],
|
||||
batches: vec![Ok(batch)],
|
||||
},
|
||||
)
|
||||
})
|
||||
@@ -2109,7 +2141,13 @@ mod tests {
|
||||
}
|
||||
|
||||
fn with_responses(
|
||||
responses: impl IntoIterator<Item = (RegionId, Arc<Schema>, Vec<RecordBatch>)>,
|
||||
responses: impl IntoIterator<
|
||||
Item = (
|
||||
RegionId,
|
||||
Arc<Schema>,
|
||||
Vec<common_recordbatch::error::Result<RecordBatch>>,
|
||||
),
|
||||
>,
|
||||
) -> Self {
|
||||
let responses = responses
|
||||
.into_iter()
|
||||
@@ -2129,19 +2167,17 @@ mod tests {
|
||||
|
||||
struct TestRecordBatchStream {
|
||||
schema: Arc<Schema>,
|
||||
batches: Vec<RecordBatch>,
|
||||
index: usize,
|
||||
batches: Vec<common_recordbatch::error::Result<RecordBatch>>,
|
||||
}
|
||||
|
||||
impl Stream for TestRecordBatchStream {
|
||||
type Item = common_recordbatch::error::Result<RecordBatch>;
|
||||
|
||||
fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
||||
if let Some(batch) = self.batches.get(self.index).cloned() {
|
||||
self.index += 1;
|
||||
Poll::Ready(Some(Ok(batch)))
|
||||
} else {
|
||||
if self.batches.is_empty() {
|
||||
Poll::Ready(None)
|
||||
} else {
|
||||
Poll::Ready(Some(self.batches.remove(0)))
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -2202,10 +2238,13 @@ mod tests {
|
||||
}),
|
||||
Ordering::SeqCst,
|
||||
);
|
||||
crate::error::UnimplementedSnafu {
|
||||
operation: "test do_get failure",
|
||||
}
|
||||
.fail()
|
||||
Err(crate::error::Error::QueryExecution {
|
||||
source: BoxedError::new(PlainError::new(
|
||||
"neutral do_get error".to_string(),
|
||||
StatusCode::RequestOutdated,
|
||||
)),
|
||||
location: snafu::Location::default(),
|
||||
})
|
||||
}
|
||||
|
||||
async fn handle_remote_dyn_filter_update(
|
||||
@@ -2523,8 +2562,19 @@ mod tests {
|
||||
.expect("test handler needs a response for every requested region");
|
||||
Ok(Box::pin(TestRecordBatchStream {
|
||||
schema: response.advertised_schema.clone(),
|
||||
batches: response.batches.clone(),
|
||||
index: 0,
|
||||
batches: response
|
||||
.batches
|
||||
.iter()
|
||||
.map(|batch| match batch {
|
||||
Ok(batch) => Ok(batch.clone()),
|
||||
Err(error) => Err(common_recordbatch::error::ExternalSnafu.into_error(
|
||||
BoxedError::new(PlainError::new(
|
||||
error.to_string(),
|
||||
error.status_code(),
|
||||
)),
|
||||
)),
|
||||
})
|
||||
.collect(),
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -3085,7 +3135,7 @@ mod tests {
|
||||
Arc::new(TestRegionQueryHandler::with_responses(vec![(
|
||||
RegionId::new(1024, 1),
|
||||
remote_schema,
|
||||
vec![batch],
|
||||
vec![Ok(batch)],
|
||||
)])),
|
||||
1,
|
||||
))
|
||||
@@ -3122,7 +3172,7 @@ mod tests {
|
||||
Arc::new(TestRegionQueryHandler::with_responses(vec![(
|
||||
region_id,
|
||||
advertised_schema,
|
||||
vec![batch],
|
||||
vec![Ok(batch)],
|
||||
)])),
|
||||
1,
|
||||
);
|
||||
@@ -3155,7 +3205,7 @@ mod tests {
|
||||
Arc::new(TestRegionQueryHandler::with_responses(vec![(
|
||||
region_id,
|
||||
advertised_schema,
|
||||
vec![batch],
|
||||
vec![Ok(batch)],
|
||||
)])),
|
||||
1,
|
||||
);
|
||||
@@ -3190,7 +3240,7 @@ mod tests {
|
||||
Arc::new(TestRegionQueryHandler::with_responses(vec![(
|
||||
region_id,
|
||||
advertised_schema,
|
||||
vec![batch],
|
||||
vec![Ok(batch)],
|
||||
)])),
|
||||
1,
|
||||
))
|
||||
|
||||
@@ -498,8 +498,23 @@ impl From<Error> for DataFusionError {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use common_error::ext::PlainError;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_datafusion_external_boxed_error_status_code() {
|
||||
let error = Error::DataFusion {
|
||||
error: DataFusionError::External(Box::new(BoxedError::new(PlainError::new(
|
||||
"neutral error".to_string(),
|
||||
StatusCode::RequestOutdated,
|
||||
)))),
|
||||
location: Location::default(),
|
||||
};
|
||||
|
||||
assert_eq!(error.status_code(), StatusCode::RequestOutdated);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_backend_delegates_error_metadata() {
|
||||
let source = common_datasource::error::LocalFileAccessDisabledSnafu {
|
||||
|
||||
@@ -745,7 +745,7 @@ impl ErrorExt for Error {
|
||||
#[cfg(not(windows))]
|
||||
UpdateJemallocMetrics { .. } => StatusCode::Internal,
|
||||
|
||||
CollectRecordbatch { .. } => StatusCode::EngineExecuteQuery,
|
||||
CollectRecordbatch { source, .. } => source.status_code(),
|
||||
|
||||
ExecuteQuery { source, .. }
|
||||
| ExecutePlan { source, .. }
|
||||
@@ -986,3 +986,37 @@ pub fn status_code_to_http_status(status_code: &StatusCode) -> HttpStatusCode {
|
||||
| StatusCode::EngineExecuteQuery => HttpStatusCode::INTERNAL_SERVER_ERROR,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use common_error::GREPTIME_DB_HEADER_ERROR_CODE;
|
||||
use common_error::ext::PlainError;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn collect_recordbatch_preserves_poll_stream_status_in_tonic_status() {
|
||||
let error = Error::CollectRecordbatch {
|
||||
source: common_recordbatch::error::Error::PollStream {
|
||||
error: DataFusionError::External(Box::new(BoxedError::new(PlainError::new(
|
||||
"neutral error".to_string(),
|
||||
StatusCode::RequestOutdated,
|
||||
)))),
|
||||
location: Location::default(),
|
||||
},
|
||||
location: Location::default(),
|
||||
};
|
||||
|
||||
let status: tonic::Status = error.into();
|
||||
assert_eq!(status.code(), tonic::Code::InvalidArgument);
|
||||
assert_eq!(
|
||||
status
|
||||
.metadata()
|
||||
.get(GREPTIME_DB_HEADER_ERROR_CODE)
|
||||
.unwrap()
|
||||
.to_str()
|
||||
.unwrap(),
|
||||
(StatusCode::RequestOutdated as u32).to_string()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -33,6 +33,8 @@ mod test {
|
||||
use client::region::RegionRequester;
|
||||
use client::{Client, Database};
|
||||
use common_catalog::consts::{DEFAULT_CATALOG_NAME, DEFAULT_SCHEMA_NAME};
|
||||
use common_error::ext::ErrorExt;
|
||||
use common_error::status_code::StatusCode;
|
||||
use common_grpc::channel_manager::{ChannelConfig, ChannelManager};
|
||||
use common_grpc::flight::do_put::{DoPutMetadata, DoPutResponse};
|
||||
use common_grpc::flight::{FlightDecoder, FlightEncoder, FlightMessage};
|
||||
@@ -720,12 +722,14 @@ mod test {
|
||||
let mut stale_fence_error = None;
|
||||
while let Some(batch) = stream.next().await {
|
||||
if let Err(err) = batch {
|
||||
stale_fence_error = Some(format!("{err:?}"));
|
||||
stale_fence_error = Some(err);
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
let err_msg = stale_fence_error.expect("expected stale snapshot fence rejection");
|
||||
let stale_fence_error = stale_fence_error.expect("expected stale snapshot fence rejection");
|
||||
assert_eq!(stale_fence_error.status_code(), StatusCode::RequestOutdated);
|
||||
let err_msg = format!("{stale_fence_error:?}");
|
||||
assert!(
|
||||
err_msg.contains("STALE_SNAPSHOT_FENCE")
|
||||
|| err_msg.contains("RequestOutdated")
|
||||
|
||||
Reference in New Issue
Block a user