1use 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 #[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
208async 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}