use std::future::Future;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::watch;
use tokio::task::JoinHandle;
use super::AntigravityError;
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct TriggerConfig {
pub message: String,
pub interval: Duration,
}
impl TriggerConfig {
#[must_use]
pub fn new(message: impl Into<String>, interval: Duration) -> Self {
Self {
message: message.into(),
interval,
}
}
}
#[derive(Debug, Default)]
pub(crate) struct TriggerTasks {
handles: Vec<JoinHandle<()>>,
}
impl TriggerTasks {
pub(crate) fn push(&mut self, handle: JoinHandle<()>) {
self.handles.push(handle);
}
pub(crate) fn abort_all(&mut self) {
for handle in self.handles.drain(..) {
handle.abort();
}
}
}
impl Drop for TriggerTasks {
fn drop(&mut self) {
self.abort_all();
}
}
pub(crate) fn spawn_trigger_task<F, Fut>(
config: TriggerConfig,
mut idle: watch::Receiver<bool>,
turn_sync: Arc<tokio::sync::Mutex<()>>,
send: F,
) -> JoinHandle<()>
where
F: Fn(String) -> Fut + Send + 'static,
Fut: Future<Output = Result<(), AntigravityError>> + Send,
{
tokio::spawn(async move {
loop {
tokio::time::sleep(config.interval).await;
loop {
if idle.wait_for(|is_idle| *is_idle).await.is_err() {
tracing::debug!("Trigger task exiting: agent dropped");
return;
}
let sync = turn_sync.lock().await;
if !*idle.borrow() {
drop(sync);
continue;
}
if let Err(e) = send(config.message.clone()).await {
tracing::warn!("Trigger delivery failed; stopping trigger task: {e}");
return;
}
break;
}
tracing::debug!(
interval = ?config.interval,
"Delivered automated trigger message"
);
}
})
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::sync::mpsc;
#[test]
fn test_trigger_config_construction() {
let trigger = TriggerConfig::new("check the queue", Duration::from_secs(60));
assert_eq!(trigger.message, "check the queue");
assert_eq!(trigger.interval, Duration::from_secs(60));
}
fn observable_trigger_with_sync(
config: TriggerConfig,
idle: watch::Receiver<bool>,
turn_sync: Arc<tokio::sync::Mutex<()>>,
) -> (JoinHandle<()>, mpsc::UnboundedReceiver<String>) {
let (tx, rx) = mpsc::unbounded_channel();
let handle = spawn_trigger_task(config, idle, turn_sync, move |message| {
let tx = tx.clone();
async move {
tx.send(message)
.map_err(|e| AntigravityError::WebSocket(e.to_string()))
}
});
(handle, rx)
}
fn observable_trigger(
config: TriggerConfig,
idle: watch::Receiver<bool>,
) -> (JoinHandle<()>, mpsc::UnboundedReceiver<String>) {
observable_trigger_with_sync(config, idle, Arc::new(tokio::sync::Mutex::new(())))
}
#[tokio::test(start_paused = true)]
async fn test_trigger_fires_repeatedly_when_idle() {
let (_idle_tx, idle_rx) = watch::channel(true);
let (handle, mut rx) =
observable_trigger(TriggerConfig::new("ping", Duration::from_secs(60)), idle_rx);
assert_eq!(rx.recv().await.as_deref(), Some("ping"));
assert_eq!(rx.recv().await.as_deref(), Some("ping"));
assert_eq!(rx.recv().await.as_deref(), Some("ping"));
handle.abort();
}
#[tokio::test(start_paused = true)]
async fn test_trigger_defers_while_busy_and_collapses_missed_intervals() {
let (idle_tx, idle_rx) = watch::channel(false); let (handle, mut rx) = observable_trigger(
TriggerConfig::new("check", Duration::from_millis(10)),
idle_rx,
);
tokio::time::sleep(Duration::from_millis(105)).await;
assert!(rx.try_recv().is_err(), "must not deliver while busy");
idle_tx.send_replace(true);
assert_eq!(rx.recv().await.as_deref(), Some("check"));
tokio::time::sleep(Duration::from_millis(5)).await;
assert!(
rx.try_recv().is_err(),
"missed intervals must collapse into one delivery"
);
handle.abort();
}
#[tokio::test(start_paused = true)]
async fn test_trigger_defers_when_turn_begins_during_delivery_window() {
let (idle_tx, idle_rx) = watch::channel(true);
let turn_sync = Arc::new(tokio::sync::Mutex::new(()));
let (handle, mut rx) = observable_trigger_with_sync(
TriggerConfig::new("check", Duration::from_millis(10)),
idle_rx,
Arc::clone(&turn_sync),
);
let sync = turn_sync.lock().await;
tokio::time::sleep(Duration::from_millis(15)).await;
for _ in 0..20 {
tokio::task::yield_now().await;
}
assert!(
rx.try_recv().is_err(),
"must not deliver while the turn lock is held"
);
idle_tx.send_replace(false);
drop(sync);
for _ in 0..20 {
tokio::task::yield_now().await;
}
assert!(
rx.try_recv().is_err(),
"must re-check idleness under the lock and defer"
);
idle_tx.send_replace(true);
assert_eq!(rx.recv().await.as_deref(), Some("check"));
handle.abort();
}
#[tokio::test(start_paused = true)]
async fn test_trigger_task_exits_when_idle_channel_closes() {
let (idle_tx, idle_rx) = watch::channel(false);
let (handle, _rx) =
observable_trigger(TriggerConfig::new("x", Duration::from_millis(1)), idle_rx);
drop(idle_tx);
handle.await.expect("task exits cleanly, not aborted");
}
#[tokio::test(start_paused = true)]
async fn test_trigger_task_exits_when_send_fails() {
let (_idle_tx, idle_rx) = watch::channel(true);
let handle = spawn_trigger_task(
TriggerConfig::new("x", Duration::from_millis(1)),
idle_rx,
Arc::new(tokio::sync::Mutex::new(())),
|_message| async { Err(AntigravityError::WebSocket("closed".to_string())) },
);
handle.await.expect("task exits cleanly after send failure");
}
#[tokio::test(start_paused = true)]
async fn test_trigger_tasks_abort_on_drop() {
let (_idle_tx, idle_rx) = watch::channel(false); let (handle, _rx) =
observable_trigger(TriggerConfig::new("x", Duration::from_secs(1)), idle_rx);
let abort_handle = handle.abort_handle();
let mut tasks = TriggerTasks::default();
tasks.push(handle);
drop(tasks);
while !abort_handle.is_finished() {
tokio::task::yield_now().await;
}
}
}