Skip to main content

ferrum_sampler/
structured_output.rs

1//! Tokenizer-aware hard constraints for structured output.
2//!
3//! A factory owns the tokenizer trie and grammar compiler and is shared by an
4//! engine. Each request gets an independent matcher with no shared mutable
5//! grammar state.
6
7use std::{
8    collections::{HashMap, HashSet},
9    str,
10    sync::Arc,
11};
12
13use ferrum_interfaces::tokenizer::Tokenizer;
14use ferrum_types::{FerrumError, ResponseFormat, Result, StructuredOutputStart, TokenId};
15use llguidance::{
16    api::TopLevelGrammar,
17    toktrie::{InferenceCapabilities, TokEnv, TokRxInfo, TokTrie, TokenizerEnv},
18    JsonCompileOptions, Matcher, ParserFactory,
19};
20use parking_lot::Mutex;
21use serde_json::json;
22
23const MAX_CACHED_GRAMMARS: usize = 64;
24// Structured output is the requested product result; hidden reasoning may use
25// at most the other half of a normal-sized completion budget.
26const AUTO_STRUCTURED_RESERVE_DIVISOR: usize = 2;
27const MIN_AUTO_STRUCTURED_RESERVE_TOKENS: usize = 32;
28const MAX_AUTO_STRUCTURED_RESERVE_TOKENS: usize = 1024;
29const MAX_IDENTICAL_TOKEN_RUN: usize = MAX_AUTO_STRUCTURED_RESERVE_TOKENS / 2;
30
31/// Immutable per-request output budget used when a structured grammar starts
32/// after a reasoning delimiter.
33#[derive(Debug, Clone, Copy, PartialEq, Eq)]
34pub struct StructuredOutputBudgetPlan {
35    pub total_output_tokens: usize,
36    pub reasoning_token_limit: usize,
37    pub boundary_token_count: usize,
38    pub structured_reserve_tokens: usize,
39}
40
41impl StructuredOutputBudgetPlan {
42    fn automatic(total_output_tokens: usize, boundary_token_count: usize) -> Result<Self> {
43        if boundary_token_count == 0 || total_output_tokens <= boundary_token_count {
44            return Err(FerrumError::invalid_request(format!(
45                "structured output requires max_tokens greater than its {boundary_token_count}-token delimiter"
46            )));
47        }
48        let available_after_boundary = total_output_tokens - boundary_token_count;
49        let proportional_reserve = total_output_tokens.div_ceil(AUTO_STRUCTURED_RESERVE_DIVISOR);
50        let structured_reserve_tokens = proportional_reserve
51            .clamp(
52                MIN_AUTO_STRUCTURED_RESERVE_TOKENS,
53                MAX_AUTO_STRUCTURED_RESERVE_TOKENS,
54            )
55            .min(available_after_boundary);
56        Ok(Self {
57            total_output_tokens,
58            reasoning_token_limit: total_output_tokens
59                - boundary_token_count
60                - structured_reserve_tokens,
61            boundary_token_count,
62            structured_reserve_tokens,
63        })
64    }
65}
66
67#[derive(Debug, Clone, Copy, PartialEq, Eq)]
68struct StructuredOutputLivenessPolicy {
69    max_identical_token_run: usize,
70}
71
72impl StructuredOutputLivenessPolicy {
73    fn for_request(max_output_tokens: usize, budget: Option<StructuredOutputBudgetPlan>) -> Self {
74        let guaranteed_structured_tokens = budget
75            .map(|plan| plan.structured_reserve_tokens)
76            .unwrap_or(max_output_tokens);
77        Self {
78            max_identical_token_run: guaranteed_structured_tokens
79                .div_ceil(2)
80                .clamp(1, MAX_IDENTICAL_TOKEN_RUN),
81        }
82    }
83}
84
85/// Shared, immutable tokenizer and grammar compilation state.
86pub struct StructuredOutputFactory {
87    parser_factory: ParserFactory,
88    tokenizer: Arc<dyn Tokenizer + Send + Sync>,
89    vocab_size: usize,
90    defined_token_ids: Arc<[bool]>,
91    json_token_classes: Arc<[StructuredOutputTokenClass]>,
92    grammar_templates: Mutex<HashMap<String, Matcher>>,
93}
94
95impl std::fmt::Debug for StructuredOutputFactory {
96    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
97        f.debug_struct("StructuredOutputFactory")
98            .field("vocab_size", &self.vocab_size)
99            .finish_non_exhaustive()
100    }
101}
102
103impl StructuredOutputFactory {
104    /// Build the tokenizer trie once for this engine.
105    pub fn new(tokenizer: Arc<dyn Tokenizer + Send + Sync>) -> Result<Self> {
106        Self::new_with_model_vocab_size(tokenizer, None)
107    }
108
109    /// Build against the executor's logits width when it is larger than the
110    /// tokenizer base vocabulary (for example added EOS/control tokens).
111    pub fn new_with_model_vocab_size(
112        tokenizer: Arc<dyn Tokenizer + Send + Sync>,
113        model_vocab_size: Option<usize>,
114    ) -> Result<Self> {
115        let eos = tokenizer.special_tokens().eos_token.ok_or_else(|| {
116            FerrumError::config("structured output requires a tokenizer EOS token")
117        })?;
118        let vocab_size = model_vocab_size
119            .unwrap_or_else(|| tokenizer.vocab_size())
120            .max(tokenizer.vocab_size());
121        if vocab_size == 0 || eos.get() as usize >= vocab_size {
122            return Err(FerrumError::config(format!(
123                "structured output tokenizer has invalid vocab/EOS: vocab_size={vocab_size}, eos={}",
124                eos.get()
125            )));
126        }
127
128        let special_ids = tokenizer_special_ids(tokenizer.as_ref());
129        let mut defined_token_ids = Vec::with_capacity(vocab_size);
130        let mut json_token_classes = Vec::with_capacity(vocab_size);
131        let token_bytes = (0..vocab_size)
132            .map(|idx| {
133                let token = TokenId::new(idx as u32);
134                if special_ids.contains(&token.get()) {
135                    defined_token_ids.push(true);
136                    json_token_classes.push(StructuredOutputTokenClass::Control);
137                    special_token_marker(token)
138                } else if let Some(bytes) = tokenizer
139                    .token_bytes(token)
140                    .filter(|bytes| !bytes.is_empty())
141                {
142                    defined_token_ids.push(true);
143                    json_token_classes.push(classify_json_token_bytes(&bytes));
144                    bytes
145                } else {
146                    // Keep vocabulary holes out of the trie. The explicit
147                    // eligibility mask below is still required because
148                    // llguidance's wildcard slice represents its root as an
149                    // all-token bitset, including IDs with no trie node.
150                    defined_token_ids.push(false);
151                    json_token_classes.push(StructuredOutputTokenClass::Undefined);
152                    Vec::new()
153                }
154            })
155            .collect::<Vec<_>>();
156
157        let mut eos_tokens = vec![eos.get()];
158        eos_tokens.extend(
159            tokenizer
160                .special_tokens()
161                .extra_eos_tokens
162                .iter()
163                .map(|token| token.get())
164                .filter(|token| *token < vocab_size as u32),
165        );
166        eos_tokens.sort_unstable();
167        eos_tokens.dedup();
168        if let Some(position) = eos_tokens.iter().position(|token| *token == eos.get()) {
169            eos_tokens.swap(0, position);
170        }
171
172        let info = TokRxInfo::new(vocab_size as u32, eos.get());
173        let trie = TokTrie::from(&info, &token_bytes).with_eos_tokens(&eos_tokens);
174        let tok_env: TokEnv = Arc::new(FerrumTokenizerEnv {
175            tokenizer: Arc::clone(&tokenizer),
176            trie,
177        });
178        let mut parser_factory = ParserFactory::new(
179            &tok_env,
180            InferenceCapabilities {
181                ff_tokens: false,
182                conditional_ff_tokens: false,
183                backtrack: false,
184                fork: false,
185            },
186            &llguidance::earley::SlicedBiasComputer::general_slices(),
187        )
188        .map_err(|error| {
189            FerrumError::config(format!("build structured-output parser factory: {error}"))
190        })?;
191        parser_factory.quiet();
192
193        Ok(Self {
194            parser_factory,
195            tokenizer,
196            vocab_size,
197            defined_token_ids: defined_token_ids.into(),
198            json_token_classes: json_token_classes.into(),
199            grammar_templates: Mutex::new(HashMap::new()),
200        })
201    }
202
203    /// Compile one request's grammar while reusing the tokenizer trie.
204    pub fn create_processor(
205        &self,
206        response_format: &ResponseFormat,
207        start: &StructuredOutputStart,
208        max_output_tokens: usize,
209        stop_token_ids: &HashSet<u32>,
210        stop_text_sequences: &[String],
211    ) -> Result<Option<StructuredOutputProcessor>> {
212        let schema = match response_format {
213            ResponseFormat::Text => return Ok(None),
214            ResponseFormat::JsonObject => json!({"type": "object"}),
215            ResponseFormat::JsonSchema(schema) => {
216                serde_json::from_str(schema).map_err(|error| {
217                    FerrumError::invalid_request(format!(
218                        "response_format.schema is not valid JSON: {error}"
219                    ))
220                })?
221            }
222        };
223        let schema = compact_json_schema(schema)?;
224        let grammar_key = serde_json::to_string(&schema).map_err(|error| {
225            FerrumError::invalid_request(format!("serialize structured-output schema: {error}"))
226        })?;
227        let matcher = {
228            let mut templates = self.grammar_templates.lock();
229            if let Some(template) = templates.get(&grammar_key) {
230                template.deep_clone()
231            } else {
232                let grammar = TopLevelGrammar::from_json_schema(schema);
233                let parser = self
234                    .parser_factory
235                    .create_parser(grammar)
236                    .map_err(|error| {
237                        FerrumError::invalid_request(format!(
238                            "unsupported structured-output grammar: {error}"
239                        ))
240                    })?;
241                let matcher = Matcher::new(Ok(parser));
242                if templates.len() >= MAX_CACHED_GRAMMARS {
243                    templates.clear();
244                }
245                templates.insert(grammar_key, matcher.deep_clone());
246                matcher
247            }
248        };
249        let (activation, budget) = match start {
250            StructuredOutputStart::Immediate => (Activation::Active, None),
251            StructuredOutputStart::AfterDelimiter(delimiter) => {
252                if delimiter.is_empty() {
253                    return Err(FerrumError::invalid_request(
254                        "structured-output delimiter must not be empty",
255                    ));
256                }
257                let delimiter_tokens = if let Some(token) = self.tokenizer.token_id(delimiter) {
258                    vec![token.get()]
259                } else {
260                    self.tokenizer
261                        .encode(delimiter, false)?
262                        .into_iter()
263                        .map(|token| token.get())
264                        .collect::<Vec<_>>()
265                };
266                if delimiter_tokens.is_empty() {
267                    return Err(FerrumError::invalid_request(format!(
268                        "structured-output delimiter {delimiter:?} did not tokenize"
269                    )));
270                }
271                if let Some(token) = delimiter_tokens
272                    .iter()
273                    .find(|token| stop_token_ids.contains(token))
274                {
275                    return Err(FerrumError::invalid_request(format!(
276                        "structured-output delimiter token {token} conflicts with a stop token"
277                    )));
278                }
279                if let Some(stop) = stop_text_sequences
280                    .iter()
281                    .find(|stop| !stop.is_empty() && delimiter.contains(stop.as_str()))
282                {
283                    return Err(FerrumError::invalid_request(format!(
284                        "structured-output delimiter {delimiter:?} conflicts with stop sequence {stop:?}"
285                    )));
286                }
287                let budget = StructuredOutputBudgetPlan::automatic(
288                    max_output_tokens,
289                    delimiter_tokens.len(),
290                )?;
291                (
292                    Activation::Boundary {
293                        delimiter_tokens,
294                        forcing: false,
295                    },
296                    Some(budget),
297                )
298            }
299        };
300
301        let grammar_start = matches!(activation, Activation::Active).then_some(0);
302        let liveness = StructuredOutputLivenessPolicy::for_request(max_output_tokens, budget);
303        Ok(Some(StructuredOutputProcessor {
304            state: Mutex::new(ProcessorState {
305                matcher,
306                activation: activation.clone(),
307                initial_activation: activation,
308                consumed: 0,
309                boundary_forced: false,
310                boundary_start: None,
311                grammar_start,
312                trailing_grammar_token_id: None,
313                trailing_identical_token_count: 0,
314                liveness_intervention_count: 0,
315                last_liveness_intervention_at: None,
316            }),
317            vocab_size: self.vocab_size,
318            defined_token_ids: Arc::clone(&self.defined_token_ids),
319            json_token_classes: Arc::clone(&self.json_token_classes),
320            budget,
321            liveness,
322        }))
323    }
324}
325
326fn compact_json_schema(schema: serde_json::Value) -> Result<serde_json::Value> {
327    let mut schema = match schema {
328        schema @ serde_json::Value::Object(_) => schema,
329        schema @ serde_json::Value::Bool(_) => json!({"allOf": [schema]}),
330        _ => {
331            return Err(FerrumError::invalid_request(
332                "response_format.schema must be a JSON Schema object or boolean",
333            ));
334        }
335    };
336
337    // x-guidance is an llguidance compiler extension, not a JSON Schema
338    // constraint. Keep compiler policy owned by Ferrum so a request cannot
339    // re-enable an unbounded whitespace loop or change JSON separators.
340    JsonCompileOptions {
341        whitespace_flexible: false,
342        ..JsonCompileOptions::default()
343    }
344    .apply_to(&mut schema);
345    Ok(schema)
346}
347
348/// Per-request structured-output parser state.
349pub struct StructuredOutputProcessor {
350    state: Mutex<ProcessorState>,
351    vocab_size: usize,
352    defined_token_ids: Arc<[bool]>,
353    json_token_classes: Arc<[StructuredOutputTokenClass]>,
354    budget: Option<StructuredOutputBudgetPlan>,
355    liveness: StructuredOutputLivenessPolicy,
356}
357
358/// Typed phase returned after applying a structured-output constraint.
359#[derive(Debug, Clone, Copy, PartialEq, Eq)]
360pub enum StructuredOutputPhase {
361    WaitingForDelimiter,
362    ForcingDelimiter,
363    EnforcingGrammar,
364}
365
366/// Privacy-safe lexical class for the tail of an incomplete grammar.
367///
368/// This deliberately exposes neither decoded text nor a token history. It is
369/// used only by terminal diagnostics to distinguish an unbounded whitespace,
370/// number, string, or control-token run from a parser activation failure.
371#[derive(Debug, Clone, Copy, PartialEq, Eq)]
372#[repr(u8)]
373pub enum StructuredOutputTokenClass {
374    Whitespace,
375    Number,
376    Structural,
377    StringBoundary,
378    Literal,
379    Other,
380    Control,
381    Undefined,
382}
383
384/// Allocation-free hot-path result of one structured-output mask operation.
385#[derive(Debug, Clone, Copy, PartialEq, Eq)]
386pub struct StructuredOutputMaskOutcome {
387    pub phase: StructuredOutputPhase,
388    pub accepting: bool,
389    pub liveness_intervention: bool,
390    /// Generated-token index at which the visible grammar-owned output starts.
391    /// `None` means the processor is still in the hidden pre-grammar domain.
392    /// The execution engine uses this boundary to scope request-local sampling
393    /// history without teaching the grammar about penalty policy.
394    pub grammar_start_token_index: Option<usize>,
395    /// Exact delimiter token authorized for the next sampling step while the
396    /// processor is waiting to activate. The engine uses this typed grant to
397    /// avoid rejecting an intentionally hidden special token during output
398    /// quality filtering.
399    pub required_delimiter_token_id: Option<u32>,
400}
401
402/// Terminal/debug snapshot that distinguishes activation failures from an
403/// incomplete grammar without retaining generated text.
404#[derive(Debug, Clone, Copy, PartialEq, Eq)]
405pub struct StructuredOutputProgress {
406    pub phase: StructuredOutputPhase,
407    pub generated_token_count: usize,
408    pub consumed_token_count: usize,
409    pub delimiter_token_count: Option<usize>,
410    pub delimiter_prefix_token_count: usize,
411    pub reasoning_token_count: Option<usize>,
412    pub boundary_forced: bool,
413    pub budget: Option<StructuredOutputBudgetPlan>,
414    pub grammar_token_count: usize,
415    pub trailing_token_class: Option<StructuredOutputTokenClass>,
416    pub trailing_token_class_count: usize,
417    pub trailing_token_id: Option<u32>,
418    pub trailing_identical_token_count: usize,
419    pub liveness_identical_token_limit: usize,
420    pub liveness_intervention_count: usize,
421    pub accepting: bool,
422}
423
424impl std::fmt::Debug for StructuredOutputProcessor {
425    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
426        let state = self.state.lock();
427        f.debug_struct("StructuredOutputProcessor")
428            .field("vocab_size", &self.vocab_size)
429            .field("consumed", &state.consumed)
430            .field("active", &matches!(state.activation, Activation::Active))
431            .field("budget", &self.budget)
432            .finish()
433    }
434}
435
436struct ProcessorState {
437    matcher: Matcher,
438    activation: Activation,
439    initial_activation: Activation,
440    consumed: usize,
441    boundary_forced: bool,
442    boundary_start: Option<usize>,
443    grammar_start: Option<usize>,
444    trailing_grammar_token_id: Option<u32>,
445    trailing_identical_token_count: usize,
446    liveness_intervention_count: usize,
447    last_liveness_intervention_at: Option<usize>,
448}
449
450#[derive(Clone)]
451enum Activation {
452    Active,
453    Boundary {
454        delimiter_tokens: Vec<u32>,
455        forcing: bool,
456    },
457}
458
459impl StructuredOutputProcessor {
460    /// Consume newly generated tokens and hard-mask every illegal next token.
461    /// Waiting-for-reasoning mode leaves normal logits untouched until the
462    /// typed delimiter has been observed.
463    pub fn mask_logits(&self, logits: &mut [f32], generated: &[TokenId]) -> Result<()> {
464        self.mask_logits_inner(logits, generated, None, None)
465            .map(|_| ())
466    }
467
468    /// Apply the grammar mask while allowing engine-resolved stop tokens once
469    /// the grammar accepts. Some model templates use an end-of-turn token that
470    /// is not the tokenizer's primary EOS, so it cannot be represented by the
471    /// grammar parser's EOS set alone.
472    pub fn mask_logits_with_terminals(
473        &self,
474        logits: &mut [f32],
475        generated: &[TokenId],
476        terminal_token_ids: &HashSet<u32>,
477        hidden_control_token_ids: &HashSet<u32>,
478    ) -> Result<StructuredOutputMaskOutcome> {
479        self.mask_logits_inner(
480            logits,
481            generated,
482            Some(terminal_token_ids),
483            Some(hidden_control_token_ids),
484        )
485    }
486
487    fn mask_logits_inner(
488        &self,
489        logits: &mut [f32],
490        generated: &[TokenId],
491        terminal_token_ids: Option<&HashSet<u32>>,
492        hidden_control_token_ids: Option<&HashSet<u32>>,
493    ) -> Result<StructuredOutputMaskOutcome> {
494        let mut state = self.state.lock();
495        advance_state(&mut state, generated, terminal_token_ids)?;
496        activate_forcing_if_due(&mut state, generated, self.budget);
497        if let Activation::Boundary {
498            delimiter_tokens,
499            forcing,
500        } = &state.activation
501        {
502            self.mask_undefined_token_ids(logits);
503            let delimiter_prefix_token_count =
504                delimiter_prefix_token_count(generated, delimiter_tokens);
505            let required_delimiter_token = delimiter_tokens
506                .get(delimiter_prefix_token_count)
507                .copied()
508                .ok_or_else(|| {
509                    FerrumError::internal("structured-output delimiter state has no next token")
510                })?;
511            if *forcing {
512                force_exact_token(logits, required_delimiter_token)?;
513            } else if let Some(hidden_control_token_ids) = hidden_control_token_ids {
514                for token_id in hidden_control_token_ids {
515                    if required_delimiter_token == *token_id {
516                        continue;
517                    }
518                    if let Some(logit) = logits.get_mut(*token_id as usize) {
519                        *logit = f32::NEG_INFINITY;
520                    }
521                }
522            }
523            return Ok(StructuredOutputMaskOutcome {
524                phase: if *forcing {
525                    StructuredOutputPhase::ForcingDelimiter
526                } else {
527                    StructuredOutputPhase::WaitingForDelimiter
528                },
529                accepting: false,
530                liveness_intervention: false,
531                grammar_start_token_index: None,
532                required_delimiter_token_id: Some(required_delimiter_token),
533            });
534        }
535
536        let grammar_start_token_index = state.grammar_start.ok_or_else(|| {
537            FerrumError::internal(
538                "active structured-output processor has no grammar start token index",
539            )
540        })?;
541
542        let accepting = state.matcher.is_accepting().map_err(|error| {
543            FerrumError::model(format!(
544                "structured-output acceptance check failed: {error}"
545            ))
546        })?;
547        let mask = state.matcher.compute_mask_or_eos().map_err(|error| {
548            FerrumError::model(format!("structured-output mask failed: {error}"))
549        })?;
550        let mut finite_allowed = 0usize;
551        for (idx, logit) in logits.iter_mut().enumerate() {
552            let token = idx as u32;
553            let hidden_non_terminal_control = hidden_control_token_ids
554                .is_some_and(|controls| controls.contains(&token))
555                && !terminal_token_ids.is_some_and(|terminals| terminals.contains(&token));
556            let allowed = idx < self.vocab_size
557                && self.defined_token_ids.get(idx).copied().unwrap_or(false)
558                && !hidden_non_terminal_control
559                && (mask.is_allowed(token)
560                    || (accepting
561                        && terminal_token_ids.is_some_and(|terminals| terminals.contains(&token))));
562            if !allowed {
563                *logit = f32::NEG_INFINITY;
564            } else if logit.is_finite() {
565                finite_allowed += 1;
566            }
567        }
568        if finite_allowed == 0 {
569            return Err(FerrumError::model(
570                "structured-output grammar has no legal finite token",
571            ));
572        }
573        let liveness_intervention = if !accepting
574            && state.trailing_identical_token_count >= self.liveness.max_identical_token_run
575        {
576            state
577                .trailing_grammar_token_id
578                .and_then(|token| logits.get_mut(token as usize))
579                .is_some_and(|logit| {
580                    if logit.is_finite() && finite_allowed > 1 {
581                        *logit = f32::NEG_INFINITY;
582                        if state.last_liveness_intervention_at != Some(generated.len()) {
583                            state.liveness_intervention_count += 1;
584                            state.last_liveness_intervention_at = Some(generated.len());
585                        }
586                        true
587                    } else {
588                        false
589                    }
590                })
591        } else {
592            false
593        };
594        Ok(StructuredOutputMaskOutcome {
595            phase: StructuredOutputPhase::EnforcingGrammar,
596            accepting,
597            liveness_intervention,
598            grammar_start_token_index: Some(grammar_start_token_index),
599            required_delimiter_token_id: None,
600        })
601    }
602
603    fn mask_undefined_token_ids(&self, logits: &mut [f32]) {
604        for (idx, logit) in logits.iter_mut().enumerate() {
605            if !self.defined_token_ids.get(idx).copied().unwrap_or(false) {
606                *logit = f32::NEG_INFINITY;
607            }
608        }
609    }
610
611    /// True only when reasoning has closed and the grammar accepts the full
612    /// generated structured value.
613    pub fn is_accepting(&self, generated: &[TokenId]) -> Result<bool> {
614        self.is_accepting_inner(generated, None)
615    }
616
617    /// Completion check that treats an engine-resolved terminal sampled after
618    /// grammar acceptance as framing rather than part of the JSON value.
619    pub fn is_accepting_with_terminals(
620        &self,
621        generated: &[TokenId],
622        terminal_token_ids: &HashSet<u32>,
623    ) -> Result<bool> {
624        self.is_accepting_inner(generated, Some(terminal_token_ids))
625    }
626
627    fn is_accepting_inner(
628        &self,
629        generated: &[TokenId],
630        terminal_token_ids: Option<&HashSet<u32>>,
631    ) -> Result<bool> {
632        Ok(self
633            .progress_inner(generated, terminal_token_ids)?
634            .accepting)
635    }
636
637    /// Inspect the typed activation/grammar state after consuming `generated`.
638    pub fn progress_with_terminals(
639        &self,
640        generated: &[TokenId],
641        terminal_token_ids: &HashSet<u32>,
642    ) -> Result<StructuredOutputProgress> {
643        self.progress_inner(generated, Some(terminal_token_ids))
644    }
645
646    fn progress_inner(
647        &self,
648        generated: &[TokenId],
649        terminal_token_ids: Option<&HashSet<u32>>,
650    ) -> Result<StructuredOutputProgress> {
651        let mut state = self.state.lock();
652        advance_state(&mut state, generated, terminal_token_ids)?;
653        activate_forcing_if_due(&mut state, generated, self.budget);
654        let (phase, delimiter_token_count, delimiter_prefix_token_count, accepting) =
655            match &state.activation {
656                Activation::Boundary {
657                    delimiter_tokens,
658                    forcing,
659                } => (
660                    if *forcing {
661                        StructuredOutputPhase::ForcingDelimiter
662                    } else {
663                        StructuredOutputPhase::WaitingForDelimiter
664                    },
665                    Some(delimiter_tokens.len()),
666                    delimiter_prefix_token_count(generated, delimiter_tokens),
667                    false,
668                ),
669                Activation::Active => (
670                    StructuredOutputPhase::EnforcingGrammar,
671                    None,
672                    0,
673                    state.matcher.is_accepting().map_err(|error| {
674                        FerrumError::model(format!(
675                            "structured-output acceptance check failed: {error}"
676                        ))
677                    })?,
678                ),
679            };
680        let grammar_tokens = state
681            .grammar_start
682            .and_then(|start| generated.get(start..))
683            .unwrap_or_default();
684        let trailing_token_id = grammar_tokens.last().map(|token| token.get());
685        let trailing_token_class = trailing_token_id.map(|token| {
686            self.json_token_classes
687                .get(token as usize)
688                .copied()
689                .unwrap_or(StructuredOutputTokenClass::Undefined)
690        });
691        let trailing_token_class_count = trailing_token_class.map_or(0, |class| {
692            grammar_tokens
693                .iter()
694                .rev()
695                .take_while(|token| {
696                    self.json_token_classes
697                        .get(token.get() as usize)
698                        .copied()
699                        .unwrap_or(StructuredOutputTokenClass::Undefined)
700                        == class
701                })
702                .count()
703        });
704        let trailing_identical_token_count = trailing_token_id.map_or(0, |token_id| {
705            grammar_tokens
706                .iter()
707                .rev()
708                .take_while(|token| token.get() == token_id)
709                .count()
710        });
711        Ok(StructuredOutputProgress {
712            phase,
713            generated_token_count: generated.len(),
714            consumed_token_count: state.consumed,
715            delimiter_token_count: delimiter_token_count
716                .or(self.budget.map(|budget| budget.boundary_token_count)),
717            delimiter_prefix_token_count,
718            reasoning_token_count: self
719                .budget
720                .map(|_| state.boundary_start.unwrap_or(generated.len())),
721            boundary_forced: state.boundary_forced,
722            budget: self.budget,
723            grammar_token_count: grammar_tokens.len(),
724            trailing_token_class,
725            trailing_token_class_count,
726            trailing_token_id,
727            trailing_identical_token_count,
728            liveness_identical_token_limit: self.liveness.max_identical_token_run,
729            liveness_intervention_count: state.liveness_intervention_count,
730            accepting,
731        })
732    }
733
734    pub fn reset(&self) -> Result<()> {
735        let mut state = self.state.lock();
736        state
737            .matcher
738            .reset()
739            .map_err(|error| FerrumError::internal(format!("reset structured output: {error}")))?;
740        state.activation = state.initial_activation.clone();
741        state.consumed = 0;
742        state.boundary_forced = false;
743        state.boundary_start = None;
744        state.grammar_start = matches!(state.initial_activation, Activation::Active).then_some(0);
745        state.trailing_grammar_token_id = None;
746        state.trailing_identical_token_count = 0;
747        state.liveness_intervention_count = 0;
748        state.last_liveness_intervention_at = None;
749        Ok(())
750    }
751}
752
753fn activate_forcing_if_due(
754    state: &mut ProcessorState,
755    generated: &[TokenId],
756    budget: Option<StructuredOutputBudgetPlan>,
757) {
758    let Some(budget) = budget else {
759        return;
760    };
761    let should_force = matches!(
762        state.activation,
763        Activation::Boundary { forcing: false, .. }
764    ) && generated.len() >= budget.reasoning_token_limit;
765    if should_force {
766        let delimiter_prefix_token_count = match &state.activation {
767            Activation::Boundary {
768                delimiter_tokens, ..
769            } => delimiter_prefix_token_count(generated, delimiter_tokens),
770            Activation::Active => 0,
771        };
772        if let Activation::Boundary { forcing, .. } = &mut state.activation {
773            *forcing = true;
774        }
775        state.boundary_forced = true;
776        state.boundary_start = Some(generated.len() - delimiter_prefix_token_count);
777    }
778}
779
780fn force_exact_token(logits: &mut [f32], required_token: u32) -> Result<()> {
781    let required_index = required_token as usize;
782    if required_index >= logits.len() {
783        return Err(FerrumError::model(format!(
784            "structured-output delimiter token {required_token} is outside logits width {}",
785            logits.len()
786        )));
787    }
788    logits.fill(f32::NEG_INFINITY);
789    logits[required_index] = 0.0;
790    Ok(())
791}
792
793fn delimiter_prefix_token_count(generated: &[TokenId], delimiter_tokens: &[u32]) -> usize {
794    let max_prefix = generated
795        .len()
796        .min(delimiter_tokens.len().saturating_sub(1));
797    (1..=max_prefix)
798        .rev()
799        .find(|prefix_len| {
800            generated[generated.len() - prefix_len..]
801                .iter()
802                .zip(&delimiter_tokens[..*prefix_len])
803                .all(|(token, expected)| token.get() == *expected)
804        })
805        .unwrap_or(0)
806}
807
808fn advance_state(
809    state: &mut ProcessorState,
810    generated: &[TokenId],
811    terminal_token_ids: Option<&HashSet<u32>>,
812) -> Result<()> {
813    if state.consumed > generated.len() {
814        return Err(FerrumError::internal(
815            "structured-output token history moved backwards without reset",
816        ));
817    }
818
819    if let Activation::Boundary {
820        delimiter_tokens, ..
821    } = &state.activation
822    {
823        let search_from = state.consumed.saturating_sub(delimiter_tokens.len());
824        if let Some(offset) = generated[search_from..]
825            .windows(delimiter_tokens.len())
826            .position(|window| {
827                window
828                    .iter()
829                    .zip(delimiter_tokens)
830                    .all(|(token, expected)| token.get() == *expected)
831            })
832        {
833            let grammar_start = search_from + offset + delimiter_tokens.len();
834            state.boundary_start = Some(grammar_start - delimiter_tokens.len());
835            state.grammar_start = Some(grammar_start);
836            state.activation = Activation::Active;
837            state.consumed = grammar_start;
838            state.trailing_grammar_token_id = None;
839            state.trailing_identical_token_count = 0;
840        } else {
841            state.consumed = generated.len();
842            return Ok(());
843        }
844    }
845
846    for token in &generated[state.consumed..] {
847        if terminal_token_ids.is_some_and(|terminals| terminals.contains(&token.get()))
848            && state.matcher.is_accepting().map_err(|error| {
849                FerrumError::model(format!(
850                    "structured-output acceptance check failed: {error}"
851                ))
852            })?
853        {
854            continue;
855        }
856        state.matcher.consume_token(token.get()).map_err(|error| {
857            FerrumError::model(format!(
858                "structured-output token {} violated the grammar: {error}",
859                token.get()
860            ))
861        })?;
862        if state.trailing_grammar_token_id == Some(token.get()) {
863            state.trailing_identical_token_count += 1;
864        } else {
865            state.trailing_grammar_token_id = Some(token.get());
866            state.trailing_identical_token_count = 1;
867        }
868    }
869    state.consumed = generated.len();
870    Ok(())
871}
872
873struct FerrumTokenizerEnv {
874    tokenizer: Arc<dyn Tokenizer + Send + Sync>,
875    trie: TokTrie,
876}
877
878impl TokenizerEnv for FerrumTokenizerEnv {
879    fn tok_trie(&self) -> &TokTrie {
880        &self.trie
881    }
882
883    fn tokenize_bytes(&self, bytes: &[u8]) -> Vec<u32> {
884        str::from_utf8(bytes)
885            .ok()
886            .and_then(|text| self.tokenizer.encode(text, false).ok())
887            .map(|tokens| tokens.into_iter().map(|token| token.get()).collect())
888            .unwrap_or_else(|| self.trie.greedy_tokenize(bytes))
889    }
890
891    fn tokenize_is_canonical(&self) -> bool {
892        false
893    }
894}
895
896fn tokenizer_special_ids(tokenizer: &(dyn Tokenizer + Send + Sync)) -> HashSet<u32> {
897    let special = tokenizer.special_tokens();
898    [
899        special.bos_token,
900        special.eos_token,
901        special.unk_token,
902        special.pad_token,
903        special.sep_token,
904        special.cls_token,
905        special.mask_token,
906    ]
907    .into_iter()
908    .flatten()
909    .chain(special.extra_eos_tokens.iter().copied())
910    .map(|token| token.get())
911    .collect()
912}
913
914fn special_token_marker(token: TokenId) -> Vec<u8> {
915    let mut marker = vec![TokTrie::SPECIAL_TOKEN_MARKER];
916    marker.extend_from_slice(format!("[{}]", token.get()).as_bytes());
917    marker
918}
919
920fn classify_json_token_bytes(bytes: &[u8]) -> StructuredOutputTokenClass {
921    if bytes.is_empty() {
922        StructuredOutputTokenClass::Undefined
923    } else if bytes
924        .iter()
925        .all(|byte| matches!(byte, b' ' | b'\n' | b'\r' | b'\t'))
926    {
927        StructuredOutputTokenClass::Whitespace
928    } else if bytes
929        .iter()
930        .all(|byte| byte.is_ascii_digit() || matches!(byte, b'-' | b'+' | b'.' | b'e' | b'E'))
931    {
932        StructuredOutputTokenClass::Number
933    } else if bytes
934        .iter()
935        .all(|byte| matches!(byte, b'{' | b'}' | b'[' | b']' | b',' | b':'))
936    {
937        StructuredOutputTokenClass::Structural
938    } else if bytes.iter().all(|byte| matches!(byte, b'"' | b'\\')) {
939        StructuredOutputTokenClass::StringBoundary
940    } else if bytes.iter().all(u8::is_ascii_alphabetic) {
941        StructuredOutputTokenClass::Literal
942    } else {
943        StructuredOutputTokenClass::Other
944    }
945}
946
947#[cfg(test)]
948mod tests {
949    use super::*;
950    use ferrum_interfaces::tokenizer::{ChatMessage, TokenizerInfo, TokenizerType};
951    use ferrum_types::SpecialTokens;
952
953    const EOS: u32 = 256;
954    const TEST_MAX_OUTPUT_TOKENS: usize = 128;
955
956    struct ByteTokenizer {
957        special: SpecialTokens,
958        token_text: Vec<String>,
959    }
960
961    impl ByteTokenizer {
962        fn new() -> Self {
963            let mut token_text = (0u16..=255)
964                .map(|byte| char::from_u32(byte as u32).unwrap().to_string())
965                .collect::<Vec<_>>();
966            token_text.push("<eos>".to_string());
967            Self {
968                special: SpecialTokens {
969                    eos_token: Some(TokenId::new(EOS)),
970                    ..SpecialTokens::default()
971                },
972                token_text,
973            }
974        }
975    }
976
977    impl Tokenizer for ByteTokenizer {
978        fn encode(&self, text: &str, _add_special: bool) -> Result<Vec<TokenId>> {
979            Ok(text
980                .as_bytes()
981                .iter()
982                .map(|byte| TokenId::new(*byte as u32))
983                .collect())
984        }
985
986        fn decode(&self, tokens: &[TokenId], _skip_special: bool) -> Result<String> {
987            Ok(tokens
988                .iter()
989                .filter(|token| token.get() < 256)
990                .map(|token| token.get() as u8 as char)
991                .collect())
992        }
993
994        fn decode_incremental(&self, _prev: &[TokenId], next: TokenId) -> Result<String> {
995            self.decode(&[next], true)
996        }
997
998        fn vocab_size(&self) -> usize {
999            self.token_text.len()
1000        }
1001
1002        fn special_tokens(&self) -> &SpecialTokens {
1003            &self.special
1004        }
1005
1006        fn token_id(&self, text: &str) -> Option<TokenId> {
1007            (text.len() == 1).then(|| TokenId::new(text.as_bytes()[0] as u32))
1008        }
1009
1010        fn token_text(&self, token_id: TokenId) -> Option<&str> {
1011            self.token_text
1012                .get(token_id.get() as usize)
1013                .map(String::as_str)
1014        }
1015
1016        fn apply_chat_template(&self, messages: &[ChatMessage]) -> Result<String> {
1017            Ok(messages
1018                .iter()
1019                .map(|message| message.content.as_str())
1020                .collect::<Vec<_>>()
1021                .join("\n"))
1022        }
1023
1024        fn info(&self) -> TokenizerInfo {
1025            TokenizerInfo {
1026                tokenizer_type: TokenizerType::BPE,
1027                vocab_size: self.vocab_size(),
1028                special_tokens: self.special.clone(),
1029                supports_incremental: true,
1030                supports_chat_template: false,
1031                max_token_length: Some(1),
1032                model_name: Some("byte-test".to_string()),
1033            }
1034        }
1035    }
1036
1037    struct MergedObjectTokenizer {
1038        inner: ByteTokenizer,
1039    }
1040
1041    impl MergedObjectTokenizer {
1042        const OBJECT: u32 = 256;
1043        const EOS: u32 = 257;
1044
1045        fn new() -> Self {
1046            let mut inner = ByteTokenizer::new();
1047            inner.token_text[EOS as usize] = "{}".to_string();
1048            inner.token_text.push("<eos>".to_string());
1049            inner.special.eos_token = Some(TokenId::new(Self::EOS));
1050            Self { inner }
1051        }
1052    }
1053
1054    impl Tokenizer for MergedObjectTokenizer {
1055        fn encode(&self, text: &str, add_special: bool) -> Result<Vec<TokenId>> {
1056            self.inner.encode(text, add_special)
1057        }
1058
1059        fn decode(&self, tokens: &[TokenId], skip_special: bool) -> Result<String> {
1060            let mut decoded = String::new();
1061            for token in tokens {
1062                match token.get() {
1063                    Self::OBJECT => decoded.push_str("{}"),
1064                    Self::EOS if skip_special => {}
1065                    Self::EOS => decoded.push_str("<eos>"),
1066                    _ => decoded.push_str(&self.inner.decode(&[*token], skip_special)?),
1067                }
1068            }
1069            Ok(decoded)
1070        }
1071
1072        fn decode_incremental(&self, _prev: &[TokenId], next: TokenId) -> Result<String> {
1073            self.decode(&[next], true)
1074        }
1075
1076        fn vocab_size(&self) -> usize {
1077            self.inner.token_text.len()
1078        }
1079
1080        fn special_tokens(&self) -> &SpecialTokens {
1081            &self.inner.special
1082        }
1083
1084        fn token_id(&self, text: &str) -> Option<TokenId> {
1085            (text == "{}")
1086                .then(|| TokenId::new(Self::OBJECT))
1087                .or_else(|| self.inner.token_id(text))
1088        }
1089
1090        fn token_text(&self, token_id: TokenId) -> Option<&str> {
1091            self.inner.token_text(token_id)
1092        }
1093
1094        fn info(&self) -> TokenizerInfo {
1095            TokenizerInfo {
1096                vocab_size: self.vocab_size(),
1097                max_token_length: Some(2),
1098                model_name: Some("merged-object-test".to_string()),
1099                ..self.inner.info()
1100            }
1101        }
1102    }
1103
1104    fn factory() -> StructuredOutputFactory {
1105        StructuredOutputFactory::new(Arc::new(ByteTokenizer::new())).unwrap()
1106    }
1107
1108    fn assert_and_append(
1109        processor: &StructuredOutputProcessor,
1110        generated: &mut Vec<TokenId>,
1111        text: &str,
1112    ) {
1113        for byte in text.bytes() {
1114            let mut logits = vec![0.0; EOS as usize + 1];
1115            processor.mask_logits(&mut logits, generated).unwrap();
1116            assert!(
1117                logits[byte as usize].is_finite(),
1118                "byte {byte:?} rejected after {:?}",
1119                generated
1120            );
1121            generated.push(TokenId::new(byte as u32));
1122        }
1123    }
1124
1125    #[test]
1126    fn json_object_hard_masks_non_object_roots() {
1127        let processor = factory()
1128            .create_processor(
1129                &ResponseFormat::JsonObject,
1130                &StructuredOutputStart::Immediate,
1131                TEST_MAX_OUTPUT_TOKENS,
1132                &HashSet::new(),
1133                &[],
1134            )
1135            .unwrap()
1136            .unwrap();
1137        let mut logits = vec![0.0; EOS as usize + 1];
1138        processor.mask_logits(&mut logits, &[]).unwrap();
1139        assert!(logits[b'{' as usize].is_finite());
1140        assert!(!logits[b'[' as usize].is_finite());
1141        assert!(!logits[b'`' as usize].is_finite());
1142        assert!(!logits[EOS as usize].is_finite());
1143    }
1144
1145    #[test]
1146    fn json_object_uses_compact_separators_without_unbounded_whitespace() {
1147        let processor = factory()
1148            .create_processor(
1149                &ResponseFormat::JsonObject,
1150                &StructuredOutputStart::Immediate,
1151                TEST_MAX_OUTPUT_TOKENS,
1152                &HashSet::new(),
1153                &[],
1154            )
1155            .unwrap()
1156            .unwrap();
1157        let generated = vec![TokenId::new(b'{' as u32)];
1158        let mut logits = vec![0.0; EOS as usize + 1];
1159        processor.mask_logits(&mut logits, &generated).unwrap();
1160
1161        assert!(!logits[b' ' as usize].is_finite());
1162        assert!(logits[b'}' as usize].is_finite());
1163        assert!(logits[b'"' as usize].is_finite());
1164    }
1165
1166    #[test]
1167    fn undefined_model_vocab_ids_are_masked_inside_wildcard_strings() {
1168        let tokenizer = Arc::new(ByteTokenizer::new());
1169        let undefined_token = tokenizer.vocab_size() as u32;
1170        assert_eq!(
1171            tokenizer.token_bytes(TokenId::new(undefined_token)),
1172            Some(Vec::new()),
1173            "the test tokenizer must reproduce a decoder that returns empty text for an unknown id"
1174        );
1175        let processor = StructuredOutputFactory::new_with_model_vocab_size(
1176            tokenizer,
1177            Some(undefined_token as usize + 1),
1178        )
1179        .unwrap()
1180        .create_processor(
1181            &ResponseFormat::JsonObject,
1182            &StructuredOutputStart::Immediate,
1183            TEST_MAX_OUTPUT_TOKENS,
1184            &HashSet::new(),
1185            &[],
1186        )
1187        .unwrap()
1188        .unwrap();
1189
1190        let mut generated = Vec::new();
1191        assert_and_append(&processor, &mut generated, r#"{"value":"Ferrum "#);
1192        let mut logits = vec![0.0; undefined_token as usize + 1];
1193        processor.mask_logits(&mut logits, &generated).unwrap();
1194        assert!(logits[b'x' as usize].is_finite());
1195        assert!(!logits[undefined_token as usize].is_finite());
1196    }
1197
1198    #[test]
1199    fn json_object_accepts_nested_unicode_escape_and_eos_only_after_close() {
1200        let processor = factory()
1201            .create_processor(
1202                &ResponseFormat::JsonObject,
1203                &StructuredOutputStart::Immediate,
1204                TEST_MAX_OUTPUT_TOKENS,
1205                &HashSet::new(),
1206                &[],
1207            )
1208            .unwrap()
1209            .unwrap();
1210        let mut generated = Vec::new();
1211        assert_and_append(
1212            &processor,
1213            &mut generated,
1214            r#"{"items":[true,null,{"name":"line\u000A"}],"n":-1.2e+3}"#,
1215        );
1216        assert!(processor.is_accepting(&generated).unwrap());
1217        let mut logits = vec![0.0; EOS as usize + 1];
1218        processor.mask_logits(&mut logits, &generated).unwrap();
1219        assert!(logits[EOS as usize].is_finite());
1220        assert!(!logits[b'x' as usize].is_finite());
1221    }
1222
1223    #[test]
1224    fn json_object_accepts_a_complete_root_from_one_merged_token() {
1225        let processor = StructuredOutputFactory::new(Arc::new(MergedObjectTokenizer::new()))
1226            .unwrap()
1227            .create_processor(
1228                &ResponseFormat::JsonObject,
1229                &StructuredOutputStart::Immediate,
1230                TEST_MAX_OUTPUT_TOKENS,
1231                &HashSet::new(),
1232                &[],
1233            )
1234            .unwrap()
1235            .unwrap();
1236        let generated = vec![TokenId::new(MergedObjectTokenizer::OBJECT)];
1237        let terminals = HashSet::from([MergedObjectTokenizer::EOS]);
1238
1239        let progress = processor
1240            .progress_with_terminals(&generated, &terminals)
1241            .unwrap();
1242        assert!(progress.accepting);
1243
1244        let mut logits = vec![0.0; MergedObjectTokenizer::EOS as usize + 1];
1245        let outcome = processor
1246            .mask_logits_with_terminals(&mut logits, &generated, &terminals, &HashSet::new())
1247            .unwrap();
1248        assert!(outcome.accepting);
1249        assert_eq!(outcome.grammar_start_token_index, Some(0));
1250        assert!(!logits[MergedObjectTokenizer::OBJECT as usize].is_finite());
1251        assert!(logits[MergedObjectTokenizer::EOS as usize].is_finite());
1252    }
1253
1254    #[test]
1255    fn json_object_breaks_an_unbounded_identical_token_run_when_closure_is_legal() {
1256        let processor = StructuredOutputFactory::new(Arc::new(MergedObjectTokenizer::new()))
1257            .unwrap()
1258            .create_processor(
1259                &ResponseFormat::JsonObject,
1260                &StructuredOutputStart::Immediate,
1261                64,
1262                &HashSet::new(),
1263                &[],
1264            )
1265            .unwrap()
1266            .unwrap();
1267        let mut generated =
1268            br#"{"marker":""#.iter().map(|byte| TokenId::new(*byte as u32)).collect::<Vec<_>>();
1269        generated.extend(std::iter::repeat_n(
1270            TokenId::new(MergedObjectTokenizer::OBJECT),
1271            32,
1272        ));
1273
1274        let mut logits = vec![0.0; MergedObjectTokenizer::EOS as usize + 1];
1275        let outcome = processor
1276            .mask_logits_with_terminals(
1277                &mut logits,
1278                &generated,
1279                &HashSet::from([MergedObjectTokenizer::EOS]),
1280                &HashSet::new(),
1281            )
1282            .unwrap();
1283
1284        assert!(!outcome.accepting);
1285        assert!(outcome.liveness_intervention);
1286        assert!(!logits[MergedObjectTokenizer::OBJECT as usize].is_finite());
1287        assert!(logits[b'"' as usize].is_finite());
1288
1289        generated.extend([TokenId::new(b'"' as u32), TokenId::new(b'}' as u32)]);
1290        let progress = processor
1291            .progress_with_terminals(&generated, &HashSet::from([MergedObjectTokenizer::EOS]))
1292            .unwrap();
1293        assert!(progress.accepting);
1294        assert_eq!(progress.liveness_identical_token_limit, 32);
1295        assert_eq!(progress.liveness_intervention_count, 1);
1296    }
1297
1298    #[test]
1299    fn structured_liveness_guard_preserves_the_only_finite_grammar_candidate() {
1300        let processor = StructuredOutputFactory::new(Arc::new(MergedObjectTokenizer::new()))
1301            .unwrap()
1302            .create_processor(
1303                &ResponseFormat::JsonObject,
1304                &StructuredOutputStart::Immediate,
1305                64,
1306                &HashSet::new(),
1307                &[],
1308            )
1309            .unwrap()
1310            .unwrap();
1311        let mut generated =
1312            br#"{"marker":""#.iter().map(|byte| TokenId::new(*byte as u32)).collect::<Vec<_>>();
1313        generated.extend(std::iter::repeat_n(
1314            TokenId::new(MergedObjectTokenizer::OBJECT),
1315            32,
1316        ));
1317        let mut logits = vec![f32::NEG_INFINITY; MergedObjectTokenizer::EOS as usize + 1];
1318        logits[MergedObjectTokenizer::OBJECT as usize] = 0.0;
1319
1320        let outcome = processor
1321            .mask_logits_with_terminals(
1322                &mut logits,
1323                &generated,
1324                &HashSet::from([MergedObjectTokenizer::EOS]),
1325                &HashSet::new(),
1326            )
1327            .unwrap();
1328
1329        assert!(!outcome.liveness_intervention);
1330        assert!(logits[MergedObjectTokenizer::OBJECT as usize].is_finite());
1331    }
1332
1333    #[test]
1334    fn structured_liveness_limit_is_derived_from_the_guaranteed_result_budget() {
1335        let budget = StructuredOutputBudgetPlan::automatic(4096, 1).unwrap();
1336        assert_eq!(budget.structured_reserve_tokens, 1024);
1337        assert_eq!(
1338            StructuredOutputLivenessPolicy::for_request(4096, Some(budget)).max_identical_token_run,
1339            512
1340        );
1341        assert_eq!(
1342            StructuredOutputLivenessPolicy::for_request(64, None).max_identical_token_run,
1343            32
1344        );
1345    }
1346
1347    #[test]
1348    fn strict_schema_rejects_wrong_property_and_accepts_required_value() {
1349        let schema = r#"{
1350            "type":"object",
1351            "properties":{"answer":{"const":42}},
1352            "required":["answer"],
1353            "additionalProperties":false
1354        }"#;
1355        let processor = factory()
1356            .create_processor(
1357                &ResponseFormat::JsonSchema(schema.to_string()),
1358                &StructuredOutputStart::Immediate,
1359                TEST_MAX_OUTPUT_TOKENS,
1360                &HashSet::new(),
1361                &[],
1362            )
1363            .unwrap()
1364            .unwrap();
1365        let mut generated = Vec::new();
1366        assert_and_append(&processor, &mut generated, r#"{"answer":42}"#);
1367        assert!(processor.is_accepting(&generated).unwrap());
1368    }
1369
1370    #[test]
1371    fn request_cannot_override_compact_json_compiler_policy() {
1372        let schema = r#"{
1373            "type":"object",
1374            "properties":{"answer":{"const":42}},
1375            "required":["answer"],
1376            "additionalProperties":false,
1377            "x-guidance":{
1378                "item_separator":", ",
1379                "key_separator":": ",
1380                "whitespace_flexible":true
1381            }
1382        }"#;
1383        let processor = factory()
1384            .create_processor(
1385                &ResponseFormat::JsonSchema(schema.to_string()),
1386                &StructuredOutputStart::Immediate,
1387                TEST_MAX_OUTPUT_TOKENS,
1388                &HashSet::new(),
1389                &[],
1390            )
1391            .unwrap()
1392            .unwrap();
1393        let mut generated = Vec::new();
1394        assert_and_append(&processor, &mut generated, r#"{"answer":"#);
1395        let mut logits = vec![0.0; EOS as usize + 1];
1396        processor.mask_logits(&mut logits, &generated).unwrap();
1397
1398        assert!(!logits[b' ' as usize].is_finite());
1399        assert!(logits[b'4' as usize].is_finite());
1400        assert_and_append(&processor, &mut generated, "42}");
1401        assert!(processor.is_accepting(&generated).unwrap());
1402    }
1403
1404    #[test]
1405    fn boolean_json_schema_keeps_its_semantics_under_compact_policy() {
1406        let processor = factory()
1407            .create_processor(
1408                &ResponseFormat::JsonSchema("true".to_string()),
1409                &StructuredOutputStart::Immediate,
1410                TEST_MAX_OUTPUT_TOKENS,
1411                &HashSet::new(),
1412                &[],
1413            )
1414            .unwrap()
1415            .unwrap();
1416        let mut generated = Vec::new();
1417        assert_and_append(&processor, &mut generated, "true");
1418        assert!(processor.is_accepting(&generated).unwrap());
1419    }
1420
1421    #[test]
1422    fn terminal_progress_classifies_an_unclosed_number_without_retaining_text() {
1423        let processor = factory()
1424            .create_processor(
1425                &ResponseFormat::JsonSchema(
1426                    r#"{"type":"object","properties":{"value":{"type":"integer"}},"required":["value"],"additionalProperties":false}"#
1427                        .to_string(),
1428                ),
1429                &StructuredOutputStart::Immediate,
1430                TEST_MAX_OUTPUT_TOKENS,
1431                &HashSet::new(),
1432                &[],
1433            )
1434            .unwrap()
1435            .unwrap();
1436        let generated =
1437            r#"{"value":123777"#.bytes().map(|byte| TokenId::new(byte as u32)).collect::<Vec<_>>();
1438        let progress = processor
1439            .progress_with_terminals(&generated, &HashSet::new())
1440            .unwrap();
1441
1442        assert_eq!(progress.phase, StructuredOutputPhase::EnforcingGrammar);
1443        assert_eq!(progress.grammar_token_count, generated.len());
1444        assert_eq!(
1445            progress.trailing_token_class,
1446            Some(StructuredOutputTokenClass::Number)
1447        );
1448        assert_eq!(progress.trailing_token_class_count, 6);
1449        assert_eq!(progress.trailing_token_id, Some(b'7' as u32));
1450        assert_eq!(progress.trailing_identical_token_count, 3);
1451        assert!(!progress.accepting);
1452    }
1453
1454    #[test]
1455    fn lexical_diagnostics_classify_json_bytes_without_decoding_content() {
1456        assert_eq!(
1457            classify_json_token_bytes(b" \n\t"),
1458            StructuredOutputTokenClass::Whitespace
1459        );
1460        assert_eq!(
1461            classify_json_token_bytes(b"-12.5e+3"),
1462            StructuredOutputTokenClass::Number
1463        );
1464        assert_eq!(
1465            classify_json_token_bytes(br#"{}[],:"#),
1466            StructuredOutputTokenClass::Structural
1467        );
1468        assert_eq!(
1469            classify_json_token_bytes(br#"\""#),
1470            StructuredOutputTokenClass::StringBoundary
1471        );
1472        assert_eq!(
1473            classify_json_token_bytes(b"true"),
1474            StructuredOutputTokenClass::Literal
1475        );
1476        assert_eq!(
1477            classify_json_token_bytes(br#""value":"#),
1478            StructuredOutputTokenClass::Other
1479        );
1480    }
1481
1482    struct FragmentedUtf8Tokenizer {
1483        special: SpecialTokens,
1484        token_text: Vec<String>,
1485    }
1486
1487    impl FragmentedUtf8Tokenizer {
1488        const FIRE_HEAD: u32 = 128;
1489        const FIRE_TAIL: u32 = 129;
1490        const EOS: u32 = 130;
1491
1492        fn new() -> Self {
1493            let mut token_text = (0u8..=127)
1494                .map(|byte| (byte as char).to_string())
1495                .collect::<Vec<_>>();
1496            token_text.extend(["\u{fffd}".to_string(), "\u{fffd}".to_string()]);
1497            token_text.push("<eos>".to_string());
1498            Self {
1499                special: SpecialTokens {
1500                    eos_token: Some(TokenId::new(Self::EOS)),
1501                    ..SpecialTokens::default()
1502                },
1503                token_text,
1504            }
1505        }
1506
1507        fn raw_bytes(token: TokenId) -> Option<Vec<u8>> {
1508            match token.get() {
1509                byte @ 0..=127 => Some(vec![byte as u8]),
1510                Self::FIRE_HEAD => Some(vec![0xf0, 0x9f]),
1511                Self::FIRE_TAIL => Some(vec![0x94, 0xa5]),
1512                _ => None,
1513            }
1514        }
1515    }
1516
1517    impl Tokenizer for FragmentedUtf8Tokenizer {
1518        fn encode(&self, text: &str, _add_special: bool) -> Result<Vec<TokenId>> {
1519            let mut tokens = Vec::new();
1520            let mut bytes = text.as_bytes();
1521            while let Some((&byte, remaining)) = bytes.split_first() {
1522                if bytes.starts_with(&[0xf0, 0x9f, 0x94, 0xa5]) {
1523                    tokens.push(TokenId::new(Self::FIRE_HEAD));
1524                    tokens.push(TokenId::new(Self::FIRE_TAIL));
1525                    bytes = &bytes[4..];
1526                } else if byte <= 127 {
1527                    tokens.push(TokenId::new(byte as u32));
1528                    bytes = remaining;
1529                } else {
1530                    return Err(FerrumError::tokenizer(
1531                        "fragmented UTF-8 test tokenizer received unsupported input",
1532                    ));
1533                }
1534            }
1535            Ok(tokens)
1536        }
1537
1538        fn decode(&self, tokens: &[TokenId], _skip_special: bool) -> Result<String> {
1539            let bytes = tokens
1540                .iter()
1541                .filter_map(|token| Self::raw_bytes(*token))
1542                .flatten()
1543                .collect::<Vec<_>>();
1544            Ok(String::from_utf8_lossy(&bytes).into_owned())
1545        }
1546
1547        fn decode_incremental(&self, _prev: &[TokenId], next: TokenId) -> Result<String> {
1548            self.decode(&[next], true)
1549        }
1550
1551        fn vocab_size(&self) -> usize {
1552            self.token_text.len()
1553        }
1554
1555        fn special_tokens(&self) -> &SpecialTokens {
1556            &self.special
1557        }
1558
1559        fn token_id(&self, text: &str) -> Option<TokenId> {
1560            (text.len() == 1 && text.is_ascii()).then(|| TokenId::new(text.as_bytes()[0] as u32))
1561        }
1562
1563        fn token_text(&self, token_id: TokenId) -> Option<&str> {
1564            self.token_text
1565                .get(token_id.get() as usize)
1566                .map(String::as_str)
1567        }
1568
1569        fn token_bytes(&self, token_id: TokenId) -> Option<Vec<u8>> {
1570            Self::raw_bytes(token_id)
1571        }
1572
1573        fn info(&self) -> TokenizerInfo {
1574            TokenizerInfo {
1575                tokenizer_type: TokenizerType::BPE,
1576                vocab_size: self.vocab_size(),
1577                special_tokens: self.special.clone(),
1578                supports_incremental: true,
1579                supports_chat_template: false,
1580                max_token_length: Some(2),
1581                model_name: Some("fragmented-utf8-test".to_string()),
1582            }
1583        }
1584    }
1585
1586    #[test]
1587    fn strict_schema_accepts_utf8_split_across_byte_level_tokens() {
1588        let tokenizer = Arc::new(FragmentedUtf8Tokenizer::new());
1589        assert!(tokenizer
1590            .decode(&[TokenId::new(FragmentedUtf8Tokenizer::FIRE_HEAD)], false)
1591            .unwrap()
1592            .contains('\u{fffd}'));
1593        let processor = StructuredOutputFactory::new(tokenizer)
1594            .unwrap()
1595            .create_processor(
1596                &ResponseFormat::JsonSchema(
1597                    r#"{"type":"object","properties":{"value":{"const":"\ud83d\udd25"}},"required":["value"],"additionalProperties":false}"#
1598                        .to_string(),
1599                ),
1600                &StructuredOutputStart::Immediate,
1601                TEST_MAX_OUTPUT_TOKENS,
1602                &HashSet::new(),
1603                &[],
1604            )
1605            .unwrap()
1606            .unwrap();
1607
1608        let mut generated = Vec::new();
1609        for token in r#"{"value":""#
1610            .bytes()
1611            .map(|byte| TokenId::new(byte as u32))
1612            .chain([
1613                TokenId::new(FragmentedUtf8Tokenizer::FIRE_HEAD),
1614                TokenId::new(FragmentedUtf8Tokenizer::FIRE_TAIL),
1615            ])
1616            .chain(r#""}"#.bytes().map(|byte| TokenId::new(byte as u32)))
1617        {
1618            let mut logits = vec![0.0; FragmentedUtf8Tokenizer::EOS as usize + 1];
1619            processor.mask_logits(&mut logits, &generated).unwrap();
1620            assert!(
1621                logits[token.get() as usize].is_finite(),
1622                "token {} rejected after {:?}",
1623                token.get(),
1624                generated
1625            );
1626            generated.push(token);
1627        }
1628        assert!(processor.is_accepting(&generated).unwrap());
1629    }
1630
1631    #[test]
1632    fn reasoning_delimiter_defers_then_activates_the_grammar() {
1633        let processor = factory()
1634            .create_processor(
1635                &ResponseFormat::JsonObject,
1636                &StructuredOutputStart::AfterDelimiter("</think>".to_string()),
1637                TEST_MAX_OUTPUT_TOKENS,
1638                &HashSet::new(),
1639                &[],
1640            )
1641            .unwrap()
1642            .unwrap();
1643        let mut generated = Vec::new();
1644        let controls = HashSet::from([b'<' as u32, b'>' as u32]);
1645        let mut waiting_logits = vec![0.0; EOS as usize + 1];
1646        let waiting = processor
1647            .mask_logits_with_terminals(&mut waiting_logits, &generated, &HashSet::new(), &controls)
1648            .unwrap();
1649        assert_eq!(waiting.phase, StructuredOutputPhase::WaitingForDelimiter);
1650        assert!(!waiting.accepting);
1651        assert_eq!(waiting.grammar_start_token_index, None);
1652        assert_eq!(waiting.required_delimiter_token_id, Some(b'<' as u32));
1653        assert!(waiting_logits[b'<' as usize].is_finite());
1654        assert!(!waiting_logits[b'>' as usize].is_finite());
1655
1656        let delimiter_prefix = "</think"
1657            .bytes()
1658            .map(|byte| TokenId::new(byte as u32))
1659            .collect::<Vec<_>>();
1660        let mut partial_logits = vec![0.0; EOS as usize + 1];
1661        let partial = processor
1662            .mask_logits_with_terminals(
1663                &mut partial_logits,
1664                &delimiter_prefix,
1665                &HashSet::new(),
1666                &controls,
1667            )
1668            .unwrap();
1669        assert_eq!(partial.phase, StructuredOutputPhase::WaitingForDelimiter);
1670        assert_eq!(partial.grammar_start_token_index, None);
1671        assert_eq!(partial.required_delimiter_token_id, Some(b'>' as u32));
1672        assert!(!partial_logits[b'<' as usize].is_finite());
1673        assert!(partial_logits[b'>' as usize].is_finite());
1674        let partial_progress = processor
1675            .progress_with_terminals(&delimiter_prefix, &HashSet::new())
1676            .unwrap();
1677        assert_eq!(partial_progress.delimiter_token_count, Some(8));
1678        assert_eq!(partial_progress.delimiter_prefix_token_count, 7);
1679
1680        processor.reset().unwrap();
1681
1682        assert_and_append(&processor, &mut generated, "reasoning [is free]</think>");
1683        let mut logits = vec![0.0; EOS as usize + 1];
1684        let active = processor
1685            .mask_logits_with_terminals(&mut logits, &generated, &HashSet::new(), &HashSet::new())
1686            .unwrap();
1687        assert_eq!(active.phase, StructuredOutputPhase::EnforcingGrammar);
1688        assert_eq!(active.grammar_start_token_index, Some(27));
1689        assert!(logits[b'{' as usize].is_finite());
1690        assert!(!logits[b'[' as usize].is_finite());
1691        assert!(!logits[EOS as usize].is_finite());
1692
1693        assert_and_append(&processor, &mut generated, r#"{"ok":true}"#);
1694        assert!(processor.is_accepting(&generated).unwrap());
1695        let progress = processor
1696            .progress_with_terminals(&generated, &HashSet::new())
1697            .unwrap();
1698        assert_eq!(progress.phase, StructuredOutputPhase::EnforcingGrammar);
1699        assert!(progress.accepting);
1700        assert_eq!(progress.generated_token_count, generated.len());
1701        assert!(!progress.boundary_forced);
1702        assert_eq!(progress.reasoning_token_count, Some(19));
1703    }
1704
1705    #[test]
1706    fn reasoning_budget_forces_exact_delimiter_and_preserves_structured_reserve() {
1707        let processor = factory()
1708            .create_processor(
1709                &ResponseFormat::JsonObject,
1710                &StructuredOutputStart::AfterDelimiter("</think>".to_string()),
1711                48,
1712                &HashSet::new(),
1713                &[],
1714            )
1715            .unwrap()
1716            .unwrap();
1717        let mut generated = "reason!!"
1718            .bytes()
1719            .map(|byte| TokenId::new(byte as u32))
1720            .collect::<Vec<_>>();
1721
1722        for expected in "</think>".bytes() {
1723            let mut logits = vec![f32::NEG_INFINITY; EOS as usize + 1];
1724            let outcome = processor
1725                .mask_logits_with_terminals(
1726                    &mut logits,
1727                    &generated,
1728                    &HashSet::new(),
1729                    &HashSet::new(),
1730                )
1731                .unwrap();
1732            assert_eq!(outcome.phase, StructuredOutputPhase::ForcingDelimiter);
1733            assert_eq!(outcome.grammar_start_token_index, None);
1734            assert_eq!(outcome.required_delimiter_token_id, Some(expected as u32));
1735            assert_eq!(
1736                logits
1737                    .iter()
1738                    .enumerate()
1739                    .filter(|(_, logit)| logit.is_finite())
1740                    .map(|(token, _)| token)
1741                    .collect::<Vec<_>>(),
1742                vec![expected as usize]
1743            );
1744            generated.push(TokenId::new(expected as u32));
1745        }
1746
1747        let mut grammar_logits = vec![0.0; EOS as usize + 1];
1748        let outcome = processor
1749            .mask_logits_with_terminals(
1750                &mut grammar_logits,
1751                &generated,
1752                &HashSet::new(),
1753                &HashSet::new(),
1754            )
1755            .unwrap();
1756        assert_eq!(outcome.phase, StructuredOutputPhase::EnforcingGrammar);
1757        assert!(grammar_logits[b'{' as usize].is_finite());
1758        assert!(!grammar_logits[b'[' as usize].is_finite());
1759
1760        let progress = processor
1761            .progress_with_terminals(&generated, &HashSet::new())
1762            .unwrap();
1763        assert_eq!(progress.reasoning_token_count, Some(8));
1764        assert!(progress.boundary_forced);
1765        assert_eq!(
1766            progress.budget,
1767            Some(StructuredOutputBudgetPlan {
1768                total_output_tokens: 48,
1769                reasoning_token_limit: 8,
1770                boundary_token_count: 8,
1771                structured_reserve_tokens: 32,
1772            })
1773        );
1774
1775        processor.reset().unwrap();
1776        let reset_progress = processor
1777            .progress_with_terminals(&[], &HashSet::new())
1778            .unwrap();
1779        assert_eq!(
1780            reset_progress.phase,
1781            StructuredOutputPhase::WaitingForDelimiter
1782        );
1783        assert!(!reset_progress.boundary_forced);
1784        assert_eq!(reset_progress.reasoning_token_count, Some(0));
1785    }
1786
1787    #[test]
1788    fn reasoning_budget_reserves_half_of_a_normal_completion_for_structure() {
1789        assert_eq!(
1790            StructuredOutputBudgetPlan::automatic(1024, 1).unwrap(),
1791            StructuredOutputBudgetPlan {
1792                total_output_tokens: 1024,
1793                reasoning_token_limit: 511,
1794                boundary_token_count: 1,
1795                structured_reserve_tokens: 512,
1796            }
1797        );
1798    }
1799
1800    #[test]
1801    fn reasoning_delimiter_requires_room_beyond_the_boundary() {
1802        let error = factory()
1803            .create_processor(
1804                &ResponseFormat::JsonObject,
1805                &StructuredOutputStart::AfterDelimiter("</think>".to_string()),
1806                8,
1807                &HashSet::new(),
1808                &[],
1809            )
1810            .unwrap_err();
1811        assert!(error
1812            .to_string()
1813            .contains("max_tokens greater than its 8-token delimiter"));
1814    }
1815
1816    #[test]
1817    fn forcing_accounts_for_an_existing_delimiter_prefix() {
1818        let processor = factory()
1819            .create_processor(
1820                &ResponseFormat::JsonObject,
1821                &StructuredOutputStart::AfterDelimiter("</think>".to_string()),
1822                48,
1823                &HashSet::new(),
1824                &[],
1825            )
1826            .unwrap()
1827            .unwrap();
1828        let generated = "reason</"
1829            .bytes()
1830            .map(|byte| TokenId::new(byte as u32))
1831            .collect::<Vec<_>>();
1832        let mut logits = vec![0.0; EOS as usize + 1];
1833
1834        let outcome = processor
1835            .mask_logits_with_terminals(&mut logits, &generated, &HashSet::new(), &HashSet::new())
1836            .unwrap();
1837        assert_eq!(outcome.phase, StructuredOutputPhase::ForcingDelimiter);
1838        assert_eq!(outcome.required_delimiter_token_id, Some(b't' as u32));
1839        let progress = processor
1840            .progress_with_terminals(&generated, &HashSet::new())
1841            .unwrap();
1842        assert_eq!(progress.delimiter_prefix_token_count, 2);
1843        assert_eq!(progress.reasoning_token_count, Some(6));
1844    }
1845
1846    #[test]
1847    fn delimiter_rejects_any_conflicting_stop_condition_up_front() {
1848        let token_error = factory()
1849            .create_processor(
1850                &ResponseFormat::JsonObject,
1851                &StructuredOutputStart::AfterDelimiter("</think>".to_string()),
1852                TEST_MAX_OUTPUT_TOKENS,
1853                &HashSet::from([b'/' as u32]),
1854                &[],
1855            )
1856            .unwrap_err();
1857        assert!(token_error
1858            .to_string()
1859            .contains("conflicts with a stop token"));
1860
1861        let text_error = factory()
1862            .create_processor(
1863                &ResponseFormat::JsonObject,
1864                &StructuredOutputStart::AfterDelimiter("</think>".to_string()),
1865                TEST_MAX_OUTPUT_TOKENS,
1866                &HashSet::new(),
1867                &["think".to_string()],
1868            )
1869            .unwrap_err();
1870        assert!(text_error
1871            .to_string()
1872            .contains("conflicts with stop sequence"));
1873    }
1874}