use std::sync::Arc;
use axum::response::sse::Event;
use systemprompt_models::RequestContext;
use tokio_stream::wrappers::ReceiverStream;
use crate::error::AgentResult;
use crate::models::a2a::Message;
use crate::models::a2a::jsonrpc::NumberOrString;
use crate::services::a2a_server::handlers::AgentHandlerState;
use crate::services::a2a_server::processing::message::ProcessMessageStreamParams;
use crate::services::registry::AgentRegistry;
use super::event_loop::{ProcessEventsParams, process_events};
use super::event_loop_lifecycle::handle_stream_creation_error;
use super::initialization::setup_stream;
use super::types::StreamInput;
use super::webhook_client::WebhookContext;
pub struct CreateSseStreamParams {
pub message: Message,
pub agent_name: String,
pub state: Arc<AgentHandlerState>,
pub request_id: NumberOrString,
pub context: RequestContext,
}
impl std::fmt::Debug for CreateSseStreamParams {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CreateSseStreamParams")
.field("message", &self.message)
.field("agent_name", &self.agent_name)
.field("request_id", &self.request_id)
.field("context", &self.context)
.finish_non_exhaustive()
}
}
#[derive(Debug, Clone, Copy)]
pub struct StreamRejected;
pub async fn create_sse_stream(
params: CreateSseStreamParams,
) -> Result<ReceiverStream<Event>, StreamRejected> {
create_sse_stream_with_registry(params, AgentRegistry::new()).await
}
pub async fn create_sse_stream_with_registry(
params: CreateSseStreamParams,
registry: AgentResult<AgentRegistry>,
) -> Result<ReceiverStream<Event>, StreamRejected> {
let CreateSseStreamParams {
message,
agent_name,
state,
request_id,
context,
} = params;
let Ok(permit) = Arc::clone(&state.stream_semaphore).try_acquire_owned() else {
tracing::warn!(
event = "a2a.stream.rejected",
agent = %agent_name,
available_permits = state.stream_semaphore.available_permits(),
"A2A stream rejected: global concurrency cap reached"
);
return Err(StreamRejected);
};
let (tx, rx) = tokio::sync::mpsc::channel(1024);
let active_tasks = state.active_tasks.clone();
let input = StreamInput {
message,
agent_name,
state,
request_id,
context,
registry,
};
active_tasks.tracker().spawn(async move {
let _permit = permit;
run_stream(input, tx).await;
});
Ok(ReceiverStream::new(rx))
}
async fn run_stream(input: StreamInput, tx: tokio::sync::mpsc::Sender<Event>) {
let active_tasks = input.state.active_tasks.clone();
let webhooks = input.state.agent_state.webhooks();
let Ok(setup) = setup_stream(input, &tx).await else {
return;
};
tracing::info!(agent = %setup.agent_name, "Starting message stream processing for agent");
let guard = active_tasks.register(setup.task_id.clone());
match setup
.processor
.process_message_stream(ProcessMessageStreamParams {
a2a_message: &setup.message,
agent_runtime: &setup.agent_runtime,
agent_name: &setup.agent_name,
context: &setup.context,
task_id: setup.task_id.clone(),
cancel: guard.token(),
})
.await
{
Ok(stream) => {
let params = ProcessEventsParams {
tx,
stream,
task_id: setup.task_id,
context_id: setup.context_id,
message_id: setup.message_id,
original_message: setup.message,
agent_name: setup.agent_name,
context: setup.context,
task_repo: setup.task_repo,
processor: setup.processor,
};
process_events(params).await;
},
Err(e) => {
let webhook_context = WebhookContext::for_request(webhooks, &setup.context);
handle_stream_creation_error(
&webhook_context,
e,
&setup.task_id,
&setup.context_id,
&setup.task_repo,
)
.await;
},
}
}