systemprompt_database/admin/
query_executor.rs1use std::sync::Arc;
12
13use sqlx::postgres::{PgPool, PgRow};
14use sqlx::{Column, Row};
15use thiserror::Error;
16
17use crate::admin::admin_sql::{AdminSql, AdminSqlError, DEFAULT_READONLY_ROW_LIMIT};
18use crate::models::QueryResult;
19use crate::services::postgres::conversion::row_to_json;
20
21#[derive(Error, Debug)]
22pub enum QueryExecutorError {
23 #[error("Invalid admin SQL: {0}")]
24 InvalidSql(#[from] AdminSqlError),
25
26 #[error("Query execution failed: {0}")]
27 ExecutionFailed(#[from] sqlx::Error),
28}
29
30#[derive(Debug)]
31pub struct QueryExecutor {
32 pool: Arc<PgPool>,
33}
34
35impl QueryExecutor {
36 pub const fn new(pool: Arc<PgPool>) -> Self {
37 Self { pool }
38 }
39
40 pub async fn execute_readonly(
41 &self,
42 raw_sql: &str,
43 row_limit: Option<usize>,
44 ) -> Result<QueryResult, QueryExecutorError> {
45 let sql = AdminSql::parse_readonly(raw_sql)?;
46 let start = std::time::Instant::now();
47 let mut tx = self.pool.begin().await?;
48 sqlx::query("SET TRANSACTION READ ONLY")
49 .execute(tx.as_mut())
50 .await?;
51 let rows = sqlx::query(sqlx::AssertSqlSafe(sql.as_str()))
52 .fetch_all(tx.as_mut())
53 .await?;
54 tx.rollback().await?;
55 Ok(to_result(
56 &rows,
57 row_limit.unwrap_or(DEFAULT_READONLY_ROW_LIMIT),
58 start,
59 ))
60 }
61
62 pub async fn execute_write(&self, raw_sql: &str) -> Result<QueryResult, QueryExecutorError> {
63 let sql = AdminSql::parse_unrestricted(raw_sql)?;
64 let start = std::time::Instant::now();
65 let rows = sqlx::query(sqlx::AssertSqlSafe(sql.as_str()))
66 .fetch_all(&*self.pool)
67 .await?;
68 Ok(to_result(&rows, usize::MAX, start))
69 }
70}
71
72fn to_result(rows: &[PgRow], row_limit: usize, start: std::time::Instant) -> QueryResult {
73 let execution_time = u64::try_from(start.elapsed().as_millis()).unwrap_or(u64::MAX);
74 let columns = rows.first().map_or_else(Vec::new, |first_row| {
75 first_row
76 .columns()
77 .iter()
78 .map(|c| c.name().to_owned())
79 .collect()
80 });
81 let total_rows = rows.len();
82 let result_rows = rows.iter().take(row_limit).map(row_to_json).collect();
83 QueryResult {
84 columns,
85 rows: result_rows,
86 row_count: total_rows,
87 execution_time_ms: execution_time,
88 }
89}