systemprompt_agent/services/
message.rs1use 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}