Skip to main content

systemprompt_agent/repository/task/
queries.rs

1//! Task read queries.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use crate::models::TaskRow;
7use sqlx::PgPool;
8use std::sync::Arc;
9use systemprompt_database::DbPool;
10use systemprompt_identifiers::{AgentName, ContextId, SessionId, TaskId, TraceId, UserId};
11use systemprompt_traits::RepositoryError;
12
13use super::constructor::TaskConstructor;
14use crate::models::a2a::Task;
15
16pub async fn get_task(
17    pool: &Arc<PgPool>,
18    db_pool: &DbPool,
19    task_id: &TaskId,
20) -> Result<Option<Task>, RepositoryError> {
21    let task_id_str = task_id.as_str();
22    let row = sqlx::query_as!(
23        TaskRow,
24        r#"SELECT
25            task_id as "task_id!: TaskId",
26            context_id as "context_id!: ContextId",
27            status as "status!",
28            status_timestamp,
29            user_id as "user_id?: UserId",
30            session_id as "session_id?: SessionId",
31            trace_id as "trace_id?: TraceId",
32            agent_name as "agent_name?: AgentName",
33            started_at,
34            completed_at,
35            execution_time_ms,
36            error_message,
37            metadata,
38            created_at as "created_at!",
39            updated_at as "updated_at!"
40        FROM agent_tasks WHERE task_id = $1"#,
41        task_id_str
42    )
43    .fetch_optional(pool.as_ref())
44    .await
45    .map_err(RepositoryError::database)?;
46
47    let Some(_row) = row else {
48        return Ok(None);
49    };
50
51    let constructor = TaskConstructor::new(db_pool)?;
52    let task = constructor.construct_task_from_task_id(task_id).await?;
53
54    Ok(Some(task))
55}
56
57pub async fn list_tasks_by_context(
58    pool: &Arc<PgPool>,
59    db_pool: &DbPool,
60    context_id: &ContextId,
61) -> Result<Vec<Task>, RepositoryError> {
62    let context_id_str = context_id.as_str();
63    let rows = sqlx::query_as!(
64        TaskRow,
65        r#"SELECT
66            task_id as "task_id!: TaskId",
67            context_id as "context_id!: ContextId",
68            status as "status!",
69            status_timestamp,
70            user_id as "user_id?: UserId",
71            session_id as "session_id?: SessionId",
72            trace_id as "trace_id?: TraceId",
73            agent_name as "agent_name?: AgentName",
74            started_at,
75            completed_at,
76            execution_time_ms,
77            error_message,
78            metadata,
79            created_at as "created_at!",
80            updated_at as "updated_at!"
81        FROM agent_tasks WHERE context_id = $1 ORDER BY created_at ASC"#,
82        context_id_str
83    )
84    .fetch_all(pool.as_ref())
85    .await
86    .map_err(RepositoryError::database)?;
87
88    let constructor = TaskConstructor::new(db_pool)?;
89    let task_ids: Vec<TaskId> = rows.iter().map(|r| r.task_id.clone()).collect();
90    let tasks = constructor.construct_tasks_batch(&task_ids).await?;
91
92    Ok(tasks)
93}
94
95pub async fn get_tasks_by_user_id(
96    pool: &Arc<PgPool>,
97    db_pool: &DbPool,
98    user_id: &UserId,
99    limit: Option<i32>,
100    offset: Option<i32>,
101) -> Result<Vec<Task>, RepositoryError> {
102    let lim = limit.map_or(1000, i64::from);
103    let off = offset.map_or(0, i64::from);
104    let user_id_str = user_id.as_str();
105
106    let rows = sqlx::query_as!(
107        TaskRow,
108        r#"SELECT
109            task_id as "task_id!: TaskId",
110            context_id as "context_id!: ContextId",
111            status as "status!",
112            status_timestamp,
113            user_id as "user_id?: UserId",
114            session_id as "session_id?: SessionId",
115            trace_id as "trace_id?: TraceId",
116            agent_name as "agent_name?: AgentName",
117            started_at,
118            completed_at,
119            execution_time_ms,
120            error_message,
121            metadata,
122            created_at as "created_at!",
123            updated_at as "updated_at!"
124        FROM agent_tasks WHERE user_id = $1 ORDER BY created_at DESC LIMIT $2 OFFSET $3"#,
125        user_id_str,
126        lim,
127        off
128    )
129    .fetch_all(pool.as_ref())
130    .await
131    .map_err(RepositoryError::database)?;
132
133    let constructor = TaskConstructor::new(db_pool)?;
134    let task_ids: Vec<TaskId> = rows.iter().map(|r| r.task_id.clone()).collect();
135    let tasks = constructor.construct_tasks_batch(&task_ids).await?;
136
137    Ok(tasks)
138}
139
140#[derive(Debug, Clone)]
141pub struct TaskContextInfo {
142    pub context_id: ContextId,
143    pub user_id: Option<UserId>,
144}
145
146pub async fn get_task_context_info(
147    pool: &Arc<PgPool>,
148    task_id: &TaskId,
149) -> Result<Option<TaskContextInfo>, RepositoryError> {
150    let task_id_str = task_id.as_str();
151    let row = sqlx::query!(
152        r#"SELECT
153            context_id as "context_id!: ContextId",
154            user_id as "user_id?: UserId"
155        FROM agent_tasks WHERE task_id = $1"#,
156        task_id_str
157    )
158    .fetch_optional(pool.as_ref())
159    .await
160    .map_err(RepositoryError::database)?;
161
162    Ok(row.map(|r| TaskContextInfo {
163        context_id: r.context_id,
164        user_id: r.user_id,
165    }))
166}