1use 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;
24const 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#[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
85pub 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 pub fn new(tokenizer: Arc<dyn Tokenizer + Send + Sync>) -> Result<Self> {
106 Self::new_with_model_vocab_size(tokenizer, None)
107 }
108
109 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 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 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 JsonCompileOptions {
341 whitespace_flexible: false,
342 ..JsonCompileOptions::default()
343 }
344 .apply_to(&mut schema);
345 Ok(schema)
346}
347
348pub 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
360pub enum StructuredOutputPhase {
361 WaitingForDelimiter,
362 ForcingDelimiter,
363 EnforcingGrammar,
364}
365
366#[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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
386pub struct StructuredOutputMaskOutcome {
387 pub phase: StructuredOutputPhase,
388 pub accepting: bool,
389 pub liveness_intervention: bool,
390 pub grammar_start_token_index: Option<usize>,
395 pub required_delimiter_token_id: Option<u32>,
400}
401
402#[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 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 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 pub fn is_accepting(&self, generated: &[TokenId]) -> Result<bool> {
614 self.is_accepting_inner(generated, None)
615 }
616
617 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 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}