agentforge_db/
shadow_repo.rs1use crate::db_err;
2use agentforge_core::{AgentForgeError, Result, ShadowComparison, ShadowRun, ShadowRunStatus};
3use chrono::{DateTime, Utc};
4use sqlx::{FromRow, PgPool};
5use uuid::Uuid;
6
7#[derive(FromRow)]
8struct ShadowRunRow {
9 id: Uuid,
10 champion_agent_id: Uuid,
11 candidate_agent_id: Uuid,
12 traffic_percent: i16,
13 status: String,
14 comparison_result: Option<serde_json::Value>,
15 error_message: Option<String>,
16 created_at: DateTime<Utc>,
17 started_at: Option<DateTime<Utc>>,
18 completed_at: Option<DateTime<Utc>>,
19}
20
21impl ShadowRunRow {
22 fn into_shadow_run(self) -> Result<ShadowRun> {
23 let status = match self.status.as_str() {
24 "pending" => ShadowRunStatus::Pending,
25 "running" => ShadowRunStatus::Running,
26 "complete" => ShadowRunStatus::Complete,
27 _ => ShadowRunStatus::Error,
28 };
29 let comparison: Option<ShadowComparison> = self
30 .comparison_result
31 .map(serde_json::from_value)
32 .transpose()
33 .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
34 Ok(ShadowRun {
35 id: self.id,
36 champion_agent_id: self.champion_agent_id,
37 candidate_agent_id: self.candidate_agent_id,
38 traffic_percent: self.traffic_percent as u8,
39 status,
40 comparison,
41 error_message: self.error_message,
42 created_at: self.created_at,
43 started_at: self.started_at,
44 completed_at: self.completed_at,
45 })
46 }
47}
48
49pub struct ShadowRepo {
50 pool: PgPool,
51}
52
53impl ShadowRepo {
54 pub fn new(pool: PgPool) -> Self {
55 Self { pool }
56 }
57
58 pub async fn insert(&self, run: &ShadowRun) -> Result<ShadowRun> {
59 let status_str = run.status.to_string();
60 let comparison_json = run
61 .comparison
62 .as_ref()
63 .map(serde_json::to_value)
64 .transpose()
65 .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
66
67 sqlx::query(
68 "INSERT INTO shadow_runs \
69 (id, champion_agent_id, candidate_agent_id, traffic_percent, \
70 status, comparison_result, error_message, created_at) \
71 VALUES ($1, $2, $3, $4, $5, $6, $7, $8)",
72 )
73 .bind(run.id)
74 .bind(run.champion_agent_id)
75 .bind(run.candidate_agent_id)
76 .bind(run.traffic_percent as i16)
77 .bind(&status_str)
78 .bind(&comparison_json)
79 .bind(&run.error_message)
80 .bind(run.created_at)
81 .execute(&self.pool)
82 .await
83 .map_err(db_err)?;
84
85 self.find_by_id(run.id).await
86 }
87
88 pub async fn find_by_id(&self, id: Uuid) -> Result<ShadowRun> {
89 sqlx::query_as::<_, ShadowRunRow>(
90 "SELECT id, champion_agent_id, candidate_agent_id, traffic_percent, \
91 status, comparison_result, error_message, \
92 created_at, started_at, completed_at \
93 FROM shadow_runs WHERE id = $1",
94 )
95 .bind(id)
96 .fetch_optional(&self.pool)
97 .await
98 .map_err(db_err)?
99 .ok_or_else(|| AgentForgeError::NotFound {
100 resource: "ShadowRun",
101 id: id.to_string(),
102 })?
103 .into_shadow_run()
104 }
105
106 pub async fn update_status(
107 &self,
108 id: Uuid,
109 status: &ShadowRunStatus,
110 comparison: Option<&ShadowComparison>,
111 error_message: Option<&str>,
112 ) -> Result<()> {
113 let status_str = status.to_string();
114 let comparison_json = comparison
115 .map(serde_json::to_value)
116 .transpose()
117 .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
118 let completed_at: Option<DateTime<Utc>> =
119 if *status == ShadowRunStatus::Complete || *status == ShadowRunStatus::Error {
120 Some(Utc::now())
121 } else {
122 None
123 };
124 let started_at: Option<DateTime<Utc>> = if *status == ShadowRunStatus::Running {
125 Some(Utc::now())
126 } else {
127 None
128 };
129
130 sqlx::query(
131 "UPDATE shadow_runs \
132 SET status = $2, \
133 comparison_result = $3, \
134 error_message = $4, \
135 started_at = COALESCE(started_at, $5), \
136 completed_at = $6 \
137 WHERE id = $1",
138 )
139 .bind(id)
140 .bind(&status_str)
141 .bind(&comparison_json)
142 .bind(error_message)
143 .bind(started_at)
144 .bind(completed_at)
145 .execute(&self.pool)
146 .await
147 .map_err(db_err)?;
148
149 Ok(())
150 }
151
152 pub async fn list(&self, limit: i64, offset: i64) -> Result<Vec<ShadowRun>> {
153 let rows = sqlx::query_as::<_, ShadowRunRow>(
154 "SELECT id, champion_agent_id, candidate_agent_id, traffic_percent, \
155 status, comparison_result, error_message, \
156 created_at, started_at, completed_at \
157 FROM shadow_runs ORDER BY created_at DESC LIMIT $1 OFFSET $2",
158 )
159 .bind(limit)
160 .bind(offset)
161 .fetch_all(&self.pool)
162 .await
163 .map_err(db_err)?;
164
165 rows.into_iter().map(|r| r.into_shadow_run()).collect()
166 }
167}