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;
28const MUTATION_DIFF_ROW_CAP: i64 = 100;
29
30pub 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
101pub 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, ¶m_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, ¶m_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, ¶ms).await?;
330 }
331 let before = if include_diff {
332 snapshot_rows(client, schema, table, &where_sql, ¶ms).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, ¶m_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
362pub 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, ¶m_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
485pub 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
524pub 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
556pub 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
597pub 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("e_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("e_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, ¶m_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, ¶m_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}