Skip to main content

tokio_prompt_orchestrator/
chain_of_thought.rs

1//! Chain-of-thought prompting with structured scratchpad and step decomposition.
2//!
3//! Provides [`CotPromptBuilder`] to construct prompts that elicit step-by-step
4//! reasoning, and [`CotParser`] to extract structured [`ThoughtStep`]s from
5//! model responses.
6
7// ---------------------------------------------------------------------------
8// Data structures
9// ---------------------------------------------------------------------------
10
11/// A single reasoning step extracted from a chain-of-thought response.
12#[derive(Debug, Clone, PartialEq)]
13pub struct ThoughtStep {
14    /// 1-based ordinal index within the chain.
15    pub step_id: u32,
16    /// The reasoning / thinking performed at this step.
17    pub thought: String,
18    /// The conclusion drawn from this step's reasoning.
19    pub conclusion: String,
20    /// Heuristic confidence score in `[0.0, 1.0]`.
21    pub confidence: f64,
22    /// Whether this step indicated that a tool call is required.
23    pub requires_tool: bool,
24}
25
26/// A complete chain-of-thought response with all steps and a final answer.
27#[derive(Debug, Clone)]
28pub struct ChainOfThought {
29    /// Ordered reasoning steps.
30    pub steps: Vec<ThoughtStep>,
31    /// The final distilled answer extracted from the response.
32    pub final_answer: String,
33    /// Aggregate confidence across all steps.
34    pub total_confidence: f64,
35}
36
37// ---------------------------------------------------------------------------
38// CotStrategy
39// ---------------------------------------------------------------------------
40
41/// Strategy used to construct the chain-of-thought prompt.
42#[derive(Debug, Clone)]
43pub enum CotStrategy {
44    /// Zero-shot: ask the model to reason step-by-step without examples.
45    ZeroShot,
46    /// Few-shot: provide example (question, CoT answer) strings before the
47    /// target question.
48    FewShot(Vec<String>),
49    /// Tree-of-thought: instruct the model to explore branching solution paths.
50    TreeOfThought {
51        /// Number of parallel branches to explore.
52        branching_factor: u8,
53    },
54    /// Self-consistency: ask the model to produce multiple independent
55    /// reasoning chains and select the most consistent answer.
56    SelfConsistency {
57        /// Number of independent reasoning samples to request.
58        samples: u8,
59    },
60}
61
62// ---------------------------------------------------------------------------
63// CotPromptBuilder
64// ---------------------------------------------------------------------------
65
66/// Builds prompts that elicit chain-of-thought reasoning.
67#[derive(Debug, Default)]
68pub struct CotPromptBuilder;
69
70impl CotPromptBuilder {
71    /// Construct a new builder.
72    pub fn new() -> Self {
73        Self
74    }
75
76    /// Build a CoT prompt for `question` using the given `strategy`.
77    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    /// Return a set of (question, CoT answer) example pairs for common
130    /// reasoning types.
131    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    /// Build a tree-of-thought prompt that instructs the model to explore
162    /// `branching` parallel solution paths.
163    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// ---------------------------------------------------------------------------
181// CotParser
182// ---------------------------------------------------------------------------
183
184/// Parses chain-of-thought responses into structured [`ThoughtStep`]s.
185#[derive(Debug, Default)]
186pub struct CotParser;
187
188impl CotParser {
189    /// Create a new parser.
190    pub fn new() -> Self {
191        Self
192    }
193
194    /// Extract numbered steps from a response.
195    ///
196    /// Recognises lines of the form:
197    /// ```text
198    /// Step N: <thought> -> <conclusion>
199    /// Step N: <thought> → <conclusion>
200    /// ```
201    /// If no arrow separator is present the entire text is treated as the
202    /// thought and the conclusion is left empty.
203    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            // Match "Step N:" prefix (case-insensitive).
210            let lower = trimmed.to_lowercase();
211            if !lower.starts_with("step ") {
212                continue;
213            }
214            // Find the colon after the step number.
215            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            // Split on "->" or "→".
228            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                // Unicode arrow →
232                (
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    /// Extract the final answer from a response.
259    ///
260    /// Looks for lines beginning with `Therefore:`, `Answer:`, or
261    /// `Conclusion:` (case-insensitive) and returns the text that follows.
262    /// If none is found, returns an empty string.
263    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    /// Compute average confidence across all steps.
278    ///
279    /// Heuristic: longer thought text implies more thorough reasoning, which
280    /// maps to a higher confidence, capped at `1.0`.
281    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    /// Convenience: parse steps and final answer, then produce a full
290    /// [`ChainOfThought`].
291    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    // -----------------------------------------------------------------------
303    // Private helpers
304    // -----------------------------------------------------------------------
305
306    /// Heuristic confidence based on word count of the thought text.
307    ///
308    /// - 0 words  → 0.0
309    /// - 1 word   → 0.2
310    /// - 5 words  → 0.5
311    /// - 10+ words → approaches 1.0 (asymptotic)
312    fn heuristic_confidence(thought: &str) -> f64 {
313        let words = thought.split_whitespace().count();
314        if words == 0 {
315            return 0.0;
316        }
317        // Sigmoid-like: score = words / (words + 10) scaled to [0.2, 1.0].
318        let raw = words as f64 / (words as f64 + 10.0);
319        // Map [0, 1) → [0.2, 1.0).
320        0.2 + raw * 0.8
321    }
322}
323
324// ---------------------------------------------------------------------------
325// Tests
326// ---------------------------------------------------------------------------
327
328#[cfg(test)]
329mod tests {
330    use super::*;
331
332    // --- CotPromptBuilder ---
333
334    #[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        // branching_factor of 0 should be clamped to 2.
365        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    // --- CotParser ---
392
393    #[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        // Very long thought should not exceed 1.0.
499        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}