codewhale_workflow_js/
testing.rs1use std::sync::Mutex;
10use std::sync::atomic::{AtomicUsize, Ordering};
11use std::time::Duration;
12
13use async_trait::async_trait;
14use tokio::sync::oneshot;
15
16use crate::driver::{
17 BudgetSnapshot, ProgressEvent, SpawnedTask, TaskCompletion, TaskRequest, WorkflowDriver,
18};
19use crate::error::DriverError;
20
21#[derive(Debug, Clone)]
23pub enum FakeReply {
24 Complete(String),
26 Fail(String),
28 Cancelled,
30 BudgetExhausted(String),
32 Reject(String),
34 Never,
37}
38
39#[derive(Debug)]
40struct ReplyRule {
41 needle: String,
42 delay: Option<Duration>,
43 reply: FakeReply,
44}
45
46#[derive(Debug, Default)]
47struct Inner {
48 rules: Vec<ReplyRule>,
49 requests: Vec<TaskRequest>,
50 events: Vec<ProgressEvent>,
51 budget: BudgetSnapshot,
52 spend_per_task: u64,
53 next_id: u64,
54 held: Vec<oneshot::Sender<TaskCompletion>>,
55}
56
57#[derive(Debug, Default)]
62pub struct FakeDriver {
63 inner: Mutex<Inner>,
64 cancel_calls: AtomicUsize,
65}
66
67impl FakeDriver {
68 pub fn new() -> Self {
70 Self::default()
71 }
72
73 pub fn on(&self, needle: &str, reply: FakeReply) {
76 self.on_with_delay_opt(needle, reply, None);
77 }
78
79 pub fn on_with_delay(&self, needle: &str, reply: FakeReply, delay: Duration) {
82 self.on_with_delay_opt(needle, reply, Some(delay));
83 }
84
85 fn on_with_delay_opt(&self, needle: &str, reply: FakeReply, delay: Option<Duration>) {
86 self.lock().rules.push(ReplyRule {
87 needle: needle.to_string(),
88 delay,
89 reply,
90 });
91 }
92
93 pub fn set_budget(&self, total: Option<u64>, spend_per_task: u64) {
96 let mut inner = self.lock();
97 inner.budget = BudgetSnapshot { total, spent: 0 };
98 inner.spend_per_task = spend_per_task;
99 }
100
101 pub fn requests(&self) -> Vec<TaskRequest> {
103 self.lock().requests.clone()
104 }
105
106 pub fn request_descriptions(&self) -> Vec<String> {
108 self.lock()
109 .requests
110 .iter()
111 .map(|request| request.description.clone())
112 .collect()
113 }
114
115 pub fn spawn_count(&self) -> usize {
117 self.lock().requests.len()
118 }
119
120 pub fn events(&self) -> Vec<ProgressEvent> {
122 self.lock().events.clone()
123 }
124
125 pub fn cancel_all_calls(&self) -> usize {
127 self.cancel_calls.load(Ordering::SeqCst)
128 }
129
130 fn lock(&self) -> std::sync::MutexGuard<'_, Inner> {
131 self.inner.lock().expect("FakeDriver mutex poisoned")
132 }
133}
134
135#[async_trait]
136impl WorkflowDriver for FakeDriver {
137 async fn spawn_task(&self, request: TaskRequest) -> Result<SpawnedTask, DriverError> {
138 let (task_id, reply, delay) = {
139 let mut inner = self.lock();
140 let matched = inner
141 .rules
142 .iter()
143 .find(|rule| request.description.contains(&rule.needle))
144 .map(|rule| (rule.reply.clone(), rule.delay));
145 let (reply, delay) = matched.unwrap_or_else(|| {
146 (
147 FakeReply::Complete(format!("done:{}", request.description)),
148 None,
149 )
150 });
151 if let FakeReply::Reject(message) = reply {
152 return Err(DriverError::Rejected(message));
153 }
154 inner.requests.push(request);
155 inner.budget.spent += inner.spend_per_task;
156 inner.next_id += 1;
157 (format!("agent_{:04}", inner.next_id), reply, delay)
158 };
159
160 let (tx, rx) = oneshot::channel();
161 match reply {
162 FakeReply::Never => self.lock().held.push(tx),
163 reply => {
164 let completion = match reply {
165 FakeReply::Complete(text) => TaskCompletion::Completed { text },
166 FakeReply::Fail(message) => TaskCompletion::Failed { message },
167 FakeReply::Cancelled => TaskCompletion::Cancelled,
168 FakeReply::BudgetExhausted(message) => {
169 TaskCompletion::BudgetExhausted { message }
170 }
171 FakeReply::Reject(_) | FakeReply::Never => unreachable!("handled above"),
172 };
173 match delay {
174 None => {
175 let _ = tx.send(completion);
176 }
177 Some(delay) => {
178 tokio::spawn(async move {
179 tokio::time::sleep(delay).await;
180 let _ = tx.send(completion);
181 });
182 }
183 }
184 }
185 }
186 Ok(SpawnedTask {
187 task_id,
188 completion: rx,
189 })
190 }
191
192 fn cancel_all(&self) {
193 self.cancel_calls.fetch_add(1, Ordering::SeqCst);
194 }
195
196 fn budget(&self) -> BudgetSnapshot {
197 self.lock().budget
198 }
199
200 fn progress(&self, event: ProgressEvent) {
201 self.lock().events.push(event);
202 }
203}