use std::collections::{HashMap, HashSet};
use crate::lexer::text;
use crate::token::TokenKind;
#[derive(Clone, Debug, PartialEq, Eq)]
enum Sym {
Ref(String),
Lit(String),
Kind(TokenKind),
Repeat(Box<Sym>, Quant),
Group(Vec<Vec<Sym>>),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Quant {
Star,
Plus,
Opt,
}
impl Quant {
fn bounds(self) -> (usize, usize) {
match self {
Quant::Star => (0, usize::MAX),
Quant::Plus => (1, usize::MAX),
Quant::Opt => (0, 1),
}
}
fn glyph(self) -> char {
match self {
Quant::Star => '*',
Quant::Plus => '+',
Quant::Opt => '?',
}
}
}
#[derive(Clone, Debug)]
struct Rule {
name: String,
alts: Vec<Vec<Sym>>,
weights: Vec<f64>,
}
#[derive(Clone, Debug)]
pub struct Grammar {
rules: Vec<Rule>,
start: String,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct GrammarError {
pub line: usize,
pub msg: String,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Node {
Terminal(String),
Branch(String, Vec<Node>),
}
type Memo = HashMap<(String, usize), Option<(Node, usize)>>;
struct ParseCtx<'a> {
lr: &'a HashSet<String>,
dep: &'a HashSet<String>,
seed: Memo,
packrat: Memo,
}
impl Grammar {
pub fn parse(src: &str) -> Result<Grammar, GrammarError> {
let mut rules: Vec<Rule> = Vec::new();
for (idx, raw) in src.lines().enumerate() {
let line_no = idx + 1;
let line = strip_comment(raw).trim();
if line.is_empty() {
continue;
}
let Some((name, rhs)) = line.split_once(":=") else {
return Err(GrammarError { line: line_no, msg: "expected `name := ...`".into() });
};
let name = name.trim().to_string();
if name.is_empty() {
return Err(GrammarError { line: line_no, msg: "empty rule name".into() });
}
let (alts, weights) = parse_rhs(rhs, line_no)?;
rules.push(Rule { name, alts, weights });
}
if rules.is_empty() {
return Err(GrammarError { line: 0, msg: "grammar has no rules".into() });
}
let start = rules[0].name.clone();
Ok(Grammar { rules, start })
}
#[must_use]
pub fn with_start(mut self, start: &str) -> Self {
self.start = start.to_string();
self
}
#[must_use]
pub fn start(&self) -> &str {
&self.start
}
#[must_use]
pub fn rule_names(&self) -> Vec<&str> {
self.rules.iter().map(|r| r.name.as_str()).collect()
}
fn rule(&self, name: &str) -> Option<&Rule> {
self.rules.iter().find(|r| r.name == name)
}
#[must_use]
pub fn parse_input(&self, input: &[u8]) -> Option<Node> {
let toks: Vec<Term> = crate::parallel_lex::lex_parallel(input)
.iter()
.filter(|t| t.kind != TokenKind::Whitespace)
.map(|t| Term {
kind: t.kind,
text: String::from_utf8_lossy(text(input, t)).into_owned(),
})
.collect();
let lr = self.left_recursive_rules();
let dep = self.lr_dependent(&lr);
let mut ctx =
ParseCtx { lr: &lr, dep: &dep, seed: HashMap::new(), packrat: HashMap::new() };
let (node, pos) = self.parse_rule(&self.start, &toks, 0, &mut ctx)?;
(pos == toks.len()).then_some(node)
}
fn lr_dependent(&self, lr: &HashSet<String>) -> HashSet<String> {
let mut out = lr.clone();
loop {
let mut changed = false;
for rule in &self.rules {
if out.contains(&rule.name) {
continue;
}
let mut refs: Vec<String> = Vec::new();
for alt in &rule.alts {
for s in alt {
sym_all_refs(s, &mut refs);
}
}
if refs.iter().any(|r| out.contains(r.as_str())) {
out.insert(rule.name.clone());
changed = true;
}
}
if !changed {
return out;
}
}
}
fn left_recursive_rules(&self) -> HashSet<String> {
let nullable = self.nullable_rules();
self.rules
.iter()
.filter(|r| self.reaches_self_leftmost(&r.name, &nullable))
.map(|r| r.name.clone())
.collect()
}
fn nullable_rules(&self) -> HashSet<String> {
let mut nullable: HashSet<String> = HashSet::new();
loop {
let mut changed = false;
for rule in &self.rules {
if !nullable.contains(&rule.name)
&& rule.alts.iter().any(|alt| self.seq_nullable(alt, &nullable))
{
nullable.insert(rule.name.clone());
changed = true;
}
}
if !changed {
return nullable;
}
}
}
fn seq_nullable(&self, seq: &[Sym], nullable: &HashSet<String>) -> bool {
seq.iter().all(|s| self.sym_nullable(s, nullable))
}
fn sym_nullable(&self, sym: &Sym, nullable: &HashSet<String>) -> bool {
match sym {
Sym::Ref(r) => nullable.contains(r.as_str()),
Sym::Lit(_) | Sym::Kind(_) => false,
Sym::Repeat(inner, q) => {
matches!(q, Quant::Star | Quant::Opt) || self.sym_nullable(inner, nullable)
}
Sym::Group(alts) => alts.iter().any(|alt| self.seq_nullable(alt, nullable)),
}
}
fn reaches_self_leftmost(&self, name: &str, nullable: &HashSet<String>) -> bool {
let mut stack: Vec<String> = vec![name.to_string()];
let mut seen: HashSet<String> = HashSet::new();
let mut refs: Vec<String> = Vec::new();
while let Some(cur) = stack.pop() {
let Some(rule) = self.rule(&cur) else { continue };
refs.clear();
for alt in &rule.alts {
self.seq_leftmost_refs(alt, nullable, &mut refs);
}
for r in refs.drain(..) {
if r.as_str() == name {
return true;
}
if seen.insert(r.clone()) {
stack.push(r);
}
}
}
false
}
fn seq_leftmost_refs(&self, seq: &[Sym], nullable: &HashSet<String>, out: &mut Vec<String>) {
for s in seq {
self.sym_leftmost_refs(s, nullable, out);
if !self.sym_nullable(s, nullable) {
break;
}
}
}
fn sym_leftmost_refs(&self, sym: &Sym, nullable: &HashSet<String>, out: &mut Vec<String>) {
match sym {
Sym::Ref(r) => out.push(r.clone()),
Sym::Lit(_) | Sym::Kind(_) => {}
Sym::Repeat(inner, _) => self.sym_leftmost_refs(inner, nullable, out),
Sym::Group(alts) => {
for alt in alts {
self.seq_leftmost_refs(alt, nullable, out);
}
}
}
}
fn parse_rule(
&self,
name: &str,
toks: &[Term],
pos: usize,
ctx: &mut ParseCtx<'_>,
) -> Option<(Node, usize)> {
let key = (name.to_string(), pos);
if let Some(seed) = ctx.seed.get(&key) {
return seed.clone();
}
if !ctx.lr.contains(name) {
let cacheable = !ctx.dep.contains(name);
if cacheable && let Some(hit) = ctx.packrat.get(&key) {
return hit.clone();
}
let out = self.parse_alts(name, toks, pos, ctx);
if cacheable {
ctx.packrat.insert(key, out.clone());
}
return out;
}
ctx.seed.insert(key.clone(), None);
let mut seed = self.parse_alts(name, toks, pos, ctx);
while let Some((_, end)) = seed.clone() {
ctx.seed.insert(key.clone(), seed.clone());
match self.parse_alts(name, toks, pos, ctx) {
Some((node, np)) if np > end => seed = Some((node, np)),
_ => break,
}
}
ctx.seed.remove(&key);
seed
}
fn parse_alts(
&self,
name: &str,
toks: &[Term],
pos: usize,
ctx: &mut ParseCtx<'_>,
) -> Option<(Node, usize)> {
let rule = self.rule(name)?;
for alt in &rule.alts {
if let Some((children, np)) = self.match_seq(alt, toks, pos, ctx) {
return Some((Node::Branch(name.to_string(), children), np));
}
}
None
}
fn match_seq(
&self,
syms: &[Sym],
toks: &[Term],
pos: usize,
ctx: &mut ParseCtx<'_>,
) -> Option<(Vec<Node>, usize)> {
let mut children: Vec<Node> = Vec::new();
let mut p = pos;
for sym in syms {
p = self.match_sym(sym, toks, p, ctx, &mut children)?;
}
Some((children, p))
}
fn match_sym(
&self,
sym: &Sym,
toks: &[Term],
pos: usize,
ctx: &mut ParseCtx<'_>,
children: &mut Vec<Node>,
) -> Option<usize> {
match sym {
Sym::Ref(name) => {
let (node, np) = self.parse_rule(name, toks, pos, ctx)?;
children.push(node);
Some(np)
}
Sym::Lit(lit) => {
let t = toks.get(pos)?;
if &t.text != lit {
return None;
}
children.push(Node::Terminal(t.text.clone()));
Some(pos + 1)
}
Sym::Kind(k) => {
let t = toks.get(pos)?;
if t.kind != *k {
return None;
}
children.push(Node::Terminal(t.text.clone()));
Some(pos + 1)
}
Sym::Repeat(inner, quant) => {
let (min, max) = quant.bounds();
let mut local: Vec<Node> = Vec::new();
let mut p = pos;
let mut count = 0;
while count < max {
let mut round: Vec<Node> = Vec::new();
match self.match_sym(inner, toks, p, ctx, &mut round) {
Some(np) if np > p => {
local.append(&mut round);
p = np;
count += 1;
}
_ => break,
}
}
if count < min {
return None;
}
children.append(&mut local);
Some(p)
}
Sym::Group(alts) => {
for alt in alts {
if let Some((mut kids, np)) = self.match_seq(alt, toks, pos, ctx) {
children.append(&mut kids);
return Some(np);
}
}
None
}
}
}
}
struct Term {
kind: TokenKind,
text: String,
}
impl Node {
#[must_use]
pub fn sexpr(&self) -> String {
match self {
Node::Terminal(t) => t.clone(),
Node::Branch(_, kids) if kids.len() == 1 => kids[0].sexpr(),
Node::Branch(name, kids) => {
let parts: Vec<String> = kids.iter().map(Node::sexpr).collect();
format!("({name} {})", parts.join(" "))
}
}
}
}
fn strip_comment(line: &str) -> &str {
line.split_once('#').map_or(line, |(before, _)| before)
}
fn sym_all_refs(sym: &Sym, out: &mut Vec<String>) {
match sym {
Sym::Ref(r) => out.push(r.clone()),
Sym::Lit(_) | Sym::Kind(_) => {}
Sym::Repeat(inner, _) => sym_all_refs(inner, out),
Sym::Group(alts) => {
for alt in alts {
for s in alt {
sym_all_refs(s, out);
}
}
}
}
}
#[derive(Clone, Debug, PartialEq)]
enum RhsTok {
Ref(String),
Lit(String),
Word(String),
LParen,
RParen,
Pipe,
Quant(Quant),
Weight(f64),
}
fn tokenize_rhs(rhs: &str, line: usize) -> Result<Vec<RhsTok>, GrammarError> {
let bytes = rhs.as_bytes();
let mut out = Vec::new();
let mut i = 0;
while i < bytes.len() {
let c = bytes[i];
if c.is_ascii_whitespace() {
i += 1;
} else if c == b'<' {
let start = i + 1;
let end = rhs[start..]
.find('>')
.map(|o| start + o)
.ok_or_else(|| GrammarError { line, msg: "unterminated `<reference>`".into() })?;
if start == end {
return Err(GrammarError { line, msg: "empty `<reference>`".into() });
}
out.push(RhsTok::Ref(rhs[start..end].to_string()));
i = end + 1;
} else if c == b'"' {
let start = i + 1;
let end = rhs[start..]
.find('"')
.map(|o| start + o)
.ok_or_else(|| GrammarError { line, msg: "unterminated `\"literal\"`".into() })?;
out.push(RhsTok::Lit(rhs[start..end].to_string()));
i = end + 1;
} else if c == b'(' {
out.push(RhsTok::LParen);
i += 1;
} else if c == b')' {
out.push(RhsTok::RParen);
i += 1;
} else if c == b'|' {
out.push(RhsTok::Pipe);
i += 1;
} else if let Some(q) = quant_of(c) {
out.push(RhsTok::Quant(q));
i += 1;
} else if c == b'@' {
let start = i + 1;
let mut j = start;
while j < bytes.len() && (bytes[j].is_ascii_digit() || bytes[j] == b'.') {
j += 1;
}
let num = &rhs[start..j];
let w: f64 = num
.parse()
.map_err(|_| GrammarError { line, msg: format!("`@{num}` is not a weight") })?;
out.push(RhsTok::Weight(w));
i = j;
} else {
let start = i;
while i < bytes.len() {
let d = bytes[i];
if d.is_ascii_whitespace()
|| matches!(d, b'<' | b'"' | b'(' | b')' | b'|' | b'*' | b'+' | b'?' | b'@')
{
break;
}
i += 1;
}
out.push(RhsTok::Word(rhs[start..i].to_string()));
}
}
Ok(out)
}
fn quant_of(c: u8) -> Option<Quant> {
match c {
b'*' => Some(Quant::Star),
b'+' => Some(Quant::Plus),
b'?' => Some(Quant::Opt),
_ => None,
}
}
fn parse_rhs(rhs: &str, line: usize) -> Result<(Vec<Vec<Sym>>, Vec<f64>), GrammarError> {
let toks = tokenize_rhs(rhs, line)?;
let mut pos = 0;
let mut alts = Vec::new();
let mut weights = Vec::new();
loop {
alts.push(parse_seq(&toks, &mut pos, line)?);
let w = if let Some(RhsTok::Weight(w)) = toks.get(pos) {
pos += 1;
*w
} else {
1.0
};
weights.push(w);
if matches!(toks.get(pos), Some(RhsTok::Pipe)) {
pos += 1;
} else {
break;
}
}
if pos != toks.len() {
return Err(GrammarError { line, msg: "unexpected `)` in rule body".into() });
}
Ok((alts, weights))
}
fn parse_alts(toks: &[RhsTok], pos: &mut usize, line: usize) -> Result<Vec<Vec<Sym>>, GrammarError> {
let mut alts = vec![parse_seq(toks, pos, line)?];
while matches!(toks.get(*pos), Some(RhsTok::Pipe)) {
*pos += 1;
alts.push(parse_seq(toks, pos, line)?);
}
Ok(alts)
}
fn parse_seq(toks: &[RhsTok], pos: &mut usize, line: usize) -> Result<Vec<Sym>, GrammarError> {
let mut seq = Vec::new();
while let Some(t) = toks.get(*pos) {
if matches!(t, RhsTok::Pipe | RhsTok::RParen | RhsTok::Weight(_)) {
break;
}
seq.push(parse_item(toks, pos, line)?);
}
Ok(seq)
}
fn parse_item(toks: &[RhsTok], pos: &mut usize, line: usize) -> Result<Sym, GrammarError> {
let mut sym = parse_primary(toks, pos, line)?;
while let Some(RhsTok::Quant(q)) = toks.get(*pos) {
sym = Sym::Repeat(Box::new(sym), *q);
*pos += 1;
}
Ok(sym)
}
fn parse_primary(toks: &[RhsTok], pos: &mut usize, line: usize) -> Result<Sym, GrammarError> {
let tok = toks
.get(*pos)
.ok_or_else(|| GrammarError { line, msg: "unexpected end of rule body".into() })?;
match tok {
RhsTok::Ref(r) => {
*pos += 1;
Ok(Sym::Ref(r.clone()))
}
RhsTok::Lit(l) => {
*pos += 1;
Ok(Sym::Lit(l.clone()))
}
RhsTok::Word(w) => {
let sym = kind_from_word(w, line)?;
*pos += 1;
Ok(sym)
}
RhsTok::LParen => {
*pos += 1;
let alts = parse_alts(toks, pos, line)?;
if !matches!(toks.get(*pos), Some(RhsTok::RParen)) {
return Err(GrammarError { line, msg: "unterminated `(` group".into() });
}
*pos += 1;
Ok(Sym::Group(alts))
}
RhsTok::Quant(q) => {
Err(GrammarError { line, msg: format!("`{}` has no symbol to repeat", q.glyph()) })
}
RhsTok::RParen | RhsTok::Pipe | RhsTok::Weight(_) => {
Err(GrammarError { line, msg: "expected a symbol".into() })
}
}
}
fn kind_from_word(word: &str, line: usize) -> Result<Sym, GrammarError> {
let kind = match word {
"number" => TokenKind::Number,
"ident" => TokenKind::Word,
"string" => TokenKind::Quoted,
"ip" => TokenKind::Ip,
"url" => TokenKind::Url,
"email" => TokenKind::Email,
"time" => TokenKind::Timestamp,
"punct" => TokenKind::Punct,
other => {
return Err(GrammarError {
line,
msg: format!("unknown symbol `{other}` (use <ref>, \"lit\", a group `(...)`, or a terminal: number/ident/string/ip/url/email/time/punct)"),
});
}
};
Ok(Sym::Kind(kind))
}
pub trait Semiring: Clone + PartialEq {
fn zero() -> Self;
fn one() -> Self;
fn add(&self, other: &Self) -> Self;
fn mul(&self, other: &Self) -> Self;
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Recognize(pub bool);
impl Semiring for Recognize {
fn zero() -> Self {
Recognize(false)
}
fn one() -> Self {
Recognize(true)
}
fn add(&self, other: &Self) -> Self {
Recognize(self.0 || other.0)
}
fn mul(&self, other: &Self) -> Self {
Recognize(self.0 && other.0)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Count(pub u64);
impl Semiring for Count {
fn zero() -> Self {
Count(0)
}
fn one() -> Self {
Count(1)
}
fn add(&self, other: &Self) -> Self {
Count(self.0.saturating_add(other.0))
}
fn mul(&self, other: &Self) -> Self {
Count(self.0.saturating_mul(other.0))
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Viterbi(pub f64);
impl Semiring for Viterbi {
fn zero() -> Self {
Viterbi(0.0)
}
fn one() -> Self {
Viterbi(1.0)
}
fn add(&self, other: &Self) -> Self {
Viterbi(self.0.max(other.0))
}
fn mul(&self, other: &Self) -> Self {
Viterbi(self.0 * other.0)
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Prob(pub f64);
impl Semiring for Prob {
fn zero() -> Self {
Prob(0.0)
}
fn one() -> Self {
Prob(1.0)
}
fn add(&self, other: &Self) -> Self {
Prob(self.0 + other.0)
}
fn mul(&self, other: &Self) -> Self {
Prob(self.0 * other.0)
}
}
#[derive(Clone, Debug)]
enum BSym {
Lit(String),
Kind(TokenKind),
Nt(usize),
}
struct Bnf {
alts: Vec<Vec<Vec<BSym>>>,
weights: Vec<Vec<f64>>,
start: usize,
}
fn desugar_sym(
sym: &Sym,
names: &HashMap<&str, usize>,
weights: &mut Vec<Vec<f64>>,
alts: &mut Vec<Vec<Vec<BSym>>>,
) -> Option<BSym> {
match sym {
Sym::Lit(l) => Some(BSym::Lit(l.clone())),
Sym::Kind(k) => Some(BSym::Kind(*k)),
Sym::Ref(r) => names.get(r.as_str()).map(|&i| BSym::Nt(i)),
Sym::Repeat(inner, quant) => {
let inner_sym = desugar_sym(inner, names, weights, alts)?;
let nt = alts.len();
alts.push(Vec::new());
weights.push(Vec::new());
alts[nt] = match quant {
Quant::Star => vec![vec![inner_sym, BSym::Nt(nt)], vec![]],
Quant::Plus => vec![vec![inner_sym.clone(), BSym::Nt(nt)], vec![inner_sym]],
Quant::Opt => vec![vec![inner_sym], vec![]],
};
weights[nt] = vec![1.0; alts[nt].len()];
Some(BSym::Nt(nt))
}
Sym::Group(galts) => {
let nt = alts.len();
alts.push(Vec::new());
weights.push(Vec::new());
let mut bnf_alts = Vec::with_capacity(galts.len());
for ga in galts {
let mut seq = Vec::with_capacity(ga.len());
for s in ga {
seq.push(desugar_sym(s, names, weights, alts)?);
}
bnf_alts.push(seq);
}
let count = bnf_alts.len();
alts[nt] = bnf_alts;
weights[nt] = vec![1.0; count];
Some(BSym::Nt(nt))
}
}
}
struct LatEdge {
to: usize,
text: String,
kind: TokenKind,
}
pub struct Lattice {
n: usize,
edges: Vec<Vec<LatEdge>>,
}
impl Lattice {
fn linear(toks: &[Term]) -> Lattice {
let mut edges: Vec<Vec<LatEdge>> = (0..toks.len()).map(|_| Vec::new()).collect();
for (i, t) in toks.iter().enumerate() {
edges[i].push(LatEdge { to: i + 1, text: t.text.clone(), kind: t.kind });
}
Lattice { n: toks.len(), edges }
}
#[must_use]
pub fn segment(input: &str, dict: &[String]) -> Lattice {
let chars: Vec<char> = input.chars().collect();
let n = chars.len();
let mut edges: Vec<Vec<LatEdge>> = (0..n).map(|_| Vec::new()).collect();
for i in 0..n {
for w in dict {
let wlen = w.chars().count();
if wlen > 0 && i + wlen <= n && chars[i..i + wlen].iter().copied().eq(w.chars()) {
edges[i].push(LatEdge { to: i + wlen, text: w.clone(), kind: TokenKind::Word });
}
}
}
Lattice { n, edges }
}
fn edge_matches(&self, i: usize, m: usize, term: &BSym) -> bool {
self.edges.get(i).is_some_and(|es| {
es.iter().any(|e| {
e.to == m
&& match term {
BSym::Lit(l) => e.text == *l,
BSym::Kind(k) => e.kind == *k,
BSym::Nt(_) => false,
}
})
})
}
}
struct Inside<'a, S: Semiring, W: Fn(f64) -> S> {
bnf: &'a Bnf,
lattice: &'a Lattice,
weight_of: &'a W,
nt_memo: HashMap<(usize, usize, usize), Option<S>>,
seq_memo: HashMap<(usize, usize, usize, usize, usize), S>,
}
impl<S: Semiring, W: Fn(f64) -> S> Inside<'_, S, W> {
fn nt(&mut self, nt: usize, i: usize, j: usize) -> S {
if let Some(entry) = self.nt_memo.get(&(nt, i, j)) {
return entry.clone().unwrap_or_else(S::zero);
}
self.nt_memo.insert((nt, i, j), None);
let mut val = S::zero();
for ai in 0..self.bnf.alts[nt].len() {
let tail = self.seq(nt, ai, 0, i, j);
let weighted = (self.weight_of)(self.bnf.weights[nt][ai]).mul(&tail);
val = val.add(&weighted);
}
self.nt_memo.insert((nt, i, j), Some(val.clone()));
val
}
fn seq(&mut self, nt: usize, ai: usize, dot: usize, i: usize, j: usize) -> S {
if dot == self.bnf.alts[nt][ai].len() {
return if i == j { S::one() } else { S::zero() };
}
if let Some(v) = self.seq_memo.get(&(nt, ai, dot, i, j)) {
return v.clone();
}
let mut val = S::zero();
for m in i..=j {
let head = self.sym(nt, ai, dot, i, m);
if head == S::zero() {
continue;
}
let tail = self.seq(nt, ai, dot + 1, m, j);
val = val.add(&head.mul(&tail));
}
self.seq_memo.insert((nt, ai, dot, i, j), val.clone());
val
}
fn sym(&mut self, nt: usize, ai: usize, dot: usize, i: usize, m: usize) -> S {
match self.bnf.alts[nt][ai][dot].clone() {
BSym::Nt(b) => self.nt(b, i, m),
ref term => {
if self.lattice.edge_matches(i, m, term) {
S::one()
} else {
S::zero()
}
}
}
}
}
impl Grammar {
fn to_bnf(&self) -> Option<Bnf> {
let mut names: HashMap<&str, usize> = HashMap::new();
for (i, r) in self.rules.iter().enumerate() {
names.entry(r.name.as_str()).or_insert(i);
}
let start = *names.get(self.start.as_str())?;
let mut alts: Vec<Vec<Vec<BSym>>> = vec![Vec::new(); self.rules.len()];
let mut weights: Vec<Vec<f64>> = vec![Vec::new(); self.rules.len()];
for (i, rule) in self.rules.iter().enumerate() {
let mut rule_alts = Vec::with_capacity(rule.alts.len());
for alt in &rule.alts {
let mut seq = Vec::with_capacity(alt.len());
for sym in alt {
seq.push(desugar_sym(sym, &names, &mut weights, &mut alts)?);
}
rule_alts.push(seq);
}
alts[i] = rule_alts;
weights[i] = rule.weights.clone();
}
Some(Bnf { alts, weights, start })
}
#[must_use]
pub fn value<S, W>(&self, input: &[u8], weight_of: W) -> Option<S>
where
S: Semiring,
W: Fn(f64) -> S,
{
let toks: Vec<Term> = crate::parallel_lex::lex_parallel(input)
.iter()
.filter(|t| t.kind != TokenKind::Whitespace)
.map(|t| Term {
kind: t.kind,
text: String::from_utf8_lossy(text(input, t)).into_owned(),
})
.collect();
self.value_over_lattice(&Lattice::linear(&toks), weight_of)
}
#[must_use]
pub fn value_over_lattice<S, W>(&self, lattice: &Lattice, weight_of: W) -> Option<S>
where
S: Semiring,
W: Fn(f64) -> S,
{
let bnf = self.to_bnf()?;
let mut inside = Inside {
bnf: &bnf,
lattice,
weight_of: &weight_of,
nt_memo: HashMap::new(),
seq_memo: HashMap::new(),
};
Some(inside.nt(bnf.start, 0, lattice.n))
}
#[must_use]
pub fn recognizes(&self, input: &[u8]) -> bool {
self.value::<Recognize, _>(input, |w| Recognize(w > 0.0)).is_some_and(|r| r.0)
}
#[must_use]
pub fn count_parses(&self, input: &[u8]) -> u64 {
self.value::<Count, _>(input, |_| Count::one()).map_or(0, |c| c.0)
}
#[must_use]
pub fn best_parse_prob(&self, input: &[u8]) -> f64 {
self.value::<Viterbi, _>(input, Viterbi).map_or(0.0, |v| v.0)
}
#[must_use]
pub fn total_prob(&self, input: &[u8]) -> f64 {
self.value::<Prob, _>(input, Prob).map_or(0.0, |p| p.0)
}
#[must_use]
pub fn count_segmentations(&self, input: &str, dict: &[String]) -> u64 {
let lattice = Lattice::segment(input, dict);
self.value_over_lattice::<Count, _>(&lattice, |_| Count::one()).map_or(0, |c| c.0)
}
#[must_use]
pub fn best_segmentation_prob(&self, input: &str, dict: &[String]) -> f64 {
let lattice = Lattice::segment(input, dict);
self.value_over_lattice::<Viterbi, _>(&lattice, Viterbi).map_or(0.0, |v| v.0)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn deep_non_left_recursive_grammar_parses_in_bounded_time() {
let depth = 22;
let mut src = String::new();
for i in 0..depth {
src.push_str(&format!("r{i} := <r{}> \"a\" | <r{}> \"b\"\n", i + 1, i + 1));
}
src.push_str(&format!("r{depth} := \"x\"\n"));
let g = Grammar::parse(&src).expect("grammar parses");
let start = std::time::Instant::now();
assert!(g.parse_input(b"x").is_none(), "no alternative can complete");
let ms = start.elapsed().as_secs_f64() * 1000.0;
assert!(ms < 1000.0, "depth-{depth} descent took {ms:.1} ms; memoization is not engaged");
}
#[test]
fn left_recursion_still_grows_under_the_packrat_memo() {
let g = Grammar::parse(
"expr := <expr> \"+\" <term> | <term>\nterm := number\n",
)
.expect("grammar parses");
let node = g.parse_input(b"1 + 2 + 3").expect("parses");
assert_eq!(node.sexpr(), "(expr (expr 1 + 2) + 3)");
}
const ARITH: &str = r#"
expr := <expr> "+" <term> | <expr> "-" <term> | <term>
term := <term> "*" <factor> | <term> "/" <factor> | <factor>
factor := number | ident | "(" <expr> ")"
"#;
fn parse(grammar: &str, input: &str) -> Option<String> {
Grammar::parse(grammar).unwrap().parse_input(input.as_bytes()).map(|n| n.sexpr())
}
#[test]
fn precedence_is_respected() {
assert_eq!(parse(ARITH, "2 + 3 * 4").unwrap(), "(expr 2 + (term 3 * 4))");
}
#[test]
fn left_associativity() {
assert_eq!(parse(ARITH, "2 - 3 - 4").unwrap(), "(expr (expr 2 - 3) - 4)");
}
#[test]
fn parenthesized_grouping_overrides_precedence() {
assert_eq!(parse(ARITH, "(2 + 3) * 4").unwrap(), "(term (factor ( (expr 2 + 3) )) * 4)");
}
#[test]
fn the_same_grammar_parses_identifiers_any_language() {
assert_eq!(
parse(ARITH, "width * 2 + height").unwrap(),
"(expr (term width * 2) + height)"
);
}
#[test]
fn recursive_call_grammar() {
let g = r#"
call := ident "(" <args> ")"
args := <arg> "," <args> | <arg>
arg := <call> | ident | number
"#;
let tree = parse(g, "f(g(x), 3)").unwrap();
assert!(tree.contains("call"), "expected a call node, got {tree}");
assert!(tree.contains("g"), "expected the nested call, got {tree}");
}
#[test]
fn indirect_left_recursion_left_associates() {
let g = r#"
expr := <sum>
sum := <expr> "+" number | number
"#;
assert_eq!(parse(g, "1 + 2 + 3").unwrap(), "(sum (sum 1 + 2) + 3)");
assert_eq!(parse(g, "7").unwrap(), "7");
assert!(parse(g, "1 +").is_none());
}
#[test]
fn indirect_left_recursion_through_a_three_rule_cycle() {
let g = r#"
a := <b>
b := <c>
c := <a> "-" number | number
"#;
assert_eq!(parse(g, "9 - 4 - 1").unwrap(), "(c (c 9 - 4) - 1)");
}
#[test]
fn quantifiers_star_plus_opt() {
assert_eq!(parse("items := number*", "1 2 3").unwrap(), "(items 1 2 3)");
assert!(parse("items := number*", "x").is_none());
assert_eq!(parse("nums := number+", "7 8").unwrap(), "(nums 7 8)");
assert!(parse("nums := number+", "x").is_none());
assert_eq!(parse("signed := \"-\"? number", "-5").unwrap(), "(signed - 5)");
assert_eq!(parse("signed := \"-\"? number", "5").unwrap(), "5");
assert!(parse("signed := \"-\"? number", "5 6").is_none());
let call = "call := ident \"(\" number* \")\"";
assert_eq!(parse(call, "f(1 2 3)").unwrap(), "(call f ( 1 2 3 ))");
}
#[test]
fn bare_quantifier_is_a_grammar_error() {
let e = Grammar::parse("r := *").unwrap_err();
assert!(e.msg.contains("no symbol to repeat"), "got {}", e.msg);
}
#[test]
fn group_alternation_and_repetition() {
assert_eq!(parse("x := ( \"a\" | \"b\" )", "a").unwrap(), "a");
assert_eq!(parse("x := ( \"a\" | \"b\" )", "b").unwrap(), "b");
assert!(parse("x := ( \"a\" | \"b\" )", "c").is_none());
let call = "call := ident \"(\" ( number | ident )* \")\"";
assert_eq!(parse(call, "f(1 x 2)").unwrap(), "(call f ( 1 x 2 ))");
assert_eq!(parse(call, "g()").unwrap(), "(call g ( ))");
let pairs = "pairs := ( ident number )+";
assert_eq!(parse(pairs, "a 1 b 2").unwrap(), "(pairs a 1 b 2)");
assert!(parse(pairs, "a 1 b").is_none()); }
#[test]
fn left_recursion_behind_a_nullable_prefix_terminates() {
let g = "expr := <pre> <expr> \"+\" number | number\npre := \"neg\"?";
let tree = parse(g, "1 + 2 + 3").expect("must parse without stack overflow");
assert!(tree.contains('+'), "got {tree}");
}
#[test]
fn unterminated_or_stray_group_is_a_grammar_error() {
assert!(Grammar::parse("r := ( <a>").is_err());
assert!(Grammar::parse("r := <a> )").is_err());
}
#[test]
fn input_that_does_not_parse_returns_none() {
assert!(parse(ARITH, "2 +").is_none());
}
#[test]
fn unknown_symbol_is_a_grammar_error() {
let e = Grammar::parse("r := frobnicate").unwrap_err();
assert!(e.msg.contains("unknown symbol"));
}
#[test]
fn semiring_recognition_matches_the_parser() {
let g = Grammar::parse(ARITH).unwrap();
for input in ["2 + 3 * 4", "(2 + 3) * 4", "width * 2 + height"] {
assert!(g.recognizes(input.as_bytes()), "{input} should recognize");
assert!(g.parse_input(input.as_bytes()).is_some());
}
for input in ["2 +", "+ 3", ""] {
assert!(!g.recognizes(input.as_bytes()), "{input} should not recognize");
}
}
#[test]
fn semiring_counts_ambiguous_derivations() {
let g = Grammar::parse("expr := <expr> \"+\" <expr> | number").unwrap();
assert_eq!(g.count_parses(b"1"), 1);
assert_eq!(g.count_parses(b"1 + 2"), 1);
assert_eq!(g.count_parses(b"1 + 2 + 3"), 2); assert_eq!(g.count_parses(b"1 + 2 + 3 + 4"), 5); assert_eq!(g.count_parses(b"1 + 2 + 3 + 4 + 5"), 14); assert_eq!(g.count_parses(b"1 +"), 0); assert_eq!(g.recognizes(b"1 + 2 + 3"), g.count_parses(b"1 + 2 + 3") > 0);
assert_eq!(g.recognizes(b"1 +"), g.count_parses(b"1 +") > 0);
}
#[test]
fn semiring_counts_through_quantifiers_and_groups() {
let g = Grammar::parse("seq := ( ident )*").unwrap();
assert_eq!(g.count_parses(b"a b c"), 1);
assert_eq!(g.count_parses(b""), 1); assert_eq!(g.count_parses(b"1"), 0); }
#[test]
fn weighted_grammar_scores_parses() {
let g = Grammar::parse("expr := <expr> \"+\" <expr> @0.5 | number @0.5").unwrap();
let close = |a: f64, b: f64| (a - b).abs() < 1e-9;
assert!(close(g.total_prob(b"1"), 0.5));
assert!(close(g.best_parse_prob(b"1"), 0.5));
assert!(close(g.total_prob(b"1 + 2"), 0.125));
assert!(close(g.best_parse_prob(b"1 + 2 + 3"), 0.031_25));
assert!(close(g.total_prob(b"1 + 2 + 3"), 0.062_5));
assert!(close(g.total_prob(b"1 +"), 0.0));
}
#[test]
fn weighted_grammar_picks_the_more_probable_parse() {
let g = Grammar::parse(
"s := <add> @0.8 | <mul> @0.2\nadd := number \"x\" number\nmul := number \"x\" number",
)
.unwrap();
let close = |a: f64, b: f64| (a - b).abs() < 1e-9;
assert!(close(g.total_prob(b"1 x 2"), 1.0));
assert!(close(g.best_parse_prob(b"1 x 2"), 0.8));
assert_eq!(g.count_parses(b"1 x 2"), 2);
}
#[test]
fn bare_weight_with_no_number_is_a_grammar_error() {
assert!(Grammar::parse("r := number @x").is_err());
}
#[test]
fn grammar_parses_over_a_segmentation_lattice() {
let dict: Vec<String> = ["a", "b", "ab"].iter().map(|s| (*s).to_string()).collect();
let words = Grammar::parse("s := <s> ident | ident").unwrap();
assert_eq!(words.count_segmentations("abab", &dict), 4);
assert_eq!(words.count_segmentations("aa", &dict), 1); assert_eq!(words.count_segmentations("abc", &dict), 0);
let only_ab = Grammar::parse("s := \"ab\" <s> | \"ab\"").unwrap();
assert_eq!(only_ab.count_segmentations("abab", &dict), 1);
assert!(only_ab.best_segmentation_prob("abab", &dict) > 0.0);
}
}