1use 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 pub is_insert_hint_delete: bool,
37}
38
39#[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 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 => {
126 let text = other.to_string();
127 let escaped = text.replace('\\', "\\\\").replace('\'', "''");
128 format!("'{escaped}'")
129 }
130 }
131}
132
133pub 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 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
296pub 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 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}