use std::fmt::{self};
use common::fmt::Fmt;
use surrealdb_types::{SqlFormat, ToSql, write_sql};
use crate::statements::{
AccessStatement, KillStatement, LiveStatement, OptionStatement, ShowStatement, UseStatement,
};
use crate::{Expr, Literal, Param};
#[derive(Clone, Copy, Eq, PartialEq, Debug, Default)]
#[cfg_attr(feature = "arbitrary", derive(arbitrary::Arbitrary))]
pub enum ExplainFormat {
#[default]
Text,
Json,
}
#[derive(Debug, PartialEq, Clone)]
pub struct Ast {
pub expressions: Vec<TopLevelExpr>,
}
impl Ast {
pub fn single_expr(expr: Expr) -> Self {
Ast {
expressions: vec![TopLevelExpr::Expr(expr)],
}
}
pub fn num_statements(&self) -> usize {
self.expressions.len()
}
pub fn is_value_expression(&self) -> bool {
if self.expressions.len() != 1 {
return false;
}
let TopLevelExpr::Expr(expr) = &self.expressions[0] else {
return false;
};
is_value_expr(expr)
}
pub fn get_let_statements(&self) -> Vec<String> {
let mut let_var_names = Vec::new();
for expr in &self.expressions {
if let TopLevelExpr::Expr(Expr::Let(stmt)) = expr {
let_var_names.push(stmt.name.as_str().to_owned());
}
}
let_var_names
}
pub fn add_param(&mut self, name: String) {
self.expressions.push(TopLevelExpr::Expr(Expr::Param(Param::new(name))));
}
pub fn contains_cancel(&self) -> bool {
self.expressions.iter().any(|e| matches!(e, TopLevelExpr::Cancel))
}
pub fn is_sole_commit(&self) -> bool {
matches!(self.expressions.as_slice(), [TopLevelExpr::Commit])
}
pub fn is_sole_cancel(&self) -> bool {
matches!(self.expressions.as_slice(), [TopLevelExpr::Cancel])
}
pub fn into_execution_units(self) -> Vec<Ast> {
let mut units = Vec::new();
let mut block: Option<Vec<TopLevelExpr>> = None;
for expr in self.expressions {
match expr {
TopLevelExpr::Begin => {
block.get_or_insert_with(Vec::new).push(expr);
}
TopLevelExpr::Commit | TopLevelExpr::Cancel => match block.take() {
Some(mut statements) => {
statements.push(expr);
units.push(Ast {
expressions: statements,
});
}
None => units.push(Ast {
expressions: vec![expr],
}),
},
other => match block.as_mut() {
Some(statements) => statements.push(other),
None => units.push(Ast {
expressions: vec![other],
}),
},
}
}
if let Some(statements) = block {
units.push(Ast {
expressions: statements,
});
}
units
}
}
fn is_value_expr(expr: &Expr) -> bool {
match expr {
Expr::Param(_) | Expr::Constant(_) => true,
Expr::Literal(lit) => is_value_literal(lit),
_ => false,
}
}
fn is_value_literal(lit: &Literal) -> bool {
match lit {
Literal::Array(items) | Literal::Set(items) => items.iter().all(is_value_expr),
Literal::Object(entries) => entries.iter().all(|e| is_value_expr(&e.value)),
Literal::None
| Literal::Null
| Literal::UnboundedRange
| Literal::Bool(_)
| Literal::Float(_)
| Literal::Integer(_)
| Literal::Decimal(_)
| Literal::Duration(_)
| Literal::String(_)
| Literal::RecordId(_)
| Literal::Datetime(_)
| Literal::Uuid(_)
| Literal::Regex(_)
| Literal::Geometry(_)
| Literal::File(_)
| Literal::Bytes(_) => true,
}
}
impl ToSql for Ast {
fn fmt_sql(&self, f: &mut String, fmt: SqlFormat) {
write_sql!(
f,
fmt,
"{}",
&Fmt::one_line_separated(
self.expressions
.iter()
.map(|v| Fmt::new(v, |v, f, fmt| write_sql!(f, fmt, "{v};"))),
),
)
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
#[cfg_attr(feature = "arbitrary", derive(arbitrary::Arbitrary))]
pub enum TopLevelExpr {
Begin,
Cancel,
Commit,
Access(Box<AccessStatement>),
Kill(KillStatement),
Live(Box<LiveStatement>),
Option(OptionStatement),
Use(UseStatement),
Show(ShowStatement),
Expr(Expr),
}
impl fmt::Display for TopLevelExpr {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
if f.alternate() {
write!(f, "{}", self.to_sql_pretty())
} else {
write!(f, "{}", self.to_sql())
}
}
}
impl ToSql for TopLevelExpr {
fn fmt_sql(&self, f: &mut String, fmt: SqlFormat) {
match self {
TopLevelExpr::Begin => f.push_str("BEGIN"),
TopLevelExpr::Cancel => f.push_str("CANCEL"),
TopLevelExpr::Commit => f.push_str("COMMIT"),
TopLevelExpr::Access(s) => s.fmt_sql(f, fmt),
TopLevelExpr::Kill(s) => s.fmt_sql(f, fmt),
TopLevelExpr::Live(s) => s.fmt_sql(f, fmt),
TopLevelExpr::Option(s) => s.fmt_sql(f, fmt),
TopLevelExpr::Use(s) => s.fmt_sql(f, fmt),
TopLevelExpr::Show(s) => s.fmt_sql(f, fmt),
TopLevelExpr::Expr(e) => e.fmt_sql(f, fmt),
}
}
}