1use std::collections::VecDeque;
2
3use llama_cpp_bindings_sys::llama_pos;
4use llama_cpp_bindings_sys::llama_seq_id;
5
6use llama_cpp_bindings_types::TokenUsage;
7use llama_cpp_bindings_types::TokenUsageError;
8
9use crate::batch_add_error::BatchAddError;
10use crate::context::LlamaContext;
11use crate::error::EvalMultimodalChunksError;
12use crate::error::SampleError;
13use crate::error::TokenToStringError;
14use crate::eval_multimodal_chunks_params::EvalMultimodalChunksParams;
15use crate::llama_batch::LlamaBatch;
16use crate::marker_kind::MarkerKind;
17use crate::model::LlamaModel;
18use crate::mtmd::MtmdContext;
19use crate::mtmd::MtmdInputChunks;
20use crate::sampled_token::SampledToken;
21use crate::sampling::LlamaSampler;
22use crate::streaming_json_probe::JsonProbeOutcome;
23use crate::streaming_markers::StreamingMarkers;
24use crate::token::LlamaToken;
25
26pub use crate::ingest_outcome::IngestOutcome;
27pub use crate::sampled_token_section::SampledTokenSection;
28
29#[derive(Clone, Debug)]
30struct PendingToken {
31 token: LlamaToken,
32 decoded: String,
33 section: SampledTokenSection,
34 is_boundary: bool,
35 is_from_prompt: bool,
36 is_held_for_probe: bool,
37}
38
39#[derive(Clone, Debug, Eq, PartialEq)]
40struct JsonProbeState {
41 held_text: String,
42}
43
44#[derive(Clone, Debug, Eq, PartialEq)]
45enum ProbeMode {
46 Idle,
47 Active(JsonProbeState),
48}
49
50pub struct SampledTokenClassifier<'model> {
51 model: &'model LlamaModel,
52 markers: StreamingMarkers,
53 decoder: encoding_rs::Decoder,
54 pending: VecDeque<PendingToken>,
55 section: SampledTokenSection,
56 pending_prompt_tokens: u64,
57 usage: TokenUsage,
58 probe_mode: ProbeMode,
59}
60
61impl<'model> SampledTokenClassifier<'model> {
62 #[must_use]
63 pub fn new(model: &'model LlamaModel, markers: StreamingMarkers) -> Self {
64 Self {
65 model,
66 markers,
67 decoder: encoding_rs::UTF_8.new_decoder(),
68 pending: VecDeque::new(),
69 section: SampledTokenSection::Pending,
70 pending_prompt_tokens: 0,
71 usage: TokenUsage::new(),
72 probe_mode: ProbeMode::Idle,
73 }
74 }
75
76 pub fn ingest(&mut self, token: LlamaToken) -> Result<Vec<IngestOutcome>, TokenToStringError> {
81 if !self.markers.has_any() {
82 self.usage.record_undeterminable_token();
83 let piece = self.decode(token)?;
84 return Ok(vec![IngestOutcome {
85 sampled_token: SampledToken::Undeterminable(token),
86 visible_piece: piece.clone(),
87 raw_piece: piece,
88 }]);
89 }
90
91 let decoded = self.decode(token)?;
92 self.pending.push_back(PendingToken {
93 token,
94 decoded: decoded.clone(),
95 section: self.section,
96 is_boundary: false,
97 is_from_prompt: false,
98 is_held_for_probe: false,
99 });
100
101 self.try_consume_marker_at_tail();
102
103 let mut outcomes = self.classify_pending_tail(&decoded);
104
105 outcomes.extend(self.drain_overflow());
106 Ok(outcomes)
107 }
108
109 fn classify_pending_tail(&mut self, decoded: &str) -> Vec<IngestOutcome> {
110 let probe_was_active = matches!(self.probe_mode, ProbeMode::Active(_));
111 if probe_was_active && self.section_disengages_probe() {
112 self.abandon_probe()
113 } else {
114 self.update_probe(decoded)
115 }
116 }
117
118 const fn section_disengages_probe(&self) -> bool {
119 matches!(
120 self.section,
121 SampledTokenSection::ToolCall | SampledTokenSection::Reasoning
122 )
123 }
124
125 pub fn ingest_prompt_token(&mut self, token: LlamaToken) {
126 if !self.markers.has_any() {
127 return;
128 }
129
130 self.pending.push_back(PendingToken {
131 token,
132 decoded: String::new(),
133 section: self.section,
134 is_boundary: false,
135 is_from_prompt: true,
136 is_held_for_probe: false,
137 });
138
139 self.try_consume_marker_at_tail();
140 self.drain_overflow();
141 }
142
143 pub fn ingest_prompt_tokens(&mut self, tokens: &[LlamaToken]) {
144 if !self.markers.has_any() {
145 return;
146 }
147 for &token in tokens {
148 self.ingest_prompt_token(token);
149 }
150 }
151
152 pub fn flush(&mut self) -> Vec<IngestOutcome> {
153 self.probe_mode = ProbeMode::Idle;
154 let mut outcomes = Vec::with_capacity(self.pending.len());
155 while let Some(entry) = self.pending.pop_front() {
156 if entry.is_from_prompt {
157 continue;
158 }
159 outcomes.push(self.finalize_entry(entry));
160 }
161 outcomes
162 }
163
164 fn decode(&mut self, token: LlamaToken) -> Result<String, TokenToStringError> {
165 self.model
166 .token_to_piece(&SampledToken::Content(token), &mut self.decoder, true, None)
167 }
168
169 fn try_consume_marker_at_tail(&mut self) {
170 const PROBE_KINDS: &[MarkerKind] = &[
171 MarkerKind::ReasoningOpen,
172 MarkerKind::ReasoningClose,
173 MarkerKind::ToolCallOpen,
174 MarkerKind::ToolCallClose,
175 ];
176
177 for &kind in PROBE_KINDS {
178 let Some(marker) = self.markers.lookup(kind) else {
179 continue;
180 };
181 if marker.is_empty() || self.pending.len() < marker.len() {
182 continue;
183 }
184 let span_start = self.pending.len() - marker.len();
185 let matches = self
186 .pending
187 .iter()
188 .skip(span_start)
189 .zip(marker)
190 .all(|(entry, marker_token)| entry.token == *marker_token);
191 if matches {
192 self.mark_marker_span(span_start, kind);
193 return;
194 }
195 }
196 }
197
198 fn mark_marker_span(&mut self, span_start: usize, kind: MarkerKind) {
199 let next_section = match kind {
200 MarkerKind::ReasoningOpen => SampledTokenSection::Reasoning,
201 MarkerKind::ReasoningClose | MarkerKind::ToolCallClose => SampledTokenSection::Content,
202 MarkerKind::ToolCallOpen => SampledTokenSection::ToolCall,
203 };
204 let span_section = match kind {
205 MarkerKind::ReasoningOpen => SampledTokenSection::Reasoning,
206 MarkerKind::ToolCallOpen => SampledTokenSection::ToolCall,
207 MarkerKind::ReasoningClose => {
208 if self.section == SampledTokenSection::Reasoning {
209 SampledTokenSection::Reasoning
210 } else {
211 SampledTokenSection::Content
212 }
213 }
214 MarkerKind::ToolCallClose => {
215 if self.section == SampledTokenSection::ToolCall {
216 SampledTokenSection::ToolCall
217 } else {
218 SampledTokenSection::Content
219 }
220 }
221 };
222
223 for entry in self.pending.iter_mut().skip(span_start) {
224 entry.is_boundary = true;
225 entry.section = span_section;
226 }
227
228 self.section = next_section;
229 }
230
231 fn drain_overflow(&mut self) -> Vec<IngestOutcome> {
232 let lookback = self.markers.max_token_len().saturating_sub(1);
233 let mut outcomes = Vec::new();
234
235 while let Some(front) = self.pending.front() {
236 if front.is_held_for_probe {
237 break;
238 }
239 let probe_held = self
240 .pending
241 .iter()
242 .filter(|entry| entry.is_held_for_probe)
243 .count();
244 let drainable = self.pending.len().saturating_sub(probe_held);
245 let beyond_lookback = drainable > lookback;
246 if !front.is_boundary && !beyond_lookback {
247 break;
248 }
249 let Some(entry) = self.pending.pop_front() else {
250 break;
251 };
252 if entry.is_from_prompt {
253 continue;
254 }
255 outcomes.push(self.finalize_entry(entry));
256 }
257
258 outcomes
259 }
260
261 fn update_probe(&mut self, piece: &str) -> Vec<IngestOutcome> {
262 let probe_active = matches!(self.probe_mode, ProbeMode::Active(_));
263 if !probe_active {
264 if !self.section_allows_probe_engagement() {
265 return Vec::new();
266 }
267 if !piece.trim_start().starts_with('{') {
268 return Vec::new();
269 }
270 if let Some(entry) = self.pending.back_mut() {
271 entry.is_held_for_probe = true;
272 }
273 self.probe_mode = ProbeMode::Active(JsonProbeState {
274 held_text: piece.to_owned(),
275 });
276 return self.evaluate_probe();
277 }
278
279 if let Some(entry) = self.pending.back_mut() {
280 entry.is_held_for_probe = true;
281 }
282 if let ProbeMode::Active(state) = &mut self.probe_mode {
283 state.held_text.push_str(piece);
284 }
285 self.evaluate_probe()
286 }
287
288 const fn section_allows_probe_engagement(&self) -> bool {
289 matches!(
290 self.section,
291 SampledTokenSection::Content | SampledTokenSection::Pending
292 )
293 }
294
295 fn evaluate_probe(&mut self) -> Vec<IngestOutcome> {
296 let outcome = match &self.probe_mode {
297 ProbeMode::Active(state) => JsonProbeOutcome::validate_prefix(&state.held_text),
298 ProbeMode::Idle => return Vec::new(),
299 };
300 match outcome {
301 JsonProbeOutcome::StillPossiblyValid => Vec::new(),
302 JsonProbeOutcome::CompletedValid => self.commit_probe_as_tool_call(),
303 JsonProbeOutcome::Failed => self.abandon_probe(),
304 }
305 }
306
307 fn commit_probe_as_tool_call(&mut self) -> Vec<IngestOutcome> {
308 if !matches!(self.probe_mode, ProbeMode::Active(_)) {
309 return Vec::new();
310 }
311 self.probe_mode = ProbeMode::Idle;
312 self.section = SampledTokenSection::Content;
313
314 let drained: Vec<_> = self.pending.drain(..).collect();
315 let mut outcomes = Vec::new();
316 for mut entry in drained {
317 if entry.is_held_for_probe {
318 entry.section = SampledTokenSection::ToolCall;
319 entry.is_held_for_probe = false;
320 if !entry.is_from_prompt {
321 outcomes.push(self.finalize_entry(entry));
322 }
323 } else {
324 self.pending.push_back(entry);
325 }
326 }
327 outcomes
328 }
329
330 fn abandon_probe(&mut self) -> Vec<IngestOutcome> {
331 if !matches!(self.probe_mode, ProbeMode::Active(_)) {
332 return Vec::new();
333 }
334 self.probe_mode = ProbeMode::Idle;
335
336 let drained: Vec<_> = self.pending.drain(..).collect();
337 let mut outcomes = Vec::new();
338 for mut entry in drained {
339 if entry.is_held_for_probe {
340 entry.is_held_for_probe = false;
341 if !entry.is_from_prompt {
342 outcomes.push(self.finalize_entry(entry));
343 }
344 } else {
345 self.pending.push_back(entry);
346 }
347 }
348 outcomes
349 }
350
351 fn finalize_entry(&mut self, entry: PendingToken) -> IngestOutcome {
352 let section = entry.section;
353 match section {
354 SampledTokenSection::Reasoning => self.usage.record_reasoning_token(),
355 SampledTokenSection::Content => self.usage.record_content_token(),
356 SampledTokenSection::ToolCall => self.usage.record_tool_call_token(),
357 SampledTokenSection::Pending => self.usage.record_undeterminable_token(),
358 }
359
360 let sampled_token = match section {
361 SampledTokenSection::Reasoning => SampledToken::Reasoning(entry.token),
362 SampledTokenSection::Content => SampledToken::Content(entry.token),
363 SampledTokenSection::ToolCall => SampledToken::ToolCall(entry.token),
364 SampledTokenSection::Pending => SampledToken::Undeterminable(entry.token),
365 };
366
367 let visible_piece = if entry.is_boundary {
368 String::new()
369 } else {
370 entry.decoded.clone()
371 };
372
373 IngestOutcome {
374 sampled_token,
375 visible_piece,
376 raw_piece: entry.decoded,
377 }
378 }
379
380 pub fn sample(
387 &mut self,
388 sampler: &mut LlamaSampler,
389 context: &LlamaContext,
390 idx: i32,
391 ) -> Result<(LlamaToken, Vec<IngestOutcome>), SampleError> {
392 let raw = sampler.sample(context, idx)?;
393 let outcomes = self.ingest(raw)?;
394
395 Ok((raw, outcomes))
396 }
397
398 pub fn feed_prompt_to_batch(
401 &mut self,
402 batch: &mut LlamaBatch,
403 token: LlamaToken,
404 position: llama_pos,
405 seq_ids: &[llama_seq_id],
406 logits: bool,
407 ) -> Result<(), BatchAddError> {
408 batch.add(&SampledToken::Content(token), position, seq_ids, logits)?;
409 self.ingest_prompt_token(token);
410 self.pending_prompt_tokens = self.pending_prompt_tokens.saturating_add(1);
411
412 Ok(())
413 }
414
415 pub fn feed_prompt_sequence_to_batch(
418 &mut self,
419 batch: &mut LlamaBatch,
420 tokens: &[LlamaToken],
421 seq_id: llama_seq_id,
422 logits_all: bool,
423 ) -> Result<(), BatchAddError> {
424 batch.add_sequence(tokens, seq_id, logits_all)?;
425 self.ingest_prompt_tokens(tokens);
426 self.pending_prompt_tokens = self
427 .pending_prompt_tokens
428 .saturating_add(tokens.len() as u64);
429
430 Ok(())
431 }
432
433 pub const fn commit_prompt_tokens(&mut self) -> u64 {
434 let promoted = self.pending_prompt_tokens;
435 self.usage.record_prompt_tokens(promoted);
436 self.pending_prompt_tokens = 0;
437
438 promoted
439 }
440
441 pub const fn discard_pending_prompt_tokens(&mut self) -> u64 {
442 let discarded = self.pending_prompt_tokens;
443 self.pending_prompt_tokens = 0;
444
445 discarded
446 }
447
448 #[must_use]
449 pub const fn pending_prompt_tokens(&self) -> u64 {
450 self.pending_prompt_tokens
451 }
452
453 pub fn eval_multimodal_chunks(
461 &mut self,
462 chunks: &MtmdInputChunks,
463 mtmd_ctx: &MtmdContext,
464 llama_ctx: &LlamaContext,
465 params: EvalMultimodalChunksParams,
466 ) -> Result<llama_pos, EvalMultimodalChunksError> {
467 let chunk_count = chunks.len();
468 let mut next_position = params.start_position;
469
470 for index in 0..chunk_count {
471 let chunk = chunks
472 .get(index)
473 .ok_or(EvalMultimodalChunksError::ChunkOutOfBounds(index))?;
474 let logits_for_this_chunk = params.logits_last && index + 1 == chunk_count;
475
476 next_position = chunk.eval_single(
477 mtmd_ctx,
478 llama_ctx,
479 next_position,
480 params.seq_id,
481 params.n_batch,
482 logits_for_this_chunk,
483 )?;
484 crate::ingest_prompt_chunk::ingest_prompt_chunk(self, &chunk)?;
485 }
486
487 Ok(next_position)
488 }
489
490 pub const fn record_prompt_tokens(&mut self, count: u64) {
491 self.usage.record_prompt_tokens(count);
492 }
493
494 pub const fn record_input_image_tokens(&mut self, count: u64) {
495 self.usage.record_input_image_tokens(count);
496 }
497
498 pub const fn record_input_audio_tokens(&mut self, count: u64) {
499 self.usage.record_input_audio_tokens(count);
500 }
501
502 pub const fn record_cached_prompt_tokens(&mut self, count: u64) -> Result<(), TokenUsageError> {
506 self.usage.record_cached_prompt_tokens(count)
507 }
508
509 #[must_use]
510 pub const fn usage(&self) -> &TokenUsage {
511 &self.usage
512 }
513
514 #[must_use]
515 pub fn into_usage(self) -> TokenUsage {
516 self.usage
517 }
518
519 #[must_use]
520 pub const fn current_section(&self) -> SampledTokenSection {
521 self.section
522 }
523
524 #[must_use]
525 pub const fn markers(&self) -> &StreamingMarkers {
526 &self.markers
527 }
528}
529
530#[cfg(test)]
531mod tests {
532 use super::JsonProbeState;
533 use super::PendingToken;
534 use super::ProbeMode;
535 use super::SampledTokenClassifier;
536 use crate::ingest_outcome::IngestOutcome;
537 use crate::sampled_token::SampledToken;
538 use crate::sampled_token_section::SampledTokenSection;
539 use crate::streaming_markers::StreamingMarkers;
540 use crate::token::LlamaToken;
541
542 fn token(id: i32) -> LlamaToken {
543 LlamaToken::new(id)
544 }
545
546 fn markers_with(
547 reasoning_open: Option<Vec<LlamaToken>>,
548 reasoning_close: Option<Vec<LlamaToken>>,
549 ) -> StreamingMarkers {
550 StreamingMarkers {
551 reasoning_open,
552 reasoning_close,
553 tool_call_open: None,
554 tool_call_close: None,
555 }
556 }
557
558 fn synthetic_classifier(markers: StreamingMarkers) -> SampledTokenClassifier<'static> {
559 SampledTokenClassifier {
560 model: unsafe { &*std::ptr::NonNull::<crate::model::LlamaModel>::dangling().as_ptr() },
561 markers,
562 decoder: encoding_rs::UTF_8.new_decoder(),
563 pending: std::collections::VecDeque::new(),
564 section: SampledTokenSection::Pending,
565 pending_prompt_tokens: 0,
566 usage: llama_cpp_bindings_types::TokenUsage::new(),
567 probe_mode: ProbeMode::Idle,
568 }
569 }
570
571 fn push_pending(classifier: &mut SampledTokenClassifier<'_>, token_id: i32, decoded: &str) {
572 classifier.pending.push_back(PendingToken {
573 token: token(token_id),
574 decoded: decoded.to_owned(),
575 section: classifier.section,
576 is_boundary: false,
577 is_from_prompt: false,
578 is_held_for_probe: false,
579 });
580 }
581
582 fn push_pending_from_prompt(classifier: &mut SampledTokenClassifier<'_>, token_id: i32) {
583 classifier.pending.push_back(PendingToken {
584 token: token(token_id),
585 decoded: String::new(),
586 section: classifier.section,
587 is_boundary: false,
588 is_from_prompt: true,
589 is_held_for_probe: false,
590 });
591 }
592
593 fn push_and_probe(
594 classifier: &mut SampledTokenClassifier<'_>,
595 token_id: i32,
596 decoded: &str,
597 ) -> Vec<IngestOutcome> {
598 push_pending(classifier, token_id, decoded);
599 classifier.try_consume_marker_at_tail();
600 let mut outcomes = classifier.classify_pending_tail(decoded);
601 outcomes.extend(classifier.drain_overflow());
602 outcomes
603 }
604
605 fn outcome_pieces(outcomes: &[IngestOutcome]) -> Vec<&str> {
606 outcomes
607 .iter()
608 .map(|outcome| outcome.visible_piece.as_str())
609 .collect()
610 }
611
612 fn outcome_sections(outcomes: &[IngestOutcome]) -> Vec<SampledTokenSection> {
613 outcomes
614 .iter()
615 .map(|outcome| match outcome.sampled_token {
616 SampledToken::Reasoning(_) => SampledTokenSection::Reasoning,
617 SampledToken::Content(_) => SampledTokenSection::Content,
618 SampledToken::ToolCall(_) => SampledTokenSection::ToolCall,
619 SampledToken::Undeterminable(_) => SampledTokenSection::Pending,
620 })
621 .collect()
622 }
623
624 #[test]
625 fn single_token_close_marker_when_already_in_reasoning_emits_empty_piece_for_marker() {
626 let markers = markers_with(Some(vec![token(100)]), Some(vec![token(200)]));
627 let mut classifier = synthetic_classifier(markers);
628 classifier.section = SampledTokenSection::Reasoning;
629
630 push_pending(&mut classifier, 7, "step");
631 classifier.try_consume_marker_at_tail();
632 let mut outcomes = classifier.drain_overflow();
633
634 push_pending(&mut classifier, 200, "</think>");
635 classifier.try_consume_marker_at_tail();
636 outcomes.extend(classifier.drain_overflow());
637
638 push_pending(&mut classifier, 9, "Hi");
639 classifier.try_consume_marker_at_tail();
640 outcomes.extend(classifier.drain_overflow());
641
642 outcomes.extend(classifier.flush());
643
644 assert_eq!(
645 outcome_sections(&outcomes),
646 vec![
647 SampledTokenSection::Reasoning,
648 SampledTokenSection::Reasoning,
649 SampledTokenSection::Content,
650 ],
651 );
652 assert_eq!(outcome_pieces(&outcomes), vec!["step", "", "Hi"]);
653 assert_eq!(classifier.section, SampledTokenSection::Content);
654 }
655
656 #[test]
657 fn multi_token_close_marker_suppresses_every_marker_token() {
658 let markers = markers_with(
659 Some(vec![token(100)]),
660 Some(vec![token(200), token(201), token(202)]),
661 );
662 let mut classifier = synthetic_classifier(markers);
663 classifier.section = SampledTokenSection::Reasoning;
664
665 let mut outcomes = Vec::new();
666 for (id, decoded) in [(7, "r"), (200, "</"), (201, "thi"), (202, "nk>"), (9, "OK")] {
667 push_pending(&mut classifier, id, decoded);
668 classifier.try_consume_marker_at_tail();
669 outcomes.extend(classifier.drain_overflow());
670 }
671 outcomes.extend(classifier.flush());
672
673 assert_eq!(outcome_pieces(&outcomes), vec!["r", "", "", "", "OK"]);
674 assert_eq!(classifier.section, SampledTokenSection::Content);
675 }
676
677 #[test]
678 fn marker_prefix_that_diverges_does_not_suppress_buffered_tokens() {
679 let markers = markers_with(
680 Some(vec![token(100)]),
681 Some(vec![token(200), token(201), token(202)]),
682 );
683 let mut classifier = synthetic_classifier(markers);
684 classifier.section = SampledTokenSection::Reasoning;
685
686 let mut outcomes = Vec::new();
687 for (id, decoded) in [(7, "r"), (200, "a"), (201, "b"), (300, "x")] {
688 push_pending(&mut classifier, id, decoded);
689 classifier.try_consume_marker_at_tail();
690 outcomes.extend(classifier.drain_overflow());
691 }
692 outcomes.extend(classifier.flush());
693
694 assert_eq!(outcome_pieces(&outcomes), vec!["r", "a", "b", "x"]);
695 assert!(outcomes.iter().all(|outcome| {
696 std::mem::discriminant(&outcome.sampled_token)
697 == std::mem::discriminant(&SampledToken::Reasoning(LlamaToken::new(0)))
698 }));
699 assert_eq!(classifier.section, SampledTokenSection::Reasoning);
700 }
701
702 #[test]
703 fn open_then_close_back_to_back_emits_two_empty_pieces_around_zero_content() {
704 let markers = markers_with(Some(vec![token(100)]), Some(vec![token(200)]));
705 let mut classifier = synthetic_classifier(markers);
706 classifier.section = SampledTokenSection::Content;
707
708 let mut outcomes = Vec::new();
709 for (id, decoded) in [(100, "<think>"), (200, "</think>"), (9, "Hi")] {
710 push_pending(&mut classifier, id, decoded);
711 classifier.try_consume_marker_at_tail();
712 outcomes.extend(classifier.drain_overflow());
713 }
714 outcomes.extend(classifier.flush());
715
716 assert_eq!(
717 outcome_sections(&outcomes),
718 vec![
719 SampledTokenSection::Reasoning,
720 SampledTokenSection::Reasoning,
721 SampledTokenSection::Content,
722 ],
723 );
724 assert_eq!(outcome_pieces(&outcomes), vec!["", "", "Hi"]);
725 assert_eq!(classifier.section, SampledTokenSection::Content);
726 }
727
728 #[test]
729 fn spurious_reasoning_close_in_content_section_classifies_as_content() {
730 let markers = markers_with(Some(vec![token(100)]), Some(vec![token(200)]));
731 let mut classifier = synthetic_classifier(markers);
732 classifier.section = SampledTokenSection::Content;
733
734 push_pending(&mut classifier, 200, "</think>");
735 classifier.try_consume_marker_at_tail();
736 let outcomes = classifier.drain_overflow();
737
738 assert_eq!(
739 outcome_sections(&outcomes),
740 vec![SampledTokenSection::Content],
741 );
742 assert_eq!(classifier.section, SampledTokenSection::Content);
743 }
744
745 #[test]
746 fn spurious_tool_call_close_in_reasoning_section_classifies_as_tool_call() {
747 let markers = StreamingMarkers {
748 reasoning_open: Some(vec![token(100)]),
749 reasoning_close: Some(vec![token(200)]),
750 tool_call_open: Some(vec![token(300)]),
751 tool_call_close: Some(vec![token(400)]),
752 };
753 let mut classifier = synthetic_classifier(markers);
754 classifier.section = SampledTokenSection::ToolCall;
755
756 push_pending(&mut classifier, 400, "</tool_call>");
757 classifier.try_consume_marker_at_tail();
758 let outcomes = classifier.drain_overflow();
759
760 assert_eq!(
761 outcome_sections(&outcomes),
762 vec![SampledTokenSection::ToolCall],
763 );
764 assert_eq!(classifier.section, SampledTokenSection::Content);
765 }
766
767 #[test]
768 fn flush_drains_remaining_pending_at_eog() {
769 let markers = markers_with(
770 Some(vec![token(100)]),
771 Some(vec![token(200), token(201), token(202)]),
772 );
773 let mut classifier = synthetic_classifier(markers);
774 classifier.section = SampledTokenSection::Reasoning;
775
776 push_pending(&mut classifier, 7, "abc");
777 push_pending(&mut classifier, 200, "</");
778 push_pending(&mut classifier, 201, "th");
779
780 let outcomes = classifier.flush();
781
782 assert_eq!(outcome_pieces(&outcomes), vec!["abc", "</", "th"]);
783 assert!(classifier.pending.is_empty());
784 }
785
786 #[test]
787 fn no_markers_marks_each_token_undeterminable_with_visible_piece() {
788 let markers = StreamingMarkers::default();
789 let mut classifier = synthetic_classifier(markers);
790
791 push_pending(&mut classifier, 1, "h");
792 push_pending(&mut classifier, 2, "i");
793 let outcomes = classifier.flush();
794
795 assert_eq!(outcome_pieces(&outcomes), vec!["h", "i"]);
796 assert_eq!(
797 outcome_sections(&outcomes),
798 vec![SampledTokenSection::Pending, SampledTokenSection::Pending],
799 );
800 }
801
802 #[test]
803 fn ingest_prompt_tokens_without_markers_is_noop() {
804 let markers = StreamingMarkers::default();
805 let mut classifier = synthetic_classifier(markers);
806
807 push_pending_from_prompt(&mut classifier, 7);
808 push_pending_from_prompt(&mut classifier, 8);
809
810 assert_eq!(classifier.section, SampledTokenSection::Pending);
811 assert_eq!(classifier.usage().reasoning_tokens, 0);
812 assert_eq!(classifier.usage().content_tokens, 0);
813 assert_eq!(classifier.usage().tool_call_tokens, 0);
814 assert_eq!(classifier.usage().undeterminable_tokens, 0);
815 }
816
817 #[test]
818 fn ingest_prompt_tokens_through_open_close_pair_ends_in_content() {
819 let markers = markers_with(Some(vec![token(100)]), Some(vec![token(200)]));
820 let mut classifier = synthetic_classifier(markers);
821
822 for token_id in [100, 7, 200] {
823 push_pending_from_prompt(&mut classifier, token_id);
824 classifier.try_consume_marker_at_tail();
825 classifier.drain_overflow();
826 }
827
828 assert_eq!(classifier.section, SampledTokenSection::Content);
829 assert_eq!(classifier.usage().reasoning_tokens, 0);
830 assert_eq!(classifier.usage().content_tokens, 0);
831 assert_eq!(classifier.usage().tool_call_tokens, 0);
832 assert_eq!(classifier.usage().undeterminable_tokens, 0);
833 }
834
835 #[test]
836 fn ingest_prompt_tokens_through_open_only_ends_in_reasoning() {
837 let markers = markers_with(Some(vec![token(100)]), Some(vec![token(200)]));
838 let mut classifier = synthetic_classifier(markers);
839
840 for token_id in [100, 7] {
841 push_pending_from_prompt(&mut classifier, token_id);
842 classifier.try_consume_marker_at_tail();
843 classifier.drain_overflow();
844 }
845
846 assert_eq!(classifier.section, SampledTokenSection::Reasoning);
847 assert_eq!(classifier.usage().reasoning_tokens, 0);
848 assert_eq!(classifier.usage().content_tokens, 0);
849 }
850
851 #[test]
852 fn ingest_prompt_tokens_does_not_record_usage() {
853 let markers = markers_with(
854 Some(vec![token(100)]),
855 Some(vec![token(200), token(201), token(202)]),
856 );
857 let mut classifier = synthetic_classifier(markers);
858
859 for token_id in [100, 7, 8, 9, 200, 201, 202, 11] {
860 push_pending_from_prompt(&mut classifier, token_id);
861 classifier.try_consume_marker_at_tail();
862 classifier.drain_overflow();
863 }
864 let drained = classifier.flush();
865 assert!(drained.is_empty());
866
867 assert_eq!(classifier.usage().reasoning_tokens, 0);
868 assert_eq!(classifier.usage().content_tokens, 0);
869 assert_eq!(classifier.usage().tool_call_tokens, 0);
870 assert_eq!(classifier.usage().undeterminable_tokens, 0);
871 }
872
873 #[test]
874 fn prompt_token_completing_marker_with_generated_token_is_suppressed_correctly() {
875 let markers = markers_with(
876 Some(vec![token(100)]),
877 Some(vec![token(200), token(201), token(202)]),
878 );
879 let mut classifier = synthetic_classifier(markers);
880 classifier.section = SampledTokenSection::Reasoning;
881
882 for token_id in [200, 201] {
883 push_pending_from_prompt(&mut classifier, token_id);
884 classifier.try_consume_marker_at_tail();
885 classifier.drain_overflow();
886 }
887
888 assert_eq!(classifier.section, SampledTokenSection::Reasoning);
889 assert_eq!(classifier.pending.len(), 2);
890
891 classifier.pending.push_back(PendingToken {
892 token: token(202),
893 decoded: "k>".to_owned(),
894 section: classifier.section,
895 is_boundary: false,
896 is_from_prompt: false,
897 is_held_for_probe: false,
898 });
899 classifier.try_consume_marker_at_tail();
900 let outcomes = classifier.drain_overflow();
901
902 assert_eq!(outcomes.len(), 1);
903 assert_eq!(
904 std::mem::discriminant(&outcomes[0].sampled_token),
905 std::mem::discriminant(&SampledToken::Reasoning(LlamaToken::new(0)))
906 );
907 assert_eq!(outcomes[0].visible_piece, "");
908 assert_eq!(outcomes[0].raw_piece, "k>");
909
910 assert_eq!(classifier.section, SampledTokenSection::Content);
911 assert_eq!(classifier.usage().reasoning_tokens, 1);
912 assert_eq!(classifier.usage().content_tokens, 0);
913 }
914
915 #[test]
916 fn ingest_prompt_tokens_with_multiple_round_trips_ends_in_content() {
917 let markers = markers_with(Some(vec![token(100)]), Some(vec![token(200)]));
918 let mut classifier = synthetic_classifier(markers);
919
920 for token_id in [100, 7, 200, 100, 8, 200] {
921 push_pending_from_prompt(&mut classifier, token_id);
922 classifier.try_consume_marker_at_tail();
923 classifier.drain_overflow();
924 }
925
926 assert_eq!(classifier.section, SampledTokenSection::Content);
927 assert_eq!(classifier.usage().reasoning_tokens, 0);
928 assert_eq!(classifier.usage().content_tokens, 0);
929 assert_eq!(classifier.usage().tool_call_tokens, 0);
930 assert_eq!(classifier.usage().undeterminable_tokens, 0);
931 }
932
933 #[test]
934 fn ingest_prompt_tokens_initial_section_is_always_pending() {
935 let markers = markers_with(Some(vec![token(100)]), Some(vec![token(200)]));
936 let classifier = synthetic_classifier(markers);
937
938 assert_eq!(classifier.section, SampledTokenSection::Pending);
939 }
940
941 #[test]
942 fn ingest_prompt_tokens_then_drain_for_generated_token_classifies_correctly() {
943 let markers = markers_with(Some(vec![token(100)]), Some(vec![token(200)]));
944 let mut classifier = synthetic_classifier(markers);
945
946 for token_id in [100, 7, 200] {
947 push_pending_from_prompt(&mut classifier, token_id);
948 classifier.try_consume_marker_at_tail();
949 classifier.drain_overflow();
950 }
951
952 assert_eq!(classifier.section, SampledTokenSection::Content);
953 assert_eq!(classifier.usage().reasoning_tokens, 0);
954 assert_eq!(classifier.usage().content_tokens, 0);
955
956 classifier.pending.push_back(PendingToken {
957 token: token(50),
958 decoded: "hi".to_owned(),
959 section: classifier.section,
960 is_boundary: false,
961 is_from_prompt: false,
962 is_held_for_probe: false,
963 });
964 classifier.try_consume_marker_at_tail();
965 let outcomes = classifier.drain_overflow();
966
967 assert_eq!(outcomes.len(), 1);
968 assert_eq!(
969 std::mem::discriminant(&outcomes[0].sampled_token),
970 std::mem::discriminant(&SampledToken::Content(LlamaToken::new(0)))
971 );
972 assert_eq!(outcomes[0].visible_piece, "hi");
973 assert_eq!(classifier.usage().content_tokens, 1);
974 assert_eq!(classifier.usage().reasoning_tokens, 0);
975 assert_eq!(classifier.usage().undeterminable_tokens, 0);
976 }
977
978 #[test]
979 fn close_marker_in_content_section_is_suppressed_as_boundary() {
980 let markers = markers_with(Some(vec![token(100)]), Some(vec![token(200)]));
981 let mut classifier = synthetic_classifier(markers);
982 classifier.section = SampledTokenSection::Content;
983
984 let mut outcomes = Vec::new();
985 for (id, decoded) in [(7, "hi"), (200, "</think>"), (8, "ok")] {
986 push_pending(&mut classifier, id, decoded);
987 classifier.try_consume_marker_at_tail();
988 outcomes.extend(classifier.drain_overflow());
989 }
990 outcomes.extend(classifier.flush());
991
992 assert_eq!(
993 outcome_sections(&outcomes),
994 vec![
995 SampledTokenSection::Content,
996 SampledTokenSection::Content,
997 SampledTokenSection::Content,
998 ],
999 );
1000 assert_eq!(outcome_pieces(&outcomes), vec!["hi", "", "ok"]);
1001 assert_eq!(classifier.section, SampledTokenSection::Content);
1002 }
1003
1004 #[test]
1005 fn open_marker_in_reasoning_section_is_suppressed_as_boundary() {
1006 let markers = markers_with(Some(vec![token(100)]), Some(vec![token(200)]));
1007 let mut classifier = synthetic_classifier(markers);
1008 classifier.section = SampledTokenSection::Reasoning;
1009
1010 let mut outcomes = Vec::new();
1011 for (id, decoded) in [(7, "step1"), (100, "<think>"), (8, "step2")] {
1012 push_pending(&mut classifier, id, decoded);
1013 classifier.try_consume_marker_at_tail();
1014 outcomes.extend(classifier.drain_overflow());
1015 }
1016 outcomes.extend(classifier.flush());
1017
1018 assert_eq!(outcome_pieces(&outcomes), vec!["step1", "", "step2"]);
1019 assert_eq!(classifier.section, SampledTokenSection::Reasoning);
1020 }
1021
1022 #[test]
1023 fn record_prompt_tokens_updates_usage() {
1024 let markers = markers_with(None, None);
1025 let mut classifier = synthetic_classifier(markers);
1026
1027 classifier.record_prompt_tokens(7);
1028
1029 assert_eq!(classifier.usage().prompt_tokens, 7);
1030 }
1031
1032 #[test]
1033 fn record_cached_prompt_tokens_updates_usage_when_under_limit() {
1034 let markers = markers_with(None, None);
1035 let mut classifier = synthetic_classifier(markers);
1036 classifier.record_prompt_tokens(10);
1037
1038 classifier.record_cached_prompt_tokens(3).unwrap();
1039
1040 assert_eq!(classifier.usage().cached_prompt_tokens, 3);
1041 }
1042
1043 #[test]
1044 fn record_cached_prompt_tokens_returns_error_when_over_prompt_total() {
1045 let markers = markers_with(None, None);
1046 let mut classifier = synthetic_classifier(markers);
1047 classifier.record_prompt_tokens(2);
1048
1049 let result = classifier.record_cached_prompt_tokens(5);
1050
1051 assert!(result.is_err());
1052 }
1053
1054 #[test]
1055 fn markers_accessor_returns_configured_markers() {
1056 let configured = markers_with(Some(vec![token(1)]), Some(vec![token(2)]));
1057 let classifier = synthetic_classifier(configured);
1058
1059 let returned = classifier.markers();
1060
1061 assert_eq!(returned.reasoning_open.as_deref(), Some(&[token(1)][..]));
1062 assert_eq!(returned.reasoning_close.as_deref(), Some(&[token(2)][..]));
1063 }
1064
1065 #[test]
1066 fn into_usage_consumes_classifier_and_yields_usage_snapshot() {
1067 let markers = markers_with(None, None);
1068 let mut classifier = synthetic_classifier(markers);
1069 classifier.record_prompt_tokens(11);
1070
1071 let usage = classifier.into_usage();
1072
1073 assert_eq!(usage.prompt_tokens, 11);
1074 }
1075
1076 #[test]
1077 fn spurious_tool_call_close_in_content_section_classifies_as_content() {
1078 let mut markers = markers_with(None, None);
1079 markers.tool_call_close = Some(vec![token(300)]);
1080 let mut classifier = synthetic_classifier(markers);
1081 classifier.section = SampledTokenSection::Content;
1082
1083 push_pending(&mut classifier, 300, "</tool_call>");
1084 classifier.try_consume_marker_at_tail();
1085 let outcomes = classifier.drain_overflow();
1086
1087 assert_eq!(
1088 outcome_sections(&outcomes),
1089 vec![SampledTokenSection::Content],
1090 );
1091 assert_eq!(classifier.section, SampledTokenSection::Content);
1092 }
1093
1094 fn markers_with_tool_call_open(tool_call_open: Vec<LlamaToken>) -> StreamingMarkers {
1095 StreamingMarkers {
1096 reasoning_open: None,
1097 reasoning_close: None,
1098 tool_call_open: Some(tool_call_open),
1099 tool_call_close: None,
1100 }
1101 }
1102
1103 fn feed_json_string(
1104 classifier: &mut SampledTokenClassifier<'_>,
1105 text: &str,
1106 starting_token_id: i32,
1107 ) -> Vec<IngestOutcome> {
1108 let mut outcomes = Vec::new();
1109 for (offset, ch) in text.char_indices() {
1110 let token_id = starting_token_id + i32::try_from(offset).unwrap_or(i32::MAX);
1111 let mut buffer = [0_u8; 4];
1112 let chunk = ch.encode_utf8(&mut buffer);
1113 outcomes.extend(push_and_probe(classifier, token_id, chunk));
1114 }
1115 outcomes
1116 }
1117
1118 #[test]
1119 fn json_probe_engages_when_first_non_whitespace_is_open_brace_in_content() {
1120 let markers = markers_with_tool_call_open(vec![token(900)]);
1121 let mut classifier = synthetic_classifier(markers);
1122 classifier.section = SampledTokenSection::Content;
1123
1124 push_and_probe(&mut classifier, 1, "{");
1125
1126 assert_ne!(classifier.probe_mode, ProbeMode::Idle);
1127 }
1128
1129 #[test]
1130 fn json_probe_releases_tokens_as_tool_call_when_signature_matches() {
1131 let markers = markers_with_tool_call_open(vec![token(900)]);
1132 let mut classifier = synthetic_classifier(markers);
1133 classifier.section = SampledTokenSection::Content;
1134
1135 let outcomes = feed_json_string(&mut classifier, r#"{"name":"f","arguments":{}}"#, 100);
1136
1137 assert!(!outcomes.is_empty());
1138 let sections = outcome_sections(&outcomes);
1139 assert!(
1140 sections
1141 .iter()
1142 .all(|section| *section == SampledTokenSection::ToolCall),
1143 "every emitted outcome should be ToolCall, got {sections:?}",
1144 );
1145 assert_eq!(classifier.probe_mode, ProbeMode::Idle);
1146 }
1147
1148 #[test]
1149 fn json_probe_releases_tokens_as_content_when_signature_does_not_match() {
1150 let markers = markers_with_tool_call_open(vec![token(900)]);
1151 let mut classifier = synthetic_classifier(markers);
1152 classifier.section = SampledTokenSection::Content;
1153
1154 let outcomes = feed_json_string(&mut classifier, r#"{"foo":"bar"}"#, 100);
1155
1156 let sections = outcome_sections(&outcomes);
1157 assert!(
1158 sections
1159 .iter()
1160 .all(|section| *section == SampledTokenSection::Content),
1161 "every emitted outcome should be Content, got {sections:?}",
1162 );
1163 assert_eq!(classifier.probe_mode, ProbeMode::Idle);
1164 }
1165
1166 #[test]
1167 fn json_probe_releases_tokens_as_content_when_extra_top_level_key() {
1168 let markers = markers_with_tool_call_open(vec![token(900)]);
1169 let mut classifier = synthetic_classifier(markers);
1170 classifier.section = SampledTokenSection::Content;
1171
1172 let outcomes = feed_json_string(
1173 &mut classifier,
1174 r#"{"name":"f","arguments":{},"extra":1}"#,
1175 100,
1176 );
1177
1178 assert!(outcomes.iter().all(|outcome| {
1179 std::mem::discriminant(&outcome.sampled_token)
1180 == std::mem::discriminant(&SampledToken::Content(LlamaToken::new(0)))
1181 }));
1182 }
1183
1184 #[test]
1185 fn json_probe_releases_tokens_as_content_when_arguments_is_not_object() {
1186 let markers = markers_with_tool_call_open(vec![token(900)]);
1187 let mut classifier = synthetic_classifier(markers);
1188 classifier.section = SampledTokenSection::Content;
1189
1190 let outcomes = feed_json_string(&mut classifier, r#"{"name":"f","arguments":"hi"}"#, 100);
1191
1192 assert!(outcomes.iter().all(|outcome| {
1193 std::mem::discriminant(&outcome.sampled_token)
1194 == std::mem::discriminant(&SampledToken::Content(LlamaToken::new(0)))
1195 }));
1196 }
1197
1198 #[test]
1199 fn json_probe_handles_strings_with_quoted_braces_in_arguments() {
1200 let markers = markers_with_tool_call_open(vec![token(900)]);
1201 let mut classifier = synthetic_classifier(markers);
1202 classifier.section = SampledTokenSection::Content;
1203
1204 let outcomes = feed_json_string(
1205 &mut classifier,
1206 r#"{"name":"f","arguments":{"q":"a } b"}}"#,
1207 100,
1208 );
1209
1210 assert!(outcomes.iter().all(|outcome| {
1211 std::mem::discriminant(&outcome.sampled_token)
1212 == std::mem::discriminant(&SampledToken::ToolCall(LlamaToken::new(0)))
1213 }));
1214 }
1215
1216 #[test]
1217 fn json_probe_handles_escaped_quotes_in_string_values() {
1218 let markers = markers_with_tool_call_open(vec![token(900)]);
1219 let mut classifier = synthetic_classifier(markers);
1220 classifier.section = SampledTokenSection::Content;
1221
1222 let outcomes = feed_json_string(
1223 &mut classifier,
1224 r#"{"name":"f","arguments":{"q":"he said \"hi\""}}"#,
1225 100,
1226 );
1227
1228 assert!(outcomes.iter().all(|outcome| {
1229 std::mem::discriminant(&outcome.sampled_token)
1230 == std::mem::discriminant(&SampledToken::ToolCall(LlamaToken::new(0)))
1231 }));
1232 }
1233
1234 #[test]
1235 fn json_probe_handles_unicode_letters_in_strings() {
1236 let markers = markers_with_tool_call_open(vec![token(900)]);
1237 let mut classifier = synthetic_classifier(markers);
1238 classifier.section = SampledTokenSection::Content;
1239
1240 let outcomes = feed_json_string(
1241 &mut classifier,
1242 r#"{"name":"日本語","arguments":{"city":"パリ"}}"#,
1243 100,
1244 );
1245
1246 assert!(outcomes.iter().all(|outcome| {
1247 std::mem::discriminant(&outcome.sampled_token)
1248 == std::mem::discriminant(&SampledToken::ToolCall(LlamaToken::new(0)))
1249 }));
1250 }
1251
1252 #[test]
1253 fn json_probe_handles_nested_objects() {
1254 let markers = markers_with_tool_call_open(vec![token(900)]);
1255 let mut classifier = synthetic_classifier(markers);
1256 classifier.section = SampledTokenSection::Content;
1257
1258 let outcomes = feed_json_string(
1259 &mut classifier,
1260 r#"{"name":"f","arguments":{"a":{"b":{"c":1}}}}"#,
1261 100,
1262 );
1263
1264 assert!(outcomes.iter().all(|outcome| {
1265 std::mem::discriminant(&outcome.sampled_token)
1266 == std::mem::discriminant(&SampledToken::ToolCall(LlamaToken::new(0)))
1267 }));
1268 }
1269
1270 #[test]
1271 fn json_probe_handles_arrays_inside_arguments() {
1272 let markers = markers_with_tool_call_open(vec![token(900)]);
1273 let mut classifier = synthetic_classifier(markers);
1274 classifier.section = SampledTokenSection::Content;
1275
1276 let outcomes = feed_json_string(
1277 &mut classifier,
1278 r#"{"name":"f","arguments":{"items":[1,2,3]}}"#,
1279 100,
1280 );
1281
1282 assert!(outcomes.iter().all(|outcome| {
1283 std::mem::discriminant(&outcome.sampled_token)
1284 == std::mem::discriminant(&SampledToken::ToolCall(LlamaToken::new(0)))
1285 }));
1286 }
1287
1288 #[test]
1289 fn json_probe_does_not_engage_when_first_byte_is_close_brace() {
1290 let markers = markers_with_tool_call_open(vec![token(900)]);
1291 let mut classifier = synthetic_classifier(markers);
1292 classifier.section = SampledTokenSection::Content;
1293
1294 let outcomes = feed_json_string(&mut classifier, "}}", 100);
1295
1296 assert_eq!(classifier.probe_mode, ProbeMode::Idle);
1297 assert!(outcomes.iter().all(|outcome| {
1298 std::mem::discriminant(&outcome.sampled_token)
1299 == std::mem::discriminant(&SampledToken::Content(LlamaToken::new(0)))
1300 }));
1301 }
1302
1303 #[test]
1304 fn json_probe_does_not_engage_in_reasoning_section() {
1305 let markers = StreamingMarkers {
1306 reasoning_open: Some(vec![token(800)]),
1307 reasoning_close: Some(vec![token(801)]),
1308 tool_call_open: Some(vec![token(900)]),
1309 tool_call_close: None,
1310 };
1311 let mut classifier = synthetic_classifier(markers);
1312 classifier.section = SampledTokenSection::Reasoning;
1313
1314 push_and_probe(&mut classifier, 1, "{");
1315
1316 assert_eq!(classifier.probe_mode, ProbeMode::Idle);
1317 }
1318
1319 #[test]
1320 fn json_probe_does_not_engage_in_tool_call_section() {
1321 let markers = markers_with_tool_call_open(vec![token(900)]);
1322 let mut classifier = synthetic_classifier(markers);
1323 classifier.section = SampledTokenSection::ToolCall;
1324
1325 push_and_probe(&mut classifier, 1, "{");
1326
1327 assert_eq!(classifier.probe_mode, ProbeMode::Idle);
1328 }
1329
1330 #[test]
1331 fn marker_probe_takes_precedence_when_both_could_match() {
1332 let markers = markers_with_tool_call_open(vec![token(900)]);
1333 let mut classifier = synthetic_classifier(markers);
1334 classifier.section = SampledTokenSection::Content;
1335
1336 let mut outcomes = Vec::new();
1337 outcomes.extend(push_and_probe(&mut classifier, 1, "{"));
1338 outcomes.extend(push_and_probe(&mut classifier, 900, r#"""#));
1339
1340 assert_eq!(classifier.section, SampledTokenSection::ToolCall);
1341 assert_eq!(outcome_pieces(&outcomes), vec!["{", ""]);
1342 assert_eq!(
1343 outcome_sections(&outcomes),
1344 vec![SampledTokenSection::Content, SampledTokenSection::ToolCall],
1345 );
1346 }
1347
1348 #[test]
1349 fn json_probe_consumes_two_consecutive_objects_separately() {
1350 let markers = markers_with_tool_call_open(vec![token(900)]);
1351 let mut classifier = synthetic_classifier(markers);
1352 classifier.section = SampledTokenSection::Content;
1353
1354 let mut outcomes = Vec::new();
1355 outcomes.extend(feed_json_string(
1356 &mut classifier,
1357 r#"{"name":"a","arguments":{}}"#,
1358 100,
1359 ));
1360 outcomes.extend(feed_json_string(
1361 &mut classifier,
1362 r#"{"name":"b","arguments":{"x":1}}"#,
1363 200,
1364 ));
1365
1366 let sections = outcome_sections(&outcomes);
1367 assert!(
1368 sections
1369 .iter()
1370 .all(|section| *section == SampledTokenSection::ToolCall),
1371 "two consecutive markerless tool calls must both classify as ToolCall, got {sections:?}",
1372 );
1373 }
1374
1375 #[test]
1376 fn json_probe_with_leading_whitespace_then_open_brace_classifies_whitespace_as_content_and_json_as_tool_call()
1377 {
1378 let markers = markers_with_tool_call_open(vec![token(900)]);
1379 let mut classifier = synthetic_classifier(markers);
1380 classifier.section = SampledTokenSection::Content;
1381
1382 let outcomes = feed_json_string(
1383 &mut classifier,
1384 "\n {\"name\":\"f\",\"arguments\":{}}",
1385 100,
1386 );
1387
1388 let tool_call_count = outcomes
1389 .iter()
1390 .filter(|outcome| {
1391 std::mem::discriminant(&outcome.sampled_token)
1392 == std::mem::discriminant(&SampledToken::ToolCall(LlamaToken::new(0)))
1393 })
1394 .count();
1395 let content_count = outcomes
1396 .iter()
1397 .filter(|outcome| {
1398 std::mem::discriminant(&outcome.sampled_token)
1399 == std::mem::discriminant(&SampledToken::Content(LlamaToken::new(0)))
1400 })
1401 .count();
1402 assert_eq!(
1403 content_count, 3,
1404 "leading `\\n ` should classify as content"
1405 );
1406 assert!(
1407 tool_call_count > 0,
1408 "the JSON object should classify as ToolCall",
1409 );
1410 assert_eq!(content_count + tool_call_count, outcomes.len());
1411 }
1412
1413 #[test]
1414 fn json_probe_records_tool_call_token_usage_on_commit() {
1415 let markers = markers_with_tool_call_open(vec![token(900)]);
1416 let mut classifier = synthetic_classifier(markers);
1417 classifier.section = SampledTokenSection::Content;
1418
1419 let json = r#"{"name":"f","arguments":{}}"#;
1420 let outcomes = feed_json_string(&mut classifier, json, 100);
1421
1422 let emitted = outcomes.len();
1423 let usage = classifier.usage();
1424 assert_eq!(usage.tool_call_tokens, emitted as u64);
1425 assert_eq!(usage.content_tokens, 0);
1426 }
1427
1428 #[test]
1429 fn json_probe_records_content_token_usage_on_abandon() {
1430 let markers = markers_with_tool_call_open(vec![token(900)]);
1431 let mut classifier = synthetic_classifier(markers);
1432 classifier.section = SampledTokenSection::Content;
1433
1434 let json = r#"{"foo":"bar"}"#;
1435 let outcomes = feed_json_string(&mut classifier, json, 100);
1436
1437 let emitted = outcomes.len();
1438 let usage = classifier.usage();
1439 assert_eq!(usage.content_tokens, emitted as u64);
1440 assert_eq!(usage.tool_call_tokens, 0);
1441 }
1442
1443 #[test]
1444 fn flush_during_active_json_probe_releases_held_tokens_as_content() {
1445 let markers = markers_with_tool_call_open(vec![token(900)]);
1446 let mut classifier = synthetic_classifier(markers);
1447 classifier.section = SampledTokenSection::Content;
1448
1449 push_and_probe(&mut classifier, 1, "{");
1450 push_and_probe(&mut classifier, 2, r#""name""#);
1451 assert_ne!(classifier.probe_mode, ProbeMode::Idle);
1452
1453 let outcomes = classifier.flush();
1454
1455 let sections = outcome_sections(&outcomes);
1456 assert!(
1457 sections
1458 .iter()
1459 .all(|section| *section == SampledTokenSection::Content),
1460 "mid-probe flush must release held tokens as Content, got {sections:?}",
1461 );
1462 assert_eq!(classifier.probe_mode, ProbeMode::Idle);
1463 }
1464
1465 #[test]
1466 fn evaluate_probe_while_idle_returns_no_outcomes() {
1467 let markers = markers_with_tool_call_open(vec![token(900)]);
1468 let mut classifier = synthetic_classifier(markers);
1469
1470 let outcomes = classifier.evaluate_probe();
1471
1472 assert!(outcomes.is_empty());
1473 }
1474
1475 #[test]
1476 fn commit_probe_as_tool_call_while_idle_returns_no_outcomes() {
1477 let markers = markers_with_tool_call_open(vec![token(900)]);
1478 let mut classifier = synthetic_classifier(markers);
1479
1480 let outcomes = classifier.commit_probe_as_tool_call();
1481
1482 assert!(outcomes.is_empty());
1483 }
1484
1485 #[test]
1486 fn abandon_probe_while_idle_returns_no_outcomes() {
1487 let markers = markers_with_tool_call_open(vec![token(900)]);
1488 let mut classifier = synthetic_classifier(markers);
1489
1490 let outcomes = classifier.abandon_probe();
1491
1492 assert!(outcomes.is_empty());
1493 }
1494
1495 #[test]
1496 fn commit_probe_as_tool_call_requeues_non_held_entries_and_releases_held_as_tool_call() {
1497 let markers = markers_with_tool_call_open(vec![token(900)]);
1498 let mut classifier = synthetic_classifier(markers);
1499 classifier.section = SampledTokenSection::Content;
1500
1501 classifier.pending.push_back(PendingToken {
1502 token: token(1),
1503 decoded: "before".to_owned(),
1504 section: SampledTokenSection::Content,
1505 is_boundary: false,
1506 is_from_prompt: false,
1507 is_held_for_probe: false,
1508 });
1509 classifier.pending.push_back(PendingToken {
1510 token: token(2),
1511 decoded: "{}".to_owned(),
1512 section: SampledTokenSection::Content,
1513 is_boundary: false,
1514 is_from_prompt: false,
1515 is_held_for_probe: true,
1516 });
1517 classifier.probe_mode = ProbeMode::Active(JsonProbeState {
1518 held_text: "{}".to_owned(),
1519 });
1520
1521 let outcomes = classifier.commit_probe_as_tool_call();
1522
1523 let sections = outcome_sections(&outcomes);
1524 assert_eq!(sections, vec![SampledTokenSection::ToolCall]);
1525 assert_eq!(classifier.pending.len(), 1);
1526 assert_eq!(classifier.pending[0].token, token(1));
1527 assert_eq!(classifier.probe_mode, ProbeMode::Idle);
1528 }
1529}