Skip to main content

luft_core/
testing.rs

1//! Test utilities for resume and integration testing.
2//!
3//! Provides instrumented backends that can simulate crashes, blocking,
4//! and call recording — useful for testing crash-and-resume scenarios.
5
6use crate::contract::backend::{
7    AgentBackend, AgentCapabilities, AgentResult, AgentStatus, AgentTask, BackendError, LogRef,
8    RunContext,
9};
10use crate::contract::ids::TokenUsage;
11use serde_json::Value;
12use std::sync::atomic::{AtomicU64, Ordering};
13use std::sync::{Arc, Mutex};
14
15#[derive(Debug, Clone)]
16pub struct CallRecord {
17    pub seq: u64,
18    pub agent_name: Option<String>,
19    pub thread_id: Option<String>,
20    pub prompt: String,
21}
22
23/// Backend that calls `std::process::exit(1)` after N calls.
24pub struct CrashBackend {
25    canned: Value,
26    crash_after: u64,
27    count: AtomicU64,
28}
29
30impl CrashBackend {
31    pub fn new(canned: Value, crash_after: u64) -> Self {
32        Self {
33            canned,
34            crash_after,
35            count: AtomicU64::new(0),
36        }
37    }
38}
39
40#[async_trait::async_trait]
41impl AgentBackend for CrashBackend {
42    fn id(&self) -> &'static str {
43        "crash"
44    }
45    fn capabilities(&self) -> AgentCapabilities {
46        AgentCapabilities {
47            streaming: true,
48            mcp_injection: false,
49            workflow_validate_schema: false,
50            models: vec![],
51        }
52    }
53    async fn run(&self, task: AgentTask, _ctx: RunContext) -> Result<AgentResult, BackendError> {
54        let n = self.count.fetch_add(1, Ordering::SeqCst) + 1;
55        eprintln!(
56            "[crash-backend] agent #{n}: {}",
57            task.name.as_deref().unwrap_or("?")
58        );
59        if n >= self.crash_after {
60            eprintln!("[crash-backend] exiting after {n} calls");
61            std::process::exit(1);
62        }
63        Ok(AgentResult {
64            agent_id: task.agent_id,
65            status: AgentStatus::Ok,
66            output: self.canned.clone(),
67            thread_id: task.thread_id.clone(),
68            findings: vec![],
69            tokens_used: TokenUsage::default(),
70            artifacts: vec![],
71            logs: LogRef::default(),
72        })
73    }
74    fn as_any(&self) -> &dyn std::any::Any {
75        self
76    }
77}
78
79/// Backend that records all dispatched agent names.
80#[derive(Clone)]
81pub struct CountingBackend {
82    canned: Value,
83    calls: Arc<Mutex<Vec<String>>>,
84}
85
86impl CountingBackend {
87    pub fn new(canned: Value) -> Self {
88        Self {
89            canned,
90            calls: Arc::new(Mutex::new(Vec::new())),
91        }
92    }
93    pub fn total_calls(&self) -> usize {
94        self.calls.lock().unwrap().len()
95    }
96}
97
98#[async_trait::async_trait]
99impl AgentBackend for CountingBackend {
100    fn id(&self) -> &'static str {
101        "counting"
102    }
103    fn capabilities(&self) -> AgentCapabilities {
104        AgentCapabilities {
105            streaming: true,
106            mcp_injection: false,
107            workflow_validate_schema: false,
108            models: vec![],
109        }
110    }
111    async fn run(&self, task: AgentTask, _ctx: RunContext) -> Result<AgentResult, BackendError> {
112        self.calls
113            .lock()
114            .unwrap()
115            .push(task.name.clone().unwrap_or_default());
116        Ok(AgentResult {
117            agent_id: task.agent_id,
118            status: AgentStatus::Ok,
119            output: self.canned.clone(),
120            thread_id: task.thread_id.clone(),
121            findings: vec![],
122            tokens_used: TokenUsage::default(),
123            artifacts: vec![],
124            logs: LogRef::default(),
125        })
126    }
127    fn as_any(&self) -> &dyn std::any::Any {
128        self
129    }
130}
131
132/// Backend with shared state via Arc for resume tests.
133#[derive(Clone)]
134pub struct SharedBackend {
135    canned: Value,
136    call_count: Arc<AtomicU64>,
137    pub block_on: Arc<Mutex<Option<u64>>>,
138    pub fail_on: Arc<Mutex<Option<u64>>>,
139    calls: Arc<Mutex<Vec<CallRecord>>>,
140}
141
142impl SharedBackend {
143    pub fn new(canned: Value) -> Self {
144        Self {
145            canned,
146            call_count: Arc::new(AtomicU64::new(0)),
147            block_on: Arc::new(Mutex::new(None)),
148            fail_on: Arc::new(Mutex::new(None)),
149            calls: Arc::new(Mutex::new(Vec::new())),
150        }
151    }
152    pub fn total_calls(&self) -> usize {
153        self.calls.lock().unwrap().len()
154    }
155}
156
157#[async_trait::async_trait]
158impl AgentBackend for SharedBackend {
159    fn id(&self) -> &'static str {
160        "shared"
161    }
162    fn capabilities(&self) -> AgentCapabilities {
163        AgentCapabilities {
164            streaming: true,
165            mcp_injection: false,
166            workflow_validate_schema: false,
167            models: vec![],
168        }
169    }
170    async fn run(&self, task: AgentTask, ctx: RunContext) -> Result<AgentResult, BackendError> {
171        let seq = self.call_count.fetch_add(1, Ordering::SeqCst) + 1;
172        self.calls.lock().unwrap().push(CallRecord {
173            seq,
174            agent_name: task.name.clone(),
175            thread_id: task.thread_id.clone(),
176            prompt: task.prompt.clone(),
177        });
178        if self
179            .fail_on
180            .lock()
181            .unwrap()
182            .map(|n| n == seq)
183            .unwrap_or(false)
184        {
185            return Err(BackendError::Execution("simulated failure".into()));
186        }
187        if self
188            .block_on
189            .lock()
190            .unwrap()
191            .map(|n| n == seq)
192            .unwrap_or(false)
193        {
194            ctx.cancel.cancelled().await;
195            return Err(BackendError::Cancelled);
196        }
197        Ok(AgentResult {
198            agent_id: task.agent_id,
199            status: AgentStatus::Ok,
200            output: self.canned.clone(),
201            thread_id: task.thread_id.clone(),
202            findings: vec![],
203            tokens_used: TokenUsage::default(),
204            artifacts: vec![],
205            logs: LogRef::default(),
206        })
207    }
208    fn as_any(&self) -> &dyn std::any::Any {
209        self
210    }
211}
212
213
214