1use super::{
4 map_db_error, parse_datetime_row, parse_decimal_row, parse_enum_row, parse_uuid_row,
5 with_immediate_transaction,
6};
7use chrono::Utc;
8use r2d2::Pool;
9use r2d2_sqlite::SqliteConnectionManager;
10use rust_decimal::Decimal;
11use stateset_core::{
12 CommerceError, CreatePriceLevel, CurrencyCode, PriceAdjustmentType, PriceLevel,
13 PriceLevelEntry, PriceLevelFilter, PriceLevelId, PriceLevelRepository, ProductId, Result,
14 UpdatePriceLevel,
15};
16
17#[derive(Debug)]
18pub struct SqlitePriceLevelRepository {
19 pool: Pool<SqliteConnectionManager>,
20}
21
22impl SqlitePriceLevelRepository {
23 #[must_use]
24 pub const fn new(pool: Pool<SqliteConnectionManager>) -> Self {
25 Self { pool }
26 }
27
28 fn conn(&self) -> Result<r2d2::PooledConnection<SqliteConnectionManager>> {
29 self.pool.get().map_err(|e| CommerceError::DatabaseError(e.to_string()))
30 }
31
32 fn row_to_level(row: &rusqlite::Row<'_>) -> rusqlite::Result<PriceLevel> {
33 Ok(PriceLevel {
34 id: parse_uuid_row(&row.get::<_, String>("id")?, "price_level", "id")?.into(),
35 name: row.get("name")?,
36 code: row.get("code")?,
37 description: row.get("description")?,
38 adjustment_type: parse_enum_row::<PriceAdjustmentType>(
39 &row.get::<_, String>("adjustment_type")?,
40 "price_level",
41 "adjustment_type",
42 )?,
43 adjustment_value: parse_decimal_row(
44 &row.get::<_, String>("adjustment_value")?,
45 "price_level",
46 "adjustment_value",
47 )?,
48 currency: parse_enum_row::<CurrencyCode>(
49 &row.get::<_, String>("currency")?,
50 "price_level",
51 "currency",
52 )?,
53 is_active: row.get::<_, i32>("is_active")? != 0,
54 created_at: parse_datetime_row(
55 &row.get::<_, String>("created_at")?,
56 "price_level",
57 "created_at",
58 )?,
59 updated_at: parse_datetime_row(
60 &row.get::<_, String>("updated_at")?,
61 "price_level",
62 "updated_at",
63 )?,
64 })
65 }
66
67 fn row_to_entry(row: &rusqlite::Row<'_>) -> rusqlite::Result<PriceLevelEntry> {
68 Ok(PriceLevelEntry {
69 price_level_id: parse_uuid_row(
70 &row.get::<_, String>("price_level_id")?,
71 "price_level_entry",
72 "price_level_id",
73 )?
74 .into(),
75 product_id: parse_uuid_row(
76 &row.get::<_, String>("product_id")?,
77 "price_level_entry",
78 "product_id",
79 )?
80 .into(),
81 price: parse_decimal_row(
82 &row.get::<_, String>("price")?,
83 "price_level_entry",
84 "price",
85 )?,
86 created_at: parse_datetime_row(
87 &row.get::<_, String>("created_at")?,
88 "price_level_entry",
89 "created_at",
90 )?,
91 updated_at: parse_datetime_row(
92 &row.get::<_, String>("updated_at")?,
93 "price_level_entry",
94 "updated_at",
95 )?,
96 })
97 }
98}
99
100impl PriceLevelRepository for SqlitePriceLevelRepository {
101 fn create(&self, input: CreatePriceLevel) -> Result<PriceLevel> {
102 let id = PriceLevelId::new();
103 let id_str = id.to_string();
104 let now_str = Utc::now().to_rfc3339();
105 let currency = input.currency.unwrap_or(CurrencyCode::USD);
106 with_immediate_transaction(&self.pool, |tx| {
107 tx.execute(
108 "INSERT INTO price_levels (id, name, code, description, adjustment_type, adjustment_value, currency, is_active, created_at, updated_at)
109 VALUES (?, ?, ?, ?, ?, ?, ?, 1, ?, ?)",
110 rusqlite::params![
111 &id_str,
112 &input.name,
113 &input.code,
114 &input.description,
115 input.adjustment_type.to_string(),
116 input.adjustment_value.to_string(),
117 currency.to_string(),
118 &now_str,
119 &now_str,
120 ],
121 )?;
122 tx.query_row("SELECT * FROM price_levels WHERE id = ?", [&id_str], Self::row_to_level)
123 })
124 }
125
126 fn get(&self, id: PriceLevelId) -> Result<Option<PriceLevel>> {
127 let conn = self.conn()?;
128 match conn.query_row(
129 "SELECT * FROM price_levels WHERE id = ?",
130 [id.to_string()],
131 Self::row_to_level,
132 ) {
133 Ok(l) => Ok(Some(l)),
134 Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None),
135 Err(e) => Err(map_db_error(e)),
136 }
137 }
138
139 fn update(&self, id: PriceLevelId, input: UpdatePriceLevel) -> Result<PriceLevel> {
140 let id_str = id.to_string();
141 let now_str = Utc::now().to_rfc3339();
142 with_immediate_transaction(&self.pool, |tx| {
143 let mut sets = vec!["updated_at = ?".to_string()];
144 let mut params: Vec<Box<dyn rusqlite::types::ToSql>> = vec![Box::new(now_str.clone())];
145
146 if let Some(ref name) = input.name {
147 sets.push("name = ?".into());
148 params.push(Box::new(name.clone()));
149 }
150 if let Some(ref description) = input.description {
151 sets.push("description = ?".into());
152 params.push(Box::new(description.clone()));
153 }
154 if let Some(adjustment_type) = input.adjustment_type {
155 sets.push("adjustment_type = ?".into());
156 params.push(Box::new(adjustment_type.to_string()));
157 }
158 if let Some(adjustment_value) = input.adjustment_value {
159 sets.push("adjustment_value = ?".into());
160 params.push(Box::new(adjustment_value.to_string()));
161 }
162 if let Some(is_active) = input.is_active {
163 sets.push("is_active = ?".into());
164 params.push(Box::new(is_active as i32));
165 }
166
167 let sql = format!("UPDATE price_levels SET {} WHERE id = ?", sets.join(", "));
168 params.push(Box::new(id_str.clone()));
169 let param_refs: Vec<&dyn rusqlite::types::ToSql> =
170 params.iter().map(|p| p.as_ref()).collect();
171 tx.execute(&sql, param_refs.as_slice())?;
172
173 tx.query_row("SELECT * FROM price_levels WHERE id = ?", [&id_str], Self::row_to_level)
174 })
175 }
176
177 fn list(&self, filter: PriceLevelFilter) -> Result<Vec<PriceLevel>> {
178 let conn = self.conn()?;
179 let mut sql = "SELECT * FROM price_levels WHERE 1=1".to_string();
180 let mut params: Vec<Box<dyn rusqlite::types::ToSql>> = vec![];
181 if let Some(active) = filter.is_active {
182 sql.push_str(" AND is_active = ?");
183 params.push(Box::new(active as i32));
184 }
185 sql.push_str(" ORDER BY name ASC");
186 crate::sqlite::append_limit_offset(&mut sql, filter.limit, filter.offset);
187 let param_refs: Vec<&dyn rusqlite::types::ToSql> =
188 params.iter().map(|p| p.as_ref()).collect();
189 let mut stmt = conn.prepare(&sql).map_err(map_db_error)?;
190 let rows = stmt
191 .query_map(param_refs.as_slice(), Self::row_to_level)
192 .map_err(map_db_error)?
193 .collect::<std::result::Result<Vec<_>, _>>()
194 .map_err(map_db_error)?;
195 Ok(rows)
196 }
197
198 fn delete(&self, id: PriceLevelId) -> Result<()> {
199 let id_str = id.to_string();
200 with_immediate_transaction(&self.pool, |tx| {
201 tx.execute("DELETE FROM price_level_entries WHERE price_level_id = ?", [&id_str])?;
202 tx.execute("DELETE FROM price_levels WHERE id = ?", [&id_str])?;
203 Ok(())
204 })
205 }
206
207 fn set_entry(
208 &self,
209 id: PriceLevelId,
210 product_id: ProductId,
211 price: Decimal,
212 ) -> Result<PriceLevelEntry> {
213 let id_str = id.to_string();
214 let product_str = product_id.to_string();
215 let now_str = Utc::now().to_rfc3339();
216 with_immediate_transaction(&self.pool, |tx| {
217 tx.execute(
218 "INSERT INTO price_level_entries (price_level_id, product_id, price, created_at, updated_at)
219 VALUES (?, ?, ?, ?, ?)
220 ON CONFLICT(price_level_id, product_id) DO UPDATE SET
221 price = excluded.price,
222 updated_at = excluded.updated_at",
223 rusqlite::params![&id_str, &product_str, price.to_string(), &now_str, &now_str],
224 )?;
225 tx.query_row(
226 "SELECT * FROM price_level_entries WHERE price_level_id = ? AND product_id = ?",
227 rusqlite::params![&id_str, &product_str],
228 Self::row_to_entry,
229 )
230 })
231 }
232
233 fn delete_entry(&self, id: PriceLevelId, product_id: ProductId) -> Result<()> {
234 let conn = self.conn()?;
235 conn.execute(
236 "DELETE FROM price_level_entries WHERE price_level_id = ? AND product_id = ?",
237 rusqlite::params![id.to_string(), product_id.to_string()],
238 )
239 .map_err(map_db_error)?;
240 Ok(())
241 }
242
243 fn list_entries(&self, id: PriceLevelId) -> Result<Vec<PriceLevelEntry>> {
244 let conn = self.conn()?;
245 let mut stmt = conn
246 .prepare(
247 "SELECT * FROM price_level_entries WHERE price_level_id = ? ORDER BY product_id",
248 )
249 .map_err(map_db_error)?;
250 let rows = stmt
251 .query_map([id.to_string()], Self::row_to_entry)
252 .map_err(map_db_error)?
253 .collect::<std::result::Result<Vec<_>, _>>()
254 .map_err(map_db_error)?;
255 Ok(rows)
256 }
257}
258
259#[cfg(test)]
260mod tests {
261 use super::*;
262 use crate::DatabaseConfig;
263 use crate::sqlite::SqliteDatabase;
264 use rust_decimal_macros::dec;
265
266 fn test_repo() -> SqlitePriceLevelRepository {
267 let db = SqliteDatabase::new(&DatabaseConfig::in_memory()).expect("in-memory db");
268 SqlitePriceLevelRepository::new(db.pool().clone())
269 }
270
271 fn new_level(repo: &SqlitePriceLevelRepository, code: &str) -> PriceLevel {
272 repo.create(CreatePriceLevel {
273 name: "Wholesale".into(),
274 code: code.into(),
275 description: None,
276 adjustment_type: PriceAdjustmentType::PercentageDiscount,
277 adjustment_value: dec!(10),
278 currency: Some(CurrencyCode::USD),
279 })
280 .expect("create level")
281 }
282
283 #[test]
284 fn create_get_update() {
285 let repo = test_repo();
286 let l = new_level(&repo, "WHOLESALE");
287 assert_eq!(l.adjust(dec!(100)), dec!(90));
288 let fetched = repo.get(l.id).expect("get").expect("found");
289 assert_eq!(fetched.code, "WHOLESALE");
290
291 let updated = repo
292 .update(
293 l.id,
294 UpdatePriceLevel { adjustment_value: Some(dec!(25)), ..Default::default() },
295 )
296 .expect("update");
297 assert_eq!(updated.adjust(dec!(100)), dec!(75));
298 }
299
300 #[test]
301 fn entries_upsert_and_resolve() {
302 let repo = test_repo();
303 let l = new_level(&repo, "VIP");
304 let product = ProductId::new();
305 let entry = repo.set_entry(l.id, product, dec!(42)).expect("set entry");
306 assert_eq!(entry.price, dec!(42));
307 let entry = repo.set_entry(l.id, product, dec!(40)).expect("upsert entry");
309 assert_eq!(entry.price, dec!(40));
310 assert_eq!(repo.list_entries(l.id).expect("entries").len(), 1);
311
312 repo.delete_entry(l.id, product).expect("delete entry");
313 assert_eq!(repo.list_entries(l.id).expect("entries").len(), 0);
314 }
315
316 #[test]
317 fn list_and_delete() {
318 let repo = test_repo();
319 let a = new_level(&repo, "A");
320 new_level(&repo, "B");
321 assert_eq!(repo.list(PriceLevelFilter::default()).expect("list").len(), 2);
322 repo.set_entry(a.id, ProductId::new(), dec!(5)).expect("entry");
323 repo.delete(a.id).expect("delete");
324 assert_eq!(repo.list(PriceLevelFilter::default()).expect("list").len(), 1);
325 assert_eq!(repo.list_entries(a.id).expect("entries").len(), 0);
327 }
328}