Skip to main content

systemprompt_agent/repository/task/
task_updates.rs

1//! Task status/state update mutations.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use 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}