Skip to main content

sequel_mcp/sql/
sqlite.rs

1//! SQLite execution: read-only handles for reads, BEGIN IMMEDIATE for
2//! writes, interrupt-handle cancellation, streaming with caps, symlink
3//! checks, and the shared backup pipeline.
4
5use 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    /// DDL no-op flag (MySQL preflight path); SQLite defaults false.
42    pub ddl_no_op: bool,
43    /// Absent targets of a MySQL Mixed DROP (D4A); SQLite defaults empty.
44    pub ddl_absent_targets: Vec<String>,
45    /// Targets named by the rewritten MySQL Mixed DROP; SQLite defaults
46    /// empty.
47    pub ddl_executed_targets: Vec<String>,
48    /// Protection-model warnings (MySQL DDL path); SQLite defaults empty.
49    pub warnings: Vec<&'static str>,
50    /// Operation-journal row id (D3); MySQL mutations create journals,
51    /// SQLite ones currently do not (single-file local transactions).
52    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
64/// Open a SQLite database with the legacy semantics: reads use a read-only
65/// handle that must already exist; writes create if needed.
66pub 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    // Fail-closed test-mode path gate (before any file is touched).
73    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
97/// Reads must not follow a substituted symlink outside the resolved tree.
98fn verify_no_symlink_escape(canonical: &Path) -> Result<(), SqliteError> {
99    // canonicalize already resolved symlinks; refuse when the final
100    // component itself is a link (raced) by checking the parent chain.
101    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
133/// Stream rows with an early stop at the row cap (never materialize then
134/// slice). The byte cap stops iteration once accumulated JSON size
135/// exceeds it.
136fn 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    // A statement exposing columns returns rows (SELECT/PRAGMA/RETURNING);
146    // otherwise it is a mutation: execute it and read change stats.
147    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
224/// Execute one classified statement against a SQLite database, using the
225/// shared backup pipeline. Backup capture failure denies the mutation.
226/// A wall-clock timer armed on this connection interrupts runaway
227/// statements at `policy.stmt_timeout_ms`.
228pub fn execute_sqlite_statement(
229    params: SqliteExecuteParams<'_>,
230) -> Result<ExecuteResult, SqliteError> {
231    let start = Instant::now();
232    let is_read = READ_CATEGORIES.contains(&params.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        // TRUE wall-clock deadline: a fixed iteration count of 1 ms
240        // sleeps stretches arbitrarily under scheduler contention
241        // (found by the CI runner: 16 process-heavy tests in parallel
242        // turned "5 s" into 25 s+ and hung EOF shutdown past its
243        // bound). Sleep in slices only so `done` is noticed promptly.
244        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, &params, 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(&params.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                        &params.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                &params.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
360/// Convenience classifier wrapper matching the legacy entry point.
361pub 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        // DDL
391        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        // INSERT with backup hint
404        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        // UPDATE with row backup
418        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        // Read via read-only handle
432        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}