ferrum_types/
reasoning.rs1use 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
18pub 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 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
53pub 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
73pub 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
101pub 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}