1use std::sync::Arc;
26
27use ferrum_interfaces::sampler::{LogitsProcessor, ProcessorPriority, SamplingContext};
28use ferrum_interfaces::tokenizer::Tokenizer;
29use ferrum_types::{FerrumError, Result, TokenId};
30use parking_lot::Mutex;
31use regex_automata::{
32 dfa::{dense::DFA, Automaton, StartKind},
33 util::{primitives::StateID, start::Config as StartConfig},
34 Anchored,
35};
36
37pub struct RegexGuidedProcessor {
43 dfa: DFA<Vec<u32>>,
44 state: Mutex<DfaPosition>,
47 token_bytes: Vec<Vec<u8>>,
49 eos_token: Option<TokenId>,
51 consumed: Mutex<usize>,
53}
54
55impl std::fmt::Debug for RegexGuidedProcessor {
56 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
57 f.debug_struct("RegexGuidedProcessor")
58 .field("vocab_size", &self.token_bytes.len())
59 .field("consumed", &*self.consumed.lock())
60 .finish()
61 }
62}
63
64#[derive(Copy, Clone, Debug)]
65struct DfaPosition {
66 state: StateID,
67 dead: bool,
71}
72
73impl RegexGuidedProcessor {
74 pub fn new(
77 pattern: &str,
78 tokenizer: Arc<dyn Tokenizer + Send + Sync>,
79 eos_token: Option<TokenId>,
80 ) -> Result<Self> {
81 let wrapped = format!(r"(?:{pattern})\z");
86 let dfa = DFA::builder()
87 .configure(DFA::config().start_kind(StartKind::Anchored))
88 .build(&wrapped)
89 .map_err(|e| FerrumError::invalid_request(format!("regex compile: {e}")))?;
90
91 let start = dfa
92 .start_state(&StartConfig::new().anchored(Anchored::Yes))
93 .map_err(|e| FerrumError::invalid_request(format!("regex start state: {e}")))?;
94
95 let vocab_size = tokenizer.vocab_size();
100 let mut token_bytes = Vec::with_capacity(vocab_size);
101 for i in 0..vocab_size {
102 let id = TokenId::new(i as u32);
103 let bytes = tokenizer
104 .decode(&[id], false)
105 .ok()
106 .or_else(|| tokenizer.token_text(id).map(str::to_string))
107 .map(String::into_bytes)
108 .unwrap_or_default();
109 token_bytes.push(bytes);
110 }
111
112 Ok(Self {
113 dfa,
114 state: Mutex::new(DfaPosition {
115 state: start,
116 dead: false,
117 }),
118 token_bytes,
119 eos_token,
120 consumed: Mutex::new(0),
121 })
122 }
123
124 pub fn reset(&self) -> Result<()> {
126 let start = self
127 .dfa
128 .start_state(&StartConfig::new().anchored(Anchored::Yes))
129 .map_err(|e| FerrumError::internal(format!("regex start state: {e}")))?;
130 *self.state.lock() = DfaPosition {
131 state: start,
132 dead: false,
133 };
134 *self.consumed.lock() = 0;
135 Ok(())
136 }
137
138 pub fn can_accept(&self) -> bool {
140 let pos = *self.state.lock();
141 !pos.dead && self.dfa.is_match_state(self.dfa.next_eoi_state(pos.state))
142 }
143
144 fn advance(&self, mut state: StateID, bytes: &[u8]) -> Option<StateID> {
148 for &b in bytes {
149 state = self.dfa.next_state(state, b);
150 if self.dfa.is_dead_state(state) {
151 return None;
152 }
153 }
154 Some(state)
155 }
156
157 pub fn mask_logits(&self, logits: &mut [f32]) {
159 let pos = *self.state.lock();
160 if pos.dead {
161 return;
166 }
167
168 let pattern_done = self.dfa.is_match_state(self.dfa.next_eoi_state(pos.state));
169
170 let vocab = logits.len().min(self.token_bytes.len());
174 let mut any_allowed = false;
175 let mut fallback_token: Option<(usize, f32)> = None;
176 for idx in 0..vocab {
177 let is_eos = self.eos_token.map_or(false, |e| e.get() as usize == idx);
178 let bytes = &self.token_bytes[idx];
179 let original_logit = logits[idx];
180 if !is_eos
181 && original_logit.is_finite()
182 && fallback_token
183 .map(|(_, best)| original_logit > best)
184 .unwrap_or(true)
185 {
186 fallback_token = Some((idx, original_logit));
187 }
188
189 let allowed = if is_eos {
190 pattern_done
191 } else if bytes.is_empty() {
192 pattern_done
196 } else {
197 self.advance(pos.state, bytes).is_some()
198 };
199
200 if allowed {
201 any_allowed = true;
202 } else {
203 logits[idx] = f32::NEG_INFINITY;
204 }
205 }
206
207 if !any_allowed {
214 if pattern_done {
215 self.force_eos(logits);
216 } else if let Some((token, _)) = fallback_token {
217 self.force_only_token(logits, token);
218 }
219 }
220 }
221
222 fn force_eos(&self, logits: &mut [f32]) {
223 if let Some(eos) = self.eos_token {
224 let eos_idx = eos.get() as usize;
225 for (i, l) in logits.iter_mut().enumerate() {
226 *l = if i == eos_idx { 0.0 } else { f32::NEG_INFINITY };
227 }
228 }
229 }
230
231 fn force_only_token(&self, logits: &mut [f32], token: usize) {
232 for (i, l) in logits.iter_mut().enumerate() {
233 *l = if i == token { 0.0 } else { f32::NEG_INFINITY };
234 }
235 }
236
237 pub fn advance_with_tokens_public(&self, tokens: &[TokenId]) {
241 self.advance_with_tokens(tokens);
242 }
243
244 fn advance_with_tokens(&self, tokens: &[TokenId]) {
248 let mut consumed = self.consumed.lock();
249 if *consumed >= tokens.len() {
250 return;
251 }
252 let mut pos = self.state.lock();
253 for &tok in &tokens[*consumed..] {
254 if pos.dead {
255 break;
256 }
257 let idx = tok.get() as usize;
258 if idx >= self.token_bytes.len() {
259 continue;
260 }
261 let bytes = &self.token_bytes[idx];
262 if self.eos_token.map_or(false, |e| e == tok) {
264 continue;
265 }
266 if let Some(next) = self.advance(pos.state, bytes) {
267 pos.state = next;
268 } else {
269 pos.dead = true;
270 }
271 }
272 *consumed = tokens.len();
273 }
274}
275
276impl LogitsProcessor for RegexGuidedProcessor {
277 fn process(&self, ctx: &mut SamplingContext) -> Result<()> {
278 self.advance_with_tokens(ctx.previous_tokens);
279 self.mask_logits(ctx.logits);
280 Ok(())
281 }
282
283 fn name(&self) -> &str {
284 "regex_guided"
285 }
286
287 fn priority(&self) -> ProcessorPriority {
288 ProcessorPriority::High
291 }
292}
293
294#[cfg(test)]
295mod tests {
296 use super::*;
297 use ferrum_interfaces::tokenizer::{ChatMessage, TokenizerInfo, TokenizerType};
298 use ferrum_types::{SpecialTokens, TokenId};
299
300 struct ByteTokenizer {
304 special: SpecialTokens,
305 byte_strings: Vec<String>,
306 }
307
308 impl ByteTokenizer {
309 fn new() -> Self {
310 let mut byte_strings = Vec::with_capacity(257);
311 for b in 0u8..=255 {
312 byte_strings.push(String::from_utf8(vec![b]).unwrap_or_default());
313 }
314 byte_strings.push("</s>".to_string());
315 Self {
316 special: SpecialTokens {
317 bos_token: None,
318 eos_token: Some(TokenId::new(256)),
319 unk_token: None,
320 pad_token: None,
321 sep_token: None,
322 cls_token: None,
323 mask_token: None,
324 extra_eos_tokens: Vec::new(),
325 },
326 byte_strings,
327 }
328 }
329 }
330
331 impl Tokenizer for ByteTokenizer {
332 fn encode(&self, text: &str, _add_special: bool) -> Result<Vec<TokenId>> {
333 Ok(text.bytes().map(|b| TokenId::new(b as u32)).collect())
334 }
335 fn decode(&self, tokens: &[TokenId], _skip_special: bool) -> Result<String> {
336 let mut out = String::new();
337 for t in tokens {
338 let idx = t.get() as usize;
339 if idx < 256 {
340 out.push(idx as u8 as char);
341 }
342 }
343 Ok(out)
344 }
345 fn decode_incremental(&self, _prev: &[TokenId], next: TokenId) -> Result<String> {
346 self.decode(&[next], false)
347 }
348 fn vocab_size(&self) -> usize {
349 257
350 }
351 fn special_tokens(&self) -> &SpecialTokens {
352 &self.special
353 }
354 fn token_id(&self, text: &str) -> Option<TokenId> {
355 if text.len() == 1 {
356 Some(TokenId::new(text.bytes().next().unwrap() as u32))
357 } else {
358 None
359 }
360 }
361 fn token_text(&self, token_id: TokenId) -> Option<&str> {
362 self.byte_strings
363 .get(token_id.get() as usize)
364 .map(|s| s.as_str())
365 }
366 fn apply_chat_template(&self, _messages: &[ChatMessage]) -> Result<String> {
367 Ok(String::new())
368 }
369 fn info(&self) -> TokenizerInfo {
370 TokenizerInfo {
371 tokenizer_type: TokenizerType::Custom,
372 vocab_size: 257,
373 special_tokens: self.special.clone(),
374 supports_incremental: true,
375 supports_chat_template: false,
376 max_token_length: Some(1),
377 model_name: Some("byte-tokenizer-test".into()),
378 }
379 }
380 }
381
382 fn processor(pattern: &str) -> RegexGuidedProcessor {
383 let tok: Arc<dyn Tokenizer> = Arc::new(ByteTokenizer::new());
384 RegexGuidedProcessor::new(pattern, tok, Some(TokenId::new(256))).unwrap()
385 }
386
387 struct TinyTokenizer {
388 special: SpecialTokens,
389 strings: Vec<String>,
390 }
391
392 impl TinyTokenizer {
393 fn new(strings: Vec<&str>, eos: u32) -> Self {
394 Self {
395 special: SpecialTokens {
396 bos_token: None,
397 eos_token: Some(TokenId::new(eos)),
398 unk_token: None,
399 pad_token: None,
400 sep_token: None,
401 cls_token: None,
402 mask_token: None,
403 extra_eos_tokens: Vec::new(),
404 },
405 strings: strings.into_iter().map(str::to_string).collect(),
406 }
407 }
408 }
409
410 impl Tokenizer for TinyTokenizer {
411 fn encode(&self, text: &str, _add_special: bool) -> Result<Vec<TokenId>> {
412 Ok(text
413 .chars()
414 .filter_map(|ch| self.token_id(&ch.to_string()))
415 .collect())
416 }
417
418 fn decode(&self, tokens: &[TokenId], _skip_special: bool) -> Result<String> {
419 Ok(tokens
420 .iter()
421 .filter_map(|token| self.token_text(*token))
422 .collect::<Vec<_>>()
423 .join(""))
424 }
425
426 fn decode_incremental(&self, _prev: &[TokenId], next: TokenId) -> Result<String> {
427 self.decode(&[next], false)
428 }
429
430 fn vocab_size(&self) -> usize {
431 self.strings.len()
432 }
433
434 fn special_tokens(&self) -> &SpecialTokens {
435 &self.special
436 }
437
438 fn token_id(&self, text: &str) -> Option<TokenId> {
439 self.strings
440 .iter()
441 .position(|value| value == text)
442 .map(|idx| TokenId::new(idx as u32))
443 }
444
445 fn token_text(&self, token_id: TokenId) -> Option<&str> {
446 self.strings
447 .get(token_id.get() as usize)
448 .map(|value| value.as_str())
449 }
450
451 fn apply_chat_template(&self, _messages: &[ChatMessage]) -> Result<String> {
452 Ok(String::new())
453 }
454
455 fn info(&self) -> TokenizerInfo {
456 TokenizerInfo {
457 tokenizer_type: TokenizerType::Custom,
458 vocab_size: self.strings.len(),
459 special_tokens: self.special.clone(),
460 supports_incremental: true,
461 supports_chat_template: false,
462 max_token_length: Some(1),
463 model_name: Some("tiny-tokenizer-test".into()),
464 }
465 }
466 }
467
468 #[test]
469 fn digits_only_allows_digits_at_start() {
470 let p = processor(r"[0-9]+");
471 let mut logits = vec![0.0f32; 257];
472 p.mask_logits(&mut logits);
473 for b in 0u8..=255 {
474 let expected_allowed = b.is_ascii_digit();
475 let got = logits[b as usize].is_finite();
476 assert_eq!(
477 got, expected_allowed,
478 "byte {b:?} ({}): expected allowed={expected_allowed}, got={got}",
479 b as char
480 );
481 }
482 assert!(logits[256].is_infinite() && logits[256].is_sign_negative());
484 }
485
486 #[test]
487 fn digits_only_allows_eos_after_a_digit() {
488 let p = processor(r"[0-9]+");
489 p.advance_with_tokens(&[TokenId::new(b'3' as u32)]);
490 let mut logits = vec![0.0f32; 257];
491 p.mask_logits(&mut logits);
492 assert!(
493 logits[256].is_finite(),
494 "EOS should be allowed after a digit"
495 );
496 assert!(
497 logits[b'7' as usize].is_finite(),
498 "another digit still allowed"
499 );
500 assert!(logits[b'a' as usize].is_infinite(), "alpha still forbidden");
501 }
502
503 #[test]
504 fn dead_state_disables_masking_instead_of_forcing_eos() {
505 let p = processor(r"[0-9]+");
506 p.advance_with_tokens(&[TokenId::new(b'a' as u32)]);
508 let mut logits = vec![0.0f32; 257];
509 logits[256] = 7.0;
510 p.mask_logits(&mut logits);
511 assert_eq!(logits[256], 7.0, "EOS should not be forced");
512 assert_eq!(logits[b'a' as usize], 0.0, "mask should be disabled");
513 }
514
515 #[test]
516 fn no_extension_before_accept_uses_best_non_eos_token() {
517 let tok: Arc<dyn Tokenizer> = Arc::new(TinyTokenizer::new(vec!["a", "b", "</s>"], 2));
518 let p = RegexGuidedProcessor::new("z", tok, Some(TokenId::new(2))).unwrap();
519 let mut logits = vec![0.1, 3.0, 99.0];
520 p.mask_logits(&mut logits);
521
522 assert!(
523 logits[0].is_infinite(),
524 "lower non-EOS token should be masked"
525 );
526 assert!(logits[1].is_finite(), "best non-EOS token should remain");
527 assert!(
528 logits[2].is_infinite(),
529 "EOS must stay masked before the regex accepts"
530 );
531 }
532
533 #[test]
534 fn reset_restores_initial_state() {
535 let p = processor(r"[0-9]+");
536 p.advance_with_tokens(&[TokenId::new(b'3' as u32)]);
537 assert!(p.can_accept());
538 p.reset().unwrap();
539 assert!(!p.can_accept(), "fresh state should not accept empty input");
540 }
541
542 #[test]
543 fn hex_prefix_pattern() {
544 let p = processor(r"0x[0-9a-f]+");
545 let mut logits = vec![0.0f32; 257];
546 p.mask_logits(&mut logits);
547 assert!(logits[b'0' as usize].is_finite(), "'0' starts the pattern");
548 assert!(logits[b'1' as usize].is_infinite(), "'1' can't start");
549 p.advance_with_tokens(&[TokenId::new(b'0' as u32), TokenId::new(b'x' as u32)]);
550 let mut logits = vec![0.0f32; 257];
551 p.mask_logits(&mut logits);
552 assert!(logits[b'a' as usize].is_finite());
553 assert!(logits[b'f' as usize].is_finite());
554 assert!(logits[b'g' as usize].is_infinite());
555 assert!(logits[b'9' as usize].is_finite());
556 }
557}