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 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 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 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}