use std::fmt;
#[cfg(feature = "structured-output")]
mod llg;
#[cfg(feature = "structured-output")]
pub use llg::{LlgConstraint, LlgEnv};
pub trait TokenConstraint: Send {
fn mask(&mut self, logits: &mut [f32]) -> Result<(), ConstraintError>;
fn accept(&mut self, token: u32) -> Result<(), ConstraintError>;
fn is_complete(&mut self) -> bool;
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ConstraintError {
SchemaInvalid(String),
SchemaUnsupported(String),
DeadEnd {
position: usize,
reason: String,
},
Rejected {
token: u32,
position: usize,
reason: String,
},
Vocab(String),
NotCompiled,
}
impl fmt::Display for ConstraintError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::SchemaInvalid(why) => write!(f, "SchemaInvalid: {why}"),
Self::SchemaUnsupported(why) => write!(f, "SchemaUnsupported: {why}"),
Self::DeadEnd { position, reason } => write!(
f,
"ConstraintDeadEnd: no token is allowed at output position {position}: {reason}"
),
Self::Rejected {
token,
position,
reason,
} => write!(
f,
"ConstraintRejected: token {token} is not allowed at output position {position}: {reason}"
),
Self::Vocab(why) => write!(f, "ConstraintVocab: {why}"),
Self::NotCompiled => write!(
f,
"StructuredOutputNotCompiled: this build has no `structured-output` feature, so a \
schema cannot be enforced (refused rather than ignored)"
),
}
}
}
impl std::error::Error for ConstraintError {}
pub fn load_schema(arg: &str) -> Result<serde_json::Value, ConstraintError> {
let (text, from) = match arg.strip_prefix('@') {
Some(path) => (
std::fs::read_to_string(path).map_err(|e| {
ConstraintError::SchemaInvalid(format!("cannot read the schema file {path}: {e}"))
})?,
format!("the schema file {path}"),
),
None => (arg.to_string(), "the inline schema".to_string()),
};
let value: serde_json::Value = serde_json::from_str(&text)
.map_err(|e| ConstraintError::SchemaInvalid(format!("{from} is not JSON: {e}")))?;
if !(value.is_object() || value.is_boolean()) {
return Err(ConstraintError::SchemaInvalid(format!(
"{from} is JSON but not a schema: a JSON Schema is an object or a boolean"
)));
}
Ok(value)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ConstraintVocab {
pub token_bytes: Vec<Vec<u8>>,
pub special: Vec<bool>,
pub eos: u32,
}
impl ConstraintVocab {
pub fn validate(&self) -> Result<(), ConstraintError> {
if self.token_bytes.is_empty() {
return Err(ConstraintError::Vocab(
"the vocabulary is empty".to_string(),
));
}
if self.special.len() != self.token_bytes.len() {
return Err(ConstraintError::Vocab(format!(
"{} token byte strings but {} special flags",
self.token_bytes.len(),
self.special.len()
)));
}
if self.eos as usize >= self.token_bytes.len() {
return Err(ConstraintError::Vocab(format!(
"eos {} is outside a vocabulary of {}",
self.eos,
self.token_bytes.len()
)));
}
Ok(())
}
}
pub struct ConstraintEnv {
#[cfg(feature = "structured-output")]
inner: LlgEnv,
}
impl ConstraintEnv {
pub fn new(vocab: &ConstraintVocab) -> Result<Self, ConstraintError> {
vocab.validate()?;
#[cfg(feature = "structured-output")]
{
Ok(Self {
inner: LlgEnv::new(vocab)?,
})
}
#[cfg(not(feature = "structured-output"))]
{
Err(ConstraintError::NotCompiled)
}
}
pub fn json_schema(
&self,
schema: &serde_json::Value,
) -> Result<Box<dyn TokenConstraint>, ConstraintError> {
#[cfg(feature = "structured-output")]
{
Ok(Box::new(self.inner.json_schema(schema)?))
}
#[cfg(not(feature = "structured-output"))]
{
let _ = schema;
Err(ConstraintError::NotCompiled)
}
}
pub fn lark(&self, grammar: &str) -> Result<Box<dyn TokenConstraint>, ConstraintError> {
#[cfg(feature = "structured-output")]
{
Ok(Box::new(self.inner.lark(grammar)?))
}
#[cfg(not(feature = "structured-output"))]
{
let _ = grammar;
Err(ConstraintError::NotCompiled)
}
}
}
#[cfg(test)]
mod tests;