use super::*;
use std::sync::atomic::{AtomicBool, Ordering};
struct GatedExecutor {
proceed: Arc<Notify>,
started: Arc<AtomicBool>,
}
impl AgentExecutor for GatedExecutor {
fn execute<'a>(
&'a self,
ctx: &'a RequestContext,
queue: &'a dyn EventQueueWriter,
) -> Pin<Box<dyn Future<Output = A2aResult<()>> + Send + 'a>> {
Box::pin(async move {
self.started.store(true, Ordering::SeqCst);
self.proceed.notified().await;
queue
.write(StreamResponse::StatusUpdate(TaskStatusUpdateEvent {
task_id: ctx.task_id.clone(),
context_id: ContextId::new(ctx.context_id.clone()),
status: TaskStatus::new(TaskState::Completed),
metadata: None,
}))
.await?;
Ok(())
})
}
}
#[tokio::test]
async fn idle_handler_reports_zero_queues_and_tokens() {
let handler = RequestHandlerBuilder::new(GatedExecutor {
proceed: Arc::new(Notify::new()),
started: Arc::new(AtomicBool::new(false)),
})
.with_task_store(InMemoryTaskStore::new())
.build()
.expect("build handler");
assert_eq!(
handler.active_queue_count().await,
0,
"a handler that has served nothing must report no live queues"
);
assert_eq!(
handler.cancellation_token_count().await,
0,
"a handler that has served nothing must report no cancellation tokens"
);
}
#[tokio::test]
async fn one_in_flight_task_is_counted_exactly_once() {
let proceed = Arc::new(Notify::new());
let started = Arc::new(AtomicBool::new(false));
let handler = Arc::new(
RequestHandlerBuilder::new(GatedExecutor {
proceed: proceed.clone(),
started: started.clone(),
})
.with_task_store(InMemoryTaskStore::new())
.build()
.expect("build handler"),
);
let driver = Arc::clone(&handler);
let send_handle = tokio::spawn(async move {
let result = driver
.on_send_message(make_send_params(), true, None)
.await
.expect("send message");
if let SendMessageResult::Stream(mut reader) = result {
while let Some(_event) = reader.read().await {}
} else {
panic!("expected Stream");
}
});
while !started.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
}
assert_eq!(
handler.active_queue_count().await,
1,
"exactly one task is mid-flight, so exactly one event queue must be live"
);
assert_eq!(
handler.cancellation_token_count().await,
1,
"exactly one task is mid-flight, so exactly one cancellation token must be registered"
);
proceed.notify_one();
send_handle.await.expect("send handle");
}
#[derive(Debug)]
struct DeltaRecordingStore {
inner: InMemoryTaskStore,
deltas: Arc<Mutex<Vec<ArtifactDelta>>>,
}
impl DeltaRecordingStore {
fn new(deltas: Arc<Mutex<Vec<ArtifactDelta>>>) -> Self {
Self {
inner: InMemoryTaskStore::new(),
deltas,
}
}
}
impl TaskStore for DeltaRecordingStore {
fn save<'a>(
&'a self,
task: &'a Task,
) -> Pin<Box<dyn Future<Output = A2aResult<()>> + Send + 'a>> {
self.inner.save(task)
}
fn get<'a>(
&'a self,
id: &'a TaskId,
) -> Pin<Box<dyn Future<Output = A2aResult<Option<Task>>> + Send + 'a>> {
self.inner.get(id)
}
fn insert_if_absent<'a>(
&'a self,
task: &'a Task,
) -> Pin<Box<dyn Future<Output = A2aResult<bool>> + Send + 'a>> {
self.inner.insert_if_absent(task)
}
fn delete<'a>(
&'a self,
id: &'a TaskId,
) -> Pin<Box<dyn Future<Output = A2aResult<()>> + Send + 'a>> {
self.inner.delete(id)
}
fn list<'a>(
&'a self,
params: &'a ListTasksParams,
) -> Pin<Box<dyn Future<Output = A2aResult<TaskListResponse>> + Send + 'a>> {
self.inner.list(params)
}
fn save_artifact_delta<'a>(
&'a self,
task: &'a Task,
delta: ArtifactDelta,
) -> Pin<Box<dyn Future<Output = A2aResult<()>> + Send + 'a>> {
self.deltas.lock().expect("delta log").push(delta);
self.inner.save_artifact_delta(task, delta)
}
}
struct PushingExecutor {
count: usize,
}
impl AgentExecutor for PushingExecutor {
fn execute<'a>(
&'a self,
ctx: &'a RequestContext,
queue: &'a dyn EventQueueWriter,
) -> Pin<Box<dyn Future<Output = A2aResult<()>> + Send + 'a>> {
Box::pin(async move {
for i in 0..self.count {
queue
.write(StreamResponse::ArtifactUpdate(TaskArtifactUpdateEvent {
task_id: ctx.task_id.clone(),
context_id: ContextId::new(ctx.context_id.clone()),
artifact: Artifact::new(
format!("art-{i}"),
vec![Part::text(format!("chunk {i}"))],
),
append: None,
last_chunk: Some(true),
metadata: None,
}))
.await?;
}
queue
.write(StreamResponse::StatusUpdate(TaskStatusUpdateEvent {
task_id: ctx.task_id.clone(),
context_id: ContextId::new(ctx.context_id.clone()),
status: TaskStatus::new(TaskState::Completed),
metadata: None,
}))
.await?;
Ok(())
})
}
}
#[tokio::test]
async fn pushed_artifacts_are_reported_at_their_own_index() {
const COUNT: usize = 3;
let deltas = Arc::new(Mutex::new(Vec::new()));
let handler = Arc::new(
RequestHandlerBuilder::new(PushingExecutor { count: COUNT })
.with_task_store(DeltaRecordingStore::new(Arc::clone(&deltas)))
.build()
.expect("build handler"),
);
let result = handler
.on_send_message(make_send_params(), true, None)
.await
.expect("send message");
if let SendMessageResult::Stream(mut reader) = result {
while let Some(_event) = reader.read().await {}
} else {
panic!("expected Stream");
}
for _ in 0..1000 {
if deltas.lock().expect("delta log").len() >= COUNT {
break;
}
tokio::task::yield_now().await;
}
let seen = deltas.lock().expect("delta log").clone();
let pushed: Vec<usize> = seen
.iter()
.filter_map(|d| match d {
ArtifactDelta::Pushed { index } => Some(*index),
ArtifactDelta::AppendedParts { .. } => None,
})
.collect();
assert_eq!(
pushed,
(0..COUNT).collect::<Vec<_>>(),
"each pushed artifact must be reported at the position it occupies; got {pushed:?}"
);
}