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#[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 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 match possibilities.len() {
310 1 => possibilities.iter().next().unwrap().clone(),
311 2 if possibilities.contains(&DataType::Int64)
312 && possibilities.contains(&DataType::Float64) =>
313 {
314 DataType::Float64
316 },
317 #[cfg(feature = "dtype-i128")]
318 2 if possibilities.contains(&DataType::Int64)
319 && possibilities.contains(&DataType::Int128) =>
320 {
321 DataType::Int128
323 },
324 #[cfg(feature = "dtype-i128")]
325 2 if possibilities.contains(&DataType::Int128)
326 && possibilities.contains(&DataType::Float64) =>
327 {
328 DataType::Float64
330 },
331 _ => DataType::String,
333 }
334}
335
336pub fn infer_field_schema(string: &str, try_parse_dates: bool, decimal_comma: bool) -> DataType {
338 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 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 assert_eq!(
441 infer_field_schema("9223372036854775807", false, false),
442 DataType::Int64,
443 );
444
445 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}