Skip to main content

agentforge_db/
agent_repo.rs

1use crate::db_err;
2use agentforge_core::{AgentFile, AgentFileFormat, AgentForgeError, AgentVersion, Result};
3use chrono::Utc;
4use sqlx::PgPool;
5use uuid::Uuid;
6
7pub struct AgentRepo {
8    pool: PgPool,
9}
10
11impl AgentRepo {
12    pub fn new(pool: PgPool) -> Self {
13        Self { pool }
14    }
15
16    /// Insert a new agent version. Returns error if SHA already exists.
17    pub async fn insert(&self, version: &AgentVersion) -> Result<AgentVersion> {
18        let format_str = version.format.to_string();
19        let file_content_json = serde_json::to_value(&version.file_content)
20            .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
21
22        let row = sqlx::query!(
23            r#"
24            INSERT INTO agent_versions
25                (id, name, version, sha, file_content, raw_content, format,
26                 promoted, is_champion, changelog, parent_sha, created_at, updated_at)
27            VALUES
28                ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13)
29            RETURNING id, name, version, sha, file_content, raw_content, format,
30                      promoted, is_champion, changelog, parent_sha, created_at, updated_at
31            "#,
32            version.id,
33            version.name,
34            version.version,
35            version.sha,
36            file_content_json,
37            version.raw_content,
38            format_str,
39            version.promoted,
40            version.is_champion,
41            version.changelog,
42            version.parent_sha,
43            Utc::now(),
44            Utc::now(),
45        )
46        .fetch_one(&self.pool)
47        .await
48        .map_err(db_err)?;
49
50        self.row_to_version(
51            row.id,
52            row.name,
53            row.version,
54            row.sha,
55            row.file_content,
56            row.raw_content,
57            row.format,
58            row.promoted,
59            row.is_champion,
60            row.changelog,
61            row.parent_sha,
62            row.created_at,
63            row.updated_at,
64        )
65    }
66
67    pub async fn find_by_id(&self, id: Uuid) -> Result<AgentVersion> {
68        let row = sqlx::query!(
69            r#"
70            SELECT id, name, version, sha, file_content, raw_content, format,
71                   promoted, is_champion, changelog, parent_sha, created_at, updated_at
72            FROM agent_versions
73            WHERE id = $1
74            "#,
75            id
76        )
77        .fetch_optional(&self.pool)
78        .await
79        .map_err(db_err)?
80        .ok_or_else(|| AgentForgeError::NotFound {
81            resource: "AgentVersion",
82            id: id.to_string(),
83        })?;
84
85        self.row_to_version(
86            row.id,
87            row.name,
88            row.version,
89            row.sha,
90            row.file_content,
91            row.raw_content,
92            row.format,
93            row.promoted,
94            row.is_champion,
95            row.changelog,
96            row.parent_sha,
97            row.created_at,
98            row.updated_at,
99        )
100    }
101
102    pub async fn find_by_sha(&self, sha: &str) -> Result<Option<AgentVersion>> {
103        let row = sqlx::query!(
104            r#"
105            SELECT id, name, version, sha, file_content, raw_content, format,
106                   promoted, is_champion, changelog, parent_sha, created_at, updated_at
107            FROM agent_versions
108            WHERE sha = $1
109            "#,
110            sha
111        )
112        .fetch_optional(&self.pool)
113        .await
114        .map_err(db_err)?;
115
116        match row {
117            Some(r) => Ok(Some(self.row_to_version(
118                r.id,
119                r.name,
120                r.version,
121                r.sha,
122                r.file_content,
123                r.raw_content,
124                r.format,
125                r.promoted,
126                r.is_champion,
127                r.changelog,
128                r.parent_sha,
129                r.created_at,
130                r.updated_at,
131            )?)),
132            None => Ok(None),
133        }
134    }
135
136    pub async fn list_by_name(&self, name: &str) -> Result<Vec<AgentVersion>> {
137        let rows = sqlx::query!(
138            r#"
139            SELECT id, name, version, sha, file_content, raw_content, format,
140                   promoted, is_champion, changelog, parent_sha, created_at, updated_at
141            FROM agent_versions
142            WHERE name = $1
143            ORDER BY created_at DESC
144            "#,
145            name
146        )
147        .fetch_all(&self.pool)
148        .await
149        .map_err(db_err)?;
150
151        rows.into_iter()
152            .map(|r| {
153                self.row_to_version(
154                    r.id,
155                    r.name,
156                    r.version,
157                    r.sha,
158                    r.file_content,
159                    r.raw_content,
160                    r.format,
161                    r.promoted,
162                    r.is_champion,
163                    r.changelog,
164                    r.parent_sha,
165                    r.created_at,
166                    r.updated_at,
167                )
168            })
169            .collect()
170    }
171
172    pub async fn list_all(&self, limit: i64, offset: i64) -> Result<Vec<AgentVersion>> {
173        let rows = sqlx::query!(
174            r#"
175            SELECT id, name, version, sha, file_content, raw_content, format,
176                   promoted, is_champion, changelog, parent_sha, created_at, updated_at
177            FROM agent_versions
178            ORDER BY created_at DESC
179            LIMIT $1 OFFSET $2
180            "#,
181            limit,
182            offset
183        )
184        .fetch_all(&self.pool)
185        .await
186        .map_err(db_err)?;
187
188        rows.into_iter()
189            .map(|r| {
190                self.row_to_version(
191                    r.id,
192                    r.name,
193                    r.version,
194                    r.sha,
195                    r.file_content,
196                    r.raw_content,
197                    r.format,
198                    r.promoted,
199                    r.is_champion,
200                    r.changelog,
201                    r.parent_sha,
202                    r.created_at,
203                    r.updated_at,
204                )
205            })
206            .collect()
207    }
208
209    pub async fn set_champion(&self, agent_id: Uuid, agent_name: &str) -> Result<()> {
210        let mut tx = self.pool.begin().await.map_err(db_err)?;
211
212        // Clear existing champion for this agent name
213        sqlx::query!(
214            "UPDATE agent_versions SET is_champion = FALSE WHERE name = $1",
215            agent_name
216        )
217        .execute(&mut *tx)
218        .await
219        .map_err(db_err)?;
220
221        // Set the new champion
222        sqlx::query!(
223            "UPDATE agent_versions SET is_champion = TRUE, promoted = TRUE WHERE id = $1",
224            agent_id
225        )
226        .execute(&mut *tx)
227        .await
228        .map_err(db_err)?;
229
230        tx.commit().await.map_err(db_err)?;
231        Ok(())
232    }
233
234    pub async fn update_changelog(&self, agent_id: Uuid, changelog: &str) -> Result<()> {
235        sqlx::query!(
236            "UPDATE agent_versions SET changelog = $1 WHERE id = $2",
237            changelog,
238            agent_id
239        )
240        .execute(&self.pool)
241        .await
242        .map_err(db_err)?;
243        Ok(())
244    }
245
246    pub async fn delete(&self, id: Uuid) -> Result<bool> {
247        let result = sqlx::query("DELETE FROM agent_versions WHERE id = $1")
248            .bind(id)
249            .execute(&self.pool)
250            .await
251            .map_err(db_err)?;
252        Ok(result.rows_affected() > 0)
253    }
254
255    #[allow(clippy::too_many_arguments)]
256    fn row_to_version(
257        &self,
258        id: Uuid,
259        name: String,
260        version: String,
261        sha: String,
262        file_content: serde_json::Value,
263        raw_content: String,
264        format: String,
265        promoted: bool,
266        is_champion: bool,
267        changelog: Option<String>,
268        parent_sha: Option<String>,
269        created_at: chrono::DateTime<Utc>,
270        updated_at: chrono::DateTime<Utc>,
271    ) -> Result<AgentVersion> {
272        use std::str::FromStr;
273        let parsed_format = AgentFileFormat::from_str(&format)?;
274        let parsed_file: AgentFile = serde_json::from_value(file_content)
275            .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
276
277        Ok(AgentVersion {
278            id,
279            name,
280            version,
281            sha,
282            file_content: parsed_file,
283            raw_content,
284            format: parsed_format,
285            promoted,
286            is_champion,
287            changelog,
288            parent_sha,
289            created_at,
290            updated_at,
291        })
292    }
293}