use std::{ pin::Pin, sync::Arc, task::{Context, Poll}, }; use api::{ query::{ ast::ExtStatement, dispatcher::QueryDispatcher, execution::{Output, QueryStateMachine}, function::FuncMetaManagerRef, logical_planner::{LogicalPlanner, Plan}, parser::Parser, session::{SessionCtx, SessionCtxFactory}, Query, }, QueryError, QueryResult, }; use async_trait::async_trait; use datafusion::{ arrow::{datatypes::SchemaRef, record_batch::RecordBatch}, config::CsvOptions, datasource::{ file_format::{csv::CsvFormat, json::JsonFormat, parquet::ParquetFormat}, listing::{ListingOptions, ListingTable, ListingTableConfig, ListingTableUrl}, }, error::Result as DFResult, execution::{RecordBatchStream, SendableRecordBatchStream}, }; use futures::{Stream, StreamExt}; use s3s::dto::SelectObjectContentInput; use crate::{ execution::factory::QueryExecutionFactoryRef, metadata::{base_table::BaseTableProvider, ContextProviderExtension, MetadataProvider, TableHandleProviderRef}, sql::logical::planner::DefaultLogicalPlanner, }; #[derive(Clone)] pub struct SimpleQueryDispatcher { input: SelectObjectContentInput, // client for default tenant _default_table_provider: TableHandleProviderRef, session_factory: Arc, // parser parser: Arc, // get query execution factory query_execution_factory: QueryExecutionFactoryRef, func_manager: FuncMetaManagerRef, } #[async_trait] impl QueryDispatcher for SimpleQueryDispatcher { async fn execute_query(&self, query: &Query) -> QueryResult { let query_state_machine = { self.build_query_state_machine(query.clone()).await? }; let logical_plan = self.build_logical_plan(query_state_machine.clone()).await?; let logical_plan = match logical_plan { Some(plan) => plan, None => return Ok(Output::Nil(())), }; let result = self.execute_logical_plan(logical_plan, query_state_machine).await?; Ok(result) } async fn build_logical_plan(&self, query_state_machine: Arc) -> QueryResult> { let session = &query_state_machine.session; let query = &query_state_machine.query; let scheme_provider = self.build_scheme_provider(session).await?; let logical_planner = DefaultLogicalPlanner::new(&scheme_provider); let statements = self.parser.parse(query.content())?; // not allow multi statement if statements.len() > 1 { return Err(QueryError::MultiStatement { num: statements.len(), sql: query_state_machine.query.content().to_string(), }); } let stmt = match statements.front() { Some(stmt) => stmt.clone(), None => return Ok(None), }; let logical_plan = self .statement_to_logical_plan(stmt, &logical_planner, query_state_machine) .await?; Ok(Some(logical_plan)) } async fn execute_logical_plan(&self, logical_plan: Plan, query_state_machine: Arc) -> QueryResult { self.execute_logical_plan(logical_plan, query_state_machine).await } async fn build_query_state_machine(&self, query: Query) -> QueryResult> { let session = self.session_factory.create_session_ctx(query.context()).await?; let query_state_machine = Arc::new(QueryStateMachine::begin(query, session)); Ok(query_state_machine) } } impl SimpleQueryDispatcher { async fn statement_to_logical_plan( &self, stmt: ExtStatement, logical_planner: &DefaultLogicalPlanner<'_, S>, query_state_machine: Arc, ) -> QueryResult { // begin analyze query_state_machine.begin_analyze(); let logical_plan = logical_planner .create_logical_plan(stmt, &query_state_machine.session) .await?; query_state_machine.end_analyze(); Ok(logical_plan) } async fn execute_logical_plan(&self, logical_plan: Plan, query_state_machine: Arc) -> QueryResult { let execution = self .query_execution_factory .create_query_execution(logical_plan, query_state_machine.clone()) .await?; match execution.start().await { Ok(Output::StreamData(stream)) => Ok(Output::StreamData(Box::pin(TrackedRecordBatchStream { inner: stream }))), Ok(nil @ Output::Nil(_)) => Ok(nil), Err(err) => Err(err), } } async fn build_scheme_provider(&self, session: &SessionCtx) -> QueryResult { let path = format!("s3://{}/{}", self.input.bucket, self.input.key); let table_path = ListingTableUrl::parse(path)?; let listing_options = if self.input.request.input_serialization.csv.is_some() { let file_format = CsvFormat::default().with_options(CsvOptions::default().with_has_header(true)); ListingOptions::new(Arc::new(file_format)).with_file_extension(".csv") } else if self.input.request.input_serialization.parquet.is_some() { let file_format = ParquetFormat::new(); ListingOptions::new(Arc::new(file_format)).with_file_extension(".parquet") } else if self.input.request.input_serialization.json.is_some() { let file_format = JsonFormat::default(); ListingOptions::new(Arc::new(file_format)).with_file_extension(".json") } else { return Err(QueryError::NotImplemented { err: "not support this file type".to_string(), }); }; let resolve_schema = listing_options.infer_schema(session.inner(), &table_path).await?; let config = ListingTableConfig::new(table_path) .with_listing_options(listing_options) .with_schema(resolve_schema); let provider = Arc::new(ListingTable::try_new(config)?); let current_session_table_provider = self.build_table_handle_provider()?; let metadata_provider = MetadataProvider::new(provider, current_session_table_provider, self.func_manager.clone(), session.clone()); Ok(metadata_provider) } fn build_table_handle_provider(&self) -> QueryResult { let current_session_table_provider: Arc = Arc::new(BaseTableProvider::default()); Ok(current_session_table_provider) } } pub struct TrackedRecordBatchStream { inner: SendableRecordBatchStream, } impl RecordBatchStream for TrackedRecordBatchStream { fn schema(&self) -> SchemaRef { self.inner.schema() } } impl Stream for TrackedRecordBatchStream { type Item = DFResult; fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { self.inner.poll_next_unpin(cx) } } #[derive(Default, Clone)] pub struct SimpleQueryDispatcherBuilder { input: Option, default_table_provider: Option, session_factory: Option>, parser: Option>, query_execution_factory: Option, func_manager: Option, } impl SimpleQueryDispatcherBuilder { pub fn with_input(mut self, input: SelectObjectContentInput) -> Self { self.input = Some(input); self } pub fn with_default_table_provider(mut self, default_table_provider: TableHandleProviderRef) -> Self { self.default_table_provider = Some(default_table_provider); self } pub fn with_session_factory(mut self, session_factory: Arc) -> Self { self.session_factory = Some(session_factory); self } pub fn with_parser(mut self, parser: Arc) -> Self { self.parser = Some(parser); self } pub fn with_query_execution_factory(mut self, query_execution_factory: QueryExecutionFactoryRef) -> Self { self.query_execution_factory = Some(query_execution_factory); self } pub fn with_func_manager(mut self, func_manager: FuncMetaManagerRef) -> Self { self.func_manager = Some(func_manager); self } pub fn build(self) -> QueryResult> { let input = self.input.ok_or_else(|| QueryError::BuildQueryDispatcher { err: "lost of input".to_string(), })?; let session_factory = self.session_factory.ok_or_else(|| QueryError::BuildQueryDispatcher { err: "lost of session_factory".to_string(), })?; let parser = self.parser.ok_or_else(|| QueryError::BuildQueryDispatcher { err: "lost of parser".to_string(), })?; let query_execution_factory = self.query_execution_factory.ok_or_else(|| QueryError::BuildQueryDispatcher { err: "lost of query_execution_factory".to_string(), })?; let func_manager = self.func_manager.ok_or_else(|| QueryError::BuildQueryDispatcher { err: "lost of func_manager".to_string(), })?; let default_table_provider = self.default_table_provider.ok_or_else(|| QueryError::BuildQueryDispatcher { err: "lost of default_table_provider".to_string(), })?; let dispatcher = Arc::new(SimpleQueryDispatcher { input, _default_table_provider: default_table_provider, session_factory, parser, query_execution_factory, func_manager, }); Ok(dispatcher) } }