Skip to main content

tokio_prompt_orchestrator/
cascade.rs

1//! Cascading multi-turn inference engine.
2//!
3//! Enables multi-turn reasoning loops where a model's output can trigger
4//! additional pipeline passes — tool calls, function execution, chain-of-thought
5//! continuation — until a termination condition is met.
6//!
7//! ## Architecture
8//!
9//! ```text
10//! PromptRequest
11//!       │
12//!       ▼
13//! ┌─────────────────────────────────────────┐
14//! │          CascadeEngine                  │
15//! │                                         │
16//! │  Turn 1: infer → parse tools            │
17//! │  Turn 2: execute tools → infer again    │
18//! │  Turn N: model returns DONE / no tools  │
19//! └─────────────────────────────────────────┘
20//!       │
21//!       ▼
22//! CascadeResult (all turns + final answer)
23//! ```
24//!
25//! ## Termination Conditions
26//!
27//! The loop exits when ANY of the following is true:
28//! - The model output contains no tool calls (`ToolCallParser` finds nothing)
29//! - The output text contains the DONE sentinel string
30//! - `max_turns` is reached (default: 10)
31//! - A pipeline error occurs (propagated to caller)
32//!
33//! ## Tool Call Format
34//!
35//! The engine recognises tool calls in this JSON block format:
36//!
37//! ```text
38//! <tool_call>
39//! {"name": "search", "arguments": {"query": "rust tokio"}}
40//! </tool_call>
41//! ```
42//!
43//! Custom parsers can be registered via [`CascadeEngine::with_tool_parser`].
44
45use crate::{OrchestratorError, SessionId};
46use async_trait::async_trait;
47use serde::{Deserialize, Serialize};
48use std::collections::HashMap;
49use std::sync::Arc;
50use std::time::{Duration, Instant};
51use tokio::sync::Mutex;
52use tracing::{debug, info, warn};
53
54/// Default maximum number of cascade turns before halting.
55pub const DEFAULT_MAX_TURNS: usize = 10;
56
57/// Sentinel string that signals the model is done cascading.
58pub const DONE_SENTINEL: &str = "[DONE]";
59
60/// A single parsed tool call extracted from a model response.
61#[derive(Debug, Clone, Serialize, Deserialize)]
62pub struct ToolCall {
63    /// The name of the tool to invoke.
64    pub name: String,
65    /// Arguments as a JSON-compatible map.
66    pub arguments: HashMap<String, serde_json::Value>,
67    /// Position in the response text where this call was found.
68    pub position: usize,
69}
70
71/// The result of a single cascade turn.
72#[derive(Debug, Clone, Serialize, Deserialize)]
73pub struct CascadeTurn {
74    /// Turn index (0-based).
75    pub turn: usize,
76    /// The prompt sent in this turn.
77    pub prompt: String,
78    /// The model's raw response.
79    pub response: String,
80    /// Tool calls parsed from the response, if any.
81    pub tool_calls: Vec<ToolCall>,
82    /// Tool execution results injected into the next turn's context.
83    pub tool_results: Vec<ToolResult>,
84    /// Elapsed time for this turn.
85    pub elapsed_ms: u64,
86}
87
88/// The result of a tool execution.
89#[derive(Debug, Clone, Serialize, Deserialize)]
90pub struct ToolResult {
91    /// Name of the tool that was called.
92    pub tool_name: String,
93    /// JSON-serializable result.
94    pub result: serde_json::Value,
95    /// Whether execution succeeded.
96    pub success: bool,
97    /// Error message if `success` is false.
98    pub error: Option<String>,
99}
100
101/// Final output of a complete cascade run.
102#[derive(Debug, Clone, Serialize, Deserialize)]
103pub struct CascadeResult {
104    /// Session this cascade belongs to.
105    pub session_id: SessionId,
106    /// All turns executed, in order.
107    pub turns: Vec<CascadeTurn>,
108    /// The final model response text (last turn's response).
109    pub final_answer: String,
110    /// Why the cascade stopped.
111    pub termination_reason: TerminationReason,
112    /// Total wall-clock time across all turns.
113    pub total_elapsed_ms: u64,
114    /// Number of tool calls executed.
115    pub total_tool_calls: usize,
116}
117
118/// Why a cascade loop exited.
119#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
120pub enum TerminationReason {
121    /// Model produced no tool calls — natural completion.
122    NoToolCalls,
123    /// Model emitted the DONE sentinel explicitly.
124    DoneSentinel,
125    /// Maximum turn limit was reached.
126    MaxTurnsReached,
127    /// A pipeline error aborted the cascade.
128    Error(String),
129}
130
131/// Trait for parsing tool calls out of a raw model response.
132///
133/// The default implementation looks for `<tool_call>…</tool_call>` JSON blocks.
134/// Register custom parsers via [`CascadeEngine::with_tool_parser`].
135#[async_trait]
136pub trait ToolCallParser: Send + Sync {
137    /// Parse zero or more tool calls from a model response string.
138    async fn parse(&self, response: &str) -> Vec<ToolCall>;
139}
140
141/// Default parser: finds `<tool_call>…</tool_call>` JSON blocks.
142pub struct XmlStyleToolParser;
143
144#[async_trait]
145impl ToolCallParser for XmlStyleToolParser {
146    async fn parse(&self, response: &str) -> Vec<ToolCall> {
147        let mut calls = Vec::new();
148        let mut search_from = 0usize;
149
150        while let Some(start_rel) = response[search_from..].find("<tool_call>") {
151            let abs_start = search_from + start_rel;
152            let content_start = abs_start + "<tool_call>".len();
153
154            if let Some(end_rel) = response[content_start..].find("</tool_call>") {
155                let abs_end = content_start + end_rel;
156                let json_str = response[content_start..abs_end].trim();
157
158                match serde_json::from_str::<HashMap<String, serde_json::Value>>(json_str) {
159                    Ok(mut map) => {
160                        let name = map
161                            .remove("name")
162                            .and_then(|v| v.as_str().map(String::from))
163                            .unwrap_or_else(|| "unknown".to_string());
164                        let arguments = map
165                            .remove("arguments")
166                            .and_then(|v| v.as_object().cloned())
167                            .map(|o| o.into_iter().collect())
168                            .unwrap_or_default();
169                        calls.push(ToolCall {
170                            name,
171                            arguments,
172                            position: abs_start,
173                        });
174                    }
175                    Err(e) => {
176                        warn!(error = %e, "failed to parse tool_call JSON block");
177                    }
178                }
179                search_from = abs_end + "</tool_call>".len();
180            } else {
181                break;
182            }
183        }
184
185        calls
186    }
187}
188
189/// Trait for executing tool calls and returning results.
190///
191/// Implement this to wire real tools (web search, code execution, DB queries)
192/// into the cascade loop. The default [`NoopToolExecutor`] returns a stub result.
193#[async_trait]
194pub trait ToolExecutor: Send + Sync {
195    /// Execute a single tool call and return its result.
196    async fn execute(&self, call: &ToolCall) -> ToolResult;
197}
198
199/// Stub executor that returns a placeholder result without doing real work.
200///
201/// Useful for testing the cascade loop logic without external dependencies.
202pub struct NoopToolExecutor;
203
204#[async_trait]
205impl ToolExecutor for NoopToolExecutor {
206    async fn execute(&self, call: &ToolCall) -> ToolResult {
207        debug!(tool = %call.name, "noop executor called");
208        ToolResult {
209            tool_name: call.name.clone(),
210            result: serde_json::json!({
211                "status": "ok",
212                "note": "noop executor — register a real ToolExecutor via CascadeEngine::with_executor"
213            }),
214            success: true,
215            error: None,
216        }
217    }
218}
219
220/// Builds the next-turn prompt by injecting tool results into the conversation.
221fn build_next_prompt(
222    original_prompt: &str,
223    turns: &[CascadeTurn],
224    pending_results: &[ToolResult],
225) -> String {
226    let mut parts = Vec::new();
227
228    // Original user prompt
229    parts.push(format!("User: {original_prompt}"));
230
231    // Prior turns
232    for turn in turns {
233        parts.push(format!("Assistant: {}", turn.response));
234        for result in &turn.tool_results {
235            let result_str = serde_json::to_string(&result.result).unwrap_or_default();
236            parts.push(format!(
237                "Tool[{}]: {}",
238                result.tool_name, result_str
239            ));
240        }
241    }
242
243    // Pending results from last turn's tool calls
244    for result in pending_results {
245        let result_str = serde_json::to_string(&result.result).unwrap_or_default();
246        parts.push(format!(
247            "Tool[{}]: {}",
248            result.tool_name, result_str
249        ));
250    }
251
252    parts.push("Assistant:".to_string());
253    parts.join("\n\n")
254}
255
256/// Configuration for a cascade run.
257#[derive(Debug, Clone)]
258pub struct CascadeConfig {
259    /// Maximum number of turns before forcibly stopping (default: 10).
260    pub max_turns: usize,
261    /// Maximum wall-clock duration for the entire cascade.
262    pub timeout: Duration,
263    /// System prompt injected at the start of every turn.
264    pub system_prompt: Option<String>,
265}
266
267impl Default for CascadeConfig {
268    fn default() -> Self {
269        Self {
270            max_turns: DEFAULT_MAX_TURNS,
271            timeout: Duration::from_secs(300),
272            system_prompt: None,
273        }
274    }
275}
276
277/// Infer function type: takes a prompt string, returns the model response.
278///
279/// In production this is wired to [`spawn_pipeline`](crate::spawn_pipeline) or
280/// any [`ModelWorker`](crate::ModelWorker). In tests it can be a simple closure.
281pub type InferFn = Arc<
282    dyn Fn(String) -> std::pin::Pin<
283            Box<dyn std::future::Future<Output = Result<String, OrchestratorError>> + Send>,
284        > + Send
285        + Sync,
286>;
287
288/// Multi-turn cascading inference engine.
289///
290/// Drives a model through a tool-call loop until a termination condition is met.
291///
292/// # Example
293///
294/// ```no_run
295/// use std::sync::Arc;
296/// use tokio_prompt_orchestrator::cascade::{CascadeEngine, CascadeConfig, NoopToolExecutor};
297///
298/// # async fn example() {
299/// let engine = CascadeEngine::new(
300///     Arc::new(|prompt: String| Box::pin(async move {
301///         Ok(format!("Echo: {prompt}"))
302///     })),
303///     Arc::new(NoopToolExecutor),
304/// );
305///
306/// let result = engine
307///     .run("Summarise the Rust docs", &Default::default())
308///     .await
309///     .unwrap();
310///
311/// println!("Final answer: {}", result.final_answer);
312/// println!("Turns taken: {}", result.turns.len());
313/// # }
314/// ```
315pub struct CascadeEngine {
316    infer_fn: InferFn,
317    executor: Arc<dyn ToolExecutor>,
318    parser: Arc<dyn ToolCallParser>,
319}
320
321impl CascadeEngine {
322    /// Create a new engine with the given infer function and tool executor.
323    pub fn new(infer_fn: InferFn, executor: Arc<dyn ToolExecutor>) -> Self {
324        Self {
325            infer_fn,
326            executor,
327            parser: Arc::new(XmlStyleToolParser),
328        }
329    }
330
331    /// Override the default `<tool_call>` XML parser with a custom one.
332    pub fn with_tool_parser(mut self, parser: Arc<dyn ToolCallParser>) -> Self {
333        self.parser = parser;
334        self
335    }
336
337    /// Run the cascade for the given prompt, returning all turns and the final answer.
338    ///
339    /// # Errors
340    ///
341    /// Returns an error only if the first inference call fails. Subsequent
342    /// failures are recorded in [`CascadeResult::termination_reason`].
343    pub async fn run(
344        &self,
345        prompt: &str,
346        config: &CascadeConfig,
347    ) -> Result<CascadeResult, OrchestratorError> {
348        let session = SessionId::new(uuid::Uuid::new_v4().to_string());
349        let start = Instant::now();
350        let deadline = start + config.timeout;
351
352        let mut turns: Vec<CascadeTurn> = Vec::new();
353        let mut current_prompt = if let Some(ref sys) = config.system_prompt {
354            format!("{sys}\n\n{prompt}")
355        } else {
356            prompt.to_string()
357        };
358
359        info!(
360            session_id = %session.as_str(),
361            max_turns = config.max_turns,
362            "cascade started"
363        );
364
365        for turn_idx in 0..config.max_turns {
366            // Hard timeout guard
367            if Instant::now() >= deadline {
368                warn!(turn = turn_idx, "cascade hit wall-clock timeout");
369                let total_ms = start.elapsed().as_millis() as u64;
370                let total_calls: usize = turns.iter().map(|t| t.tool_calls.len()).sum();
371                return Ok(CascadeResult {
372                    session_id: session,
373                    final_answer: turns.last().map(|t| t.response.clone()).unwrap_or_default(),
374                    termination_reason: TerminationReason::MaxTurnsReached,
375                    total_elapsed_ms: total_ms,
376                    total_tool_calls: total_calls,
377                    turns,
378                });
379            }
380
381            let turn_start = Instant::now();
382
383            // Infer
384            let response = (self.infer_fn)(current_prompt.clone()).await.map_err(|e| {
385                if turn_idx == 0 {
386                    // First turn failure is a hard error
387                    e
388                } else {
389                    // Subsequent turn failures are soft — we stop the cascade
390                    OrchestratorError::Other(format!("cascade turn {turn_idx} failed: {e}"))
391                }
392            })?;
393
394            let elapsed_ms = turn_start.elapsed().as_millis() as u64;
395
396            // Check DONE sentinel
397            if response.contains(DONE_SENTINEL) {
398                let resp_clean = response.replace(DONE_SENTINEL, "").trim().to_string();
399                let total_ms = start.elapsed().as_millis() as u64;
400                let total_calls: usize = turns.iter().map(|t| t.tool_calls.len()).sum();
401
402                turns.push(CascadeTurn {
403                    turn: turn_idx,
404                    prompt: current_prompt,
405                    response: resp_clean.clone(),
406                    tool_calls: vec![],
407                    tool_results: vec![],
408                    elapsed_ms,
409                });
410
411                info!(
412                    session_id = %session.as_str(),
413                    turns = turn_idx + 1,
414                    reason = "done_sentinel",
415                    "cascade complete"
416                );
417
418                return Ok(CascadeResult {
419                    session_id: session,
420                    final_answer: resp_clean,
421                    termination_reason: TerminationReason::DoneSentinel,
422                    total_elapsed_ms: total_ms,
423                    total_tool_calls: total_calls,
424                    turns,
425                });
426            }
427
428            // Parse tool calls
429            let tool_calls = self.parser.parse(&response).await;
430
431            if tool_calls.is_empty() {
432                // Natural completion — no tools requested
433                let total_ms = start.elapsed().as_millis() as u64;
434                let total_calls: usize = turns.iter().map(|t| t.tool_calls.len()).sum();
435
436                turns.push(CascadeTurn {
437                    turn: turn_idx,
438                    prompt: current_prompt,
439                    response: response.clone(),
440                    tool_calls: vec![],
441                    tool_results: vec![],
442                    elapsed_ms,
443                });
444
445                info!(
446                    session_id = %session.as_str(),
447                    turns = turn_idx + 1,
448                    reason = "no_tool_calls",
449                    "cascade complete"
450                );
451
452                return Ok(CascadeResult {
453                    session_id: session,
454                    final_answer: response,
455                    termination_reason: TerminationReason::NoToolCalls,
456                    total_elapsed_ms: total_ms,
457                    total_tool_calls: total_calls,
458                    turns,
459                });
460            }
461
462            // Execute all tool calls concurrently
463            let executor = Arc::clone(&self.executor);
464            let calls_for_exec = tool_calls.clone();
465            let mut exec_futures = Vec::with_capacity(calls_for_exec.len());
466            for call in &calls_for_exec {
467                let exec = Arc::clone(&executor);
468                let c = call.clone();
469                exec_futures.push(async move { exec.execute(&c).await });
470            }
471
472            let tool_results = futures::future::join_all(exec_futures).await;
473
474            debug!(
475                turn = turn_idx,
476                tool_count = tool_results.len(),
477                "tool calls executed"
478            );
479
480            // Build next prompt
481            let next_prompt = build_next_prompt(
482                prompt,
483                &turns,
484                &tool_results,
485            );
486
487            turns.push(CascadeTurn {
488                turn: turn_idx,
489                prompt: current_prompt,
490                response,
491                tool_calls,
492                tool_results,
493                elapsed_ms,
494            });
495
496            current_prompt = next_prompt;
497        }
498
499        // Max turns reached
500        let total_ms = start.elapsed().as_millis() as u64;
501        let total_calls: usize = turns.iter().map(|t| t.tool_calls.len()).sum();
502        let final_answer = turns.last().map(|t| t.response.clone()).unwrap_or_default();
503
504        warn!(
505            session_id = %session.as_str(),
506            max_turns = config.max_turns,
507            "cascade reached max turn limit"
508        );
509
510        Ok(CascadeResult {
511            session_id: session,
512            final_answer,
513            termination_reason: TerminationReason::MaxTurnsReached,
514            total_elapsed_ms: total_ms,
515            total_tool_calls: total_calls,
516            turns,
517        })
518    }
519}
520
521/// A thread-safe store of active cascade sessions for introspection.
522///
523/// Useful for monitoring dashboards and the web API to surface
524/// how many cascades are running and their turn counts.
525pub struct CascadeMonitor {
526    active: Arc<Mutex<HashMap<String, usize>>>,
527}
528
529impl Default for CascadeMonitor {
530    fn default() -> Self {
531        Self {
532            active: Arc::new(Mutex::new(HashMap::new())),
533        }
534    }
535}
536
537impl CascadeMonitor {
538    /// Record that a cascade has started.
539    pub async fn on_start(&self, session_id: &str) {
540        self.active.lock().await.insert(session_id.to_string(), 0);
541    }
542
543    /// Record that a cascade has completed another turn.
544    pub async fn on_turn(&self, session_id: &str, turn: usize) {
545        if let Some(entry) = self.active.lock().await.get_mut(session_id) {
546            *entry = turn;
547        }
548    }
549
550    /// Record that a cascade has finished.
551    pub async fn on_complete(&self, session_id: &str) {
552        self.active.lock().await.remove(session_id);
553    }
554
555    /// Return the number of active cascade sessions.
556    pub async fn active_count(&self) -> usize {
557        self.active.lock().await.len()
558    }
559
560    /// Return a snapshot of active sessions and their current turn number.
561    pub async fn snapshot(&self) -> HashMap<String, usize> {
562        self.active.lock().await.clone()
563    }
564}
565
566#[cfg(test)]
567mod tests {
568    use super::*;
569
570    fn make_echo_infer() -> InferFn {
571        Arc::new(|prompt: String| {
572            Box::pin(async move { Ok(format!("Echo: {}", &prompt[..prompt.len().min(50)])) })
573        })
574    }
575
576    fn make_one_tool_infer() -> InferFn {
577        use std::sync::atomic::{AtomicUsize, Ordering};
578        let calls = Arc::new(AtomicUsize::new(0));
579        Arc::new(move |_prompt: String| {
580            let n = calls.fetch_add(1, Ordering::SeqCst);
581            Box::pin(async move {
582                if n == 0 {
583                    // First call: request a tool
584                    Ok(r#"<tool_call>{"name":"search","arguments":{"q":"rust"}}</tool_call>"#
585                        .to_string())
586                } else {
587                    // Second call: done
588                    Ok("Final answer based on search results.".to_string())
589                }
590            })
591        })
592    }
593
594    #[tokio::test]
595    async fn test_no_tool_calls_terminates_immediately() {
596        let engine = CascadeEngine::new(make_echo_infer(), Arc::new(NoopToolExecutor));
597        let result = engine.run("Hello", &Default::default()).await.unwrap();
598        assert_eq!(result.termination_reason, TerminationReason::NoToolCalls);
599        assert_eq!(result.turns.len(), 1);
600        assert_eq!(result.total_tool_calls, 0);
601    }
602
603    #[tokio::test]
604    async fn test_done_sentinel_terminates() {
605        let infer: InferFn = Arc::new(|_p: String| {
606            Box::pin(async { Ok(format!("I am done. {DONE_SENTINEL}")) })
607        });
608        let engine = CascadeEngine::new(infer, Arc::new(NoopToolExecutor));
609        let result = engine.run("Test", &Default::default()).await.unwrap();
610        assert_eq!(result.termination_reason, TerminationReason::DoneSentinel);
611    }
612
613    #[tokio::test]
614    async fn test_tool_call_then_answer() {
615        let engine = CascadeEngine::new(make_one_tool_infer(), Arc::new(NoopToolExecutor));
616        let result = engine.run("Search for Rust", &Default::default()).await.unwrap();
617        assert_eq!(result.termination_reason, TerminationReason::NoToolCalls);
618        assert_eq!(result.turns.len(), 2);
619        assert_eq!(result.total_tool_calls, 1);
620        assert_eq!(result.turns[0].tool_calls[0].name, "search");
621    }
622
623    #[tokio::test]
624    async fn test_max_turns_respected() {
625        let infer: InferFn = Arc::new(|_p: String| {
626            Box::pin(async {
627                Ok(r#"<tool_call>{"name":"loop","arguments":{}}</tool_call>"#.to_string())
628            })
629        });
630        let engine = CascadeEngine::new(infer, Arc::new(NoopToolExecutor));
631        let config = CascadeConfig {
632            max_turns: 3,
633            ..Default::default()
634        };
635        let result = engine.run("Loop forever", &config).await.unwrap();
636        assert_eq!(result.termination_reason, TerminationReason::MaxTurnsReached);
637        assert_eq!(result.turns.len(), 3);
638    }
639
640    #[tokio::test]
641    async fn test_xml_parser_extracts_tool_calls() {
642        let parser = XmlStyleToolParser;
643        let text = r#"
644            I'll search for this.
645            <tool_call>{"name": "search", "arguments": {"query": "tokio async"}}</tool_call>
646            And also look this up:
647            <tool_call>{"name": "lookup", "arguments": {"id": 42}}</tool_call>
648        "#;
649        let calls = parser.parse(text).await;
650        assert_eq!(calls.len(), 2);
651        assert_eq!(calls[0].name, "search");
652        assert_eq!(calls[1].name, "lookup");
653    }
654
655    #[tokio::test]
656    async fn test_cascade_monitor() {
657        let monitor = CascadeMonitor::default();
658        monitor.on_start("sess-1").await;
659        monitor.on_start("sess-2").await;
660        assert_eq!(monitor.active_count().await, 2);
661        monitor.on_turn("sess-1", 3).await;
662        let snap = monitor.snapshot().await;
663        assert_eq!(snap["sess-1"], 3);
664        monitor.on_complete("sess-1").await;
665        assert_eq!(monitor.active_count().await, 1);
666    }
667}