diff --git a/src/cli/src/data.rs b/src/cli/src/data.rs index c880bd81c73..3e1924f12ee 100644 --- a/src/cli/src/data.rs +++ b/src/cli/src/data.rs @@ -18,6 +18,7 @@ mod import; pub mod import_v2; pub(crate) mod path; pub(crate) mod progress; +mod schema_export; pub mod snapshot_storage; pub(crate) mod sql; mod storage_export; diff --git a/src/cli/src/data/export.rs b/src/cli/src/data/export.rs index 148c27316e7..055c5d4ee02 100644 --- a/src/cli/src/data/export.rs +++ b/src/cli/src/data/export.rs @@ -27,6 +27,7 @@ use tokio::sync::Semaphore; use tokio::time::Instant; use crate::common::{ObjectStoreConfig, new_fs_object_store}; +use crate::data::schema_export::append_schema_ddl; use crate::data::storage_export::{ AzblobBackend, FsBackend, GcsBackend, OssBackend, S3Backend, StorageExport, StorageType, }; @@ -335,32 +336,6 @@ impl Export { )) } - async fn show_create( - &self, - show_type: &str, - catalog: &str, - schema: &str, - table: Option<&str>, - ) -> Result { - let sql = match table { - Some(table) => format!( - r#"SHOW CREATE {} "{}"."{}"."{}""#, - show_type, catalog, schema, table - ), - None => format!(r#"SHOW CREATE {} "{}"."{}""#, show_type, catalog, schema), - }; - let records = self - .database_client - .sql_in_public(&sql) - .await? - .context(EmptyResultSnafu)?; - let Value::String(create) = &records[0][1] else { - unreachable!() - }; - - Ok(format!("{};\n", create)) - } - async fn export_create_database(&self) -> Result<()> { let timer = Instant::now(); let db_names = self.get_db_names().await?; @@ -368,9 +343,15 @@ impl Export { let operator = self.build_prefer_fs_operator().await?; for schema in db_names { - let create_database = self - .show_create("DATABASE", &self.catalog, &schema, None) - .await?; + let mut create_database = String::new(); + append_schema_ddl( + &self.database_client, + &self.catalog, + &schema, + std::iter::once(("DATABASE", None)), + &mut create_database, + ) + .await?; let file_path = self.get_file_path(&schema, "create_database.sql"); self.write_to_storage(&operator, &file_path, create_database.into_bytes()) @@ -415,23 +396,28 @@ impl Export { } let file_path = export_self.get_file_path(&schema, "create_tables.sql"); - let mut content = Vec::new(); - - // Add table creation SQL - for (c, s, t) in metric_physical_tables.iter().chain(&remaining_tables) { - let create_table = export_self.show_create("TABLE", c, s, Some(t)).await?; - content.extend_from_slice(create_table.as_bytes()); - } - - // Add view creation SQL - for (c, s, v) in &views { - let create_view = export_self.show_create("VIEW", c, s, Some(v)).await?; - content.extend_from_slice(create_view.as_bytes()); - } + let mut content = String::new(); + let objects = metric_physical_tables + .iter() + .chain(&remaining_tables) + .map(|(_, _, table)| ("TABLE", Some(table.as_str()))) + .chain( + views + .iter() + .map(|(_, _, view)| ("VIEW", Some(view.as_str()))), + ); + append_schema_ddl( + &export_self.database_client, + &export_self.catalog, + &schema, + objects, + &mut content, + ) + .await?; // Write to storage export_self - .write_to_storage(&operator, &file_path, content) + .write_to_storage(&operator, &file_path, content.into_bytes()) .await?; info!( @@ -445,9 +431,11 @@ impl Export { }); } - let success = self.execute_tasks(tasks).await; + for result in futures::future::join_all(tasks).await { + result?; + } let elapsed = timer.elapsed(); - info!("Success {success}/{db_count} jobs, cost: {elapsed:?}"); + info!("Success {db_count}/{db_count} jobs, cost: {elapsed:?}"); Ok(()) } diff --git a/src/cli/src/data/export_v2/command.rs b/src/cli/src/data/export_v2/command.rs index 20e69ef5a99..48779f3c435 100644 --- a/src/cli/src/data/export_v2/command.rs +++ b/src/cli/src/data/export_v2/command.rs @@ -25,15 +25,15 @@ use common_error::ext::BoxedError; use common_telemetry::info; use serde_json::Value; use servers::http::{ColumnSchema, GreptimeQueryOutput, OutputSchema}; -use snafu::{OptionExt, ResultExt}; +use snafu::ResultExt; use crate::Tool; use crate::common::ObjectStoreConfig; use crate::data::export_v2::coordinator::{ExportDataOptions, export_data}; use crate::data::export_v2::error::{ - ChunkTimeWindowRequiresBoundsSnafu, DatabaseSnafu, EmptyResultSnafu, IoSnafu, - ManifestVersionMismatchSnafu, Result, ResumeConfigMismatchSnafu, SchemaOnlyArgsNotAllowedSnafu, - SchemaOnlyModeMismatchSnafu, SnapshotVerifyFailedSnafu, UnexpectedValueTypeSnafu, + ChunkTimeWindowRequiresBoundsSnafu, DatabaseSnafu, IoSnafu, ManifestVersionMismatchSnafu, + Result, ResumeConfigMismatchSnafu, SchemaOnlyArgsNotAllowedSnafu, SchemaOnlyModeMismatchSnafu, + SnapshotVerifyFailedSnafu, UnexpectedValueTypeSnafu, }; use crate::data::export_v2::extractor::SchemaExtractor; use crate::data::export_v2::manifest::{ @@ -42,10 +42,11 @@ use crate::data::export_v2::manifest::{ use crate::data::export_v2::schema::{DDL_DIR, SCHEMA_DIR, SCHEMAS_FILE}; use crate::data::path::{data_dir_for_schema_chunk, ddl_path_for_schema}; use crate::data::progress::{ProgressMode, build_progress_reporter}; +use crate::data::schema_export::append_schema_ddl; use crate::data::snapshot_storage::{ OpenDalStorage, SnapshotStorage, validate_snapshot_uri, validate_uri, }; -use crate::data::sql::{escape_sql_identifier, escape_sql_literal}; +use crate::data::sql::escape_sql_literal; use crate::database::{DatabaseClient, parse_proxy_opts}; /// Export V2 commands. @@ -620,8 +621,10 @@ impl ExportCreate { info!("Exported {} schemas", schema_snapshot.schemas.len()); // 5. Export DDL files for import recovery. - let ddl_by_schema = self.build_ddl_by_schema(&schema_names).await?; - for (schema, ddl) in ddl_by_schema { + let mut sorted_schemas = schema_names; + sorted_schemas.sort(); + for schema in sorted_schemas { + let ddl = self.build_schema_ddl(&schema).await?; let ddl_path = ddl_path_for_schema(&schema); self.storage.write_text(&ddl_path, &ddl).await?; info!("Exported DDL for schema {} to {}", schema, ddl_path); @@ -659,45 +662,31 @@ impl ExportCreate { Ok(()) } - async fn build_ddl_by_schema(&self, schema_names: &[String]) -> Result> { - let mut schemas = schema_names.to_vec(); - schemas.sort(); - - let mut ddl_by_schema = Vec::with_capacity(schemas.len()); - for schema in schemas { - let create_database = self.show_create("DATABASE", &schema, None).await?; - - let (mut physical_tables, mut tables, mut views) = - self.get_schema_objects(&schema).await?; - physical_tables.sort(); - let mut physical_ddls = Vec::with_capacity(physical_tables.len()); - for table in physical_tables { - physical_ddls.push(self.show_create("TABLE", &schema, Some(&table)).await?); - } - - tables.sort(); - let mut table_ddls = Vec::with_capacity(tables.len()); - for table in tables { - table_ddls.push(self.show_create("TABLE", &schema, Some(&table)).await?); - } - - views.sort(); - let mut view_ddls = Vec::with_capacity(views.len()); - for view in views { - view_ddls.push(self.show_create("VIEW", &schema, Some(&view)).await?); - } - - let ddl = build_schema_ddl( - &schema, - create_database, - physical_ddls, - table_ddls, - view_ddls, - ); - ddl_by_schema.push((schema, ddl)); - } - - Ok(ddl_by_schema) + async fn build_schema_ddl(&self, schema: &str) -> Result { + let (mut physical_tables, mut tables, mut views) = self.get_schema_objects(schema).await?; + physical_tables.sort(); + tables.sort(); + views.sort(); + let objects = std::iter::once(("DATABASE", None)) + .chain( + physical_tables + .iter() + .chain(&tables) + .map(|table| ("TABLE", Some(table.as_str()))), + ) + .chain(views.iter().map(|view| ("VIEW", Some(view.as_str())))); + let mut ddl = format!("-- Schema: {schema}\n"); + append_schema_ddl( + &self.database_client, + &self.config.catalog, + schema, + objects, + &mut ddl, + ) + .await + .context(DatabaseSnafu)?; + ddl.push('\n'); + Ok(ddl) } async fn get_schema_objects( @@ -770,65 +759,6 @@ impl ExportCreate { Ok(tables.into_iter().collect()) } - - async fn show_create( - &self, - show_type: &str, - schema: &str, - table: Option<&str>, - ) -> Result { - let sql = match table { - Some(table) => format!( - r#"SHOW CREATE {} "{}"."{}"."{}""#, - show_type, - escape_sql_identifier(&self.config.catalog), - escape_sql_identifier(schema), - escape_sql_identifier(table) - ), - None => format!( - r#"SHOW CREATE {} "{}"."{}""#, - show_type, - escape_sql_identifier(&self.config.catalog), - escape_sql_identifier(schema) - ), - }; - - let records: Option>> = self - .database_client - .sql_in_public(&sql) - .await - .context(DatabaseSnafu)?; - let rows = records.context(EmptyResultSnafu)?; - let row = rows.first().context(EmptyResultSnafu)?; - let Some(Value::String(create)) = row.get(1) else { - return UnexpectedValueTypeSnafu.fail(); - }; - - Ok(format!("{};\n", create)) - } -} - -fn build_schema_ddl( - schema: &str, - create_database: String, - physical_tables: Vec, - tables: Vec, - views: Vec, -) -> String { - let mut ddl = String::new(); - ddl.push_str(&format!("-- Schema: {}\n", schema)); - ddl.push_str(&create_database); - for stmt in physical_tables { - ddl.push_str(&stmt); - } - for stmt in tables { - ddl.push_str(&stmt); - } - for stmt in views { - ddl.push_str(&stmt); - } - ddl.push('\n'); - ddl } fn validate_resume_config(manifest: &Manifest, config: &ExportConfig) -> Result<()> { @@ -1637,25 +1567,6 @@ mod tests { ); } - #[test] - fn test_build_schema_ddl_order() { - let ddl = build_schema_ddl( - "public", - "CREATE DATABASE public;\n".to_string(), - vec!["PHYSICAL;\n".to_string()], - vec!["TABLE;\n".to_string()], - vec!["VIEW;\n".to_string()], - ); - - let db_pos = ddl.find("CREATE DATABASE").unwrap(); - let physical_pos = ddl.find("PHYSICAL;").unwrap(); - let table_pos = ddl.find("TABLE;").unwrap(); - let view_pos = ddl.find("VIEW;").unwrap(); - assert!(db_pos < physical_pos); - assert!(physical_pos < table_pos); - assert!(table_pos < view_pos); - } - #[tokio::test] async fn test_build_rejects_chunk_window_without_bounds() { let cmd = ExportCreateCommand::parse_from([ diff --git a/src/cli/src/data/schema_export.rs b/src/cli/src/data/schema_export.rs new file mode 100644 index 00000000000..8bb99121d12 --- /dev/null +++ b/src/cli/src/data/schema_export.rs @@ -0,0 +1,508 @@ +// Copyright 2023 Greptime Team +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//! Bounded SHOW CREATE requests shared by both export commands. + +use common_catalog::consts::DEFAULT_SCHEMA_NAME; +use servers::http::GreptimeQueryOutput; +use servers::http::result::greptime_result_v1::GreptimedbV1Response; + +use crate::data::sql::escape_sql_identifier; +use crate::database::DatabaseClient; +use crate::error::{InvalidArgumentsSnafu, Result, UnexpectedSnafu}; + +const MAX_STATEMENTS: usize = 128; +const MAX_SQL_BYTES: usize = 128 * 1024; + +/// Appends DDL in dependency order supplied by the caller using bounded requests. +pub(crate) fn append_schema_ddl<'a>( + client: &'a DatabaseClient, + catalog: &'a str, + schema: &'a str, + objects: impl Iterator)> + Send + 'a, + ddl: &'a mut String, +) -> impl std::future::Future> + Send + 'a { + let mut statements = objects + .map(move |(kind, table)| show_create_sql(kind, catalog, schema, table)) + .peekable(); + async move { + loop { + let batch = next_batch(&mut statements)?; + if batch.is_empty() { + return Ok(()); + } + let response = client + .sql_response(&batch.concat(), DEFAULT_SCHEMA_NAME) + .await + .map_err(|error| { + UnexpectedSnafu { + msg: format!("SHOW CREATE batch failed for {}: {error}", batch.concat()), + } + .build() + })?; + append_results(&batch, &response, ddl)?; + } + } +} + +fn show_create_sql(kind: &str, catalog: &str, schema: &str, table: Option<&str>) -> String { + let mut sql = format!( + "SHOW CREATE {kind} \"{}\".\"{}\"", + escape_sql_identifier(catalog), + escape_sql_identifier(schema), + ); + if let Some(table) = table { + sql.push_str(".\""); + sql.push_str(&escape_sql_identifier(table)); + sql.push('"'); + } + sql.push_str(";\n"); + sql +} + +fn next_batch( + statements: &mut std::iter::Peekable>, +) -> Result> { + let mut batch = Vec::new(); + let mut bytes = 0; + while let Some(sql) = statements.peek() { + if sql.len() > MAX_SQL_BYTES { + return InvalidArgumentsSnafu { + msg: format!( + "SHOW CREATE exceeds {MAX_SQL_BYTES} SQL bytes ({}): {sql}", + sql.len() + ), + } + .fail(); + } + if batch.len() == MAX_STATEMENTS || bytes + sql.len() > MAX_SQL_BYTES { + break; + } + bytes += sql.len(); + if let Some(sql) = statements.next() { + batch.push(sql); + } + } + Ok(batch) +} + +fn append_results( + batch: &[String], + response: &GreptimedbV1Response, + ddl: &mut String, +) -> Result<()> { + let outputs = response.output(); + if outputs.len() != batch.len() { + return UnexpectedSnafu { + msg: format!( + "SHOW CREATE expected {} results, received {} for {}", + batch.len(), + outputs.len(), + batch.concat() + ), + } + .fail(); + } + for (sql, output) in batch.iter().zip(outputs) { + let GreptimeQueryOutput::Records(records) = output else { + return UnexpectedSnafu { + msg: format!("expected SHOW CREATE records for {sql}"), + } + .fail(); + }; + if records.num_cols() != 2 || records.rows().len() != 1 || records.rows()[0].len() != 2 { + return UnexpectedSnafu { + msg: format!("expected one two-column SHOW CREATE row for {sql}"), + } + .fail(); + } + let row = &records.rows()[0]; + let Some(create) = row[1].as_str().filter(|value| !value.is_empty()) else { + return UnexpectedSnafu { + msg: format!("expected SHOW CREATE DDL string for {sql}"), + } + .fail(); + }; + if !row[0].is_string() { + return UnexpectedSnafu { + msg: format!("expected SHOW CREATE object name for {sql}"), + } + .fail(); + } + ddl.push_str(create); + ddl.push_str(";\n"); + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use serde_json::{Value, json}; + use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader}; + + use super::*; + + fn records(name: &str, ddl: &str) -> Value { + json!({"records": {"schema": {"column_schemas": [ + {"name": "Table", "data_type": "String"}, + {"name": "Create Table", "data_type": "String"} + ]}, "rows": [[name, ddl]], "total_rows": 1}}) + } + + #[test] + fn batch_limits_include_quoted_utf8_and_separators() { + let special = show_create_sql("TABLE", "cat.name", "s\";库", Some("ta.\";表")); + assert_eq!( + special, + "SHOW CREATE TABLE \"cat.name\".\"s\"\";库\".\"ta.\"\";表\";\n" + ); + let mut statements = std::iter::repeat_n(special.clone(), 129).peekable(); + assert_eq!( + next_batch(&mut statements).unwrap(), + vec![special.clone(); 128] + ); + assert_eq!(next_batch(&mut statements).unwrap(), vec![special]); + assert!(next_batch(&mut statements).unwrap().is_empty()); + + let prefix = show_create_sql("TABLE", "c", "s", Some("")); + // Every quote expands to two bytes; this also exercises a multibyte name. + let name = format!("表{}", "\"".repeat((MAX_SQL_BYTES - prefix.len() - 3) / 2)); + let mut exact = show_create_sql("TABLE", "c", "s", Some(&name)); + let padding = MAX_SQL_BYTES - exact.len(); + exact = show_create_sql( + "TABLE", + "c", + "s", + Some(&format!("{name}{}", "a".repeat(padding))), + ); + assert_eq!(exact.len(), MAX_SQL_BYTES); + let small = show_create_sql("TABLE", "c", "s", Some("next")); + let mut statements = [exact.clone(), small.clone()].into_iter().peekable(); + assert_eq!(next_batch(&mut statements).unwrap(), vec![exact.clone()]); + assert_eq!(next_batch(&mut statements).unwrap(), vec![small]); + let mut oversized = std::iter::once(format!("{exact} ")).peekable(); + assert!( + next_batch(&mut oversized) + .unwrap_err() + .to_string() + .contains("exceeds 131072") + ); + + let half = show_create_sql( + "TABLE", + "c", + "s", + Some(&"a".repeat(MAX_SQL_BYTES / 2 - prefix.len())), + ); + let mut statements = [half.clone(), half.clone(), prefix.clone()] + .into_iter() + .peekable(); + assert_eq!( + next_batch(&mut statements).unwrap(), + vec![half.clone(), half] + ); + assert_eq!(next_batch(&mut statements).unwrap(), vec![prefix]); + } + + #[test] + fn show_create_requires_exact_record_shape_and_count() { + let batch = vec![ + "SHOW CREATE TABLE first".into(), + "SHOW CREATE TABLE middle".into(), + ]; + let valid = records("first", "CREATE TABLE first"); + for output in [ + vec![], + vec![valid.clone()], + vec![valid.clone(); 3], + vec![valid.clone(), json!({"affectedrows": 1})], + vec![valid.clone(), records("middle", "")], + vec![valid.clone(), { + let mut r = valid.clone(); + r["records"]["rows"] = json!([]); + r + }], + vec![valid.clone(), { + let mut r = valid.clone(); + r["records"]["rows"] = json!([["middle", 4]]); + r + }], + vec![valid.clone(), { + let mut r = valid.clone(); + r["records"]["rows"] = json!([["middle", "ddl", "extra"]]); + r + }], + ] { + let response = + serde_json::from_value(json!({"output": output, "execution_time_ms": 0})).unwrap(); + let error = append_results(&batch, &response, &mut String::new()).unwrap_err(); + assert!(error.to_string().contains("middle"), "{error}"); + } + } + + async fn start_mock_sql_server( + bodies: Vec<(u16, String)>, + ) -> (DatabaseClient, tokio::task::JoinHandle>) { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + tokio::time::timeout(std::time::Duration::from_secs(5), async move { + let (socket, _) = listener.accept().await.unwrap(); + let mut socket = BufReader::new(socket); + let mut queries = Vec::new(); + for (status, body) in bodies { + let mut headers = String::new(); + loop { + let mut line = String::new(); + if socket.read_line(&mut line).await.unwrap() == 0 { + assert!(headers.is_empty()); + return queries; + } + headers.push_str(&line); + if line == "\r\n" { + break; + } + } + let headers = headers.to_lowercase(); + assert!(headers.starts_with("post /v1/sql ")); + assert!(headers.contains("authorization: basic dxnlcjpwyxnzd29yza==")); + assert!(headers.contains("x-greptime-timeout: 7s")); + let length: usize = headers + .lines() + .find_map(|line| line.strip_prefix("content-length: ")) + .unwrap() + .parse() + .unwrap(); + let mut form = vec![0; length]; + socket.read_exact(&mut form).await.unwrap(); + let params: std::collections::HashMap<_, _> = + url::form_urlencoded::parse(&form).into_owned().collect(); + assert_eq!(params["db"], "catalog-public"); + queries.push(params["sql"].clone()); + socket + .get_mut() + .write_all( + format!( + "HTTP/1.1 {status} Response\r\nContent-Length: {}\r\n\r\n{body}", + body.len() + ) + .as_bytes(), + ) + .await + .unwrap(); + } + queries + }) + .await + .expect("mock SQL exchange timed out") + }); + let client = DatabaseClient::new( + address.to_string(), + "catalog".into(), + Some("user:password".into()), + std::time::Duration::from_secs(7), + None, + true, + ); + (client, server) + } + + #[tokio::test] + async fn batches_preserve_order_context_and_reuse_connection() { + let names: Vec<_> = (0..129).map(|i| format!("表.{i}\";x")).collect(); + let outputs: Vec<_> = names + .iter() + .map(|name| records(name, &format!("DDL {name}"))) + .collect(); + let bodies = outputs + .chunks(128) + .map(|chunk| { + ( + 200, + json!({"output": chunk, "execution_time_ms": 0}).to_string(), + ) + }) + .collect(); + let (client, server) = start_mock_sql_server(bodies).await; + let mut ddl = String::new(); + append_schema_ddl( + &client.clone(), + "catalog", + "schema", + names.iter().map(|name| ("TABLE", Some(name.as_str()))), + &mut ddl, + ) + .await + .unwrap(); + assert_eq!( + ddl, + names + .iter() + .map(|name| format!("DDL {name};\n")) + .collect::() + ); + let queries = server.await.unwrap(); + let expected: Vec<_> = (0..129) + .map(|i| format!("SHOW CREATE TABLE \"catalog\".\"schema\".\"表.{i}\"\";x\";\n")) + .collect(); + assert_eq!( + queries, + vec![expected[..128].concat(), expected[128..].concat()] + ); + } + + #[tokio::test] + async fn legacy_command_stops_before_data_after_schema_failure() { + use clap::Parser; + + use crate::data::export::ExportCommand; + + let output = |rows: Value, columns: usize| { + json!({"records": {"schema": {"column_schemas": (0..columns) + .map(|i| json!({"name": i.to_string(), "data_type": "String"})) + .collect::>()}, "rows": rows, "total_rows": rows.as_array().unwrap().len()}}) + }; + let databases = output(json!([["public"]]), 1); + let mut malformed = records("middle", "CREATE TABLE middle"); + malformed["records"]["rows"][0][1] = json!(42); + let bodies = vec![ + vec![databases.clone()], + vec![records("public", "CREATE DATABASE public")], + vec![databases.clone()], + vec![output(json!([]), 3)], + vec![output( + json!([ + ["catalog", "public", "first", "BASE TABLE"], + ["catalog", "public", "middle", "BASE TABLE"], + ["catalog", "public", "last", "BASE TABLE"] + ]), + 4, + )], + vec![ + records("first", "CREATE TABLE first"), + malformed, + records("last", "CREATE TABLE last"), + ], + vec![databases], + vec![json!({"affectedrows": 0})], + ] + .into_iter() + .map(|output| { + ( + 200, + json!({"output": output, "execution_time_ms": 0}).to_string(), + ) + }) + .collect::>(); + for target in ["schema", "all"] { + let (client, server) = start_mock_sql_server(bodies.clone()).await; + let dir = tempfile::tempdir().unwrap(); + let command = ExportCommand::parse_from([ + "export", + "--addr", + client.addr(), + "--database", + "catalog-public", + "--output-dir", + dir.path().to_str().unwrap(), + "--target", + target, + "--auth-basic", + "user:password", + "--timeout", + "7s", + "--no-proxy", + ]); + let tool = command.build().await.unwrap(); + let result = tool.do_work().await; + drop(tool); + let queries = server.await.unwrap(); + assert!(result.unwrap_err().to_string().contains("middle")); + assert!( + dir.path() + .join("catalog/public/create_database.sql") + .is_file() + ); + assert!(!dir.path().join("catalog/public/create_tables.sql").exists()); + assert_eq!(queries.len(), 6); + assert!(queries[5].contains("SHOW CREATE TABLE \"catalog\".\"public\".\"middle\"")); + assert!(!queries.iter().any(|sql| sql.starts_with("COPY"))); + } + } + + #[tokio::test] + async fn sql_response_errors_do_not_expose_copy_credentials() { + let secret = "sentinel-connection-secret"; + let echoed = "sentinel-response-secret"; + let sql = format!( + "COPY DATABASE public TO 's3://bucket/' CONNECTION (SECRET_ACCESS_KEY='{secret}')" + ); + let bodies = vec![ + (200, format!("not json {echoed}")), + (200, json!({"output": [{"records": echoed}], "execution_time_ms": 0}).to_string()), + (400, json!({"code": 3001, "error": echoed}).to_string()), + (200, json!({"error": echoed}).to_string()), + (200, json!({"code": 3001, "output": [records("first", echoed)], "execution_time_ms": 0}).to_string()), + (200, json!({"output": [{"error": echoed}]}).to_string()), + (200, json!({"output": [{"code": 3001, "records": records("first", echoed)["records"]}], "execution_time_ms": 0}).to_string()), + ]; + let count = bodies.len(); + let (client, server) = start_mock_sql_server(bodies).await; + for _ in 0..count { + let error = client + .sql_response(&sql, "public") + .await + .unwrap_err() + .to_string(); + assert!(error.contains("SQL"), "{error}"); + assert!(!error.contains(secret), "{error}"); + assert!(!error.contains(echoed), "{error}"); + assert!(!error.contains("COPY DATABASE"), "{error}"); + } + assert_eq!(server.await.unwrap().len(), count); + } + + #[tokio::test] + async fn rejects_http_and_statement_errors_without_retry() { + let valid = records("first", "CREATE TABLE first"); + let bodies = vec![ + (400, json!({"code": 3001, "error": "middle not found", "execution_time_ms": 0}).to_string()), + (200, json!({"code": 3001, "output": [valid.clone(), valid.clone(), valid.clone()], "execution_time_ms": 0}).to_string()), + (200, json!({"output": [valid.clone(), {"error": "middle denied"}, valid.clone()], "execution_time_ms": 0}).to_string()), + (200, json!({"output": [valid.clone(), {"error": "middle denied", "records": valid["records"]}, valid], "execution_time_ms": 0}).to_string()), + (200, json!({"output": [records("first", "DDL"), {"code": 3001, "records": records("middle", "DDL")["records"]}, records("last", "DDL")], "execution_time_ms": 0}).to_string()), + (200, r#"{"execution_time_ms":0}"#.into()), + (200, "not json".into()), + ]; + let count = bodies.len(); + let (client, server) = start_mock_sql_server(bodies).await; + for _ in 0..count { + let mut ddl = String::new(); + let error = append_schema_ddl( + &client, + "catalog", + "schema", + ["first", "middle", "last"] + .into_iter() + .map(|name| ("TABLE", Some(name))), + &mut ddl, + ) + .await + .unwrap_err(); + assert!(error.to_string().contains("middle"), "{error}"); + assert!(ddl.is_empty()); + } + assert_eq!(server.await.unwrap().len(), count); + } +} diff --git a/src/cli/src/database.rs b/src/cli/src/database.rs index ac314bc5899..da277e166be 100644 --- a/src/cli/src/database.rs +++ b/src/cli/src/database.rs @@ -12,6 +12,7 @@ // See the License for the specific language governing permissions and // limitations under the License. +use std::sync::Arc; use std::time::Duration; use base64::Engine; @@ -37,6 +38,7 @@ pub struct DatabaseClient { timeout: Duration, proxy: Option, no_proxy: bool, + client: Arc>, } pub fn parse_proxy_opts( @@ -86,6 +88,7 @@ impl DatabaseClient { timeout, proxy, no_proxy, + client: Arc::new(tokio::sync::OnceCell::new()), } } @@ -109,15 +112,7 @@ impl DatabaseClient { async fn require_packed_capability(&self, capability: &str) -> Result<()> { let url = format!("http://{}/v1/capabilities", self.addr); - let mut builder = reqwest::Client::builder().timeout(self.timeout); - if let Some(proxy) = self.proxy.clone() { - builder = builder.proxy(proxy); - } - if self.no_proxy { - builder = builder.no_proxy(); - } - let client = builder.build().context(BuildClientSnafu)?; - let mut request = client.get(&url); + let mut request = self.http_client().await?.get(&url).timeout(self.timeout); if let Some(auth) = &self.auth_header { request = request.header("Authorization", auth); } @@ -173,15 +168,9 @@ impl DatabaseClient { ("db", format!("{}-{}", self.catalog, schema)), ("sql", sql.to_string()), ]; - let mut builder = reqwest::Client::builder(); - if let Some(proxy) = self.proxy.clone() { - builder = builder.proxy(proxy); - } - if self.no_proxy { - builder = builder.no_proxy(); - } - let client = builder.build().context(BuildClientSnafu)?; - let mut request = client + let mut request = self + .http_client() + .await? .post(&url) .form(¶ms) .header("Content-Type", "application/x-www-form-urlencoded"); @@ -197,17 +186,69 @@ impl DatabaseClient { let response = request.send().await.with_context(|_| HttpQuerySqlSnafu { reason: format!("bad url: {}", url), })?; - let response = response - .error_for_status() - .with_context(|_| HttpQuerySqlSnafu { - reason: format!("query failed: {}", sql), - })?; - + let status = response.status(); let text = response.text().await.with_context(|_| HttpQuerySqlSnafu { reason: "cannot get response text".to_string(), })?; + let value: Value = serde_json::from_str(&text).map_err(|_| { + crate::error::UnexpectedSnafu { + msg: format!("invalid SQL response ({status})"), + } + .build() + })?; + if !status.is_success() + || value.get("error").is_some() + || value + .get("code") + .is_some_and(|code| code.as_u64() != Some(0)) + { + return crate::error::UnexpectedSnafu { + msg: format!( + "SQL request failed ({status}, code {:?})", + value.get("code").and_then(Value::as_u64) + ), + } + .fail(); + } + if let Some(outputs) = value.get("output").and_then(Value::as_array) { + for (index, output) in outputs.iter().enumerate() { + if output.get("error").is_some() + || output + .get("code") + .is_some_and(|code| code.as_u64() != Some(0)) + { + return crate::error::UnexpectedSnafu { + msg: format!( + "SQL statement {} failed (code {:?})", + index + 1, + output.get("code").and_then(Value::as_u64) + ), + } + .fail(); + } + } + } + serde_json::from_value(value).map_err(|_| { + crate::error::UnexpectedSnafu { + msg: format!("invalid SQL response ({status})"), + } + .build() + }) + } - serde_json::from_str::(&text).context(SerdeJsonSnafu) + async fn http_client(&self) -> Result<&reqwest::Client> { + self.client + .get_or_try_init(|| async { + let mut builder = reqwest::Client::builder(); + if let Some(proxy) = self.proxy.clone() { + builder = builder.proxy(proxy); + } + if self.no_proxy { + builder = builder.no_proxy(); + } + builder.build().context(BuildClientSnafu) + }) + .await } } diff --git a/tests-integration/tests/export_logical_tables.rs b/tests-integration/tests/export_logical_tables.rs index 80169a391cd..da1232a243b 100644 --- a/tests-integration/tests/export_logical_tables.rs +++ b/tests-integration/tests/export_logical_tables.rs @@ -1541,6 +1541,20 @@ async fn metric_export_v2_cli_roundtrip(s3: bool, packed: bool) { .await; sql(instance, "INSERT INTO audit VALUES ('a',1,1),('z',NULL,3)").await; sql(instance, "CREATE VIEW dashboard AS SELECT * FROM audit").await; + let special_name = "audit.dashboard"; + let quoted_special = special_name.replace('"', "\"\""); + sql( + instance, + &format!("CREATE VIEW \"{quoted_special}\" AS SELECT * FROM audit"), + ) + .await; + if packed && !s3 { + for id in 0..128 { + let name = format!("batch_{id:03}"); + sql(instance, &format!("CREATE TABLE {name} (host STRING PRIMARY KEY, val DOUBLE, ts TIMESTAMP TIME INDEX) ENGINE=metric WITH (on_physical_table='v2_a')")).await; + names.push(name); + } + } sql(instance, "CREATE DATABASE z_later").await; if packed { sql(instance, "CREATE DATABASE empty_schema").await; @@ -1573,6 +1587,21 @@ async fn metric_export_v2_cli_roundtrip(s3: bool, packed: bool) { names.push("bulk".into()); } let (addr, server) = export_http(instance.clone()).await; + if packed && !s3 { + let client = cli::DatabaseClient::new( + addr.clone(), + "greptime".into(), + None, + std::time::Duration::from_secs(60), + None, + true, + ); + let error = client.sql_in_public( + "SHOW CREATE TABLE audit; SHOW CREATE TABLE missing_middle; SHOW CREATE VIEW dashboard" + ).await.unwrap_err(); + assert!(error.to_string().contains("SQL request failed"), "{error}"); + } + let destination = tempfile::tempdir_in(common_test_util::find_workspace_path(".")).unwrap(); for experimental in [false, true] { for layout in ["packed", "invalid"] { @@ -1876,6 +1905,23 @@ async fn metric_export_v2_cli_roundtrip(s3: bool, packed: bool) { ) .await ); + for view in ["dashboard", special_name] { + let query = format!( + "SELECT * FROM \"{}\" ORDER BY ts,host", + view.replace('"', "\"\"") + ); + assert_eq!( + values(instance, &query).await, + values(target.fe_instance(), &query).await + ); + assert_eq!( + table(instance, view).await.schema().column_schemas(), + table(target.fe_instance(), view) + .await + .schema() + .column_schemas() + ); + } if packed { let schema_uri = format!("{uri}/schema-only"); let mut schema_args = vec![