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 },
325 byte_strings,
326 }
327 }
328 }
329
330 impl Tokenizer for ByteTokenizer {
331 fn encode(&self, text: &str, _add_special: bool) -> Result<Vec<TokenId>> {
332 Ok(text.bytes().map(|b| TokenId::new(b as u32)).collect())
333 }
334 fn decode(&self, tokens: &[TokenId], _skip_special: bool) -> Result<String> {
335 let mut out = String::new();
336 for t in tokens {
337 let idx = t.get() as usize;
338 if idx < 256 {
339 out.push(idx as u8 as char);
340 }
341 }
342 Ok(out)
343 }
344 fn decode_incremental(&self, _prev: &[TokenId], next: TokenId) -> Result<String> {
345 self.decode(&[next], false)
346 }
347 fn vocab_size(&self) -> usize {
348 257
349 }
350 fn special_tokens(&self) -> &SpecialTokens {
351 &self.special
352 }
353 fn token_id(&self, text: &str) -> Option<TokenId> {
354 if text.len() == 1 {
355 Some(TokenId::new(text.bytes().next().unwrap() as u32))
356 } else {
357 None
358 }
359 }
360 fn token_text(&self, token_id: TokenId) -> Option<&str> {
361 self.byte_strings
362 .get(token_id.get() as usize)
363 .map(|s| s.as_str())
364 }
365 fn apply_chat_template(&self, _messages: &[ChatMessage]) -> Result<String> {
366 Ok(String::new())
367 }
368 fn info(&self) -> TokenizerInfo {
369 TokenizerInfo {
370 tokenizer_type: TokenizerType::Custom,
371 vocab_size: 257,
372 special_tokens: self.special.clone(),
373 supports_incremental: true,
374 supports_chat_template: false,
375 max_token_length: Some(1),
376 model_name: Some("byte-tokenizer-test".into()),
377 }
378 }
379 }
380
381 fn processor(pattern: &str) -> RegexGuidedProcessor {
382 let tok: Arc<dyn Tokenizer> = Arc::new(ByteTokenizer::new());
383 RegexGuidedProcessor::new(pattern, tok, Some(TokenId::new(256))).unwrap()
384 }
385
386 struct TinyTokenizer {
387 special: SpecialTokens,
388 strings: Vec<String>,
389 }
390
391 impl TinyTokenizer {
392 fn new(strings: Vec<&str>, eos: u32) -> Self {
393 Self {
394 special: SpecialTokens {
395 bos_token: None,
396 eos_token: Some(TokenId::new(eos)),
397 unk_token: None,
398 pad_token: None,
399 sep_token: None,
400 cls_token: None,
401 mask_token: None,
402 },
403 strings: strings.into_iter().map(str::to_string).collect(),
404 }
405 }
406 }
407
408 impl Tokenizer for TinyTokenizer {
409 fn encode(&self, text: &str, _add_special: bool) -> Result<Vec<TokenId>> {
410 Ok(text
411 .chars()
412 .filter_map(|ch| self.token_id(&ch.to_string()))
413 .collect())
414 }
415
416 fn decode(&self, tokens: &[TokenId], _skip_special: bool) -> Result<String> {
417 Ok(tokens
418 .iter()
419 .filter_map(|token| self.token_text(*token))
420 .collect::<Vec<_>>()
421 .join(""))
422 }
423
424 fn decode_incremental(&self, _prev: &[TokenId], next: TokenId) -> Result<String> {
425 self.decode(&[next], false)
426 }
427
428 fn vocab_size(&self) -> usize {
429 self.strings.len()
430 }
431
432 fn special_tokens(&self) -> &SpecialTokens {
433 &self.special
434 }
435
436 fn token_id(&self, text: &str) -> Option<TokenId> {
437 self.strings
438 .iter()
439 .position(|value| value == text)
440 .map(|idx| TokenId::new(idx as u32))
441 }
442
443 fn token_text(&self, token_id: TokenId) -> Option<&str> {
444 self.strings
445 .get(token_id.get() as usize)
446 .map(|value| value.as_str())
447 }
448
449 fn apply_chat_template(&self, _messages: &[ChatMessage]) -> Result<String> {
450 Ok(String::new())
451 }
452
453 fn info(&self) -> TokenizerInfo {
454 TokenizerInfo {
455 tokenizer_type: TokenizerType::Custom,
456 vocab_size: self.strings.len(),
457 special_tokens: self.special.clone(),
458 supports_incremental: true,
459 supports_chat_template: false,
460 max_token_length: Some(1),
461 model_name: Some("tiny-tokenizer-test".into()),
462 }
463 }
464 }
465
466 #[test]
467 fn digits_only_allows_digits_at_start() {
468 let p = processor(r"[0-9]+");
469 let mut logits = vec![0.0f32; 257];
470 p.mask_logits(&mut logits);
471 for b in 0u8..=255 {
472 let expected_allowed = b.is_ascii_digit();
473 let got = logits[b as usize].is_finite();
474 assert_eq!(
475 got, expected_allowed,
476 "byte {b:?} ({}): expected allowed={expected_allowed}, got={got}",
477 b as char
478 );
479 }
480 assert!(logits[256].is_infinite() && logits[256].is_sign_negative());
482 }
483
484 #[test]
485 fn digits_only_allows_eos_after_a_digit() {
486 let p = processor(r"[0-9]+");
487 p.advance_with_tokens(&[TokenId::new(b'3' as u32)]);
488 let mut logits = vec![0.0f32; 257];
489 p.mask_logits(&mut logits);
490 assert!(
491 logits[256].is_finite(),
492 "EOS should be allowed after a digit"
493 );
494 assert!(
495 logits[b'7' as usize].is_finite(),
496 "another digit still allowed"
497 );
498 assert!(logits[b'a' as usize].is_infinite(), "alpha still forbidden");
499 }
500
501 #[test]
502 fn dead_state_disables_masking_instead_of_forcing_eos() {
503 let p = processor(r"[0-9]+");
504 p.advance_with_tokens(&[TokenId::new(b'a' as u32)]);
506 let mut logits = vec![0.0f32; 257];
507 logits[256] = 7.0;
508 p.mask_logits(&mut logits);
509 assert_eq!(logits[256], 7.0, "EOS should not be forced");
510 assert_eq!(logits[b'a' as usize], 0.0, "mask should be disabled");
511 }
512
513 #[test]
514 fn no_extension_before_accept_uses_best_non_eos_token() {
515 let tok: Arc<dyn Tokenizer> = Arc::new(TinyTokenizer::new(vec!["a", "b", "</s>"], 2));
516 let p = RegexGuidedProcessor::new("z", tok, Some(TokenId::new(2))).unwrap();
517 let mut logits = vec![0.1, 3.0, 99.0];
518 p.mask_logits(&mut logits);
519
520 assert!(
521 logits[0].is_infinite(),
522 "lower non-EOS token should be masked"
523 );
524 assert!(logits[1].is_finite(), "best non-EOS token should remain");
525 assert!(
526 logits[2].is_infinite(),
527 "EOS must stay masked before the regex accepts"
528 );
529 }
530
531 #[test]
532 fn reset_restores_initial_state() {
533 let p = processor(r"[0-9]+");
534 p.advance_with_tokens(&[TokenId::new(b'3' as u32)]);
535 assert!(p.can_accept());
536 p.reset().unwrap();
537 assert!(!p.can_accept(), "fresh state should not accept empty input");
538 }
539
540 #[test]
541 fn hex_prefix_pattern() {
542 let p = processor(r"0x[0-9a-f]+");
543 let mut logits = vec![0.0f32; 257];
544 p.mask_logits(&mut logits);
545 assert!(logits[b'0' as usize].is_finite(), "'0' starts the pattern");
546 assert!(logits[b'1' as usize].is_infinite(), "'1' can't start");
547 p.advance_with_tokens(&[TokenId::new(b'0' as u32), TokenId::new(b'x' as u32)]);
548 let mut logits = vec![0.0f32; 257];
549 p.mask_logits(&mut logits);
550 assert!(logits[b'a' as usize].is_finite());
551 assert!(logits[b'f' as usize].is_finite());
552 assert!(logits[b'g' as usize].is_infinite());
553 assert!(logits[b'9' as usize].is_finite());
554 }
555}