refactor(flight): add request builder and defer DoGet execution (#8953)

* refactor: add flight request builder

Signed-off-by: WenyXu <wenymedia@gmail.com>

* refactor: use flight request builder in flow

Signed-off-by: WenyXu <wenymedia@gmail.com>

* refactor: defer frontend flight query execution

Signed-off-by: WenyXu <wenymedia@gmail.com>

* fix(grpc): avoid cloning requests during auth

Signed-off-by: WenyXu <wenymedia@gmail.com>

* test(grpc): cover flight request timeout

Signed-off-by: WenyXu <wenymedia@gmail.com>

* refactor(client): share Flight message reader

Signed-off-by: WenyXu <wenymedia@gmail.com>

* feat(flow): add Flight DoGet timeout

Signed-off-by: WenyXu <wenymedia@gmail.com>

* refactor(client): gate Flight DDL helpers for testing

Signed-off-by: WenyXu <wenymedia@gmail.com>

* fix(client): restore Flight stream error semantics

Signed-off-by: WenyXu <wenymedia@gmail.com>

* docs(grpc): document Flight stream input constructors

Signed-off-by: WenyXu <wenymedia@gmail.com>

* fix(flight): preserve deferred stream context

Signed-off-by: WenyXu <wenymedia@gmail.com>

* fix(client): use Flight stream SNAFU context

Signed-off-by: WenyXu <wenymedia@gmail.com>

---------

Signed-off-by: WenyXu <wenymedia@gmail.com>
This commit is contained in:
Weny Xu
2026-08-26 14:00:39 +00:00
committed by GitHub
parent 35ea88a4ef
commit 1409e66837
12 changed files with 923 additions and 399 deletions
+336 -154
View File
@@ -12,20 +12,25 @@
// See the License for the specific language governing permissions and
// limitations under the License.
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::task::{Context, Poll};
use std::time::Duration;
use api::v1::auth_header::AuthScheme;
#[cfg(feature = "testing")]
use api::v1::ddl_request::Expr as DdlExpr;
use api::v1::greptime_database_client::GreptimeDatabaseClient;
use api::v1::greptime_request::Request;
use api::v1::query_request::Query;
#[cfg(feature = "testing")]
use api::v1::{AlterTableExpr, CreateTableExpr, DdlRequest};
use api::v1::{
AlterTableExpr, AuthHeader, Basic, CreateTableExpr, DdlRequest, GreptimeRequest,
InsertRequests, QueryRequest, RequestHeader, RowInsertRequests,
AuthHeader, Basic, GreptimeRequest, InsertRequests, QueryRequest, RequestHeader,
RowInsertRequests,
};
use arc_swap::ArcSwapOption;
use arrow_flight::{FlightData, Ticket};
@@ -34,7 +39,7 @@ use base64::Engine;
use base64::prelude::BASE64_STANDARD;
use common_catalog::build_db_string;
use common_catalog::consts::{DEFAULT_CATALOG_NAME, DEFAULT_SCHEMA_NAME};
use common_error::ext::BoxedError;
use common_error::ext::{BoxedError, ErrorExt};
use common_grpc::flight::do_put::DoPutResponse;
use common_grpc::flight::{
FLOW_EXTENSIONS_METADATA_KEY, FlightDecoder, FlightMessage, SNAPSHOT_SEQS_METADATA_KEY,
@@ -49,21 +54,23 @@ use common_telemetry::{error, warn};
use futures::future;
use futures_util::{Stream, StreamExt, TryStreamExt};
use prost::Message;
use snafu::ResultExt;
use snafu::{IntoError, ResultExt};
use tonic::metadata::{AsciiMetadataKey, AsciiMetadataValue, MetadataMap, MetadataValue};
use tonic::transport::Channel;
use crate::error::{
ConvertFlightDataSnafu, Error, FlightGetSnafu, IllegalFlightMessagesSnafu,
ConvertFlightDataSnafu, Error, FlightGetSnafu, FlightStreamSnafu, IllegalFlightMessagesSnafu,
InvalidTonicMetadataValueSnafu,
};
use crate::flight::decode_flight_data;
use crate::flight::{FlightMessageReader, decode_flight_data};
use crate::{Client, Result, error, from_grpc_response};
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";
/// Terminal metrics associated with a query output.
///
/// For streaming outputs, metrics are only final after the stream is fully
@@ -257,19 +264,17 @@ fn attach_terminal_metrics(output: Output, terminal_metrics: &OutputMetrics) ->
}
async fn output_from_flight_message_stream<S>(
mut flight_message_stream: S,
remote_addr: String,
flight_message_stream: S,
) -> Result<OutputWithMetrics>
where
S: Stream<Item = Result<FlightMessage>> + Send + Unpin + 'static,
{
let Some(first_flight_message) = flight_message_stream.next().await else {
return IllegalFlightMessagesSnafu {
reason: "Expect the response not to be empty",
}
.fail();
};
let first_flight_message = first_flight_message?;
let mut reader = FlightMessageReader::new(remote_addr, flight_message_stream);
let first_flight_message = reader
.read_first()
.await
.map_err(|error| flight_stream_error(reader.remote_addr(), error))?;
match first_flight_message {
FlightMessage::AffectedRows { rows, metrics } => {
@@ -277,7 +282,10 @@ where
if let Some(metrics) = metrics {
terminal_metrics.update(Some(parse_terminal_metrics(&metrics)?));
}
let next_message = flight_message_stream.next().await.transpose()?;
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() => {
@@ -316,15 +324,19 @@ where
);
let schema_cloned = schema.clone();
let stream = Box::pin(stream!({
while let Some(flight_message_item) = flight_message_stream.next().await {
let flight_message = match flight_message_item {
Ok(message) => message,
Err(e) => {
yield Err(BoxedError::new(e)).context(ExternalSnafu);
loop {
let flight_message = match reader.read_next().await {
Ok(Some(message)) => message,
Ok(None) => break,
Err(error) => {
yield Err(BoxedError::new(flight_stream_error(
reader.remote_addr(),
error,
)))
.context(ExternalSnafu);
break;
}
};
match flight_message {
FlightMessage::RecordBatch(arrow_batch) => {
yield Ok(RecordBatch::from_df_record_batch(
@@ -371,6 +383,25 @@ where
}
}
fn flight_stream_error(addr: &str, error: Error) -> Error {
let tonic_code = error.tonic_code().unwrap_or(tonic::Code::Unknown);
let message = error.to_string();
if error.status_code().should_log_error() {
error!(
error; "Failed to receive Flight data, addr: {}, code: {}",
addr,
tonic_code
);
}
FlightStreamSnafu {
addr: addr.to_string(),
tonic_code,
message,
}
.into_error(BoxedError::new(error))
}
#[derive(Clone, Debug, Default)]
pub struct Database {
// The "catalog" and "schema" to be used in processing the requests at the server side.
@@ -389,6 +420,42 @@ pub struct Database {
ctx: FlightContext,
}
#[derive(Default)]
struct FlightRequestOptions {
hints: Option<String>,
flow_extensions: Option<String>,
snapshot_seqs: Option<String>,
timeout: Option<Duration>,
}
impl FlightRequestOptions {
fn apply_to<T>(self, request: &mut tonic::Request<T>) -> Result<()> {
let metadata = request.metadata_mut();
if let Some(hints) = self.hints {
Database::put_metadata_value(metadata, HINTS_METADATA_KEY, hints)?;
}
if let Some(flow_extensions) = self.flow_extensions {
Database::put_metadata_value(metadata, FLOW_EXTENSIONS_METADATA_KEY, flow_extensions)?;
}
if let Some(snapshot_seqs) = self.snapshot_seqs {
Database::put_metadata_value(metadata, SNAPSHOT_SEQS_METADATA_KEY, snapshot_seqs)?;
}
if let Some(timeout) = self.timeout {
request.set_timeout(timeout);
}
Ok(())
}
}
/// A single Flight DoGet request to a [`Database`].
///
/// The builder carries request-scoped metadata and options. It does not modify
/// the underlying [`Database`], so its configuration cannot affect later RPCs.
pub struct DatabaseFlightRequest<'a> {
database: &'a Database,
options: FlightRequestOptions,
}
pub struct DatabaseClient {
pub addr: String,
pub inner: GreptimeDatabaseClient<Channel>,
@@ -483,6 +550,14 @@ impl Database {
});
}
/// Creates a builder for a single Flight DoGet request.
pub fn flight_request(&self) -> DatabaseFlightRequest<'_> {
DatabaseFlightRequest {
database: self,
options: FlightRequestOptions::default(),
}
}
/// Make an InsertRequests request to the database.
pub async fn insert(&self, requests: InsertRequests) -> Result<u32> {
self.handle(Request::Inserts(requests)).await
@@ -538,44 +613,31 @@ impl Database {
}
fn put_hints(metadata: &mut MetadataMap, hints: &[(&str, &str)]) -> Result<()> {
let Some(value) = hints
.iter()
.map(|(k, v)| format!("{}={}", k, v))
.reduce(|a, b| format!("{},{}", a, b))
else {
let Some(value) = Self::encode_hints(hints) else {
return Ok(());
};
let key = AsciiMetadataKey::from_static("x-greptime-hints");
let value = AsciiMetadataValue::from_str(&value).context(InvalidTonicMetadataValueSnafu)?;
metadata.insert(key, value);
Ok(())
Self::put_metadata_value(metadata, HINTS_METADATA_KEY, value)
}
fn put_flow_extensions(
metadata: &mut MetadataMap,
flow_extensions: &[(&str, &str)],
) -> Result<()> {
if flow_extensions.is_empty() {
return Ok(());
}
let value = serde_json::to_string(&flow_extensions.to_vec())
.expect("flow extension pairs should serialize");
Self::put_metadata_value(metadata, FLOW_EXTENSIONS_METADATA_KEY, value)
fn encode_hints(hints: &[(&str, &str)]) -> Option<String> {
hints
.iter()
.map(|(k, v)| format!("{}={}", k, v))
.reduce(|a, b| format!("{},{}", a, b))
}
fn put_snapshot_seqs(
metadata: &mut MetadataMap,
snapshot_seqs: &std::collections::HashMap<u64, u64>,
) -> Result<()> {
if snapshot_seqs.is_empty() {
return Ok(());
}
fn encode_flow_extensions(flow_extensions: &[(&str, &str)]) -> Option<String> {
(!flow_extensions.is_empty()).then(|| {
serde_json::to_string(&flow_extensions.to_vec())
.expect("flow extension pairs should serialize")
})
}
let value =
serde_json::to_string(snapshot_seqs).expect("snapshot sequence map should serialize");
Self::put_metadata_value(metadata, SNAPSHOT_SEQS_METADATA_KEY, value)
fn encode_snapshot_seqs(snapshot_seqs: &HashMap<u64, u64>) -> Option<String> {
(!snapshot_seqs.is_empty()).then(|| {
serde_json::to_string(snapshot_seqs).expect("snapshot sequence map should serialize")
})
}
fn put_metadata_value(
@@ -673,7 +735,7 @@ impl Database {
where
S: AsRef<str>,
{
self.sql_with_hint(sql, &[]).await
self.flight_request().sql(sql).await
}
/// Executes a SQL query with optional hints for query optimization.
@@ -681,12 +743,7 @@ impl Database {
where
S: AsRef<str>,
{
let request = Request::Query(QueryRequest {
query: Some(Query::Sql(sql.as_ref().to_string())),
});
self.do_get(request, hints, &[], &Default::default())
.await
.map(OutputWithMetrics::into_output)
self.flight_request().with_hints(hints).sql(sql).await
}
/// Executes a SQL query and returns the output with terminal metrics.
@@ -701,77 +758,33 @@ impl Database {
where
S: AsRef<str>,
{
self.query_with_terminal_metrics_and_flow_extensions(
QueryRequest {
query: Some(Query::Sql(sql.as_ref().to_string())),
},
hints,
&[],
&Default::default(),
)
.await
self.flight_request()
.with_hints(hints)
.sql_with_terminal_metrics(sql)
.await
}
/// Executes a logical plan directly without SQL parsing.
pub async fn logical_plan(&self, logical_plan: Vec<u8>) -> Result<Output> {
self.query_with_terminal_metrics_and_flow_extensions(
QueryRequest {
query: Some(Query::LogicalPlan(logical_plan)),
},
&[],
&[],
&Default::default(),
)
.await
.map(OutputWithMetrics::into_output)
}
/// Executes a query and carries flow extensions through Flight metadata.
///
/// This is the lower-level terminal-metrics API for Flow callers that need
/// to pass JSON-bearing flow extensions without going through hint metadata.
pub async fn query_with_terminal_metrics_and_flow_extensions(
&self,
request: QueryRequest,
hints: &[(&str, &str)],
flow_extensions: &[(&str, &str)],
snapshot_seqs: &std::collections::HashMap<u64, u64>,
) -> Result<OutputWithMetrics> {
self.do_get(
Request::Query(request),
hints,
flow_extensions,
snapshot_seqs,
)
.await
self.flight_request().logical_plan(logical_plan).await
}
/// Creates a new table using the provided table expression.
#[cfg(feature = "testing")]
pub async fn create(&self, expr: CreateTableExpr) -> Result<Output> {
let request = Request::Ddl(DdlRequest {
expr: Some(DdlExpr::CreateTable(expr)),
});
self.do_get(request, &[], &[], &Default::default())
.await
.map(OutputWithMetrics::into_output)
self.flight_request().create(expr).await
}
/// Alters an existing table using the provided alter expression.
#[cfg(feature = "testing")]
pub async fn alter(&self, expr: AlterTableExpr) -> Result<Output> {
let request = Request::Ddl(DdlRequest {
expr: Some(DdlExpr::AlterTable(expr)),
});
self.do_get(request, &[], &[], &Default::default())
.await
.map(OutputWithMetrics::into_output)
self.flight_request().alter(expr).await
}
async fn do_get(
&self,
request: Request,
hints: &[(&str, &str)],
flow_extensions: &[(&str, &str)],
snapshot_seqs: &std::collections::HashMap<u64, u64>,
options: FlightRequestOptions,
) -> Result<OutputWithMetrics> {
let request = self.to_rpc_request(request);
let request = Ticket {
@@ -779,12 +792,10 @@ impl Database {
};
let mut request = tonic::Request::new(request);
let metadata = request.metadata_mut();
Self::put_hints(metadata, hints)?;
Self::put_flow_extensions(metadata, flow_extensions)?;
Self::put_snapshot_seqs(metadata, snapshot_seqs)?;
options.apply_to(&mut request)?;
let mut client = self.client.make_flight_client(false, false)?;
let remote_addr = client.addr().to_string();
let response = client.mut_inner().do_get(request).await.or_else(|e| {
let tonic_code = e.code();
@@ -796,7 +807,7 @@ impl Database {
e
);
Err(BoxedError::new(e)).with_context(|_| FlightGetSnafu {
addr: client.addr().to_string(),
addr: remote_addr.clone(),
tonic_code,
})
})?;
@@ -808,7 +819,7 @@ impl Database {
future::ready(decode_flight_data(&mut decoder, flight_data))
});
output_from_flight_message_stream(flight_message_stream).await
output_from_flight_message_stream(remote_addr, flight_message_stream).await
}
/// Ingest a stream of [RecordBatch]es that belong to a table, using Arrow Flight's "`DoPut`"
@@ -847,6 +858,98 @@ impl Database {
}
}
impl<'a> DatabaseFlightRequest<'a> {
/// Adds query optimization hints to this Flight request.
pub fn with_hints(mut self, hints: &[(&str, &str)]) -> Self {
self.options.hints = Database::encode_hints(hints);
self
}
/// Adds Flow extensions to this Flight request.
pub fn with_flow_extensions(mut self, flow_extensions: &[(&str, &str)]) -> Self {
self.options.flow_extensions = Database::encode_flow_extensions(flow_extensions);
self
}
/// Adds snapshot sequence fences to this Flight request.
pub fn with_snapshot_seqs(mut self, snapshot_seqs: &HashMap<u64, u64>) -> Self {
self.options.snapshot_seqs = Database::encode_snapshot_seqs(snapshot_seqs);
self
}
/// Sets a timeout for this Flight request only.
pub fn with_timeout(mut self, timeout: Duration) -> Self {
self.options.timeout = Some(timeout);
self
}
/// Executes a SQL query.
pub async fn sql<S>(self, sql: S) -> Result<Output>
where
S: AsRef<str>,
{
let request = Request::Query(QueryRequest {
query: Some(Query::Sql(sql.as_ref().to_string())),
});
self.do_get(request)
.await
.map(OutputWithMetrics::into_output)
}
/// Executes a SQL query and returns terminal metrics.
pub async fn sql_with_terminal_metrics<S>(self, sql: S) -> Result<OutputWithMetrics>
where
S: AsRef<str>,
{
self.query_with_terminal_metrics(QueryRequest {
query: Some(Query::Sql(sql.as_ref().to_string())),
})
.await
}
/// Executes a logical plan directly without SQL parsing.
pub async fn logical_plan(self, logical_plan: Vec<u8>) -> Result<Output> {
self.query_with_terminal_metrics(QueryRequest {
query: Some(Query::LogicalPlan(logical_plan)),
})
.await
.map(OutputWithMetrics::into_output)
}
/// Executes a query and returns terminal metrics.
pub async fn query_with_terminal_metrics(
self,
request: QueryRequest,
) -> Result<OutputWithMetrics> {
self.do_get(Request::Query(request)).await
}
/// Creates a new table using the provided table expression.
#[cfg(feature = "testing")]
pub async fn create(self, expr: CreateTableExpr) -> Result<Output> {
self.do_get(Request::Ddl(DdlRequest {
expr: Some(DdlExpr::CreateTable(expr)),
}))
.await
.map(OutputWithMetrics::into_output)
}
/// Alters an existing table using the provided alter expression.
#[cfg(feature = "testing")]
pub async fn alter(self, expr: AlterTableExpr) -> Result<Output> {
self.do_get(Request::Ddl(DdlRequest {
expr: Some(DdlExpr::AlterTable(expr)),
}))
.await
.map(OutputWithMetrics::into_output)
}
async fn do_get(self, request: Request) -> Result<OutputWithMetrics> {
let Self { database, options } = self;
database.do_get(request, options).await
}
}
/// by grpc standard, only `Unavailable` is retryable, see: https://github.com/grpc/grpc/blob/master/doc/statuscodes.md#status-codes-and-their-use-in-grpc
pub fn is_grpc_retryable(err: &tonic::Status) -> bool {
matches!(err.code(), tonic::Code::Unavailable)
@@ -934,12 +1037,14 @@ mod tests {
#[test]
fn test_put_flow_extensions_preserves_comma_bearing_values() {
let mut metadata = MetadataMap::new();
Database::put_flow_extensions(
Database::put_metadata_value(
&mut metadata,
&[
FLOW_EXTENSIONS_METADATA_KEY,
Database::encode_flow_extensions(&[
("flow.return_region_seq", "true"),
("flow.incremental_after_seqs", r#"{"1":10,"2":20}"#),
],
])
.unwrap(),
)
.unwrap();
@@ -969,7 +1074,12 @@ mod tests {
(9_007_199_254_740_993_u64, 9_007_199_254_740_995_u64),
]);
Database::put_snapshot_seqs(&mut metadata, &snapshot_seqs).unwrap();
Database::put_metadata_value(
&mut metadata,
SNAPSHOT_SEQS_METADATA_KEY,
Database::encode_snapshot_seqs(&snapshot_seqs).unwrap(),
)
.unwrap();
let value = metadata
.get(SNAPSHOT_SEQS_METADATA_KEY)
@@ -980,6 +1090,50 @@ mod tests {
assert_eq!(decoded, snapshot_seqs);
}
#[test]
fn test_flight_request_builder_applies_request_options() {
let database = Database::new("greptime", "public", Client::default());
let snapshot_seqs = HashMap::from([(42, 99)]);
let request = database
.flight_request()
.with_hints(&[("query_parallelism", "1")])
.with_flow_extensions(&[("flow.return_region_seq", "true")])
.with_snapshot_seqs(&snapshot_seqs)
.with_timeout(Duration::from_millis(50));
let mut tonic_request = tonic::Request::new(());
request.options.apply_to(&mut tonic_request).unwrap();
let metadata = tonic_request.metadata();
assert_eq!(
metadata.get(HINTS_METADATA_KEY).unwrap(),
"query_parallelism=1"
);
assert_eq!(
serde_json::from_str::<Vec<(String, String)>>(
metadata
.get(FLOW_EXTENSIONS_METADATA_KEY)
.unwrap()
.to_str()
.unwrap(),
)
.unwrap(),
vec![("flow.return_region_seq".to_string(), "true".to_string())]
);
assert_eq!(
serde_json::from_str::<HashMap<u64, u64>>(
metadata
.get(SNAPSHOT_SEQS_METADATA_KEY)
.unwrap()
.to_str()
.unwrap(),
)
.unwrap(),
snapshot_seqs
);
assert!(metadata.get("grpc-timeout").is_some());
}
#[test]
fn test_flight_ctx() {
let mut ctx = FlightContext::default();
@@ -1020,6 +1174,28 @@ mod tests {
assert_eq!(expected.should_retry(), actual.should_retry());
}
#[test]
fn test_flight_stream_error_preserves_addr_and_message() {
let error = flight_stream_error(
"127.0.0.1:4001",
Status::out_of_range("message length too large").into(),
);
assert!(matches!(
&error,
Error::FlightStream {
addr,
tonic_code: Code::OutOfRange,
message,
..
} if addr == "127.0.0.1:4001" && message == "message length too large"
));
assert_eq!(
"Failed to receive Flight data from 127.0.0.1:4001, code: Operation was attempted past the valid range: message length too large",
error.to_string(),
);
}
#[test]
fn test_from_tonic_status_with_retry_hint() {
let mut headers = HeaderMap::new();
@@ -1117,13 +1293,13 @@ mod tests {
#[tokio::test]
async fn test_affected_rows_inline_metrics_are_parsed() {
let output = output_from_flight_message_stream(futures_util::stream::iter(vec![Ok(
FlightMessage::AffectedRows {
let output = output_from_flight_message_stream(
"test-peer".to_string(),
futures_util::stream::iter(vec![Ok(FlightMessage::AffectedRows {
rows: 3,
metrics: Some(terminal_metrics_json()),
},
)]
as Vec<Result<FlightMessage>>))
})] as Vec<Result<FlightMessage>>),
)
.await
.unwrap();
@@ -1138,14 +1314,16 @@ mod tests {
#[tokio::test]
async fn test_affected_rows_inline_metrics_rejects_trailing_metrics() {
let metrics_json = terminal_metrics_json();
let err = output_from_flight_message_stream(futures_util::stream::iter(vec![
Ok(FlightMessage::AffectedRows {
rows: 3,
metrics: Some(metrics_json.clone()),
}),
Ok(FlightMessage::Metrics(metrics_json)),
]
as Vec<Result<FlightMessage>>))
let err = output_from_flight_message_stream(
"test-peer".to_string(),
futures_util::stream::iter(vec![
Ok(FlightMessage::AffectedRows {
rows: 3,
metrics: Some(metrics_json.clone()),
}),
Ok(FlightMessage::Metrics(metrics_json)),
] as Vec<Result<FlightMessage>>),
)
.await
.unwrap_err();
@@ -1167,12 +1345,14 @@ mod tests {
vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
)
.unwrap();
let output = output_from_flight_message_stream(futures_util::stream::iter(vec![
Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
Ok(FlightMessage::RecordBatch(batch.into_df_record_batch())),
Ok(FlightMessage::Metrics("{not-json}".to_string())),
]
as Vec<Result<FlightMessage>>))
let output = output_from_flight_message_stream(
"test-peer".to_string(),
futures_util::stream::iter(vec![
Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
Ok(FlightMessage::RecordBatch(batch.into_df_record_batch())),
Ok(FlightMessage::Metrics("{not-json}".to_string())),
] as Vec<Result<FlightMessage>>),
)
.await
.unwrap();
let terminal_metrics = output.metrics.clone();
@@ -1211,18 +1391,20 @@ mod tests {
vec![Arc::new(Int32Vector::from_slice([2])) as VectorRef],
)
.unwrap();
let output = output_from_flight_message_stream(futures_util::stream::iter(vec![
Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
Ok(FlightMessage::RecordBatch(
first_batch.into_df_record_batch(),
)),
Ok(FlightMessage::Metrics(terminal_metrics_json_with_seq(1))),
Ok(FlightMessage::RecordBatch(
second_batch.into_df_record_batch(),
)),
Ok(FlightMessage::Metrics(terminal_metrics_json_with_seq(2))),
]
as Vec<Result<FlightMessage>>))
let output = output_from_flight_message_stream(
"test-peer".to_string(),
futures_util::stream::iter(vec![
Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
Ok(FlightMessage::RecordBatch(
first_batch.into_df_record_batch(),
)),
Ok(FlightMessage::Metrics(terminal_metrics_json_with_seq(1))),
Ok(FlightMessage::RecordBatch(
second_batch.into_df_record_batch(),
)),
Ok(FlightMessage::Metrics(terminal_metrics_json_with_seq(2))),
] as Vec<Result<FlightMessage>>),
)
.await
.unwrap();
let terminal_metrics = output.metrics.clone();
+18
View File
@@ -40,6 +40,21 @@ pub enum Error {
source: BoxedError,
},
#[snafu(display(
"Failed to receive Flight data from {}, code: {}: {}",
addr,
tonic_code,
message
))]
FlightStream {
addr: String,
tonic_code: Code,
message: String,
source: BoxedError,
#[snafu(implicit)]
location: Location,
},
#[snafu(display("Failed to convert FlightData"))]
ConvertFlightData {
#[snafu(implicit)]
@@ -154,6 +169,7 @@ impl ErrorExt for Error {
Error::Server { code, .. } | Error::Tonic { code, .. } => *code,
Error::FlightGet { source, .. }
| Error::FlightStream { source, .. }
| Error::RegionServer { source, .. }
| Error::FlowServer { source, .. } => source.status_code(),
Error::CreateChannel { source, .. }
@@ -174,6 +190,7 @@ impl ErrorExt for Error {
match self {
Error::Tonic { retry_hint, .. } => *retry_hint,
Error::FlightGet { source, .. }
| Error::FlightStream { source, .. }
| Error::RegionServer { source, .. }
| Error::FlowServer { source, .. }
| Error::External { source, .. } => source.retry_hint(),
@@ -193,6 +210,7 @@ impl Error {
pub fn tonic_code(&self) -> Option<Code> {
match self {
Self::FlightGet { tonic_code, .. }
| Self::FlightStream { tonic_code, .. }
| Self::RegionServer {
code: tonic_code, ..
}
+72 -2
View File
@@ -12,12 +12,82 @@
// 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 snafu::ResultExt;
use futures_util::stream::Peekable;
use futures_util::{Stream, StreamExt};
use snafu::{OptionExt, ResultExt};
use crate::Result;
use crate::error::{ConvertFlightDataSnafu, Error};
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>,
}
impl<S> FlightMessageReader<S>
where
S: Stream<Item = Result<FlightMessage>> + Unpin,
{
pub(crate) fn new(remote_addr: impl Into<String>, messages: S) -> Self {
Self {
remote_addr: remote_addr.into(),
messages: messages.peekable(),
}
}
pub(crate) fn remote_addr(&self) -> &str {
&self.remote_addr
}
pub(crate) async fn read_first(&mut self) -> Result<FlightMessage> {
self.read_next().await?.context(IllegalFlightMessagesSnafu {
reason: "Expect the response not to be empty",
})
}
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(
decoder: &mut FlightDecoder,
+135 -54
View File
@@ -47,7 +47,7 @@ use crate::error::{
self, FlightGetSnafu, IllegalDatabaseResponseSnafu, IllegalFlightMessagesSnafu,
MissingFieldSnafu, Result, ServerSnafu,
};
use crate::flight::decode_flight_data;
use crate::flight::{FlightMessageKind, FlightMessageReader, decode_flight_data};
use crate::{Client, metrics};
const FLIGHT_DO_GET_TIMEOUT: Duration = Duration::from_secs(10);
@@ -203,26 +203,22 @@ impl RegionRequester {
async fn recordbatches_from_flight_message_stream<S>(
addr: String,
mut flight_message_stream: S,
flight_message_stream: S,
) -> Result<SendableRecordBatchStream>
where
S: Stream<Item = Result<FlightMessage>> + Send + Unpin + 'static,
{
let Some(first_flight_message) = flight_message_stream.next().await else {
return IllegalFlightMessagesSnafu {
reason: "Expect the response not to be empty",
}
.fail()
.map_err(|error| flight_stream_error(&addr, error));
};
let FlightMessage::Schema(schema) =
first_flight_message.map_err(|e| flight_stream_error(&addr, e))?
let mut reader = FlightMessageReader::new(addr.clone(), flight_message_stream);
let FlightMessage::Schema(schema) = reader
.read_first()
.await
.map_err(|error| flight_stream_error(reader.remote_addr(), error))?
else {
return IllegalFlightMessagesSnafu {
reason: "Expect schema to be the first flight message",
}
.fail()
.map_err(|error| flight_stream_error(&addr, error));
.map_err(|error| flight_stream_error(reader.remote_addr(), error));
};
let metrics = Arc::new(ArcSwapOption::from(None));
@@ -233,31 +229,23 @@ where
let schema =
Arc::new(datatypes::schema::Schema::try_from(schema).context(error::ConvertSchemaSnafu)?);
let schema_cloned = schema.clone();
let stream_addr = addr.clone();
let stream_addr = addr;
let stream = Box::pin(stream!({
let _span = tracing_context.attach(common_telemetry::tracing::info_span!(
"poll_flight_data_stream"
));
let mut buffered_message: Option<FlightMessage> = None;
let mut stream_ended = false;
while !stream_ended {
// get the next message from the buffered message or read from the flight message stream
let flight_message_item = if let Some(msg) = buffered_message.take() {
Some(Ok(msg))
} else {
flight_message_stream.next().await
};
let flight_message = match flight_message_item {
Some(Ok(message)) => message,
Some(Err(e)) => {
yield Err(BoxedError::new(flight_stream_error(&stream_addr, e)))
let flight_message = match reader.read_next().await {
Ok(Some(message)) => message,
Ok(None) => break,
Err(error) => {
yield Err(BoxedError::new(flight_stream_error(&stream_addr, error)))
.context(ExternalSnafu);
break;
}
None => break,
};
match flight_message {
@@ -265,42 +253,54 @@ where
let result_to_yield =
RecordBatch::from_df_record_batch(schema_cloned.clone(), record_batch);
// get the next message from the stream. normally it should be a metrics message.
if let Some(next_flight_message_result) = flight_message_stream.next().await {
match next_flight_message_result {
Ok(FlightMessage::Metrics(s)) => {
let m = serde_json::from_str(&s).ok().map(Arc::new);
metrics_ref.swap(m);
}
Ok(FlightMessage::RecordBatch(rb)) => {
// for some reason it's not a metrics message, so we need to buffer this record batch
// and yield it in the next iteration.
buffered_message = Some(FlightMessage::RecordBatch(rb));
}
Ok(_) => {
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;
}
Err(e) => {
yield Err(BoxedError::new(flight_stream_error(&stream_addr, e)))
// 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;
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;
}
} else {
// the stream has ended
stream_ended = true;
}
yield Ok(result_to_yield);
}
FlightMessage::Metrics(s) => {
// just a branch in case of some metrics message comes after other things.
// Metrics may arrive before the next RecordBatch.
let m = serde_json::from_str(&s).ok().map(Arc::new);
metrics_ref.swap(m);
continue;
@@ -412,6 +412,7 @@ mod test {
RemoteDynFilterUnregister, RemoteDynFilterUpdate, region_request, remote_dyn_filter_request,
};
use common_recordbatch::adapter::RecordBatchMetrics;
use datatypes::arrow::array::Int32Array;
use datatypes::prelude::{ConcreteDataType, VectorRef};
use datatypes::schema::{ColumnSchema, Schema};
use datatypes::vectors::Int32Vector;
@@ -598,6 +599,86 @@ mod test {
assert_eq!(metrics.elapsed_compute, 7);
}
#[tokio::test]
async fn test_record_batch_stream_updates_following_metrics_before_yielding_batch() {
let schema = test_schema();
let batch = RecordBatch::new(
schema.clone(),
vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
)
.unwrap();
let mut recordbatches = recordbatches_from_flight_message_stream(
"test-peer".to_string(),
stream::iter(vec![
Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
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!(recordbatches.next().await.is_none());
}
#[tokio::test]
async fn test_record_batch_stream_preserves_peeked_record_batch() {
let schema = test_schema();
let first_batch = RecordBatch::new(
schema.clone(),
vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
)
.unwrap();
let second_batch = RecordBatch::new(
schema.clone(),
vec![Arc::new(Int32Vector::from_slice([2])) as VectorRef],
)
.unwrap();
let mut recordbatches = recordbatches_from_flight_message_stream(
"test-peer".to_string(),
stream::iter(vec![
Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
Ok(FlightMessage::RecordBatch(
first_batch.into_df_record_batch(),
)),
Ok(FlightMessage::RecordBatch(
second_batch.into_df_record_batch(),
)),
]),
)
.await
.unwrap();
let first_batch = recordbatches.next().await.unwrap().unwrap();
let second_batch = recordbatches.next().await.unwrap().unwrap();
assert_eq!(
first_batch
.column(0)
.as_any()
.downcast_ref::<Int32Array>()
.unwrap()
.value(0),
1
);
assert_eq!(
second_batch
.column(0)
.as_any()
.downcast_ref::<Int32Array>()
.unwrap()
.value(0),
2
);
assert!(recordbatches.next().await.is_none());
}
#[tokio::test]
async fn test_record_batch_stream_exposes_error_after_pre_batch_metrics() {
let schema = test_schema();
+7 -2
View File
@@ -67,7 +67,8 @@ use servers::error::{
};
use servers::grpc::FlightCompression;
use servers::grpc::flight::{
FlightCraft, FlightRecordBatchSource, FlightRecordBatchStream, TonicStream,
FlightCraft, FlightRecordBatchSource, FlightRecordBatchStream, FlightRecordBatchStreamInput,
TonicStream,
};
use servers::grpc::region_server::RegionServerHandler;
use session::context::{
@@ -990,7 +991,11 @@ impl FlightCraft for RegionServer {
};
let stream = Box::pin(FlightRecordBatchStream::new(
FlightRecordBatchSource::initializer(initializer),
FlightRecordBatchStreamInput::initializer(async move {
initializer
.await
.map(FlightRecordBatchSource::RecordBatches)
}),
tracing_context,
self.flight_compression,
query_ctx,
+3
View File
@@ -44,6 +44,8 @@ pub struct BatchingModeOptions {
/// The gRPC connection timeout
#[serde(with = "humantime_serde")]
pub grpc_conn_timeout: Duration,
#[serde(with = "humantime_serde")]
pub experimental_flight_do_get_timeout: Duration,
/// The gRPC max retry number
pub experimental_grpc_max_retries: u32,
/// Flow wait for available frontend timeout,
@@ -72,6 +74,7 @@ impl Default for BatchingModeOptions {
slow_query_threshold: Duration::from_secs(60),
experimental_min_refresh_duration: Duration::new(5, 0),
grpc_conn_timeout: Duration::from_secs(5),
experimental_flight_do_get_timeout: Duration::from_secs(10),
experimental_grpc_max_retries: 3,
experimental_frontend_scan_timeout: Duration::from_secs(30),
experimental_max_filter_num_per_query: 20,
+7 -54
View File
@@ -18,7 +18,6 @@ use std::collections::HashMap;
use std::sync::{Arc, Mutex, RwLock, Weak};
use api::v1::greptime_request::Request;
use api::v1::query_request::Query;
use api::v1::{CreateTableExpr, QueryRequest};
use client::{Client, DEFAULT_CATALOG_NAME, DEFAULT_SCHEMA_NAME, Database, OutputWithMetrics};
use common_error::ext::BoxedError;
@@ -340,53 +339,6 @@ impl FrontendClient {
})
}
/// Execute a SQL statement on the frontend.
pub async fn sql(&self, catalog: &str, schema: &str, sql: &str) -> Result<Output, Error> {
match self {
FrontendClient::Distributed { .. } => {
let db = self.get_random_active_frontend(catalog, schema).await?;
db.database
.sql(sql)
.await
.map_err(BoxedError::new)
.context(ExternalSnafu)
}
FrontendClient::Standalone {
database_client, ..
} => {
let ctx = QueryContextBuilder::default()
.current_catalog(catalog.to_string())
.current_schema(schema.to_string())
.build();
let ctx = Arc::new(ctx);
{
let database_client = {
database_client
.handler
.lock()
.unwrap()
.as_ref()
.context(UnexpectedSnafu {
reason: "Standalone's frontend instance is not set",
})?
.upgrade()
.context(UnexpectedSnafu {
reason: "Failed to upgrade database client",
})?
};
let req = Request::Query(QueryRequest {
query: Some(Query::Sql(sql.to_string())),
});
database_client
.do_query(req, ctx)
.await
.map_err(BoxedError::new)
.context(ExternalSnafu)
}
}
}
}
/// Execute a flow query and return terminal metrics. `snapshot_seqs` are
/// optional read upper bounds used only by snapshot-fenced repair chunks.
pub(crate) async fn query_with_terminal_metrics(
@@ -413,12 +365,12 @@ impl FrontendClient {
peer: db.peer.clone(),
});
db.database
.query_with_terminal_metrics_and_flow_extensions(
request,
&hints,
extensions,
snapshot_seqs,
)
.flight_request()
.with_hints(&hints)
.with_flow_extensions(extensions)
.with_snapshot_seqs(snapshot_seqs)
.with_timeout(batch_opts.experimental_flight_do_get_timeout)
.query_with_terminal_metrics(request)
.await
.map_err(BoxedError::new)
.context(ExternalSnafu)
@@ -623,6 +575,7 @@ mod tests {
use std::task::{Context, Poll};
use std::time::Duration;
use api::v1::query_request::Query;
use arrow_flight::flight_service_server::FlightServiceServer;
use arrow_flight::{FlightData, Ticket};
use common_query::{Output, OutputData};
+40 -35
View File
@@ -30,8 +30,7 @@ use bytes::{self, Bytes};
use common_error::ext::ErrorExt;
use common_grpc::flight::do_put::{DoPutMetadata, DoPutResponse};
use common_grpc::flight::{
FLOW_EXTENSIONS_METADATA_KEY, FlightDecoder, FlightEncoder, FlightMessage,
SNAPSHOT_SEQS_METADATA_KEY,
FLOW_EXTENSIONS_METADATA_KEY, FlightDecoder, FlightMessage, SNAPSHOT_SEQS_METADATA_KEY,
};
use common_memory_manager::MemoryGuard;
use common_query::{Output, OutputData};
@@ -46,14 +45,16 @@ use prost::Message;
use query::metrics::terminal_recordbatch_metrics_from_plan_if_requested;
use query::options::FlowQueryExtensions;
use session::context::{Channel, QueryContextRef};
use snafu::{IntoError, ResultExt, ensure};
use snafu::{IntoError, OptionExt, ResultExt, ensure};
use table::table_name::TableName;
use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream;
use tonic::{Request, Response, Status, Streaming};
use crate::error::{InvalidParameterSnafu, Result, ToJsonSnafu};
pub use crate::grpc::flight::stream::{FlightRecordBatchSource, FlightRecordBatchStream};
use crate::error::{InvalidParameterSnafu, InvalidQuerySnafu, Result, ToJsonSnafu};
pub use crate::grpc::flight::stream::{
FlightRecordBatchSource, FlightRecordBatchStream, FlightRecordBatchStreamInput,
};
use crate::grpc::greptime_handler::{
GreptimeRequestHandler, create_query_context, get_request_type,
};
@@ -221,9 +222,12 @@ impl FlightCraft for GreptimeRequestHandler {
);
let flight_compression = self.flight_compression;
async {
let output = self
.handle_request_with_query_ctx(request, query_ctx.clone())
let query = request.request.context(InvalidQuerySnafu {
reason: "Expecting non-empty GreptimeRequest.",
})?;
self.authenticate_request_with_query_ctx(request.header.as_ref(), &query_ctx)
.await?;
let output = self.handle_request_with_query_ctx(query, query_ctx.clone());
let stream = to_flight_data_stream(
output,
TracingContext::from_current_span(),
@@ -573,32 +577,36 @@ fn extract_json_metadata<T: serde::de::DeserializeOwned>(
Ok(Some(parsed))
}
fn to_flight_data_stream(
output: Output,
fn to_flight_data_stream<F>(
output: F,
tracing_context: TracingContext,
flight_compression: FlightCompression,
query_ctx: QueryContextRef,
should_emit_terminal_metrics: bool,
) -> TonicStream<FlightData> {
) -> TonicStream<FlightData>
where
F: std::future::Future<Output = Result<Output>> + Send + 'static,
{
let initializer = async move {
let output = output.await.map_err(Status::from)?;
output_to_flight_record_batch_source(output, should_emit_terminal_metrics)
};
let stream = FlightRecordBatchStream::new(
FlightRecordBatchStreamInput::initializer(initializer),
tracing_context,
flight_compression,
query_ctx,
);
Box::pin(stream) as _
}
fn output_to_flight_record_batch_source(
output: Output,
should_emit_terminal_metrics: bool,
) -> TonicResult<FlightRecordBatchSource> {
match output.data {
OutputData::Stream(stream) => {
let stream = FlightRecordBatchStream::new(
FlightRecordBatchSource::RecordBatches(stream),
tracing_context,
flight_compression,
query_ctx,
);
Box::pin(stream) as _
}
OutputData::RecordBatches(x) => {
let stream = FlightRecordBatchStream::new(
FlightRecordBatchSource::RecordBatches(x.as_stream()),
tracing_context,
flight_compression,
query_ctx,
);
Box::pin(stream) as _
}
OutputData::Stream(stream) => Ok(FlightRecordBatchSource::RecordBatches(stream)),
OutputData::RecordBatches(x) => Ok(FlightRecordBatchSource::RecordBatches(x.as_stream())),
OutputData::AffectedRows(rows) => {
let terminal_metrics = match terminal_recordbatch_metrics_from_plan_if_requested(
output.meta.plan,
@@ -607,20 +615,17 @@ fn to_flight_data_stream(
Some(metrics) => match serde_json::to_string(&metrics) {
Ok(metrics) => Some(metrics),
Err(e) => {
let stream = tokio_stream::once(Err(Status::internal(format!(
return Err(Status::internal(format!(
"Failed to serialize terminal metrics: {e}"
))));
return Box::pin(stream) as _;
)));
}
},
None => None,
};
let affected_rows = FlightEncoder::default().encode(FlightMessage::AffectedRows {
Ok(FlightRecordBatchSource::AffectedRows {
rows,
metrics: terminal_metrics,
});
let stream = tokio_stream::iter(affected_rows.into_iter().map(Ok));
Box::pin(stream) as _
})
}
}
}
+106 -44
View File
@@ -41,27 +41,34 @@ use crate::error;
use crate::grpc::FlightCompression;
use crate::grpc::flight::TonicResult;
pub type FlightRecordBatchStreamInitializer = Pin<
Box<
dyn Future<Output = std::result::Result<SendableRecordBatchStream, tonic::Status>>
+ Send
+ 'static,
>,
>;
pub enum FlightRecordBatchSource {
RecordBatches(SendableRecordBatchStream),
Initializer(FlightRecordBatchStreamInitializer),
AffectedRows {
rows: usize,
metrics: Option<String>,
},
}
impl FlightRecordBatchSource {
pub fn initializer<F>(initializer: F) -> Self
where
F: Future<Output = std::result::Result<SendableRecordBatchStream, tonic::Status>>
+ Send
+ 'static,
{
Self::Initializer(Box::pin(initializer))
/// Determines whether a Flight result is ready now or initialized asynchronously.
pub enum FlightRecordBatchStreamInput<F = std::future::Ready<TonicResult<FlightRecordBatchSource>>>
{
Ready(FlightRecordBatchSource),
Initializer(F),
}
impl FlightRecordBatchStreamInput {
/// Creates an input from a source that is already available.
pub fn ready(source: FlightRecordBatchSource) -> Self {
Self::Ready(source)
}
}
impl<F> FlightRecordBatchStreamInput<F> {
/// Creates an input that obtains its source asynchronously.
///
/// Errors from the initializer are returned through the Flight response stream.
pub fn initializer(initializer: F) -> Self {
Self::Initializer(initializer)
}
}
@@ -160,31 +167,32 @@ impl FlightRecordBatchStream {
Self::send_metrics(tx, metrics, metrics_str).await
}
pub fn new(
source: FlightRecordBatchSource,
pub fn new<F>(
input: FlightRecordBatchStreamInput<F>,
tracing_context: TracingContext,
compression: FlightCompression,
query_ctx: QueryContextRef,
) -> Self {
let should_send_partial_metrics = query_ctx.explain_verbose();
let can_send_metrics_before_batch = query_ctx.explain_verbose()
&& query_ctx.live_analyze_metrics_enabled()
&& query_ctx
.remote_query_id()
.zip(query_ctx.extension(SUPPORT_FLIGHT_METRICS_BEFORE_BATCH_EXTENSION_KEY))
.is_some_and(|(remote_query_id, capability)| capability == remote_query_id);
) -> Self
where
F: Future<Output = TonicResult<FlightRecordBatchSource>> + Send + 'static,
{
let (mut tx, rx) = mpsc::channel::<TonicResult<FlightMessage>>(1);
let source_type = match &source {
FlightRecordBatchSource::RecordBatches(_) => "record_batches",
FlightRecordBatchSource::Initializer(_) => "initializer",
let source_type = match &input {
FlightRecordBatchStreamInput::Ready(FlightRecordBatchSource::RecordBatches(_)) => {
"record_batches"
}
FlightRecordBatchStreamInput::Ready(FlightRecordBatchSource::AffectedRows {
..
}) => "affected_rows",
FlightRecordBatchStreamInput::Initializer(_) => "initializer",
};
let initializer_tracing_context = tracing_context.clone();
let join_handle = common_runtime::spawn_global(
async move {
let recordbatches = async move {
match source {
FlightRecordBatchSource::RecordBatches(recordbatches) => Ok(recordbatches),
FlightRecordBatchSource::Initializer(initializer) => initializer.await,
let source = async move {
match input {
FlightRecordBatchStreamInput::Ready(source) => Ok(source),
FlightRecordBatchStreamInput::Initializer(initializer) => initializer.await,
}
}
.trace(
@@ -193,8 +201,20 @@ impl FlightRecordBatchStream {
)
.await;
match recordbatches {
Ok(recordbatches) => {
match source {
Ok(FlightRecordBatchSource::RecordBatches(recordbatches)) => {
let should_send_partial_metrics = query_ctx.explain_verbose();
let can_send_metrics_before_batch =
query_ctx.explain_verbose()
&& query_ctx.live_analyze_metrics_enabled()
&& query_ctx
.remote_query_id()
.zip(query_ctx.extension(
SUPPORT_FLIGHT_METRICS_BEFORE_BATCH_EXTENSION_KEY,
))
.is_some_and(|(remote_query_id, capability)| {
capability == remote_query_id
});
Self::flight_data_stream(
recordbatches,
tx,
@@ -203,6 +223,11 @@ impl FlightRecordBatchStream {
)
.await;
}
Ok(FlightRecordBatchSource::AffectedRows { rows, metrics }) => {
let _ = tx
.send(Ok(FlightMessage::AffectedRows { rows, metrics }))
.await;
}
Err(status) => {
let _ = tx.send(Err(status)).await;
}
@@ -485,7 +510,9 @@ mod test {
.unwrap()
.as_stream();
let mut stream = FlightRecordBatchStream::new(
FlightRecordBatchSource::RecordBatches(recordbatches),
FlightRecordBatchStreamInput::ready(FlightRecordBatchSource::RecordBatches(
recordbatches,
)),
TracingContext::default(),
FlightCompression::default(),
QueryContext::arc(),
@@ -519,10 +546,34 @@ mod test {
}
}
#[tokio::test]
async fn test_flight_record_batch_stream_encodes_affected_rows() {
let mut stream = FlightRecordBatchStream::new(
FlightRecordBatchStreamInput::ready(FlightRecordBatchSource::AffectedRows {
rows: 42,
metrics: Some(r#"{"region_watermarks":[]}"#.to_string()),
}),
TracingContext::default(),
FlightCompression::default(),
QueryContext::arc(),
);
let data = stream.next().await.unwrap().unwrap();
let message = FlightDecoder::default().try_decode(&data).unwrap().unwrap();
assert!(matches!(
message,
FlightMessage::AffectedRows {
rows: 42,
metrics: Some(_),
}
));
assert!(stream.next().await.is_none());
}
#[tokio::test]
async fn test_flight_record_batch_stream_forwards_initializer_error() {
let mut stream = FlightRecordBatchStream::new(
FlightRecordBatchSource::initializer(async {
FlightRecordBatchStreamInput::initializer(async {
Err(tonic::Status::unavailable(
"remote read initialization failed",
))
@@ -552,9 +603,12 @@ mod test {
metrics,
});
let query_ctx = query_context_with_live_metrics_and_matching_capability();
query_ctx.set_explain_verbose(true);
let initializer_query_ctx = query_ctx.clone();
let mut stream = FlightRecordBatchStream::new(
FlightRecordBatchSource::RecordBatches(recordbatches),
FlightRecordBatchStreamInput::initializer(async move {
initializer_query_ctx.set_explain_verbose(true);
Ok(FlightRecordBatchSource::RecordBatches(recordbatches))
}),
TracingContext::default(),
FlightCompression::default(),
query_ctx,
@@ -609,7 +663,9 @@ mod test {
let query_ctx = query_context_with_live_metrics_and_matching_capability();
query_ctx.set_explain_verbose(true);
let mut stream = FlightRecordBatchStream::new(
FlightRecordBatchSource::RecordBatches(recordbatches),
FlightRecordBatchStreamInput::ready(FlightRecordBatchSource::RecordBatches(
recordbatches,
)),
TracingContext::default(),
FlightCompression::default(),
query_ctx,
@@ -663,7 +719,9 @@ mod test {
let query_ctx = query_context_with_matching_capability();
query_ctx.set_explain_verbose(true);
let mut stream = FlightRecordBatchStream::new(
FlightRecordBatchSource::RecordBatches(recordbatches),
FlightRecordBatchStreamInput::ready(FlightRecordBatchSource::RecordBatches(
recordbatches,
)),
TracingContext::default(),
FlightCompression::default(),
query_ctx,
@@ -706,7 +764,9 @@ mod test {
let query_ctx = Arc::new(query_ctx);
query_ctx.set_explain_verbose(true);
let mut stream = FlightRecordBatchStream::new(
FlightRecordBatchSource::RecordBatches(recordbatches),
FlightRecordBatchStreamInput::ready(FlightRecordBatchSource::RecordBatches(
recordbatches,
)),
TracingContext::default(),
FlightCompression::default(),
query_ctx,
@@ -745,7 +805,9 @@ mod test {
});
let query_ctx = query_context_with_matching_capability();
let mut stream = FlightRecordBatchStream::new(
FlightRecordBatchSource::RecordBatches(recordbatches),
FlightRecordBatchStreamInput::ready(FlightRecordBatchSource::RecordBatches(
recordbatches,
)),
TracingContext::default(),
FlightCompression::default(),
query_ctx,
+42 -30
View File
@@ -15,11 +15,13 @@
//! Handler for Greptime Database service. It's implemented by frontend.
use std::collections::HashMap;
use std::future::Future;
use std::str::FromStr;
use std::sync::{Arc, RwLock};
use std::time::Instant;
use api::helper::request_type;
use api::v1::greptime_request::Request as QueryRequest;
use api::v1::{GreptimeRequest, RequestHeader};
use auth::UserProviderRef;
use common_catalog::consts::{DEFAULT_CATALOG_NAME, DEFAULT_SCHEMA_NAME};
@@ -79,30 +81,39 @@ impl GreptimeRequestHandler {
) -> Result<Output> {
let header = request.header.as_ref();
let query_ctx = create_query_context(Channel::Grpc, header, hints, HashMap::new())?;
self.handle_request_with_query_ctx(request, query_ctx).await
}
pub(crate) async fn handle_request_with_query_ctx(
&self,
request: GreptimeRequest,
query_ctx: QueryContextRef,
) -> Result<Output> {
let query = request.request.context(InvalidQuerySnafu {
reason: "Expecting non-empty GreptimeRequest.",
})?;
self.authenticate_request_with_query_ctx(request.header.as_ref(), &query_ctx)
.await?;
self.handle_request_with_query_ctx(query, query_ctx).await
}
let header = request.header.as_ref();
let user_info = context_auth::auth(self.user_provider.clone(), header, &query_ctx).await?;
pub(crate) async fn authenticate_request_with_query_ctx(
&self,
header: Option<&RequestHeader>,
query_ctx: &QueryContextRef,
) -> Result<()> {
let user_info = context_auth::auth(self.user_provider.clone(), header, query_ctx).await?;
query_ctx.set_current_user(user_info);
Ok(())
}
pub(crate) fn handle_request_with_query_ctx(
&self,
query: QueryRequest,
query_ctx: QueryContextRef,
) -> impl Future<Output = Result<Output>> + Send + 'static {
let handler = self.handler.clone();
let runtime = self.runtime.clone();
let request_type = request_type(&query).to_string();
let db = query_ctx.get_db_string();
let timer = RequestTimer::new(db.clone(), request_type);
let tracing_context = TracingContext::from_current_span();
let result_future = async move {
handler
async move {
let result_future = async move {
handler
.do_query(query, query_ctx)
.trace(tracing_context.attach(tracing::info_span!(
"GreptimeRequestHandler::handle_request_runtime"
@@ -118,26 +129,27 @@ impl GreptimeRequestHandler {
}
e
})
};
};
match &self.runtime {
Some(runtime) => {
// Executes requests in another runtime to
// 1. prevent the execution from being cancelled unexpected by Tonic runtime;
// - Refer to our blog for the rational behind it:
// https://www.greptime.com/blogs/2023-01-12-hidden-control-flow.html
// - Obtaining a `JoinHandle` to get the panic message (if there's any).
// From its docs, `JoinHandle` is cancel safe. The task keeps running even it's handle been dropped.
// 2. avoid the handler blocks the gRPC runtime incidentally.
runtime
.spawn(result_future)
.await
.context(JoinTaskSnafu)
.inspect_err(|e| {
timer.record(e.status_code());
})?
match runtime {
Some(runtime) => {
// Executes requests in another runtime to
// 1. prevent the execution from being cancelled unexpected by Tonic runtime;
// - Refer to our blog for the rational behind it:
// https://www.greptime.com/blogs/2023-01-12-hidden-control-flow.html
// - Obtaining a `JoinHandle` to get the panic message (if there's any).
// From its docs, `JoinHandle` is cancel safe. The task keeps running even it's handle been dropped.
// 2. avoid the handler blocks the gRPC runtime incidentally.
runtime
.spawn(result_future)
.await
.context(JoinTaskSnafu)
.inspect_err(|e| {
timer.record(e.status_code());
})?
}
None => result_future.await,
}
None => result_future.await,
}
}
+156 -24
View File
@@ -16,40 +16,45 @@
mod test {
use std::collections::HashMap;
use std::net::SocketAddr;
use std::pin::Pin;
use std::sync::Arc;
use std::time::{Duration, Instant};
use api::v1::auth_header::AuthScheme;
use api::v1::greptime_request::Request as GreptimeQueryRequest;
use api::v1::query_request::Query;
use api::v1::{Basic, ColumnDataType, ColumnDef, CreateTableExpr, QueryRequest, SemanticType};
use arrow_flight::flight_service_server::FlightServiceServer;
use arrow_flight::{FlightData, FlightDescriptor, Ticket};
use auth::user_provider_from_option;
use client::region::RegionRequester;
use client::{Client, Database};
use common_catalog::consts::{DEFAULT_CATALOG_NAME, DEFAULT_SCHEMA_NAME};
use common_grpc::channel_manager::{ChannelConfig, ChannelManager};
use common_grpc::flight::do_put::DoPutMetadata;
use common_grpc::flight::do_put::{DoPutMetadata, DoPutResponse};
use common_grpc::flight::{FlightEncoder, FlightMessage};
use common_query::OutputData;
use common_query::{Output, OutputData};
use common_recordbatch::adapter::RegionWatermarkEntry;
use common_recordbatch::{RecordBatch, RecordBatches, SendableRecordBatchStream};
use common_telemetry::tracing_context::TracingContext;
use datatypes::prelude::{ConcreteDataType, ScalarVector, VectorRef};
use datatypes::schema::{ColumnSchema, Schema};
use datatypes::vectors::{Int32Vector, StringVector, TimestampMillisecondVector};
use futures_util::StreamExt;
use futures_util::{Stream, StreamExt};
use hyper_util::rt::TokioIo;
use itertools::Itertools;
use servers::grpc::builder::GrpcServerBuilder;
use servers::grpc::flight::{
FlightCraft, FlightCraftWrapper, FlightRecordBatchSource, FlightRecordBatchStream,
TonicStream,
FlightRecordBatchStreamInput, PutRecordBatchRequestStream, TonicStream,
};
use servers::grpc::greptime_handler::GreptimeRequestHandler;
use servers::grpc::{FlightCompression, GrpcServerConfig};
use servers::query_handler::grpc::GrpcQueryHandler;
use servers::server::Server;
use tonic::Response;
use session::context::QueryContextRef;
use tonic::transport::Server as TonicServer;
use tonic::{Response, Status};
use tower::service_fn;
use crate::cluster::GreptimeDbClusterBuilder;
@@ -59,6 +64,12 @@ mod test {
struct SlowFlightCraft;
struct ErrorFlightCraft;
struct SlowRemoteQueryHandler {
region_requester: RegionRequester,
}
fn slow_recordbatch_stream() -> SendableRecordBatchStream {
let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
"value",
@@ -83,9 +94,11 @@ mod test {
_: tonic::Request<Ticket>,
) -> std::result::Result<Response<TonicStream<FlightData>>, tonic::Status> {
let stream = FlightRecordBatchStream::new(
FlightRecordBatchSource::initializer(async {
FlightRecordBatchStreamInput::initializer(async {
tokio::time::sleep(Duration::from_secs(2)).await;
Ok(slow_recordbatch_stream())
Ok(FlightRecordBatchSource::RecordBatches(
slow_recordbatch_stream(),
))
}),
TracingContext::default(),
FlightCompression::default(),
@@ -96,14 +109,72 @@ mod test {
}
}
#[tokio::test(flavor = "multi_thread")]
async fn test_do_get_timeout_does_not_cancel_slow_flight_stream() {
#[async_trait::async_trait]
impl FlightCraft for ErrorFlightCraft {
async fn do_get(
&self,
_: tonic::Request<Ticket>,
) -> std::result::Result<Response<TonicStream<FlightData>>, tonic::Status> {
let stream = FlightRecordBatchStream::new(
FlightRecordBatchStreamInput::initializer(async {
Err(Status::internal("deferred initializer detail"))
}),
TracingContext::default(),
FlightCompression::default(),
session::context::QueryContext::arc(),
);
Ok(Response::new(Box::pin(stream)))
}
}
#[async_trait::async_trait]
impl GrpcQueryHandler for SlowRemoteQueryHandler {
async fn do_query(
&self,
_: GreptimeQueryRequest,
_: QueryContextRef,
) -> servers::error::Result<Output> {
let stream = self
.region_requester
.do_get_inner(Ticket::default())
.await
.unwrap();
Ok(Output::new_with_stream(stream))
}
fn handle_put_record_batch_stream(
&self,
_: PutRecordBatchRequestStream,
_: QueryContextRef,
) -> Pin<Box<dyn Stream<Item = servers::error::Result<DoPutResponse>> + Send>> {
Box::pin(futures::stream::empty())
}
}
fn client_for_flight_craft<T>(addr: &'static str, craft: T) -> Client
where
T: FlightCraft,
{
client_for_flight_craft_with_max_encoding(addr, craft, None)
}
fn client_for_flight_craft_with_max_encoding<T>(
addr: &'static str,
craft: T,
max_encoding_message_size: Option<usize>,
) -> Client
where
T: FlightCraft,
{
let (client_io, server_io) = tokio::io::duplex(1024);
tokio::spawn(async move {
let flight_service = FlightServiceServer::new(FlightCraftWrapper(craft));
let flight_service = match max_encoding_message_size {
Some(size) => flight_service.max_encoding_message_size(size),
None => flight_service,
};
TonicServer::builder()
.add_service(FlightServiceServer::new(FlightCraftWrapper(
SlowFlightCraft,
)))
.add_service(flight_service)
.serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>(
server_io,
)]))
@@ -115,7 +186,7 @@ mod test {
let mut client_io = Some(client_io);
channel_manager
.reset_with_connector(
"slow-flight",
addr,
service_fn(move |_| {
let client_io = client_io.take();
@@ -127,7 +198,12 @@ mod test {
}),
)
.unwrap();
let client = Client::with_manager_and_urls(channel_manager, ["slow-flight"]);
Client::with_manager_and_urls(channel_manager, [addr])
}
#[tokio::test(flavor = "multi_thread")]
async fn test_do_get_timeout_does_not_cancel_slow_flight_stream() {
let client = client_for_flight_craft("slow-flight", SlowFlightCraft);
let mut flight_client = client.make_flight_client(false, false).unwrap();
let start = Instant::now();
@@ -143,6 +219,64 @@ mod test {
assert!(stream.message().await.unwrap().is_some());
}
#[tokio::test(flavor = "multi_thread")]
async fn test_deferred_flight_initializer_error_preserves_message() {
let client = client_for_flight_craft("error-flight", ErrorFlightCraft);
let mut flight_client = client.make_flight_client(false, false).unwrap();
let mut stream = flight_client
.mut_inner()
.do_get(tonic::Request::new(Ticket::default()))
.await
.unwrap()
.into_inner();
let error = stream.message().await.unwrap_err();
assert_eq!(tonic::Code::Internal, error.code());
assert_eq!("deferred initializer detail", error.message());
}
#[tokio::test(flavor = "multi_thread")]
async fn test_deferred_flight_encode_error_preserves_message() {
let client =
client_for_flight_craft_with_max_encoding("limited-flight", SlowFlightCraft, Some(1));
let mut flight_client = client.make_flight_client(false, false).unwrap();
let mut stream = flight_client
.mut_inner()
.do_get(tonic::Request::new(Ticket::default()))
.await
.unwrap()
.into_inner();
let error = stream.message().await.unwrap_err();
assert_eq!(tonic::Code::OutOfRange, error.code());
assert!(error.message().contains("message length too large"));
}
#[tokio::test(flavor = "multi_thread")]
async fn test_flight_request_timeout_does_not_cancel_slow_datanode_stream() {
let datanode_client = client_for_flight_craft("slow-datanode", SlowFlightCraft);
let frontend_handler = GreptimeRequestHandler::new(
Arc::new(SlowRemoteQueryHandler {
region_requester: RegionRequester::new(datanode_client, false, false),
}),
None,
None,
FlightCompression::default(),
);
let frontend_client = client_for_flight_craft("slow-frontend", frontend_handler);
let database = Database::new(DEFAULT_CATALOG_NAME, DEFAULT_SCHEMA_NAME, frontend_client);
let start = Instant::now();
let output = database
.flight_request()
.with_timeout(Duration::from_secs(1))
.sql("select 1")
.await
.unwrap();
// `sql()` waits for the first Flight message. Succeeding after this delay
// proves the request timeout applied only to the response header.
assert!(start.elapsed() >= Duration::from_secs(1));
assert!(matches!(output.data, OutputData::Stream(_)));
}
#[tokio::test(flavor = "multi_thread")]
async fn test_standalone_flight_do_put() {
common_telemetry::init_default_ut_logging();
@@ -446,16 +580,14 @@ mod test {
);
let result = client
.query_with_terminal_metrics_and_flow_extensions(
QueryRequest {
query: Some(Query::Sql(
"select ts, a, `B` from foo order by ts".to_string(),
)),
},
&[],
&[("flow.return_region_seq", "true")],
&HashMap::from([(region_id.as_u64(), stale_snapshot_seq)]),
)
.flight_request()
.with_flow_extensions(&[("flow.return_region_seq", "true")])
.with_snapshot_seqs(&HashMap::from([(region_id.as_u64(), stale_snapshot_seq)]))
.query_with_terminal_metrics(QueryRequest {
query: Some(Query::Sql(
"select ts, a, `B` from foo order by ts".to_string(),
)),
})
.await
.unwrap();
+1
View File
@@ -2252,6 +2252,7 @@ query_timeout = "10m"
slow_query_threshold = "1m"
experimental_min_refresh_duration = "5s"
grpc_conn_timeout = "5s"
experimental_flight_do_get_timeout = "10s"
experimental_grpc_max_retries = 3
experimental_frontend_scan_timeout = "30s"
experimental_max_filter_num_per_query = 20