Skip to main content

systemprompt_agent/services/
message.rs

1//! Persisting A2A conversation messages, including transactional writes and
2//! synthetic messages for MCP tool executions.
3
4use crate::services::shared::{AgentServiceError, Result};
5use serde_json::json;
6use uuid::Uuid;
7
8use crate::models::a2a::{Message, MessageRole, Part, TextPart};
9use crate::repository::context::message::PersistMessageWithTxParams;
10use crate::repository::task::TaskRepository;
11use systemprompt_database::{DatabaseProvider, DatabaseTransaction, DbPool};
12use systemprompt_identifiers::{ContextId, MessageId, TaskId};
13use systemprompt_models::RequestContext;
14
15pub struct PersistMessageInTxParams<'a> {
16    pub tx: &'a mut dyn DatabaseTransaction,
17    pub message: &'a Message,
18    pub task_id: &'a TaskId,
19    pub context_id: &'a ContextId,
20    pub user_id: Option<&'a systemprompt_identifiers::UserId>,
21    pub session_id: &'a systemprompt_identifiers::SessionId,
22    pub trace_id: &'a systemprompt_identifiers::TraceId,
23}
24
25impl std::fmt::Debug for PersistMessageInTxParams<'_> {
26    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
27        f.debug_struct("PersistMessageInTxParams")
28            .field("message", &self.message)
29            .field("task_id", &self.task_id)
30            .field("context_id", &self.context_id)
31            .field("user_id", &self.user_id)
32            .field("session_id", &self.session_id)
33            .field("trace_id", &self.trace_id)
34            .finish_non_exhaustive()
35    }
36}
37
38#[derive(Debug)]
39pub struct PersistMessagesParams<'a> {
40    pub task_id: &'a TaskId,
41    pub context_id: &'a ContextId,
42    pub messages: Vec<Message>,
43    pub user_id: Option<&'a systemprompt_identifiers::UserId>,
44    pub session_id: &'a systemprompt_identifiers::SessionId,
45    pub trace_id: &'a systemprompt_identifiers::TraceId,
46}
47
48#[derive(Debug)]
49pub struct CreateToolExecutionMessageParams<'a> {
50    pub task_id: &'a TaskId,
51    pub context_id: &'a ContextId,
52    pub tool_name: &'a str,
53    pub tool_args: &'a serde_json::Value,
54    pub request_context: &'a RequestContext,
55}
56
57pub struct MessageService {
58    task_repo: TaskRepository,
59}
60
61impl std::fmt::Debug for MessageService {
62    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
63        f.debug_struct("MessageService").finish_non_exhaustive()
64    }
65}
66
67impl MessageService {
68    pub fn new(db_pool: &DbPool) -> Result<Self> {
69        Ok(Self {
70            task_repo: TaskRepository::new(db_pool)?,
71        })
72    }
73
74    pub async fn persist_message_in_tx(&self, params: PersistMessageInTxParams<'_>) -> Result<i32> {
75        let PersistMessageInTxParams {
76            tx,
77            message,
78            task_id,
79            context_id,
80            user_id,
81            session_id,
82            trace_id,
83        } = params;
84        let sequence_number = self
85            .task_repo
86            .get_next_sequence_number_in_tx(tx, task_id)
87            .await?;
88
89        self.task_repo
90            .persist_message_with_tx(PersistMessageWithTxParams {
91                tx,
92                message,
93                task_id,
94                context_id,
95                sequence_number,
96                user_id,
97                session_id,
98                trace_id,
99            })
100            .await
101            .map_err(|e| {
102                AgentServiceError::Internal(format!("Failed to persist message: {}", e))
103            })?;
104
105        tracing::info!(
106            message_id = %message.message_id,
107            task_id = %task_id,
108            sequence_number = sequence_number,
109            "Message persisted"
110        );
111
112        Ok(sequence_number)
113    }
114
115    pub async fn persist_messages(&self, params: PersistMessagesParams<'_>) -> Result<Vec<i32>> {
116        let PersistMessagesParams {
117            task_id,
118            context_id,
119            messages,
120            user_id,
121            session_id,
122            trace_id,
123        } = params;
124
125        if messages.is_empty() {
126            return Ok(Vec::new());
127        }
128
129        let mut tx = self
130            .task_repo
131            .db_pool()
132            .as_ref()
133            .begin_transaction()
134            .await?;
135        let mut sequence_numbers = Vec::new();
136
137        tracing::info!(
138            task_id = %task_id,
139            message_count = messages.len(),
140            "Persisting multiple messages"
141        );
142
143        for message in messages {
144            let seq = self
145                .persist_message_in_tx(PersistMessageInTxParams {
146                    tx: &mut *tx,
147                    message: &message,
148                    task_id,
149                    context_id,
150                    user_id,
151                    session_id,
152                    trace_id,
153                })
154                .await?;
155            sequence_numbers.push(seq);
156        }
157
158        tx.commit().await?;
159
160        tracing::info!(
161            task_id = %task_id,
162            sequence_numbers = ?sequence_numbers,
163            "Messages persisted successfully"
164        );
165
166        Ok(sequence_numbers)
167    }
168
169    pub async fn create_tool_execution_message(
170        &self,
171        params: CreateToolExecutionMessageParams<'_>,
172    ) -> Result<(String, i32)> {
173        let CreateToolExecutionMessageParams {
174            task_id,
175            context_id,
176            tool_name,
177            tool_args,
178            request_context,
179        } = params;
180        let message_id = Uuid::new_v4().to_string();
181
182        let tool_args_display =
183            serde_json::to_string_pretty(tool_args).unwrap_or_else(|_| tool_args.to_string());
184
185        let timestamp = chrono::Utc::now().to_rfc3339();
186
187        let message = Message {
188            role: MessageRole::User,
189            message_id: MessageId::new(message_id.clone()),
190            task_id: Some(task_id.clone()),
191            context_id: context_id.clone(),
192            parts: vec![Part::Text(TextPart {
193                text: format!(
194                    "Executed MCP tool: {} with arguments:\n{}\n\nExecution ID: {} at {}",
195                    tool_name,
196                    tool_args_display,
197                    task_id.as_str(),
198                    timestamp
199                ),
200            })],
201            metadata: Some(json!({
202                "source": "mcp_direct_call",
203                "tool_name": tool_name,
204                "is_synthetic": true,
205                "tool_args": tool_args,
206                "execution_timestamp": timestamp,
207            })),
208            extensions: None,
209            reference_task_ids: None,
210        };
211
212        let mut tx = self
213            .task_repo
214            .db_pool()
215            .as_ref()
216            .begin_transaction()
217            .await?;
218
219        let sequence_number = self
220            .persist_message_in_tx(PersistMessageInTxParams {
221                tx: &mut *tx,
222                message: &message,
223                task_id,
224                context_id,
225                user_id: Some(request_context.user_id()),
226                session_id: request_context.session_id(),
227                trace_id: request_context.trace_id(),
228            })
229            .await?;
230
231        tx.commit().await?;
232
233        tracing::info!(
234            message_id = %message_id,
235            task_id = %task_id,
236            tool_name = %tool_name,
237            sequence_number = sequence_number,
238            "Created synthetic tool execution message"
239        );
240
241        Ok((message_id, sequence_number))
242    }
243}