Skip to main content

systemprompt_agent/repository/task/
task_updates.rs

1//! Transactional task completion: message history and the guarded state
2//! transition land in one transaction, so a task is never marked complete
3//! without its messages.
4//!
5//! Copyright (c) systemprompt.io — Business Source License 1.1.
6//! See <https://systemprompt.io> for licensing details.
7
8use 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, &timestamp).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, &params).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}