Skip to main content

ferrum_types/
reasoning.rs

1//! Shared parsing for model outputs that carry `<think>` reasoning blocks.
2
3pub const THINK_START_TAG: &str = "<think>";
4pub const THINK_END_TAG: &str = "</think>";
5
6#[derive(Debug, Clone, PartialEq, Eq)]
7pub struct ParsedReasoningResponse {
8    pub content: String,
9    pub reasoning: Option<String>,
10}
11
12pub fn has_unclosed_thinking_block(prompt: &str) -> bool {
13    match (prompt.rfind(THINK_START_TAG), prompt.rfind(THINK_END_TAG)) {
14        (Some(start), Some(end)) => start > end,
15        (Some(_), None) => true,
16        _ => false,
17    }
18}
19
20/// Parse generated text when the rendered prompt already opened `<think>`.
21/// Some reasoning templates emit only the closing tag in generated text.
22pub fn parse_reasoning_response_started_in_think(text: &str) -> ParsedReasoningResponse {
23    if text.contains(THINK_START_TAG) {
24        return parse_reasoning_response(text);
25    }
26    let Some(end) = text.find(THINK_END_TAG) else {
27        return ParsedReasoningResponse {
28            content: String::new(),
29            reasoning: (!text.is_empty()).then(|| text.to_string()),
30        };
31    };
32    let reasoning = text[..end].to_string();
33    let content = text[end + THINK_END_TAG.len()..]
34        .trim_start_matches(['\r', '\n'])
35        .to_string();
36    ParsedReasoningResponse {
37        content,
38        reasoning: (!reasoning.is_empty()).then_some(reasoning),
39    }
40}
41
42pub fn parse_reasoning_response(text: &str) -> ParsedReasoningResponse {
43    let Some(start) = text.find(THINK_START_TAG) else {
44        if let Some(end) = text.find(THINK_END_TAG) {
45            let reasoning = text[..end].to_string();
46            let content = text[end + THINK_END_TAG.len()..]
47                .trim_start_matches(['\r', '\n'])
48                .to_string();
49            return ParsedReasoningResponse {
50                content,
51                reasoning: (!reasoning.is_empty()).then_some(reasoning),
52            };
53        }
54        return ParsedReasoningResponse {
55            content: text.to_string(),
56            reasoning: None,
57        };
58    };
59
60    let before = &text[..start];
61    let after_start = &text[start + THINK_START_TAG.len()..];
62    let Some(end) = after_start.find(THINK_END_TAG) else {
63        return ParsedReasoningResponse {
64            content: before.to_string(),
65            reasoning: Some(after_start.to_string()),
66        };
67    };
68
69    let reasoning = after_start[..end].to_string();
70    let after_end = &after_start[end + THINK_END_TAG.len()..];
71    let mut content = String::with_capacity(before.len() + after_end.len());
72    content.push_str(before);
73    content.push_str(after_end.trim_start_matches(['\r', '\n']));
74    ParsedReasoningResponse {
75        content,
76        reasoning: (!reasoning.is_empty()).then_some(reasoning),
77    }
78}
79
80pub fn parse_reasoning_response_for_prompt(
81    text: &str,
82    prompt_opened_thinking: bool,
83) -> ParsedReasoningResponse {
84    if prompt_opened_thinking {
85        parse_reasoning_response_started_in_think(text)
86    } else {
87        parse_reasoning_response(text)
88    }
89}
90
91#[cfg(test)]
92mod tests {
93    use super::*;
94
95    #[test]
96    fn parses_explicit_reasoning_block() {
97        let parsed = parse_reasoning_response("<think>reason</think>\nanswer");
98        assert_eq!(parsed.reasoning.as_deref(), Some("reason"));
99        assert_eq!(parsed.content, "answer");
100    }
101
102    #[test]
103    fn parses_prompt_opened_reasoning_block() {
104        let parsed = parse_reasoning_response_started_in_think("reason</think>\nanswer");
105        assert_eq!(parsed.reasoning.as_deref(), Some("reason"));
106        assert_eq!(parsed.content, "answer");
107    }
108
109    #[test]
110    fn preserves_plain_content() {
111        let parsed = parse_reasoning_response("answer");
112        assert_eq!(parsed.reasoning, None);
113        assert_eq!(parsed.content, "answer");
114    }
115
116    #[test]
117    fn detects_prompt_opened_reasoning() {
118        assert!(has_unclosed_thinking_block("assistant:<think>\n"));
119        assert!(!has_unclosed_thinking_block(
120            "assistant:<think>reason</think>\n"
121        ));
122    }
123}