Skip to main content

nexql_tools/
write.rs

1// SPDX-License-Identifier: GPL-3.0-only
2// Copyright (C) 2026 NexQL-OSS Team
3
4//! Write/admin MCP tool executors (Phase 9).
5
6use std::collections::HashMap;
7use std::sync::Arc;
8
9use chrono::{DateTime, NaiveDate, NaiveDateTime};
10use deadpool_postgres::Object;
11use nexql_policy::{
12    AccessMode, ObjectRef, SqlDecision, enforce_read_table_policy, select_table_refs,
13    validate_readonly_sql, validate_write_sql,
14};
15use rust_decimal::Decimal;
16use serde_json::{Map, Value, json};
17use tokio_postgres::SimpleQueryMessage;
18use tokio_postgres::types::{Json, ToSql};
19use uuid::Uuid;
20
21use crate::cell_json::{redact_pii_in_rows, rows_to_json_vec};
22use crate::error::ToolError;
23use crate::exec::ToolOutcome;
24use crate::session::ToolSession;
25use crate::sql::{is_safe_ident, parse_ref, quote_ident, quote_ref};
26
27const IMPORT_BATCH_SIZE: usize = 100;
28const MUTATION_DIFF_ROW_CAP: i64 = 100;
29
30/// Run validated SQL inside an explicit transaction; roll back on error or `dry_run`.
31pub async fn execute_sql(
32    session: &Arc<ToolSession>,
33    sql: &str,
34    dry_run: bool,
35    include_diff: bool,
36) -> Result<ToolOutcome, ToolError> {
37    let mode = session.access_mode();
38    match validate_write_sql(mode, sql)? {
39        SqlDecision::Allow => {}
40        SqlDecision::Reject => {
41            return Err(ToolError::Execution(format!(
42                "Security Error: SQL is not permitted in {:?} mode.",
43                mode
44            )));
45        }
46    }
47    if matches!(validate_readonly_sql(sql)?, SqlDecision::Allow) {
48        enforce_read_table_policy(&session.filter(), sql)?;
49    }
50
51    let client = session.checkout().await?;
52    client.batch_execute("BEGIN").await?;
53    let outcome = async {
54        let (rows, command_tag) = run_simple_query(&client, sql).await?;
55        Ok::<_, ToolError>((rows, command_tag))
56    }
57    .await;
58
59    let rolled_back = dry_run || outcome.is_err();
60    if rolled_back {
61        let _ = client.batch_execute("ROLLBACK").await;
62    } else {
63        let _ = client.batch_execute("COMMIT").await;
64    }
65
66    match outcome {
67        Ok((rows, rows_affected)) => {
68            let rows = redact_row_results(session, Some(sql), None, None, rows);
69            let mut payload = json!({
70                "dry_run": dry_run,
71                "rolled_back": rolled_back,
72                "rows_affected": rows_affected,
73                "rows": rows,
74            });
75            if include_diff && !rows.is_empty() {
76                payload["after"] = json!(rows.clone());
77            }
78            Ok(ToolOutcome::ok_json(payload))
79        }
80        Err(e) => Err(append_constraint_hint(e)),
81    }
82}
83
84fn append_constraint_hint(err: ToolError) -> ToolError {
85    if let ToolError::Execution(ref msg) = err
86        && let Some(name) = extract_constraint_name(msg)
87    {
88        return ToolError::Execution(format!("{msg} (constraint: {name})"));
89    }
90    err
91}
92
93fn extract_constraint_name(message: &str) -> Option<String> {
94    message
95        .split("constraint \"")
96        .nth(1)
97        .and_then(|rest| rest.split('"').next())
98        .map(str::to_owned)
99}
100
101/// Structured insert/update/delete by primary key (parameterized).
102pub async fn edit_row(session: &Arc<ToolSession>, args: &Value) -> Result<ToolOutcome, ToolError> {
103    let table_ref = args
104        .get("table")
105        .and_then(|v| v.as_str())
106        .ok_or_else(|| ToolError::InvalidArgs("table is required (schema.name)".into()))?;
107    let action = args.get("action").and_then(|v| v.as_str()).ok_or_else(|| {
108        ToolError::InvalidArgs("action is required (insert|update|delete)".into())
109    })?;
110    let dry_run = args
111        .get("dry_run")
112        .and_then(|v| v.as_bool())
113        .unwrap_or(false);
114    let include_diff = args
115        .get("include_diff")
116        .and_then(|v| v.as_bool())
117        .unwrap_or(dry_run);
118
119    let (schema, table) = parse_ref(table_ref).map_err(ToolError::InvalidArgs)?;
120    if !session.filter().allows_table(&schema, &table) {
121        return Err(ToolError::Execution(format!(
122            "Table \"{schema}.{table}\" is denied by policy filter."
123        )));
124    }
125
126    let client = session.checkout().await?;
127    client.batch_execute("BEGIN").await?;
128
129    let result = async {
130        match action.to_ascii_lowercase().as_str() {
131            "insert" => {
132                edit_row_insert(session, &client, &schema, &table, args, include_diff).await
133            }
134            "update" => {
135                edit_row_update(session, &client, &schema, &table, args, include_diff).await
136            }
137            "delete" => {
138                edit_row_delete(session, &client, &schema, &table, args, include_diff).await
139            }
140            other => Err(ToolError::InvalidArgs(format!(
141                "Unsupported action \"{other}\". Use insert, update, or delete."
142            ))),
143        }
144    }
145    .await;
146
147    let rolled_back = dry_run || result.is_err();
148    if rolled_back {
149        let _ = client.batch_execute("ROLLBACK").await;
150    } else {
151        let _ = client.batch_execute("COMMIT").await;
152        let (connection_id, database) = session.active_context().await;
153        session.mark_index_stale(&connection_id, &database);
154    }
155
156    match result {
157        Ok(mut outcome) => {
158            if let Some(obj) = outcome.structured.as_mut().and_then(|v| v.as_object_mut()) {
159                obj.insert("dry_run".into(), json!(dry_run));
160                obj.insert("rolled_back".into(), json!(rolled_back));
161            }
162            Ok(outcome)
163        }
164        Err(e) => Err(append_constraint_hint(e)),
165    }
166}
167
168async fn edit_row_insert(
169    session: &Arc<ToolSession>,
170    client: &Object,
171    schema: &str,
172    table: &str,
173    args: &Value,
174    include_diff: bool,
175) -> Result<ToolOutcome, ToolError> {
176    let values = args
177        .get("values")
178        .and_then(|v| v.as_object())
179        .ok_or_else(|| ToolError::InvalidArgs("values object is required for insert".into()))?;
180    if values.is_empty() {
181        return Err(ToolError::InvalidArgs(
182            "values must contain at least one column".into(),
183        ));
184    }
185    let column_types = load_column_types(client, schema, table).await?;
186    let mut columns = Vec::new();
187    let mut params: Vec<Box<dyn ToSql + Sync + Send>> = Vec::new();
188    for (col, val) in values {
189        validate_column_name(col)?;
190        columns.push(quote_ident(col));
191        params.push(json_to_sql_param(
192            column_types.get(col).map(String::as_str),
193            val,
194        )?);
195    }
196    let placeholders: Vec<String> = (1..=params.len()).map(|i| format!("${i}")).collect();
197    let sql = format!(
198        "INSERT INTO {} ({}) VALUES ({}) RETURNING *",
199        quote_ref(schema, table),
200        columns.join(", "),
201        placeholders.join(", ")
202    );
203    let param_refs: Vec<&(dyn ToSql + Sync)> = params
204        .iter()
205        .map(|p| p.as_ref() as &(dyn ToSql + Sync))
206        .collect();
207    let rows = client.query(&sql, &param_refs[..]).await?;
208    let rows = redact_row_results(session, None, Some(schema), Some(table), simple_rows_to_json(&rows));
209    let mut payload = json!({
210        "action": "insert",
211        "table": format!("{schema}.{table}"),
212        "rows_affected": rows.len(),
213        "rows": rows,
214    });
215    if include_diff {
216        payload["after"] = json!(rows);
217    }
218    Ok(ToolOutcome::ok_json(payload))
219}
220
221async fn edit_row_update(
222    session: &Arc<ToolSession>,
223    client: &Object,
224    schema: &str,
225    table: &str,
226    args: &Value,
227    include_diff: bool,
228) -> Result<ToolOutcome, ToolError> {
229    let pk = args
230        .get("pk")
231        .and_then(|v| v.as_object())
232        .ok_or_else(|| ToolError::InvalidArgs("pk object is required for update".into()))?;
233    if pk.is_empty() {
234        return Err(ToolError::InvalidArgs(
235            "pk must contain at least one primary-key column".into(),
236        ));
237    }
238    let values = args
239        .get("values")
240        .and_then(|v| v.as_object())
241        .ok_or_else(|| ToolError::InvalidArgs("values object is required for update".into()))?;
242    if values.is_empty() {
243        return Err(ToolError::InvalidArgs(
244            "values must contain at least one column to update".into(),
245        ));
246    }
247
248    let column_types = load_column_types(client, schema, table).await?;
249    let (where_sql, pk_params) = pk_where_clause(pk, &column_types)?;
250    if include_diff {
251        ensure_pk_row_cap(client, schema, table, &where_sql, &pk_params).await?;
252    }
253    let before = if include_diff {
254        snapshot_rows(client, schema, table, &where_sql, &pk_params).await?
255    } else {
256        Vec::new()
257    };
258
259    let mut set_cols = Vec::new();
260    let mut params: Vec<Box<dyn ToSql + Sync + Send>> = Vec::new();
261    for (col, val) in values {
262        validate_column_name(col)?;
263        let idx = params.len() + 1;
264        set_cols.push(format!("{} = ${idx}", quote_ident(col)));
265        params.push(json_to_sql_param(
266            column_types.get(col).map(String::as_str),
267            val,
268        )?);
269    }
270    let pk_offset = params.len();
271    let mut where_cols = Vec::new();
272    for (col, val) in pk {
273        validate_column_name(col)?;
274        let idx = params.len() + 1;
275        where_cols.push(format!("{} = ${idx}", quote_ident(col)));
276        params.push(json_to_sql_param(
277            column_types.get(col).map(String::as_str),
278            val,
279        )?);
280    }
281    let _ = pk_offset;
282    let sql = format!(
283        "UPDATE {} SET {} WHERE {} RETURNING *",
284        quote_ref(schema, table),
285        set_cols.join(", "),
286        where_cols.join(" AND ")
287    );
288    let param_refs: Vec<&(dyn ToSql + Sync)> = params
289        .iter()
290        .map(|p| p.as_ref() as &(dyn ToSql + Sync))
291        .collect();
292    let rows = client.query(&sql, &param_refs[..]).await?;
293    let after =
294        redact_row_results(session, None, Some(schema), Some(table), simple_rows_to_json(&rows));
295    let mut payload = json!({
296        "action": "update",
297        "table": format!("{schema}.{table}"),
298        "rows_affected": after.len(),
299        "rows": after,
300    });
301    if include_diff {
302        payload["before"] = json!(before);
303        payload["after"] = json!(after);
304        payload["diff"] = json!(compute_row_diff(&before, &after));
305    }
306    Ok(ToolOutcome::ok_json(payload))
307}
308
309async fn edit_row_delete(
310    session: &Arc<ToolSession>,
311    client: &Object,
312    schema: &str,
313    table: &str,
314    args: &Value,
315    include_diff: bool,
316) -> Result<ToolOutcome, ToolError> {
317    let pk = args
318        .get("pk")
319        .and_then(|v| v.as_object())
320        .ok_or_else(|| ToolError::InvalidArgs("pk object is required for delete".into()))?;
321    if pk.is_empty() {
322        return Err(ToolError::InvalidArgs(
323            "pk must contain at least one primary-key column".into(),
324        ));
325    }
326    let column_types = load_column_types(client, schema, table).await?;
327    let (where_sql, params) = pk_where_clause(pk, &column_types)?;
328    if include_diff {
329        ensure_pk_row_cap(client, schema, table, &where_sql, &params).await?;
330    }
331    let before = if include_diff {
332        snapshot_rows(client, schema, table, &where_sql, &params).await?
333    } else {
334        Vec::new()
335    };
336    let sql = format!(
337        "DELETE FROM {} WHERE {} RETURNING *",
338        quote_ref(schema, table),
339        where_sql
340    );
341    let param_refs: Vec<&(dyn ToSql + Sync)> = params
342        .iter()
343        .map(|p| p.as_ref() as &(dyn ToSql + Sync))
344        .collect();
345    let rows = client.query(&sql, &param_refs[..]).await?;
346    let after =
347        redact_row_results(session, None, Some(schema), Some(table), simple_rows_to_json(&rows));
348    let mut payload = json!({
349        "action": "delete",
350        "table": format!("{schema}.{table}"),
351        "rows_affected": after.len(),
352        "rows": after,
353    });
354    if include_diff {
355        payload["before"] = json!(before);
356        payload["after"] = json!(after);
357        payload["diff"] = json!(compute_row_diff(&before, &after));
358    }
359    Ok(ToolOutcome::ok_json(payload))
360}
361
362/// Batched INSERT from a JSON rows array.
363pub async fn import_data(
364    session: &Arc<ToolSession>,
365    args: &Value,
366) -> Result<ToolOutcome, ToolError> {
367    let table_ref = args
368        .get("table")
369        .and_then(|v| v.as_str())
370        .ok_or_else(|| ToolError::InvalidArgs("table is required (schema.name)".into()))?;
371    let rows_val = args
372        .get("rows")
373        .and_then(|v| v.as_array())
374        .ok_or_else(|| ToolError::InvalidArgs("rows array is required".into()))?;
375    if rows_val.is_empty() {
376        return Ok(ToolOutcome::ok_json(json!({
377            "table": table_ref,
378            "rows_imported": 0,
379            "batches": 0,
380        })));
381    }
382
383    let (schema, table) = parse_ref(table_ref).map_err(ToolError::InvalidArgs)?;
384    if !session.filter().allows_table(&schema, &table) {
385        return Err(ToolError::Execution(format!(
386            "Table \"{schema}.{table}\" is denied by policy filter."
387        )));
388    }
389
390    let columns: Vec<String> = if let Some(cols) = args.get("columns").and_then(|v| v.as_array()) {
391        cols.iter()
392            .map(|c| {
393                let s = c
394                    .as_str()
395                    .ok_or_else(|| ToolError::InvalidArgs("columns must be strings".into()))?;
396                validate_column_name(s)?;
397                Ok(s.to_string())
398            })
399            .collect::<Result<Vec<_>, ToolError>>()?
400    } else {
401        let first = rows_val[0]
402            .as_object()
403            .ok_or_else(|| ToolError::InvalidArgs("each row must be a JSON object".into()))?;
404        let mut cols: Vec<String> = first.keys().cloned().collect();
405        cols.sort();
406        for col in &cols {
407            validate_column_name(col)?;
408        }
409        cols
410    };
411
412    let client = session.checkout().await?;
413    client.batch_execute("BEGIN").await?;
414
415    let mut total_imported = 0u64;
416    let mut batches = 0u32;
417    let result = async {
418        for chunk in rows_val.chunks(IMPORT_BATCH_SIZE) {
419            let (sql, params) = build_batch_insert(&schema, &table, &columns, chunk)?;
420            let param_refs: Vec<&(dyn ToSql + Sync)> = params
421                .iter()
422                .map(|p| p.as_ref() as &(dyn ToSql + Sync))
423                .collect();
424            let affected = client.execute(&sql, &param_refs[..]).await?;
425            total_imported += affected;
426            batches += 1;
427        }
428        Ok::<_, ToolError>(())
429    }
430    .await;
431
432    match &result {
433        Ok(_) => {
434            let _ = client.batch_execute("COMMIT").await;
435        }
436        Err(_) => {
437            let _ = client.batch_execute("ROLLBACK").await;
438        }
439    }
440    result?;
441
442    Ok(ToolOutcome::ok_json(json!({
443        "table": format!("{schema}.{table}"),
444        "rows_imported": total_imported,
445        "batches": batches,
446        "columns": columns,
447    })))
448}
449
450fn build_batch_insert(
451    schema: &str,
452    table: &str,
453    columns: &[String],
454    rows: &[Value],
455) -> Result<(String, Vec<Box<dyn ToSql + Sync + Send>>), ToolError> {
456    let quoted_cols = columns
457        .iter()
458        .map(|c| quote_ident(c))
459        .collect::<Vec<_>>()
460        .join(", ");
461    let mut params: Vec<Box<dyn ToSql + Sync + Send>> = Vec::new();
462    let mut value_groups = Vec::new();
463    for row in rows {
464        let obj = row
465            .as_object()
466            .ok_or_else(|| ToolError::InvalidArgs("each row must be a JSON object".into()))?;
467        let mut placeholders = Vec::new();
468        for col in columns {
469            let val = obj.get(col).unwrap_or(&Value::Null);
470            let idx = params.len() + 1;
471            placeholders.push(format!("${idx}"));
472            params.push(json_to_sql_param(None, val)?);
473        }
474        value_groups.push(format!("({})", placeholders.join(", ")));
475    }
476    let sql = format!(
477        "INSERT INTO {} ({}) VALUES {}",
478        quote_ref(schema, table),
479        quoted_cols,
480        value_groups.join(", ")
481    );
482    Ok((sql, params))
483}
484
485/// Run DDL validated for Admin mode inside a transaction.
486pub async fn apply_ddl(
487    session: &Arc<ToolSession>,
488    sql: &str,
489    dry_run: bool,
490) -> Result<ToolOutcome, ToolError> {
491    assert_ddl_statement(sql)?;
492    match validate_write_sql(AccessMode::Admin, sql)? {
493        SqlDecision::Allow => {}
494        SqlDecision::Reject => {
495            return Err(ToolError::Execution(
496                "Security Error: DDL statement is not permitted.".into(),
497            ));
498        }
499    }
500
501    let client = session.checkout().await?;
502    client.batch_execute("BEGIN").await?;
503    let outcome = run_simple_query(&client, sql).await;
504    let rolled_back = dry_run || outcome.is_err();
505    if rolled_back {
506        let _ = client.batch_execute("ROLLBACK").await;
507    } else {
508        let _ = client.batch_execute("COMMIT").await;
509    }
510    let (rows, rows_affected) = outcome?;
511    if !rolled_back {
512        let (connection_id, database) = session.active_context().await;
513        session.mark_index_stale(&connection_id, &database);
514    }
515    Ok(ToolOutcome::ok_json(json!({
516        "dry_run": dry_run,
517        "rolled_back": rolled_back,
518        "rows_affected": rows_affected,
519        "rows": rows,
520        "indexStale": !rolled_back,
521    })))
522}
523
524/// `CREATE INDEX CONCURRENTLY` — must run outside a transaction.
525pub async fn create_index_concurrently(
526    session: &Arc<ToolSession>,
527    sql: &str,
528) -> Result<ToolOutcome, ToolError> {
529    let upper = sql.trim().to_ascii_uppercase();
530    if !upper.contains("CREATE INDEX") || !upper.contains("CONCURRENTLY") {
531        return Err(ToolError::InvalidArgs(
532            "sql must be a CREATE INDEX CONCURRENTLY statement".into(),
533        ));
534    }
535    match validate_write_sql(AccessMode::Admin, sql)? {
536        SqlDecision::Allow => {}
537        SqlDecision::Reject => {
538            return Err(ToolError::Execution(
539                "Security Error: index statement is not permitted.".into(),
540            ));
541        }
542    }
543
544    let client = session.checkout().await?;
545    let (rows, rows_affected) = run_simple_query(&client, sql).await?;
546    let (connection_id, database) = session.active_context().await;
547    session.mark_index_stale(&connection_id, &database);
548    Ok(ToolOutcome::ok_json(json!({
549        "rows_affected": rows_affected,
550        "rows": rows,
551        "note": "CREATE INDEX CONCURRENTLY runs outside a transaction.",
552        "indexStale": true,
553    })))
554}
555
556/// VACUUM / ANALYZE / REINDEX — cannot run inside a transaction.
557pub async fn run_maintenance(
558    session: &Arc<ToolSession>,
559    args: &Value,
560) -> Result<ToolOutcome, ToolError> {
561    let action = args.get("action").and_then(|v| v.as_str()).ok_or_else(|| {
562        ToolError::InvalidArgs("action is required (vacuum|analyze|reindex)".into())
563    })?;
564    let full = args.get("full").and_then(|v| v.as_bool()).unwrap_or(false);
565    let table_ref = args.get("table").and_then(|v| v.as_str());
566
567    let sql = match action.to_ascii_lowercase().as_str() {
568        "vacuum" => build_vacuum_sql(table_ref, full)?,
569        "analyze" => build_analyze_sql(table_ref)?,
570        "reindex" => build_reindex_sql(table_ref)?,
571        other => {
572            return Err(ToolError::InvalidArgs(format!(
573                "Unsupported action \"{other}\". Use vacuum, analyze, or reindex."
574            )));
575        }
576    };
577
578    match validate_write_sql(AccessMode::Admin, &sql)? {
579        SqlDecision::Allow => {}
580        SqlDecision::Reject => {
581            return Err(ToolError::Execution(
582                "Security Error: maintenance statement is not permitted.".into(),
583            ));
584        }
585    }
586
587    let client = session.checkout().await?;
588    let (rows, rows_affected) = run_simple_query(&client, &sql).await?;
589    Ok(ToolOutcome::ok_json(json!({
590        "action": action,
591        "sql": sql,
592        "rows_affected": rows_affected,
593        "rows": rows,
594    })))
595}
596
597/// Cancel or terminate a backend by pid.
598pub async fn terminate_query(
599    session: &Arc<ToolSession>,
600    args: &Value,
601) -> Result<ToolOutcome, ToolError> {
602    let pid = args
603        .get("pid")
604        .and_then(|v| v.as_i64())
605        .ok_or_else(|| ToolError::InvalidArgs("pid is required".into()))?;
606    if pid <= 0 {
607        return Err(ToolError::InvalidArgs(
608            "pid must be a positive integer".into(),
609        ));
610    }
611    let force = args.get("force").and_then(|v| v.as_bool()).unwrap_or(false);
612
613    let client = session.checkout().await?;
614    let own_pid: i32 = client
615        .query_one("SELECT pg_backend_pid()", &[])
616        .await?
617        .get(0);
618    if pid == i64::from(own_pid) {
619        return Err(ToolError::Execution(
620            "refusing to cancel/terminate the current session backend".into(),
621        ));
622    }
623
624    let target = client
625        .query_opt(
626            "SELECT pid, usename, state, query, usesuper FROM pg_stat_activity WHERE pid = $1",
627            &[&(pid as i32)],
628        )
629        .await?;
630    let Some(row) = target else {
631        return Err(ToolError::Execution(format!(
632            "No backend found with pid {pid}"
633        )));
634    };
635    let usesuper: bool = row.get("usesuper");
636    if usesuper {
637        return Err(ToolError::Execution(
638            "refusing to cancel/terminate a superuser backend — use a direct superuser session if required"
639                .into(),
640        ));
641    }
642
643    let fn_name = if force {
644        "pg_terminate_backend"
645    } else {
646        "pg_cancel_backend"
647    };
648    let sql = format!("SELECT {fn_name}($1)");
649    let result: bool = client.query_one(&sql, &[&(pid as i32)]).await?.get(0);
650
651    Ok(ToolOutcome::ok_json(json!({
652        "pid": pid,
653        "force": force,
654        "success": result,
655        "target": {
656            "usename": row.get::<_, Option<String>>("usename"),
657            "state": row.get::<_, Option<String>>("state"),
658            "query": row.get::<_, Option<String>>("query"),
659        },
660    })))
661}
662
663fn build_vacuum_sql(table_ref: Option<&str>, full: bool) -> Result<String, ToolError> {
664    let mut sql = String::from("VACUUM");
665    if full {
666        sql.push_str(" FULL");
667    }
668    if let Some(ref_) = table_ref {
669        let (schema, table) = parse_ref(ref_).map_err(ToolError::InvalidArgs)?;
670        sql.push(' ');
671        sql.push_str(&quote_ref(&schema, &table));
672    }
673    Ok(sql)
674}
675
676fn build_analyze_sql(table_ref: Option<&str>) -> Result<String, ToolError> {
677    let mut sql = String::from("ANALYZE");
678    if let Some(ref_) = table_ref {
679        let (schema, table) = parse_ref(ref_).map_err(ToolError::InvalidArgs)?;
680        sql.push(' ');
681        sql.push_str(&quote_ref(&schema, &table));
682    }
683    Ok(sql)
684}
685
686fn build_reindex_sql(table_ref: Option<&str>) -> Result<String, ToolError> {
687    let ref_ = table_ref.ok_or_else(|| {
688        ToolError::InvalidArgs("table (schema.name) is required for reindex".into())
689    })?;
690    let (schema, table) = parse_ref(ref_).map_err(ToolError::InvalidArgs)?;
691    Ok(format!("REINDEX TABLE {}", quote_ref(&schema, &table)))
692}
693
694fn assert_ddl_statement(sql: &str) -> Result<(), ToolError> {
695    let upper = sql.trim().to_ascii_uppercase();
696    let ddl_prefixes = [
697        "CREATE ",
698        "ALTER ",
699        "DROP ",
700        "TRUNCATE ",
701        "COMMENT ON ",
702        "GRANT ",
703        "REVOKE ",
704        "RENAME ",
705    ];
706    if !ddl_prefixes.iter().any(|p| upper.starts_with(p)) {
707        return Err(ToolError::Execution(
708            "apply_ddl only accepts DDL statements (CREATE, ALTER, DROP, TRUNCATE, …). Use execute_sql for DML."
709                .into(),
710        ));
711    }
712    Ok(())
713}
714
715fn pk_where_clause(
716    pk: &Map<String, Value>,
717    column_types: &HashMap<String, String>,
718) -> Result<(String, Vec<Box<dyn ToSql + Sync + Send>>), ToolError> {
719    let mut where_cols = Vec::new();
720    let mut params: Vec<Box<dyn ToSql + Sync + Send>> = Vec::new();
721    for (col, val) in pk {
722        validate_column_name(col)?;
723        let idx = params.len() + 1;
724        where_cols.push(format!("{} = ${idx}", quote_ident(col)));
725        params.push(json_to_sql_param(
726            column_types.get(col).map(String::as_str),
727            val,
728        )?);
729    }
730    Ok((where_cols.join(" AND "), params))
731}
732
733async fn ensure_pk_row_cap(
734    client: &Object,
735    schema: &str,
736    table: &str,
737    where_sql: &str,
738    params: &[Box<dyn ToSql + Sync + Send>],
739) -> Result<(), ToolError> {
740    let sql = format!(
741        "SELECT COUNT(*)::bigint FROM {} WHERE {where_sql}",
742        quote_ref(schema, table)
743    );
744    let param_refs: Vec<&(dyn ToSql + Sync)> = params
745        .iter()
746        .map(|p| p.as_ref() as &(dyn ToSql + Sync))
747        .collect();
748    let row = client.query_one(&sql, &param_refs[..]).await?;
749    let count: i64 = row.get(0);
750    if count > MUTATION_DIFF_ROW_CAP {
751        return Err(ToolError::Execution(format!(
752            "Mutation would affect {count} rows — diff capture is capped at {MUTATION_DIFF_ROW_CAP}. Narrow the primary key predicate."
753        )));
754    }
755    Ok(())
756}
757
758async fn snapshot_rows(
759    client: &Object,
760    schema: &str,
761    table: &str,
762    where_sql: &str,
763    params: &[Box<dyn ToSql + Sync + Send>],
764) -> Result<Vec<Value>, ToolError> {
765    let sql = format!(
766        "SELECT * FROM {} WHERE {where_sql}",
767        quote_ref(schema, table)
768    );
769    let param_refs: Vec<&(dyn ToSql + Sync)> = params
770        .iter()
771        .map(|p| p.as_ref() as &(dyn ToSql + Sync))
772        .collect();
773    let rows = client.query(&sql, &param_refs[..]).await?;
774    Ok(simple_rows_to_json(&rows))
775}
776
777fn compute_row_diff(before: &[Value], after: &[Value]) -> Vec<Value> {
778    let mut diffs = Vec::new();
779    let pairs = before.len().min(after.len());
780    for idx in 0..pairs {
781        let (Some(b_obj), Some(a_obj)) = (before[idx].as_object(), after[idx].as_object()) else {
782            continue;
783        };
784        for (col, before_val) in b_obj {
785            let after_val = a_obj.get(col).unwrap_or(&Value::Null);
786            if before_val != after_val {
787                diffs.push(json!({
788                    "row": idx,
789                    "column": col,
790                    "before": before_val,
791                    "after": after_val,
792                }));
793            }
794        }
795    }
796    diffs
797}
798
799fn validate_column_name(col: &str) -> Result<(), ToolError> {
800    if !is_safe_ident(col) {
801        return Err(ToolError::InvalidArgs(format!(
802            "Invalid column name \"{col}\"."
803        )));
804    }
805    Ok(())
806}
807
808fn json_to_sql_param(
809    pg_type: Option<&str>,
810    val: &Value,
811) -> Result<Box<dyn ToSql + Sync + Send>, ToolError> {
812    let typ = pg_type.unwrap_or("text").to_ascii_lowercase();
813    match val {
814        Value::Null => Ok(Box::new(None::<String>)),
815        Value::Bool(b) => Ok(Box::new(*b)),
816        Value::String(s) => {
817            if typ.contains("json") {
818                return Ok(Box::new(Json(val.clone())));
819            }
820            if typ.contains("uuid") {
821                let parsed = Uuid::parse_str(s).map_err(|e| {
822                    ToolError::InvalidArgs(format!("invalid uuid for column: {e}"))
823                })?;
824                return Ok(Box::new(parsed));
825            }
826            if typ.contains("int") || typ == "bigint" || typ == "smallint" {
827                let n: i64 = s.parse().map_err(|e| {
828                    ToolError::InvalidArgs(format!("invalid integer for column: {e}"))
829                })?;
830                return Ok(Box::new(n));
831            }
832            if typ.contains("numeric") || typ.contains("decimal") {
833                let d = Decimal::from_str_exact(s).or_else(|_| s.parse::<Decimal>()).map_err(
834                    |e| ToolError::InvalidArgs(format!("invalid numeric for column: {e}")),
835                )?;
836                return Ok(Box::new(d));
837            }
838            if typ.contains("timestamp") {
839                if let Ok(dt) = DateTime::parse_from_rfc3339(s) {
840                    return Ok(Box::new(dt));
841                }
842                if let Ok(dt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%dT%H:%M:%S%.f") {
843                    return Ok(Box::new(dt));
844                }
845            }
846            if typ == "date" {
847                let d = NaiveDate::parse_from_str(s, "%Y-%m-%d").map_err(|e| {
848                    ToolError::InvalidArgs(format!("invalid date for column: {e}"))
849                })?;
850                return Ok(Box::new(d));
851            }
852            Ok(Box::new(s.clone()))
853        }
854        Value::Number(n) => {
855            if typ.contains("json") {
856                return Ok(Box::new(Json(val.clone())));
857            }
858            if typ.contains("int") || typ == "bigint" || typ == "smallint" {
859                let i = n
860                    .as_i64()
861                    .ok_or_else(|| ToolError::InvalidArgs("integer out of range".into()))?;
862                return Ok(Box::new(i));
863            }
864            if typ.contains("numeric") || typ.contains("decimal") {
865                let d = Decimal::from_str_exact(&n.to_string()).map_err(|e| {
866                    ToolError::InvalidArgs(format!("invalid numeric for column: {e}"))
867                })?;
868                return Ok(Box::new(d));
869            }
870            if let Some(f) = n.as_f64() {
871                return Ok(Box::new(f));
872            }
873            Ok(Box::new(n.to_string()))
874        }
875        Value::Array(_) | Value::Object(_) => Ok(Box::new(Json(val.clone()))),
876    }
877}
878
879async fn load_column_types(
880    client: &Object,
881    schema: &str,
882    table: &str,
883) -> Result<HashMap<String, String>, ToolError> {
884    let rows = client
885        .query(
886            r#"SELECT a.attname AS name, format_type(a.atttypid, a.atttypmod) AS typ
887               FROM pg_attribute a
888               JOIN pg_class c ON c.oid = a.attrelid
889               JOIN pg_namespace n ON n.oid = c.relnamespace
890               WHERE n.nspname = $1 AND c.relname = $2
891                 AND a.attnum > 0 AND NOT a.attisdropped"#,
892            &[&schema, &table],
893        )
894        .await?;
895    Ok(rows
896        .iter()
897        .map(|r| (r.get::<_, String>("name"), r.get::<_, String>("typ")))
898        .collect())
899}
900
901async fn run_simple_query(
902    client: &Object,
903    sql: &str,
904) -> Result<(Vec<Value>, Option<u64>), ToolError> {
905    let messages = client.simple_query(sql).await?;
906    Ok(collect_simple_query(messages))
907}
908
909fn collect_simple_query(messages: Vec<SimpleQueryMessage>) -> (Vec<Value>, Option<u64>) {
910    let mut rows = Vec::new();
911    let mut rows_affected = None;
912    for msg in messages {
913        match msg {
914            SimpleQueryMessage::Row(row) => {
915                let mut map = Map::new();
916                for col in row.columns() {
917                    let cell = row
918                        .try_get(col.name())
919                        .ok()
920                        .flatten()
921                        .map(|s| Value::String(s.to_string()))
922                        .unwrap_or(Value::Null);
923                    map.insert(col.name().to_string(), cell);
924                }
925                rows.push(Value::Object(map));
926            }
927            SimpleQueryMessage::CommandComplete(n) => rows_affected = Some(n),
928            SimpleQueryMessage::RowDescription(_) => {}
929            _ => {}
930        }
931    }
932    (rows, rows_affected)
933}
934
935fn simple_rows_to_json(rows: &[tokio_postgres::Row]) -> Vec<Value> {
936    rows_to_json_vec(rows)
937}
938
939fn redact_row_results(
940    session: &ToolSession,
941    sql: Option<&str>,
942    schema: Option<&str>,
943    table: Option<&str>,
944    rows: Vec<Value>,
945) -> Vec<Value> {
946    let filter = session.filter();
947    if filter.pii_columns.is_empty() {
948        return rows;
949    }
950    let tables = if let (Some(schema), Some(table)) = (schema, table) {
951        vec![ObjectRef::new(schema, table)]
952    } else if let Some(sql) = sql {
953        select_table_refs(sql).unwrap_or_default()
954    } else {
955        Vec::new()
956    };
957    redact_pii_in_rows(rows, &filter.pii_columns, &tables).0
958}
959
960#[cfg(test)]
961mod tests {
962    use super::*;
963
964    #[test]
965    fn assert_ddl_rejects_select() {
966        assert!(assert_ddl_statement("SELECT 1").is_err());
967    }
968
969    #[test]
970    fn assert_ddl_accepts_create() {
971        assert!(assert_ddl_statement("CREATE TABLE t (id int)").is_ok());
972    }
973
974    #[test]
975    fn execute_sql_enforces_read_table_policy() {
976        use nexql_policy::{PolicyFilter, SqlDecision, enforce_read_table_policy, validate_readonly_sql};
977
978        let sql = "SELECT * FROM auth.credentials";
979        assert_eq!(validate_readonly_sql(sql).unwrap(), SqlDecision::Allow);
980        let filter = PolicyFilter {
981            deny_schemas: vec!["auth".into()],
982            ..Default::default()
983        };
984        assert!(enforce_read_table_policy(&filter, sql).is_err());
985    }
986
987    #[test]
988    fn build_vacuum_table() {
989        let sql = build_vacuum_sql(Some("public.users"), false).unwrap();
990        assert_eq!(sql, "VACUUM \"public\".\"users\"");
991    }
992
993    #[test]
994    fn build_batch_insert_sql() {
995        let rows = [json!({"id": 1, "name": "a"}), json!({"id": 2, "name": "b"})];
996        let (sql, params) =
997            build_batch_insert("public", "users", &["id".into(), "name".into()], &rows).unwrap();
998        assert!(sql.starts_with("INSERT INTO \"public\".\"users\""));
999        assert_eq!(params.len(), 4);
1000    }
1001}