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.into_inner();
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,
}
}
})
}
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.into_inner();
let param = inner.next().unwrap().as_str().to_string();
let body_pair = inner.next().unwrap();
let body = if body_pair.as_rule() == Rule::block {
body_pair
.into_inner()
.next()
.map(|st| st.into_inner().flat_map(parse_pair).collect())
.unwrap_or_default()
} else {
vec![Statement::Expression {
value: parse_expr(body_pair),
}]
};
Expr::Lambda { param, body }
}
_ => 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 => Expr::String(pair.as_str().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::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::stmt_list => body_pair.into_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 }
}
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))
));
}
}