use serde::{Deserialize, Serialize};
use serde_json::Value;
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CteClause {
pub name: String,
pub columns: Vec<String>,
pub recursive: bool,
pub query: Box<Query>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct WithClause {
pub ctes: Vec<CteClause>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub enum SetOperator {
Union,
UnionAll,
Intersect,
Except,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct SetOperationClause {
pub op: SetOperator,
pub query: Box<Query>,
}
impl SetOperator {
pub fn is_all(&self) -> bool {
matches!(self, SetOperator::UnionAll)
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Query {
pub with_clause: Option<WithClause>,
pub let_clauses: Vec<LetClause>,
pub for_clauses: Vec<ForClause>,
pub join_clauses: Vec<JoinClause>,
pub filter_clauses: Vec<FilterClause>,
pub sort_clause: Option<SortClause>,
pub limit_clause: Option<LimitClause>,
pub return_clause: Option<ReturnClause>,
pub create_stream_clause: Option<CreateStreamClause>,
pub create_materialized_view_clause: Option<CreateMaterializedViewClause>,
pub refresh_materialized_view_clause: Option<RefreshMaterializedViewClause>,
pub window_clause: Option<WindowClause>,
pub body_clauses: Vec<BodyClause>,
#[serde(default)]
pub set_operations: Vec<SetOperationClause>,
}
impl Query {
pub fn has_mutations(&self) -> bool {
self.body_clauses.iter().any(|clause| {
matches!(
clause,
BodyClause::Insert(_)
| BodyClause::Update(_)
| BodyClause::Upsert(_)
| BodyClause::Remove(_)
)
}) || self.create_stream_clause.is_some()
|| self.create_materialized_view_clause.is_some()
|| self.refresh_materialized_view_clause.is_some()
|| self
.set_operations
.iter()
.any(|op| op.query.has_mutations())
|| self
.with_clause
.as_ref()
.is_some_and(|with| with.ctes.iter().any(|cte| cte.query.has_mutations()))
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub enum BodyClause {
For(ForClause),
Let(LetClause),
Filter(FilterClause),
Insert(InsertClause),
Update(UpdateClause),
Upsert(UpsertClause),
Remove(RemoveClause),
Join(JoinClause),
GraphTraversal(GraphTraversalClause),
ShortestPath(ShortestPathClause),
Collect(CollectClause),
Window(WindowClause),
Search(FilterClause),
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub enum EdgeDirection {
Outbound,
Inbound,
Any,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct GraphTraversalClause {
pub vertex_var: String,
pub edge_var: Option<String>,
pub direction: EdgeDirection,
pub start_vertex: Expression,
pub edge_collection: String,
pub min_depth: usize,
pub max_depth: usize,
#[serde(default)]
pub path_var: Option<String>,
#[serde(default)]
pub prune: Option<Expression>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ShortestPathClause {
pub vertex_var: String,
pub edge_var: Option<String>,
pub start_vertex: Expression,
pub end_vertex: Expression,
pub direction: EdgeDirection,
pub edge_collection: String,
#[serde(default)]
pub weight: Option<String>,
#[serde(default)]
pub path_var: Option<String>,
#[serde(default)]
pub mode: PathFindMode,
#[serde(default)]
pub k: Option<usize>,
#[serde(default)]
pub min_len: Option<usize>,
#[serde(default)]
pub max_len: Option<usize>,
#[serde(default)]
pub limit: Option<usize>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)]
pub enum PathFindMode {
#[default]
Shortest,
AllShortest,
KShortest,
KPaths,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CreateStreamClause {
pub name: String,
pub if_not_exists: bool,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CreateMaterializedViewClause {
pub name: String,
pub if_not_exists: bool,
pub query: Box<Query>,
#[serde(default)]
pub refresh_schedule: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct RefreshMaterializedViewClause {
pub name: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub enum WindowType {
Tumbling,
Sliding,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct WindowClause {
pub window_type: WindowType,
pub duration: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct LetClause {
pub variable: String,
pub expression: Expression,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ForClause {
pub variable: String,
pub collection: String,
pub source_variable: Option<String>,
pub source_expression: Option<Expression>,
#[serde(default)]
pub system_time: Option<Expression>,
#[serde(default)]
pub valid_time: Option<ValidTimeSpec>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub enum ValidTimeSpec {
AsOf(Expression),
Range { from: Expression, to: Expression },
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct FilterClause {
pub expression: Expression,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct InsertClause {
pub document: Expression,
pub collection: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct UpdateClause {
pub selector: Expression,
pub changes: Expression,
pub collection: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct UpsertClause {
pub search: Expression,
pub insert: Expression,
pub update: Expression,
pub collection: String,
pub replace: bool,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct RemoveClause {
pub selector: Expression,
pub collection: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub enum JoinType {
Inner,
Left,
Right,
FullOuter,
Asof,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub enum AsofStrategy {
Backward,
Forward,
Nearest,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct AsofSpec {
pub left_time: Expression,
pub right_time: Expression,
pub strategy: AsofStrategy,
pub tolerance: Option<Expression>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct JoinClause {
pub join_type: JoinType,
pub variable: String,
pub collection: String,
pub condition: Expression,
#[serde(default)]
pub asof: Option<AsofSpec>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CollectClause {
pub group_vars: Vec<(String, Expression)>,
pub into_var: Option<String>,
#[serde(default)]
pub keep_vars: Vec<String>,
pub count_var: Option<String>,
pub aggregates: Vec<AggregateExpr>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct AggregateExpr {
pub variable: String,
pub function: String,
pub argument: Option<Expression>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct SortClause {
pub fields: Vec<(Expression, bool)>, }
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct LimitClause {
pub offset: Expression,
pub count: Option<Expression>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ReturnClause {
pub expression: Expression,
#[serde(default)]
pub distinct: bool,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub enum TemplateStringPart {
Literal(String),
Expression(Box<Expression>),
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub enum Expression {
Variable(String),
BindVariable(String),
FieldAccess(Box<Expression>, String),
OptionalFieldAccess(Box<Expression>, String),
DynamicFieldAccess(Box<Expression>, Box<Expression>),
ArrayAccess(Box<Expression>, Box<Expression>),
ArraySpreadAccess(Box<Expression>, Option<String>),
Literal(Value),
BinaryOp {
left: Box<Expression>,
op: BinaryOperator,
right: Box<Expression>,
},
UnaryOp {
op: UnaryOperator,
operand: Box<Expression>,
},
Object(Vec<(String, Expression)>),
Array(Vec<Expression>),
Range(Box<Expression>, Box<Expression>),
FunctionCall { name: String, args: Vec<Expression> },
Subquery(Box<Query>),
Ternary {
condition: Box<Expression>,
true_expr: Box<Expression>,
false_expr: Box<Expression>,
},
Case {
operand: Option<Box<Expression>>,
when_clauses: Vec<(Expression, Expression)>,
else_clause: Option<Box<Expression>>,
},
Pipeline {
left: Box<Expression>,
right: Box<Expression>,
},
Lambda {
params: Vec<String>,
body: Box<Expression>,
},
WindowFunctionCall {
function: String,
arguments: Vec<Expression>,
over_clause: WindowSpec,
},
TemplateString { parts: Vec<TemplateStringPart> },
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct WindowSpec {
pub partition_by: Vec<Expression>,
pub order_by: Vec<(Expression, bool)>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub enum BinaryOperator {
Equal,
NotEqual,
LessThan,
LessThanOrEqual,
GreaterThan,
GreaterThanOrEqual,
Spaceship,
SemanticMatch,
In,
NotIn,
And,
Or,
Add,
Subtract,
Multiply,
Divide,
Modulus, Exponent,
Like,
NotLike,
RegEx,
NotRegEx,
FuzzyEqual,
BitwiseAnd,
BitwiseOr,
BitwiseXor,
LeftShift,
RightShift,
NullCoalesce,
LogicalOr,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub enum UnaryOperator {
Not,
Negate,
BitwiseNot,
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn test_expression_literal() {
let expr = Expression::Literal(json!(42));
assert_eq!(expr, Expression::Literal(json!(42)));
}
#[test]
fn test_expression_variable() {
let expr = Expression::Variable("doc".to_string());
if let Expression::Variable(name) = expr {
assert_eq!(name, "doc");
} else {
panic!("Expected Variable");
}
}
#[test]
fn test_expression_field_access() {
let expr = Expression::FieldAccess(
Box::new(Expression::Variable("doc".to_string())),
"name".to_string(),
);
if let Expression::FieldAccess(base, field) = expr {
assert_eq!(*base, Expression::Variable("doc".to_string()));
assert_eq!(field, "name");
} else {
panic!("Expected FieldAccess");
}
}
#[test]
fn test_expression_binary_op() {
let expr = Expression::BinaryOp {
left: Box::new(Expression::Variable("a".to_string())),
op: BinaryOperator::Add,
right: Box::new(Expression::Literal(json!(1))),
};
if let Expression::BinaryOp { left, op, right } = expr {
assert_eq!(*left, Expression::Variable("a".to_string()));
assert_eq!(op, BinaryOperator::Add);
assert_eq!(*right, Expression::Literal(json!(1)));
} else {
panic!("Expected BinaryOp");
}
}
#[test]
fn test_for_clause() {
let clause = ForClause {
variable: "doc".to_string(),
collection: "users".to_string(),
source_variable: None,
source_expression: None,
system_time: None,
valid_time: None,
};
assert_eq!(clause.variable, "doc");
assert_eq!(clause.collection, "users");
}
#[test]
fn test_filter_clause() {
let clause = FilterClause {
expression: Expression::Literal(json!(true)),
};
assert_eq!(clause.expression, Expression::Literal(json!(true)));
}
#[test]
fn test_limit_clause() {
let clause = LimitClause {
offset: Expression::Literal(json!(0)),
count: Some(Expression::Literal(json!(10))),
};
assert_eq!(clause.offset, Expression::Literal(json!(0)));
assert_eq!(clause.count, Some(Expression::Literal(json!(10))));
let unbounded = LimitClause {
offset: Expression::Literal(json!(5)),
count: None,
};
assert!(unbounded.count.is_none());
}
#[test]
fn test_sort_clause() {
let clause = SortClause {
fields: vec![(
Expression::FieldAccess(
Box::new(Expression::Variable("doc".to_string())),
"age".to_string(),
),
true,
)],
};
assert_eq!(clause.fields.len(), 1);
assert!(clause.fields[0].1); }
#[test]
fn test_let_clause() {
let clause = LetClause {
variable: "x".to_string(),
expression: Expression::Literal(json!(42)),
};
assert_eq!(clause.variable, "x");
}
#[test]
fn test_insert_clause() {
let clause = InsertClause {
document: Expression::Object(vec![]),
collection: "users".to_string(),
};
assert_eq!(clause.collection, "users");
}
#[test]
fn test_edge_direction() {
assert_ne!(EdgeDirection::Inbound, EdgeDirection::Outbound);
assert_ne!(EdgeDirection::Any, EdgeDirection::Inbound);
}
#[test]
fn test_binary_operators() {
assert_eq!(BinaryOperator::Equal.clone(), BinaryOperator::Equal);
assert_ne!(BinaryOperator::Equal, BinaryOperator::NotEqual);
assert_ne!(BinaryOperator::Add, BinaryOperator::Subtract);
}
#[test]
fn test_unary_operators() {
assert_eq!(UnaryOperator::Not.clone(), UnaryOperator::Not);
assert_ne!(UnaryOperator::Not, UnaryOperator::Negate);
}
#[test]
fn test_expression_clone() {
let expr = Expression::Variable("test".to_string());
let cloned = expr.clone();
assert_eq!(expr, cloned);
}
#[test]
fn test_query_default() {
let query = Query {
with_clause: None,
let_clauses: vec![],
for_clauses: vec![],
join_clauses: vec![],
filter_clauses: vec![],
sort_clause: None,
limit_clause: None,
return_clause: None,
create_stream_clause: None,
create_materialized_view_clause: None,
refresh_materialized_view_clause: None,
window_clause: None,
body_clauses: vec![],
set_operations: vec![],
};
assert!(query.for_clauses.is_empty());
assert!(query.return_clause.is_none());
}
#[test]
fn test_collect_clause() {
let clause = CollectClause {
group_vars: vec![(
"category".to_string(),
Expression::FieldAccess(
Box::new(Expression::Variable("doc".to_string())),
"cat".to_string(),
),
)],
into_var: Some("items".to_string()),
keep_vars: vec![],
count_var: Some("cnt".to_string()),
aggregates: vec![],
};
assert_eq!(clause.group_vars.len(), 1);
assert_eq!(clause.into_var, Some("items".to_string()));
assert_eq!(clause.count_var, Some("cnt".to_string()));
}
#[test]
fn test_aggregate_expr() {
let agg = AggregateExpr {
variable: "total".to_string(),
function: "SUM".to_string(),
argument: Some(Expression::FieldAccess(
Box::new(Expression::Variable("doc".to_string())),
"price".to_string(),
)),
};
assert_eq!(agg.variable, "total");
assert_eq!(agg.function, "SUM");
assert!(agg.argument.is_some());
}
#[test]
fn test_expression_array() {
let expr = Expression::Array(vec![
Expression::Literal(json!(1)),
Expression::Literal(json!(2)),
Expression::Literal(json!(3)),
]);
if let Expression::Array(items) = expr {
assert_eq!(items.len(), 3);
} else {
panic!("Expected Array");
}
}
#[test]
fn test_expression_object() {
let expr = Expression::Object(vec![
("name".to_string(), Expression::Literal(json!("test"))),
("value".to_string(), Expression::Literal(json!(42))),
]);
if let Expression::Object(fields) = expr {
assert_eq!(fields.len(), 2);
assert_eq!(fields[0].0, "name");
} else {
panic!("Expected Object");
}
}
#[test]
fn test_expression_range() {
let expr = Expression::Range(
Box::new(Expression::Literal(json!(1))),
Box::new(Expression::Literal(json!(5))),
);
if let Expression::Range(start, end) = expr {
assert_eq!(*start, Expression::Literal(json!(1)));
assert_eq!(*end, Expression::Literal(json!(5)));
} else {
panic!("Expected Range");
}
}
#[test]
fn test_expression_function_call() {
let expr = Expression::FunctionCall {
name: "LENGTH".to_string(),
args: vec![Expression::Variable("arr".to_string())],
};
if let Expression::FunctionCall { name, args } = expr {
assert_eq!(name, "LENGTH");
assert_eq!(args.len(), 1);
} else {
panic!("Expected FunctionCall");
}
}
#[test]
fn test_expression_ternary() {
let expr = Expression::Ternary {
condition: Box::new(Expression::Variable("flag".to_string())),
true_expr: Box::new(Expression::Literal(json!(1))),
false_expr: Box::new(Expression::Literal(json!(0))),
};
if let Expression::Ternary {
condition,
true_expr,
false_expr,
} = expr
{
assert_eq!(*condition, Expression::Variable("flag".to_string()));
assert_eq!(*true_expr, Expression::Literal(json!(1)));
assert_eq!(*false_expr, Expression::Literal(json!(0)));
} else {
panic!("Expected Ternary");
}
}
}