Skip to main content

polars_io/csv/read/
schema_inference.rs

1use polars_buffer::Buffer;
2#[cfg(any(
3    feature = "dtype-date",
4    feature = "dtype-datetime",
5    feature = "dtype-time"
6))]
7use polars_core::chunked_array::temporal::string::infer as date_infer;
8#[cfg(any(
9    feature = "dtype-date",
10    feature = "dtype-datetime",
11    feature = "dtype-time"
12))]
13use polars_core::chunked_array::temporal::string::patterns::Pattern;
14use polars_core::prelude::*;
15use polars_utils::format_pl_smallstr;
16
17use super::splitfields::SplitFields;
18use super::{CsvParseOptions, NullValues};
19use crate::utils::{BOOLEAN_RE, FLOAT_RE, FLOAT_RE_DECIMAL, INTEGER_RE};
20
21/// Low-level CSV schema inference function.
22///
23/// Use `read_until_start_and_infer_schema` instead.
24#[allow(clippy::too_many_arguments)]
25pub(super) fn infer_file_schema_impl(
26    header_line: &Option<Buffer<u8>>,
27    content_lines: &[Buffer<u8>],
28    infer_all_as_str: bool,
29    parse_options: &CsvParseOptions,
30    column_names_overwrite: Option<&[PlSmallStr]>,
31    schema_overwrite: Option<&Schema>,
32    ignore_extra_columns: bool,
33    insert_missing_columns: bool,
34) -> PolarsResult<Schema> {
35    let mut headers = if let Some(header_line) = header_line {
36        infer_headers(header_line, parse_options)?
37    } else {
38        Vec::with_capacity(8)
39    };
40
41    let extend_header_with_unknown_column = header_line.is_none();
42
43    let mut column_types = vec![PlIndexSet::<DataType>::with_capacity(4); headers.len()];
44    let mut nulls = vec![false; headers.len()];
45
46    for content_line in content_lines {
47        infer_types_from_line(
48            content_line,
49            infer_all_as_str,
50            &mut headers,
51            extend_header_with_unknown_column,
52            parse_options,
53            &mut column_types,
54            &mut nulls,
55        );
56    }
57
58    if let Some(column_names_overwrite) = column_names_overwrite {
59        let mut err_hint: String = String::new();
60
61        if column_names_overwrite.len() < headers.len() && !ignore_extra_columns {
62            let n = headers.len() - column_names_overwrite.len();
63            err_hint = format!("pass extra_columns='ignore' to ignore ({n}) extra columns.")
64        }
65
66        if column_names_overwrite.len() > headers.len() && !insert_missing_columns {
67            let n = column_names_overwrite.len() - headers.len();
68            err_hint = format!(
69                "pass missing_columns='insert' to create ({n}) missing columns with all-NULL values."
70            );
71        }
72
73        if !err_hint.is_empty() {
74            polars_bail!(
75                SchemaMismatch:
76                "provided `new_columns` does not match number of columns in file ({} != {} in file). \
77                Ensure the number of names match, or {err_hint}",
78                column_names_overwrite.len(),
79                headers.len(),
80            )
81        }
82
83        headers.truncate(column_names_overwrite.len());
84        column_types.truncate(column_names_overwrite.len());
85
86        for (i, name) in column_names_overwrite.iter().cloned().enumerate() {
87            if i < headers.len() {
88                headers[i] = name
89            } else {
90                headers.push(name)
91            }
92
93            if i >= column_types.len() {
94                column_types.push(PlIndexSet::from_iter(Some(DataType::Null)))
95            }
96        }
97    }
98
99    Ok(build_schema(&headers, &column_types, schema_overwrite))
100}
101
102fn infer_headers(
103    mut header_line: &[u8],
104    parse_options: &CsvParseOptions,
105) -> PolarsResult<Vec<PlSmallStr>> {
106    let len = header_line.len();
107
108    if header_line.last().copied() == Some(b'\r') {
109        header_line = &header_line[..len - 1];
110    }
111
112    let byterecord = SplitFields::new(
113        header_line,
114        parse_options.separator,
115        parse_options.quote_char,
116        parse_options.eol_char,
117    );
118
119    let headers = byterecord
120        .map(|(slice, needs_escaping)| {
121            let slice_escaped = if needs_escaping && (slice.len() >= 2) {
122                &slice[1..(slice.len() - 1)]
123            } else {
124                slice
125            };
126            String::from_utf8_lossy(slice_escaped)
127        })
128        .collect::<Vec<_>>();
129
130    let mut deduplicated_headers = PlIndexSet::with_capacity(headers.len());
131    let mut header_names = PlHashMap::with_capacity(headers.len());
132
133    for name in &headers {
134        let count = header_names.entry(name.as_ref()).or_insert(0usize);
135        let duplicated = *count != 0;
136        let deduplicated_name = if duplicated {
137            format_pl_smallstr!("{}_duplicated_{}", name, *count - 1)
138        } else {
139            PlSmallStr::from_str(name)
140        };
141
142        if !deduplicated_headers.insert(deduplicated_name.clone()) {
143            let (deduplicated_from, nth_duplicated) = if duplicated {
144                (name.as_ref(), 1 + *count)
145            } else {
146                let i = deduplicated_name.rfind("_duplicated_").unwrap();
147                (
148                    &deduplicated_name[..i],
149                    2 + deduplicated_name[i + 12..].parse::<usize>().unwrap(),
150                )
151            };
152
153            polars_bail!(
154                Duplicate:
155                "de-duplication of occurrence #{nth_duplicated} of column name '{deduplicated_from}' \
156                failed; the name '{deduplicated_name}' also exists in the file."
157            )
158        }
159
160        *count += 1;
161    }
162
163    Ok(Vec::from_iter(deduplicated_headers))
164}
165
166fn infer_types_from_line(
167    mut line: &[u8],
168    infer_all_as_str: bool,
169    headers: &mut Vec<PlSmallStr>,
170    extend_header_with_unknown_column: bool,
171    parse_options: &CsvParseOptions,
172    column_types: &mut Vec<PlIndexSet<DataType>>,
173    nulls: &mut Vec<bool>,
174) {
175    let line_len = line.len();
176    if line.last().copied() == Some(b'\r') {
177        line = &line[..line_len - 1];
178    }
179
180    let record = SplitFields::new(
181        line,
182        parse_options.separator,
183        parse_options.quote_char,
184        parse_options.eol_char,
185    );
186
187    for (i, (slice, needs_escaping)) in record.enumerate() {
188        if i >= headers.len() {
189            if extend_header_with_unknown_column {
190                headers.push(column_name(i));
191                column_types.push(Default::default());
192                nulls.push(false);
193            } else {
194                break;
195            }
196        }
197
198        if infer_all_as_str {
199            column_types[i].insert(DataType::String);
200            continue;
201        }
202
203        if slice.is_empty() {
204            nulls[i] = true;
205        } else {
206            let slice_escaped = if needs_escaping && (slice.len() >= 2) {
207                &slice[1..(slice.len() - 1)]
208            } else {
209                slice
210            };
211            let s = String::from_utf8_lossy(slice_escaped);
212            let dtype = match &parse_options.null_values {
213                None => Some(infer_field_schema(
214                    &s,
215                    parse_options.try_parse_dates,
216                    parse_options.decimal_comma,
217                )),
218                Some(NullValues::AllColumns(names)) => {
219                    if !names.iter().any(|nv| nv == s.as_ref()) {
220                        Some(infer_field_schema(
221                            &s,
222                            parse_options.try_parse_dates,
223                            parse_options.decimal_comma,
224                        ))
225                    } else {
226                        None
227                    }
228                },
229                Some(NullValues::AllColumnsSingle(name)) => {
230                    if s.as_ref() != name.as_str() {
231                        Some(infer_field_schema(
232                            &s,
233                            parse_options.try_parse_dates,
234                            parse_options.decimal_comma,
235                        ))
236                    } else {
237                        None
238                    }
239                },
240                Some(NullValues::Named(names)) => {
241                    let current_name = &headers[i];
242                    let null_name = &names.iter().find(|name| name.0 == current_name);
243
244                    if let Some(null_name) = null_name {
245                        if null_name.1.as_str() != s.as_ref() {
246                            Some(infer_field_schema(
247                                &s,
248                                parse_options.try_parse_dates,
249                                parse_options.decimal_comma,
250                            ))
251                        } else {
252                            None
253                        }
254                    } else {
255                        Some(infer_field_schema(
256                            &s,
257                            parse_options.try_parse_dates,
258                            parse_options.decimal_comma,
259                        ))
260                    }
261                },
262            };
263            if let Some(dtype) = dtype {
264                column_types[i].insert(dtype);
265            }
266        }
267    }
268}
269
270fn build_schema(
271    headers: &[PlSmallStr],
272    column_types: &[PlIndexSet<DataType>],
273    schema_overwrite: Option<&Schema>,
274) -> Schema {
275    assert!(headers.len() == column_types.len());
276
277    let get_schema_overwrite = |field_name| {
278        if let Some(schema_overwrite) = schema_overwrite {
279            // Apply schema_overwrite by column name only. Positional overrides are handled
280            // separately via dtype_overwrite.
281            if let Some((_, name, dtype)) = schema_overwrite.get_full(field_name) {
282                return Some((name.clone(), dtype.clone()));
283            }
284        }
285
286        None
287    };
288
289    Schema::from_iter(
290        headers
291            .iter()
292            .zip(column_types)
293            .map(|(field_name, type_possibilities)| {
294                let (name, dtype) = get_schema_overwrite(field_name).unwrap_or_else(|| {
295                    (
296                        field_name.clone(),
297                        finish_infer_field_schema(type_possibilities),
298                    )
299                });
300
301                Field::new(name, dtype)
302            }),
303    )
304}
305
306pub fn finish_infer_field_schema(possibilities: &PlIndexSet<DataType>) -> DataType {
307    // determine data type based on possible types
308    // if there are incompatible types, use DataType::String
309    match possibilities.len() {
310        1 => possibilities.iter().next().unwrap().clone(),
311        2 if possibilities.contains(&DataType::Int64)
312            && possibilities.contains(&DataType::Float64) =>
313        {
314            // we have an integer and double, fall down to double
315            DataType::Float64
316        },
317        #[cfg(feature = "dtype-i128")]
318        2 if possibilities.contains(&DataType::Int64)
319            && possibilities.contains(&DataType::Int128) =>
320        {
321            // all values fit within i128
322            DataType::Int128
323        },
324        #[cfg(feature = "dtype-i128")]
325        2 if possibilities.contains(&DataType::Int128)
326            && possibilities.contains(&DataType::Float64) =>
327        {
328            // fall down to double for mixed int128 and float
329            DataType::Float64
330        },
331        // default to String for conflicting datatypes (e.g bool and int)
332        _ => DataType::String,
333    }
334}
335
336/// Infer the data type of a record
337pub fn infer_field_schema(string: &str, try_parse_dates: bool, decimal_comma: bool) -> DataType {
338    // when quoting is enabled in the reader, these quotes aren't escaped, we default to
339    // String for them
340    let bytes = string.as_bytes();
341    if bytes.len() >= 2 && *bytes.first().unwrap() == b'"' && *bytes.last().unwrap() == b'"' {
342        if try_parse_dates {
343            #[cfg(any(
344                feature = "dtype-date",
345                feature = "dtype-datetime",
346                feature = "dtype-time"
347            ))]
348            {
349                match date_infer::infer_pattern_single(&string[1..string.len() - 1]) {
350                    Some(pattern_with_offset) => match pattern_with_offset {
351                        Pattern::DatetimeYMD | Pattern::DatetimeDMY => {
352                            DataType::Datetime(TimeUnit::Microseconds, None)
353                        },
354                        Pattern::DateYMD | Pattern::DateDMY => DataType::Date,
355                        Pattern::DatetimeYMDZ => {
356                            DataType::Datetime(TimeUnit::Microseconds, Some(TimeZone::UTC))
357                        },
358                        Pattern::Time => DataType::Time,
359                    },
360                    None => DataType::String,
361                }
362            }
363            #[cfg(not(any(
364                feature = "dtype-date",
365                feature = "dtype-datetime",
366                feature = "dtype-time"
367            )))]
368            {
369                panic!("activate one of {{'dtype-date', 'dtype-datetime', dtype-time'}} features")
370            }
371        } else {
372            DataType::String
373        }
374    }
375    // match regex in a particular order
376    else if BOOLEAN_RE.is_match(string) {
377        DataType::Boolean
378    } else if !decimal_comma && FLOAT_RE.is_match(string)
379        || decimal_comma && FLOAT_RE_DECIMAL.is_match(string)
380    {
381        DataType::Float64
382    } else if INTEGER_RE.is_match(string) {
383        if string.parse::<i64>().is_ok() {
384            DataType::Int64
385        } else {
386            #[cfg(feature = "dtype-i128")]
387            {
388                DataType::Int128
389            }
390            #[cfg(not(feature = "dtype-i128"))]
391            {
392                DataType::Int64
393            }
394        }
395    } else if try_parse_dates {
396        #[cfg(any(
397            feature = "dtype-date",
398            feature = "dtype-datetime",
399            feature = "dtype-time"
400        ))]
401        {
402            match date_infer::infer_pattern_single(string) {
403                Some(pattern_with_offset) => match pattern_with_offset {
404                    Pattern::DatetimeYMD | Pattern::DatetimeDMY => {
405                        DataType::Datetime(TimeUnit::Microseconds, None)
406                    },
407                    Pattern::DateYMD | Pattern::DateDMY => DataType::Date,
408                    Pattern::DatetimeYMDZ => {
409                        DataType::Datetime(TimeUnit::Microseconds, Some(TimeZone::UTC))
410                    },
411                    Pattern::Time => DataType::Time,
412                },
413                None => DataType::String,
414            }
415        }
416        #[cfg(not(any(
417            feature = "dtype-date",
418            feature = "dtype-datetime",
419            feature = "dtype-time"
420        )))]
421        {
422            panic!("activate one of {{'dtype-date', 'dtype-datetime', dtype-time'}} features")
423        }
424    } else {
425        DataType::String
426    }
427}
428
429fn column_name(i: usize) -> PlSmallStr {
430    format_pl_smallstr!("column_{}", i)
431}
432
433#[cfg(test)]
434mod tests {
435    use super::*;
436
437    #[test]
438    fn test_infer_field_schema_i64_overflow() {
439        // Values within i64 range should infer as Int64.
440        assert_eq!(
441            infer_field_schema("9223372036854775807", false, false),
442            DataType::Int64,
443        );
444
445        // Values exceeding i64::MAX should infer as Int128 when the feature is enabled,
446        // otherwise as String.
447        let large = "12345678901234567890";
448        #[cfg(feature = "dtype-i128")]
449        assert_eq!(infer_field_schema(large, false, false), DataType::Int128,);
450        #[cfg(not(feature = "dtype-i128"))]
451        assert_eq!(infer_field_schema(large, false, false), DataType::Int64,);
452    }
453
454    #[test]
455    #[cfg(feature = "dtype-i128")]
456    fn test_finish_infer_field_schema_i64_and_i128() {
457        let mut possibilities = PlIndexSet::new();
458        possibilities.insert(DataType::Int64);
459        possibilities.insert(DataType::Int128);
460        assert_eq!(finish_infer_field_schema(&possibilities), DataType::Int128);
461    }
462}