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}