systemprompt_agent/services/a2a_server/processing/message/message_handler/
mod.rs1mod helpers;
12
13use std::sync::Arc;
14
15use uuid::Uuid;
16
17use self::helpers::{
18 BroadcastAguiLifecycleParams, broadcast_agui_lifecycle, collect_stream_response,
19};
20use crate::models::a2a::{Message, MessageRole, Part, Task, TaskState, TaskStatus, TextPart};
21use crate::services::a2a_server::processing::message::persistence::{
22 broadcast_completion, persist_completed_task,
23};
24use crate::services::a2a_server::processing::message::stream_processor::StreamProcessor;
25use crate::services::a2a_server::processing::message::{
26 MessageProcessor, ProcessMessageStreamParams,
27};
28use crate::services::a2a_server::processing::task_builder::build_completed_task;
29use crate::services::a2a_server::streaming::broadcast::{
30 BroadcastTaskCreatedParams, broadcast_task_created,
31};
32use crate::services::shared::{AgentServiceError, Result};
33use systemprompt_identifiers::{ContextId, MessageId, SessionId, TaskId, TraceId, UserId};
34use systemprompt_models::{RequestContext, TaskMetadata};
35
36impl MessageProcessor {
37 pub(in crate::services::a2a_server) async fn handle_message(
38 &self,
39 message: Message,
40 agent_name: &str,
41 context: &RequestContext,
42 ) -> Result<Task> {
43 let agent_runtime = self.load_agent_runtime(agent_name).await?;
44 self.handle_message_with_runtime(message, &agent_runtime, agent_name, context)
45 .await
46 }
47
48 pub async fn handle_message_with_runtime(
49 &self,
50 message: Message,
51 agent_runtime: &crate::models::AgentRuntimeInfo,
52 agent_name: &str,
53 context: &RequestContext,
54 ) -> Result<Task> {
55 tracing::info!(agent_name = %agent_name, "Handling non-streaming message");
56
57 let context_id = &message.context_id;
58
59 self.repositories
60 .contexts
61 .get_context(context_id, context.user_id())
62 .await
63 .map_err(|e| {
64 AgentServiceError::Internal(format!(
65 "Context validation failed - context_id: {}, user_id: {}, error: {}",
66 context_id,
67 context.user_id(),
68 e
69 ))
70 })?;
71
72 tracing::info!(
73 context_id = %context_id,
74 user_id = %context.user_id(),
75 "Context validated"
76 );
77
78 let task_id = resolve_task_id(&message);
79 let task = new_submitted_task(&task_id, context_id, agent_name);
80
81 self.persist_and_announce(&task, &message, agent_name, context)
82 .await?;
83
84 let stream_processor = StreamProcessor {
85 ai_service: Arc::clone(&self.ai_service),
86 context_service: self.context_service.clone(),
87 skill_service: Arc::clone(&self.skill_service),
88 execution_step_repo: Arc::clone(&self.execution_step_repo),
89 };
90
91 let chunk_rx = stream_processor
92 .process_message_stream(ProcessMessageStreamParams {
93 a2a_message: &message,
94 agent_runtime,
95 agent_name,
96 context,
97 task_id: task_id.clone(),
98 })
99 .await?;
100
101 let (response_text, tool_artifacts) = collect_stream_response(chunk_rx, context).await?;
102
103 let task = build_completed_task(
104 task_id,
105 context_id.clone(),
106 response_text.clone(),
107 message.clone(),
108 tool_artifacts,
109 );
110
111 let agent_message = resolve_agent_message(&task, &message, &response_text);
112
113 if context.user_type() == systemprompt_models::auth::UserType::Anon {
114 tracing::warn!(
115 context_id = %context_id,
116 session_id = %context.session_id(),
117 "Saving messages for anonymous user"
118 );
119 }
120
121 self.persist_or_mark_failed(&task, &message, &agent_message, context)
122 .await?;
123
124 broadcast_completion(&task, context).await;
125
126 broadcast_agui_lifecycle(BroadcastAguiLifecycleParams {
127 context,
128 context_id,
129 task: &task,
130 agent_message: &agent_message,
131 response_text: &response_text,
132 })
133 .await;
134
135 Ok(task)
136 }
137
138 async fn persist_and_announce(
139 &self,
140 task: &Task,
141 message: &Message,
142 agent_name: &str,
143 context: &RequestContext,
144 ) -> Result<()> {
145 if let Err(e) = self
146 .repositories
147 .tasks
148 .create_task(crate::repository::task::RepoCreateTaskParams {
149 task,
150 user_id: &UserId::new(context.user_id().as_str()),
151 session_id: &SessionId::new(context.session_id().as_str()),
152 trace_id: &TraceId::new(context.trace_id().as_str()),
153 agent_name,
154 })
155 .await
156 {
157 return Err(AgentServiceError::Internal(format!(
158 "Failed to persist task at start: {e}"
159 )));
160 }
161
162 tracing::info!(task_id = %task.id, "Task persisted to database");
163
164 broadcast_task_created(BroadcastTaskCreatedParams {
165 task_id: &task.id,
166 context_id: &task.context_id,
167 user_id: context.user_id().as_str(),
168 user_message: message,
169 agent_name,
170 token: context.auth_token().as_str(),
171 })
172 .await;
173
174 let working_timestamp = chrono::Utc::now();
175 if let Err(e) = self
176 .repositories
177 .tasks
178 .update_task_state(&task.id, TaskState::Working, &working_timestamp)
179 .await
180 {
181 tracing::error!(task_id = %task.id, error = %e, "Failed to mark task as working");
182 }
183
184 Ok(())
185 }
186
187 async fn persist_or_mark_failed(
188 &self,
189 task: &Task,
190 user_message: &Message,
191 agent_message: &Message,
192 context: &RequestContext,
193 ) -> Result<()> {
194 let Err(e) = persist_completed_task(
195 crate::services::a2a_server::processing::message::persistence::PersistCompletedTaskParams {
196 task,
197 user_message,
198 agent_message,
199 context,
200 repositories: &self.repositories,
201 artifacts_already_published: false,
202 },
203 )
204 .await
205 else {
206 return Ok(());
207 };
208
209 let error_msg = format!("Failed to persist completed task: {}", e);
210 tracing::error!(task_id = %task.id, error = %e, "Failed to persist completed task");
211
212 let failed_timestamp = chrono::Utc::now();
213 if let Err(update_err) = self
214 .repositories
215 .tasks
216 .update_task_failed_with_error(&task.id, &error_msg, &failed_timestamp)
217 .await
218 {
219 tracing::error!(task_id = %task.id, error = %update_err, "Failed to update task to failed state");
220 }
221
222 Err(e)
223 }
224}
225
226fn resolve_task_id(message: &Message) -> TaskId {
227 message.task_id.clone().map_or_else(
228 || {
229 let new_task_id = TaskId::new(Uuid::new_v4().to_string());
230 tracing::info!(task_id = %new_task_id, "Starting NEW task with generated ID");
231 new_task_id
232 },
233 |existing_task_id| {
234 tracing::info!(task_id = %existing_task_id, "Continuing existing task");
235 existing_task_id
236 },
237 )
238}
239
240fn new_submitted_task(task_id: &TaskId, context_id: &ContextId, agent_name: &str) -> Task {
241 Task {
242 id: task_id.clone(),
243 context_id: context_id.clone(),
244 status: TaskStatus {
245 state: TaskState::Submitted,
246 message: None,
247 timestamp: Some(chrono::Utc::now()),
248 },
249 history: None,
250 artifacts: None,
251 metadata: Some(TaskMetadata::new_agent_message(agent_name.to_owned())),
252 created_at: Some(chrono::Utc::now()),
253 last_modified: Some(chrono::Utc::now()),
254 }
255}
256
257fn resolve_agent_message(task: &Task, user_message: &Message, response_text: &str) -> Message {
258 task.status.message.clone().unwrap_or_else(|| {
259 let client_message_id = user_message
260 .metadata
261 .as_ref()
262 .and_then(|m| m.get("clientMessageId"))
263 .cloned();
264
265 let metadata = client_message_id.map(|id| serde_json::json!({"clientMessageId": id}));
266
267 Message {
268 role: MessageRole::Agent,
269 parts: vec![Part::Text(TextPart {
270 text: response_text.to_owned(),
271 })],
272 message_id: MessageId::generate(),
273 task_id: Some(task.id.clone()),
274 context_id: task.context_id.clone(),
275 metadata,
276 extensions: None,
277 reference_task_ids: None,
278 }
279 })
280}