Skip to main content

operator/statement/
copy_table_from.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::collections::{BTreeSet, HashMap, HashSet};
16use std::future::Future;
17use std::pin::Pin;
18use std::sync::Arc;
19use std::task::{Context, Poll};
20
21use client::{Output, OutputData, OutputMeta};
22use common_base::readable_size::ReadableSize;
23use common_datasource::file_format::csv::{
24    CsvFormat, is_skippable_arrow_error, tolerant_csv_stream,
25};
26use common_datasource::file_format::json::JsonFormat;
27use common_datasource::file_format::orc::{ReaderAdapter, infer_orc_schema, new_orc_stream_reader};
28use common_datasource::file_format::{FileFormat, Format, file_to_stream};
29use common_datasource::lister::{Lister, Source};
30use common_datasource::object_store::{LocalFileAccess, build_backend_with_path};
31use common_query::{OutputCost, OutputRows};
32use common_recordbatch::DfSendableRecordBatchStream;
33use common_recordbatch::adapter::RecordBatchStreamTypeAdapter;
34use common_telemetry::{debug, tracing};
35use datafusion::datasource::physical_plan::{CsvSource, FileSource, JsonSource};
36use datafusion::parquet::arrow::ParquetRecordBatchStreamBuilder;
37use datafusion::parquet::arrow::arrow_reader::ArrowReaderMetadata;
38use datafusion_common::DataFusionError;
39use datafusion_common::arrow::error::ArrowError;
40use datafusion_common::config::CsvOptions;
41use datafusion_expr::Expr;
42use datatypes::arrow::compute::can_cast_types;
43use datatypes::arrow::datatypes::{DataType as ArrowDataType, Schema, SchemaRef};
44use datatypes::arrow::record_batch::RecordBatch;
45use datatypes::vectors::Helper;
46use futures_util::StreamExt;
47use object_store::{EntryMode, ObjectStore};
48use regex::Regex;
49use session::context::QueryContextRef;
50use snafu::{ResultExt, ensure};
51use table::requests::{CopyTableRequest, InsertRequest};
52use table::table_reference::TableReference;
53use tokio_util::compat::FuturesAsyncReadCompatExt;
54
55use crate::error::{self, IntoVectorsSnafu, Result};
56use crate::statement::StatementExecutor;
57
58const DEFAULT_BATCH_SIZE: usize = 8192;
59const DEFAULT_READ_BUFFER: usize = 256 * 1024;
60
61enum FileMetadata {
62    Parquet {
63        schema: SchemaRef,
64        metadata: ArrowReaderMetadata,
65        path: String,
66    },
67    Orc {
68        schema: SchemaRef,
69        path: String,
70    },
71    Json {
72        schema: SchemaRef,
73        format: JsonFormat,
74        path: String,
75    },
76    Csv {
77        schema: SchemaRef,
78        format: CsvFormat,
79        path: String,
80    },
81}
82
83impl FileMetadata {
84    /// Returns the [SchemaRef]
85    pub fn schema(&self) -> &SchemaRef {
86        match self {
87            FileMetadata::Parquet { schema, .. } => schema,
88            FileMetadata::Orc { schema, .. } => schema,
89            FileMetadata::Json { schema, .. } => schema,
90            FileMetadata::Csv { schema, .. } => schema,
91        }
92    }
93}
94
95async fn list_copy_from_paths(
96    req: &CopyTableRequest,
97    local_file_access: &LocalFileAccess,
98) -> Result<(ObjectStore, Vec<String>)> {
99    let backend = build_backend_with_path(&req.location, &req.connection, local_file_access)
100        .await
101        .context(error::BuildBackendSnafu)?;
102    let regex = req
103        .pattern
104        .as_ref()
105        .map(|x| Regex::new(x))
106        .transpose()
107        .context(error::BuildRegexSnafu)?;
108
109    // Listing a known file's parent for every table makes COPY DATABASE do
110    // quadratic directory work.
111    if let Some(filename) = &backend.object_path {
112        let is_file = backend
113            .is_file(filename)
114            .await
115            .with_context(|_| common_datasource::error::ListObjectsSnafu {
116                path: req.location.clone(),
117            })
118            .context(error::ListObjectsSnafu)?;
119        let paths = if is_file {
120            vec![filename.clone()]
121        } else {
122            vec![]
123        };
124        return Ok((backend.object_store, paths));
125    }
126
127    let lister = Lister::new(
128        backend.object_store.clone(),
129        Source::Dir,
130        req.location.clone(),
131        regex,
132    );
133
134    let entries = lister.list().await.context(error::ListObjectsSnafu)?;
135    debug!(
136        "Copy from location: {:?}, entries: {entries:?}",
137        req.location
138    );
139    let paths = entries
140        .into_iter()
141        .filter(|entry| entry.metadata().mode() == EntryMode::FILE)
142        .map(|entry| entry.path().to_string())
143        .collect();
144    Ok((backend.object_store, paths))
145}
146
147impl StatementExecutor {
148    async fn collect_metadata(
149        &self,
150        object_store: &ObjectStore,
151        format: Format,
152        path: String,
153    ) -> Result<FileMetadata> {
154        match format {
155            Format::Csv(format) => Ok(FileMetadata::Csv {
156                schema: Arc::new(
157                    format
158                        .infer_schema(object_store, &path)
159                        .await
160                        .context(error::InferSchemaSnafu { path: &path })?,
161                ),
162                format,
163                path,
164            }),
165            Format::Json(format) => Ok(FileMetadata::Json {
166                schema: Arc::new(
167                    format
168                        .infer_schema(object_store, &path)
169                        .await
170                        .context(error::InferSchemaSnafu { path: &path })?,
171                ),
172                format,
173                path,
174            }),
175            Format::Parquet(_) => {
176                let meta = object_store
177                    .stat(&path)
178                    .await
179                    .context(error::ReadObjectSnafu { path: &path })?;
180                let mut reader = object_store
181                    .reader(&path)
182                    .await
183                    .context(error::ReadObjectSnafu { path: &path })?
184                    .into_futures_async_read(0..meta.content_length())
185                    .await
186                    .context(error::ReadObjectSnafu { path: &path })?
187                    .compat();
188                let metadata = ArrowReaderMetadata::load_async(&mut reader, Default::default())
189                    .await
190                    .context(error::ReadParquetMetadataSnafu)?;
191
192                Ok(FileMetadata::Parquet {
193                    schema: metadata.schema().clone(),
194                    metadata,
195                    path,
196                })
197            }
198            Format::Orc(_) => {
199                let meta = object_store
200                    .stat(&path)
201                    .await
202                    .context(error::ReadObjectSnafu { path: &path })?;
203
204                let reader = object_store
205                    .reader(&path)
206                    .await
207                    .context(error::ReadObjectSnafu { path: &path })?;
208
209                let schema = infer_orc_schema(ReaderAdapter::new(reader, meta.content_length()))
210                    .await
211                    .context(error::ReadOrcSnafu)?;
212
213                Ok(FileMetadata::Orc {
214                    schema: Arc::new(schema),
215                    path,
216                })
217            }
218        }
219    }
220
221    async fn build_read_stream(
222        &self,
223        compat_schema: SchemaRef,
224        object_store: &ObjectStore,
225        file_metadata: &FileMetadata,
226        projection: Vec<usize>,
227        filters: Vec<Expr>,
228    ) -> Result<DfSendableRecordBatchStream> {
229        match file_metadata {
230            FileMetadata::Csv {
231                format,
232                path,
233                schema,
234            } => {
235                let output_schema = Arc::new(
236                    compat_schema
237                        .project(&projection)
238                        .context(error::ProjectSchemaSnafu)?,
239                );
240
241                let options = CsvOptions::default()
242                    .with_has_header(format.has_header)
243                    .with_delimiter(format.delimiter);
244                let csv_source = CsvSource::new(schema.clone())
245                    .with_csv_options(options)
246                    .with_batch_size(DEFAULT_BATCH_SIZE);
247                let stream = if format.skip_bad_records {
248                    let reader_schema =
249                        csv_reader_schema_for_skip_bad_records(schema, &compat_schema);
250                    tolerant_csv_stream(
251                        object_store,
252                        path,
253                        Arc::new(reader_schema),
254                        projection.clone(),
255                        format,
256                    )
257                    .await
258                    .context(error::BuildFileStreamSnafu)?
259                } else {
260                    file_to_stream(
261                        object_store,
262                        path,
263                        csv_source,
264                        Some(projection),
265                        format.compression_type,
266                    )
267                    .await
268                    .context(error::BuildFileStreamSnafu)?
269                };
270
271                let stream = Box::pin(
272                    // The projection is already applied in the CSV reader when we created the stream,
273                    // so we pass None here to avoid double projection which would cause schema mismatch errors.
274                    RecordBatchStreamTypeAdapter::new(output_schema, stream, None)
275                        .with_filter(filters)
276                        .context(error::PhysicalExprSnafu)?,
277                );
278                if format.skip_bad_records {
279                    Ok(Box::pin(SkipBadRecordsStream::new(stream, path)))
280                } else {
281                    Ok(stream)
282                }
283            }
284            FileMetadata::Json {
285                path,
286                format,
287                schema,
288            } => {
289                let output_schema = Arc::new(
290                    compat_schema
291                        .project(&projection)
292                        .context(error::ProjectSchemaSnafu)?,
293                );
294
295                let json_source =
296                    JsonSource::new(schema.clone()).with_batch_size(DEFAULT_BATCH_SIZE);
297                let stream = file_to_stream(
298                    object_store,
299                    path,
300                    json_source,
301                    Some(projection),
302                    format.compression_type,
303                )
304                .await
305                .context(error::BuildFileStreamSnafu)?;
306
307                Ok(Box::pin(
308                    // The projection is already applied in the JSON reader when we created the stream,
309                    // so we pass None here to avoid double projection which would cause schema mismatch errors.
310                    RecordBatchStreamTypeAdapter::new(output_schema, stream, None)
311                        .with_filter(filters)
312                        .context(error::PhysicalExprSnafu)?,
313                ))
314            }
315            FileMetadata::Parquet { metadata, path, .. } => {
316                let meta = object_store
317                    .stat(path)
318                    .await
319                    .context(error::ReadObjectSnafu { path })?;
320                let reader = object_store
321                    .reader_with(path)
322                    .chunk(DEFAULT_READ_BUFFER)
323                    .await
324                    .context(error::ReadObjectSnafu { path })?
325                    .into_futures_async_read(0..meta.content_length())
326                    .await
327                    .context(error::ReadObjectSnafu { path })?
328                    .compat();
329                let builder =
330                    ParquetRecordBatchStreamBuilder::new_with_metadata(reader, metadata.clone());
331                let stream = builder
332                    .build()
333                    .context(error::BuildParquetRecordBatchStreamSnafu)?;
334
335                let output_schema = Arc::new(
336                    compat_schema
337                        .project(&projection)
338                        .context(error::ProjectSchemaSnafu)?,
339                );
340                Ok(Box::pin(
341                    RecordBatchStreamTypeAdapter::new(output_schema, stream, Some(projection))
342                        .with_filter(filters)
343                        .context(error::PhysicalExprSnafu)?,
344                ))
345            }
346            FileMetadata::Orc { path, .. } => {
347                let meta = object_store
348                    .stat(path)
349                    .await
350                    .context(error::ReadObjectSnafu { path })?;
351
352                let reader = object_store
353                    .reader_with(path)
354                    .chunk(DEFAULT_READ_BUFFER)
355                    .await
356                    .context(error::ReadObjectSnafu { path })?;
357                let stream =
358                    new_orc_stream_reader(ReaderAdapter::new(reader, meta.content_length()))
359                        .await
360                        .context(error::ReadOrcSnafu)?;
361
362                let output_schema = Arc::new(
363                    compat_schema
364                        .project(&projection)
365                        .context(error::ProjectSchemaSnafu)?,
366                );
367
368                Ok(Box::pin(
369                    RecordBatchStreamTypeAdapter::new(output_schema, stream, Some(projection))
370                        .with_filter(filters)
371                        .context(error::PhysicalExprSnafu)?,
372                ))
373            }
374        }
375    }
376
377    /// Imports one indexed stream with the ordinary COPY schema mapping and inserter.
378    pub(crate) async fn copy_indexed_parquet<
379        R: datafusion::parquet::arrow::async_reader::AsyncFileReader + Send + Unpin + 'static,
380    >(
381        &self,
382        reader: R,
383        table: table::TableRef,
384        expected_rows: u64,
385        pending: &mut crate::statement::import_packed::PendingPackedInserts,
386        cancellation: &tokio_util::sync::CancellationToken,
387        query_ctx: QueryContextRef,
388    ) -> Result<()> {
389        let builder = tokio::select! {
390            _ = cancellation.cancelled() => return error::PackedImportCancelledSnafu.fail(),
391            result = ParquetRecordBatchStreamBuilder::new(reader) => result.context(error::ReadParquetMetadataSnafu)?,
392        };
393        let table_schema = table.schema().arrow_schema().clone();
394        let (file_projection, table_projection, _) =
395            generated_schema_projection_and_compatible_file_schema(builder.schema(), &table_schema);
396        let file_schema = Arc::new(
397            builder
398                .schema()
399                .project(&file_projection)
400                .context(error::ProjectSchemaSnafu)?,
401        );
402        let target_schema = Arc::new(
403            table_schema
404                .project(&table_projection)
405                .context(error::ProjectSchemaSnafu)?,
406        );
407        ensure_schema_compatible(&file_schema, &target_schema)?;
408        let stream = builder
409            .with_batch_size(DEFAULT_BATCH_SIZE)
410            .build()
411            .context(error::BuildParquetRecordBatchStreamSnafu)?;
412        let mut stream = Box::pin(RecordBatchStreamTypeAdapter::new(
413            target_schema.clone(),
414            stream,
415            Some(file_projection),
416        ));
417        let info = table.table_info();
418        let mut rows = 0u64;
419        loop {
420            pending.before_decode().await?;
421            if cancellation.is_cancelled() {
422                return error::PackedImportCancelledSnafu.fail();
423            }
424            let batch = tokio::select! {
425                _ = cancellation.cancelled() => return error::PackedImportCancelledSnafu.fail(),
426                batch = stream.next() => batch,
427            };
428            let Some(batch) = batch else {
429                break;
430            };
431            let batch = batch.context(error::ReadDfRecordBatchSnafu)?;
432            rows += batch.num_rows() as u64;
433            if rows > expected_rows {
434                return error::InvalidCopyParameterSnafu {
435                    key: "row_count",
436                    value: expected_rows.to_string(),
437                }
438                .fail();
439            }
440            let vectors = Helper::try_into_vectors(batch.columns()).context(IntoVectorsSnafu)?;
441            let retained_bytes = vectors.iter().map(|v| v.memory_size()).sum::<usize>();
442            let columns_values = target_schema
443                .fields()
444                .iter()
445                .map(|f| f.name().clone())
446                .zip(vectors)
447                .collect();
448            let bytes = retained_bytes.max(batch.get_array_memory_size());
449            let inserter = self.inserter.clone();
450            let request = InsertRequest {
451                catalog_name: info.catalog_name.clone(),
452                schema_name: info.schema_name.clone(),
453                table_name: info.name.clone(),
454                columns_values,
455                skip_wal: query_ctx.skip_wal(),
456            };
457            let ctx = query_ctx.clone();
458            pending
459                .admit(
460                    bytes,
461                    async move { inserter.handle_table_insert(request, ctx).await },
462                    cancellation,
463                )
464                .await?;
465        }
466        if rows != expected_rows {
467            return error::InvalidCopyParameterSnafu {
468                key: "row_count",
469                value: expected_rows.to_string(),
470            }
471            .fail();
472        }
473        Ok(())
474    }
475
476    #[tracing::instrument(skip_all)]
477    pub async fn copy_table_from(
478        &self,
479        req: CopyTableRequest,
480        query_ctx: QueryContextRef,
481    ) -> Result<Output> {
482        let table_ref = TableReference {
483            catalog: &req.catalog_name,
484            schema: &req.schema_name,
485            table: &req.table_name,
486        };
487        let table = self.get_table(&table_ref).await?;
488        let format = Format::try_from(&req.with).context(error::ParseFileFormatSnafu)?;
489        let (object_store, paths) = list_copy_from_paths(&req, &self.local_file_access).await?;
490        let mut files = Vec::with_capacity(paths.len());
491        let table_schema = table.schema().arrow_schema().clone();
492        let filters = table
493            .schema()
494            .timestamp_column()
495            .and_then(|c| {
496                common_query::logical_plan::build_same_type_ts_filter(c, req.timestamp_range)
497            })
498            .into_iter()
499            .collect::<Vec<_>>();
500
501        for path in paths {
502            let file_metadata = self
503                .collect_metadata(&object_store, format.clone(), path)
504                .await?;
505
506            validate_csv_headers_if_required(&file_metadata, &table_schema)?;
507            let schema_mapping = copy_from_schema_mapping(&file_metadata, &table_schema);
508            let projected_file_schema = Arc::new(
509                file_metadata
510                    .schema()
511                    .project(&schema_mapping.file_projection)
512                    .context(error::ProjectSchemaSnafu)?,
513            );
514            let projected_table_schema = Arc::new(
515                table_schema
516                    .project(&schema_mapping.table_projection)
517                    .context(error::ProjectSchemaSnafu)?,
518            );
519            ensure_schema_compatible(&projected_file_schema, &projected_table_schema)?;
520
521            files.push((
522                Arc::new(schema_mapping.compat_file_schema),
523                schema_mapping.file_projection,
524                projected_table_schema,
525                file_metadata,
526            ))
527        }
528
529        let mut rows_inserted = 0;
530        let mut insert_cost = 0;
531        let max_insert_rows = req
532            .limit
533            .map(|n| {
534                usize::try_from(n).map_err(|_| {
535                    error::InvalidCopyParameterSnafu {
536                        key: "limit".to_string(),
537                        value: n.to_string(),
538                    }
539                    .build()
540                })
541            })
542            .transpose()?;
543        if max_insert_rows == Some(0) {
544            return Ok(gen_insert_output(rows_inserted, insert_cost));
545        }
546
547        let mut accepted_rows = 0;
548        for (compat_schema, file_schema_projection, projected_table_schema, file_metadata) in files
549        {
550            let mut stream = self
551                .build_read_stream(
552                    compat_schema,
553                    &object_store,
554                    &file_metadata,
555                    file_schema_projection,
556                    filters.clone(),
557                )
558                .await?;
559
560            let fields = projected_table_schema
561                .fields()
562                .iter()
563                .map(|f| f.name().clone())
564                .collect::<Vec<_>>();
565
566            // TODO(hl): make this configurable through options.
567            let pending_mem_threshold = ReadableSize::mb(32).as_bytes();
568            let mut pending_mem_size = 0;
569            let mut pending = vec![];
570
571            while let Some(r) = stream.next().await {
572                let record_batch = r.context(error::ReadDfRecordBatchSnafu)?;
573                let record_batch = if let Some(max_insert_rows) = max_insert_rows {
574                    let remaining_rows = max_insert_rows - accepted_rows;
575                    if record_batch.num_rows() > remaining_rows {
576                        record_batch.slice(0, remaining_rows)
577                    } else {
578                        record_batch
579                    }
580                } else {
581                    record_batch
582                };
583                let record_batch_rows = record_batch.num_rows();
584                let vectors =
585                    Helper::try_into_vectors(record_batch.columns()).context(IntoVectorsSnafu)?;
586
587                pending_mem_size += vectors.iter().map(|v| v.memory_size()).sum::<usize>();
588
589                let columns_values = fields
590                    .iter()
591                    .cloned()
592                    .zip(vectors)
593                    .collect::<HashMap<_, _>>();
594
595                pending.push(self.inserter.handle_table_insert(
596                    InsertRequest {
597                        catalog_name: req.catalog_name.clone(),
598                        schema_name: req.schema_name.clone(),
599                        table_name: req.table_name.clone(),
600                        columns_values,
601                        skip_wal: query_ctx.skip_wal(),
602                    },
603                    query_ctx.clone(),
604                ));
605                accepted_rows += record_batch_rows;
606
607                if pending_mem_size as u64 >= pending_mem_threshold {
608                    let (rows, cost) = batch_insert(&mut pending, &mut pending_mem_size).await?;
609                    rows_inserted += rows;
610                    insert_cost += cost;
611                }
612
613                if let Some(max_insert_rows) = max_insert_rows
614                    && accepted_rows == max_insert_rows
615                {
616                    if !pending.is_empty() {
617                        let (rows, cost) =
618                            batch_insert(&mut pending, &mut pending_mem_size).await?;
619                        rows_inserted += rows;
620                        insert_cost += cost;
621                    }
622                    return Ok(gen_insert_output(rows_inserted, insert_cost));
623                }
624            }
625
626            if !pending.is_empty() {
627                let (rows, cost) = batch_insert(&mut pending, &mut pending_mem_size).await?;
628                rows_inserted += rows;
629                insert_cost += cost;
630            }
631        }
632
633        Ok(gen_insert_output(rows_inserted, insert_cost))
634    }
635}
636
637fn gen_insert_output(rows_inserted: usize, insert_cost: usize) -> Output {
638    Output::new(
639        OutputData::AffectedRows(rows_inserted),
640        OutputMeta::new_with_cost(insert_cost),
641    )
642}
643
644struct SkipBadRecordsStream {
645    inner: DfSendableRecordBatchStream,
646    path: String,
647}
648
649impl SkipBadRecordsStream {
650    fn new(inner: DfSendableRecordBatchStream, path: impl Into<String>) -> Self {
651        Self {
652            inner,
653            path: path.into(),
654        }
655    }
656}
657
658impl datafusion::physical_plan::RecordBatchStream for SkipBadRecordsStream {
659    fn schema(&self) -> SchemaRef {
660        self.inner.schema()
661    }
662}
663
664impl futures::Stream for SkipBadRecordsStream {
665    type Item = datafusion_common::Result<RecordBatch>;
666
667    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
668        let this = self.get_mut();
669        loop {
670            match this.inner.as_mut().poll_next(cx) {
671                Poll::Ready(Some(Err(error))) if is_skippable_record_error(&error) => {
672                    common_telemetry::warn!(
673                        "Skipping bad record while copying from {}: {}",
674                        this.path,
675                        error
676                    );
677                    continue;
678                }
679                other => return other,
680            }
681        }
682    }
683}
684
685fn is_skippable_record_error(error: &DataFusionError) -> bool {
686    match error {
687        DataFusionError::ArrowError(error, _) => is_skippable_arrow_error(error),
688        DataFusionError::External(error) => error
689            .downcast_ref::<ArrowError>()
690            .is_some_and(is_skippable_arrow_error),
691        DataFusionError::Context(_, error) => is_skippable_record_error(error),
692        _ => false,
693    }
694}
695
696/// Executes all pending inserts all at once, drain pending requests and reset pending bytes.
697async fn batch_insert(
698    pending: &mut Vec<impl Future<Output = Result<Output>>>,
699    pending_bytes: &mut usize,
700) -> Result<(OutputRows, OutputCost)> {
701    let batch = pending.drain(..);
702    let result = futures::future::try_join_all(batch)
703        .await?
704        .iter()
705        .map(|o| o.extract_rows_and_cost())
706        .reduce(|(a, b), (c, d)| (a + c, b + d))
707        .unwrap_or((0, 0));
708    *pending_bytes = 0;
709    Ok(result)
710}
711
712/// Custom type compatibility check for GreptimeDB that handles Map -> Binary (JSON) conversion
713fn can_cast_types_for_greptime(from: &ArrowDataType, to: &ArrowDataType) -> bool {
714    // Handle Map -> Binary conversion for JSON types
715    if let ArrowDataType::Map(_, _) = from
716        && let ArrowDataType::Binary = to
717    {
718        return true;
719    }
720
721    // For all other cases, use Arrow's built-in can_cast_types
722    can_cast_types(from, to)
723}
724
725fn csv_reader_schema_for_skip_bad_records(file: &SchemaRef, compat: &SchemaRef) -> Schema {
726    let fields = file
727        .fields()
728        .iter()
729        .enumerate()
730        .map(|(idx, file_field)| match compat.fields().get(idx) {
731            Some(compat_field) if can_csv_reader_parse_type(compat_field.data_type()) => {
732                compat_field.clone()
733            }
734            _ => file_field.clone(),
735        })
736        .collect::<Vec<_>>();
737
738    Schema::new_with_metadata(fields, file.metadata().clone())
739}
740
741fn can_csv_reader_parse_type(data_type: &ArrowDataType) -> bool {
742    match data_type {
743        ArrowDataType::Boolean
744        | ArrowDataType::Decimal32(_, _)
745        | ArrowDataType::Decimal64(_, _)
746        | ArrowDataType::Decimal128(_, _)
747        | ArrowDataType::Decimal256(_, _)
748        | ArrowDataType::Int8
749        | ArrowDataType::Int16
750        | ArrowDataType::Int32
751        | ArrowDataType::Int64
752        | ArrowDataType::UInt8
753        | ArrowDataType::UInt16
754        | ArrowDataType::UInt32
755        | ArrowDataType::UInt64
756        | ArrowDataType::Float32
757        | ArrowDataType::Float64
758        | ArrowDataType::Date32
759        | ArrowDataType::Date64
760        | ArrowDataType::Time32(_)
761        | ArrowDataType::Time64(_)
762        | ArrowDataType::Timestamp(_, _)
763        | ArrowDataType::Null
764        | ArrowDataType::Utf8
765        | ArrowDataType::Utf8View => true,
766        ArrowDataType::Dictionary(_, value_type) => value_type.as_ref() == &ArrowDataType::Utf8,
767        _ => false,
768    }
769}
770
771fn ensure_schema_compatible(from: &SchemaRef, to: &SchemaRef) -> Result<()> {
772    let not_match = from
773        .fields
774        .iter()
775        .zip(to.fields.iter())
776        .map(|(l, r)| (l.data_type(), r.data_type()))
777        .enumerate()
778        .find(|(_, (l, r))| !can_cast_types_for_greptime(l, r));
779
780    if let Some((index, _)) = not_match {
781        error::InvalidSchemaSnafu {
782            index,
783            table_schema: to.to_string(),
784            file_schema: from.to_string(),
785        }
786        .fail()
787    } else {
788        Ok(())
789    }
790}
791
792fn validate_csv_headers_if_required(file_metadata: &FileMetadata, table: &SchemaRef) -> Result<()> {
793    let FileMetadata::Csv {
794        schema,
795        format,
796        path,
797    } = file_metadata
798    else {
799        return Ok(());
800    };
801
802    if !format.strict_headers {
803        return Ok(());
804    }
805
806    let mut seen_file_columns = HashSet::with_capacity(schema.fields().len());
807    let duplicate_columns = schema
808        .fields()
809        .iter()
810        .filter_map(|field| {
811            if seen_file_columns.insert(field.name().clone()) {
812                None
813            } else {
814                Some(field.name().clone())
815            }
816        })
817        .collect::<BTreeSet<_>>()
818        .into_iter()
819        .collect::<Vec<_>>();
820    let file_columns = seen_file_columns.into_iter().collect::<BTreeSet<_>>();
821    let table_columns = table
822        .fields()
823        .iter()
824        .map(|field| field.name().clone())
825        .collect::<BTreeSet<_>>();
826    let unknown_columns = file_columns
827        .difference(&table_columns)
828        .cloned()
829        .collect::<Vec<_>>();
830    let missing_columns = table_columns
831        .difference(&file_columns)
832        .cloned()
833        .collect::<Vec<_>>();
834
835    ensure!(
836        unknown_columns.is_empty() && missing_columns.is_empty() && duplicate_columns.is_empty(),
837        error::CsvHeaderMismatchSnafu {
838            path,
839            unknown_columns,
840            missing_columns,
841            duplicate_columns,
842        }
843    );
844
845    Ok(())
846}
847
848/// Generates a maybe compatible schema of the file schema.
849///
850/// If there is a field is found in table schema,
851/// copy the field data type to maybe compatible schema(`compatible_fields`).
852fn generated_schema_projection_and_compatible_file_schema(
853    file: &SchemaRef,
854    table: &SchemaRef,
855) -> (Vec<usize>, Vec<usize>, Schema) {
856    let mut file_projection = Vec::with_capacity(file.fields.len());
857    let mut table_projection = Vec::with_capacity(file.fields.len());
858    let mut compatible_fields = file.fields.iter().cloned().collect::<Vec<_>>();
859    for (file_idx, file_field) in file.fields.iter().enumerate() {
860        if let Some((table_idx, table_field)) = table.fields.find(file_field.name()) {
861            file_projection.push(file_idx);
862            table_projection.push(table_idx);
863
864            // Safety: the compatible_fields has same length as file schema
865            compatible_fields[file_idx] = table_field.clone();
866        }
867    }
868
869    (
870        file_projection,
871        table_projection,
872        Schema::new(compatible_fields),
873    )
874}
875
876struct CopyFromSchemaMapping {
877    file_projection: Vec<usize>,
878    table_projection: Vec<usize>,
879    compat_file_schema: Schema,
880}
881
882fn copy_from_schema_mapping(
883    file_metadata: &FileMetadata,
884    table: &SchemaRef,
885) -> CopyFromSchemaMapping {
886    match file_metadata {
887        FileMetadata::Csv { schema, format, .. } if !format.has_header => {
888            generated_positional_schema_projection_and_compatible_file_schema(schema, table)
889        }
890        _ => {
891            let (file_projection, table_projection, compat_file_schema) =
892                generated_schema_projection_and_compatible_file_schema(
893                    file_metadata.schema(),
894                    table,
895                );
896            CopyFromSchemaMapping {
897                file_projection,
898                table_projection,
899                compat_file_schema,
900            }
901        }
902    }
903}
904
905fn generated_positional_schema_projection_and_compatible_file_schema(
906    file: &SchemaRef,
907    table: &SchemaRef,
908) -> CopyFromSchemaMapping {
909    let len = file.fields.len().min(table.fields.len());
910    let file_projection = (0..len).collect::<Vec<_>>();
911    let table_projection = (0..len).collect::<Vec<_>>();
912    let compatible_fields = file
913        .fields
914        .iter()
915        .enumerate()
916        .map(|(idx, file_field)| {
917            if idx < len {
918                table.fields[idx].clone()
919            } else {
920                file_field.clone()
921            }
922        })
923        .collect::<Vec<_>>();
924
925    CopyFromSchemaMapping {
926        file_projection,
927        table_projection,
928        compat_file_schema: Schema::new(compatible_fields),
929    }
930}
931
932#[cfg(test)]
933mod tests {
934    use std::sync::Arc;
935
936    use datatypes::arrow::datatypes::{DataType, Field, Schema};
937
938    use super::*;
939
940    fn copy_from_request(location: &str, pattern: Option<&str>) -> CopyTableRequest {
941        CopyTableRequest {
942            catalog_name: "greptime".into(),
943            schema_name: "public".into(),
944            table_name: "test".into(),
945            location: location.into(),
946            with: HashMap::new(),
947            connection: HashMap::new(),
948            pattern: pattern.map(String::from),
949            direction: table::requests::CopyDirection::Import,
950            timestamp_range: None,
951            limit: None,
952        }
953    }
954
955    #[tokio::test]
956    async fn test_copy_from_local_paths() {
957        use common_test_util::temp_dir::create_temp_dir;
958
959        let dir = create_temp_dir("copy_from_paths");
960        let nested = dir.path().join("nested");
961        std::fs::create_dir_all(nested.join("data.dir")).unwrap();
962        std::fs::write(nested.join("data file.parquet"), b"data").unwrap();
963        std::fs::write(nested.join("other.csv"), b"other").unwrap();
964        let access = LocalFileAccess::sandboxed(dir.path()).unwrap();
965        for location in [
966            "nested/data file.parquet".to_string(),
967            url::Url::from_file_path(nested.join("data file.parquet"))
968                .unwrap()
969                .to_string(),
970        ] {
971            // PATTERN applies to directories, not an explicitly selected file.
972            let req = copy_from_request(&location, Some("does-not-match"));
973            let (store, paths) = list_copy_from_paths(&req, &access).await.unwrap();
974            assert_eq!(paths, ["data file.parquet"]);
975            assert_eq!(store.read(&paths[0]).await.unwrap().to_vec(), b"data");
976        }
977
978        let req = copy_from_request("nested/", Some("^data"));
979        let (_, paths) = list_copy_from_paths(&req, &access).await.unwrap();
980        assert_eq!(paths, ["data file.parquet"]);
981        let req = copy_from_request("nested/", None);
982        let (_, mut paths) = list_copy_from_paths(&req, &access).await.unwrap();
983        paths.sort();
984        assert_eq!(paths, ["data file.parquet", "other.csv"]);
985        let req = copy_from_request("nested/data.dir", None);
986        assert!(
987            list_copy_from_paths(&req, &access)
988                .await
989                .unwrap()
990                .1
991                .is_empty()
992        );
993        let req = copy_from_request("nested/", Some("does-not-match"));
994        assert!(
995            list_copy_from_paths(&req, &access)
996                .await
997                .unwrap()
998                .1
999                .is_empty()
1000        );
1001
1002        let req = copy_from_request("nested/missing.parquet", None);
1003        assert!(matches!(
1004            list_copy_from_paths(&req, &access).await,
1005            Err(error::Error::ListObjects {
1006                source: common_datasource::error::Error::ListObjects { error, .. }, ..
1007            }) if error.kind() == object_store::ErrorKind::NotFound
1008        ));
1009        let req = copy_from_request("nested/data file.parquet", Some("["));
1010        assert!(matches!(
1011            list_copy_from_paths(&req, &access).await,
1012            Err(error::Error::BuildRegex { .. })
1013        ));
1014        let req = copy_from_request("nested/data file.parquet", None);
1015        assert!(matches!(
1016            list_copy_from_paths(&req, &LocalFileAccess::Disabled).await,
1017            Err(error::Error::BuildBackend {
1018                source: common_datasource::error::Error::LocalFileAccessDisabled { .. },
1019                ..
1020            })
1021        ));
1022        let req = copy_from_request("../outside.parquet", None);
1023        assert!(matches!(
1024            list_copy_from_paths(&req, &access).await,
1025            Err(error::Error::BuildBackend {
1026                source: common_datasource::error::Error::LocalFileAccessDenied { .. },
1027                ..
1028            })
1029        ));
1030    }
1031
1032    #[cfg(unix)]
1033    #[tokio::test]
1034    async fn test_copy_from_local_symlinks() {
1035        use common_test_util::temp_dir::create_temp_dir;
1036
1037        let dir = create_temp_dir("copy_from_symlinks");
1038        std::fs::write(dir.path().join("data.parquet"), b"data").unwrap();
1039        std::os::unix::fs::symlink("data.parquet", dir.path().join("link.parquet")).unwrap();
1040        std::os::unix::fs::symlink("missing.parquet", dir.path().join("dangling.parquet")).unwrap();
1041        let access = LocalFileAccess::sandboxed(dir.path()).unwrap();
1042
1043        let req = copy_from_request("link.parquet", None);
1044        assert!(
1045            list_copy_from_paths(&req, &access)
1046                .await
1047                .unwrap()
1048                .1
1049                .is_empty()
1050        );
1051        let req = copy_from_request("./", None);
1052        let (_, paths) = list_copy_from_paths(&req, &access).await.unwrap();
1053        assert_eq!(paths, ["data.parquet"]);
1054        let req = copy_from_request("dangling.parquet", None);
1055        assert!(matches!(
1056            list_copy_from_paths(&req, &access).await,
1057            Err(error::Error::ListObjects {
1058                source: common_datasource::error::Error::ListObjects { error, .. }, ..
1059            }) if error.kind() == object_store::ErrorKind::NotFound
1060        ));
1061    }
1062
1063    #[tokio::test]
1064    async fn test_copy_from_s3_paths_without_relisting() {
1065        use std::sync::Mutex;
1066
1067        use axum::Router;
1068        use axum::http::{Method, StatusCode, Uri};
1069        use axum::response::IntoResponse;
1070
1071        let requests = Arc::new(Mutex::new(Vec::new()));
1072        let recorded = requests.clone();
1073        let app = Router::new().fallback(move |method: Method, uri: Uri| {
1074            let recorded = recorded.clone();
1075            async move {
1076                recorded.lock().unwrap().push((method.clone(), uri.clone()));
1077                if method == Method::HEAD {
1078                    let status = if uri.path().ends_with("missing.parquet") {
1079                        StatusCode::NOT_FOUND
1080                    } else if uri.path().ends_with("denied.parquet") {
1081                        StatusCode::FORBIDDEN
1082                    } else {
1083                        StatusCode::OK
1084                    };
1085                    return (status, [("content-length", "0")]).into_response();
1086                }
1087                assert_eq!(method, Method::GET);
1088                let query =
1089                    axum::extract::Query::<HashMap<String, String>>::try_from_uri(&uri).unwrap();
1090                assert_eq!(query.get("prefix").unwrap(), "backup/schema/1/");
1091                assert_eq!(query.get("list-type").unwrap(), "2");
1092                let files = (0..16).map(|i| format!(
1093                    "<Contents><Key>backup/schema/1/data.{i}.parquet</Key><LastModified>2026-09-01T00:00:00Z</LastModified><Size>0</Size></Contents>"
1094                )).collect::<String>();
1095                format!(
1096                    "<ListBucketResult><IsTruncated>false</IsTruncated>{files}\
1097                    <Contents><Key>backup/schema/1/other.csv</Key><LastModified>2026-09-01T00:00:00Z</LastModified><Size>0</Size></Contents>\
1098                    <CommonPrefixes><Prefix>backup/schema/1/data.dir/</Prefix></CommonPrefixes>\
1099                    </ListBucketResult>"
1100                )
1101                .into_response()
1102            }
1103        });
1104        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1105        let endpoint = format!("http://{}", listener.local_addr().unwrap());
1106        let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
1107        let connection = HashMap::from([
1108            ("endpoint".into(), endpoint),
1109            ("region".into(), "us-east-1".into()),
1110            ("access_key_id".into(), "test".into()),
1111            ("secret_access_key".into(), "test".into()),
1112            ("disable_ec2_metadata".into(), "true".into()),
1113        ]);
1114        for i in 0..16 {
1115            let filename = format!("data.{i}.parquet");
1116            let mut req = copy_from_request(
1117                &format!("s3://bucket/backup/schema/1/{filename}"),
1118                Some("does-not-match"),
1119            );
1120            req.connection = connection.clone();
1121            let (_, paths) = list_copy_from_paths(&req, &LocalFileAccess::Disabled)
1122                .await
1123                .unwrap();
1124            assert_eq!(paths, [filename]);
1125        }
1126        {
1127            let requests = requests.lock().unwrap();
1128            assert_eq!(requests.len(), 16);
1129            assert!(requests.iter().all(|(method, uri)| method == Method::HEAD
1130                && uri.path().starts_with("/bucket/backup/schema/1/")));
1131        }
1132        let mut req = copy_from_request("s3://bucket/backup/schema/1/", Some(r"^data\."));
1133        req.connection = connection.clone();
1134        let (_, mut paths) = list_copy_from_paths(&req, &LocalFileAccess::Disabled)
1135            .await
1136            .unwrap();
1137        paths.sort();
1138        let mut expected = (0..16)
1139            .map(|i| format!("data.{i}.parquet"))
1140            .collect::<Vec<_>>();
1141        expected.sort();
1142        assert_eq!(paths, expected);
1143        assert_eq!(
1144            requests
1145                .lock()
1146                .unwrap()
1147                .iter()
1148                .filter(|(method, _)| *method == Method::GET)
1149                .count(),
1150            1
1151        );
1152
1153        for (name, kind) in [
1154            ("missing", object_store::ErrorKind::NotFound),
1155            ("denied", object_store::ErrorKind::PermissionDenied),
1156        ] {
1157            let mut req =
1158                copy_from_request(&format!("s3://bucket/backup/schema/1/{name}.parquet"), None);
1159            req.connection = connection.clone();
1160            assert!(matches!(
1161                list_copy_from_paths(&req, &LocalFileAccess::Disabled).await,
1162                Err(error::Error::ListObjects {
1163                    source: common_datasource::error::Error::ListObjects { error, .. }, ..
1164                }) if error.kind() == kind
1165            ));
1166        }
1167        server.abort();
1168    }
1169
1170    fn test_schema_matches(from: (DataType, bool), to: (DataType, bool), matches: bool) {
1171        let s1 = Arc::new(Schema::new(vec![Field::new("col", from.0.clone(), from.1)]));
1172        let s2 = Arc::new(Schema::new(vec![Field::new("col", to.0.clone(), to.1)]));
1173        let res = ensure_schema_compatible(&s1, &s2);
1174        assert_eq!(
1175            matches,
1176            res.is_ok(),
1177            "from data type: {}, to data type: {}, expected: {}, but got: {}",
1178            from.0,
1179            to.0,
1180            matches,
1181            res.is_ok()
1182        )
1183    }
1184
1185    #[test]
1186    fn test_ensure_datatype_matches_ignore_timezone() {
1187        test_schema_matches(
1188            (
1189                DataType::Timestamp(datatypes::arrow::datatypes::TimeUnit::Second, None),
1190                true,
1191            ),
1192            (
1193                DataType::Timestamp(datatypes::arrow::datatypes::TimeUnit::Second, None),
1194                true,
1195            ),
1196            true,
1197        );
1198
1199        test_schema_matches(
1200            (
1201                DataType::Timestamp(
1202                    datatypes::arrow::datatypes::TimeUnit::Second,
1203                    Some("UTC".into()),
1204                ),
1205                true,
1206            ),
1207            (
1208                DataType::Timestamp(datatypes::arrow::datatypes::TimeUnit::Second, None),
1209                true,
1210            ),
1211            true,
1212        );
1213
1214        test_schema_matches(
1215            (
1216                DataType::Timestamp(
1217                    datatypes::arrow::datatypes::TimeUnit::Second,
1218                    Some("UTC".into()),
1219                ),
1220                true,
1221            ),
1222            (
1223                DataType::Timestamp(
1224                    datatypes::arrow::datatypes::TimeUnit::Second,
1225                    Some("PDT".into()),
1226                ),
1227                true,
1228            ),
1229            true,
1230        );
1231
1232        test_schema_matches(
1233            (
1234                DataType::Timestamp(
1235                    datatypes::arrow::datatypes::TimeUnit::Second,
1236                    Some("UTC".into()),
1237                ),
1238                true,
1239            ),
1240            (
1241                DataType::Timestamp(
1242                    datatypes::arrow::datatypes::TimeUnit::Millisecond,
1243                    Some("UTC".into()),
1244                ),
1245                true,
1246            ),
1247            true,
1248        );
1249
1250        test_schema_matches((DataType::Int8, true), (DataType::Int8, true), true);
1251
1252        test_schema_matches((DataType::Int8, true), (DataType::Int16, true), true);
1253    }
1254
1255    #[test]
1256    fn test_data_type_equals_ignore_timezone_with_options() {
1257        test_schema_matches(
1258            (
1259                DataType::Timestamp(
1260                    datatypes::arrow::datatypes::TimeUnit::Microsecond,
1261                    Some("UTC".into()),
1262                ),
1263                true,
1264            ),
1265            (
1266                DataType::Timestamp(
1267                    datatypes::arrow::datatypes::TimeUnit::Millisecond,
1268                    Some("PDT".into()),
1269                ),
1270                true,
1271            ),
1272            true,
1273        );
1274
1275        test_schema_matches(
1276            (DataType::Utf8, true),
1277            (
1278                DataType::Timestamp(
1279                    datatypes::arrow::datatypes::TimeUnit::Millisecond,
1280                    Some("PDT".into()),
1281                ),
1282                true,
1283            ),
1284            true,
1285        );
1286
1287        test_schema_matches(
1288            (
1289                DataType::Timestamp(
1290                    datatypes::arrow::datatypes::TimeUnit::Millisecond,
1291                    Some("PDT".into()),
1292                ),
1293                true,
1294            ),
1295            (DataType::Utf8, true),
1296            true,
1297        );
1298    }
1299
1300    #[test]
1301    fn test_map_to_binary_json_compatibility() {
1302        // Test Map -> Binary conversion for JSON types
1303        let map_type = DataType::Map(
1304            Arc::new(Field::new(
1305                "key_value",
1306                DataType::Struct(
1307                    vec![
1308                        Field::new("key", DataType::Utf8, false),
1309                        Field::new("value", DataType::Utf8, false),
1310                    ]
1311                    .into(),
1312                ),
1313                false,
1314            )),
1315            false,
1316        );
1317
1318        test_schema_matches((map_type, false), (DataType::Binary, true), true);
1319
1320        test_schema_matches((DataType::Int8, true), (DataType::Int16, true), true);
1321        test_schema_matches((DataType::Utf8, true), (DataType::Binary, true), true);
1322    }
1323
1324    fn make_test_schema(v: &[Field]) -> Arc<Schema> {
1325        Arc::new(Schema::new(v.to_vec()))
1326    }
1327
1328    #[test]
1329    fn test_compatible_file_schema() {
1330        let file_schema0 = make_test_schema(&[
1331            Field::new("c1", DataType::UInt8, true),
1332            Field::new("c2", DataType::UInt8, true),
1333        ]);
1334
1335        let table_schema = make_test_schema(&[
1336            Field::new("c1", DataType::Int16, true),
1337            Field::new("c2", DataType::Int16, true),
1338            Field::new("c3", DataType::Int16, true),
1339        ]);
1340
1341        let compat_schema = make_test_schema(&[
1342            Field::new("c1", DataType::Int16, true),
1343            Field::new("c2", DataType::Int16, true),
1344        ]);
1345
1346        let (_, tp, _) =
1347            generated_schema_projection_and_compatible_file_schema(&file_schema0, &table_schema);
1348
1349        assert_eq!(table_schema.project(&tp).unwrap(), *compat_schema);
1350    }
1351
1352    #[test]
1353    fn test_schema_projection() {
1354        let file_schema0 = make_test_schema(&[
1355            Field::new("c1", DataType::UInt8, true),
1356            Field::new("c2", DataType::UInt8, true),
1357            Field::new("c3", DataType::UInt8, true),
1358        ]);
1359
1360        let file_schema1 = make_test_schema(&[
1361            Field::new("c3", DataType::UInt8, true),
1362            Field::new("c4", DataType::UInt8, true),
1363        ]);
1364
1365        let file_schema2 = make_test_schema(&[
1366            Field::new("c3", DataType::UInt8, true),
1367            Field::new("c4", DataType::UInt8, true),
1368            Field::new("c5", DataType::UInt8, true),
1369        ]);
1370
1371        let file_schema3 = make_test_schema(&[
1372            Field::new("c1", DataType::UInt8, true),
1373            Field::new("c2", DataType::UInt8, true),
1374        ]);
1375
1376        let table_schema = make_test_schema(&[
1377            Field::new("c3", DataType::UInt8, true),
1378            Field::new("c4", DataType::UInt8, true),
1379            Field::new("c5", DataType::UInt8, true),
1380        ]);
1381
1382        let tests = [
1383            (&file_schema0, &table_schema, true), // intersection
1384            (&file_schema1, &table_schema, true), // subset
1385            (&file_schema2, &table_schema, true), // full-eq
1386            (&file_schema3, &table_schema, true), // non-intersection
1387        ];
1388
1389        for test in tests {
1390            let (fp, tp, _) =
1391                generated_schema_projection_and_compatible_file_schema(test.0, test.1);
1392            assert_eq!(test.0.project(&fp).unwrap(), test.1.project(&tp).unwrap());
1393        }
1394    }
1395
1396    #[test]
1397    fn test_csv_reader_schema_for_skip_bad_records() {
1398        let file_schema = make_test_schema(&[
1399            Field::new("id", DataType::Utf8, true),
1400            Field::new("jsons", DataType::Utf8, true),
1401            Field::new("ts", DataType::Utf8, true),
1402        ]);
1403        let compat_schema = make_test_schema(&[
1404            Field::new("id", DataType::UInt32, true),
1405            Field::new("jsons", DataType::Binary, true),
1406            Field::new(
1407                "ts",
1408                DataType::Timestamp(datatypes::arrow::datatypes::TimeUnit::Millisecond, None),
1409                true,
1410            ),
1411        ]);
1412
1413        let reader_schema = csv_reader_schema_for_skip_bad_records(&file_schema, &compat_schema);
1414
1415        assert_eq!(reader_schema.field(0).data_type(), &DataType::UInt32);
1416        assert_eq!(reader_schema.field(1).data_type(), &DataType::Utf8);
1417        assert_eq!(
1418            reader_schema.field(2).data_type(),
1419            compat_schema.field(2).data_type()
1420        );
1421    }
1422
1423    fn make_csv_metadata(schema: Arc<Schema>, has_header: bool) -> FileMetadata {
1424        FileMetadata::Csv {
1425            schema,
1426            format: CsvFormat {
1427                has_header,
1428                ..CsvFormat::default()
1429            },
1430            path: "test.csv".to_string(),
1431        }
1432    }
1433
1434    fn make_strict_csv_metadata(schema: Arc<Schema>) -> FileMetadata {
1435        FileMetadata::Csv {
1436            schema,
1437            format: CsvFormat {
1438                strict_headers: true,
1439                ..CsvFormat::default()
1440            },
1441            path: "test.csv".to_string(),
1442        }
1443    }
1444
1445    fn assert_field(schema: &Schema, idx: usize, name: &str, data_type: &DataType) {
1446        let field = schema.field(idx);
1447        assert_eq!(field.name(), name);
1448        assert_eq!(field.data_type(), data_type);
1449    }
1450
1451    #[test]
1452    fn test_strict_csv_headers_allows_reordered_columns() {
1453        let file_schema = make_test_schema(&[
1454            Field::new("ts", DataType::Utf8, true),
1455            Field::new("host_id", DataType::UInt8, true),
1456            Field::new("reading_value", DataType::Float64, true),
1457        ]);
1458        let table_schema = make_test_schema(&[
1459            Field::new("host_id", DataType::UInt32, true),
1460            Field::new("reading_value", DataType::Float64, true),
1461            Field::new("ts", DataType::Utf8, true),
1462        ]);
1463
1464        validate_csv_headers_if_required(&make_strict_csv_metadata(file_schema), &table_schema)
1465            .unwrap();
1466    }
1467
1468    #[test]
1469    fn test_strict_csv_headers_rejects_unknown_columns() {
1470        let file_schema = make_test_schema(&[
1471            Field::new("host_id", DataType::UInt8, true),
1472            Field::new("reading_value", DataType::Float64, true),
1473            Field::new("ts", DataType::Utf8, true),
1474            Field::new("extra", DataType::Utf8, true),
1475        ]);
1476        let table_schema = make_test_schema(&[
1477            Field::new("host_id", DataType::UInt32, true),
1478            Field::new("reading_value", DataType::Float64, true),
1479            Field::new("ts", DataType::Utf8, true),
1480        ]);
1481
1482        let err =
1483            validate_csv_headers_if_required(&make_strict_csv_metadata(file_schema), &table_schema)
1484                .unwrap_err();
1485
1486        assert!(matches!(
1487            err,
1488            error::Error::CsvHeaderMismatch {
1489                unknown_columns,
1490                missing_columns,
1491                duplicate_columns,
1492                ..
1493            } if unknown_columns == vec!["extra".to_string()]
1494                && missing_columns.is_empty()
1495                && duplicate_columns.is_empty()
1496        ));
1497    }
1498
1499    #[test]
1500    fn test_strict_csv_headers_rejects_missing_columns() {
1501        let file_schema = make_test_schema(&[
1502            Field::new("host_id", DataType::UInt8, true),
1503            Field::new("ts", DataType::Utf8, true),
1504        ]);
1505        let table_schema = make_test_schema(&[
1506            Field::new("host_id", DataType::UInt32, true),
1507            Field::new("reading_value", DataType::Float64, true),
1508            Field::new("ts", DataType::Utf8, true),
1509        ]);
1510
1511        let err =
1512            validate_csv_headers_if_required(&make_strict_csv_metadata(file_schema), &table_schema)
1513                .unwrap_err();
1514
1515        assert!(matches!(
1516            err,
1517            error::Error::CsvHeaderMismatch {
1518                unknown_columns,
1519                missing_columns,
1520                duplicate_columns,
1521                ..
1522            } if unknown_columns.is_empty()
1523                && missing_columns == vec!["reading_value".to_string()]
1524                && duplicate_columns.is_empty()
1525        ));
1526    }
1527
1528    #[test]
1529    fn test_strict_csv_headers_rejects_duplicate_columns() {
1530        let file_schema = make_test_schema(&[
1531            Field::new("host_id", DataType::UInt8, true),
1532            Field::new("reading_value", DataType::Float64, true),
1533            Field::new("ts", DataType::Utf8, true),
1534            Field::new("host_id", DataType::UInt16, true),
1535        ]);
1536        let table_schema = make_test_schema(&[
1537            Field::new("host_id", DataType::UInt32, true),
1538            Field::new("reading_value", DataType::Float64, true),
1539            Field::new("ts", DataType::Utf8, true),
1540        ]);
1541
1542        let err =
1543            validate_csv_headers_if_required(&make_strict_csv_metadata(file_schema), &table_schema)
1544                .unwrap_err();
1545
1546        assert!(matches!(
1547            err,
1548            error::Error::CsvHeaderMismatch {
1549                unknown_columns,
1550                missing_columns,
1551                duplicate_columns,
1552                ..
1553            } if unknown_columns.is_empty()
1554                && missing_columns.is_empty()
1555                && duplicate_columns == vec!["host_id".to_string()]
1556        ));
1557    }
1558
1559    #[test]
1560    fn test_headerless_csv_schema_projection_is_positional() {
1561        let file_schema = make_test_schema(&[
1562            Field::new("column_1", DataType::UInt8, true),
1563            Field::new("column_2", DataType::Float64, true),
1564            Field::new("column_3", DataType::Utf8, true),
1565        ]);
1566        let table_schema = make_test_schema(&[
1567            Field::new("host_id", DataType::UInt32, true),
1568            Field::new("reading_value", DataType::Float64, true),
1569            Field::new(
1570                "ts",
1571                DataType::Timestamp(datatypes::arrow::datatypes::TimeUnit::Millisecond, None),
1572                true,
1573            ),
1574        ]);
1575
1576        let mapping =
1577            copy_from_schema_mapping(&make_csv_metadata(file_schema, false), &table_schema);
1578
1579        assert_eq!(mapping.file_projection, vec![0, 1, 2]);
1580        assert_eq!(mapping.table_projection, vec![0, 1, 2]);
1581        assert_field(&mapping.compat_file_schema, 0, "host_id", &DataType::UInt32);
1582        assert_field(
1583            &mapping.compat_file_schema,
1584            1,
1585            "reading_value",
1586            &DataType::Float64,
1587        );
1588        assert_field(
1589            &mapping.compat_file_schema,
1590            2,
1591            "ts",
1592            table_schema.field(2).data_type(),
1593        );
1594        assert_eq!(
1595            mapping
1596                .compat_file_schema
1597                .project(&mapping.file_projection)
1598                .unwrap(),
1599            table_schema.project(&mapping.table_projection).unwrap()
1600        );
1601    }
1602
1603    #[test]
1604    fn test_headerless_csv_schema_projection_ignores_extra_file_columns() {
1605        let file_schema = make_test_schema(&[
1606            Field::new("column_1", DataType::UInt8, true),
1607            Field::new("column_2", DataType::Float64, true),
1608            Field::new("column_3", DataType::Utf8, true),
1609            Field::new("column_4", DataType::Utf8, true),
1610        ]);
1611        let table_schema = make_test_schema(&[
1612            Field::new("host_id", DataType::UInt32, true),
1613            Field::new("reading_value", DataType::Float64, true),
1614            Field::new("ts", DataType::Utf8, true),
1615        ]);
1616
1617        let mapping =
1618            copy_from_schema_mapping(&make_csv_metadata(file_schema, false), &table_schema);
1619
1620        assert_eq!(mapping.file_projection, vec![0, 1, 2]);
1621        assert_eq!(mapping.table_projection, vec![0, 1, 2]);
1622        assert_eq!(mapping.compat_file_schema.fields().len(), 4);
1623        assert_field(&mapping.compat_file_schema, 0, "host_id", &DataType::UInt32);
1624        assert_field(
1625            &mapping.compat_file_schema,
1626            1,
1627            "reading_value",
1628            &DataType::Float64,
1629        );
1630        assert_field(&mapping.compat_file_schema, 2, "ts", &DataType::Utf8);
1631        assert_field(&mapping.compat_file_schema, 3, "column_4", &DataType::Utf8);
1632    }
1633
1634    #[test]
1635    fn test_headerless_csv_schema_projection_supports_prefix_import() {
1636        let file_schema = make_test_schema(&[
1637            Field::new("column_1", DataType::UInt8, true),
1638            Field::new("column_2", DataType::Float64, true),
1639        ]);
1640        let table_schema = make_test_schema(&[
1641            Field::new("host_id", DataType::UInt32, true),
1642            Field::new("reading_value", DataType::Float64, true),
1643            Field::new("ts", DataType::Utf8, true),
1644        ]);
1645
1646        let mapping =
1647            copy_from_schema_mapping(&make_csv_metadata(file_schema, false), &table_schema);
1648
1649        assert_eq!(mapping.file_projection, vec![0, 1]);
1650        assert_eq!(mapping.table_projection, vec![0, 1]);
1651        assert_field(&mapping.compat_file_schema, 0, "host_id", &DataType::UInt32);
1652        assert_field(
1653            &mapping.compat_file_schema,
1654            1,
1655            "reading_value",
1656            &DataType::Float64,
1657        );
1658        assert_eq!(
1659            mapping
1660                .compat_file_schema
1661                .project(&mapping.file_projection)
1662                .unwrap(),
1663            table_schema.project(&mapping.table_projection).unwrap()
1664        );
1665    }
1666
1667    #[test]
1668    fn test_csv_reader_schema_for_skip_bad_records_uses_positional_mapping() {
1669        let file_schema = make_test_schema(&[
1670            Field::new("column_1", DataType::Utf8, true),
1671            Field::new("column_2", DataType::Utf8, true),
1672            Field::new("column_3", DataType::Utf8, true),
1673        ]);
1674        let table_schema = make_test_schema(&[
1675            Field::new("host_id", DataType::UInt32, true),
1676            Field::new("jsons", DataType::Binary, true),
1677            Field::new(
1678                "ts",
1679                DataType::Timestamp(datatypes::arrow::datatypes::TimeUnit::Millisecond, None),
1680                true,
1681            ),
1682        ]);
1683        let mapping = copy_from_schema_mapping(
1684            &make_csv_metadata(file_schema.clone(), false),
1685            &table_schema,
1686        );
1687        let compat_schema = Arc::new(mapping.compat_file_schema);
1688
1689        let reader_schema = csv_reader_schema_for_skip_bad_records(&file_schema, &compat_schema);
1690
1691        assert_eq!(reader_schema.field(0).data_type(), &DataType::UInt32);
1692        assert_eq!(reader_schema.field(1).data_type(), &DataType::Utf8);
1693        assert_eq!(
1694            reader_schema.field(2).data_type(),
1695            table_schema.field(2).data_type()
1696        );
1697    }
1698}