1use ag_agent::{ReasoningLevel, SpeedMode};
4use ag_session::SettingName;
5use async_trait::async_trait;
6use sqlx::SqlitePool;
7
8use crate::DbError;
9
10#[cfg_attr(test, mockall::automock)]
12#[async_trait]
13pub trait SettingRepository: Send + Sync {
14 async fn get_project_setting(
16 &self,
17 project_id: i64,
18 name: SettingName,
19 ) -> Result<Option<String>, DbError>;
20
21 async fn get_setting(&self, name: SettingName) -> Result<Option<String>, DbError>;
23
24 async fn load_active_project_id(&self) -> Result<Option<i64>, DbError>;
26
27 async fn load_project_reasoning_level(
29 &self,
30 project_id: i64,
31 name: SettingName,
32 ) -> Result<ReasoningLevel, DbError>;
33
34 async fn load_project_speed_mode(
36 &self,
37 project_id: i64,
38 name: SettingName,
39 ) -> Result<SpeedMode, DbError>;
40
41 async fn set_active_project_id(&self, project_id: i64) -> Result<(), DbError>;
43
44 async fn upsert_project_setting(
46 &self,
47 project_id: i64,
48 name: SettingName,
49 value: &str,
50 ) -> Result<(), DbError>;
51
52 async fn upsert_project_settings(
54 &self,
55 project_id: i64,
56 settings: Vec<(SettingName, String)>,
57 ) -> Result<(), DbError>;
58
59 async fn upsert_setting(&self, name: SettingName, value: &str) -> Result<(), DbError>;
61}
62
63#[derive(Clone)]
65pub(crate) struct SqliteSettingRepository(SqlitePool);
66
67impl SqliteSettingRepository {
68 pub(crate) fn new(pool: SqlitePool) -> Self {
70 Self(pool)
71 }
72}
73
74struct RequiredSettingValueRow {
76 value: String,
77}
78
79#[async_trait]
80impl SettingRepository for SqliteSettingRepository {
81 async fn get_project_setting(
82 &self,
83 project_id: i64,
84 name: SettingName,
85 ) -> Result<Option<String>, DbError> {
86 let setting_name = name.as_str();
87 let row = sqlx::query_as!(
88 RequiredSettingValueRow,
89 r#"
90SELECT value AS "value!: _"
91FROM project_setting
92WHERE project_id = ?
93 AND name = ?
94"#,
95 project_id,
96 setting_name
97 )
98 .fetch_optional(&self.0)
99 .await?;
100
101 Ok(row.map(|row| row.value))
102 }
103
104 async fn get_setting(&self, name: SettingName) -> Result<Option<String>, DbError> {
105 let setting_name = name.as_str();
106 let row = sqlx::query_as!(
107 RequiredSettingValueRow,
108 r#"
109SELECT value AS "value!: _"
110FROM setting
111WHERE name = ?
112"#,
113 setting_name
114 )
115 .fetch_optional(&self.0)
116 .await?;
117
118 Ok(row.map(|row| row.value))
119 }
120
121 async fn load_active_project_id(&self) -> Result<Option<i64>, DbError> {
122 let setting_value = self.get_setting(SettingName::ActiveProjectId).await?;
123
124 Ok(setting_value.and_then(|value| value.parse::<i64>().ok()))
125 }
126
127 async fn load_project_reasoning_level(
128 &self,
129 project_id: i64,
130 name: SettingName,
131 ) -> Result<ReasoningLevel, DbError> {
132 let setting_value = self.get_project_setting(project_id, name).await?;
133
134 let reasoning_level = setting_value
135 .as_deref()
136 .and_then(|value| value.parse::<ReasoningLevel>().ok())
137 .unwrap_or_default();
138
139 Ok(reasoning_level)
140 }
141
142 async fn load_project_speed_mode(
143 &self,
144 project_id: i64,
145 name: SettingName,
146 ) -> Result<SpeedMode, DbError> {
147 let setting_value = self.get_project_setting(project_id, name).await?;
148
149 let speed_mode = setting_value
150 .as_deref()
151 .and_then(|value| value.parse::<SpeedMode>().ok())
152 .unwrap_or_default();
153
154 Ok(speed_mode)
155 }
156
157 async fn set_active_project_id(&self, project_id: i64) -> Result<(), DbError> {
158 self.upsert_setting(SettingName::ActiveProjectId, &project_id.to_string())
159 .await
160 }
161
162 async fn upsert_project_setting(
163 &self,
164 project_id: i64,
165 name: SettingName,
166 value: &str,
167 ) -> Result<(), DbError> {
168 sqlx::query!(
169 r"
170INSERT INTO project_setting (project_id, name, value)
171VALUES (?, ?, ?)
172ON CONFLICT(project_id, name) DO UPDATE
173SET value = excluded.value
174",
175 project_id,
176 name.as_str(),
177 value
178 )
179 .execute(&self.0)
180 .await?;
181
182 Ok(())
183 }
184
185 async fn upsert_project_settings(
186 &self,
187 project_id: i64,
188 settings: Vec<(SettingName, String)>,
189 ) -> Result<(), DbError> {
190 let mut transaction = self.0.begin().await?;
191
192 for (name, value) in settings {
193 sqlx::query!(
194 r"
195INSERT INTO project_setting (project_id, name, value)
196VALUES (?, ?, ?)
197ON CONFLICT(project_id, name) DO UPDATE
198SET value = excluded.value
199",
200 project_id,
201 name.as_str(),
202 value
203 )
204 .execute(&mut *transaction)
205 .await?;
206 }
207
208 transaction.commit().await?;
209
210 Ok(())
211 }
212
213 async fn upsert_setting(&self, name: SettingName, value: &str) -> Result<(), DbError> {
214 sqlx::query!(
215 r"
216INSERT INTO setting (name, value)
217VALUES (?, ?)
218ON CONFLICT(name) DO UPDATE
219SET value = excluded.value
220",
221 name.as_str(),
222 value
223 )
224 .execute(&self.0)
225 .await?;
226
227 Ok(())
228 }
229}
230
231#[cfg(test)]
232mod tests {
233 use ag_agent::AgentModel;
234
235 use super::*;
236 use crate::AppRepositories;
237
238 #[tokio::test]
239 async fn test_upsert_project_settings_rolls_back_partial_batch() {
241 let (repositories, pool) = AppRepositories::in_memory_with_pool()
243 .await
244 .expect("db should open");
245 let project_id = repositories
246 .projects()
247 .upsert_project("/tmp/project", Some("main".to_string()))
248 .await
249 .expect("failed to insert project");
250 sqlx::query!(
251 r"
252CREATE TRIGGER fail_default_fast_agent_insert
253BEFORE INSERT ON project_setting
254WHEN NEW.name = 'DefaultFastAgent'
255BEGIN
256 SELECT RAISE(FAIL, 'forced setting failure');
257END;
258"
259 )
260 .execute(&pool)
261 .await
262 .expect("failed to install failure trigger");
263
264 let result = repositories
266 .settings()
267 .upsert_project_settings(
268 project_id,
269 vec![
270 (
271 SettingName::DefaultFastModel,
272 AgentModel::Gemini31Pro.as_str().to_string(),
273 ),
274 (SettingName::DefaultFastAgent, "antigravity".to_string()),
275 ],
276 )
277 .await;
278
279 assert!(result.is_err());
281 assert_eq!(
282 repositories
283 .settings()
284 .get_project_setting(project_id, SettingName::DefaultFastModel)
285 .await
286 .expect("failed to load fast model setting"),
287 None
288 );
289 assert_eq!(
290 repositories
291 .settings()
292 .get_project_setting(project_id, SettingName::DefaultFastAgent)
293 .await
294 .expect("failed to load fast agent setting"),
295 None
296 );
297 }
298}