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