yggdrasil/src/api/database/sql_console.rs
xfy 5426458f4a
Some checks failed
CI / check (push) Failing after 5m18s
CI / build (push) Has been skipped
fix(docs): resolve rustdoc broken links and dead_code warnings
doc 构建报 8 个 broken_intra_doc_links + dx 构建报 7 个 dead_code 警告。
两者根源相同:被引用的常量/类型/模块要么是 server-only(消费方都在
#[cfg(feature="server")] 块里),要么是跨 cfg 引用(host doc 构建看不到
wasm32-only 类型)。

dead_code:给 sql_console 的 MAX_ROWS/ABSOLUTELY_FORBIDDEN/GuardResult 与
backup 的 BACKUP_DIR/FILENAME_RE/BACKUP_SIGNATURE 加 #[cfg(feature="server")]
gate——这些都是安全护栏常量,消费方全在 server fn 体里,WASM 构建剥掉 fn 体后
常量即成死代码。sysinfo_sampler 的 WASM 桩 read_snapshot 加 allow(dead_code)
(与 codemirror_bridge 既有处理一致)。

rustdoc 链接:跨模块引用改全路径(crate::xxx),跨 cfg 引用(EditorHandle 在
wasm32 mod 内、host doc 构建不可见)转义为纯文本。

验证:cargo doc(0 warning)+ cargo check(default)+ wasm32 check 全部干净。
2026-06-30 17:42:52 +08:00

