1use 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
54pub const DEFAULT_MAX_TURNS: usize = 10;
56
57pub const DONE_SENTINEL: &str = "[DONE]";
59
60#[derive(Debug, Clone, Serialize, Deserialize)]
62pub struct ToolCall {
63 pub name: String,
65 pub arguments: HashMap<String, serde_json::Value>,
67 pub position: usize,
69}
70
71#[derive(Debug, Clone, Serialize, Deserialize)]
73pub struct CascadeTurn {
74 pub turn: usize,
76 pub prompt: String,
78 pub response: String,
80 pub tool_calls: Vec<ToolCall>,
82 pub tool_results: Vec<ToolResult>,
84 pub elapsed_ms: u64,
86}
87
88#[derive(Debug, Clone, Serialize, Deserialize)]
90pub struct ToolResult {
91 pub tool_name: String,
93 pub result: serde_json::Value,
95 pub success: bool,
97 pub error: Option<String>,
99}
100
101#[derive(Debug, Clone, Serialize, Deserialize)]
103pub struct CascadeResult {
104 pub session_id: SessionId,
106 pub turns: Vec<CascadeTurn>,
108 pub final_answer: String,
110 pub termination_reason: TerminationReason,
112 pub total_elapsed_ms: u64,
114 pub total_tool_calls: usize,
116}
117
118#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
120pub enum TerminationReason {
121 NoToolCalls,
123 DoneSentinel,
125 MaxTurnsReached,
127 Error(String),
129}
130
131#[async_trait]
136pub trait ToolCallParser: Send + Sync {
137 async fn parse(&self, response: &str) -> Vec<ToolCall>;
139}
140
141pub 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#[async_trait]
194pub trait ToolExecutor: Send + Sync {
195 async fn execute(&self, call: &ToolCall) -> ToolResult;
197}
198
199pub 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
220fn build_next_prompt(
222 original_prompt: &str,
223 turns: &[CascadeTurn],
224 pending_results: &[ToolResult],
225) -> String {
226 let mut parts = Vec::new();
227
228 parts.push(format!("User: {original_prompt}"));
230
231 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 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#[derive(Debug, Clone)]
258pub struct CascadeConfig {
259 pub max_turns: usize,
261 pub timeout: Duration,
263 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
277pub 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
288pub struct CascadeEngine {
316 infer_fn: InferFn,
317 executor: Arc<dyn ToolExecutor>,
318 parser: Arc<dyn ToolCallParser>,
319}
320
321impl CascadeEngine {
322 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 pub fn with_tool_parser(mut self, parser: Arc<dyn ToolCallParser>) -> Self {
333 self.parser = parser;
334 self
335 }
336
337 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 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 let response = (self.infer_fn)(current_prompt.clone()).await.map_err(|e| {
385 if turn_idx == 0 {
386 e
388 } else {
389 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 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 let tool_calls = self.parser.parse(&response).await;
430
431 if tool_calls.is_empty() {
432 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 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 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 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
521pub 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 pub async fn on_start(&self, session_id: &str) {
540 self.active.lock().await.insert(session_id.to_string(), 0);
541 }
542
543 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 pub async fn on_complete(&self, session_id: &str) {
552 self.active.lock().await.remove(session_id);
553 }
554
555 pub async fn active_count(&self) -> usize {
557 self.active.lock().await.len()
558 }
559
560 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 Ok(r#"<tool_call>{"name":"search","arguments":{"q":"rust"}}</tool_call>"#
585 .to_string())
586 } else {
587 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}