1use 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
28pub 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
40fn 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
78pub 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
90fn 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
111const BLOCK_BYTES: usize = 1 << 20;
114
115pub fn write(df: &mut DataFrame, mut writer: impl Write) -> PolarsResult<()> {
119 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 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 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
163fn 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 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
195pub 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 #[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 #[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 #[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 #[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 #[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 #[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 #[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}