Skip to main content

arrow_csv/reader/
mod.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18//! CSV Reading: [`Reader`] and [`ReaderBuilder`]
19//!
20//! # Basic Usage
21//!
22//! This CSV reader allows CSV files to be read into the Arrow memory model. Records are
23//! loaded in batches and are then converted from row-based data to columnar data.
24//!
25//! Example:
26//!
27//! ```
28//! # use arrow_schema::*;
29//! # use arrow_csv::{Reader, ReaderBuilder};
30//! # use std::fs::File;
31//! # use std::sync::Arc;
32//!
33//! let schema = Schema::new(vec![
34//!     Field::new("city", DataType::Utf8, false),
35//!     Field::new("lat", DataType::Float64, false),
36//!     Field::new("lng", DataType::Float64, false),
37//! ]);
38//!
39//! let file = File::open("test/data/uk_cities.csv").unwrap();
40//!
41//! let mut csv = ReaderBuilder::new(Arc::new(schema)).build(file).unwrap();
42//! let batch = csv.next().unwrap().unwrap();
43//! ```
44//!
45//! # Example: Numeric calculations on CSV
46//! This code finds the maximum value in column 0 of a CSV file containing
47//! ```csv
48//! c1,c2,c3,c4
49//! 1,1.1,"hong kong",true
50//! 3,323.12,"XiAn",false
51//! 10,131323.12,"cheng du",false
52//! ```
53//!
54//! ```
55//! # use arrow_array::cast::AsArray;
56//! # use arrow_array::types::Int16Type;
57//! # use arrow_csv::ReaderBuilder;
58//! # use arrow_schema::{DataType, Field, Schema};
59//! # use std::fs::File;
60//! # use std::sync::Arc;
61//! // Open the example file
62//! let file = File::open("test/data/example.csv").unwrap();
63//! let csv_schema = Schema::new(vec![
64//!     Field::new("c1", DataType::Int16, true),
65//!     Field::new("c2", DataType::Float32, true),
66//!     Field::new("c3", DataType::Utf8, true),
67//!     Field::new("c4", DataType::Boolean, true),
68//! ]);
69//! let mut reader = ReaderBuilder::new(Arc::new(csv_schema))
70//!     .with_header(true)
71//!     .build(file)
72//!     .unwrap();
73//! // find the maximum value in column 0 across all batches
74//! let mut max_c0 = 0;
75//! while let Some(r) = reader.next() {
76//!   let r = r.unwrap(); // handle error
77//!   // get the max value in column(0) for this batch
78//!   let col = r.column(0).as_primitive::<Int16Type>();
79//!   let batch_max = col.iter().max().flatten().unwrap_or_default();
80//!   max_c0 = max_c0.max(batch_max);
81//! }
82//! assert_eq!(max_c0, 10);
83//!```
84//!
85//! # Async Usage
86//!
87//! The lower-level [`Decoder`] can be integrated with various forms of async data streams,
88//! and is designed to be agnostic to the various different kinds of async IO primitives found
89//! within the Rust ecosystem.
90//!
91//! For example, see below for how it can be used with an arbitrary `Stream` of `Bytes`
92//!
93//! ```
94//! # use std::task::{Poll, ready};
95//! # use bytes::{Buf, Bytes};
96//! # use arrow_schema::ArrowError;
97//! # use futures::stream::{Stream, StreamExt};
98//! # use arrow_array::RecordBatch;
99//! # use arrow_csv::reader::Decoder;
100//! #
101//! fn decode_stream<S: Stream<Item = Bytes> + Unpin>(
102//!     mut decoder: Decoder,
103//!     mut input: S,
104//! ) -> impl Stream<Item = Result<RecordBatch, ArrowError>> {
105//!     let mut buffered = Bytes::new();
106//!     futures::stream::poll_fn(move |cx| {
107//!         loop {
108//!             if buffered.is_empty() {
109//!                 if let Some(b) = ready!(input.poll_next_unpin(cx)) {
110//!                     buffered = b;
111//!                 }
112//!                 // Note: don't break on `None` as the decoder needs
113//!                 // to be called with an empty array to delimit the
114//!                 // final record
115//!             }
116//!             let decoded = match decoder.decode(buffered.as_ref()) {
117//!                 Ok(0) => break,
118//!                 Ok(decoded) => decoded,
119//!                 Err(e) => return Poll::Ready(Some(Err(e))),
120//!             };
121//!             buffered.advance(decoded);
122//!         }
123//!
124//!         Poll::Ready(decoder.flush().transpose())
125//!     })
126//! }
127//!
128//! ```
129//!
130//! In a similar vein, it can also be used with tokio-based IO primitives
131//!
132//! ```
133//! # use std::pin::Pin;
134//! # use std::task::{Poll, ready};
135//! # use futures::Stream;
136//! # use tokio::io::AsyncBufRead;
137//! # use arrow_array::RecordBatch;
138//! # use arrow_csv::reader::Decoder;
139//! # use arrow_schema::ArrowError;
140//! fn decode_stream<R: AsyncBufRead + Unpin>(
141//!     mut decoder: Decoder,
142//!     mut reader: R,
143//! ) -> impl Stream<Item = Result<RecordBatch, ArrowError>> {
144//!     futures::stream::poll_fn(move |cx| {
145//!         loop {
146//!             let b = match ready!(Pin::new(&mut reader).poll_fill_buf(cx)) {
147//!                 Ok(b) => b,
148//!                 Err(e) => return Poll::Ready(Some(Err(e.into()))),
149//!             };
150//!             let decoded = match decoder.decode(b) {
151//!                 // Note: the decoder needs to be called with an empty
152//!                 // array to delimit the final record
153//!                 Ok(0) => break,
154//!                 Ok(decoded) => decoded,
155//!                 Err(e) => return Poll::Ready(Some(Err(e))),
156//!             };
157//!             Pin::new(&mut reader).consume(decoded);
158//!         }
159//!
160//!         Poll::Ready(decoder.flush().transpose())
161//!     })
162//! }
163//! ```
164//!
165
166mod 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
185/// Order should match [`InferredDataType`]
186static REGEX_SET: LazyLock<RegexSet> = LazyLock::new(|| {
187    RegexSet::new([
188        r"(?i)^(true)$|^(false)$(?-i)", //BOOLEAN
189        r"^[+-]?(\d+)$",                //INTEGER
190        r"^[+-]?((\d*\.\d+|\d+\.\d*)([eE][-+]?\d+)?|\d+([eE][-+]?\d+))$", //DECIMAL
191        r"^\d{4}-\d\d-\d\d$",           //DATE32
192        r"^\d{4}-\d\d-\d\d[T ]\d\d:\d\d:\d\d(?:[^\d\.].*)?$", //Timestamp(Second)
193        r"^\d{4}-\d\d-\d\d[T ]\d\d:\d\d:\d\d\.\d{1,3}(?:[^\d].*)?$", //Timestamp(Millisecond)
194        r"^\d{4}-\d\d-\d\d[T ]\d\d:\d\d:\d\d\.\d{1,6}(?:[^\d].*)?$", //Timestamp(Microsecond)
195        r"^\d{4}-\d\d-\d\d[T ]\d\d:\d\d:\d\d\.\d{1,9}(?:[^\d].*)?$", //Timestamp(Nanosecond)
196    ])
197    .unwrap()
198});
199
200/// A wrapper over `Option<Regex>` to check if the value is `NULL`.
201#[derive(Debug, Clone, Default)]
202struct NullRegex(Option<Regex>);
203
204impl NullRegex {
205    /// Returns true if the value should be considered as `NULL` according to
206    /// the provided regular expression.
207    #[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 booleans indicating type
219    ///
220    /// 0 - Boolean
221    /// 1 - Integer
222    /// 2 - Float64
223    /// 3 - Date32
224    /// 4 - Timestamp(Second)
225    /// 5 - Timestamp(Millisecond)
226    /// 6 - Timestamp(Microsecond)
227    /// 7 - Timestamp(Nanosecond)
228    /// 8 - Utf8
229    packed: u16,
230}
231
232impl InferredDataType {
233    /// Returns the inferred data type
234    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, // Promote Int64 to Float64
240            b if b != 0 && (b & !0b11111000) == 0 => match b.leading_zeros() {
241                // Promote to highest precision temporal type
242                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    /// Updates the [`InferredDataType`] with the given string
254    fn update(&mut self, string: &str) {
255        self.packed |= if string.starts_with('"') {
256            1 << 8 // Utf8
257        } 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                // if overflow i64, fallback to utf8
260                1 << 8
261            } else {
262                1 << m
263            }
264        } else if string == "NaN" || string == "nan" || string == "inf" || string == "-inf" {
265            1 << 2 // Float64
266        } else {
267            1 << 8 // Utf8
268        }
269    }
270}
271
272/// The format specification for the CSV file
273#[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    /// Specify whether the CSV file has a header, defaults to `false`
288    ///
289    /// When `true`, the first row of the CSV file is treated as a header row
290    pub fn with_header(mut self, has_header: bool) -> Self {
291        self.header = has_header;
292        self
293    }
294
295    /// Specify whether to validate the CSV header against the schema, defaults to `false`
296    ///
297    /// When `true`, the first row gets validated against the schema before any data is read
298    ///
299    /// Only applies when [`Self::with_header`] is set to `true`
300    pub fn with_header_validation(mut self, validate_header: bool) -> Self {
301        self.header_validation = validate_header;
302        self
303    }
304
305    /// Specify a custom delimiter character, defaults to comma `','`
306    pub fn with_delimiter(mut self, delimiter: u8) -> Self {
307        self.delimiter = Some(delimiter);
308        self
309    }
310
311    /// Specify an escape character, defaults to `None`
312    pub fn with_escape(mut self, escape: u8) -> Self {
313        self.escape = Some(escape);
314        self
315    }
316
317    /// Specify a custom quote character, defaults to double quote `'"'`
318    pub fn with_quote(mut self, quote: u8) -> Self {
319        self.quote = Some(quote);
320        self
321    }
322
323    /// Specify a custom terminator character, defaults to CRLF
324    pub fn with_terminator(mut self, terminator: u8) -> Self {
325        self.terminator = Some(terminator);
326        self
327    }
328
329    /// Specify a comment character, defaults to `None`
330    ///
331    /// Lines starting with this character will be ignored
332    pub fn with_comment(mut self, comment: u8) -> Self {
333        self.comment = Some(comment);
334        self
335    }
336
337    /// Provide a regex to match null values, defaults to `^$`
338    pub fn with_null_regex(mut self, null_regex: Regex) -> Self {
339        self.null_regex = NullRegex(Some(null_regex));
340        self
341    }
342
343    /// Whether to allow truncated rows when parsing.
344    ///
345    /// By default this is set to `false` and will error if the CSV rows have different lengths.
346    /// When set to true then it will allow records with less than the expected number of columns
347    /// and fill the missing columns with nulls. If the record's schema is not nullable, then it
348    /// will still return an error.
349    pub fn with_truncated_rows(mut self, allow: bool) -> Self {
350        self.truncated_rows = allow;
351        self
352    }
353
354    /// Infer format settings from the CSV records in `reader`
355    ///
356    /// This currently infers whether the first record is a header. Up to
357    /// `max_records` records after the first record are inspected; if `None`, all
358    /// records are read. Detection is conservative and returns no header when the
359    /// sampled records do not provide type evidence. Returns the updated format
360    /// and the number of records read, including the first header candidate.
361    ///
362    /// # Example
363    ///
364    /// ```
365    /// use arrow_csv::reader::Format;
366    /// use std::io::Cursor;
367    ///
368    /// let csv = "name,count\nalice,1\nbob,2\n";
369    /// let (format, format_records_read) =
370    ///     Format::default().infer_format(Cursor::new(csv), Some(10))?;
371    /// let (schema, records_read) = format.infer_schema(Cursor::new(csv), None)?;
372    ///
373    /// assert_eq!(schema.field(0).name(), "name");
374    /// assert_eq!(format_records_read, 3);
375    /// assert_eq!(records_read, 2);
376    /// # Ok::<_, arrow_schema::ArrowError>(())
377    /// ```
378    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    /// Infer whether the first CSV record is a header
389    ///
390    /// Inspects up to `max_records` records after the first record. Returns `true`
391    /// when a value in the first record is text while the remaining values in the
392    /// same column have a consistent non-text type.
393    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                // Numeric-looking values (e.g. +1) are not header evidence, even
438                // when ordinary schema inference conservatively treats them as text.
439                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    /// Infer schema of CSV records from the provided `reader`
447    ///
448    /// If `max_records` is `None`, all records will be read, otherwise up to `max_records`
449    /// records are read to infer the schema
450    ///
451    /// Returns inferred schema and number of records read
452    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        // get or create header names
460        // when has_header is false, creates default column names with column_ prefix
461        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        // keep track of inferred field types
473        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            // Note since we may be looking at a sample of the data, we make the safe assumption that
486            // they could be nullable
487            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        // build schema from inference results
497        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    /// Build a [`csv::Reader`] for this [`Format`]
507    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    /// Build a [`csv_core::Reader`] for this [`Format`]
529    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
547/// Infer schema from a list of CSV files by reading through first n records
548/// with `max_read_records` controlling the maximum number of records to read.
549///
550/// Files will be read in the given order until n records have been reached.
551///
552/// If `max_read_records` is not set, all files will be read fully to infer the schema.
553pub 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
583// optional bounds of the reader, of the form (min line, max line).
584type Bounds = Option<(usize, usize)>;
585
586/// CSV file reader using [`std::io::BufReader`]
587///
588/// See [`ReaderBuilder`] to construct a CSV reader with options and  the
589/// [module-level documentation](crate::reader) for more details and examples
590pub type Reader<R> = BufReader<StdBufReader<R>>;
591
592/// CSV file reader implementation. See [`Reader`] for usage
593///
594/// Despite having the same name as [`std::io::BufReader`, this structure does
595/// not buffer reads itself
596pub struct BufReader<R> {
597    /// File reader
598    reader: R,
599    /// The decoder
600    decoder: Decoder,
601    /// Schema of the record batches produced by this reader
602    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    /// The number of rows padded because they had fewer fields than the schema
618    ///
619    /// Always 0 unless [`ReaderBuilder::with_truncated_rows`] was set to `true`.
620    ///
621    /// The count is cumulative over the lifetime of this reader, so reading it
622    /// between batches yields a running total of the rows read so far, and reading it
623    /// once the reader is exhausted yields the total for the whole input. Rows that
624    /// are skipped rather than read into a batch, such as a header row or rows before
625    /// the start bound, do not contribute.
626    ///
627    /// A padded row is indistinguishable from a row with genuinely empty trailing
628    /// fields once it has been read, so this counter is the only way to tell the two
629    /// apart.
630    ///
631    /// ```
632    /// # use std::io::Cursor;
633    /// # use std::sync::Arc;
634    /// # use arrow_csv::ReaderBuilder;
635    /// # use arrow_schema::{DataType, Field, Schema};
636    /// #
637    /// let schema = Arc::new(Schema::new(vec![
638    ///     Field::new("a", DataType::Int32, true),
639    ///     Field::new("b", DataType::Int32, true),
640    /// ]));
641    ///
642    /// let mut reader = ReaderBuilder::new(schema)
643    ///     .with_truncated_rows(true)
644    ///     .build(Cursor::new("1,2\n3\n"))
645    ///     .unwrap();
646    ///
647    /// let batches = reader.by_ref().collect::<Result<Vec<_>, _>>().unwrap();
648    /// assert_eq!(batches[0].num_rows(), 2);
649    /// assert_eq!(reader.truncated_row_count(), 1);
650    /// ```
651    pub fn truncated_row_count(&self) -> usize {
652        self.decoder.truncated_row_count()
653    }
654}
655
656impl<R: Read> Reader<R> {
657    /// Returns the schema of the reader, useful for getting the schema without reading
658    /// record batches
659    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            // Yield if decoded no bytes or the decoder is full
671            //
672            // The capacity check avoids looping around and potentially
673            // blocking reading data in fill_buf that isn't needed
674            // to flush the next batch
675            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/// A push-based interface for decoding CSV data from an arbitrary byte stream
699///
700/// See [`Reader`] for a higher-level interface for interface with [`Read`]
701///
702/// The push-based interface facilitates integration with sources that yield arbitrarily
703/// delimited bytes ranges, such as [`BufRead`], or a chunked byte stream received from
704/// object storage
705///
706/// ```
707/// # use std::io::BufRead;
708/// # use arrow_array::RecordBatch;
709/// # use arrow_csv::ReaderBuilder;
710/// # use arrow_schema::{ArrowError, SchemaRef};
711/// #
712/// fn read_from_csv<R: BufRead>(
713///     mut reader: R,
714///     schema: SchemaRef,
715///     batch_size: usize,
716/// ) -> Result<impl Iterator<Item = Result<RecordBatch, ArrowError>>, ArrowError> {
717///     let mut decoder = ReaderBuilder::new(schema)
718///         .with_batch_size(batch_size)
719///         .build_decoder();
720///
721///     let mut next = move || {
722///         loop {
723///             let buf = reader.fill_buf()?;
724///             let decoded = decoder.decode(buf)?;
725///             if decoded == 0 {
726///                 break;
727///             }
728///
729///             // Consume the number of bytes read
730///             reader.consume(decoded);
731///         }
732///         decoder.flush()
733///     };
734///     Ok(std::iter::from_fn(move || next().transpose()))
735/// }
736/// ```
737#[derive(Debug)]
738pub struct Decoder {
739    /// Explicit schema for the CSV file
740    schema: SchemaRef,
741
742    /// Optional projection for which columns to load (zero-based column indices)
743    projection: Option<Vec<usize>>,
744
745    /// Number of records per batch
746    batch_size: usize,
747
748    /// Rows to skip
749    to_skip: usize,
750
751    /// Whether to validate the first skipped row against the schema
752    header_validation: bool,
753
754    /// Current line number
755    line_number: usize,
756
757    /// End line number
758    end: usize,
759
760    /// A decoder for [`StringRecords`]
761    record_decoder: RecordDecoder,
762
763    /// Check if the string matches this pattern for `NULL`.
764    null_regex: NullRegex,
765}
766
767impl Decoder {
768    /// Decode records from `buf` returning the number of bytes read
769    ///
770    /// This method returns once `batch_size` objects have been parsed since the
771    /// last call to [`Self::flush`], or `buf` is exhausted. Any remaining bytes
772    /// should be included in the next call to [`Self::decode`]
773    ///
774    /// There is no requirement that `buf` contains a whole number of records, facilitating
775    /// integration with arbitrary byte streams, such as that yielded by [`BufRead`] or
776    /// network sources such as object storage
777    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            // Skip in units of `to_read` to avoid over-allocating buffers
794            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    /// Flushes the currently buffered data to a [`RecordBatch`]
807    ///
808    /// This should only be called after [`Self::decode`] has returned `Ok(0)`,
809    /// otherwise may return an error if part way through decoding a record
810    ///
811    /// Returns `Ok(None)` if no buffered data
812    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    /// Returns the number of records that can be read before requiring a call to [`Self::flush`]
830    pub fn capacity(&self) -> usize {
831        self.batch_size - self.record_decoder.len()
832    }
833
834    /// The number of rows padded because they had fewer fields than the schema
835    ///
836    /// Always 0 unless [`ReaderBuilder::with_truncated_rows`] was set to `true`.
837    ///
838    /// The count is cumulative over the lifetime of this decoder and is not reset by
839    /// [`Self::flush`], so reading it between batches yields a running total of the
840    /// rows decoded so far, and reading it once the input is exhausted yields the
841    /// total for the whole stream. Rows that are skipped rather than decoded into a
842    /// batch, such as a header row or rows before the start bound, do not contribute.
843    ///
844    /// A padded row is indistinguishable from a row with genuinely empty trailing
845    /// fields once it has been decoded, so this counter is the only way to tell the
846    /// two apart.
847    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
870/// Parses a slice of [`StringRecords`] into a [RecordBatch]
871fn 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
1131// parse the column string to an Arrow Array
1132fn 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            // append null
1145            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
1165// parses a specific column (col_idx) into an Arrow Array.
1166fn 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                    // TODO: we should surface the underlying error here.
1184                    "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
1253// parses a specific column (col_idx) into an Arrow Array.
1254fn 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                    // TODO: we should surface the underlying error here.
1272                    "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/// Builder for CSV [`Reader`]s
1286#[derive(Debug)]
1287pub struct ReaderBuilder {
1288    /// Schema of the CSV file
1289    schema: SchemaRef,
1290    /// Format of the CSV file
1291    format: Format,
1292    /// Batch size (number of records to load each time)
1293    ///
1294    /// The default batch size when using the `ReaderBuilder` is 1024 records
1295    batch_size: usize,
1296    /// The bounds over which to scan the reader. `None` starts from 0 and runs until EOF.
1297    bounds: Bounds,
1298    /// Optional projection for which columns to load (zero-based column indices)
1299    projection: Option<Vec<usize>>,
1300}
1301
1302impl ReaderBuilder {
1303    /// Create a new builder for configuring [`Reader`] CSV parsing options.
1304    ///
1305    /// To convert a builder into a reader, call [`ReaderBuilder::build`]. See
1306    /// the [module-level documentation](crate::reader) for more details and examples.
1307    ///
1308    /// # Example
1309    ///
1310    /// ```
1311    /// # use arrow_csv::{Reader, ReaderBuilder};
1312    /// # use std::fs::File;
1313    /// # use std::io::Seek;
1314    /// # use std::sync::Arc;
1315    /// # use arrow_csv::reader::Format;
1316    /// #
1317    /// let mut file = File::open("test/data/uk_cities_with_headers.csv").unwrap();
1318    /// // Infer the schema with the first 100 records
1319    /// let (schema, _) = Format::default().infer_schema(&mut file, Some(100)).unwrap();
1320    /// file.rewind().unwrap();
1321    ///
1322    /// // create a builder
1323    /// ReaderBuilder::new(Arc::new(schema)).build(file).unwrap();
1324    /// ```
1325    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    /// Set whether the CSV file has a header
1336    pub fn with_header(mut self, has_header: bool) -> Self {
1337        self.format.header = has_header;
1338        self
1339    }
1340
1341    /// Set whether to validate the CSV header against the schema
1342    ///
1343    /// This option only applies when [`Self::with_header`] is set to `true`, and defaults to `false`
1344    pub fn with_header_validation(mut self, validate_header: bool) -> Self {
1345        self.format.header_validation = validate_header;
1346        self
1347    }
1348
1349    /// Overrides the [Format] of this [ReaderBuilder]
1350    pub fn with_format(mut self, format: Format) -> Self {
1351        self.format = format;
1352        self
1353    }
1354
1355    /// Set the CSV file's column delimiter as a byte character
1356    pub fn with_delimiter(mut self, delimiter: u8) -> Self {
1357        self.format.delimiter = Some(delimiter);
1358        self
1359    }
1360
1361    /// Set the given character as the CSV file's escape character
1362    pub fn with_escape(mut self, escape: u8) -> Self {
1363        self.format.escape = Some(escape);
1364        self
1365    }
1366
1367    /// Set the given character as the CSV file's quote character, by default it is double quote
1368    pub fn with_quote(mut self, quote: u8) -> Self {
1369        self.format.quote = Some(quote);
1370        self
1371    }
1372
1373    /// Provide a custom terminator character, defaults to CRLF
1374    pub fn with_terminator(mut self, terminator: u8) -> Self {
1375        self.format.terminator = Some(terminator);
1376        self
1377    }
1378
1379    /// Provide a comment character, lines starting with this character will be ignored
1380    pub fn with_comment(mut self, comment: u8) -> Self {
1381        self.format.comment = Some(comment);
1382        self
1383    }
1384
1385    /// Provide a regex to match null values, defaults to `^$`
1386    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    /// Set the batch size (number of records to load at one time)
1392    pub fn with_batch_size(mut self, batch_size: usize) -> Self {
1393        self.batch_size = batch_size;
1394        self
1395    }
1396
1397    /// Set the bounds over which to scan the reader.
1398    /// `start` and `end` are line numbers.
1399    pub fn with_bounds(mut self, start: usize, end: usize) -> Self {
1400        self.bounds = Some((start, end));
1401        self
1402    }
1403
1404    /// Set the reader's column projection
1405    pub fn with_projection(mut self, projection: Vec<usize>) -> Self {
1406        self.projection = Some(projection);
1407        self
1408    }
1409
1410    /// Whether to allow truncated rows when parsing.
1411    ///
1412    /// By default this is set to `false` and will error if the CSV rows have different lengths.
1413    /// When set to true then it will allow records with less than the expected number of columns
1414    /// and fill the missing columns with nulls. If the record's schema is not nullable, then it
1415    /// will still return an error.
1416    pub fn with_truncated_rows(mut self, allow: bool) -> Self {
1417        self.format.truncated_rows = allow;
1418        self
1419    }
1420
1421    /// Create a new `Reader` from a non-buffered reader
1422    ///
1423    /// If `R: BufRead` consider using [`Self::build_buffered`] to avoid unnecessary additional
1424    /// buffering, as internally this method wraps `reader` in [`std::io::BufReader`]
1425    pub fn build<R: Read>(self, reader: R) -> Result<Reader<R>, ArrowError> {
1426        self.build_buffered(StdBufReader::new(reader))
1427    }
1428
1429    /// Create a new `BufReader` from a buffered reader
1430    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    /// Builds a decoder that can be used to decode CSV from an arbitrary byte stream
1444    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            // Inferred numeric types must also be accepted by the CSV decoder.
1499            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        // access data from a primitive array
1523        let lat = batch.column(1).as_primitive::<Float64Type>();
1524        assert_eq!(57.653484, lat.value(0));
1525
1526        // access data from a string array (ListArray<u8>)
1527        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        // access data from a primitive array
1569        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        // Rounding half away from zero, surrounding whitespace, exponent
1607        // notation and negative scales are all accepted
1608        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        // Invalid and out-of-range values are errors, never panics
1631        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        // access data from a primitive array
1672        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)] // Takes too long
1835    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        // access data from a primitive array
1858        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        // access data from a string array (ListArray<u8>)
1866        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)] // Takes too long
1877    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        // csv field names should be 'column_{number}'
1886        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        // access data from a primitive array
1898        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        // access data from a string array (ListArray<u8>)
1906        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)] // Takes too long
1917    fn test_csv_builder_with_bounds() {
1918        let mut file = File::open("test/data/uk_cities.csv").unwrap();
1919
1920        // Set the bounds to the lines 0, 1 and 2.
1921        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        // access data from a string array (ListArray<u8>)
1930        let city = batch
1931            .column(0)
1932            .as_any()
1933            .downcast_ref::<StringArray>()
1934            .unwrap();
1935
1936        // The value on line 0 is within the bounds
1937        assert_eq!("Elgin, Scotland, the UK", city.value(0));
1938
1939        // The value on line 13 is outside of the bounds. Therefore
1940        // the call to .value() will panic.
1941        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)] // Takes too long
2157    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        // "nil"s should be NULL
2209        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)] // Takes too long
2218    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)] // Takes too long
2277    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)] // Takes too long
2314    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    /// Infer the data type of a record
2388    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)] // Takes too long
2396    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)] // Takes too long
2540    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(); // empty csv file should be skipped
2544        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        // reading csv2 will set c2 to optional
2550        writeln!(csv2, "c1,c2,c3,c4").unwrap();
2551        writeln!(csv2, "10,,3.14,true").unwrap();
2552        // reading csv4 will set c3 to optional
2553        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), // only csv1 and csv2 should be read
2565            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        // Encode the expected behavior of boolean parsing
2650        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'~'); // default is ", change to ~
2695
2696        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'\\'); // default is None, change to \
2727
2728        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'); // default is CRLF, change to LF
2759
2760        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        // Empty rows are skipped by the underlying csv parser
2945        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    /// Schema used by the `truncated_row_count` tests below
3031    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        // "Carol,35," has all three fields, the last one just happens to be empty, so it
3057        // parses to the same null as a padded row would. The count must not be inferred
3058        // from the nulls in the batch
3059        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        // The short row is an error rather than something to count
3100        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        // Six short rows read two at a time
3111        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        // A running total, not a per batch count
3127        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        // The header is one field short of the schema, so skipping it pads it. Skipped
3134        // rows never reach a batch and must not be counted
3135        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)] // Takes too long
3191    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        // Expect 4 calls to fill_buf
3342        // 1. Read first 3 rows
3343        // 2. Read final row
3344        // 3. Delimit and flush final row
3345        // 4. Iterator finished
3346        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)] // Takes too long
3352    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)] // Takes too long
3421    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        // Include line number in the error message to help locate and fix the issue
3433        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)] // Unsupported inline assembly
3537    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}