Skip to main content

datui_lib/
avro_types.rs

1//! Columns cast to types Polars' Avro writer can hold, for an Avro export only.
2//!
3//! The writer knows booleans, 32- and 64-bit integers and floats, strings,
4//! binary, dates, naive millisecond and microsecond datetimes, and lists and
5//! structs of those. Anything else fails the whole export with "not yet
6//! implemented", so each is cast to the nearest type it does know. The casts
7//! are strict: a value that does not fit (a `u64` past `i64::MAX`) fails the
8//! export by column name instead of turning null.
9//!
10//! It also writes decimals, but wrongly: it drops the sign byte of a positive
11//! value whose leading byte is 0x80 or more, so every reader sees 327.68 as
12//! -327.68. Decimals are written as their exact text instead.
13//!
14//! And it writes names as they are, with an empty record name, but an Avro name
15//! is `[A-Za-z_][A-Za-z0-9_]*` and strict readers refuse the file. [`write`]
16//! names the record and gives each column and struct field a valid name in the
17//! file's schema, with the original as the field's `doc`. It also writes the
18//! header once, where Polars' writer repeats it for every chunk, and cuts
19//! blocks by size rather than one per chunk.
20
21use std::io::Write;
22
23use polars::prelude::*;
24use polars_arrow::io::avro::avro_schema::file::CompressedBlock;
25use polars_arrow::io::avro::avro_schema::schema::{Field as AvroField, Schema as AvroSchema};
26use polars_arrow::io::avro::{avro_schema, write as avro_write};
27
28/// The record name of an export. Polars' default is empty, which strict readers
29/// refuse; its nested records are `r1`, `r2`, ..., so this never clashes.
30pub const RECORD_NAME: &str = "Row";
31
32fn is_avro_name(name: &str) -> bool {
33    let mut chars = name.chars();
34    chars
35        .next()
36        .is_some_and(|c| c.is_ascii_alphabetic() || c == '_')
37        && chars.all(|c| c.is_ascii_alphanumeric() || c == '_')
38}
39
40/// Valid Avro names for `names`, in order. A valid name is kept; in any other,
41/// each character outside `[A-Za-z0-9_]` becomes `_`, a leading digit gets a
42/// `_` in front, and a name then taken gets the first free `_2`, `_3`, ...
43/// Valid names are reserved first, so a column already valid is never renamed.
44fn avro_names<'a>(names: impl IntoIterator<Item = &'a str>) -> Vec<String> {
45    let names: Vec<&str> = names.into_iter().collect();
46    let mut taken: PlHashSet<String> = names
47        .iter()
48        .filter(|name| is_avro_name(name))
49        .map(|name| name.to_string())
50        .collect();
51    names
52        .iter()
53        .map(|&name| {
54            if is_avro_name(name) {
55                return name.to_string();
56            }
57            let mut base: String = name
58                .chars()
59                .map(|c| if c.is_ascii_alphanumeric() { c } else { '_' })
60                .collect();
61            if !base.starts_with(|c: char| c.is_ascii_alphabetic() || c == '_') {
62                base.insert(0, '_');
63            }
64            let name = if taken.contains(&base) {
65                (2..)
66                    .map(|n| format!("{base}_{n}"))
67                    .find(|candidate| !taken.contains(candidate))
68                    .expect("some suffix is free")
69            } else {
70                base
71            };
72            taken.insert(name.clone());
73            name
74        })
75        .collect()
76}
77
78/// Whether an Avro export renames this column or a struct field inside it.
79pub fn renames(name: &str, dtype: &DataType) -> bool {
80    fn inside(dtype: &DataType) -> bool {
81        match dtype {
82            DataType::List(inner) | DataType::Array(inner, _) => inside(inner),
83            DataType::Struct(fields) => fields.iter().any(|f| renames(f.name(), f.dtype())),
84            _ => false,
85        }
86    }
87    !is_avro_name(name) || inside(dtype)
88}
89
90/// `fields` under valid Avro names at any depth, each renamed one keeping its
91/// original name as its `doc`. Only the schema changes: values are written by
92/// position.
93fn name_fields(fields: &mut [AvroField]) {
94    fn name_schema(schema: &mut AvroSchema) {
95        match schema {
96            AvroSchema::Union(branches) => branches.iter_mut().for_each(name_schema),
97            AvroSchema::Array(items) | AvroSchema::Map(items) => name_schema(items),
98            AvroSchema::Record(record) => name_fields(&mut record.fields),
99            _ => {}
100        }
101    }
102    let names = avro_names(fields.iter().map(|f| f.name.as_str()));
103    for (field, name) in fields.iter_mut().zip(names) {
104        if field.name != name {
105            field.doc = Some(std::mem::replace(&mut field.name, name));
106        }
107        name_schema(&mut field.schema);
108    }
109}
110
111/// A block is cut once it holds this many bytes, however the frame is chunked:
112/// a reader holds a whole block in memory, and Java's refuses one past 2 GiB.
113const BLOCK_BYTES: usize = 1 << 20;
114
115/// Write `df`, prepared by [`lazy_for_avro`], as an uncompressed Avro file:
116/// Polars' `AvroWriter`'s encoding, with valid names, one header, and blocks
117/// of about [`BLOCK_BYTES`].
118pub fn write(df: &mut DataFrame, mut writer: impl Write) -> PolarsResult<()> {
119    // Serializing walks the columns' chunks together, so they must line up.
120    df.align_chunks_par();
121    let schema = df.schema().to_arrow(CompatLevel::oldest());
122    let mut record = avro_write::to_record(&schema, RECORD_NAME.to_string())?;
123    name_fields(&mut record.fields);
124    // avro-schema turns any I/O error into "OutOfSpec", so it only ever writes
125    // to `out` and the file's errors, such as a full disk, come from here.
126    let mut out = Vec::new();
127    avro_schema::write::write_metadata(&mut out, record.clone(), None)?;
128    writer.write_all(&out)?;
129
130    let mut block = CompressedBlock::default();
131    let mut flush = |block: &mut CompressedBlock| -> PolarsResult<()> {
132        out.clear();
133        avro_schema::write::write_block(&mut out, block)?;
134        writer.write_all(&out)?;
135        block.data.clear();
136        block.number_of_rows = 0;
137        Ok(())
138    };
139    for chunk in df.iter_chunks(CompatLevel::oldest(), true) {
140        let mut serializers: Vec<_> = chunk
141            .iter()
142            .zip(&record.fields)
143            .map(|(array, field)| avro_write::new_serializer(array.as_ref(), &field.schema))
144            .collect();
145        // Row by row, as Polars' `serialize` does, cutting a block by size.
146        for _ in 0..chunk.height() {
147            for serializer in &mut serializers {
148                let value = serializer.next().expect("a value for every row");
149                block.data.extend_from_slice(value);
150            }
151            block.number_of_rows += 1;
152            if block.data.len() >= BLOCK_BYTES {
153                flush(&mut block)?;
154            }
155        }
156    }
157    if block.number_of_rows > 0 {
158        flush(&mut block)?;
159    }
160    Ok(())
161}
162
163/// What `dtype` is written as. `durations` keeps times and durations as
164/// microsecond durations, the step before they become plain integers: a
165/// direct cast would count in each type's own unit, nanoseconds for a time.
166fn writable(dtype: &DataType, durations: bool) -> DataType {
167    use DataType as D;
168    match dtype {
169        D::Int8 | D::Int16 | D::UInt8 | D::UInt16 => D::Int32,
170        D::UInt32 | D::UInt64 => D::Int64,
171        D::Int128 | D::UInt128 | D::Decimal(..) => D::String,
172        D::Float16 => D::Float32,
173        // A time zone goes, the instant stays: the values read as UTC.
174        D::Datetime(unit, _) => D::Datetime(
175            match unit {
176                TimeUnit::Nanoseconds => TimeUnit::Microseconds,
177                unit => *unit,
178            },
179            None,
180        ),
181        D::Time | D::Duration(_) if durations => D::Duration(TimeUnit::Microseconds),
182        D::Time | D::Duration(_) => D::Int64,
183        D::Categorical(..) | D::Enum(..) | D::Null => D::String,
184        D::List(inner) | D::Array(inner, _) => D::List(Box::new(writable(inner, durations))),
185        D::Struct(fields) => D::Struct(
186            fields
187                .iter()
188                .map(|f| Field::new(f.name().clone(), writable(f.dtype(), durations)))
189                .collect(),
190        ),
191        other => other.clone(),
192    }
193}
194
195/// `lf` with every column Avro cannot hold cast to one it can, under its own
196/// name. Planned, not run.
197pub fn lazy_for_avro(mut lf: LazyFrame) -> PolarsResult<LazyFrame> {
198    let schema = lf.collect_schema()?;
199    let exprs: Vec<Expr> = schema
200        .iter()
201        .filter_map(|(name, dtype)| {
202            let target = writable(dtype, false);
203            if &target == dtype {
204                return None;
205            }
206            let step = writable(dtype, true);
207            let expr = col(name.clone());
208            let expr = if step == target {
209                expr
210            } else {
211                expr.strict_cast(step)
212            };
213            Some(expr.strict_cast(target))
214        })
215        .collect();
216    Ok(if exprs.is_empty() {
217        lf
218    } else {
219        lf.with_columns(exprs)
220    })
221}
222
223#[cfg(test)]
224mod tests {
225    use super::*;
226    use polars::io::avro::AvroReader;
227
228    fn written(lf: LazyFrame) -> Vec<u8> {
229        let mut df = lazy_for_avro(lf).unwrap().collect().unwrap();
230        let mut bytes = Vec::new();
231        write(&mut df, &mut bytes).unwrap();
232        bytes
233    }
234
235    fn round_trip(lf: LazyFrame) -> DataFrame {
236        AvroReader::new(std::io::Cursor::new(written(lf)))
237            .finish()
238            .unwrap()
239    }
240
241    fn category() -> DataType {
242        DataType::from_categories(Categories::global())
243    }
244
245    /// Every type the writer lacks comes back as the one it was cast to, with
246    /// its values; the types it has come back untouched.
247    #[test]
248    fn every_type_avro_lacks_is_written() {
249        let lf = df!(
250            "n" => [Some(3_600_000_000i64), None],
251            "s" => [Some("a"), Some("b")],
252        )
253        .unwrap()
254        .lazy()
255        .select([
256            col("n").alias("kept"),
257            col("n").cast(DataType::Int16).alias("i16"),
258            col("n").cast(DataType::UInt32).alias("u32"),
259            col("n").cast(DataType::UInt64).alias("u64"),
260            col("n").cast(DataType::Int128).alias("i128"),
261            lit(327.68).cast(DataType::Decimal(10, 2)).alias("dec"),
262            col("n")
263                .cast(DataType::Datetime(TimeUnit::Nanoseconds, None))
264                .alias("ns"),
265            col("n")
266                .cast(DataType::Datetime(
267                    TimeUnit::Microseconds,
268                    TimeZone::opt_try_new(Some("America/New_York")).unwrap(),
269                ))
270                .alias("zoned"),
271            col("n").cast(DataType::Time).alias("time"),
272            col("n")
273                .cast(DataType::Duration(TimeUnit::Milliseconds))
274                .alias("ms"),
275            col("s").cast(category()).alias("cat"),
276            lit(NULL).alias("null"),
277            col("n")
278                .fill_null(0)
279                .implode(true)
280                .cast(DataType::Array(Box::new(DataType::Int64), 2))
281                .alias("arr"),
282            col("s").cast(category()).implode(true).alias("cats"),
283            as_struct(vec![
284                col("s").cast(category()),
285                col("n").cast(DataType::Time),
286            ])
287            .alias("point"),
288        ]);
289        let back = round_trip(lf);
290        let dtype = |name: &str| back.column(name).unwrap().dtype().clone();
291        let first = |name: &str| back.column(name).unwrap().get(0).unwrap().into_static();
292        assert_eq!(dtype("kept"), DataType::Int64);
293        assert_eq!(dtype("i16"), DataType::Int32);
294        assert_eq!(dtype("u32"), DataType::Int64);
295        assert_eq!(dtype("u64"), DataType::Int64);
296        assert_eq!(first("i128"), AnyValue::StringOwned("3600000000".into()));
297        assert_eq!(
298            first("dec"),
299            AnyValue::StringOwned("327.68".into()),
300            "the writer's own decimal reads back as -327.68"
301        );
302        assert_eq!(
303            first("ns"),
304            AnyValue::Datetime(3_600_000, TimeUnit::Microseconds, None)
305        );
306        assert_eq!(
307            first("zoned"),
308            AnyValue::Datetime(3_600_000_000, TimeUnit::Microseconds, None),
309            "the UTC instant, not the wall time in New York"
310        );
311        assert_eq!(first("time"), AnyValue::Int64(3_600_000), "microseconds");
312        assert_eq!(
313            first("ms"),
314            AnyValue::Int64(3_600_000_000_000),
315            "microseconds, whatever the unit"
316        );
317        assert_eq!(dtype("cat"), DataType::String);
318        assert_eq!(first("cat"), AnyValue::StringOwned("a".into()));
319        assert_eq!(dtype("null"), DataType::String);
320        assert_eq!(dtype("arr"), DataType::List(Box::new(DataType::Int64)));
321        assert_eq!(dtype("cats"), DataType::List(Box::new(DataType::String)));
322        assert_eq!(
323            dtype("point"),
324            DataType::Struct(vec![
325                Field::new("s".into(), DataType::String),
326                Field::new("n".into(), DataType::Int64),
327            ])
328        );
329    }
330
331    /// A name Avro refuses is made valid; a valid one never moves, even when a
332    /// renamed one would land on it.
333    #[test]
334    fn names_are_made_valid_and_unique() {
335        let names = avro_names(["my col", "2024", "a-b", "a_b", "délai", "", "_2024", "ok_1"]);
336        assert_eq!(
337            names,
338            [
339                "my_col", "_2024_2", "a_b_2", "a_b", "d_lai", "_", "_2024", "ok_1"
340            ]
341        );
342        assert!(!renames("ok_1", &DataType::Int64));
343        assert!(renames("my col", &DataType::Int64));
344        let fields = |name: &str| {
345            DataType::List(Box::new(DataType::Struct(vec![Field::new(
346                name.into(),
347                DataType::Int64,
348            )])))
349        };
350        assert!(renames("ok", &fields("x y")));
351        assert!(!renames("ok", &fields("x_y")));
352    }
353
354    /// Struct fields are renamed inside a list too, with their values and
355    /// nulls, and the columns under them.
356    #[test]
357    fn struct_fields_are_renamed_at_any_depth() {
358        let lf = df!("n" => [Some(1i64), None, Some(3)])
359            .unwrap()
360            .lazy()
361            .select([
362                as_struct(vec![
363                    col("n").alias("x y"),
364                    (col("n") * lit(10)).alias("x-y"),
365                ])
366                .implode(true)
367                .alias("my list"),
368                when(col("n").is_null())
369                    .then(lit(NULL).cast(DataType::Struct(vec![Field::new(
370                        "1st".into(),
371                        DataType::Int64,
372                    )])))
373                    .otherwise(as_struct(vec![col("n").alias("1st")]))
374                    .alias("point"),
375            ]);
376        let back = round_trip(lf);
377        assert_eq!(
378            back.column("my_list").unwrap().dtype(),
379            &DataType::List(Box::new(DataType::Struct(vec![
380                Field::new("x_y".into(), DataType::Int64),
381                Field::new("x_y_2".into(), DataType::Int64),
382            ])))
383        );
384        let items = back
385            .column("my_list")
386            .unwrap()
387            .list()
388            .unwrap()
389            .get_as_series(0)
390            .unwrap();
391        let field = |name: &str| {
392            let values = items.struct_().unwrap().field_by_name(name).unwrap();
393            values.i64().unwrap().iter().collect::<Vec<_>>()
394        };
395        assert_eq!(field("x_y"), [Some(1), None, Some(3)]);
396        assert_eq!(field("x_y_2"), [Some(10), None, Some(30)]);
397        let point = back.column("point").unwrap();
398        assert_eq!(point.null_count(), 1, "{point:?}");
399        let first = point.struct_().unwrap().field_by_name("_1st").unwrap();
400        assert_eq!(first.i64().unwrap().get(2), Some(3));
401    }
402
403    /// A renamed field keeps its original name as its doc, nested ones too, and
404    /// a frame of several chunks is one header and a block for each.
405    #[test]
406    fn originals_are_docs_and_chunks_share_one_header() {
407        let part = df!("my col" => [1i64], "ok" => [2i64])
408            .unwrap()
409            .lazy()
410            .with_column(as_struct(vec![col("ok").alias("x y")]).alias("point"))
411            .collect()
412            .unwrap();
413        let mut df = part.clone();
414        df.vstack_mut(&part).unwrap();
415        assert_eq!(df.first_col_n_chunks(), 2);
416        let mut bytes = Vec::new();
417        write(&mut df, &mut bytes).unwrap();
418
419        let record = avro_schema::read::read_metadata(&mut std::io::Cursor::new(&bytes))
420            .unwrap()
421            .record;
422        assert_eq!(record.name, RECORD_NAME);
423        let docs = |fields: &[AvroField]| -> Vec<(String, Option<String>)> {
424            fields
425                .iter()
426                .map(|f| (f.name.clone(), f.doc.clone()))
427                .collect()
428        };
429        assert_eq!(
430            docs(&record.fields),
431            [
432                ("my_col".to_string(), Some("my col".to_string())),
433                ("ok".to_string(), None),
434                ("point".to_string(), None),
435            ]
436        );
437        let AvroSchema::Union(branches) = &record.fields[2].schema else {
438            panic!("{:?}", record.fields[2].schema);
439        };
440        let AvroSchema::Record(point) = &branches[1] else {
441            panic!("{branches:?}");
442        };
443        assert_eq!(
444            docs(&point.fields),
445            [("x_y".to_string(), Some("x y".to_string()))]
446        );
447
448        let back = AvroReader::new(std::io::Cursor::new(bytes))
449            .finish()
450            .unwrap();
451        let my_col = back.column("my_col").unwrap().i64().unwrap();
452        assert_eq!(my_col.iter().collect::<Vec<_>>(), [Some(1), Some(1)]);
453    }
454
455    /// Blocks are cut by size, not by chunk: a hundred one-row chunks are one
456    /// block, and one chunk of about 3 MiB is three, and every row reads back.
457    #[test]
458    fn blocks_are_cut_by_size() {
459        use avro_schema::read::fallible_streaming_iterator::FallibleStreamingIterator;
460        fn blocks(df: &mut DataFrame) -> Vec<(usize, usize)> {
461            let mut bytes = Vec::new();
462            write(df, &mut bytes).unwrap();
463            let back = AvroReader::new(std::io::Cursor::new(&bytes))
464                .finish()
465                .unwrap();
466            assert!(back.equals(df), "{back:?}");
467            let mut reader = std::io::Cursor::new(bytes);
468            let marker = avro_schema::read::read_metadata(&mut reader)
469                .unwrap()
470                .marker;
471            let mut iter = avro_schema::read::block_iterator(reader, None, marker);
472            let mut sizes = Vec::new();
473            while let Some(block) = iter.next().unwrap() {
474                sizes.push((block.number_of_rows, block.data.len()));
475            }
476            sizes
477        }
478        let text = "x".repeat(1000);
479        let row = df!("s" => [text.as_str()]).unwrap();
480        let mut many = row.clone();
481        for _ in 0..99 {
482            many.vstack_mut(&row).unwrap();
483        }
484        assert_eq!(many.first_col_n_chunks(), 100);
485        let sizes = blocks(&mut many);
486        assert_eq!(sizes.len(), 1, "{sizes:?}");
487        assert_eq!(sizes[0].0, 100);
488
489        let mut big = df!("s" => vec![text.as_str(); 3000]).unwrap();
490        assert_eq!(big.first_col_n_chunks(), 1);
491        let sizes = blocks(&mut big);
492        assert_eq!(sizes.len(), 3, "{sizes:?}");
493        assert_eq!(sizes.iter().map(|(rows, _)| rows).sum::<usize>(), 3000);
494        for (_, size) in &sizes[..2] {
495            assert!(
496                (BLOCK_BYTES..BLOCK_BYTES + 1100).contains(size),
497                "{sizes:?}"
498            );
499        }
500    }
501
502    /// A failed write reports the I/O error, not avro-schema's "OutOfSpec".
503    #[test]
504    fn a_write_error_says_what_failed() {
505        struct Full;
506        impl Write for Full {
507            fn write(&mut self, _: &[u8]) -> std::io::Result<usize> {
508                Err(std::io::Error::other("disk full"))
509            }
510            fn flush(&mut self) -> std::io::Result<()> {
511                Ok(())
512            }
513        }
514        let mut df = df!("n" => [1i64]).unwrap();
515        let err = write(&mut df, Full).unwrap_err().to_string();
516        assert!(err.contains("disk full"), "{err}");
517    }
518
519    /// A value that does not fit fails by column name rather than turning null.
520    #[test]
521    fn an_unsigned_value_past_i64_fails_by_name() {
522        let lf = df!("big" => [u64::MAX]).unwrap().lazy();
523        let err = lazy_for_avro(lf)
524            .unwrap()
525            .collect()
526            .unwrap_err()
527            .to_string();
528        assert!(err.contains("'big'"), "{err}");
529    }
530}