Skip to main content

systemprompt_database/admin/
admin_sql.rs

1//! Parser/validator for admin-supplied SQL strings.
2//!
3//! Both modes parse with `pg_query` and accept exactly one statement.
4//! [`AdminSql::parse_readonly`] additionally requires the statement root to
5//! be a `SELECT`, `EXPLAIN` of a `SELECT`, or `SHOW`, and refuses any
6//! data-modifying, DDL or utility node anywhere in the tree — including a
7//! CTE that writes. The executor runs read-only statements inside a
8//! `READ ONLY` transaction so Postgres refuses what the parse cannot see
9//! (a volatile function that writes).
10//!
11//! Copyright (c) systemprompt.io — Business Source License 1.1.
12//! See <https://systemprompt.io> for licensing details.
13
14use pg_query::NodeEnum;
15use thiserror::Error;
16
17pub const DEFAULT_READONLY_ROW_LIMIT: usize = 1000;
18
19#[derive(Debug, Error)]
20pub enum AdminSqlError {
21    #[error("SQL query is empty")]
22    Empty,
23    #[error("SQL query could not be parsed: {0}")]
24    Parse(#[from] pg_query::Error),
25    #[error("SQL query syntax tree could not be inspected: {0}")]
26    Inspect(#[source] serde_json::Error),
27    #[error("SQL query contains multiple statements; only one is allowed")]
28    MultipleStatements,
29    #[error("SQL query must be a SELECT, an EXPLAIN of a SELECT, or SHOW")]
30    NotReadOnly,
31    #[error("SQL query contains a data-modifying, DDL or utility statement in read-only mode")]
32    WriteInReadOnly,
33}
34
35#[derive(Debug, Clone, PartialEq, Eq)]
36pub struct AdminSql(String);
37
38impl AdminSql {
39    pub fn parse_readonly(raw: &str) -> Result<Self, AdminSqlError> {
40        let (text, root) = single_statement(raw)?;
41        let inner = match &root {
42            NodeEnum::ExplainStmt(explain) => explain
43                .query
44                .as_ref()
45                .and_then(|q| q.node.as_ref())
46                .ok_or(AdminSqlError::NotReadOnly)?,
47            other => other,
48        };
49        match inner {
50            NodeEnum::SelectStmt(_) | NodeEnum::VariableShowStmt(_) => {},
51            _ => return Err(AdminSqlError::NotReadOnly),
52        }
53        let nested = statement_kinds(inner)?;
54        if nested
55            .iter()
56            .any(|kind| !READONLY_ROOTS.contains(&kind.as_str()))
57        {
58            return Err(AdminSqlError::WriteInReadOnly);
59        }
60        Ok(Self(text))
61    }
62
63    pub fn parse_unrestricted(raw: &str) -> Result<Self, AdminSqlError> {
64        let (text, _) = single_statement(raw)?;
65        Ok(Self(text))
66    }
67
68    pub fn as_str(&self) -> &str {
69        &self.0
70    }
71}
72
73fn single_statement(raw: &str) -> Result<(String, NodeEnum), AdminSqlError> {
74    let parsed = pg_query::parse(raw)?;
75    let mut stmts = parsed.protobuf.stmts.into_iter();
76    let Some(first) = stmts.next() else {
77        return Err(AdminSqlError::Empty);
78    };
79    if stmts.next().is_some() {
80        return Err(AdminSqlError::MultipleStatements);
81    }
82    let node = first
83        .stmt
84        .and_then(|s| s.node)
85        .ok_or(AdminSqlError::Empty)?;
86    let start = usize::try_from(first.stmt_location).unwrap_or(0);
87    let end = if first.stmt_len > 0 {
88        start.saturating_add(usize::try_from(first.stmt_len).unwrap_or(0))
89    } else {
90        raw.len()
91    };
92    let text = raw.get(start..end).unwrap_or(raw).trim();
93    let text = text.strip_suffix(';').unwrap_or(text).trim_end();
94    Ok((text.to_owned(), node))
95}
96
97const READONLY_ROOTS: &[&str] = &["SelectStmt", "VariableShowStmt"];
98
99// Why: the typed `nodes()` walk skips node kinds it does not model, so the
100// exhaustive check serialises the protobuf tree and looks at every statement
101// node wherever it sits — a CTE, a subquery, a function argument.
102fn statement_kinds(root: &NodeEnum) -> Result<Vec<String>, AdminSqlError> {
103    // JSON: pg_query protobuf AST, externally tagged by node variant name
104    let tree = serde_json::to_value(root).map_err(AdminSqlError::Inspect)?;
105    let mut kinds = Vec::new();
106    collect_statement_kinds(&tree, &mut kinds);
107    Ok(kinds)
108}
109
110fn collect_statement_kinds(value: &serde_json::Value, kinds: &mut Vec<String>) {
111    match value {
112        serde_json::Value::Object(map) => {
113            for (key, child) in map {
114                if key.ends_with("Stmt") {
115                    kinds.push(key.clone());
116                }
117                collect_statement_kinds(child, kinds);
118            }
119        },
120        serde_json::Value::Array(items) => {
121            for item in items {
122                collect_statement_kinds(item, kinds);
123            }
124        },
125        _ => {},
126    }
127}