1use crate::{FerrumError, ModelOutputProtocol, Result};
4
5mod gemma;
6
7#[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
36pub 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 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
71pub 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
91pub 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
119pub fn parse_reasoning_response_started_in_think(text: &str) -> ParsedReasoningResponse {
122 let end = text.find(THINK_END_TAG);
123 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}