1use 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}