1use std::io::Write;
2
3use arrow::datatypes::SchemaRef;
4use arrow::record_batch::RecordBatch;
5use parquet::arrow::ArrowWriter;
6use parquet::basic::{Compression, GzipLevel, ZstdLevel};
7use parquet::file::properties::WriterProperties;
8
9use crate::config::CompressionType;
10use crate::error::Result;
11
12pub struct ParquetFormat {
13 compression: CompressionType,
14 compression_level: Option<u32>,
15 row_group_rows: Option<usize>,
17}
18
19impl ParquetFormat {
20 pub fn new(
21 compression: CompressionType,
22 compression_level: Option<u32>,
23 row_group_rows: Option<usize>,
24 ) -> Self {
25 Self {
26 compression,
27 compression_level,
28 row_group_rows,
29 }
30 }
31
32 fn build_compression(&self) -> Compression {
33 match self.compression {
34 CompressionType::Zstd => {
35 let level = self.compression_level.unwrap_or(3) as i32;
36 Compression::ZSTD(ZstdLevel::try_new(level).unwrap_or_default())
37 }
38 CompressionType::Snappy => Compression::SNAPPY,
39 CompressionType::Gzip => {
40 let level = self.compression_level.unwrap_or(6);
41 Compression::GZIP(GzipLevel::try_new(level).unwrap_or_default())
42 }
43 CompressionType::Lz4 => Compression::LZ4_RAW,
48 CompressionType::None => Compression::UNCOMPRESSED,
49 }
50 }
51}
52
53pub struct ParquetFormatWriter {
54 inner: ArrowWriter<Box<dyn Write + Send>>,
55}
56
57impl super::Format for ParquetFormat {
58 fn create_writer(
59 &self,
60 schema: &SchemaRef,
61 writer: Box<dyn Write + Send>,
62 ) -> Result<Box<dyn super::FormatWriter + Send>> {
63 if let Some(field) = schema.fields().iter().find(|f| {
70 matches!(
71 f.data_type(),
72 arrow::datatypes::DataType::Decimal128(_, s)
73 | arrow::datatypes::DataType::Decimal256(_, s) if *s < 0
74 )
75 }) {
76 anyhow::bail!(
77 "Parquet cannot write column '{}' ({:?}): a NEGATIVE decimal scale is valid in \
78 the source and in Arrow but not in the Parquet decimal type (scale must be >= 0). \
79 Cast the column to a non-negative scale in the query (e.g. \
80 `round(col)::numeric(p,0)`), or use `format: csv`.",
81 field.name(),
82 field.data_type()
83 );
84 }
85 let mut builder = WriterProperties::builder()
93 .set_compression(self.build_compression())
94 .set_created_by("rivet".to_string());
95 if self.row_group_rows.is_some() {
96 builder = builder.set_max_row_group_row_count(self.row_group_rows);
97 }
98 let props = builder.build();
99
100 let inner = ArrowWriter::try_new(writer, schema.clone(), Some(props))?;
101 Ok(Box::new(ParquetFormatWriter { inner }))
102 }
103
104 fn file_extension(&self) -> &str {
105 "parquet"
106 }
107}
108
109impl super::FormatWriter for ParquetFormatWriter {
110 fn write_batch(&mut self, batch: &RecordBatch) -> Result<()> {
111 self.inner.write(batch)?;
112 Ok(())
113 }
114
115 fn finish(self: Box<Self>) -> Result<()> {
116 self.inner.close()?;
117 Ok(())
118 }
119
120 fn bytes_written(&self) -> u64 {
121 self.inner.bytes_written() as u64
122 }
123}
124
125#[cfg(test)]
126mod tests {
127 use super::*;
128 use crate::format::Format;
129 use arrow::array::Int64Array;
130 use arrow::datatypes::{DataType, Field, Schema};
131 use std::sync::Arc;
132
133 fn int64_schema() -> Arc<Schema> {
134 Arc::new(Schema::new(vec![Field::new("id", DataType::Int64, false)]))
135 }
136
137 fn one_batch(schema: &Arc<Schema>) -> arrow::record_batch::RecordBatch {
138 arrow::record_batch::RecordBatch::try_new(
139 schema.clone(),
140 vec![Arc::new(Int64Array::from(vec![1i64, 2, 3]))],
141 )
142 .unwrap()
143 }
144
145 fn make_writer(
146 compression: CompressionType,
147 level: Option<u32>,
148 ) -> Box<dyn crate::format::FormatWriter> {
149 let schema = int64_schema();
150 ParquetFormat::new(compression, level, None)
151 .create_writer(&schema, Box::new(Vec::<u8>::new()))
152 .expect("create_writer should succeed")
153 }
154
155 #[test]
158 fn file_extension_is_parquet() {
159 assert_eq!(
160 ParquetFormat::new(CompressionType::None, None, None).file_extension(),
161 "parquet"
162 );
163 }
164
165 #[test]
168 fn create_writer_zstd_default_level_succeeds() {
169 let _ = make_writer(CompressionType::Zstd, None);
170 }
171
172 #[test]
173 fn create_writer_zstd_explicit_level_succeeds() {
174 let _ = make_writer(CompressionType::Zstd, Some(9));
175 }
176
177 #[test]
178 fn create_writer_snappy_succeeds() {
179 let _ = make_writer(CompressionType::Snappy, None);
180 }
181
182 #[test]
183 fn create_writer_gzip_succeeds() {
184 let _ = make_writer(CompressionType::Gzip, None);
185 }
186
187 #[test]
188 fn create_writer_lz4_succeeds() {
189 let _ = make_writer(CompressionType::Lz4, None);
190 }
191
192 #[test]
193 fn create_writer_uncompressed_succeeds() {
194 let _ = make_writer(CompressionType::None, None);
195 }
196
197 #[test]
200 fn write_batch_and_finish_returns_ok() {
201 let schema = int64_schema();
202 let fmt = ParquetFormat::new(CompressionType::Zstd, None, None);
203 let mut writer = fmt
205 .create_writer(&schema, Box::new(Vec::<u8>::new()))
206 .unwrap();
207 writer.write_batch(&one_batch(&schema)).unwrap();
208 writer.finish().unwrap(); }
210
211 #[test]
212 fn finish_without_write_produces_valid_empty_parquet() {
213 let schema = int64_schema();
214 let fmt = ParquetFormat::new(CompressionType::None, None, None);
215 let writer = fmt
217 .create_writer(&schema, Box::new(Vec::<u8>::new()))
218 .unwrap();
219 writer.finish().unwrap();
220 }
221
222 #[test]
225 fn row_group_rows_none_uses_library_default() {
226 let schema = int64_schema();
227 let fmt = ParquetFormat::new(CompressionType::None, None, None);
228 let mut writer = fmt
229 .create_writer(&schema, Box::new(Vec::<u8>::new()))
230 .unwrap();
231 writer.write_batch(&one_batch(&schema)).unwrap();
232 writer.finish().unwrap();
233 }
234
235 #[test]
236 fn row_group_rows_some_succeeds() {
237 let schema = int64_schema();
238 let fmt = ParquetFormat::new(CompressionType::None, None, Some(100));
239 let mut writer = fmt
240 .create_writer(&schema, Box::new(Vec::<u8>::new()))
241 .unwrap();
242 writer.write_batch(&one_batch(&schema)).unwrap();
243 writer.finish().unwrap();
244 }
245
246 #[test]
247 fn lz4_maps_to_the_standard_raw_codec_not_hadoop_framed() {
248 assert_eq!(
252 ParquetFormat::new(CompressionType::Lz4, None, None).build_compression(),
253 Compression::LZ4_RAW
254 );
255 }
256
257 #[test]
258 fn negative_scale_decimal_is_refused_loudly_at_writer_creation() {
259 use arrow::datatypes::{DataType, Field, Schema};
263 let schema = std::sync::Arc::new(Schema::new(vec![Field::new(
264 "amount",
265 DataType::Decimal128(10, -2),
266 true,
267 )]));
268 let result = ParquetFormat::new(CompressionType::None, None, None)
269 .create_writer(&schema, Box::new(Vec::<u8>::new()));
270 assert!(
271 result.is_err(),
272 "a negative-scale decimal must be refused, not crash mid-export"
273 );
274 let msg = result.err().unwrap().to_string();
275 assert!(
276 msg.contains("NEGATIVE decimal scale") && msg.contains("amount"),
277 "error must name the column + the negative-scale cause: {msg}"
278 );
279 }
280
281 fn write_batch_to_bytes(compression: CompressionType) -> Vec<u8> {
284 let schema = int64_schema();
285 let tmp = tempfile::NamedTempFile::new().unwrap();
286 let file = std::fs::File::create(tmp.path()).unwrap();
287 let mut w = ParquetFormat::new(compression, None, None)
288 .create_writer(&schema, Box::new(file))
289 .unwrap();
290 w.write_batch(&one_batch(&schema)).unwrap();
291 w.finish().unwrap();
292 std::fs::read(tmp.path()).unwrap()
293 }
294
295 #[test]
296 fn output_is_byte_deterministic_for_identical_rows() {
297 let a = write_batch_to_bytes(CompressionType::Zstd);
300 let b = write_batch_to_bytes(CompressionType::Zstd);
301 assert_eq!(a, b, "identical rows must yield byte-identical parquet");
302 }
303
304 #[test]
305 fn created_by_is_pinned_and_version_free() {
306 use parquet::file::reader::{FileReader, SerializedFileReader};
307 let bytes = write_batch_to_bytes(CompressionType::None);
308 let reader = SerializedFileReader::new(bytes::Bytes::from(bytes)).unwrap();
309 let created_by = reader.metadata().file_metadata().created_by();
310 assert_eq!(
311 created_by,
312 Some("rivet"),
313 "created_by must be the pinned constant"
314 );
315 let cb = created_by.unwrap();
318 assert!(
319 !cb.contains("version") && !cb.contains("parquet"),
320 "created_by must not embed the library version: {cb:?}"
321 );
322 }
323}