Skip to main content

stateset_db/sqlite/
segments.rs

1//! SQLite implementation of customer segment repository
2
3use 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            // Update cached member count
215            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        // Guards against the tiers-class bug: a modeled Vec field that is
380        // persisted but never asserted round-trip. Segments store rules as a
381        // JSON column; verify a non-empty rule set survives create + re-read.
382        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}