andygrove commented on code in PR #2416: URL: https://github.com/apache/datafusion-ballista/pull/2416#discussion_r4116283899
########## ballista/flight-sql/src/service.rs: ########## @@ -0,0 +1,852 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! The Arrow Flight SQL frontend itself. + +use std::pin::Pin; +use std::sync::Arc; +use std::time::Duration; + +use arrow::array::RecordBatch; +use arrow::datatypes::{Schema, SchemaRef}; +use arrow::ipc::writer::IpcWriteOptions; +use arrow_flight::encode::FlightDataEncoderBuilder; +use arrow_flight::flight_service_server::FlightService; +use arrow_flight::sql::server::FlightSqlService; +use arrow_flight::sql::{ + ActionCancelQueryRequest, ActionCancelQueryResult, + ActionClosePreparedStatementRequest, ActionCreatePreparedStatementRequest, + ActionCreatePreparedStatementResult, Any, CommandGetCatalogs, CommandGetDbSchemas, + CommandGetSqlInfo, CommandGetTableTypes, CommandGetTables, CommandGetXdbcTypeInfo, + CommandPreparedStatementQuery, CommandStatementQuery, CommandStatementUpdate, + ProstMessageExt, SqlInfo, TicketStatementQuery, + metadata::{SqlInfoData, XdbcTypeInfoData}, + server::PeekableFlightDataStream, +}; +use arrow_flight::{ + Action, FlightData, FlightDescriptor, FlightEndpoint, FlightInfo, HandshakeRequest, + HandshakeResponse, IpcMessage, SchemaAsIpc, Ticket, +}; +use ballista_core::error::BallistaError; +use ballista_core::flight_proxy_service::BallistaFlightProxyService; +use ballista_core::planner::scans_only_local_tables; +use ballista_core::serde::protobuf::PartitionLocation; +use ballista_core::serde::scheduler::{Action as BallistaAction, ShuffleFileKind}; +use ballista_core::serde::{decode_protobuf, protobuf}; +use datafusion::logical_expr::{DdlStatement, LogicalPlan}; +use datafusion::prelude::SessionContext; +use futures::{Stream, TryStreamExt}; +use prost::Message; +use tonic::metadata::MetadataMap; +use tonic::{Request, Response, Status, Streaming}; +use uuid::Uuid; + +use crate::auth::{AnonymousAuthenticator, Authenticator}; +use crate::backend::QueryBackend; +use crate::metadata; +use crate::session::{LocalResult, Prepared, SessionStore}; +use crate::ticket::StatementHandle; + +/// Session shared by every client that connects without authenticating. +/// +/// Anonymous clients cannot be told apart, so they necessarily share catalog +/// state. Configure an [`Authenticator`] to get a session per connection. +pub const ANONYMOUS_SESSION: &str = "flight-sql-anonymous"; + +/// How long a session, prepared statement, or unredeemed local result may sit +/// idle before it is discarded. +const DEFAULT_TTL: Duration = Duration::from_secs(30 * 60); + +/// How often expired handles are swept. +const REAP_INTERVAL: Duration = Duration::from_secs(60); + +type DoGetStream = + Pin<Box<dyn Stream<Item = Result<FlightData, Status>> + Send + 'static>>; + +/// Serves Arrow Flight SQL on behalf of a Ballista cluster. +/// +/// Clients send SQL text; the frontend plans it against the session's catalog, +/// submits the plan through a [`QueryBackend`], and hands back one +/// `FlightEndpoint` per output partition. `DoGet` on those tickets is proxied +/// to the executor holding the partition, so clients never need to reach +/// executors themselves — the failure mode that made the pre-46.0.0 +/// implementation unusable behind NAT, Docker, and Kubernetes. +pub struct BallistaFlightSqlService<B: QueryBackend> { + backend: Arc<B>, + proxy: BallistaFlightProxyService, + auth: Arc<dyn Authenticator>, + store: Arc<SessionStore>, + sql_info: SqlInfoData, + xdbc_info: XdbcTypeInfoData, +} + +impl<B: QueryBackend> BallistaFlightSqlService<B> { + /// Builds a frontend over `backend`, using `proxy` to stream partition + /// data back from executors. + /// + /// The service authenticates nobody until an [`Authenticator`] is supplied + /// via [`with_authenticator`](Self::with_authenticator). + pub fn new(backend: Arc<B>, proxy: BallistaFlightProxyService) -> Self { + let store = Arc::new(SessionStore::new(DEFAULT_TTL)); + + let backend_for_reaper = backend.clone(); + store.spawn_reaper(REAP_INTERVAL, move |session_id| { + let backend = backend_for_reaper.clone(); + async move { + if let Err(e) = backend.close_session(&session_id).await { + log::warn!("flight-sql: failed to close session {session_id}: {e}"); + } + } + }); + + Self { + backend, + proxy, + auth: Arc::new(AnonymousAuthenticator), + store, + sql_info: metadata::sql_info(), + xdbc_info: metadata::xdbc_type_info(), + } + } + + /// Installs an authenticator. Without one, every handshake is accepted and + /// unauthenticated clients share a single session. + pub fn with_authenticator(mut self, auth: Arc<dyn Authenticator>) -> Self { + self.auth = auth; + self + } + + /// True when the service will accept unauthenticated clients, which the + /// scheduler logs at startup. + pub fn allows_anonymous(&self) -> bool { + self.auth.allows_anonymous() + } + + /// Resolves the Ballista session for a request from its bearer token. + fn session_id(&self, metadata: &MetadataMap) -> Result<String, Status> { + match bearer_token(metadata) { + Some(token) => self.store.session(&token).ok_or_else(|| { + Status::unauthenticated( + "unknown or expired session token; re-run the Flight handshake", + ) + }), + None if self.auth.allows_anonymous() => Ok(ANONYMOUS_SESSION.to_string()), + None => Err(Status::unauthenticated( + "missing bearer token; authenticate with the Flight handshake first", + )), + } + } + + /// Returns the context for `session_id`, building it on first use. + /// + /// The cache is what makes a session a session: [`QueryBackend::session`] + /// builds a fresh `SessionContext` every call, so without it a table + /// created by one request would be invisible to the next — and every + /// request would pay for a full DataFusion session to be constructed. + async fn open_session( + &self, + session_id: &str, + ) -> Result<Arc<SessionContext>, Status> { + if let Some(ctx) = self.store.context(session_id) { + return Ok(ctx); + } + + let ctx = self + .backend + .session(session_id) + .await + .map_err(|e| Status::internal(format!("failed to open session: {e}")))?; + + Ok(self.store.insert_context(session_id.to_string(), ctx)) + } + + /// Resolves a request to the context it should be planned against. + async fn context( + &self, + metadata: &MetadataMap, + ) -> Result<Arc<SessionContext>, Status> { + self.open_session(&self.session_id(metadata)?).await + } + + /// Plans `sql` against the session's catalog. + async fn plan(ctx: &SessionContext, sql: &str) -> Result<LogicalPlan, Status> { + ctx.state() + .create_logical_plan(sql) + .await + .map_err(|e| Status::invalid_argument(format!("failed to plan query: {e}"))) + } + + /// Runs a planned statement and describes where to collect its results. + async fn flight_info_for( + &self, + ctx: Arc<SessionContext>, + plan: LogicalPlan, + descriptor: FlightDescriptor, + job_name: &str, + ) -> Result<FlightInfo, Status> { + if let Disposition::Unsupported(reason) = disposition(&plan) { + return Err(Status::unimplemented(reason)); + } + + if disposition(&plan) == Disposition::RunOnScheduler { + let (schema, batches) = execute_locally(&ctx, plan).await?; + let handle = Uuid::new_v4().to_string(); + self.store.insert_result( + handle.clone(), + LocalResult { + schema: schema.clone(), + batches, + }, + ); + + let endpoint = Self::endpoint(StatementHandle::Local(handle)); + return build_flight_info(&schema, vec![endpoint], descriptor); + } + + let result = self + .backend + .execute(job_name, ctx, plan) + .await + .map_err(query_failed)?; + + let endpoints = result + .partitions + .into_iter() + .map(|location| partition_handle(location).map(Self::endpoint)) + .collect::<Result<Vec<_>, _>>() + .map_err(|e| Status::internal(format!("invalid partition location: {e}")))?; + + log::debug!( + "flight-sql: job {} produced {} endpoint(s)", + result.job_id, + endpoints.len() + ); + + build_flight_info(&result.schema, endpoints, descriptor) + } + + /// Builds the endpoint a client redeems for one slice of the result. + /// + /// Endpoints carry no location, which Flight defines as "fetch from the + /// server that gave you this FlightInfo". That keeps every cluster-internal + /// address off the wire and lets the frontend work unchanged behind NAT, a + /// load balancer, or an ingress. + fn endpoint(handle: StatementHandle) -> FlightEndpoint { + let ticket = TicketStatementQuery { + statement_handle: handle.encode().into(), + }; + FlightEndpoint::new().with_ticket(Ticket { + ticket: ticket.as_any().encode_to_vec().into(), + }) + } + + /// Serves a metadata command by round-tripping the command itself as the + /// ticket, so `DoGet` lands back on the matching handler. + fn metadata_info<C: ProstMessageExt>( + command: C, + schema: SchemaRef, + descriptor: FlightDescriptor, + ) -> Result<Response<FlightInfo>, Status> { + let endpoint = FlightEndpoint::new().with_ticket(Ticket { + ticket: command.as_any().encode_to_vec().into(), + }); + build_flight_info(&schema, vec![endpoint], descriptor).map(Response::new) + } +} + +/// What the frontend should do with a planned statement. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum Disposition { + /// Execute on the scheduler: it mutates session state and produces no data + /// worth distributing. + RunOnScheduler, + /// Submit to the cluster. + Distribute, + /// Refuse, with an explanation for the client. + Unsupported(&'static str), +} + +/// Classifies a plan. +/// +/// One function rather than two predicates, because the interesting cases are +/// the ones where the answers overlap: `CREATE TABLE AS SELECT` is DDL, and +/// DDL runs on the scheduler, so a caller that asked "is this DDL?" before +/// asking "is this supported?" would silently execute its query on one node. +/// Returning a single verdict makes that ordering impossible to get wrong. +fn disposition(plan: &LogicalPlan) -> Disposition { + match plan { + LogicalPlan::Dml(_) => Disposition::Unsupported( + "Ballista Flight SQL does not support INSERT/UPDATE/DELETE; \ + the distributed write path is not implemented", + ), + LogicalPlan::Copy(_) => Disposition::Unsupported( + "Ballista Flight SQL does not support COPY; \ + the distributed write path is not implemented", + ), + LogicalPlan::Ddl(DdlStatement::CreateMemoryTable(_)) => Disposition::Unsupported( + "Ballista Flight SQL does not support CREATE TABLE AS SELECT, \ + because it would execute on the scheduler rather than the cluster; \ + use CREATE EXTERNAL TABLE over data the executors can read", + ), + // Other DDL only edits the session catalog, and `SET`-style statements + // only edit session config; neither has anything to distribute. + LogicalPlan::Ddl(_) | LogicalPlan::Statement(_) => Disposition::RunOnScheduler, + // `SHOW ...` and anything else reading `information_schema` describes + // the catalog held here. There is nothing to distribute, and the + // physical form of those scans cannot be serialized for an executor, + // so distributing one would hand the client a job that never runs. + _ if scans_only_local_tables(plan) => Disposition::RunOnScheduler, + _ => Disposition::Distribute, + } +} + +async fn execute_locally( + ctx: &SessionContext, + plan: LogicalPlan, +) -> Result<(SchemaRef, Vec<RecordBatch>), Status> { + let df = ctx + .execute_logical_plan(plan) + .await + .map_err(|e| Status::internal(format!("failed to execute statement: {e}")))?; + let planned_schema: SchemaRef = Arc::new(df.schema().as_arrow().clone()); + let batches = df + .collect() + .await + .map_err(|e| Status::internal(format!("failed to execute statement: {e}")))?; + + // Prefer the schema the data actually carries; DataFusion's DDL results + // are empty and their frame schema is not always the same object. + let schema = batches + .first() + .map(|batch| batch.schema()) + .unwrap_or(planned_schema); + + Ok((schema, batches)) +} + +/// Turns a shuffle partition into the ticket payload the Flight proxy already +/// knows how to redeem. +fn partition_handle( + location: PartitionLocation, +) -> Result<StatementHandle, BallistaError> { + let layout = location.layout(); + let partition_id = location.partition_id.ok_or_else(|| { + BallistaError::Internal("partition location has no partition id".to_string()) + })?; + let executor = location.executor_meta.ok_or_else(|| { + BallistaError::Internal("partition location has no executor metadata".to_string()) + })?; + + let action = BallistaAction::FetchPartition { + job_id: partition_id.job_id.into(), + stage_id: partition_id.stage_id as usize, + partition_id: partition_id.partition_id as usize, + host: executor.host, + port: executor.port as u16, + file_id: location.file_id, + layout, + // A Flight client wants the whole partition: the data file, in full. + file_kind: ShuffleFileKind::Data, + byte_ranges: vec![], + }; + + let encoded: protobuf::Action = action.try_into()?; + Ok(StatementHandle::Partition(encoded.encode_to_vec())) +} + +fn build_flight_info( + schema: &Schema, + endpoints: Vec<FlightEndpoint>, + descriptor: FlightDescriptor, +) -> Result<FlightInfo, Status> { + FlightInfo::new() + .try_with_schema(schema) + .map_err(|e| Status::internal(format!("failed to encode result schema: {e}"))) + .map(|info| info.with_descriptor(descriptor).with_endpoints(endpoints)) +} + +/// Streams a single metadata batch back to the client. +fn one_batch_response(batch: RecordBatch) -> Response<DoGetStream> { + batch_response(batch.schema(), vec![batch]) +} + +fn batch_response(schema: SchemaRef, batches: Vec<RecordBatch>) -> Response<DoGetStream> { + let stream = FlightDataEncoderBuilder::new() + .with_schema(schema) + .build(futures::stream::iter(batches.into_iter().map(Ok))) + .map_err(|e| Status::internal(format!("failed to encode results: {e}"))); + + Response::new(Box::pin(stream) as DoGetStream) +} + +fn bearer_token(metadata: &MetadataMap) -> Option<String> { + let value = metadata.get("authorization")?.to_str().ok()?; + value + .strip_prefix("Bearer ") + .or_else(|| value.strip_prefix("bearer ")) + .map(str::to_string) +} + +fn invalid_ticket(e: impl std::fmt::Display) -> Status { + Status::invalid_argument(format!("invalid ticket: {e}")) +} + +fn query_failed(e: BallistaError) -> Status { + // A failed job is the client's problem to see, not an opaque 500. + Status::internal(format!("query execution failed: {e}")) +} + +fn encode_schema(schema: &Schema) -> Result<Vec<u8>, Status> { + let message: IpcMessage = SchemaAsIpc::new(schema, &IpcWriteOptions::default()) + .try_into() + .map_err(|e| Status::internal(format!("failed to encode schema: {e}")))?; + Ok(message.0.to_vec()) +} + +#[tonic::async_trait] +impl<B: QueryBackend> FlightSqlService for BallistaFlightSqlService<B> { + type FlightService = BallistaFlightSqlService<B>; + + async fn do_handshake( + &self, + request: Request<Streaming<HandshakeRequest>>, + ) -> Result< + Response<Pin<Box<dyn Stream<Item = Result<HandshakeResponse, Status>> + Send>>>, + Status, + > { + let identity = self.auth.authenticate(request.metadata()).await?; + + let token = Uuid::new_v4().to_string(); + let session_id = format!("flight-sql-{}", Uuid::new_v4()); + + // Build the session eagerly so a failure surfaces at handshake time + // rather than on the client's first query, and so the first query does + // not pay for it. + self.open_session(&session_id).await?; + self.store.insert_session(token.clone(), session_id.clone()); + + log::debug!( + "flight-sql: handshake for {:?} bound to session {session_id}", + identity.user + ); + + let result = HandshakeResponse { + protocol_version: 0, + payload: token.clone().into(), + }; + let stream = futures::stream::once(async move { Ok(result) }); + + let mut response = Response::new(Box::pin(stream) as _); + response.metadata_mut().insert( + "authorization", + format!("Bearer {token}") + .parse() + .map_err(|_| Status::internal("failed to encode session token"))?, + ); + Ok(response) + } + + /// Redeems tickets that are not Flight SQL commands. + /// + /// Ballista's own Rust client fetches shuffle output with a + /// `ballista.protobuf.Action` ticket. When the Flight SQL frontend is + /// mounted it replaces the standalone proxy on the scheduler's port, so it + /// has to keep serving those tickets. + async fn do_get_fallback( + &self, + request: Request<Ticket>, + _message: Any, + ) -> Result<Response<DoGetStream>, Status> { + decode_protobuf(&request.get_ref().ticket).map_err(invalid_ticket)?; + self.proxy.do_get(request).await + } + + async fn get_flight_info_statement( + &self, + query: CommandStatementQuery, + request: Request<FlightDescriptor>, + ) -> Result<Response<FlightInfo>, Status> { + let ctx = self.context(request.metadata()).await?; + let plan = Self::plan(&ctx, &query.query).await?; + let descriptor = request.into_inner(); + + self.flight_info_for(ctx, plan, descriptor, &job_name(&query.query)) + .await + .map(Response::new) + } + + async fn get_flight_info_prepared_statement( + &self, + query: CommandPreparedStatementQuery, + request: Request<FlightDescriptor>, + ) -> Result<Response<FlightInfo>, Status> { + let handle = prepared_handle(&query.prepared_statement_handle)?; + let prepared = self.store.prepared(&handle).ok_or_else(|| { + Status::not_found("unknown or expired prepared statement handle") + })?; + + let ctx = self.open_session(&prepared.session_id).await?; + let descriptor = request.into_inner(); + + self.flight_info_for(ctx, prepared.plan, descriptor, "flight-sql prepared") + .await + .map(Response::new) + } + + async fn get_flight_info_catalogs( + &self, + query: CommandGetCatalogs, + request: Request<FlightDescriptor>, + ) -> Result<Response<FlightInfo>, Status> { + let schema = query.into_builder().schema(); + Self::metadata_info(query, schema, request.into_inner()) + } + + async fn get_flight_info_schemas( + &self, + query: CommandGetDbSchemas, + request: Request<FlightDescriptor>, + ) -> Result<Response<FlightInfo>, Status> { + let schema = query.clone().into_builder().schema(); + Self::metadata_info(query, schema, request.into_inner()) + } + + async fn get_flight_info_tables( + &self, + query: CommandGetTables, + request: Request<FlightDescriptor>, + ) -> Result<Response<FlightInfo>, Status> { + let schema = query.clone().into_builder().schema(); + Self::metadata_info(query, schema, request.into_inner()) + } + + async fn get_flight_info_table_types( + &self, + query: CommandGetTableTypes, + request: Request<FlightDescriptor>, + ) -> Result<Response<FlightInfo>, Status> { + let schema = query.into_builder().schema(); + Self::metadata_info(query, schema, request.into_inner()) + } + + async fn get_flight_info_sql_info( + &self, + query: CommandGetSqlInfo, + request: Request<FlightDescriptor>, + ) -> Result<Response<FlightInfo>, Status> { + let schema = query.clone().into_builder(&self.sql_info).schema(); + Self::metadata_info(query, schema, request.into_inner()) + } + + async fn get_flight_info_xdbc_type_info( + &self, + query: CommandGetXdbcTypeInfo, + request: Request<FlightDescriptor>, + ) -> Result<Response<FlightInfo>, Status> { + let schema = query.into_builder(&self.xdbc_info).schema(); + Self::metadata_info(query, schema, request.into_inner()) + } + + async fn do_get_statement( + &self, + ticket: TicketStatementQuery, + _request: Request<Ticket>, + ) -> Result<Response<DoGetStream>, Status> { + let handle = + StatementHandle::decode(&ticket.statement_handle).map_err(invalid_ticket)?; + + match handle { + StatementHandle::Partition(action) => { + // Unwrap back to the executor-facing ticket and let the proxy + // redeem it. The proxy dials the executor with a fresh request, + // so the client's credentials are not forwarded onwards. + self.proxy + .do_get(Request::new(Ticket { + ticket: action.into(), + })) + .await + } + StatementHandle::Local(handle) => { + let result = self.store.take_result(&handle).ok_or_else(|| { + Status::not_found("result already consumed or expired") + })?; + Ok(batch_response(result.schema, result.batches)) + } + } + } + + async fn do_get_catalogs( + &self, + query: CommandGetCatalogs, + request: Request<Ticket>, + ) -> Result<Response<DoGetStream>, Status> { + let ctx = self.context(request.metadata()).await?; + Ok(one_batch_response(metadata::catalogs(&ctx, query)?)) + } + + async fn do_get_schemas( + &self, + query: CommandGetDbSchemas, + request: Request<Ticket>, + ) -> Result<Response<DoGetStream>, Status> { + let ctx = self.context(request.metadata()).await?; + Ok(one_batch_response(metadata::db_schemas(&ctx, query)?)) + } + + async fn do_get_tables( + &self, + query: CommandGetTables, + request: Request<Ticket>, + ) -> Result<Response<DoGetStream>, Status> { + let ctx = self.context(request.metadata()).await?; + Ok(one_batch_response(metadata::tables(&ctx, query).await?)) + } + + async fn do_get_table_types( + &self, + query: CommandGetTableTypes, + _request: Request<Ticket>, + ) -> Result<Response<DoGetStream>, Status> { + Ok(one_batch_response(metadata::table_types(query)?)) + } + + async fn do_get_sql_info( + &self, + query: CommandGetSqlInfo, + _request: Request<Ticket>, + ) -> Result<Response<DoGetStream>, Status> { + let batch = query + .into_builder(&self.sql_info) + .build() + .map_err(|e| Status::internal(format!("failed to build SqlInfo: {e}")))?; + Ok(one_batch_response(batch)) + } + + async fn do_get_xdbc_type_info( + &self, + query: CommandGetXdbcTypeInfo, + _request: Request<Ticket>, + ) -> Result<Response<DoGetStream>, Status> { + let batch = query + .into_builder(&self.xdbc_info) + .build() + .map_err(|e| Status::internal(format!("failed to build type info: {e}")))?; + Ok(one_batch_response(batch)) + } + + async fn do_put_statement_update( + &self, + ticket: CommandStatementUpdate, + request: Request<PeekableFlightDataStream>, + ) -> Result<i64, Status> { + let ctx = self.context(request.metadata()).await?; + let plan = Self::plan(&ctx, &ticket.query).await?; + + match disposition(&plan) { + Disposition::RunOnScheduler => {} + Disposition::Unsupported(reason) => { + return Err(Status::unimplemented(reason)); + } + Disposition::Distribute => { + return Err(Status::unimplemented( + "Ballista Flight SQL supports DDL and session statements on this \ + path; the distributed write path is not implemented", + )); + } + } + + execute_locally(&ctx, plan).await?; + // Only DDL and session statements reach here, and neither reports an + // affected-row count. + Ok(0) + } + + async fn do_action_create_prepared_statement( + &self, + query: ActionCreatePreparedStatementRequest, + request: Request<Action>, + ) -> Result<ActionCreatePreparedStatementResult, Status> { + let session_id = self.session_id(request.metadata())?; + let ctx = self.open_session(&session_id).await?; + let plan = Self::plan(&ctx, &query.query).await?; + let schema = plan.schema().as_arrow().clone(); + + let handle = Uuid::new_v4().to_string(); + self.store + .insert_prepared(handle.clone(), Prepared { session_id, plan }); + + Ok(ActionCreatePreparedStatementResult { + prepared_statement_handle: handle.into_bytes().into(), + dataset_schema: encode_schema(&schema)?.into(), + // Bound parameters are not supported yet, so the parameter schema + // is empty rather than absent: clients read "no parameters". + parameter_schema: encode_schema(&Schema::empty())?.into(), + }) + } + + async fn do_action_close_prepared_statement( + &self, + query: ActionClosePreparedStatementRequest, + _request: Request<Action>, + ) -> Result<(), Status> { + let handle = prepared_handle(&query.prepared_statement_handle)?; + self.store.remove_prepared(&handle); + Ok(()) + } + + async fn do_action_cancel_query( Review Comment: Removed. `do_action_cancel_query`, `job_id_from_flight_info`, `QueryBackend::cancel` and the test are gone, and `FlightSqlServerCancel` now reports `false`. The crate README lists query cancellation under not implemented. It can come back with `PollFlightInfo`. ########## ballista/flight-sql/src/metadata.rs: ########## @@ -0,0 +1,295 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! Driver-facing metadata: server capabilities, type info, and catalog +//! introspection. +//! +//! Catalog answers come from the session's real DataFusion catalog rather than +//! hand-built record batches, so what a driver sees in its schema browser is +//! what a query in the same session can actually reference. + +use arrow::array::RecordBatch; +use arrow_flight::sql::metadata::{ + GetCatalogsBuilder, GetDbSchemasBuilder, GetTablesBuilder, SqlInfoData, + SqlInfoDataBuilder, XdbcTypeInfo, XdbcTypeInfoData, XdbcTypeInfoDataBuilder, +}; +use arrow_flight::sql::{ + CommandGetCatalogs, CommandGetDbSchemas, CommandGetTableTypes, CommandGetTables, + Nullable, Searchable, SqlInfo, SqlSupportedTransaction, XdbcDataType, +}; +use ballista_core::BALLISTA_VERSION; +use datafusion::catalog::TableProvider; +use datafusion::logical_expr::TableType; +use datafusion::prelude::SessionContext; +use tonic::Status; + +/// Table type strings reported to clients, matching the JDBC vocabulary that +/// Flight SQL drivers expect. +const TABLE: &str = "TABLE"; +const VIEW: &str = "VIEW"; +const LOCAL_TEMPORARY: &str = "LOCAL TEMPORARY"; + +/// Describes the server to drivers deciding what SQL they may emit. +/// +/// The old implementation left `CommandGetSqlInfo` unimplemented with a +/// `// TODO: implement for FlightSQL JDBC to work` comment, which is precisely +/// why the JDBC driver could not connect. +pub(crate) fn sql_info() -> SqlInfoData { + let mut builder = SqlInfoDataBuilder::new(); + + builder.append(SqlInfo::FlightSqlServerName, "Apache DataFusion Ballista"); + builder.append(SqlInfo::FlightSqlServerVersion, BALLISTA_VERSION); + // Arrow IPC format version, per format/Schema.fbs. + builder.append(SqlInfo::FlightSqlServerArrowVersion, "1.3"); + builder.append(SqlInfo::FlightSqlServerReadOnly, false); + builder.append(SqlInfo::FlightSqlServerSql, true); + builder.append(SqlInfo::FlightSqlServerSubstrait, false); + builder.append( + SqlInfo::FlightSqlServerTransaction, + SqlSupportedTransaction::None as i32, + ); + builder.append(SqlInfo::FlightSqlServerCancel, true); + builder.append(SqlInfo::FlightSqlServerBulkIngestion, false); + + builder.append(SqlInfo::SqlDdlCatalog, false); + builder.append(SqlInfo::SqlDdlSchema, true); + builder.append(SqlInfo::SqlDdlTable, true); + // DataFusion folds unquoted identifiers to lower case and preserves the + // case of quoted ones: SQL_CASE_SENSITIVITY_LOWERCASE / _CASE_INSENSITIVE. + builder.append(SqlInfo::SqlIdentifierCase, 2i32); Review Comment: Fixed. Both use the named `SqlSupportedCaseSensitivity` values now. Unquoted identifiers report `SqlCaseSensitivityLowercase`. The enum has no value for quoted identifiers stored as written, so I followed the Arrow reference server (`FlightSqlExample`) and used `SqlCaseSensitivityCaseInsensitive`. The comment explains the choice. ########## ballista/flight-sql/src/service.rs: ########## @@ -0,0 +1,852 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! The Arrow Flight SQL frontend itself. + +use std::pin::Pin; +use std::sync::Arc; +use std::time::Duration; + +use arrow::array::RecordBatch; +use arrow::datatypes::{Schema, SchemaRef}; +use arrow::ipc::writer::IpcWriteOptions; +use arrow_flight::encode::FlightDataEncoderBuilder; +use arrow_flight::flight_service_server::FlightService; +use arrow_flight::sql::server::FlightSqlService; +use arrow_flight::sql::{ + ActionCancelQueryRequest, ActionCancelQueryResult, + ActionClosePreparedStatementRequest, ActionCreatePreparedStatementRequest, + ActionCreatePreparedStatementResult, Any, CommandGetCatalogs, CommandGetDbSchemas, + CommandGetSqlInfo, CommandGetTableTypes, CommandGetTables, CommandGetXdbcTypeInfo, + CommandPreparedStatementQuery, CommandStatementQuery, CommandStatementUpdate, + ProstMessageExt, SqlInfo, TicketStatementQuery, + metadata::{SqlInfoData, XdbcTypeInfoData}, + server::PeekableFlightDataStream, +}; +use arrow_flight::{ + Action, FlightData, FlightDescriptor, FlightEndpoint, FlightInfo, HandshakeRequest, + HandshakeResponse, IpcMessage, SchemaAsIpc, Ticket, +}; +use ballista_core::error::BallistaError; +use ballista_core::flight_proxy_service::BallistaFlightProxyService; +use ballista_core::planner::scans_only_local_tables; +use ballista_core::serde::protobuf::PartitionLocation; +use ballista_core::serde::scheduler::{Action as BallistaAction, ShuffleFileKind}; +use ballista_core::serde::{decode_protobuf, protobuf}; +use datafusion::logical_expr::{DdlStatement, LogicalPlan}; +use datafusion::prelude::SessionContext; +use futures::{Stream, TryStreamExt}; +use prost::Message; +use tonic::metadata::MetadataMap; +use tonic::{Request, Response, Status, Streaming}; +use uuid::Uuid; + +use crate::auth::{AnonymousAuthenticator, Authenticator}; +use crate::backend::QueryBackend; +use crate::metadata; +use crate::session::{LocalResult, Prepared, SessionStore}; +use crate::ticket::StatementHandle; + +/// Session shared by every client that connects without authenticating. +/// +/// Anonymous clients cannot be told apart, so they necessarily share catalog +/// state. Configure an [`Authenticator`] to get a session per connection. +pub const ANONYMOUS_SESSION: &str = "flight-sql-anonymous"; + +/// How long a session, prepared statement, or unredeemed local result may sit +/// idle before it is discarded. +const DEFAULT_TTL: Duration = Duration::from_secs(30 * 60); + +/// How often expired handles are swept. +const REAP_INTERVAL: Duration = Duration::from_secs(60); + +type DoGetStream = + Pin<Box<dyn Stream<Item = Result<FlightData, Status>> + Send + 'static>>; + +/// Serves Arrow Flight SQL on behalf of a Ballista cluster. +/// +/// Clients send SQL text; the frontend plans it against the session's catalog, +/// submits the plan through a [`QueryBackend`], and hands back one +/// `FlightEndpoint` per output partition. `DoGet` on those tickets is proxied +/// to the executor holding the partition, so clients never need to reach +/// executors themselves — the failure mode that made the pre-46.0.0 +/// implementation unusable behind NAT, Docker, and Kubernetes. +pub struct BallistaFlightSqlService<B: QueryBackend> { + backend: Arc<B>, + proxy: BallistaFlightProxyService, + auth: Arc<dyn Authenticator>, + store: Arc<SessionStore>, + sql_info: SqlInfoData, + xdbc_info: XdbcTypeInfoData, +} + +impl<B: QueryBackend> BallistaFlightSqlService<B> { + /// Builds a frontend over `backend`, using `proxy` to stream partition + /// data back from executors. + /// + /// The service authenticates nobody until an [`Authenticator`] is supplied + /// via [`with_authenticator`](Self::with_authenticator). + pub fn new(backend: Arc<B>, proxy: BallistaFlightProxyService) -> Self { + let store = Arc::new(SessionStore::new(DEFAULT_TTL)); + + let backend_for_reaper = backend.clone(); + store.spawn_reaper(REAP_INTERVAL, move |session_id| { + let backend = backend_for_reaper.clone(); + async move { + if let Err(e) = backend.close_session(&session_id).await { + log::warn!("flight-sql: failed to close session {session_id}: {e}"); + } + } + }); + + Self { + backend, + proxy, + auth: Arc::new(AnonymousAuthenticator), + store, + sql_info: metadata::sql_info(), + xdbc_info: metadata::xdbc_type_info(), + } + } + + /// Installs an authenticator. Without one, every handshake is accepted and + /// unauthenticated clients share a single session. + pub fn with_authenticator(mut self, auth: Arc<dyn Authenticator>) -> Self { + self.auth = auth; + self + } + + /// True when the service will accept unauthenticated clients, which the + /// scheduler logs at startup. + pub fn allows_anonymous(&self) -> bool { + self.auth.allows_anonymous() + } + + /// Resolves the Ballista session for a request from its bearer token. + fn session_id(&self, metadata: &MetadataMap) -> Result<String, Status> { + match bearer_token(metadata) { + Some(token) => self.store.session(&token).ok_or_else(|| { + Status::unauthenticated( + "unknown or expired session token; re-run the Flight handshake", + ) + }), + None if self.auth.allows_anonymous() => Ok(ANONYMOUS_SESSION.to_string()), + None => Err(Status::unauthenticated( + "missing bearer token; authenticate with the Flight handshake first", + )), + } + } + + /// Returns the context for `session_id`, building it on first use. + /// + /// The cache is what makes a session a session: [`QueryBackend::session`] + /// builds a fresh `SessionContext` every call, so without it a table + /// created by one request would be invisible to the next — and every + /// request would pay for a full DataFusion session to be constructed. + async fn open_session( + &self, + session_id: &str, + ) -> Result<Arc<SessionContext>, Status> { + if let Some(ctx) = self.store.context(session_id) { + return Ok(ctx); + } + + let ctx = self + .backend + .session(session_id) + .await + .map_err(|e| Status::internal(format!("failed to open session: {e}")))?; + + Ok(self.store.insert_context(session_id.to_string(), ctx)) + } + + /// Resolves a request to the context it should be planned against. + async fn context( + &self, + metadata: &MetadataMap, + ) -> Result<Arc<SessionContext>, Status> { + self.open_session(&self.session_id(metadata)?).await + } + + /// Plans `sql` against the session's catalog. + async fn plan(ctx: &SessionContext, sql: &str) -> Result<LogicalPlan, Status> { + ctx.state() + .create_logical_plan(sql) + .await + .map_err(|e| Status::invalid_argument(format!("failed to plan query: {e}"))) + } + + /// Runs a planned statement and describes where to collect its results. + async fn flight_info_for( + &self, + ctx: Arc<SessionContext>, + plan: LogicalPlan, + descriptor: FlightDescriptor, + job_name: &str, + ) -> Result<FlightInfo, Status> { + if let Disposition::Unsupported(reason) = disposition(&plan) { + return Err(Status::unimplemented(reason)); + } + + if disposition(&plan) == Disposition::RunOnScheduler { + let (schema, batches) = execute_locally(&ctx, plan).await?; + let handle = Uuid::new_v4().to_string(); + self.store.insert_result( + handle.clone(), + LocalResult { + schema: schema.clone(), + batches, + }, + ); + + let endpoint = Self::endpoint(StatementHandle::Local(handle)); + return build_flight_info(&schema, vec![endpoint], descriptor); + } + + let result = self + .backend + .execute(job_name, ctx, plan) + .await + .map_err(query_failed)?; + + let endpoints = result + .partitions + .into_iter() + .map(|location| partition_handle(location).map(Self::endpoint)) + .collect::<Result<Vec<_>, _>>() + .map_err(|e| Status::internal(format!("invalid partition location: {e}")))?; + + log::debug!( + "flight-sql: job {} produced {} endpoint(s)", + result.job_id, + endpoints.len() + ); + + build_flight_info(&result.schema, endpoints, descriptor) + } + + /// Builds the endpoint a client redeems for one slice of the result. + /// + /// Endpoints carry no location, which Flight defines as "fetch from the + /// server that gave you this FlightInfo". That keeps every cluster-internal + /// address off the wire and lets the frontend work unchanged behind NAT, a + /// load balancer, or an ingress. + fn endpoint(handle: StatementHandle) -> FlightEndpoint { + let ticket = TicketStatementQuery { + statement_handle: handle.encode().into(), + }; + FlightEndpoint::new().with_ticket(Ticket { + ticket: ticket.as_any().encode_to_vec().into(), + }) + } + + /// Serves a metadata command by round-tripping the command itself as the + /// ticket, so `DoGet` lands back on the matching handler. + fn metadata_info<C: ProstMessageExt>( + command: C, + schema: SchemaRef, + descriptor: FlightDescriptor, + ) -> Result<Response<FlightInfo>, Status> { + let endpoint = FlightEndpoint::new().with_ticket(Ticket { + ticket: command.as_any().encode_to_vec().into(), + }); + build_flight_info(&schema, vec![endpoint], descriptor).map(Response::new) + } +} + +/// What the frontend should do with a planned statement. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum Disposition { + /// Execute on the scheduler: it mutates session state and produces no data + /// worth distributing. + RunOnScheduler, + /// Submit to the cluster. + Distribute, + /// Refuse, with an explanation for the client. + Unsupported(&'static str), +} + +/// Classifies a plan. +/// +/// One function rather than two predicates, because the interesting cases are +/// the ones where the answers overlap: `CREATE TABLE AS SELECT` is DDL, and +/// DDL runs on the scheduler, so a caller that asked "is this DDL?" before +/// asking "is this supported?" would silently execute its query on one node. +/// Returning a single verdict makes that ordering impossible to get wrong. +fn disposition(plan: &LogicalPlan) -> Disposition { + match plan { + LogicalPlan::Dml(_) => Disposition::Unsupported( + "Ballista Flight SQL does not support INSERT/UPDATE/DELETE; \ + the distributed write path is not implemented", + ), + LogicalPlan::Copy(_) => Disposition::Unsupported( + "Ballista Flight SQL does not support COPY; \ + the distributed write path is not implemented", + ), + LogicalPlan::Ddl(DdlStatement::CreateMemoryTable(_)) => Disposition::Unsupported( + "Ballista Flight SQL does not support CREATE TABLE AS SELECT, \ + because it would execute on the scheduler rather than the cluster; \ + use CREATE EXTERNAL TABLE over data the executors can read", + ), + // Other DDL only edits the session catalog, and `SET`-style statements + // only edit session config; neither has anything to distribute. + LogicalPlan::Ddl(_) | LogicalPlan::Statement(_) => Disposition::RunOnScheduler, Review Comment: Fixed. `refusal` now walks the whole plan with `apply_with_subqueries` and refuses `Dml`, `Copy`, `CreateMemoryTable` and `Statement::Execute` wherever they appear. `EXECUTE q` and `EXPLAIN ANALYZE INSERT ...` are both refused now, and `statements_that_cannot_be_distributed_are_refused` covers them. ########## ballista/scheduler/src/config.rs: ########## @@ -397,6 +405,19 @@ pub struct SchedulerConfig { pub task_max_failures: usize, /// Number of failures attempts before stage is considered failed pub stage_max_failures: usize, + /// Whether to serve Arrow Flight SQL on the scheduler's gRPC port. + /// + /// The frontend subsumes the plain Flight result proxy when enabled, so + /// Ballista's own clients keep working on the same port. + #[cfg(feature = "flight-sql")] + pub flight_sql: bool, + /// Authenticates Arrow Flight SQL handshakes. + /// + /// Without one the frontend accepts every client, so this is the only way Review Comment: Reworded. The field doc now says the authenticator isolates clients' sessions but doesn't secure the port, and to keep the scheduler on a trusted network. The crate README says the same. I also just fixed the `--flight-sql` help text, which still implied the authenticator secures the port. ########## ballista/flight-sql/src/service.rs: ########## @@ -0,0 +1,852 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! The Arrow Flight SQL frontend itself. + +use std::pin::Pin; +use std::sync::Arc; +use std::time::Duration; + +use arrow::array::RecordBatch; +use arrow::datatypes::{Schema, SchemaRef}; +use arrow::ipc::writer::IpcWriteOptions; +use arrow_flight::encode::FlightDataEncoderBuilder; +use arrow_flight::flight_service_server::FlightService; +use arrow_flight::sql::server::FlightSqlService; +use arrow_flight::sql::{ + ActionCancelQueryRequest, ActionCancelQueryResult, + ActionClosePreparedStatementRequest, ActionCreatePreparedStatementRequest, + ActionCreatePreparedStatementResult, Any, CommandGetCatalogs, CommandGetDbSchemas, + CommandGetSqlInfo, CommandGetTableTypes, CommandGetTables, CommandGetXdbcTypeInfo, + CommandPreparedStatementQuery, CommandStatementQuery, CommandStatementUpdate, + ProstMessageExt, SqlInfo, TicketStatementQuery, + metadata::{SqlInfoData, XdbcTypeInfoData}, + server::PeekableFlightDataStream, +}; +use arrow_flight::{ + Action, FlightData, FlightDescriptor, FlightEndpoint, FlightInfo, HandshakeRequest, + HandshakeResponse, IpcMessage, SchemaAsIpc, Ticket, +}; +use ballista_core::error::BallistaError; +use ballista_core::flight_proxy_service::BallistaFlightProxyService; +use ballista_core::planner::scans_only_local_tables; +use ballista_core::serde::protobuf::PartitionLocation; +use ballista_core::serde::scheduler::{Action as BallistaAction, ShuffleFileKind}; +use ballista_core::serde::{decode_protobuf, protobuf}; +use datafusion::logical_expr::{DdlStatement, LogicalPlan}; +use datafusion::prelude::SessionContext; +use futures::{Stream, TryStreamExt}; +use prost::Message; +use tonic::metadata::MetadataMap; +use tonic::{Request, Response, Status, Streaming}; +use uuid::Uuid; + +use crate::auth::{AnonymousAuthenticator, Authenticator}; +use crate::backend::QueryBackend; +use crate::metadata; +use crate::session::{LocalResult, Prepared, SessionStore}; +use crate::ticket::StatementHandle; + +/// Session shared by every client that connects without authenticating. +/// +/// Anonymous clients cannot be told apart, so they necessarily share catalog +/// state. Configure an [`Authenticator`] to get a session per connection. +pub const ANONYMOUS_SESSION: &str = "flight-sql-anonymous"; + +/// How long a session, prepared statement, or unredeemed local result may sit +/// idle before it is discarded. +const DEFAULT_TTL: Duration = Duration::from_secs(30 * 60); + +/// How often expired handles are swept. +const REAP_INTERVAL: Duration = Duration::from_secs(60); + +type DoGetStream = + Pin<Box<dyn Stream<Item = Result<FlightData, Status>> + Send + 'static>>; + +/// Serves Arrow Flight SQL on behalf of a Ballista cluster. +/// +/// Clients send SQL text; the frontend plans it against the session's catalog, +/// submits the plan through a [`QueryBackend`], and hands back one +/// `FlightEndpoint` per output partition. `DoGet` on those tickets is proxied +/// to the executor holding the partition, so clients never need to reach +/// executors themselves — the failure mode that made the pre-46.0.0 +/// implementation unusable behind NAT, Docker, and Kubernetes. +pub struct BallistaFlightSqlService<B: QueryBackend> { + backend: Arc<B>, + proxy: BallistaFlightProxyService, + auth: Arc<dyn Authenticator>, + store: Arc<SessionStore>, + sql_info: SqlInfoData, + xdbc_info: XdbcTypeInfoData, +} + +impl<B: QueryBackend> BallistaFlightSqlService<B> { + /// Builds a frontend over `backend`, using `proxy` to stream partition + /// data back from executors. + /// + /// The service authenticates nobody until an [`Authenticator`] is supplied + /// via [`with_authenticator`](Self::with_authenticator). + pub fn new(backend: Arc<B>, proxy: BallistaFlightProxyService) -> Self { + let store = Arc::new(SessionStore::new(DEFAULT_TTL)); + + let backend_for_reaper = backend.clone(); + store.spawn_reaper(REAP_INTERVAL, move |session_id| { + let backend = backend_for_reaper.clone(); + async move { + if let Err(e) = backend.close_session(&session_id).await { + log::warn!("flight-sql: failed to close session {session_id}: {e}"); + } + } + }); + + Self { + backend, + proxy, + auth: Arc::new(AnonymousAuthenticator), + store, + sql_info: metadata::sql_info(), + xdbc_info: metadata::xdbc_type_info(), + } + } + + /// Installs an authenticator. Without one, every handshake is accepted and + /// unauthenticated clients share a single session. + pub fn with_authenticator(mut self, auth: Arc<dyn Authenticator>) -> Self { + self.auth = auth; + self + } + + /// True when the service will accept unauthenticated clients, which the + /// scheduler logs at startup. + pub fn allows_anonymous(&self) -> bool { + self.auth.allows_anonymous() + } + + /// Resolves the Ballista session for a request from its bearer token. + fn session_id(&self, metadata: &MetadataMap) -> Result<String, Status> { + match bearer_token(metadata) { + Some(token) => self.store.session(&token).ok_or_else(|| { + Status::unauthenticated( + "unknown or expired session token; re-run the Flight handshake", + ) + }), + None if self.auth.allows_anonymous() => Ok(ANONYMOUS_SESSION.to_string()), + None => Err(Status::unauthenticated( + "missing bearer token; authenticate with the Flight handshake first", + )), + } + } + + /// Returns the context for `session_id`, building it on first use. + /// + /// The cache is what makes a session a session: [`QueryBackend::session`] + /// builds a fresh `SessionContext` every call, so without it a table + /// created by one request would be invisible to the next — and every + /// request would pay for a full DataFusion session to be constructed. + async fn open_session( + &self, + session_id: &str, + ) -> Result<Arc<SessionContext>, Status> { + if let Some(ctx) = self.store.context(session_id) { + return Ok(ctx); + } + + let ctx = self + .backend + .session(session_id) + .await + .map_err(|e| Status::internal(format!("failed to open session: {e}")))?; + + Ok(self.store.insert_context(session_id.to_string(), ctx)) + } + + /// Resolves a request to the context it should be planned against. + async fn context( + &self, + metadata: &MetadataMap, + ) -> Result<Arc<SessionContext>, Status> { + self.open_session(&self.session_id(metadata)?).await + } + + /// Plans `sql` against the session's catalog. + async fn plan(ctx: &SessionContext, sql: &str) -> Result<LogicalPlan, Status> { + ctx.state() + .create_logical_plan(sql) + .await + .map_err(|e| Status::invalid_argument(format!("failed to plan query: {e}"))) + } + + /// Runs a planned statement and describes where to collect its results. + async fn flight_info_for( + &self, + ctx: Arc<SessionContext>, + plan: LogicalPlan, + descriptor: FlightDescriptor, + job_name: &str, + ) -> Result<FlightInfo, Status> { + if let Disposition::Unsupported(reason) = disposition(&plan) { Review Comment: Fixed. `flight_info_for` now does a single `match disposition(&plan)` over the three variants, like `do_put_statement_update`. ########## ballista/flight-sql/src/session.rs: ########## @@ -0,0 +1,277 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! Server-side state the Flight SQL frontend keeps between requests. +//! +//! Everything in here is keyed by an opaque handle the client is given and +//! hands back, and everything expires. A client that disconnects without +//! closing its prepared statements (or that never redeems a ticket) must not +//! pin memory forever, which is what the previous Flight SQL implementation +//! got wrong: its plan cache only shed entries on an explicit close. + +use std::sync::Arc; +use std::time::{Duration, Instant}; + +use arrow::array::RecordBatch; +use arrow::datatypes::SchemaRef; +use dashmap::DashMap; +use datafusion::logical_expr::LogicalPlan; +use datafusion::prelude::SessionContext; + +/// A prepared statement's server-side state. +#[derive(Clone)] +pub(crate) struct Prepared { + /// Session the statement was prepared in; it is planned and executed + /// against that session's catalog. + pub session_id: String, + /// The plan as prepared; execution never re-plans it. + pub plan: LogicalPlan, +} + +/// Result of a statement the frontend ran locally rather than distributing +/// (DDL, and DML whose only output is an affected-row count). +#[derive(Clone)] +pub(crate) struct LocalResult { + pub schema: SchemaRef, + pub batches: Vec<RecordBatch>, +} + +struct Tracked<T> { + value: T, + last_used: Instant, +} + +impl<T> Tracked<T> { + fn new(value: T) -> Self { + Self { + value, + last_used: Instant::now(), + } + } +} + +/// Handle stores for sessions, prepared statements, and locally-computed +/// results, all with a shared idle TTL. +pub(crate) struct SessionStore { + ttl: Duration, + /// Bearer token -> Ballista session id. + sessions: DashMap<String, Tracked<String>>, + /// Ballista session id -> its context. + contexts: DashMap<String, Tracked<Arc<SessionContext>>>, + /// Prepared statement handle -> prepared plan. + prepared: DashMap<String, Tracked<Prepared>>, + /// Local result handle -> materialized batches. + results: DashMap<String, Tracked<LocalResult>>, +} + +impl SessionStore { + pub(crate) fn new(ttl: Duration) -> Self { + Self { + ttl, + sessions: DashMap::new(), + contexts: DashMap::new(), + prepared: DashMap::new(), + results: DashMap::new(), + } + } + + pub(crate) fn insert_session(&self, token: String, session_id: String) { + self.sessions.insert(token, Tracked::new(session_id)); + } + + /// Resolves a bearer token to its session id, refreshing its idle timer. + pub(crate) fn session(&self, token: &str) -> Option<String> { + self.sessions.get_mut(token).map(|mut entry| { + entry.last_used = Instant::now(); + entry.value.clone() + }) + } + + /// Drops a token and the session it named, returning that session id. + /// + /// Each token gets its own session, so removing the token always releases + /// the session with it. + pub(crate) fn remove_session(&self, token: &str) -> Option<String> { + let (_, entry) = self.sessions.remove(token)?; + self.contexts.remove(&entry.value); + Some(entry.value) + } + + /// Returns the cached context for a session, refreshing its idle timer. + pub(crate) fn context(&self, session_id: &str) -> Option<Arc<SessionContext>> { + self.contexts.get_mut(session_id).map(|mut entry| { + entry.last_used = Instant::now(); + entry.value.clone() + }) + } + + /// Caches a context, or returns the one another caller cached first. + /// + /// Two requests for a cold session can both build one; letting the first + /// insertion win keeps them looking at the same catalog. + pub(crate) fn insert_context( + &self, + session_id: String, + ctx: Arc<SessionContext>, + ) -> Arc<SessionContext> { + self.contexts + .entry(session_id) + .or_insert_with(|| Tracked::new(ctx)) + .value + .clone() + } + + pub(crate) fn insert_prepared(&self, handle: String, prepared: Prepared) { + self.prepared.insert(handle, Tracked::new(prepared)); + } + + pub(crate) fn prepared(&self, handle: &str) -> Option<Prepared> { + self.prepared.get_mut(handle).map(|mut entry| { + entry.last_used = Instant::now(); + entry.value.clone() + }) + } + + pub(crate) fn remove_prepared(&self, handle: &str) { + self.prepared.remove(handle); + } + + pub(crate) fn insert_result(&self, handle: String, result: LocalResult) { + self.results.insert(handle, Tracked::new(result)); + } + + /// Takes a local result. Results are single-use: a ticket is redeemed once. + pub(crate) fn take_result(&self, handle: &str) -> Option<LocalResult> { + self.results.remove(handle).map(|(_, entry)| entry.value) + } + + /// Evicts everything idle for longer than the TTL. + /// + /// Returns the session ids that no longer have any live token, so the + /// caller can release them in the backend. + pub(crate) fn sweep(&self) -> Vec<String> { Review Comment: Done. `sweep` collects the released session ids in a single `retain` pass and drops their contexts afterwards. `remove_session` is gone. ########## ballista/flight-sql/src/session.rs: ########## @@ -0,0 +1,277 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! Server-side state the Flight SQL frontend keeps between requests. +//! +//! Everything in here is keyed by an opaque handle the client is given and +//! hands back, and everything expires. A client that disconnects without +//! closing its prepared statements (or that never redeems a ticket) must not +//! pin memory forever, which is what the previous Flight SQL implementation +//! got wrong: its plan cache only shed entries on an explicit close. + +use std::sync::Arc; +use std::time::{Duration, Instant}; + +use arrow::array::RecordBatch; +use arrow::datatypes::SchemaRef; +use dashmap::DashMap; +use datafusion::logical_expr::LogicalPlan; +use datafusion::prelude::SessionContext; + +/// A prepared statement's server-side state. +#[derive(Clone)] +pub(crate) struct Prepared { + /// Session the statement was prepared in; it is planned and executed + /// against that session's catalog. + pub session_id: String, + /// The plan as prepared; execution never re-plans it. + pub plan: LogicalPlan, +} + +/// Result of a statement the frontend ran locally rather than distributing +/// (DDL, and DML whose only output is an affected-row count). +#[derive(Clone)] +pub(crate) struct LocalResult { + pub schema: SchemaRef, + pub batches: Vec<RecordBatch>, +} + +struct Tracked<T> { + value: T, + last_used: Instant, +} + +impl<T> Tracked<T> { + fn new(value: T) -> Self { + Self { + value, + last_used: Instant::now(), + } + } +} + +/// Handle stores for sessions, prepared statements, and locally-computed +/// results, all with a shared idle TTL. +pub(crate) struct SessionStore { + ttl: Duration, + /// Bearer token -> Ballista session id. + sessions: DashMap<String, Tracked<String>>, + /// Ballista session id -> its context. + contexts: DashMap<String, Tracked<Arc<SessionContext>>>, + /// Prepared statement handle -> prepared plan. + prepared: DashMap<String, Tracked<Prepared>>, + /// Local result handle -> materialized batches. + results: DashMap<String, Tracked<LocalResult>>, +} + +impl SessionStore { + pub(crate) fn new(ttl: Duration) -> Self { + Self { + ttl, + sessions: DashMap::new(), + contexts: DashMap::new(), + prepared: DashMap::new(), + results: DashMap::new(), + } + } + + pub(crate) fn insert_session(&self, token: String, session_id: String) { + self.sessions.insert(token, Tracked::new(session_id)); + } + + /// Resolves a bearer token to its session id, refreshing its idle timer. + pub(crate) fn session(&self, token: &str) -> Option<String> { + self.sessions.get_mut(token).map(|mut entry| { + entry.last_used = Instant::now(); + entry.value.clone() + }) + } + + /// Drops a token and the session it named, returning that session id. + /// + /// Each token gets its own session, so removing the token always releases + /// the session with it. + pub(crate) fn remove_session(&self, token: &str) -> Option<String> { + let (_, entry) = self.sessions.remove(token)?; + self.contexts.remove(&entry.value); + Some(entry.value) + } + + /// Returns the cached context for a session, refreshing its idle timer. + pub(crate) fn context(&self, session_id: &str) -> Option<Arc<SessionContext>> { + self.contexts.get_mut(session_id).map(|mut entry| { + entry.last_used = Instant::now(); + entry.value.clone() + }) + } + + /// Caches a context, or returns the one another caller cached first. + /// + /// Two requests for a cold session can both build one; letting the first + /// insertion win keeps them looking at the same catalog. + pub(crate) fn insert_context( + &self, + session_id: String, + ctx: Arc<SessionContext>, + ) -> Arc<SessionContext> { + self.contexts + .entry(session_id) + .or_insert_with(|| Tracked::new(ctx)) + .value + .clone() + } + + pub(crate) fn insert_prepared(&self, handle: String, prepared: Prepared) { + self.prepared.insert(handle, Tracked::new(prepared)); + } + + pub(crate) fn prepared(&self, handle: &str) -> Option<Prepared> { + self.prepared.get_mut(handle).map(|mut entry| { + entry.last_used = Instant::now(); + entry.value.clone() + }) + } + + pub(crate) fn remove_prepared(&self, handle: &str) { + self.prepared.remove(handle); + } + + pub(crate) fn insert_result(&self, handle: String, result: LocalResult) { + self.results.insert(handle, Tracked::new(result)); + } + + /// Takes a local result. Results are single-use: a ticket is redeemed once. + pub(crate) fn take_result(&self, handle: &str) -> Option<LocalResult> { + self.results.remove(handle).map(|(_, entry)| entry.value) + } + + /// Evicts everything idle for longer than the TTL. + /// + /// Returns the session ids that no longer have any live token, so the + /// caller can release them in the backend. + pub(crate) fn sweep(&self) -> Vec<String> { + let ttl = self.ttl; + let expired: Vec<String> = self + .sessions + .iter() + .filter(|entry| entry.value().last_used.elapsed() > ttl) + .map(|entry| entry.key().clone()) + .collect(); + + let mut released = Vec::new(); + for token in expired { + if let Some(session_id) = self.remove_session(&token) { + released.push(session_id); + } + } + + self.contexts + .retain(|_, entry| entry.last_used.elapsed() <= ttl); + self.prepared + .retain(|_, entry| entry.last_used.elapsed() <= ttl); + self.results + .retain(|_, entry| entry.last_used.elapsed() <= ttl); + + released + } + + /// Starts a background task that sweeps at `interval`, closing sessions it + /// evicts. The task ends when the last reference to the store is dropped. + pub(crate) fn spawn_reaper<F, Fut>(self: &Arc<Self>, interval: Duration, close: F) + where + F: Fn(String) -> Fut + Send + Sync + 'static, + Fut: std::future::Future<Output = ()> + Send, + { + let store = Arc::downgrade(self); + tokio::spawn(async move { + let mut ticker = tokio::time::interval(interval); + // The first tick completes immediately; skip it so we do not sweep + // a store that was created a moment ago. + ticker.tick().await; + loop { + ticker.tick().await; + let Some(store) = store.upgrade() else { + return; + }; + for session_id in store.sweep() { + log::debug!("flight-sql: expiring idle session {session_id}"); + close(session_id).await; + } + } + }); + } +} + +#[cfg(test)] +mod test { + use super::*; + + #[test] + fn session_survives_while_touched_and_expires_when_idle() { + let store = SessionStore::new(Duration::from_millis(50)); + store.insert_session("token".to_string(), "session".to_string()); + + assert_eq!(store.session("token"), Some("session".to_string())); + assert!(store.sweep().is_empty()); + + std::thread::sleep(Duration::from_millis(60)); Review Comment: Fixed. The test backdates `last_used` instead of sleeping, like the idle-eviction tests in `client_pool.rs`. ########## .github/workflows/rust.yml: ########## @@ -99,6 +99,12 @@ jobs: export ARROW_TEST_DATA=$(pwd)/testing/data export PARQUET_TEST_DATA=$(pwd)/parquet-testing/data cargo test --profile ci --features=testcontainers + - name: Run Arrow Flight SQL tests + # The frontend is a non-default feature, so the workspace run above + # compiles its integration test away. + run: | + export PATH=$PATH:$HOME/d/protoc/bin + cargo test --profile ci -p ballista-scheduler --features flight-sql --test flight_sql Review Comment: Both are sorted. The separate step is gone and `ballista-scheduler/flight-sql` is part of `--features` in the main test step, so it's one build. Separately, #2473 switched `rust_clippy.sh` to `--workspace`, so `ballista-flight-sql` gets linted now. -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected] --------------------------------------------------------------------- To unsubscribe, e-mail: [email protected] For additional commands, e-mail: [email protected]
