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