From 37d7ec49eb72306d1872b2978d397689944d61bd Mon Sep 17 00:00:00 2001 From: xfy Date: Mon, 29 Jun 2026 19:24:06 +0800 Subject: [PATCH] fix(database): harden SQL console guards + review fixes MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 最终 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 字段。 --- src/api/database/backup.rs | 36 ++++++++++++++++++++++++---- src/api/database/export.rs | 21 ++++++++--------- src/api/database/sql_console.rs | 42 ++++++++++++++++++++++++--------- 3 files changed, 73 insertions(+), 26 deletions(-) diff --git a/src/api/database/backup.rs b/src/api/database/backup.rs index 84cc60a..5600482 100644 --- a/src/api/database/backup.rs +++ b/src/api/database/backup.rs @@ -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 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() diff --git a/src/api/database/export.rs b/src/api/database/export.rs index c292ec2..f73d65b 100644 --- a/src/api/database/export.rs +++ b/src/api/database/export.rs @@ -50,11 +50,11 @@ pub async fn export_data( // 2. 解析来源 + 白名单/只读校验 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() { "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 { +/// 解析导出来源,返回(内部 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 { "导出查询必须是只读(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 = 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(", ") )); diff --git a/src/api/database/sql_console.rs b/src/api/database/sql_console.rs index b74744d..e7ff071 100644 --- a/src/api/database/sql_console.rs +++ b/src/api/database/sql_console.rs @@ -40,8 +40,6 @@ pub struct SqlResult { pub explain: Option, /// 是否因 500 行上限截断。 pub truncated: bool, - /// 截断时的估算总行数(重查用)。 - pub total_estimate: Option, } /// 结果行数上限(超出截断 + 提示)。 @@ -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 Result std::time::Instant + Copy, ) -> Result { @@ -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 = 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,