Skip to main content

sequel_mcp/backup/
extractor.rs

1//! Backup spec extraction (`backup/extractor.ts` port): what to capture
2//! before a mutation so it can be undone.
3
4use crate::policy::classifier::Dialect;
5use sqlparser::ast::{Expr, Statement};
6use sqlparser::parser::Parser;
7use thiserror::Error;
8
9#[derive(Debug, Clone, PartialEq, serde::Serialize)]
10pub struct BackupTable {
11    pub db: Option<String>,
12    pub table: String,
13    pub select_sql: String,
14    pub locking: &'static str, // "FOR UPDATE" | "NONE"
15}
16
17#[derive(Debug, Clone, PartialEq, serde::Serialize)]
18#[serde(tag = "kind", rename_all = "kebab-case")]
19pub enum BackupSpec {
20    None {
21        #[serde(skip_serializing_if = "Option::is_none")]
22        reason: Option<String>,
23    },
24    Rows {
25        tables: Vec<BackupTable>,
26    },
27    Schema {
28        tables: Vec<SchemaTable>,
29    },
30    Combined {
31        tables: Vec<BackupTable>,
32    },
33    InsertHint {
34        table: SchemaTable,
35        columns: Vec<String>,
36        explicit_pk_values: Option<Vec<Vec<serde_json::Value>>>,
37    },
38}
39
40#[derive(Debug, Clone, PartialEq, serde::Serialize)]
41pub struct SchemaTable {
42    pub db: Option<String>,
43    pub table: String,
44}
45
46#[derive(Debug, Error)]
47pub enum ExtractError {
48    #[error("parse error: {0}")]
49    Parse(String),
50}
51
52pub fn quote_ident(id: &str) -> String {
53    format!("`{}`", id.replace('`', "``"))
54}
55
56pub fn table_ref_sql(db: &Option<String>, table: &str) -> String {
57    match db {
58        Some(db) => format!("{}.{}", quote_ident(db), quote_ident(table)),
59        None => quote_ident(table),
60    }
61}
62
63fn object_parts(name: &sqlparser::ast::ObjectName) -> (Option<String>, String) {
64    let mut parts = Vec::new();
65    for p in &name.0 {
66        if let sqlparser::ast::ObjectNamePart::Identifier(i) = p {
67            parts.push(i.value.clone());
68        }
69    }
70    match parts.len() {
71        0 => (None, String::new()),
72        1 => (None, parts.remove(0)),
73        _ => {
74            let table = parts.pop().unwrap_or_default();
75            (parts.pop(), table)
76        }
77    }
78}
79
80fn expr_sql(e: &Expr) -> String {
81    e.to_string()
82}
83
84fn lock_suffix(dialect: Dialect) -> &'static str {
85    if dialect == Dialect::MySql {
86        " FOR UPDATE"
87    } else {
88        ""
89    }
90}
91
92/// Render the FROM-side of the statement (tables + joins) exactly as
93/// written, so backup SELECTs keep join semantics.
94fn from_sql(twj: &sqlparser::ast::TableWithJoins) -> String {
95    twj.to_string()
96}
97
98pub fn is_backup_required(ast_type: &str) -> bool {
99    matches!(
100        ast_type,
101        "update" | "delete" | "replace" | "insert" | "truncate" | "drop" | "alter" | "rename"
102    )
103}
104
105/// Lock suffixes recognised after stripping; each must move *after* the
106/// appended LIMIT to stay valid MySQL.
107const TRAILING_LOCK_SUFFIXES: [&str; 5] = [
108    " FOR UPDATE NOWAIT",
109    " FOR UPDATE SKIP LOCKED",
110    " FOR UPDATE",
111    " FOR SHARE",
112    " LOCK IN SHARE MODE",
113];
114
115/// Append `LIMIT n` to a backup SELECT, preserving any trailing lock
116/// clause. Returns `None` when the query cannot be safely rewritten
117/// (existing LIMIT/OFFSET, set operations, CTE, trailing semicolon or
118/// comment, parenthesized tail, or an unmatched lock clause) — callers
119/// must DENY in that case rather than execute the unbounded original.
120pub fn with_limit(select_sql: &str, n: u64) -> Option<String> {
121    if select_sql.contains(';') {
122        return None; // trailing semicolon or statement separator: refuse
123    }
124    let lower = select_sql.to_ascii_lowercase();
125    if lower.contains(" limit ") || lower.ends_with(" limit") {
126        return None; // existing LIMIT: composing another is invalid
127    }
128    if lower.contains(" offset ") {
129        return None;
130    }
131    if lower.starts_with("with ") || lower.contains(" union ") {
132        return None; // set operations/CTEs: not safely suffix-rewritable
133    }
134    if select_sql.contains("--") || select_sql.contains("/*") {
135        return None; // comments could hide a tail we fail to see
136    }
137    if select_sql.ends_with(')') {
138        return None; // parenthesized query expression: refuse
139    }
140    for suffix in TRAILING_LOCK_SUFFIXES {
141        if let Some(base) = select_sql.strip_suffix(suffix) {
142            return Some(format!("{base} LIMIT {n}{suffix}"));
143        }
144    }
145    if lower.contains(" for update")
146        || lower.contains(" for share")
147        || lower.contains(" for key share")
148        || lower.contains(" lock in share mode")
149    {
150        return None; // a lock clause we failed to match exactly: refuse
151    }
152    Some(format!("{select_sql} LIMIT {n}"))
153}
154
155const PK_GUESS_NAMES: [&str; 3] = ["id", "uuid", "pk"];
156
157fn guess_pk_column(columns: &[String]) -> Option<usize> {
158    for guess in PK_GUESS_NAMES {
159        if let Some(i) = columns.iter().position(|c| c.eq_ignore_ascii_case(guess)) {
160            return Some(i);
161        }
162    }
163    None
164}
165
166fn literal_json(v: &Expr) -> serde_json::Value {
167    match v {
168        Expr::Value(val) => match &val.value {
169            sqlparser::ast::Value::Number(n, _) => {
170                serde_json::from_str(n).unwrap_or_else(|_| serde_json::json!(n))
171            }
172            sqlparser::ast::Value::SingleQuotedString(s)
173            | sqlparser::ast::Value::DoubleQuotedString(s) => serde_json::json!(s),
174            sqlparser::ast::Value::NationalStringLiteral(s) => serde_json::json!(s),
175            sqlparser::ast::Value::HexStringLiteral(s) => serde_json::json!(s),
176            sqlparser::ast::Value::Null => serde_json::Value::Null,
177            sqlparser::ast::Value::Placeholder(p) => serde_json::json!(p),
178            sqlparser::ast::Value::EscapedStringLiteral(s) => serde_json::json!(s),
179            sqlparser::ast::Value::SingleQuotedByteStringLiteral(s)
180            | sqlparser::ast::Value::DoubleQuotedByteStringLiteral(s) => serde_json::json!(s),
181            _ => serde_json::Value::Null,
182        },
183        _ => serde_json::Value::Null,
184    }
185}
186
187/// Build the backup spec for a statement. `ast_type` is the legacy name
188/// from the classifier.
189pub fn extract_backup_spec(
190    sql: &str,
191    ast_type: &str,
192    dialect: Dialect,
193) -> Result<BackupSpec, ExtractError> {
194    let stmts = Parser::parse_sql(
195        match dialect {
196            Dialect::MySql => {
197                &sqlparser::dialect::MySqlDialect {} as &dyn sqlparser::dialect::Dialect
198            }
199            Dialect::SQLite => {
200                &sqlparser::dialect::SQLiteDialect {} as &dyn sqlparser::dialect::Dialect
201            }
202        },
203        sql,
204    )
205    .map_err(|e| ExtractError::Parse(e.to_string()))?;
206    let Some(stmt) = stmts.into_iter().next() else {
207        return Ok(BackupSpec::None {
208            reason: Some("no statement".into()),
209        });
210    };
211
212    match ast_type {
213        "update" => update_spec(&stmt, dialect),
214        "delete" => delete_spec(&stmt, dialect),
215        "replace" => replace_spec(&stmt, dialect),
216        "insert" => insert_spec(&stmt),
217        "truncate" => truncate_spec(&stmt),
218        "drop" => drop_spec(&stmt),
219        "alter" | "rename" => schema_spec(&stmt),
220        _ => Ok(BackupSpec::None {
221            reason: Some(format!("no backup strategy for AST type \"{ast_type}\"")),
222        }),
223    }
224}
225
226fn where_suffix(where_expr: Option<&Expr>) -> String {
227    where_expr
228        .map(|w| format!(" WHERE {}", expr_sql(w)))
229        .unwrap_or_default()
230}
231
232fn update_spec(stmt: &Statement, dialect: Dialect) -> Result<BackupSpec, ExtractError> {
233    let Statement::Update(u) = stmt else {
234        return Ok(none("not an UPDATE"));
235    };
236    let from = &u.table;
237    let where_sql = where_suffix(u.selection.as_ref());
238    let lock = lock_suffix(dialect);
239
240    // SET-target qualification decides which joined tables get row backups
241    // (legacy `inferMutatedTables`); unqualified SET falls back to the
242    // first table.
243    let mut qualified: Vec<String> = Vec::new();
244    for a in &u.assignments {
245        if let sqlparser::ast::AssignmentTarget::ColumnName(id) = &a.target {
246            let parts = &id.0;
247            if parts.len() == 2
248                && let sqlparser::ast::ObjectNamePart::Identifier(tbl) = &parts[0]
249                && !qualified.contains(&tbl.value)
250            {
251                qualified.push(tbl.value.clone());
252            }
253        }
254    }
255
256    let all_tables = collect_tables_with_aliases(from);
257    let targets: Vec<(Option<String>, String, String)> = if qualified.is_empty() {
258        let (db, table, rendered) = all_tables.first().cloned().unwrap_or_default();
259        vec![(db, table, rendered)]
260    } else {
261        all_tables
262            .iter()
263            .filter(|(_, table, _)| qualified.iter().any(|q| table.eq_ignore_ascii_case(q)))
264            .cloned()
265            .collect()
266    };
267
268    let full_from = from_sql(from);
269    let mut tables = Vec::new();
270    for (db, table, _rendered) in targets {
271        let select = if all_tables.len() == 1 && where_sql.is_empty() {
272            format!("SELECT * FROM {}{}", table_ref_sql(&db, &table), lock)
273        } else {
274            let alias_col = format!("{}.*", quote_ident(&table));
275            format!(
276                "SELECT {} FROM {}{}{}",
277                alias_col, full_from, where_sql, lock
278            )
279        };
280        tables.push(BackupTable {
281            db,
282            table,
283            select_sql: select,
284            locking: if lock.is_empty() {
285                "NONE"
286            } else {
287                "FOR UPDATE"
288            },
289        });
290    }
291    if tables.is_empty() {
292        return Ok(none("no mutated tables identified in multi-table UPDATE"));
293    }
294    Ok(BackupSpec::Rows { tables })
295}
296
297/// (db, table, rendered-alias-or-name) for every table factor in FROM.
298fn collect_tables_with_aliases(
299    twj: &sqlparser::ast::TableWithJoins,
300) -> Vec<(Option<String>, String, String)> {
301    let mut out = Vec::new();
302    let mut visit = |f: &sqlparser::ast::TableFactor| {
303        if let sqlparser::ast::TableFactor::Table { name, alias, .. } = f {
304            let (db, table) = object_parts(name);
305            let rendered = alias
306                .as_ref()
307                .map(|a| a.name.value.clone())
308                .unwrap_or_else(|| table.clone());
309            out.push((db, table, rendered));
310        }
311    };
312    visit(&twj.relation);
313    for j in &twj.joins {
314        visit(&j.relation);
315    }
316    out
317}
318
319fn delete_spec(stmt: &Statement, dialect: Dialect) -> Result<BackupSpec, ExtractError> {
320    let Statement::Delete(d) = stmt else {
321        return Ok(none("not a DELETE"));
322    };
323    let where_sql = where_suffix(d.selection.as_ref());
324    let lock = lock_suffix(dialect);
325
326    let mut from_tables: Vec<(Option<String>, String, String)> = Vec::new();
327    let collect = |twjs: &Vec<sqlparser::ast::TableWithJoins>, out: &mut Vec<_>| {
328        for twj in twjs {
329            out.extend(collect_tables_with_aliases(twj));
330        }
331    };
332    match &d.from {
333        sqlparser::ast::FromTable::WithFromKeyword(t) => collect(t, &mut from_tables),
334        sqlparser::ast::FromTable::WithoutKeyword(t) => collect(t, &mut from_tables),
335    }
336    if let Some(using) = &d.using {
337        collect(using, &mut from_tables);
338    }
339
340    // DELETE targets: `DELETE a, b FROM …` names them explicitly.
341    let explicit: Vec<(Option<String>, String)> = d.tables.iter().map(object_parts).collect();
342
343    let targets: Vec<(Option<String>, String)> = if !explicit.is_empty() {
344        explicit
345            .into_iter()
346            .map(|(db, table)| {
347                let resolved = from_tables
348                    .iter()
349                    .find(|(_, t, _)| t.eq_ignore_ascii_case(&table))
350                    .map(|(fdb, _, _)| fdb.clone())
351                    .unwrap_or(db);
352                (resolved, table)
353            })
354            .collect()
355    } else {
356        from_tables
357            .iter()
358            .map(|(db, t, _)| (db.clone(), t.clone()))
359            .collect()
360    };
361
362    if from_tables.is_empty() {
363        return Ok(none("DELETE has no source table"));
364    }
365
366    let mut tables = Vec::new();
367    for (db, table) in targets {
368        let select = if from_tables.len() == 1 && where_sql.is_empty() {
369            format!("SELECT * FROM {}{}", table_ref_sql(&db, &table), lock)
370        } else {
371            let rendered = from_tables
372                .iter()
373                .find(|(_, t, _)| t.eq_ignore_ascii_case(&table))
374                .map(|(_, _, r)| r.clone())
375                .unwrap_or_else(|| table.clone());
376            format!(
377                "SELECT {}.* FROM {}{}{}",
378                quote_ident(&rendered),
379                from_tables
380                    .iter()
381                    .map(|(db, t, _)| table_ref_sql(db, t))
382                    .collect::<Vec<_>>()
383                    .join(", "),
384                where_sql,
385                lock
386            )
387        };
388        tables.push(BackupTable {
389            db,
390            table,
391            select_sql: select,
392            locking: if lock.is_empty() {
393                "NONE"
394            } else {
395                "FOR UPDATE"
396            },
397        });
398    }
399    Ok(BackupSpec::Rows { tables })
400}
401
402fn replace_spec(stmt: &Statement, dialect: Dialect) -> Result<BackupSpec, ExtractError> {
403    let Statement::Insert(ins) = stmt else {
404        return Ok(none("not a REPLACE"));
405    };
406    if !ins.replace_into {
407        return Ok(none("not a REPLACE"));
408    }
409    let Some((db, table)) = table_of(ins) else {
410        return Ok(none("REPLACE target unclear"));
411    };
412    let columns: Vec<String> = ins
413        .columns
414        .iter()
415        .map(|c| {
416            c.0.iter()
417                .filter_map(|p| match p {
418                    sqlparser::ast::ObjectNamePart::Identifier(i) => Some(i.value.clone()),
419                    _ => None,
420                })
421                .next()
422                .unwrap_or_default()
423        })
424        .collect();
425    let rows = literal_rows(ins);
426    if let (Some(pk_idx), Some(all_rows)) = (guess_pk_column(&columns), rows)
427        && !all_rows.is_empty()
428    {
429        let values: Vec<String> = all_rows
430            .iter()
431            .map(|row| sql_literal(&row[pk_idx]))
432            .collect();
433        let lock = lock_suffix(dialect);
434        let select = format!(
435            "SELECT * FROM {} WHERE {} IN ({}){}",
436            table_ref_sql(&db, &table),
437            quote_ident(&columns[pk_idx]),
438            values.join(", "),
439            lock
440        );
441        return Ok(BackupSpec::Rows {
442            tables: vec![BackupTable {
443                db,
444                table,
445                select_sql: select,
446                locking: if lock.is_empty() {
447                    "NONE"
448                } else {
449                    "FOR UPDATE"
450                },
451            }],
452        });
453    }
454    Ok(none(
455        "REPLACE without identifiable PK column; no backup taken",
456    ))
457}
458
459fn insert_spec(stmt: &Statement) -> Result<BackupSpec, ExtractError> {
460    let Statement::Insert(ins) = stmt else {
461        return Ok(none("not an INSERT"));
462    };
463    let Some((db, table)) = table_of(ins) else {
464        return Ok(none("INSERT target unclear"));
465    };
466    let columns: Vec<String> = ins
467        .columns
468        .iter()
469        .map(|c| {
470            c.0.iter()
471                .filter_map(|p| match p {
472                    sqlparser::ast::ObjectNamePart::Identifier(i) => Some(i.value.clone()),
473                    _ => None,
474                })
475                .next()
476                .unwrap_or_default()
477        })
478        .collect();
479    let rows = literal_rows(ins);
480    let explicit = match (guess_pk_column(&columns), rows) {
481        (Some(idx), Some(all)) if !all.is_empty() => {
482            Some(all.iter().map(|r| vec![r[idx].clone()]).collect())
483        }
484        _ => None,
485    };
486    Ok(BackupSpec::InsertHint {
487        table: SchemaTable { db, table },
488        columns,
489        explicit_pk_values: explicit,
490    })
491}
492
493fn table_of(ins: &sqlparser::ast::Insert) -> Option<(Option<String>, String)> {
494    match &ins.table {
495        sqlparser::ast::TableObject::TableName(name) => Some(object_parts(name)),
496        _ => None,
497    }
498}
499
500fn literal_rows(ins: &sqlparser::ast::Insert) -> Option<Vec<Vec<serde_json::Value>>> {
501    let source = ins.source.as_ref()?;
502    match source.body.as_ref() {
503        sqlparser::ast::SetExpr::Values(values) => Some(
504            values
505                .rows
506                .iter()
507                .map(|row| row.iter().map(literal_json).collect())
508                .collect(),
509        ),
510        _ => None,
511    }
512}
513
514fn sql_literal(v: &serde_json::Value) -> String {
515    match v {
516        serde_json::Value::Null => "NULL".into(),
517        serde_json::Value::Bool(b) => if *b { "1" } else { "0" }.into(),
518        serde_json::Value::Number(n) => n.to_string(),
519        serde_json::Value::String(s) => format!("'{}'", s.replace('\'', "''")),
520        other => format!("'{}'", other.to_string().replace('\'', "''")),
521    }
522}
523
524fn truncate_spec(stmt: &Statement) -> Result<BackupSpec, ExtractError> {
525    let Statement::Truncate(t) = stmt else {
526        return Ok(none("not a TRUNCATE"));
527    };
528    let Some(target) = t.table_names.first() else {
529        return Ok(none("TRUNCATE target unclear"));
530    };
531    let (db, table) = object_parts(&target.name);
532    let select_sql = format!("SELECT * FROM {}", table_ref_sql(&db, &table));
533    Ok(BackupSpec::Combined {
534        tables: vec![BackupTable {
535            db,
536            table,
537            select_sql,
538            locking: "NONE",
539        }],
540    })
541}
542
543fn drop_spec(stmt: &Statement) -> Result<BackupSpec, ExtractError> {
544    let Statement::Drop { names, .. } = stmt else {
545        return Ok(none("not a DROP"));
546    };
547    let Some(name) = names.first() else {
548        return Ok(none("DROP target unclear"));
549    };
550    let (db, table) = object_parts(name);
551    let select_sql = format!("SELECT * FROM {}", table_ref_sql(&db, &table));
552    Ok(BackupSpec::Combined {
553        tables: vec![BackupTable {
554            db,
555            table,
556            select_sql,
557            locking: "NONE",
558        }],
559    })
560}
561
562fn schema_spec(stmt: &Statement) -> Result<BackupSpec, ExtractError> {
563    match stmt {
564        Statement::AlterTable(a) => {
565            let (db, table) = object_parts(&a.name);
566            Ok(BackupSpec::Schema {
567                tables: vec![SchemaTable { db, table }],
568            })
569        }
570        Statement::RenameTable(renames) => {
571            let mut tables = Vec::new();
572            for rn in renames {
573                let (odb, otable) = object_parts(&rn.old_name);
574                let (ndb, ntable) = object_parts(&rn.new_name);
575                tables.push(SchemaTable {
576                    db: odb,
577                    table: otable,
578                });
579                let _ = (ndb, ntable);
580            }
581            Ok(BackupSpec::Schema { tables })
582        }
583        _ => Ok(none("not ALTER/RENAME")),
584    }
585}
586
587fn none(reason: &str) -> BackupSpec {
588    BackupSpec::None {
589        reason: Some(reason.to_string()),
590    }
591}
592
593#[cfg(test)]
594mod tests {
595    use super::*;
596
597    fn spec(sql: &str, ast_type: &str) -> BackupSpec {
598        extract_backup_spec(sql, ast_type, Dialect::MySql).unwrap()
599    }
600
601    #[test]
602    fn update_with_where_selects_for_update() {
603        let s = spec("UPDATE users SET name = 'x' WHERE id = 1", "update");
604        match s {
605            BackupSpec::Rows { tables } => {
606                assert_eq!(tables.len(), 1);
607                assert_eq!(tables[0].table, "users");
608                assert!(tables[0].select_sql.contains("FOR UPDATE"));
609                assert!(
610                    tables[0].select_sql.contains("WHERE"),
611                    "select keeps the WHERE: {}",
612                    tables[0].select_sql
613                );
614                assert!(tables[0].select_sql.contains("id"));
615            }
616            other => panic!("{other:?}"),
617        }
618    }
619
620    #[test]
621    fn update_without_where_backs_up_all_rows() {
622        let s = spec("UPDATE users SET name = 'x'", "update");
623        match s {
624            BackupSpec::Rows { tables } => {
625                assert!(tables[0].select_sql.starts_with("SELECT * FROM `users`"));
626                assert!(!tables[0].select_sql.contains("WHERE"));
627            }
628            other => panic!("{other:?}"),
629        }
630    }
631
632    #[test]
633    fn multi_table_update_backs_up_set_targets() {
634        let s = spec(
635            "UPDATE a JOIN b ON a.id = b.id SET a.x = 1, b.y = 2",
636            "update",
637        );
638        match s {
639            BackupSpec::Rows { tables } => assert_eq!(tables.len(), 2),
640            other => panic!("{other:?}"),
641        }
642    }
643
644    #[test]
645    fn delete_multi_target() {
646        let s = spec(
647            "DELETE a FROM a JOIN b ON a.id = b.id WHERE b.flag = 1",
648            "delete",
649        );
650        match s {
651            BackupSpec::Rows { tables } => {
652                assert_eq!(tables.len(), 1);
653                assert_eq!(tables[0].table, "a");
654                assert!(tables[0].select_sql.contains("WHERE"));
655            }
656            other => panic!("{other:?}"),
657        }
658    }
659
660    #[test]
661    fn replace_with_pk_preselects() {
662        let s = spec("REPLACE INTO users (id, name) VALUES (1, 'a')", "replace");
663        match s {
664            BackupSpec::Rows { tables } => {
665                assert!(tables[0].select_sql.contains("`id` IN (1)"));
666                assert!(tables[0].select_sql.contains("FOR UPDATE"));
667            }
668            other => panic!("{other:?}"),
669        }
670    }
671
672    #[test]
673    fn replace_without_pk_is_none() {
674        let s = spec("REPLACE INTO users (name) VALUES ('a')", "replace");
675        assert!(matches!(s, BackupSpec::None { .. }));
676    }
677
678    #[test]
679    fn insert_hint_explicit_pk() {
680        let s = spec("INSERT INTO users (id, name) VALUES (5, 'a')", "insert");
681        match s {
682            BackupSpec::InsertHint {
683                explicit_pk_values: Some(v),
684                ..
685            } => assert_eq!(v, vec![vec![serde_json::json!(5)]]),
686            other => panic!("{other:?}"),
687        }
688    }
689
690    #[test]
691    fn insert_hint_without_pk() {
692        let s = spec("INSERT INTO users (name) VALUES ('a')", "insert");
693        match s {
694            BackupSpec::InsertHint {
695                explicit_pk_values: None,
696                ..
697            } => {}
698            other => panic!("{other:?}"),
699        }
700    }
701
702    #[test]
703    fn truncate_and_drop_are_combined() {
704        assert!(matches!(
705            spec("TRUNCATE TABLE users", "truncate"),
706            BackupSpec::Combined { .. }
707        ));
708        assert!(matches!(
709            spec("DROP TABLE users", "drop"),
710            BackupSpec::Combined { .. }
711        ));
712    }
713
714    #[test]
715    fn alter_is_schema_only() {
716        assert!(matches!(
717            spec("ALTER TABLE users ADD COLUMN email TEXT", "alter"),
718            BackupSpec::Schema { .. }
719        ));
720    }
721
722    #[test]
723    fn d6_limit_forms() {
724        use super::with_limit as wl;
725        // Rewritable forms.
726        assert_eq!(
727            wl("SELECT * FROM t WHERE id = 1 FOR UPDATE", 11).unwrap(),
728            "SELECT * FROM t WHERE id = 1 LIMIT 11 FOR UPDATE"
729        );
730        assert_eq!(
731            wl("SELECT * FROM t FOR UPDATE NOWAIT", 11).unwrap(),
732            "SELECT * FROM t LIMIT 11 FOR UPDATE NOWAIT"
733        );
734        assert_eq!(
735            wl("SELECT * FROM t FOR UPDATE SKIP LOCKED", 11).unwrap(),
736            "SELECT * FROM t LIMIT 11 FOR UPDATE SKIP LOCKED"
737        );
738        assert_eq!(
739            wl("SELECT * FROM t FOR SHARE", 11).unwrap(),
740            "SELECT * FROM t LIMIT 11 FOR SHARE"
741        );
742        assert_eq!(
743            wl("SELECT * FROM t LOCK IN SHARE MODE", 11).unwrap(),
744            "SELECT * FROM t LIMIT 11 LOCK IN SHARE MODE"
745        );
746        // Optimizer hints are comment-delimited; a suffix rewriter cannot
747        // distinguish them from tail-hiding comments, so they are denied.
748        assert!(
749            wl("SELECT /*+ hint */ * FROM t", 11).is_none(),
750            "optimizer hints denied (comment-shaped)"
751        );
752        assert_eq!(
753            wl("SELECT * FROM t", 11).unwrap(),
754            "SELECT * FROM t LIMIT 11"
755        );
756        // Unrewritable forms must return None (deny), never a bad rewrite.
757        assert!(
758            wl("SELECT * FROM t LIMIT 5", 11).is_none(),
759            "existing LIMIT"
760        );
761        assert!(
762            wl("SELECT * FROM t LIMIT 5 OFFSET 2", 11).is_none(),
763            "offset"
764        );
765        assert!(
766            wl("SELECT * FROM a UNION SELECT * FROM b", 11).is_none(),
767            "union"
768        );
769        assert!(
770            wl("WITH x AS (SELECT 1) SELECT * FROM x", 11).is_none(),
771            "cte"
772        );
773        assert!(wl("SELECT * FROM t;", 11).is_none(), "semicolon");
774        assert!(
775            wl("SELECT * FROM t -- comment", 11).is_none(),
776            "line comment"
777        );
778        assert!(wl("SELECT * FROM /* c */ t", 11).is_none(), "block comment");
779        assert!(wl("(SELECT * FROM t)", 11).is_none(), "parenthesized");
780        assert!(
781            wl("SELECT * FROM t FOR KEY SHARE", 11).is_none(),
782            "unknown lock"
783        );
784    }
785
786    #[test]
787    fn create_is_none() {
788        assert!(matches!(
789            spec("CREATE TABLE t (id INT)", "create"),
790            BackupSpec::None { .. }
791        ));
792    }
793}