1use surrealdb_types::{
2 Number as PublicNumber, RecordId as PublicRecordId, SqlFormat, ToSql, Value as PublicValue,
3 write_sql,
4};
5
6use crate::ast::ExplainFormat;
7use crate::literal::ObjectEntry;
8use crate::lookup::LookupKind;
9use crate::operator::BindingPower;
10use crate::statements::{
11 AlterStatement, CreateStatement, DefineStatement, DeleteStatement, ForeachStatement,
12 IfelseStatement, InfoStatement, InsertStatement, OutputStatement, RebuildStatement,
13 RelateStatement, RemoveStatement, SelectStatement, SetStatement, SleepStatement,
14 UpdateStatement, UpsertStatement,
15};
16use crate::{
17 BinaryOperator, Block, Closure, Constant, CoverStmts, Dir, FunctionCall, Idiom, Literal, Mock,
18 Param, Part, PostfixOperator, PrefixOperator, RecordIdKeyLit, RecordIdLit, TableName,
19};
20
21#[derive(Clone, Debug, Eq, PartialEq)]
22#[cfg_attr(feature = "arbitrary", derive(arbitrary::Arbitrary))]
23pub enum Expr {
24 Literal(Literal),
25
26 Param(Param),
27 Idiom(Idiom),
28 Table(TableName),
29 Mock(Mock),
30 Block(Box<Block>),
31 Constant(Constant),
32 Prefix {
33 op: PrefixOperator,
34 expr: Box<Expr>,
35 },
36 Postfix {
37 expr: Box<Expr>,
38 op: PostfixOperator,
39 },
40 Binary {
41 left: Box<Expr>,
42 op: BinaryOperator,
43 right: Box<Expr>,
44 },
45 FunctionCall(Box<FunctionCall>),
47 Closure(Box<Closure>),
48
49 Break,
50 Continue,
51 Throw(Box<Expr>),
52
53 Return(Box<OutputStatement>),
54 IfElse(Box<IfelseStatement>),
55 Select(Box<SelectStatement>),
56 Create(Box<CreateStatement>),
57 Update(Box<UpdateStatement>),
58 Delete(Box<DeleteStatement>),
59 Relate(Box<RelateStatement>),
60 Insert(Box<InsertStatement>),
61 Define(Box<DefineStatement>),
62 Remove(Box<RemoveStatement>),
63 Rebuild(Box<RebuildStatement>),
64 Upsert(Box<UpsertStatement>),
65 Alter(Box<AlterStatement>),
66 Info(Box<InfoStatement>),
67 Foreach(Box<ForeachStatement>),
68 Let(Box<SetStatement>),
69 Sleep(Box<SleepStatement>),
70 Explain {
71 format: ExplainFormat,
72 analyze: bool,
73 statement: Box<Expr>,
74 },
75}
76
77impl Expr {
78 pub fn to_idiom(&self) -> Idiom {
79 match self {
80 Expr::Idiom(i) => i.simplify(),
81 Expr::Param(i) => Idiom::field(i.clone().into_strand()),
82 Expr::FunctionCall(x) => x.receiver.to_idiom(),
83 Expr::Literal(l) => match l {
84 Literal::String(s) => Idiom::field(s.clone()),
85 Literal::Datetime(d) => Idiom::field(surrealdb_types::fmt_datetime_sql(*d)),
86 x => Idiom::field(x.to_sql()),
87 },
88 x => Idiom::field(x.to_sql()),
89 }
90 }
91
92 pub fn from_public_value(value: PublicValue) -> Self {
93 match value {
94 PublicValue::None => Expr::Literal(Literal::None),
95 PublicValue::Null => Expr::Literal(Literal::Null),
96 PublicValue::Bool(x) => Expr::Literal(Literal::Bool(x)),
97 PublicValue::Number(PublicNumber::Float(x)) => Expr::Literal(Literal::Float(x)),
98 PublicValue::Number(PublicNumber::Int(x)) => Expr::Literal(Literal::Integer(x)),
99 PublicValue::Number(PublicNumber::Decimal(x)) => Expr::Literal(Literal::Decimal(x)),
100 PublicValue::String(x) => Expr::Literal(Literal::String(x.into())),
101 PublicValue::Bytes(x) => Expr::Literal(Literal::Bytes(x.into_inner())),
102 PublicValue::Regex(x) => Expr::Literal(Literal::Regex(x.into_inner())),
103 PublicValue::Table(x) => Expr::Table(TableName::new(x.into_string())),
104 PublicValue::RecordId(PublicRecordId {
105 table,
106 key,
107 }) => Expr::Literal(Literal::RecordId(RecordIdLit {
108 table: TableName::new(table.into_string()),
109 key: RecordIdKeyLit::from_record_id_key(key),
110 })),
111 PublicValue::Array(x) => {
112 Expr::Literal(Literal::Array(x.into_iter().map(Expr::from_public_value).collect()))
113 }
114 PublicValue::Set(x) => {
115 Expr::Literal(Literal::Set(x.into_iter().map(Expr::from_public_value).collect()))
116 }
117 PublicValue::Object(x) => Expr::Literal(Literal::Object(
118 x.into_iter()
119 .map(|(k, v)| ObjectEntry {
120 key: k.into(),
121 value: Expr::from_public_value(v),
122 })
123 .collect(),
124 )),
125 PublicValue::Duration(x) => Expr::Literal(Literal::Duration(x.into_inner())),
126 PublicValue::Datetime(x) => Expr::Literal(Literal::Datetime(x.into_inner())),
127 PublicValue::Uuid(x) => Expr::Literal(Literal::Uuid(x.into_inner())),
128 PublicValue::Geometry(x) => Expr::Literal(Literal::Geometry(geo::Geometry::from(x))),
129 PublicValue::File(x) => Expr::Literal(Literal::File(x.into())),
130 PublicValue::Range(x) => convert_public_range_to_literal(*x),
131 }
132 }
133
134 pub fn needs_parentheses(&self) -> bool {
138 match self {
139 Expr::Literal(Literal::UnboundedRange | Literal::RecordId(_))
140 | Expr::Closure(_)
141 | Expr::Break
142 | Expr::Continue
143 | Expr::Throw(_)
144 | Expr::Return(_)
145 | Expr::IfElse(_)
146 | Expr::Select(_)
147 | Expr::Create(_)
148 | Expr::Update(_)
149 | Expr::Delete(_)
150 | Expr::Relate(_)
151 | Expr::Insert(_)
152 | Expr::Define(_)
153 | Expr::Remove(_)
154 | Expr::Rebuild(_)
155 | Expr::Upsert(_)
156 | Expr::Alter(_)
157 | Expr::Info(_)
158 | Expr::Foreach(_)
159 | Expr::Let(_)
160 | Expr::Sleep(_)
161 | Expr::Explain {
162 ..
163 } => true,
164
165 Expr::Postfix {
166 op,
167 ..
168 } => matches!(
169 op,
170 PostfixOperator::Range
171 | PostfixOperator::RangeSkip
172 | PostfixOperator::MethodCall(_, _)
173 | PostfixOperator::Call(_)
174 ),
175
176 Expr::Literal(_)
177 | Expr::Param(_)
178 | Expr::Idiom(_)
179 | Expr::Table(_)
180 | Expr::Mock(_)
181 | Expr::Block(_)
182 | Expr::Constant(_)
183 | Expr::Prefix {
184 ..
185 }
186 | Expr::Binary {
187 ..
188 }
189 | Expr::FunctionCall(_) => false,
190 }
191 }
192
193 pub fn has_left_none_null(&self) -> bool {
198 match self {
199 Expr::Literal(Literal::None) | Expr::Literal(Literal::Null) => true,
200 Expr::Binary {
201 left: expr,
202 ..
203 }
204 | Expr::Postfix {
205 expr,
206 ..
207 } => expr.has_left_none_null(),
208 Expr::Idiom(x) => {
209 if let Some(Part::Start(x)) = x.0.first() {
210 x.has_left_none_null()
211 } else {
212 false
213 }
214 }
215 _ => false,
216 }
217 }
218
219 pub fn has_left_minus(&self) -> bool {
220 match self {
221 Expr::Prefix {
222 op: PrefixOperator::Negate,
223 ..
224 } => true,
225 Expr::Postfix {
226 expr,
227 ..
228 }
229 | Expr::Binary {
230 left: expr,
231 ..
232 } => expr.has_left_minus(),
233 Expr::Literal(Literal::Integer(x)) => x.is_negative(),
234 Expr::Literal(Literal::Float(x)) => x.is_sign_negative(),
235 Expr::Literal(Literal::Decimal(x)) => x.is_sign_negative(),
236 Expr::Idiom(x) => {
237 if let Some(x) = x.0.first()
238 && let Part::Graph(lookup) = x
239 && let LookupKind::Graph(Dir::Out) = lookup.kind
240 {
241 return true;
242 }
243 false
244 }
245 _ => false,
246 }
247 }
248
249 pub fn has_left_idiom(&self) -> bool {
250 match self {
251 Expr::Idiom(_) => true,
252
253 Expr::Postfix {
254 expr,
255 ..
256 }
257 | Expr::Binary {
258 left: expr,
259 ..
260 } => expr.has_left_idiom(),
261 _ => false,
262 }
263 }
264}
265
266fn convert_public_range_to_literal(range: surrealdb_types::Range) -> Expr {
267 use crate::literal::Literal;
268 use crate::operator::BinaryOperator;
269
270 let range = range.into_inner();
271
272 let op = match (&range.0, &range.1) {
274 (std::ops::Bound::Included(_), std::ops::Bound::Included(_)) => {
275 BinaryOperator::RangeInclusive
276 }
277 _ => BinaryOperator::Range,
278 };
279
280 let start_expr = match range.0 {
281 std::ops::Bound::Included(v) => Expr::from_public_value(v),
282 std::ops::Bound::Excluded(v) => Expr::from_public_value(v),
283 std::ops::Bound::Unbounded => Expr::Literal(Literal::None),
284 };
285
286 let end_expr = match range.1 {
287 std::ops::Bound::Included(v) => Expr::from_public_value(v),
288 std::ops::Bound::Excluded(v) => Expr::from_public_value(v),
289 std::ops::Bound::Unbounded => Expr::Literal(Literal::None),
290 };
291
292 Expr::Binary {
293 left: Box::new(start_expr),
294 op,
295 right: Box::new(end_expr),
296 }
297}
298
299impl ToSql for Expr {
300 fn fmt_sql(&self, f: &mut String, fmt: SqlFormat) {
301 match self {
302 Expr::Literal(literal) => literal.fmt_sql(f, fmt),
303 Expr::Param(param) => param.fmt_sql(f, fmt),
304 Expr::Idiom(idiom) => idiom.fmt_sql(f, fmt),
305 Expr::Table(ident) => write_sql!(f, fmt, "{ident}"),
306 Expr::Mock(mock) => mock.fmt_sql(f, fmt),
307 Expr::Block(block) => block.fmt_sql(f, fmt),
308 Expr::Constant(constant) => constant.fmt_sql(f, fmt),
309 Expr::Prefix {
310 op,
311 expr,
312 } => {
313 let expr_bp = BindingPower::for_expr(expr);
314 let op_bp = BindingPower::for_prefix_operator(op);
315 if expr.needs_parentheses()
316 || expr_bp < op_bp
317 || expr_bp == op_bp && matches!(expr_bp, BindingPower::Range)
318 || *op == PrefixOperator::Negate && expr.has_left_minus()
321 {
322 write_sql!(f, fmt, "{op}({expr})");
323 } else {
324 write_sql!(f, fmt, "{op}{expr}");
325 }
326 }
327 Expr::Postfix {
328 expr,
329 op,
330 } => {
331 let expr_bp = BindingPower::for_expr(expr);
332 let op_bp = BindingPower::for_postfix_operator(op);
333 if expr.needs_parentheses()
334 || expr_bp < op_bp
335 || expr_bp == op_bp && matches!(expr_bp, BindingPower::Range)
336 || matches!(op, PostfixOperator::Call(_))
337 {
338 write_sql!(f, fmt, "({expr}){op}");
339 } else {
340 write_sql!(f, fmt, "{expr}{op}");
341 }
342 }
343 Expr::Binary {
344 left,
345 op,
346 right,
347 } => {
348 let op_bp = BindingPower::for_binary_operator(op);
349 let left_bp = BindingPower::for_expr(left);
350 let right_bp = BindingPower::for_expr(right);
351
352 if left.needs_parentheses()
360 || left_bp < op_bp
361 || left_bp == op_bp
362 && matches!(
363 left_bp,
364 BindingPower::Range
365 | BindingPower::Relation | BindingPower::Equality
366 | BindingPower::Power
367 ) {
368 write_sql!(f, fmt, "({left})");
369 } else {
370 write_sql!(f, fmt, "{left}");
371 }
372
373 if matches!(
374 op,
375 BinaryOperator::Range
376 | BinaryOperator::RangeSkip
377 | BinaryOperator::RangeInclusive
378 | BinaryOperator::RangeSkipInclusive
379 ) {
380 op.fmt_sql(f, fmt);
381 } else {
382 f.push(' ');
383 op.fmt_sql(f, fmt);
384 f.push(' ');
385 }
386
387 if right.needs_parentheses()
397 || right_bp < op_bp
398 || right_bp == op_bp
399 && matches!(
400 right_bp,
401 BindingPower::Range
402 | BindingPower::Relation | BindingPower::Equality
403 | BindingPower::AddSub | BindingPower::MulDiv
404 ) {
405 write_sql!(f, fmt, "({right})");
406 } else {
407 write_sql!(f, fmt, "{right}");
408 }
409 }
410 Expr::FunctionCall(function_call) => function_call.fmt_sql(f, fmt),
411 Expr::Closure(closure) => closure.fmt_sql(f, fmt),
412 Expr::Break => f.push_str("BREAK"),
413 Expr::Continue => f.push_str("CONTINUE"),
414 Expr::Return(x) => x.fmt_sql(f, fmt),
415 Expr::Throw(expr) => write_sql!(f, fmt, "THROW {}", CoverStmts(expr.as_ref())),
416 Expr::IfElse(s) => s.fmt_sql(f, fmt),
417 Expr::Select(s) => s.fmt_sql(f, fmt),
418 Expr::Create(s) => s.fmt_sql(f, fmt),
419 Expr::Update(s) => s.fmt_sql(f, fmt),
420 Expr::Delete(s) => s.fmt_sql(f, fmt),
421 Expr::Relate(s) => s.fmt_sql(f, fmt),
422 Expr::Insert(s) => s.fmt_sql(f, fmt),
423 Expr::Define(s) => s.fmt_sql(f, fmt),
424 Expr::Remove(s) => s.fmt_sql(f, fmt),
425 Expr::Rebuild(s) => s.fmt_sql(f, fmt),
426 Expr::Upsert(s) => s.fmt_sql(f, fmt),
427 Expr::Alter(s) => s.fmt_sql(f, fmt),
428 Expr::Info(s) => s.fmt_sql(f, fmt),
429 Expr::Foreach(s) => s.fmt_sql(f, fmt),
430 Expr::Let(s) => s.fmt_sql(f, fmt),
431 Expr::Sleep(s) => s.fmt_sql(f, fmt),
432 Expr::Explain {
433 format: explain_format,
434 analyze,
435 statement,
436 } => {
437 f.push_str("EXPLAIN");
438 if *analyze {
439 f.push_str(" ANALYZE");
440 }
441 match explain_format {
442 ExplainFormat::Text => f.push_str(" FORMAT TEXT"),
443 ExplainFormat::Json => f.push_str(" FORMAT JSON"),
444 }
445 f.push(' ');
446 statement.fmt_sql(f, fmt);
447 }
448 }
449 }
450}