1use crate::backup::capture::{capture_backup_sqlite, capture_insert_hint};
6use crate::backup::extractor::{BackupSpec, extract_backup_spec};
7use crate::config::SqliteConnection;
8use crate::policy::classifier::{ClassifiedStatement, ClassifyError, Dialect, classify_statement};
9use crate::policy::model::{Policy, SqlCategory};
10use rusqlite::types::ValueRef;
11use rusqlite::{Connection, OpenFlags, Rows, Statement};
12use std::path::{Path, PathBuf};
13use std::sync::Arc;
14use std::time::{Duration, Instant};
15use thiserror::Error;
16
17use crate::audit::AuditDb;
18
19#[derive(Debug, Error)]
20pub enum SqliteError {
21 #[error("{0}")]
22 Db(#[from] rusqlite::Error),
23 #[error("cannot classify statement: {0}")]
24 Classify(String),
25 #[error("read-only connection: database file {0} not found")]
26 NotFound(String),
27 #[error("read-only connection: path {0} is a symlink — refusing")]
28 Symlink(String),
29 #[error("test mode: {0}")]
30 TestModeRefused(String),
31 #[error("statement timed out after {0}ms")]
32 Timeout(u64),
33 #[error("backup overflow: {0}")]
34 BackupOverflow(String),
35 #[error("backup capture failed: {0} — mutation denied")]
36 BackupFailed(String),
37}
38
39#[derive(Debug)]
40pub struct ExecuteResult {
41 pub ddl_no_op: bool,
43 pub ddl_absent_targets: Vec<String>,
45 pub ddl_executed_targets: Vec<String>,
48 pub warnings: Vec<&'static str>,
50 pub journal_id: Option<i64>,
53 pub rows: Vec<serde_json::Value>,
54 pub fields: Vec<String>,
55 pub affected_rows: u64,
56 pub truncated: bool,
57 pub duration_ms: u64,
58 pub backup_id: Option<i64>,
59 pub backup_row_count: u64,
60}
61
62const READ_CATEGORIES: [SqlCategory; 1] = [SqlCategory::Read];
63
64pub fn open_sqlite_database(
67 conn: &SqliteConnection,
68 readonly: bool,
69 timeout_ms: u32,
70) -> Result<Connection, SqliteError> {
71 let filename = crate::app::paths::expand_tilde(&conn.path);
72 crate::app::test_mode::check_sqlite_path(&filename).map_err(SqliteError::TestModeRefused)?;
74 if readonly {
75 let canonical = std::fs::canonicalize(&filename)
76 .map_err(|_| SqliteError::NotFound(filename.display().to_string()))?;
77 verify_no_symlink_escape(&canonical)?;
78 let db = Connection::open_with_flags(
79 &canonical,
80 OpenFlags::SQLITE_OPEN_READ_ONLY | OpenFlags::SQLITE_OPEN_NO_MUTEX,
81 )?;
82 db.busy_timeout(Duration::from_millis(timeout_ms.max(1) as u64))?;
83 Ok(db)
84 } else {
85 let db = Connection::open_with_flags(
86 &filename,
87 OpenFlags::SQLITE_OPEN_READ_WRITE
88 | OpenFlags::SQLITE_OPEN_CREATE
89 | OpenFlags::SQLITE_OPEN_NO_MUTEX,
90 )?;
91 db.busy_timeout(Duration::from_millis(timeout_ms.max(1) as u64))?;
92 db.pragma_update(None, "foreign_keys", true)?;
93 Ok(db)
94 }
95}
96
97fn verify_no_symlink_escape(canonical: &Path) -> Result<(), SqliteError> {
99 let mut dir = canonical.parent().map(PathBuf::from);
102 while let Some(d) = dir {
103 let meta = std::fs::symlink_metadata(&d);
104 if matches!(meta, Ok(m) if m.file_type().is_symlink()) {
105 return Err(SqliteError::Symlink(d.display().to_string()));
106 }
107 dir = d.parent().map(PathBuf::from);
108 }
109 Ok(())
110}
111
112fn value_to_json(vr: ValueRef<'_>) -> serde_json::Value {
113 match vr {
114 ValueRef::Null => serde_json::Value::Null,
115 ValueRef::Integer(i) => serde_json::json!(i),
116 ValueRef::Real(f) => serde_json::json!(f),
117 ValueRef::Text(t) => serde_json::json!(String::from_utf8_lossy(t)),
118 ValueRef::Blob(b) => {
119 use base64::Engine;
120 serde_json::json!(base64::engine::general_purpose::STANDARD.encode(b))
121 }
122 }
123}
124
125struct StreamStats {
126 rows: Vec<serde_json::Value>,
127 fields: Vec<String>,
128 affected_rows: u64,
129 truncated: bool,
130 insert_id: Option<i64>,
131}
132
133fn run_streaming(
137 db: &Connection,
138 sql: &str,
139 row_cap: u32,
140 byte_cap: u64,
141) -> Result<StreamStats, SqliteError> {
142 let interrupt = db.get_interrupt_handle();
143 let mut stmt: Statement = db.prepare(sql)?;
144 let fields: Vec<String> = stmt.column_names().iter().map(|s| s.to_string()).collect();
145 let is_reader = stmt.column_count() > 0;
148 if !is_reader {
149 stmt.execute([])?;
150 drop(stmt);
151 let _ = interrupt;
152 let stats = get_change_stats(db)?;
153 return Ok(StreamStats {
154 rows: Vec::new(),
155 fields: Vec::new(),
156 affected_rows: stats.0,
157 truncated: false,
158 insert_id: stats.1,
159 });
160 }
161 let mut rows: Rows = stmt.query([])?;
162 let mut out = Vec::new();
163 let mut truncated = false;
164 let mut bytes: u64 = 0;
165 while let Some(row) = rows.next()? {
166 if out.len() >= row_cap as usize {
167 truncated = true;
168 break;
169 }
170 let mut obj = serde_json::Map::with_capacity(fields.len());
171 for (i, name) in fields.iter().enumerate() {
172 let v = value_to_json(row.get_ref(i)?);
173 obj.insert(name.clone(), v);
174 }
175 let val = serde_json::Value::Object(obj);
176 bytes += val.encoded_len() as u64;
177 out.push(val);
178 if bytes > byte_cap {
179 truncated = true;
180 break;
181 }
182 }
183 drop(rows);
184 drop(stmt);
185 let _ = interrupt;
186 let stats = get_change_stats(db)?;
187 Ok(StreamStats {
188 rows: out,
189 fields,
190 affected_rows: stats.0,
191 truncated,
192 insert_id: stats.1,
193 })
194}
195
196trait JsonLen {
197 fn encoded_len(&self) -> usize;
198}
199
200impl JsonLen for serde_json::Value {
201 fn encoded_len(&self) -> usize {
202 serde_json::to_string(self).map(|s| s.len()).unwrap_or(0)
203 }
204}
205
206fn get_change_stats(db: &Connection) -> Result<(u64, Option<i64>), SqliteError> {
207 let affected: i64 = db.query_row("SELECT changes()", [], |r| r.get(0))?;
208 let insert_id: i64 = db.query_row("SELECT last_insert_rowid()", [], |r| r.get(0))?;
209 Ok((
210 affected.max(0) as u64,
211 if insert_id > 0 { Some(insert_id) } else { None },
212 ))
213}
214
215pub struct SqliteExecuteParams<'a> {
216 pub connection: &'a SqliteConnection,
217 pub sql: &'a str,
218 pub classified: &'a ClassifiedStatement,
219 pub policy: &'a Policy,
220 pub database: Option<&'a str>,
221 pub audit: Option<Arc<AuditDb>>,
222}
223
224pub fn execute_sqlite_statement(
229 params: SqliteExecuteParams<'_>,
230) -> Result<ExecuteResult, SqliteError> {
231 let start = Instant::now();
232 let is_read = READ_CATEGORIES.contains(¶ms.classified.category);
233 let db = open_sqlite_database(params.connection, is_read, params.policy.stmt_timeout_ms)?;
234 let interrupt = db.get_interrupt_handle();
235 let done = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
236 let done_flag = done.clone();
237 let timeout_ms = params.policy.stmt_timeout_ms;
238 let timer = std::thread::spawn(move || {
239 let deadline = Instant::now() + Duration::from_millis(timeout_ms.max(1) as u64);
245 while !done_flag.load(std::sync::atomic::Ordering::SeqCst) {
246 let now = Instant::now();
247 if now >= deadline {
248 break;
249 }
250 std::thread::sleep((deadline - now).min(Duration::from_millis(50)));
251 }
252 if !done_flag.load(std::sync::atomic::Ordering::SeqCst) {
253 interrupt.interrupt();
254 }
255 });
256
257 let result = execute_on_connection(&db, ¶ms, start);
258 done.store(true, std::sync::atomic::Ordering::SeqCst);
259 let _ = timer.join();
260 match result {
261 Err(SqliteError::Db(rusqlite::Error::SqliteFailure(ffi, _)))
262 if ffi.code == rusqlite::ErrorCode::OperationInterrupted =>
263 {
264 Err(SqliteError::Timeout(timeout_ms as u64))
265 }
266 other => other,
267 }
268}
269
270fn execute_on_connection(
271 db: &Connection,
272 params: &SqliteExecuteParams<'_>,
273 start: Instant,
274) -> Result<ExecuteResult, SqliteError> {
275 let is_read = READ_CATEGORIES.contains(¶ms.classified.category);
276 let mut in_txn = false;
277 if !is_read && params.classified.category != SqlCategory::TxCtrl {
278 db.execute_batch("BEGIN IMMEDIATE")?;
279 in_txn = true;
280 }
281
282 let result = (|| -> Result<ExecuteResult, SqliteError> {
283 let mut backup_id: Option<i64> = None;
284 let mut backup_row_count: u64 = 0;
285 let mut pending_insert_spec: Option<BackupSpec> = None;
286
287 if crate::backup::extractor::is_backup_required(params.classified.ast_type) {
288 let spec = extract_backup_spec(params.sql, params.classified.ast_type, Dialect::SQLite)
289 .map_err(|e| SqliteError::BackupFailed(e.to_string()))?;
290 match &spec {
291 BackupSpec::InsertHint { .. } => pending_insert_spec = Some(spec),
292 BackupSpec::None { .. } => {}
293 _ => {
294 match capture_backup_sqlite(
295 db,
296 &spec,
297 ¶ms.connection.name,
298 params
299 .database
300 .or(Some(params.connection.database.as_str())),
301 params.policy,
302 params.audit.as_ref(),
303 ) {
304 Ok(Some(captured)) => {
305 backup_id = Some(captured.backup_id);
306 backup_row_count = captured.total_rows;
307 }
308 Ok(None) => {}
309 Err(e) => return Err(SqliteError::BackupFailed(e.to_string())),
310 }
311 }
312 }
313 }
314
315 let stats = run_streaming(db, params.sql, params.policy.row_cap, 4 * 1024 * 1024)?;
316
317 if let Some(spec) = pending_insert_spec
318 && let Some(id) = capture_insert_hint(
319 &spec,
320 ¶ms.connection.name,
321 params
322 .database
323 .or(Some(params.connection.database.as_str())),
324 stats.insert_id,
325 stats.affected_rows,
326 params.audit.as_ref(),
327 )
328 {
329 backup_id = Some(id);
330 backup_row_count = stats.affected_rows;
331 }
332
333 if in_txn {
334 db.execute_batch("COMMIT")?;
335 in_txn = false;
336 }
337
338 Ok(ExecuteResult {
339 journal_id: None,
340 ddl_no_op: false,
341 ddl_absent_targets: Vec::new(),
342 ddl_executed_targets: Vec::new(),
343 warnings: Vec::new(),
344 rows: stats.rows,
345 fields: stats.fields,
346 affected_rows: stats.affected_rows,
347 truncated: stats.truncated,
348 duration_ms: start.elapsed().as_millis() as u64,
349 backup_id,
350 backup_row_count,
351 })
352 })();
353
354 if result.is_err() && in_txn {
355 let _ = db.execute_batch("ROLLBACK");
356 }
357 result
358}
359
360pub fn classify_for_sqlite(sql: &str) -> Result<ClassifiedStatement, ClassifyError> {
362 classify_statement(sql, Dialect::SQLite)
363}
364
365#[cfg(test)]
366mod tests {
367 use super::*;
368 use crate::policy::model::{PolicyPresetName, policy_from_preset};
369
370 fn sqlite_conn(path: &std::path::Path) -> SqliteConnection {
371 SqliteConnection {
372 name: "local-sqlite".into(),
373 path: path.display().to_string(),
374 ..SqliteConnection::default()
375 }
376 }
377
378 fn classify(sql: &str) -> ClassifiedStatement {
379 classify_for_sqlite(sql).unwrap()
380 }
381
382 #[test]
383 fn ddl_write_read_roundtrip_with_backups() {
384 let dir = tempfile::tempdir().unwrap();
385 let file = dir.path().join("app.sqlite");
386 let conn = sqlite_conn(&file);
387 let policy = policy_from_preset(PolicyPresetName::Development);
388 let audit = Arc::new(AuditDb::at_path(&dir.path().join("audit.sqlite")).unwrap());
389
390 let ddl = classify("CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT)");
392 let r = execute_sqlite_statement(SqliteExecuteParams {
393 connection: &conn,
394 sql: "CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT)",
395 classified: &ddl,
396 policy: &policy,
397 database: None,
398 audit: Some(audit.clone()),
399 })
400 .unwrap();
401 assert_eq!(r.affected_rows, 0);
402
403 let ins = classify("INSERT INTO users (id, name) VALUES (1, 'a'), (2, 'b')");
405 let r = execute_sqlite_statement(SqliteExecuteParams {
406 connection: &conn,
407 sql: "INSERT INTO users (id, name) VALUES (1, 'a'), (2, 'b')",
408 classified: &ins,
409 policy: &policy,
410 database: None,
411 audit: Some(audit.clone()),
412 })
413 .unwrap();
414 assert_eq!(r.affected_rows, 2);
415 assert!(r.backup_id.is_some(), "insert-hint backup expected");
416
417 let upd = classify("UPDATE users SET name = 'x' WHERE id = 1");
419 let r = execute_sqlite_statement(SqliteExecuteParams {
420 connection: &conn,
421 sql: "UPDATE users SET name = 'x' WHERE id = 1",
422 classified: &upd,
423 policy: &policy,
424 database: None,
425 audit: Some(audit.clone()),
426 })
427 .unwrap();
428 assert_eq!(r.affected_rows, 1);
429 assert!(r.backup_id.is_some());
430
431 let read = classify("SELECT id, name FROM users");
433 let r = execute_sqlite_statement(SqliteExecuteParams {
434 connection: &conn,
435 sql: "SELECT id, name FROM users",
436 classified: &read,
437 policy: &policy,
438 database: None,
439 audit: Some(audit.clone()),
440 })
441 .unwrap();
442 assert_eq!(r.rows.len(), 2);
443 assert_eq!(r.fields, vec!["id", "name"]);
444 assert_eq!(r.rows[0]["name"], serde_json::json!("x"));
445 }
446
447 #[test]
448 fn row_cap_truncates_without_materializing() {
449 let dir = tempfile::tempdir().unwrap();
450 let file = dir.path().join("app.sqlite");
451 let conn = sqlite_conn(&file);
452 let mut policy = policy_from_preset(PolicyPresetName::Development);
453 policy.row_cap = 5;
454 let audit = Arc::new(AuditDb::at_path(&dir.path().join("audit.sqlite")).unwrap());
455 let ddl = classify("CREATE TABLE t (n INTEGER)");
456 execute_sqlite_statement(SqliteExecuteParams {
457 connection: &conn,
458 sql: "CREATE TABLE t (n INTEGER)",
459 classified: &ddl,
460 policy: &policy,
461 database: None,
462 audit: Some(audit.clone()),
463 })
464 .unwrap();
465 let ins = classify("INSERT INTO t (n) VALUES (1),(2),(3),(4),(5),(6),(7),(8)");
466 execute_sqlite_statement(SqliteExecuteParams {
467 connection: &conn,
468 sql: "INSERT INTO t (n) VALUES (1),(2),(3),(4),(5),(6),(7),(8)",
469 classified: &ins,
470 policy: &policy,
471 database: None,
472 audit: Some(audit.clone()),
473 })
474 .unwrap();
475 let read = classify("SELECT n FROM t");
476 let r = execute_sqlite_statement(SqliteExecuteParams {
477 connection: &conn,
478 sql: "SELECT n FROM t",
479 classified: &read,
480 policy: &policy,
481 database: None,
482 audit: Some(audit.clone()),
483 })
484 .unwrap();
485 assert_eq!(r.rows.len(), 5);
486 assert!(r.truncated);
487 }
488
489 #[test]
490 fn readonly_missing_file_is_an_error() {
491 let dir = tempfile::tempdir().unwrap();
492 let conn = sqlite_conn(&dir.path().join("missing.sqlite"));
493 let err = open_sqlite_database(&conn, true, 5000).unwrap_err();
494 assert!(matches!(err, SqliteError::NotFound(_)));
495 }
496
497 #[test]
498 fn statement_timeout_interrupts() {
499 let dir = tempfile::tempdir().unwrap();
500 let file = dir.path().join("app.sqlite");
501 let conn = sqlite_conn(&file);
502 let mut policy = policy_from_preset(PolicyPresetName::Development);
503 policy.stmt_timeout_ms = 150;
504 let audit = Arc::new(AuditDb::at_path(&dir.path().join("audit.sqlite")).unwrap());
505 let ddl = classify("CREATE TABLE t (n INTEGER)");
506 execute_sqlite_statement(SqliteExecuteParams {
507 connection: &conn,
508 sql: "CREATE TABLE t (n INTEGER)",
509 classified: &ddl,
510 policy: &policy,
511 database: None,
512 audit: Some(audit.clone()),
513 })
514 .unwrap();
515 let read = classify(
516 "WITH RECURSIVE c(n) AS (SELECT 1 UNION ALL SELECT n+1 FROM c) SELECT count(*) FROM c",
517 );
518 let err = execute_sqlite_statement(SqliteExecuteParams {
519 connection: &conn,
520 sql: "WITH RECURSIVE c(n) AS (SELECT 1 UNION ALL SELECT n+1 FROM c) SELECT count(*) FROM c",
521 classified: &read,
522 policy: &policy,
523 database: None,
524 audit: Some(audit.clone()),
525 })
526 .unwrap_err();
527 assert!(matches!(err, SqliteError::Timeout(_)), "{err:?}");
528 }
529}