1use crate::db_err;
2use agentforge_core::{
3 AgentForgeError, DimensionScores, EvalRun, EvalRunStatus, FailureClusterSummary, Result,
4};
5use chrono::{DateTime, Utc};
6use sqlx::{PgPool, Row};
7use uuid::Uuid;
8
9pub struct EvalRepo {
10 pool: PgPool,
11}
12
13impl EvalRepo {
14 pub fn new(pool: PgPool) -> Self {
15 Self { pool }
16 }
17
18 pub async fn insert(&self, run: &EvalRun) -> Result<EvalRun> {
19 let status_str = run.status.to_string();
20 let _clusters_json = run
21 .failure_clusters
22 .as_ref()
23 .map(serde_json::to_value)
24 .transpose()
25 .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
26
27 sqlx::query(
28 r#"
29 INSERT INTO eval_runs
30 (id, agent_id, scenario_set_id, status, scenario_count,
31 completed_count, error_count, seed, concurrency, created_at, updated_at)
32 VALUES ($1, $2, $3, $4::eval_run_status, $5, $6, $7, $8, $9, $10, $11)
33 "#,
34 )
35 .bind(run.id)
36 .bind(run.agent_id)
37 .bind(run.scenario_set_id)
38 .bind(status_str)
39 .bind(run.scenario_count as i32)
40 .bind(run.completed_count as i32)
41 .bind(run.error_count as i32)
42 .bind(run.seed as i32)
43 .bind(run.concurrency as i32)
44 .bind(Utc::now())
45 .bind(Utc::now())
46 .execute(&self.pool)
47 .await
48 .map_err(db_err)?;
49
50 self.find_by_id(run.id).await
51 }
52
53 pub async fn find_by_id(&self, id: Uuid) -> Result<EvalRun> {
54 let row = sqlx::query(
56 r#"
57 SELECT id, agent_id, scenario_set_id, status::TEXT AS status,
58 scenario_count, completed_count, error_count,
59 aggregate_score, pass_rate,
60 task_completion, tool_selection, argument_correctness,
61 path_efficiency, schema_compliance, instruction_adherence,
62 failure_clusters, seed, concurrency, error_message,
63 started_at, completed_at, created_at, updated_at,
64 opt_status, opt_rounds, opt_best_score, opt_best_agent_id
65 FROM eval_runs WHERE id = $1
66 "#,
67 )
68 .bind(id)
69 .fetch_optional(&self.pool)
70 .await
71 .map_err(db_err)?
72 .ok_or_else(|| AgentForgeError::NotFound {
73 resource: "EvalRun",
74 id: id.to_string(),
75 })?;
76
77 let tc: Option<f64> = row.get("task_completion");
78 let ts: Option<f64> = row.get("tool_selection");
79 let ac: Option<f64> = row.get("argument_correctness");
80 let pe: Option<f64> = row.get("path_efficiency");
81 let sc: Option<f64> = row.get("schema_compliance");
82 let ia: Option<f64> = row.get("instruction_adherence");
83
84 let scores = if let (Some(tc), Some(ts), Some(ac), Some(pe), Some(sc), Some(ia)) =
85 (tc, ts, ac, pe, sc, ia)
86 {
87 Some(DimensionScores {
88 task_completion: tc,
89 tool_selection: ts,
90 argument_correctness: ac,
91 path_efficiency: pe,
92 schema_compliance: sc,
93 instruction_adherence: ia,
94 })
95 } else {
96 None
97 };
98
99 let clusters_json: Option<serde_json::Value> = row.get("failure_clusters");
100 let failure_clusters: Option<Vec<FailureClusterSummary>> = clusters_json
101 .map(serde_json::from_value)
102 .transpose()
103 .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
104
105 let status_str: String = row.get("status");
106 Ok(EvalRun {
107 id: row.get("id"),
108 agent_id: row.get("agent_id"),
109 scenario_set_id: row.get("scenario_set_id"),
110 status: parse_status(&status_str),
111 scenario_count: row.get::<i32, _>("scenario_count") as u32,
112 completed_count: row.get::<i32, _>("completed_count") as u32,
113 error_count: row.get::<i32, _>("error_count") as u32,
114 aggregate_score: row.get("aggregate_score"),
115 pass_rate: row.get("pass_rate"),
116 scores,
117 failure_clusters,
118 seed: row.get::<i32, _>("seed") as u32,
119 concurrency: row.get::<i32, _>("concurrency") as u32,
120 error_message: row.get("error_message"),
121 started_at: row.get::<Option<DateTime<Utc>>, _>("started_at"),
122 completed_at: row.get::<Option<DateTime<Utc>>, _>("completed_at"),
123 created_at: row.get("created_at"),
124 updated_at: row.get("updated_at"),
125 opt_status: row.get("opt_status"),
126 opt_rounds: row.get::<i32, _>("opt_rounds"),
127 opt_best_score: row.get("opt_best_score"),
128 opt_best_agent_id: row.get("opt_best_agent_id"),
129 })
130 }
131
132 pub async fn update_opt_tracking(
134 &self,
135 id: Uuid,
136 status: &str,
137 rounds: i32,
138 best_score: Option<f64>,
139 best_agent_id: Option<Uuid>,
140 ) -> Result<()> {
141 sqlx::query(
142 r#"
143 UPDATE eval_runs
144 SET opt_status = $2, opt_rounds = $3,
145 opt_best_score = COALESCE($4, opt_best_score),
146 opt_best_agent_id = COALESCE($5, opt_best_agent_id),
147 updated_at = NOW()
148 WHERE id = $1
149 "#,
150 )
151 .bind(id)
152 .bind(status)
153 .bind(rounds)
154 .bind(best_score)
155 .bind(best_agent_id)
156 .execute(&self.pool)
157 .await
158 .map_err(db_err)?;
159 Ok(())
160 }
161
162 pub async fn update_status(&self, id: Uuid, status: &EvalRunStatus) -> Result<()> {
163 let status_str = status.to_string();
164 let started_at = if *status == EvalRunStatus::Running {
165 Some(Utc::now())
166 } else {
167 None
168 };
169 let completed_at = if matches!(
170 status,
171 EvalRunStatus::Complete | EvalRunStatus::Error | EvalRunStatus::Cancelled
172 ) {
173 Some(Utc::now())
174 } else {
175 None
176 };
177
178 sqlx::query(
179 r#"
180 UPDATE eval_runs
181 SET status = $2::eval_run_status,
182 started_at = COALESCE($3, started_at),
183 completed_at = COALESCE($4, completed_at),
184 updated_at = NOW()
185 WHERE id = $1
186 "#,
187 )
188 .bind(id)
189 .bind(status_str)
190 .bind(started_at)
191 .bind(completed_at)
192 .execute(&self.pool)
193 .await
194 .map_err(db_err)?;
195
196 Ok(())
197 }
198
199 pub async fn update_progress(&self, id: Uuid, completed: u32, errors: u32) -> Result<()> {
200 sqlx::query!(
201 r#"
202 UPDATE eval_runs
203 SET completed_count = $2, error_count = $3, updated_at = NOW()
204 WHERE id = $1
205 "#,
206 id,
207 completed as i32,
208 errors as i32,
209 )
210 .execute(&self.pool)
211 .await
212 .map_err(db_err)?;
213 Ok(())
214 }
215
216 pub async fn save_scores(
217 &self,
218 id: Uuid,
219 scores: &DimensionScores,
220 aggregate_score: f64,
221 pass_rate: f64,
222 failure_clusters: &[FailureClusterSummary],
223 ) -> Result<()> {
224 let clusters_json = serde_json::to_value(failure_clusters)
225 .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
226
227 sqlx::query!(
228 r#"
229 UPDATE eval_runs
230 SET aggregate_score = $2, pass_rate = $3,
231 task_completion = $4, tool_selection = $5,
232 argument_correctness = $6, path_efficiency = $7,
233 schema_compliance = $8, instruction_adherence = $9,
234 failure_clusters = $10, updated_at = NOW()
235 WHERE id = $1
236 "#,
237 id,
238 aggregate_score,
239 pass_rate,
240 scores.task_completion,
241 scores.tool_selection,
242 scores.argument_correctness,
243 scores.path_efficiency,
244 scores.schema_compliance,
245 scores.instruction_adherence,
246 clusters_json,
247 )
248 .execute(&self.pool)
249 .await
250 .map_err(db_err)?;
251 Ok(())
252 }
253
254 pub async fn list_by_agent(&self, agent_id: Uuid, limit: i64) -> Result<Vec<EvalRun>> {
255 let rows = sqlx::query!(
256 r#"
257 SELECT id FROM eval_runs
258 WHERE agent_id = $1
259 ORDER BY created_at DESC
260 LIMIT $2
261 "#,
262 agent_id,
263 limit,
264 )
265 .fetch_all(&self.pool)
266 .await
267 .map_err(db_err)?;
268
269 let mut results = Vec::new();
270 for r in rows {
271 results.push(self.find_by_id(r.id).await?);
272 }
273 Ok(results)
274 }
275
276 pub async fn list_all(&self, limit: i64, offset: i64) -> Result<Vec<EvalRun>> {
277 let rows: Vec<(uuid::Uuid,)> =
278 sqlx::query_as("SELECT id FROM eval_runs ORDER BY created_at DESC LIMIT $1 OFFSET $2")
279 .bind(limit)
280 .bind(offset)
281 .fetch_all(&self.pool)
282 .await
283 .map_err(db_err)?;
284
285 let mut results = Vec::new();
286 for (id,) in rows {
287 results.push(self.find_by_id(id).await?);
288 }
289 Ok(results)
290 }
291
292 pub async fn save_error(&self, id: Uuid, message: &str) -> Result<()> {
293 sqlx::query!(
294 r#"
295 UPDATE eval_runs
296 SET status = 'error'::eval_run_status, error_message = $2,
297 completed_at = NOW(), updated_at = NOW()
298 WHERE id = $1
299 "#,
300 id,
301 message,
302 )
303 .execute(&self.pool)
304 .await
305 .map_err(db_err)?;
306 Ok(())
307 }
308
309 pub async fn set_error_message(&self, id: Uuid, message: &str) -> Result<()> {
313 sqlx::query!(
314 r#"
315 UPDATE eval_runs
316 SET error_message = $2, updated_at = NOW()
317 WHERE id = $1
318 "#,
319 id,
320 message,
321 )
322 .execute(&self.pool)
323 .await
324 .map_err(db_err)?;
325 Ok(())
326 }
327
328 pub async fn cancel_or_delete(&self, id: Uuid) -> Result<bool> {
331 let result = sqlx::query(
334 r#"
335 UPDATE eval_runs
336 SET status = CASE
337 WHEN status IN ('pending'::eval_run_status, 'running'::eval_run_status)
338 THEN 'cancelled'::eval_run_status
339 ELSE status
340 END,
341 updated_at = NOW()
342 WHERE id = $1
343 "#,
344 )
345 .bind(id)
346 .execute(&self.pool)
347 .await
348 .map_err(db_err)?;
349 Ok(result.rows_affected() > 0)
350 }
351}
352
353fn parse_status(s: &str) -> EvalRunStatus {
354 match s {
355 "running" => EvalRunStatus::Running,
356 "complete" => EvalRunStatus::Complete,
357 "error" => EvalRunStatus::Error,
358 "cancelled" => EvalRunStatus::Cancelled,
359 _ => EvalRunStatus::Pending,
360 }
361}