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