use std::sync::Arc;
use llguidance::api::TopLevelGrammar;
use llguidance::toktrie::{SimpleVob, TokEnv, TokRxInfo, TokTrie, TokenId, TokenizerEnv};
use llguidance::{Matcher, ParserFactory};
use memra_tokenizer::Tokenizer;
#[derive(Debug, Clone)]
pub enum GrammarSpec {
JsonObject,
JsonSchema(serde_json::Value),
}
pub fn parse_response_format(v: Option<&serde_json::Value>)
-> Result<Option<GrammarSpec>, String>
{
let Some(v) = v else { return Ok(None) };
let ty = v.get("type").and_then(|t| t.as_str())
.ok_or("response_format.type must be a string")?;
match ty {
"text" => Ok(None),
"json_object" => Ok(Some(GrammarSpec::JsonObject)),
"json_schema" => {
let js = v.get("json_schema")
.ok_or("response_format.json_schema is required for type json_schema")?;
if !js.is_object() {
return Err("response_format.json_schema must be an object".into());
}
let schema = js.get("schema").unwrap_or(js).clone();
Ok(Some(GrammarSpec::JsonSchema(schema)))
}
other => Err(format!("response_format type {other:?} is not supported \
(text | json_object | json_schema)")),
}
}
struct MemraTokEnv {
trie: TokTrie,
}
impl TokenizerEnv for MemraTokEnv {
fn tok_trie(&self) -> &TokTrie {
&self.trie
}
fn tokenize_bytes(&self, s: &[u8]) -> Vec<TokenId> {
self.trie.greedy_tokenize(s)
}
fn tokenize_is_canonical(&self) -> bool {
false
}
}
pub struct ConstraintFactory {
factory: ParserFactory,
}
impl ConstraintFactory {
pub fn new(tok: &Tokenizer) -> Result<Self, String> {
let n = tok.vocab_size();
let mut words: Vec<Vec<u8>> = Vec::with_capacity(n);
for id in 0..n as u32 {
if tok.token_is_control(id) {
let mut w = vec![TokTrie::SPECIAL_TOKEN_MARKER];
w.extend_from_slice(format!("[{id}]").as_bytes());
words.push(w);
} else {
words.push(tok.decode_bytes_special(&[id], true));
}
}
let info = TokRxInfo::new(n as u32, tok.eos_id());
let trie = TokTrie::from(&info, &words);
let env: TokEnv = Arc::new(MemraTokEnv { trie });
let mut factory = ParserFactory::new_simple(&env)
.map_err(|e| format!("constraint factory: {e}"))?;
factory.quiet();
Ok(Self { factory })
}
pub fn matcher(&self, spec: &GrammarSpec) -> SessionConstraint {
let schema = match spec {
GrammarSpec::JsonObject => serde_json::json!({"type": "object"}),
GrammarSpec::JsonSchema(s) => s.clone(),
};
let grammar = TopLevelGrammar::from_json_schema(schema);
SessionConstraint::new(Matcher::new(self.factory.create_parser(grammar)))
}
}
pub fn apply_mask(mask: &SimpleVob, logits: &mut [f32]) {
let n = logits.len();
mask.iter_unset_entries(|i| {
if i < n {
logits[i] = f32::NEG_INFINITY;
}
});
if mask.len() < n {
for l in &mut logits[mask.len()..] {
*l = f32::NEG_INFINITY;
}
}
}
pub struct SessionConstraint {
m: Matcher,
pub steps: u64,
pub mask_ns: u128,
pub spec_clones: u64,
pub spec_ns: u128,
pub draft_masks: u64,
pub draft_mask_ns: u128,
}
impl SessionConstraint {
pub fn new(m: Matcher) -> Self {
Self { m, steps: 0, mask_ns: 0,
spec_clones: 0, spec_ns: 0, draft_masks: 0, draft_mask_ns: 0 }
}
pub fn error(&self) -> Option<String> {
self.m.get_error()
}
pub fn compute_mask(&mut self) -> Result<SimpleVob, String> {
let t0 = std::time::Instant::now();
let mask = self.m.compute_mask_or_eos().map_err(|e| e.to_string())?;
self.steps += 1;
self.mask_ns += t0.elapsed().as_nanos();
Ok(mask)
}
pub fn mask_logits(&mut self, logits: &mut [f32]) -> Result<(), String> {
let mask = self.compute_mask()?;
apply_mask(&mask, logits);
Ok(())
}
pub fn consume(&mut self, tok: u32) -> Result<(), String> {
self.m.consume_token(tok).map_err(|e| e.to_string())
}
pub fn clone_matcher(&mut self) -> Matcher {
let t0 = std::time::Instant::now();
let m = self.m.clone();
self.spec_clones += 1;
self.spec_ns += t0.elapsed().as_nanos();
m
}
}
pub struct SpecGrammar<'a> {
c: &'a mut SessionConstraint,
eos: u32,
cur: Option<SimpleVob>,
spec: Option<Matcher>,
on: bool,
}
pub fn draft_mask_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_DRAFT_MASK").map(|v| v != "0").unwrap_or(true))
}
impl<'a> SpecGrammar<'a> {
pub fn new(c: &'a mut SessionConstraint, eos: u32) -> Self {
Self { c, eos, cur: None, spec: None, on: draft_mask_on() }
}
fn cur_mask(&mut self) -> Result<&SimpleVob, String> {
if self.cur.is_none() {
self.cur = Some(self.c.compute_mask()?);
}
Ok(self.cur.as_ref().unwrap())
}
}
impl memra_engine::spec::SpecConstraint for SpecGrammar<'_> {
fn mask_logits(&mut self, logits: &mut [f32]) -> Result<(), String> {
let mask = self.cur_mask()?;
apply_mask(mask, logits);
Ok(())
}
fn mask_words(&mut self) -> Result<Vec<u32>, String> {
Ok(self.cur_mask()?.as_slice().to_vec())
}
fn is_allowed(&mut self, tok: u32) -> Result<bool, String> {
let mask = self.cur_mask()?;
Ok((tok as usize) < mask.len() && mask.is_allowed(tok))
}
fn consume(&mut self, tok: u32) -> Result<(), String> {
self.spec = None;
if tok == self.eos {
return Ok(()); }
self.c.consume(tok)?;
self.cur = None;
Ok(())
}
fn draft_mask_enabled(&self) -> bool {
self.on
}
fn draft_begin(&mut self) -> Result<(), String> {
if !self.on {
self.spec = None;
return Ok(());
}
self.spec = Some(self.c.clone_matcher());
Ok(())
}
fn draft_mask_words(&mut self) -> Result<Option<Vec<u32>>, String> {
if !self.on {
return Ok(None);
}
let Some(spec) = self.spec.as_mut() else { return Ok(None) };
let t0 = std::time::Instant::now();
let mask = spec.compute_mask_or_eos().map_err(|e| e.to_string())?;
self.c.draft_masks += 1;
self.c.draft_mask_ns += t0.elapsed().as_nanos();
Ok(Some(mask.as_slice().to_vec()))
}
fn draft_advance(&mut self, tok: u32) -> Result<bool, String> {
if !self.on {
return Ok(false);
}
let Some(spec) = self.spec.as_mut() else { return Ok(false) };
if tok == self.eos {
return Ok(false); }
match spec.consume_token(tok) {
Ok(()) => Ok(true),
Err(_) => Ok(false),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use llguidance::toktrie::ApproximateTokEnv;
#[test]
fn parse_response_format_forms() {
assert!(parse_response_format(None).unwrap().is_none());
let text = serde_json::json!({"type": "text"});
assert!(parse_response_format(Some(&text)).unwrap().is_none());
let jo = serde_json::json!({"type": "json_object"});
assert!(matches!(parse_response_format(Some(&jo)).unwrap(),
Some(GrammarSpec::JsonObject)));
let js = serde_json::json!({"type": "json_schema", "json_schema": {
"name": "x", "schema": {"type": "object", "required": ["a"]}}});
match parse_response_format(Some(&js)).unwrap() {
Some(GrammarSpec::JsonSchema(s)) => assert_eq!(s["required"][0], "a"),
other => panic!("wrong parse: {other:?}"),
}
let js2 = serde_json::json!({"type": "json_schema",
"json_schema": {"type": "object"}});
match parse_response_format(Some(&js2)).unwrap() {
Some(GrammarSpec::JsonSchema(s)) => assert_eq!(s["type"], "object"),
other => panic!("wrong parse: {other:?}"),
}
let bad = serde_json::json!({"type": "yaml"});
assert!(parse_response_format(Some(&bad)).is_err());
let bad2 = serde_json::json!({"type": "json_schema"});
assert!(parse_response_format(Some(&bad2)).is_err());
let bad3 = serde_json::json!({"type": 3});
assert!(parse_response_format(Some(&bad3)).is_err());
}
#[test]
fn apply_mask_bans_unset_and_padding_tail() {
let mut vob = SimpleVob::alloc(8);
vob.allow_token(2);
vob.allow_token(5);
let mut logits = vec![1.0f32; 10];
apply_mask(&vob, &mut logits);
for (i, &l) in logits.iter().enumerate() {
if i == 2 || i == 5 {
assert_eq!(l, 1.0, "allowed token {i} must be untouched");
} else {
assert_eq!(l, f32::NEG_INFINITY, "banned token {i} must be -inf");
}
}
}
#[test]
fn schema_mask_forced_sequence() {
let env = ApproximateTokEnv::single_byte_env();
let factory = ParserFactory::new_simple(&env).unwrap();
let schema = serde_json::json!({
"type": "object",
"properties": {"a": {"type": "integer"}},
"required": ["a"],
"additionalProperties": false
});
let mut m = Matcher::new(
factory.create_parser(TopLevelGrammar::from_json_schema(schema)));
assert!(m.get_error().is_none(), "{:?}", m.get_error());
let eos = env.tok_trie().eos_token();
let mut out: Vec<u8> = Vec::new();
for _ in 0..256 {
let mask = m.compute_mask_or_eos().unwrap();
assert!(mask.num_set() > 0, "empty mask");
let mut pick: Option<u32> = None;
mask.iter_set_entries(|i| {
let ws = matches!(i as u8, b'\t' | b'\n' | b'\r' | b' ') && i < 128;
if !ws && pick.is_none() {
pick = Some(i as u32);
}
});
let t = pick.expect("only whitespace allowed — walker stuck");
if t == eos {
break;
}
m.consume_token(t).unwrap();
out.extend_from_slice(env.tok_trie().token(t));
}
let text = String::from_utf8(out).unwrap();
let v: serde_json::Value = serde_json::from_str(&text)
.unwrap_or_else(|e| panic!("forced output is not JSON: {e}: {text:?}"));
assert!(v.is_object(), "not an object: {text:?}");
let a = v.get("a").unwrap_or_else(|| panic!("required key missing: {text:?}"));
assert!(a.as_f64().is_some_and(|f| f.fract() == 0.0),
"required integer key not an integer: {text:?}");
}
#[test]
fn speculative_clone_masks_illegal_draft_and_leaves_real_state() {
use memra_engine::spec::SpecConstraint;
let env = ApproximateTokEnv::single_byte_env();
let factory = ParserFactory::new_simple(&env).unwrap();
let schema = serde_json::json!({
"type": "object",
"properties": {"a": {"type": "integer"}},
"required": ["a"],
"additionalProperties": false
});
let mut sc = SessionConstraint::new(Matcher::new(
factory.create_parser(TopLevelGrammar::from_json_schema(schema))));
assert!(sc.error().is_none());
let eos = env.tok_trie().eos_token();
let mut g = SpecGrammar::new(&mut sc, eos);
assert!(g.on, "draft masking must default ON");
g.draft_begin().unwrap();
let w0 = g.draft_mask_words().unwrap().expect("draft mask must be present when ON");
let allowed = |words: &[u32], t: u32| -> bool {
let w = (t >> 5) as usize;
w < words.len() && (words[w] >> (t & 31)) & 1 == 1
};
assert!(allowed(&w0, b'{' as u32), "'{{' must be legal at draft pos 0");
assert!(!allowed(&w0, b'x' as u32), "'x' must be MASKED at draft pos 0");
assert!(!allowed(&w0, b'a' as u32), "bare 'a' (unquoted key) must be masked at pos 0");
assert!(g.draft_advance(b'{' as u32).unwrap(), "legal draft must extend the chain");
let w1 = g.draft_mask_words().unwrap().unwrap();
assert!(allowed(&w1, b'"' as u32), "after '{{' a quoted key must be legal");
assert!(!allowed(&w1, b'{' as u32), "a second '{{' must be masked at draft pos 1");
let real: Vec<u32> = SpecConstraint::mask_words(&mut g).unwrap();
assert_eq!(real, w0, "real matcher moved during a draft chain (byte-identity break)");
assert!(!g.draft_advance(b'{' as u32).unwrap(),
"illegal draft token must end the chain, not error");
let real2: Vec<u32> = SpecConstraint::mask_words(&mut g).unwrap();
assert_eq!(real2, w0, "real matcher moved after a dead speculative chain");
SpecConstraint::consume(&mut g, b'{' as u32).unwrap();
g.draft_begin().unwrap();
let w2 = g.draft_mask_words().unwrap().unwrap();
assert_eq!(w2, w1, "a fresh chain after emitting '{{' must match the pos-1 mask");
assert!(sc.spec_clones >= 2, "clone meter must count each chain start");
assert!(sc.draft_masks >= 3, "draft-mask meter must count each masked position");
}
#[test]
fn consume_outside_mask_is_error() {
let env = ApproximateTokEnv::single_byte_env();
let factory = ParserFactory::new_simple(&env).unwrap();
let mut m = Matcher::new(factory.create_parser(
TopLevelGrammar::from_json_schema(serde_json::json!({"type": "object"}))));
assert!(m.consume_token(b'x' as u32).is_err());
}
}