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