systemprompt_database/admin/
admin_sql.rs1use 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
99fn statement_kinds(root: &NodeEnum) -> Result<Vec<String>, AdminSqlError> {
103 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}