Skip to main content

ferrum_types/
reasoning.rs

1//! Shared parsing for declared model reasoning protocols.
2
3use crate::{FerrumError, ModelOutputProtocol, Result};
4
5mod gemma;
6
7pub const THINK_START_TAG: &str = "<think>";
8pub const THINK_END_TAG: &str = "</think>";
9pub const GEMMA_THOUGHT_START_TAG: &str = "<|channel>thought\n";
10pub const GEMMA_THOUGHT_END_TAG: &str = "<channel|>";
11
12#[derive(Debug, Clone, PartialEq, Eq)]
13pub struct ParsedReasoningResponse {
14    pub content: String,
15    pub reasoning: Option<String>,
16}
17
18/// The full generated opening header and closing marker of a reasoning block.
19/// Harmony owns a separate message protocol rather than a delimited block.
20pub const fn model_reasoning_markers(
21    protocol: ModelOutputProtocol,
22) -> Option<(&'static str, &'static str)> {
23    match protocol {
24        ModelOutputProtocol::Text => Some((THINK_START_TAG, THINK_END_TAG)),
25        ModelOutputProtocol::GemmaThought => Some((GEMMA_THOUGHT_START_TAG, GEMMA_THOUGHT_END_TAG)),
26        ModelOutputProtocol::HarmonyGptOss => None,
27    }
28}
29
30pub fn has_unclosed_model_reasoning_block(protocol: ModelOutputProtocol, prompt: &str) -> bool {
31    match protocol {
32        ModelOutputProtocol::Text => has_unclosed_thinking_block(prompt),
33        ModelOutputProtocol::HarmonyGptOss => false,
34        ModelOutputProtocol::GemmaThought => {
35            // The generated turn, including a tool-response continuation, owns
36            // the prefill state. A marker mentioned in an earlier user turn
37            // must not make an ordinary model turn start inside reasoning.
38            let turn = prompt
39                .rsplit_once("<|turn>")
40                .map_or(prompt, |(_, turn)| turn);
41            match (
42                turn.rfind(GEMMA_THOUGHT_START_TAG),
43                turn.rfind(GEMMA_THOUGHT_END_TAG),
44            ) {
45                (Some(start), Some(end)) => start > end,
46                (Some(_), None) => true,
47                _ => false,
48            }
49        }
50    }
51}
52
53/// Parse a cumulative generated prefix without exposing partial protocol
54/// markers. Gemma accepts only the declared thought channel; ordinary words
55/// such as "thought" remain ordinary content outside that channel.
56pub fn parse_model_reasoning_response(
57    protocol: ModelOutputProtocol,
58    text: &str,
59    prompt_opened_thinking: bool,
60) -> Result<ParsedReasoningResponse> {
61    match protocol {
62        ModelOutputProtocol::Text => Ok(parse_reasoning_response_for_prompt(
63            text,
64            prompt_opened_thinking,
65        )),
66        ModelOutputProtocol::GemmaThought => gemma::parse(text, prompt_opened_thinking),
67        ModelOutputProtocol::HarmonyGptOss => Err(FerrumError::invalid_request(
68            "Harmony output requires its message-protocol parser",
69        )),
70    }
71}
72
73/// Hold only prefixes that can still become framing, not complete responses.
74/// Callers retain the cumulative input and parse it again when more arrives.
75pub fn should_defer_model_reasoning_stream_delta(
76    protocol: ModelOutputProtocol,
77    text: &str,
78) -> bool {
79    match protocol {
80        ModelOutputProtocol::Text => {
81            let candidate = text.trim_start_matches(['\r', '\n']);
82            candidate.is_empty()
83                || THINK_START_TAG.starts_with(candidate)
84                || THINK_END_TAG.starts_with(candidate)
85        }
86        ModelOutputProtocol::GemmaThought => {
87            text.is_empty() || gemma::has_partial_marker_suffix(text)
88        }
89        ModelOutputProtocol::HarmonyGptOss => false,
90    }
91}
92
93pub fn has_unclosed_thinking_block(prompt: &str) -> bool {
94    match (prompt.rfind(THINK_START_TAG), prompt.rfind(THINK_END_TAG)) {
95        (Some(start), Some(end)) => start > end,
96        (Some(_), None) => true,
97        _ => false,
98    }
99}
100
101/// Parse generated text when the rendered prompt already opened `<think>`.
102/// Some reasoning templates emit only the closing tag in generated text.
103pub fn parse_reasoning_response_started_in_think(text: &str) -> ParsedReasoningResponse {
104    if text.contains(THINK_START_TAG) {
105        return parse_reasoning_response(text);
106    }
107    let Some(end) = text.find(THINK_END_TAG) else {
108        return ParsedReasoningResponse {
109            content: String::new(),
110            reasoning: (!text.is_empty()).then(|| text.to_string()),
111        };
112    };
113    let reasoning = text[..end].to_string();
114    let content = text[end + THINK_END_TAG.len()..]
115        .trim_start_matches(['\r', '\n'])
116        .to_string();
117    ParsedReasoningResponse {
118        content,
119        reasoning: (!reasoning.is_empty()).then_some(reasoning),
120    }
121}
122
123pub fn parse_reasoning_response(text: &str) -> ParsedReasoningResponse {
124    let Some(start) = text.find(THINK_START_TAG) else {
125        if let Some(end) = text.find(THINK_END_TAG) {
126            let reasoning = text[..end].to_string();
127            let content = text[end + THINK_END_TAG.len()..]
128                .trim_start_matches(['\r', '\n'])
129                .to_string();
130            return ParsedReasoningResponse {
131                content,
132                reasoning: (!reasoning.is_empty()).then_some(reasoning),
133            };
134        }
135        return ParsedReasoningResponse {
136            content: text.to_string(),
137            reasoning: None,
138        };
139    };
140
141    let before = &text[..start];
142    let after_start = &text[start + THINK_START_TAG.len()..];
143    let Some(end) = after_start.find(THINK_END_TAG) else {
144        return ParsedReasoningResponse {
145            content: before.to_string(),
146            reasoning: Some(after_start.to_string()),
147        };
148    };
149
150    let reasoning = after_start[..end].to_string();
151    let after_end = &after_start[end + THINK_END_TAG.len()..];
152    let mut content = String::with_capacity(before.len() + after_end.len());
153    content.push_str(before);
154    content.push_str(after_end.trim_start_matches(['\r', '\n']));
155    ParsedReasoningResponse {
156        content,
157        reasoning: (!reasoning.is_empty()).then_some(reasoning),
158    }
159}
160
161pub fn parse_reasoning_response_for_prompt(
162    text: &str,
163    prompt_opened_thinking: bool,
164) -> ParsedReasoningResponse {
165    if prompt_opened_thinking {
166        parse_reasoning_response_started_in_think(text)
167    } else {
168        parse_reasoning_response(text)
169    }
170}
171
172#[cfg(test)]
173mod tests {
174    use super::*;
175
176    #[test]
177    fn parses_explicit_reasoning_block() {
178        let parsed = parse_reasoning_response("<think>reason</think>\nanswer");
179        assert_eq!(parsed.reasoning.as_deref(), Some("reason"));
180        assert_eq!(parsed.content, "answer");
181    }
182
183    #[test]
184    fn parses_prompt_opened_reasoning_block() {
185        let parsed = parse_reasoning_response_started_in_think("reason</think>\nanswer");
186        assert_eq!(parsed.reasoning.as_deref(), Some("reason"));
187        assert_eq!(parsed.content, "answer");
188    }
189
190    #[test]
191    fn preserves_plain_content() {
192        let parsed = parse_reasoning_response("answer");
193        assert_eq!(parsed.reasoning, None);
194        assert_eq!(parsed.content, "answer");
195    }
196
197    #[test]
198    fn detects_prompt_opened_reasoning() {
199        assert!(has_unclosed_thinking_block("assistant:<think>\n"));
200        assert!(!has_unclosed_thinking_block(
201            "assistant:<think>reason</think>\n"
202        ));
203    }
204}