1use 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
29pub 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
78pub 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, ¶m_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, ¶m_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, ¶m_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
281pub 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, ¶m_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
404pub 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
443pub 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
475pub 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
516pub 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("e_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("e_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}