systemprompt_agent/repository/task/
task_updates.rs1use super::TaskRepository;
9use super::state::transition_in_tx;
10use crate::models::a2a::{Message, Task};
11use crate::repository::context::message::{
12 PersistMessageSqlxParams, get_next_sequence_number_sqlx, persist_message_sqlx,
13};
14use systemprompt_identifiers::{ContextId, SessionId, TaskId, TraceId, UserId};
15use systemprompt_traits::RepositoryError;
16
17#[expect(
18 missing_debug_implementations,
19 reason = "params struct holds non-Debug references"
20)]
21pub struct UpdateTaskAndSaveMessagesParams<'a> {
22 pub task: &'a Task,
23 pub user_message: &'a Message,
24 pub agent_message: &'a Message,
25 pub user_id: Option<&'a UserId>,
26 pub session_id: &'a SessionId,
27 pub trace_id: &'a TraceId,
28}
29
30#[expect(
31 missing_debug_implementations,
32 reason = "params struct holds non-Debug references"
33)]
34pub struct PersistMessagesTxParams<'a> {
35 pub task_id: &'a TaskId,
36 pub context_id: &'a ContextId,
37 pub messages: &'a [Message],
38 pub user_id: Option<&'a UserId>,
39 pub session_id: &'a SessionId,
40 pub trace_id: &'a TraceId,
41}
42
43impl TaskRepository {
44 pub async fn update_task_and_save_messages(
45 &self,
46 params: UpdateTaskAndSaveMessagesParams<'_>,
47 ) -> Result<Task, RepositoryError> {
48 let UpdateTaskAndSaveMessagesParams {
49 task,
50 user_message,
51 agent_message,
52 user_id,
53 session_id,
54 trace_id,
55 } = params;
56 let messages = [user_message.clone(), agent_message.clone()];
57 let mut tx = self
58 .write_pool
59 .begin()
60 .await
61 .map_err(RepositoryError::database)?;
62
63 persist_messages_in_tx(
64 &mut tx,
65 &PersistMessagesTxParams {
66 task_id: &task.id,
67 context_id: &task.context_id,
68 messages: &messages,
69 user_id,
70 session_id,
71 trace_id,
72 },
73 )
74 .await?;
75 update_task_metadata(&mut tx, task).await?;
76 let timestamp = task.status.timestamp.unwrap_or_else(chrono::Utc::now);
77 transition_in_tx(&mut tx, &task.id, task.status.state, ×tamp).await?;
78
79 tx.commit().await.map_err(RepositoryError::database)?;
80
81 self.count_messages(session_id, messages.len()).await;
82
83 self.get_task(&task.id).await?.ok_or_else(|| {
84 RepositoryError::NotFound(format!("Task not found after update: {}", task.id))
85 })
86 }
87
88 pub async fn persist_messages(
89 &self,
90 params: PersistMessagesTxParams<'_>,
91 ) -> Result<Vec<i32>, RepositoryError> {
92 let mut tx = self
93 .write_pool
94 .begin()
95 .await
96 .map_err(RepositoryError::database)?;
97 let sequence_numbers = persist_messages_in_tx(&mut tx, ¶ms).await?;
98 tx.commit().await.map_err(RepositoryError::database)?;
99
100 self.count_messages(params.session_id, params.messages.len())
101 .await;
102 Ok(sequence_numbers)
103 }
104
105 async fn count_messages(&self, session_id: &SessionId, count: usize) {
106 for _ in 0..count {
107 if let Err(e) = self.sessions.increment_message_count(session_id).await {
108 tracing::warn!(error = %e, session_id = %session_id, "Failed to increment session message count");
109 }
110 }
111 }
112
113 pub async fn delete_task(&self, task_id: &TaskId) -> Result<(), RepositoryError> {
114 let task_id_str = task_id.as_str();
115 let mut tx = self
116 .write_pool
117 .begin()
118 .await
119 .map_err(RepositoryError::database)?;
120
121 sqlx::query!(
122 "DELETE FROM message_parts WHERE message_id IN (SELECT message_id FROM task_messages \
123 WHERE task_id = $1)",
124 task_id_str
125 )
126 .execute(&mut *tx)
127 .await
128 .map_err(RepositoryError::database)?;
129
130 sqlx::query!("DELETE FROM task_messages WHERE task_id = $1", task_id_str)
131 .execute(&mut *tx)
132 .await
133 .map_err(RepositoryError::database)?;
134
135 sqlx::query!(
136 "DELETE FROM task_execution_steps WHERE task_id = $1",
137 task_id_str
138 )
139 .execute(&mut *tx)
140 .await
141 .map_err(RepositoryError::database)?;
142
143 sqlx::query!("DELETE FROM agent_tasks WHERE task_id = $1", task_id_str)
144 .execute(&mut *tx)
145 .await
146 .map_err(RepositoryError::database)?;
147
148 tx.commit().await.map_err(RepositoryError::database)
149 }
150}
151
152async fn persist_messages_in_tx(
153 tx: &mut sqlx::Transaction<'static, sqlx::Postgres>,
154 params: &PersistMessagesTxParams<'_>,
155) -> Result<Vec<i32>, RepositoryError> {
156 let mut sequence_numbers = Vec::with_capacity(params.messages.len());
157 for message in params.messages {
158 let sequence_number = get_next_sequence_number_sqlx(tx, params.task_id).await?;
159 persist_message_sqlx(PersistMessageSqlxParams {
160 tx,
161 message,
162 task_id: params.task_id,
163 context_id: params.context_id,
164 sequence_number,
165 user_id: params.user_id,
166 session_id: params.session_id,
167 trace_id: params.trace_id,
168 })
169 .await?;
170 sequence_numbers.push(sequence_number);
171 }
172 Ok(sequence_numbers)
173}
174
175async fn update_task_metadata(
176 tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
177 task: &Task,
178) -> Result<(), RepositoryError> {
179 let metadata_json = match &task.metadata {
180 Some(metadata) => serde_json::to_value(metadata).map_err(RepositoryError::Serialization)?,
181 None => serde_json::json!({}),
182 };
183
184 let result = sqlx::query!(
185 r#"UPDATE agent_tasks SET metadata = $1, updated_at = CURRENT_TIMESTAMP WHERE task_id = $2"#,
186 metadata_json,
187 task.id.as_str()
188 )
189 .execute(&mut **tx)
190 .await
191 .map_err(RepositoryError::database)?;
192
193 if result.rows_affected() == 0 {
194 return Err(RepositoryError::NotFound(format!(
195 "Task not found for update: {}",
196 task.id
197 )));
198 }
199
200 Ok(())
201}