use super::element::{GrammarRule, GrammarStack, GreType};
use super::error::GrammarError;
use super::machine::{advance_stack, elem, match_char, match_partial_char, match_token, Grammar};
use super::utf8::{decode_piece, PartialUtf8};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Candidate<'a> {
pub index: usize,
pub id: u32,
pub piece: &'a [u8],
}
impl<'a> Candidate<'a> {
pub fn new(index: usize, id: u32, piece: &'a [u8]) -> Self {
Self { index, id, piece }
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct DecodedPiece {
pub code_points: Vec<u32>,
pub partial_utf8: PartialUtf8,
}
pub fn decode_for(grammar: &Grammar, piece: &[u8]) -> DecodedPiece {
let (code_points, partial_utf8) = decode_piece(piece, grammar.partial_utf8());
DecodedPiece {
code_points,
partial_utf8,
}
}
#[derive(Debug, Clone, Copy)]
struct Cand {
index: usize,
id: u32,
slot: usize,
off: usize,
partial_utf8: PartialUtf8,
}
pub fn reject_candidates(
grammar: &Grammar,
candidates: &[Candidate<'_>],
) -> Result<Vec<usize>, GrammarError> {
if grammar.is_awaiting_trigger() {
return Err(GrammarError::AwaitingTrigger);
}
if candidates.is_empty() {
return Ok(Vec::new());
}
if grammar.stacks().is_empty() {
return Ok(candidates.iter().map(|c| c.index).collect());
}
let mut arena: Vec<Vec<u32>> = Vec::with_capacity(candidates.len());
let mut cands: Vec<Cand> = Vec::with_capacity(candidates.len());
for c in candidates {
let (code_points, partial_utf8) = decode_piece(c.piece, grammar.partial_utf8());
arena.push(code_points);
cands.push(Cand {
index: c.index,
id: c.id,
slot: arena.len() - 1,
off: 0,
partial_utf8,
});
}
let rejects = reject_over_stacks(grammar.rules(), grammar.stacks(), &arena, &cands)?;
Ok(rejects.into_iter().map(|c| c.index).collect())
}
pub fn accepts_token(grammar: &Grammar, id: u32, piece: &[u8]) -> Result<bool, GrammarError> {
let c = [Candidate::new(0, id, piece)];
Ok(reject_candidates(grammar, &c)?.is_empty())
}
fn reject_over_stacks(
rules: &[GrammarRule],
stacks: &[GrammarStack],
arena: &[Vec<u32>],
candidates: &[Cand],
) -> Result<Vec<Cand>, GrammarError> {
if stacks.is_empty() {
return Err(GrammarError::Internal(
"reject_over_stacks called with no stacks",
));
}
if candidates.is_empty() {
return Ok(Vec::new());
}
let mut rejects = reject_for_stack(rules, &stacks[0], arena, candidates)?;
for stack in &stacks[1..] {
if rejects.is_empty() {
break;
}
rejects = reject_for_stack(rules, stack, arena, &rejects)?;
}
Ok(rejects)
}
fn reject_for_stack(
rules: &[GrammarRule],
stack: &GrammarStack,
arena: &[Vec<u32>],
candidates: &[Cand],
) -> Result<Vec<Cand>, GrammarError> {
let mut rejects: Vec<Cand> = Vec::with_capacity(candidates.len());
let Some(&stack_pos) = stack.last() else {
for tok in candidates {
if arena[tok.slot][tok.off] != 0 || tok.partial_utf8.n_remain != 0 {
rejects.push(*tok);
}
}
return Ok(rejects);
};
let stack_elem = elem(rules, stack_pos);
if matches!(stack_elem.gtype, GreType::Token | GreType::TokenNot) {
for tok in candidates {
if arena[tok.slot][tok.off] == 0 {
if tok.partial_utf8.n_remain != 0 {
rejects.push(*tok);
}
} else if !match_token(stack_elem, tok.id) {
rejects.push(*tok);
}
}
return Ok(rejects);
}
let mut next_candidates: Vec<Cand> = Vec::with_capacity(candidates.len());
for tok in candidates {
if arena[tok.slot][tok.off] == 0 {
if tok.partial_utf8.n_remain != 0
&& !match_partial_char(rules, stack_pos, tok.partial_utf8)?
{
rejects.push(*tok);
}
} else if match_char(rules, stack_pos, arena[tok.slot][tok.off])?.0 {
let mut advanced = *tok;
advanced.off += 1;
next_candidates.push(advanced);
} else {
rejects.push(*tok);
}
}
let (_, stack_pos_after) = match_char(rules, stack_pos, 0)?;
let mut stack_after = stack[..stack.len() - 1].to_vec();
if !elem(rules, stack_pos_after).is_end_of_sequence() {
stack_after.push(stack_pos_after);
}
let mut next_stacks: Vec<GrammarStack> = Vec::new();
advance_stack(rules, &stack_after, &mut next_stacks)?;
let next_rejects = reject_over_stacks(rules, &next_stacks, arena, &next_candidates)?;
for tok in next_rejects {
let mut back = tok;
back.off -= 1;
rejects.push(back);
}
Ok(rejects)
}