1use 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
23pub 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 structured_output: 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#[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 structured_output: 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#[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 structured_output: 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