Skip to main content

stateset_db/sqlite/
rewards.rs

1//! SQLite implementation of reward catalog repository
2
3use super::{
4    map_db_error, parse_datetime_row, parse_decimal_opt_row, parse_enum_row, parse_uuid_row,
5    with_immediate_transaction,
6};
7use chrono::Utc;
8use r2d2::Pool;
9use r2d2_sqlite::SqliteConnectionManager;
10use stateset_core::{
11    CommerceError, CreateReward, Result, Reward, RewardFilter, RewardId, RewardRepository,
12};
13
14#[derive(Debug)]
15pub struct SqliteRewardRepository {
16    pool: Pool<SqliteConnectionManager>,
17}
18
19impl SqliteRewardRepository {
20    #[must_use]
21    pub const fn new(pool: Pool<SqliteConnectionManager>) -> Self {
22        Self { pool }
23    }
24
25    fn conn(&self) -> Result<r2d2::PooledConnection<SqliteConnectionManager>> {
26        self.pool.get().map_err(|e| CommerceError::DatabaseError(e.to_string()))
27    }
28
29    fn row_to_reward(row: &rusqlite::Row<'_>) -> rusqlite::Result<Reward> {
30        Ok(Reward {
31            id: parse_uuid_row(&row.get::<_, String>("id")?, "reward", "id")?.into(),
32            program_id: parse_uuid_row(
33                &row.get::<_, String>("program_id")?,
34                "reward",
35                "program_id",
36            )?
37            .into(),
38            name: row.get("name")?,
39            description: row.get("description")?,
40            points_cost: row.get::<_, i64>("points_cost")? as u64,
41            reward_type: parse_enum_row(
42                &row.get::<_, String>("reward_type")?,
43                "reward",
44                "reward_type",
45            )?,
46            value: parse_decimal_opt_row(row.get("value")?, "reward", "value")?,
47            is_active: row.get::<_, i32>("is_active")? != 0,
48            created_at: parse_datetime_row(
49                &row.get::<_, String>("created_at")?,
50                "reward",
51                "created_at",
52            )?,
53            updated_at: parse_datetime_row(
54                &row.get::<_, String>("updated_at")?,
55                "reward",
56                "updated_at",
57            )?,
58        })
59    }
60}
61
62impl RewardRepository for SqliteRewardRepository {
63    fn create(&self, input: CreateReward) -> Result<Reward> {
64        let id = RewardId::new();
65        let now = Utc::now();
66        let id_str = id.to_string();
67        let now_str = now.to_rfc3339();
68
69        with_immediate_transaction(&self.pool, |tx| {
70            tx.execute(
71                "INSERT INTO rewards (id, program_id, name, description, points_cost, reward_type, value, is_active, created_at, updated_at)
72                 VALUES (?, ?, ?, ?, ?, ?, ?, 1, ?, ?)",
73                rusqlite::params![
74                    &id_str,
75                    input.program_id.to_string(),
76                    &input.name,
77                    &input.description,
78                    input.points_cost as i64,
79                    input.reward_type.to_string(),
80                    input.value.map(|v| v.to_string()),
81                    &now_str,
82                    &now_str,
83                ],
84            )?;
85
86            tx.query_row("SELECT * FROM rewards WHERE id = ?", [&id_str], Self::row_to_reward)
87        })
88    }
89
90    fn get(&self, id: RewardId) -> Result<Option<Reward>> {
91        let conn = self.conn()?;
92        match conn.query_row(
93            "SELECT * FROM rewards WHERE id = ?",
94            [id.to_string()],
95            Self::row_to_reward,
96        ) {
97            Ok(r) => Ok(Some(r)),
98            Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None),
99            Err(e) => Err(map_db_error(e)),
100        }
101    }
102
103    fn list(&self, filter: RewardFilter) -> Result<Vec<Reward>> {
104        let conn = self.conn()?;
105        let mut sql = "SELECT * FROM rewards WHERE 1=1".to_string();
106        let mut params: Vec<Box<dyn rusqlite::types::ToSql>> = vec![];
107
108        if let Some(program_id) = filter.program_id {
109            sql.push_str(" AND program_id = ?");
110            params.push(Box::new(program_id.to_string()));
111        }
112        if let Some(reward_type) = filter.reward_type {
113            sql.push_str(" AND reward_type = ?");
114            params.push(Box::new(reward_type.to_string()));
115        }
116        if let Some(is_active) = filter.is_active {
117            sql.push_str(" AND is_active = ?");
118            params.push(Box::new(is_active as i32));
119        }
120
121        sql.push_str(" ORDER BY created_at DESC");
122
123        crate::sqlite::append_limit_offset(&mut sql, filter.limit, filter.offset);
124
125        let param_refs: Vec<&dyn rusqlite::types::ToSql> =
126            params.iter().map(|p| p.as_ref()).collect();
127        let mut stmt = conn.prepare(&sql).map_err(map_db_error)?;
128        let rows = stmt
129            .query_map(param_refs.as_slice(), Self::row_to_reward)
130            .map_err(map_db_error)?
131            .collect::<std::result::Result<Vec<_>, _>>()
132            .map_err(map_db_error)?;
133        Ok(rows)
134    }
135
136    fn delete(&self, id: RewardId) -> Result<()> {
137        let conn = self.conn()?;
138        conn.execute("DELETE FROM rewards WHERE id = ?", [id.to_string()]).map_err(map_db_error)?;
139        Ok(())
140    }
141}
142
143#[cfg(test)]
144mod tests {
145    use super::*;
146    use crate::DatabaseConfig;
147    use crate::sqlite::SqliteDatabase;
148    use rust_decimal_macros::dec;
149    use stateset_core::{LoyaltyProgramId, RewardType};
150
151    fn test_repo() -> SqliteRewardRepository {
152        let db = SqliteDatabase::new(&DatabaseConfig::in_memory()).expect("in-memory db");
153        let conn = db.conn().expect("conn");
154        conn.execute_batch(
155            "CREATE TABLE IF NOT EXISTS rewards (
156                id TEXT PRIMARY KEY,
157                program_id TEXT NOT NULL,
158                name TEXT NOT NULL,
159                description TEXT,
160                points_cost INTEGER NOT NULL DEFAULT 0,
161                reward_type TEXT NOT NULL DEFAULT 'discount',
162                value TEXT,
163                is_active INTEGER NOT NULL DEFAULT 1,
164                created_at TEXT NOT NULL DEFAULT (datetime('now')),
165                updated_at TEXT NOT NULL DEFAULT (datetime('now'))
166            );",
167        )
168        .expect("create table");
169        SqliteRewardRepository::new(db.pool().clone())
170    }
171
172    #[test]
173    fn create_and_get_reward() {
174        let repo = test_repo();
175        let program_id = LoyaltyProgramId::new();
176        let reward = repo
177            .create(CreateReward {
178                program_id,
179                name: "10% Off Coupon".into(),
180                description: Some("Get 10% off your next order".into()),
181                points_cost: 500,
182                reward_type: RewardType::Discount,
183                value: Some(dec!(10.00)),
184            })
185            .expect("create");
186
187        assert_eq!(reward.name, "10% Off Coupon");
188        assert_eq!(reward.points_cost, 500);
189        assert_eq!(reward.value, Some(dec!(10.00)));
190        assert!(reward.is_active);
191
192        let fetched = repo.get(reward.id).expect("get").expect("found");
193        assert_eq!(fetched.id, reward.id);
194        assert_eq!(fetched.program_id, program_id);
195    }
196
197    #[test]
198    fn list_and_delete_rewards() {
199        let repo = test_repo();
200        let program_id = LoyaltyProgramId::new();
201
202        for i in 0..3 {
203            repo.create(CreateReward {
204                program_id,
205                name: format!("Reward {i}"),
206                description: None,
207                points_cost: (i + 1) * 100,
208                reward_type: RewardType::Discount,
209                value: None,
210            })
211            .expect("create");
212        }
213
214        let all = repo.list(RewardFilter::default()).expect("list");
215        assert_eq!(all.len(), 3);
216
217        repo.delete(all[0].id).expect("delete");
218        let remaining = repo.list(RewardFilter::default()).expect("list after delete");
219        assert_eq!(remaining.len(), 2);
220    }
221}