Skip to main content

systemprompt_agent/services/a2a_server/processing/message/message_handler/
mod.rs

1//! Non-streaming message handling for [`MessageProcessor`].
2//!
3//! Implements [`MessageProcessor::handle_message`]: it validates the context,
4//! persists a submitted task, runs the stream pipeline to completion, builds
5//! the finished [`Task`](crate::models::a2a::Task), persists it, and broadcasts
6//! the completion and AG-UI lifecycle events.
7//!
8//! Copyright (c) systemprompt.io — Business Source License 1.1.
9//! See <https://systemprompt.io> for licensing details.
10
11mod 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}