use std::collections::HashMap;
use std::sync::Arc;
use anyhow::{Result, bail, ensure};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Elem {
End,
Alt,
RuleRef(u32),
Char(u8),
CharRngUpper(u8),
CharNot(u8),
CharAlt(u8),
}
impl Elem {
fn is_end_of_seq(self) -> bool {
matches!(self, Elem::End | Elem::Alt)
}
fn char_value(self) -> u8 {
match self {
Elem::Char(c) | Elem::CharRngUpper(c) | Elem::CharNot(c) | Elem::CharAlt(c) => c,
_ => 0,
}
}
}
#[derive(Debug, Clone)]
pub struct Grammar {
rules: Vec<Vec<Elem>>,
root: u32,
}
impl Grammar {
pub fn parse(src: &str) -> Result<Grammar> {
Parser::new(src).parse_grammar()
}
#[inline]
fn elem(&self, p: Pos) -> Elem {
self.rules[p.rule as usize][p.idx as usize]
}
#[inline]
fn is_end_of_seq(&self, p: Pos) -> bool {
let rule = &self.rules[p.rule as usize];
p.idx as usize >= rule.len() || rule[p.idx as usize].is_end_of_seq()
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct Pos {
rule: u32,
idx: u32,
}
#[derive(Clone)]
pub struct GrammarState {
grammar: Arc<Grammar>,
stacks: Vec<Vec<Pos>>,
}
impl GrammarState {
pub fn new(grammar: Arc<Grammar>) -> Self {
let mut stacks = Vec::new();
let root = grammar.root;
let mut i = 0u32;
loop {
let pos = Pos { rule: root, idx: i };
let mut stack = Vec::new();
if !grammar.is_end_of_seq(pos) {
stack.push(pos);
}
advance_stack(&grammar, stack, &mut stacks);
while !grammar.is_end_of_seq(Pos { rule: root, idx: i }) {
i += 1;
}
if grammar.elem(Pos { rule: root, idx: i }) == Elem::Alt {
i += 1; } else {
break; }
}
GrammarState { grammar, stacks }
}
pub fn is_complete(&self) -> bool {
self.stacks.iter().any(|s| s.is_empty())
}
pub fn is_dead(&self) -> bool {
self.stacks.is_empty()
}
pub fn accepts(&self, bytes: &[u8]) -> bool {
if bytes.is_empty() {
return false;
}
let mut stacks = self.stacks.clone();
for &b in bytes {
stacks = step(&self.grammar, &stacks, b);
if stacks.is_empty() {
return false;
}
}
true
}
pub fn accept(&mut self, bytes: &[u8]) {
for &b in bytes {
let next = step(&self.grammar, &self.stacks, b);
if next.is_empty() {
self.stacks = next;
return;
}
self.stacks = next;
}
}
}
fn advance_stack(g: &Grammar, stack: Vec<Pos>, out: &mut Vec<Vec<Pos>>) {
let Some(&top) = stack.last() else {
if !out.contains(&stack) {
out.push(stack);
}
return;
};
match g.elem(top) {
Elem::RuleRef(rule_id) => {
let mut i = 0u32;
loop {
let alt_start = Pos {
rule: rule_id,
idx: i,
};
let mut new_stack = stack[..stack.len() - 1].to_vec();
let cont = Pos {
rule: top.rule,
idx: top.idx + 1,
};
if !g.is_end_of_seq(cont) {
new_stack.push(cont);
}
if !g.is_end_of_seq(alt_start) {
new_stack.push(alt_start);
}
advance_stack(g, new_stack, out);
while !g.is_end_of_seq(Pos {
rule: rule_id,
idx: i,
}) {
i += 1;
}
if g.elem(Pos {
rule: rule_id,
idx: i,
}) == Elem::Alt
{
i += 1;
} else {
break;
}
}
}
Elem::Char(_) | Elem::CharNot(_) if !out.contains(&stack) => out.push(stack),
_ => {}
}
}
fn step(g: &Grammar, stacks: &[Vec<Pos>], b: u8) -> Vec<Vec<Pos>> {
let mut out = Vec::new();
for stack in stacks {
let Some(&top) = stack.last() else {
continue; };
if let Some(next_idx) = match_char(&g.rules[top.rule as usize], top.idx, b) {
let mut new_stack = stack[..stack.len() - 1].to_vec();
let next = Pos {
rule: top.rule,
idx: next_idx,
};
if !g.is_end_of_seq(next) {
new_stack.push(next);
}
advance_stack(g, new_stack, &mut out);
}
}
out
}
fn match_char(rule: &[Elem], idx: u32, b: u8) -> Option<u32> {
let mut i = idx as usize;
let negated = matches!(rule[i], Elem::CharNot(_));
let mut found = false;
loop {
let lo = rule[i].char_value();
let hi = if i + 1 < rule.len()
&& let Elem::CharRngUpper(h) = rule[i + 1]
{
i += 1;
h
} else {
lo
};
if b >= lo && b <= hi {
found = true;
}
if i + 1 < rule.len() && matches!(rule[i + 1], Elem::CharAlt(_)) {
i += 1;
continue;
}
break;
}
let matched = found != negated;
if matched { Some(i as u32 + 1) } else { None }
}
pub struct GrammarMask {
token_bytes: Vec<Vec<u8>>,
eos: Option<u32>,
special: Vec<bool>,
}
impl GrammarMask {
pub fn new(token_bytes: Vec<Vec<u8>>, eos: Option<u32>, special: Vec<bool>) -> Self {
assert_eq!(
token_bytes.len(),
special.len(),
"GrammarMask: token_bytes and special must have the same (vocab) length"
);
GrammarMask {
token_bytes,
eos,
special,
}
}
pub fn vocab_size(&self) -> usize {
self.token_bytes.len()
}
pub fn token_bytes(&self, id: u32) -> &[u8] {
self.token_bytes
.get(id as usize)
.map_or(&[], |v| v.as_slice())
}
pub fn apply(&self, state: &GrammarState, logits: &mut [f32]) -> usize {
let n = self.token_bytes.len().min(logits.len());
let mut allowed = 0usize;
let complete = state.is_complete();
for (id, logit) in logits.iter_mut().take(n).enumerate() {
let ok = if Some(id as u32) == self.eos {
complete
} else if self.special[id] {
false
} else {
state.accepts(&self.token_bytes[id])
};
if ok {
allowed += 1;
} else {
*logit = f32::NEG_INFINITY;
}
}
for l in logits.iter_mut().skip(n) {
*l = f32::NEG_INFINITY;
}
allowed
}
}
struct Parser<'a> {
src: &'a [u8],
pos: usize,
rules: Vec<Vec<Elem>>,
symbol_ids: HashMap<String, u32>,
defined: Vec<bool>,
}
impl<'a> Parser<'a> {
fn new(src: &'a str) -> Self {
Parser {
src: src.as_bytes(),
pos: 0,
rules: Vec::new(),
symbol_ids: HashMap::new(),
defined: Vec::new(),
}
}
fn peek(&self) -> Option<u8> {
self.src.get(self.pos).copied()
}
fn bump(&mut self) -> Option<u8> {
let c = self.peek();
if c.is_some() {
self.pos += 1;
}
c
}
fn symbol_id(&mut self, name: &str) -> u32 {
if let Some(&id) = self.symbol_ids.get(name) {
return id;
}
let id = self.rules.len() as u32;
self.rules.push(Vec::new());
self.defined.push(false);
self.symbol_ids.insert(name.to_string(), id);
id
}
fn anon_rule(&mut self) -> u32 {
let id = self.rules.len() as u32;
self.rules.push(Vec::new());
self.defined.push(true);
id
}
fn skip_ws(&mut self) {
loop {
match self.peek() {
Some(b' ' | b'\t' | b'\r' | b'\n') => {
self.pos += 1;
}
Some(b'#') => {
while let Some(c) = self.peek() {
self.pos += 1;
if c == b'\n' {
break;
}
}
}
_ => break,
}
}
}
fn is_word_byte(c: u8) -> bool {
c.is_ascii_alphanumeric() || c == b'-' || c == b'_'
}
fn parse_name(&mut self) -> Result<String> {
let start = self.pos;
while let Some(c) = self.peek() {
if Self::is_word_byte(c) {
self.pos += 1;
} else {
break;
}
}
ensure!(self.pos > start, "expected a rule name at byte {}", start);
Ok(String::from_utf8_lossy(&self.src[start..self.pos]).into_owned())
}
fn looks_like_rule_start(&self) -> bool {
let mut i = self.pos;
let s = self.src;
if i >= s.len() || !Self::is_word_byte(s[i]) {
return false;
}
while i < s.len() && Self::is_word_byte(s[i]) {
i += 1;
}
while i < s.len() && matches!(s[i], b' ' | b'\t' | b'\r' | b'\n') {
i += 1;
}
s[i..].starts_with(b"::=")
}
fn parse_grammar(mut self) -> Result<Grammar> {
self.skip_ws();
while self.peek().is_some() {
self.parse_rule()?;
self.skip_ws();
}
let root = *self
.symbol_ids
.get("root")
.ok_or_else(|| anyhow::anyhow!("grammar has no `root` rule"))?;
for (name, &id) in &self.symbol_ids {
ensure!(self.defined[id as usize], "undefined rule: `{name}`");
}
self.check_no_left_recursion()?;
Ok(Grammar {
rules: self.rules,
root,
})
}
fn alternates(&self, rule: usize) -> Vec<&[Elem]> {
let elems = &self.rules[rule];
let mut alts = Vec::new();
let mut start = 0usize;
for (i, e) in elems.iter().enumerate() {
match e {
Elem::Alt => {
alts.push(&elems[start..i]);
start = i + 1;
}
Elem::End => {
alts.push(&elems[start..i]);
break;
}
_ => {}
}
}
alts
}
fn compute_nullable(&self) -> Vec<bool> {
let n = self.rules.len();
let mut nullable = vec![false; n];
loop {
let mut changed = false;
for r in 0..n {
if nullable[r] {
continue;
}
let any = self.alternates(r).iter().any(|alt| {
alt.iter().all(|e| match e {
Elem::RuleRef(b) => nullable[*b as usize],
_ => false,
})
});
if any {
nullable[r] = true;
changed = true;
}
}
if !changed {
break;
}
}
nullable
}
fn rule_label(&self, id: u32) -> String {
match self.symbol_ids.iter().find(|&(_, &v)| v == id) {
Some((name, _)) => format!("rule `{name}`"),
None => "an anonymous (…)/repetition subrule".to_string(),
}
}
fn check_no_left_recursion(&self) -> Result<()> {
let n = self.rules.len();
let nullable = self.compute_nullable();
let mut adj: Vec<Vec<u32>> = vec![Vec::new(); n];
for (a, adj_a) in adj.iter_mut().enumerate() {
for alt in self.alternates(a) {
for &e in alt {
match e {
Elem::RuleRef(b) => {
adj_a.push(b);
if !nullable[b as usize] {
break; }
}
_ => break,
}
}
}
}
#[derive(Clone, Copy, PartialEq)]
enum Color {
White,
Gray,
Black,
}
let mut color = vec![Color::White; n];
for s in 0..n {
if color[s] != Color::White {
continue;
}
color[s] = Color::Gray;
let mut stack: Vec<(usize, usize)> = vec![(s, 0)];
while let Some(&(node, ci)) = stack.last() {
if ci < adj[node].len() {
stack.last_mut().unwrap().1 += 1;
let b = adj[node][ci] as usize;
match color[b] {
Color::Gray => bail!(
"grammar is left-recursive: {} can recurse without consuming input",
self.rule_label(b as u32)
),
Color::White => {
color[b] = Color::Gray;
stack.push((b, 0));
}
Color::Black => {}
}
} else {
color[node] = Color::Black;
stack.pop();
}
}
}
Ok(())
}
fn parse_rule(&mut self) -> Result<()> {
let name = self.parse_name()?;
self.skip_ws();
ensure!(
self.src[self.pos..].starts_with(b"::="),
"expected `::=` after rule `{name}`"
);
self.pos += 3;
self.skip_ws();
let id = self.symbol_id(&name);
ensure!(!self.defined[id as usize], "rule `{name}` defined twice");
let elems = self.parse_alternates(&name)?;
self.rules[id as usize] = elems;
self.defined[id as usize] = true;
Ok(())
}
fn parse_alternates(&mut self, rule_name: &str) -> Result<Vec<Elem>> {
let mut out = Vec::new();
self.parse_sequence(&mut out, rule_name)?;
self.skip_ws();
while self.peek() == Some(b'|') {
self.pos += 1;
self.skip_ws();
out.push(Elem::Alt);
self.parse_sequence(&mut out, rule_name)?;
self.skip_ws();
}
out.push(Elem::End);
Ok(out)
}
fn parse_sequence(&mut self, out: &mut Vec<Elem>, rule_name: &str) -> Result<()> {
loop {
self.skip_ws();
match self.peek() {
None => break,
Some(b'|') | Some(b')') => break,
_ if self.looks_like_rule_start() => break,
_ => {}
}
let last_start = out.len();
match self.peek().unwrap() {
b'"' => self.parse_string_literal(out)?,
b'[' => self.parse_char_class(out)?,
b'(' => {
self.pos += 1; let sub = self.anon_rule();
let body = self.parse_alternates(rule_name)?;
self.rules[sub as usize] = body;
self.skip_ws();
ensure!(
self.peek() == Some(b')'),
"unclosed `(` in rule `{rule_name}`"
);
self.pos += 1; out.push(Elem::RuleRef(sub));
}
c if Self::is_word_byte(c) => {
let name = self.parse_name()?;
let id = self.symbol_id(&name);
out.push(Elem::RuleRef(id));
}
c => bail!("unexpected `{}` in rule `{rule_name}`", c as char),
}
if let Some(op @ (b'*' | b'+' | b'?')) = self.peek() {
self.pos += 1;
self.apply_repetition(out, last_start, op)?;
}
}
Ok(())
}
fn apply_repetition(&mut self, out: &mut Vec<Elem>, start: usize, op: u8) -> Result<()> {
ensure!(
start < out.len(),
"repetition operator `{}` has nothing to repeat",
op as char
);
let unit: Vec<Elem> = out.split_off(start);
let sub = self.anon_rule();
let mut body = Vec::new();
match op {
b'*' => {
body.extend_from_slice(&unit);
body.push(Elem::RuleRef(sub));
body.push(Elem::Alt);
body.push(Elem::End);
}
b'+' => {
body.extend_from_slice(&unit);
body.push(Elem::RuleRef(sub));
body.push(Elem::Alt);
body.extend_from_slice(&unit);
body.push(Elem::End);
}
b'?' => {
body.extend_from_slice(&unit);
body.push(Elem::Alt);
body.push(Elem::End);
}
_ => unreachable!(),
}
self.rules[sub as usize] = body;
out.push(Elem::RuleRef(sub));
Ok(())
}
fn parse_string_literal(&mut self, out: &mut Vec<Elem>) -> Result<()> {
self.pos += 1; loop {
match self.peek() {
None => bail!("unterminated string literal"),
Some(b'"') => {
self.pos += 1;
break;
}
Some(b'\\') => {
self.pos += 1;
for b in self.parse_escape()? {
out.push(Elem::Char(b));
}
}
Some(_) => {
let b = self.bump().unwrap();
out.push(Elem::Char(b));
}
}
}
Ok(())
}
fn parse_char_class(&mut self, out: &mut Vec<Elem>) -> Result<()> {
self.pos += 1; let negated = self.peek() == Some(b'^');
if negated {
self.pos += 1;
}
let mut first = true;
loop {
match self.peek() {
None => bail!("unterminated char class"),
Some(b']') => {
self.pos += 1;
break;
}
_ => {}
}
let lo = self.parse_class_char()?;
if first {
out.push(if negated {
Elem::CharNot(lo)
} else {
Elem::Char(lo)
});
first = false;
} else {
out.push(Elem::CharAlt(lo));
}
if self.peek() == Some(b'-') && self.src.get(self.pos + 1) != Some(&b']') {
self.pos += 1; let hi = self.parse_class_char()?;
out.push(Elem::CharRngUpper(hi));
}
}
ensure!(!first, "empty char class `[]`");
Ok(())
}
fn parse_class_char(&mut self) -> Result<u8> {
match self.peek() {
Some(b'\\') => {
self.pos += 1;
let bytes = self.parse_escape()?;
ensure!(
bytes.len() == 1,
"multi-byte escape not allowed inside a char class (byte-level v1)"
);
Ok(bytes[0])
}
Some(c) if c < 0x80 => {
self.pos += 1;
Ok(c)
}
Some(c) => {
bail!("non-ASCII byte 0x{c:02x} in char class not supported (byte-level v1)")
}
None => bail!("unterminated char class"),
}
}
fn parse_escape(&mut self) -> Result<Vec<u8>> {
let c = self
.bump()
.ok_or_else(|| anyhow::anyhow!("dangling escape `\\`"))?;
Ok(match c {
b'n' => vec![b'\n'],
b'r' => vec![b'\r'],
b't' => vec![b'\t'],
b'\\' => vec![b'\\'],
b'"' => vec![b'"'],
b'\'' => vec![b'\''],
b']' => vec![b']'],
b'[' => vec![b'['],
b'-' => vec![b'-'],
b'x' => {
let h = self.take_hex(2)?;
vec![h as u8]
}
b'u' => {
let cp = self.take_hex(4)?;
let ch = char::from_u32(cp)
.ok_or_else(|| anyhow::anyhow!("invalid \\u escape: {cp:#x}"))?;
ch.to_string().into_bytes()
}
other => bail!("unknown escape `\\{}`", other as char),
})
}
fn take_hex(&mut self, n: usize) -> Result<u32> {
let mut v = 0u32;
for _ in 0..n {
let c = self
.bump()
.ok_or_else(|| anyhow::anyhow!("short hex escape"))?;
let d = (c as char)
.to_digit(16)
.ok_or_else(|| anyhow::anyhow!("bad hex digit `{}`", c as char))?;
v = v * 16 + d;
}
Ok(v)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn grammar(src: &str) -> Arc<Grammar> {
Arc::new(Grammar::parse(src).expect("grammar should parse"))
}
fn run(g: &Arc<Grammar>, input: &[u8]) -> Option<GrammarState> {
let mut st = GrammarState::new(g.clone());
for &b in input {
if !st.accepts(&[b]) {
return None;
}
st.accept(&[b]);
}
Some(st)
}
#[test]
fn literal_sequence() {
let g = grammar(r#"root ::= "ab""#);
let st = run(&g, b"ab").unwrap();
assert!(st.is_complete());
assert!(run(&g, b"ac").is_none());
let part = run(&g, b"a").unwrap();
assert!(!part.is_complete());
}
#[test]
fn alternation_and_class() {
let g = grammar(r#"root ::= "yes" | "no" | [0-9]"#);
assert!(run(&g, b"yes").unwrap().is_complete());
assert!(run(&g, b"no").unwrap().is_complete());
assert!(run(&g, b"7").unwrap().is_complete());
assert!(run(&g, b"maybe").is_none());
}
#[test]
fn repetition_star_plus_opt() {
let star = grammar(r#"root ::= "a"*"#);
assert!(GrammarState::new(star.clone()).is_complete()); assert!(run(&star, b"aaaa").unwrap().is_complete());
let plus = grammar(r#"root ::= "a"+"#);
assert!(!GrammarState::new(plus.clone()).is_complete()); assert!(run(&plus, b"a").unwrap().is_complete());
assert!(run(&plus, b"aaa").unwrap().is_complete());
let opt = grammar(r#"root ::= "a"? "b""#);
assert!(run(&opt, b"ab").unwrap().is_complete());
assert!(run(&opt, b"b").unwrap().is_complete());
assert!(run(&opt, b"aab").is_none());
}
#[test]
fn negated_class_and_groups() {
let g = grammar(r#"root ::= "\"" ([^"\\])* "\"""#);
assert!(run(&g, br#""hello""#).unwrap().is_complete());
assert!(run(&g, br#""""#).unwrap().is_complete());
assert!(run(&g, br#""a"b""#).is_none());
}
#[test]
fn rule_references_and_recursion() {
let g = grammar(
r#"
root ::= list
list ::= "[" items "]"
items ::= digit ("," digit)*
digit ::= [0-9]
"#,
);
assert!(run(&g, b"[1,2,3]").unwrap().is_complete());
assert!(run(&g, b"[5]").unwrap().is_complete());
assert!(run(&g, b"[1,]").is_none());
assert!(run(&g, b"[]").is_none());
}
#[test]
fn errors_not_panics() {
assert!(Grammar::parse("foo ::= bar").is_err()); assert!(Grammar::parse(r#"root ::= "unterminated"#).is_err());
assert!(Grammar::parse("root ::= ").is_ok()); assert!(Grammar::parse(r#"root ::= ("a""#).is_err()); }
#[test]
fn rejects_left_recursion_no_stack_overflow() {
assert!(Grammar::parse(r#"root ::= ""*"#).is_err());
assert!(Grammar::parse(r#"root ::= ""+"#).is_err());
assert!(Grammar::parse(r#"root ::= root "x" | "y""#).is_err());
assert!(Grammar::parse("root ::= a\na ::= b\nb ::= a").is_err());
assert!(Grammar::parse("root ::= n*\nn ::= \"a\"?").is_err());
assert!(Grammar::parse(r#"root ::= "a" root | "a""#).is_ok());
assert!(Grammar::parse("root ::= ws \"x\"\nws ::= \" \"*").is_ok());
}
#[test]
fn mask_gates_eos_and_specials() {
let g = grammar(r#"root ::= "hi""#);
let token_bytes = vec![b"h".to_vec(), b"i".to_vec(), b"x".to_vec(), vec![], vec![]];
let special = vec![false, false, false, true, true];
let mask = GrammarMask::new(token_bytes, Some(3), special);
let mut st = GrammarState::new(g.clone());
let mut logits = vec![0.0f32; 5];
let allowed = mask.apply(&st, &mut logits);
assert_eq!(allowed, 1); assert_eq!(logits[0], 0.0);
assert_eq!(logits[1], f32::NEG_INFINITY); assert_eq!(logits[3], f32::NEG_INFINITY);
st.accept(b"h");
st.accept(b"i");
let mut logits = vec![0.0f32; 5];
let allowed = mask.apply(&st, &mut logits);
assert_eq!(allowed, 1); assert_eq!(logits[3], 0.0); assert_eq!(logits[0], f32::NEG_INFINITY);
}
}