Skip to main content

systemprompt_agent/repository/agent_service/
mod.rs

1//! Repository for declared agent services (named processes registered with the
2//! platform).
3
4use 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}