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 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 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 pub fn contains_cancel(&self) -> bool {
73 self.expressions.iter().any(|e| matches!(e, TopLevelExpr::Cancel))
74 }
75
76 pub fn is_sole_commit(&self) -> bool {
78 matches!(self.expressions.as_slice(), [TopLevelExpr::Commit])
79 }
80
81 pub fn is_sole_cancel(&self) -> bool {
83 matches!(self.expressions.as_slice(), [TopLevelExpr::Cancel])
84 }
85
86 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}