Skip to main content

systemprompt_agent/repository/task/
task_updates.rs

1use 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}