Skip to main content

stateset_db/sqlite/
price_levels.rs

1//! SQLite implementation of the price level repository
2
3use 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        // upsert overrides
308        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        // entries cascade-deleted
326        assert_eq!(repo.list_entries(a.id).expect("entries").len(), 0);
327    }
328}