Skip to main content

agentforge_db/
shadow_repo.rs

1use 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}