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 session_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 workflow_validate_schema: false,
50 session_resume: false,
51 models: vec![],
52 }
53 }
54 async fn run(&self, task: AgentTask, _ctx: RunContext) -> Result<AgentResult, BackendError> {
55 let n = self.count.fetch_add(1, Ordering::SeqCst) + 1;
56 eprintln!(
57 "[crash-backend] agent #{n}: {}",
58 task.name.as_deref().unwrap_or("?")
59 );
60 if n >= self.crash_after {
61 eprintln!("[crash-backend] exiting after {n} calls");
62 std::process::exit(1);
63 }
64 Ok(AgentResult {
65 agent_id: task.agent_id,
66 status: AgentStatus::Ok,
67 output: self.canned.clone(),
68 session_id: task.session_id.clone(),
69 findings: vec![],
70 tokens_used: TokenUsage::default(),
71 artifacts: vec![],
72 logs: LogRef::default(),
73 })
74 }
75 fn as_any(&self) -> &dyn std::any::Any {
76 self
77 }
78}
79
80#[derive(Clone)]
82pub struct CountingBackend {
83 canned: Value,
84 calls: Arc<Mutex<Vec<String>>>,
85}
86
87impl CountingBackend {
88 pub fn new(canned: Value) -> Self {
89 Self {
90 canned,
91 calls: Arc::new(Mutex::new(Vec::new())),
92 }
93 }
94 pub fn total_calls(&self) -> usize {
95 self.calls.lock().unwrap().len()
96 }
97}
98
99#[async_trait::async_trait]
100impl AgentBackend for CountingBackend {
101 fn id(&self) -> &'static str {
102 "counting"
103 }
104 fn capabilities(&self) -> AgentCapabilities {
105 AgentCapabilities {
106 streaming: true,
107 mcp_injection: false,
108 workflow_validate_schema: false,
109 session_resume: false,
110 models: vec![],
111 }
112 }
113 async fn run(&self, task: AgentTask, _ctx: RunContext) -> Result<AgentResult, BackendError> {
114 self.calls
115 .lock()
116 .unwrap()
117 .push(task.name.clone().unwrap_or_default());
118 Ok(AgentResult {
119 agent_id: task.agent_id,
120 status: AgentStatus::Ok,
121 output: self.canned.clone(),
122 session_id: task.session_id.clone(),
123 findings: vec![],
124 tokens_used: TokenUsage::default(),
125 artifacts: vec![],
126 logs: LogRef::default(),
127 })
128 }
129 fn as_any(&self) -> &dyn std::any::Any {
130 self
131 }
132}
133
134#[derive(Clone)]
136pub struct SharedBackend {
137 canned: Value,
138 call_count: Arc<AtomicU64>,
139 pub block_on: Arc<Mutex<Option<u64>>>,
140 pub fail_on: Arc<Mutex<Option<u64>>>,
141 calls: Arc<Mutex<Vec<CallRecord>>>,
142}
143
144impl SharedBackend {
145 pub fn new(canned: Value) -> Self {
146 Self {
147 canned,
148 call_count: Arc::new(AtomicU64::new(0)),
149 block_on: Arc::new(Mutex::new(None)),
150 fail_on: Arc::new(Mutex::new(None)),
151 calls: Arc::new(Mutex::new(Vec::new())),
152 }
153 }
154 pub fn total_calls(&self) -> usize {
155 self.calls.lock().unwrap().len()
156 }
157}
158
159#[async_trait::async_trait]
160impl AgentBackend for SharedBackend {
161 fn id(&self) -> &'static str {
162 "shared"
163 }
164 fn capabilities(&self) -> AgentCapabilities {
165 AgentCapabilities {
166 streaming: true,
167 mcp_injection: false,
168 workflow_validate_schema: false,
169 session_resume: false,
170 models: vec![],
171 }
172 }
173 async fn run(&self, task: AgentTask, ctx: RunContext) -> Result<AgentResult, BackendError> {
174 let seq = self.call_count.fetch_add(1, Ordering::SeqCst) + 1;
175 self.calls.lock().unwrap().push(CallRecord {
176 seq,
177 agent_name: task.name.clone(),
178 session_id: task.session_id.clone(),
179 prompt: task.prompt.clone(),
180 });
181 if self
182 .fail_on
183 .lock()
184 .unwrap()
185 .map(|n| n == seq)
186 .unwrap_or(false)
187 {
188 return Err(BackendError::Execution("simulated failure".into()));
189 }
190 if self
191 .block_on
192 .lock()
193 .unwrap()
194 .map(|n| n == seq)
195 .unwrap_or(false)
196 {
197 ctx.cancel.cancelled().await;
198 return Err(BackendError::Cancelled);
199 }
200 Ok(AgentResult {
201 agent_id: task.agent_id,
202 status: AgentStatus::Ok,
203 output: self.canned.clone(),
204 session_id: task.session_id.clone(),
205 findings: vec![],
206 tokens_used: TokenUsage::default(),
207 artifacts: vec![],
208 logs: LogRef::default(),
209 })
210 }
211 fn as_any(&self) -> &dyn std::any::Any {
212 self
213 }
214}
215
216
217