use std::sync::Arc;
use std::time::Duration;
use axum::response::sse::{Event, Sse};
use futures_util::stream::{Stream, StreamExt as _};
use serde_json::{Value, json};
use crate::core::{RunId, Seq};
use crate::journal::RecordKind;
use crate::runtime::Runtime;
use super::a2a::{
A2aArtifact, A2aTask, ArtifactCache, StreamSlot, TaskState, cached_artifacts, sealed_state,
task_artifacts,
};
const POLL: Duration = Duration::from_millis(200);
fn stream_response(id: &Value, payload: &Value) -> Event {
Event::default().data(json!({"jsonrpc": "2.0", "id": id, "result": payload}).to_string())
}
pub(super) fn status_update(
run: RunId,
case: Option<&str>,
state: TaskState,
detail: &str,
) -> Value {
json!({
"statusUpdate": {
"taskId": run.to_string(),
"contextId": case.unwrap_or(&run.to_string()),
"status": {
"state": state,
"message": {
"messageId": format!("{run}-{detail}"),
"role": "ROLE_AGENT",
"parts": [{"text": detail}],
"taskId": run.to_string(),
},
},
}
})
}
pub(super) fn artifact_update(run: RunId, case: Option<&str>, artifact: &A2aArtifact) -> Value {
json!({
"artifactUpdate": {
"taskId": run.to_string(),
"contextId": case.unwrap_or(&run.to_string()),
"artifact": artifact,
"append": false,
"lastChunk": true,
}
})
}
pub(super) fn progress_of(kind: &RecordKind) -> Option<(TaskState, String)> {
match kind {
RecordKind::StepStarted { skill } => Some((TaskState::Working, format!("started {skill}"))),
RecordKind::StepFinished { outcome } => {
Some((TaskState::Working, format!("finished: {outcome}")))
}
RecordKind::RunSuspended { reason } => {
Some((TaskState::InputRequired, format!("waiting: {reason}")))
}
RecordKind::RunConcluded { outcome, .. } => Some((sealed_state(outcome), outcome.clone())),
_ => None,
}
}
pub(super) async fn payloads_for_record(
runtime: &Runtime,
record: &crate::journal::Record,
case: Option<&str>,
) -> Result<Vec<Value>, crate::core::RuntimeError> {
let run = record.body.run;
if let RecordKind::RunConcluded { outcome, .. } = record.kind() {
let state = sealed_state(outcome);
let mut payloads = Vec::new();
if state == TaskState::Completed
&& let Some(artifacts) = task_artifacts(runtime, run, state).await?
{
payloads.extend(
artifacts
.iter()
.map(|artifact| artifact_update(run, case, artifact)),
);
}
payloads.push(status_update(run, case, state, outcome));
return Ok(payloads);
}
Ok(progress_of(record.kind())
.map(|(state, detail)| vec![status_update(run, case, state, &detail)])
.unwrap_or_default())
}
pub(super) const fn closes(state: TaskState) -> bool {
matches!(
state,
TaskState::Completed | TaskState::Failed | TaskState::Canceled | TaskState::Rejected
)
}
#[allow(clippy::too_many_arguments)]
pub(super) fn tail(
runtime: Arc<Runtime>,
cache: Arc<std::sync::Mutex<ArtifactCache>>,
slot: StreamSlot,
run: RunId,
case: Option<String>,
id: Value,
first: A2aTask,
from: Seq,
) -> Sse<impl Stream<Item = Result<Event, std::convert::Infallible>>> {
Sse::new(
frames(runtime, cache, run, case, id, first, from).map(move |event| {
let _held = &slot;
event
}),
)
}
fn frames(
runtime: Arc<Runtime>,
cache: Arc<std::sync::Mutex<ArtifactCache>>,
run: RunId,
case: Option<String>,
id: Value,
first: A2aTask,
from: Seq,
) -> impl Stream<Item = Result<Event, std::convert::Infallible>> {
let already_over = closes(first.status.state);
let stream = async_stream::stream! {
yield Ok(stream_response(&id, &json!({ "task": first })));
if already_over {
return;
}
let mut next = from;
loop {
let Ok(records) = runtime.journal().read(run, next).await else {
return;
};
let mut done = false;
for record in &records {
next = record.body.seq + 1;
if let RecordKind::RunConcluded { outcome, .. } = record.kind() {
let state = sealed_state(outcome);
if state == TaskState::Completed {
match cached_artifacts(&runtime, &cache, run, state).await {
Ok(Some(artifacts)) => {
for artifact in artifacts {
yield Ok(stream_response(
&id,
&artifact_update(run, case.as_deref(), &artifact),
));
}
}
Ok(None) => {}
Err(_) => return,
}
}
yield Ok(stream_response(
&id,
&status_update(run, case.as_deref(), state, outcome),
));
done = true;
continue;
}
if let Some((state, detail)) = progress_of(record.kind()) {
yield Ok(stream_response(
&id,
&status_update(run, case.as_deref(), state, &detail),
));
done |= closes(state);
}
}
if done {
return;
}
if runtime.is_draining() {
return;
}
tokio::time::sleep(POLL).await;
}
};
stream
}
#[cfg(all(test, feature = "redb"))]
mod tests {
use super::{Event, RunId, TaskState, frames};
use futures_util::StreamExt as _;
use serde_json::json;
use std::sync::Arc;
#[tokio::test]
async fn a_stream_on_an_already_finished_task_ends() {
let store = Arc::new(crate::store::RedbStore::open_in_memory().expect("store"));
let runtime = crate::runtime::Runtime::builder(
Arc::clone(&store) as Arc<dyn crate::journal::JournalStore>
)
.build();
let run = RunId::generate();
let finished = crate::api::a2a::task_of(run, TaskState::Completed, "succeeded", None);
let collected: Vec<Result<Event, std::convert::Infallible>> = tokio::time::timeout(
std::time::Duration::from_secs(5),
frames(runtime, Arc::default(), run, None, json!(1), finished, 1).collect(),
)
.await
.expect(
"a stream on an already-finished task never ended: its closing record was \
written before this subscriber existed, so the loop will never see it",
);
assert_eq!(
collected.len(),
1,
"a stream on a finished task must yield the snapshot and stop; it \
produced {} events",
collected.len()
);
}
#[tokio::test]
async fn a_stream_ends_when_the_instance_is_draining() {
let store = Arc::new(crate::store::RedbStore::open_in_memory().expect("store"));
let runtime = crate::runtime::Runtime::builder(
Arc::clone(&store) as Arc<dyn crate::journal::JournalStore>
)
.build();
let run = RunId::generate();
let working = crate::api::a2a::task_of(run, TaskState::Working, "accepted", None);
runtime.drain(std::time::Duration::ZERO).await;
let collected: Vec<Result<Event, std::convert::Infallible>> = tokio::time::timeout(
std::time::Duration::from_secs(5),
frames(runtime, Arc::default(), run, None, json!(1), working, 1).collect(),
)
.await
.expect(
"a stream did not end while the instance was draining: it would hold the \
server's graceful shutdown open for as long as the run it watches",
);
assert_eq!(
collected.len(),
1,
"a draining instance must yield the snapshot and stop; it produced {} events",
collected.len()
);
}
}