Skip to main content

surrealdb_sql/
ast.rs

1use std::fmt::{self};
2
3use common::fmt::Fmt;
4use surrealdb_types::{SqlFormat, ToSql, write_sql};
5
6use crate::statements::{
7	AccessStatement, KillStatement, LiveStatement, OptionStatement, ShowStatement, UseStatement,
8};
9use crate::{Expr, Literal, Param};
10
11#[derive(Clone, Copy, Eq, PartialEq, Debug, Default)]
12#[cfg_attr(feature = "arbitrary", derive(arbitrary::Arbitrary))]
13pub enum ExplainFormat {
14	#[default]
15	Text,
16	Json,
17}
18
19#[derive(Debug, PartialEq, Clone)]
20pub struct Ast {
21	pub expressions: Vec<TopLevelExpr>,
22}
23
24impl Ast {
25	/// Creates an ast with a signle expression
26	pub fn single_expr(expr: Expr) -> Self {
27		Ast {
28			expressions: vec![TopLevelExpr::Expr(expr)],
29		}
30	}
31
32	pub fn num_statements(&self) -> usize {
33		self.expressions.len()
34	}
35
36	/// Returns `true` if this AST is exactly one top-level expression and that
37	/// expression describes inert data only: a literal (including objects,
38	/// arrays and sets of inert data), a `$param` reference, or a built-in
39	/// constant.
40	///
41	/// Use this to validate request payloads that are intended to supply
42	/// values rather than to execute SurrealQL (e.g. the body of a REST
43	/// `/key` request that is bound to `$data`). Function calls, idioms,
44	/// statements (CREATE/UPDATE/...), blocks, binary expressions and other
45	/// executable forms are all rejected, including when nested inside
46	/// object or array literals.
47	pub fn is_value_expression(&self) -> bool {
48		if self.expressions.len() != 1 {
49			return false;
50		}
51		let TopLevelExpr::Expr(expr) = &self.expressions[0] else {
52			return false;
53		};
54		is_value_expr(expr)
55	}
56
57	pub fn get_let_statements(&self) -> Vec<String> {
58		let mut let_var_names = Vec::new();
59		for expr in &self.expressions {
60			if let TopLevelExpr::Expr(Expr::Let(stmt)) = expr {
61				let_var_names.push(stmt.name.as_str().to_owned());
62			}
63		}
64		let_var_names
65	}
66
67	pub fn add_param(&mut self, name: String) {
68		self.expressions.push(TopLevelExpr::Expr(Expr::Param(Param::new(name))));
69	}
70
71	/// Whether a top-level `CANCEL` (transaction rollback) appears in this AST.
72	pub fn contains_cancel(&self) -> bool {
73		self.expressions.iter().any(|e| matches!(e, TopLevelExpr::Cancel))
74	}
75
76	/// Whether this AST is exactly a single top-level `COMMIT`.
77	pub fn is_sole_commit(&self) -> bool {
78		matches!(self.expressions.as_slice(), [TopLevelExpr::Commit])
79	}
80
81	/// Whether this AST is exactly a single top-level `CANCEL`.
82	pub fn is_sole_cancel(&self) -> bool {
83		matches!(self.expressions.as_slice(), [TopLevelExpr::Cancel])
84	}
85
86	/// Split this query into sequentially executable units: each
87	/// non-transactional top-level statement becomes a unit of its own, while
88	/// the statements of a `BEGIN`..`COMMIT`/`CANCEL` block stay together as
89	/// one unit. A `BEGIN` without a matching `COMMIT`/`CANCEL` keeps the
90	/// remainder of the query in its unit.
91	pub fn into_execution_units(self) -> Vec<Ast> {
92		let mut units = Vec::new();
93		let mut block: Option<Vec<TopLevelExpr>> = None;
94		for expr in self.expressions {
95			match expr {
96				TopLevelExpr::Begin => {
97					block.get_or_insert_with(Vec::new).push(expr);
98				}
99				TopLevelExpr::Commit | TopLevelExpr::Cancel => match block.take() {
100					Some(mut statements) => {
101						statements.push(expr);
102						units.push(Ast {
103							expressions: statements,
104						});
105					}
106					None => units.push(Ast {
107						expressions: vec![expr],
108					}),
109				},
110				other => match block.as_mut() {
111					Some(statements) => statements.push(other),
112					None => units.push(Ast {
113						expressions: vec![other],
114					}),
115				},
116			}
117		}
118		if let Some(statements) = block {
119			units.push(Ast {
120				expressions: statements,
121			});
122		}
123		units
124	}
125}
126
127fn is_value_expr(expr: &Expr) -> bool {
128	match expr {
129		Expr::Param(_) | Expr::Constant(_) => true,
130		Expr::Literal(lit) => is_value_literal(lit),
131		_ => false,
132	}
133}
134
135fn is_value_literal(lit: &Literal) -> bool {
136	match lit {
137		Literal::Array(items) | Literal::Set(items) => items.iter().all(is_value_expr),
138		Literal::Object(entries) => entries.iter().all(|e| is_value_expr(&e.value)),
139		Literal::None
140		| Literal::Null
141		| Literal::UnboundedRange
142		| Literal::Bool(_)
143		| Literal::Float(_)
144		| Literal::Integer(_)
145		| Literal::Decimal(_)
146		| Literal::Duration(_)
147		| Literal::String(_)
148		| Literal::RecordId(_)
149		| Literal::Datetime(_)
150		| Literal::Uuid(_)
151		| Literal::Regex(_)
152		| Literal::Geometry(_)
153		| Literal::File(_)
154		| Literal::Bytes(_) => true,
155	}
156}
157
158impl ToSql for Ast {
159	fn fmt_sql(&self, f: &mut String, fmt: SqlFormat) {
160		write_sql!(
161			f,
162			fmt,
163			"{}",
164			&Fmt::one_line_separated(
165				self.expressions
166					.iter()
167					.map(|v| Fmt::new(v, |v, f, fmt| write_sql!(f, fmt, "{v};"))),
168			),
169		)
170	}
171}
172
173#[derive(Clone, Debug, Eq, PartialEq)]
174#[cfg_attr(feature = "arbitrary", derive(arbitrary::Arbitrary))]
175pub enum TopLevelExpr {
176	Begin,
177	Cancel,
178	Commit,
179	Access(Box<AccessStatement>),
180	Kill(KillStatement),
181	Live(Box<LiveStatement>),
182	Option(OptionStatement),
183	Use(UseStatement),
184	Show(ShowStatement),
185	Expr(Expr),
186}
187
188impl fmt::Display for TopLevelExpr {
189	fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
190		if f.alternate() {
191			write!(f, "{}", self.to_sql_pretty())
192		} else {
193			write!(f, "{}", self.to_sql())
194		}
195	}
196}
197
198impl ToSql for TopLevelExpr {
199	fn fmt_sql(&self, f: &mut String, fmt: SqlFormat) {
200		match self {
201			TopLevelExpr::Begin => f.push_str("BEGIN"),
202			TopLevelExpr::Cancel => f.push_str("CANCEL"),
203			TopLevelExpr::Commit => f.push_str("COMMIT"),
204			TopLevelExpr::Access(s) => s.fmt_sql(f, fmt),
205			TopLevelExpr::Kill(s) => s.fmt_sql(f, fmt),
206			TopLevelExpr::Live(s) => s.fmt_sql(f, fmt),
207			TopLevelExpr::Option(s) => s.fmt_sql(f, fmt),
208			TopLevelExpr::Use(s) => s.fmt_sql(f, fmt),
209			TopLevelExpr::Show(s) => s.fmt_sql(f, fmt),
210			TopLevelExpr::Expr(e) => e.fmt_sql(f, fmt),
211		}
212	}
213}