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