Skip to main content

agentforge_db/
benchmark_repo.rs

1use crate::db_err;
2use agentforge_core::{AgentForgeError, BenchmarkResult, BenchmarkRun, BenchmarkSuite, Result};
3use chrono::{DateTime, Utc};
4use sqlx::{FromRow, PgPool};
5use uuid::Uuid;
6
7#[derive(FromRow)]
8struct BenchmarkRunRow {
9    id: Uuid,
10    agent_id: Uuid,
11    suite: String,
12    total_tasks: i32,
13    correct: i32,
14    accuracy: Option<f64>,
15    percentile_rank: Option<f64>,
16    #[allow(dead_code)]
17    status: String,
18    created_at: DateTime<Utc>,
19    started_at: Option<DateTime<Utc>>,
20    completed_at: Option<DateTime<Utc>>,
21}
22
23impl BenchmarkRunRow {
24    fn into_run(self) -> BenchmarkRun {
25        BenchmarkRun {
26            id: self.id,
27            agent_id: self.agent_id,
28            suite: parse_suite(&self.suite),
29            total_tasks: self.total_tasks as u32,
30            correct: self.correct as u32,
31            accuracy: self.accuracy.unwrap_or(0.0),
32            percentile_rank: self.percentile_rank,
33            results: vec![],
34            started_at: self.started_at.unwrap_or(self.created_at),
35            completed_at: self.completed_at,
36        }
37    }
38}
39
40pub struct BenchmarkRepo {
41    pool: PgPool,
42}
43
44impl BenchmarkRepo {
45    pub fn new(pool: PgPool) -> Self {
46        Self { pool }
47    }
48
49    pub async fn insert_run(&self, run: &BenchmarkRun) -> Result<BenchmarkRun> {
50        let suite_str = run.suite.to_string();
51
52        sqlx::query(
53            "INSERT INTO benchmark_runs \
54                (id, agent_id, suite, total_tasks, correct, accuracy, \
55                 percentile_rank, status, created_at, started_at, completed_at) \
56             VALUES ($1, $2, $3, $4, $5, $6, $7, 'pending', $8, $9, $10)",
57        )
58        .bind(run.id)
59        .bind(run.agent_id)
60        .bind(&suite_str)
61        .bind(run.total_tasks as i32)
62        .bind(run.correct as i32)
63        .bind(run.accuracy)
64        .bind(run.percentile_rank)
65        .bind(run.started_at)
66        .bind(run.started_at)
67        .bind(run.completed_at)
68        .execute(&self.pool)
69        .await
70        .map_err(db_err)?;
71
72        for result in &run.results {
73            self.insert_result(run.id, result).await?;
74        }
75
76        self.find_run_by_id(run.id).await
77    }
78
79    pub async fn insert_result(&self, run_id: Uuid, result: &BenchmarkResult) -> Result<()> {
80        let suite_str = result.suite.to_string();
81
82        sqlx::query(
83            "INSERT INTO benchmark_results \
84                (id, benchmark_run_id, task_id, suite, agent_answer, \
85                 expected_answer, correct, score, latency_ms, token_cost_usd, created_at) \
86             VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11)",
87        )
88        .bind(Uuid::new_v4())
89        .bind(run_id)
90        .bind(&result.task_id)
91        .bind(&suite_str)
92        .bind(&result.agent_answer)
93        .bind(None::<String>)
94        .bind(result.correct)
95        .bind(result.score)
96        .bind(result.latency_ms as i64)
97        .bind(result.token_cost_usd)
98        .bind(Utc::now())
99        .execute(&self.pool)
100        .await
101        .map_err(db_err)?;
102
103        Ok(())
104    }
105
106    pub async fn find_run_by_id(&self, id: Uuid) -> Result<BenchmarkRun> {
107        sqlx::query_as::<_, BenchmarkRunRow>(
108            "SELECT id, agent_id, suite, total_tasks, correct, accuracy, \
109                    percentile_rank, status, created_at, started_at, completed_at \
110             FROM benchmark_runs WHERE id = $1",
111        )
112        .bind(id)
113        .fetch_optional(&self.pool)
114        .await
115        .map_err(db_err)?
116        .ok_or_else(|| AgentForgeError::NotFound {
117            resource: "BenchmarkRun",
118            id: id.to_string(),
119        })
120        .map(|r| r.into_run())
121    }
122
123    pub async fn update_run_complete(
124        &self,
125        id: Uuid,
126        total_tasks: u32,
127        correct: u32,
128        accuracy: f64,
129        percentile_rank: Option<f64>,
130    ) -> Result<()> {
131        sqlx::query(
132            "UPDATE benchmark_runs \
133             SET total_tasks     = $2, \
134                 correct         = $3, \
135                 accuracy        = $4, \
136                 percentile_rank = $5, \
137                 status          = 'complete', \
138                 completed_at    = $6 \
139             WHERE id = $1",
140        )
141        .bind(id)
142        .bind(total_tasks as i32)
143        .bind(correct as i32)
144        .bind(accuracy)
145        .bind(percentile_rank)
146        .bind(Utc::now())
147        .execute(&self.pool)
148        .await
149        .map_err(db_err)?;
150
151        Ok(())
152    }
153
154    pub async fn list_by_agent(&self, agent_id: Uuid) -> Result<Vec<BenchmarkRun>> {
155        let rows = sqlx::query_as::<_, BenchmarkRunRow>(
156            "SELECT id, agent_id, suite, total_tasks, correct, accuracy, \
157                    percentile_rank, status, created_at, started_at, completed_at \
158             FROM benchmark_runs WHERE agent_id = $1 ORDER BY created_at DESC",
159        )
160        .bind(agent_id)
161        .fetch_all(&self.pool)
162        .await
163        .map_err(db_err)?;
164
165        Ok(rows.into_iter().map(|r| r.into_run()).collect())
166    }
167}
168
169fn parse_suite(s: &str) -> BenchmarkSuite {
170    match s {
171        "agentbench" => BenchmarkSuite::AgentBench,
172        "webarena" => BenchmarkSuite::WebArena,
173        _ => BenchmarkSuite::Gaia,
174    }
175}