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}