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