Skip to main content

agentforge_db/
eval_repo.rs

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        // Use non-macro sqlx::query so new opt_* columns don't require .sqlx/ regeneration.
55        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    /// Update the iterative optimization loop tracking state on an eval run.
133    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    /// Set the run-level error_message without changing status.
310    /// Used to surface a sample trace failure reason on runs that completed
311    /// but had all (or partial) traces erroring due to LLM API issues.
312    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    /// Cancel a run that is still pending/running, or hard-delete a completed/errored run.
329    /// Returns `true` if a row was affected, `false` if the ID did not exist.
330    pub async fn cancel_or_delete(&self, id: Uuid) -> Result<bool> {
331        // If the run is still active, transition it to cancelled first so any
332        // in-flight background tasks can observe the status change.
333        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}