Skip to main content

agentforge_db/
trace_repo.rs

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    /// Paginated version of list_by_run.
158    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    /// Returns all scenario IDs that passed in a given run.
171    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}