Skip to main content

agentforge_db/
scenario_repo.rs

1use crate::db_err;
2use agentforge_core::{
3    AgentForgeError, DifficultyTier, Result, Scenario, ScenarioExpected, ScenarioInput,
4    ScenarioSource,
5};
6use chrono::Utc;
7use sqlx::PgPool;
8use uuid::Uuid;
9
10pub struct ScenarioRepo {
11    pool: PgPool,
12}
13
14impl ScenarioRepo {
15    pub fn new(pool: PgPool) -> Self {
16        Self { pool }
17    }
18
19    pub async fn insert_batch(&self, scenarios: &[Scenario]) -> Result<Vec<Uuid>> {
20        let mut ids = Vec::new();
21        for s in scenarios {
22            let id = self.insert(s).await?;
23            ids.push(id);
24        }
25        Ok(ids)
26    }
27
28    pub async fn insert(&self, s: &Scenario) -> Result<Uuid> {
29        let input_json = serde_json::to_value(&s.input)
30            .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
31        let expected_json = serde_json::to_value(&s.expected)
32            .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
33        let diff_str = s.difficulty.to_string();
34        let source_str = s.source.to_string();
35
36        sqlx::query(
37            r#"
38            INSERT INTO scenarios (id, agent_id, input, expected, difficulty, domain, source, tags, created_at)
39            VALUES ($1, $2, $3, $4, $5::difficulty_tier, $6, $7::scenario_source, $8, $9)
40            "#,
41        )
42        .bind(s.id)
43        .bind(s.agent_id)
44        .bind(input_json)
45        .bind(expected_json)
46        .bind(diff_str)
47        .bind(s.domain.clone())
48        .bind(source_str)
49        .bind(&s.tags)
50        .bind(Utc::now())
51        .execute(&self.pool)
52        .await
53        .map_err(db_err)?;
54
55        Ok(s.id)
56    }
57
58    pub async fn find_by_id(&self, id: Uuid) -> Result<Scenario> {
59        let row = sqlx::query!(
60            r#"
61            SELECT id, agent_id, input, expected,
62                   difficulty as "difficulty: String",
63                   domain,
64                   source as "source: String",
65                   tags, created_at
66            FROM scenarios WHERE id = $1
67            "#,
68            id
69        )
70        .fetch_optional(&self.pool)
71        .await
72        .map_err(db_err)?
73        .ok_or_else(|| AgentForgeError::NotFound {
74            resource: "Scenario",
75            id: id.to_string(),
76        })?;
77
78        self.row_to_scenario(
79            row.id,
80            row.agent_id,
81            row.input,
82            row.expected,
83            row.difficulty,
84            row.domain,
85            row.source,
86            row.tags,
87            row.created_at,
88        )
89    }
90
91    pub async fn list_by_agent(&self, agent_id: Uuid, limit: i64) -> Result<Vec<Scenario>> {
92        let rows = sqlx::query!(
93            r#"
94            SELECT id, agent_id, input, expected,
95                   difficulty as "difficulty: String",
96                   domain,
97                   source as "source: String",
98                   tags, created_at
99            FROM scenarios
100            WHERE agent_id = $1
101            ORDER BY created_at DESC
102            LIMIT $2
103            "#,
104            agent_id,
105            limit,
106        )
107        .fetch_all(&self.pool)
108        .await
109        .map_err(db_err)?;
110
111        rows.into_iter()
112            .map(|r| {
113                self.row_to_scenario(
114                    r.id,
115                    r.agent_id,
116                    r.input,
117                    r.expected,
118                    r.difficulty,
119                    r.domain,
120                    r.source,
121                    r.tags,
122                    r.created_at,
123                )
124            })
125            .collect()
126    }
127
128    #[allow(clippy::too_many_arguments)]
129    fn row_to_scenario(
130        &self,
131        id: Uuid,
132        agent_id: Uuid,
133        input: serde_json::Value,
134        expected: serde_json::Value,
135        difficulty: String,
136        domain: Option<String>,
137        source: String,
138        tags: Vec<String>,
139        created_at: chrono::DateTime<Utc>,
140    ) -> Result<Scenario> {
141        let parsed_input: ScenarioInput = serde_json::from_value(input)
142            .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
143        let parsed_expected: ScenarioExpected = serde_json::from_value(expected)
144            .map_err(|e| AgentForgeError::SerializationError(e.to_string()))?;
145
146        Ok(Scenario {
147            id,
148            agent_id,
149            input: parsed_input,
150            expected: parsed_expected,
151            difficulty: parse_difficulty(&difficulty),
152            domain,
153            source: parse_source(&source),
154            tags,
155            created_at,
156        })
157    }
158}
159
160fn parse_difficulty(s: &str) -> DifficultyTier {
161    match s {
162        "easy" => DifficultyTier::Easy,
163        "hard" => DifficultyTier::Hard,
164        "edge" => DifficultyTier::Edge,
165        _ => DifficultyTier::Medium,
166    }
167}
168
169fn parse_source(s: &str) -> ScenarioSource {
170    match s {
171        "adversarial" => ScenarioSource::Adversarial,
172        "domain_seeded" => ScenarioSource::DomainSeeded,
173        "manual" => ScenarioSource::Manual,
174        _ => ScenarioSource::SchemaDerived,
175    }
176}