Skip to main content

systemprompt_agent/repository/context/
queries.rs

1use chrono::{DateTime, Utc};
2
3use super::ContextRepository;
4use crate::models::context::{ContextStateEvent, UserContext, UserContextWithStats};
5use crate::repository::task::constructor::TaskConstructor;
6use systemprompt_identifiers::{ContextId, SessionId, TaskId, UserId};
7use systemprompt_traits::RepositoryError;
8
9impl ContextRepository {
10    pub async fn find_user_id_for_context(
11        &self,
12        context_id: &ContextId,
13    ) -> Result<Option<UserId>, RepositoryError> {
14        let row = sqlx::query_scalar!(
15            r#"SELECT user_id FROM user_contexts WHERE context_id = $1"#,
16            context_id.as_str(),
17        )
18        .fetch_optional(&*self.pool)
19        .await
20        .map_err(RepositoryError::database)?;
21        Ok(row.map(UserId::new))
22    }
23
24    pub async fn get_context(
25        &self,
26        context_id: &ContextId,
27        user_id: &UserId,
28    ) -> Result<UserContext, RepositoryError> {
29        let row = sqlx::query!(
30            r#"SELECT
31                context_id as "context_id!",
32                user_id as "user_id!",
33                name as "name!",
34                created_at as "created_at!",
35                updated_at as "updated_at!"
36            FROM user_contexts WHERE context_id = $1 AND user_id = $2"#,
37            context_id.as_str(),
38            user_id.as_str()
39        )
40        .fetch_one(&*self.pool)
41        .await
42        .map_err(|e| match e {
43            sqlx::Error::RowNotFound => RepositoryError::NotFound(format!(
44                "Context {} not found for user {}",
45                context_id, user_id
46            )),
47            _ => RepositoryError::database(e),
48        })?;
49
50        Ok(UserContext {
51            context_id: ContextId::new(row.context_id),
52            user_id: UserId::new(row.user_id),
53            name: row.name,
54            created_at: row.created_at,
55            updated_at: row.updated_at,
56        })
57    }
58
59    pub async fn list_contexts_basic(
60        &self,
61        user_id: &UserId,
62    ) -> Result<Vec<UserContext>, RepositoryError> {
63        let rows = sqlx::query!(
64            r#"SELECT
65                context_id as "context_id!",
66                user_id as "user_id!",
67                name as "name!",
68                created_at as "created_at!",
69                updated_at as "updated_at!"
70            FROM user_contexts WHERE user_id = $1 ORDER BY updated_at DESC"#,
71            user_id.as_str()
72        )
73        .fetch_all(&*self.pool)
74        .await
75        .map_err(RepositoryError::database)?;
76
77        Ok(rows
78            .into_iter()
79            .map(|r| UserContext {
80                context_id: ContextId::new(r.context_id),
81                user_id: UserId::new(r.user_id),
82                name: r.name,
83                created_at: r.created_at,
84                updated_at: r.updated_at,
85            })
86            .collect())
87    }
88
89    pub async fn list_contexts_with_stats(
90        &self,
91        user_id: &UserId,
92    ) -> Result<Vec<UserContextWithStats>, RepositoryError> {
93        let rows = sqlx::query!(
94            r#"SELECT
95                c.context_id as "context_id!",
96                c.user_id as "user_id!",
97                c.name as "name!",
98                c.created_at as "created_at!",
99                c.updated_at as "updated_at!",
100                COALESCE(COUNT(DISTINCT t.task_id), 0)::bigint as "task_count!",
101                COALESCE(COUNT(DISTINCT m.id), 0)::bigint as "message_count!",
102                MAX(m.created_at) as last_message_at
103            FROM user_contexts c
104            LEFT JOIN agent_tasks t ON t.context_id = c.context_id
105            LEFT JOIN task_messages m ON m.task_id = t.task_id
106            WHERE c.user_id = $1
107            GROUP BY c.context_id
108            ORDER BY c.updated_at DESC"#,
109            user_id.as_str()
110        )
111        .fetch_all(&*self.pool)
112        .await
113        .map_err(RepositoryError::database)?;
114
115        Ok(rows
116            .into_iter()
117            .map(|r| UserContextWithStats {
118                context_id: ContextId::new(r.context_id),
119                user_id: UserId::new(r.user_id),
120                name: r.name,
121                created_at: r.created_at,
122                updated_at: r.updated_at,
123                task_count: r.task_count,
124                message_count: r.message_count,
125                last_message_at: r.last_message_at,
126            })
127            .collect())
128    }
129
130    pub async fn find_by_session_id(
131        &self,
132        session_id: &SessionId,
133    ) -> Result<Option<UserContext>, RepositoryError> {
134        let row = sqlx::query!(
135            r#"SELECT
136                context_id as "context_id!",
137                user_id as "user_id!",
138                name as "name!",
139                created_at as "created_at!",
140                updated_at as "updated_at!"
141            FROM user_contexts WHERE session_id = $1
142            ORDER BY created_at DESC LIMIT 1"#,
143            session_id.as_str()
144        )
145        .fetch_optional(&*self.pool)
146        .await
147        .map_err(RepositoryError::database)?;
148
149        Ok(row.map(|r| UserContext {
150            context_id: ContextId::new(r.context_id),
151            user_id: UserId::new(r.user_id),
152            name: r.name,
153            created_at: r.created_at,
154            updated_at: r.updated_at,
155        }))
156    }
157
158    pub async fn get_context_events_since(
159        &self,
160        context_id: &ContextId,
161        last_seen: DateTime<Utc>,
162    ) -> Result<Vec<ContextStateEvent>, RepositoryError> {
163        let mut events = Vec::new();
164
165        let task_ids: Vec<String> = sqlx::query_scalar!(
166            r#"SELECT t.task_id as "task_id!" FROM agent_tasks t
167             WHERE t.context_id = $1 AND t.updated_at > $2
168             ORDER BY t.updated_at ASC"#,
169            context_id.as_str(),
170            last_seen
171        )
172        .fetch_all(&*self.pool)
173        .await
174        .map_err(RepositoryError::database)?;
175
176        if !task_ids.is_empty() {
177            let constructor = TaskConstructor::new(&self.db_pool)?;
178            let task_ids_typed: Vec<TaskId> = task_ids.iter().map(TaskId::new).collect();
179            let tasks = constructor.construct_tasks_batch(&task_ids_typed).await?;
180
181            for task in tasks {
182                events.push(ContextStateEvent::TaskStatusChanged {
183                    task,
184                    context_id: context_id.clone(),
185                    timestamp: Utc::now(),
186                });
187            }
188        }
189
190        let context_updates = sqlx::query!(
191            r#"SELECT
192                context_id as "context_id!",
193                name as "name!",
194                updated_at as "updated_at!"
195            FROM user_contexts
196            WHERE context_id = $1 AND updated_at > $2
197            ORDER BY updated_at ASC"#,
198            context_id.as_str(),
199            last_seen
200        )
201        .fetch_all(&*self.pool)
202        .await
203        .map_err(RepositoryError::database)?;
204
205        for row in context_updates {
206            events.push(ContextStateEvent::ContextUpdated {
207                context_id: ContextId::new(row.context_id),
208                name: row.name,
209                timestamp: row.updated_at,
210            });
211        }
212
213        events.sort_by_key(ContextStateEvent::timestamp);
214
215        Ok(events)
216    }
217}