1mod records;
167
168use arrow_array::builder::{NullBuilder, PrimitiveBuilder};
169use arrow_array::types::*;
170use arrow_array::*;
171use arrow_cast::parse::{Parser, parse_decimal, string_to_datetime};
172use arrow_schema::*;
173use chrono::{TimeZone, Utc};
174use csv::StringRecord;
175use regex::{Regex, RegexSet};
176use std::fmt::{self, Debug};
177use std::fs::File;
178use std::io::{BufRead, BufReader as StdBufReader, Read};
179use std::sync::{Arc, LazyLock};
180
181use crate::map_csv_error;
182use crate::reader::records::{RecordDecoder, StringRecords};
183use arrow_array::timezone::Tz;
184
185static REGEX_SET: LazyLock<RegexSet> = LazyLock::new(|| {
187 RegexSet::new([
188 r"(?i)^(true)$|^(false)$(?-i)", r"^[+-]?(\d+)$", r"^[+-]?((\d*\.\d+|\d+\.\d*)([eE][-+]?\d+)?|\d+([eE][-+]?\d+))$", r"^\d{4}-\d\d-\d\d$", r"^\d{4}-\d\d-\d\d[T ]\d\d:\d\d:\d\d(?:[^\d\.].*)?$", r"^\d{4}-\d\d-\d\d[T ]\d\d:\d\d:\d\d\.\d{1,3}(?:[^\d].*)?$", r"^\d{4}-\d\d-\d\d[T ]\d\d:\d\d:\d\d\.\d{1,6}(?:[^\d].*)?$", r"^\d{4}-\d\d-\d\d[T ]\d\d:\d\d:\d\d\.\d{1,9}(?:[^\d].*)?$", ])
197 .unwrap()
198});
199
200#[derive(Debug, Clone, Default)]
202struct NullRegex(Option<Regex>);
203
204impl NullRegex {
205 #[inline]
208 fn is_null(&self, s: &str) -> bool {
209 match &self.0 {
210 Some(r) => r.is_match(s),
211 None => s.is_empty(),
212 }
213 }
214}
215
216#[derive(Default, Copy, Clone)]
217struct InferredDataType {
218 packed: u16,
230}
231
232impl InferredDataType {
233 fn get(&self) -> DataType {
235 match self.packed {
236 0 => DataType::Null,
237 1 => DataType::Boolean,
238 2 => DataType::Int64,
239 4 | 6 => DataType::Float64, b if b != 0 && (b & !0b11111000) == 0 => match b.leading_zeros() {
241 8 => DataType::Timestamp(TimeUnit::Nanosecond, None),
243 9 => DataType::Timestamp(TimeUnit::Microsecond, None),
244 10 => DataType::Timestamp(TimeUnit::Millisecond, None),
245 11 => DataType::Timestamp(TimeUnit::Second, None),
246 12 => DataType::Date32,
247 _ => unreachable!(),
248 },
249 _ => DataType::Utf8,
250 }
251 }
252
253 fn update(&mut self, string: &str) {
255 self.packed |= if string.starts_with('"') {
256 1 << 8 } else if let Some(m) = REGEX_SET.matches(string).into_iter().next() {
258 if m == 1 && string.len() >= 19 && string.parse::<i64>().is_err() {
259 1 << 8
261 } else {
262 1 << m
263 }
264 } else if string == "NaN" || string == "nan" || string == "inf" || string == "-inf" {
265 1 << 2 } else {
267 1 << 8 }
269 }
270}
271
272#[derive(Debug, Clone, Default)]
274pub struct Format {
275 header: bool,
276 header_validation: bool,
277 delimiter: Option<u8>,
278 escape: Option<u8>,
279 quote: Option<u8>,
280 terminator: Option<u8>,
281 comment: Option<u8>,
282 null_regex: NullRegex,
283 truncated_rows: bool,
284}
285
286impl Format {
287 pub fn with_header(mut self, has_header: bool) -> Self {
291 self.header = has_header;
292 self
293 }
294
295 pub fn with_header_validation(mut self, validate_header: bool) -> Self {
301 self.header_validation = validate_header;
302 self
303 }
304
305 pub fn with_delimiter(mut self, delimiter: u8) -> Self {
307 self.delimiter = Some(delimiter);
308 self
309 }
310
311 pub fn with_escape(mut self, escape: u8) -> Self {
313 self.escape = Some(escape);
314 self
315 }
316
317 pub fn with_quote(mut self, quote: u8) -> Self {
319 self.quote = Some(quote);
320 self
321 }
322
323 pub fn with_terminator(mut self, terminator: u8) -> Self {
325 self.terminator = Some(terminator);
326 self
327 }
328
329 pub fn with_comment(mut self, comment: u8) -> Self {
333 self.comment = Some(comment);
334 self
335 }
336
337 pub fn with_null_regex(mut self, null_regex: Regex) -> Self {
339 self.null_regex = NullRegex(Some(null_regex));
340 self
341 }
342
343 pub fn with_truncated_rows(mut self, allow: bool) -> Self {
350 self.truncated_rows = allow;
351 self
352 }
353
354 pub fn infer_format<R: Read>(
379 mut self,
380 reader: R,
381 max_records: Option<usize>,
382 ) -> Result<(Self, usize), ArrowError> {
383 let (header, records_read) = self.infer_header(reader, max_records)?;
384 self.header = header;
385 Ok((self, records_read))
386 }
387
388 fn infer_header<R: Read>(
394 &self,
395 reader: R,
396 max_records: Option<usize>,
397 ) -> Result<(bool, usize), ArrowError> {
398 let mut format = self.clone();
399 format.header = false;
400 let mut csv_reader = format.build_reader(reader);
401
402 let mut first_record = StringRecord::new();
403 if !csv_reader
404 .read_record(&mut first_record)
405 .map_err(map_csv_error)?
406 {
407 return Ok((false, 0));
408 }
409
410 let mut first_types = vec![InferredDataType::default(); first_record.len()];
411 for (value, inferred) in first_record.iter().zip(&mut first_types) {
412 if !self.null_regex.is_null(value) {
413 inferred.update(value);
414 }
415 }
416
417 let mut column_types = vec![InferredDataType::default(); first_record.len()];
418 let mut record = StringRecord::new();
419 let mut records_count = 0;
420 let max_records = max_records.unwrap_or(usize::MAX);
421 while records_count < max_records
422 && csv_reader.read_record(&mut record).map_err(map_csv_error)?
423 {
424 records_count += 1;
425 for (value, inferred) in record.iter().zip(&mut column_types) {
426 if !self.null_regex.is_null(value) {
427 inferred.update(value);
428 }
429 }
430 }
431
432 let has_header = first_types
433 .iter()
434 .zip(&column_types)
435 .zip(first_record.iter())
436 .any(|((first, rest), value)| {
437 first.get() == DataType::Utf8
440 && value.parse::<f64>().is_err()
441 && !matches!(rest.get(), DataType::Utf8 | DataType::Null)
442 });
443 Ok((has_header, records_count + 1))
444 }
445
446 pub fn infer_schema<R: Read>(
453 &self,
454 reader: R,
455 max_records: Option<usize>,
456 ) -> Result<(Schema, usize), ArrowError> {
457 let mut csv_reader = self.build_reader(reader);
458
459 let headers: Vec<String> = if self.header {
462 let headers = &csv_reader.headers().map_err(map_csv_error)?.clone();
463 headers.iter().map(|s| s.to_string()).collect()
464 } else {
465 let first_record_count = &csv_reader.headers().map_err(map_csv_error)?.len();
466 (0..*first_record_count)
467 .map(|i| format!("column_{}", i + 1))
468 .collect()
469 };
470
471 let header_length = headers.len();
472 let mut column_types: Vec<InferredDataType> = vec![Default::default(); header_length];
474
475 let mut records_count = 0;
476
477 let mut record = StringRecord::new();
478 let max_records = max_records.unwrap_or(usize::MAX);
479 while records_count < max_records {
480 if !csv_reader.read_record(&mut record).map_err(map_csv_error)? {
481 break;
482 }
483 records_count += 1;
484
485 for (i, column_type) in column_types.iter_mut().enumerate().take(header_length) {
488 if let Some(string) = record.get(i)
489 && !self.null_regex.is_null(string)
490 {
491 column_type.update(string)
492 }
493 }
494 }
495
496 let fields: Fields = column_types
498 .iter()
499 .zip(&headers)
500 .map(|(inferred, field_name)| Field::new(field_name, inferred.get(), true))
501 .collect();
502
503 Ok((Schema::new(fields), records_count))
504 }
505
506 fn build_reader<R: Read>(&self, reader: R) -> csv::Reader<R> {
508 let mut builder = csv::ReaderBuilder::new();
509 builder.has_headers(self.header);
510 builder.flexible(self.truncated_rows);
511
512 if let Some(c) = self.delimiter {
513 builder.delimiter(c);
514 }
515 builder.escape(self.escape);
516 if let Some(c) = self.quote {
517 builder.quote(c);
518 }
519 if let Some(t) = self.terminator {
520 builder.terminator(csv::Terminator::Any(t));
521 }
522 if let Some(comment) = self.comment {
523 builder.comment(Some(comment));
524 }
525 builder.from_reader(reader)
526 }
527
528 fn build_parser(&self) -> csv_core::Reader {
530 let mut builder = csv_core::ReaderBuilder::new();
531 builder.escape(self.escape);
532 builder.comment(self.comment);
533
534 if let Some(c) = self.delimiter {
535 builder.delimiter(c);
536 }
537 if let Some(c) = self.quote {
538 builder.quote(c);
539 }
540 if let Some(t) = self.terminator {
541 builder.terminator(csv_core::Terminator::Any(t));
542 }
543 builder.build()
544 }
545}
546
547pub fn infer_schema_from_files(
554 files: &[String],
555 delimiter: u8,
556 max_read_records: Option<usize>,
557 has_header: bool,
558) -> Result<Schema, ArrowError> {
559 let mut schemas = vec![];
560 let mut records_to_read = max_read_records.unwrap_or(usize::MAX);
561 let format = Format {
562 delimiter: Some(delimiter),
563 header: has_header,
564 ..Default::default()
565 };
566
567 for fname in files {
568 let f = File::open(fname)?;
569 let (schema, records_read) = format.infer_schema(f, Some(records_to_read))?;
570 if records_read == 0 {
571 continue;
572 }
573 schemas.push(schema.clone());
574 records_to_read -= records_read;
575 if records_to_read == 0 {
576 break;
577 }
578 }
579
580 Schema::try_merge(schemas)
581}
582
583type Bounds = Option<(usize, usize)>;
585
586pub type Reader<R> = BufReader<StdBufReader<R>>;
591
592pub struct BufReader<R> {
597 reader: R,
599 decoder: Decoder,
601 schema: SchemaRef,
603}
604
605impl<R> fmt::Debug for BufReader<R>
606where
607 R: BufRead,
608{
609 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
610 f.debug_struct("Reader")
611 .field("decoder", &self.decoder)
612 .finish()
613 }
614}
615
616impl<R> BufReader<R> {
617 pub fn truncated_row_count(&self) -> usize {
652 self.decoder.truncated_row_count()
653 }
654}
655
656impl<R: Read> Reader<R> {
657 pub fn schema(&self) -> SchemaRef {
660 self.schema.clone()
661 }
662}
663
664impl<R: BufRead> BufReader<R> {
665 fn read(&mut self) -> Result<Option<RecordBatch>, ArrowError> {
666 loop {
667 let buf = self.reader.fill_buf()?;
668 let decoded = self.decoder.decode(buf)?;
669 self.reader.consume(decoded);
670 if decoded == 0 || self.decoder.capacity() == 0 {
676 break;
677 }
678 }
679
680 self.decoder.flush()
681 }
682}
683
684impl<R: BufRead> Iterator for BufReader<R> {
685 type Item = Result<RecordBatch, ArrowError>;
686
687 fn next(&mut self) -> Option<Self::Item> {
688 self.read().transpose()
689 }
690}
691
692impl<R: BufRead> RecordBatchReader for BufReader<R> {
693 fn schema(&self) -> SchemaRef {
694 self.schema.clone()
695 }
696}
697
698#[derive(Debug)]
738pub struct Decoder {
739 schema: SchemaRef,
741
742 projection: Option<Vec<usize>>,
744
745 batch_size: usize,
747
748 to_skip: usize,
750
751 header_validation: bool,
753
754 line_number: usize,
756
757 end: usize,
759
760 record_decoder: RecordDecoder,
762
763 null_regex: NullRegex,
765}
766
767impl Decoder {
768 pub fn decode(&mut self, buf: &[u8]) -> Result<usize, ArrowError> {
778 if self.to_skip != 0 {
779 if self.header_validation {
780 let (skipped, bytes) = self.record_decoder.decode(buf, 1)?;
781
782 if skipped == 0 {
783 return Ok(bytes);
784 }
785
786 let rows = self.record_decoder.flush()?;
787 validate_header(&rows, self.schema.fields())?;
788 self.header_validation = false;
789 self.to_skip -= 1;
790 return Ok(bytes);
791 }
792
793 let to_skip = self.to_skip.min(self.batch_size);
795 let (skipped, bytes) = self.record_decoder.decode(buf, to_skip)?;
796 self.to_skip -= skipped;
797 self.record_decoder.clear();
798 return Ok(bytes);
799 }
800
801 let to_read = self.batch_size.min(self.end - self.line_number) - self.record_decoder.len();
802 let (_, bytes) = self.record_decoder.decode(buf, to_read)?;
803 Ok(bytes)
804 }
805
806 pub fn flush(&mut self) -> Result<Option<RecordBatch>, ArrowError> {
813 if self.record_decoder.is_empty() {
814 return Ok(None);
815 }
816
817 let rows = self.record_decoder.flush()?;
818 let batch = parse(
819 &rows,
820 &self.schema,
821 self.projection.as_ref(),
822 self.line_number,
823 &self.null_regex,
824 )?;
825 self.line_number += rows.len();
826 Ok(Some(batch))
827 }
828
829 pub fn capacity(&self) -> usize {
831 self.batch_size - self.record_decoder.len()
832 }
833
834 pub fn truncated_row_count(&self) -> usize {
848 self.record_decoder.truncated_row_count()
849 }
850}
851
852fn validate_header(rows: &StringRecords<'_>, fields: &Fields) -> Result<(), ArrowError> {
853 let header = rows.iter().next().ok_or_else(|| {
854 ArrowError::CsvError("CSV header validation failed: no header row found".to_string())
855 })?;
856
857 for (idx, field) in fields.iter().enumerate() {
858 let actual = header.get(idx);
859 let expected = field.name();
860 if actual != expected {
861 return Err(ArrowError::CsvError(format!(
862 "CSV header does not match schema at column {idx}: expected {expected:?} but found {actual:?}"
863 )));
864 }
865 }
866
867 Ok(())
868}
869
870fn parse(
872 rows: &StringRecords<'_>,
873 schema: &Schema,
874 projection: Option<&Vec<usize>>,
875 line_number: usize,
876 null_regex: &NullRegex,
877) -> Result<RecordBatch, ArrowError> {
878 let fields = schema.fields();
879 let projection: Vec<usize> = match projection {
880 Some(v) => v.clone(),
881 None => fields.iter().enumerate().map(|(i, _)| i).collect(),
882 };
883 let projected_schema = Arc::new(schema.project(&projection)?);
884
885 let arrays: Result<Vec<ArrayRef>, _> = projection
886 .iter()
887 .map(|i| {
888 let i = *i;
889 let field = &fields[i];
890 match field.data_type() {
891 DataType::Boolean => build_boolean_array(line_number, rows, i, null_regex),
892 DataType::Decimal32(precision, scale) => build_decimal_array::<Decimal32Type>(
893 line_number,
894 rows,
895 i,
896 *precision,
897 *scale,
898 null_regex,
899 ),
900 DataType::Decimal64(precision, scale) => build_decimal_array::<Decimal64Type>(
901 line_number,
902 rows,
903 i,
904 *precision,
905 *scale,
906 null_regex,
907 ),
908 DataType::Decimal128(precision, scale) => build_decimal_array::<Decimal128Type>(
909 line_number,
910 rows,
911 i,
912 *precision,
913 *scale,
914 null_regex,
915 ),
916 DataType::Decimal256(precision, scale) => build_decimal_array::<Decimal256Type>(
917 line_number,
918 rows,
919 i,
920 *precision,
921 *scale,
922 null_regex,
923 ),
924 DataType::Int8 => {
925 build_primitive_array::<Int8Type>(line_number, rows, i, null_regex)
926 }
927 DataType::Int16 => {
928 build_primitive_array::<Int16Type>(line_number, rows, i, null_regex)
929 }
930 DataType::Int32 => {
931 build_primitive_array::<Int32Type>(line_number, rows, i, null_regex)
932 }
933 DataType::Int64 => {
934 build_primitive_array::<Int64Type>(line_number, rows, i, null_regex)
935 }
936 DataType::UInt8 => {
937 build_primitive_array::<UInt8Type>(line_number, rows, i, null_regex)
938 }
939 DataType::UInt16 => {
940 build_primitive_array::<UInt16Type>(line_number, rows, i, null_regex)
941 }
942 DataType::UInt32 => {
943 build_primitive_array::<UInt32Type>(line_number, rows, i, null_regex)
944 }
945 DataType::UInt64 => {
946 build_primitive_array::<UInt64Type>(line_number, rows, i, null_regex)
947 }
948 DataType::Float16 => {
949 build_primitive_array::<Float16Type>(line_number, rows, i, null_regex)
950 }
951 DataType::Float32 => {
952 build_primitive_array::<Float32Type>(line_number, rows, i, null_regex)
953 }
954 DataType::Float64 => {
955 build_primitive_array::<Float64Type>(line_number, rows, i, null_regex)
956 }
957 DataType::Date32 => {
958 build_primitive_array::<Date32Type>(line_number, rows, i, null_regex)
959 }
960 DataType::Date64 => {
961 build_primitive_array::<Date64Type>(line_number, rows, i, null_regex)
962 }
963 DataType::Time32(TimeUnit::Second) => {
964 build_primitive_array::<Time32SecondType>(line_number, rows, i, null_regex)
965 }
966 DataType::Time32(TimeUnit::Millisecond) => {
967 build_primitive_array::<Time32MillisecondType>(line_number, rows, i, null_regex)
968 }
969 DataType::Time64(TimeUnit::Microsecond) => {
970 build_primitive_array::<Time64MicrosecondType>(line_number, rows, i, null_regex)
971 }
972 DataType::Time64(TimeUnit::Nanosecond) => {
973 build_primitive_array::<Time64NanosecondType>(line_number, rows, i, null_regex)
974 }
975 DataType::Timestamp(TimeUnit::Second, tz) => {
976 build_timestamp_array::<TimestampSecondType>(
977 line_number,
978 rows,
979 i,
980 tz.as_deref(),
981 null_regex,
982 )
983 }
984 DataType::Timestamp(TimeUnit::Millisecond, tz) => {
985 build_timestamp_array::<TimestampMillisecondType>(
986 line_number,
987 rows,
988 i,
989 tz.as_deref(),
990 null_regex,
991 )
992 }
993 DataType::Timestamp(TimeUnit::Microsecond, tz) => {
994 build_timestamp_array::<TimestampMicrosecondType>(
995 line_number,
996 rows,
997 i,
998 tz.as_deref(),
999 null_regex,
1000 )
1001 }
1002 DataType::Timestamp(TimeUnit::Nanosecond, tz) => {
1003 build_timestamp_array::<TimestampNanosecondType>(
1004 line_number,
1005 rows,
1006 i,
1007 tz.as_deref(),
1008 null_regex,
1009 )
1010 }
1011 DataType::Null => Ok(Arc::new({
1012 let mut builder = NullBuilder::new();
1013 builder.append_nulls(rows.len());
1014 builder.finish()
1015 }) as ArrayRef),
1016 DataType::Utf8 => Ok(Arc::new(
1017 rows.iter()
1018 .map(|row| {
1019 let s = row.get(i);
1020 (!null_regex.is_null(s)).then_some(s)
1021 })
1022 .collect::<StringArray>(),
1023 ) as ArrayRef),
1024 DataType::Utf8View => Ok(Arc::new(
1025 rows.iter()
1026 .map(|row| {
1027 let s = row.get(i);
1028 (!null_regex.is_null(s)).then_some(s)
1029 })
1030 .collect::<StringViewArray>(),
1031 ) as ArrayRef),
1032 DataType::Dictionary(key_type, value_type)
1033 if value_type.as_ref() == &DataType::Utf8 =>
1034 {
1035 match key_type.as_ref() {
1036 DataType::Int8 => Ok(Arc::new(
1037 rows.iter()
1038 .map(|row| {
1039 let s = row.get(i);
1040 (!null_regex.is_null(s)).then_some(s)
1041 })
1042 .collect::<DictionaryArray<Int8Type>>(),
1043 ) as ArrayRef),
1044 DataType::Int16 => Ok(Arc::new(
1045 rows.iter()
1046 .map(|row| {
1047 let s = row.get(i);
1048 (!null_regex.is_null(s)).then_some(s)
1049 })
1050 .collect::<DictionaryArray<Int16Type>>(),
1051 ) as ArrayRef),
1052 DataType::Int32 => Ok(Arc::new(
1053 rows.iter()
1054 .map(|row| {
1055 let s = row.get(i);
1056 (!null_regex.is_null(s)).then_some(s)
1057 })
1058 .collect::<DictionaryArray<Int32Type>>(),
1059 ) as ArrayRef),
1060 DataType::Int64 => Ok(Arc::new(
1061 rows.iter()
1062 .map(|row| {
1063 let s = row.get(i);
1064 (!null_regex.is_null(s)).then_some(s)
1065 })
1066 .collect::<DictionaryArray<Int64Type>>(),
1067 ) as ArrayRef),
1068 DataType::UInt8 => Ok(Arc::new(
1069 rows.iter()
1070 .map(|row| {
1071 let s = row.get(i);
1072 (!null_regex.is_null(s)).then_some(s)
1073 })
1074 .collect::<DictionaryArray<UInt8Type>>(),
1075 ) as ArrayRef),
1076 DataType::UInt16 => Ok(Arc::new(
1077 rows.iter()
1078 .map(|row| {
1079 let s = row.get(i);
1080 (!null_regex.is_null(s)).then_some(s)
1081 })
1082 .collect::<DictionaryArray<UInt16Type>>(),
1083 ) as ArrayRef),
1084 DataType::UInt32 => Ok(Arc::new(
1085 rows.iter()
1086 .map(|row| {
1087 let s = row.get(i);
1088 (!null_regex.is_null(s)).then_some(s)
1089 })
1090 .collect::<DictionaryArray<UInt32Type>>(),
1091 ) as ArrayRef),
1092 DataType::UInt64 => Ok(Arc::new(
1093 rows.iter()
1094 .map(|row| {
1095 let s = row.get(i);
1096 (!null_regex.is_null(s)).then_some(s)
1097 })
1098 .collect::<DictionaryArray<UInt64Type>>(),
1099 ) as ArrayRef),
1100 _ => Err(ArrowError::ParseError(format!(
1101 "Unsupported dictionary key type {key_type}"
1102 ))),
1103 }
1104 }
1105 other => Err(ArrowError::ParseError(format!(
1106 "Unsupported data type {other:?}"
1107 ))),
1108 }
1109 })
1110 .collect();
1111
1112 RecordBatch::try_new_with_options(
1113 projected_schema,
1114 arrays?,
1115 &RecordBatchOptions::new()
1116 .with_match_field_names(true)
1117 .with_row_count(Some(rows.len())),
1118 )
1119}
1120
1121fn parse_bool(string: &str) -> Option<bool> {
1122 if string.eq_ignore_ascii_case("false") {
1123 Some(false)
1124 } else if string.eq_ignore_ascii_case("true") {
1125 Some(true)
1126 } else {
1127 None
1128 }
1129}
1130
1131fn build_decimal_array<T: DecimalType>(
1133 _line_number: usize,
1134 rows: &StringRecords<'_>,
1135 col_idx: usize,
1136 precision: u8,
1137 scale: i8,
1138 null_regex: &NullRegex,
1139) -> Result<ArrayRef, ArrowError> {
1140 let mut decimal_builder = PrimitiveBuilder::<T>::with_capacity(rows.len());
1141 for row in rows.iter() {
1142 let s = row.get(col_idx);
1143 if null_regex.is_null(s) {
1144 decimal_builder.append_null();
1146 } else {
1147 let decimal_value: Result<T::Native, _> = parse_decimal::<T>(s, precision, scale);
1148 match decimal_value {
1149 Ok(v) => {
1150 decimal_builder.append_value(v);
1151 }
1152 Err(e) => {
1153 return Err(e);
1154 }
1155 }
1156 }
1157 }
1158 Ok(Arc::new(
1159 decimal_builder
1160 .finish()
1161 .with_precision_and_scale(precision, scale)?,
1162 ))
1163}
1164
1165fn build_primitive_array<T: ArrowPrimitiveType + Parser>(
1167 line_number: usize,
1168 rows: &StringRecords<'_>,
1169 col_idx: usize,
1170 null_regex: &NullRegex,
1171) -> Result<ArrayRef, ArrowError> {
1172 rows.iter()
1173 .enumerate()
1174 .map(|(row_index, row)| {
1175 let s = row.get(col_idx);
1176 if null_regex.is_null(s) {
1177 return Ok(None);
1178 }
1179
1180 match T::parse(s) {
1181 Some(e) => Ok(Some(e)),
1182 None => Err(ArrowError::ParseError(format!(
1183 "Error while parsing value '{}' as type '{}' for column {} at line {}. Row data: '{}'",
1185 s,
1186 T::DATA_TYPE,
1187 col_idx,
1188 line_number + row_index,
1189 row
1190 ))),
1191 }
1192 })
1193 .collect::<Result<PrimitiveArray<T>, ArrowError>>()
1194 .map(|e| Arc::new(e) as ArrayRef)
1195}
1196
1197fn build_timestamp_array<T: ArrowTimestampType>(
1198 line_number: usize,
1199 rows: &StringRecords<'_>,
1200 col_idx: usize,
1201 timezone: Option<&str>,
1202 null_regex: &NullRegex,
1203) -> Result<ArrayRef, ArrowError> {
1204 Ok(Arc::new(match timezone {
1205 Some(timezone) => {
1206 let tz: Tz = timezone.parse()?;
1207 build_timestamp_array_impl::<T, _>(line_number, rows, col_idx, &tz, null_regex)?
1208 .with_timezone(timezone)
1209 }
1210 None => build_timestamp_array_impl::<T, _>(line_number, rows, col_idx, &Utc, null_regex)?,
1211 }))
1212}
1213
1214fn build_timestamp_array_impl<T: ArrowTimestampType, Tz: TimeZone>(
1215 line_number: usize,
1216 rows: &StringRecords<'_>,
1217 col_idx: usize,
1218 timezone: &Tz,
1219 null_regex: &NullRegex,
1220) -> Result<PrimitiveArray<T>, ArrowError> {
1221 rows.iter()
1222 .enumerate()
1223 .map(|(row_index, row)| {
1224 let s = row.get(col_idx);
1225 if null_regex.is_null(s) {
1226 return Ok(None);
1227 }
1228
1229 let date = string_to_datetime(timezone, s)
1230 .and_then(|date| match T::UNIT {
1231 TimeUnit::Second => Ok(date.timestamp()),
1232 TimeUnit::Millisecond => Ok(date.timestamp_millis()),
1233 TimeUnit::Microsecond => Ok(date.timestamp_micros()),
1234 TimeUnit::Nanosecond => date.timestamp_nanos_opt().ok_or_else(|| {
1235 ArrowError::ParseError(format!(
1236 "{} would overflow 64-bit signed nanoseconds",
1237 date.to_rfc3339(),
1238 ))
1239 }),
1240 })
1241 .map_err(|e| {
1242 ArrowError::ParseError(format!(
1243 "Error parsing column {col_idx} at line {}: {}",
1244 line_number + row_index,
1245 e
1246 ))
1247 })?;
1248 Ok(Some(date))
1249 })
1250 .collect()
1251}
1252
1253fn build_boolean_array(
1255 line_number: usize,
1256 rows: &StringRecords<'_>,
1257 col_idx: usize,
1258 null_regex: &NullRegex,
1259) -> Result<ArrayRef, ArrowError> {
1260 rows.iter()
1261 .enumerate()
1262 .map(|(row_index, row)| {
1263 let s = row.get(col_idx);
1264 if null_regex.is_null(s) {
1265 return Ok(None);
1266 }
1267 let parsed = parse_bool(s);
1268 match parsed {
1269 Some(e) => Ok(Some(e)),
1270 None => Err(ArrowError::ParseError(format!(
1271 "Error while parsing value '{}' as type '{}' for column {} at line {}. Row data: '{}'",
1273 s,
1274 "Boolean",
1275 col_idx,
1276 line_number + row_index,
1277 row
1278 ))),
1279 }
1280 })
1281 .collect::<Result<BooleanArray, _>>()
1282 .map(|e| Arc::new(e) as ArrayRef)
1283}
1284
1285#[derive(Debug)]
1287pub struct ReaderBuilder {
1288 schema: SchemaRef,
1290 format: Format,
1292 batch_size: usize,
1296 bounds: Bounds,
1298 projection: Option<Vec<usize>>,
1300}
1301
1302impl ReaderBuilder {
1303 pub fn new(schema: SchemaRef) -> ReaderBuilder {
1326 Self {
1327 schema,
1328 format: Format::default(),
1329 batch_size: 1024,
1330 bounds: None,
1331 projection: None,
1332 }
1333 }
1334
1335 pub fn with_header(mut self, has_header: bool) -> Self {
1337 self.format.header = has_header;
1338 self
1339 }
1340
1341 pub fn with_header_validation(mut self, validate_header: bool) -> Self {
1345 self.format.header_validation = validate_header;
1346 self
1347 }
1348
1349 pub fn with_format(mut self, format: Format) -> Self {
1351 self.format = format;
1352 self
1353 }
1354
1355 pub fn with_delimiter(mut self, delimiter: u8) -> Self {
1357 self.format.delimiter = Some(delimiter);
1358 self
1359 }
1360
1361 pub fn with_escape(mut self, escape: u8) -> Self {
1363 self.format.escape = Some(escape);
1364 self
1365 }
1366
1367 pub fn with_quote(mut self, quote: u8) -> Self {
1369 self.format.quote = Some(quote);
1370 self
1371 }
1372
1373 pub fn with_terminator(mut self, terminator: u8) -> Self {
1375 self.format.terminator = Some(terminator);
1376 self
1377 }
1378
1379 pub fn with_comment(mut self, comment: u8) -> Self {
1381 self.format.comment = Some(comment);
1382 self
1383 }
1384
1385 pub fn with_null_regex(mut self, null_regex: Regex) -> Self {
1387 self.format.null_regex = NullRegex(Some(null_regex));
1388 self
1389 }
1390
1391 pub fn with_batch_size(mut self, batch_size: usize) -> Self {
1393 self.batch_size = batch_size;
1394 self
1395 }
1396
1397 pub fn with_bounds(mut self, start: usize, end: usize) -> Self {
1400 self.bounds = Some((start, end));
1401 self
1402 }
1403
1404 pub fn with_projection(mut self, projection: Vec<usize>) -> Self {
1406 self.projection = Some(projection);
1407 self
1408 }
1409
1410 pub fn with_truncated_rows(mut self, allow: bool) -> Self {
1417 self.format.truncated_rows = allow;
1418 self
1419 }
1420
1421 pub fn build<R: Read>(self, reader: R) -> Result<Reader<R>, ArrowError> {
1426 self.build_buffered(StdBufReader::new(reader))
1427 }
1428
1429 pub fn build_buffered<R: BufRead>(self, reader: R) -> Result<BufReader<R>, ArrowError> {
1431 let schema = match &self.projection {
1432 Some(projection) => Arc::new(self.schema.project(projection)?),
1433 None => self.schema.clone(),
1434 };
1435
1436 Ok(BufReader {
1437 reader,
1438 decoder: self.build_decoder(),
1439 schema,
1440 })
1441 }
1442
1443 pub fn build_decoder(self) -> Decoder {
1445 let delimiter = self.format.build_parser();
1446 let record_decoder = RecordDecoder::new(
1447 delimiter,
1448 self.schema.fields().len(),
1449 self.format.truncated_rows,
1450 );
1451
1452 let header = self.format.header as usize;
1453
1454 let (start, end) = match self.bounds {
1455 Some((start, end)) => (start + header, end + header),
1456 None => (header, usize::MAX),
1457 };
1458
1459 Decoder {
1460 schema: self.schema,
1461 to_skip: start,
1462 header_validation: self.format.header && self.format.header_validation,
1463 record_decoder,
1464 line_number: start,
1465 end,
1466 projection: self.projection,
1467 batch_size: self.batch_size,
1468 null_regex: self.format.null_regex,
1469 }
1470 }
1471}
1472
1473#[cfg(test)]
1474mod tests {
1475 use super::*;
1476
1477 use std::io::{Cursor, Seek, SeekFrom, Write};
1478 use tempfile::NamedTempFile;
1479
1480 use arrow_array::cast::AsArray;
1481 use arrow_cast::display::array_value_to_string;
1482
1483 #[test]
1484 fn test_infer_schema_leading_plus_numbers() {
1485 for (csv, expected_type) in [
1486 ("+1\n2\n-3\n", DataType::Int64),
1487 ("+1.5\n2.5\n-3.5\n", DataType::Float64),
1488 ("+1e3\n+2.5e-2\n-3E+2\n", DataType::Float64),
1489 ("+9223372036854775807\n0\n", DataType::Int64),
1490 ("+9223372036854775808\n0\n", DataType::Utf8),
1491 ("+-1\n2\n", DataType::Utf8),
1492 ("+\n2\n", DataType::Utf8),
1493 ] {
1494 let (schema, records_read) = Format::default()
1495 .infer_schema(Cursor::new(csv), None)
1496 .unwrap();
1497 assert_eq!(schema.field(0).data_type(), &expected_type, "CSV: {csv:?}");
1498 let reader = ReaderBuilder::new(Arc::new(schema))
1500 .build(Cursor::new(csv))
1501 .unwrap();
1502 let rows: usize = reader.map(|batch| batch.unwrap().num_rows()).sum();
1503 assert_eq!(rows, records_read, "CSV: {csv:?}");
1504 }
1505 }
1506
1507 #[test]
1508 fn test_csv() {
1509 let schema = Arc::new(Schema::new(vec![
1510 Field::new("city", DataType::Utf8, false),
1511 Field::new("lat", DataType::Float64, false),
1512 Field::new("lng", DataType::Float64, false),
1513 ]));
1514
1515 let file = File::open("test/data/uk_cities.csv").unwrap();
1516 let mut csv = ReaderBuilder::new(schema.clone()).build(file).unwrap();
1517 assert_eq!(schema, csv.schema());
1518 let batch = csv.next().unwrap().unwrap();
1519 assert_eq!(37, batch.num_rows());
1520 assert_eq!(3, batch.num_columns());
1521
1522 let lat = batch.column(1).as_primitive::<Float64Type>();
1524 assert_eq!(57.653484, lat.value(0));
1525
1526 let city = batch.column(0).as_string::<i32>();
1528
1529 assert_eq!("Aberdeen, Aberdeen City, UK", city.value(13));
1530 }
1531
1532 #[test]
1533 fn test_csv_schema_metadata() {
1534 let mut metadata = std::collections::HashMap::new();
1535 metadata.insert("foo".to_owned(), "bar".to_owned());
1536 let schema = Arc::new(Schema::new_with_metadata(
1537 vec![
1538 Field::new("city", DataType::Utf8, false),
1539 Field::new("lat", DataType::Float64, false),
1540 Field::new("lng", DataType::Float64, false),
1541 ],
1542 metadata.clone(),
1543 ));
1544
1545 let file = File::open("test/data/uk_cities.csv").unwrap();
1546
1547 let mut csv = ReaderBuilder::new(schema.clone()).build(file).unwrap();
1548 assert_eq!(schema, csv.schema());
1549 let batch = csv.next().unwrap().unwrap();
1550 assert_eq!(37, batch.num_rows());
1551 assert_eq!(3, batch.num_columns());
1552
1553 assert_eq!(batch.schema().metadata(), &metadata);
1554 }
1555
1556 #[test]
1557 fn test_csv_reader_with_decimal() {
1558 let schema = Arc::new(Schema::new(vec![
1559 Field::new("city", DataType::Utf8, false),
1560 Field::new("lat", DataType::Decimal128(38, 6), false),
1561 Field::new("lng", DataType::Decimal256(76, 6), false),
1562 ]));
1563
1564 let file = File::open("test/data/decimal_test.csv").unwrap();
1565
1566 let mut csv = ReaderBuilder::new(schema).build(file).unwrap();
1567 let batch = csv.next().unwrap().unwrap();
1568 let lat = batch
1570 .column(1)
1571 .as_any()
1572 .downcast_ref::<Decimal128Array>()
1573 .unwrap();
1574
1575 assert_eq!("57.653484", lat.value_as_string(0));
1576 assert_eq!("53.002666", lat.value_as_string(1));
1577 assert_eq!("52.412811", lat.value_as_string(2));
1578 assert_eq!("51.481583", lat.value_as_string(3));
1579 assert_eq!("12.123457", lat.value_as_string(4));
1580 assert_eq!("50.760000", lat.value_as_string(5));
1581 assert_eq!("0.123000", lat.value_as_string(6));
1582 assert_eq!("123.000000", lat.value_as_string(7));
1583 assert_eq!("123.000000", lat.value_as_string(8));
1584 assert_eq!("-50.760000", lat.value_as_string(9));
1585
1586 let lng = batch
1587 .column(2)
1588 .as_any()
1589 .downcast_ref::<Decimal256Array>()
1590 .unwrap();
1591
1592 assert_eq!("-3.335724", lng.value_as_string(0));
1593 assert_eq!("-2.179404", lng.value_as_string(1));
1594 assert_eq!("-1.778197", lng.value_as_string(2));
1595 assert_eq!("-3.179090", lng.value_as_string(3));
1596 assert_eq!("-3.179090", lng.value_as_string(4));
1597 assert_eq!("0.290472", lng.value_as_string(5));
1598 assert_eq!("0.290472", lng.value_as_string(6));
1599 assert_eq!("0.290472", lng.value_as_string(7));
1600 assert_eq!("0.290472", lng.value_as_string(8));
1601 assert_eq!("0.290472", lng.value_as_string(9));
1602 }
1603
1604 #[test]
1605 fn test_csv_reader_decimal_parsing() {
1606 let data = " 1.995 ,1.5e2,1234.5,0e0\n-0.005,-1.5E-2,-150,1E+2\n123,+.5,5,-7\n";
1609 let schema = Arc::new(Schema::new(vec![
1610 Field::new("a", DataType::Decimal128(10, 2), false),
1611 Field::new("b", DataType::Decimal64(18, 2), false),
1612 Field::new("c", DataType::Decimal128(10, -2), false),
1613 Field::new("d", DataType::Decimal32(9, 0), false),
1614 ]));
1615 let mut csv = ReaderBuilder::new(schema).build(Cursor::new(data)).unwrap();
1616 let batch = csv.next().unwrap().unwrap();
1617 let column = |i: usize| {
1618 (0..batch.num_rows())
1619 .map(|row| array_value_to_string(batch.column(i), row).unwrap())
1620 .collect::<Vec<_>>()
1621 };
1622 assert_eq!(column(0), ["2.00", "-0.01", "123.00"]);
1623 assert_eq!(column(1), ["150.00", "-0.02", "0.50"]);
1624 assert_eq!(
1625 batch.column(2).as_primitive::<Decimal128Type>().values(),
1626 &[12, -2, 0]
1627 );
1628 assert_eq!(column(3), ["0", "100", "-7"]);
1629
1630 for (data, expected) in [
1632 ("abc\n", "Invalid decimal format: \"abc\""),
1633 ("1.2.3\n", "Invalid decimal format: \"1.2.3\""),
1634 (
1635 "123456789\n",
1636 "\"123456789\" does not fit in Decimal128(5, 2)",
1637 ),
1638 ("1e99999\n", "does not fit in Decimal128(5, 2)"),
1639 (
1640 &format!("{}\n", "1".repeat(300)),
1641 "does not fit in Decimal128(5, 2)",
1642 ),
1643 (
1644 "4825037936439135476.2609835314269495255615E-14\n",
1645 "does not fit in Decimal128(5, 2)",
1646 ),
1647 ] {
1648 let schema = Arc::new(Schema::new(vec![Field::new(
1649 "a",
1650 DataType::Decimal128(5, 2),
1651 false,
1652 )]));
1653 let mut csv = ReaderBuilder::new(schema).build(Cursor::new(data)).unwrap();
1654 let err = csv.next().unwrap().unwrap_err().to_string();
1655 assert!(err.contains(expected), "{data:?}: {err}");
1656 }
1657 }
1658
1659 #[test]
1660 fn test_csv_reader_with_decimal_3264() {
1661 let schema = Arc::new(Schema::new(vec![
1662 Field::new("city", DataType::Utf8, false),
1663 Field::new("lat", DataType::Decimal32(9, 6), false),
1664 Field::new("lng", DataType::Decimal64(16, 6), false),
1665 ]));
1666
1667 let file = File::open("test/data/decimal_test.csv").unwrap();
1668
1669 let mut csv = ReaderBuilder::new(schema).build(file).unwrap();
1670 let batch = csv.next().unwrap().unwrap();
1671 let lat = batch
1673 .column(1)
1674 .as_any()
1675 .downcast_ref::<Decimal32Array>()
1676 .unwrap();
1677
1678 assert_eq!("57.653484", lat.value_as_string(0));
1679 assert_eq!("53.002666", lat.value_as_string(1));
1680 assert_eq!("52.412811", lat.value_as_string(2));
1681 assert_eq!("51.481583", lat.value_as_string(3));
1682 assert_eq!("12.123457", lat.value_as_string(4));
1683 assert_eq!("50.760000", lat.value_as_string(5));
1684 assert_eq!("0.123000", lat.value_as_string(6));
1685 assert_eq!("123.000000", lat.value_as_string(7));
1686 assert_eq!("123.000000", lat.value_as_string(8));
1687 assert_eq!("-50.760000", lat.value_as_string(9));
1688
1689 let lng = batch
1690 .column(2)
1691 .as_any()
1692 .downcast_ref::<Decimal64Array>()
1693 .unwrap();
1694
1695 assert_eq!("-3.335724", lng.value_as_string(0));
1696 assert_eq!("-2.179404", lng.value_as_string(1));
1697 assert_eq!("-1.778197", lng.value_as_string(2));
1698 assert_eq!("-3.179090", lng.value_as_string(3));
1699 assert_eq!("-3.179090", lng.value_as_string(4));
1700 assert_eq!("0.290472", lng.value_as_string(5));
1701 assert_eq!("0.290472", lng.value_as_string(6));
1702 assert_eq!("0.290472", lng.value_as_string(7));
1703 assert_eq!("0.290472", lng.value_as_string(8));
1704 assert_eq!("0.290472", lng.value_as_string(9));
1705 }
1706
1707 #[test]
1708 fn test_csv_from_buf_reader() {
1709 let schema = Schema::new(vec![
1710 Field::new("city", DataType::Utf8, false),
1711 Field::new("lat", DataType::Float64, false),
1712 Field::new("lng", DataType::Float64, false),
1713 ]);
1714
1715 let file_with_headers = File::open("test/data/uk_cities_with_headers.csv").unwrap();
1716 let file_without_headers = File::open("test/data/uk_cities.csv").unwrap();
1717 let both_files = file_with_headers
1718 .chain(Cursor::new("\n".to_string()))
1719 .chain(file_without_headers);
1720 let mut csv = ReaderBuilder::new(Arc::new(schema))
1721 .with_header(true)
1722 .build(both_files)
1723 .unwrap();
1724 let batch = csv.next().unwrap().unwrap();
1725 assert_eq!(74, batch.num_rows());
1726 assert_eq!(3, batch.num_columns());
1727 }
1728
1729 #[test]
1730 fn test_infer_format_with_typed_columns() {
1731 let csv = "name,count,active\nalice,1,true\nbob,2,false\n";
1732
1733 let (format, format_records_read) = Format::default()
1734 .infer_format(Cursor::new(csv), None)
1735 .unwrap();
1736 let (schema, records_read) = format.infer_schema(Cursor::new(csv), None).unwrap();
1737
1738 assert_eq!(schema.field(0).name(), "name");
1739 assert_eq!(schema.field(1).name(), "count");
1740 assert_eq!(schema.field(2).name(), "active");
1741 assert_eq!(format_records_read, 3);
1742 assert_eq!(records_read, 2);
1743 }
1744
1745 #[test]
1746 fn test_infer_format_without_header() {
1747 let csv = "1,true\n2,false\n";
1748
1749 let (format, format_records_read) = Format::default()
1750 .infer_format(Cursor::new(csv), None)
1751 .unwrap();
1752 let (schema, records_read) = format.infer_schema(Cursor::new(csv), None).unwrap();
1753
1754 assert_eq!(schema.field(0).name(), "column_1");
1755 assert_eq!(schema.field(1).name(), "column_2");
1756 assert_eq!(format_records_read, 2);
1757 assert_eq!(records_read, 2);
1758 }
1759
1760 #[test]
1761 fn test_infer_format_returns_no_header_when_ambiguous() {
1762 for csv in ["name,count\n", "alice,london\nbob,paris\n"] {
1763 let (format, _) = Format::default()
1764 .infer_format(Cursor::new(csv), None)
1765 .unwrap();
1766 let (schema, _) = format.infer_schema(Cursor::new(csv), None).unwrap();
1767 assert_eq!(schema.field(0).name(), "column_1", "CSV: {csv:?}");
1768 }
1769
1770 let (format, format_records_read) = Format::default()
1771 .infer_format(Cursor::new(""), None)
1772 .unwrap();
1773 let (schema, records_read) = format.infer_schema(Cursor::new(""), None).unwrap();
1774 assert!(schema.fields().is_empty());
1775 assert_eq!(format_records_read, 0);
1776 assert_eq!(records_read, 0);
1777 }
1778
1779 #[test]
1780 fn test_infer_format_honors_format_options() {
1781 let csv = "name;count\nalice;1\nbob;2\n";
1782 let (format, _) = Format::default()
1783 .with_delimiter(b';')
1784 .infer_format(Cursor::new(csv), None)
1785 .unwrap();
1786 let (schema, records_read) = format.infer_schema(Cursor::new(csv), None).unwrap();
1787
1788 assert_eq!(schema.field(0).name(), "name");
1789 assert_eq!(schema.field(1).name(), "count");
1790 assert_eq!(records_read, 2);
1791 }
1792
1793 #[test]
1794 fn test_infer_format_respects_max_records() {
1795 let csv = "name,count\nalice,1\nbob,unknown\n";
1796 let infer = |max_records| {
1797 let (format, records_read) = Format::default()
1798 .infer_format(Cursor::new(csv), max_records)
1799 .unwrap();
1800 let (schema, _) = format.infer_schema(Cursor::new(csv), None).unwrap();
1801 (schema.field(0).name().clone(), records_read)
1802 };
1803
1804 assert_eq!(infer(Some(1)), ("name".to_string(), 2));
1805 assert_eq!(infer(None), ("column_1".to_string(), 3));
1806 assert_eq!(infer(Some(0)), ("column_1".to_string(), 1));
1807 }
1808
1809 #[test]
1810 fn test_infer_format_numeric_text_is_not_header() {
1811 for csv in [
1812 "+1\n2\n3\n",
1813 "+1.5\n2.5\n3.5\n",
1814 "+1e3\n2e3\n3e3\n",
1815 "9223372036854775808\n2\n3\n",
1816 ] {
1817 let (format, records_read) = Format::default()
1818 .infer_format(Cursor::new(csv), None)
1819 .unwrap();
1820 let (schema, schema_records_read) =
1821 format.infer_schema(Cursor::new(csv), None).unwrap();
1822 let (ordinary_schema, _) = Format::default()
1823 .infer_schema(Cursor::new(csv), None)
1824 .unwrap();
1825
1826 assert_eq!(schema.field(0).name(), "column_1", "CSV: {csv:?}");
1827 assert_eq!(schema, ordinary_schema, "CSV: {csv:?}");
1828 assert_eq!(records_read, 3);
1829 assert_eq!(schema_records_read, 3);
1830 }
1831 }
1832
1833 #[test]
1834 #[cfg_attr(miri, ignore)] fn test_csv_with_schema_inference() {
1836 let mut file = File::open("test/data/uk_cities_with_headers.csv").unwrap();
1837
1838 let (schema, _) = Format::default()
1839 .with_header(true)
1840 .infer_schema(&mut file, None)
1841 .unwrap();
1842
1843 file.rewind().unwrap();
1844 let builder = ReaderBuilder::new(Arc::new(schema)).with_header(true);
1845
1846 let mut csv = builder.build(file).unwrap();
1847 let expected_schema = Schema::new(vec![
1848 Field::new("city", DataType::Utf8, true),
1849 Field::new("lat", DataType::Float64, true),
1850 Field::new("lng", DataType::Float64, true),
1851 ]);
1852 assert_eq!(Arc::new(expected_schema), csv.schema());
1853 let batch = csv.next().unwrap().unwrap();
1854 assert_eq!(37, batch.num_rows());
1855 assert_eq!(3, batch.num_columns());
1856
1857 let lat = batch
1859 .column(1)
1860 .as_any()
1861 .downcast_ref::<Float64Array>()
1862 .unwrap();
1863 assert_eq!(57.653484, lat.value(0));
1864
1865 let city = batch
1867 .column(0)
1868 .as_any()
1869 .downcast_ref::<StringArray>()
1870 .unwrap();
1871
1872 assert_eq!("Aberdeen, Aberdeen City, UK", city.value(13));
1873 }
1874
1875 #[test]
1876 #[cfg_attr(miri, ignore)] fn test_csv_with_schema_inference_no_headers() {
1878 let mut file = File::open("test/data/uk_cities.csv").unwrap();
1879
1880 let (schema, _) = Format::default().infer_schema(&mut file, None).unwrap();
1881 file.rewind().unwrap();
1882
1883 let mut csv = ReaderBuilder::new(Arc::new(schema)).build(file).unwrap();
1884
1885 let schema = csv.schema();
1887 assert_eq!("column_1", schema.field(0).name());
1888 assert_eq!("column_2", schema.field(1).name());
1889 assert_eq!("column_3", schema.field(2).name());
1890 let batch = csv.next().unwrap().unwrap();
1891 let batch_schema = batch.schema();
1892
1893 assert_eq!(schema, batch_schema);
1894 assert_eq!(37, batch.num_rows());
1895 assert_eq!(3, batch.num_columns());
1896
1897 let lat = batch
1899 .column(1)
1900 .as_any()
1901 .downcast_ref::<Float64Array>()
1902 .unwrap();
1903 assert_eq!(57.653484, lat.value(0));
1904
1905 let city = batch
1907 .column(0)
1908 .as_any()
1909 .downcast_ref::<StringArray>()
1910 .unwrap();
1911
1912 assert_eq!("Aberdeen, Aberdeen City, UK", city.value(13));
1913 }
1914
1915 #[test]
1916 #[cfg_attr(miri, ignore)] fn test_csv_builder_with_bounds() {
1918 let mut file = File::open("test/data/uk_cities.csv").unwrap();
1919
1920 let (schema, _) = Format::default().infer_schema(&mut file, None).unwrap();
1922 file.rewind().unwrap();
1923 let mut csv = ReaderBuilder::new(Arc::new(schema))
1924 .with_bounds(0, 2)
1925 .build(file)
1926 .unwrap();
1927 let batch = csv.next().unwrap().unwrap();
1928
1929 let city = batch
1931 .column(0)
1932 .as_any()
1933 .downcast_ref::<StringArray>()
1934 .unwrap();
1935
1936 assert_eq!("Elgin, Scotland, the UK", city.value(0));
1938
1939 let result = std::panic::catch_unwind(|| city.value(13));
1942 assert!(result.is_err());
1943 }
1944
1945 #[test]
1946 fn test_csv_with_projection() {
1947 let schema = Arc::new(Schema::new(vec![
1948 Field::new("city", DataType::Utf8, false),
1949 Field::new("lat", DataType::Float64, false),
1950 Field::new("lng", DataType::Float64, false),
1951 ]));
1952
1953 let file = File::open("test/data/uk_cities.csv").unwrap();
1954
1955 let mut csv = ReaderBuilder::new(schema)
1956 .with_projection(vec![0, 1])
1957 .build(file)
1958 .unwrap();
1959
1960 let projected_schema = Arc::new(Schema::new(vec![
1961 Field::new("city", DataType::Utf8, false),
1962 Field::new("lat", DataType::Float64, false),
1963 ]));
1964 assert_eq!(projected_schema, csv.schema());
1965 let batch = csv.next().unwrap().unwrap();
1966 assert_eq!(projected_schema, batch.schema());
1967 assert_eq!(37, batch.num_rows());
1968 assert_eq!(2, batch.num_columns());
1969 }
1970
1971 #[test]
1972 fn test_csv_record_batch_reader_schema() {
1973 let schema = Arc::new(Schema::new(vec![
1974 Field::new("a", DataType::Int32, false),
1975 Field::new("b", DataType::Int32, false),
1976 ]));
1977
1978 let cases = [
1979 None,
1980 Some(vec![]),
1981 Some(vec![1]),
1982 Some(vec![1, 0]),
1983 Some(vec![1, 1]),
1984 ];
1985 for projection in cases {
1986 let builder = ReaderBuilder::new(schema.clone());
1987 let builder = match projection {
1988 Some(projection) => builder.with_projection(projection),
1989 None => builder,
1990 };
1991 let mut reader = builder.build(Cursor::new(b"1,2\n")).unwrap();
1992
1993 let reader_schema = RecordBatchReader::schema(&reader);
1994 let batch = reader.next().unwrap().unwrap();
1995
1996 assert_eq!(reader_schema, batch.schema());
1997 }
1998 }
1999
2000 #[test]
2001 fn test_csv_reader_rejects_invalid_projection() {
2002 let schema = Arc::new(Schema::new(vec![
2003 Field::new("a", DataType::Int32, false),
2004 Field::new("b", DataType::Int32, false),
2005 ]));
2006
2007 let result = ReaderBuilder::new(schema)
2008 .with_projection(vec![2])
2009 .build(Cursor::new(b"1,2\n"));
2010
2011 assert!(matches!(
2012 result,
2013 Err(ArrowError::SchemaError(message))
2014 if message == "project index 2 out of bounds, max field 2"
2015 ));
2016 }
2017
2018 #[test]
2019 fn test_csv_decoder_rejects_invalid_projection() {
2020 let schema = Arc::new(Schema::new(vec![
2021 Field::new("a", DataType::Int32, false),
2022 Field::new("b", DataType::Int32, false),
2023 ]));
2024 let mut decoder = ReaderBuilder::new(schema)
2025 .with_projection(vec![2])
2026 .build_decoder();
2027
2028 decoder.decode(b"1,2\n").unwrap();
2029 let result = decoder.flush();
2030
2031 assert!(matches!(
2032 result,
2033 Err(ArrowError::SchemaError(message))
2034 if message == "project index 2 out of bounds, max field 2"
2035 ));
2036 }
2037
2038 #[test]
2039 fn test_csv_with_dictionary() {
2040 let schema = Arc::new(Schema::new(vec![
2041 Field::new_dictionary("city", DataType::Int32, DataType::Utf8, false),
2042 Field::new("lat", DataType::Float64, false),
2043 Field::new("lng", DataType::Float64, false),
2044 ]));
2045
2046 let file = File::open("test/data/uk_cities.csv").unwrap();
2047
2048 let mut csv = ReaderBuilder::new(schema)
2049 .with_projection(vec![0, 1])
2050 .build(file)
2051 .unwrap();
2052
2053 let projected_schema = Arc::new(Schema::new(vec![
2054 Field::new_dictionary("city", DataType::Int32, DataType::Utf8, false),
2055 Field::new("lat", DataType::Float64, false),
2056 ]));
2057 assert_eq!(projected_schema, csv.schema());
2058 let batch = csv.next().unwrap().unwrap();
2059 assert_eq!(projected_schema, batch.schema());
2060 assert_eq!(37, batch.num_rows());
2061 assert_eq!(2, batch.num_columns());
2062
2063 let strings = arrow_cast::cast(batch.column(0), &DataType::Utf8).unwrap();
2064 let strings = strings.as_string::<i32>();
2065
2066 assert_eq!(strings.value(0), "Elgin, Scotland, the UK");
2067 assert_eq!(strings.value(4), "Eastbourne, East Sussex, UK");
2068 assert_eq!(strings.value(29), "Uckfield, East Sussex, UK");
2069 }
2070
2071 #[test]
2072 fn test_csv_with_nullable_dictionary() {
2073 let offset_type = vec![
2074 DataType::Int8,
2075 DataType::Int16,
2076 DataType::Int32,
2077 DataType::Int64,
2078 DataType::UInt8,
2079 DataType::UInt16,
2080 DataType::UInt32,
2081 DataType::UInt64,
2082 ];
2083 for data_type in offset_type {
2084 let file = File::open("test/data/dictionary_nullable_test.csv").unwrap();
2085 let dictionary_type =
2086 DataType::Dictionary(Box::new(data_type), Box::new(DataType::Utf8));
2087 let schema = Arc::new(Schema::new(vec![
2088 Field::new("id", DataType::Utf8, false),
2089 Field::new("name", dictionary_type.clone(), true),
2090 ]));
2091
2092 let mut csv = ReaderBuilder::new(schema)
2093 .build(file.try_clone().unwrap())
2094 .unwrap();
2095
2096 let batch = csv.next().unwrap().unwrap();
2097 assert_eq!(3, batch.num_rows());
2098 assert_eq!(2, batch.num_columns());
2099
2100 let names = arrow_cast::cast(batch.column(1), &dictionary_type).unwrap();
2101 assert!(!names.is_null(2));
2102 assert!(names.is_null(1));
2103 }
2104 }
2105 #[test]
2106 fn test_nulls() {
2107 let schema = Arc::new(Schema::new(vec![
2108 Field::new("c_int", DataType::UInt64, false),
2109 Field::new("c_float", DataType::Float32, true),
2110 Field::new("c_string", DataType::Utf8, true),
2111 Field::new("c_bool", DataType::Boolean, false),
2112 ]));
2113
2114 let file = File::open("test/data/null_test.csv").unwrap();
2115
2116 let mut csv = ReaderBuilder::new(schema)
2117 .with_header(true)
2118 .build(file)
2119 .unwrap();
2120
2121 let batch = csv.next().unwrap().unwrap();
2122
2123 assert!(!batch.column(1).is_null(0));
2124 assert!(!batch.column(1).is_null(1));
2125 assert!(batch.column(1).is_null(2));
2126 assert!(!batch.column(1).is_null(3));
2127 assert!(!batch.column(1).is_null(4));
2128 }
2129
2130 #[test]
2131 fn test_init_nulls() {
2132 let schema = Arc::new(Schema::new(vec![
2133 Field::new("c_int", DataType::UInt64, true),
2134 Field::new("c_float", DataType::Float32, true),
2135 Field::new("c_string", DataType::Utf8, true),
2136 Field::new("c_bool", DataType::Boolean, true),
2137 Field::new("c_null", DataType::Null, true),
2138 ]));
2139 let file = File::open("test/data/init_null_test.csv").unwrap();
2140
2141 let mut csv = ReaderBuilder::new(schema)
2142 .with_header(true)
2143 .build(file)
2144 .unwrap();
2145
2146 let batch = csv.next().unwrap().unwrap();
2147
2148 assert!(batch.column(1).is_null(0));
2149 assert!(!batch.column(1).is_null(1));
2150 assert!(batch.column(1).is_null(2));
2151 assert!(!batch.column(1).is_null(3));
2152 assert!(!batch.column(1).is_null(4));
2153 }
2154
2155 #[test]
2156 #[cfg_attr(miri, ignore)] fn test_init_nulls_with_inference() {
2158 let format = Format::default().with_header(true).with_delimiter(b',');
2159
2160 let mut file = File::open("test/data/init_null_test.csv").unwrap();
2161 let (schema, _) = format.infer_schema(&mut file, None).unwrap();
2162 file.rewind().unwrap();
2163
2164 let expected_schema = Schema::new(vec![
2165 Field::new("c_int", DataType::Int64, true),
2166 Field::new("c_float", DataType::Float64, true),
2167 Field::new("c_string", DataType::Utf8, true),
2168 Field::new("c_bool", DataType::Boolean, true),
2169 Field::new("c_null", DataType::Null, true),
2170 ]);
2171 assert_eq!(schema, expected_schema);
2172
2173 let mut csv = ReaderBuilder::new(Arc::new(schema))
2174 .with_format(format)
2175 .build(file)
2176 .unwrap();
2177
2178 let batch = csv.next().unwrap().unwrap();
2179
2180 assert!(batch.column(1).is_null(0));
2181 assert!(!batch.column(1).is_null(1));
2182 assert!(batch.column(1).is_null(2));
2183 assert!(!batch.column(1).is_null(3));
2184 assert!(!batch.column(1).is_null(4));
2185 }
2186
2187 #[test]
2188 fn test_custom_nulls() {
2189 let schema = Arc::new(Schema::new(vec![
2190 Field::new("c_int", DataType::UInt64, true),
2191 Field::new("c_float", DataType::Float32, true),
2192 Field::new("c_string", DataType::Utf8, true),
2193 Field::new("c_bool", DataType::Boolean, true),
2194 ]));
2195
2196 let file = File::open("test/data/custom_null_test.csv").unwrap();
2197
2198 let null_regex = Regex::new("^nil$").unwrap();
2199
2200 let mut csv = ReaderBuilder::new(schema)
2201 .with_header(true)
2202 .with_null_regex(null_regex)
2203 .build(file)
2204 .unwrap();
2205
2206 let batch = csv.next().unwrap().unwrap();
2207
2208 assert!(batch.column(0).is_null(1));
2210 assert!(batch.column(1).is_null(2));
2211 assert!(batch.column(3).is_null(4));
2212 assert!(batch.column(2).is_null(3));
2213 assert!(!batch.column(2).is_null(4));
2214 }
2215
2216 #[test]
2217 #[cfg_attr(miri, ignore)] fn test_nulls_with_inference() {
2219 let mut file = File::open("test/data/various_types.csv").unwrap();
2220 let format = Format::default().with_header(true).with_delimiter(b'|');
2221
2222 let (schema, _) = format.infer_schema(&mut file, None).unwrap();
2223 file.rewind().unwrap();
2224
2225 let builder = ReaderBuilder::new(Arc::new(schema))
2226 .with_format(format)
2227 .with_batch_size(512)
2228 .with_projection(vec![0, 1, 2, 3, 4, 5]);
2229
2230 let mut csv = builder.build(file).unwrap();
2231 let batch = csv.next().unwrap().unwrap();
2232
2233 assert_eq!(10, batch.num_rows());
2234 assert_eq!(6, batch.num_columns());
2235
2236 let schema = batch.schema();
2237
2238 assert_eq!(&DataType::Int64, schema.field(0).data_type());
2239 assert_eq!(&DataType::Float64, schema.field(1).data_type());
2240 assert_eq!(&DataType::Float64, schema.field(2).data_type());
2241 assert_eq!(&DataType::Boolean, schema.field(3).data_type());
2242 assert_eq!(&DataType::Date32, schema.field(4).data_type());
2243 assert_eq!(
2244 &DataType::Timestamp(TimeUnit::Second, None),
2245 schema.field(5).data_type()
2246 );
2247
2248 let names: Vec<&str> = schema.fields().iter().map(|x| x.name().as_str()).collect();
2249 assert_eq!(
2250 names,
2251 vec![
2252 "c_int",
2253 "c_float",
2254 "c_string",
2255 "c_bool",
2256 "c_date",
2257 "c_datetime"
2258 ]
2259 );
2260
2261 assert!(schema.field(0).is_nullable());
2262 assert!(schema.field(1).is_nullable());
2263 assert!(schema.field(2).is_nullable());
2264 assert!(schema.field(3).is_nullable());
2265 assert!(schema.field(4).is_nullable());
2266 assert!(schema.field(5).is_nullable());
2267
2268 assert!(!batch.column(1).is_null(0));
2269 assert!(!batch.column(1).is_null(1));
2270 assert!(batch.column(1).is_null(2));
2271 assert!(!batch.column(1).is_null(3));
2272 assert!(!batch.column(1).is_null(4));
2273 }
2274
2275 #[test]
2276 #[cfg_attr(miri, ignore)] fn test_custom_nulls_with_inference() {
2278 let mut file = File::open("test/data/custom_null_test.csv").unwrap();
2279
2280 let null_regex = Regex::new("^nil$").unwrap();
2281
2282 let format = Format::default()
2283 .with_header(true)
2284 .with_null_regex(null_regex);
2285
2286 let (schema, _) = format.infer_schema(&mut file, None).unwrap();
2287 file.rewind().unwrap();
2288
2289 let expected_schema = Schema::new(vec![
2290 Field::new("c_int", DataType::Int64, true),
2291 Field::new("c_float", DataType::Float64, true),
2292 Field::new("c_string", DataType::Utf8, true),
2293 Field::new("c_bool", DataType::Boolean, true),
2294 ]);
2295
2296 assert_eq!(schema, expected_schema);
2297
2298 let builder = ReaderBuilder::new(Arc::new(schema))
2299 .with_format(format)
2300 .with_batch_size(512)
2301 .with_projection(vec![0, 1, 2, 3]);
2302
2303 let mut csv = builder.build(file).unwrap();
2304 let batch = csv.next().unwrap().unwrap();
2305
2306 assert_eq!(5, batch.num_rows());
2307 assert_eq!(4, batch.num_columns());
2308
2309 assert_eq!(batch.schema().as_ref(), &expected_schema);
2310 }
2311
2312 #[test]
2313 #[cfg_attr(miri, ignore)] fn test_scientific_notation_with_inference() {
2315 let mut file = File::open("test/data/scientific_notation_test.csv").unwrap();
2316 let format = Format::default().with_header(false).with_delimiter(b',');
2317
2318 let (schema, _) = format.infer_schema(&mut file, None).unwrap();
2319 file.rewind().unwrap();
2320
2321 let builder = ReaderBuilder::new(Arc::new(schema))
2322 .with_format(format)
2323 .with_batch_size(512)
2324 .with_projection(vec![0, 1]);
2325
2326 let mut csv = builder.build(file).unwrap();
2327 let batch = csv.next().unwrap().unwrap();
2328
2329 let schema = batch.schema();
2330
2331 assert_eq!(&DataType::Float64, schema.field(0).data_type());
2332 }
2333
2334 fn invalid_csv_helper(file_name: &str) -> String {
2335 let file = File::open(file_name).unwrap();
2336 let schema = Schema::new(vec![
2337 Field::new("c_int", DataType::UInt64, false),
2338 Field::new("c_float", DataType::Float32, false),
2339 Field::new("c_string", DataType::Utf8, false),
2340 Field::new("c_bool", DataType::Boolean, false),
2341 ]);
2342
2343 let builder = ReaderBuilder::new(Arc::new(schema))
2344 .with_header(true)
2345 .with_delimiter(b'|')
2346 .with_batch_size(512)
2347 .with_projection(vec![0, 1, 2, 3]);
2348
2349 let mut csv = builder.build(file).unwrap();
2350
2351 csv.next().unwrap().unwrap_err().to_string()
2352 }
2353
2354 #[test]
2355 fn test_parse_invalid_csv_float() {
2356 let file_name = "test/data/various_invalid_types/invalid_float.csv";
2357
2358 let error = invalid_csv_helper(file_name);
2359 assert_eq!(
2360 "Parser error: Error while parsing value '4.x4' as type 'Float32' for column 1 at line 4. Row data: '[4,4.x4,,false]'",
2361 error
2362 );
2363 }
2364
2365 #[test]
2366 fn test_parse_invalid_csv_int() {
2367 let file_name = "test/data/various_invalid_types/invalid_int.csv";
2368
2369 let error = invalid_csv_helper(file_name);
2370 assert_eq!(
2371 "Parser error: Error while parsing value '2.3' as type 'UInt64' for column 0 at line 2. Row data: '[2.3,2.2,2.22,false]'",
2372 error
2373 );
2374 }
2375
2376 #[test]
2377 fn test_parse_invalid_csv_bool() {
2378 let file_name = "test/data/various_invalid_types/invalid_bool.csv";
2379
2380 let error = invalid_csv_helper(file_name);
2381 assert_eq!(
2382 "Parser error: Error while parsing value 'none' as type 'Boolean' for column 3 at line 2. Row data: '[2,2.2,2.22,none]'",
2383 error
2384 );
2385 }
2386
2387 fn infer_field_schema(string: &str) -> DataType {
2389 let mut v = InferredDataType::default();
2390 v.update(string);
2391 v.get()
2392 }
2393
2394 #[test]
2395 #[cfg_attr(miri, ignore)] fn test_infer_field_schema() {
2397 assert_eq!(infer_field_schema("A"), DataType::Utf8);
2398 assert_eq!(infer_field_schema("\"123\""), DataType::Utf8);
2399 assert_eq!(infer_field_schema("10"), DataType::Int64);
2400 assert_eq!(infer_field_schema("10.2"), DataType::Float64);
2401 assert_eq!(infer_field_schema(".2"), DataType::Float64);
2402 assert_eq!(infer_field_schema("2."), DataType::Float64);
2403 assert_eq!(infer_field_schema("NaN"), DataType::Float64);
2404 assert_eq!(infer_field_schema("nan"), DataType::Float64);
2405 assert_eq!(infer_field_schema("inf"), DataType::Float64);
2406 assert_eq!(infer_field_schema("-inf"), DataType::Float64);
2407 assert_eq!(infer_field_schema("true"), DataType::Boolean);
2408 assert_eq!(infer_field_schema("trUe"), DataType::Boolean);
2409 assert_eq!(infer_field_schema("false"), DataType::Boolean);
2410 assert_eq!(infer_field_schema("2020-11-08"), DataType::Date32);
2411 assert_eq!(
2412 infer_field_schema("2020-11-08T14:20:01"),
2413 DataType::Timestamp(TimeUnit::Second, None)
2414 );
2415 assert_eq!(
2416 infer_field_schema("2020-11-08 14:20:01"),
2417 DataType::Timestamp(TimeUnit::Second, None)
2418 );
2419 assert_eq!(
2420 infer_field_schema("2020-11-08 14:20:01"),
2421 DataType::Timestamp(TimeUnit::Second, None)
2422 );
2423 assert_eq!(infer_field_schema("-5.13"), DataType::Float64);
2424 assert_eq!(infer_field_schema("0.1300"), DataType::Float64);
2425 assert_eq!(
2426 infer_field_schema("2021-12-19 13:12:30.921"),
2427 DataType::Timestamp(TimeUnit::Millisecond, None)
2428 );
2429 assert_eq!(
2430 infer_field_schema("2021-12-19T13:12:30.123456789"),
2431 DataType::Timestamp(TimeUnit::Nanosecond, None)
2432 );
2433 assert_eq!(infer_field_schema("–9223372036854775809"), DataType::Utf8);
2434 assert_eq!(infer_field_schema("9223372036854775808"), DataType::Utf8);
2435 }
2436
2437 #[test]
2438 fn parse_date32() {
2439 assert_eq!(Date32Type::parse("1970-01-01").unwrap(), 0);
2440 assert_eq!(Date32Type::parse("2020-03-15").unwrap(), 18336);
2441 assert_eq!(Date32Type::parse("1945-05-08").unwrap(), -9004);
2442 }
2443
2444 #[test]
2445 fn parse_time() {
2446 assert_eq!(
2447 Time64NanosecondType::parse("12:10:01.123456789 AM"),
2448 Some(601_123_456_789)
2449 );
2450 assert_eq!(
2451 Time64MicrosecondType::parse("12:10:01.123456 am"),
2452 Some(601_123_456)
2453 );
2454 assert_eq!(
2455 Time32MillisecondType::parse("2:10:01.12 PM"),
2456 Some(51_001_120)
2457 );
2458 assert_eq!(Time32SecondType::parse("2:10:01 pm"), Some(51_001));
2459 }
2460
2461 #[test]
2462 fn parse_date64() {
2463 assert_eq!(Date64Type::parse("1970-01-01T00:00:00").unwrap(), 0);
2464 assert_eq!(
2465 Date64Type::parse("2018-11-13T17:11:10").unwrap(),
2466 1542129070000
2467 );
2468 assert_eq!(
2469 Date64Type::parse("2018-11-13T17:11:10.011").unwrap(),
2470 1542129070011
2471 );
2472 assert_eq!(
2473 Date64Type::parse("1900-02-28T12:34:56").unwrap(),
2474 -2203932304000
2475 );
2476 assert_eq!(
2477 Date64Type::parse_formatted("1900-02-28 12:34:56", "%Y-%m-%d %H:%M:%S").unwrap(),
2478 -2203932304000
2479 );
2480 assert_eq!(
2481 Date64Type::parse_formatted("1900-02-28 12:34:56+0030", "%Y-%m-%d %H:%M:%S%z").unwrap(),
2482 -2203932304000 - (30 * 60 * 1000)
2483 );
2484 }
2485
2486 fn test_parse_timestamp_impl<T: ArrowTimestampType>(
2487 timezone: Option<Arc<str>>,
2488 expected: &[i64],
2489 ) {
2490 let csv = [
2491 "1970-01-01T00:00:00",
2492 "1970-01-01T00:00:00Z",
2493 "1970-01-01T00:00:00+02:00",
2494 ]
2495 .join("\n");
2496 let schema = Arc::new(Schema::new(vec![Field::new(
2497 "field",
2498 DataType::Timestamp(T::UNIT, timezone.clone()),
2499 true,
2500 )]));
2501
2502 let mut decoder = ReaderBuilder::new(schema).build_decoder();
2503
2504 let decoded = decoder.decode(csv.as_bytes()).unwrap();
2505 assert_eq!(decoded, csv.len());
2506 decoder.decode(&[]).unwrap();
2507
2508 let batch = decoder.flush().unwrap().unwrap();
2509 assert_eq!(batch.num_columns(), 1);
2510 assert_eq!(batch.num_rows(), 3);
2511 let col = batch.column(0).as_primitive::<T>();
2512 assert_eq!(col.values(), expected);
2513 assert_eq!(col.data_type(), &DataType::Timestamp(T::UNIT, timezone));
2514 }
2515
2516 #[test]
2517 fn test_parse_timestamp() {
2518 test_parse_timestamp_impl::<TimestampNanosecondType>(None, &[0, 0, -7_200_000_000_000]);
2519 test_parse_timestamp_impl::<TimestampNanosecondType>(
2520 Some("+00:00".into()),
2521 &[0, 0, -7_200_000_000_000],
2522 );
2523 test_parse_timestamp_impl::<TimestampNanosecondType>(
2524 Some("-05:00".into()),
2525 &[18_000_000_000_000, 0, -7_200_000_000_000],
2526 );
2527 test_parse_timestamp_impl::<TimestampMicrosecondType>(
2528 Some("-03".into()),
2529 &[10_800_000_000, 0, -7_200_000_000],
2530 );
2531 test_parse_timestamp_impl::<TimestampMillisecondType>(
2532 Some("-03".into()),
2533 &[10_800_000, 0, -7_200_000],
2534 );
2535 test_parse_timestamp_impl::<TimestampSecondType>(Some("-03".into()), &[10_800, 0, -7_200]);
2536 }
2537
2538 #[test]
2539 #[cfg_attr(miri, ignore)] fn test_infer_schema_from_multiple_files() {
2541 let mut csv1 = NamedTempFile::new().unwrap();
2542 let mut csv2 = NamedTempFile::new().unwrap();
2543 let csv3 = NamedTempFile::new().unwrap(); let mut csv4 = NamedTempFile::new().unwrap();
2545 writeln!(csv1, "c1,c2,c3").unwrap();
2546 writeln!(csv1, "1,\"foo\",0.5").unwrap();
2547 writeln!(csv1, "3,\"bar\",1").unwrap();
2548 writeln!(csv1, "3,\"bar\",2e-06").unwrap();
2549 writeln!(csv2, "c1,c2,c3,c4").unwrap();
2551 writeln!(csv2, "10,,3.14,true").unwrap();
2552 writeln!(csv4, "c1,c2,c3").unwrap();
2554 writeln!(csv4, "10,\"foo\",").unwrap();
2555
2556 let schema = infer_schema_from_files(
2557 &[
2558 csv3.path().to_str().unwrap().to_string(),
2559 csv1.path().to_str().unwrap().to_string(),
2560 csv2.path().to_str().unwrap().to_string(),
2561 csv4.path().to_str().unwrap().to_string(),
2562 ],
2563 b',',
2564 Some(4), true,
2566 )
2567 .unwrap();
2568
2569 assert_eq!(schema.fields().len(), 4);
2570 assert!(schema.field(0).is_nullable());
2571 assert!(schema.field(1).is_nullable());
2572 assert!(schema.field(2).is_nullable());
2573 assert!(schema.field(3).is_nullable());
2574
2575 assert_eq!(&DataType::Int64, schema.field(0).data_type());
2576 assert_eq!(&DataType::Utf8, schema.field(1).data_type());
2577 assert_eq!(&DataType::Float64, schema.field(2).data_type());
2578 assert_eq!(&DataType::Boolean, schema.field(3).data_type());
2579 }
2580
2581 #[test]
2582 fn test_bounded() {
2583 let schema = Schema::new(vec![Field::new("int", DataType::UInt32, false)]);
2584 let data = [
2585 vec!["0"],
2586 vec!["1"],
2587 vec!["2"],
2588 vec!["3"],
2589 vec!["4"],
2590 vec!["5"],
2591 vec!["6"],
2592 ];
2593
2594 let data = data
2595 .iter()
2596 .map(|x| x.join(","))
2597 .collect::<Vec<_>>()
2598 .join("\n");
2599 let data = data.as_bytes();
2600
2601 let reader = std::io::Cursor::new(data);
2602
2603 let mut csv = ReaderBuilder::new(Arc::new(schema))
2604 .with_batch_size(2)
2605 .with_projection(vec![0])
2606 .with_bounds(2, 6)
2607 .build_buffered(reader)
2608 .unwrap();
2609
2610 let batch = csv.next().unwrap().unwrap();
2611 let a = batch.column(0);
2612 let a = a.as_any().downcast_ref::<UInt32Array>().unwrap();
2613 assert_eq!(a, &UInt32Array::from(vec![2, 3]));
2614
2615 let batch = csv.next().unwrap().unwrap();
2616 let a = batch.column(0);
2617 let a = a.as_any().downcast_ref::<UInt32Array>().unwrap();
2618 assert_eq!(a, &UInt32Array::from(vec![4, 5]));
2619
2620 assert!(csv.next().is_none());
2621 }
2622
2623 #[test]
2624 fn test_empty_projection() {
2625 let schema = Schema::new(vec![Field::new("int", DataType::UInt32, false)]);
2626 let data = [vec!["0"], vec!["1"]];
2627
2628 let data = data
2629 .iter()
2630 .map(|x| x.join(","))
2631 .collect::<Vec<_>>()
2632 .join("\n");
2633
2634 let mut csv = ReaderBuilder::new(Arc::new(schema))
2635 .with_batch_size(2)
2636 .with_projection(vec![])
2637 .build_buffered(Cursor::new(data.as_bytes()))
2638 .unwrap();
2639
2640 let batch = csv.next().unwrap().unwrap();
2641 assert_eq!(batch.columns().len(), 0);
2642 assert_eq!(batch.num_rows(), 2);
2643
2644 assert!(csv.next().is_none());
2645 }
2646
2647 #[test]
2648 fn test_parsing_bool() {
2649 assert_eq!(Some(true), parse_bool("true"));
2651 assert_eq!(Some(true), parse_bool("tRUe"));
2652 assert_eq!(Some(true), parse_bool("True"));
2653 assert_eq!(Some(true), parse_bool("TRUE"));
2654 assert_eq!(None, parse_bool("t"));
2655 assert_eq!(None, parse_bool("T"));
2656 assert_eq!(None, parse_bool(""));
2657
2658 assert_eq!(Some(false), parse_bool("false"));
2659 assert_eq!(Some(false), parse_bool("fALse"));
2660 assert_eq!(Some(false), parse_bool("False"));
2661 assert_eq!(Some(false), parse_bool("FALSE"));
2662 assert_eq!(None, parse_bool("f"));
2663 assert_eq!(None, parse_bool("F"));
2664 assert_eq!(None, parse_bool(""));
2665 }
2666
2667 #[test]
2668 fn test_parsing_float() {
2669 assert_eq!(Some(12.34), Float64Type::parse("12.34"));
2670 assert_eq!(Some(-12.34), Float64Type::parse("-12.34"));
2671 assert_eq!(Some(12.0), Float64Type::parse("12"));
2672 assert_eq!(Some(0.0), Float64Type::parse("0"));
2673 assert_eq!(Some(2.0), Float64Type::parse("2."));
2674 assert_eq!(Some(0.2), Float64Type::parse(".2"));
2675 assert!(Float64Type::parse("nan").unwrap().is_nan());
2676 assert!(Float64Type::parse("NaN").unwrap().is_nan());
2677 assert!(Float64Type::parse("inf").unwrap().is_infinite());
2678 assert!(Float64Type::parse("inf").unwrap().is_sign_positive());
2679 assert!(Float64Type::parse("-inf").unwrap().is_infinite());
2680 assert!(Float64Type::parse("-inf").unwrap().is_sign_negative());
2681 assert_eq!(None, Float64Type::parse(""));
2682 assert_eq!(None, Float64Type::parse("dd"));
2683 assert_eq!(None, Float64Type::parse("12.34.56"));
2684 }
2685
2686 #[test]
2687 fn test_non_std_quote() {
2688 let schema = Schema::new(vec![
2689 Field::new("text1", DataType::Utf8, false),
2690 Field::new("text2", DataType::Utf8, false),
2691 ]);
2692 let builder = ReaderBuilder::new(Arc::new(schema))
2693 .with_header(false)
2694 .with_quote(b'~'); let mut csv_text = Vec::new();
2697 let mut csv_writer = std::io::Cursor::new(&mut csv_text);
2698 for index in 0..10 {
2699 let text1 = format!("id{index:}");
2700 let text2 = format!("value{index:}");
2701 csv_writer
2702 .write_fmt(format_args!("~{text1}~,~{text2}~\r\n"))
2703 .unwrap();
2704 }
2705 let mut csv_reader = std::io::Cursor::new(&csv_text);
2706 let mut reader = builder.build(&mut csv_reader).unwrap();
2707 let batch = reader.next().unwrap().unwrap();
2708 let col0 = batch.column(0);
2709 assert_eq!(col0.len(), 10);
2710 let col0_arr = col0.as_any().downcast_ref::<StringArray>().unwrap();
2711 assert_eq!(col0_arr.value(0), "id0");
2712 let col1 = batch.column(1);
2713 assert_eq!(col1.len(), 10);
2714 let col1_arr = col1.as_any().downcast_ref::<StringArray>().unwrap();
2715 assert_eq!(col1_arr.value(5), "value5");
2716 }
2717
2718 #[test]
2719 fn test_non_std_escape() {
2720 let schema = Schema::new(vec![
2721 Field::new("text1", DataType::Utf8, false),
2722 Field::new("text2", DataType::Utf8, false),
2723 ]);
2724 let builder = ReaderBuilder::new(Arc::new(schema))
2725 .with_header(false)
2726 .with_escape(b'\\'); let mut csv_text = Vec::new();
2729 let mut csv_writer = std::io::Cursor::new(&mut csv_text);
2730 for index in 0..10 {
2731 let text1 = format!("id{index:}");
2732 let text2 = format!("value\\\"{index:}");
2733 csv_writer
2734 .write_fmt(format_args!("\"{text1}\",\"{text2}\"\r\n"))
2735 .unwrap();
2736 }
2737 let mut csv_reader = std::io::Cursor::new(&csv_text);
2738 let mut reader = builder.build(&mut csv_reader).unwrap();
2739 let batch = reader.next().unwrap().unwrap();
2740 let col0 = batch.column(0);
2741 assert_eq!(col0.len(), 10);
2742 let col0_arr = col0.as_any().downcast_ref::<StringArray>().unwrap();
2743 assert_eq!(col0_arr.value(0), "id0");
2744 let col1 = batch.column(1);
2745 assert_eq!(col1.len(), 10);
2746 let col1_arr = col1.as_any().downcast_ref::<StringArray>().unwrap();
2747 assert_eq!(col1_arr.value(5), "value\"5");
2748 }
2749
2750 #[test]
2751 fn test_non_std_terminator() {
2752 let schema = Schema::new(vec![
2753 Field::new("text1", DataType::Utf8, false),
2754 Field::new("text2", DataType::Utf8, false),
2755 ]);
2756 let builder = ReaderBuilder::new(Arc::new(schema))
2757 .with_header(false)
2758 .with_terminator(b'\n'); let mut csv_text = Vec::new();
2761 let mut csv_writer = std::io::Cursor::new(&mut csv_text);
2762 for index in 0..10 {
2763 let text1 = format!("id{index:}");
2764 let text2 = format!("value{index:}");
2765 csv_writer
2766 .write_fmt(format_args!("\"{text1}\",\"{text2}\"\n"))
2767 .unwrap();
2768 }
2769 let mut csv_reader = std::io::Cursor::new(&csv_text);
2770 let mut reader = builder.build(&mut csv_reader).unwrap();
2771 let batch = reader.next().unwrap().unwrap();
2772 let col0 = batch.column(0);
2773 assert_eq!(col0.len(), 10);
2774 let col0_arr = col0.as_any().downcast_ref::<StringArray>().unwrap();
2775 assert_eq!(col0_arr.value(0), "id0");
2776 let col1 = batch.column(1);
2777 assert_eq!(col1.len(), 10);
2778 let col1_arr = col1.as_any().downcast_ref::<StringArray>().unwrap();
2779 assert_eq!(col1_arr.value(5), "value5");
2780 }
2781
2782 #[test]
2783 fn test_header_bounds() {
2784 let csv = "a,b\na,b\na,b\na,b\na,b\n";
2785 let tests = [
2786 (None, false, 5),
2787 (None, true, 4),
2788 (Some((0, 4)), false, 4),
2789 (Some((1, 4)), false, 3),
2790 (Some((0, 4)), true, 4),
2791 (Some((1, 4)), true, 3),
2792 ];
2793 let schema = Arc::new(Schema::new(vec![
2794 Field::new("a", DataType::Utf8, false),
2795 Field::new("a", DataType::Utf8, false),
2796 ]));
2797
2798 for (idx, (bounds, has_header, expected)) in tests.into_iter().enumerate() {
2799 let mut reader = ReaderBuilder::new(schema.clone()).with_header(has_header);
2800 if let Some((start, end)) = bounds {
2801 reader = reader.with_bounds(start, end);
2802 }
2803 let b = reader
2804 .build_buffered(Cursor::new(csv.as_bytes()))
2805 .unwrap()
2806 .next()
2807 .unwrap()
2808 .unwrap();
2809 assert_eq!(b.num_rows(), expected, "{idx}");
2810 }
2811 }
2812
2813 #[test]
2814 fn test_header_validation() {
2815 let schema = Arc::new(Schema::new(vec![
2816 Field::new("a", DataType::Int32, false),
2817 Field::new("b", DataType::Int32, false),
2818 ]));
2819
2820 let csv = "a,c\n1,2\n";
2821 let err = ReaderBuilder::new(schema.clone())
2822 .with_header(true)
2823 .with_header_validation(true)
2824 .build_buffered(Cursor::new(csv.as_bytes()))
2825 .unwrap()
2826 .next()
2827 .unwrap()
2828 .unwrap_err()
2829 .to_string();
2830 assert_eq!(
2831 err,
2832 "Csv error: CSV header does not match schema at column 1: expected \"b\" but found \"c\""
2833 );
2834
2835 let batch = ReaderBuilder::new(schema)
2836 .with_header(true)
2837 .with_header_validation(false)
2838 .build_buffered(Cursor::new(csv.as_bytes()))
2839 .unwrap()
2840 .next()
2841 .unwrap()
2842 .unwrap();
2843 assert_eq!(batch.num_rows(), 1);
2844 }
2845
2846 #[test]
2847 fn test_header_validation_with_buffered_reader() {
2848 let schema = Arc::new(Schema::new(vec![
2849 Field::new("a", DataType::Int32, false),
2850 Field::new("b", DataType::Int32, false),
2851 ]));
2852
2853 let csv = "a,b\n1,2\n";
2854 let buffered = std::io::BufReader::with_capacity(1, Cursor::new(csv.as_bytes()));
2855 let batch = ReaderBuilder::new(schema)
2856 .with_header(true)
2857 .with_header_validation(true)
2858 .build_buffered(buffered)
2859 .unwrap()
2860 .next()
2861 .unwrap()
2862 .unwrap();
2863
2864 assert_eq!(batch.num_rows(), 1);
2865 let a = batch.column(0).as_primitive::<Int32Type>();
2866 assert_eq!(a.value(0), 1);
2867 }
2868
2869 #[test]
2870 fn test_header_validation_with_truncated_rows() {
2871 let schema = Arc::new(Schema::new(vec![
2872 Field::new("a", DataType::Int32, true),
2873 Field::new("b", DataType::Int32, true),
2874 ]));
2875
2876 let csv = "a\n1\n";
2877 let err = ReaderBuilder::new(schema.clone())
2878 .with_header(true)
2879 .with_header_validation(true)
2880 .with_truncated_rows(true)
2881 .build_buffered(Cursor::new(csv.as_bytes()))
2882 .unwrap()
2883 .next()
2884 .unwrap()
2885 .unwrap_err()
2886 .to_string();
2887 assert_eq!(
2888 err,
2889 "Csv error: CSV header does not match schema at column 1: expected \"b\" but found \"\"",
2890 )
2891 }
2892
2893 #[test]
2894 fn test_null_boolean() {
2895 let csv = "true,false\nFalse,True\n,True\nFalse,";
2896 let schema = Arc::new(Schema::new(vec![
2897 Field::new("a", DataType::Boolean, true),
2898 Field::new("a", DataType::Boolean, true),
2899 ]));
2900
2901 let b = ReaderBuilder::new(schema)
2902 .build_buffered(Cursor::new(csv.as_bytes()))
2903 .unwrap()
2904 .next()
2905 .unwrap()
2906 .unwrap();
2907
2908 assert_eq!(b.num_rows(), 4);
2909 assert_eq!(b.num_columns(), 2);
2910
2911 let c = b.column(0).as_boolean();
2912 assert_eq!(c.null_count(), 1);
2913 assert!(c.value(0));
2914 assert!(!c.value(1));
2915 assert!(c.is_null(2));
2916 assert!(!c.value(3));
2917
2918 let c = b.column(1).as_boolean();
2919 assert_eq!(c.null_count(), 1);
2920 assert!(!c.value(0));
2921 assert!(c.value(1));
2922 assert!(c.value(2));
2923 assert!(c.is_null(3));
2924 }
2925
2926 #[test]
2927 fn test_truncated_rows() {
2928 let data = "a,b,c\n1,2,3\n4,5\n\n6,7,8";
2929 let schema = Arc::new(Schema::new(vec![
2930 Field::new("a", DataType::Int32, true),
2931 Field::new("b", DataType::Int32, true),
2932 Field::new("c", DataType::Int32, true),
2933 ]));
2934
2935 let reader = ReaderBuilder::new(schema.clone())
2936 .with_header(true)
2937 .with_truncated_rows(true)
2938 .build(Cursor::new(data))
2939 .unwrap();
2940
2941 let batches = reader.collect::<Result<Vec<_>, _>>();
2942 assert!(batches.is_ok());
2943 let batch = batches.unwrap().into_iter().next().unwrap();
2944 assert_eq!(batch.num_rows(), 3);
2946
2947 let reader = ReaderBuilder::new(schema.clone())
2948 .with_header(true)
2949 .with_truncated_rows(false)
2950 .build(Cursor::new(data))
2951 .unwrap();
2952
2953 let batches = reader.collect::<Result<Vec<_>, _>>();
2954 assert!(match batches {
2955 Err(ArrowError::CsvError(e)) => e.contains("incorrect number of fields"),
2956 _ => false,
2957 });
2958 }
2959
2960 #[test]
2961 fn test_truncated_rows_csv() {
2962 let file = File::open("test/data/truncated_rows.csv").unwrap();
2963 let schema = Arc::new(Schema::new(vec![
2964 Field::new("Name", DataType::Utf8, true),
2965 Field::new("Age", DataType::UInt32, true),
2966 Field::new("Occupation", DataType::Utf8, true),
2967 Field::new("DOB", DataType::Date32, true),
2968 ]));
2969 let reader = ReaderBuilder::new(schema.clone())
2970 .with_header(true)
2971 .with_batch_size(24)
2972 .with_truncated_rows(true);
2973 let csv = reader.build(file).unwrap();
2974 let batches = csv.collect::<Result<Vec<_>, _>>().unwrap();
2975
2976 assert_eq!(batches.len(), 1);
2977 let batch = &batches[0];
2978 assert_eq!(batch.num_rows(), 6);
2979 assert_eq!(batch.num_columns(), 4);
2980 let name = batch
2981 .column(0)
2982 .as_any()
2983 .downcast_ref::<StringArray>()
2984 .unwrap();
2985 let age = batch
2986 .column(1)
2987 .as_any()
2988 .downcast_ref::<UInt32Array>()
2989 .unwrap();
2990 let occupation = batch
2991 .column(2)
2992 .as_any()
2993 .downcast_ref::<StringArray>()
2994 .unwrap();
2995 let dob = batch
2996 .column(3)
2997 .as_any()
2998 .downcast_ref::<Date32Array>()
2999 .unwrap();
3000
3001 assert_eq!(name.value(0), "A1");
3002 assert_eq!(name.value(1), "B2");
3003 assert!(name.is_null(2));
3004 assert_eq!(name.value(3), "C3");
3005 assert_eq!(name.value(4), "D4");
3006 assert_eq!(name.value(5), "E5");
3007
3008 assert_eq!(age.value(0), 34);
3009 assert_eq!(age.value(1), 29);
3010 assert!(age.is_null(2));
3011 assert_eq!(age.value(3), 45);
3012 assert!(age.is_null(4));
3013 assert_eq!(age.value(5), 31);
3014
3015 assert_eq!(occupation.value(0), "Engineer");
3016 assert_eq!(occupation.value(1), "Doctor");
3017 assert!(occupation.is_null(2));
3018 assert_eq!(occupation.value(3), "Artist");
3019 assert!(occupation.is_null(4));
3020 assert!(occupation.is_null(5));
3021
3022 assert_eq!(dob.value(0), 5675);
3023 assert!(dob.is_null(1));
3024 assert!(dob.is_null(2));
3025 assert_eq!(dob.value(3), -1858);
3026 assert!(dob.is_null(4));
3027 assert!(dob.is_null(5));
3028 }
3029
3030 fn truncated_row_count_schema() -> SchemaRef {
3032 Arc::new(Schema::new(vec![
3033 Field::new("name", DataType::Utf8, true),
3034 Field::new("age", DataType::Int32, true),
3035 Field::new("city", DataType::Utf8, true),
3036 ]))
3037 }
3038
3039 #[test]
3040 fn test_truncated_row_count_counts_padded_rows() {
3041 let data = "name,age,city\nAlice,25,Rome\nBob,30\n";
3042
3043 let mut reader = ReaderBuilder::new(truncated_row_count_schema())
3044 .with_header(true)
3045 .with_truncated_rows(true)
3046 .build(Cursor::new(data))
3047 .unwrap();
3048
3049 let batches = reader.by_ref().collect::<Result<Vec<_>, _>>().unwrap();
3050 assert_eq!(batches[0].num_rows(), 2);
3051 assert_eq!(reader.truncated_row_count(), 1);
3052 }
3053
3054 #[test]
3055 fn test_truncated_row_count_ignores_empty_trailing_field() {
3056 let data = "name,age,city\nAlice,25,Rome\nCarol,35,\n";
3060
3061 let mut reader = ReaderBuilder::new(truncated_row_count_schema())
3062 .with_header(true)
3063 .with_truncated_rows(true)
3064 .build(Cursor::new(data))
3065 .unwrap();
3066
3067 let batches = reader.by_ref().collect::<Result<Vec<_>, _>>().unwrap();
3068 let batch = &batches[0];
3069 assert_eq!(batch.num_rows(), 2);
3070 assert!(batch.column(2).is_null(1));
3071 assert_eq!(reader.truncated_row_count(), 0);
3072 }
3073
3074 #[test]
3075 fn test_truncated_row_count_clean_file() {
3076 let data = "name,age,city\nAlice,25,Rome\nBob,30,Milan\n";
3077
3078 let mut reader = ReaderBuilder::new(truncated_row_count_schema())
3079 .with_header(true)
3080 .with_truncated_rows(true)
3081 .build(Cursor::new(data))
3082 .unwrap();
3083
3084 let batches = reader.by_ref().collect::<Result<Vec<_>, _>>().unwrap();
3085 assert_eq!(batches[0].num_rows(), 2);
3086 assert_eq!(reader.truncated_row_count(), 0);
3087 }
3088
3089 #[test]
3090 fn test_truncated_row_count_without_truncated_rows() {
3091 let data = "name,age,city\nAlice,25,Rome\nBob,30\n";
3092
3093 let mut reader = ReaderBuilder::new(truncated_row_count_schema())
3094 .with_header(true)
3095 .with_truncated_rows(false)
3096 .build(Cursor::new(data))
3097 .unwrap();
3098
3099 let err = reader.by_ref().collect::<Result<Vec<_>, _>>().unwrap_err();
3101 assert!(
3102 err.to_string().contains("incorrect number of fields"),
3103 "{err}"
3104 );
3105 assert_eq!(reader.truncated_row_count(), 0);
3106 }
3107
3108 #[test]
3109 fn test_truncated_row_count_accumulates_across_batches() {
3110 let data = "name,age,city\nn0,0\nn1,1\nn2,2\nn3,3\nn4,4\nn5,5\n";
3112
3113 let mut reader = ReaderBuilder::new(truncated_row_count_schema())
3114 .with_header(true)
3115 .with_truncated_rows(true)
3116 .with_batch_size(2)
3117 .build(Cursor::new(data))
3118 .unwrap();
3119
3120 let mut running = vec![];
3121 while let Some(batch) = reader.next().transpose().unwrap() {
3122 assert_eq!(batch.num_rows(), 2);
3123 running.push(reader.truncated_row_count());
3124 }
3125
3126 assert_eq!(running, vec![2, 4, 6]);
3128 assert_eq!(reader.truncated_row_count(), 6);
3129 }
3130
3131 #[test]
3132 fn test_truncated_row_count_excludes_skipped_rows() {
3133 let data = "name,age\nAlice,25,Rome\nBob,30,Milan\n";
3136
3137 let mut reader = ReaderBuilder::new(truncated_row_count_schema())
3138 .with_header(true)
3139 .with_truncated_rows(true)
3140 .build(Cursor::new(data))
3141 .unwrap();
3142
3143 let batches = reader.by_ref().collect::<Result<Vec<_>, _>>().unwrap();
3144 assert_eq!(batches[0].num_rows(), 2);
3145 assert_eq!(reader.truncated_row_count(), 0);
3146 }
3147
3148 #[test]
3149 fn test_truncated_row_count_on_decoder() {
3150 let data = "1,2\n3\n";
3151 let schema = Arc::new(Schema::new(vec![
3152 Field::new("a", DataType::Int32, true),
3153 Field::new("b", DataType::Int32, true),
3154 ]));
3155
3156 let mut decoder = ReaderBuilder::new(schema)
3157 .with_truncated_rows(true)
3158 .build_decoder();
3159
3160 assert_eq!(decoder.truncated_row_count(), 0);
3161 let decoded = decoder.decode(data.as_bytes()).unwrap();
3162 assert_eq!(decoded, data.len());
3163 decoder.flush().unwrap().unwrap();
3164 assert_eq!(decoder.truncated_row_count(), 1);
3165 }
3166
3167 #[test]
3168 fn test_truncated_rows_not_nullable_error() {
3169 let data = "a,b,c\n1,2,3\n4,5";
3170 let schema = Arc::new(Schema::new(vec![
3171 Field::new("a", DataType::Int32, false),
3172 Field::new("b", DataType::Int32, false),
3173 Field::new("c", DataType::Int32, false),
3174 ]));
3175
3176 let reader = ReaderBuilder::new(schema.clone())
3177 .with_header(true)
3178 .with_truncated_rows(true)
3179 .build(Cursor::new(data))
3180 .unwrap();
3181
3182 let batches = reader.collect::<Result<Vec<_>, _>>();
3183 assert!(match batches {
3184 Err(ArrowError::InvalidArgumentError(e)) => e.contains("contains null values"),
3185 _ => false,
3186 });
3187 }
3188
3189 #[test]
3190 #[cfg_attr(miri, ignore)] fn test_buffered() {
3192 let tests = [
3193 ("test/data/uk_cities.csv", false, 37),
3194 ("test/data/various_types.csv", true, 10),
3195 ("test/data/decimal_test.csv", false, 10),
3196 ];
3197
3198 for (path, has_header, expected_rows) in tests {
3199 let (schema, _) = Format::default()
3200 .infer_schema(File::open(path).unwrap(), None)
3201 .unwrap();
3202 let schema = Arc::new(schema);
3203
3204 for batch_size in [1, 4] {
3205 for capacity in [1, 3, 7, 100] {
3206 let reader = ReaderBuilder::new(schema.clone())
3207 .with_batch_size(batch_size)
3208 .with_header(has_header)
3209 .build(File::open(path).unwrap())
3210 .unwrap();
3211
3212 let expected = reader.collect::<Result<Vec<_>, _>>().unwrap();
3213
3214 assert_eq!(
3215 expected.iter().map(|x| x.num_rows()).sum::<usize>(),
3216 expected_rows
3217 );
3218
3219 let buffered =
3220 std::io::BufReader::with_capacity(capacity, File::open(path).unwrap());
3221
3222 let reader = ReaderBuilder::new(schema.clone())
3223 .with_batch_size(batch_size)
3224 .with_header(has_header)
3225 .build_buffered(buffered)
3226 .unwrap();
3227
3228 let actual = reader.collect::<Result<Vec<_>, _>>().unwrap();
3229 assert_eq!(expected, actual)
3230 }
3231 }
3232 }
3233 }
3234
3235 fn err_test(csv: &[u8], expected: &str) {
3236 fn err_test_with_schema(csv: &[u8], expected: &str, schema: Arc<Schema>) {
3237 let buffer = std::io::BufReader::with_capacity(2, Cursor::new(csv));
3238 let b = ReaderBuilder::new(schema)
3239 .with_batch_size(2)
3240 .build_buffered(buffer)
3241 .unwrap();
3242 let err = b.collect::<Result<Vec<_>, _>>().unwrap_err().to_string();
3243 assert_eq!(err, expected)
3244 }
3245
3246 let schema_utf8 = Arc::new(Schema::new(vec![
3247 Field::new("text1", DataType::Utf8, true),
3248 Field::new("text2", DataType::Utf8, true),
3249 ]));
3250 err_test_with_schema(csv, expected, schema_utf8);
3251
3252 let schema_utf8view = Arc::new(Schema::new(vec![
3253 Field::new("text1", DataType::Utf8View, true),
3254 Field::new("text2", DataType::Utf8View, true),
3255 ]));
3256 err_test_with_schema(csv, expected, schema_utf8view);
3257 }
3258
3259 #[test]
3260 fn test_invalid_utf8() {
3261 err_test(
3262 b"sdf,dsfg\ndfd,hgh\xFFue\n,sds\nFalhghse,",
3263 "Csv error: Encountered invalid UTF-8 data for line 2 and field 2",
3264 );
3265
3266 err_test(
3267 b"sdf,dsfg\ndksdk,jf\nd\xFFfd,hghue\n,sds\nFalhghse,",
3268 "Csv error: Encountered invalid UTF-8 data for line 3 and field 1",
3269 );
3270
3271 err_test(
3272 b"sdf,dsfg\ndksdk,jf\ndsdsfd,hghue\n,sds\nFalhghse,\xFF",
3273 "Csv error: Encountered invalid UTF-8 data for line 5 and field 2",
3274 );
3275
3276 err_test(
3277 b"\xFFsdf,dsfg\ndksdk,jf\ndsdsfd,hghue\n,sds\nFalhghse,\xFF",
3278 "Csv error: Encountered invalid UTF-8 data for line 1 and field 1",
3279 );
3280 }
3281
3282 struct InstrumentedRead<R> {
3283 r: R,
3284 fill_count: usize,
3285 fill_sizes: Vec<usize>,
3286 }
3287
3288 impl<R> InstrumentedRead<R> {
3289 fn new(r: R) -> Self {
3290 Self {
3291 r,
3292 fill_count: 0,
3293 fill_sizes: vec![],
3294 }
3295 }
3296 }
3297
3298 impl<R: Seek> Seek for InstrumentedRead<R> {
3299 fn seek(&mut self, pos: SeekFrom) -> std::io::Result<u64> {
3300 self.r.seek(pos)
3301 }
3302 }
3303
3304 impl<R: BufRead> Read for InstrumentedRead<R> {
3305 fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
3306 self.r.read(buf)
3307 }
3308 }
3309
3310 impl<R: BufRead> BufRead for InstrumentedRead<R> {
3311 fn fill_buf(&mut self) -> std::io::Result<&[u8]> {
3312 self.fill_count += 1;
3313 let buf = self.r.fill_buf()?;
3314 self.fill_sizes.push(buf.len());
3315 Ok(buf)
3316 }
3317
3318 fn consume(&mut self, amt: usize) {
3319 self.r.consume(amt)
3320 }
3321 }
3322
3323 #[test]
3324 fn test_io() {
3325 let schema = Arc::new(Schema::new(vec![
3326 Field::new("a", DataType::Utf8, false),
3327 Field::new("b", DataType::Utf8, false),
3328 ]));
3329 let csv = "foo,bar\nbaz,foo\na,b\nc,d";
3330 let mut read = InstrumentedRead::new(Cursor::new(csv.as_bytes()));
3331 let reader = ReaderBuilder::new(schema)
3332 .with_batch_size(3)
3333 .build_buffered(&mut read)
3334 .unwrap();
3335
3336 let batches = reader.collect::<Result<Vec<_>, _>>().unwrap();
3337 assert_eq!(batches.len(), 2);
3338 assert_eq!(batches[0].num_rows(), 3);
3339 assert_eq!(batches[1].num_rows(), 1);
3340
3341 assert_eq!(&read.fill_sizes, &[23, 3, 0, 0]);
3347 assert_eq!(read.fill_count, 4);
3348 }
3349
3350 #[test]
3351 #[cfg_attr(miri, ignore)] fn test_inference() {
3353 let cases: &[(&[&str], DataType)] = &[
3354 (&[], DataType::Null),
3355 (&["false", "12"], DataType::Utf8),
3356 (&["12", "cupcakes"], DataType::Utf8),
3357 (&["12", "12.4"], DataType::Float64),
3358 (&["14050", "24332"], DataType::Int64),
3359 (&["14050.0", "true"], DataType::Utf8),
3360 (&["14050", "2020-03-19 00:00:00"], DataType::Utf8),
3361 (&["14050", "2340.0", "2020-03-19 00:00:00"], DataType::Utf8),
3362 (
3363 &["2020-03-19 02:00:00", "2020-03-19 00:00:00"],
3364 DataType::Timestamp(TimeUnit::Second, None),
3365 ),
3366 (&["2020-03-19", "2020-03-20"], DataType::Date32),
3367 (
3368 &["2020-03-19", "2020-03-19 02:00:00", "2020-03-19 00:00:00"],
3369 DataType::Timestamp(TimeUnit::Second, None),
3370 ),
3371 (
3372 &[
3373 "2020-03-19",
3374 "2020-03-19 02:00:00",
3375 "2020-03-19 00:00:00.000",
3376 ],
3377 DataType::Timestamp(TimeUnit::Millisecond, None),
3378 ),
3379 (
3380 &[
3381 "2020-03-19",
3382 "2020-03-19 02:00:00",
3383 "2020-03-19 00:00:00.000000",
3384 ],
3385 DataType::Timestamp(TimeUnit::Microsecond, None),
3386 ),
3387 (
3388 &["2020-03-19 02:00:00+02:00", "2020-03-19 02:00:00Z"],
3389 DataType::Timestamp(TimeUnit::Second, None),
3390 ),
3391 (
3392 &[
3393 "2020-03-19",
3394 "2020-03-19 02:00:00+02:00",
3395 "2020-03-19 02:00:00Z",
3396 "2020-03-19 02:00:00.12Z",
3397 ],
3398 DataType::Timestamp(TimeUnit::Millisecond, None),
3399 ),
3400 (
3401 &[
3402 "2020-03-19",
3403 "2020-03-19 02:00:00.000000000",
3404 "2020-03-19 00:00:00.000000",
3405 ],
3406 DataType::Timestamp(TimeUnit::Nanosecond, None),
3407 ),
3408 ];
3409
3410 for (values, expected) in cases {
3411 let mut t = InferredDataType::default();
3412 for v in *values {
3413 t.update(v)
3414 }
3415 assert_eq!(&t.get(), expected, "{values:?}")
3416 }
3417 }
3418
3419 #[test]
3420 #[cfg_attr(miri, ignore)] fn test_record_length_mismatch() {
3422 let csv = "\
3423 a,b,c\n\
3424 1,2,3\n\
3425 4,5\n\
3426 6,7,8";
3427 let mut read = Cursor::new(csv.as_bytes());
3428 let result = Format::default()
3429 .with_header(true)
3430 .infer_schema(&mut read, None);
3431 assert!(result.is_err());
3432 assert_eq!(
3434 result.err().unwrap().to_string(),
3435 "Csv error: Encountered unequal lengths between records on CSV file. Expected 3 records, found 2 records at line 3"
3436 );
3437 }
3438
3439 #[test]
3440 fn test_comment() {
3441 let schema = Schema::new(vec![
3442 Field::new("a", DataType::Int8, false),
3443 Field::new("b", DataType::Int8, false),
3444 ]);
3445
3446 let csv = "# comment1 \n1,2\n#comment2\n11,22";
3447 let mut read = Cursor::new(csv.as_bytes());
3448 let reader = ReaderBuilder::new(Arc::new(schema))
3449 .with_comment(b'#')
3450 .build(&mut read)
3451 .unwrap();
3452
3453 let batches = reader.collect::<Result<Vec<_>, _>>().unwrap();
3454 assert_eq!(batches.len(), 1);
3455 let b = batches.first().unwrap();
3456 assert_eq!(b.num_columns(), 2);
3457 assert_eq!(
3458 b.column(0)
3459 .as_any()
3460 .downcast_ref::<Int8Array>()
3461 .unwrap()
3462 .values(),
3463 &vec![1, 11]
3464 );
3465 assert_eq!(
3466 b.column(1)
3467 .as_any()
3468 .downcast_ref::<Int8Array>()
3469 .unwrap()
3470 .values(),
3471 &vec![2, 22]
3472 );
3473 }
3474
3475 #[test]
3476 fn test_parse_string_view_single_column() {
3477 let csv = ["foo", "something_cannot_be_inlined", "foobar"].join("\n");
3478 let schema = Arc::new(Schema::new(vec![Field::new(
3479 "c1",
3480 DataType::Utf8View,
3481 true,
3482 )]));
3483
3484 let mut decoder = ReaderBuilder::new(schema).build_decoder();
3485
3486 let decoded = decoder.decode(csv.as_bytes()).unwrap();
3487 assert_eq!(decoded, csv.len());
3488 decoder.decode(&[]).unwrap();
3489
3490 let batch = decoder.flush().unwrap().unwrap();
3491 assert_eq!(batch.num_columns(), 1);
3492 assert_eq!(batch.num_rows(), 3);
3493 let col = batch.column(0).as_string_view();
3494 assert_eq!(col.data_type(), &DataType::Utf8View);
3495 assert_eq!(col.value(0), "foo");
3496 assert_eq!(col.value(1), "something_cannot_be_inlined");
3497 assert_eq!(col.value(2), "foobar");
3498 }
3499
3500 #[test]
3501 fn test_parse_string_view_multi_column() {
3502 let csv = ["foo,", ",something_cannot_be_inlined", "foobarfoobar,bar"].join("\n");
3503 let schema = Arc::new(Schema::new(vec![
3504 Field::new("c1", DataType::Utf8View, true),
3505 Field::new("c2", DataType::Utf8View, true),
3506 ]));
3507
3508 let mut decoder = ReaderBuilder::new(schema).build_decoder();
3509
3510 let decoded = decoder.decode(csv.as_bytes()).unwrap();
3511 assert_eq!(decoded, csv.len());
3512 decoder.decode(&[]).unwrap();
3513
3514 let batch = decoder.flush().unwrap().unwrap();
3515 assert_eq!(batch.num_columns(), 2);
3516 assert_eq!(batch.num_rows(), 3);
3517 let c1 = batch.column(0).as_string_view();
3518 let c2 = batch.column(1).as_string_view();
3519 assert_eq!(c1.data_type(), &DataType::Utf8View);
3520 assert_eq!(c2.data_type(), &DataType::Utf8View);
3521
3522 assert!(!c1.is_null(0));
3523 assert!(c1.is_null(1));
3524 assert!(!c1.is_null(2));
3525 assert_eq!(c1.value(0), "foo");
3526 assert_eq!(c1.value(2), "foobarfoobar");
3527
3528 assert!(c2.is_null(0));
3529 assert!(!c2.is_null(1));
3530 assert!(!c2.is_null(2));
3531 assert_eq!(c2.value(1), "something_cannot_be_inlined");
3532 assert_eq!(c2.value(2), "bar");
3533 }
3534
3535 #[test]
3536 #[cfg_attr(miri, ignore)] fn test_float_precision() {
3538 let data = [
3539 "f16,f32,f64",
3540 "1.5,1.5,1.5",
3541 "0.25,0.25,0.25",
3542 "1.23456789,1.23456789,1.23456789",
3543 "1.234567890123456,1.234567890123456,1.234567890123456",
3544 "-2.5,-2.5,-2.5",
3545 "0,0,0",
3546 ",,",
3547 ]
3548 .join("\n");
3549
3550 let schema = Schema::new(vec![
3551 Field::new("f16", DataType::Float16, true),
3552 Field::new("f32", DataType::Float32, true),
3553 Field::new("f64", DataType::Float64, true),
3554 ]);
3555
3556 let mut reader = ReaderBuilder::new(Arc::new(schema))
3557 .with_header(true)
3558 .build(Cursor::new(data))
3559 .unwrap();
3560
3561 let batch = reader.next().unwrap().unwrap();
3562 assert_eq!(batch.num_rows(), 7);
3563
3564 let f16_col = batch.column(0).as_primitive::<Float16Type>();
3565 let f32_col = batch.column(1).as_primitive::<Float32Type>();
3566 let f64_col = batch.column(2).as_primitive::<Float64Type>();
3567
3568 assert_eq!(f16_col.value(0), half::f16::from_f32(1.5));
3569 assert_eq!(f32_col.value(0), 1.5f32);
3570 assert_eq!(f64_col.value(0), 1.5f64);
3571
3572 assert_eq!(f16_col.value(1), half::f16::from_f32(0.25));
3573 assert_eq!(f32_col.value(1), 0.25f32);
3574 assert_eq!(f64_col.value(1), 0.25f64);
3575
3576 assert_eq!(f16_col.value(2), half::f16::from_f32(1.234_567_9));
3577 assert_eq!(f32_col.value(2), 1.234_567_9_f32);
3578 assert_eq!(f64_col.value(2), 1.23456789f64);
3579
3580 assert_eq!(f16_col.value(3), half::f16::from_f64(1.234567890123456f64));
3581 assert_eq!(f32_col.value(3), 1.234_567_9_f32);
3582 assert_eq!(f64_col.value(3), 1.234567890123456f64);
3583
3584 assert_eq!(f16_col.value(4), half::f16::from_f32(-2.5));
3585 assert_eq!(f32_col.value(4), -2.5f32);
3586 assert_eq!(f64_col.value(4), -2.5f64);
3587
3588 assert_eq!(f16_col.value(5), half::f16::from_f32(0.0));
3589 assert_eq!(f32_col.value(5), 0.0f32);
3590 assert_eq!(f64_col.value(5), 0.0f64);
3591
3592 assert!(f16_col.is_null(6));
3593 assert!(f32_col.is_null(6));
3594 assert!(f64_col.is_null(6));
3595 }
3596}