fix(query): preserve typed errors through distributed Flight streams

Signed-off-by: discord9 <discord9@163.com>
This commit is contained in:
discord9
2026-09-08 13:57:38 +08:00
parent 6ddb2d2ec4
commit d5ecdd3361
6 changed files with 231 additions and 28 deletions
+48
View File
@@ -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
);
}
}
+52
View File
@@ -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);
}
}
}
+75 -25
View File
@@ -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,
))
+15
View File
@@ -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 {
+35 -1
View File
@@ -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()
);
}
}
+6 -2
View File
@@ -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")