use std::collections::VecDeque;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ReasoningBudgetPlan {
pub start_seqs: Vec<Vec<usize>>,
pub end_seqs: Vec<Vec<usize>>,
pub forced: Vec<usize>,
pub budget: u32,
pub prefill: Vec<usize>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BudgetState {
Idle,
Counting,
WaitingUtf8,
Forcing,
Done,
}
#[derive(Debug, Clone)]
struct SeqMatcher {
seqs: Vec<Vec<usize>>,
window: VecDeque<usize>,
longest: usize,
}
impl SeqMatcher {
fn new(seqs: &[Vec<usize>]) -> Self {
let mut kept: Vec<Vec<usize>> = Vec::new();
for seq in seqs {
if !seq.is_empty() && !kept.contains(seq) {
kept.push(seq.clone());
}
}
let longest = kept.iter().map(Vec::len).max().unwrap_or(0);
Self {
seqs: kept,
window: VecDeque::with_capacity(longest),
longest,
}
}
fn advance(&mut self, token: usize) -> Option<usize> {
if self.longest == 0 {
return None;
}
if self.window.len() == self.longest {
self.window.pop_front();
}
self.window.push_back(token);
let hit = self
.seqs
.iter()
.enumerate()
.filter(|(_, seq)| {
seq.len() <= self.window.len()
&& self
.window
.iter()
.skip(self.window.len() - seq.len())
.eq(seq.iter())
})
.max_by_key(|(_, seq)| seq.len())
.map(|(i, _)| i);
if hit.is_some() {
self.reset();
}
hit
}
fn reset(&mut self) {
self.window.clear();
}
}
#[derive(Debug, Clone)]
pub struct ReasoningBudget {
start: SeqMatcher,
end: SeqMatcher,
forced: Vec<usize>,
budget: u32,
remaining: u32,
state: BudgetState,
force_pos: usize,
}
impl ReasoningBudget {
pub fn new(plan: &ReasoningBudgetPlan) -> Self {
let mut machine = Self {
start: SeqMatcher::new(&plan.start_seqs),
end: SeqMatcher::new(&plan.end_seqs),
forced: plan.forced.clone(),
budget: plan.budget,
remaining: plan.budget,
state: BudgetState::Idle,
force_pos: 0,
};
for &token in &plan.prefill {
machine.accept(token, true);
}
machine
}
pub fn state(&self) -> BudgetState {
self.state
}
pub fn forced_token(&self) -> Option<usize> {
match self.state {
BudgetState::Forcing => self.forced.get(self.force_pos).copied(),
_ => None,
}
}
pub fn mask(&self, logits: &mut [f32]) {
let Some(forced) = self.forced_token() else {
return;
};
for (id, logit) in logits.iter_mut().enumerate() {
if id != forced {
*logit = f32::NEG_INFINITY;
}
}
}
fn arm(&mut self) {
self.state = BudgetState::Counting;
self.remaining = self.budget;
self.end.reset();
if self.remaining == 0 {
self.state = BudgetState::Forcing;
self.force_pos = 0;
}
}
fn force(&mut self) {
self.state = BudgetState::Forcing;
self.force_pos = 0;
self.end.reset();
}
pub fn accept(&mut self, token: usize, piece_complete: bool) {
match self.state {
BudgetState::Idle | BudgetState::Done => {
if self.start.advance(token).is_some() {
self.arm();
}
}
BudgetState::Counting | BudgetState::WaitingUtf8 => {
if self.end.advance(token).is_some() {
self.state = BudgetState::Done;
return;
}
if self.state == BudgetState::WaitingUtf8 {
if piece_complete {
self.force();
}
return;
}
self.remaining = self.remaining.saturating_sub(1);
if self.remaining == 0 {
if piece_complete {
self.force();
} else {
self.state = BudgetState::WaitingUtf8;
self.end.reset();
}
}
}
BudgetState::Forcing => {
self.end.advance(token);
if self.forced.get(self.force_pos) == Some(&token) {
self.force_pos += 1;
} else {
self.force_pos = 0;
}
if self.force_pos >= self.forced.len() {
self.state = BudgetState::Done;
}
}
}
}
}
pub fn lossy_piece_is_complete(piece: &str) -> bool {
!piece.ends_with('\u{FFFD}')
}
#[cfg(test)]
mod tests {
use super::*;
const OPEN: usize = 100;
const CLOSE: usize = 101;
const TOOL_A: usize = 102;
const TOOL_B: usize = 103;
fn plan(budget: u32, prefill: &[usize]) -> ReasoningBudgetPlan {
ReasoningBudgetPlan {
start_seqs: vec![vec![OPEN]],
end_seqs: vec![vec![CLOSE], vec![TOOL_A, TOOL_B]],
forced: vec![CLOSE],
budget,
prefill: prefill.to_vec(),
}
}
fn logits(n: usize) -> Vec<f32> {
(0..n).map(|i| i as f32).collect()
}
fn trace(m: &mut ReasoningBudget, tokens: &[usize]) -> Vec<BudgetState> {
tokens
.iter()
.map(|&t| {
m.accept(t, true);
m.state()
})
.collect()
}
#[test]
fn n_tokens_of_thought_pass_and_the_next_draw_is_the_closer() {
let mut m = ReasoningBudget::new(&plan(3, &[]));
assert_eq!(m.state(), BudgetState::Idle);
assert_eq!(m.forced_token(), None);
assert_eq!(
trace(&mut m, &[OPEN, 1, 2]),
[
BudgetState::Counting,
BudgetState::Counting,
BudgetState::Counting
]
);
assert_eq!(m.forced_token(), None, "two of three spent: passthrough");
m.accept(3, true);
assert_eq!(m.state(), BudgetState::Forcing);
assert_eq!(m.forced_token(), Some(CLOSE));
let mut l = logits(200);
m.mask(&mut l);
assert!(l[CLOSE].is_finite());
assert!(
l.iter()
.enumerate()
.all(|(i, v)| i == CLOSE || *v == f32::NEG_INFINITY),
"every logit but the closer's is -inf while forcing"
);
m.accept(CLOSE, true);
assert_eq!(m.state(), BudgetState::Done);
assert_eq!(m.forced_token(), None, "the answer is unconstrained");
let mut l = logits(200);
m.mask(&mut l);
assert_eq!(l, logits(200), "mask is a no-op once done");
}
#[test]
fn budget_zero_forces_the_closer_right_after_the_opener() {
let mut m = ReasoningBudget::new(&plan(0, &[]));
m.accept(OPEN, true);
assert_eq!(m.state(), BudgetState::Forcing);
assert_eq!(m.forced_token(), Some(CLOSE));
}
#[test]
fn a_natural_close_ends_counting_without_forcing() {
let mut m = ReasoningBudget::new(&plan(10, &[]));
trace(&mut m, &[OPEN, 1, 2, CLOSE]);
assert_eq!(m.state(), BudgetState::Done);
assert_eq!(m.forced_token(), None);
}
#[test]
fn a_multi_token_end_sequence_closes_only_when_complete() {
let mut m = ReasoningBudget::new(&plan(10, &[]));
trace(&mut m, &[OPEN, TOOL_A]);
assert_eq!(m.state(), BudgetState::Counting);
m.accept(TOOL_B, true);
assert_eq!(m.state(), BudgetState::Done);
}
#[test]
fn a_prompt_that_opened_the_block_starts_counting_before_the_first_draw() {
const NEWLINE: usize = 7;
let mut m = ReasoningBudget::new(&plan(2, &[OPEN, NEWLINE]));
assert_eq!(m.state(), BudgetState::Counting);
m.accept(1, true);
assert_eq!(m.state(), BudgetState::Forcing);
}
#[test]
fn a_prefill_token_after_the_opener_at_budget_zero_does_not_stand_in_for_the_closer() {
const NEWLINE: usize = 7;
let m = ReasoningBudget::new(&plan(0, &[OPEN, NEWLINE]));
assert_eq!(m.state(), BudgetState::Forcing);
assert_eq!(m.forced_token(), Some(CLOSE));
}
#[test]
fn a_multi_token_forced_sequence_is_forced_in_order() {
let mut p = plan(0, &[]);
p.forced = vec![50, 51, CLOSE];
let mut m = ReasoningBudget::new(&p);
m.accept(OPEN, true);
for &expect in &[50, 51, CLOSE] {
assert_eq!(m.forced_token(), Some(expect));
m.accept(expect, true);
}
assert_eq!(m.state(), BudgetState::Done);
}
#[test]
fn a_split_character_delays_forcing_until_it_is_whole() {
let mut m = ReasoningBudget::new(&plan(1, &[]));
m.accept(OPEN, true);
m.accept(1, false);
assert_eq!(m.state(), BudgetState::WaitingUtf8);
assert_eq!(
m.forced_token(),
None,
"the closer must not split a character"
);
m.accept(2, false);
assert_eq!(m.state(), BudgetState::WaitingUtf8);
m.accept(3, true);
assert_eq!(m.state(), BudgetState::Forcing);
}
#[test]
fn done_re_arms_on_a_new_opener() {
let mut m = ReasoningBudget::new(&plan(1, &[]));
trace(&mut m, &[OPEN, CLOSE, 5, 6, OPEN]);
assert_eq!(m.state(), BudgetState::Counting);
m.accept(9, true);
assert_eq!(m.state(), BudgetState::Forcing);
}
#[test]
fn a_multi_token_opener_must_arrive_whole() {
let mut p = plan(5, &[]);
p.start_seqs = vec![vec![OPEN, 1, 2]];
let mut m = ReasoningBudget::new(&p);
trace(&mut m, &[OPEN, 1, 9, 1, 2]);
assert_eq!(m.state(), BudgetState::Idle, "OPEN 1 9 broke the sequence");
trace(&mut m, &[OPEN, 1, 2]);
assert_eq!(m.state(), BudgetState::Counting);
}
#[test]
fn empty_sequences_never_match() {
let mut p = plan(5, &[]);
p.start_seqs = vec![vec![]];
let mut m = ReasoningBudget::new(&p);
trace(&mut m, &[1, 2, 3]);
assert_eq!(m.state(), BudgetState::Idle);
}
#[test]
fn a_split_multibyte_piece_reads_incomplete_until_its_last_byte() {
let euro = "€".as_bytes(); let first = String::from_utf8_lossy(&euro[..1]);
let middle = String::from_utf8_lossy(&euro[1..2]);
let whole = String::from_utf8_lossy(euro);
assert!(!lossy_piece_is_complete(&first));
assert!(!lossy_piece_is_complete(&middle));
assert!(lossy_piece_is_complete(&whole));
assert!(lossy_piece_is_complete("plain ascii"));
assert!(lossy_piece_is_complete(""));
}
}