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 contains multiple statements; only one is allowed")]
26    MultipleStatements,
27    #[error("SQL query must be a SELECT, an EXPLAIN of a SELECT, or SHOW")]
28    NotReadOnly,
29    #[error("SQL query contains a data-modifying, DDL or utility statement in read-only mode")]
30    WriteInReadOnly,
31}
32
33#[derive(Debug, Clone, PartialEq, Eq)]
34pub struct AdminSql(String);
35
36impl AdminSql {
37    pub fn parse_readonly(raw: &str) -> Result<Self, AdminSqlError> {
38        let (text, root) = single_statement(raw)?;
39        let inner = match &root {
40            NodeEnum::ExplainStmt(explain) => explain
41                .query
42                .as_ref()
43                .and_then(|q| q.node.as_ref())
44                .ok_or(AdminSqlError::NotReadOnly)?,
45            other => other,
46        };
47        match inner {
48            NodeEnum::SelectStmt(_) | NodeEnum::VariableShowStmt(_) => {},
49            _ => return Err(AdminSqlError::NotReadOnly),
50        }
51        let nested = statement_kinds(inner)?;
52        if nested
53            .iter()
54            .any(|kind| !READONLY_ROOTS.contains(&kind.as_str()))
55        {
56            return Err(AdminSqlError::WriteInReadOnly);
57        }
58        Ok(Self(text))
59    }
60
61    pub fn parse_unrestricted(raw: &str) -> Result<Self, AdminSqlError> {
62        let (text, _) = single_statement(raw)?;
63        Ok(Self(text))
64    }
65
66    pub fn as_str(&self) -> &str {
67        &self.0
68    }
69}
70
71fn single_statement(raw: &str) -> Result<(String, NodeEnum), AdminSqlError> {
72    let parsed = pg_query::parse(raw)?;
73    let mut stmts = parsed.protobuf.stmts.into_iter();
74    let Some(first) = stmts.next() else {
75        return Err(AdminSqlError::Empty);
76    };
77    if stmts.next().is_some() {
78        return Err(AdminSqlError::MultipleStatements);
79    }
80    let node = first
81        .stmt
82        .and_then(|s| s.node)
83        .ok_or(AdminSqlError::Empty)?;
84    let start = usize::try_from(first.stmt_location).unwrap_or(0);
85    let end = if first.stmt_len > 0 {
86        start.saturating_add(usize::try_from(first.stmt_len).unwrap_or(0))
87    } else {
88        raw.len()
89    };
90    let text = raw.get(start..end).unwrap_or(raw).trim();
91    let text = text.strip_suffix(';').unwrap_or(text).trim_end();
92    Ok((text.to_owned(), node))
93}
94
95const READONLY_ROOTS: &[&str] = &["SelectStmt", "VariableShowStmt"];
96
97// Why: the typed `nodes()` walk skips node kinds it does not model, so the
98// exhaustive check serialises the protobuf tree and looks at every statement
99// node wherever it sits — a CTE, a subquery, a function argument.
100fn statement_kinds(root: &NodeEnum) -> Result<Vec<String>, AdminSqlError> {
101    // JSON: pg_query protobuf AST, externally tagged by node variant name
102    let tree = serde_json::to_value(root)
103        .map_err(|e| AdminSqlError::Parse(pg_query::Error::InvalidJson(e.to_string())))?;
104    let mut kinds = Vec::new();
105    collect_statement_kinds(&tree, &mut kinds);
106    Ok(kinds)
107}
108
109fn collect_statement_kinds(value: &serde_json::Value, kinds: &mut Vec<String>) {
110    match value {
111        serde_json::Value::Object(map) => {
112            for (key, child) in map {
113                if key.ends_with("Stmt") {
114                    kinds.push(key.clone());
115                }
116                collect_statement_kinds(child, kinds);
117            }
118        },
119        serde_json::Value::Array(items) => {
120            for item in items {
121                collect_statement_kinds(item, kinds);
122            }
123        },
124        _ => {},
125    }
126}