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//!
4//! Copyright (c) systemprompt.io — Business Source License 1.1.
5//! See <https://systemprompt.io> for licensing details.
6
7use 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}