use super::NormResult;
use super::resolve;
use crate::grammar::parse_tree::{
Alternative, ExprSymbol, Grammar, GrammarItem, NonterminalData, NonterminalString, Symbol,
SymbolKind,
};
use std::fmt;
use std::str::FromStr;
use string_cache::DefaultAtom as Atom;
#[cfg(test)]
mod test;
pub const PREC_ATTR: &str = "precedence";
pub const LVL_ARG: &str = "level";
pub const ASSOC_ATTR: &str = "assoc";
pub const SIDE_ARG: &str = "side";
#[allow(clippy::enum_variant_names)]
#[derive(Clone, Copy, Eq, PartialEq, Default)]
pub enum Assoc {
Left,
Right,
NonAssoc,
#[default]
FullyAssoc,
}
#[derive(Clone, Copy, Eq, PartialEq)]
pub enum Substitution<'a> {
OneThen(&'a SymbolKind, &'a SymbolKind),
Every(&'a SymbolKind),
}
#[derive(Clone, Copy, Eq, PartialEq)]
pub enum Direction {
Forward,
Backward,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ParseAssocError {
_priv: (),
}
impl fmt::Display for ParseAssocError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
"provided value was neither `left`, `right` nor `none`".fmt(f)
}
}
impl FromStr for Assoc {
type Err = ParseAssocError;
fn from_str(s: &str) -> Result<Assoc, ParseAssocError> {
match s {
"left" => Ok(Assoc::Left),
"right" => Ok(Assoc::Right),
"none" => Ok(Assoc::NonAssoc),
"all" => Ok(Assoc::FullyAssoc),
_ => Err(ParseAssocError { _priv: () }),
}
}
}
pub fn expand_precedence(input: Grammar) -> NormResult<Grammar> {
let input = resolve::resolve(input)?;
let mut result: Vec<GrammarItem> = Vec::with_capacity(input.items.len());
for item in input.items.into_iter() {
match item {
GrammarItem::Nonterminal(d) if has_prec_attr(&d) => result.extend(expand_nonterm(d)?),
item => result.push(item),
};
}
Ok(Grammar {
items: result,
..input
})
}
pub fn has_prec_attr(non_term: &NonterminalData) -> bool {
non_term
.alternatives
.first()
.map(|alt| {
alt.attributes
.iter()
.any(|attr| attr.id == *PREC_ATTR || attr.id == *ASSOC_ATTR)
})
.unwrap_or(false)
}
fn expand_nonterm(mut nonterm: NonterminalData) -> NormResult<Vec<GrammarItem>> {
let mut lvls: Vec<u32> = Vec::new();
let mut alts_with_attr: Vec<(u32, Assoc, Alternative)> =
Vec::with_capacity(nonterm.alternatives.len());
let _ = nonterm.alternatives.drain(..).fold(
(0, Assoc::default()),
|(last_lvl, last_assoc): (u32, Assoc), mut alt| {
let (lvl, last_assoc) = alt
.attributes
.iter()
.position(|attr| attr.id == *PREC_ATTR)
.map(|index| {
let attr = alt.attributes.remove(index);
let (_, val) = attr.get_arg_equal().unwrap();
(val.parse().unwrap(), Assoc::default())
})
.unwrap_or((last_lvl, last_assoc));
let assoc = alt
.attributes
.iter()
.position(|attr| attr.id == *ASSOC_ATTR)
.map(|index| {
let attr = alt.attributes.remove(index);
let (_, val) = attr.get_arg_equal().unwrap();
val.parse().unwrap()
})
.unwrap_or(last_assoc);
alts_with_attr.push((lvl, assoc, alt));
lvls.push(lvl);
(lvl, assoc)
},
);
lvls.sort_unstable();
lvls.dedup();
let rest = &mut alts_with_attr.into_iter();
let lvl_max = *lvls.last().unwrap();
let result = Some(None)
.into_iter()
.chain(lvls.iter().map(Some))
.zip(lvls.iter())
.map(|(lvl_prec_opt, lvl)| {
let name = NonterminalString(Atom::from(if *lvl == lvl_max {
format!("{}", nonterm.name)
} else {
format!("{}{}", nonterm.name, lvl)
}));
let nonterm_prev = lvl_prec_opt.map(|lvl_prec| {
SymbolKind::Nonterminal(NonterminalString(Atom::from(format!(
"{}{}",
nonterm.name, lvl_prec
))))
});
let (alts_with_prec, new_rest): (Vec<_>, Vec<_>) =
rest.partition(|(l, _, _)| *l == *lvl);
*rest = new_rest.into_iter();
let mut alts_with_assoc: Vec<_> = alts_with_prec
.into_iter()
.map(|(_, assoc, alt)| (assoc, alt))
.collect();
let symbol_kind = &SymbolKind::Nonterminal(name.clone());
for (assoc, alt) in &mut alts_with_assoc {
let err_msg = "unexpected associativity attribute on the first precedence level";
let (subst, dir) = match assoc {
Assoc::Left => (
Substitution::OneThen(symbol_kind, nonterm_prev.as_ref().expect(err_msg)),
Direction::Forward,
),
Assoc::Right => (
Substitution::OneThen(symbol_kind, nonterm_prev.as_ref().expect(err_msg)),
Direction::Backward,
),
Assoc::NonAssoc => (
Substitution::Every(nonterm_prev.as_ref().expect(err_msg)),
Direction::Forward,
),
Assoc::FullyAssoc => (Substitution::Every(symbol_kind), Direction::Forward),
};
replace_nonterm(alt, &nonterm.name, subst, dir)
}
let mut alternatives: Vec<_> =
alts_with_assoc.into_iter().map(|(_, alt)| alt).collect();
if let Some(kind) = nonterm_prev {
alternatives.push(Alternative {
span: nonterm.span,
expr: ExprSymbol {
symbols: vec![Symbol {
kind,
span: nonterm.span,
}],
},
condition: None,
action: None,
attributes: vec![],
});
}
GrammarItem::Nonterminal(NonterminalData {
visibility: nonterm.visibility.clone(),
name,
attributes: nonterm.attributes.clone(),
span: nonterm.span,
args: nonterm.args.clone(), type_decl: nonterm.type_decl.clone(),
alternatives,
})
});
let items = result.collect();
assert!(rest.next().is_none());
Ok(items)
}
fn replace_nonterm(
alt: &mut Alternative,
target: &NonterminalString,
subst: Substitution<'_>,
dir: Direction,
) {
replace_symbols(&mut alt.expr.symbols, target, subst, dir);
}
fn replace_symbols<'a>(
symbols: &mut [Symbol],
target: &NonterminalString,
subst: Substitution<'a>,
dir: Direction,
) -> Substitution<'a> {
match dir {
Direction::Forward => symbols.iter_mut().fold(subst, |subst, symbol| {
replace_symbol(symbol, target, subst, dir)
}),
Direction::Backward => symbols.iter_mut().rev().fold(subst, |subst, symbol| {
replace_symbol(symbol, target, subst, dir)
}),
}
}
fn replace_symbol<'a>(
symbol: &mut Symbol,
target: &NonterminalString,
subst: Substitution<'a>,
dir: Direction,
) -> Substitution<'a> {
match symbol.kind {
SymbolKind::AmbiguousId(ref id) => {
panic!("ambiguous id `{id}` encountered after name resolution")
}
SymbolKind::Nonterminal(ref name) if name == target => match subst {
Substitution::Every(sym_kind) => {
symbol.kind = sym_kind.clone();
subst
}
Substitution::OneThen(fst, snd) => {
symbol.kind = fst.clone();
Substitution::Every(snd)
}
},
SymbolKind::Macro(ref mut m) => {
if dir == Direction::Forward {
m.args
.iter_mut()
.fold(subst, |subst, sym| replace_symbol(sym, target, subst, dir))
} else {
m.args
.iter_mut()
.rev()
.fold(subst, |subst, sym| replace_symbol(sym, target, subst, dir))
}
}
SymbolKind::Expr(ref mut expr) => replace_symbols(&mut expr.symbols, target, subst, dir),
SymbolKind::Repeat(ref mut repeat) => {
replace_symbol(&mut repeat.symbol, target, subst, dir)
}
SymbolKind::Choose(ref mut sym)
| SymbolKind::Name(_, ref mut sym)
| SymbolKind::Tuple(_, ref mut sym) => replace_symbol(sym, target, subst, dir),
SymbolKind::Terminal(_)
| SymbolKind::Nonterminal(_)
| SymbolKind::Error
| SymbolKind::Lookahead
| SymbolKind::Lookbehind => subst,
}
}