diff --git a/src/common/query/src/error.rs b/src/common/query/src/error.rs index dd2b29adf3..538c18b704 100644 --- a/src/common/query/src/error.rs +++ b/src/common/query/src/error.rs @@ -277,6 +277,8 @@ pub fn datafusion_status_code( DataFusionError::External(e) => { if let Some(ext) = (*e).downcast_ref::() { ext.status_code() + } else if let Some(ext) = (*e).downcast_ref::() { + ext.status_code() } else { default_status.unwrap_or(StatusCode::EngineExecuteQuery) } @@ -285,3 +287,49 @@ pub fn datafusion_status_code( _ => 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::(&boxed_error(), None), + StatusCode::RequestOutdated + ); + assert_eq!( + datafusion_status_code::(&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::(&direct_error, None), + StatusCode::PlanQuery + ); + + let unknown_error = + || DataFusionError::External(Box::new(std::io::Error::other("neutral error"))); + assert_eq!( + datafusion_status_code::(&unknown_error(), None), + StatusCode::EngineExecuteQuery + ); + assert_eq!( + datafusion_status_code::(&unknown_error(), Some(StatusCode::PlanQuery)), + StatusCode::PlanQuery + ); + } +} diff --git a/src/common/recordbatch/src/error.rs b/src/common/recordbatch/src/error.rs index 469320c19c..f1425492fb 100644 --- a/src/common/recordbatch/src/error.rs +++ b/src/common/recordbatch/src/error.rs @@ -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::() + .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); + } + } +} diff --git a/src/query/src/dist_plan/merge_scan.rs b/src/query/src/dist_plan/merge_scan.rs index 26474b7ca2..ffe154ade0 100644 --- a/src/query/src/dist_plan/merge_scan.rs +++ b/src/query/src/dist_plan/merge_scan.rs @@ -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, - batches: Vec, + batches: Vec>, } #[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, Vec)>, + responses: impl IntoIterator< + Item = ( + RegionId, + Arc, + Vec>, + ), + >, ) -> Self { let responses = responses .into_iter() @@ -2129,19 +2167,17 @@ mod tests { struct TestRecordBatchStream { schema: Arc, - batches: Vec, - index: usize, + batches: Vec>, } impl Stream for TestRecordBatchStream { type Item = common_recordbatch::error::Result; fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { - 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, )) diff --git a/src/query/src/error.rs b/src/query/src/error.rs index fdfeaafde9..a294641557 100644 --- a/src/query/src/error.rs +++ b/src/query/src/error.rs @@ -498,8 +498,23 @@ impl From 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 { diff --git a/src/servers/src/error.rs b/src/servers/src/error.rs index fa5957759c..fe1d226a10 100644 --- a/src/servers/src/error.rs +++ b/src/servers/src/error.rs @@ -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() + ); + } +} diff --git a/tests-integration/src/grpc/flight.rs b/tests-integration/src/grpc/flight.rs index f37d5ee638..d4b8927e5b 100644 --- a/tests-integration/src/grpc/flight.rs +++ b/tests-integration/src/grpc/flight.rs @@ -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")