systemprompt_agent/repository/agent_service/
mod.rs1use sqlx::PgPool;
8use std::sync::Arc;
9use systemprompt_database::DbPool;
10use systemprompt_traits::RepositoryError;
11
12use crate::error::AgentError;
13
14#[derive(Debug)]
15pub struct AgentServiceRow {
16 pub name: String,
17 pub pid: Option<i32>,
18 pub port: i32,
19 pub status: String,
20}
21
22#[derive(Debug)]
23pub struct AgentServerIdRow {
24 pub name: String,
25}
26
27#[derive(Debug)]
28pub struct AgentServerIdPidRow {
29 pub name: String,
30 pub pid: i32,
31}
32
33#[derive(Debug, Clone)]
34pub struct AgentServiceRepository {
35 pool: Arc<PgPool>,
36 write_pool: Arc<PgPool>,
37}
38
39impl AgentServiceRepository {
40 pub fn new(db: &DbPool) -> Result<Self, AgentError> {
41 let pool = db.pool_arc().map_err(|e| AgentError::Init(e.to_string()))?;
42 let write_pool = db
43 .write_pool_arc()
44 .map_err(|e| AgentError::Init(e.to_string()))?;
45 Ok(Self { pool, write_pool })
46 }
47
48 pub async fn register_agent(
49 &self,
50 name: &str,
51 pid: u32,
52 port: u16,
53 ) -> Result<String, RepositoryError> {
54 self.remove_agent_service(name).await?;
55
56 let pool = &self.write_pool;
57 let pid_i32 = pid as i32;
58 let port_i32 = i32::from(port);
59
60 sqlx::query!(
61 "INSERT INTO services (name, module_name, pid, port, status, updated_at)
62 VALUES ($1, 'agent', $2, $3, 'running', CURRENT_TIMESTAMP)
63 ON CONFLICT (name) DO UPDATE SET pid = $2, port = $3, status = 'running', updated_at \
64 = CURRENT_TIMESTAMP",
65 name,
66 pid_i32,
67 port_i32
68 )
69 .execute(pool.as_ref())
70 .await
71 .map_err(RepositoryError::database)?;
72
73 Ok(name.to_owned())
74 }
75
76 pub async fn register_agent_starting(
77 &self,
78 name: &str,
79 pid: u32,
80 port: u16,
81 ) -> Result<String, RepositoryError> {
82 self.remove_agent_service(name).await?;
83
84 let pool = &self.write_pool;
85 let pid_i32 = pid as i32;
86 let port_i32 = i32::from(port);
87
88 sqlx::query!(
89 "INSERT INTO services (name, module_name, pid, port, status, updated_at)
90 VALUES ($1, 'agent', $2, $3, 'starting', CURRENT_TIMESTAMP)
91 ON CONFLICT (name) DO UPDATE SET pid = $2, port = $3, status = 'starting', updated_at \
92 = CURRENT_TIMESTAMP",
93 name,
94 pid_i32,
95 port_i32
96 )
97 .execute(pool.as_ref())
98 .await
99 .map_err(RepositoryError::database)?;
100
101 Ok(name.to_owned())
102 }
103
104 pub async fn mark_running(&self, agent_name: &str) -> Result<(), RepositoryError> {
105 let pool = &self.write_pool;
106
107 sqlx::query!(
108 "UPDATE services SET status = 'running', updated_at = CURRENT_TIMESTAMP WHERE name = \
109 $1",
110 agent_name
111 )
112 .execute(pool.as_ref())
113 .await
114 .map_err(RepositoryError::database)?;
115
116 Ok(())
117 }
118
119 pub async fn get_agent_status(
120 &self,
121 agent_name: &str,
122 ) -> Result<Option<AgentServiceRow>, RepositoryError> {
123 let pool = &self.pool;
124
125 let row = sqlx::query!(
126 "SELECT name, pid, port, status FROM services WHERE name = $1",
127 agent_name
128 )
129 .fetch_optional(pool.as_ref())
130 .await
131 .map_err(RepositoryError::database)?;
132
133 Ok(row.map(|r| AgentServiceRow {
134 name: r.name,
135 pid: r.pid,
136 port: r.port,
137 status: r.status,
138 }))
139 }
140
141 pub async fn mark_crashed(&self, agent_name: &str) -> Result<(), RepositoryError> {
142 let pool = &self.write_pool;
143
144 sqlx::query!(
145 "UPDATE services SET status = 'error', pid = NULL, updated_at = CURRENT_TIMESTAMP \
146 WHERE name = $1",
147 agent_name
148 )
149 .execute(pool.as_ref())
150 .await
151 .map_err(RepositoryError::database)?;
152
153 Ok(())
154 }
155
156 pub async fn mark_stopped(&self, agent_name: &str) -> Result<(), RepositoryError> {
157 let pool = &self.write_pool;
158
159 sqlx::query!(
160 "UPDATE services SET status = 'stopped', pid = NULL, updated_at = CURRENT_TIMESTAMP \
161 WHERE name = $1",
162 agent_name
163 )
164 .execute(pool.as_ref())
165 .await
166 .map_err(RepositoryError::database)?;
167
168 Ok(())
169 }
170
171 pub async fn mark_error(&self, agent_name: &str) -> Result<(), RepositoryError> {
172 let pool = &self.write_pool;
173
174 sqlx::query!(
175 "UPDATE services SET status = 'error', pid = NULL, updated_at = CURRENT_TIMESTAMP \
176 WHERE name = $1",
177 agent_name
178 )
179 .execute(pool.as_ref())
180 .await
181 .map_err(RepositoryError::database)?;
182
183 Ok(())
184 }
185
186 pub async fn list_running_agents(&self) -> Result<Vec<AgentServerIdRow>, RepositoryError> {
187 let pool = &self.pool;
188
189 let rows = sqlx::query!("SELECT name FROM services WHERE status = 'running'")
190 .fetch_all(pool.as_ref())
191 .await
192 .map_err(RepositoryError::database)?;
193
194 Ok(rows
195 .into_iter()
196 .map(|r| AgentServerIdRow { name: r.name })
197 .collect())
198 }
199
200 pub async fn list_running_agent_pids(
201 &self,
202 ) -> Result<Vec<AgentServerIdPidRow>, RepositoryError> {
203 let pool = &self.pool;
204
205 let rows = sqlx::query!(
206 "SELECT name, pid FROM services WHERE status = 'running' AND pid IS NOT NULL"
207 )
208 .fetch_all(pool.as_ref())
209 .await
210 .map_err(RepositoryError::database)?;
211
212 Ok(rows
213 .into_iter()
214 .filter_map(|r| r.pid.map(|pid| AgentServerIdPidRow { name: r.name, pid }))
215 .collect())
216 }
217
218 pub async fn remove_agent_service(&self, agent_name: &str) -> Result<(), RepositoryError> {
219 let pool = &self.write_pool;
220
221 sqlx::query!("DELETE FROM services WHERE name = $1", agent_name)
222 .execute(pool.as_ref())
223 .await
224 .map_err(RepositoryError::database)?;
225
226 Ok(())
227 }
228
229 pub async fn update_health_status(
230 &self,
231 agent_name: &str,
232 health_status: &str,
233 ) -> Result<(), RepositoryError> {
234 let pool = &self.write_pool;
235
236 sqlx::query!(
237 "UPDATE services SET status = $1, updated_at = CURRENT_TIMESTAMP WHERE name = $2",
238 health_status,
239 agent_name
240 )
241 .execute(pool.as_ref())
242 .await
243 .map_err(RepositoryError::database)?;
244
245 Ok(())
246 }
247}