Skip to main content

systemprompt_agent/services/a2a_server/streaming/
messages.rs

1//! SSE stream creation for A2A messages, capped by the global semaphore.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use std::sync::Arc;
7
8use axum::response::sse::Event;
9use systemprompt_models::RequestContext;
10use tokio_stream::wrappers::ReceiverStream;
11
12use crate::error::AgentResult;
13use crate::models::a2a::Message;
14use crate::models::a2a::jsonrpc::NumberOrString;
15use crate::models::a2a::protocol::PushNotificationConfig;
16use crate::services::a2a_server::handlers::AgentHandlerState;
17use crate::services::a2a_server::processing::message::ProcessMessageStreamParams;
18use crate::services::registry::AgentRegistry;
19
20use super::event_loop::{ProcessEventsParams, process_events};
21use super::event_loop_lifecycle::handle_stream_creation_error;
22use super::initialization::setup_stream;
23use super::types::StreamInput;
24use super::webhook_client::WebhookContext;
25
26pub struct CreateSseStreamParams {
27    pub message: Message,
28    pub agent_name: String,
29    pub state: Arc<AgentHandlerState>,
30    pub request_id: NumberOrString,
31    pub context: RequestContext,
32    pub callback_config: Option<PushNotificationConfig>,
33}
34
35impl std::fmt::Debug for CreateSseStreamParams {
36    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
37        f.debug_struct("CreateSseStreamParams")
38            .field("message", &self.message)
39            .field("agent_name", &self.agent_name)
40            .field("request_id", &self.request_id)
41            .field("context", &self.context)
42            .field("callback_config", &self.callback_config)
43            .finish_non_exhaustive()
44    }
45}
46
47/// Returned by [`create_sse_stream`] when the global stream-concurrency cap
48/// is exhausted, so the streaming handler can answer with HTTP 503.
49#[derive(Debug, Clone, Copy)]
50pub struct StreamRejected;
51
52pub async fn create_sse_stream(
53    params: CreateSseStreamParams,
54) -> Result<ReceiverStream<Event>, StreamRejected> {
55    create_sse_stream_with_registry(params, AgentRegistry::new()).await
56}
57
58pub async fn create_sse_stream_with_registry(
59    params: CreateSseStreamParams,
60    registry: AgentResult<AgentRegistry>,
61) -> Result<ReceiverStream<Event>, StreamRejected> {
62    let CreateSseStreamParams {
63        message,
64        agent_name,
65        state,
66        request_id,
67        context,
68        callback_config,
69    } = params;
70
71    let Ok(permit) = Arc::clone(&state.stream_semaphore).try_acquire_owned() else {
72        tracing::warn!(
73            event = "a2a.stream.rejected",
74            agent = %agent_name,
75            available_permits = state.stream_semaphore.available_permits(),
76            "A2A stream rejected: global concurrency cap reached"
77        );
78        return Err(StreamRejected);
79    };
80
81    let (tx, rx) = tokio::sync::mpsc::channel(1024);
82
83    tracing::info!("create_sse_stream() called - spawning tokio task");
84
85    let input = StreamInput {
86        message,
87        agent_name,
88        state,
89        request_id,
90        context,
91        callback_config,
92        registry,
93    };
94
95    tokio::spawn(async move {
96        let _permit = permit;
97        tracing::info!("Inside tokio::spawn - task execution started");
98        run_stream(input, tx).await;
99    });
100
101    Ok(ReceiverStream::new(rx))
102}
103
104async fn run_stream(input: StreamInput, tx: tokio::sync::mpsc::Sender<Event>) {
105    let Ok(setup) = setup_stream(input, &tx).await else {
106        return;
107    };
108
109    tracing::info!(agent = %setup.agent_name, "Starting message stream processing for agent");
110
111    match setup
112        .processor
113        .process_message_stream(ProcessMessageStreamParams {
114            a2a_message: &setup.message,
115            agent_runtime: &setup.agent_runtime,
116            agent_name: &setup.agent_name,
117            context: &setup.context,
118            task_id: setup.task_id.clone(),
119        })
120        .await
121    {
122        Ok(chunk_rx) => {
123            let params = ProcessEventsParams {
124                tx,
125                chunk_rx,
126                task_id: setup.task_id,
127                context_id: setup.context_id,
128                message_id: setup.message_id,
129                original_message: setup.message,
130                agent_name: setup.agent_name,
131                context: setup.context,
132                task_repo: setup.task_repo,
133                processor: setup.processor,
134                request_id: setup.request_id,
135            };
136            process_events(params).await;
137        },
138        Err(e) => {
139            let webhook_context = WebhookContext::new(
140                setup.context.user_id().clone(),
141                setup.context.auth_token().as_str(),
142            );
143            handle_stream_creation_error(
144                &webhook_context,
145                e,
146                &setup.task_id,
147                &setup.context_id,
148                &setup.task_repo,
149            )
150            .await;
151        },
152    }
153}