diff --git a/src/common/datasource/src/parquet_writer.rs b/src/common/datasource/src/parquet_writer.rs index ae1fcec04a9..ed35c314339 100644 --- a/src/common/datasource/src/parquet_writer.rs +++ b/src/common/datasource/src/parquet_writer.rs @@ -29,6 +29,13 @@ use tokio_util::sync::CancellationToken; use crate::DEFAULT_WRITE_BUFFER_SIZE; use crate::error::{self, Result}; +/// Destination creation policy; conditional failures never authorize path deletion. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum ParquetCreationPolicy { + Overwrite, + IfNotExists, +} + /// Limits for one Parquet file. Flush thresholds are not hard memory caps. #[derive(Clone, Copy, Debug)] pub struct ParquetWriterLimits { @@ -49,6 +56,8 @@ pub struct ParquetFileWriter { store: ObjectStore, path: String, limits: Option, + creation: ParquetCreationPolicy, + close_started: bool, } impl ParquetFileWriter { @@ -60,6 +69,26 @@ impl ParquetFileWriter { path: &str, concurrency: usize, limits: Option, + ) -> Result { + Self::open_with_creation( + schema, + store, + path, + concurrency, + limits, + ParquetCreationPolicy::Overwrite, + ) + .await + } + + /// Open with a conditional policy only when the caller verified backend support. + pub async fn open_with_creation( + schema: SchemaRef, + store: ObjectStore, + path: &str, + concurrency: usize, + limits: Option, + creation: ParquetCreationPolicy, ) -> Result { let mut props = WriterProperties::builder() .set_compression(Compression::ZSTD(ZstdLevel::default())) @@ -90,6 +119,7 @@ impl ParquetFileWriter { .writer_with(path) .concurrent(concurrency) .chunk(DEFAULT_WRITE_BUFFER_SIZE.as_bytes() as usize) + .if_not_exists(creation == ParquetCreationPolicy::IfNotExists) .await .context(error::WriteObjectSnafu { path })?; Ok(Self { @@ -98,6 +128,8 @@ impl ParquetFileWriter { store, path: path.to_owned(), limits, + creation, + close_started: false, }) } @@ -156,9 +188,12 @@ impl ParquetFileWriter { } async fn write_bytes(&mut self, bytes: Vec) -> Result<()> { - if !bytes.is_empty() { + let bytes = Bytes::from(bytes); + let chunk = DEFAULT_WRITE_BUFFER_SIZE.as_bytes() as usize; + // Slices retain the complete encoded allocation until its last submission. + for offset in (0..bytes.len()).step_by(chunk) { self.sink - .write(bytes) + .write(bytes.slice(offset..(offset + chunk).min(bytes.len()))) .await .context(error::WriteObjectSnafu { path: &self.path })?; } @@ -183,6 +218,7 @@ impl ParquetFileWriter { .context(error::JoinHandleSnafu)??; self.write_bytes(bytes).await?; check_cancelled(cancellation)?; + self.close_started = true; self.sink .close() .await @@ -191,13 +227,16 @@ impl ParquetFileWriter { Ok(()) } - /// Abort an exclusively owned file after all in-flight operations have completed. - /// If the backend cannot abort, deletion assumes this attempt owns the path. + /// Abort after all in-flight operations complete. Preserve ambiguous commits; + /// conditional callers delegate cleanup exclusively to the backend. pub async fn abort(mut self) -> Result<()> { let result = self.sink.abort().await; - if result - .as_ref() - .is_err_and(|e| e.kind() == object_store::ErrorKind::Unsupported) + if self.creation == ParquetCreationPolicy::Overwrite + && result.as_ref().is_err_and(|error| { + error.kind() == object_store::ErrorKind::Unsupported + && (!self.close_started + || object_store::secure_fs::is_unsynced_overwrite_abort(error)) + }) { let store = self.store.clone(); let path = self.path.clone(); @@ -288,6 +327,196 @@ mod tests { .unwrap() } + #[tokio::test] + async fn abort_preserves_collisions_and_ambiguous_commits() { + use object_store::layers::mock::{MockLayerBuilder, MockWriterFactory, oio}; + + struct AmbiguousCommit(oio::Writer); + impl oio::Write for AmbiguousCommit { + async fn write(&mut self, bytes: object_store::Buffer) -> object_store::Result<()> { + self.0.write(bytes).await + } + async fn close( + &mut self, + ) -> object_store::Result { + self.0.close().await?; + Err(object_store::Error::new( + object_store::ErrorKind::Unexpected, + "lost close reply", + )) + } + async fn abort(&mut self) -> object_store::Result<()> { + Err(object_store::Error::new( + object_store::ErrorKind::Unsupported, + "cannot abort", + )) + } + } + + for (creation, existing) in [ + (ParquetCreationPolicy::Overwrite, false), + (ParquetCreationPolicy::IfNotExists, false), + (ParquetCreationPolicy::IfNotExists, true), + ] { + let directory = common_test_util::temp_dir::create_temp_dir("conditional_parquet"); + let store = object_store::secure_fs::SecureFsRoot::open(directory.path()) + .unwrap() + .build_operator(); + let path = "conditional.parquet"; + if existing { + store.write(path, "original").await.unwrap(); + } + let factory: MockWriterFactory = + Arc::new(|_, _, writer| Box::new(AmbiguousCommit(writer))); + let store = store.layer( + MockLayerBuilder::default() + .writer_factory(factory) + .build() + .unwrap(), + ); + let mut writer = ParquetFileWriter::open_with_creation( + batch().schema(), + store.clone(), + path, + 1, + None, + creation, + ) + .await + .unwrap(); + writer.write(batch(), None).await.unwrap(); + assert!(writer.finish(None).await.is_err()); + assert!(writer.abort().await.is_err()); + if existing { + assert_eq!( + store.read(path).await.unwrap().to_bytes(), + Bytes::from_static(b"original") + ); + } else { + assert_eq!( + read(&store, path) + .await + .metadata() + .file_metadata() + .num_rows(), + 4 + ); + } + } + } + + #[tokio::test] + async fn overwrite_abort_deletes_after_unsynced_close() { + use object_store::layers::mock::{MockLayerBuilder, MockWriterFactory, oio}; + + struct FailedClose(oio::Writer); + impl oio::Write for FailedClose { + async fn write(&mut self, bytes: object_store::Buffer) -> object_store::Result<()> { + self.0.write(bytes).await + } + async fn close( + &mut self, + ) -> object_store::Result { + Err(object_store::Error::new( + object_store::ErrorKind::Unexpected, + "close failed before sync", + )) + } + async fn abort(&mut self) -> object_store::Result<()> { + self.0.abort().await + } + } + + let directory = common_test_util::temp_dir::create_temp_dir("unsynced_parquet"); + let store = object_store::secure_fs::SecureFsRoot::open(directory.path()) + .unwrap() + .build_operator(); + let path = "partial.parquet"; + let factory: MockWriterFactory = Arc::new(|_, _, writer| Box::new(FailedClose(writer))); + let store = store.layer( + MockLayerBuilder::default() + .writer_factory(factory) + .build() + .unwrap(), + ); + let mut writer = ParquetFileWriter::open(batch().schema(), store.clone(), path, 1, None) + .await + .unwrap(); + writer.write(batch(), None).await.unwrap(); + assert!(writer.finish(None).await.is_err()); + assert!(store.exists(path).await.unwrap()); + writer.abort().await.unwrap(); + assert!(!store.exists(path).await.unwrap()); + } + + #[tokio::test] + async fn large_footer_flush_uses_bounded_submissions_in_one_parquet_stream() { + use std::sync::Mutex; + + use object_store::layers::mock::{Metadata, MockLayerBuilder, MockWriterFactory, oio}; + struct ObservedWriter(oio::Writer, Arc>>); + impl oio::Write for ObservedWriter { + async fn write(&mut self, bytes: object_store::Buffer) -> object_store::Result<()> { + self.1.lock().unwrap().push(bytes.len()); + self.0.write(bytes).await + } + async fn close(&mut self) -> object_store::Result { + self.0.close().await + } + async fn abort(&mut self) -> object_store::Result<()> { + self.0.abort().await + } + } + let directory = common_test_util::temp_dir::create_temp_dir("bounded_parquet"); + let store = object_store::secure_fs::SecureFsRoot::open(directory.path()) + .unwrap() + .build_operator(); + let sizes = Arc::new(Mutex::new(Vec::new())); + let factory: MockWriterFactory = Arc::new({ + let sizes = sizes.clone(); + move |_, _, writer| Box::new(ObservedWriter(writer, sizes.clone())) + }); + let store = store.layer( + MockLayerBuilder::default() + .writer_factory(factory) + .build() + .unwrap(), + ); + let mut state = 17u64; + let values = (0..600_000) + .map(|_| { + state ^= state << 13; + state ^= state >> 7; + state ^= state << 17; + state as i64 + }) + .collect::>(); + let array = Arc::new(Int64Array::from(values)) as arrow::array::ArrayRef; + let batch = RecordBatch::try_from_iter([("a", array.clone()), ("b", array)]).unwrap(); + let mut writer = + ParquetFileWriter::open(batch.schema(), store.clone(), "large.parquet", 1, None) + .await + .unwrap(); + // No storage-layer chunking: observe the application's actual submissions. + writer.sink = store.writer("large.parquet").await.unwrap(); + writer.write(batch.clone(), None).await.unwrap(); + writer.finish(None).await.unwrap(); + let sizes = sizes.lock().unwrap().clone(); + let limit = DEFAULT_WRITE_BUFFER_SIZE.as_bytes() as usize; + assert!(sizes.iter().sum::() > limit); + assert!(sizes.iter().all(|size| *size <= limit), "{sizes:?}"); + let actual = read(&store, "large.parquet") + .await + .build() + .unwrap() + .collect::, _>>() + .unwrap(); + assert_eq!( + arrow::compute::concat_batches(&batch.schema(), &actual).unwrap(), + batch + ); + } + #[tokio::test] async fn row_and_byte_limits_split_batches() { let store = ObjectStore::new(object_store::services::Memory::default()).unwrap(); diff --git a/src/frontend/src/instance.rs b/src/frontend/src/instance.rs index 92bc74ca10b..6e531214127 100644 --- a/src/frontend/src/instance.rs +++ b/src/frontend/src/instance.rs @@ -2763,12 +2763,12 @@ mod tests { ) } - struct ReleasedExportSource { + struct CancellableExportSource { schema: GtSchemaRef, channels: std::sync::Mutex, oneshot::Receiver<()>)>>, } - impl DataSource for ReleasedExportSource { + impl DataSource for CancellableExportSource { fn get_stream( &self, _request: ScanRequest, @@ -2788,15 +2788,15 @@ mod tests { } #[tokio::test] - async fn test_metric_export_http_timeout_drains_ordinary_writer() { + async fn test_metric_export_http_timeout_cancels_ordinary_source() { let destination = common_test_util::temp_dir::create_temp_dir("metric_export_timeout"); let (started_tx, started_rx) = oneshot::channel(); - let (release_tx, release_rx) = oneshot::channel(); + let (mut release_tx, release_rx) = oneshot::channel(); let info = test_table_info(1024, "source").unwrap(); let source = Arc::new(Table::new( Arc::new(info.clone()), FilterPushDownType::Unsupported, - Arc::new(ReleasedExportSource { + Arc::new(CancellableExportSource { schema: info.meta.schema.clone(), channels: std::sync::Mutex::new(Some((started_tx, release_rx))), }), @@ -2834,9 +2834,9 @@ mod tests { let server_task = tokio::spawn(async move { axum::serve(listener, app).await.unwrap(); }); + let uri = reqwest::Url::from_directory_path(destination.path()).unwrap(); let sql = format!( - "COPY DATABASE greptime.public TO '{}/' WITH (experimental_metric_export='true')", - destination.path().display() + "COPY DATABASE greptime.public TO '{uri}' WITH (experimental_metric_export='true')" ); let request = tokio::spawn(async move { reqwest::Client::builder() @@ -2855,21 +2855,10 @@ mod tests { .unwrap(); let response = request.await.unwrap(); assert_eq!(response.status(), reqwest::StatusCode::REQUEST_TIMEOUT); - // A dropped ordinary stream would close this receiver before release. - release_tx.send(()).unwrap(); - let file = destination.path().join("source.parquet"); - tokio::time::timeout(Duration::from_secs(5), async { - loop { - if std::fs::read(&file) - .is_ok_and(|bytes| bytes.len() > 8 && bytes.ends_with(b"PAR1")) - { - break; - } - tokio::time::sleep(Duration::from_millis(10)).await; - } - }) - .await - .unwrap(); + tokio::time::timeout(Duration::from_secs(5), release_tx.closed()) + .await + .unwrap(); + assert!(!destination.path().join("source.parquet").exists()); server_task.abort(); } diff --git a/src/object-store/src/secure_fs.rs b/src/object-store/src/secure_fs.rs index 18afa3ac820..feaaa42fdf2 100644 --- a/src/object-store/src/secure_fs.rs +++ b/src/object-store/src/secure_fs.rs @@ -274,6 +274,7 @@ impl Service for SecureFsBackend { path, args, file: None, + synced: false, }) } @@ -387,6 +388,29 @@ struct SecureFsWriter { path: PathBuf, args: OpWrite, file: Option, + synced: bool, +} + +#[derive(Debug)] +struct UnsyncedOverwrite { + flush_error: Option, +} + +impl fmt::Display for UnsyncedOverwrite { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "overwrite file was not synced") + } +} + +impl std::error::Error for UnsyncedOverwrite { + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + self.flush_error.as_ref().map(|error| error as _) + } +} + +/// Whether an abort error proves an opened overwrite has not been synced. +pub fn is_unsynced_overwrite_abort(error: &Error) -> bool { + std::error::Error::source(error).is_some_and(|source| source.is::()) } impl SecureFsWriter { @@ -438,10 +462,18 @@ impl oio::Write for SecureFsWriter { } async fn close(&mut self) -> Result { - let file = self.ensure_file().await?; - file.flush().await.map_err(new_std_io_error)?; - file.sync_all().await.map_err(new_std_io_error)?; - let metadata = file.metadata().await.map_err(new_std_io_error)?; + { + let file = self.ensure_file().await?; + file.flush().await.map_err(new_std_io_error)?; + file.sync_all().await.map_err(new_std_io_error)?; + } + self.synced = true; + let metadata = self + .ensure_file() + .await? + .metadata() + .await + .map_err(new_std_io_error)?; let mut builder = MetadataBuilder::file(metadata.len()); builder.last_modified(Timestamp::try_from( metadata.modified().map_err(new_std_io_error)?, @@ -450,10 +482,43 @@ impl oio::Write for SecureFsWriter { } async fn abort(&mut self) -> Result<()> { - Err(Error::new( + // Tokio writes may finish in the blocking pool after write_all returns. + let flush = match self.file.as_mut() { + Some(file) => file.flush().await.map_err(new_std_io_error), + None => Ok(()), + }; + if self.args.if_not_exists() { + // A failed exclusive create owns no file. Once data is synced, preserve + // potentially committed output for the caller's deliberate retry. + if let Some(file) = self.file.take() { + drop(file); + if !self.synced { + let root = self.root.clone(); + let path = self.path.clone(); + let cleanup = + common_runtime::spawn_blocking_global(move || root.dir.remove_file(path)) + .await + .map_err(new_task_join_error) + .and_then(|result| result.map_err(new_std_io_error)); + return flush.and(cleanup); + } + } + return flush; + } + if self.file.is_none() { + return flush; + } + let error = Error::new( ErrorKind::Unsupported, "filesystem writes cannot be aborted without atomic writes", - )) + ); + if !self.synced { + return Err(error.set_source(UnsyncedOverwrite { + flush_error: flush.err(), + })); + } + flush?; + Err(error) } } @@ -800,16 +865,181 @@ mod tests { } #[tokio::test] - async fn test_writer_abort_is_unsupported_without_atomic_write() { - let temp_dir = create_temp_dir("secure_fs_writer_abort"); + async fn test_conditional_abort_only_removes_owned_partial_file() { + let temp_dir = create_temp_dir("secure_fs_conditional_abort"); let operator = SecureFsRoot::open(temp_dir.path()) .unwrap() .build_operator(); - let mut writer = operator.writer("partial").await.unwrap(); - writer.write(Bytes::from_static(b"partial")).await.unwrap(); + for started in [false, true] { + let mut writer = operator + .writer_with("partial") + .if_not_exists(true) + .await + .unwrap(); + if started { + writer.write(Bytes::from_static(b"partial")).await.unwrap(); + } + writer.abort().await.unwrap(); + assert!(!operator.exists("partial").await.unwrap()); + } + } - let error = writer.abort().await.unwrap_err(); + #[tokio::test] + async fn test_overwrite_abort_preserves_unopened_destination() { + use std::path::PathBuf; - assert_eq!(ErrorKind::Unsupported, error.kind()); + use opendal::raw::OpWrite; + use opendal::raw::oio::Write; + + let temp_dir = create_temp_dir("secure_fs_unopened_abort"); + std::fs::write(temp_dir.path().join("existing"), b"original").unwrap(); + let mut writer = super::SecureFsWriter { + root: SecureFsRoot::open(temp_dir.path()).unwrap(), + path: PathBuf::from("existing"), + args: OpWrite::default(), + file: None, + synced: false, + }; + + writer.abort().await.unwrap(); + assert_eq!( + std::fs::read(temp_dir.path().join("existing")).unwrap(), + b"original" + ); + } + + #[cfg(target_os = "linux")] + #[tokio::test] + async fn test_abort_after_background_write_error() { + use std::path::PathBuf; + + use opendal::options::WriteOptions; + use opendal::raw::OpWrite; + use opendal::raw::oio::Write; + use tokio::io::AsyncWriteExt; + + for if_not_exists in [true, false] { + let temp_dir = create_temp_dir("secure_fs_flush_error_abort"); + let root = SecureFsRoot::open(temp_dir.path()).unwrap(); + let (args, _) = OpWrite::from_options( + &root.build_operator().info().capability(), + WriteOptions { + if_not_exists, + ..Default::default() + }, + ) + .unwrap(); + let mut writer = super::SecureFsWriter { + root, + path: PathBuf::from("partial"), + args, + file: None, + synced: false, + }; + writer.ensure_file().await.unwrap(); + let mut failing_file = tokio::fs::OpenOptions::new() + .write(true) + .open("/dev/full") + .await + .unwrap(); + failing_file.write_all(b"partial").await.unwrap(); + writer.file = Some(failing_file); + + let error = writer.abort().await.unwrap_err(); + if if_not_exists { + assert!(!temp_dir.path().join("partial").exists()); + } else { + assert_eq!(error.kind(), opendal::ErrorKind::Unsupported); + assert!(std::error::Error::source(&error).is_some()); + } + } + } + + #[cfg(target_os = "linux")] + #[tokio::test] + async fn test_conditional_abort_after_close_flush_error() { + use std::path::PathBuf; + + use opendal::options::WriteOptions; + use opendal::raw::OpWrite; + use opendal::raw::oio::Write; + use tokio::io::AsyncWriteExt; + + let temp_dir = create_temp_dir("secure_fs_close_flush_error"); + let root = SecureFsRoot::open(temp_dir.path()).unwrap(); + let (args, _) = OpWrite::from_options( + &root.build_operator().info().capability(), + WriteOptions { + if_not_exists: true, + ..Default::default() + }, + ) + .unwrap(); + let mut writer = super::SecureFsWriter { + root, + path: PathBuf::from("partial"), + args, + file: None, + synced: false, + }; + writer.ensure_file().await.unwrap(); + let mut failing_file = tokio::fs::OpenOptions::new() + .write(true) + .open("/dev/full") + .await + .unwrap(); + failing_file.write_all(b"partial").await.unwrap(); + writer.file = Some(failing_file); + + assert!(writer.close().await.is_err()); + writer.abort().await.unwrap(); + assert!(!temp_dir.path().join("partial").exists()); + } + + #[test] + fn test_writer_abort_drains_before_reporting_unsupported() { + use std::path::PathBuf; + + use opendal::Buffer; + use opendal::raw::OpWrite; + use opendal::raw::oio::Write; + + use super::SecureFsWriter; + + tokio::runtime::Builder::new_current_thread() + .enable_all() + .max_blocking_threads(1) + .build() + .unwrap() + .block_on(async { + let temp_dir = create_temp_dir("secure_fs_writer_abort"); + let mut writer = SecureFsWriter { + root: SecureFsRoot::open(temp_dir.path()).unwrap(), + path: PathBuf::from("partial"), + args: OpWrite::default(), + file: None, + synced: false, + }; + writer.ensure_file().await.unwrap(); + let (started, ready) = tokio::sync::oneshot::channel(); + let (release, blocked) = std::sync::mpsc::channel(); + let blocking = tokio::task::spawn_blocking(move || { + started.send(()).unwrap(); + blocked.recv().unwrap(); + }); + ready.await.unwrap(); + writer.write(Buffer::from("partial")).await.unwrap(); + let abort = writer.abort(); + tokio::pin!(abort); + let pending = futures::poll!(&mut abort).is_pending(); + release.send(()).unwrap(); + assert!(pending); + assert_eq!(ErrorKind::Unsupported, abort.await.unwrap_err().kind()); + blocking.await.unwrap(); + assert_eq!( + std::fs::read(temp_dir.path().join("partial")).unwrap(), + b"partial" + ); + }); } } diff --git a/src/object-store/tests/object_store_test.rs b/src/object-store/tests/object_store_test.rs index 735d2c56f15..dc2482d1722 100644 --- a/src/object-store/tests/object_store_test.rs +++ b/src/object-store/tests/object_store_test.rs @@ -290,6 +290,9 @@ async fn test_s3_backend() -> Result<()> { .secret_access_key(&env::var("GT_S3_ACCESS_KEY")?) .region(&env::var("GT_S3_REGION")?) .bucket(&bucket); + if let Ok(endpoint) = env::var("GT_S3_ENDPOINT_URL") { + builder = builder.endpoint(&endpoint); + } // Honors an S3-compatible endpoint (MinIO in CI) so the test does not // fall back to resolving the bucket against real AWS. @@ -306,6 +309,7 @@ async fn test_s3_backend() -> Result<()> { test_object_crud(&store).await?; test_object_list(&store).await?; test_object_list_start_after(&store).await?; + test_conditional_creation(&store).await?; assert_opendal_metrics(); guard.remove_all().await?; } @@ -313,6 +317,74 @@ async fn test_s3_backend() -> Result<()> { Ok(()) } +async fn test_conditional_creation(store: &ObjectStore) -> Result<()> { + const PART: usize = 8 * 1024 * 1024; + for size in [32, 2 * PART + 17] { + let payload = Bytes::from(vec![42; size]); + for preexisting in [false, true] { + let path = format!("conditional-{size}-{preexisting}"); + if preexisting { + store.write(&path, "original").await?; + } + let attempt = |value: Bytes| { + let path = &path; + async move { + let mut writer = store + .writer_with(path) + .if_not_exists(true) + .concurrent(1) + .chunk(PART) + .await + .unwrap(); + let result = async { + for offset in (0..value.len()).step_by(PART) { + writer + .write(value.slice(offset..(offset + PART).min(value.len()))) + .await?; + } + writer.close().await + } + .await; + if result.is_err() { + writer.abort().await.unwrap(); + } + result + } + }; + let (a, b) = tokio::join!(attempt(payload.clone()), attempt(payload.clone())); + assert_eq!( + usize::from(a.is_ok()) + usize::from(b.is_ok()), + usize::from(!preexisting) + ); + for error in [a.err(), b.err()].into_iter().flatten() { + assert!( + matches!( + error.kind(), + object_store::ErrorKind::ConditionNotMatch + | object_store::ErrorKind::AlreadyExists + ), + "{error:?}" + ); + } + let expected = if preexisting { + Bytes::from_static(b"original") + } else { + payload.clone() + }; + assert_eq!(store.read(&path).await?.to_bytes(), expected); + store.delete(&path).await?; + } + } + Ok(()) +} + +#[tokio::test] +async fn test_secure_fs_conditional_creation() -> Result<()> { + let dir = TempDir::new()?; + let store = object_store::secure_fs::SecureFsRoot::open(dir.path())?.build_operator(); + test_conditional_creation(&store).await +} + #[tokio::test] async fn test_oss_backend() -> Result<()> { common_telemetry::init_default_ut_logging(); diff --git a/src/operator/src/statement/copy_table_to.rs b/src/operator/src/statement/copy_table_to.rs index 21b82f51a3b..d3d44f7655f 100644 --- a/src/operator/src/statement/copy_table_to.rs +++ b/src/operator/src/statement/copy_table_to.rs @@ -15,6 +15,8 @@ use std::collections::HashMap; use std::sync::Arc; +use arrow::array::{Array, ArrayRef, AsArray, UInt64Array, make_array}; +use arrow::datatypes::DataType; use client::OutputData; use common_base::readable_size::ReadableSize; use common_datasource::file_format::Format; @@ -22,6 +24,7 @@ use common_datasource::file_format::csv::stream_to_csv; use common_datasource::file_format::json::stream_to_json; use common_datasource::file_format::parquet::stream_to_parquet; use common_datasource::object_store::build_backend_for_write_with_path; +use common_datasource::parquet_writer::ParquetFileWriter; use common_query::Output; use common_recordbatch::adapter::DfRecordBatchStreamAdapter; use common_recordbatch::{ @@ -32,16 +35,22 @@ use common_telemetry::{debug, tracing}; use datafusion::datasource::DefaultTableSource; use datafusion_common::TableReference as DfTableReference; use datafusion_expr::LogicalPlanBuilder; +use futures::StreamExt; use object_store::ObjectStore; use session::context::QueryContextRef; -use snafu::{OptionExt, ResultExt}; +use snafu::{OptionExt, ResultExt, ensure}; use table::TableRef; use table::requests::CopyTableRequest; use table::table::adapter::DfTableProviderAdapter; use table::table_reference::TableReference; +use tokio_util::sync::CancellationToken; use crate::error::{self, BuildDfLogicalPlanSnafu, ExecLogicalPlanSnafu, Result}; use crate::statement::StatementExecutor; +use crate::statement::export_logical_tables::writers::ExportWriteBudget; +use crate::statement::export_logical_tables::{ + expand_export_batch, map_writer_error, rows_within_budget, +}; // The buffer size should be greater than 5MB (minimum multipart upload size). /// Buffer size to flush data to object stores. @@ -117,6 +126,17 @@ impl StatementExecutor { table: TableRef, req: CopyTableRequest, query_ctx: QueryContextRef, + ) -> Result { + self.copy_captured_table_to_managed(table, req, query_ctx, None) + .await + } + + pub(crate) async fn copy_captured_table_to_managed( + &self, + table: TableRef, + req: CopyTableRequest, + query_ctx: QueryContextRef, + managed: Option<(&ExportWriteBudget, &CancellationToken)>, ) -> Result { let info = table.table_info(); let table_ref = TableReference::full(&info.catalog_name, &info.schema_name, &info.name); @@ -165,7 +185,7 @@ impl StatementExecutor { } = &req; debug!("Copy table: {table_id} to location: {location}"); - self.copy_to_file(&format, output, location, connection) + self.copy_to_file_managed(&format, output, location, connection, managed) .await } @@ -176,9 +196,25 @@ impl StatementExecutor { location: &str, connection: &HashMap, ) -> Result { - let output = output - .map_dictionary_to_values() - .context(error::BuildRecordBatchSnafu)?; + self.copy_to_file_managed(format, output, location, connection, None) + .await + } + + async fn copy_to_file_managed( + &self, + format: &Format, + output: Output, + location: &str, + connection: &HashMap, + managed: Option<(&ExportWriteBudget, &CancellationToken)>, + ) -> Result { + let output = if managed.is_none() { + output + .map_dictionary_to_values() + .context(error::BuildRecordBatchSnafu)? + } else { + output + }; let stream = match output.data { OutputData::Stream(stream) => stream, OutputData::RecordBatches(record_batches) => record_batches.as_stream(), @@ -192,7 +228,408 @@ impl StatementExecutor { let filename = backend.object_path.context(error::UnexpectedSnafu { violated: format!("Expected filename, path: {location}"), })?; - self.stream_to_file(stream, format, backend.object_store, &filename) + if let Some((budget, token)) = managed { + stream_to_managed_parquet(stream, backend.object_store, &filename, budget, token).await + } else { + self.stream_to_file(stream, format, backend.object_store, &filename) + .await + } + } +} + +pub(crate) async fn stream_to_managed_parquet( + mut stream: SendableRecordBatchStream, + store: ObjectStore, + path: &str, + budget: &ExportWriteBudget, + token: &CancellationToken, +) -> Result { + use common_recordbatch::{RecordBatch, map_dictionary_to_values_schema}; + let original = stream.schema(); + let (expanded_schema, expand) = map_dictionary_to_values_schema(original.clone()); + let json_columns = expanded_schema + .column_schemas() + .iter() + .enumerate() + .filter_map(|(index, column)| { + (column.data_type.is_json() && !column.data_type.is_json2()).then_some(index) + }) + .collect::>(); + let (mapped_schema, json) = map_json_type_to_string_schema(expanded_schema.clone()); + let output_schema = if json { + mapped_schema + } else { + expanded_schema.clone() + }; + let mut writer = + ParquetFileWriter::open(output_schema.arrow_schema().clone(), store, path, 1, None) .await + .context(error::WriteStreamToFileSnafu { path })?; + let mut started = false; + let result = async { + let mut rows = 0; + loop { + let batch = tokio::select! { + biased; + _ = token.cancelled() => return error::LogicalTableExportCancelledSnafu.fail(), + batch = stream.next() => batch, + }; + let Some(batch) = batch else { break }; + let batch = batch + .context(error::BuildRecordBatchSnafu)? + .into_df_record_batch(); + let mut offset = 0; + while offset < batch.num_rows() { + // The scan batch remains query-owned. Bound the downstream slice and + // detach its buffers below if it would retain more than its reservation. + let (conversion, retained) = + ExportWriteBudget::conversion_budget(0, batch.num_columns(), usize::MAX)?; + let input = batch.clone(); + let json_columns = json_columns.clone(); + let (len, estimated) = common_runtime::spawn_blocking_global(move || { + rows_within_budget(&input, offset, input.num_rows(), conversion, &json_columns) + }) + .await + .context(error::JoinTaskSnafu)??; + let reservation = retained.saturating_add(estimated.saturating_mul(4)); + let permit = budget.reserve(reservation, token).await?; + let batch = batch.clone(); + let (original, expanded_schema, output_schema) = ( + original.clone(), + expanded_schema.clone(), + output_schema.clone(), + ); + let (batch, len) = common_runtime::spawn_blocking_global(move || { + let mut batch = RecordBatch::from_df_record_batch( + original.clone(), + batch.slice(offset, len), + ); + if expand { + batch = RecordBatch::from_df_record_batch( + expanded_schema.clone(), + expand_export_batch( + &batch.into_df_record_batch(), + expanded_schema.arrow_schema().clone(), + )?, + ); + } + if json { + batch = map_json_type_to_string(batch, &expanded_schema, &output_schema) + .context(error::BuildRecordBatchSnafu)?; + } + let mut batch = batch.into_df_record_batch(); + if batch.get_array_memory_size() > reservation { + let indices = UInt64Array::from_iter_values(0..batch.num_rows() as u64); + batch = arrow::compute::take_record_batch(&batch, &indices) + .context(error::ComputeArrowSnafu)?; + let arrays = batch + .columns() + .iter() + .map(compact_view_buffers) + .collect::>>()?; + batch = arrow::record_batch::RecordBatch::try_new(batch.schema(), arrays) + .context(error::ComputeArrowSnafu)?; + } + ensure!( + batch.get_array_memory_size() <= reservation, + error::LogicalTableExportResourceSnafu { + reason: "converted backing buffers exceed reservation" + } + ); + Ok::<_, error::Error>((batch, len)) + }) + .await + .context(error::JoinTaskSnafu)??; + started = true; + let write = writer.write(batch, Some(token)).await; + drop(permit); + write.map_err(|error| map_writer_error(error, path))?; + rows += len; + offset += len; + } + } + started = true; + writer + .finish(Some(token)) + .await + .map_err(|error| map_writer_error(error, path))?; + Ok(rows) + } + .await; + if result.is_err() { + token.cancel(); + if started && let Err(error) = writer.abort().await { + common_telemetry::warn!(error; "Failed to abort ordinary export file"); + } + } + result +} + +// Arrow take copies ordinary buffers but shares view data buffers, including +// views nested inside lists or structs. Reclaim those unselected values too. +fn compact_view_buffers(array: &ArrayRef) -> Result { + match array.data_type() { + DataType::Utf8View => Ok(Arc::new(array.as_string_view().gc())), + DataType::BinaryView => Ok(Arc::new(array.as_binary_view().gc())), + _ => { + let data = array.to_data(); + if data.child_data().is_empty() { + return Ok(array.clone()); + } + let children = data + .child_data() + .iter() + .map(|child| { + compact_view_buffers(&make_array(child.clone())).map(|array| array.to_data()) + }) + .collect::>>()?; + Ok(make_array( + data.into_builder() + .child_data(children) + .build() + .context(error::ComputeArrowSnafu)?, + )) + } + } +} + +#[cfg(test)] +mod tests { + use arrow::array::{ + ArrayRef, BinaryArray, BinaryViewArray, DictionaryArray, Int32Array, StringArray, + StringViewArray, + }; + use arrow::datatypes::Int32Type; + use common_recordbatch::{RecordBatch, RecordBatches}; + use datafusion::parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder; + use datatypes::prelude::ConcreteDataType; + use datatypes::schema::{ColumnSchema, Schema}; + + use super::*; + + #[tokio::test] + async fn managed_copy_preserves_existing_file_before_sink_open() { + let temp_dir = common_test_util::temp_dir::create_temp_dir("managed_copy_existing_file"); + let store = object_store::secure_fs::SecureFsRoot::open(temp_dir.path()) + .unwrap() + .build_operator(); + let path = "existing.parquet"; + store.write(path, "original").await.unwrap(); + let schema = Arc::new(Schema::new(vec![ColumnSchema::new( + "value", + ConcreteDataType::int32_datatype(), + false, + )])); + let batch = arrow::record_batch::RecordBatch::try_new( + schema.arrow_schema().clone(), + vec![Arc::new(Int32Array::from(vec![1]))], + ) + .unwrap(); + let batches = async_stream::stream! { + yield Ok(batch); + yield Err(datafusion::error::DataFusionError::Execution("source failed".into())); + }; + let stream = datafusion::physical_plan::stream::RecordBatchStreamAdapter::new( + schema.arrow_schema().clone(), + batches, + ); + let stream = + common_recordbatch::adapter::RecordBatchStreamAdapter::try_new(Box::pin(stream)) + .unwrap(); + + let result = stream_to_managed_parquet( + Box::pin(stream), + store.clone(), + path, + &ExportWriteBudget::new(1), + &CancellationToken::new(), + ) + .await; + + assert!(result.is_err()); + assert_eq!( + store.read(path).await.unwrap().to_bytes().as_ref(), + b"original" + ); + } + + #[tokio::test] + async fn managed_copy_rechunks_large_backing_buffers() { + let value = "x".repeat(9 * 1024); + for view in [false, true] { + let data_type = if view { + ConcreteDataType::utf8_view_datatype() + } else { + ConcreteDataType::string_datatype() + }; + let schema = Arc::new(Schema::new(vec![ColumnSchema::new( + "value", data_type, false, + )])); + let values = vec![value.as_str(); 8192]; + let array: ArrayRef = if view { + Arc::new(StringViewArray::from(values)) + } else { + Arc::new(StringArray::from(values)) + }; + let batch = arrow::record_batch::RecordBatch::try_new( + schema.arrow_schema().clone(), + vec![array], + ) + .unwrap(); + assert!(batch.get_array_memory_size() > 64 * 1024 * 1024); + // Also exercise a small slice that still references the large allocation. + for input in [batch.clone(), batch.slice(200, 3)] { + let rows = input.num_rows(); + let batches = RecordBatches::try_new( + schema.clone(), + vec![RecordBatch::from_df_record_batch(schema.clone(), input)], + ) + .unwrap(); + let store = ObjectStore::new(object_store::services::Memory::default()).unwrap(); + let budget = ExportWriteBudget::new(1); + assert_eq!( + stream_to_managed_parquet( + batches.as_stream(), + store.clone(), + "large.parquet", + &budget, + &CancellationToken::new(), + ) + .await + .unwrap(), + rows + ); + let reader = ParquetRecordBatchReaderBuilder::try_new( + store.read("large.parquet").await.unwrap().to_bytes(), + ) + .unwrap() + .build() + .unwrap(); + let mut actual_rows = 0; + for batch in reader { + let batch = batch.unwrap(); + let strings = arrow::compute::cast(batch.column(0), &DataType::Utf8).unwrap(); + assert!( + strings + .as_string::() + .iter() + .all(|item| item == Some(value.as_str())) + ); + actual_rows += batch.num_rows(); + } + assert_eq!(actual_rows, rows); + assert_eq!(budget.available(), (1, 64 * 1024 * 1024)); + } + } + } + + #[tokio::test] + async fn managed_copy_preserves_dictionary_json_views_and_empty_schema() { + let schema = Arc::new(Schema::new(vec![ + ColumnSchema::new( + "host", + ConcreteDataType::dictionary_datatype( + ConcreteDataType::int32_datatype(), + ConcreteDataType::string_datatype(), + ), + true, + ), + ColumnSchema::new("json", ConcreteDataType::json_datatype(), true), + ColumnSchema::new("text_view", ConcreteDataType::utf8_view_datatype(), true), + ColumnSchema::new( + "binary_view", + ConcreteDataType::binary_view_datatype(), + true, + ), + ])); + let dictionary = DictionaryArray::::new( + Int32Array::from(vec![Some(0), None, Some(1), Some(1)]), + Arc::new(StringArray::from(vec!["tag", &"x".repeat(3 * 1024 * 1024)])), + ); + let json = datatypes::types::parse_string_to_jsonb(&format!( + r#"{{"value":"escaped\ntext{}","n":123456789}}"#, + "y".repeat(160_000) + )) + .unwrap(); + let arrays = vec![ + Arc::new(dictionary) as ArrayRef, + Arc::new(BinaryArray::from(vec![ + Some(json.as_slice()), + None, + None, + Some(json.as_slice()), + ])), + Arc::new(StringViewArray::from(vec![ + Some("short"), + None, + Some("long view"), + Some("long view"), + ])), + Arc::new(BinaryViewArray::from(vec![ + Some(&b"short"[..]), + None, + Some(&b"long view"[..]), + Some(&b"long view"[..]), + ])), + ]; + let batch = + arrow::record_batch::RecordBatch::try_new(schema.arrow_schema().clone(), arrays) + .unwrap(); + let batch = RecordBatch::from_df_record_batch(schema.clone(), batch); + for empty in [false, true] { + let batches = RecordBatches::try_new( + schema.clone(), + if empty { vec![] } else { vec![batch.clone()] }, + ) + .unwrap(); + let store = ObjectStore::new(object_store::services::Memory::default()).unwrap(); + let budget = ExportWriteBudget::new(1); + let token = CancellationToken::new(); + let managed = stream_to_managed_parquet( + batches.as_stream(), + store.clone(), + "managed.parquet", + &budget, + &token, + ) + .await + .unwrap(); + let output = Output::new_with_record_batches(batches) + .map_dictionary_to_values() + .unwrap(); + let OutputData::RecordBatches(batches) = output.data else { + panic!("expected batches"); + }; + let stream = SendableRecordBatchMapper::new( + batches.as_stream(), + map_json_type_to_string, + map_json_type_to_string_schema, + ); + let ordinary = stream_to_parquet( + Box::pin(DfRecordBatchStreamAdapter::new(Box::pin(stream))), + store.clone(), + "ordinary.parquet", + WRITE_CONCURRENCY, + ) + .await + .unwrap(); + assert_eq!(managed, ordinary); + let mut outputs = Vec::new(); + for name in ["managed.parquet", "ordinary.parquet"] { + let reader = ParquetRecordBatchReaderBuilder::try_new( + store.read(name).await.unwrap().to_bytes(), + ) + .unwrap(); + let schema = reader.schema().clone(); + let values = reader + .build() + .unwrap() + .collect::, _>>() + .unwrap(); + outputs.push((schema, values)); + } + assert_eq!(outputs[0], outputs[1]); + assert_eq!(budget.available(), (1, 64 * 1024 * 1024)); + } } } diff --git a/src/operator/src/statement/database_copy.rs b/src/operator/src/statement/database_copy.rs index 117104b9d8f..9878560afaf 100644 --- a/src/operator/src/statement/database_copy.rs +++ b/src/operator/src/statement/database_copy.rs @@ -27,6 +27,7 @@ use table::TableRef; use table::metadata::TableType; use table::requests::CopyDatabaseRequest; use table::table_reference::TableReference; +use tokio::sync::Semaphore; use url::Url; use crate::error::{self, Result}; @@ -56,7 +57,7 @@ pub(crate) fn parse_parallelism_from_option_map(options: &HashMap().ok()) .unwrap_or_else(get_total_cpu_cores) - .max(1) + .clamp(1, Semaphore::MAX_PERMITS) } /// Rejects import-only layouts before either database export path creates output. @@ -321,5 +322,11 @@ mod tests { let options = HashMap::from([("parallelism".to_string(), "0".to_string())]); assert_eq!(parse_parallelism_from_option_map(&options), 1); + + let options = HashMap::from([("parallelism".to_string(), usize::MAX.to_string())]); + assert_eq!( + parse_parallelism_from_option_map(&options), + Semaphore::MAX_PERMITS + ); } } diff --git a/src/operator/src/statement/export_database.rs b/src/operator/src/statement/export_database.rs index 7c5de554487..dbe18ba763a 100644 --- a/src/operator/src/statement/export_database.rs +++ b/src/operator/src/statement/export_database.rs @@ -37,6 +37,7 @@ use crate::statement::database_copy::{ DatabaseExportFile, parse_parallelism_from_option_map, validate_database_directory, validate_database_export_layout, }; +use crate::statement::export_logical_tables::writers::{ExportWriteBudget, retain_error}; use crate::statement::export_logical_tables::{LogicalTableExport, LogicalTableExportLimits}; /// A validated request-scoped selection, not a metadata snapshot or an ACL token. @@ -202,46 +203,52 @@ impl StatementExecutor { } output_files.sort(); let req = &plan.request; - let rows = run_database_export_jobs( - plan.jobs, - parse_parallelism_from_option_map(&req.with), - cancellation, - |job, token| { - let ctx = ctx.clone(); - async move { - match job { - DatabaseExportJob::Metric(unit) => self - .export_logical_tables( - &unit, - &req.location, - &req.connection, - req.time_range.as_ref(), - LogicalTableExportLimits::default(), - &token, - ctx, - ) - .await - .map(|summary| summary.rows), - DatabaseExportJob::Ordinary { table, output } => { - let info = table.table_info(); - let copy = CopyTableRequest { - catalog_name: info.catalog_name.clone(), - schema_name: info.schema_name.clone(), - table_name: info.name.clone(), - location: output.location, - with: req.with.clone(), - connection: req.connection.clone(), - pattern: None, - direction: CopyDirection::Export, - timestamp_range: req.time_range, - limit: None, - }; - self.copy_captured_table_to(table, copy, ctx).await - } + let parallelism = parse_parallelism_from_option_map(&req.with); + let budget = ExportWriteBudget::new(parallelism); + let rows = run_database_export_jobs(plan.jobs, parallelism, cancellation, |job, token| { + let ctx = ctx.clone(); + let budget = budget.clone(); + async move { + match job { + DatabaseExportJob::Metric(unit) => self + .export_logical_tables_managed( + &unit, + &req.location, + &req.connection, + req.time_range.as_ref(), + LogicalTableExportLimits::default(), + &token, + ctx, + budget, + ) + .await + .map(|summary| summary.rows), + DatabaseExportJob::Ordinary { table, output } => { + let info = table.table_info(); + let copy = CopyTableRequest { + catalog_name: info.catalog_name.clone(), + schema_name: info.schema_name.clone(), + table_name: info.name.clone(), + location: output.location, + with: req.with.clone(), + connection: req.connection.clone(), + pattern: None, + direction: CopyDirection::Export, + timestamp_range: req.time_range, + limit: None, + }; + let _permit = budget.writer(&token).await?; + self.copy_captured_table_to_managed( + table, + copy, + ctx, + Some((&budget, &token)), + ) + .await } } - }, - ) + } + }) .await?; Ok(DatabaseExportSummary { rows, output_files }) } @@ -261,6 +268,7 @@ async fn run_database_export_jobs>>( loop { while first_error.is_none() && !cancellation.is_cancelled() + && !token.is_cancelled() && active.len() < parallelism.max(1) { let Some(job) = jobs.next() else { break }; @@ -284,11 +292,12 @@ async fn run_database_export_jobs>>( }; match result { Some(Ok(count)) => rows += count, - Some(Err(err)) if first_error.is_none() => { - first_error = Some(err); + Some(Err(err)) => { + if !cancellation.is_cancelled() || first_error.is_none() { + retain_error(&mut first_error, err); + } token.cancel(); } - Some(Err(err)) => common_telemetry::warn!(err; "Failed to drain database export job"), None => break, } } @@ -339,7 +348,7 @@ mod tests { } .fail(); } - // Ordinary COPY continues its I/O even when Metric jobs cancel. + // Already-started ordinary I/O drains after cancellation. token.cancelled().await; started.send(10).unwrap(); finish_io.acquire().await.unwrap().forget(); diff --git a/src/operator/src/statement/export_logical_tables.rs b/src/operator/src/statement/export_logical_tables.rs index 7547a3f54de..4a2e186d477 100644 --- a/src/operator/src/statement/export_logical_tables.rs +++ b/src/operator/src/statement/export_logical_tables.rs @@ -19,16 +19,21 @@ //! exclusively by the export attempt. Completion and cleanup of closed files //! belong to the caller's chunk protocol, not individual file completion. +pub(crate) mod writers; + use std::collections::{BTreeMap, BTreeSet, HashMap}; use std::sync::Arc; -use arrow::array::{Array, AsArray, UInt32Array}; -use arrow::compute::cast; +use arrow::array::{Array, ArrayRef, AsArray, ListArray, StructArray, UInt32Array}; +use arrow::buffer::OffsetBuffer; +use arrow::compute::{cast, take}; use arrow::datatypes::{DataType, SchemaRef}; use arrow::downcast_dictionary_array; use arrow::record_batch::RecordBatch; use common_datasource::object_store::build_backend_for_write; -use common_datasource::parquet_writer::{ParquetFileWriter, ParquetWriterLimits}; +use common_datasource::parquet_writer::{ + ParquetCreationPolicy, ParquetFileWriter, ParquetWriterLimits, +}; use common_meta::key::table_route::{TableRouteManager, TableRouteValue}; use common_query::OutputData; use common_recordbatch::SendableRecordBatchStream; @@ -52,6 +57,7 @@ use tokio_util::sync::CancellationToken; use crate::error::{self, InvalidLogicalTableExportSnafu, LogicalTableExportResourceSnafu, Result}; use crate::statement::StatementExecutor; use crate::statement::database_copy::DatabaseExportFile; +use crate::statement::export_logical_tables::writers::{ExportWriteBudget, Payload, TableWriters}; /// Export preprocessing and per-file limits. Query memory and spill remain /// governed by the query engine. @@ -101,7 +107,7 @@ pub struct LogicalTableExport { logical_tables: BTreeMap, } -struct LogicalTableProjection { +pub(crate) struct LogicalTableProjection { output: DatabaseExportFile, schema: SchemaRef, projection: Vec, @@ -306,6 +312,31 @@ impl StatementExecutor { limits: LogicalTableExportLimits, cancellation: &CancellationToken, query_ctx: QueryContextRef, + ) -> Result { + self.export_logical_tables_managed( + unit, + directory, + connection, + time_range, + limits, + &cancellation.child_token(), + query_ctx, + ExportWriteBudget::new(1), + ) + .await + } + + #[allow(clippy::too_many_arguments)] + pub(crate) async fn export_logical_tables_managed( + &self, + unit: &LogicalTableExport, + directory: &str, + connection: &HashMap, + time_range: Option<&TimestampRange>, + limits: LogicalTableExportLimits, + cancellation: &CancellationToken, + query_ctx: QueryContextRef, + budget: Arc, ) -> Result { limits.validate()?; let (store, stream) = tokio::select! { @@ -325,10 +356,11 @@ impl StatementExecutor { Ok((store, stream)) } => result?, }; - export_stream(unit, stream, &store, limits, cancellation).await + export_stream_managed(unit, stream, &store, limits, cancellation, budget).await } } +#[cfg(test)] async fn export_stream( unit: &LogicalTableExport, stream: SendableRecordBatchStream, @@ -336,19 +368,42 @@ async fn export_stream( limits: LogicalTableExportLimits, cancellation: &CancellationToken, ) -> Result { - let mut active = None; - let result = write_tables(unit, stream, store, limits, cancellation, &mut active).await; - if result.is_err() - && let Some(writer) = active - && let Err(cleanup_error) = writer - .writer - .abort() - .await - .map_err(|error| map_writer_error(error, &writer.path)) - { - common_telemetry::warn!(cleanup_error; "Failed to clean up incomplete Metric export file"); - } - result + export_stream_managed( + unit, + stream, + store, + limits, + &cancellation.child_token(), + ExportWriteBudget::new(1), + ) + .await +} + +async fn export_stream_managed( + unit: &LogicalTableExport, + stream: SendableRecordBatchStream, + store: &ObjectStore, + limits: LogicalTableExportLimits, + cancellation: &CancellationToken, + budget: Arc, +) -> Result { + let mut writers = TableWriters::new(budget.clone()); + let result = write_tables( + unit, + stream, + store, + limits, + cancellation, + &mut writers, + &budget, + ) + .await; + let (summary, result) = match result { + Ok(summary) => (summary, Ok(())), + Err(error) => (LogicalTableExportSummary::default(), Err(error)), + }; + writers.drain(result, cancellation).await?; + Ok(summary) } fn check_cancelled(cancellation: &CancellationToken) -> Result<()> { @@ -365,7 +420,8 @@ async fn write_tables( store: &ObjectStore, limits: LogicalTableExportLimits, cancellation: &CancellationToken, - active: &mut Option, + writers: &mut TableWriters, + budget: &ExportWriteBudget, ) -> Result { let id_index = stream @@ -423,22 +479,19 @@ async fn write_tables( while end < batch.num_rows() && ids.value(end) == id { end += 1; } - if active.as_ref().is_some_and(|writer| writer.table_id != id) { - finish_active(active, cancellation).await?; + if writers.table_id().is_some_and(|last| last != id) { + writers.close_input(); } check_cancelled(cancellation)?; if let Some(file) = unit.logical_tables.get(&id) { - if active.is_none() { - *active = Some(ActiveWriter::open(id, file, store, limits).await?); + if writers.table_id().is_none() { + writers.open(id, file, store, limits, cancellation).await?; written.insert(id); summary.files += 1; } let projected = batch .project(&file.projection) .context(error::ProjectSchemaSnafu)?; - let writer = active.as_mut().context(error::UnexpectedSnafu { - violated: "missing logical writer", - })?; let mut offset = start; while offset < end { let (expanded, consumed) = expand_bounded_slice( @@ -447,15 +500,12 @@ async fn write_tables( offset, end, limits.conversion_bytes, + budget, + cancellation, ) .await?; check_cancelled(cancellation)?; - // Do not drop an in-flight file operation before cleanup. - writer - .writer - .write(expanded, Some(cancellation)) - .await - .map_err(|error| map_writer_error(error, &writer.path))?; + writers.send(expanded, cancellation).await?; check_cancelled(cancellation)?; offset += consumed; summary.rows += consumed; @@ -466,12 +516,12 @@ async fn write_tables( start = end; } } - finish_active(active, cancellation).await?; + writers.close_input(); for (&id, file) in &unit.logical_tables { if !written.contains(&id) { check_cancelled(cancellation)?; - *active = Some(ActiveWriter::open(id, file, store, limits).await?); - finish_active(active, cancellation).await?; + writers.open(id, file, store, limits, cancellation).await?; + writers.close_input(); summary.files += 1; } } @@ -480,85 +530,158 @@ async fn write_tables( } struct ActiveWriter { - table_id: u32, path: String, writer: ParquetFileWriter, } impl ActiveWriter { async fn open( - id: u32, table: &LogicalTableProjection, store: &ObjectStore, limits: LogicalTableExportLimits, ) -> Result { let path = table.output.path.clone(); - ensure!( - !store - .exists(&path) - .await - .context(error::ReadObjectSnafu { path: &path })?, - InvalidLogicalTableExportSnafu { - reason: format!("output already exists: {path}") - } - ); - let writer = ParquetFileWriter::open( + let conditional = store.info().capability().write_with_if_not_exists; + if !conditional { + ensure!( + !store + .exists(&path) + .await + .context(error::ReadObjectSnafu { path: &path })?, + InvalidLogicalTableExportSnafu { + reason: format!("output already exists: {path}") + } + ); + } + let writer = ParquetFileWriter::open_with_creation( table.schema.clone(), store.clone(), &path, 1, Some(limits.writer), + if conditional { + ParquetCreationPolicy::IfNotExists + } else { + ParquetCreationPolicy::Overwrite + }, ) .await .map_err(|error| map_writer_error(error, &path))?; - Ok(Self { - table_id: id, - path, - writer, - }) + Ok(Self { path, writer }) } } -async fn finish_active( - active: &mut Option, - cancellation: &CancellationToken, -) -> Result<()> { - if let Some(writer) = active.as_mut() { - writer - .writer - .finish(Some(cancellation)) - .await - .map_err(|error| map_writer_error(error, &writer.path))?; - check_cancelled(cancellation)?; - *active = None; - } - Ok(()) -} - async fn expand_bounded_slice( batch: RecordBatch, schema: SchemaRef, start: usize, end: usize, - budget: usize, -) -> Result<(RecordBatch, usize)> { + requested: usize, + budget: &ExportWriteBudget, + cancellation: &CancellationToken, +) -> Result<(Payload, usize)> { + let (conversion, retained) = ExportWriteBudget::conversion_budget( + batch.get_array_memory_size(), + batch.num_columns(), + requested, + )?; + let input = batch.clone(); + let (len, estimated) = common_runtime::spawn_blocking_global(move || { + rows_within_budget(&input, start, end, conversion, &[]) + }) + .await + .context(error::JoinTaskSnafu)??; + let reservation = retained.saturating_add(estimated.saturating_mul(4)); + let permit = budget.reserve(reservation, cancellation).await?; common_runtime::spawn_blocking_global(move || { - let len = rows_within_budget(&batch, start, end, budget)?; let slice = batch.slice(start, len); - let arrays = slice - .columns() - .iter() - .zip(schema.fields()) - .map(|(array, field)| cast(array, field.data_type()).context(error::ComputeArrowSnafu)) - .collect::>>()?; - let expanded = RecordBatch::try_new(schema, arrays).context(error::ComputeArrowSnafu)?; - Ok((expanded, len)) + let expanded = expand_export_batch(&slice, schema)?; + ensure!( + expanded.get_array_memory_size() <= reservation, + LogicalTableExportResourceSnafu { + reason: "converted backing buffers exceed reservation" + } + ); + Ok(( + Payload { + batch: expanded, + permit, + }, + len, + )) }) .await .context(error::JoinTaskSnafu)? } -fn map_writer_error(source: common_datasource::error::Error, path: &str) -> error::Error { +/// Limits nested casts to selected children: Arrow's List cast otherwise expands +/// the entire values array, even when the parent has been sliced. +pub(crate) fn expand_export_batch(batch: &RecordBatch, schema: SchemaRef) -> Result { + let arrays = batch + .columns() + .iter() + .zip(schema.fields()) + .map(|(array, field)| expand_export_array(array, field.data_type())) + .collect::>>()?; + RecordBatch::try_new(schema, arrays).context(error::ComputeArrowSnafu) +} + +fn expand_export_array(array: &ArrayRef, target: &DataType) -> Result { + if array.data_type() == target { + return Ok(array.clone()); + } + match (array.data_type(), target) { + (DataType::Dictionary(_, _), _) => { + downcast_dictionary_array! { + array => { + // Select dictionary entries before recursively converting nested values. + let selected = take(array.values().as_ref(), array.keys(), None) + .context(error::ComputeArrowSnafu)?; + expand_export_array(&selected, target) + }, + _ => error::UnexpectedSnafu { violated: "invalid dictionary array" }.fail(), + } + } + (DataType::List(_), DataType::List(field)) => { + let list = array.as_list::(); + let offsets = list.value_offsets(); + let start = offsets[0]; + let end = offsets[list.len()]; + let values = list.values().slice(start as usize, (end - start) as usize); + let values = expand_export_array(&values, field.data_type())?; + let offsets = OffsetBuffer::new( + offsets + .iter() + .map(|offset| offset - start) + .collect::>() + .into(), + ); + Ok(Arc::new( + ListArray::try_new(field.clone(), offsets, values, list.nulls().cloned()) + .context(error::ComputeArrowSnafu)?, + )) + } + (DataType::Struct(_), DataType::Struct(fields)) => { + let array = array.as_struct(); + let columns = array + .columns() + .iter() + .zip(fields) + .map(|(array, field)| expand_export_array(array, field.data_type())) + .collect::>>()?; + Ok(Arc::new( + StructArray::try_new(fields.clone(), columns, array.nulls().cloned()) + .context(error::ComputeArrowSnafu)?, + )) + } + _ => cast(array, target).context(error::ComputeArrowSnafu), + } +} + +pub(crate) fn map_writer_error( + source: common_datasource::error::Error, + path: &str, +) -> error::Error { match source { common_datasource::error::Error::ParquetWriteCancelled {} => { error::LogicalTableExportCancelledSnafu.build() @@ -572,6 +695,12 @@ fn map_writer_error(source: common_datasource::error::Error, path: &str) -> erro common_datasource::error::Error::ParquetWriterResource { reason } => { LogicalTableExportResourceSnafu { reason }.build() } + common_datasource::error::Error::WriteObject { error, .. } + if error.kind() == object_store::ErrorKind::ConditionNotMatch => + { + let reason = format!("output already exists: {path}"); + InvalidLogicalTableExportSnafu { reason }.build() + } source => error::WriteStreamToFileSnafu { path }.into_error(source), } } @@ -579,7 +708,7 @@ fn map_writer_error(source: common_datasource::error::Error, path: &str) -> erro // Count only selected logical values before dictionary expansion, including nested // histogram lists/structs. Offset and validity overhead is charged per value. fn estimate_value_size(array: &dyn Array, row: usize) -> Result { - if array.is_null(row) { + if array.is_null(row) && !matches!(array.data_type(), DataType::Struct(_) | DataType::List(_)) { return Ok(32); } let bytes = match array.data_type() { @@ -587,8 +716,10 @@ fn estimate_value_size(array: &dyn Array, row: usize) -> Result { DataType::Null => 0, DataType::Utf8 => array.as_string::().value(row).len(), DataType::LargeUtf8 => array.as_string::().value(row).len(), + DataType::Utf8View => array.as_string_view().value(row).len(), DataType::Binary => array.as_binary::().value(row).len(), DataType::LargeBinary => array.as_binary::().value(row).len(), + DataType::BinaryView => array.as_binary_view().value(row).len(), DataType::Struct(_) => { array .as_struct() @@ -629,18 +760,31 @@ fn estimate_value_size(array: &dyn Array, row: usize) -> Result { Ok(bytes.saturating_add(16)) } -fn rows_within_budget( +pub(crate) fn rows_within_budget( batch: &RecordBatch, start: usize, end: usize, budget: usize, -) -> Result { + json_columns: &[usize], +) -> Result<(usize, usize)> { let mut bytes = 0usize; let mut row = start; while row < end { - let row_bytes = batch.columns().iter().try_fold(0usize, |sum, array| { - Ok::<_, error::Error>(sum.saturating_add(estimate_value_size(array.as_ref(), row)?)) - })?; + let row_bytes = + batch + .columns() + .iter() + .enumerate() + .try_fold(0usize, |sum, (index, array)| { + // JSON escaping and number formatting only expand JSON values. + let expansion = if json_columns.contains(&index) && !array.is_null(row) { + 8 + } else { + 1 + }; + let size = estimate_value_size(array.as_ref(), row)?; + Ok::<_, error::Error>(sum.saturating_add(size.saturating_mul(expansion))) + })?; if row_bytes > budget.saturating_sub(bytes) { break; } @@ -653,7 +797,7 @@ fn rows_within_budget( reason: "one expanded logical row exceeds conversion byte budget" } ); - Ok(row - start) + Ok((row - start, bytes)) } #[cfg(test)] diff --git a/src/operator/src/statement/export_logical_tables/tests.rs b/src/operator/src/statement/export_logical_tables/tests.rs index 4576c1e395d..a03fdf1a80a 100644 --- a/src/operator/src/statement/export_logical_tables/tests.rs +++ b/src/operator/src/statement/export_logical_tables/tests.rs @@ -18,6 +18,7 @@ use arrow::array::{ use arrow::buffer::OffsetBuffer; use arrow::datatypes::{Field, Schema, UInt32Type}; use bytes::Bytes; +use common_error::status_code::StatusCode; use common_recordbatch::{RecordBatch as GreptimeRecordBatch, RecordBatches}; use datafusion::parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder; use table::test_util::EmptyTable; @@ -137,13 +138,14 @@ async fn routes_across_batches_and_writes_empty_files() { vec![Some("a"), Some("b"), Some("unselected")], ), ]; - let result = write_tables( + let budget = ExportWriteBudget::new(4); + let result = export_stream_managed( &unit, stream(batches), &store, export_limits(), &CancellationToken::new(), - &mut None, + budget.clone(), ) .await .unwrap(); @@ -174,6 +176,7 @@ async fn routes_across_batches_and_writes_empty_files() { let (schema, requests) = read(&store, "requests.parquet").await; assert_eq!(schema.fields(), unit.logical_tables[&1027].schema.fields()); assert_eq!(requests[0].column(1).null_count(), 1); + assert_eq!(budget.available(), (4, 64 * 1024 * 1024)); } #[tokio::test] @@ -211,21 +214,16 @@ async fn rejects_invalid_order_ids_and_resource_exhaustion() { ), ] { let store = ObjectStore::new(object_store::services::Memory::default()).unwrap(); - let mut active = None; - let err = write_tables( + let err = export_stream( &unit, stream(batches), &store, limits, &CancellationToken::new(), - &mut active, ) .await .unwrap_err(); assert!(err.to_string().contains(message), "{err}"); - if let Some(writer) = active { - writer.writer.abort().await.unwrap(); - } } } @@ -233,17 +231,17 @@ async fn rejects_invalid_order_ids_and_resource_exhaustion() { async fn existing_outputs_are_not_overwritten() { let store = ObjectStore::new(object_store::services::Memory::default()).unwrap(); store.write("cpu.v1.parquet", "keep").await.unwrap(); - let err = write_tables( + let err = export_stream( &unit(), stream(vec![batch(vec![Some(1025)], vec![None])]), &store, LogicalTableExportLimits::default(), &CancellationToken::new(), - &mut None, ) .await .unwrap_err(); - assert!(err.to_string().contains("already exists")); + let status = common_error::ext::ErrorExt::status_code(&err); + assert_eq!(status, StatusCode::InvalidArguments); assert_eq!( store.read("cpu.v1.parquet").await.unwrap().to_bytes(), Bytes::from_static(b"keep") @@ -272,10 +270,10 @@ fn dictionary_and_nested_histogram_values_are_bounded_before_expansion() { ("histogram", Arc::new(histogram)), ]) .unwrap(); - assert_eq!(rows_within_budget(&batch, 0, 3, 4300).unwrap(), 1); + assert_eq!(rows_within_budget(&batch, 0, 3, 4300, &[]).unwrap().0, 1); // The dictionary and container overhead fit; the nested list elements do not. - assert!(rows_within_budget(&batch, 0, 3, 4180).is_err()); - assert_eq!(rows_within_budget(&batch, 0, 3, 15000).unwrap(), 3); + assert!(rows_within_budget(&batch, 0, 3, 4180, &[]).is_err()); + assert_eq!(rows_within_budget(&batch, 0, 3, 15000, &[]).unwrap().0, 3); } #[test] @@ -440,64 +438,195 @@ impl object_store::layers::mock::oio::Write for PausedFileWriter { } #[tokio::test] -async fn cancellation_waits_for_file_creation_before_cleanup() { - use object_store::layers::mock::{MockLayerBuilder, MockWriterFactory}; - let directory = common_test_util::temp_dir::create_temp_dir("metric_export_pending_open"); - let access = - common_datasource::object_store::LocalFileAccess::sandboxed(directory.path()).unwrap(); - let store = build_backend_for_write( - &format!("{}/", directory.path().display()), - &HashMap::new(), - &access, - ) - .await - .unwrap(); - let started = Arc::new(tokio::sync::Notify::new()); - let release = Arc::new(tokio::sync::Notify::new()); - let factory: MockWriterFactory = Arc::new({ - let started = started.clone(); - let release = release.clone(); - move |_, _, inner| { - Box::new(PausedFileWriter { - inner: Some(inner), - started: started.clone(), - release: release.clone(), - }) +async fn cancellation_drains_storage_and_preserves_committed_files() { + for large in [false, true] { + use object_store::layers::mock::{MockLayerBuilder, MockWriterFactory}; + let directory = common_test_util::temp_dir::create_temp_dir("metric_export_pending_open"); + let access = + common_datasource::object_store::LocalFileAccess::sandboxed(directory.path()).unwrap(); + let store = build_backend_for_write( + &format!("{}/", directory.path().display()), + &HashMap::new(), + &access, + ) + .await + .unwrap(); + let started = Arc::new(tokio::sync::Notify::new()); + let release = Arc::new(tokio::sync::Notify::new()); + let factory: MockWriterFactory = Arc::new({ + let started = started.clone(); + let release = release.clone(); + move |_, _, inner| { + Box::new(PausedFileWriter { + inner: Some(inner), + started: started.clone(), + release: release.clone(), + }) + } + }); + let store = store.layer( + MockLayerBuilder::default() + .writer_factory(factory) + .build() + .unwrap(), + ); + let cancellation = CancellationToken::new(); + let (unit, input) = if large { + let field = Field::new("value", DataType::Utf8, true); + let unit = LogicalTableExport::try_new( + table( + 1024, + "phy", + vec![Field::new(TABLE_ID, DataType::UInt32, false), field.clone()], + true, + ), + &[table(1025, "cpu.v1", vec![field], false)], + ) + .unwrap(); + let mut state = 17u64; + let strings = (0..512) + .map(|_| { + (0..32768) + .map(|_| { + state ^= state << 13; + state ^= state >> 7; + state ^= state << 17; + char::from(b' ' + (state % 95) as u8) + }) + .collect::() + }) + .collect::>(); + let input = RecordBatch::try_from_iter([ + ( + TABLE_ID, + Arc::new(UInt32Array::from(vec![1025; 512])) as ArrayRef, + ), + ("value", Arc::new(StringArray::from(strings))), + ]) + .unwrap(); + (unit, input) + } else { + (unit(), batch(vec![Some(1025)], vec![Some("a")])) + }; + let budget = ExportWriteBudget::new(1); + let mut limits = LogicalTableExportLimits::default(); + limits.writer.flush_threshold_bytes = 1024 * 1024; + let export = export_stream_managed( + &unit, + stream(vec![input]), + &store, + limits, + &cancellation, + budget.clone(), + ); + tokio::pin!(export); + tokio::select! { + result = &mut export => panic!("export completed before the file open: {result:?}"), + _ = started.notified() => {}, } - }); - let store = store.layer( - MockLayerBuilder::default() - .writer_factory(factory) - .build() - .unwrap(), - ); - let cancellation = CancellationToken::new(); - let unit = unit(); - let export = export_stream( - &unit, - stream(vec![batch(vec![Some(1025)], vec![Some("a")])]), - &store, - LogicalTableExportLimits::default(), - &cancellation, - ); - tokio::pin!(export); - tokio::select! { - result = &mut export => panic!("export completed before the file open: {result:?}"), - _ = started.notified() => {}, + if large { + assert!(budget.available().1 < 64 * 1024 * 1024); + } + let held = budget.available(); + cancellation.cancel(); + assert!(futures::poll!(&mut export).is_pending()); + assert_eq!(budget.available().0, held.0); + if large { + // Cancelled admission releases its pending reservation; the worker + // still owns payload capacity while the storage operation is paused. + assert!(budget.available().1 < 64 * 1024 * 1024); + } + release.notify_one(); + let result = export.await; + assert!(matches!( + result, + Err(error::Error::LogicalTableExportCancelled { .. }) + )); + assert_eq!(store.exists("cpu.v1.parquet").await.unwrap(), !large); + if !large { + let (_, batches) = read(&store, "cpu.v1.parquet").await; + assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::(), 1); + } + assert_eq!(budget.available(), (1, 64 * 1024 * 1024)); } - cancellation.cancel(); - assert!(futures::poll!(&mut export).is_pending()); - release.notify_one(); - let result = export.await; - assert!(matches!( - result, - Err(error::Error::LogicalTableExportCancelled { .. }) - )); - assert!(!store.exists("cpu.v1.parquet").await.unwrap()); } struct FailedAbortWriter(object_store::layers::mock::oio::Writer); +#[tokio::test] +async fn metric_fallback_and_ordinary_preserve_ambiguous_commits() { + use object_store::layers::CapabilityOverrideLayer; + use object_store::layers::mock::{Metadata, MockLayerBuilder, MockWriterFactory, oio}; + + use crate::statement::copy_table_to::stream_to_managed_parquet; + + struct AmbiguousCommit(oio::Writer); + impl oio::Write for AmbiguousCommit { + async fn write(&mut self, bytes: object_store::Buffer) -> object_store::Result<()> { + self.0.write(bytes).await + } + async fn close(&mut self) -> object_store::Result { + self.0.close().await?; + Err(object_store::Error::new( + object_store::ErrorKind::Unexpected, + "lost close reply", + )) + } + async fn abort(&mut self) -> object_store::Result<()> { + Err(object_store::Error::new( + object_store::ErrorKind::Unsupported, + "cannot abort", + )) + } + } + let factory: MockWriterFactory = Arc::new(|_, args, inner| { + assert!(!args.if_not_exists()); + Box::new(AmbiguousCommit(inner)) + }); + let store = ObjectStore::new(object_store::services::Memory::default()) + .unwrap() + .layer(CapabilityOverrideLayer::new(|mut capability| { + capability.write_with_if_not_exists = false; + capability + })) + .layer( + MockLayerBuilder::default() + .writer_factory(factory) + .build() + .unwrap(), + ); + assert!(!store.info().capability().write_with_if_not_exists); + let rows = || stream(vec![batch(vec![Some(1025)], vec![Some("a")])]); + assert!( + export_stream( + &unit(), + rows(), + &store, + export_limits(), + &CancellationToken::new() + ) + .await + .is_err() + ); + let budget = ExportWriteBudget::new(1); + assert!( + stream_to_managed_parquet( + rows(), + store.clone(), + "ordinary.parquet", + &budget, + &CancellationToken::new() + ) + .await + .is_err() + ); + for path in ["cpu.v1.parquet", "ordinary.parquet"] { + let (_, batches) = read(&store, path).await; + assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::(), 1); + } + assert_eq!(budget.available(), (1, 64 * 1024 * 1024)); +} + impl object_store::layers::mock::oio::Write for FailedAbortWriter { async fn write(&mut self, bytes: object_store::Buffer) -> object_store::Result<()> { self.0.write(bytes).await @@ -574,3 +703,243 @@ async fn validates_membership_by_table_route() { } } } + +#[tokio::test] +async fn retained_backing_is_reserved_before_conversion_and_until_payload_drop() { + let input = batch(vec![Some(1025); 8192], vec![Some("shared"); 8192]); + let tiny = input.slice(0, 1); + let schema = Arc::new(Schema::new(vec![Field::new("host", DataType::Utf8, true)])); + let projected = tiny.project(&[2]).unwrap(); + let full = input.project(&[2]).unwrap().get_array_memory_size(); + assert_eq!(projected.get_array_memory_size(), full); + let budget = ExportWriteBudget::new(1); + let token = CancellationToken::new(); + let blocker = budget.reserve(64 * 1024 * 1024, &token).await.unwrap(); + let convert = expand_bounded_slice(projected, schema, 0, 1, 1024, &budget, &token); + tokio::pin!(convert); + assert!(futures::poll!(&mut convert).is_pending()); + drop(blocker); + let (payload, rows) = convert.await.unwrap(); + assert_eq!(rows, 1); + assert!( + 64 * 1024 * 1024 - budget.available().1 >= full + payload.batch.get_array_memory_size() + ); + drop(payload); + assert_eq!(budget.available(), (1, 64 * 1024 * 1024)); + assert!(budget.reserve(64 * 1024 * 1024 + 1, &token).await.is_err()); +} + +#[tokio::test] +async fn groups_and_ordinary_files_share_writer_admission_and_drain() { + use std::sync::atomic::{AtomicUsize, Ordering}; + + use object_store::layers::mock::{Metadata, MockLayerBuilder, MockWriterFactory, oio}; + use tokio::sync::Semaphore; + + use crate::statement::copy_table_to::stream_to_managed_parquet; + + struct CountedWriter { + inner: oio::Writer, + active: Arc, + started: Arc, + release: Arc, + } + impl Drop for CountedWriter { + fn drop(&mut self) { + self.active.fetch_sub(1, Ordering::SeqCst); + } + } + impl oio::Write for CountedWriter { + async fn write(&mut self, bytes: object_store::Buffer) -> object_store::Result<()> { + self.inner.write(bytes).await + } + async fn close(&mut self) -> object_store::Result { + self.started.add_permits(1); + self.release.acquire().await.unwrap().forget(); + self.inner.close().await + } + async fn abort(&mut self) -> object_store::Result<()> { + self.inner.abort().await + } + } + for parallelism in [1, 4] { + for cancel in [false, true] { + let active = Arc::new(AtomicUsize::new(0)); + let peak = Arc::new(AtomicUsize::new(0)); + let started = Arc::new(Semaphore::new(0)); + let release = Arc::new(Semaphore::new(0)); + let factory: MockWriterFactory = Arc::new({ + let (active, peak, started, release) = ( + active.clone(), + peak.clone(), + started.clone(), + release.clone(), + ); + move |_, args, inner| { + assert_eq!(args.concurrent(), 1); + let now = active.fetch_add(1, Ordering::SeqCst) + 1; + peak.fetch_max(now, Ordering::SeqCst); + Box::new(CountedWriter { + inner, + active: active.clone(), + started: started.clone(), + release: release.clone(), + }) + } + }); + let store = ObjectStore::new(object_store::services::Memory::default()) + .unwrap() + .layer( + MockLayerBuilder::default() + .writer_factory(factory) + .build() + .unwrap(), + ); + let budget = ExportWriteBudget::new(parallelism); + let token = CancellationToken::new(); + let mut a = unit(); + let mut b = unit(); + for (group, unit) in [("a", &mut a), ("b", &mut b)] { + for table in unit.logical_tables.values_mut() { + table.output.path = format!("{group}/{}", table.output.path); + } + } + let rows = || { + stream(vec![batch( + vec![Some(1025), Some(1027)], + vec![Some("a"), Some("b")], + )]) + }; + let ordinary = async { + let _permit = budget.writer(&token).await?; + stream_to_managed_parquet( + rows(), + store.clone(), + "ordinary.parquet", + &budget, + &token, + ) + .await + }; + let work = async { + tokio::join!( + export_stream_managed( + &a, + rows(), + &store, + export_limits(), + &token, + budget.clone() + ), + export_stream_managed( + &b, + rows(), + &store, + export_limits(), + &token, + budget.clone() + ), + ordinary, + ) + }; + tokio::pin!(work); + tokio::select! { + _ = started.acquire_many(parallelism as u32) => {}, + result = &mut work => panic!("completed while close paused: {result:?}"), + } + assert_eq!(budget.available().0, 0); + assert_eq!(peak.load(Ordering::SeqCst), parallelism); + if cancel { + token.cancel(); + } + assert!(futures::poll!(&mut work).is_pending()); + release.add_permits(10); + let (a, b, ordinary) = work.await; + assert_eq!(a.is_err(), cancel); + assert_eq!(b.is_err(), cancel); + assert_eq!(ordinary.is_err(), cancel); + assert_eq!(active.load(Ordering::SeqCst), 0); + assert!(peak.load(Ordering::SeqCst) <= parallelism); + assert_eq!(budget.available(), (parallelism, 64 * 1024 * 1024)); + } + } +} + +#[tokio::test] +async fn nested_dictionary_conversion_limits_child_ranges_and_charges_null_parents() { + let values = Arc::new(DictionaryArray::::new( + UInt32Array::from(vec![0; 1024]), + Arc::new(StringArray::from(vec!["x".repeat(4096)])), + )) as ArrayRef; + let list = ListArray::new( + Arc::new(Field::new("item", values.data_type().clone(), true)), + OffsetBuffer::new(vec![0i32, 1023, 1024].into()), + values, + None, + ); + let input = RecordBatch::try_from_iter([("list", Arc::new(list) as ArrayRef)]).unwrap(); + let schema = Arc::new(Schema::new(vec![Field::new( + "list", + DataType::List(Arc::new(Field::new("item", DataType::Utf8, true))), + true, + )])); + let budget = ExportWriteBudget::new(1); + let (payload, rows) = expand_bounded_slice( + input, + schema, + 1, + 2, + 8192, + &budget, + &CancellationToken::new(), + ) + .await + .unwrap(); + assert_eq!(rows, 1); + assert!(payload.batch.get_array_memory_size() < 16384); + let list = payload.batch.column(0).as_list::(); + assert_eq!(list.value_offsets(), &[0, 1]); + assert_eq!(list.values().as_string::().value(0), "x".repeat(4096)); + drop(payload); + assert_eq!(budget.available(), (1, 64 * 1024 * 1024)); + + let values = Arc::new(DictionaryArray::::new( + UInt32Array::from(vec![0]), + Arc::new(StringArray::from(vec!["x".repeat(4096)])), + )) as ArrayRef; + let structure = StructArray::new( + vec![Arc::new(Field::new( + "child", + values.data_type().clone(), + true, + ))] + .into(), + vec![values], + Some(arrow::buffer::NullBuffer::from(vec![false])), + ); + let input = RecordBatch::try_from_iter([("struct", Arc::new(structure) as ArrayRef)]).unwrap(); + assert!(rows_within_budget(&input, 0, 1, 128, &[]).is_err()); +} + +#[tokio::test] +async fn completed_table_workers_are_reaped_during_admission() { + for parallelism in [1, 4] { + let mut unit = unit(); + let file = unit.logical_tables.get_mut(&1025).unwrap(); + let store = ObjectStore::new(object_store::services::Memory::default()).unwrap(); + let budget = ExportWriteBudget::new(parallelism); + let token = CancellationToken::new(); + let mut writers = TableWriters::new(budget.clone()); + for id in 0..64 { + file.output.path = format!("empty-{id}.parquet"); + writers + .open(id, file, &store, export_limits(), &token) + .await + .unwrap(); + assert!(writers.pending_tasks() <= parallelism + 1); + } + writers.drain(Ok(()), &token).await.unwrap(); + assert_eq!(writers.pending_tasks(), 0); + assert_eq!(budget.available(), (parallelism, 64 * 1024 * 1024)); + } +} diff --git a/src/operator/src/statement/export_logical_tables/writers.rs b/src/operator/src/statement/export_logical_tables/writers.rs new file mode 100644 index 00000000000..38a63936a5b --- /dev/null +++ b/src/operator/src/statement/export_logical_tables/writers.rs @@ -0,0 +1,310 @@ +// Copyright 2023 Greptime Team +// +// Licensed 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. + +use std::sync::Arc; + +use arrow::record_batch::RecordBatch; +use futures::stream::FuturesUnordered; +use futures::{FutureExt, StreamExt}; +use object_store::ObjectStore; +use snafu::{OptionExt, ResultExt, ensure}; +use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc}; +use tokio::task::JoinHandle; +use tokio_util::sync::CancellationToken; + +use crate::error::{self, Result}; +use crate::statement::export_logical_tables::{ + ActiveWriter, LogicalTableExportLimits, LogicalTableProjection, check_cancelled, + map_writer_error, +}; + +const PAYLOAD_BYTES: usize = 64 * 1024 * 1024; + +/// Request-wide downstream payload and writer admission, separate from query memory. +pub(crate) struct ExportWriteBudget { + writers: Arc, + bytes: Arc, + max_writers: usize, +} + +impl ExportWriteBudget { + pub(crate) fn new(parallelism: usize) -> Arc { + let max_writers = parallelism.max(1); + Arc::new(Self { + writers: Arc::new(Semaphore::new(max_writers)), + bytes: Arc::new(Semaphore::new(PAYLOAD_BYTES)), + max_writers, + }) + } + + #[cfg(test)] + pub(crate) fn available(&self) -> (usize, usize) { + ( + self.writers.available_permits(), + self.bytes.available_permits(), + ) + } + + pub(crate) async fn writer(&self, token: &CancellationToken) -> Result { + tokio::select! { + biased; + _ = token.cancelled() => error::LogicalTableExportCancelledSnafu.fail(), + permit = self.writers.clone().acquire_owned() => permit.map_err(|_| error::UnexpectedSnafu { violated: "writer budget closed" }.build()), + } + } + + pub(crate) async fn reserve( + &self, + size: usize, + token: &CancellationToken, + ) -> Result { + ensure!( + size <= PAYLOAD_BYTES, + error::LogicalTableExportResourceSnafu { + reason: "batch backing buffers exceed request payload budget" + } + ); + tokio::select! { + biased; + _ = token.cancelled() => error::LogicalTableExportCancelledSnafu.fail(), + permit = self.bytes.clone().acquire_many_owned(size as u32) => permit.map_err(|_| error::UnexpectedSnafu { violated: "payload budget closed" }.build()), + } + } + + /// Accounts for retained input allocations, conversion capacity growth and + /// array metadata. Pass zero backing when the caller detaches oversized slices. + pub(crate) fn conversion_budget( + backing: usize, + columns: usize, + requested: usize, + ) -> Result<(usize, usize)> { + let overhead = columns.saturating_mul(1024); + let available = PAYLOAD_BYTES + .saturating_sub(backing) + .saturating_sub(overhead) + / 4; + let conversion = requested.min(available); + ensure!( + conversion > 0, + error::LogicalTableExportResourceSnafu { + reason: "batch backing buffers exceed request payload budget" + } + ); + Ok((conversion, backing.saturating_add(overhead))) + } +} + +pub(crate) struct Payload { + pub(crate) batch: RecordBatch, + pub(crate) permit: OwnedSemaphorePermit, +} + +pub(crate) struct TableWriters { + current: Option<(u32, mpsc::Sender)>, + tasks: FuturesUnordered>>, + budget: Arc, +} + +impl TableWriters { + pub(crate) fn new(budget: Arc) -> Self { + Self { + current: None, + tasks: FuturesUnordered::new(), + budget, + } + } + + #[cfg(test)] + pub(crate) fn pending_tasks(&self) -> usize { + self.tasks.len() + } + + pub(crate) fn table_id(&self) -> Option { + self.current.as_ref().map(|(id, _)| *id) + } + + pub(crate) fn close_input(&mut self) { + self.current = None; + } + + pub(crate) async fn open( + &mut self, + id: u32, + table: &LogicalTableProjection, + store: &ObjectStore, + limits: LogicalTableExportLimits, + token: &CancellationToken, + ) -> Result<()> { + // EOF must precede acquiring the next slot, including when P is one. + self.close_input(); + let permit = self.budget.writer(token).await?; + self.reap_for_admission(token).await?; + let writer = ActiveWriter::open(table, store, limits).await?; + let (sender, receiver) = mpsc::channel(2); + self.current = Some((id, sender)); + let token = token.clone(); + self.tasks.push(common_runtime::spawn_global(async move { + let _permit = permit; + let guard = token.clone().drop_guard(); + let result = run_writer(writer, receiver, &token).await; + if result.is_ok() { + guard.disarm(); + } + result + })); + Ok(()) + } + + async fn reap_for_admission(&mut self, token: &CancellationToken) -> Result<()> { + while let Some(Some(result)) = self.tasks.next().now_or_never() { + result.context(error::JoinTaskSnafu)??; + } + // A worker can release its permit before its JoinHandle becomes ready. + while self.tasks.len() > self.budget.max_writers { + let result = tokio::select! { + biased; + _ = token.cancelled() => return error::LogicalTableExportCancelledSnafu.fail(), + result = self.tasks.next() => result.context(error::UnexpectedSnafu { violated: "writer task queue unexpectedly empty" })?, + }; + result.context(error::JoinTaskSnafu)??; + } + Ok(()) + } + + pub(crate) async fn send(&self, payload: Payload, token: &CancellationToken) -> Result<()> { + let (_, sender) = self.current.as_ref().context(error::UnexpectedSnafu { + violated: "missing logical writer", + })?; + tokio::select! { + biased; + _ = token.cancelled() => error::LogicalTableExportCancelledSnafu.fail(), + result = sender.send(payload) => result.map_err(|_| error::LogicalTableExportCancelledSnafu.build()), + } + } + + pub(crate) async fn drain( + &mut self, + result: Result<()>, + token: &CancellationToken, + ) -> Result<()> { + self.close_input(); + let mut first_error = result.err(); + if first_error.is_some() { + token.cancel(); + } + while let Some(result) = self.tasks.next().await { + if let Err(err) = result.context(error::JoinTaskSnafu).and_then(|r| r) { + retain_error(&mut first_error, err); + token.cancel(); + } + } + match first_error { + Some(err) => Err(err), + None => Ok(()), + } + } +} + +pub(crate) fn retain_error(first: &mut Option, error: error::Error) { + if first.as_ref().is_none_or(|err| { + matches!( + err, + error::Error::LogicalTableExportCancelled { .. } + | error::Error::DatabaseExportCancelled { .. } + ) + }) { + *first = Some(error); + } +} + +async fn run_writer( + mut writer: ActiveWriter, + mut receiver: mpsc::Receiver, + token: &CancellationToken, +) -> Result<()> { + let result = async { + loop { + let payload = tokio::select! { + biased; + _ = token.cancelled() => return error::LogicalTableExportCancelledSnafu.fail(), + payload = receiver.recv() => payload, + }; + let Some(Payload { batch, permit }) = payload else { + break; + }; + let result = writer.writer.write(batch, Some(token)).await; + drop(permit); + result.map_err(|error| map_writer_error(error, &writer.path))?; + } + writer + .writer + .finish(Some(token)) + .await + .map_err(|error| map_writer_error(error, &writer.path))?; + check_cancelled(token) + } + .await; + if result.is_err() { + token.cancel(); + receiver.close(); + drop(receiver); + if let Err(error) = writer.writer.abort().await { + common_telemetry::warn!(error; "Failed to abort Metric export file"); + } + } + result +} + +#[cfg(test)] +mod tests { + use tokio::sync::oneshot; + + use super::*; + + #[tokio::test] + async fn released_permit_does_not_hide_pending_join_handles() { + for parallelism in [1, 4] { + let budget = ExportWriteBudget::new(parallelism); + let token = CancellationToken::new(); + let mut writers = TableWriters::new(budget.clone()); + let mut resume = Vec::new(); + for _ in 0..=parallelism { + let permit = budget.writer(&token).await.unwrap(); + let (released_tx, released_rx) = oneshot::channel(); + let (resume_tx, resume_rx) = oneshot::channel(); + writers.tasks.push(common_runtime::spawn_global(async move { + drop(permit); + released_tx.send(()).unwrap(); + resume_rx.await.unwrap(); + Ok(()) + })); + released_rx.await.unwrap(); + resume.push(resume_tx); + } + assert_eq!(budget.available().0, parallelism); + + let mut reap = Box::pin(writers.reap_for_admission(&token)); + assert!(reap.as_mut().now_or_never().is_none()); + resume.pop().unwrap().send(()).unwrap(); + reap.await.unwrap(); + assert_eq!(writers.pending_tasks(), parallelism); + + for tx in resume { + tx.send(()).unwrap(); + } + writers.drain(Ok(()), &token).await.unwrap(); + assert_eq!(budget.available().0, parallelism); + } + } +} diff --git a/tests-integration/tests/export_logical_tables.rs b/tests-integration/tests/export_logical_tables.rs index 41fe8382c82..1f9161d970f 100644 --- a/tests-integration/tests/export_logical_tables.rs +++ b/tests-integration/tests/export_logical_tables.rs @@ -228,7 +228,7 @@ fn database_export_request(directory: &std::path::Path) -> table::requests::Copy } } -async fn database_export_roundtrip(instance: &Arc) { +async fn database_export_roundtrip(instance: &Arc, parallelism: usize) { let destination = tempfile::tempdir_in(common_test_util::find_workspace_path(".")).unwrap(); let (first_logical_table_names, _, renamed_physical_table) = create_metric_export_source_tables(instance, "db_a", "dense").await; @@ -253,7 +253,9 @@ async fn database_export_roundtrip(instance: &Arc) { ]; let mut names = selected.clone(); names.extend([renamed_physical_table, "dashboard".into()]); - let req = database_export_request(&destination.path().join("data")); + let mut req = database_export_request(&destination.path().join("data")); + req.with + .insert("parallelism".into(), parallelism.to_string()); let executor = instance.statement_executor(); let captured = executor .capture_database_export_tables(&req, None, &QueryContext::arc()) @@ -448,7 +450,7 @@ async fn database_export_standalone_roundtrip() { let standalone = GreptimeDbStandaloneBuilder::new("database_export") .build() .await; - database_export_roundtrip(standalone.fe_instance()).await; + database_export_roundtrip(standalone.fe_instance(), 1).await; } #[tokio::test(flavor = "multi_thread")] @@ -464,7 +466,7 @@ async fn database_export_distributed_roundtrip() { ) .build(false) .await; - database_export_roundtrip(cluster.fe_instance()).await; + database_export_roundtrip(cluster.fe_instance(), 4).await; } #[tokio::test(flavor = "multi_thread")]