use std::sync::Arc;
use llguidance::api::TopLevelGrammar;
use llguidance::toktrie::{ApproximateTokEnv, TokEnv, TokRxInfo, TokTrie};
use llguidance::{Matcher, ParserFactory};
use super::{ConstraintError, ConstraintVocab, TokenConstraint};
pub struct LlgEnv {
factory: Arc<ParserFactory>,
}
impl LlgEnv {
pub fn new(vocab: &ConstraintVocab) -> Result<Self, ConstraintError> {
vocab.validate()?;
let words: Vec<Vec<u8>> = vocab
.token_bytes
.iter()
.zip(&vocab.special)
.map(|(bytes, &special)| {
if special {
let mut marked = Vec::with_capacity(bytes.len() + 1);
marked.push(TokTrie::SPECIAL_TOKEN_MARKER);
marked.extend_from_slice(bytes);
marked
} else {
bytes.clone()
}
})
.collect();
let vocab_size = u32::try_from(words.len()).map_err(|_| {
ConstraintError::Vocab(format!("{} tokens do not fit a u32 id", words.len()))
})?;
let info = TokRxInfo::new(vocab_size, vocab.eos);
let trie = TokTrie::from(&info, &words);
let env: TokEnv = Arc::new(ApproximateTokEnv::new(trie));
let mut factory = ParserFactory::new_simple(&env).map_err(|e| {
ConstraintError::Vocab(format!("llguidance refused the vocabulary: {e}"))
})?;
factory.quiet();
Ok(Self {
factory: Arc::new(factory),
})
}
pub fn json_schema(
&self,
schema: &serde_json::Value,
) -> Result<LlgConstraint, ConstraintError> {
let mut schema = match schema {
serde_json::Value::Bool(true) => serde_json::json!({}),
serde_json::Value::Bool(false) => {
return Err(ConstraintError::SchemaInvalid(
"the schema `false` admits no document, so nothing could ever be generated"
.to_string(),
))
},
other => other.clone(),
};
if schema.is_object() && schema.get("x-guidance").is_none() {
llguidance::JsonCompileOptions {
whitespace_flexible: false,
..Default::default()
}
.apply_to(&mut schema);
}
self.compile(TopLevelGrammar::from_json_schema(schema), "JSON schema")
}
pub fn lark(&self, grammar: &str) -> Result<LlgConstraint, ConstraintError> {
self.compile(
TopLevelGrammar::from_lark(grammar.to_string()),
"Lark grammar",
)
}
fn compile(
&self,
grammar: TopLevelGrammar,
what: &str,
) -> Result<LlgConstraint, ConstraintError> {
let parser = self.factory.create_parser(grammar).map_err(|e| {
ConstraintError::SchemaUnsupported(format!(
"the {what} does not compile for constrained decoding: {e}"
))
})?;
let mut matcher = Matcher::new(Ok(parser));
if let Some(e) = matcher.get_error() {
return Err(ConstraintError::SchemaUnsupported(format!(
"the {what} does not start: {e}"
)));
}
let warnings = matcher.grammar_warnings();
if !warnings.is_empty() {
return Err(ConstraintError::SchemaUnsupported(format!(
"the {what} compiles only approximately: {}",
warnings.join("; ")
)));
}
Ok(LlgConstraint {
matcher,
position: 0,
})
}
}
pub struct LlgConstraint {
matcher: Matcher,
position: usize,
}
impl TokenConstraint for LlgConstraint {
fn mask(&mut self, logits: &mut [f32]) -> Result<(), ConstraintError> {
let allowed = self
.matcher
.compute_mask_or_eos()
.map_err(|e| ConstraintError::DeadEnd {
position: self.position,
reason: e
.to_string()
.lines()
.next()
.unwrap_or("no reason given")
.to_string(),
})?;
let mut any = false;
for (id, logit) in logits.iter_mut().enumerate() {
let ok = id < allowed.len() && u32::try_from(id).is_ok_and(|t| allowed.is_allowed(t));
if ok {
any = true;
} else {
*logit = f32::NEG_INFINITY;
}
}
if any {
Ok(())
} else {
Err(ConstraintError::DeadEnd {
position: self.position,
reason: "the mask allows no token".to_string(),
})
}
}
fn accept(&mut self, token: u32) -> Result<(), ConstraintError> {
self.matcher
.consume_token(token)
.map_err(|e| ConstraintError::Rejected {
token,
position: self.position,
reason: e.to_string(),
})?;
self.position += 1;
Ok(())
}
fn is_complete(&mut self) -> bool {
self.matcher.is_stopped() || self.matcher.is_accepting().unwrap_or(false)
}
}