Skip to main content

systemprompt_database/admin/
query_executor.rs

1//! Query executor used by the CLI's `infra db query` and `db exec` commands.
2//!
3//! Part of the documented sqlx allowlist: SQL is supplied dynamically by
4//! the operator and validated through [`AdminSql`]. Read-only statements run
5//! inside a `READ ONLY` transaction so Postgres refuses the writes the parse
6//! cannot see (a volatile function that writes).
7//!
8//! Copyright (c) systemprompt.io — Business Source License 1.1.
9//! See <https://systemprompt.io> for licensing details.
10
11use 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}