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 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
97fn statement_kinds(root: &NodeEnum) -> Result<Vec<String>, AdminSqlError> {
101 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}