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
7/// Resolved reasoning behavior of the model-owned template and output protocol.
8/// Unknown is not equivalent to a template that explicitly has no reasoning mode.
9#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
10#[serde(rename_all = "snake_case")]
11pub enum ModelReasoningProtocol {
12    #[default]
13    Unknown,
14    None,
15    PromptOpened,
16    ModelGenerated,
17}
18
19impl ModelReasoningProtocol {
20    pub const fn supports_reasoning(self) -> bool {
21        matches!(self, Self::PromptOpened | Self::ModelGenerated)
22    }
23}
24
25pub const THINK_START_TAG: &str = "<think>";
26pub const THINK_END_TAG: &str = "</think>";
27pub const GEMMA_THOUGHT_START_TAG: &str = "<|channel>thought\n";
28pub const GEMMA_THOUGHT_END_TAG: &str = "<channel|>";
29
30#[derive(Debug, Clone, PartialEq, Eq)]
31pub struct ParsedReasoningResponse {
32    pub content: String,
33    pub reasoning: Option<String>,
34}
35
36/// The full generated opening header and closing marker of a reasoning block.
37/// Harmony owns a separate message protocol rather than a delimited block.
38pub const fn model_reasoning_markers(
39    protocol: ModelOutputProtocol,
40) -> Option<(&'static str, &'static str)> {
41    match protocol {
42        ModelOutputProtocol::Text => Some((THINK_START_TAG, THINK_END_TAG)),
43        ModelOutputProtocol::GemmaThought => Some((GEMMA_THOUGHT_START_TAG, GEMMA_THOUGHT_END_TAG)),
44        ModelOutputProtocol::HarmonyGptOss => None,
45    }
46}
47
48pub fn has_unclosed_model_reasoning_block(protocol: ModelOutputProtocol, prompt: &str) -> bool {
49    match protocol {
50        ModelOutputProtocol::Text => has_unclosed_thinking_block(prompt),
51        ModelOutputProtocol::HarmonyGptOss => false,
52        ModelOutputProtocol::GemmaThought => {
53            // The generated turn, including a tool-response continuation, owns
54            // the prefill state. A marker mentioned in an earlier user turn
55            // must not make an ordinary model turn start inside reasoning.
56            let turn = prompt
57                .rsplit_once("<|turn>")
58                .map_or(prompt, |(_, turn)| turn);
59            match (
60                turn.rfind(GEMMA_THOUGHT_START_TAG),
61                turn.rfind(GEMMA_THOUGHT_END_TAG),
62            ) {
63                (Some(start), Some(end)) => start > end,
64                (Some(_), None) => true,
65                _ => false,
66            }
67        }
68    }
69}
70
71/// Parse a cumulative generated prefix without exposing partial protocol
72/// markers. Gemma accepts only the declared thought channel; ordinary words
73/// such as "thought" remain ordinary content outside that channel.
74pub fn parse_model_reasoning_response(
75    protocol: ModelOutputProtocol,
76    text: &str,
77    prompt_opened_thinking: bool,
78) -> Result<ParsedReasoningResponse> {
79    match protocol {
80        ModelOutputProtocol::Text => Ok(parse_reasoning_response_for_prompt(
81            text,
82            prompt_opened_thinking,
83        )),
84        ModelOutputProtocol::GemmaThought => gemma::parse(text, prompt_opened_thinking),
85        ModelOutputProtocol::HarmonyGptOss => Err(FerrumError::invalid_request(
86            "Harmony output requires its message-protocol parser",
87        )),
88    }
89}
90
91/// Hold only prefixes that can still become framing, not complete responses.
92/// Callers retain the cumulative input and parse it again when more arrives.
93pub fn should_defer_model_reasoning_stream_delta(
94    protocol: ModelOutputProtocol,
95    text: &str,
96) -> bool {
97    match protocol {
98        ModelOutputProtocol::Text => {
99            let candidate = text.trim_start_matches(['\r', '\n']);
100            candidate.is_empty()
101                || THINK_START_TAG.starts_with(candidate)
102                || THINK_END_TAG.starts_with(candidate)
103        }
104        ModelOutputProtocol::GemmaThought => {
105            text.is_empty() || gemma::has_partial_marker_suffix(text)
106        }
107        ModelOutputProtocol::HarmonyGptOss => false,
108    }
109}
110
111pub fn has_unclosed_thinking_block(prompt: &str) -> bool {
112    match (prompt.rfind(THINK_START_TAG), prompt.rfind(THINK_END_TAG)) {
113        (Some(start), Some(end)) => start > end,
114        (Some(_), None) => true,
115        _ => false,
116    }
117}
118
119/// Parse generated text when the rendered prompt already opened `<think>`.
120/// Some reasoning templates emit only the closing tag in generated text.
121pub fn parse_reasoning_response_started_in_think(text: &str) -> ParsedReasoningResponse {
122    let end = text.find(THINK_END_TAG);
123    // A generated opener before the closer can repeat the template's opener.
124    // Once the prompt-opened block closes, further tags belong to the content.
125    if text
126        .find(THINK_START_TAG)
127        .is_some_and(|start| end.is_none_or(|end| start < end))
128    {
129        return parse_reasoning_response(text);
130    }
131    let Some(end) = end else {
132        return ParsedReasoningResponse {
133            content: String::new(),
134            reasoning: (!text.is_empty()).then(|| text.to_string()),
135        };
136    };
137    let reasoning = text[..end].to_string();
138    let content = text[end + THINK_END_TAG.len()..]
139        .trim_start_matches(['\r', '\n'])
140        .to_string();
141    ParsedReasoningResponse {
142        content,
143        reasoning: (!reasoning.is_empty()).then_some(reasoning),
144    }
145}
146
147pub fn parse_reasoning_response(text: &str) -> ParsedReasoningResponse {
148    let Some(start) = text.find(THINK_START_TAG) else {
149        if let Some(end) = text.find(THINK_END_TAG) {
150            let reasoning = text[..end].to_string();
151            let content = text[end + THINK_END_TAG.len()..]
152                .trim_start_matches(['\r', '\n'])
153                .to_string();
154            return ParsedReasoningResponse {
155                content,
156                reasoning: (!reasoning.is_empty()).then_some(reasoning),
157            };
158        }
159        return ParsedReasoningResponse {
160            content: text.to_string(),
161            reasoning: None,
162        };
163    };
164
165    let before = &text[..start];
166    let after_start = &text[start + THINK_START_TAG.len()..];
167    let Some(end) = after_start.find(THINK_END_TAG) else {
168        return ParsedReasoningResponse {
169            content: before.to_string(),
170            reasoning: Some(after_start.to_string()),
171        };
172    };
173
174    let reasoning = after_start[..end].to_string();
175    let after_end = &after_start[end + THINK_END_TAG.len()..];
176    let mut content = String::with_capacity(before.len() + after_end.len());
177    content.push_str(before);
178    content.push_str(after_end.trim_start_matches(['\r', '\n']));
179    ParsedReasoningResponse {
180        content,
181        reasoning: (!reasoning.is_empty()).then_some(reasoning),
182    }
183}
184
185pub fn parse_reasoning_response_for_prompt(
186    text: &str,
187    prompt_opened_thinking: bool,
188) -> ParsedReasoningResponse {
189    if prompt_opened_thinking {
190        parse_reasoning_response_started_in_think(text)
191    } else {
192        parse_reasoning_response(text)
193    }
194}
195
196#[cfg(test)]
197mod tests {
198    use super::*;
199
200    #[test]
201    fn parses_explicit_reasoning_block() {
202        let parsed = parse_reasoning_response("<think>reason</think>\nanswer");
203        assert_eq!(parsed.reasoning.as_deref(), Some("reason"));
204        assert_eq!(parsed.content, "answer");
205    }
206
207    #[test]
208    fn parses_prompt_opened_reasoning_block() {
209        let parsed = parse_reasoning_response_started_in_think("reason</think>\nanswer");
210        assert_eq!(parsed.reasoning.as_deref(), Some("reason"));
211        assert_eq!(parsed.content, "answer");
212    }
213
214    #[test]
215    fn prompt_opened_thinking_preserves_literal_tags_in_final_content() {
216        for content in [
217            r#"{"text":"<think>literal</think>"}"#,
218            r#"{"text":"an unpaired <think> marker"}"#,
219            "The tags `<think>` and `</think>` are ordinary text here.",
220        ] {
221            let raw = format!("reason</think>\r\n{content}");
222            let parsed = parse_model_reasoning_response(ModelOutputProtocol::Text, &raw, true)
223                .expect("declared text reasoning protocol");
224            assert_eq!(parsed.reasoning.as_deref(), Some("reason"), "{content}");
225            assert_eq!(parsed.content, content);
226        }
227    }
228
229    #[test]
230    fn empty_prompt_opened_thinking_preserves_final_literal_tags() {
231        let parsed = parse_reasoning_response_started_in_think(
232            "</think>\n{\"text\":\"<think>literal</think>\"}",
233        );
234        assert_eq!(parsed.reasoning, None);
235        assert_eq!(parsed.content, r#"{"text":"<think>literal</think>"}"#);
236    }
237
238    #[test]
239    fn prompt_opened_thinking_keeps_closed_boundary_across_stream_prefixes() {
240        let content = r#"{"text":"答案 🦀 <think>literal</think>"}"#;
241        for end in content
242            .char_indices()
243            .map(|(index, _)| index)
244            .chain(std::iter::once(content.len()))
245        {
246            let prefix = &content[..end];
247            let raw = format!("推理</think>\n{prefix}");
248            let parsed = parse_model_reasoning_response(ModelOutputProtocol::Text, &raw, true)
249                .expect("declared text reasoning protocol");
250            assert_eq!(parsed.reasoning.as_deref(), Some("推理"), "{prefix}");
251            assert_eq!(parsed.content, prefix);
252        }
253    }
254
255    #[test]
256    fn prompt_opened_thinking_accepts_repeated_opening_and_unfinished_reasoning() {
257        let parsed = parse_reasoning_response_started_in_think(
258            "<think>reason</think>\nThe literal <think> tag remains content.",
259        );
260        assert_eq!(parsed.reasoning.as_deref(), Some("reason"));
261        assert_eq!(parsed.content, "The literal <think> tag remains content.");
262
263        for raw in ["reason", "<think>reason"] {
264            let parsed = parse_reasoning_response_started_in_think(raw);
265            assert_eq!(parsed.reasoning.as_deref(), Some("reason"));
266            assert!(parsed.content.is_empty());
267        }
268    }
269
270    #[test]
271    fn preserves_plain_content() {
272        let parsed = parse_reasoning_response("answer");
273        assert_eq!(parsed.reasoning, None);
274        assert_eq!(parsed.content, "answer");
275    }
276
277    #[test]
278    fn detects_prompt_opened_reasoning() {
279        assert!(has_unclosed_thinking_block("assistant:<think>\n"));
280        assert!(!has_unclosed_thinking_block(
281            "assistant:<think>reason</think>\n"
282        ));
283    }
284}