Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

117 changes: 101 additions & 16 deletions crates/aura-test-utils/src/mock_agent.rs
Original file line number Diff line number Diff line change
Expand Up @@ -84,6 +84,8 @@ enum Script {
pub struct MockAgent {
on_stream_start: Option<StartHook>,
script: Script,
run_id: aura::RunId,
stream_claim: aura::streaming::StreamClaim,
}

impl MockAgent {
Expand All @@ -92,6 +94,8 @@ impl MockAgent {
Self {
on_stream_start: None,
script: Script::Pending,
run_id: aura::RunId::mint(),
stream_claim: aura::streaming::StreamClaim::default(),
}
}

Expand All @@ -103,13 +107,13 @@ impl MockAgent {
Self::scripted(items.into_iter().map(Step::Item).collect())
}

/// Runs the given steps in order, then ends the stream. The script is
/// consumed by the first `stream` call; a second call
/// on the same agent yields an empty stream.
/// Runs the given steps in order, then ends the stream.
pub fn scripted(steps: Vec<Step>) -> Self {
Self {
on_stream_start: None,
script: Script::Steps(Mutex::new(Some(steps))),
run_id: aura::RunId::mint(),
stream_claim: aura::streaming::StreamClaim::default(),
}
}

Expand Down Expand Up @@ -168,13 +172,22 @@ impl StreamingAgent for MockAgent {
("test", "fake")
}

fn run_id(&self) -> aura::RunId {
self.run_id
}
Comment thread
justintime4tea marked this conversation as resolved.

async fn stream(
&self,
_query: &str,
_chat_history: Vec<Message>,
options: aura::streaming::RunOptions,
request_id: &str,
) -> AgentRun {
// Held to the contract a real agent keeps, so a test that names the
// wrong run or streams twice fails the way production would.
if let Err(refused) = self.stream_claim.claim(self.run_id, request_id) {
return AgentRun::refused(refused);
}
let stream = self.start(request_id).await;
// Carries a caller-supplied token so `cancel_token()` returns the one the
// caller named. The scripts do not race it, so cancelling does not end a
Expand All @@ -200,7 +213,12 @@ mod tests {
async fn a_pending_agent_never_yields() {
let agent = MockAgent::pending();
let mut stream = agent
.stream("q", vec![], aura::streaming::RunOptions::default(), "req_1")
.stream(
"q",
vec![],
aura::streaming::RunOptions::default(),
&agent.run_id().to_string(),
)
.await
.into_events();
assert!(
Expand All @@ -216,7 +234,12 @@ mod tests {
async fn a_yielding_agent_produces_its_items_then_ends() {
let agent = MockAgent::yielding(vec![items::text("hello "), items::text("world")]);
let mut stream = agent
.stream("q", vec![], aura::streaming::RunOptions::default(), "req_1")
.stream(
"q",
vec![],
aura::streaming::RunOptions::default(),
&agent.run_id().to_string(),
)
.await
.into_events();

Expand All @@ -234,21 +257,31 @@ mod tests {
#[tokio::test]
async fn the_start_hook_runs_before_the_stream() {
let ran = Arc::new(AtomicBool::new(false));
let flag = Arc::clone(&ran);
let seen = Arc::new(Mutex::new(None::<String>));
let (flag, seen_by_hook) = (Arc::clone(&ran), Arc::clone(&seen));
let agent = MockAgent::pending().on_stream_start(move |request_id| {
let flag = Arc::clone(&flag);
let (flag, seen) = (Arc::clone(&flag), Arc::clone(&seen_by_hook));
async move {
assert_eq!(request_id, "req_1");
*seen.lock().expect("seen lock") = Some(request_id);
flag.store(true, Ordering::SeqCst);
}
});

let _ = agent
.stream("q", vec![], aura::streaming::RunOptions::default(), "req_1")
.stream(
"q",
vec![],
aura::streaming::RunOptions::default(),
&agent.run_id().to_string(),
)
.await
.into_events();

assert!(ran.load(Ordering::SeqCst), "hook should run");
assert_eq!(
seen.lock().expect("seen lock").as_deref(),
Some(agent.run_id().to_string().as_str())
);
}

#[tokio::test(start_paused = true)]
Expand All @@ -270,35 +303,82 @@ mod tests {
"q",
vec![],
aura::streaming::RunOptions::default(),
"req_42",
&agent.run_id().to_string(),
)
.await
.into_events();
let items: Vec<_> = stream.collect().await;

assert_eq!(order.lock().expect("order lock").as_slice(), ["req_42"]);
assert_eq!(
order.lock().expect("order lock").as_slice(),
[agent.run_id().to_string()]
);
assert_eq!(items.len(), 1, "effects do not yield stream items");
}

/// The mock refuses what a real agent refuses: a stream naming another
/// run, which never reaches the start hook.
#[tokio::test]
async fn a_stream_naming_another_run_is_refused_before_the_hook() {
let ran = Arc::new(AtomicBool::new(false));
let flag = Arc::clone(&ran);
let agent = MockAgent::yielding([items::text("never")]).on_stream_start(move |_| {
let flag = Arc::clone(&flag);
async move { flag.store(true, Ordering::SeqCst) }
});

let items: Vec<_> = agent
.stream("q", vec![], aura::streaming::RunOptions::default(), "req_1")
.await
.into_events()
.collect()
.await;

let [Err(error)] = items.as_slice() else {
panic!("expected one refusal, got {items:?}");
};
assert!(matches!(
error.downcast_ref::<aura::streaming::StreamRefused>(),
Some(aura::streaming::StreamRefused::OtherId { .. })
));
assert!(!ran.load(Ordering::SeqCst), "a refused stream runs nothing");
}

#[tokio::test(start_paused = true)]
async fn a_script_is_consumed_by_the_first_stream_call() {
async fn a_second_stream_is_refused() {
let agent = MockAgent::yielding([items::text("once")]);

let first: Vec<_> = agent
.stream("q", vec![], aura::streaming::RunOptions::default(), "req_1")
.stream(
"q",
vec![],
aura::streaming::RunOptions::default(),
&agent.run_id().to_string(),
)
.await
.into_events()
.collect()
.await;
let second: Vec<_> = agent
.stream("q", vec![], aura::streaming::RunOptions::default(), "req_1")
.stream(
"q",
vec![],
aura::streaming::RunOptions::default(),
&agent.run_id().to_string(),
)
.await
.into_events()
.collect()
.await;

assert_eq!(first.len(), 1);
assert!(second.is_empty());
let [Err(error)] = second.as_slice() else {
panic!("expected one refusal, got {second:?}");
};
assert!(matches!(
error.downcast_ref::<aura::streaming::StreamRefused>(),
Some(aura::streaming::StreamRefused::AlreadyStreamed { .. })
));
}

#[tokio::test(start_paused = true)]
Expand All @@ -318,7 +398,12 @@ mod tests {
]);

let mut stream = agent
.stream("q", vec![], aura::streaming::RunOptions::default(), "req_1")
.stream(
"q",
vec![],
aura::streaming::RunOptions::default(),
&agent.run_id().to_string(),
)
.await
.into_events();
while stream.next().await.is_some() {
Expand Down
1 change: 1 addition & 0 deletions crates/aura-web-server/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,7 @@ http-body-util = "0.1"
tempfile = "3"
aura-config = { path = "../aura-config", features = ["test_util"] }
wiremock = "0.6"
tracing-subscriber = { workspace = true }

# Exhaustively explores interleavings of the task-cancel map, which a stress
# test cannot reach: the window is one mutex release and reacquire wide.
Expand Down
78 changes: 74 additions & 4 deletions crates/aura-web-server/src/a2a/agent_executor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -243,7 +243,18 @@ impl AgentExecutor for AuraAgentExecutor {
let run_tools = self.app_state.run_tools(&config);
let mut append_tracker: HashMap<(String, String, String), bool> = HashMap::new();

Box::pin(async_stream::stream! {
// The run is polled inside this span, so its log lines, the agent's
// included, carry the run's id and the task it executes.
let run_id = aura::RunId::mint();
let span = tracing::info_span!(
parent: None,
"agent.stream",
run.id = %run_id,
a2a.task_id = %ctx.task_id,
a2a.context_id = %ctx.context_id,
);

let execution = async_stream::stream! {
let task_id = ctx.task_id.clone();
let context_id = ctx.context_id.clone();

Expand All @@ -269,7 +280,9 @@ impl AgentExecutor for AuraAgentExecutor {
metadata: None,
}));

let request_id = format!("a2a_{}", task_id);
// In string form, the run's id is the request id everything
// request-keyed reads: MCP cancellation, approvals, the cancel map.
let request_id = run_id.to_string();

// Registered before the agent build and history fetch, both of which
// await, so a cancelTask during those has a token to cancel. Its
Expand All @@ -295,7 +308,7 @@ impl AgentExecutor for AuraAgentExecutor {
Some(&req_headers),
session_id,
None,
Some(request_id.clone()),
Some(run_id),
run_tools,
)
.await
Expand Down Expand Up @@ -558,7 +571,8 @@ impl AgentExecutor for AuraAgentExecutor {
metadata: None,
}));
}
})
};
in_span(span, Box::pin(execution))
}

fn cancel(&self, ctx: ExecutorContext) -> BoxStream<'static, Result<StreamResponse, A2AError>> {
Expand Down Expand Up @@ -621,6 +635,18 @@ fn extract_text(parts: Vec<Part>) -> Result<String, A2AError> {
Ok(strings.join("\n"))
}

/// `stream`, polled inside `span`, so its log lines carry the span's fields.
fn in_span<T: 'static>(
span: tracing::Span,
mut stream: BoxStream<'static, T>,
) -> BoxStream<'static, T> {
futures_util::stream::poll_fn(move |cx| {
let _entered = span.enter();
stream.poll_next_unpin(cx)
})
.boxed()
}

pub(super) fn fail_status(task_id: &str, context_id: &str, error_msg: &str) -> StreamResponse {
StreamResponse::StatusUpdate(TaskStatusUpdateEvent {
task_id: task_id.to_string(),
Expand Down Expand Up @@ -1046,6 +1072,50 @@ mod tests {
assert!(claim_agent(&state, &task_id, &agent) == CancelRaced::Yes);
}

/// A line the run logs carries its span's fields, which is what ties a
/// run's id to the task it executes.
#[tokio::test]
async fn a_line_logged_inside_the_span_carries_its_fields() {
let written = Arc::new(std::sync::Mutex::new(Vec::<u8>::new()));
let writer = {
let written = Arc::clone(&written);
move || SharedWriter(Arc::clone(&written))
};
let subscriber = tracing_subscriber::fmt()
.with_writer(writer)
.with_ansi(false)
.finish();
let _default = tracing::subscriber::set_default(subscriber);

let span = tracing::info_span!("agent.stream", a2a.task_id = %"t_1");
let run = futures_util::stream::once(async {
tracing::warn!("inside the run");
1
})
.boxed();
assert_eq!(in_span(span, run).collect::<Vec<_>>().await, vec![1]);

let written = String::from_utf8(written.lock().unwrap().clone()).unwrap();
let line = written
.lines()
.find(|line| line.contains("inside the run"))
.expect("the line is written");
assert!(line.contains("agent.stream{a2a.task_id=t_1}"), "{line}");
}

struct SharedWriter(Arc<std::sync::Mutex<Vec<u8>>>);

impl std::io::Write for SharedWriter {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.0.lock().unwrap().extend_from_slice(buf);
Ok(buf.len())
}

fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}

/// The guard is created before the entry, because the agent build and the
/// history fetch can both return early and would otherwise leave it behind.
#[test]
Expand Down
Loading
Loading