systemprompt_agent/repository/task/
task_updates.rs1use super::{TaskRepository, task_state_to_db_string};
2use crate::models::a2a::{Message, Task, TaskState};
3use crate::repository::context::message::{
4 FileUploadContext, PersistMessageSqlxParams, get_next_sequence_number_sqlx,
5 persist_message_sqlx,
6};
7use systemprompt_traits::RepositoryError;
8
9#[expect(
10 missing_debug_implementations,
11 reason = "params struct holds non-Debug references"
12)]
13pub struct UpdateTaskAndSaveMessagesParams<'a> {
14 pub task: &'a Task,
15 pub user_message: &'a Message,
16 pub agent_message: &'a Message,
17 pub user_id: Option<&'a systemprompt_identifiers::UserId>,
18 pub session_id: &'a systemprompt_identifiers::SessionId,
19 pub trace_id: &'a systemprompt_identifiers::TraceId,
20}
21
22impl TaskRepository {
23 pub async fn update_task_and_save_messages(
24 &self,
25 params: UpdateTaskAndSaveMessagesParams<'_>,
26 ) -> Result<Task, RepositoryError> {
27 let UpdateTaskAndSaveMessagesParams {
28 task,
29 user_message,
30 agent_message,
31 user_id,
32 session_id,
33 trace_id,
34 } = params;
35 let mut tx = self
36 .write_pool
37 .begin()
38 .await
39 .map_err(RepositoryError::database)?;
40
41 let status = task_state_to_db_string(task.status.state);
42 let metadata_json = task
43 .metadata
44 .as_ref().map_or_else(|| serde_json::json!({}), |m| {
45 serde_json::to_value(m).unwrap_or_else(|e| {
46 tracing::warn!(error = %e, task_id = %task.id, "Failed to serialize task metadata");
47 serde_json::json!({})
48 })
49 });
50
51 let task_id_str = task.id.as_str();
52 let is_completed = task.status.state == TaskState::Completed;
53
54 let result = if is_completed {
55 sqlx::query!(
56 r#"UPDATE agent_tasks SET
57 status = $1,
58 status_timestamp = $2,
59 metadata = $3,
60 updated_at = CURRENT_TIMESTAMP,
61 completed_at = CURRENT_TIMESTAMP,
62 started_at = COALESCE(started_at, CURRENT_TIMESTAMP),
63 execution_time_ms = EXTRACT(EPOCH FROM (CURRENT_TIMESTAMP - COALESCE(started_at, CURRENT_TIMESTAMP))) * 1000
64 WHERE task_id = $4"#,
65 status,
66 task.status.timestamp,
67 metadata_json,
68 task_id_str
69 )
70 .execute(&mut *tx)
71 .await
72 .map_err(RepositoryError::database)?
73 } else {
74 sqlx::query!(
75 r#"UPDATE agent_tasks SET status = $1, status_timestamp = $2, metadata = $3, updated_at = CURRENT_TIMESTAMP WHERE task_id = $4"#,
76 status,
77 task.status.timestamp,
78 metadata_json,
79 task_id_str
80 )
81 .execute(&mut *tx)
82 .await
83 .map_err(RepositoryError::database)?
84 };
85
86 if result.rows_affected() == 0 {
87 return Err(RepositoryError::NotFound(format!(
88 "Task not found for update: {}",
89 task.id
90 )));
91 }
92
93 let context_id_ref = &task.context_id;
94 let upload_ctx = self
95 .file_upload_provider
96 .as_ref()
97 .map(|svc| FileUploadContext {
98 upload_provider: svc,
99 context_id: context_id_ref,
100 user_id,
101 session_id: Some(session_id),
102 trace_id: Some(trace_id),
103 });
104
105 let user_seq = get_next_sequence_number_sqlx(&mut tx, &task.id).await?;
106 persist_message_sqlx(PersistMessageSqlxParams {
107 tx: &mut tx,
108 message: user_message,
109 task_id: &task.id,
110 context_id: context_id_ref,
111 sequence_number: user_seq,
112 user_id,
113 session_id,
114 trace_id,
115 upload_ctx: upload_ctx.as_ref(),
116 })
117 .await?;
118
119 let agent_seq = get_next_sequence_number_sqlx(&mut tx, &task.id).await?;
120 persist_message_sqlx(PersistMessageSqlxParams {
121 tx: &mut tx,
122 message: agent_message,
123 task_id: &task.id,
124 context_id: context_id_ref,
125 sequence_number: agent_seq,
126 user_id,
127 session_id,
128 trace_id,
129 upload_ctx: upload_ctx.as_ref(),
130 })
131 .await?;
132
133 tx.commit().await.map_err(RepositoryError::database)?;
134
135 if let Some(ref analytics_provider) = self.session_analytics_provider {
136 for _ in 0..2 {
137 if let Err(e) = analytics_provider.increment_message_count(session_id).await {
138 tracing::warn!(error = %e, "Failed to increment analytics message count");
139 }
140 }
141 }
142
143 let updated_task = self.get_task(&task.id).await?.ok_or_else(|| {
144 RepositoryError::NotFound(format!("Task not found after update: {}", task.id))
145 })?;
146
147 Ok(updated_task)
148 }
149
150 pub async fn delete_task(
151 &self,
152 task_id: &systemprompt_identifiers::TaskId,
153 ) -> Result<(), RepositoryError> {
154 let task_id_str = task_id.as_str();
155
156 sqlx::query!(
157 "DELETE FROM message_parts WHERE message_id IN (SELECT message_id FROM task_messages \
158 WHERE task_id = $1)",
159 task_id_str
160 )
161 .execute(&*self.write_pool)
162 .await
163 .map_err(RepositoryError::database)?;
164
165 sqlx::query!("DELETE FROM task_messages WHERE task_id = $1", task_id_str)
166 .execute(&*self.write_pool)
167 .await
168 .map_err(RepositoryError::database)?;
169
170 sqlx::query!(
171 "DELETE FROM task_execution_steps WHERE task_id = $1",
172 task_id_str
173 )
174 .execute(&*self.write_pool)
175 .await
176 .map_err(RepositoryError::database)?;
177
178 sqlx::query!("DELETE FROM agent_tasks WHERE task_id = $1", task_id_str)
179 .execute(&*self.write_pool)
180 .await
181 .map_err(RepositoryError::database)?;
182
183 Ok(())
184 }
185}