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