use std::fmt;
use fancy_regex::Regex;
use super::error::GrammarError;
#[derive(Clone)]
pub struct TriggerPattern {
pattern: String,
search: Regex,
full: Option<Regex>,
}
impl TriggerPattern {
pub fn new(pattern: &str) -> Result<Self, GrammarError> {
let search = compile(pattern)?;
let full = if pattern.starts_with('^') && pattern.ends_with('$') {
Some(compile(&format!(r"\A(?:{pattern})\z"))?)
} else {
None
};
Ok(Self {
pattern: pattern.to_string(),
search,
full,
})
}
pub fn pattern(&self) -> &str {
&self.pattern
}
pub fn find(&self, input: &str) -> Result<Option<usize>, GrammarError> {
if let Some(full) = &self.full {
if let Some(caps) = self.captures(full, input)? {
return Ok(Some(start_of_match(&caps)));
}
}
match self.captures(&self.search, input)? {
Some(caps) => Ok(Some(start_of_match(&caps))),
None => Ok(None),
}
}
fn captures<'t>(
&self,
re: &Regex,
input: &'t str,
) -> Result<Option<fancy_regex::Captures<'t>>, GrammarError> {
re.captures(input)
.map_err(|e| GrammarError::TriggerPatternFailed {
pattern: self.pattern.clone(),
reason: e.to_string(),
})
}
}
impl fmt::Debug for TriggerPattern {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_tuple("TriggerPattern")
.field(&self.pattern)
.finish()
}
}
impl PartialEq for TriggerPattern {
fn eq(&self, other: &Self) -> bool {
self.pattern == other.pattern
}
}
impl Eq for TriggerPattern {}
fn compile(pattern: &str) -> Result<Regex, GrammarError> {
Regex::new(pattern).map_err(|e| GrammarError::TriggerPatternInvalid {
pattern: pattern.to_string(),
reason: e.to_string(),
})
}
fn start_of_match(caps: &fancy_regex::Captures<'_>) -> usize {
for i in 1..caps.len() {
if let Some(m) = caps.get(i) {
if m.end() > m.start() {
return m.start();
}
}
}
caps.get(0).map(|m| m.start()).unwrap_or(0)
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct LazyTriggers {
tokens: Vec<u32>,
patterns: Vec<TriggerPattern>,
mandatory: bool,
}
impl LazyTriggers {
pub fn new() -> Self {
Self::default()
}
pub fn mandatory(mut self) -> Self {
self.mandatory = true;
self
}
pub fn is_mandatory(&self) -> bool {
self.mandatory
}
pub fn with_token(mut self, id: u32) -> Self {
self.tokens.push(id);
self
}
pub fn with_pattern(mut self, pattern: &str) -> Result<Self, GrammarError> {
self.patterns.push(TriggerPattern::new(pattern)?);
Ok(self)
}
pub fn with_word(mut self, word: &str) -> Result<Self, GrammarError> {
self.patterns
.push(TriggerPattern::new(&fancy_regex::escape(word))?);
Ok(self)
}
pub fn with_full_pattern(mut self, pattern: &str) -> Result<Self, GrammarError> {
let anchored = if pattern.is_empty() {
"^$".to_string()
} else {
let head = if pattern.starts_with('^') { "" } else { "^" };
let tail = if pattern.ends_with('$') { "" } else { "$" };
format!("{head}{pattern}{tail}")
};
self.patterns.push(TriggerPattern::new(&anchored)?);
Ok(self)
}
pub fn is_empty(&self) -> bool {
self.tokens.is_empty() && self.patterns.is_empty()
}
pub fn tokens(&self) -> &[u32] {
&self.tokens
}
pub fn patterns(&self) -> &[TriggerPattern] {
&self.patterns
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum TriggerStep {
Awaiting,
Fired(Vec<(u32, Vec<u8>)>),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LazyState {
triggers: LazyTriggers,
awaiting: bool,
buffer: Vec<u8>,
spans: Vec<(u32, usize, usize)>,
}
impl LazyState {
pub fn new(triggers: LazyTriggers) -> Self {
Self {
triggers,
awaiting: true,
buffer: Vec::new(),
spans: Vec::new(),
}
}
pub fn awaiting(&self) -> bool {
self.awaiting
}
pub fn buffer(&self) -> &[u8] {
&self.buffer
}
pub fn triggers(&self) -> &LazyTriggers {
&self.triggers
}
pub fn is_mandatory(&self) -> bool {
self.triggers.mandatory
}
pub fn observe(&mut self, token: u32, piece: &[u8]) -> Result<TriggerStep, GrammarError> {
debug_assert!(self.awaiting, "observe on a grammar that already fired");
if self.triggers.tokens.contains(&token) {
self.fire();
return Ok(TriggerStep::Fired(vec![(token, piece.to_vec())]));
}
let start = self.buffer.len();
self.buffer.extend_from_slice(piece);
self.spans.push((token, start, self.buffer.len()));
let Some(at) = self.find_trigger()? else {
return Ok(TriggerStep::Awaiting);
};
let mut replay = Vec::new();
for &(tok, tok_start, tok_end) in &self.spans {
if tok_end <= at {
continue;
}
let from = tok_start.max(at);
replay.push((tok, self.buffer[from..tok_end].to_vec()));
}
self.fire();
Ok(TriggerStep::Fired(replay))
}
fn find_trigger(&self) -> Result<Option<usize>, GrammarError> {
if self.triggers.patterns.is_empty() {
return Ok(None);
}
let text = match std::str::from_utf8(&self.buffer) {
Ok(s) => s,
Err(e) => {
let valid = e.valid_up_to();
std::str::from_utf8(&self.buffer[..valid]).unwrap_or("")
}
};
for pattern in &self.triggers.patterns {
if let Some(at) = pattern.find(text)? {
return Ok(Some(at));
}
}
Ok(None)
}
fn fire(&mut self) {
self.awaiting = false;
self.buffer.clear();
self.spans.clear();
}
}
#[cfg(test)]
mod tests {
use super::*;
fn triggers_with(pattern: &str) -> LazyState {
LazyState::new(LazyTriggers::new().with_pattern(pattern).unwrap())
}
#[test]
fn a_pattern_matches_across_token_boundaries() {
let mut s = triggers_with("<tool_call>");
assert_eq!(s.observe(1, b"<tool").unwrap(), TriggerStep::Awaiting);
assert_eq!(s.observe(2, b"_ca").unwrap(), TriggerStep::Awaiting);
let TriggerStep::Fired(replay) = s.observe(3, b"ll>").unwrap() else {
panic!("the buffer now holds the whole trigger word");
};
assert_eq!(
replay,
vec![
(1, b"<tool".to_vec()),
(2, b"_ca".to_vec()),
(3, b"ll>".to_vec())
]
);
assert!(!s.awaiting());
assert!(s.buffer().is_empty());
}
#[test]
fn text_before_the_trigger_is_dropped_and_the_straddling_token_is_truncated() {
let mut s = triggers_with("<tool_call>");
assert_eq!(
s.observe(7, b"sure, let me look").unwrap(),
TriggerStep::Awaiting
);
let TriggerStep::Fired(replay) = s.observe(8, b" up<tool_call>").unwrap() else {
panic!("trigger is complete");
};
assert_eq!(
replay,
vec![(8, b"<tool_call>".to_vec())],
"the prose token must not reach the grammar, and token 8 must lose its \" up\""
);
}
#[test]
fn a_trigger_token_matches_an_id_and_discards_the_prose() {
let mut s = LazyState::new(LazyTriggers::new().with_token(42));
assert_eq!(s.observe(1, b"thinking...").unwrap(), TriggerStep::Awaiting);
let step = s.observe(42, b"<tool_call>").unwrap();
assert_eq!(
step,
TriggerStep::Fired(vec![(42, b"<tool_call>".to_vec())]),
"only the trigger token itself is fed to the grammar"
);
}
#[test]
fn a_trigger_token_does_not_match_by_piece() {
let mut s = LazyState::new(LazyTriggers::new().with_token(42));
assert_eq!(s.observe(9, b"<tool_call>").unwrap(), TriggerStep::Awaiting);
}
#[test]
fn the_grammar_starts_at_the_first_non_empty_capture_group() {
let mut s = triggers_with(r"<\|start\|>assistant(\s+to)");
let TriggerStep::Fired(replay) = s.observe(1, b"<|start|>assistant to").unwrap() else {
panic!("trigger matches");
};
assert_eq!(
replay,
vec![(1, b" to".to_vec())],
"the grammar must be fed from the capture group, not the match"
);
}
#[test]
fn an_empty_capture_group_is_skipped() {
let p = TriggerPattern::new(r"ab(x?)(c)").unwrap();
assert_eq!(
p.find("zzabc").unwrap(),
Some(4),
"group 1 matched empty, so group 2 decides"
);
let p = TriggerPattern::new(r"ab(?:c)").unwrap();
assert_eq!(
p.find("zzabc").unwrap(),
Some(2),
"no group: the match start"
);
}
#[test]
fn an_anchored_pattern_matches_only_the_whole_buffer() {
let p = TriggerPattern::new(r"^\s+to$").unwrap();
assert_eq!(p.find(" to").unwrap(), Some(0));
assert_eq!(
p.find(" to ").unwrap(),
None,
"a trailing space means the buffer is no longer the whole match"
);
assert_eq!(p.find("x to").unwrap(), None);
}
#[test]
fn a_word_trigger_is_matched_literally() {
let t = LazyTriggers::new().with_word("[TOOL_CALLS]").unwrap();
let p = &t.patterns()[0];
assert_eq!(p.find("say [TOOL_CALLS] now").unwrap(), Some(4));
assert_eq!(
p.find("say TOOL_CALLS now").unwrap(),
None,
"unescaped, the brackets would be a character class"
);
}
#[test]
fn a_full_pattern_is_anchored_at_both_ends() {
let t = LazyTriggers::new().with_full_pattern("to").unwrap();
assert_eq!(t.patterns()[0].pattern(), "^to$");
let t = LazyTriggers::new().with_full_pattern("^to").unwrap();
assert_eq!(t.patterns()[0].pattern(), "^to$");
let t = LazyTriggers::new().with_full_pattern("").unwrap();
assert_eq!(t.patterns()[0].pattern(), "^$");
}
#[test]
fn a_partial_codepoint_does_not_break_matching() {
let mut s = triggers_with("é!");
assert_eq!(s.observe(1, b"\xc3").unwrap(), TriggerStep::Awaiting);
let TriggerStep::Fired(replay) = s.observe(2, b"\xa9!").unwrap() else {
panic!("the character is complete now");
};
assert_eq!(
replay,
vec![(1, b"\xc3".to_vec()), (2, b"\xa9!".to_vec())],
"the half-character token is replayed too: it is inside the match"
);
}
#[test]
fn an_uncompilable_pattern_is_refused() {
let err = LazyTriggers::new().with_pattern("(unclosed").unwrap_err();
assert!(
matches!(err, GrammarError::TriggerPatternInvalid { .. }),
"{err}"
);
}
#[test]
fn a_lookahead_pattern_compiles_and_matches() {
let p = TriggerPattern::new(r">>>(?!all)").unwrap();
assert_eq!(p.find(">>>get_weather").unwrap(), Some(0));
assert_eq!(p.find(">>>all").unwrap(), None);
}
}
#[cfg(test)]
mod grammar_tests {
use super::*;
use crate::grammar::candidates::{reject_candidates, Candidate};
use crate::grammar::machine::Grammar;
fn lazy_grammar(src: &str, triggers: LazyTriggers) -> Grammar {
Grammar::from_str_with_root(src, "root")
.expect("grammar parses")
.into_lazy(triggers)
.expect("triggers are not empty")
}
#[test]
fn prose_the_grammar_forbids_is_accepted_while_awaiting() {
let src = r#"root ::= "<t>" "{}""#;
let mut lazy = lazy_grammar(src, LazyTriggers::new().with_word("<t>").unwrap());
lazy.accept_token(1, b"sure!")
.expect("an untriggered grammar accepts anything");
assert!(lazy.is_awaiting_trigger());
assert_eq!(lazy.trigger_buffer(), b"sure!");
let mut eager = Grammar::from_str_with_root(src, "root").unwrap();
eager
.accept_token(1, b"sure!")
.expect_err("without a trigger the same prose kills the parse");
}
#[test]
fn the_replay_leaves_the_grammar_where_the_trigger_text_put_it() {
let mut g = lazy_grammar(
r#"root ::= "<t>" "{}""#,
LazyTriggers::new().with_word("<t>").unwrap(),
);
g.accept_token(1, b"hmm <").unwrap();
g.accept_token(2, b"t>").unwrap();
assert!(!g.is_awaiting_trigger(), "the trigger word is complete");
assert!(g.trigger_buffer().is_empty());
assert!(
g.accept_token(3, b"{").is_ok(),
"the grammar should be past \"<t>\" and expecting \"{{\""
);
g.accept_token(4, b"}").unwrap();
assert!(g.allows_eog(), "the parse is complete");
}
#[test]
fn a_replay_the_grammar_rejects_is_an_error() {
let mut g = lazy_grammar(
r#"root ::= "{}""#,
LazyTriggers::new().with_word("<t>").unwrap(),
);
let err = g
.accept_token(1, b"<t>")
.expect_err("this grammar cannot consume its own trigger word");
assert!(matches!(err, GrammarError::NoViableStack { .. }), "{err}");
}
#[test]
fn an_untriggered_grammar_allows_end_of_generation() {
let mut g = lazy_grammar(
r#"root ::= "<t>" "{}""#,
LazyTriggers::new().with_word("<t>").unwrap(),
);
assert!(g.allows_eog());
assert!(g.accept_eog().is_ok());
g.accept_token(1, b"<t>").unwrap();
assert!(
!g.allows_eog(),
"once triggered it is an ordinary unsatisfied grammar"
);
}
#[test]
fn the_candidate_walk_refuses_an_untriggered_grammar() {
let g = lazy_grammar(
r#"root ::= "<t>" "{}""#,
LazyTriggers::new().with_word("<t>").unwrap(),
);
let cands = [Candidate::new(0, 0, b"zzz")];
let err = reject_candidates(&g, &cands).expect_err("no answer to give yet");
assert!(matches!(err, GrammarError::AwaitingTrigger), "{err}");
}
#[test]
fn a_lazy_grammar_with_no_triggers_is_refused() {
let err = Grammar::from_str_with_root(r#"root ::= "a""#, "root")
.unwrap()
.into_lazy(LazyTriggers::new())
.expect_err("nothing could ever switch this on");
assert!(matches!(err, GrammarError::LazyWithoutTriggers), "{err}");
}
#[test]
fn an_eager_grammar_is_not_awaiting_anything() {
let g = Grammar::from_str_with_root(r#"root ::= "a""#, "root").unwrap();
assert!(!g.is_lazy());
assert!(!g.is_awaiting_trigger());
assert!(g.trigger_buffer().is_empty());
}
}