tokio_prompt_orchestrator/
chain_of_thought.rs1#[derive(Debug, Clone, PartialEq)]
13pub struct ThoughtStep {
14 pub step_id: u32,
16 pub thought: String,
18 pub conclusion: String,
20 pub confidence: f64,
22 pub requires_tool: bool,
24}
25
26#[derive(Debug, Clone)]
28pub struct ChainOfThought {
29 pub steps: Vec<ThoughtStep>,
31 pub final_answer: String,
33 pub total_confidence: f64,
35}
36
37#[derive(Debug, Clone)]
43pub enum CotStrategy {
44 ZeroShot,
46 FewShot(Vec<String>),
49 TreeOfThought {
51 branching_factor: u8,
53 },
54 SelfConsistency {
57 samples: u8,
59 },
60}
61
62#[derive(Debug, Default)]
68pub struct CotPromptBuilder;
69
70impl CotPromptBuilder {
71 pub fn new() -> Self {
73 Self
74 }
75
76 pub fn build_cot_prompt(&self, question: &str, strategy: &CotStrategy) -> String {
78 match strategy {
79 CotStrategy::ZeroShot => {
80 format!(
81 "Answer the following question by thinking step-by-step.\n\
82 For each step, write:\n\
83 Step N: <your reasoning> -> <conclusion>\n\n\
84 After all steps, write:\n\
85 Therefore: <final answer>\n\n\
86 Question: {question}"
87 )
88 }
89
90 CotStrategy::FewShot(examples) => {
91 let mut prompt = String::from(
92 "Answer the following question by thinking step-by-step, \
93 following the examples below.\n\n",
94 );
95 for (i, ex) in examples.iter().enumerate() {
96 prompt.push_str(&format!("--- Example {} ---\n{}\n\n", i + 1, ex));
97 }
98 prompt.push_str("--- Your turn ---\n");
99 prompt.push_str(&format!(
100 "For each step write:\n\
101 Step N: <your reasoning> -> <conclusion>\n\n\
102 After all steps write:\n\
103 Therefore: <final answer>\n\n\
104 Question: {question}"
105 ));
106 prompt
107 }
108
109 CotStrategy::TreeOfThought { branching_factor } => {
110 self.tree_of_thought_prompt(question, *branching_factor)
111 }
112
113 CotStrategy::SelfConsistency { samples } => {
114 format!(
115 "Produce {samples} independent reasoning chains for the question below.\n\
116 Number each chain as \"Chain 1:\", \"Chain 2:\", etc.\n\
117 Within each chain, write each step as:\n\
118 Step N: <reasoning> -> <conclusion>\n\
119 End each chain with:\n\
120 Therefore: <answer>\n\n\
121 After all chains, identify the most consistent answer and write:\n\
122 Answer: <consensus answer>\n\n\
123 Question: {question}"
124 )
125 }
126 }
127 }
128
129 pub fn few_shot_examples(&self) -> Vec<(String, String)> {
132 vec![
133 (
134 "If a train travels 60 km/h for 2.5 hours, how far does it go?".to_string(),
135 "Step 1: Identify the formula. Distance = speed × time. -> Formula is d = v × t.\n\
136 Step 2: Substitute values. d = 60 km/h × 2.5 h. -> d = 150 km.\n\
137 Step 3: Check units. km/h × h = km. -> Units are consistent.\n\
138 Therefore: The train travels 150 km."
139 .to_string(),
140 ),
141 (
142 "Is 97 a prime number?".to_string(),
143 "Step 1: Check divisibility by 2. 97 is odd. -> Not divisible by 2.\n\
144 Step 2: Check divisibility by 3. 9+7=16, not divisible by 3. -> Not divisible by 3.\n\
145 Step 3: Check up to sqrt(97) ≈ 9.8. Test 5, 7. Neither divides 97. -> No factors found.\n\
146 Therefore: 97 is a prime number."
147 .to_string(),
148 ),
149 (
150 "What is the capital of Australia?".to_string(),
151 "Step 1: Recall common misconception. Many assume Sydney is the capital. -> \
152 Sydney is the largest city but not the capital.\n\
153 Step 2: Recall founding history. Canberra was purpose-built as a compromise \
154 between Sydney and Melbourne. -> Canberra is the capital.\n\
155 Therefore: The capital of Australia is Canberra."
156 .to_string(),
157 ),
158 ]
159 }
160
161 pub fn tree_of_thought_prompt(&self, question: &str, branching: u8) -> String {
164 let branching = branching.max(2);
165 format!(
166 "Use the Tree-of-Thought method to answer the question below.\n\
167 Explore exactly {branching} distinct solution approaches in parallel.\n\n\
168 For each approach, label it \"Branch N:\" (N = 1..{branching}) and \
169 within each branch write each step as:\n\
170 Step N: <reasoning> -> <conclusion>\n\n\
171 After exploring all branches, evaluate them:\n\
172 Evaluation: <which branch is strongest and why>\n\n\
173 Then write the final answer:\n\
174 Therefore: <final answer>\n\n\
175 Question: {question}"
176 )
177 }
178}
179
180#[derive(Debug, Default)]
186pub struct CotParser;
187
188impl CotParser {
189 pub fn new() -> Self {
191 Self
192 }
193
194 pub fn parse_steps(&self, response: &str) -> Vec<ThoughtStep> {
204 let mut steps = Vec::new();
205
206 for line in response.lines() {
207 let trimmed = line.trim();
208
209 let lower = trimmed.to_lowercase();
211 if !lower.starts_with("step ") {
212 continue;
213 }
214 let colon_pos = match trimmed.find(':') {
216 Some(p) => p,
217 None => continue,
218 };
219 let step_num_str = trimmed[5..colon_pos].trim();
220 let step_id: u32 = match step_num_str.parse() {
221 Ok(n) => n,
222 Err(_) => continue,
223 };
224
225 let body = trimmed[colon_pos + 1..].trim();
226
227 let (thought, conclusion) = if let Some(pos) = body.find("->") {
229 (body[..pos].trim().to_string(), body[pos + 2..].trim().to_string())
230 } else if let Some(pos) = body.find('\u{2192}') {
231 (
233 body[..pos].trim().to_string(),
234 body[pos + '\u{2192}'.len_utf8()..].trim().to_string(),
235 )
236 } else {
237 (body.to_string(), String::new())
238 };
239
240 let requires_tool = thought.to_lowercase().contains("tool")
241 || thought.to_lowercase().contains("search")
242 || thought.to_lowercase().contains("lookup");
243
244 let confidence = Self::heuristic_confidence(&thought);
245
246 steps.push(ThoughtStep {
247 step_id,
248 thought,
249 conclusion,
250 confidence,
251 requires_tool,
252 });
253 }
254
255 steps
256 }
257
258 pub fn extract_final_answer(&self, response: &str) -> String {
264 let markers = ["therefore:", "answer:", "conclusion:"];
265 for line in response.lines() {
266 let lower = line.trim().to_lowercase();
267 for marker in &markers {
268 if lower.starts_with(marker) {
269 let after = line.trim()[marker.len()..].trim();
270 return after.to_string();
271 }
272 }
273 }
274 String::new()
275 }
276
277 pub fn compute_confidence(&self, steps: &[ThoughtStep]) -> f64 {
282 if steps.is_empty() {
283 return 0.0;
284 }
285 let sum: f64 = steps.iter().map(|s| s.confidence).sum();
286 (sum / steps.len() as f64).min(1.0)
287 }
288
289 pub fn parse(&self, response: &str) -> ChainOfThought {
292 let steps = self.parse_steps(response);
293 let final_answer = self.extract_final_answer(response);
294 let total_confidence = self.compute_confidence(&steps);
295 ChainOfThought {
296 steps,
297 final_answer,
298 total_confidence,
299 }
300 }
301
302 fn heuristic_confidence(thought: &str) -> f64 {
313 let words = thought.split_whitespace().count();
314 if words == 0 {
315 return 0.0;
316 }
317 let raw = words as f64 / (words as f64 + 10.0);
319 0.2 + raw * 0.8
321 }
322}
323
324#[cfg(test)]
329mod tests {
330 use super::*;
331
332 #[test]
335 fn test_zero_shot_contains_question() {
336 let builder = CotPromptBuilder::new();
337 let prompt = builder.build_cot_prompt("Why is the sky blue?", &CotStrategy::ZeroShot);
338 assert!(prompt.contains("Why is the sky blue?"));
339 assert!(prompt.contains("Step N:"));
340 assert!(prompt.contains("Therefore:"));
341 }
342
343 #[test]
344 fn test_few_shot_contains_examples() {
345 let builder = CotPromptBuilder::new();
346 let examples = vec!["Q: 2+2? A: Step 1: 2+2=4 -> 4.\nTherefore: 4".to_string()];
347 let prompt =
348 builder.build_cot_prompt("What is 3+3?", &CotStrategy::FewShot(examples));
349 assert!(prompt.contains("Example 1"));
350 assert!(prompt.contains("What is 3+3?"));
351 }
352
353 #[test]
354 fn test_tree_of_thought_prompt_has_branches() {
355 let builder = CotPromptBuilder::new();
356 let prompt = builder.tree_of_thought_prompt("Solve X", 3);
357 assert!(prompt.contains("Branch N:"));
358 assert!(prompt.contains("3 distinct solution"));
359 }
360
361 #[test]
362 fn test_tree_of_thought_minimum_branching() {
363 let builder = CotPromptBuilder::new();
364 let prompt = builder.tree_of_thought_prompt("Q", 0);
366 assert!(prompt.contains("2 distinct solution"));
367 }
368
369 #[test]
370 fn test_self_consistency_prompt() {
371 let builder = CotPromptBuilder::new();
372 let prompt = builder.build_cot_prompt(
373 "Is 11 prime?",
374 &CotStrategy::SelfConsistency { samples: 3 },
375 );
376 assert!(prompt.contains("3 independent reasoning chains"));
377 assert!(prompt.contains("Chain 1:"));
378 }
379
380 #[test]
381 fn test_few_shot_examples_non_empty() {
382 let builder = CotPromptBuilder::new();
383 let examples = builder.few_shot_examples();
384 assert!(!examples.is_empty());
385 for (q, a) in &examples {
386 assert!(!q.is_empty());
387 assert!(!a.is_empty());
388 }
389 }
390
391 #[test]
394 fn test_parse_steps_arrow() {
395 let parser = CotParser::new();
396 let response = "Step 1: Identify values -> values are 2 and 3.\n\
397 Step 2: Add them -> result is 5.\n\
398 Therefore: 5";
399 let steps = parser.parse_steps(response);
400 assert_eq!(steps.len(), 2);
401 assert_eq!(steps[0].step_id, 1);
402 assert_eq!(steps[0].conclusion, "values are 2 and 3.");
403 assert_eq!(steps[1].step_id, 2);
404 assert_eq!(steps[1].conclusion, "result is 5.");
405 }
406
407 #[test]
408 fn test_parse_steps_no_arrow() {
409 let parser = CotParser::new();
410 let response = "Step 1: Just thinking here.";
411 let steps = parser.parse_steps(response);
412 assert_eq!(steps.len(), 1);
413 assert_eq!(steps[0].thought, "Just thinking here.");
414 assert!(steps[0].conclusion.is_empty());
415 }
416
417 #[test]
418 fn test_parse_steps_requires_tool() {
419 let parser = CotParser::new();
420 let response = "Step 1: Need to tool-call the API -> result pending.";
421 let steps = parser.parse_steps(response);
422 assert_eq!(steps.len(), 1);
423 assert!(steps[0].requires_tool);
424 }
425
426 #[test]
427 fn test_extract_final_answer_therefore() {
428 let parser = CotParser::new();
429 let resp = "Step 1: x -> y.\nTherefore: The answer is 42.";
430 assert_eq!(parser.extract_final_answer(resp), "The answer is 42.");
431 }
432
433 #[test]
434 fn test_extract_final_answer_answer_marker() {
435 let parser = CotParser::new();
436 let resp = "Step 1: x -> y.\nAnswer: Yes, it is.";
437 assert_eq!(parser.extract_final_answer(resp), "Yes, it is.");
438 }
439
440 #[test]
441 fn test_extract_final_answer_conclusion_marker() {
442 let parser = CotParser::new();
443 let resp = "Step 1: x -> y.\nConclusion: Definitely true.";
444 assert_eq!(parser.extract_final_answer(resp), "Definitely true.");
445 }
446
447 #[test]
448 fn test_extract_final_answer_missing() {
449 let parser = CotParser::new();
450 let resp = "Step 1: x -> y.";
451 assert!(parser.extract_final_answer(resp).is_empty());
452 }
453
454 #[test]
455 fn test_compute_confidence_empty() {
456 let parser = CotParser::new();
457 assert_eq!(parser.compute_confidence(&[]), 0.0);
458 }
459
460 #[test]
461 fn test_compute_confidence_increases_with_length() {
462 let parser = CotParser::new();
463 let short = ThoughtStep {
464 step_id: 1,
465 thought: "hi".to_string(),
466 conclusion: String::new(),
467 confidence: CotParser::heuristic_confidence("hi"),
468 requires_tool: false,
469 };
470 let long = ThoughtStep {
471 step_id: 2,
472 thought: "a b c d e f g h i j k l m n o p q r s t u v w x y z".to_string(),
473 conclusion: String::new(),
474 confidence: CotParser::heuristic_confidence(
475 "a b c d e f g h i j k l m n o p q r s t u v w x y z",
476 ),
477 requires_tool: false,
478 };
479 assert!(long.confidence > short.confidence);
480 assert!(long.confidence <= 1.0);
481 }
482
483 #[test]
484 fn test_parse_full_chain() {
485 let parser = CotParser::new();
486 let response = "Step 1: Consider the problem carefully -> problem is addition.\n\
487 Step 2: Compute 3 + 4 -> result is 7.\n\
488 Therefore: 7";
489 let cot = parser.parse(response);
490 assert_eq!(cot.steps.len(), 2);
491 assert_eq!(cot.final_answer, "7");
492 assert!(cot.total_confidence > 0.0);
493 assert!(cot.total_confidence <= 1.0);
494 }
495
496 #[test]
497 fn test_confidence_capped_at_one() {
498 let very_long = "word ".repeat(100);
500 let conf = CotParser::heuristic_confidence(&very_long);
501 assert!(conf <= 1.0);
502 assert!(conf > 0.0);
503 }
504}