use std::sync::Mutex;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use async_trait::async_trait;
use tokio::sync::oneshot;
use crate::driver::{
BudgetSnapshot, ProgressEvent, SpawnedTask, TaskCompletion, TaskRequest, WorkflowDriver,
};
use crate::error::DriverError;
#[derive(Debug, Clone)]
pub enum FakeReply {
Complete(String),
Fail(String),
Cancelled,
BudgetExhausted(String),
Reject(String),
Never,
}
#[derive(Debug)]
struct ReplyRule {
needle: String,
delay: Option<Duration>,
reply: FakeReply,
}
#[derive(Debug, Default)]
struct Inner {
rules: Vec<ReplyRule>,
requests: Vec<TaskRequest>,
events: Vec<ProgressEvent>,
budget: BudgetSnapshot,
spend_per_task: u64,
next_id: u64,
held: Vec<oneshot::Sender<TaskCompletion>>,
}
#[derive(Debug, Default)]
pub struct FakeDriver {
inner: Mutex<Inner>,
cancel_calls: AtomicUsize,
}
impl FakeDriver {
pub fn new() -> Self {
Self::default()
}
pub fn on(&self, needle: &str, reply: FakeReply) {
self.on_with_delay_opt(needle, reply, None);
}
pub fn on_with_delay(&self, needle: &str, reply: FakeReply, delay: Duration) {
self.on_with_delay_opt(needle, reply, Some(delay));
}
fn on_with_delay_opt(&self, needle: &str, reply: FakeReply, delay: Option<Duration>) {
self.lock().rules.push(ReplyRule {
needle: needle.to_string(),
delay,
reply,
});
}
pub fn set_budget(&self, total: Option<u64>, spend_per_task: u64) {
let mut inner = self.lock();
inner.budget = BudgetSnapshot { total, spent: 0 };
inner.spend_per_task = spend_per_task;
}
pub fn requests(&self) -> Vec<TaskRequest> {
self.lock().requests.clone()
}
pub fn request_descriptions(&self) -> Vec<String> {
self.lock()
.requests
.iter()
.map(|request| request.description.clone())
.collect()
}
pub fn spawn_count(&self) -> usize {
self.lock().requests.len()
}
pub fn events(&self) -> Vec<ProgressEvent> {
self.lock().events.clone()
}
pub fn cancel_all_calls(&self) -> usize {
self.cancel_calls.load(Ordering::SeqCst)
}
fn lock(&self) -> std::sync::MutexGuard<'_, Inner> {
self.inner.lock().expect("FakeDriver mutex poisoned")
}
}
#[async_trait]
impl WorkflowDriver for FakeDriver {
async fn spawn_task(&self, request: TaskRequest) -> Result<SpawnedTask, DriverError> {
let (task_id, reply, delay) = {
let mut inner = self.lock();
let matched = inner
.rules
.iter()
.find(|rule| request.description.contains(&rule.needle))
.map(|rule| (rule.reply.clone(), rule.delay));
let (reply, delay) = matched.unwrap_or_else(|| {
(
FakeReply::Complete(format!("done:{}", request.description)),
None,
)
});
if let FakeReply::Reject(message) = reply {
return Err(DriverError::Rejected(message));
}
inner.requests.push(request);
inner.budget.spent += inner.spend_per_task;
inner.next_id += 1;
(format!("agent_{:04}", inner.next_id), reply, delay)
};
let (tx, rx) = oneshot::channel();
match reply {
FakeReply::Never => self.lock().held.push(tx),
reply => {
let completion = match reply {
FakeReply::Complete(text) => TaskCompletion::Completed { text },
FakeReply::Fail(message) => TaskCompletion::Failed { message },
FakeReply::Cancelled => TaskCompletion::Cancelled,
FakeReply::BudgetExhausted(message) => {
TaskCompletion::BudgetExhausted { message }
}
FakeReply::Reject(_) | FakeReply::Never => unreachable!("handled above"),
};
match delay {
None => {
let _ = tx.send(completion);
}
Some(delay) => {
tokio::spawn(async move {
tokio::time::sleep(delay).await;
let _ = tx.send(completion);
});
}
}
}
}
Ok(SpawnedTask {
task_id,
completion: rx,
})
}
fn cancel_all(&self) {
self.cancel_calls.fetch_add(1, Ordering::SeqCst);
}
fn budget(&self) -> BudgetSnapshot {
self.lock().budget
}
fn progress(&self, event: ProgressEvent) {
self.lock().events.push(event);
}
}