use crate::ast::{BinaryOpKind, Expr, Statement};
use std::collections::HashMap;
use pest::Parser;
use pest_derive::Parser;
#[derive(Parser)]
#[grammar = "src/vex.grammar"]
pub struct DSLParser;
pub fn parse(src: &str) -> Result<Vec<Statement>, Box<pest::error::Error<Rule>>> {
let pairs = DSLParser::parse(Rule::program, src)?;
Ok(pairs.flat_map(|pair| parse_pair(pair)).collect())
}
fn parse_pair(pair: pest::iterators::Pair<'_, Rule>) -> Vec<Statement> {
match pair.as_rule() {
Rule::program => parse_program(pair),
Rule::stmt_list => parse_statements(pair),
Rule::stmt => vec![parse_statement(pair)],
Rule::EOI => Vec::new(),
other => unreachable!("Unexpected outer rule: {:?}", other),
}
}
fn parse_program(pair: pest::iterators::Pair<'_, Rule>) -> Vec<Statement> {
let stmt_list_pair = pair.into_inner().next().unwrap();
parse_statements(stmt_list_pair)
}
fn parse_statements(pair: pest::iterators::Pair<'_, Rule>) -> Vec<Statement> {
pair.into_inner()
.map(|stmt_pair| parse_statement(stmt_pair))
.collect()
}
fn parse_statement(pair: pest::iterators::Pair<'_, Rule>) -> Statement {
let inner_pair = pair.into_inner().next().unwrap();
match inner_pair.as_rule() {
Rule::let_stmt => {
let mut inner_rules = inner_pair.into_inner();
let name = inner_rules.next().unwrap().as_str().to_string();
let value = parse_expr(inner_rules.next().unwrap());
Statement::LetExpression { name, value }
}
Rule::call_stmt => {
let mut inner_rules = inner_pair.into_inner();
let function = inner_rules.next().unwrap().as_str().to_string();
let args = inner_rules.map(parse_expr).collect();
Statement::Call { function, args }
}
Rule::print_stmt => {
let var = inner_pair.into_inner().next().unwrap().as_str().to_string();
Statement::Print { name: var }
}
Rule::fn_def => {
let mut parts = inner_pair.into_inner();
let name = parts.next().unwrap().as_str().to_owned();
let (params, body_pair) = if let Some(p) = parts.next() {
if p.as_rule() == Rule::param_list {
let params = p.into_inner().map(|i| i.as_str().to_owned()).collect();
let body = parts.next().unwrap();
(params, body)
} else {
(Vec::new(), p)
}
} else {
panic!("fn_def without block")
};
let body = body_pair.into_inner().flat_map(parse_pair).collect();
Statement::FunctionDef { name, params, body }
}
_ => {
let expr = parse_expr(inner_pair);
Statement::Expression { value: expr }
}
}
}
fn parse_expr(mut pair: pest::iterators::Pair<Rule>) -> Expr {
while matches!(pair.as_rule(), Rule::expr | Rule::method_arg) {
pair = pair.into_inner().next().unwrap();
}
if pair.as_rule() == Rule::atom {
pair = pair.into_inner().next().unwrap();
}
match pair.as_rule() {
Rule::sum => {
let mut inner = pair.into_inner();
let mut lhs = parse_expr(inner.next().unwrap());
while let Some(op) = inner.next() {
let rhs = parse_expr(inner.next().unwrap());
lhs = Expr::BinaryOp {
lhs: Box::new(lhs),
op: BinaryOpKind::try_from(op.as_str()).unwrap(),
rhs: Box::new(rhs),
};
}
lhs
}
Rule::product => {
let mut inner = pair.into_inner();
let mut lhs = parse_expr(inner.next().unwrap());
while let Some(op) = inner.next() {
let rhs = parse_expr(inner.next().unwrap());
lhs = Expr::BinaryOp {
lhs: Box::new(lhs),
op: BinaryOpKind::try_from(op.as_str()).unwrap(),
rhs: Box::new(rhs),
};
}
lhs
}
Rule::postfix_expr => {
let mut inner = pair.into_inner();
let base = parse_atom(inner.next().unwrap());
inner.fold(base, |expr, pfx| {
let mut parts = pfx.clone().into_inner();
match pfx.as_str().chars().next().unwrap() {
'.' => {
let ident = parts.next().unwrap().as_str().to_string();
if parts.peek().is_some() {
let arg_expr = parse_expr(parts.next().unwrap());
Expr::MethodCall {
target: Box::new(expr),
method: ident,
arg: Box::new(arg_expr),
}
} else {
Expr::Field {
target: Box::new(expr),
name: ident,
}
}
}
'[' => {
let idx_expr = parse_expr(parts.next().unwrap());
Expr::Index {
target: Box::new(expr),
index: Box::new(idx_expr),
}
}
_ => unreachable!(),
}
})
}
Rule::call_expr => {
let mut inner = pair.into_inner();
let func = inner.next().unwrap().as_str().to_string();
let args = inner
.next()
.map(|alist| alist.into_inner().map(parse_expr).collect())
.unwrap_or_default();
Expr::Call {
function: func,
args,
}
}
Rule::lambda_expr => {
let mut inner = pair.clone().into_inner();
let param = inner.next().unwrap().as_str().to_string();
let body_pair = inner.next().unwrap();
let body: Vec<Statement> = match body_pair.as_rule() {
Rule::block => {
let mut inner = pair.into_inner();
let only = inner.next();
if let Some(stmt_list) = only {
if stmt_list.as_rule() == Rule::key_val_pairs {
let hashmap_expr = Expr::HashMap(
stmt_list
.into_inner()
.map(|kv| {
let mut kv_inner = kv.into_inner();
let key = kv_inner.next().unwrap().as_str().to_string();
let value = parse_expr(kv_inner.next().unwrap());
(key, value)
})
.collect(),
);
return Expr::Block(vec![Statement::Expression {
value: hashmap_expr,
}]);
}
}
if inner.clone().count() == 1 {
let first = inner.next().unwrap();
if first.as_rule() == Rule::expr {
vec![Statement::Expression {
value: parse_expr(first),
}]
} else {
first.into_inner().flat_map(parse_pair).collect()
}
} else {
inner.flat_map(parse_pair).collect()
}
}
_ => vec![Statement::Expression {
value: parse_expr(body_pair),
}],
};
Expr::Lambda { param, body }
}
Rule::paren_expr => {
let inner = pair.into_inner().next().unwrap();
parse_expr(inner)
}
_ => parse_atom(pair),
}
}
fn parse_atom(pair: pest::iterators::Pair<'_, Rule>) -> Expr {
let pair = if pair.as_rule() == Rule::atom {
pair.into_inner().next().unwrap()
} else {
pair
};
match pair.as_rule() {
Rule::boolean => match pair.as_str() {
"true" => Expr::Boolean(true),
"false" => Expr::Boolean(false),
_ => unreachable!("Vex is not built for a quantum computer bro"),
},
Rule::string => {
let raw = pair.as_str();
let stripped = &raw[1..raw.len() - 1];
Expr::String(stripped.to_string())
}
Rule::number => {
let raw = pair.as_str();
let value: f64 = raw.parse().unwrap();
if raw.contains('.') {
Expr::Float(value)
} else {
Expr::Integer(value as i32)
}
}
Rule::array => {
let values = pair
.into_inner()
.next()
.map(|list| list.into_inner().map(parse_expr).collect())
.unwrap_or_default();
Expr::Array(values)
}
Rule::hashmap => {
let mut hashmap = HashMap::new();
let kv_pairs = pair.into_inner().next().unwrap();
assert_eq!(kv_pairs.as_rule(), Rule::key_val_pairs);
for kv_pair in kv_pairs.into_inner() {
assert_eq!(kv_pair.as_rule(), Rule::key_val_pair);
let mut inner = kv_pair.into_inner();
let key = inner.next().unwrap().as_str().trim_matches('"').to_string();
let value = parse_expr(inner.next().unwrap());
hashmap.insert(key, value);
}
Expr::HashMap(hashmap)
}
Rule::ident => Expr::Identifier(pair.as_str().to_string()),
Rule::block => {
let stmts = pair.into_inner().flat_map(parse_pair).collect::<Vec<_>>();
if stmts.len() == 1 {
if let Statement::Expression {
value: Expr::HashMap(_),
} = &stmts[0]
{
return Expr::Block(stmts);
}
}
Expr::Block(stmts)
}
Rule::lambda_expr => {
let mut inner = pair.into_inner();
let param = inner.next().unwrap().as_str().to_string();
let body_pair = inner.next().unwrap();
let body: Vec<Statement> = match body_pair.as_rule() {
Rule::block => {
let block_inner = body_pair.into_inner();
let only = block_inner.clone().collect::<Vec<_>>();
if only.len() == 1 {
let expr_candidate = only.into_iter().next().unwrap();
let value = parse_expr(expr_candidate);
vec![Statement::Expression { value }]
} else {
block_inner.flat_map(parse_pair).collect()
}
}
_ => vec![Statement::Expression {
value: parse_expr(body_pair),
}],
};
Expr::Lambda { param, body }
}
Rule::postfix_expr => parse_expr(pair),
Rule::call_stmt => {
let mut inner_rules = pair.into_inner();
let function = inner_rules.next().unwrap().as_str().to_string();
let args = inner_rules.map(parse_expr).collect();
Expr::Call { function, args }
}
Rule::call_expr => {
let mut inner = pair.into_inner();
let function = inner.next().unwrap().as_str().to_owned();
let args = if let Some(arg_list) = inner.next() {
arg_list.into_inner().map(parse_expr).collect()
} else {
Vec::new()
};
Expr::Call { function, args }
}
Rule::paren_expr => {
let mut inner = pair.into_inner();
let start_expr = parse_expr(inner.next().unwrap());
let end_expr = parse_expr(inner.next().unwrap());
Expr::Range {
start: Box::new(start_expr),
end: Box::new(end_expr),
}
}
other => unreachable!("Unexpected expression rule: {:?}", other),
}
}
#[cfg(test)]
mod tests {
use super::*;
fn one_stmt(src: &str) -> Statement {
let mut ast = parse(src).expect("parse ok");
assert_eq!(ast.len(), 1, "expected exactly one top level statement");
ast.remove(0)
}
fn one_expr(src: &str) -> Expr {
match one_stmt(src) {
Statement::Expression { value } => value,
other => panic!("expected `Expression` statement, got {other:?}"),
}
}
#[test]
fn parses_true_literal() {
let expr = one_expr("true");
assert!(matches!(expr, Expr::Boolean(true)));
}
#[test]
fn parses_false_literal() {
let expr = one_expr("false");
assert!(matches!(expr, Expr::Boolean(false)));
}
#[test]
fn parses_let_with_true() {
let parsed = parse("let flag = true").unwrap();
assert_eq!(parsed.len(), 1);
match &parsed[0] {
Statement::LetExpression { name, value } => {
assert_eq!(name, "flag");
assert!(matches!(value, &Expr::Boolean(true)));
}
_ => panic!("Expected let binding"),
}
}
#[test]
fn parses_let_with_false() {
let parsed = parse("let flag = false").unwrap();
assert_eq!(parsed.len(), 1);
match &parsed[0] {
Statement::LetExpression { name, value } => {
assert_eq!(name, "flag");
assert!(matches!(value, &Expr::Boolean(false)));
}
_ => panic!("Expected let binding"),
}
}
#[test]
fn parses_call_without_args() {
let e = one_expr("clock()");
assert!(matches!(e,
Expr::Call { ref function, ref args }
if function == "clock" && args.is_empty()
));
}
#[test]
fn parses_call_with_args() {
let e = one_expr("plus(1, 2)");
assert!(matches!(e,
Expr::Call { ref function, ref args }
if function == "plus"
&& matches!(&args[..],
[Expr::Integer(1), Expr::Integer(2)]
)
));
}
#[test]
fn parses_basic_call_expression() {
let input = "plus(1, 2)";
let ast = parse(input).unwrap();
assert!(matches!(ast.as_slice(), [Statement::Expression { .. }]));
}
#[test]
fn parses_field_access() {
let input = "person.name";
let ast = parse(input).unwrap();
assert!(matches!(
ast.as_slice(),
[Statement::Expression {
value:
Expr::Field {
target,
name
}
}]
if matches!(&**target, Expr::Identifier(id) if id == "person")
&& name == "name"
));
}
#[test]
fn parses_field_access_chain() {
let input = "user.profile.name";
let ast = parse(input).unwrap();
let expr = match &ast[..] {
[Statement::Expression { value }] => value,
_ => panic!("unexpected AST shape"),
};
if let Expr::Field {
name: last,
target: outer,
} = expr
{
if last != "name" {
panic!("expected last field to be `name`, got `{last}`");
}
if let Expr::Field {
name: mid,
target: inner,
} = &**outer
{
if mid != "profile" {
panic!("expected mid field to be `profile`, got `{mid}`");
}
if let Expr::Identifier(first) = &**inner {
assert_eq!(first, "user");
return;
}
}
}
panic!("did not match field‑access chain user.profile.name");
}
#[test]
fn parses_simple_function_definition() {
let s = r#"
fn inc(x) {
x + 1
}
"#;
let ast = parse(s).unwrap();
assert!(matches!(
ast.as_slice(),
[Statement::FunctionDef { name, params, .. }]
if name.as_str() == "inc"
&& params == &vec!["x".to_string()]
));
}
#[test]
fn precedence_mul_beats_add() {
let e = one_expr("1 + 2 * 3");
assert!(matches!(e,
Expr::BinaryOp { op: BinaryOpKind::Add, rhs, .. }
if matches!(&*rhs,
Expr::BinaryOp { op: BinaryOpKind::Mul, .. })
));
}
#[test]
fn parses_inline_lambda_as_arg() {
let e = one_expr("arr.map(|x| x * x)");
match e {
Expr::MethodCall { method, arg, .. } => {
assert_eq!(method, "map");
assert!(matches!(*arg,
Expr::Lambda { ref param, .. } if param == "x"
));
}
_ => panic!("not a MethodCall"),
}
}
#[test]
fn parses_hashmap_literal_and_field() {
let e = one_expr(r#"{ msg: "hi", n: 3 }.msg"#);
match e {
Expr::Field { name, .. } => assert_eq!(name, "msg"),
_ => panic!("expected Field access"),
}
}
#[test]
fn parses_let_statement() {
let stmt = one_stmt("let a = 42");
assert!(matches!(stmt,
Statement::LetExpression { ref name, ref value }
if name == "a" && matches!(value, Expr::Integer(42))
));
}
#[test]
fn parses_array_indexing() {
let e = one_expr("arr[2]");
match e {
Expr::Index { target, index } => {
assert!(matches!(*target, Expr::Identifier(ref name) if name == "arr"));
assert!(matches!(*index, Expr::Integer(2)));
}
_ => panic!("Expected Expr::Index"),
}
}
#[test]
fn parses_nested_indexing() {
let e = one_expr("matrix[1][0]");
if let Expr::Index {
target: outer_target,
index: outer_index,
} = e
{
assert!(matches!(*outer_index, Expr::Integer(0)));
if let Expr::Index {
target: inner_target,
index: inner_index,
} = *outer_target
{
assert!(matches!(*inner_index, Expr::Integer(1)));
assert!(matches!(*inner_target, Expr::Identifier(ref name) if name == "matrix"));
} else {
panic!("Expected inner Expr::Index");
}
} else {
panic!("Expected outer Expr::Index");
}
}
#[test]
fn parses_array_length_field() {
let expr = one_expr("arr.length");
match expr {
Expr::Field { target, name } => {
assert_eq!(name, "length");
match *target {
Expr::Identifier(ref id) => assert_eq!(id, "arr"),
_ => panic!("Expected `arr` as field target"),
}
}
_ => panic!("Expected Expr::Field"),
}
}
#[test]
fn parses_simple_range() {
let expr = one_expr("(0..5)");
match expr {
Expr::Range { start, end } => match (*start, *end) {
(Expr::Integer(a), Expr::Integer(b)) => {
assert_eq!(a, 0);
assert_eq!(b, 5);
}
_ => panic!("Expected integer range bounds"),
},
other => panic!("Expected Expr::Range, got {other:?}"),
}
}
#[test]
fn parses_let_binding_with_range() {
let ast = parse("let r = (1..10)").unwrap();
assert_eq!(ast.len(), 1);
match &ast[0] {
Statement::LetExpression { name, value } => {
assert_eq!(name, "r");
match value {
Expr::Range { start, end } => match (&**start, &**end) {
(Expr::Integer(a), Expr::Integer(b)) => {
assert_eq!(*a, 1);
assert_eq!(*b, 10);
}
_ => panic!("Expected integer bounds in range"),
},
_ => panic!("Expected Expr::Range inside let binding"),
}
}
_ => panic!("Expected let binding"),
}
}
#[test]
fn parses_range_in_method_chain() {
let expr = one_expr("(0..3).for_each(|i| i)");
match expr {
Expr::MethodCall {
target,
method,
arg,
} => {
assert_eq!(method, "for_each");
match *target {
Expr::Range { start, end } => match (*start, *end) {
(Expr::Integer(a), Expr::Integer(b)) => {
assert_eq!(a, 0);
assert_eq!(b, 3);
}
_ => panic!("Expected integer bounds"),
},
_ => panic!("Expected Expr::Range as method target"),
}
assert!(matches!(*arg, Expr::Lambda { .. }));
}
_ => panic!("Expected Expr::MethodCall on Expr::Range"),
}
}
#[test]
fn parses_hashmap_inside_lambda_map() {
let e = one_expr("[1, 2, 3].map(|x| { value: x, label: \"hi\" })");
match e {
Expr::MethodCall { method, arg, .. } => {
assert_eq!(method, "map");
assert!(matches!(*arg, Expr::Lambda { .. }));
}
_ => panic!("Expected MethodCall"),
}
}
}