perf: batch schema export requests (#9408)

* perf: batch schema export requests

Signed-off-by: jeremyhi <fengjiachun@gmail.com>

* fix: propagate legacy schema export failures

Signed-off-by: jeremyhi <fengjiachun@gmail.com>

* fix: keep credentials out of SQL response errors

Signed-off-by: jeremyhi <fengjiachun@gmail.com>

* test: use a typo-safe dotted catalog name

Signed-off-by: jeremyhi <fengjiachun@gmail.com>

* refactor(cli): address schema export review feedback

Signed-off-by: jeremyhi <fengjiachun@gmail.com>

---------

Signed-off-by: jeremyhi <fengjiachun@gmail.com>
This commit is contained in:
jeremyhi
2026-09-30 06:23:33 +00:00
committed by GitHub
parent 8553db3f27
commit 78eefbe4c4
6 changed files with 689 additions and 194 deletions
+1
View File
@@ -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;
+33 -45
View File
@@ -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<String> {
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(())
}
+35 -124
View File
@@ -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<Vec<(String, String)>> {
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<String> {
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<String> {
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<Vec<Vec<Value>>> = 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<String>,
tables: Vec<String>,
views: Vec<String>,
) -> 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([
+508
View File
@@ -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<Item = (&'static str, Option<&'a str>)> + Send + 'a,
ddl: &'a mut String,
) -> impl std::future::Future<Output = Result<()>> + 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<impl Iterator<Item = String>>,
) -> Result<Vec<String>> {
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<Vec<String>>) {
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::<String>()
);
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::<Vec<_>>()}, "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::<Vec<_>>();
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);
}
}
+66 -25
View File
@@ -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<reqwest::Proxy>,
no_proxy: bool,
client: Arc<tokio::sync::OnceCell<reqwest::Client>>,
}
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(&params)
.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::<GreptimedbV1Response>(&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
}
}
@@ -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![