Skip to main content

surrealdb_sql/
expression.rs

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	// TODO: Factor out the call from the function expression.
46	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	// NOTE: Changes to this function also likely require changes to
135	// `surrealdb_expr::expr::Expr::needs_parentheses`.
136	/// Returns if this expression needs to be parenthesized when inside another expression.
137	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	/// Returns true if there is a `NONE` or `NULL` value in the left most spot when formatting.
194	/// returns true for `NONE + 1`, `NULL()`, `NONE`, `NULL..` etc.
195	///
196	/// Required for proper formatting when `NONE` can conflict with a clause.
197	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	// Determine the operator first before moving the values
273	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					// We need to avoid `--` from showing up so we need to cover if the expression
319					// has a left minus
320					|| *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				// A left-associative operator's default nesting already puts a
353				// same-power child on the left (`a - b + c` parses as
354				// `(a - b) + c`), so equal power there needs no parentheses.
355				// `Power` is right-associative, so the opposite holds: a
356				// same-power left child only exists because it was explicitly
357				// parenthesised (`(a ** b) ** c`), and dropping the parens
358				// would reparse it right-associatively instead.
359				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				// Mirror image of the left operand: a same-power right child
388				// under a left-associative operator only exists because it
389				// was explicitly parenthesised (`a - (b + c)`, `a / (b * c)`),
390				// since default left-associative nesting never produces one.
391				// Dropping the parentheses would reparse it left-associatively
392				// instead, changing the value for any non-associative pairing
393				// within the group (`-` vs `+`, `/` vs `*`). `Power` is
394				// right-associative, so its default nesting already matches
395				// and needs no parentheses here.
396				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}