use datafusion::sql::sqlparser::ast::{
Expr, ObjectName, PipeOperator, Query, Select, SelectItem, Statement, TableFactor, Visit,
Visitor,
};
use datafusion::sql::sqlparser::dialect::GenericDialect;
use datafusion::sql::sqlparser::parser::Parser;
use std::collections::BTreeMap;
use std::ops::ControlFlow;
#[derive(Debug, thiserror::Error)]
pub(crate) enum StatementRejected {
#[error("statement kind not allowed: {0}")]
DisallowedKind(String),
#[error("could not parse SQL: {0}")]
ParseError(String),
#[error("statement exceeds the {bound} complexity bound")]
Complexity {
bound: &'static str,
},
#[error("statement calls a function outside the closed surface")]
FunctionSurface,
}
impl StatementRejected {
pub(crate) const fn reason_key(&self) -> &'static str {
match self {
Self::DisallowedKind(_) => "statement_kind",
Self::ParseError(_) => "statement_parse",
Self::Complexity { bound } => bound,
Self::FunctionSurface => "function_surface",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum AllowedStatement {
Query,
Explain,
}
pub(crate) fn check_statement_allowed(
sql: &str,
allow_explain: bool,
) -> Result<AllowedStatement, StatementRejected> {
let statements = Parser::parse_sql(&GenericDialect {}, sql)
.map_err(|err| StatementRejected::ParseError(err.to_string()))?;
match statements.as_slice() {
[Statement::Query(query)] => {
reject_recursive_cte(query)?;
reject_ordered_information_schema(query)?;
reject_outside_surface(query)?;
reject_over_complex(query)?;
Ok(AllowedStatement::Query)
}
[Statement::Explain { analyze: true, .. }] => Err(StatementRejected::DisallowedKind(
"EXPLAIN ANALYZE executes the query and is never allowed".to_string(),
)),
[
Statement::Explain {
analyze: false,
statement,
..
},
] => {
if !allow_explain {
return Err(StatementRejected::DisallowedKind(
"EXPLAIN is not enabled on this surface".to_string(),
));
}
match statement.as_ref() {
Statement::Query(query) => {
reject_recursive_cte(query)?;
reject_outside_surface(query)?;
reject_over_complex(query)?;
Ok(AllowedStatement::Explain)
}
other => Err(StatementRejected::DisallowedKind(format!(
"EXPLAIN of a {} statement is not allowed",
statement_kind(other)
))),
}
}
[other] => Err(StatementRejected::DisallowedKind(format!(
"{} statements are not allowed",
statement_kind(other)
))),
[] => Err(StatementRejected::DisallowedKind(
"no statement found".to_string(),
)),
multiple => Err(StatementRejected::DisallowedKind(format!(
"expected exactly one statement, found {}",
multiple.len()
))),
}
}
fn reject_recursive_cte(query: &Query) -> Result<(), StatementRejected> {
let mut scan = RecursiveCteScan::default();
let _: ControlFlow<()> = query.visit(&mut scan);
if scan.has_recursive_with {
return Err(StatementRejected::DisallowedKind(
"WITH RECURSIVE is not supported".to_string(),
));
}
Ok(())
}
#[derive(Debug, Default)]
struct RecursiveCteScan {
has_recursive_with: bool,
}
impl Visitor for RecursiveCteScan {
type Break = ();
fn pre_visit_query(&mut self, query: &Query) -> ControlFlow<Self::Break> {
if query.with.as_ref().is_some_and(|with| with.recursive) {
self.has_recursive_with = true;
}
ControlFlow::Continue(())
}
}
fn reject_ordered_information_schema(query: &Query) -> Result<(), StatementRejected> {
let mut scan = InformationSchemaOrderingScan::default();
let _: ControlFlow<()> = query.visit(&mut scan);
if scan.has_information_schema_reference && scan.has_ordering {
return Err(StatementRejected::DisallowedKind(
"ordering an information_schema query is not supported: it can silently drop rows \
from the result — remove the ORDER BY (or SORT BY, or pipe ORDER BY) and query \
information_schema without ordering it"
.to_string(),
));
}
Ok(())
}
#[derive(Debug, Default)]
struct InformationSchemaOrderingScan {
has_information_schema_reference: bool,
has_ordering: bool,
}
impl Visitor for InformationSchemaOrderingScan {
type Break = ();
fn pre_visit_relation(&mut self, relation: &ObjectName) -> ControlFlow<Self::Break> {
if references_information_schema(relation) {
self.has_information_schema_reference = true;
}
ControlFlow::Continue(())
}
fn pre_visit_query(&mut self, query: &Query) -> ControlFlow<Self::Break> {
if query.order_by.is_some()
|| query
.pipe_operators
.iter()
.any(|operator| matches!(operator, PipeOperator::OrderBy { .. }))
{
self.has_ordering = true;
}
ControlFlow::Continue(())
}
fn pre_visit_select(&mut self, select: &Select) -> ControlFlow<Self::Break> {
if !select.sort_by.is_empty() {
self.has_ordering = true;
}
ControlFlow::Continue(())
}
}
fn reject_outside_surface(query: &Query) -> Result<(), StatementRejected> {
let mut scan = FunctionSurfaceScan::default();
let _: ControlFlow<()> = query.visit(&mut scan);
if scan.outside {
return Err(StatementRejected::FunctionSurface);
}
Ok(())
}
#[derive(Debug, Default)]
struct FunctionSurfaceScan {
outside: bool,
}
impl Visitor for FunctionSurfaceScan {
type Break = ();
fn pre_visit_expr(&mut self, expr: &Expr) -> ControlFlow<Self::Break> {
match expr {
Expr::Function(function) => {
let name = function
.name
.0
.last()
.and_then(|part| part.as_ident())
.map_or_else(String::new, |ident| ident.value.to_ascii_lowercase());
if !crate::function_surface::allowed_names().contains(&name) {
self.outside = true;
}
}
Expr::Overlay { .. } => self.outside = true,
_ => {}
}
ControlFlow::Continue(())
}
fn pre_visit_table_factor(&mut self, factor: &TableFactor) -> ControlFlow<Self::Break> {
match factor {
TableFactor::Table { args: Some(_), .. } | TableFactor::Function { .. } => {
self.outside = true;
}
_ => {}
}
ControlFlow::Continue(())
}
}
pub(crate) const MAX_EXPRESSION_NODES: usize = 1024;
pub(crate) const MAX_EXPRESSION_DEPTH: usize = 32;
pub(crate) const MAX_FUNCTION_CALLS: usize = 128;
pub(crate) const MAX_SELECT_ITEMS: usize = 512;
pub(crate) const MAX_QUERY_NODES: usize = 64;
pub(crate) const MAX_EXPRESSION_OUTPUT_REFERENCES: usize = 4;
pub(crate) const MAX_SELECT_COLUMN_FANOUT: usize = 4;
pub(crate) const MAX_STATEMENT_COLUMN_REFERENCES: usize = 8;
fn reject_over_complex(query: &Query) -> Result<(), StatementRejected> {
let mut scan = ComplexityScan::default();
let _: ControlFlow<()> = query.visit(&mut scan);
scan.check()
}
#[derive(Debug, Default)]
struct ComplexityScan {
expression_nodes: usize,
expression_depth: usize,
max_expression_depth: usize,
function_calls: usize,
max_select_items: usize,
query_nodes: usize,
max_expression_output_references: usize,
max_column_fanout: usize,
statement_column_references: BTreeMap<String, usize>,
}
impl ComplexityScan {
fn check(&self) -> Result<(), StatementRejected> {
let statement_fanout = self
.statement_column_references
.values()
.max()
.copied()
.unwrap_or(0);
let checks = [
(
self.expression_nodes,
MAX_EXPRESSION_NODES,
"expression_nodes",
),
(
self.max_expression_depth,
MAX_EXPRESSION_DEPTH,
"expression_depth",
),
(self.function_calls, MAX_FUNCTION_CALLS, "function_calls"),
(self.max_select_items, MAX_SELECT_ITEMS, "select_items"),
(self.query_nodes, MAX_QUERY_NODES, "query_nodes"),
(
self.max_expression_output_references,
MAX_EXPRESSION_OUTPUT_REFERENCES,
"expression_output_references",
),
(
self.max_column_fanout,
MAX_SELECT_COLUMN_FANOUT,
"column_fanout",
),
(
statement_fanout,
MAX_STATEMENT_COLUMN_REFERENCES,
"statement_column_references",
),
];
for &(observed, limit, bound) in &checks {
if observed > limit {
return Err(StatementRejected::Complexity { bound });
}
}
Ok(())
}
fn merge_emit_map(&mut self, fanout: &mut BTreeMap<String, usize>, expr: &Expr) {
let mut emit = EmitScan::default();
for (name, count) in output_references(expr, &mut emit) {
*fanout.entry(name.clone()).or_default() += count;
*self.statement_column_references.entry(name).or_default() += count;
}
self.max_expression_output_references =
self.max_expression_output_references.max(emit.max_node);
}
}
impl Visitor for ComplexityScan {
type Break = ();
fn pre_visit_expr(&mut self, expr: &Expr) -> ControlFlow<Self::Break> {
if self.expression_depth == 0 {
let mut emit = EmitScan::default();
let _ = output_references(expr, &mut emit);
self.max_expression_output_references =
self.max_expression_output_references.max(emit.max_node);
}
if matches!(expr, Expr::Function(_)) {
self.function_calls += 1;
}
self.expression_nodes += 1;
self.expression_depth += 1;
self.max_expression_depth = self.max_expression_depth.max(self.expression_depth);
ControlFlow::Continue(())
}
fn post_visit_expr(&mut self, _expr: &Expr) -> ControlFlow<Self::Break> {
self.expression_depth -= 1;
ControlFlow::Continue(())
}
fn pre_visit_select(&mut self, select: &Select) -> ControlFlow<Self::Break> {
self.max_select_items = self.max_select_items.max(select.projection.len());
let mut fanout: BTreeMap<String, usize> = BTreeMap::new();
let mut wildcard = false;
for item in &select.projection {
match item {
SelectItem::UnnamedExpr(expr)
| SelectItem::ExprWithAlias { expr, .. }
| SelectItem::ExprWithAliases { expr, .. }
| SelectItem::QualifiedWildcard(
datafusion::sql::sqlparser::ast::SelectItemQualifiedWildcardKind::Expr(expr),
_,
) => self.merge_emit_map(&mut fanout, expr),
SelectItem::QualifiedWildcard(
datafusion::sql::sqlparser::ast::SelectItemQualifiedWildcardKind::ObjectName(_),
_,
)
| SelectItem::Wildcard(_) => wildcard = true,
}
}
let select_fanout = fanout.values().max().copied().unwrap_or(0) + usize::from(wildcard);
self.max_column_fanout = self.max_column_fanout.max(select_fanout);
ControlFlow::Continue(())
}
fn pre_visit_query(&mut self, _query: &Query) -> ControlFlow<Self::Break> {
self.query_nodes += 1;
ControlFlow::Continue(())
}
}
#[derive(Debug, Default)]
struct EmitScan {
max_node: usize,
}
#[allow(clippy::too_many_lines)]
fn output_references(expr: &Expr, emit: &mut EmitScan) -> BTreeMap<String, usize> {
macro_rules! probe {
($child:expr) => {{
let _ = output_references($child, emit);
}};
}
fn merge_into(into: &mut BTreeMap<String, usize>, child: BTreeMap<String, usize>) {
for (name, count) in child {
*into.entry(name).or_default() += count;
}
}
fn one(name: &str) -> BTreeMap<String, usize> {
let mut map = BTreeMap::new();
map.insert(name.to_ascii_lowercase(), 1);
map
}
fn emit_children<'a>(
children: impl IntoIterator<Item = &'a Expr>,
emit: &mut EmitScan,
) -> BTreeMap<String, usize> {
let mut map = BTreeMap::new();
for child in children {
merge_into(&mut map, output_references(child, emit));
}
map
}
fn emit_function_args(
map: &mut BTreeMap<String, usize>,
emit: &mut EmitScan,
arguments: &datafusion::sql::sqlparser::ast::FunctionArguments,
) {
use datafusion::sql::sqlparser::ast::{FunctionArg, FunctionArgExpr, FunctionArguments};
if let FunctionArguments::List(list) = arguments {
for arg in &list.args {
match arg {
FunctionArg::Named { arg, .. } | FunctionArg::Unnamed(arg) => {
if let FunctionArgExpr::Expr(arg) = arg {
merge_into(map, output_references(arg, emit));
}
}
FunctionArg::ExprNamed { name, arg, .. } => {
let _ = output_references(name, emit);
if let FunctionArgExpr::Expr(arg) = arg {
merge_into(map, output_references(arg, emit));
}
}
}
}
}
}
let local = match expr {
Expr::Identifier(ident) => one(&ident.value),
Expr::CompoundIdentifier(idents) => idents
.last()
.map_or_else(BTreeMap::new, |last| one(&last.value)),
Expr::CompoundFieldAccess { root, access_chain } => {
let map = output_references(root, emit);
for access in access_chain {
if let datafusion::sql::sqlparser::ast::AccessExpr::Subscript(subscript) = access {
use datafusion::sql::sqlparser::ast::Subscript;
match subscript {
Subscript::Index { index } => probe!(index),
Subscript::Slice {
lower_bound,
upper_bound,
stride,
} => {
for bound in [lower_bound, upper_bound, stride].into_iter().flatten() {
probe!(bound);
}
}
}
}
}
map
}
Expr::JsonAccess { value, .. } => output_references(value, emit),
Expr::Function(function) => {
let mut map = BTreeMap::new();
emit_function_args(&mut map, emit, &function.parameters);
emit_function_args(&mut map, emit, &function.args);
if let Some(filter) = &function.filter {
probe!(filter);
}
if let Some(datafusion::sql::sqlparser::ast::WindowType::WindowSpec(spec)) =
&function.over
{
for key in &spec.partition_by {
probe!(key);
}
for order in &spec.order_by {
probe!(&order.expr);
}
}
for order in &function.within_group {
probe!(&order.expr);
}
map
}
Expr::Case {
operand,
conditions,
else_result,
..
} => {
if let Some(operand) = operand {
probe!(operand);
}
let mut map = BTreeMap::new();
for when in conditions {
probe!(&when.condition);
merge_into(&mut map, output_references(&when.result, emit));
}
if let Some(else_result) = else_result {
merge_into(&mut map, output_references(else_result, emit));
}
map
}
Expr::BinaryOp { left, op, right } => {
use datafusion::sql::sqlparser::ast::BinaryOperator as B;
match op {
B::StringConcat
| B::Arrow
| B::LongArrow
| B::HashArrow
| B::HashLongArrow
| B::HashMinus
| B::DoubleHash
| B::Custom(_)
| B::PGCustomBinaryOperator(_) => {
emit_children([left.as_ref(), right.as_ref()], emit)
}
_ => {
probe!(left);
probe!(right);
BTreeMap::new()
}
}
}
Expr::UnaryOp { expr: inner, .. }
| Expr::Ceil { expr: inner, .. }
| Expr::Floor { expr: inner, .. }
| Expr::Extract { expr: inner, .. }
| Expr::OuterJoin(inner)
| Expr::IsFalse(inner)
| Expr::IsNotFalse(inner)
| Expr::IsTrue(inner)
| Expr::IsNotTrue(inner)
| Expr::IsNull(inner)
| Expr::IsNotNull(inner)
| Expr::IsUnknown(inner)
| Expr::IsNotUnknown(inner)
| Expr::IsNormalized { expr: inner, .. }
| Expr::InSubquery { expr: inner, .. }
| Expr::InUnnest { expr: inner, .. }
| Expr::Position { expr: inner, .. } => {
probe!(inner);
BTreeMap::new()
}
Expr::Nested(inner)
| Expr::Collate { expr: inner, .. }
| Expr::Named { expr: inner, .. }
| Expr::Prior(inner)
| Expr::Cast { expr: inner, .. }
| Expr::Prefixed { value: inner, .. } => output_references(inner, emit),
Expr::Substring {
expr: inner,
substring_from,
substring_for,
..
} => {
for position in [substring_from, substring_for].into_iter().flatten() {
probe!(position);
}
output_references(inner, emit)
}
Expr::Trim {
expr: inner,
trim_what,
trim_characters,
..
} => {
if let Some(what) = trim_what {
probe!(what);
}
for character in trim_characters.iter().flatten() {
probe!(character);
}
output_references(inner, emit)
}
Expr::Overlay {
expr: inner,
overlay_what,
overlay_from,
overlay_for,
} => {
probe!(overlay_from);
if let Some(for_expr) = overlay_for {
probe!(for_expr);
}
emit_children([inner.as_ref(), overlay_what.as_ref()], emit)
}
Expr::Convert {
expr: inner,
styles,
..
} => {
for style in styles {
probe!(style);
}
output_references(inner, emit)
}
Expr::Tuple(children)
| Expr::Array(datafusion::sql::sqlparser::ast::Array { elem: children, .. })
| Expr::Struct {
values: children, ..
} => emit_children(children.iter(), emit),
Expr::Dictionary(fields) => {
emit_children(fields.iter().map(|field| field.value.as_ref()), emit)
}
Expr::Map(map) => emit_children(
map.entries
.iter()
.flat_map(|entry| [entry.key.as_ref(), entry.value.as_ref()]),
emit,
),
Expr::Subquery(_) => {
let mut map = BTreeMap::new();
map.insert(
SCALAR_SUBQUERY_KEY.to_owned(),
MAX_EXPRESSION_OUTPUT_REFERENCES,
);
map
}
Expr::IsDistinctFrom(left, right)
| Expr::IsNotDistinctFrom(left, right)
| Expr::AnyOp { left, right, .. }
| Expr::AllOp { left, right, .. } => {
probe!(left);
probe!(right);
BTreeMap::new()
}
Expr::InList {
expr: inner, list, ..
} => {
probe!(inner);
for item in list {
probe!(item);
}
BTreeMap::new()
}
Expr::Between {
expr: inner,
low,
high,
..
} => {
probe!(inner);
probe!(low);
probe!(high);
BTreeMap::new()
}
Expr::Like {
expr: inner,
pattern,
..
}
| Expr::ILike {
expr: inner,
pattern,
..
}
| Expr::SimilarTo {
expr: inner,
pattern,
..
}
| Expr::RLike {
expr: inner,
pattern,
..
} => {
probe!(inner);
probe!(pattern);
BTreeMap::new()
}
Expr::AtTimeZone {
timestamp,
time_zone,
} => {
probe!(timestamp);
probe!(time_zone);
BTreeMap::new()
}
Expr::GroupingSets(sets) | Expr::Cube(sets) | Expr::Rollup(sets) => {
for set in sets {
for member in set {
probe!(member);
}
}
BTreeMap::new()
}
Expr::Exists { .. }
| Expr::MatchAgainst { .. }
| Expr::Value(_)
| Expr::TypedString(_)
| Expr::Wildcard(_)
| Expr::QualifiedWildcard(..)
| Expr::Lambda(_) => BTreeMap::new(),
Expr::Interval(interval) => {
probe!(&interval.value);
BTreeMap::new()
}
Expr::MemberOf(member_of) => {
probe!(&member_of.value);
probe!(&member_of.array);
BTreeMap::new()
}
};
emit.max_node = emit
.max_node
.max(local.values().max().copied().unwrap_or(0));
local
}
const SCALAR_SUBQUERY_KEY: &str = "?subquery";
fn references_information_schema(relation: &ObjectName) -> bool {
relation.0.iter().any(|part| {
part.as_ident()
.is_some_and(|ident| ident.value.eq_ignore_ascii_case("information_schema"))
})
}
fn statement_kind(statement: &Statement) -> String {
statement
.to_string()
.split_whitespace()
.next()
.unwrap_or("unknown")
.trim_end_matches(|c: char| !c.is_ascii_alphanumeric())
.to_ascii_uppercase()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn select_is_allowed() {
assert_eq!(
check_statement_allowed("SELECT 1", false).expect("must be allowed"),
AllowedStatement::Query
);
}
#[test]
fn with_select_is_allowed() {
assert_eq!(
check_statement_allowed("WITH t AS (SELECT 1) SELECT * FROM t", false)
.expect("must be allowed"),
AllowedStatement::Query
);
}
#[test]
fn explain_select_allowed_when_opted_in() {
assert_eq!(
check_statement_allowed("EXPLAIN SELECT 1", true).expect("must be allowed"),
AllowedStatement::Explain
);
}
#[test]
fn explain_select_rejected_without_opt_in() {
let err = check_statement_allowed("EXPLAIN SELECT 1", false).unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn explain_analyze_rejected_with_opt_in() {
let err = check_statement_allowed("EXPLAIN ANALYZE SELECT 1", true).unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn explain_analyze_rejected_without_opt_in() {
let err = check_statement_allowed("EXPLAIN ANALYZE SELECT 1", false).unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn explain_of_insert_rejected_even_with_opt_in() {
let err = check_statement_allowed("EXPLAIN INSERT INTO t VALUES (1)", true).unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn create_table_rejected() {
let err = check_statement_allowed("CREATE TABLE t (a INT)", false).unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn drop_table_rejected() {
let err = check_statement_allowed("DROP TABLE t", false).unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn insert_rejected() {
let err = check_statement_allowed("INSERT INTO t VALUES (1)", false).unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn update_rejected() {
let err = check_statement_allowed("UPDATE t SET a = 1", false).unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn delete_rejected() {
let err = check_statement_allowed("DELETE FROM t", false).unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn show_rejected() {
let err = check_statement_allowed("SHOW TABLES", false).unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn copy_rejected() {
let err = check_statement_allowed("COPY t TO 'out.csv'", false).unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn set_rejected() {
let err = check_statement_allowed("SET timezone = 'UTC'", false).unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn multi_statement_batch_rejected() {
let err = check_statement_allowed("SELECT 1; SELECT 2", false).unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn empty_string_rejected() {
let err = check_statement_allowed("", false).unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn garbage_rejected_as_parse_error() {
let err = check_statement_allowed("not even close to sql (((", false).unwrap_err();
assert!(matches!(err, StatementRejected::ParseError(_)));
}
#[test]
fn information_schema_direct_order_by_rejected() {
let err = check_statement_allowed(
"SELECT table_name FROM information_schema.tables ORDER BY table_name",
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn catalog_qualified_information_schema_order_by_rejected() {
let err = check_statement_allowed(
"SELECT table_name FROM datafusion.information_schema.tables ORDER BY table_name",
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn quoted_mixed_case_information_schema_order_by_rejected() {
let err = check_statement_allowed(
r#"SELECT "TABLES"."TABLE_NAME" FROM "Information_Schema"."TABLES" ORDER BY "TABLES"."TABLE_NAME""#,
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn aliased_information_schema_order_by_rejected() {
let err = check_statement_allowed(
"SELECT t.table_name FROM information_schema.tables t ORDER BY t.table_name",
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn join_with_information_schema_order_by_rejected() {
let err = check_statement_allowed(
"SELECT t.table_name FROM information_schema.tables t \
JOIN information_schema.columns c ON t.table_name = c.table_name \
ORDER BY t.table_name",
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn ordering_in_subquery_over_information_schema_rejected() {
let err = check_statement_allowed(
"SELECT * FROM (SELECT table_name FROM information_schema.tables \
ORDER BY table_name) sub",
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn outer_order_by_over_cte_reading_information_schema_rejected() {
let err = check_statement_allowed(
"WITH t AS (SELECT table_name FROM information_schema.tables) \
SELECT * FROM t ORDER BY table_name",
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn pipe_style_order_by_over_information_schema_rejected() {
let err = check_statement_allowed(
"SELECT table_name FROM information_schema.tables |> ORDER BY table_name",
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn union_with_information_schema_arm_and_outer_order_by_rejected() {
let err = check_statement_allowed(
"SELECT table_name FROM information_schema.tables \
UNION SELECT name FROM usage ORDER BY table_name",
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn outer_order_by_over_derived_subquery_reading_information_schema_rejected() {
let err = check_statement_allowed(
"SELECT table_name FROM (SELECT table_name FROM information_schema.tables) sub \
ORDER BY table_name",
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn information_schema_without_order_by_allowed() {
assert_eq!(
check_statement_allowed("SELECT table_name FROM information_schema.tables", false)
.expect("must be allowed"),
AllowedStatement::Query
);
}
#[test]
fn order_by_over_ordinary_table_allowed() {
assert_eq!(
check_statement_allowed("SELECT * FROM usage ORDER BY position", false)
.expect("must be allowed"),
AllowedStatement::Query
);
}
#[test]
fn explain_over_ordered_information_schema_allowed() {
assert_eq!(
check_statement_allowed(
"EXPLAIN SELECT table_name FROM information_schema.tables ORDER BY table_name",
true,
)
.expect("must be allowed"),
AllowedStatement::Explain
);
}
#[test]
fn top_level_with_recursive_rejected() {
let err = check_statement_allowed(
"WITH RECURSIVE n AS (SELECT 1 AS v UNION ALL SELECT v + 1 FROM n) \
SELECT v FROM n",
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn with_recursive_in_derived_subquery_rejected() {
let err = check_statement_allowed(
"SELECT v FROM \
(WITH RECURSIVE n AS (SELECT 1 AS v UNION ALL SELECT v + 1 FROM n) \
SELECT v FROM n) sub",
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn with_recursive_in_exists_subquery_rejected() {
let err = check_statement_allowed(
"SELECT v FROM usage WHERE EXISTS \
(WITH RECURSIVE n AS (SELECT 1 AS v UNION ALL SELECT v + 1 FROM n) \
SELECT v FROM n)",
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn with_recursive_nested_in_cte_body_rejected() {
let err = check_statement_allowed(
"WITH outer_cte AS \
(WITH RECURSIVE n AS (SELECT 1 AS v UNION ALL SELECT v + 1 FROM n) \
SELECT v FROM n) \
SELECT v FROM outer_cte",
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn with_recursive_in_in_subquery_rejected() {
let err = check_statement_allowed(
"SELECT v FROM usage WHERE position IN \
(WITH RECURSIVE n AS (SELECT 1 AS v UNION ALL SELECT v + 1 FROM n) \
SELECT v FROM n)",
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn explain_of_with_recursive_rejected() {
let err = check_statement_allowed(
"EXPLAIN WITH RECURSIVE n AS (SELECT 1 AS v) SELECT v FROM n",
true,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn non_recursive_with_still_allowed() {
assert_eq!(
check_statement_allowed(
"WITH t AS (SELECT position FROM usage) SELECT position FROM t",
false,
)
.expect("must be allowed"),
AllowedStatement::Query
);
}
#[test]
fn with_recursive_in_union_arm_rejected() {
let err = check_statement_allowed(
"SELECT position FROM usage UNION ALL SELECT v FROM \
(WITH RECURSIVE n AS (SELECT 1 AS v UNION ALL SELECT v + 1 FROM n) \
SELECT v FROM n) sub",
false,
)
.unwrap_err();
assert!(matches!(err, StatementRejected::DisallowedKind(_)));
}
#[test]
fn refused_function_by_name_and_case_rejected() {
for sql in [
"SELECT repeat('x', 2)",
"SELECT REPEAT('x', 2)",
"SELECT Repeat('x', 2)",
] {
let err = check_statement_allowed(sql, false).unwrap_err();
assert!(
matches!(err, StatementRejected::FunctionSurface),
"{sql} must refuse as FunctionSurface, got {err:?}"
);
}
}
#[test]
fn every_refused_name_and_alias_rejected() {
for name in crate::function_surface::refused_names() {
let sql = format!("SELECT {name}('x', 1)");
let err = check_statement_allowed(&sql, false).unwrap_err();
assert!(
matches!(err, StatementRejected::FunctionSurface),
"{sql} must refuse as FunctionSurface, got {err:?}"
);
}
}
#[test]
fn qualified_refused_function_rejected() {
let err = check_statement_allowed("SELECT pg_catalog.repeat('x', 2)", false).unwrap_err();
assert!(matches!(err, StatementRejected::FunctionSurface));
}
#[test]
fn overlay_syntax_rejected() {
let err =
check_statement_allowed("SELECT OVERLAY('abc' PLACING 'x' FROM 2)", false).unwrap_err();
assert!(matches!(err, StatementRejected::FunctionSurface));
}
#[test]
fn table_function_in_from_rejected() {
for sql in [
"SELECT * FROM range(1, 10)",
"SELECT * FROM generate_series(1, 10)",
] {
let err = check_statement_allowed(sql, false).unwrap_err();
assert!(
matches!(err, StatementRejected::FunctionSurface),
"{sql} must refuse as FunctionSurface, got {err:?}"
);
}
}
#[test]
fn unclassified_function_rejected() {
let err = check_statement_allowed("SELECT a_function_nobody_listed(1)", false).unwrap_err();
assert!(matches!(err, StatementRejected::FunctionSurface));
}
#[test]
fn unnest_table_factor_allowed() {
assert_eq!(
check_statement_allowed("SELECT * FROM UNNEST(make_array(1, 2))", false)
.expect("must be allowed"),
AllowedStatement::Query
);
}
#[test]
fn allowed_function_call_allowed() {
assert_eq!(
check_statement_allowed(
"SELECT substr(name, 1, 2), count(*) FROM usage GROUP BY name",
false,
)
.expect("must be allowed"),
AllowedStatement::Query
);
}
fn wide_select(item: &str, n: usize) -> String {
let items: Vec<String> = (0..n).map(|_| item.to_owned()).collect();
format!("SELECT {}", items.join(", "))
}
fn nested_parens(n: usize) -> String {
format!("SELECT {}x{}", "(".repeat(n), ")".repeat(n))
}
#[test]
fn expression_nodes_at_and_over_bound() {
let at = format!(
"SELECT x IN ({})",
(0..MAX_EXPRESSION_NODES - 2)
.map(|v| v.to_string())
.collect::<Vec<_>>()
.join(", ")
);
let over = format!(
"SELECT x IN ({})",
(0..MAX_EXPRESSION_NODES - 1)
.map(|v| v.to_string())
.collect::<Vec<_>>()
.join(", ")
);
assert!(check_statement_allowed(&at, false).is_ok());
let err = check_statement_allowed(&over, false).unwrap_err();
assert!(
matches!(
err,
StatementRejected::Complexity {
bound: "expression_nodes"
}
),
"got {err:?}"
);
}
#[test]
fn expression_depth_at_and_over_bound() {
let at = nested_parens(MAX_EXPRESSION_DEPTH - 1);
let over = nested_parens(MAX_EXPRESSION_DEPTH);
assert!(check_statement_allowed(&at, false).is_ok());
let err = check_statement_allowed(&over, false).unwrap_err();
assert!(
matches!(
err,
StatementRejected::Complexity {
bound: "expression_depth"
}
),
"got {err:?}"
);
}
#[test]
fn function_calls_at_and_over_bound() {
let at = wide_select("abs(1)", MAX_FUNCTION_CALLS);
let over = wide_select("abs(1)", MAX_FUNCTION_CALLS + 1);
assert!(check_statement_allowed(&at, false).is_ok());
let err = check_statement_allowed(&over, false).unwrap_err();
assert!(
matches!(
err,
StatementRejected::Complexity {
bound: "function_calls"
}
),
"got {err:?}"
);
}
#[test]
fn select_items_at_and_over_bound() {
let at = wide_select("1", MAX_SELECT_ITEMS);
let over = wide_select("1", MAX_SELECT_ITEMS + 1);
assert!(check_statement_allowed(&at, false).is_ok());
let err = check_statement_allowed(&over, false).unwrap_err();
assert!(
matches!(
err,
StatementRejected::Complexity {
bound: "select_items"
}
),
"got {err:?}"
);
}
#[test]
fn query_nodes_at_and_over_bound() {
let ctes = |n: usize| {
let names: Vec<String> = (0..n).map(|i| format!("c{i} AS (SELECT 1)")).collect();
format!("WITH {} SELECT 1", names.join(", "))
};
assert!(check_statement_allowed(&ctes(MAX_QUERY_NODES - 1), false).is_ok());
let err = check_statement_allowed(&ctes(MAX_QUERY_NODES), false).unwrap_err();
assert!(
matches!(
err,
StatementRejected::Complexity {
bound: "query_nodes"
}
),
"got {err:?}"
);
}
#[test]
fn expression_output_references_at_and_over_bound() {
let at = "SELECT c || c || c || c FROM t";
let over = "SELECT c || c || c || c || c FROM t";
assert!(check_statement_allowed(at, false).is_ok());
let err = check_statement_allowed(over, false).unwrap_err();
assert!(
matches!(
err,
StatementRejected::Complexity {
bound: "expression_output_references"
}
),
"got {err:?}"
);
}
#[test]
fn column_fanout_at_and_over_bound() {
let at = wide_select("c", MAX_SELECT_COLUMN_FANOUT);
let over = wide_select("c", MAX_SELECT_COLUMN_FANOUT + 1);
assert!(check_statement_allowed(&at, false).is_ok());
let err = check_statement_allowed(&over, false).unwrap_err();
assert!(
matches!(
err,
StatementRejected::Complexity {
bound: "column_fanout"
}
),
"got {err:?}"
);
}
#[test]
fn statement_column_references_at_and_over_bound() {
let union = |n: usize| {
(0..n)
.map(|_| "SELECT c FROM t".to_owned())
.collect::<Vec<_>>()
.join(" UNION ALL ")
};
assert!(check_statement_allowed(&union(MAX_STATEMENT_COLUMN_REFERENCES), false).is_ok());
let err = check_statement_allowed(&union(MAX_STATEMENT_COLUMN_REFERENCES + 1), false)
.unwrap_err();
assert!(
matches!(
err,
StatementRejected::Complexity {
bound: "statement_column_references"
}
),
"got {err:?}"
);
}
#[test]
fn join_predicate_does_not_consume_emit_budget() {
assert!(
check_statement_allowed(
"SELECT s.position FROM l JOIN s ON \
l.a = s.a AND l.b = s.b AND l.c = s.c AND l.d = s.d AND \
l.e = s.e AND l.f = s.f AND l.g = s.g AND l.h = s.h",
false,
)
.is_ok()
);
}
#[test]
fn explain_of_over_bound_query_rejected() {
let sql = format!("EXPLAIN {}", wide_select("1", MAX_SELECT_ITEMS + 1));
let err = check_statement_allowed(&sql, true).unwrap_err();
assert!(
matches!(
err,
StatementRejected::Complexity {
bound: "select_items"
}
),
"got {err:?}"
);
}
#[test]
fn explain_of_refused_function_rejected() {
let err = check_statement_allowed("EXPLAIN SELECT repeat('x', 2)", true).unwrap_err();
assert!(matches!(err, StatementRejected::FunctionSurface));
}
#[test]
fn every_reason_key_is_a_registered_metric_label() {
use std::collections::BTreeSet;
let mut produced = BTreeSet::new();
for rejected in [
StatementRejected::DisallowedKind("x".to_owned()),
StatementRejected::ParseError("x".to_owned()),
StatementRejected::FunctionSurface,
] {
produced.insert(rejected.reason_key());
}
let union = |n: usize| {
(0..n)
.map(|_| "SELECT c FROM t".to_owned())
.collect::<Vec<_>>()
.join(" UNION ALL ")
};
let ctes = |n: usize| {
let names: Vec<String> = (0..n).map(|i| format!("c{i} AS (SELECT 1)")).collect();
format!("WITH {} SELECT 1", names.join(", "))
};
let over = [
format!(
"SELECT x IN ({})",
(0..MAX_EXPRESSION_NODES - 1)
.map(|v| v.to_string())
.collect::<Vec<_>>()
.join(", ")
),
nested_parens(MAX_EXPRESSION_DEPTH),
wide_select("abs(1)", MAX_FUNCTION_CALLS + 1),
wide_select("1", MAX_SELECT_ITEMS + 1),
ctes(MAX_QUERY_NODES),
"SELECT c || c || c || c || c FROM t".to_owned(),
wide_select("c", MAX_SELECT_COLUMN_FANOUT + 1),
union(MAX_STATEMENT_COLUMN_REFERENCES + 1),
];
for sql in over {
let err = check_statement_allowed(&sql, false).unwrap_err();
produced.insert(err.reason_key());
}
let registered: BTreeSet<&str> = crate::metrics::REFUSAL_REASONS.iter().copied().collect();
assert_eq!(produced, registered);
}
#[test]
fn fixed_statements_fit_the_surface_and_bounds() {
use polyc_query_model::statements::*;
let statements: &[(&str, &str)] = &[
("composite_trace", COMPOSITE_TRACE_SQL),
("composite_trace_memory", COMPOSITE_TRACE_MEMORY_SQL),
("composite_trace_routines", COMPOSITE_TRACE_ROUTINES_SQL),
("dashboard_attribution", DASHBOARD_ATTRIBUTION_SQL),
("dashboard_context", DASHBOARD_CONTEXT_SQL),
("dashboard_conversations", DASHBOARD_CONVERSATIONS_SQL),
("dashboard_spend", DASHBOARD_SPEND_SQL),
("fleet_usage", FLEET_USAGE_SQL),
("persona_directory_detail", PERSONA_DIRECTORY_DETAIL_SQL),
("persona_directory_index", PERSONA_DIRECTORY_INDEX_SQL),
("routine_fire", ROUTINE_FIRE_SQL),
("routine_fire_count", ROUTINE_FIRE_COUNT_SQL),
("routine_fire_last", ROUTINE_FIRE_LAST_SQL),
("routine_fire_outcome", ROUTINE_FIRE_OUTCOME_SQL),
("routine_fires", ROUTINE_FIRES_SQL),
("routine_fires_index_admin", ROUTINE_FIRES_INDEX_ADMIN_SQL),
("routine_fires_index_owner", ROUTINE_FIRES_INDEX_OWNER_SQL),
("routine_lifecycle", ROUTINE_LIFECYCLE_SQL),
("routine_overview", ROUTINE_OVERVIEW_SQL),
(
"routine_owner_active_grants",
ROUTINE_OWNER_ACTIVE_GRANTS_SQL,
),
(
"routine_owner_approval_aggregates",
ROUTINE_OWNER_APPROVAL_AGGREGATES_SQL,
),
(
"routine_owner_fire_dispatch",
ROUTINE_OWNER_FIRE_DISPATCH_SQL,
),
("routine_owner_refusals", ROUTINE_OWNER_REFUSALS_SQL),
("routine_owner_stopped_tool", ROUTINE_OWNER_STOPPED_TOOL_SQL),
];
for &(name, sql) in statements {
assert_eq!(
check_statement_allowed(sql, false)
.unwrap_or_else(|err| panic!("{name} must be allowed, got {err:?}")),
AllowedStatement::Query
);
}
}
}