Skip to main content

yggdrasil/api/database/
sql_console.rs

1#![allow(clippy::unused_unit, deprecated)]
2
3//! SQL 控制台执行(全读写 + 4 道护栏)。
4//!
5//! 护栏:
6//! 1. 高危语句闸门:`DROP DATABASE`/`DROP SCHEMA`(字符串预检)绝禁;
7//!    `DROP`/`TRUNCATE`/`ALTER` 需 `confirm_dangerous`。
8//! 2. 无 WHERE 拦截:`UPDATE`/`DELETE` 无 `selection` 拒绝。
9//! 3. 查询超时上限:复用 `STATEMENT_TIMEOUT_SECS`(pool 层已注入 GUC)。
10//! 4. 前端二次确认(前端实现)。
11//!
12//! 默认禁止多语句(`allow_multi` 放开)。
13
14use dioxus::prelude::*;
15use serde::{Deserialize, Serialize};
16
17// admin 鉴权 + DB 查询仅在 server 构建里被 server function 体引用。
18#[cfg(feature = "server")]
19use crate::api::auth::get_current_admin_user;
20#[cfg(feature = "server")]
21use crate::api::error::AppError;
22#[cfg(feature = "server")]
23use crate::db::pool::get_conn;
24
25#[derive(Deserialize, Serialize, Debug, Clone, Copy)]
26pub struct ExecuteSqlOpts {
27    /// 是否允许多语句(`;` 分隔),默认 false。
28    pub allow_multi: bool,
29    /// 是否勾选「我了解后果」(放开 DROP/TRUNCATE/ALTER 等高危)。
30    pub confirm_dangerous: bool,
31    /// 是否带 EXPLAIN 执行计划。
32    pub with_explain: bool,
33}
34
35#[derive(Serialize, Deserialize, Debug, Default, Clone, PartialEq)]
36pub struct SqlResult {
37    pub columns: Vec<String>,
38    /// 每格用 JSON 表示(text/int/timestamp/bool/null)。
39    pub rows: Vec<Vec<serde_json::Value>>,
40    pub affected_rows: u64,
41    pub elapsed_ms: u64,
42    /// 语句类型(来自 AST,如 "Select"/"Update"/"CreateTable")。
43    pub statement_type: String,
44    pub explain: Option<String>,
45    /// SQL 已解析但执行失败时的脱敏/截断错误文本(admin-only 控制台可见)。
46    pub error: Option<String>,
47    /// 是否因 500 行上限截断。
48    pub truncated: bool,
49}
50
51// 以下常量/枚举仅被 server function 体引用(WASM 构建里 server fn 体被 cfg 剥掉,
52// 故这些符号也需 gate,否则非 server 构建会报 dead_code)。
53
54/// 结果行数上限(超出截断 + 提示)。
55#[cfg(feature = "server")]
56const MAX_ROWS: usize = 500;
57
58/// 绝对禁止的语句关键词(字符串预检,sqlparser 无 ObjectType::Database/Schema)。
59/// 命中即拒,不可放行。
60///
61/// 注意:这里存的是「关键字序列」,由 [`is_absolutely_forbidden`] 做 token 级
62/// 匹配——而非原始 `contains()` 子串。这样 `DROP   DATABASE`(多空格)、
63/// `DROP\tDATABASE`、`DROP\nDATABASE` 等绕过单空格子串的写法都能命中。
64/// 关键背景:sqlparser 的 PostgreSqlDialect 无法解析 DROP/CREATE DATABASE,
65/// 这类语句【没有 AST 兜底】,字符串预检是唯一防线,故必须 token 级鲁棒。
66#[cfg(feature = "server")]
67const ABSOLUTELY_FORBIDDEN: &[&[&str]] = &[
68    &["drop", "database"],
69    &["drop", "schema"],
70    &["create", "database"],
71];
72
73/// 护栏 1 的字符串预检(token 序列匹配)。
74///
75/// 把 SQL 按空白拆成 token(小写、去逗号/分号尾缀),再扫描是否出现
76/// `ABSOLUTELY_FORBIDDEN` 里任一连续关键字序列。返回命中的序列描述供报错。
77///
78/// 用 token 序列而非 `contains()` 是为了拦多空格/制表符/换行绕过——
79/// `drop   database` 在子串匹配下漏,在 token 序列下命中。
80#[cfg(feature = "server")]
81fn is_absolutely_forbidden(sql: &str) -> Option<&'static str> {
82    // 规范化:小写 + 按任意空白拆分 + 去掉 token 尾部的 , ; ( )
83    let lowered = sql.to_lowercase();
84    let tokens: Vec<&str> = lowered
85        .split_whitespace()
86        .map(|t| t.trim_end_matches([',', ';', '(', ')']))
87        .collect();
88    for forbidden in ABSOLUTELY_FORBIDDEN {
89        // 在 token 流里滑窗查连续序列
90        for window in tokens.windows(forbidden.len()) {
91            if window == *forbidden {
92                return Some(match forbidden {
93                    ["drop", "database"] => "DROP DATABASE",
94                    ["drop", "schema"] => "DROP SCHEMA",
95                    ["create", "database"] => "CREATE DATABASE",
96                    _ => "未知高危操作",
97                });
98            }
99        }
100    }
101    None
102}
103
104/// 护栏检查返回值。
105#[cfg(feature = "server")]
106#[derive(Debug)]
107enum GuardResult {
108    Allowed,
109    /// 需 confirm_dangerous 才放行。
110    NeedsConfirm,
111    /// 不可放行(附带原因)。
112    Forbidden(String),
113}
114
115/// 护栏 1+2:sqlparser 解析后遍历 AST,检查高危语句与无 WHERE 的 UPDATE/DELETE。
116#[cfg(feature = "server")]
117fn check_guards(asts: &[sqlparser::ast::Statement], confirm_dangerous: bool) -> GuardResult {
118    use sqlparser::ast::{ObjectType, Statement};
119
120    for stmt in asts {
121        match stmt {
122            // 护栏 1(绝禁):DROP SCHEMA 和 DROP DATABASE 永远禁止。
123            // 这里在 AST 层结构性禁止,防 SQL 注释/空白绕过字符串预检。
124            Statement::Drop {
125                object_type: ObjectType::Schema,
126                ..
127            } => {
128                return GuardResult::Forbidden("禁止 DROP SCHEMA".to_string());
129            }
130            Statement::Drop {
131                object_type: ObjectType::Database,
132                ..
133            } => {
134                return GuardResult::Forbidden("禁止 DROP DATABASE".to_string());
135            }
136            // 护栏 1(绝禁):CREATE DATABASE 永远禁止。sqlparser 能解析它,
137            // 故在 AST 层补上结构禁止——这是字符串预检之外的第二道防线,
138            // 防多空格/注释绕过 is_absolutely_forbidden 的 token 匹配。
139            Statement::CreateDatabase { .. } => {
140                return GuardResult::Forbidden("禁止 CREATE DATABASE".to_string());
141            }
142            // 护栏 1:需确认的高危语句(DROP TABLE/VIEW/INDEX 等、TRUNCATE、ALTER)
143            Statement::Drop { .. } | Statement::Truncate { .. } | Statement::AlterTable { .. } => {
144                if !confirm_dangerous {
145                    return GuardResult::NeedsConfirm;
146                }
147            }
148            // 护栏 2:UPDATE 无 WHERE
149            Statement::Update(sqlparser::ast::Update {
150                selection: None, ..
151            }) => {
152                return GuardResult::Forbidden(
153                    "UPDATE 缺少 WHERE 子句,将影响全表。请加 WHERE 条件。".to_string(),
154                );
155            }
156            // 护栏 2:DELETE 无 WHERE
157            Statement::Delete(sqlparser::ast::Delete {
158                selection: None, ..
159            }) => {
160                return GuardResult::Forbidden(
161                    "DELETE 缺少 WHERE 子句,将影响全表。请加 WHERE 条件。".to_string(),
162                );
163            }
164            _ => {}
165        }
166    }
167    GuardResult::Allowed
168}
169
170/// 提取语句类型名(AST 变体名,如 "Select"/"Insert"/"Update")。
171#[cfg(feature = "server")]
172fn statement_type_name(stmt: &sqlparser::ast::Statement) -> String {
173    use sqlparser::ast::Statement;
174    let name = match stmt {
175        Statement::Query(_) => "Select",
176        Statement::Insert(_) => "Insert",
177        Statement::Update(_) => "Update",
178        Statement::Delete(_) => "Delete",
179        Statement::CreateTable { .. } => "CreateTable",
180        Statement::AlterTable { .. } => "AlterTable",
181        Statement::Drop { .. } => "Drop",
182        Statement::Truncate { .. } => "Truncate",
183        Statement::Explain { .. } => "Explain",
184        _ => "Other",
185    };
186    name.to_string()
187}
188
189/// 判断语句是否只读(SELECT/WITH...SELECT/EXPLAIN/SHOW 系列)。
190#[cfg(feature = "server")]
191fn is_read_only(stmt: &sqlparser::ast::Statement) -> bool {
192    use sqlparser::ast::Statement;
193    matches!(
194        stmt,
195        // SELECT 与 WITH ... SELECT 均解析为 Query。
196        Statement::Query(_)
197            | Statement::Explain { .. }
198            // SHOW 系列只读取元数据/配置参数,不会写入业务表。
199            | Statement::ShowVariable { .. }
200            | Statement::ShowVariables { .. }
201            | Statement::ShowStatus { .. }
202            | Statement::ShowCreate { .. }
203            | Statement::ShowColumns { .. }
204            | Statement::ShowCatalogs { .. }
205            | Statement::ShowDatabases { .. }
206            | Statement::ShowProcessList { .. }
207            | Statement::ShowSchemas { .. }
208            | Statement::ShowCharset(_)
209            | Statement::ShowObjects(_)
210            | Statement::ShowTables { .. }
211            | Statement::ShowViews { .. }
212            | Statement::ShowFunctions { .. }
213            | Statement::ShowCollation { .. }
214    )
215}
216
217/// 判断一组语句中是否含有写操作(非只读语句)。
218///
219/// SQL 控制台执行写操作(INSERT/UPDATE/DELETE/TRUNCATE/ALTER/DROP 等)可能直改
220/// posts/comments/tags 等业务表,绕过 server function 的正常缓存失效路径。
221/// 抽成纯函数便于单测「哪些语句集合应触发兜底失效」。
222#[cfg(feature = "server")]
223fn writes_affect_cache(stmts: &[sqlparser::ast::Statement]) -> bool {
224    stmts.iter().any(|s| !is_read_only(s))
225}
226
227/// 把一列的值转成 JSON(按 PG 类型名分发)。
228#[cfg(feature = "server")]
229fn col_to_json(row: &tokio_postgres::Row, idx: usize) -> serde_json::Value {
230    use serde_json::json;
231    let ty = row
232        .columns()
233        .get(idx)
234        .map(|c| c.type_().name())
235        .unwrap_or("");
236    match ty {
237        "int2" => row
238            .try_get::<_, Option<i16>>(idx)
239            .ok()
240            .flatten()
241            .map(|v| json!(v))
242            .unwrap_or(serde_json::Value::Null),
243        "int4" => row
244            .try_get::<_, Option<i32>>(idx)
245            .ok()
246            .flatten()
247            .map(|v| json!(v))
248            .unwrap_or(serde_json::Value::Null),
249        "int8" => row
250            .try_get::<_, Option<i64>>(idx)
251            .ok()
252            .flatten()
253            .map(|v| json!(v))
254            .unwrap_or(serde_json::Value::Null),
255        "float4" => row
256            .try_get::<_, Option<f32>>(idx)
257            .ok()
258            .flatten()
259            .map(|v| json!(v))
260            .unwrap_or(serde_json::Value::Null),
261        "float8" => row
262            .try_get::<_, Option<f64>>(idx)
263            .ok()
264            .flatten()
265            .map(|v| json!(v))
266            .unwrap_or(serde_json::Value::Null),
267        "bool" => row
268            .try_get::<_, Option<bool>>(idx)
269            .ok()
270            .flatten()
271            .map(|v| json!(v))
272            .unwrap_or(serde_json::Value::Null),
273        // 其余(text/varchar/timestamp/jsonb/...)一律按字符串取,失败则 null
274        _ => row
275            .try_get::<_, Option<String>>(idx)
276            .ok()
277            .flatten()
278            .map(|v| json!(v))
279            .unwrap_or(serde_json::Value::Null),
280    }
281}
282
283/// 执行 SQL(全读写,管理员)。护栏见模块文档。
284#[server(ExecuteSql, "/api")]
285pub async fn execute_sql(sql: String, opts: ExecuteSqlOpts) -> Result<SqlResult, ServerFnError> {
286    let _user = get_current_admin_user().await?;
287
288    #[cfg(feature = "server")]
289    {
290        use sqlparser::dialect::PostgreSqlDialect;
291        use sqlparser::parser::Parser;
292
293        // 护栏 1(绝禁):字符串预检 DROP/CREATE DATABASE、DROP SCHEMA
294        // (token 序列匹配,防多空格绕过;这类语句无 AST 兜底)
295        if let Some(name) = is_absolutely_forbidden(&sql) {
296            return Err(AppError::BadRequest(format!("禁止的操作:{}", name)).into());
297        }
298
299        // 解析 SQL
300        let dialect = PostgreSqlDialect {};
301        let asts = Parser::parse_sql(&dialect, &sql)
302            .map_err(|e| AppError::BadRequest(format!("SQL 解析失败:{e}")))?;
303        if asts.is_empty() {
304            return Err(AppError::BadRequest("空的 SQL 语句".into()).into());
305        }
306
307        // 多语句检查(默认禁止)
308        if asts.len() > 1 && !opts.allow_multi {
309            return Err(AppError::BadRequest(
310                "检测到多条语句,请勾选「允许多语句」后再执行".into(),
311            )
312            .into());
313        }
314
315        // 护栏 1+2:AST 检查
316        match check_guards(&asts, opts.confirm_dangerous) {
317            GuardResult::Forbidden(msg) => {
318                return Err(AppError::BadRequest(msg).into());
319            }
320            GuardResult::NeedsConfirm => {
321                return Err(AppError::BadRequest(
322                    "高危操作(DROP/TRUNCATE/ALTER),需勾选「我了解后果」".into(),
323                )
324                .into());
325            }
326            GuardResult::Allowed => {}
327        }
328
329        let client = get_conn().await.map_err(AppError::db_conn)?;
330        let start = std::time::Instant::now;
331
332        // 逐条执行:每条用其 AST 重序列化的形式(stmt.to_string()),
333        // 而非原始整段 SQL——保证执行的语句与护栏检查的 AST 完全一致,
334        // 避免 allow_multi 时整段 SQL 被重复执行,也杜绝读/写分类与实际执行解耦。
335        let mut last_result = SqlResult::default();
336        for stmt in &asts {
337            // 重序列化单条语句为可执行 SQL(去掉末尾分号,避免与 execute 的隐式分号冲突)
338            let stmt_sql = stmt.to_string();
339            last_result = execute_one(&client, stmt, &stmt_sql, opts.with_explain, start).await?;
340            if last_result.error.is_some() {
341                break;
342            }
343        }
344        // 是否含写语句:抽成纯函数便于单测(见 writes_affect_cache)。
345        let has_write = writes_affect_cache(&asts);
346        // SQL 控制台写操作可能直改 posts/comments/tags,绕过了 server function 的
347        // 正常失效路径(delete_post/create_post 等已在内部失效)。此处全量兜底失效,
348        // 避免最长 10 分钟(TTL_SINGLE_POST=600s)内继续吐出陈旧数据(含已删除文章)。
349        if has_write {
350            crate::cache::invalidate_all_post_caches();
351            crate::cache::invalidate_search_results();
352            crate::cache::invalidate_all_comments();
353            crate::ssr_cache::invalidate_ssr_all_public();
354            crate::ssr_cache::bump_global_generation();
355        }
356        Ok(last_result)
357    }
358    #[cfg(not(feature = "server"))]
359    {
360        let _ = (sql, opts);
361        Ok(SqlResult::default())
362    }
363}
364
365/// 执行单条语句,返回结果。
366///
367/// `stmt_sql` 必须是 `stmt` 重序列化后的**单条**语句 SQL(不含其他语句),
368/// 保证护栏检查的 AST 与实际执行的语句一致。
369/// 把 admin-only SQL 执行错误限制在可控长度,避免异常文本撑大 server function 响应。
370#[cfg(feature = "server")]
371fn format_sql_error(error: impl std::fmt::Display) -> String {
372    error.to_string().chars().take(2_000).collect()
373}
374
375#[cfg(feature = "server")]
376async fn execute_one(
377    client: &deadpool_postgres::Object,
378    stmt: &sqlparser::ast::Statement,
379    stmt_sql: &str,
380    with_explain: bool,
381    start: impl Fn() -> std::time::Instant + Copy,
382) -> Result<SqlResult, ServerFnError> {
383    let statement_type = statement_type_name(stmt);
384    let read_only = is_read_only(stmt);
385
386    if with_explain && read_only {
387        // EXPLAIN 模式:包裹单条语句取执行计划,取首列文本拼接
388        let explain_sql = format!("EXPLAIN {}", stmt_sql.trim_end_matches(';'));
389        let rows = match client.query(&explain_sql, &[]).await {
390            Ok(rows) => rows,
391            Err(error) => {
392                return Ok(SqlResult {
393                    statement_type,
394                    error: Some(format_sql_error(error)),
395                    elapsed_ms: start().elapsed().as_millis() as u64,
396                    ..Default::default()
397                });
398            }
399        };
400        let explain = rows
401            .iter()
402            .filter_map(|r| r.try_get::<_, String>(0).ok())
403            .collect::<Vec<_>>()
404            .join("\n");
405        return Ok(SqlResult {
406            statement_type,
407            explain: Some(explain),
408            elapsed_ms: start().elapsed().as_millis() as u64,
409            ..Default::default()
410        });
411    }
412
413    if read_only {
414        // 只读:取结果集。列名从第一行取(空结果集时无列名,前端容错)。
415        // 包成子查询并追加 LIMIT MAX_ROWS+1,让 DB 只回传这么多行——避免
416        // `SELECT * FROM big_table` 把整张表物化进内存(仅 statement_timeout 兜底,M11)。
417        // +1 行供客户端检测截断:实际行数 > MAX_ROWS 时 truncated=true。
418        // 内层 SQL 已带 LIMIT 时外层 LIMIT 取更小值,结果正确;列名随子查询透传。
419        let limited_sql = format!(
420            "SELECT * FROM ({}) AS q LIMIT {}",
421            stmt_sql.trim_end_matches(';'),
422            MAX_ROWS + 1
423        );
424        let rows = match client.query(&limited_sql, &[]).await {
425            Ok(rows) => rows,
426            Err(error) => {
427                return Ok(SqlResult {
428                    statement_type,
429                    error: Some(format_sql_error(error)),
430                    elapsed_ms: start().elapsed().as_millis() as u64,
431                    ..Default::default()
432                });
433            }
434        };
435        let columns: Vec<String> = rows
436            .first()
437            .map(|r| r.columns().iter().map(|c| c.name().to_string()).collect())
438            .unwrap_or_default();
439        let mut data: Vec<Vec<serde_json::Value>> = Vec::new();
440        let mut truncated = false;
441        for r in &rows {
442            if data.len() >= MAX_ROWS {
443                truncated = true;
444                break;
445            }
446            let row: Vec<serde_json::Value> = (0..r.len()).map(|i| col_to_json(r, i)).collect();
447            data.push(row);
448        }
449        Ok(SqlResult {
450            columns,
451            rows: data,
452            truncated,
453            statement_type,
454            elapsed_ms: start().elapsed().as_millis() as u64,
455            ..Default::default()
456        })
457    } else {
458        // 写操作:返回影响行数
459        let affected = match client.execute(stmt_sql, &[]).await {
460            Ok(affected) => affected,
461            Err(error) => {
462                return Ok(SqlResult {
463                    statement_type,
464                    error: Some(format_sql_error(error)),
465                    elapsed_ms: start().elapsed().as_millis() as u64,
466                    ..Default::default()
467                });
468            }
469        };
470        Ok(SqlResult {
471            affected_rows: affected,
472            statement_type,
473            elapsed_ms: start().elapsed().as_millis() as u64,
474            ..Default::default()
475        })
476    }
477}
478
479#[cfg(all(test, feature = "server"))]
480mod tests {
481    use super::*;
482
483    /// 用 PostgreSql 方言解析 SQL 为 AST 列表,供测试复用。
484    fn parse(sql: &str) -> Vec<sqlparser::ast::Statement> {
485        use sqlparser::dialect::PostgreSqlDialect;
486        use sqlparser::parser::Parser;
487        Parser::parse_sql(&PostgreSqlDialect {}, sql).unwrap_or_default()
488    }
489
490    // ── 护栏 1:绝对禁止(token 序列字符串预检层) ──────────────────
491    // is_absolutely_forbidden 做 token 级匹配,防多空格/制表符/换行绕过。
492    // 关键背景:sqlparser 无法解析 DROP/CREATE DATABASE,这一层是唯一防线。
493
494    #[test]
495    fn absolutely_forbidden_targets_database_and_schema() {
496        // 锁定关键字序列集合:任何改动都应是有意识的。
497        assert_eq!(
498            ABSOLUTELY_FORBIDDEN,
499            &[
500                &["drop", "database"],
501                &["drop", "schema"],
502                &["create", "database"],
503            ]
504        );
505    }
506
507    #[test]
508    fn precheck_catches_forbidden_regardless_of_case() {
509        for sql in [
510            "DROP DATABASE yggdrasil",
511            "drop schema public",
512            "CREATE DATABASE evil",
513            "Drop Database x",
514        ] {
515            assert!(is_absolutely_forbidden(sql).is_some(), "应拦截: {sql:?}");
516        }
517    }
518
519    #[test]
520    fn precheck_catches_multi_space_bypass() {
521        // token 序列匹配必须命中多空格/制表符/换行绕过——
522        // 这是改用 is_absolutely_forbidden(替代旧 contains)的核心动机:
523        // sqlparser 无法解析 DROP/CREATE DATABASE,这一层是唯一防线。
524        for sql in [
525            "DROP   DATABASE x",
526            "drop\tdatabase\tx",
527            "DROP\nDATABASE\nx",
528            "DROP\t\tDATABASE x;",
529        ] {
530            assert!(
531                is_absolutely_forbidden(sql).is_some(),
532                "多空格绕过应被拦截: {sql:?}"
533            );
534        }
535    }
536
537    #[test]
538    fn precheck_returns_canonical_name_for_error_message() {
539        assert_eq!(
540            is_absolutely_forbidden("DROP DATABASE x"),
541            Some("DROP DATABASE")
542        );
543        assert_eq!(
544            is_absolutely_forbidden("drop schema public"),
545            Some("DROP SCHEMA")
546        );
547        assert_eq!(
548            is_absolutely_forbidden("CREATE DATABASE evil"),
549            Some("CREATE DATABASE")
550        );
551    }
552
553    #[test]
554    fn create_database_is_blocked_at_both_string_and_ast_layers() {
555        // CREATE DATABASE 能被 sqlparser 解析(→ AST),故有双重防线:
556        // 1. is_absolutely_forbidden(字符串 token 预检)
557        // 2. check_guards 的 Statement::CreateDatabase 分支(AST 层绝禁)
558        // 本测试锁定两道防线都生效,且确认 CREATE DATABASE 确实可被解析
559        // (若未来它变得不可解析,字符串预检就成唯一防线,需另作加固)。
560        assert!(is_absolutely_forbidden("CREATE DATABASE x").is_some());
561        let asts = parse("CREATE DATABASE x");
562        assert!(!asts.is_empty(), "CREATE DATABASE 应可被 sqlparser 解析");
563        assert!(matches!(
564            check_guards(&asts, true),
565            GuardResult::Forbidden(_)
566        ));
567    }
568
569    #[test]
570    fn drop_database_is_guarded_by_both_precheck_and_ast_check() {
571        let asts = parse("DROP DATABASE x");
572        assert!(!asts.is_empty(), "DROP DATABASE 应可被 sqlparser 解析");
573        assert!(matches!(
574            check_guards(&asts, true),
575            GuardResult::Forbidden(_)
576        ));
577        assert!(is_absolutely_forbidden("DROP DATABASE x").is_some());
578    }
579
580    #[test]
581    fn precheck_ignores_benign_statements() {
582        for sql in [
583            "SELECT * FROM users",
584            "DROP TABLE old_logs",
585            "CREATE TABLE t (id int)",
586            "DELETE FROM t WHERE id = 1",
587            "CREATE INDEX idx ON t (col)",
588        ] {
589            assert!(is_absolutely_forbidden(sql).is_none(), "不应误拦: {sql:?}");
590        }
591    }
592
593    // ── 护栏 1:AST 层 DROP SCHEMA 绝禁(防注释/空白绕过字符串预检) ──
594
595    #[test]
596    fn guard_forbids_drop_schema_even_though_string_precheck_is_bypassable() {
597        // 即使字符串预检被某种方式绕过,AST 层仍绝禁 DROP SCHEMA。
598        let asts = parse("DROP SCHEMA public");
599        match check_guards(&asts, true) {
600            GuardResult::Forbidden(msg) => {
601                assert!(msg.contains("SCHEMA"), "DROP SCHEMA 应被禁止, 得到: {msg}")
602            }
603            other => panic!("DROP SCHEMA 应 Forbidden, 得到 {other:?}"),
604        }
605    }
606
607    #[test]
608    fn guard_drop_schema_ignores_confirm_flag() {
609        // 即便勾选了 confirm_dangerous,DROP SCHEMA 仍绝禁。
610        let asts = parse("DROP SCHEMA public");
611        assert!(matches!(
612            check_guards(&asts, true),
613            GuardResult::Forbidden(_)
614        ));
615    }
616
617    // ── 护栏 1:DROP/TRUNCATE/ALTER 需确认 ─────────────────────────
618
619    #[test]
620    fn guard_drop_table_needs_confirm() {
621        let asts = parse("DROP TABLE old_logs");
622        assert!(matches!(
623            check_guards(&asts, false),
624            GuardResult::NeedsConfirm
625        ));
626    }
627
628    #[test]
629    fn guard_drop_table_allowed_with_confirm() {
630        let asts = parse("DROP TABLE old_logs");
631        assert!(matches!(check_guards(&asts, true), GuardResult::Allowed));
632    }
633
634    #[test]
635    fn guard_truncate_needs_confirm() {
636        let asts = parse("TRUNCATE TABLE sessions");
637        assert!(matches!(
638            check_guards(&asts, false),
639            GuardResult::NeedsConfirm
640        ));
641    }
642
643    #[test]
644    fn guard_alter_table_needs_confirm() {
645        let asts = parse("ALTER TABLE posts ADD COLUMN foo text");
646        assert!(matches!(
647            check_guards(&asts, false),
648            GuardResult::NeedsConfirm
649        ));
650    }
651
652    #[test]
653    fn guard_alter_table_allowed_with_confirm() {
654        let asts = parse("ALTER TABLE posts ADD COLUMN foo text");
655        assert!(matches!(check_guards(&asts, true), GuardResult::Allowed));
656    }
657
658    // ── 护栏 2:UPDATE/DELETE 无 WHERE 绝禁 ────────────────────────
659
660    #[test]
661    fn guard_update_without_where_is_forbidden() {
662        let asts = parse("UPDATE posts SET title = 'x'");
663        match check_guards(&asts, true) {
664            GuardResult::Forbidden(msg) => assert!(msg.contains("WHERE")),
665            other => panic!("无 WHERE 的 UPDATE 应 Forbidden, 得到 {other:?}"),
666        }
667    }
668
669    #[test]
670    fn guard_delete_without_where_is_forbidden() {
671        let asts = parse("DELETE FROM posts");
672        match check_guards(&asts, true) {
673            GuardResult::Forbidden(msg) => assert!(msg.contains("WHERE")),
674            other => panic!("无 WHERE 的 DELETE 应 Forbidden, 得到 {other:?}"),
675        }
676    }
677
678    #[test]
679    fn guard_update_with_where_allowed() {
680        let asts = parse("UPDATE posts SET title = 'x' WHERE id = 1");
681        assert!(matches!(check_guards(&asts, false), GuardResult::Allowed));
682    }
683
684    #[test]
685    fn guard_delete_with_where_allowed() {
686        let asts = parse("DELETE FROM posts WHERE id = 1");
687        assert!(matches!(check_guards(&asts, false), GuardResult::Allowed));
688    }
689
690    #[test]
691    fn guard_update_with_where_not_rescued_by_confirm() {
692        // confirm_dangerous 不应放行无 WHERE 的 UPDATE/DELETE——这是不可协商的安全护栏。
693        let asts = parse("UPDATE posts SET title = 'x'");
694        assert!(matches!(
695            check_guards(&asts, true),
696            GuardResult::Forbidden(_)
697        ));
698    }
699
700    // ── 正常语句 ───────────────────────────────────────────────────
701
702    #[test]
703    fn guard_allows_select_insert_create() {
704        for sql in [
705            "SELECT * FROM posts",
706            "INSERT INTO posts (title) VALUES ('x')",
707            "CREATE TABLE t (id int)",
708        ] {
709            let asts = parse(sql);
710            assert!(
711                matches!(check_guards(&asts, false), GuardResult::Allowed),
712                "应放行: {sql}"
713            );
714        }
715    }
716
717    // ── 多语句:任一命中即拒 ───────────────────────────────────────
718
719    #[test]
720    fn guard_checks_all_statements_in_batch() {
721        // 第二条是无 WHERE 的 DELETE,即便第一条正常,整体也应被拦。
722        let asts = parse("SELECT 1; DELETE FROM posts");
723        assert!(matches!(
724            check_guards(&asts, false),
725            GuardResult::Forbidden(_)
726        ));
727    }
728
729    #[test]
730    fn guard_first_dangerous_short_circuits() {
731        // 第一条 DROP TABLE 未确认,应在 NeedsConfirm 处停下。
732        let asts = parse("DROP TABLE a; SELECT 1");
733        assert!(matches!(
734            check_guards(&asts, false),
735            GuardResult::NeedsConfirm
736        ));
737    }
738
739    // ── statement_type_name ───────────────────────────────────────
740
741    #[test]
742    fn statement_type_name_maps_variants() {
743        assert_eq!(statement_type_name(&parse("SELECT 1")[0]), "Select");
744        assert_eq!(
745            statement_type_name(&parse("INSERT INTO t (a) VALUES (1)")[0]),
746            "Insert"
747        );
748        assert_eq!(
749            statement_type_name(&parse("UPDATE t SET a = 1 WHERE id = 1")[0]),
750            "Update"
751        );
752        assert_eq!(
753            statement_type_name(&parse("DELETE FROM t WHERE id = 1")[0]),
754            "Delete"
755        );
756        assert_eq!(
757            statement_type_name(&parse("CREATE TABLE t (id int)")[0]),
758            "CreateTable"
759        );
760        assert_eq!(
761            statement_type_name(&parse("ALTER TABLE t ADD COLUMN x int")[0]),
762            "AlterTable"
763        );
764        assert_eq!(statement_type_name(&parse("DROP TABLE t")[0]), "Drop");
765        assert_eq!(statement_type_name(&parse("TRUNCATE t")[0]), "Truncate");
766    }
767
768    // ── is_read_only ──────────────────────────────────────────────
769
770    #[test]
771    fn is_read_only_classifies_correctly() {
772        assert!(is_read_only(&parse("SELECT 1")[0]));
773        assert!(is_read_only(&parse("EXPLAIN SELECT 1")[0]));
774        assert!(!is_read_only(&parse("UPDATE t SET a = 1 WHERE id = 1")[0]));
775        assert!(!is_read_only(&parse("DELETE FROM t WHERE id = 1")[0]));
776        assert!(!is_read_only(&parse("INSERT INTO t (a) VALUES (1)")[0]));
777    }
778
779    // ── writes_affect_cache(SQL 控制台兜底失效开关) ─────────────
780    // 任何写语句(非只读)都应触发全量缓存失效,兜底绕过 server function 的直改 DB。
781
782    #[test]
783    fn writes_affect_cache_true_for_write_statements() {
784        for sql in [
785            "UPDATE posts SET deleted_at = NOW() WHERE id = 648",
786            "DELETE FROM posts WHERE id = 648",
787            "INSERT INTO posts (title) VALUES ('x')",
788            "TRUNCATE posts",
789            "ALTER TABLE posts ADD COLUMN x int",
790            "DROP TABLE posts",
791        ] {
792            assert!(
793                writes_affect_cache(&parse(sql)),
794                "写语句应触发兜底失效:{sql:?}"
795            );
796        }
797    }
798
799    #[test]
800    fn writes_affect_cache_false_for_read_only_statements() {
801        for sql in [
802            "SELECT 1",
803            "EXPLAIN SELECT * FROM posts",
804            "SELECT id FROM posts WHERE id = 1",
805        ] {
806            assert!(
807                !writes_affect_cache(&parse(sql)),
808                "只读语句不应触发兜底失效:{sql:?}"
809            );
810        }
811    }
812
813    #[test]
814    fn writes_affect_cache_mixed_statements_flagged() {
815        // 多语句中混入任一写语句即应触发(与 allow_multi 语义一致)。
816        let stmts = parse("SELECT 1; UPDATE posts SET deleted_at = NULL WHERE id = 648");
817        assert!(writes_affect_cache(&stmts));
818    }
819    #[test]
820    fn format_sql_error_truncates_long_messages() {
821        let message = format_sql_error("x".repeat(2_500));
822        assert_eq!(message.chars().count(), 2_000);
823    }
824}