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::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::{Entry, 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
95impl StatementExecutor {
96    async fn list_copy_from_entries(
97        &self,
98        req: &CopyTableRequest,
99    ) -> Result<(ObjectStore, Vec<Entry>)> {
100        let backend =
101            build_backend_with_path(&req.location, &req.connection, &self.local_file_access)
102                .await
103                .context(error::BuildBackendSnafu)?;
104        let regex = req
105            .pattern
106            .as_ref()
107            .map(|x| Regex::new(x))
108            .transpose()
109            .context(error::BuildRegexSnafu)?;
110
111        let source = if let Some(filename) = backend.object_path {
112            Source::Filename(filename)
113        } else {
114            Source::Dir
115        };
116
117        let lister = Lister::new(
118            backend.object_store.clone(),
119            source.clone(),
120            req.location.clone(),
121            regex,
122        );
123
124        let entries = lister.list().await.context(error::ListObjectsSnafu)?;
125        debug!(
126            "Copy from location: {:?}, {source:?}, entries: {entries:?}",
127            req.location
128        );
129        Ok((backend.object_store, entries))
130    }
131
132    async fn collect_metadata(
133        &self,
134        object_store: &ObjectStore,
135        format: Format,
136        path: String,
137    ) -> Result<FileMetadata> {
138        match format {
139            Format::Csv(format) => Ok(FileMetadata::Csv {
140                schema: Arc::new(
141                    format
142                        .infer_schema(object_store, &path)
143                        .await
144                        .context(error::InferSchemaSnafu { path: &path })?,
145                ),
146                format,
147                path,
148            }),
149            Format::Json(format) => Ok(FileMetadata::Json {
150                schema: Arc::new(
151                    format
152                        .infer_schema(object_store, &path)
153                        .await
154                        .context(error::InferSchemaSnafu { path: &path })?,
155                ),
156                format,
157                path,
158            }),
159            Format::Parquet(_) => {
160                let meta = object_store
161                    .stat(&path)
162                    .await
163                    .context(error::ReadObjectSnafu { path: &path })?;
164                let mut reader = object_store
165                    .reader(&path)
166                    .await
167                    .context(error::ReadObjectSnafu { path: &path })?
168                    .into_futures_async_read(0..meta.content_length())
169                    .await
170                    .context(error::ReadObjectSnafu { path: &path })?
171                    .compat();
172                let metadata = ArrowReaderMetadata::load_async(&mut reader, Default::default())
173                    .await
174                    .context(error::ReadParquetMetadataSnafu)?;
175
176                Ok(FileMetadata::Parquet {
177                    schema: metadata.schema().clone(),
178                    metadata,
179                    path,
180                })
181            }
182            Format::Orc(_) => {
183                let meta = object_store
184                    .stat(&path)
185                    .await
186                    .context(error::ReadObjectSnafu { path: &path })?;
187
188                let reader = object_store
189                    .reader(&path)
190                    .await
191                    .context(error::ReadObjectSnafu { path: &path })?;
192
193                let schema = infer_orc_schema(ReaderAdapter::new(reader, meta.content_length()))
194                    .await
195                    .context(error::ReadOrcSnafu)?;
196
197                Ok(FileMetadata::Orc {
198                    schema: Arc::new(schema),
199                    path,
200                })
201            }
202        }
203    }
204
205    async fn build_read_stream(
206        &self,
207        compat_schema: SchemaRef,
208        object_store: &ObjectStore,
209        file_metadata: &FileMetadata,
210        projection: Vec<usize>,
211        filters: Vec<Expr>,
212    ) -> Result<DfSendableRecordBatchStream> {
213        match file_metadata {
214            FileMetadata::Csv {
215                format,
216                path,
217                schema,
218            } => {
219                let output_schema = Arc::new(
220                    compat_schema
221                        .project(&projection)
222                        .context(error::ProjectSchemaSnafu)?,
223                );
224
225                let options = CsvOptions::default()
226                    .with_has_header(format.has_header)
227                    .with_delimiter(format.delimiter);
228                let csv_source = CsvSource::new(schema.clone())
229                    .with_csv_options(options)
230                    .with_batch_size(DEFAULT_BATCH_SIZE);
231                let stream = if format.skip_bad_records {
232                    let reader_schema =
233                        csv_reader_schema_for_skip_bad_records(schema, &compat_schema);
234                    tolerant_csv_stream(
235                        object_store,
236                        path,
237                        Arc::new(reader_schema),
238                        projection.clone(),
239                        format,
240                    )
241                    .await
242                    .context(error::BuildFileStreamSnafu)?
243                } else {
244                    file_to_stream(
245                        object_store,
246                        path,
247                        csv_source,
248                        Some(projection),
249                        format.compression_type,
250                    )
251                    .await
252                    .context(error::BuildFileStreamSnafu)?
253                };
254
255                let stream = Box::pin(
256                    // The projection is already applied in the CSV reader when we created the stream,
257                    // so we pass None here to avoid double projection which would cause schema mismatch errors.
258                    RecordBatchStreamTypeAdapter::new(output_schema, stream, None)
259                        .with_filter(filters)
260                        .context(error::PhysicalExprSnafu)?,
261                );
262                if format.skip_bad_records {
263                    Ok(Box::pin(SkipBadRecordsStream::new(stream, path)))
264                } else {
265                    Ok(stream)
266                }
267            }
268            FileMetadata::Json {
269                path,
270                format,
271                schema,
272            } => {
273                let output_schema = Arc::new(
274                    compat_schema
275                        .project(&projection)
276                        .context(error::ProjectSchemaSnafu)?,
277                );
278
279                let json_source =
280                    JsonSource::new(schema.clone()).with_batch_size(DEFAULT_BATCH_SIZE);
281                let stream = file_to_stream(
282                    object_store,
283                    path,
284                    json_source,
285                    Some(projection),
286                    format.compression_type,
287                )
288                .await
289                .context(error::BuildFileStreamSnafu)?;
290
291                Ok(Box::pin(
292                    // The projection is already applied in the JSON reader when we created the stream,
293                    // so we pass None here to avoid double projection which would cause schema mismatch errors.
294                    RecordBatchStreamTypeAdapter::new(output_schema, stream, None)
295                        .with_filter(filters)
296                        .context(error::PhysicalExprSnafu)?,
297                ))
298            }
299            FileMetadata::Parquet { metadata, path, .. } => {
300                let meta = object_store
301                    .stat(path)
302                    .await
303                    .context(error::ReadObjectSnafu { path })?;
304                let reader = object_store
305                    .reader_with(path)
306                    .chunk(DEFAULT_READ_BUFFER)
307                    .await
308                    .context(error::ReadObjectSnafu { path })?
309                    .into_futures_async_read(0..meta.content_length())
310                    .await
311                    .context(error::ReadObjectSnafu { path })?
312                    .compat();
313                let builder =
314                    ParquetRecordBatchStreamBuilder::new_with_metadata(reader, metadata.clone());
315                let stream = builder
316                    .build()
317                    .context(error::BuildParquetRecordBatchStreamSnafu)?;
318
319                let output_schema = Arc::new(
320                    compat_schema
321                        .project(&projection)
322                        .context(error::ProjectSchemaSnafu)?,
323                );
324                Ok(Box::pin(
325                    RecordBatchStreamTypeAdapter::new(output_schema, stream, Some(projection))
326                        .with_filter(filters)
327                        .context(error::PhysicalExprSnafu)?,
328                ))
329            }
330            FileMetadata::Orc { path, .. } => {
331                let meta = object_store
332                    .stat(path)
333                    .await
334                    .context(error::ReadObjectSnafu { path })?;
335
336                let reader = object_store
337                    .reader_with(path)
338                    .chunk(DEFAULT_READ_BUFFER)
339                    .await
340                    .context(error::ReadObjectSnafu { path })?;
341                let stream =
342                    new_orc_stream_reader(ReaderAdapter::new(reader, meta.content_length()))
343                        .await
344                        .context(error::ReadOrcSnafu)?;
345
346                let output_schema = Arc::new(
347                    compat_schema
348                        .project(&projection)
349                        .context(error::ProjectSchemaSnafu)?,
350                );
351
352                Ok(Box::pin(
353                    RecordBatchStreamTypeAdapter::new(output_schema, stream, Some(projection))
354                        .with_filter(filters)
355                        .context(error::PhysicalExprSnafu)?,
356                ))
357            }
358        }
359    }
360
361    #[tracing::instrument(skip_all)]
362    pub async fn copy_table_from(
363        &self,
364        req: CopyTableRequest,
365        query_ctx: QueryContextRef,
366    ) -> Result<Output> {
367        let table_ref = TableReference {
368            catalog: &req.catalog_name,
369            schema: &req.schema_name,
370            table: &req.table_name,
371        };
372        let table = self.get_table(&table_ref).await?;
373        let format = Format::try_from(&req.with).context(error::ParseFileFormatSnafu)?;
374        let (object_store, entries) = self.list_copy_from_entries(&req).await?;
375        let mut files = Vec::with_capacity(entries.len());
376        let table_schema = table.schema().arrow_schema().clone();
377        let filters = table
378            .schema()
379            .timestamp_column()
380            .and_then(|c| {
381                common_query::logical_plan::build_same_type_ts_filter(c, req.timestamp_range)
382            })
383            .into_iter()
384            .collect::<Vec<_>>();
385
386        for entry in entries.iter() {
387            if entry.metadata().mode() != EntryMode::FILE {
388                continue;
389            }
390            let path = entry.path();
391            let file_metadata = self
392                .collect_metadata(&object_store, format.clone(), path.to_string())
393                .await?;
394
395            validate_csv_headers_if_required(&file_metadata, &table_schema)?;
396            let schema_mapping = copy_from_schema_mapping(&file_metadata, &table_schema);
397            let projected_file_schema = Arc::new(
398                file_metadata
399                    .schema()
400                    .project(&schema_mapping.file_projection)
401                    .context(error::ProjectSchemaSnafu)?,
402            );
403            let projected_table_schema = Arc::new(
404                table_schema
405                    .project(&schema_mapping.table_projection)
406                    .context(error::ProjectSchemaSnafu)?,
407            );
408            ensure_schema_compatible(&projected_file_schema, &projected_table_schema)?;
409
410            files.push((
411                Arc::new(schema_mapping.compat_file_schema),
412                schema_mapping.file_projection,
413                projected_table_schema,
414                file_metadata,
415            ))
416        }
417
418        let mut rows_inserted = 0;
419        let mut insert_cost = 0;
420        let max_insert_rows = req
421            .limit
422            .map(|n| {
423                usize::try_from(n).map_err(|_| {
424                    error::InvalidCopyParameterSnafu {
425                        key: "limit".to_string(),
426                        value: n.to_string(),
427                    }
428                    .build()
429                })
430            })
431            .transpose()?;
432        if max_insert_rows == Some(0) {
433            return Ok(gen_insert_output(rows_inserted, insert_cost));
434        }
435
436        let mut accepted_rows = 0;
437        for (compat_schema, file_schema_projection, projected_table_schema, file_metadata) in files
438        {
439            let mut stream = self
440                .build_read_stream(
441                    compat_schema,
442                    &object_store,
443                    &file_metadata,
444                    file_schema_projection,
445                    filters.clone(),
446                )
447                .await?;
448
449            let fields = projected_table_schema
450                .fields()
451                .iter()
452                .map(|f| f.name().clone())
453                .collect::<Vec<_>>();
454
455            // TODO(hl): make this configurable through options.
456            let pending_mem_threshold = ReadableSize::mb(32).as_bytes();
457            let mut pending_mem_size = 0;
458            let mut pending = vec![];
459
460            while let Some(r) = stream.next().await {
461                let record_batch = r.context(error::ReadDfRecordBatchSnafu)?;
462                let record_batch = if let Some(max_insert_rows) = max_insert_rows {
463                    let remaining_rows = max_insert_rows - accepted_rows;
464                    if record_batch.num_rows() > remaining_rows {
465                        record_batch.slice(0, remaining_rows)
466                    } else {
467                        record_batch
468                    }
469                } else {
470                    record_batch
471                };
472                let record_batch_rows = record_batch.num_rows();
473                let vectors =
474                    Helper::try_into_vectors(record_batch.columns()).context(IntoVectorsSnafu)?;
475
476                pending_mem_size += vectors.iter().map(|v| v.memory_size()).sum::<usize>();
477
478                let columns_values = fields
479                    .iter()
480                    .cloned()
481                    .zip(vectors)
482                    .collect::<HashMap<_, _>>();
483
484                pending.push(self.inserter.handle_table_insert(
485                    InsertRequest {
486                        catalog_name: req.catalog_name.clone(),
487                        schema_name: req.schema_name.clone(),
488                        table_name: req.table_name.clone(),
489                        columns_values,
490                    },
491                    query_ctx.clone(),
492                ));
493                accepted_rows += record_batch_rows;
494
495                if pending_mem_size as u64 >= pending_mem_threshold {
496                    let (rows, cost) = batch_insert(&mut pending, &mut pending_mem_size).await?;
497                    rows_inserted += rows;
498                    insert_cost += cost;
499                }
500
501                if let Some(max_insert_rows) = max_insert_rows
502                    && accepted_rows == max_insert_rows
503                {
504                    if !pending.is_empty() {
505                        let (rows, cost) =
506                            batch_insert(&mut pending, &mut pending_mem_size).await?;
507                        rows_inserted += rows;
508                        insert_cost += cost;
509                    }
510                    return Ok(gen_insert_output(rows_inserted, insert_cost));
511                }
512            }
513
514            if !pending.is_empty() {
515                let (rows, cost) = batch_insert(&mut pending, &mut pending_mem_size).await?;
516                rows_inserted += rows;
517                insert_cost += cost;
518            }
519        }
520
521        Ok(gen_insert_output(rows_inserted, insert_cost))
522    }
523}
524
525fn gen_insert_output(rows_inserted: usize, insert_cost: usize) -> Output {
526    Output::new(
527        OutputData::AffectedRows(rows_inserted),
528        OutputMeta::new_with_cost(insert_cost),
529    )
530}
531
532struct SkipBadRecordsStream {
533    inner: DfSendableRecordBatchStream,
534    path: String,
535}
536
537impl SkipBadRecordsStream {
538    fn new(inner: DfSendableRecordBatchStream, path: impl Into<String>) -> Self {
539        Self {
540            inner,
541            path: path.into(),
542        }
543    }
544}
545
546impl datafusion::physical_plan::RecordBatchStream for SkipBadRecordsStream {
547    fn schema(&self) -> SchemaRef {
548        self.inner.schema()
549    }
550}
551
552impl futures::Stream for SkipBadRecordsStream {
553    type Item = datafusion_common::Result<RecordBatch>;
554
555    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
556        let this = self.get_mut();
557        loop {
558            match this.inner.as_mut().poll_next(cx) {
559                Poll::Ready(Some(Err(error))) if is_skippable_record_error(&error) => {
560                    common_telemetry::warn!(
561                        "Skipping bad record while copying from {}: {}",
562                        this.path,
563                        error
564                    );
565                    continue;
566                }
567                other => return other,
568            }
569        }
570    }
571}
572
573fn is_skippable_record_error(error: &DataFusionError) -> bool {
574    match error {
575        DataFusionError::ArrowError(error, _) => is_skippable_arrow_error(error),
576        DataFusionError::External(error) => error
577            .downcast_ref::<ArrowError>()
578            .is_some_and(is_skippable_arrow_error),
579        DataFusionError::Context(_, error) => is_skippable_record_error(error),
580        _ => false,
581    }
582}
583
584/// Executes all pending inserts all at once, drain pending requests and reset pending bytes.
585async fn batch_insert(
586    pending: &mut Vec<impl Future<Output = Result<Output>>>,
587    pending_bytes: &mut usize,
588) -> Result<(OutputRows, OutputCost)> {
589    let batch = pending.drain(..);
590    let result = futures::future::try_join_all(batch)
591        .await?
592        .iter()
593        .map(|o| o.extract_rows_and_cost())
594        .reduce(|(a, b), (c, d)| (a + c, b + d))
595        .unwrap_or((0, 0));
596    *pending_bytes = 0;
597    Ok(result)
598}
599
600/// Custom type compatibility check for GreptimeDB that handles Map -> Binary (JSON) conversion
601fn can_cast_types_for_greptime(from: &ArrowDataType, to: &ArrowDataType) -> bool {
602    // Handle Map -> Binary conversion for JSON types
603    if let ArrowDataType::Map(_, _) = from
604        && let ArrowDataType::Binary = to
605    {
606        return true;
607    }
608
609    // For all other cases, use Arrow's built-in can_cast_types
610    can_cast_types(from, to)
611}
612
613fn csv_reader_schema_for_skip_bad_records(file: &SchemaRef, compat: &SchemaRef) -> Schema {
614    let fields = file
615        .fields()
616        .iter()
617        .enumerate()
618        .map(|(idx, file_field)| match compat.fields().get(idx) {
619            Some(compat_field) if can_csv_reader_parse_type(compat_field.data_type()) => {
620                compat_field.clone()
621            }
622            _ => file_field.clone(),
623        })
624        .collect::<Vec<_>>();
625
626    Schema::new_with_metadata(fields, file.metadata().clone())
627}
628
629fn can_csv_reader_parse_type(data_type: &ArrowDataType) -> bool {
630    match data_type {
631        ArrowDataType::Boolean
632        | ArrowDataType::Decimal32(_, _)
633        | ArrowDataType::Decimal64(_, _)
634        | ArrowDataType::Decimal128(_, _)
635        | ArrowDataType::Decimal256(_, _)
636        | ArrowDataType::Int8
637        | ArrowDataType::Int16
638        | ArrowDataType::Int32
639        | ArrowDataType::Int64
640        | ArrowDataType::UInt8
641        | ArrowDataType::UInt16
642        | ArrowDataType::UInt32
643        | ArrowDataType::UInt64
644        | ArrowDataType::Float32
645        | ArrowDataType::Float64
646        | ArrowDataType::Date32
647        | ArrowDataType::Date64
648        | ArrowDataType::Time32(_)
649        | ArrowDataType::Time64(_)
650        | ArrowDataType::Timestamp(_, _)
651        | ArrowDataType::Null
652        | ArrowDataType::Utf8
653        | ArrowDataType::Utf8View => true,
654        ArrowDataType::Dictionary(_, value_type) => value_type.as_ref() == &ArrowDataType::Utf8,
655        _ => false,
656    }
657}
658
659fn ensure_schema_compatible(from: &SchemaRef, to: &SchemaRef) -> Result<()> {
660    let not_match = from
661        .fields
662        .iter()
663        .zip(to.fields.iter())
664        .map(|(l, r)| (l.data_type(), r.data_type()))
665        .enumerate()
666        .find(|(_, (l, r))| !can_cast_types_for_greptime(l, r));
667
668    if let Some((index, _)) = not_match {
669        error::InvalidSchemaSnafu {
670            index,
671            table_schema: to.to_string(),
672            file_schema: from.to_string(),
673        }
674        .fail()
675    } else {
676        Ok(())
677    }
678}
679
680fn validate_csv_headers_if_required(file_metadata: &FileMetadata, table: &SchemaRef) -> Result<()> {
681    let FileMetadata::Csv {
682        schema,
683        format,
684        path,
685    } = file_metadata
686    else {
687        return Ok(());
688    };
689
690    if !format.strict_headers {
691        return Ok(());
692    }
693
694    let mut seen_file_columns = HashSet::with_capacity(schema.fields().len());
695    let duplicate_columns = schema
696        .fields()
697        .iter()
698        .filter_map(|field| {
699            if seen_file_columns.insert(field.name().clone()) {
700                None
701            } else {
702                Some(field.name().clone())
703            }
704        })
705        .collect::<BTreeSet<_>>()
706        .into_iter()
707        .collect::<Vec<_>>();
708    let file_columns = seen_file_columns.into_iter().collect::<BTreeSet<_>>();
709    let table_columns = table
710        .fields()
711        .iter()
712        .map(|field| field.name().clone())
713        .collect::<BTreeSet<_>>();
714    let unknown_columns = file_columns
715        .difference(&table_columns)
716        .cloned()
717        .collect::<Vec<_>>();
718    let missing_columns = table_columns
719        .difference(&file_columns)
720        .cloned()
721        .collect::<Vec<_>>();
722
723    ensure!(
724        unknown_columns.is_empty() && missing_columns.is_empty() && duplicate_columns.is_empty(),
725        error::CsvHeaderMismatchSnafu {
726            path,
727            unknown_columns,
728            missing_columns,
729            duplicate_columns,
730        }
731    );
732
733    Ok(())
734}
735
736/// Generates a maybe compatible schema of the file schema.
737///
738/// If there is a field is found in table schema,
739/// copy the field data type to maybe compatible schema(`compatible_fields`).
740fn generated_schema_projection_and_compatible_file_schema(
741    file: &SchemaRef,
742    table: &SchemaRef,
743) -> (Vec<usize>, Vec<usize>, Schema) {
744    let mut file_projection = Vec::with_capacity(file.fields.len());
745    let mut table_projection = Vec::with_capacity(file.fields.len());
746    let mut compatible_fields = file.fields.iter().cloned().collect::<Vec<_>>();
747    for (file_idx, file_field) in file.fields.iter().enumerate() {
748        if let Some((table_idx, table_field)) = table.fields.find(file_field.name()) {
749            file_projection.push(file_idx);
750            table_projection.push(table_idx);
751
752            // Safety: the compatible_fields has same length as file schema
753            compatible_fields[file_idx] = table_field.clone();
754        }
755    }
756
757    (
758        file_projection,
759        table_projection,
760        Schema::new(compatible_fields),
761    )
762}
763
764struct CopyFromSchemaMapping {
765    file_projection: Vec<usize>,
766    table_projection: Vec<usize>,
767    compat_file_schema: Schema,
768}
769
770fn copy_from_schema_mapping(
771    file_metadata: &FileMetadata,
772    table: &SchemaRef,
773) -> CopyFromSchemaMapping {
774    match file_metadata {
775        FileMetadata::Csv { schema, format, .. } if !format.has_header => {
776            generated_positional_schema_projection_and_compatible_file_schema(schema, table)
777        }
778        _ => {
779            let (file_projection, table_projection, compat_file_schema) =
780                generated_schema_projection_and_compatible_file_schema(
781                    file_metadata.schema(),
782                    table,
783                );
784            CopyFromSchemaMapping {
785                file_projection,
786                table_projection,
787                compat_file_schema,
788            }
789        }
790    }
791}
792
793fn generated_positional_schema_projection_and_compatible_file_schema(
794    file: &SchemaRef,
795    table: &SchemaRef,
796) -> CopyFromSchemaMapping {
797    let len = file.fields.len().min(table.fields.len());
798    let file_projection = (0..len).collect::<Vec<_>>();
799    let table_projection = (0..len).collect::<Vec<_>>();
800    let compatible_fields = file
801        .fields
802        .iter()
803        .enumerate()
804        .map(|(idx, file_field)| {
805            if idx < len {
806                table.fields[idx].clone()
807            } else {
808                file_field.clone()
809            }
810        })
811        .collect::<Vec<_>>();
812
813    CopyFromSchemaMapping {
814        file_projection,
815        table_projection,
816        compat_file_schema: Schema::new(compatible_fields),
817    }
818}
819
820#[cfg(test)]
821mod tests {
822    use std::sync::Arc;
823
824    use datatypes::arrow::datatypes::{DataType, Field, Schema};
825
826    use super::*;
827
828    fn test_schema_matches(from: (DataType, bool), to: (DataType, bool), matches: bool) {
829        let s1 = Arc::new(Schema::new(vec![Field::new("col", from.0.clone(), from.1)]));
830        let s2 = Arc::new(Schema::new(vec![Field::new("col", to.0.clone(), to.1)]));
831        let res = ensure_schema_compatible(&s1, &s2);
832        assert_eq!(
833            matches,
834            res.is_ok(),
835            "from data type: {}, to data type: {}, expected: {}, but got: {}",
836            from.0,
837            to.0,
838            matches,
839            res.is_ok()
840        )
841    }
842
843    #[test]
844    fn test_ensure_datatype_matches_ignore_timezone() {
845        test_schema_matches(
846            (
847                DataType::Timestamp(datatypes::arrow::datatypes::TimeUnit::Second, None),
848                true,
849            ),
850            (
851                DataType::Timestamp(datatypes::arrow::datatypes::TimeUnit::Second, None),
852                true,
853            ),
854            true,
855        );
856
857        test_schema_matches(
858            (
859                DataType::Timestamp(
860                    datatypes::arrow::datatypes::TimeUnit::Second,
861                    Some("UTC".into()),
862                ),
863                true,
864            ),
865            (
866                DataType::Timestamp(datatypes::arrow::datatypes::TimeUnit::Second, None),
867                true,
868            ),
869            true,
870        );
871
872        test_schema_matches(
873            (
874                DataType::Timestamp(
875                    datatypes::arrow::datatypes::TimeUnit::Second,
876                    Some("UTC".into()),
877                ),
878                true,
879            ),
880            (
881                DataType::Timestamp(
882                    datatypes::arrow::datatypes::TimeUnit::Second,
883                    Some("PDT".into()),
884                ),
885                true,
886            ),
887            true,
888        );
889
890        test_schema_matches(
891            (
892                DataType::Timestamp(
893                    datatypes::arrow::datatypes::TimeUnit::Second,
894                    Some("UTC".into()),
895                ),
896                true,
897            ),
898            (
899                DataType::Timestamp(
900                    datatypes::arrow::datatypes::TimeUnit::Millisecond,
901                    Some("UTC".into()),
902                ),
903                true,
904            ),
905            true,
906        );
907
908        test_schema_matches((DataType::Int8, true), (DataType::Int8, true), true);
909
910        test_schema_matches((DataType::Int8, true), (DataType::Int16, true), true);
911    }
912
913    #[test]
914    fn test_data_type_equals_ignore_timezone_with_options() {
915        test_schema_matches(
916            (
917                DataType::Timestamp(
918                    datatypes::arrow::datatypes::TimeUnit::Microsecond,
919                    Some("UTC".into()),
920                ),
921                true,
922            ),
923            (
924                DataType::Timestamp(
925                    datatypes::arrow::datatypes::TimeUnit::Millisecond,
926                    Some("PDT".into()),
927                ),
928                true,
929            ),
930            true,
931        );
932
933        test_schema_matches(
934            (DataType::Utf8, true),
935            (
936                DataType::Timestamp(
937                    datatypes::arrow::datatypes::TimeUnit::Millisecond,
938                    Some("PDT".into()),
939                ),
940                true,
941            ),
942            true,
943        );
944
945        test_schema_matches(
946            (
947                DataType::Timestamp(
948                    datatypes::arrow::datatypes::TimeUnit::Millisecond,
949                    Some("PDT".into()),
950                ),
951                true,
952            ),
953            (DataType::Utf8, true),
954            true,
955        );
956    }
957
958    #[test]
959    fn test_map_to_binary_json_compatibility() {
960        // Test Map -> Binary conversion for JSON types
961        let map_type = DataType::Map(
962            Arc::new(Field::new(
963                "key_value",
964                DataType::Struct(
965                    vec![
966                        Field::new("key", DataType::Utf8, false),
967                        Field::new("value", DataType::Utf8, false),
968                    ]
969                    .into(),
970                ),
971                false,
972            )),
973            false,
974        );
975
976        test_schema_matches((map_type, false), (DataType::Binary, true), true);
977
978        test_schema_matches((DataType::Int8, true), (DataType::Int16, true), true);
979        test_schema_matches((DataType::Utf8, true), (DataType::Binary, true), true);
980    }
981
982    fn make_test_schema(v: &[Field]) -> Arc<Schema> {
983        Arc::new(Schema::new(v.to_vec()))
984    }
985
986    #[test]
987    fn test_compatible_file_schema() {
988        let file_schema0 = make_test_schema(&[
989            Field::new("c1", DataType::UInt8, true),
990            Field::new("c2", DataType::UInt8, true),
991        ]);
992
993        let table_schema = make_test_schema(&[
994            Field::new("c1", DataType::Int16, true),
995            Field::new("c2", DataType::Int16, true),
996            Field::new("c3", DataType::Int16, true),
997        ]);
998
999        let compat_schema = make_test_schema(&[
1000            Field::new("c1", DataType::Int16, true),
1001            Field::new("c2", DataType::Int16, true),
1002        ]);
1003
1004        let (_, tp, _) =
1005            generated_schema_projection_and_compatible_file_schema(&file_schema0, &table_schema);
1006
1007        assert_eq!(table_schema.project(&tp).unwrap(), *compat_schema);
1008    }
1009
1010    #[test]
1011    fn test_schema_projection() {
1012        let file_schema0 = make_test_schema(&[
1013            Field::new("c1", DataType::UInt8, true),
1014            Field::new("c2", DataType::UInt8, true),
1015            Field::new("c3", DataType::UInt8, true),
1016        ]);
1017
1018        let file_schema1 = make_test_schema(&[
1019            Field::new("c3", DataType::UInt8, true),
1020            Field::new("c4", DataType::UInt8, true),
1021        ]);
1022
1023        let file_schema2 = make_test_schema(&[
1024            Field::new("c3", DataType::UInt8, true),
1025            Field::new("c4", DataType::UInt8, true),
1026            Field::new("c5", DataType::UInt8, true),
1027        ]);
1028
1029        let file_schema3 = make_test_schema(&[
1030            Field::new("c1", DataType::UInt8, true),
1031            Field::new("c2", DataType::UInt8, true),
1032        ]);
1033
1034        let table_schema = make_test_schema(&[
1035            Field::new("c3", DataType::UInt8, true),
1036            Field::new("c4", DataType::UInt8, true),
1037            Field::new("c5", DataType::UInt8, true),
1038        ]);
1039
1040        let tests = [
1041            (&file_schema0, &table_schema, true), // intersection
1042            (&file_schema1, &table_schema, true), // subset
1043            (&file_schema2, &table_schema, true), // full-eq
1044            (&file_schema3, &table_schema, true), // non-intersection
1045        ];
1046
1047        for test in tests {
1048            let (fp, tp, _) =
1049                generated_schema_projection_and_compatible_file_schema(test.0, test.1);
1050            assert_eq!(test.0.project(&fp).unwrap(), test.1.project(&tp).unwrap());
1051        }
1052    }
1053
1054    #[test]
1055    fn test_csv_reader_schema_for_skip_bad_records() {
1056        let file_schema = make_test_schema(&[
1057            Field::new("id", DataType::Utf8, true),
1058            Field::new("jsons", DataType::Utf8, true),
1059            Field::new("ts", DataType::Utf8, true),
1060        ]);
1061        let compat_schema = make_test_schema(&[
1062            Field::new("id", DataType::UInt32, true),
1063            Field::new("jsons", DataType::Binary, true),
1064            Field::new(
1065                "ts",
1066                DataType::Timestamp(datatypes::arrow::datatypes::TimeUnit::Millisecond, None),
1067                true,
1068            ),
1069        ]);
1070
1071        let reader_schema = csv_reader_schema_for_skip_bad_records(&file_schema, &compat_schema);
1072
1073        assert_eq!(reader_schema.field(0).data_type(), &DataType::UInt32);
1074        assert_eq!(reader_schema.field(1).data_type(), &DataType::Utf8);
1075        assert_eq!(
1076            reader_schema.field(2).data_type(),
1077            compat_schema.field(2).data_type()
1078        );
1079    }
1080
1081    fn make_csv_metadata(schema: Arc<Schema>, has_header: bool) -> FileMetadata {
1082        FileMetadata::Csv {
1083            schema,
1084            format: CsvFormat {
1085                has_header,
1086                ..CsvFormat::default()
1087            },
1088            path: "test.csv".to_string(),
1089        }
1090    }
1091
1092    fn make_strict_csv_metadata(schema: Arc<Schema>) -> FileMetadata {
1093        FileMetadata::Csv {
1094            schema,
1095            format: CsvFormat {
1096                strict_headers: true,
1097                ..CsvFormat::default()
1098            },
1099            path: "test.csv".to_string(),
1100        }
1101    }
1102
1103    fn assert_field(schema: &Schema, idx: usize, name: &str, data_type: &DataType) {
1104        let field = schema.field(idx);
1105        assert_eq!(field.name(), name);
1106        assert_eq!(field.data_type(), data_type);
1107    }
1108
1109    #[test]
1110    fn test_strict_csv_headers_allows_reordered_columns() {
1111        let file_schema = make_test_schema(&[
1112            Field::new("ts", DataType::Utf8, true),
1113            Field::new("host_id", DataType::UInt8, true),
1114            Field::new("reading_value", DataType::Float64, true),
1115        ]);
1116        let table_schema = make_test_schema(&[
1117            Field::new("host_id", DataType::UInt32, true),
1118            Field::new("reading_value", DataType::Float64, true),
1119            Field::new("ts", DataType::Utf8, true),
1120        ]);
1121
1122        validate_csv_headers_if_required(&make_strict_csv_metadata(file_schema), &table_schema)
1123            .unwrap();
1124    }
1125
1126    #[test]
1127    fn test_strict_csv_headers_rejects_unknown_columns() {
1128        let file_schema = make_test_schema(&[
1129            Field::new("host_id", DataType::UInt8, true),
1130            Field::new("reading_value", DataType::Float64, true),
1131            Field::new("ts", DataType::Utf8, true),
1132            Field::new("extra", DataType::Utf8, true),
1133        ]);
1134        let table_schema = make_test_schema(&[
1135            Field::new("host_id", DataType::UInt32, true),
1136            Field::new("reading_value", DataType::Float64, true),
1137            Field::new("ts", DataType::Utf8, true),
1138        ]);
1139
1140        let err =
1141            validate_csv_headers_if_required(&make_strict_csv_metadata(file_schema), &table_schema)
1142                .unwrap_err();
1143
1144        assert!(matches!(
1145            err,
1146            error::Error::CsvHeaderMismatch {
1147                unknown_columns,
1148                missing_columns,
1149                duplicate_columns,
1150                ..
1151            } if unknown_columns == vec!["extra".to_string()]
1152                && missing_columns.is_empty()
1153                && duplicate_columns.is_empty()
1154        ));
1155    }
1156
1157    #[test]
1158    fn test_strict_csv_headers_rejects_missing_columns() {
1159        let file_schema = make_test_schema(&[
1160            Field::new("host_id", DataType::UInt8, true),
1161            Field::new("ts", DataType::Utf8, true),
1162        ]);
1163        let table_schema = make_test_schema(&[
1164            Field::new("host_id", DataType::UInt32, true),
1165            Field::new("reading_value", DataType::Float64, true),
1166            Field::new("ts", DataType::Utf8, true),
1167        ]);
1168
1169        let err =
1170            validate_csv_headers_if_required(&make_strict_csv_metadata(file_schema), &table_schema)
1171                .unwrap_err();
1172
1173        assert!(matches!(
1174            err,
1175            error::Error::CsvHeaderMismatch {
1176                unknown_columns,
1177                missing_columns,
1178                duplicate_columns,
1179                ..
1180            } if unknown_columns.is_empty()
1181                && missing_columns == vec!["reading_value".to_string()]
1182                && duplicate_columns.is_empty()
1183        ));
1184    }
1185
1186    #[test]
1187    fn test_strict_csv_headers_rejects_duplicate_columns() {
1188        let file_schema = make_test_schema(&[
1189            Field::new("host_id", DataType::UInt8, true),
1190            Field::new("reading_value", DataType::Float64, true),
1191            Field::new("ts", DataType::Utf8, true),
1192            Field::new("host_id", DataType::UInt16, true),
1193        ]);
1194        let table_schema = make_test_schema(&[
1195            Field::new("host_id", DataType::UInt32, true),
1196            Field::new("reading_value", DataType::Float64, true),
1197            Field::new("ts", DataType::Utf8, true),
1198        ]);
1199
1200        let err =
1201            validate_csv_headers_if_required(&make_strict_csv_metadata(file_schema), &table_schema)
1202                .unwrap_err();
1203
1204        assert!(matches!(
1205            err,
1206            error::Error::CsvHeaderMismatch {
1207                unknown_columns,
1208                missing_columns,
1209                duplicate_columns,
1210                ..
1211            } if unknown_columns.is_empty()
1212                && missing_columns.is_empty()
1213                && duplicate_columns == vec!["host_id".to_string()]
1214        ));
1215    }
1216
1217    #[test]
1218    fn test_headerless_csv_schema_projection_is_positional() {
1219        let file_schema = make_test_schema(&[
1220            Field::new("column_1", DataType::UInt8, true),
1221            Field::new("column_2", DataType::Float64, true),
1222            Field::new("column_3", DataType::Utf8, true),
1223        ]);
1224        let table_schema = make_test_schema(&[
1225            Field::new("host_id", DataType::UInt32, true),
1226            Field::new("reading_value", DataType::Float64, true),
1227            Field::new(
1228                "ts",
1229                DataType::Timestamp(datatypes::arrow::datatypes::TimeUnit::Millisecond, None),
1230                true,
1231            ),
1232        ]);
1233
1234        let mapping =
1235            copy_from_schema_mapping(&make_csv_metadata(file_schema, false), &table_schema);
1236
1237        assert_eq!(mapping.file_projection, vec![0, 1, 2]);
1238        assert_eq!(mapping.table_projection, vec![0, 1, 2]);
1239        assert_field(&mapping.compat_file_schema, 0, "host_id", &DataType::UInt32);
1240        assert_field(
1241            &mapping.compat_file_schema,
1242            1,
1243            "reading_value",
1244            &DataType::Float64,
1245        );
1246        assert_field(
1247            &mapping.compat_file_schema,
1248            2,
1249            "ts",
1250            table_schema.field(2).data_type(),
1251        );
1252        assert_eq!(
1253            mapping
1254                .compat_file_schema
1255                .project(&mapping.file_projection)
1256                .unwrap(),
1257            table_schema.project(&mapping.table_projection).unwrap()
1258        );
1259    }
1260
1261    #[test]
1262    fn test_headerless_csv_schema_projection_ignores_extra_file_columns() {
1263        let file_schema = make_test_schema(&[
1264            Field::new("column_1", DataType::UInt8, true),
1265            Field::new("column_2", DataType::Float64, true),
1266            Field::new("column_3", DataType::Utf8, true),
1267            Field::new("column_4", DataType::Utf8, true),
1268        ]);
1269        let table_schema = make_test_schema(&[
1270            Field::new("host_id", DataType::UInt32, true),
1271            Field::new("reading_value", DataType::Float64, true),
1272            Field::new("ts", DataType::Utf8, true),
1273        ]);
1274
1275        let mapping =
1276            copy_from_schema_mapping(&make_csv_metadata(file_schema, false), &table_schema);
1277
1278        assert_eq!(mapping.file_projection, vec![0, 1, 2]);
1279        assert_eq!(mapping.table_projection, vec![0, 1, 2]);
1280        assert_eq!(mapping.compat_file_schema.fields().len(), 4);
1281        assert_field(&mapping.compat_file_schema, 0, "host_id", &DataType::UInt32);
1282        assert_field(
1283            &mapping.compat_file_schema,
1284            1,
1285            "reading_value",
1286            &DataType::Float64,
1287        );
1288        assert_field(&mapping.compat_file_schema, 2, "ts", &DataType::Utf8);
1289        assert_field(&mapping.compat_file_schema, 3, "column_4", &DataType::Utf8);
1290    }
1291
1292    #[test]
1293    fn test_headerless_csv_schema_projection_supports_prefix_import() {
1294        let file_schema = make_test_schema(&[
1295            Field::new("column_1", DataType::UInt8, true),
1296            Field::new("column_2", DataType::Float64, true),
1297        ]);
1298        let table_schema = make_test_schema(&[
1299            Field::new("host_id", DataType::UInt32, true),
1300            Field::new("reading_value", DataType::Float64, true),
1301            Field::new("ts", DataType::Utf8, true),
1302        ]);
1303
1304        let mapping =
1305            copy_from_schema_mapping(&make_csv_metadata(file_schema, false), &table_schema);
1306
1307        assert_eq!(mapping.file_projection, vec![0, 1]);
1308        assert_eq!(mapping.table_projection, vec![0, 1]);
1309        assert_field(&mapping.compat_file_schema, 0, "host_id", &DataType::UInt32);
1310        assert_field(
1311            &mapping.compat_file_schema,
1312            1,
1313            "reading_value",
1314            &DataType::Float64,
1315        );
1316        assert_eq!(
1317            mapping
1318                .compat_file_schema
1319                .project(&mapping.file_projection)
1320                .unwrap(),
1321            table_schema.project(&mapping.table_projection).unwrap()
1322        );
1323    }
1324
1325    #[test]
1326    fn test_csv_reader_schema_for_skip_bad_records_uses_positional_mapping() {
1327        let file_schema = make_test_schema(&[
1328            Field::new("column_1", DataType::Utf8, true),
1329            Field::new("column_2", DataType::Utf8, true),
1330            Field::new("column_3", DataType::Utf8, true),
1331        ]);
1332        let table_schema = make_test_schema(&[
1333            Field::new("host_id", DataType::UInt32, true),
1334            Field::new("jsons", DataType::Binary, true),
1335            Field::new(
1336                "ts",
1337                DataType::Timestamp(datatypes::arrow::datatypes::TimeUnit::Millisecond, None),
1338                true,
1339            ),
1340        ]);
1341        let mapping = copy_from_schema_mapping(
1342            &make_csv_metadata(file_schema.clone(), false),
1343            &table_schema,
1344        );
1345        let compat_schema = Arc::new(mapping.compat_file_schema);
1346
1347        let reader_schema = csv_reader_schema_for_skip_bad_records(&file_schema, &compat_schema);
1348
1349        assert_eq!(reader_schema.field(0).data_type(), &DataType::UInt32);
1350        assert_eq!(reader_schema.field(1).data_type(), &DataType::Utf8);
1351        assert_eq!(
1352            reader_schema.field(2).data_type(),
1353            table_schema.field(2).data_type()
1354        );
1355    }
1356}