use std::collections::BTreeMap;
use super::element::{GrammarElement, GrammarRule, GreType};
use super::error::GrammarError;
use super::utf8::{byte_at, decode_char};
pub const MAX_REPETITION_THRESHOLD: u64 = 2000;
pub trait GrammarVocab {
fn tokenize_special(&self, text: &str) -> Vec<u32>;
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct ParsedGrammar {
pub rules: Vec<GrammarRule>,
pub symbol_ids: BTreeMap<String, u32>,
}
impl ParsedGrammar {
pub fn symbol_id(&self, name: &str) -> Option<u32> {
self.symbol_ids.get(name).copied()
}
pub fn symbol_name(&self, rule_id: u32) -> Option<&str> {
self.symbol_ids
.iter()
.find(|(_, &v)| v == rule_id)
.map(|(k, _)| k.as_str())
}
}
pub fn parse(src: &str) -> Result<ParsedGrammar, GrammarError> {
parse_with_vocab(src, None)
}
pub fn parse_with_vocab(
src: &str,
vocab: Option<&dyn GrammarVocab>,
) -> Result<ParsedGrammar, GrammarError> {
let mut p = Parser {
src: src.as_bytes(),
pos: 0,
vocab,
out: ParsedGrammar::default(),
};
p.parse_all()?;
Ok(p.out)
}
struct Parser<'a> {
src: &'a [u8],
pos: usize,
vocab: Option<&'a dyn GrammarVocab>,
out: ParsedGrammar,
}
impl<'a> Parser<'a> {
#[inline]
fn at(&self, i: usize) -> u8 {
byte_at(self.src, i)
}
#[inline]
fn cur(&self) -> u8 {
self.at(self.pos)
}
fn err(&self, expected: impl Into<String>) -> GrammarError {
GrammarError::syntax(expected, self.src, self.pos)
}
fn err_at(&self, expected: impl Into<String>, offset: usize) -> GrammarError {
GrammarError::syntax(expected, self.src, offset)
}
fn is_digit_char(c: u8) -> bool {
c.is_ascii_digit()
}
fn is_word_char(c: u8) -> bool {
c.is_ascii_alphabetic() || c == b'-' || Self::is_digit_char(c)
}
fn parse_space(&mut self, newline_ok: bool) {
loop {
let c = self.cur();
if c == b' ' || c == b'\t' {
self.pos += 1;
} else if c == b'#' {
while self.cur() != 0 && self.cur() != b'\r' && self.cur() != b'\n' {
self.pos += 1;
}
} else if newline_ok && (c == b'\r' || c == b'\n') {
self.pos += 1;
} else {
return;
}
}
}
fn parse_name(&self) -> Result<usize, GrammarError> {
let mut end = self.pos;
while Self::is_word_char(self.at(end)) {
end += 1;
}
if end == self.pos {
return Err(self.err("expecting name"));
}
Ok(end)
}
fn parse_int_end(&self) -> Result<usize, GrammarError> {
let mut end = self.pos;
while Self::is_digit_char(self.at(end)) {
end += 1;
}
if end == self.pos {
return Err(self.err("expecting integer"));
}
Ok(end)
}
fn parse_u64(&mut self) -> Result<u64, GrammarError> {
let end = self.parse_int_end()?;
let text = std::str::from_utf8(&self.src[self.pos..end])
.map_err(|_| self.err("expecting integer"))?;
let value: u64 = text
.parse()
.map_err(|_| self.err_at("integer is too large", self.pos))?;
self.pos = end;
Ok(value)
}
fn parse_hex(&mut self, size: usize) -> Result<u32, GrammarError> {
let start = self.pos;
let end = start + size;
let mut value: u32 = 0;
let mut p = start;
while p < end && self.at(p) != 0 {
let c = self.at(p);
let digit = match c {
b'a'..=b'f' => c - b'a' + 10,
b'A'..=b'F' => c - b'A' + 10,
b'0'..=b'9' => c - b'0',
_ => break,
};
value = (value << 4) + digit as u32;
p += 1;
}
if p != end {
self.pos = p;
return Err(self.err_at(format!("expecting {size} hex chars"), start));
}
self.pos = p;
Ok(value)
}
fn parse_char(&mut self) -> Result<u32, GrammarError> {
if self.cur() == b'\\' {
let start = self.pos;
let next = self.at(self.pos + 1);
return match next {
b'x' => {
self.pos += 2;
self.parse_hex(2)
}
b'u' => {
self.pos += 2;
self.parse_hex(4)
}
b'U' => {
self.pos += 2;
self.parse_hex(8)
}
b't' => {
self.pos += 2;
Ok(u32::from(b'\t'))
}
b'r' => {
self.pos += 2;
Ok(u32::from(b'\r'))
}
b'n' => {
self.pos += 2;
Ok(u32::from(b'\n'))
}
b'\\' | b'"' | b'[' | b']' => {
self.pos += 2;
Ok(u32::from(next))
}
_ => Err(self.err_at("unknown escape", start)),
};
}
if self.cur() != 0 {
let (value, next) = decode_char(self.src, self.pos);
self.pos = next;
return Ok(value);
}
Err(self.err("unexpected end of input"))
}
fn parse_token(&mut self) -> Result<u32, GrammarError> {
let start = self.pos;
if self.cur() != b'<' {
return Err(self.err("expecting '<'"));
}
self.pos += 1;
if self.cur() == b'[' {
self.pos += 1;
let id = self.parse_u64()?;
let id = u32::try_from(id).map_err(|_| self.err_at("token id is too large", start))?;
if self.cur() != b']' {
return Err(self.err("expecting ']'"));
}
self.pos += 1;
if self.cur() != b'>' {
return Err(self.err("expecting '>'"));
}
self.pos += 1;
return Ok(id);
}
while self.cur() != 0 && self.cur() != b'>' {
self.pos += 1;
}
if self.cur() != b'>' {
return Err(self.err("expecting '>'"));
}
self.pos += 1;
let text = std::str::from_utf8(&self.src[start..self.pos])
.map_err(|_| self.err_at("token name is not valid UTF-8", start))?
.to_string();
let Some(vocab) = self.vocab else {
return Err(GrammarError::TokenNeedsVocabulary {
token: text,
offset: start,
});
};
let ids = vocab.tokenize_special(&text);
if ids.len() != 1 {
return Err(GrammarError::TokenNotSingle {
token: text,
n_tokens: ids.len(),
});
}
Ok(ids[0])
}
fn get_symbol_id(&mut self, name: &str) -> u32 {
let next_id = self.out.symbol_ids.len() as u32;
*self
.out
.symbol_ids
.entry(name.to_string())
.or_insert(next_id)
}
fn generate_symbol_id(&mut self, base_name: &str) -> u32 {
let next_id = self.out.symbol_ids.len() as u32;
self.out
.symbol_ids
.insert(format!("{base_name}_{next_id}"), next_id);
next_id
}
fn add_rule(&mut self, rule_id: u32, rule: GrammarRule) {
let idx = rule_id as usize;
if self.out.rules.len() <= idx {
self.out.rules.resize(idx + 1, GrammarRule::new());
}
self.out.rules[idx] = rule;
}
fn parse_alternates(
&mut self,
rule_name: &str,
rule_id: u32,
is_nested: bool,
) -> Result<(), GrammarError> {
let mut rule = GrammarRule::new();
self.parse_sequence(rule_name, &mut rule, is_nested)?;
while self.cur() == b'|' {
rule.push(GrammarElement::new(GreType::Alt, 0));
self.pos += 1;
self.parse_space(true);
self.parse_sequence(rule_name, &mut rule, is_nested)?;
}
rule.push(GrammarElement::new(GreType::End, 0));
self.add_rule(rule_id, rule);
Ok(())
}
fn handle_repetitions(
&mut self,
rule: &mut GrammarRule,
rule_name: &str,
last_sym_start: usize,
n_prev_rules: &mut u64,
min_times: u64,
max_times: Option<u64>,
) -> Result<(), GrammarError> {
let no_max = max_times.is_none();
if last_sym_start == rule.len() {
return Err(self.err("expecting preceding item to */+/?/{"));
}
let prev_rule: GrammarRule = rule[last_sym_start..].to_vec();
let mut total_rules: u64 = 1;
match max_times {
Some(max) if max > 0 => total_rules = max,
_ => {
if min_times > 0 {
total_rules = min_times;
}
}
}
let product = n_prev_rules.saturating_mul(total_rules);
if product >= MAX_REPETITION_THRESHOLD {
return Err(GrammarError::RepetitionTooLarge {
requested: product,
limit: MAX_REPETITION_THRESHOLD,
offset: self.pos,
});
}
if min_times == 0 {
rule.truncate(last_sym_start);
} else {
for _ in 1..min_times {
rule.extend_from_slice(&prev_rule);
}
}
let mut last_rec_rule_id: u32 = 0;
let n_opt = match max_times {
None => 1,
Some(max) => max.saturating_sub(min_times),
};
let mut rec_rule = prev_rule.clone();
for i in 0..n_opt {
rec_rule.truncate(prev_rule.len());
let rec_rule_id = self.generate_symbol_id(rule_name);
if i > 0 || no_max {
rec_rule.push(GrammarElement::new(
GreType::RuleRef,
if no_max {
rec_rule_id
} else {
last_rec_rule_id
},
));
}
rec_rule.push(GrammarElement::new(GreType::Alt, 0));
rec_rule.push(GrammarElement::new(GreType::End, 0));
self.add_rule(rec_rule_id, rec_rule.clone());
last_rec_rule_id = rec_rule_id;
}
if n_opt > 0 {
rule.push(GrammarElement::new(GreType::RuleRef, last_rec_rule_id));
}
*n_prev_rules = product;
Ok(())
}
fn parse_sequence(
&mut self,
rule_name: &str,
rule: &mut GrammarRule,
is_nested: bool,
) -> Result<(), GrammarError> {
let mut last_sym_start = rule.len();
let mut n_prev_rules: u64 = 1;
while self.cur() != 0 {
match self.cur() {
b'"' => {
self.pos += 1;
last_sym_start = rule.len();
n_prev_rules = 1;
while self.cur() != b'"' {
if self.cur() == 0 {
return Err(self.err("unexpected end of input"));
}
let value = self.parse_char()?;
rule.push(GrammarElement::new(GreType::Char, value));
}
self.pos += 1;
self.parse_space(is_nested);
}
b'[' => {
self.pos += 1;
let mut start_type = GreType::Char;
if self.cur() == b'^' {
self.pos += 1;
start_type = GreType::CharNot;
}
last_sym_start = rule.len();
n_prev_rules = 1;
while self.cur() != b']' {
if self.cur() == 0 {
return Err(self.err("unexpected end of input"));
}
let value = self.parse_char()?;
let gtype = if last_sym_start < rule.len() {
GreType::CharAlt
} else {
start_type
};
rule.push(GrammarElement::new(gtype, value));
if self.at(self.pos) == b'-' && self.at(self.pos + 1) != b']' {
if self.at(self.pos + 1) == 0 {
return Err(self.err("unexpected end of input"));
}
self.pos += 1;
let endchar = self.parse_char()?;
rule.push(GrammarElement::new(GreType::CharRngUpper, endchar));
}
}
self.pos += 1;
self.parse_space(is_nested);
}
b'<' | b'!' => {
let mut gtype = GreType::Token;
if self.cur() == b'!' {
gtype = GreType::TokenNot;
self.pos += 1;
}
let token_id = self.parse_token()?;
last_sym_start = rule.len();
n_prev_rules = 1;
rule.push(GrammarElement::new(gtype, token_id));
self.parse_space(is_nested);
}
c if Self::is_word_char(c) => {
let name_end = self.parse_name()?;
let name = std::str::from_utf8(&self.src[self.pos..name_end])
.map_err(|_| self.err("rule name is not valid UTF-8"))?
.to_string();
let ref_rule_id = self.get_symbol_id(&name);
self.pos = name_end;
self.parse_space(is_nested);
last_sym_start = rule.len();
n_prev_rules = 1;
rule.push(GrammarElement::new(GreType::RuleRef, ref_rule_id));
}
b'(' => {
self.pos += 1;
self.parse_space(true);
let n_rules_before = self.out.symbol_ids.len() as u64;
let sub_rule_id = self.generate_symbol_id(rule_name);
self.parse_alternates(rule_name, sub_rule_id, true)?;
n_prev_rules = (self.out.symbol_ids.len() as u64 - n_rules_before).max(1);
last_sym_start = rule.len();
rule.push(GrammarElement::new(GreType::RuleRef, sub_rule_id));
if self.cur() != b')' {
return Err(self.err("expecting ')'"));
}
self.pos += 1;
self.parse_space(is_nested);
}
b'.' => {
last_sym_start = rule.len();
n_prev_rules = 1;
rule.push(GrammarElement::new(GreType::CharAny, 0));
self.pos += 1;
self.parse_space(is_nested);
}
b'*' => {
self.pos += 1;
self.parse_space(is_nested);
self.handle_repetitions(
rule,
rule_name,
last_sym_start,
&mut n_prev_rules,
0,
None,
)?;
}
b'+' => {
self.pos += 1;
self.parse_space(is_nested);
self.handle_repetitions(
rule,
rule_name,
last_sym_start,
&mut n_prev_rules,
1,
None,
)?;
}
b'?' => {
self.pos += 1;
self.parse_space(is_nested);
self.handle_repetitions(
rule,
rule_name,
last_sym_start,
&mut n_prev_rules,
0,
Some(1),
)?;
}
b'{' => {
self.pos += 1;
self.parse_space(is_nested);
if !Self::is_digit_char(self.cur()) {
return Err(self.err("expecting an int"));
}
let min_times = self.parse_u64()?;
self.parse_space(is_nested);
let mut max_times: Option<u64> = None;
if self.cur() == b'}' {
max_times = Some(min_times);
self.pos += 1;
self.parse_space(is_nested);
} else if self.cur() == b',' {
self.pos += 1;
self.parse_space(is_nested);
if Self::is_digit_char(self.cur()) {
max_times = Some(self.parse_u64()?);
self.parse_space(is_nested);
}
if self.cur() != b'}' {
return Err(self.err("expecting '}'"));
}
self.pos += 1;
self.parse_space(is_nested);
} else {
return Err(self.err("expecting ','"));
}
if min_times > MAX_REPETITION_THRESHOLD
|| max_times.is_some_and(|m| m > MAX_REPETITION_THRESHOLD)
{
return Err(GrammarError::RepetitionTooLarge {
requested: max_times.unwrap_or(min_times),
limit: MAX_REPETITION_THRESHOLD,
offset: self.pos,
});
}
self.handle_repetitions(
rule,
rule_name,
last_sym_start,
&mut n_prev_rules,
min_times,
max_times,
)?;
}
_ => break,
}
}
Ok(())
}
fn parse_rule(&mut self) -> Result<(), GrammarError> {
let name_end = self.parse_name()?;
let name = std::str::from_utf8(&self.src[self.pos..name_end])
.map_err(|_| self.err("rule name is not valid UTF-8"))?
.to_string();
self.pos = name_end;
self.parse_space(false);
let rule_id = self.get_symbol_id(&name);
if !(self.cur() == b':' && self.at(self.pos + 1) == b':' && self.at(self.pos + 2) == b'=') {
return Err(self.err("expecting ::="));
}
self.pos += 3;
self.parse_space(true);
self.parse_alternates(&name, rule_id, false)?;
if self.cur() == b'\r' {
self.pos += if self.at(self.pos + 1) == b'\n' { 2 } else { 1 };
} else if self.cur() == b'\n' {
self.pos += 1;
} else if self.cur() != 0 {
return Err(self.err("expecting newline or end"));
}
self.parse_space(true);
Ok(())
}
fn parse_all(&mut self) -> Result<(), GrammarError> {
self.parse_space(true);
while self.cur() != 0 {
self.parse_rule()?;
}
for id in 0..self.out.rules.len() {
if self.out.rules[id].is_empty() {
return Err(self.undefined(id as u32));
}
}
for rule in &self.out.rules {
for elem in rule {
if elem.gtype == GreType::RuleRef {
let idx = elem.value as usize;
if idx >= self.out.rules.len() || self.out.rules[idx].is_empty() {
return Err(self.undefined(elem.value));
}
}
}
}
Ok(())
}
fn undefined(&self, rule_id: u32) -> GrammarError {
GrammarError::UndefinedRule {
name: self
.out
.symbol_name(rule_id)
.unwrap_or("<unnamed>")
.to_string(),
rule_id,
}
}
}