Skip to main content

operator/statement/
copy_database.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::str::FromStr;
16use std::sync::Arc;
17
18use client::{Output, OutputData, OutputMeta};
19use common_catalog::format_full_table_name;
20use common_datasource::file_format::Format;
21use common_datasource::lister::{Lister, Source};
22use common_datasource::object_store::{LocalFileAccess, build_backend, build_backend_for_write};
23use common_telemetry::{debug, error, info, tracing};
24use futures::future::try_join_all;
25use object_store::Entry;
26use regex::Regex;
27use session::context::QueryContextRef;
28use snafu::ResultExt;
29use table::requests::{CopyDatabaseRequest, CopyDirection, CopyTableRequest};
30use tokio::sync::Semaphore;
31
32use crate::error;
33use crate::statement::StatementExecutor;
34use crate::statement::database_copy::{
35    DatabaseExportFile, database_import_source, parse_parallelism_from_option_map,
36    validate_database_directory, validate_database_export_layout,
37};
38
39pub(crate) const COPY_DATABASE_TIME_START_KEY: &str = "start_time";
40pub(crate) const COPY_DATABASE_TIME_END_KEY: &str = "end_time";
41pub(crate) const CONTINUE_ON_ERROR_KEY: &str = "continue_on_error";
42
43impl StatementExecutor {
44    #[tracing::instrument(skip_all)]
45    pub(crate) async fn copy_database_to(
46        &self,
47        req: CopyDatabaseRequest,
48        ctx: QueryContextRef,
49    ) -> error::Result<Output> {
50        validate_database_export_layout(&req.with)?;
51        validate_database_directory(&req.location)?;
52        build_backend_for_write(&req.location, &req.connection, &self.local_file_access)
53            .await
54            .context(error::BuildBackendSnafu)?;
55
56        let parallelism = parse_parallelism_from_option_map(&req.with);
57        info!(
58            "Copy database {}.{} to dir: {}, time: {:?}, parallelism: {}",
59            req.catalog_name, req.schema_name, req.location, req.time_range, parallelism
60        );
61        let tables = self
62            .capture_database_export_tables(&req, None, &ctx)
63            .await?;
64        let num_tables = tables.len();
65
66        let suffix = Format::try_from(&req.with)
67            .context(error::ParseFileFormatSnafu)?
68            .suffix();
69
70        let mut tasks = Vec::with_capacity(num_tables);
71        let semaphore = Arc::new(Semaphore::new(parallelism));
72
73        for (i, table) in tables.into_iter().enumerate() {
74            let table_name = table.table_info().name.clone();
75            let semaphore_moved = semaphore.clone();
76            let table_file = DatabaseExportFile::new(&req.location, &table_name, suffix)?.location;
77            let table_no = i + 1;
78            let moved_ctx = ctx.clone();
79            let full_table_name =
80                format_full_table_name(&req.catalog_name, &req.schema_name, &table_name);
81            let copy_table_req = CopyTableRequest {
82                catalog_name: req.catalog_name.clone(),
83                schema_name: req.schema_name.clone(),
84                table_name,
85                location: table_file.clone(),
86                with: req.with.clone(),
87                connection: req.connection.clone(),
88                pattern: None,
89                direction: CopyDirection::Export,
90                timestamp_range: req.time_range,
91                limit: None,
92            };
93
94            tasks.push(async move {
95                let _permit = semaphore_moved.acquire().await.unwrap();
96                info!(
97                    "Copy table({}/{}): {} to {}",
98                    table_no, num_tables, full_table_name, table_file
99                );
100                self.copy_captured_table_to(table, copy_table_req, moved_ctx)
101                    .await
102            });
103        }
104
105        let results = try_join_all(tasks).await?;
106        let exported_rows = results.into_iter().sum();
107        Ok(Output::new_with_affected_rows(exported_rows))
108    }
109
110    /// Imports data to database from a given location and returns total rows imported.
111    #[tracing::instrument(skip_all)]
112    pub(crate) async fn copy_database_from(
113        &self,
114        req: CopyDatabaseRequest,
115        ctx: QueryContextRef,
116    ) -> error::Result<Output> {
117        if let Some(layout) = req.with.get("metric_data_layout") {
118            return error::InvalidCopyParameterSnafu {
119                key: "metric_data_layout",
120                value: layout,
121            }
122            .fail();
123        }
124        validate_database_directory(&req.location)?;
125
126        let parallelism = parse_parallelism_from_option_map(&req.with);
127        info!(
128            "Copy database {}.{} from dir: {}, time: {:?}, parallelism: {}",
129            req.catalog_name, req.schema_name, req.location, req.time_range, parallelism
130        );
131        let suffix = Format::try_from(&req.with)
132            .context(error::ParseFileFormatSnafu)?
133            .suffix();
134
135        let entries = list_files_to_copy(&req, suffix, &self.local_file_access).await?;
136
137        let continue_on_error = req
138            .with
139            .get(CONTINUE_ON_ERROR_KEY)
140            .and_then(|v| bool::from_str(v).ok())
141            .unwrap_or(false);
142
143        let mut tasks = Vec::with_capacity(entries.len());
144        let semaphore = Arc::new(Semaphore::new(parallelism));
145
146        for e in entries {
147            let (table_name, location) = match database_import_source(&req.location, e.path()) {
148                Ok(source) => source,
149                Err(err) => {
150                    if continue_on_error {
151                        error!(err; "Failed to import table from file: {:?}", e);
152                        continue;
153                    } else {
154                        return Err(err);
155                    }
156                }
157            };
158
159            let req = CopyTableRequest {
160                catalog_name: req.catalog_name.clone(),
161                schema_name: req.schema_name.clone(),
162                table_name: table_name.clone(),
163                location,
164                with: req.with.clone(),
165                connection: req.connection.clone(),
166                pattern: None,
167                direction: CopyDirection::Import,
168                timestamp_range: None,
169                limit: None,
170            };
171            let moved_ctx = ctx.clone();
172            let moved_table_name = table_name.clone();
173            let moved_semaphore = semaphore.clone();
174            tasks.push(async move {
175                let _permit = moved_semaphore.acquire().await.unwrap();
176                debug!("Copy table, arg: {:?}", req);
177                match self.copy_table_from(req, moved_ctx).await {
178                    Ok(o) => {
179                        let (rows, cost) = o.extract_rows_and_cost();
180                        Ok((rows, cost))
181                    }
182                    Err(err) => {
183                        if continue_on_error {
184                            error!(err; "Failed to import file to table: {}", moved_table_name);
185                            Ok((0, 0))
186                        } else {
187                            Err(err)
188                        }
189                    }
190                }
191            });
192        }
193
194        let results = try_join_all(tasks).await?;
195        let (rows_inserted, insert_cost) = results
196            .into_iter()
197            .fold((0, 0), |(acc_rows, acc_cost), (rows, cost)| {
198                (acc_rows + rows, acc_cost + cost)
199            });
200
201        Ok(Output::new(
202            OutputData::AffectedRows(rows_inserted),
203            OutputMeta::new_with_cost(insert_cost),
204        ))
205    }
206}
207
208/// Lists all files with expected suffix that can be imported to database.
209async fn list_files_to_copy(
210    req: &CopyDatabaseRequest,
211    suffix: &str,
212    local_file_access: &LocalFileAccess,
213) -> error::Result<Vec<Entry>> {
214    let object_store = build_backend(&req.location, &req.connection, local_file_access)
215        .await
216        .context(error::BuildBackendSnafu)?;
217
218    let pattern = Regex::try_from(format!(".*{}", suffix)).context(error::BuildRegexSnafu)?;
219    let lister = Lister::new(
220        object_store.clone(),
221        Source::Dir,
222        "/".to_string(),
223        Some(pattern),
224    );
225    lister.list().await.context(error::ListObjectsSnafu)
226}
227
228#[cfg(test)]
229mod tests {
230    use std::collections::HashSet;
231
232    use common_datasource::object_store::LocalFileAccess;
233    use object_store::ObjectStore;
234    use object_store::services::Fs;
235    use object_store::util::normalize_dir;
236    #[cfg(not(windows))]
237    use path_slash::PathExt;
238    use table::requests::CopyDatabaseRequest;
239
240    use crate::statement::copy_database::list_files_to_copy;
241    use crate::statement::database_copy::database_import_source;
242
243    #[tokio::test]
244    async fn test_list_files_and_parse_table_name() {
245        let dir = common_test_util::temp_dir::create_temp_dir("test_list_files_to_copy");
246        let store_dir = normalize_dir(dir.path().to_str().unwrap());
247        let builder = Fs::default().root(&store_dir);
248        let object_store = ObjectStore::new(builder).unwrap();
249        object_store.write("a.parquet", "").await.unwrap();
250        object_store.write("b.parquet", "").await.unwrap();
251        object_store.write("c.csv", "").await.unwrap();
252        object_store.write("d", "").await.unwrap();
253        object_store.write("e.f.parquet", "").await.unwrap();
254
255        #[cfg(not(windows))]
256        let location = normalize_dir(&dir.path().to_slash().unwrap());
257        #[cfg(windows)]
258        let location = format!("{}\\", dir.path().display());
259        let request = CopyDatabaseRequest {
260            catalog_name: "catalog_0".to_string(),
261            schema_name: "schema_0".to_string(),
262            location,
263            with: [("FORMAT".to_string(), "parquet".to_string())]
264                .into_iter()
265                .collect(),
266            connection: Default::default(),
267            time_range: None,
268        };
269        let local_file_access = LocalFileAccess::sandboxed(dir.path()).unwrap();
270        let listed = list_files_to_copy(&request, ".parquet", &local_file_access)
271            .await
272            .unwrap()
273            .into_iter()
274            .map(|e| {
275                database_import_source(&request.location, e.path())
276                    .unwrap()
277                    .0
278            })
279            .collect::<HashSet<_>>();
280
281        assert_eq!(
282            ["a".to_string(), "b".to_string(), "e.f".to_string()]
283                .into_iter()
284                .collect::<HashSet<_>>(),
285            listed
286        );
287    }
288}