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