use crate::auth::can_read_run;
use crate::error::WorkflowApiError;
use crate::state::WorkflowRunState;
use axum::extract::{Path, State};
use axum::response::sse::{Event, KeepAlive, Sse};
use axum::Extension;
use klieo_auth_common::Identity;
use klieo_core::agent::AgentEvent;
use std::convert::Infallible;
use std::sync::Arc;
use tokio::sync::broadcast;
use tokio_stream::wrappers::BroadcastStream;
use tokio_stream::{Stream, StreamExt};
pub(crate) async fn get_run_events(
State(state): State<Arc<WorkflowRunState>>,
Extension(identity): Extension<Identity>,
Path(run_id): Path<String>,
) -> Result<Sse<impl Stream<Item = Result<Event, Infallible>>>, WorkflowApiError> {
let author = run_author(&state, &run_id).await?;
if !can_read_run(&identity, &author) {
return Err(WorkflowApiError::RunNotFound);
}
let receiver = state
.progress
.subscribe(&run_id)
.ok_or(WorkflowApiError::RunNotFound)?;
let stream = event_json_stream(receiver).map(|data| Ok(Event::default().data(data)));
Ok(Sse::new(stream).keep_alive(KeepAlive::default()))
}
async fn run_author(state: &WorkflowRunState, run_id: &str) -> Result<String, WorkflowApiError> {
match state.run_store.get(run_id).await {
Ok(Some(view)) => Ok(view.author),
Ok(None) => Err(WorkflowApiError::RunNotFound),
Err(e) => Err(WorkflowApiError::Internal(Box::new(e))),
}
}
fn event_json_stream(receiver: broadcast::Receiver<AgentEvent>) -> impl Stream<Item = String> {
BroadcastStream::new(receiver).filter_map(|result| {
let event = result.ok()?;
serde_json::to_string(&event).ok()
})
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn lagged_receiver_skips_gap_and_keeps_streaming() {
let (sender, receiver) = broadcast::channel(2);
for _ in 0..3 {
let _ = sender.send(AgentEvent::LlmCallStarted);
}
sender
.send(AgentEvent::ToolCallStarted {
name: "post-lag".into(),
})
.unwrap();
drop(sender);
let events: Vec<String> = event_json_stream(receiver).collect().await;
assert!(
!events.is_empty(),
"the stream must continue past the Lagged gap"
);
assert!(
events.last().unwrap().contains("post-lag"),
"a post-lag event is still delivered: {events:?}"
);
}
}