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