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:
xfy 2026-06-29 19:24:06 +08:00
parent 6e94db0c2b
commit 37d7ec49eb
3 changed files with 73 additions and 26 deletions

View File

@ -88,7 +88,21 @@ async fn run_pg_dump_backup(task_id: &str, timestamp: &str) {
);
let filename = format!("backup_{}.sql", timestamp);
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();
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));
}
}
// 按表更新进度
// 按表更新进度(用 u32 避免大 schema 下的截断/溢出)
tasks::update(
task_id,
&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,
None,
None,
@ -345,7 +359,21 @@ pub async fn restore_backup(filename: String, confirm: bool) -> Result<String, S
#[cfg(feature = "server")]
async fn run_restore(task_id: &str, filename: &str) {
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")
.arg("--version")
.output()

View File

@ -50,11 +50,11 @@ pub async fn export_data(
// 2. 解析来源 + 白名单/只读校验
let include_columns = params.include_columns.unwrap_or(true);
let source_sql = parse_source(&params.source)?;
let (source_sql, table_name) = parse_source(&params.source)?;
match params.format.as_str() {
"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())),
}
.map(|(body, content_type, filename)| {
@ -74,17 +74,17 @@ pub async fn export_data(
})
}
/// 解析导出来源,返回可直接执行的内部 SQL
/// - `table:posts` → 校验表名合法后,`SELECT * FROM "posts"`(只读)
/// - `query:SELECT ...` → 校验为只读语句后原样返回
fn parse_source(source: &str) -> Result<String, (StatusCode, String)> {
/// 解析导出来源,返回(内部 SQL表名
/// - `table:posts` → 校验表名合法后,`SELECT * FROM "posts"`(只读)+ 表名 "posts"
/// - `query:SELECT ...` → 校验为只读语句后原样返回,表名为 "export"
fn parse_source(source: &str) -> Result<(String, String), (StatusCode, String)> {
if let Some(table) = source.strip_prefix("table:") {
// 表名白名单:仅允许标识符字符,防注入
let t = table.trim();
if t.is_empty() || !is_simple_ident(t) {
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:") {
// 只读校验sqlparser 解析后所有语句均为 Query/Explain
let dialect = sqlparser::dialect::PostgreSqlDialect {};
@ -96,7 +96,7 @@ fn parse_source(source: &str) -> Result<String, (StatusCode, String)> {
"导出查询必须是只读SELECT/EXPLAIN".to_string(),
));
}
Ok(query.to_string())
Ok((query.to_string(), "export".to_string()))
} else {
Err((
StatusCode::BAD_REQUEST,
@ -144,6 +144,7 @@ async fn export_csv(
async fn export_sql(
source_sql: String,
include_columns: bool,
table_name: String,
) -> Result<(Body, &'static str, String), (StatusCode, String)> {
let client = crate::db::pool::get_conn()
.await
@ -154,8 +155,6 @@ async fn export_sql(
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("查询失败:{e}")))?;
// 表名(从 SQL 简单解析,或用 "export"
let table_name = "export";
let columns: Vec<String> = rows
.first()
.map(|r| r.columns().iter().map(|c| c.name().to_string()).collect())
@ -175,7 +174,7 @@ async fn export_sql(
.collect();
out.push_str(&format!(
"INSERT INTO {} {} VALUES ({});\n",
table_name,
&table_name,
col_clause,
vals.join(", ")
));

View File

@ -40,8 +40,6 @@ pub struct SqlResult {
pub explain: Option<String>,
/// 是否因 500 行上限截断。
pub truncated: bool,
/// 截断时的估算总行数(重查用)。
pub total_estimate: Option<i64>,
}
/// 结果行数上限(超出截断 + 提示)。
@ -67,11 +65,20 @@ fn check_guards(
asts: &[sqlparser::ast::Statement],
confirm_dangerous: bool,
) -> GuardResult {
use sqlparser::ast::Statement;
use sqlparser::ast::{ObjectType, Statement};
for stmt in asts {
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 { .. } => {
if !confirm_dangerous {
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 start = std::time::Instant::now;
// 逐条执行allow_multi 时多条,否则单条)
// 逐条执行:每条用其 AST 重序列化的形式stmt.to_string()
// 而非原始整段 SQL——保证执行的语句与护栏检查的 AST 完全一致,
// 避免 allow_multi 时整段 SQL 被重复执行,也杜绝读/写分类与实际执行解耦。
let mut last_result = SqlResult::default();
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)
}
@ -246,11 +257,14 @@ pub async fn execute_sql(sql: String, opts: ExecuteSqlOpts) -> Result<SqlResult,
}
/// 执行单条语句,返回结果。
///
/// `stmt_sql` 必须是 `stmt` 重序列化后的**单条**语句 SQL不含其他语句
/// 保证护栏检查的 AST 与实际执行的语句一致。
#[cfg(feature = "server")]
async fn execute_one(
client: &deadpool_postgres::Object,
stmt: &sqlparser::ast::Statement,
sql: &str,
stmt_sql: &str,
with_explain: bool,
start: impl Fn() -> std::time::Instant + Copy,
) -> Result<SqlResult, ServerFnError> {
@ -258,8 +272,8 @@ async fn execute_one(
let read_only = is_read_only(stmt);
if with_explain && read_only {
// EXPLAIN 模式:包裹原 SQL 执行计划,取首列文本拼接
let explain_sql = format!("EXPLAIN {}", sql.trim_end_matches(';'));
// EXPLAIN 模式:包裹单条语句取执行计划,取首列文本拼接
let explain_sql = format!("EXPLAIN {}", stmt_sql.trim_end_matches(';'));
let rows = client
.query(&explain_sql, &[])
.await
@ -279,7 +293,10 @@ async fn execute_one(
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
.first()
.map(|r| {
@ -309,7 +326,10 @@ async fn execute_one(
})
} else {
// 写操作:返回影响行数
let affected = client.execute(sql, &[]).await.map_err(AppError::query)?;
let affected = client
.execute(stmt_sql, &[])
.await
.map_err(AppError::query)?;
Ok(SqlResult {
affected_rows: affected,
statement_type,