Skip to main content

parquet_derive/
lib.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18//! This crate provides a procedural macro to derive
19//! implementations of a RecordWriter and RecordReader
20
21#![doc(
22    html_logo_url = "https://raw.githubusercontent.com/apache/parquet-format/25f05e73d8cd7f5c83532ce51cb4f4de8ba5f2a2/logo/parquet-logos_1.svg",
23    html_favicon_url = "https://raw.githubusercontent.com/apache/parquet-format/25f05e73d8cd7f5c83532ce51cb4f4de8ba5f2a2/logo/parquet-logos_1.svg"
24)]
25#![cfg_attr(docsrs, feature(doc_cfg))]
26#![deny(clippy::allow_attributes)]
27#![warn(missing_docs)]
28#![recursion_limit = "128"]
29
30#[macro_use]
31extern crate quote;
32
33use ::syn::{Data, DataStruct, DeriveInput, ext::IdentExt, parse_macro_input};
34
35mod parquet_field;
36
37/// Derive flat, simple RecordWriter implementations.
38///
39/// Works by parsing a struct tagged with `#[derive(ParquetRecordWriter)]` and emitting
40/// the correct writing code for each field of the struct. Column writers
41/// are generated in the order they are defined.
42///
43/// It is up to the programmer to keep the order of the struct
44/// fields lined up with the schema.
45///
46/// Example:
47///
48/// ```rust
49/// use parquet::file::properties::WriterProperties;
50/// use parquet::file::writer::SerializedFileWriter;
51/// use parquet::record::RecordWriter;
52/// use parquet_derive::ParquetRecordWriter;
53/// use std::fs::File;
54/// use std::sync::Arc;
55///
56/// // For reader
57/// use parquet::file::reader::{FileReader, SerializedFileReader};
58/// use parquet::record::RecordReader;
59/// use parquet_derive::ParquetRecordReader;
60///
61/// #[derive(Debug, ParquetRecordWriter, ParquetRecordReader)]
62/// struct ACompleteRecord {
63///     pub a_bool: bool,
64///     pub a_string: String,
65/// }
66///
67/// fn write_some_records() {
68///     let samples = vec![
69///         ACompleteRecord {
70///             a_bool: true,
71///             a_string: "I'm true".into(),
72///         },
73///         ACompleteRecord {
74///             a_bool: false,
75///             a_string: "I'm false".into(),
76///         },
77///     ];
78///
79///     let schema = samples.as_slice().schema().unwrap();
80///
81///     let props = Arc::new(WriterProperties::builder().build());
82///
83///     let file = File::create("example.parquet").unwrap();
84///
85///     let mut writer = SerializedFileWriter::new(file, schema, props).unwrap();
86///
87///     let mut row_group = writer.next_row_group().unwrap();
88///
89///     samples
90///         .as_slice()
91///         .write_to_row_group(&mut row_group)
92///         .unwrap();
93///
94///     row_group.close().unwrap();
95///
96///     writer.close().unwrap();
97/// }
98///
99/// fn read_some_records() -> Vec<ACompleteRecord> {
100///     let mut samples: Vec<ACompleteRecord> = Vec::new();
101///     let file = File::open("example.parquet").unwrap();
102///
103///     let reader = SerializedFileReader::new(file).unwrap();
104///     let mut row_group = reader.get_row_group(0).unwrap();
105///     samples.read_from_row_group(&mut *row_group, 2).unwrap();
106///
107///     samples
108/// }
109///
110/// pub fn main() {
111///     write_some_records();
112///
113///     let records = read_some_records();
114///
115///     std::fs::remove_file("example.parquet").unwrap();
116///
117///     assert_eq!(
118///         format!("{:?}", records),
119///         "[ACompleteRecord { a_bool: true, a_string: \"I'm true\" }, ACompleteRecord { a_bool: false, a_string: \"I'm false\" }]"
120///     );
121/// }
122/// ```
123///
124#[proc_macro_derive(ParquetRecordWriter)]
125pub fn parquet_record_writer(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
126    let input: DeriveInput = parse_macro_input!(input as DeriveInput);
127    let fields = match input.data {
128        Data::Struct(DataStruct { fields, .. }) => fields,
129        Data::Enum(_) => unimplemented!("Enum currently is not supported"),
130        Data::Union(_) => unimplemented!("Union currently is not supported"),
131    };
132
133    let field_infos: Vec<_> = fields.iter().map(parquet_field::Field::from).collect();
134
135    let writer_snippets: Vec<proc_macro2::TokenStream> =
136        field_infos.iter().map(|x| x.writer_snippet()).collect();
137
138    let derived_for = input.ident;
139    let generics = input.generics;
140
141    let field_types: Vec<proc_macro2::TokenStream> =
142        field_infos.iter().map(|x| x.parquet_type()).collect();
143
144    (quote! {
145    impl #generics ::parquet::record::RecordWriter<#derived_for #generics> for &[#derived_for #generics] {
146      fn write_to_row_group<W: ::std::io::Write + Send>(
147        &self,
148        row_group_writer: &mut ::parquet::file::writer::SerializedRowGroupWriter<'_, W>
149      ) -> ::std::result::Result<(), ::parquet::errors::ParquetError> {
150        use ::parquet::column::writer::ColumnWriter;
151
152        let mut row_group_writer = row_group_writer;
153        let records = &self; // Used by all the writer snippets to be more clear
154
155        #(
156          {
157              let mut some_column_writer = row_group_writer.next_column().unwrap();
158              if let Some(mut column_writer) = some_column_writer {
159                  #writer_snippets
160                  column_writer.close()?;
161              } else {
162                  return Err(::parquet::errors::ParquetError::General("Failed to get next column".into()))
163              }
164          }
165        );*
166
167        Ok(())
168      }
169
170      fn schema(&self) -> ::std::result::Result<::parquet::schema::types::TypePtr, ::parquet::errors::ParquetError> {
171        use ::parquet::schema::types::Type as ParquetType;
172        use ::parquet::schema::types::TypePtr;
173        use ::parquet::basic::LogicalType;
174
175        let mut fields: ::std::vec::Vec<TypePtr> = ::std::vec::Vec::new();
176        #(
177          #field_types
178        );*;
179        let group = ParquetType::group_type_builder("rust_schema")
180          .with_fields(fields)
181          .build()?;
182        Ok(group.into())
183      }
184    }
185  }).into()
186}
187
188/// Derive flat, simple RecordReader implementations.
189///
190/// Works by parsing a struct tagged with `#[derive(ParquetRecordReader)]` and emitting
191/// the correct writing code for each field of the struct. Column readers
192/// are generated by matching names in the schema to the names in the struct.
193///
194/// It is up to the programmer to ensure the names in the struct
195/// fields line up with the schema.
196///
197/// Example:
198///
199/// ```rust
200/// use parquet::record::RecordReader;
201/// use parquet::file::{serialized_reader::SerializedFileReader, reader::FileReader};
202/// use parquet_derive::{ParquetRecordReader};
203/// use std::fs::File;
204///
205/// #[derive(ParquetRecordReader)]
206/// struct ACompleteRecord {
207///     pub a_bool: bool,
208///     pub a_string: String,
209/// }
210///
211/// pub fn read_some_records() -> Vec<ACompleteRecord> {
212///   let mut samples: Vec<ACompleteRecord> = Vec::new();
213///   let file = File::open("some_file.parquet").unwrap();
214///
215///   let reader = SerializedFileReader::new(file).unwrap();
216///   let mut row_group = reader.get_row_group(0).unwrap();
217///   samples.read_from_row_group(&mut *row_group, 1).unwrap();
218///   samples
219/// }
220/// ```
221///
222#[proc_macro_derive(ParquetRecordReader)]
223pub fn parquet_record_reader(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
224    let input: DeriveInput = parse_macro_input!(input as DeriveInput);
225    let fields = match input.data {
226        Data::Struct(DataStruct { fields, .. }) => fields,
227        Data::Enum(_) => unimplemented!("Enum currently is not supported"),
228        Data::Union(_) => unimplemented!("Union currently is not supported"),
229    };
230
231    let field_infos: Vec<_> = fields.iter().map(parquet_field::Field::from).collect();
232    let field_names: Vec<_> = fields.iter().map(|f| f.ident.clone()).collect();
233    // unraw the identifiers, so raw identifiers like `r#type` are looked
234    // up by their column name `type` in the parquet file
235    let field_names_str: Vec<_> = fields
236        .iter()
237        .map(|f| {
238            f.ident
239                .as_ref()
240                .expect("Only structs with named fields are currently supported")
241                .unraw()
242                .to_string()
243        })
244        .collect();
245    let reader_snippets: Vec<proc_macro2::TokenStream> =
246        field_infos.iter().map(|x| x.reader_snippet()).collect();
247
248    let derived_for = input.ident;
249    let generics = input.generics;
250
251    (quote! {
252
253    impl #generics ::parquet::record::RecordReader<#derived_for #generics> for Vec<#derived_for #generics> {
254      fn read_from_row_group(
255        &mut self,
256        row_group_reader: &mut dyn ::parquet::file::reader::RowGroupReader,
257        num_records: usize,
258      ) -> ::std::result::Result<(), ::parquet::errors::ParquetError> {
259        use ::parquet::column::reader::ColumnReader;
260
261        let mut row_group_reader = row_group_reader;
262
263        // key: parquet file column name, value: column index
264        let mut name_to_index = std::collections::HashMap::new();
265        for (idx, col) in row_group_reader.metadata().schema_descr().columns().iter().enumerate() {
266            name_to_index.insert(col.name().to_string(), idx);
267        }
268
269        for _ in 0..num_records {
270          self.push(#derived_for {
271            #(
272              #field_names: Default::default()
273            ),*
274          })
275        }
276
277        let records = self; // Used by all the reader snippets to be more clear
278
279        #(
280          {
281              let idx: usize = match name_to_index.get(#field_names_str) {
282                Some(&col_idx) => col_idx,
283                None => {
284                  let error_msg = format!("column name '{}' is not found in parquet file!", #field_names_str);
285                  return Err(::parquet::errors::ParquetError::General(error_msg));
286                }
287              };
288              if let Ok(column_reader) = row_group_reader.get_column_reader(idx) {
289                  #reader_snippets
290              } else {
291                  return Err(::parquet::errors::ParquetError::General("Failed to get next column".into()))
292              }
293          }
294        );*
295
296        Ok(())
297      }
298    }
299  }).into()
300}