systemprompt_agent/repository/context/
queries.rs1use 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}