Skip to main content

nexql_tools/
write.rs

1//! Write/admin MCP tool executors (Phase 9).
2
3use std::sync::Arc;
4
5use deadpool_postgres::Object;
6use nexql_policy::{AccessMode, SqlDecision, validate_write_sql};
7use serde_json::{Map, Value, json};
8use tokio_postgres::SimpleQueryMessage;
9use tokio_postgres::types::ToSql;
10
11use crate::error::ToolError;
12use crate::exec::ToolOutcome;
13use crate::session::ToolSession;
14use crate::sql::{is_safe_ident, parse_ref, quote_ident, quote_ref};
15
16const IMPORT_BATCH_SIZE: usize = 100;
17
18/// Run validated SQL inside an explicit transaction; roll back on error or `dry_run`.
19pub async fn execute_sql(
20    session: &Arc<ToolSession>,
21    sql: &str,
22    dry_run: bool,
23) -> Result<ToolOutcome, ToolError> {
24    let mode = session.access_mode;
25    match validate_write_sql(mode, sql)? {
26        SqlDecision::Allow => {}
27        SqlDecision::Reject => {
28            return Err(ToolError::Execution(format!(
29                "Security Error: SQL is not permitted in {:?} mode.",
30                mode
31            )));
32        }
33    }
34
35    let client = session.checkout().await?;
36    client.batch_execute("BEGIN").await?;
37    let outcome = async {
38        let (rows, command_tag) = run_simple_query(&client, sql).await?;
39        Ok::<_, ToolError>((rows, command_tag))
40    }
41    .await;
42
43    let rolled_back = dry_run || outcome.is_err();
44    if rolled_back {
45        let _ = client.batch_execute("ROLLBACK").await;
46    } else {
47        let _ = client.batch_execute("COMMIT").await;
48    }
49
50    match outcome {
51        Ok((rows, rows_affected)) => Ok(ToolOutcome::ok_json(json!({
52            "dry_run": dry_run,
53            "rolled_back": rolled_back,
54            "rows_affected": rows_affected,
55            "rows": rows,
56        }))),
57        Err(e) => Err(e),
58    }
59}
60
61/// Structured insert/update/delete by primary key (parameterized).
62pub async fn edit_row(session: &Arc<ToolSession>, args: &Value) -> Result<ToolOutcome, ToolError> {
63    let table_ref = args
64        .get("table")
65        .and_then(|v| v.as_str())
66        .ok_or_else(|| ToolError::InvalidArgs("table is required (schema.name)".into()))?;
67    let action = args.get("action").and_then(|v| v.as_str()).ok_or_else(|| {
68        ToolError::InvalidArgs("action is required (insert|update|delete)".into())
69    })?;
70
71    let (schema, table) = parse_ref(table_ref).map_err(ToolError::InvalidArgs)?;
72    if !session.filter.allows_table(&schema, &table) {
73        return Err(ToolError::Execution(format!(
74            "Table \"{schema}.{table}\" is denied by policy filter."
75        )));
76    }
77
78    let client = session.checkout().await?;
79    client.batch_execute("BEGIN").await?;
80
81    let result = async {
82        match action.to_ascii_lowercase().as_str() {
83            "insert" => edit_row_insert(&client, &schema, &table, args).await,
84            "update" => edit_row_update(&client, &schema, &table, args).await,
85            "delete" => edit_row_delete(&client, &schema, &table, args).await,
86            other => Err(ToolError::InvalidArgs(format!(
87                "Unsupported action \"{other}\". Use insert, update, or delete."
88            ))),
89        }
90    }
91    .await;
92
93    match &result {
94        Ok(_) => {
95            let _ = client.batch_execute("COMMIT").await;
96        }
97        Err(_) => {
98            let _ = client.batch_execute("ROLLBACK").await;
99        }
100    }
101    result
102}
103
104async fn edit_row_insert(
105    client: &Object,
106    schema: &str,
107    table: &str,
108    args: &Value,
109) -> Result<ToolOutcome, ToolError> {
110    let values = args
111        .get("values")
112        .and_then(|v| v.as_object())
113        .ok_or_else(|| ToolError::InvalidArgs("values object is required for insert".into()))?;
114    if values.is_empty() {
115        return Err(ToolError::InvalidArgs(
116            "values must contain at least one column".into(),
117        ));
118    }
119    let mut columns = Vec::new();
120    let mut params: Vec<Box<dyn ToSql + Sync + Send>> = Vec::new();
121    for (col, val) in values {
122        validate_column_name(col)?;
123        columns.push(quote_ident(col));
124        params.push(json_to_sql_param(val));
125    }
126    let placeholders: Vec<String> = (1..=params.len()).map(|i| format!("${i}")).collect();
127    let sql = format!(
128        "INSERT INTO {} ({}) VALUES ({}) RETURNING *",
129        quote_ref(schema, table),
130        columns.join(", "),
131        placeholders.join(", ")
132    );
133    let param_refs: Vec<&(dyn ToSql + Sync)> = params
134        .iter()
135        .map(|p| p.as_ref() as &(dyn ToSql + Sync))
136        .collect();
137    let rows = client.query(&sql, &param_refs[..]).await?;
138    Ok(ToolOutcome::ok_json(json!({
139        "action": "insert",
140        "table": format!("{schema}.{table}"),
141        "rows": simple_rows_to_json(&rows),
142    })))
143}
144
145async fn edit_row_update(
146    client: &Object,
147    schema: &str,
148    table: &str,
149    args: &Value,
150) -> Result<ToolOutcome, ToolError> {
151    let pk = args
152        .get("pk")
153        .and_then(|v| v.as_object())
154        .ok_or_else(|| ToolError::InvalidArgs("pk object is required for update".into()))?;
155    if pk.is_empty() {
156        return Err(ToolError::InvalidArgs(
157            "pk must contain at least one primary-key column".into(),
158        ));
159    }
160    let values = args
161        .get("values")
162        .and_then(|v| v.as_object())
163        .ok_or_else(|| ToolError::InvalidArgs("values object is required for update".into()))?;
164    if values.is_empty() {
165        return Err(ToolError::InvalidArgs(
166            "values must contain at least one column to update".into(),
167        ));
168    }
169
170    let mut set_cols = Vec::new();
171    let mut params: Vec<Box<dyn ToSql + Sync + Send>> = Vec::new();
172    for (col, val) in values {
173        validate_column_name(col)?;
174        let idx = params.len() + 1;
175        set_cols.push(format!("{} = ${idx}", quote_ident(col)));
176        params.push(json_to_sql_param(val));
177    }
178    let mut where_cols = Vec::new();
179    for (col, val) in pk {
180        validate_column_name(col)?;
181        let idx = params.len() + 1;
182        where_cols.push(format!("{} = ${idx}", quote_ident(col)));
183        params.push(json_to_sql_param(val));
184    }
185    let sql = format!(
186        "UPDATE {} SET {} WHERE {} RETURNING *",
187        quote_ref(schema, table),
188        set_cols.join(", "),
189        where_cols.join(" AND ")
190    );
191    let param_refs: Vec<&(dyn ToSql + Sync)> = params
192        .iter()
193        .map(|p| p.as_ref() as &(dyn ToSql + Sync))
194        .collect();
195    let rows = client.query(&sql, &param_refs[..]).await?;
196    Ok(ToolOutcome::ok_json(json!({
197        "action": "update",
198        "table": format!("{schema}.{table}"),
199        "rows": simple_rows_to_json(&rows),
200    })))
201}
202
203async fn edit_row_delete(
204    client: &Object,
205    schema: &str,
206    table: &str,
207    args: &Value,
208) -> Result<ToolOutcome, ToolError> {
209    let pk = args
210        .get("pk")
211        .and_then(|v| v.as_object())
212        .ok_or_else(|| ToolError::InvalidArgs("pk object is required for delete".into()))?;
213    if pk.is_empty() {
214        return Err(ToolError::InvalidArgs(
215            "pk must contain at least one primary-key column".into(),
216        ));
217    }
218    let mut where_cols = Vec::new();
219    let mut params: Vec<Box<dyn ToSql + Sync + Send>> = Vec::new();
220    for (col, val) in pk {
221        validate_column_name(col)?;
222        let idx = params.len() + 1;
223        where_cols.push(format!("{} = ${idx}", quote_ident(col)));
224        params.push(json_to_sql_param(val));
225    }
226    let sql = format!(
227        "DELETE FROM {} WHERE {} RETURNING *",
228        quote_ref(schema, table),
229        where_cols.join(" AND ")
230    );
231    let param_refs: Vec<&(dyn ToSql + Sync)> = params
232        .iter()
233        .map(|p| p.as_ref() as &(dyn ToSql + Sync))
234        .collect();
235    let rows = client.query(&sql, &param_refs[..]).await?;
236    Ok(ToolOutcome::ok_json(json!({
237        "action": "delete",
238        "table": format!("{schema}.{table}"),
239        "rows": simple_rows_to_json(&rows),
240    })))
241}
242
243/// Batched INSERT from a JSON rows array.
244pub async fn import_data(
245    session: &Arc<ToolSession>,
246    args: &Value,
247) -> Result<ToolOutcome, ToolError> {
248    let table_ref = args
249        .get("table")
250        .and_then(|v| v.as_str())
251        .ok_or_else(|| ToolError::InvalidArgs("table is required (schema.name)".into()))?;
252    let rows_val = args
253        .get("rows")
254        .and_then(|v| v.as_array())
255        .ok_or_else(|| ToolError::InvalidArgs("rows array is required".into()))?;
256    if rows_val.is_empty() {
257        return Ok(ToolOutcome::ok_json(json!({
258            "table": table_ref,
259            "rows_imported": 0,
260            "batches": 0,
261        })));
262    }
263
264    let (schema, table) = parse_ref(table_ref).map_err(ToolError::InvalidArgs)?;
265    if !session.filter.allows_table(&schema, &table) {
266        return Err(ToolError::Execution(format!(
267            "Table \"{schema}.{table}\" is denied by policy filter."
268        )));
269    }
270
271    let columns: Vec<String> = if let Some(cols) = args.get("columns").and_then(|v| v.as_array()) {
272        cols.iter()
273            .map(|c| {
274                let s = c
275                    .as_str()
276                    .ok_or_else(|| ToolError::InvalidArgs("columns must be strings".into()))?;
277                validate_column_name(s)?;
278                Ok(s.to_string())
279            })
280            .collect::<Result<Vec<_>, ToolError>>()?
281    } else {
282        let first = rows_val[0]
283            .as_object()
284            .ok_or_else(|| ToolError::InvalidArgs("each row must be a JSON object".into()))?;
285        let mut cols: Vec<String> = first.keys().cloned().collect();
286        cols.sort();
287        for col in &cols {
288            validate_column_name(col)?;
289        }
290        cols
291    };
292
293    let client = session.checkout().await?;
294    client.batch_execute("BEGIN").await?;
295
296    let mut total_imported = 0u64;
297    let mut batches = 0u32;
298    let result = async {
299        for chunk in rows_val.chunks(IMPORT_BATCH_SIZE) {
300            let (sql, params) = build_batch_insert(&schema, &table, &columns, chunk)?;
301            let param_refs: Vec<&(dyn ToSql + Sync)> = params
302                .iter()
303                .map(|p| p.as_ref() as &(dyn ToSql + Sync))
304                .collect();
305            let affected = client.execute(&sql, &param_refs[..]).await?;
306            total_imported += affected;
307            batches += 1;
308        }
309        Ok::<_, ToolError>(())
310    }
311    .await;
312
313    match &result {
314        Ok(_) => {
315            let _ = client.batch_execute("COMMIT").await;
316        }
317        Err(_) => {
318            let _ = client.batch_execute("ROLLBACK").await;
319        }
320    }
321    result?;
322
323    Ok(ToolOutcome::ok_json(json!({
324        "table": format!("{schema}.{table}"),
325        "rows_imported": total_imported,
326        "batches": batches,
327        "columns": columns,
328    })))
329}
330
331fn build_batch_insert(
332    schema: &str,
333    table: &str,
334    columns: &[String],
335    rows: &[Value],
336) -> Result<(String, Vec<Box<dyn ToSql + Sync + Send>>), ToolError> {
337    let quoted_cols = columns
338        .iter()
339        .map(|c| quote_ident(c))
340        .collect::<Vec<_>>()
341        .join(", ");
342    let mut params: Vec<Box<dyn ToSql + Sync + Send>> = Vec::new();
343    let mut value_groups = Vec::new();
344    for row in rows {
345        let obj = row
346            .as_object()
347            .ok_or_else(|| ToolError::InvalidArgs("each row must be a JSON object".into()))?;
348        let mut placeholders = Vec::new();
349        for col in columns {
350            let val = obj.get(col).unwrap_or(&Value::Null);
351            let idx = params.len() + 1;
352            placeholders.push(format!("${idx}"));
353            params.push(json_to_sql_param(val));
354        }
355        value_groups.push(format!("({})", placeholders.join(", ")));
356    }
357    let sql = format!(
358        "INSERT INTO {} ({}) VALUES {}",
359        quote_ref(schema, table),
360        quoted_cols,
361        value_groups.join(", ")
362    );
363    Ok((sql, params))
364}
365
366/// Run DDL validated for Admin mode inside a transaction.
367pub async fn apply_ddl(
368    session: &Arc<ToolSession>,
369    sql: &str,
370    dry_run: bool,
371) -> Result<ToolOutcome, ToolError> {
372    assert_ddl_statement(sql)?;
373    match validate_write_sql(AccessMode::Admin, sql)? {
374        SqlDecision::Allow => {}
375        SqlDecision::Reject => {
376            return Err(ToolError::Execution(
377                "Security Error: DDL statement is not permitted.".into(),
378            ));
379        }
380    }
381
382    let client = session.checkout().await?;
383    client.batch_execute("BEGIN").await?;
384    let outcome = run_simple_query(&client, sql).await;
385    let rolled_back = dry_run || outcome.is_err();
386    if rolled_back {
387        let _ = client.batch_execute("ROLLBACK").await;
388    } else {
389        let _ = client.batch_execute("COMMIT").await;
390    }
391    let (rows, rows_affected) = outcome?;
392    Ok(ToolOutcome::ok_json(json!({
393        "dry_run": dry_run,
394        "rolled_back": rolled_back,
395        "rows_affected": rows_affected,
396        "rows": rows,
397    })))
398}
399
400/// `CREATE INDEX CONCURRENTLY` — must run outside a transaction.
401pub async fn create_index_concurrently(
402    session: &Arc<ToolSession>,
403    sql: &str,
404) -> Result<ToolOutcome, ToolError> {
405    let upper = sql.trim().to_ascii_uppercase();
406    if !upper.contains("CREATE INDEX") || !upper.contains("CONCURRENTLY") {
407        return Err(ToolError::InvalidArgs(
408            "sql must be a CREATE INDEX CONCURRENTLY statement".into(),
409        ));
410    }
411    match validate_write_sql(AccessMode::Admin, sql)? {
412        SqlDecision::Allow => {}
413        SqlDecision::Reject => {
414            return Err(ToolError::Execution(
415                "Security Error: index statement is not permitted.".into(),
416            ));
417        }
418    }
419
420    let client = session.checkout().await?;
421    let (rows, rows_affected) = run_simple_query(&client, sql).await?;
422    Ok(ToolOutcome::ok_json(json!({
423        "rows_affected": rows_affected,
424        "rows": rows,
425        "note": "CREATE INDEX CONCURRENTLY runs outside a transaction.",
426    })))
427}
428
429/// VACUUM / ANALYZE / REINDEX — cannot run inside a transaction.
430pub async fn run_maintenance(
431    session: &Arc<ToolSession>,
432    args: &Value,
433) -> Result<ToolOutcome, ToolError> {
434    let action = args.get("action").and_then(|v| v.as_str()).ok_or_else(|| {
435        ToolError::InvalidArgs("action is required (vacuum|analyze|reindex)".into())
436    })?;
437    let full = args.get("full").and_then(|v| v.as_bool()).unwrap_or(false);
438    let table_ref = args.get("table").and_then(|v| v.as_str());
439
440    let sql = match action.to_ascii_lowercase().as_str() {
441        "vacuum" => build_vacuum_sql(table_ref, full)?,
442        "analyze" => build_analyze_sql(table_ref)?,
443        "reindex" => build_reindex_sql(table_ref)?,
444        other => {
445            return Err(ToolError::InvalidArgs(format!(
446                "Unsupported action \"{other}\". Use vacuum, analyze, or reindex."
447            )));
448        }
449    };
450
451    match validate_write_sql(AccessMode::Admin, &sql)? {
452        SqlDecision::Allow => {}
453        SqlDecision::Reject => {
454            return Err(ToolError::Execution(
455                "Security Error: maintenance statement is not permitted.".into(),
456            ));
457        }
458    }
459
460    let client = session.checkout().await?;
461    let (rows, rows_affected) = run_simple_query(&client, &sql).await?;
462    Ok(ToolOutcome::ok_json(json!({
463        "action": action,
464        "sql": sql,
465        "rows_affected": rows_affected,
466        "rows": rows,
467    })))
468}
469
470/// Cancel or terminate a backend by pid.
471pub async fn terminate_query(
472    session: &Arc<ToolSession>,
473    args: &Value,
474) -> Result<ToolOutcome, ToolError> {
475    let pid = args
476        .get("pid")
477        .and_then(|v| v.as_i64())
478        .ok_or_else(|| ToolError::InvalidArgs("pid is required".into()))?;
479    if pid <= 0 {
480        return Err(ToolError::InvalidArgs(
481            "pid must be a positive integer".into(),
482        ));
483    }
484    let force = args.get("force").and_then(|v| v.as_bool()).unwrap_or(false);
485
486    let client = session.checkout().await?;
487    let own_pid: i32 = client
488        .query_one("SELECT pg_backend_pid()", &[])
489        .await?
490        .get(0);
491    if pid == i64::from(own_pid) {
492        return Err(ToolError::Execution(
493            "refusing to cancel/terminate the current session backend".into(),
494        ));
495    }
496
497    let target = client
498        .query_opt(
499            "SELECT pid, usename, state, query, usesuper FROM pg_stat_activity WHERE pid = $1",
500            &[&(pid as i32)],
501        )
502        .await?;
503    let Some(row) = target else {
504        return Err(ToolError::Execution(format!(
505            "No backend found with pid {pid}"
506        )));
507    };
508    let usesuper: bool = row.get("usesuper");
509    if usesuper {
510        return Err(ToolError::Execution(
511            "refusing to cancel/terminate a superuser backend — use a direct superuser session if required"
512                .into(),
513        ));
514    }
515
516    let fn_name = if force {
517        "pg_terminate_backend"
518    } else {
519        "pg_cancel_backend"
520    };
521    let sql = format!("SELECT {fn_name}($1)");
522    let result: bool = client.query_one(&sql, &[&(pid as i32)]).await?.get(0);
523
524    Ok(ToolOutcome::ok_json(json!({
525        "pid": pid,
526        "force": force,
527        "success": result,
528        "target": {
529            "usename": row.get::<_, Option<String>>("usename"),
530            "state": row.get::<_, Option<String>>("state"),
531            "query": row.get::<_, Option<String>>("query"),
532        },
533    })))
534}
535
536fn build_vacuum_sql(table_ref: Option<&str>, full: bool) -> Result<String, ToolError> {
537    let mut sql = String::from("VACUUM");
538    if full {
539        sql.push_str(" FULL");
540    }
541    if let Some(ref_) = table_ref {
542        let (schema, table) = parse_ref(ref_).map_err(ToolError::InvalidArgs)?;
543        sql.push(' ');
544        sql.push_str(&quote_ref(&schema, &table));
545    }
546    Ok(sql)
547}
548
549fn build_analyze_sql(table_ref: Option<&str>) -> Result<String, ToolError> {
550    let mut sql = String::from("ANALYZE");
551    if let Some(ref_) = table_ref {
552        let (schema, table) = parse_ref(ref_).map_err(ToolError::InvalidArgs)?;
553        sql.push(' ');
554        sql.push_str(&quote_ref(&schema, &table));
555    }
556    Ok(sql)
557}
558
559fn build_reindex_sql(table_ref: Option<&str>) -> Result<String, ToolError> {
560    let ref_ = table_ref.ok_or_else(|| {
561        ToolError::InvalidArgs("table (schema.name) is required for reindex".into())
562    })?;
563    let (schema, table) = parse_ref(ref_).map_err(ToolError::InvalidArgs)?;
564    Ok(format!("REINDEX TABLE {}", quote_ref(&schema, &table)))
565}
566
567fn assert_ddl_statement(sql: &str) -> Result<(), ToolError> {
568    let upper = sql.trim().to_ascii_uppercase();
569    let ddl_prefixes = [
570        "CREATE ",
571        "ALTER ",
572        "DROP ",
573        "TRUNCATE ",
574        "COMMENT ON ",
575        "GRANT ",
576        "REVOKE ",
577        "RENAME ",
578    ];
579    if !ddl_prefixes.iter().any(|p| upper.starts_with(p)) {
580        return Err(ToolError::Execution(
581            "apply_ddl only accepts DDL statements (CREATE, ALTER, DROP, TRUNCATE, …). Use execute_sql for DML."
582                .into(),
583        ));
584    }
585    Ok(())
586}
587
588fn validate_column_name(col: &str) -> Result<(), ToolError> {
589    if !is_safe_ident(col) {
590        return Err(ToolError::InvalidArgs(format!(
591            "Invalid column name \"{col}\"."
592        )));
593    }
594    Ok(())
595}
596
597fn json_to_sql_param(val: &Value) -> Box<dyn ToSql + Sync + Send> {
598    match val {
599        Value::Null => Box::new(None::<String>),
600        Value::Bool(b) => Box::new(*b),
601        Value::Number(n) => {
602            if let Some(i) = n.as_i64() {
603                Box::new(i)
604            } else if let Some(u) = n.as_u64() {
605                Box::new(i64::try_from(u).unwrap_or(i64::MAX))
606            } else if let Some(f) = n.as_f64() {
607                Box::new(f)
608            } else {
609                Box::new(n.to_string())
610            }
611        }
612        Value::String(s) => Box::new(s.clone()),
613        Value::Array(_) | Value::Object(_) => Box::new(val.clone()),
614    }
615}
616
617async fn run_simple_query(
618    client: &Object,
619    sql: &str,
620) -> Result<(Vec<Value>, Option<u64>), ToolError> {
621    let messages = client.simple_query(sql).await?;
622    Ok(collect_simple_query(messages))
623}
624
625fn collect_simple_query(messages: Vec<SimpleQueryMessage>) -> (Vec<Value>, Option<u64>) {
626    let mut rows = Vec::new();
627    let mut rows_affected = None;
628    for msg in messages {
629        match msg {
630            SimpleQueryMessage::Row(row) => {
631                let mut map = Map::new();
632                for col in row.columns() {
633                    let cell = row
634                        .try_get(col.name())
635                        .ok()
636                        .flatten()
637                        .map(|s| Value::String(s.to_string()))
638                        .unwrap_or(Value::Null);
639                    map.insert(col.name().to_string(), cell);
640                }
641                rows.push(Value::Object(map));
642            }
643            SimpleQueryMessage::CommandComplete(n) => rows_affected = Some(n),
644            SimpleQueryMessage::RowDescription(_) => {}
645            _ => {}
646        }
647    }
648    (rows, rows_affected)
649}
650
651fn simple_rows_to_json(rows: &[tokio_postgres::Row]) -> Vec<Value> {
652    rows.iter()
653        .map(|row| {
654            let mut map = Map::new();
655            for (i, col) in row.columns().iter().enumerate() {
656                let val: Option<String> = row.try_get(i).ok().flatten();
657                map.insert(
658                    col.name().to_string(),
659                    val.map(Value::String).unwrap_or(Value::Null),
660                );
661            }
662            Value::Object(map)
663        })
664        .collect()
665}
666
667#[cfg(test)]
668mod tests {
669    use super::*;
670
671    #[test]
672    fn assert_ddl_rejects_select() {
673        assert!(assert_ddl_statement("SELECT 1").is_err());
674    }
675
676    #[test]
677    fn assert_ddl_accepts_create() {
678        assert!(assert_ddl_statement("CREATE TABLE t (id int)").is_ok());
679    }
680
681    #[test]
682    fn build_vacuum_table() {
683        let sql = build_vacuum_sql(Some("public.users"), false).unwrap();
684        assert_eq!(sql, "VACUUM \"public\".\"users\"");
685    }
686
687    #[test]
688    fn build_batch_insert_sql() {
689        let rows = [json!({"id": 1, "name": "a"}), json!({"id": 2, "name": "b"})];
690        let (sql, params) =
691            build_batch_insert("public", "users", &["id".into(), "name".into()], &rows).unwrap();
692        assert!(sql.starts_with("INSERT INTO \"public\".\"users\""));
693        assert_eq!(params.len(), 4);
694    }
695}