351 lines
12 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#![allow(clippy::unused_unit, deprecated)]
//! SQL 控制台执行(全读写 + 4 道护栏)。
//!
//! 护栏:
//! 1. 高危语句闸门:`DROP DATABASE`/`DROP SCHEMA`(字符串预检)绝禁;
//! `DROP`/`TRUNCATE`/`ALTER` 需 `confirm_dangerous`。
//! 2. 无 WHERE 拦截:`UPDATE`/`DELETE` 无 `selection` 拒绝。
//! 3. 查询超时上限:复用 `STATEMENT_TIMEOUT_SECS`pool 层已注入 GUC
//! 4. 前端二次确认(前端实现)。
//!
//! 默认禁止多语句(`allow_multi` 放开)。
use dioxus::prelude::*;
use serde::{Deserialize, Serialize};
// admin 鉴权 + DB 查询仅在 server 构建里被 server function 体引用。
#[cfg(feature = "server")]
use crate::api::auth::get_current_admin_user;
#[cfg(feature = "server")]
use crate::api::error::AppError;
#[cfg(feature = "server")]
use crate::db::pool::get_conn;
#[derive(Deserialize, Serialize, Debug, Clone, Copy)]
pub struct ExecuteSqlOpts {
/// 是否允许多语句(`;` 分隔),默认 false。
pub allow_multi: bool,
/// 是否勾选「我了解后果」(放开 DROP/TRUNCATE/ALTER 等高危)。
pub confirm_dangerous: bool,
/// 是否带 EXPLAIN 执行计划。
pub with_explain: bool,
}
#[derive(Serialize, Deserialize, Debug, Default, Clone)]
pub struct SqlResult {
pub columns: Vec<String>,
/// 每格用 JSON 表示text/int/timestamp/bool/null
pub rows: Vec<Vec<serde_json::Value>>,
pub affected_rows: u64,
pub elapsed_ms: u64,
/// 语句类型(来自 AST如 "Select"/"Update"/"CreateTable")。
pub statement_type: String,
pub explain: Option<String>,
/// 是否因 500 行上限截断。
pub truncated: bool,
}
// 以下常量/枚举仅被 server function 体引用WASM 构建里 server fn 体被 cfg 剥掉,
// 故这些符号也需 gate否则非 server 构建会报 dead_code
/// 结果行数上限(超出截断 + 提示)。
#[cfg(feature = "server")]
const MAX_ROWS: usize = 500;
/// 绝对禁止的语句关键词字符串预检sqlparser 无 ObjectType::Database/Schema
/// 命中即拒,不可放行。
#[cfg(feature = "server")]
const ABSOLUTELY_FORBIDDEN: &[&str] = &["drop database", "drop schema", "create database"];
/// 护栏检查返回值。
#[cfg(feature = "server")]
#[derive(Debug)]
enum GuardResult {
Allowed,
/// 需 confirm_dangerous 才放行。
NeedsConfirm,
/// 不可放行(附带原因)。
Forbidden(String),
}
/// 护栏 1+2sqlparser 解析后遍历 AST检查高危语句与无 WHERE 的 UPDATE/DELETE。
#[cfg(feature = "server")]
fn check_guards(
asts: &[sqlparser::ast::Statement],
confirm_dangerous: bool,
) -> GuardResult {
use sqlparser::ast::{ObjectType, Statement};
for stmt in asts {
match stmt {
// 护栏 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;
}
}
// 护栏 2UPDATE 无 WHERE
Statement::Update { selection: None, .. } => {
return GuardResult::Forbidden(
"UPDATE 缺少 WHERE 子句,将影响全表。请加 WHERE 条件。".to_string(),
);
}
// 护栏 2DELETE 无 WHERE
Statement::Delete { selection: None, .. } => {
return GuardResult::Forbidden(
"DELETE 缺少 WHERE 子句,将影响全表。请加 WHERE 条件。".to_string(),
);
}
_ => {}
}
}
GuardResult::Allowed
}
/// 提取语句类型名AST 变体名,如 "Select"/"Insert"/"Update")。
#[cfg(feature = "server")]
fn statement_type_name(stmt: &sqlparser::ast::Statement) -> String {
use sqlparser::ast::Statement;
let name = match stmt {
Statement::Query(_) => "Select",
Statement::Insert { .. } => "Insert",
Statement::Update { .. } => "Update",
Statement::Delete { .. } => "Delete",
Statement::CreateTable { .. } => "CreateTable",
Statement::AlterTable { .. } => "AlterTable",
Statement::Drop { .. } => "Drop",
Statement::Truncate { .. } => "Truncate",
Statement::Explain { .. } => "Explain",
_ => "Other",
};
name.to_string()
}
/// 判断语句是否只读SELECT/EXPLAIN/SHOW/WITH...SELECT
#[cfg(feature = "server")]
fn is_read_only(stmt: &sqlparser::ast::Statement) -> bool {
use sqlparser::ast::Statement;
matches!(
stmt,
Statement::Query(_) | Statement::Explain { .. }
)
}
/// 把一列的值转成 JSON按 PG 类型名分发)。
#[cfg(feature = "server")]
fn col_to_json(row: &tokio_postgres::Row, idx: usize) -> serde_json::Value {
use serde_json::json;
let ty = row.columns().get(idx).map(|c| c.type_().name()).unwrap_or("");
match ty {
"int2" => row
.try_get::<_, Option<i16>>(idx)
.ok()
.flatten()
.map(|v| json!(v))
.unwrap_or(serde_json::Value::Null),
"int4" => row
.try_get::<_, Option<i32>>(idx)
.ok()
.flatten()
.map(|v| json!(v))
.unwrap_or(serde_json::Value::Null),
"int8" => row
.try_get::<_, Option<i64>>(idx)
.ok()
.flatten()
.map(|v| json!(v))
.unwrap_or(serde_json::Value::Null),
"float4" => row
.try_get::<_, Option<f32>>(idx)
.ok()
.flatten()
.map(|v| json!(v))
.unwrap_or(serde_json::Value::Null),
"float8" => row
.try_get::<_, Option<f64>>(idx)
.ok()
.flatten()
.map(|v| json!(v))
.unwrap_or(serde_json::Value::Null),
"bool" => row
.try_get::<_, Option<bool>>(idx)
.ok()
.flatten()
.map(|v| json!(v))
.unwrap_or(serde_json::Value::Null),
// 其余text/varchar/timestamp/jsonb/...)一律按字符串取,失败则 null
_ => row
.try_get::<_, Option<String>>(idx)
.ok()
.flatten()
.map(|v| json!(v))
.unwrap_or(serde_json::Value::Null),
}
}
/// 执行 SQL全读写管理员。护栏见模块文档。
#[server(ExecuteSql, "/api")]
pub async fn execute_sql(sql: String, opts: ExecuteSqlOpts) -> Result<SqlResult, ServerFnError> {
let _user = get_current_admin_user().await?;
#[cfg(feature = "server")]
{
use sqlparser::dialect::PostgreSqlDialect;
use sqlparser::parser::Parser;
// 护栏 1绝禁字符串预检 DROP/CREATE DATABASE、DROP SCHEMA
let normalized = sql.to_lowercase();
for forbidden in ABSOLUTELY_FORBIDDEN {
if normalized.contains(forbidden) {
return Err(AppError::BadRequest(format!(
"禁止的操作:{}",
forbidden.to_uppercase()
))
.into());
}
}
// 解析 SQL
let dialect = PostgreSqlDialect {};
let asts = Parser::parse_sql(&dialect, &sql)
.map_err(|e| AppError::BadRequest(format!("SQL 解析失败:{e}")))?;
if asts.is_empty() {
return Err(AppError::BadRequest("空的 SQL 语句".into()).into());
}
// 多语句检查(默认禁止)
if asts.len() > 1 && !opts.allow_multi {
return Err(AppError::BadRequest(
"检测到多条语句,请勾选「允许多语句」后再执行".into(),
)
.into());
}
// 护栏 1+2AST 检查
match check_guards(&asts, opts.confirm_dangerous) {
GuardResult::Forbidden(msg) => {
return Err(AppError::BadRequest(msg).into());
}
GuardResult::NeedsConfirm => {
return Err(AppError::BadRequest(
"高危操作DROP/TRUNCATE/ALTER需勾选「我了解后果」".into(),
)
.into());
}
GuardResult::Allowed => {}
}
let client = get_conn().await.map_err(AppError::db_conn)?;
let start = std::time::Instant::now;
// 逐条执行:每条用其 AST 重序列化的形式stmt.to_string()
// 而非原始整段 SQL——保证执行的语句与护栏检查的 AST 完全一致,
// 避免 allow_multi 时整段 SQL 被重复执行,也杜绝读/写分类与实际执行解耦。
let mut last_result = SqlResult::default();
for stmt in &asts {
// 重序列化单条语句为可执行 SQL去掉末尾分号避免与 execute 的隐式分号冲突)
let stmt_sql = stmt.to_string();
last_result = execute_one(&client, stmt, &stmt_sql, opts.with_explain, start).await?;
}
Ok(last_result)
}
#[cfg(not(feature = "server"))]
{
let _ = (sql, opts);
Ok(SqlResult::default())
}
}
/// 执行单条语句,返回结果。
///
/// `stmt_sql` 必须是 `stmt` 重序列化后的**单条**语句 SQL不含其他语句
/// 保证护栏检查的 AST 与实际执行的语句一致。
#[cfg(feature = "server")]
async fn execute_one(
client: &deadpool_postgres::Object,
stmt: &sqlparser::ast::Statement,
stmt_sql: &str,
with_explain: bool,
start: impl Fn() -> std::time::Instant + Copy,
) -> Result<SqlResult, ServerFnError> {
let statement_type = statement_type_name(stmt);
let read_only = is_read_only(stmt);
if with_explain && read_only {
// EXPLAIN 模式:包裹单条语句取执行计划,取首列文本拼接
let explain_sql = format!("EXPLAIN {}", stmt_sql.trim_end_matches(';'));
let rows = client
.query(&explain_sql, &[])
.await
.map_err(AppError::query)?;
let explain = rows
.iter()
.filter_map(|r| r.try_get::<_, String>(0).ok())
.collect::<Vec<_>>()
.join("\n");
return Ok(SqlResult {
statement_type,
explain: Some(explain),
elapsed_ms: start().elapsed().as_millis() as u64,
..Default::default()
});
}
if read_only {
// 只读:取结果集。列名从第一行取(空结果集时无列名,前端容错)。
let rows = client
.query(stmt_sql, &[])
.await
.map_err(AppError::query)?;
let columns: Vec<String> = rows
.first()
.map(|r| {
r.columns()
.iter()
.map(|c| c.name().to_string())
.collect()
})
.unwrap_or_default();
let mut data: Vec<Vec<serde_json::Value>> = Vec::new();
let mut truncated = false;
for r in &rows {
if data.len() >= MAX_ROWS {
truncated = true;
break;
}
let row: Vec<serde_json::Value> = (0..r.len()).map(|i| col_to_json(r, i)).collect();
data.push(row);
}
Ok(SqlResult {
columns,
rows: data,
truncated,
statement_type,
elapsed_ms: start().elapsed().as_millis() as u64,
..Default::default()
})
} else {
// 写操作:返回影响行数
let affected = client
.execute(stmt_sql, &[])
.await
.map_err(AppError::query)?;
Ok(SqlResult {
affected_rows: affected,
statement_type,
elapsed_ms: start().elapsed().as_millis() as u64,
..Default::default()
})
}
}