1use crate::db_err;
2use agentforge_core::{
3 AgentForgeError, DimensionScores, FailureCluster, Result, Trace, TraceStatus, TraceStep,
4};
5use chrono::Utc;
6use sqlx::PgPool;
7use uuid::Uuid;
8
9pub struct TraceRepo {
10 pool: PgPool,
11}
12
13impl TraceRepo {
14 pub fn new(pool: PgPool) -> Self {
15 Self { pool }
16 }
17
18 pub async fn insert(&self, trace: &Trace) -> Result<()> {
19 let status_str = trace.status.to_string();
20 let cluster_str = trace.failure_cluster.to_string();
21 let steps_json = serde_json::to_value(&trace.steps)
22 .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
23 let scores_json = trace
24 .scores
25 .as_ref()
26 .map(serde_json::to_value)
27 .transpose()
28 .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
29
30 sqlx::query(
31 r#"
32 INSERT INTO traces
33 (id, run_id, scenario_id, status, steps, final_output, scores,
34 aggregate_score, failure_cluster, failure_reason, review_needed,
35 llm_calls, tool_invocations, input_tokens, output_tokens,
36 latency_ms, retry_count, seed, created_at)
37 VALUES
38 ($1, $2, $3, $4::trace_status, $5, $6, $7, $8,
39 $9::failure_cluster, $10, $11, $12, $13, $14, $15, $16, $17, $18, $19)
40 "#,
41 )
42 .bind(trace.id)
43 .bind(trace.run_id)
44 .bind(trace.scenario_id)
45 .bind(status_str)
46 .bind(steps_json)
47 .bind(trace.final_output.clone())
48 .bind(scores_json)
49 .bind(trace.aggregate_score)
50 .bind(cluster_str)
51 .bind(trace.failure_reason.clone())
52 .bind(trace.review_needed)
53 .bind(trace.llm_calls as i32)
54 .bind(trace.tool_invocations as i32)
55 .bind(trace.input_tokens as i32)
56 .bind(trace.output_tokens as i32)
57 .bind(trace.latency_ms as i32)
58 .bind(trace.retry_count as i32)
59 .bind(trace.seed as i32)
60 .bind(Utc::now())
61 .execute(&self.pool)
62 .await
63 .map_err(db_err)?;
64 Ok(())
65 }
66
67 pub async fn find_by_id(&self, id: Uuid) -> Result<Trace> {
68 let r = sqlx::query!(
69 r#"
70 SELECT id, run_id, scenario_id,
71 status as "status: String",
72 steps, final_output, scores, aggregate_score,
73 failure_cluster as "failure_cluster: String",
74 failure_reason, review_needed,
75 llm_calls, tool_invocations, input_tokens, output_tokens,
76 latency_ms, retry_count, seed, created_at
77 FROM traces WHERE id = $1
78 "#,
79 id
80 )
81 .fetch_optional(&self.pool)
82 .await
83 .map_err(db_err)?
84 .ok_or_else(|| AgentForgeError::NotFound {
85 resource: "Trace",
86 id: id.to_string(),
87 })?;
88
89 self.convert_row(
90 r.id,
91 r.run_id,
92 r.scenario_id,
93 r.status,
94 r.steps,
95 r.final_output,
96 r.scores,
97 r.aggregate_score,
98 r.failure_cluster,
99 r.failure_reason,
100 r.review_needed,
101 r.llm_calls,
102 r.tool_invocations,
103 r.input_tokens,
104 r.output_tokens,
105 r.latency_ms,
106 r.retry_count,
107 r.seed,
108 r.created_at,
109 )
110 }
111
112 pub async fn list_by_run(&self, run_id: Uuid) -> Result<Vec<Trace>> {
113 let rows = sqlx::query!(
114 r#"
115 SELECT id, run_id, scenario_id,
116 status as "status: String",
117 steps, final_output, scores, aggregate_score,
118 failure_cluster as "failure_cluster: String",
119 failure_reason, review_needed,
120 llm_calls, tool_invocations, input_tokens, output_tokens,
121 latency_ms, retry_count, seed, created_at
122 FROM traces WHERE run_id = $1 ORDER BY created_at ASC
123 "#,
124 run_id
125 )
126 .fetch_all(&self.pool)
127 .await
128 .map_err(db_err)?;
129
130 rows.into_iter()
131 .map(|r| {
132 self.convert_row(
133 r.id,
134 r.run_id,
135 r.scenario_id,
136 r.status,
137 r.steps,
138 r.final_output,
139 r.scores,
140 r.aggregate_score,
141 r.failure_cluster,
142 r.failure_reason,
143 r.review_needed,
144 r.llm_calls,
145 r.tool_invocations,
146 r.input_tokens,
147 r.output_tokens,
148 r.latency_ms,
149 r.retry_count,
150 r.seed,
151 r.created_at,
152 )
153 })
154 .collect()
155 }
156
157 pub async fn list_by_run_paginated(
159 &self,
160 run_id: Uuid,
161 limit: i64,
162 offset: i64,
163 ) -> Result<Vec<Trace>> {
164 let all = self.list_by_run(run_id).await?;
165 let offset = offset.max(0) as usize;
166 let limit = limit.max(0) as usize;
167 Ok(all.into_iter().skip(offset).take(limit).collect())
168 }
169
170 pub async fn list_passing_scenario_ids(&self, run_id: Uuid) -> Result<Vec<Uuid>> {
172 let rows = sqlx::query!(
173 "SELECT scenario_id FROM traces WHERE run_id = $1 AND status = 'pass'::trace_status",
174 run_id
175 )
176 .fetch_all(&self.pool)
177 .await
178 .map_err(db_err)?;
179 Ok(rows.into_iter().map(|r| r.scenario_id).collect())
180 }
181
182 pub async fn count_review_needed(&self, run_id: Uuid) -> Result<i64> {
183 let row = sqlx::query!(
184 "SELECT COUNT(*) as cnt FROM traces WHERE run_id = $1 AND review_needed = TRUE",
185 run_id
186 )
187 .fetch_one(&self.pool)
188 .await
189 .map_err(db_err)?;
190 Ok(row.cnt.unwrap_or(0))
191 }
192
193 #[allow(clippy::too_many_arguments)]
194 fn convert_row(
195 &self,
196 id: Uuid,
197 run_id: Uuid,
198 scenario_id: Uuid,
199 status: String,
200 steps: serde_json::Value,
201 final_output: Option<serde_json::Value>,
202 scores: Option<serde_json::Value>,
203 aggregate_score: Option<f64>,
204 failure_cluster: String,
205 failure_reason: Option<String>,
206 review_needed: bool,
207 llm_calls: i32,
208 tool_invocations: i32,
209 input_tokens: i32,
210 output_tokens: i32,
211 latency_ms: i32,
212 retry_count: i32,
213 seed: i32,
214 created_at: chrono::DateTime<Utc>,
215 ) -> Result<Trace> {
216 let steps: Vec<TraceStep> = serde_json::from_value(steps)
217 .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
218 let scores: Option<DimensionScores> = scores
219 .map(serde_json::from_value)
220 .transpose()
221 .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
222
223 Ok(Trace {
224 id,
225 run_id,
226 scenario_id,
227 status: parse_trace_status(&status),
228 steps,
229 final_output,
230 scores,
231 aggregate_score,
232 failure_cluster: parse_failure_cluster(&failure_cluster),
233 failure_reason,
234 review_needed,
235 llm_calls: llm_calls as u32,
236 tool_invocations: tool_invocations as u32,
237 input_tokens: input_tokens as u32,
238 output_tokens: output_tokens as u32,
239 latency_ms: latency_ms as u64,
240 retry_count: retry_count as u32,
241 seed: seed as u32,
242 created_at,
243 })
244 }
245}
246
247fn parse_trace_status(s: &str) -> TraceStatus {
248 match s {
249 "pass" => TraceStatus::Pass,
250 "fail" => TraceStatus::Fail,
251 "review_needed" => TraceStatus::ReviewNeeded,
252 _ => TraceStatus::Error,
253 }
254}
255
256fn parse_failure_cluster(s: &str) -> FailureCluster {
257 match s {
258 "wrong_tool" => FailureCluster::WrongTool,
259 "hallucinated_argument" => FailureCluster::HallucinatedArgument,
260 "looping" => FailureCluster::Looping,
261 "premature_stop" => FailureCluster::PrematureStop,
262 "schema_violation" => FailureCluster::SchemaViolation,
263 "constraint_breach" => FailureCluster::ConstraintBreach,
264 "no_failure" => FailureCluster::NoFailure,
265 _ => FailureCluster::Unknown,
266 }
267}