1use super::{
4 map_db_error, parse_datetime_row, parse_enum_row, parse_json_row, parse_uuid_row,
5 with_immediate_transaction,
6};
7use chrono::Utc;
8use r2d2::Pool;
9use r2d2_sqlite::SqliteConnectionManager;
10use stateset_core::{
11 CommerceError, CreateSegment, CustomerId, Result, Segment, SegmentFilter, SegmentId,
12 SegmentMembership, SegmentRepository, SegmentRule, UpdateSegment,
13};
14
15#[derive(Debug)]
16pub struct SqliteSegmentRepository {
17 pool: Pool<SqliteConnectionManager>,
18}
19
20impl SqliteSegmentRepository {
21 #[must_use]
22 pub const fn new(pool: Pool<SqliteConnectionManager>) -> Self {
23 Self { pool }
24 }
25
26 fn conn(&self) -> Result<r2d2::PooledConnection<SqliteConnectionManager>> {
27 self.pool.get().map_err(|e| CommerceError::DatabaseError(e.to_string()))
28 }
29
30 fn row_to_segment(row: &rusqlite::Row<'_>) -> rusqlite::Result<Segment> {
31 let rules_json: String = row.get("rules")?;
32 let rules: Vec<SegmentRule> = parse_json_row(&rules_json, "segment", "rules")?;
33
34 Ok(Segment {
35 id: parse_uuid_row(&row.get::<_, String>("id")?, "segment", "id")?.into(),
36 name: row.get("name")?,
37 description: row.get("description")?,
38 segment_type: parse_enum_row(
39 &row.get::<_, String>("segment_type")?,
40 "segment",
41 "segment_type",
42 )?,
43 rules,
44 member_count: row.get::<_, i64>("member_count")? as u64,
45 created_at: parse_datetime_row(
46 &row.get::<_, String>("created_at")?,
47 "segment",
48 "created_at",
49 )?,
50 updated_at: parse_datetime_row(
51 &row.get::<_, String>("updated_at")?,
52 "segment",
53 "updated_at",
54 )?,
55 })
56 }
57
58 fn row_to_membership(row: &rusqlite::Row<'_>) -> rusqlite::Result<SegmentMembership> {
59 Ok(SegmentMembership {
60 segment_id: parse_uuid_row(
61 &row.get::<_, String>("segment_id")?,
62 "segment_membership",
63 "segment_id",
64 )?
65 .into(),
66 customer_id: parse_uuid_row(
67 &row.get::<_, String>("customer_id")?,
68 "segment_membership",
69 "customer_id",
70 )?
71 .into(),
72 joined_at: parse_datetime_row(
73 &row.get::<_, String>("joined_at")?,
74 "segment_membership",
75 "joined_at",
76 )?,
77 })
78 }
79}
80
81impl SegmentRepository for SqliteSegmentRepository {
82 fn create(&self, input: CreateSegment) -> Result<Segment> {
83 let id = SegmentId::new();
84 let now = Utc::now();
85 let id_str = id.to_string();
86 let now_str = now.to_rfc3339();
87
88 let rules_json = serde_json::to_string(&input.rules)
89 .map_err(|e| CommerceError::DatabaseError(e.to_string()))?;
90
91 with_immediate_transaction(&self.pool, |tx| {
92 tx.execute(
93 "INSERT INTO segments (id, name, description, segment_type, rules, member_count, created_at, updated_at)
94 VALUES (?, ?, ?, ?, ?, 0, ?, ?)",
95 rusqlite::params![
96 &id_str,
97 &input.name,
98 &input.description,
99 input.segment_type.to_string(),
100 &rules_json,
101 &now_str,
102 &now_str,
103 ],
104 )?;
105
106 tx.query_row("SELECT * FROM segments WHERE id = ?", [&id_str], Self::row_to_segment)
107 })
108 }
109
110 fn get(&self, id: SegmentId) -> Result<Option<Segment>> {
111 let conn = self.conn()?;
112 match conn.query_row(
113 "SELECT * FROM segments WHERE id = ?",
114 [id.to_string()],
115 Self::row_to_segment,
116 ) {
117 Ok(s) => Ok(Some(s)),
118 Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None),
119 Err(e) => Err(map_db_error(e)),
120 }
121 }
122
123 fn update(&self, id: SegmentId, input: UpdateSegment) -> Result<Segment> {
124 let id_str = id.to_string();
125 let now_str = Utc::now().to_rfc3339();
126
127 with_immediate_transaction(&self.pool, |tx| {
128 let mut sets = vec!["updated_at = ?".to_string()];
129 let mut params: Vec<Box<dyn rusqlite::types::ToSql>> = vec![Box::new(now_str.clone())];
130
131 if let Some(ref name) = input.name {
132 sets.push("name = ?".into());
133 params.push(Box::new(name.clone()));
134 }
135 if let Some(ref description) = input.description {
136 sets.push("description = ?".into());
137 params.push(Box::new(description.clone()));
138 }
139 if let Some(ref rules) = input.rules {
140 let rules_json = serde_json::to_string(rules).map_err(|e| {
141 rusqlite::Error::ToSqlConversionFailure(Box::new(CommerceError::DatabaseError(
142 e.to_string(),
143 )))
144 })?;
145 sets.push("rules = ?".into());
146 params.push(Box::new(rules_json));
147 }
148
149 let sql = format!("UPDATE segments SET {} WHERE id = ?", sets.join(", "));
150 params.push(Box::new(id_str.clone()));
151
152 let param_refs: Vec<&dyn rusqlite::types::ToSql> =
153 params.iter().map(|p| p.as_ref()).collect();
154 tx.execute(&sql, param_refs.as_slice())?;
155
156 tx.query_row("SELECT * FROM segments WHERE id = ?", [&id_str], Self::row_to_segment)
157 })
158 }
159
160 fn list(&self, filter: SegmentFilter) -> Result<Vec<Segment>> {
161 let conn = self.conn()?;
162 let mut sql = "SELECT * FROM segments WHERE 1=1".to_string();
163 let mut params: Vec<Box<dyn rusqlite::types::ToSql>> = vec![];
164
165 if let Some(segment_type) = filter.segment_type {
166 sql.push_str(" AND segment_type = ?");
167 params.push(Box::new(segment_type.to_string()));
168 }
169 if let Some(ref name) = filter.name {
170 sql.push_str(" AND name LIKE ?");
171 params.push(Box::new(format!("%{name}%")));
172 }
173
174 sql.push_str(" ORDER BY created_at DESC");
175
176 crate::sqlite::append_limit_offset(&mut sql, filter.limit, filter.offset);
177
178 let param_refs: Vec<&dyn rusqlite::types::ToSql> =
179 params.iter().map(|p| p.as_ref()).collect();
180 let mut stmt = conn.prepare(&sql).map_err(map_db_error)?;
181 let rows = stmt
182 .query_map(param_refs.as_slice(), Self::row_to_segment)
183 .map_err(map_db_error)?
184 .collect::<std::result::Result<Vec<_>, _>>()
185 .map_err(map_db_error)?;
186 Ok(rows)
187 }
188
189 fn delete(&self, id: SegmentId) -> Result<()> {
190 let conn = self.conn()?;
191 conn.execute("DELETE FROM segment_memberships WHERE segment_id = ?", [id.to_string()])
192 .map_err(map_db_error)?;
193 conn.execute("DELETE FROM segments WHERE id = ?", [id.to_string()])
194 .map_err(map_db_error)?;
195 Ok(())
196 }
197
198 fn add_member(
199 &self,
200 segment_id: SegmentId,
201 customer_id: CustomerId,
202 ) -> Result<SegmentMembership> {
203 let now = Utc::now();
204 let now_str = now.to_rfc3339();
205 let seg_str = segment_id.to_string();
206 let cust_str = customer_id.to_string();
207
208 with_immediate_transaction(&self.pool, |tx| {
209 tx.execute(
210 "INSERT OR IGNORE INTO segment_memberships (segment_id, customer_id, joined_at) VALUES (?, ?, ?)",
211 rusqlite::params![&seg_str, &cust_str, &now_str],
212 )?;
213
214 tx.execute(
216 "UPDATE segments SET member_count = (SELECT COUNT(*) FROM segment_memberships WHERE segment_id = ?), updated_at = ? WHERE id = ?",
217 rusqlite::params![&seg_str, &now_str, &seg_str],
218 )?;
219
220 tx.query_row(
221 "SELECT * FROM segment_memberships WHERE segment_id = ? AND customer_id = ?",
222 rusqlite::params![&seg_str, &cust_str],
223 Self::row_to_membership,
224 )
225 })
226 }
227
228 fn remove_member(&self, segment_id: SegmentId, customer_id: CustomerId) -> Result<()> {
229 let seg_str = segment_id.to_string();
230 let now_str = Utc::now().to_rfc3339();
231
232 with_immediate_transaction(&self.pool, |tx| {
233 tx.execute(
234 "DELETE FROM segment_memberships WHERE segment_id = ? AND customer_id = ?",
235 rusqlite::params![&seg_str, customer_id.to_string()],
236 )?;
237
238 tx.execute(
239 "UPDATE segments SET member_count = (SELECT COUNT(*) FROM segment_memberships WHERE segment_id = ?), updated_at = ? WHERE id = ?",
240 rusqlite::params![&seg_str, &now_str, &seg_str],
241 )?;
242
243 Ok(())
244 })
245 }
246
247 fn list_members(
248 &self,
249 segment_id: SegmentId,
250 limit: Option<u32>,
251 offset: Option<u32>,
252 ) -> Result<Vec<SegmentMembership>> {
253 let conn = self.conn()?;
254 let mut sql =
255 "SELECT * FROM segment_memberships WHERE segment_id = ? ORDER BY joined_at DESC"
256 .to_string();
257
258 crate::sqlite::append_limit_offset(&mut sql, limit, offset);
259
260 let mut stmt = conn.prepare(&sql).map_err(map_db_error)?;
261 let rows = stmt
262 .query_map([segment_id.to_string()], Self::row_to_membership)
263 .map_err(map_db_error)?
264 .collect::<std::result::Result<Vec<_>, _>>()
265 .map_err(map_db_error)?;
266 Ok(rows)
267 }
268
269 fn is_member(&self, segment_id: SegmentId, customer_id: CustomerId) -> Result<bool> {
270 let conn = self.conn()?;
271 let count: i64 = conn
272 .query_row(
273 "SELECT COUNT(*) FROM segment_memberships WHERE segment_id = ? AND customer_id = ?",
274 rusqlite::params![segment_id.to_string(), customer_id.to_string()],
275 |row| row.get(0),
276 )
277 .map_err(map_db_error)?;
278 Ok(count > 0)
279 }
280
281 fn count_members(&self, segment_id: SegmentId) -> Result<u64> {
282 let conn = self.conn()?;
283 let count: i64 = conn
284 .query_row(
285 "SELECT COUNT(*) FROM segment_memberships WHERE segment_id = ?",
286 [segment_id.to_string()],
287 |row| row.get(0),
288 )
289 .map_err(map_db_error)?;
290 Ok(count as u64)
291 }
292}
293
294#[cfg(test)]
295mod tests {
296 use super::*;
297 use crate::DatabaseConfig;
298 use crate::sqlite::SqliteDatabase;
299 use stateset_core::SegmentType;
300
301 fn test_repo() -> SqliteSegmentRepository {
302 let db = SqliteDatabase::new(&DatabaseConfig::in_memory()).expect("in-memory db");
303 let conn = db.conn().expect("conn");
304 conn.execute_batch(
305 "CREATE TABLE IF NOT EXISTS segments (
306 id TEXT PRIMARY KEY,
307 name TEXT NOT NULL,
308 description TEXT,
309 segment_type TEXT NOT NULL DEFAULT 'static',
310 rules TEXT NOT NULL DEFAULT '[]',
311 member_count INTEGER NOT NULL DEFAULT 0,
312 created_at TEXT NOT NULL DEFAULT (datetime('now')),
313 updated_at TEXT NOT NULL DEFAULT (datetime('now'))
314 );
315 CREATE TABLE IF NOT EXISTS segment_memberships (
316 segment_id TEXT NOT NULL,
317 customer_id TEXT NOT NULL,
318 joined_at TEXT NOT NULL DEFAULT (datetime('now')),
319 PRIMARY KEY (segment_id, customer_id),
320 FOREIGN KEY (segment_id) REFERENCES segments(id)
321 );",
322 )
323 .expect("create tables");
324 SqliteSegmentRepository::new(db.pool().clone())
325 }
326
327 #[test]
328 fn create_and_get_segment() {
329 let repo = test_repo();
330 let segment = repo
331 .create(CreateSegment {
332 name: "VIP Customers".into(),
333 description: Some("High-value customers".into()),
334 segment_type: SegmentType::Static,
335 rules: vec![],
336 })
337 .expect("create");
338
339 assert_eq!(segment.name, "VIP Customers");
340 assert_eq!(segment.member_count, 0);
341
342 let fetched = repo.get(segment.id).expect("get").expect("found");
343 assert_eq!(fetched.id, segment.id);
344 assert_eq!(fetched.name, "VIP Customers");
345 }
346
347 #[test]
348 fn add_and_count_members() {
349 let repo = test_repo();
350 let segment = repo
351 .create(CreateSegment {
352 name: "Test Segment".into(),
353 description: None,
354 segment_type: SegmentType::Static,
355 rules: vec![],
356 })
357 .expect("create");
358
359 let c1 = CustomerId::new();
360 let c2 = CustomerId::new();
361
362 repo.add_member(segment.id, c1).expect("add c1");
363 repo.add_member(segment.id, c2).expect("add c2");
364
365 assert!(repo.is_member(segment.id, c1).expect("is_member c1"));
366 assert!(repo.is_member(segment.id, c2).expect("is_member c2"));
367 assert_eq!(repo.count_members(segment.id).expect("count"), 2);
368
369 let members = repo.list_members(segment.id, None, None).expect("list members");
370 assert_eq!(members.len(), 2);
371
372 repo.remove_member(segment.id, c1).expect("remove c1");
373 assert!(!repo.is_member(segment.id, c1).expect("not member"));
374 assert_eq!(repo.count_members(segment.id).expect("count after remove"), 1);
375 }
376
377 #[test]
378 fn rules_round_trip() {
379 use stateset_core::{SegmentOperator, SegmentRule};
383
384 let repo = test_repo();
385 let created = repo
386 .create(CreateSegment {
387 name: "High spenders".into(),
388 description: None,
389 segment_type: SegmentType::Dynamic,
390 rules: vec![
391 SegmentRule {
392 field: "lifetime_value".into(),
393 operator: SegmentOperator::Gte,
394 value: "1000".into(),
395 },
396 SegmentRule {
397 field: "country".into(),
398 operator: SegmentOperator::Eq,
399 value: "US".into(),
400 },
401 ],
402 })
403 .expect("create");
404 assert_eq!(created.rules.len(), 2, "rules returned on create");
405
406 let fetched = repo.get(created.id).expect("get").expect("found");
407 assert_eq!(fetched.rules.len(), 2, "rules survive a re-read");
408 assert_eq!(fetched.rules[0].field, "lifetime_value");
409 assert_eq!(fetched.rules[0].operator, SegmentOperator::Gte);
410 assert_eq!(fetched.rules[1].value, "US");
411 }
412}