systemprompt_agent/services/a2a_server/streaming/
messages.rs1use 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#[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}