Skip to main content

common_datasource/
parquet_writer.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use arrow::datatypes::{DataType, SchemaRef};
16use arrow::record_batch::RecordBatch;
17use bytes::Bytes;
18use futures::future::BoxFuture;
19use object_store::{ObjectStore, Writer};
20use parquet::arrow::ArrowWriter;
21use parquet::arrow::async_writer::AsyncFileWriter;
22use parquet::basic::{Compression, Encoding, ZstdLevel};
23use parquet::errors::ParquetError;
24use parquet::file::properties::WriterProperties;
25use parquet::schema::types::ColumnPath;
26use snafu::{IntoError, ResultExt, ensure};
27use tokio_util::sync::CancellationToken;
28
29use crate::DEFAULT_WRITE_BUFFER_SIZE;
30use crate::error::{self, Result};
31use crate::packed_writer::PackedTableWriter;
32
33/// Destination creation policy; conditional failures never authorize path deletion.
34#[derive(Clone, Copy, Debug, PartialEq, Eq)]
35pub enum ParquetCreationPolicy {
36    Overwrite,
37    IfNotExists,
38}
39
40/// Limits for one Parquet file. Flush thresholds are not hard memory caps.
41#[derive(Clone, Copy, Debug)]
42pub struct ParquetWriterLimits {
43    /// Maximum rows in a row group.
44    pub row_group_rows: usize,
45    /// Flush when encoder memory or encoded size reaches this threshold.
46    pub flush_threshold_bytes: usize,
47    /// Maximum row groups retained in the file footer.
48    pub max_row_groups: usize,
49}
50
51/// Encodes batches on a blocking runtime and writes one object-store file.
52/// Callers must await each operation before aborting; dropping an in-flight
53/// operation can leave a filesystem worker running after cleanup.
54pub struct ParquetFileWriter {
55    encoder: Option<ArrowWriter<Vec<u8>>>,
56    sink: ParquetSink,
57    rows: u64,
58    store: ObjectStore,
59    path: String,
60    limits: Option<ParquetWriterLimits>,
61    creation: ParquetCreationPolicy,
62    close_started: bool,
63}
64
65enum ParquetSink {
66    Object(Writer),
67    Packed(Box<PackedTableWriter>),
68}
69
70impl ParquetFileWriter {
71    /// Open a file using COPY's encoding settings. Destination ownership and
72    /// overwrite policy belong to the caller. None preserves Parquet's defaults.
73    pub async fn open(
74        schema: SchemaRef,
75        store: ObjectStore,
76        path: &str,
77        concurrency: usize,
78        limits: Option<ParquetWriterLimits>,
79    ) -> Result<Self> {
80        Self::open_with_creation(
81            schema,
82            store,
83            path,
84            concurrency,
85            limits,
86            ParquetCreationPolicy::Overwrite,
87        )
88        .await
89    }
90
91    /// Open with a conditional policy only when the caller verified backend support.
92    pub async fn open_with_creation(
93        schema: SchemaRef,
94        store: ObjectStore,
95        path: &str,
96        concurrency: usize,
97        limits: Option<ParquetWriterLimits>,
98        creation: ParquetCreationPolicy,
99    ) -> Result<Self> {
100        let encoder = build_encoder(schema, limits, path)?;
101        let sink = store
102            .writer_with(path)
103            .concurrent(concurrency)
104            .chunk(DEFAULT_WRITE_BUFFER_SIZE.as_bytes() as usize)
105            .if_not_exists(creation == ParquetCreationPolicy::IfNotExists)
106            .await
107            .context(error::WriteObjectSnafu { path })?;
108        Ok(Self {
109            encoder: Some(encoder),
110            sink: ParquetSink::Object(sink),
111            rows: 0,
112            store,
113            path: path.to_owned(),
114            limits,
115            creation,
116            close_started: false,
117        })
118    }
119
120    /// Encode into a request-owned pack or standalone destination.
121    pub fn open_packed(
122        schema: SchemaRef,
123        store: ObjectStore,
124        path: &str,
125        limits: Option<ParquetWriterLimits>,
126        packed: PackedTableWriter,
127    ) -> Result<Self> {
128        Ok(Self {
129            encoder: Some(build_encoder(schema, limits, path)?),
130            sink: ParquetSink::Packed(Box::new(packed)),
131            rows: 0,
132            store,
133            path: path.into(),
134            limits,
135            creation: ParquetCreationPolicy::IfNotExists,
136            close_started: false,
137        })
138    }
139
140    /// Write a batch, enforcing file limits across batch and row-group boundaries.
141    pub async fn write(
142        &mut self,
143        batch: RecordBatch,
144        cancellation: Option<&CancellationToken>,
145    ) -> Result<()> {
146        let mut offset = 0;
147        while offset < batch.num_rows() {
148            check_cancelled(cancellation)?;
149            let mut encoder = self.encoder.take().ok_or_else(|| {
150                error::WriteParquetSnafu { path: &self.path }
151                    .into_error(ParquetError::General("Parquet writer is closed".into()))
152            })?;
153            let len = self.limits.map_or(batch.num_rows() - offset, |limits| {
154                (limits.row_group_rows - encoder.in_progress_rows()).min(batch.num_rows() - offset)
155            });
156            let slice = batch.slice(offset, len);
157            let limits = self.limits;
158            let path = self.path.clone();
159            let (encoder, bytes) = common_runtime::spawn_blocking_global(move || {
160                if let Some(limits) = limits {
161                    ensure!(
162                        encoder.flushed_row_groups().len() < limits.max_row_groups,
163                        error::ParquetWriterResourceSnafu {
164                            reason: "Parquet row-group metadata budget exceeded"
165                        }
166                    );
167                }
168                encoder
169                    .write(&slice)
170                    .context(error::WriteParquetSnafu { path: &path })?;
171                if limits.is_some_and(|limits| {
172                    encoder.memory_size() >= limits.flush_threshold_bytes
173                        || encoder.in_progress_size() >= limits.flush_threshold_bytes
174                }) {
175                    encoder
176                        .flush()
177                        .context(error::WriteParquetSnafu { path: &path })?;
178                }
179                // Draining preserves ArrowWriter's cumulative file offsets.
180                let bytes = std::mem::take(encoder.inner_mut());
181                Ok::<_, error::Error>((encoder, bytes))
182            })
183            .await
184            .context(error::JoinHandleSnafu)??;
185            self.encoder = Some(encoder);
186            check_cancelled(cancellation)?;
187            self.write_bytes(bytes).await?;
188            check_cancelled(cancellation)?;
189            self.rows += len as u64;
190            offset += len;
191        }
192        Ok(())
193    }
194
195    async fn write_bytes(&mut self, bytes: Vec<u8>) -> Result<()> {
196        let bytes = Bytes::from(bytes);
197        let sink = match &mut self.sink {
198            ParquetSink::Packed(packed) => return packed.write(bytes).await,
199            ParquetSink::Object(sink) => sink,
200        };
201        let chunk = DEFAULT_WRITE_BUFFER_SIZE.as_bytes() as usize;
202        // Slices retain the complete encoded allocation until its last submission.
203        for offset in (0..bytes.len()).step_by(chunk) {
204            sink.write(bytes.slice(offset..(offset + chunk).min(bytes.len())))
205                .await
206                .context(error::WriteObjectSnafu { path: &self.path })?;
207        }
208        Ok(())
209    }
210
211    /// Write the footer and close the file. Retains the handle for cleanup on error.
212    pub async fn finish(&mut self, cancellation: Option<&CancellationToken>) -> Result<()> {
213        check_cancelled(cancellation)?;
214        let mut encoder = self.encoder.take().ok_or_else(|| {
215            error::WriteParquetSnafu { path: &self.path }
216                .into_error(ParquetError::General("Parquet writer is closed".into()))
217        })?;
218        let path = self.path.clone();
219        let bytes = common_runtime::spawn_blocking_global(move || {
220            encoder
221                .finish()
222                .context(error::WriteParquetSnafu { path: &path })?;
223            Ok::<_, error::Error>(std::mem::take(encoder.inner_mut()))
224        })
225        .await
226        .context(error::JoinHandleSnafu)??;
227        self.write_bytes(bytes).await?;
228        check_cancelled(cancellation)?;
229        let sink = match &mut self.sink {
230            ParquetSink::Packed(packed) => return packed.finish(self.rows, cancellation).await,
231            ParquetSink::Object(sink) => sink,
232        };
233        self.close_started = true;
234        sink.close()
235            .await
236            .context(error::WriteObjectSnafu { path: &self.path })?;
237        check_cancelled(cancellation)?;
238        Ok(())
239    }
240
241    /// Abort after all in-flight operations complete. Preserve ambiguous commits;
242    /// conditional callers delegate cleanup exclusively to the backend.
243    pub async fn abort(mut self) -> Result<()> {
244        let result = match &mut self.sink {
245            ParquetSink::Packed(packed) => return packed.abort().await,
246            ParquetSink::Object(sink) => sink.abort().await,
247        };
248        if self.creation == ParquetCreationPolicy::Overwrite
249            && result.as_ref().is_err_and(|error| {
250                error.kind() == object_store::ErrorKind::Unsupported
251                    && (!self.close_started
252                        || object_store::secure_fs::is_unsynced_overwrite_abort(error))
253            })
254        {
255            let store = self.store.clone();
256            let path = self.path.clone();
257            // Secure filesystem writers require dropping the handle before deletion.
258            drop(self);
259            store
260                .delete(&path)
261                .await
262                .context(error::WriteObjectSnafu { path })?;
263        } else {
264            result.context(error::WriteObjectSnafu { path: &self.path })?;
265        }
266        Ok(())
267    }
268}
269
270fn build_encoder(
271    schema: SchemaRef,
272    limits: Option<ParquetWriterLimits>,
273    path: &str,
274) -> Result<ArrowWriter<Vec<u8>>> {
275    let mut props = WriterProperties::builder()
276        .set_compression(Compression::ZSTD(ZstdLevel::default()))
277        .set_statistics_truncate_length(None)
278        .set_column_index_truncate_length(None);
279    if let Some(limits) = limits {
280        ensure!(
281            limits.row_group_rows > 0
282                && limits.flush_threshold_bytes > 0
283                && limits.max_row_groups > 0,
284            error::InvalidParquetWriterLimitsSnafu
285        );
286        props = props
287            .set_max_row_group_row_count(Some(limits.row_group_rows))
288            .set_max_row_group_bytes(None);
289    }
290    for field in schema.fields() {
291        if matches!(field.data_type(), DataType::Timestamp(_, _)) {
292            let column = ColumnPath::new(vec![field.name().clone()]);
293            props = props
294                .set_column_dictionary_enabled(column.clone(), false)
295                .set_column_encoding(column, Encoding::DELTA_BINARY_PACKED);
296        }
297    }
298    let encoder = ArrowWriter::try_new(Vec::new(), schema, Some(props.build()))
299        .context(error::WriteParquetSnafu { path })?;
300    Ok(encoder)
301}
302
303fn check_cancelled(cancellation: Option<&CancellationToken>) -> Result<()> {
304    ensure!(
305        cancellation.is_none_or(|token| !token.is_cancelled()),
306        error::ParquetWriteCancelledSnafu
307    );
308    Ok(())
309}
310
311/// Bridges opendal [Writer] with parquet [AsyncFileWriter].
312pub struct AsyncWriter {
313    inner: Writer,
314}
315
316impl AsyncWriter {
317    /// Create a [`AsyncWriter`] by given [`Writer`].
318    pub fn new(writer: Writer) -> Self {
319        Self { inner: writer }
320    }
321}
322
323impl AsyncFileWriter for AsyncWriter {
324    fn write(&mut self, bs: Bytes) -> BoxFuture<'_, parquet::errors::Result<()>> {
325        Box::pin(async move {
326            self.inner
327                .write(bs)
328                .await
329                .map_err(|err| ParquetError::External(Box::new(err)))
330        })
331    }
332
333    fn complete(&mut self) -> BoxFuture<'_, parquet::errors::Result<()>> {
334        Box::pin(async move {
335            self.inner
336                .close()
337                .await
338                .map(|_| ())
339                .map_err(|err| ParquetError::External(Box::new(err)))
340        })
341    }
342}
343
344#[cfg(test)]
345mod tests {
346    use std::sync::Arc;
347
348    use arrow::array::{Int64Array, TimestampMillisecondArray};
349    use common_error::ext::ErrorExt;
350    use common_error::status_code::StatusCode;
351    use datafusion::physical_plan::stream::RecordBatchStreamAdapter;
352    use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
353
354    use super::*;
355    use crate::file_format::parquet::stream_to_parquet;
356
357    fn batch() -> RecordBatch {
358        RecordBatch::try_from_iter([
359            (
360                "value",
361                Arc::new(Int64Array::from(vec![Some(1), None, Some(3), Some(4)]))
362                    as arrow::array::ArrayRef,
363            ),
364            (
365                "ts",
366                Arc::new(TimestampMillisecondArray::from(vec![1, 2, 3, 4])),
367            ),
368        ])
369        .unwrap()
370    }
371
372    async fn read(store: &ObjectStore, path: &str) -> ParquetRecordBatchReaderBuilder<Bytes> {
373        ParquetRecordBatchReaderBuilder::try_new(store.read(path).await.unwrap().to_bytes())
374            .unwrap()
375    }
376
377    #[tokio::test]
378    async fn abort_preserves_collisions_and_ambiguous_commits() {
379        use object_store::layers::mock::{MockLayerBuilder, MockWriterFactory, oio};
380
381        struct AmbiguousCommit(oio::Writer);
382        impl oio::Write for AmbiguousCommit {
383            async fn write(&mut self, bytes: object_store::Buffer) -> object_store::Result<()> {
384                self.0.write(bytes).await
385            }
386            async fn close(
387                &mut self,
388            ) -> object_store::Result<object_store::layers::mock::Metadata> {
389                self.0.close().await?;
390                Err(object_store::Error::new(
391                    object_store::ErrorKind::Unexpected,
392                    "lost close reply",
393                ))
394            }
395            async fn abort(&mut self) -> object_store::Result<()> {
396                Err(object_store::Error::new(
397                    object_store::ErrorKind::Unsupported,
398                    "cannot abort",
399                ))
400            }
401        }
402
403        for (creation, existing) in [
404            (ParquetCreationPolicy::Overwrite, false),
405            (ParquetCreationPolicy::IfNotExists, false),
406            (ParquetCreationPolicy::IfNotExists, true),
407        ] {
408            let directory = common_test_util::temp_dir::create_temp_dir("conditional_parquet");
409            let store = object_store::secure_fs::SecureFsRoot::open(directory.path())
410                .unwrap()
411                .build_operator();
412            let path = "conditional.parquet";
413            if existing {
414                store.write(path, "original").await.unwrap();
415            }
416            let factory: MockWriterFactory =
417                Arc::new(|_, _, writer| Box::new(AmbiguousCommit(writer)));
418            let store = store.layer(
419                MockLayerBuilder::default()
420                    .writer_factory(factory)
421                    .build()
422                    .unwrap(),
423            );
424            let mut writer = ParquetFileWriter::open_with_creation(
425                batch().schema(),
426                store.clone(),
427                path,
428                1,
429                None,
430                creation,
431            )
432            .await
433            .unwrap();
434            writer.write(batch(), None).await.unwrap();
435            assert!(writer.finish(None).await.is_err());
436            assert!(writer.abort().await.is_err());
437            if existing {
438                assert_eq!(
439                    store.read(path).await.unwrap().to_bytes(),
440                    Bytes::from_static(b"original")
441                );
442            } else {
443                assert_eq!(
444                    read(&store, path)
445                        .await
446                        .metadata()
447                        .file_metadata()
448                        .num_rows(),
449                    4
450                );
451            }
452        }
453    }
454
455    #[tokio::test]
456    async fn overwrite_abort_deletes_after_unsynced_close() {
457        use object_store::layers::mock::{MockLayerBuilder, MockWriterFactory, oio};
458
459        struct FailedClose(oio::Writer);
460        impl oio::Write for FailedClose {
461            async fn write(&mut self, bytes: object_store::Buffer) -> object_store::Result<()> {
462                self.0.write(bytes).await
463            }
464            async fn close(
465                &mut self,
466            ) -> object_store::Result<object_store::layers::mock::Metadata> {
467                Err(object_store::Error::new(
468                    object_store::ErrorKind::Unexpected,
469                    "close failed before sync",
470                ))
471            }
472            async fn abort(&mut self) -> object_store::Result<()> {
473                self.0.abort().await
474            }
475        }
476
477        let directory = common_test_util::temp_dir::create_temp_dir("unsynced_parquet");
478        let store = object_store::secure_fs::SecureFsRoot::open(directory.path())
479            .unwrap()
480            .build_operator();
481        let path = "partial.parquet";
482        let factory: MockWriterFactory = Arc::new(|_, _, writer| Box::new(FailedClose(writer)));
483        let store = store.layer(
484            MockLayerBuilder::default()
485                .writer_factory(factory)
486                .build()
487                .unwrap(),
488        );
489        let mut writer = ParquetFileWriter::open(batch().schema(), store.clone(), path, 1, None)
490            .await
491            .unwrap();
492        writer.write(batch(), None).await.unwrap();
493        assert!(writer.finish(None).await.is_err());
494        assert!(store.exists(path).await.unwrap());
495        writer.abort().await.unwrap();
496        assert!(!store.exists(path).await.unwrap());
497    }
498
499    #[tokio::test]
500    async fn large_footer_flush_uses_bounded_submissions_in_one_parquet_stream() {
501        use std::sync::Mutex;
502
503        use object_store::layers::mock::{Metadata, MockLayerBuilder, MockWriterFactory, oio};
504        struct ObservedWriter(oio::Writer, Arc<Mutex<Vec<usize>>>);
505        impl oio::Write for ObservedWriter {
506            async fn write(&mut self, bytes: object_store::Buffer) -> object_store::Result<()> {
507                self.1.lock().unwrap().push(bytes.len());
508                self.0.write(bytes).await
509            }
510            async fn close(&mut self) -> object_store::Result<Metadata> {
511                self.0.close().await
512            }
513            async fn abort(&mut self) -> object_store::Result<()> {
514                self.0.abort().await
515            }
516        }
517        let directory = common_test_util::temp_dir::create_temp_dir("bounded_parquet");
518        let store = object_store::secure_fs::SecureFsRoot::open(directory.path())
519            .unwrap()
520            .build_operator();
521        let sizes = Arc::new(Mutex::new(Vec::new()));
522        let factory: MockWriterFactory = Arc::new({
523            let sizes = sizes.clone();
524            move |_, _, writer| Box::new(ObservedWriter(writer, sizes.clone()))
525        });
526        let store = store.layer(
527            MockLayerBuilder::default()
528                .writer_factory(factory)
529                .build()
530                .unwrap(),
531        );
532        let mut state = 17u64;
533        let values = (0..600_000)
534            .map(|_| {
535                state ^= state << 13;
536                state ^= state >> 7;
537                state ^= state << 17;
538                state as i64
539            })
540            .collect::<Vec<_>>();
541        let array = Arc::new(Int64Array::from(values)) as arrow::array::ArrayRef;
542        let batch = RecordBatch::try_from_iter([("a", array.clone()), ("b", array)]).unwrap();
543        let mut writer =
544            ParquetFileWriter::open(batch.schema(), store.clone(), "large.parquet", 1, None)
545                .await
546                .unwrap();
547        // No storage-layer chunking: observe the application's actual submissions.
548        writer.sink = ParquetSink::Object(store.writer("large.parquet").await.unwrap());
549        writer.write(batch.clone(), None).await.unwrap();
550        writer.finish(None).await.unwrap();
551        let sizes = sizes.lock().unwrap().clone();
552        let limit = DEFAULT_WRITE_BUFFER_SIZE.as_bytes() as usize;
553        assert!(sizes.iter().sum::<usize>() > limit);
554        assert!(sizes.iter().all(|size| *size <= limit), "{sizes:?}");
555        let actual = read(&store, "large.parquet")
556            .await
557            .build()
558            .unwrap()
559            .collect::<std::result::Result<Vec<_>, _>>()
560            .unwrap();
561        assert_eq!(
562            arrow::compute::concat_batches(&batch.schema(), &actual).unwrap(),
563            batch
564        );
565    }
566
567    #[tokio::test]
568    async fn row_and_byte_limits_split_batches() {
569        let store = ObjectStore::new(object_store::services::Memory::default()).unwrap();
570        let batch = batch();
571        for (row_group_rows, flush_threshold_bytes, split) in [(2, usize::MAX, 3), (100, 1, 2)] {
572            let mut writer = ParquetFileWriter::open(
573                batch.schema(),
574                store.clone(),
575                "groups.parquet",
576                1,
577                Some(ParquetWriterLimits {
578                    row_group_rows,
579                    flush_threshold_bytes,
580                    max_row_groups: 2,
581                }),
582            )
583            .await
584            .unwrap();
585            writer.write(batch.slice(0, split), None).await.unwrap();
586            writer
587                .write(batch.slice(split, batch.num_rows() - split), None)
588                .await
589                .unwrap();
590            writer.finish(None).await.unwrap();
591            let reader = read(&store, "groups.parquet").await;
592            assert_eq!(reader.metadata().num_row_groups(), 2);
593            for group in reader.metadata().row_groups() {
594                assert_eq!(group.num_rows(), 2);
595                let encodings = group.column(1).encodings().collect::<Vec<_>>();
596                assert!(encodings.contains(&Encoding::DELTA_BINARY_PACKED));
597                assert!(!encodings.contains(&Encoding::RLE_DICTIONARY));
598                assert_eq!(
599                    group.column(0).compression(),
600                    Compression::ZSTD(ZstdLevel::default())
601                );
602            }
603            let actual = reader
604                .build()
605                .unwrap()
606                .collect::<std::result::Result<Vec<_>, _>>()
607                .unwrap();
608            assert_eq!(
609                arrow::compute::concat_batches(&batch.schema(), &actual).unwrap(),
610                batch
611            );
612        }
613    }
614
615    #[tokio::test]
616    async fn oversized_batch_stops_at_footer_limit_and_can_abort() {
617        let store = ObjectStore::new(object_store::services::Memory::default()).unwrap();
618        let batch = batch();
619        let mut writer = ParquetFileWriter::open(
620            batch.schema(),
621            store.clone(),
622            "limited.parquet",
623            1,
624            Some(ParquetWriterLimits {
625                row_group_rows: 2,
626                flush_threshold_bytes: usize::MAX,
627                max_row_groups: 1,
628            }),
629        )
630        .await
631        .unwrap();
632        let err = writer.write(batch, None).await.unwrap_err();
633        assert!(matches!(err, error::Error::ParquetWriterResource { .. }));
634        assert_eq!(err.status_code(), StatusCode::Suspended);
635        writer.abort().await.unwrap();
636        assert!(!store.exists("limited.parquet").await.unwrap());
637    }
638
639    #[tokio::test]
640    async fn copy_stream_preserves_values_and_empty_schema() {
641        let store = ObjectStore::new(object_store::services::Memory::default()).unwrap();
642        let batch = batch();
643        for (path, batches) in [
644            ("copy.parquet", vec![batch.clone()]),
645            ("empty.parquet", vec![]),
646        ] {
647            let stream = RecordBatchStreamAdapter::new(
648                batch.schema(),
649                futures::stream::iter(batches.clone().into_iter().map(Ok)),
650            );
651            assert_eq!(
652                stream_to_parquet(Box::pin(stream), store.clone(), path, 2)
653                    .await
654                    .unwrap(),
655                batches.iter().map(RecordBatch::num_rows).sum::<usize>()
656            );
657            let reader = read(&store, path).await;
658            assert_eq!(reader.schema().fields(), batch.schema().fields());
659            let actual = reader
660                .build()
661                .unwrap()
662                .collect::<std::result::Result<Vec<_>, _>>()
663                .unwrap();
664            assert_eq!(actual, batches);
665        }
666    }
667
668    #[tokio::test]
669    async fn failed_copy_preserves_untouched_securefs_destination() {
670        let directory = common_test_util::temp_dir::create_temp_dir("copy_existing");
671        let path = directory.path().join("existing.parquet");
672        std::fs::write(&path, b"original bytes").unwrap();
673        let access = crate::object_store::LocalFileAccess::sandboxed(directory.path()).unwrap();
674        let store = crate::object_store::build_backend_for_write(
675            &format!("{}/", directory.path().display()),
676            &Default::default(),
677            &access,
678        )
679        .await
680        .unwrap();
681        let stream = RecordBatchStreamAdapter::new(
682            batch().schema(),
683            futures::stream::iter(vec![Err(datafusion::error::DataFusionError::Execution(
684                "injected input failure".into(),
685            ))]),
686        );
687        let err = stream_to_parquet(Box::pin(stream), store, "existing.parquet", 1)
688            .await
689            .unwrap_err();
690        assert!(matches!(err, error::Error::ReadRecordBatch { .. }));
691        assert_eq!(std::fs::read(path).unwrap(), b"original bytes");
692    }
693
694    struct PausedFooterWriter {
695        inner: object_store::layers::mock::oio::Writer,
696        paused: bool,
697        started: Arc<tokio::sync::Notify>,
698        release: Arc<tokio::sync::Notify>,
699        closed: Arc<std::sync::atomic::AtomicBool>,
700    }
701
702    impl object_store::layers::mock::oio::Write for PausedFooterWriter {
703        async fn write(&mut self, bytes: object_store::Buffer) -> object_store::Result<()> {
704            if !self.paused {
705                self.paused = true;
706                self.started.notify_one();
707                self.release.notified().await;
708            }
709            self.inner.write(bytes).await
710        }
711        async fn close(&mut self) -> object_store::Result<object_store::layers::mock::Metadata> {
712            self.closed.store(true, std::sync::atomic::Ordering::SeqCst);
713            self.inner.close().await
714        }
715        async fn abort(&mut self) -> object_store::Result<()> {
716            self.inner.abort().await
717        }
718    }
719
720    #[tokio::test]
721    async fn cancellation_during_footer_write_waits_then_aborts_before_close() {
722        use object_store::layers::mock::{MockLayerBuilder, MockWriterFactory};
723        let started = Arc::new(tokio::sync::Notify::new());
724        let release = Arc::new(tokio::sync::Notify::new());
725        let closed = Arc::new(std::sync::atomic::AtomicBool::new(false));
726        let factory: MockWriterFactory = Arc::new({
727            let (started, release, closed) = (started.clone(), release.clone(), closed.clone());
728            move |_, _, inner| {
729                Box::new(PausedFooterWriter {
730                    inner,
731                    paused: false,
732                    started: started.clone(),
733                    release: release.clone(),
734                    closed: closed.clone(),
735                })
736            }
737        });
738        let store = ObjectStore::new(object_store::services::Memory::default())
739            .unwrap()
740            .layer(
741                MockLayerBuilder::default()
742                    .writer_factory(factory)
743                    .build()
744                    .unwrap(),
745            );
746        let mut writer =
747            ParquetFileWriter::open(batch().schema(), store.clone(), "cancel.parquet", 1, None)
748                .await
749                .unwrap();
750        // Force footer bytes through the sink before its close operation.
751        writer.sink =
752            ParquetSink::Object(store.writer_with("cancel.parquet").chunk(1).await.unwrap());
753        let cancellation = CancellationToken::new();
754        let result = {
755            let finish = writer.finish(Some(&cancellation));
756            tokio::pin!(finish);
757            tokio::select! {
758                result = &mut finish => panic!("finished before footer write: {result:?}"),
759                _ = started.notified() => {},
760            }
761            cancellation.cancel();
762            assert!(futures::poll!(&mut finish).is_pending());
763            release.notify_one();
764            finish.await
765        };
766        assert!(matches!(
767            result,
768            Err(error::Error::ParquetWriteCancelled {})
769        ));
770        assert!(!closed.load(std::sync::atomic::Ordering::SeqCst));
771        writer.abort().await.unwrap();
772        assert!(!store.exists("cancel.parquet").await.unwrap());
773    }
774}