use std::sync::Arc;
use adk_core::{AdkError, Agent, Event, EventStream, Result};
use futures::StreamExt;
use tokio::sync::{Notify, RwLock, Semaphore, mpsc};
use tokio::task::JoinHandle;
use super::event_source::EventSource;
pub type TriggerHandler = Arc<
dyn Fn(
super::event_source::TriggerEvent,
Arc<dyn Agent>,
)
-> std::pin::Pin<Box<dyn std::future::Future<Output = Result<EventStream>> + Send>>
+ Send
+ Sync,
>;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AmbientAgentStatus {
Running,
Paused,
Stopped,
}
const DEFAULT_MAX_CONCURRENT_TRIGGERS: usize = 4;
pub struct AmbientAgent {
agent: Arc<dyn Agent>,
source: Arc<dyn EventSource>,
trigger_handler: Option<TriggerHandler>,
status: Arc<RwLock<AmbientAgentStatus>>,
resume_notify: Arc<Notify>,
handle: Option<JoinHandle<()>>,
max_concurrent_triggers: usize,
output_tx: Option<mpsc::Sender<Result<Event>>>,
}
impl AmbientAgent {
pub fn new(agent: Arc<dyn Agent>, source: Arc<dyn EventSource>) -> Self {
Self {
agent,
source,
trigger_handler: None,
status: Arc::new(RwLock::new(AmbientAgentStatus::Stopped)),
resume_notify: Arc::new(Notify::new()),
handle: None,
max_concurrent_triggers: DEFAULT_MAX_CONCURRENT_TRIGGERS,
output_tx: None,
}
}
pub fn with_max_concurrent_triggers(mut self, max_concurrent_triggers: usize) -> Self {
self.max_concurrent_triggers = max_concurrent_triggers.max(1);
self
}
pub fn take_output(&mut self, capacity: usize) -> mpsc::Receiver<Result<Event>> {
let (tx, rx) = mpsc::channel(capacity.max(1));
self.output_tx = Some(tx);
rx
}
pub fn with_trigger_handler(mut self, handler: TriggerHandler) -> Self {
self.trigger_handler = Some(handler);
self
}
pub async fn start(&mut self) -> Result<()> {
let current = *self.status.read().await;
if current != AmbientAgentStatus::Stopped {
return Err(AdkError::agent("agent already running"));
}
if self.trigger_handler.is_none() {
return Err(AdkError::agent(
"AmbientAgent has no trigger handler, so starting it would log trigger events \
without ever invoking the agent. Call `with_trigger_handler` with a closure \
that drives the agent through a Runner.",
));
}
let stream = self.source.subscribe().await?;
let status = Arc::clone(&self.status);
let resume_notify = Arc::clone(&self.resume_notify);
let agent = Arc::clone(&self.agent);
let trigger_handler = self.trigger_handler.clone();
*self.status.write().await = AmbientAgentStatus::Running;
let permits = Arc::new(Semaphore::new(self.max_concurrent_triggers));
let output_tx = self.output_tx.clone();
let handler = trigger_handler.expect("checked above");
let handle = tokio::spawn(async move {
let mut stream = stream;
let mut running = futures::stream::FuturesUnordered::new();
loop {
loop {
let current_status = *status.read().await;
match current_status {
AmbientAgentStatus::Running => break,
AmbientAgentStatus::Paused => resume_notify.notified().await,
AmbientAgentStatus::Stopped => return,
}
}
let event = if running.is_empty() {
stream.next().await
} else {
tokio::select! {
biased;
Some(()) = running.next() => continue,
event = stream.next() => event,
}
};
let Some(event) = event else {
while running.next().await.is_some() {}
return;
};
let handler = Arc::clone(&handler);
let agent = Arc::clone(&agent);
let permits = Arc::clone(&permits);
let output_tx = output_tx.clone();
running.push(async move {
let _permit = permits.acquire_owned().await;
tracing::info!(
agent = agent.name(),
source = %event.source,
"ambient agent triggered"
);
match handler(event, Arc::clone(&agent)).await {
Ok(mut event_stream) => {
while let Some(result) = event_stream.next().await {
let failed = result.is_err();
if let Err(ref e) = result {
tracing::warn!(error = %e, "ambient agent invocation error");
}
if let Some(ref tx) = output_tx
&& tx.send(result).await.is_err()
{
tracing::debug!("ambient output receiver dropped");
return;
}
if failed {
return;
}
}
}
Err(e) => {
tracing::warn!(error = %e, "ambient agent trigger handler failed");
if let Some(ref tx) = output_tx {
let _ = tx.send(Err(e)).await;
}
}
}
});
}
});
self.handle = Some(handle);
Ok(())
}
pub async fn stop(&mut self) -> Result<()> {
let current = *self.status.read().await;
if current == AmbientAgentStatus::Stopped {
return Err(AdkError::agent("agent already stopped"));
}
*self.status.write().await = AmbientAgentStatus::Stopped;
self.resume_notify.notify_one();
if let Some(handle) = self.handle.take() {
handle.abort();
}
Ok(())
}
pub async fn pause(&mut self) -> Result<()> {
let current = *self.status.read().await;
if current != AmbientAgentStatus::Running {
return Err(AdkError::agent("can only pause a running agent"));
}
*self.status.write().await = AmbientAgentStatus::Paused;
Ok(())
}
pub async fn resume(&mut self) -> Result<()> {
let current = *self.status.read().await;
if current != AmbientAgentStatus::Paused {
return Err(AdkError::agent("can only resume a paused agent"));
}
*self.status.write().await = AmbientAgentStatus::Running;
self.resume_notify.notify_one();
Ok(())
}
pub async fn status(&self) -> AmbientAgentStatus {
*self.status.read().await
}
}
impl Drop for AmbientAgent {
fn drop(&mut self) {
if let Some(handle) = self.handle.take() {
handle.abort();
}
}
}
impl std::fmt::Debug for AmbientAgent {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AmbientAgent")
.field("agent", &self.agent.name())
.field("source", &self.source.name())
.finish()
}
}