Skip to main content

llama_cpp_bindings/
sampled_token_classifier.rs

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    /// # Errors
77    /// Returns [`TokenToStringError`] when the sampled token cannot be
78    /// detokenised. The failure is surfaced rather than substituting an empty
79    /// piece, so classification never silently drops generated text.
80    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    /// # Errors
381    /// Forwards [`LlamaSampler::sample`] errors verbatim. Nothing is recorded on failure.
382    ///
383    /// Returns the raw sampled token (for downstream `batch.add` / `is_eog_token`
384    /// calls) alongside the outcomes that finalised this turn — see
385    /// [`Self::ingest`] for buffering semantics.
386    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    /// # Errors
399    /// Forwards [`LlamaBatch::add`] errors verbatim. Nothing is staged on failure.
400    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    /// # Errors
416    /// Forwards [`LlamaBatch::add_sequence`] errors verbatim. Nothing is staged on failure.
417    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    /// # Errors
454    /// Returns [`EvalMultimodalChunksError::EvalFailed`] when the underlying
455    /// `eval_chunks` call fails (no counters move),
456    /// [`EvalMultimodalChunksError::UnknownChunkType`] when a chunk reports a
457    /// type unknown to this binding, or
458    /// [`EvalMultimodalChunksError::ChunkOutOfBounds`] when a valid index returns
459    /// `None` from `chunks.get`.
460    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    /// # Errors
503    /// Forwards [`TokenUsageError::CachedExceedsPrompt`] when the running cached total would
504    /// exceed the prompt total.
505    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}