use std::sync::Arc;
use std::time::Duration;
use axum::response::sse::{Event, Sse};
use futures_util::stream::Stream;
use serde_json::{Value, json};
use crate::core::{RunId, Seq};
use crate::journal::RecordKind;
use crate::runtime::Runtime;
use super::a2a::{A2aArtifact, A2aTask, TaskState, sealed_state, task_artifacts, task_of};
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::RunSealed { 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::RunSealed { 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())
}
const fn closes(state: TaskState) -> bool {
matches!(
state,
TaskState::Completed | TaskState::Failed | TaskState::Canceled | TaskState::Rejected
)
}
pub fn tail(
runtime: Arc<Runtime>,
run: RunId,
case: Option<String>,
id: Value,
first: A2aTask,
from: Seq,
) -> Sse<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::RunSealed { outcome, .. } = record.kind() {
let state = sealed_state(outcome);
if state == TaskState::Completed {
match task_artifacts(&runtime, 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;
}
tokio::time::sleep(POLL).await;
}
};
Sse::new(stream)
}
pub async fn current(runtime: &Runtime, run: RunId) -> Option<(A2aTask, Option<String>, Seq)> {
let records = runtime.journal().read(run, 1).await.ok()?;
let last = records.last()?;
let (state, detail) = match last.kind() {
RecordKind::RunSuspended { reason } => (TaskState::InputRequired, reason.to_string()),
RecordKind::RunSealed { outcome, .. } => (sealed_state(outcome), outcome.clone()),
_ => (TaskState::Working, "running".to_owned()),
};
let case = records
.iter()
.find_map(|r| r.body.case.map(|c| c.to_string()));
let next = last.body.seq + 1;
let mut task = task_of(run, state, &detail, case.clone());
task.artifacts = task_artifacts(runtime, run, state).await.ok()?;
Some((task, case, next))
}