Skip to main content

sequel_mcp/backup/
restore.rs

1//! Backup restore (legacy `backup/restore.ts` port): build a
2//! dialect-specific replay plan from a stored backup row, then execute
3//! it inside one transaction. Row backups replay as upserts, schema
4//! backups as their captured `CREATE TABLE`, and insert-hint backups as
5//! the DELETE of exactly the rows the original INSERT created. The
6//! restore always goes through the normal gate when invoked as a tool
7//! (it counts as a write).
8
9use crate::audit::AuditDb;
10use std::sync::Arc;
11use thiserror::Error;
12
13#[derive(Debug, Clone, Copy, PartialEq, Eq)]
14pub enum RestoreDialect {
15    MySql,
16    SQLite,
17}
18
19#[derive(Debug, Error)]
20pub enum RestoreError {
21    #[error("backup #{0} not found")]
22    NotFound(i64),
23    #[error("restore planning failed: {0}")]
24    Plan(String),
25    #[error("restore execution failed: {0}")]
26    Execution(String),
27}
28
29#[derive(Debug, Clone)]
30pub struct RestorePlan {
31    pub backup_id: i64,
32    pub statements: Vec<String>,
33    pub row_count: u64,
34    pub warnings: Vec<String>,
35    /// Set for insert-hint plans (the dangerous DELETE case).
36    pub is_insert_hint_delete: bool,
37}
38
39/// The stored detail of one backup row.
40#[derive(Debug, Clone)]
41pub struct BackupDetail {
42    pub id: i64,
43    pub ts: String,
44    pub connection: String,
45    pub database: Option<String>,
46    pub table_name: String,
47    pub backup_kind: String,
48    pub rows_json: Option<String>,
49    pub schema_sql: Option<String>,
50    pub primary_key: Option<String>,
51    pub row_count: u64,
52    pub truncated: bool,
53}
54
55pub fn get_backup(audit: &Arc<AuditDb>, id: i64) -> Option<BackupDetail> {
56    audit
57        .with(|c| {
58            c.query_row(
59                "SELECT id, ts, connection, database, table_name, backup_kind,
60                        rows_json, schema_sql, primary_key, row_count, truncated
61                   FROM backup WHERE id = ?1",
62                [id],
63                |r| {
64                    Ok(BackupDetail {
65                        id: r.get(0)?,
66                        ts: r.get(1)?,
67                        connection: r.get(2)?,
68                        database: r.get(3)?,
69                        table_name: r.get(4)?,
70                        backup_kind: r.get(5)?,
71                        rows_json: r.get(6)?,
72                        schema_sql: r.get(7)?,
73                        primary_key: r.get(8)?,
74                        row_count: r.get::<_, i64>(9)?.max(0) as u64,
75                        truncated: r.get::<_, i64>(10)? != 0,
76                    })
77                },
78            )
79        })
80        .ok()
81}
82
83fn quote_id(id: &str) -> String {
84    format!("`{}`", id.replace('`', "``"))
85}
86
87fn table_ref_sql(db: Option<&str>, table: &str) -> String {
88    match db {
89        Some(db) => format!("{}.{}", quote_id(db), quote_id(table)),
90        None => quote_id(table),
91    }
92}
93
94fn escape_value(v: &serde_json::Value, dialect: RestoreDialect) -> String {
95    match v {
96        serde_json::Value::Null => "NULL".into(),
97        serde_json::Value::Bool(b) => u8::from(*b).to_string(),
98        serde_json::Value::Number(n) => n.to_string(),
99        serde_json::Value::String(s) => {
100            let escaped = s.replace('\\', "\\\\").replace('\'', "''");
101            format!("'{escaped}'")
102        }
103        // Structured binary values come back as
104        // {"type":"binary","encoding":"base64","data":…} — restore them
105        // as SQL blob literals.
106        serde_json::Value::Object(o)
107            if o.get("type").and_then(|t| t.as_str()) == Some("binary")
108                && o.get("encoding").and_then(|e| e.as_str()) == Some("base64") =>
109        {
110            use base64::Engine;
111            let data = o.get("data").and_then(|d| d.as_str()).unwrap_or("");
112            match base64::engine::general_purpose::STANDARD.decode(data) {
113                Ok(bytes) => {
114                    let hex = bytes.iter().map(|b| format!("{b:02x}")).collect::<String>();
115                    match dialect {
116                        RestoreDialect::SQLite => format!("X'{hex}'"),
117                        RestoreDialect::MySql => format!("0x{hex}"),
118                    }
119                }
120                Err(_) => "NULL".into(),
121            }
122        }
123        // Other objects/arrays (legacy stored JSON text) round-trip as
124        // escaped JSON strings.
125        other => {
126            let text = other.to_string();
127            let escaped = text.replace('\\', "\\\\").replace('\'', "''");
128            format!("'{escaped}'")
129        }
130    }
131}
132
133/// Build the replay plan for backup `id` in the given dialect.
134pub fn plan_restore(
135    audit: &Arc<AuditDb>,
136    id: i64,
137    dialect: RestoreDialect,
138) -> Result<RestorePlan, RestoreError> {
139    let backup = get_backup(audit, id).ok_or(RestoreError::NotFound(id))?;
140    let mut statements: Vec<String> = Vec::new();
141    let mut warnings: Vec<String> = Vec::new();
142    let mut is_insert_hint_delete = false;
143
144    if backup.truncated {
145        warnings.push("backup was truncated; restore will not be complete".into());
146    }
147
148    if backup.backup_kind == "schema" {
149        if let Some(schema_sql) = &backup.schema_sql {
150            warnings.push(
151                "schema-only backup; running this will fail unless the table was dropped first"
152                    .into(),
153            );
154            statements.push(format!("{schema_sql};"));
155        }
156        return Ok(RestorePlan {
157            backup_id: id,
158            statements,
159            row_count: 0,
160            warnings,
161            is_insert_hint_delete,
162        });
163    }
164
165    if backup.backup_kind == "combined"
166        && let Some(schema_sql) = &backup.schema_sql
167    {
168        statements.push(format!("{schema_sql};"));
169    }
170
171    if backup.backup_kind == "insert-hint" {
172        let Some(pk_json) = &backup.primary_key else {
173            warnings.push(
174                "insert-hint backup has no recoverable PK metadata; nothing to restore".into(),
175            );
176            return Ok(RestorePlan {
177                backup_id: id,
178                statements,
179                row_count: 0,
180                warnings,
181                is_insert_hint_delete,
182            });
183        };
184        let tref = table_ref_sql(backup.database.as_deref(), &backup.table_name);
185        let pk: serde_json::Value = serde_json::from_str(pk_json)
186            .map_err(|e| RestoreError::Plan(format!("bad PK metadata: {e}")))?;
187        match pk["kind"].as_str() {
188            Some("range") => {
189                let column = pk["column"].as_str().unwrap_or_default();
190                let start = pk["start"].clone();
191                let end = pk["end"].clone();
192                statements.push(format!(
193                    "DELETE FROM {tref} WHERE {} BETWEEN {} AND {};",
194                    quote_id(column),
195                    escape_value(&start, dialect),
196                    escape_value(&end, dialect),
197                ));
198            }
199            Some("explicit") => {
200                let columns: Vec<String> = pk["columns"]
201                    .as_array()
202                    .map(|a| a.iter().filter_map(|c| c.as_str().map(quote_id)).collect())
203                    .unwrap_or_default();
204                let rows = pk["values"].as_array().map(|a| {
205                    a.iter()
206                        .map(|row| {
207                            let vals: Vec<String> = row
208                                .as_array()
209                                .map(|r| r.iter().map(|v| escape_value(v, dialect)).collect())
210                                .unwrap_or_default();
211                            format!("({})", vals.join(", "))
212                        })
213                        .collect::<Vec<_>>()
214                        .join(", ")
215                });
216                if !columns.is_empty()
217                    && let Some(rows) = rows
218                    && !rows.is_empty()
219                {
220                    statements.push(format!(
221                        "DELETE FROM {tref} WHERE ({}) IN ({rows});",
222                        columns.join(", ")
223                    ));
224                }
225            }
226            _ => {}
227        }
228        warnings.push(
229            "insert-hint restore deletes the rows the INSERT created — verify before running"
230                .into(),
231        );
232        is_insert_hint_delete = true;
233        return Ok(RestorePlan {
234            backup_id: id,
235            statements,
236            row_count: backup.row_count,
237            warnings,
238            is_insert_hint_delete,
239        });
240    }
241
242    // Row backup: per-row upserts with dialect-specific conflict arms.
243    if let Some(rows_json) = &backup.rows_json {
244        let rows: Vec<serde_json::Value> = serde_json::from_str(rows_json)
245            .map_err(|e| RestoreError::Plan(format!("bad rows payload: {e}")))?;
246        if let Some(first) = rows.first().and_then(|r| r.as_object()) {
247            let cols: Vec<String> = first.keys().cloned().collect();
248            let col_list = cols
249                .iter()
250                .map(|c| quote_id(c))
251                .collect::<Vec<_>>()
252                .join(", ");
253            let update_clause = cols
254                .iter()
255                .map(|c| match dialect {
256                    RestoreDialect::SQLite => format!("{} = excluded.{}", quote_id(c), quote_id(c)),
257                    RestoreDialect::MySql => format!("{} = VALUES({})", quote_id(c), quote_id(c)),
258                })
259                .collect::<Vec<_>>()
260                .join(", ");
261            let tref = table_ref_sql(backup.database.as_deref(), &backup.table_name);
262            for row in &rows {
263                let Some(obj) = row.as_object() else { continue };
264                let values = cols
265                    .iter()
266                    .map(|c| escape_value(obj.get(c).unwrap_or(&serde_json::Value::Null), dialect))
267                    .collect::<Vec<_>>()
268                    .join(", ");
269                statements.push(match dialect {
270                    RestoreDialect::SQLite => format!(
271                        "INSERT INTO {tref} ({col_list}) VALUES ({values}) ON CONFLICT DO UPDATE SET {update_clause};"
272                    ),
273                    RestoreDialect::MySql => format!(
274                        "INSERT INTO {tref} ({col_list}) VALUES ({values}) ON DUPLICATE KEY UPDATE {update_clause};"
275                    ),
276                });
277            }
278        }
279    }
280
281    Ok(RestorePlan {
282        backup_id: id,
283        statements,
284        row_count: backup.row_count,
285        warnings,
286        is_insert_hint_delete,
287    })
288}
289
290#[derive(Debug)]
291pub struct RestoreOutcome {
292    pub statements_run: usize,
293    pub affected: u64,
294}
295
296/// Execute a plan against an open SQLite handle inside the caller's
297/// transaction.
298pub fn execute_restore_sqlite(
299    db: &rusqlite::Connection,
300    plan: &RestorePlan,
301) -> Result<RestoreOutcome, RestoreError> {
302    let mut affected: u64 = 0;
303    for stmt_raw in &plan.statements {
304        let stmt_text = stmt_raw.trim_end_matches(|c: char| c == ';' || c.is_whitespace());
305        let mut stmt = db
306            .prepare(stmt_text)
307            .map_err(|e| RestoreError::Execution(e.to_string()))?;
308        if stmt.column_count() > 0 {
309            // Reader statements (rare in plans) are drained.
310            let rows = stmt
311                .query(())
312                .map_err(|e| RestoreError::Execution(e.to_string()))?;
313            drop(rows);
314        } else {
315            stmt.execute(())
316                .map_err(|e| RestoreError::Execution(e.to_string()))?;
317            affected += db.changes();
318        }
319    }
320    Ok(RestoreOutcome {
321        statements_run: plan.statements.len(),
322        affected,
323    })
324}