Skip to main content

systemprompt_agent/services/a2a_server/streaming/
messages.rs

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