fix(database): harden SQL console guards + review fixes
最终 code review 发现的安全/正确性修复: - C1 (critical): execute_one 改用 stmt.to_string() 重序列化的单条语句 SQL 而非原始整段 SQL,保证执行的语句与护栏检查的 AST 一致;杜绝 allow_multi 时整段 SQL 被重复执行、读写分类与实际执行解耦。 - C2 (critical): check_guards 在 AST 层结构性禁止 DROP SCHEMA(ObjectType 匹配),不再仅靠字符串预检——防 SQL 注释/空白绕过。 - I2: 备份/恢复在 DATABASE_URL 为空时提前失败任务,而非传空串给 pg_dump/psql。 - I4: 回退备份进度计算用 u32 避免 >255 表时截断/溢出。 - I5: SQL 导出用真实表名(table 模式)而非硬编码 "export"。 - M7: 移除 SqlResult 中未使用的 total_estimate 字段。
This commit is contained in:
parent
6e94db0c2b
commit
37d7ec49eb
@ -88,7 +88,21 @@ async fn run_pg_dump_backup(task_id: &str, timestamp: &str) {
|
|||||||
);
|
);
|
||||||
let filename = format!("backup_{}.sql", timestamp);
|
let filename = format!("backup_{}.sql", timestamp);
|
||||||
let path = backup_path(&filename);
|
let path = backup_path(&filename);
|
||||||
let db_url = std::env::var("DATABASE_URL").unwrap_or_default();
|
let db_url = match std::env::var("DATABASE_URL") {
|
||||||
|
Ok(u) if !u.is_empty() => u,
|
||||||
|
_ => {
|
||||||
|
tasks::update(
|
||||||
|
task_id,
|
||||||
|
"DATABASE_URL 未配置",
|
||||||
|
100,
|
||||||
|
TaskStatus::Failed,
|
||||||
|
None,
|
||||||
|
Some("pg_dump 备份需要 DATABASE_URL".to_string()),
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
let mut header = String::new();
|
let mut header = String::new();
|
||||||
header.push_str(&format!("{}\n", BACKUP_SIGNATURE));
|
header.push_str(&format!("{}\n", BACKUP_SIGNATURE));
|
||||||
@ -262,11 +276,11 @@ async fn run_sql_fallback_backup(task_id: &str, timestamp: &str) {
|
|||||||
out.push_str(&format!("-- 导出失败: {}\n", e));
|
out.push_str(&format!("-- 导出失败: {}\n", e));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// 按表更新进度
|
// 按表更新进度(用 u32 避免大 schema 下的截断/溢出)
|
||||||
tasks::update(
|
tasks::update(
|
||||||
task_id,
|
task_id,
|
||||||
&format!("导出表 {}/{}", i + 1, total),
|
&format!("导出表 {}/{}", i + 1, total),
|
||||||
10 + (i + 1) as u8 * 90 / total as u8,
|
(10 + (i + 1) as u32 * 90 / total as u32).min(99) as u8,
|
||||||
TaskStatus::Running,
|
TaskStatus::Running,
|
||||||
None,
|
None,
|
||||||
None,
|
None,
|
||||||
@ -345,7 +359,21 @@ pub async fn restore_backup(filename: String, confirm: bool) -> Result<String, S
|
|||||||
#[cfg(feature = "server")]
|
#[cfg(feature = "server")]
|
||||||
async fn run_restore(task_id: &str, filename: &str) {
|
async fn run_restore(task_id: &str, filename: &str) {
|
||||||
let path = backup_path(filename);
|
let path = backup_path(filename);
|
||||||
let db_url = std::env::var("DATABASE_URL").unwrap_or_default();
|
let db_url = match std::env::var("DATABASE_URL") {
|
||||||
|
Ok(u) if !u.is_empty() => u,
|
||||||
|
_ => {
|
||||||
|
tasks::update(
|
||||||
|
task_id,
|
||||||
|
"DATABASE_URL 未配置",
|
||||||
|
100,
|
||||||
|
TaskStatus::Failed,
|
||||||
|
None,
|
||||||
|
Some("恢复需要 DATABASE_URL".to_string()),
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
};
|
||||||
let psql_ok = std::process::Command::new("psql")
|
let psql_ok = std::process::Command::new("psql")
|
||||||
.arg("--version")
|
.arg("--version")
|
||||||
.output()
|
.output()
|
||||||
|
|||||||
@ -50,11 +50,11 @@ pub async fn export_data(
|
|||||||
|
|
||||||
// 2. 解析来源 + 白名单/只读校验
|
// 2. 解析来源 + 白名单/只读校验
|
||||||
let include_columns = params.include_columns.unwrap_or(true);
|
let include_columns = params.include_columns.unwrap_or(true);
|
||||||
let source_sql = parse_source(¶ms.source)?;
|
let (source_sql, table_name) = parse_source(¶ms.source)?;
|
||||||
|
|
||||||
match params.format.as_str() {
|
match params.format.as_str() {
|
||||||
"csv" => export_csv(source_sql, include_columns).await,
|
"csv" => export_csv(source_sql, include_columns).await,
|
||||||
"sql" => export_sql(source_sql, include_columns).await,
|
"sql" => export_sql(source_sql, include_columns, table_name).await,
|
||||||
_ => Err((StatusCode::BAD_REQUEST, "不支持的格式".to_string())),
|
_ => Err((StatusCode::BAD_REQUEST, "不支持的格式".to_string())),
|
||||||
}
|
}
|
||||||
.map(|(body, content_type, filename)| {
|
.map(|(body, content_type, filename)| {
|
||||||
@ -74,17 +74,17 @@ pub async fn export_data(
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 解析导出来源,返回可直接执行的内部 SQL。
|
/// 解析导出来源,返回(内部 SQL,表名)。
|
||||||
/// - `table:posts` → 校验表名合法后,`SELECT * FROM "posts"`(只读)。
|
/// - `table:posts` → 校验表名合法后,`SELECT * FROM "posts"`(只读)+ 表名 "posts"。
|
||||||
/// - `query:SELECT ...` → 校验为只读语句后原样返回。
|
/// - `query:SELECT ...` → 校验为只读语句后原样返回,表名为 "export"。
|
||||||
fn parse_source(source: &str) -> Result<String, (StatusCode, String)> {
|
fn parse_source(source: &str) -> Result<(String, String), (StatusCode, String)> {
|
||||||
if let Some(table) = source.strip_prefix("table:") {
|
if let Some(table) = source.strip_prefix("table:") {
|
||||||
// 表名白名单:仅允许标识符字符,防注入
|
// 表名白名单:仅允许标识符字符,防注入
|
||||||
let t = table.trim();
|
let t = table.trim();
|
||||||
if t.is_empty() || !is_simple_ident(t) {
|
if t.is_empty() || !is_simple_ident(t) {
|
||||||
return Err((StatusCode::BAD_REQUEST, "无效的表名".to_string()));
|
return Err((StatusCode::BAD_REQUEST, "无效的表名".to_string()));
|
||||||
}
|
}
|
||||||
Ok(format!("SELECT * FROM \"{}\"", t))
|
Ok((format!("SELECT * FROM \"{}\"", t), t.to_string()))
|
||||||
} else if let Some(query) = source.strip_prefix("query:") {
|
} else if let Some(query) = source.strip_prefix("query:") {
|
||||||
// 只读校验:sqlparser 解析后所有语句均为 Query/Explain
|
// 只读校验:sqlparser 解析后所有语句均为 Query/Explain
|
||||||
let dialect = sqlparser::dialect::PostgreSqlDialect {};
|
let dialect = sqlparser::dialect::PostgreSqlDialect {};
|
||||||
@ -96,7 +96,7 @@ fn parse_source(source: &str) -> Result<String, (StatusCode, String)> {
|
|||||||
"导出查询必须是只读(SELECT/EXPLAIN)".to_string(),
|
"导出查询必须是只读(SELECT/EXPLAIN)".to_string(),
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
Ok(query.to_string())
|
Ok((query.to_string(), "export".to_string()))
|
||||||
} else {
|
} else {
|
||||||
Err((
|
Err((
|
||||||
StatusCode::BAD_REQUEST,
|
StatusCode::BAD_REQUEST,
|
||||||
@ -144,6 +144,7 @@ async fn export_csv(
|
|||||||
async fn export_sql(
|
async fn export_sql(
|
||||||
source_sql: String,
|
source_sql: String,
|
||||||
include_columns: bool,
|
include_columns: bool,
|
||||||
|
table_name: String,
|
||||||
) -> Result<(Body, &'static str, String), (StatusCode, String)> {
|
) -> Result<(Body, &'static str, String), (StatusCode, String)> {
|
||||||
let client = crate::db::pool::get_conn()
|
let client = crate::db::pool::get_conn()
|
||||||
.await
|
.await
|
||||||
@ -154,8 +155,6 @@ async fn export_sql(
|
|||||||
.await
|
.await
|
||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("查询失败:{e}")))?;
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("查询失败:{e}")))?;
|
||||||
|
|
||||||
// 表名(从 SQL 简单解析,或用 "export")
|
|
||||||
let table_name = "export";
|
|
||||||
let columns: Vec<String> = rows
|
let columns: Vec<String> = rows
|
||||||
.first()
|
.first()
|
||||||
.map(|r| r.columns().iter().map(|c| c.name().to_string()).collect())
|
.map(|r| r.columns().iter().map(|c| c.name().to_string()).collect())
|
||||||
@ -175,7 +174,7 @@ async fn export_sql(
|
|||||||
.collect();
|
.collect();
|
||||||
out.push_str(&format!(
|
out.push_str(&format!(
|
||||||
"INSERT INTO {} {} VALUES ({});\n",
|
"INSERT INTO {} {} VALUES ({});\n",
|
||||||
table_name,
|
&table_name,
|
||||||
col_clause,
|
col_clause,
|
||||||
vals.join(", ")
|
vals.join(", ")
|
||||||
));
|
));
|
||||||
|
|||||||
@ -40,8 +40,6 @@ pub struct SqlResult {
|
|||||||
pub explain: Option<String>,
|
pub explain: Option<String>,
|
||||||
/// 是否因 500 行上限截断。
|
/// 是否因 500 行上限截断。
|
||||||
pub truncated: bool,
|
pub truncated: bool,
|
||||||
/// 截断时的估算总行数(重查用)。
|
|
||||||
pub total_estimate: Option<i64>,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 结果行数上限(超出截断 + 提示)。
|
/// 结果行数上限(超出截断 + 提示)。
|
||||||
@ -67,11 +65,20 @@ fn check_guards(
|
|||||||
asts: &[sqlparser::ast::Statement],
|
asts: &[sqlparser::ast::Statement],
|
||||||
confirm_dangerous: bool,
|
confirm_dangerous: bool,
|
||||||
) -> GuardResult {
|
) -> GuardResult {
|
||||||
use sqlparser::ast::Statement;
|
use sqlparser::ast::{ObjectType, Statement};
|
||||||
|
|
||||||
for stmt in asts {
|
for stmt in asts {
|
||||||
match stmt {
|
match stmt {
|
||||||
// 护栏 1:需确认的高危语句
|
// 护栏 1(绝禁):DROP SCHEMA 永远禁止(ObjectType 无 Database 变体,
|
||||||
|
// DROP DATABASE 由字符串预检 + AST 双重拦截——见上方 ABSOLUTELY_FORBIDDEN)。
|
||||||
|
// 这里在 AST 层结构性禁止 DROP SCHEMA,防 SQL 注释/空白绕过字符串预检。
|
||||||
|
Statement::Drop {
|
||||||
|
object_type: ObjectType::Schema,
|
||||||
|
..
|
||||||
|
} => {
|
||||||
|
return GuardResult::Forbidden("禁止 DROP SCHEMA".to_string());
|
||||||
|
}
|
||||||
|
// 护栏 1:需确认的高危语句(DROP TABLE/VIEW/INDEX 等、TRUNCATE、ALTER)
|
||||||
Statement::Drop { .. } | Statement::Truncate { .. } | Statement::AlterTable { .. } => {
|
Statement::Drop { .. } | Statement::Truncate { .. } | Statement::AlterTable { .. } => {
|
||||||
if !confirm_dangerous {
|
if !confirm_dangerous {
|
||||||
return GuardResult::NeedsConfirm;
|
return GuardResult::NeedsConfirm;
|
||||||
@ -231,10 +238,14 @@ pub async fn execute_sql(sql: String, opts: ExecuteSqlOpts) -> Result<SqlResult,
|
|||||||
let client = get_conn().await.map_err(AppError::db_conn)?;
|
let client = get_conn().await.map_err(AppError::db_conn)?;
|
||||||
let start = std::time::Instant::now;
|
let start = std::time::Instant::now;
|
||||||
|
|
||||||
// 逐条执行(allow_multi 时多条,否则单条)
|
// 逐条执行:每条用其 AST 重序列化的形式(stmt.to_string()),
|
||||||
|
// 而非原始整段 SQL——保证执行的语句与护栏检查的 AST 完全一致,
|
||||||
|
// 避免 allow_multi 时整段 SQL 被重复执行,也杜绝读/写分类与实际执行解耦。
|
||||||
let mut last_result = SqlResult::default();
|
let mut last_result = SqlResult::default();
|
||||||
for stmt in &asts {
|
for stmt in &asts {
|
||||||
last_result = execute_one(&client, stmt, &sql, opts.with_explain, start).await?;
|
// 重序列化单条语句为可执行 SQL(去掉末尾分号,避免与 execute 的隐式分号冲突)
|
||||||
|
let stmt_sql = stmt.to_string();
|
||||||
|
last_result = execute_one(&client, stmt, &stmt_sql, opts.with_explain, start).await?;
|
||||||
}
|
}
|
||||||
Ok(last_result)
|
Ok(last_result)
|
||||||
}
|
}
|
||||||
@ -246,11 +257,14 @@ pub async fn execute_sql(sql: String, opts: ExecuteSqlOpts) -> Result<SqlResult,
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// 执行单条语句,返回结果。
|
/// 执行单条语句,返回结果。
|
||||||
|
///
|
||||||
|
/// `stmt_sql` 必须是 `stmt` 重序列化后的**单条**语句 SQL(不含其他语句),
|
||||||
|
/// 保证护栏检查的 AST 与实际执行的语句一致。
|
||||||
#[cfg(feature = "server")]
|
#[cfg(feature = "server")]
|
||||||
async fn execute_one(
|
async fn execute_one(
|
||||||
client: &deadpool_postgres::Object,
|
client: &deadpool_postgres::Object,
|
||||||
stmt: &sqlparser::ast::Statement,
|
stmt: &sqlparser::ast::Statement,
|
||||||
sql: &str,
|
stmt_sql: &str,
|
||||||
with_explain: bool,
|
with_explain: bool,
|
||||||
start: impl Fn() -> std::time::Instant + Copy,
|
start: impl Fn() -> std::time::Instant + Copy,
|
||||||
) -> Result<SqlResult, ServerFnError> {
|
) -> Result<SqlResult, ServerFnError> {
|
||||||
@ -258,8 +272,8 @@ async fn execute_one(
|
|||||||
let read_only = is_read_only(stmt);
|
let read_only = is_read_only(stmt);
|
||||||
|
|
||||||
if with_explain && read_only {
|
if with_explain && read_only {
|
||||||
// EXPLAIN 模式:包裹原 SQL 执行计划,取首列文本拼接
|
// EXPLAIN 模式:包裹单条语句取执行计划,取首列文本拼接
|
||||||
let explain_sql = format!("EXPLAIN {}", sql.trim_end_matches(';'));
|
let explain_sql = format!("EXPLAIN {}", stmt_sql.trim_end_matches(';'));
|
||||||
let rows = client
|
let rows = client
|
||||||
.query(&explain_sql, &[])
|
.query(&explain_sql, &[])
|
||||||
.await
|
.await
|
||||||
@ -279,7 +293,10 @@ async fn execute_one(
|
|||||||
|
|
||||||
if read_only {
|
if read_only {
|
||||||
// 只读:取结果集。列名从第一行取(空结果集时无列名,前端容错)。
|
// 只读:取结果集。列名从第一行取(空结果集时无列名,前端容错)。
|
||||||
let rows = client.query(sql, &[]).await.map_err(AppError::query)?;
|
let rows = client
|
||||||
|
.query(stmt_sql, &[])
|
||||||
|
.await
|
||||||
|
.map_err(AppError::query)?;
|
||||||
let columns: Vec<String> = rows
|
let columns: Vec<String> = rows
|
||||||
.first()
|
.first()
|
||||||
.map(|r| {
|
.map(|r| {
|
||||||
@ -309,7 +326,10 @@ async fn execute_one(
|
|||||||
})
|
})
|
||||||
} else {
|
} else {
|
||||||
// 写操作:返回影响行数
|
// 写操作:返回影响行数
|
||||||
let affected = client.execute(sql, &[]).await.map_err(AppError::query)?;
|
let affected = client
|
||||||
|
.execute(stmt_sql, &[])
|
||||||
|
.await
|
||||||
|
.map_err(AppError::query)?;
|
||||||
Ok(SqlResult {
|
Ok(SqlResult {
|
||||||
affected_rows: affected,
|
affected_rows: affected,
|
||||||
statement_type,
|
statement_type,
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user